Coverage for hopwise/model/general_recommender/lightgcn.py: 89%
71 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
1# @Time : 2020/8/31
2# @Author : Changxin Tian
3# @Email : cx.tian@outlook.com
5# UPDATE:
6# @Time : 2020/9/16, 2021/12/22
7# @Author : Shanlei Mu, Gaowei Zhang
8# @Email : slmu@ruc.edu.cn, 1462034631@qq.com
10r"""LightGCN
11################################################
13Reference:
14 Xiangnan He et al. "LightGCN: Simplifying and Powering Graph Convolution Network for Recommendation." in SIGIR 2020.
16Reference code:
17 https://github.com/kuandeng/LightGCN
18""" # noqa: E501
20import torch
22from hopwise.model.abstract_recommender import GeneralRecommender
23from hopwise.model.init import xavier_uniform_initialization
24from hopwise.model.loss import BPRLoss, EmbLoss
25from hopwise.utils import InputType
28class LightGCN(GeneralRecommender):
29 r"""LightGCN is a GCN-based recommender model.
31 LightGCN includes only the most essential component in GCN — neighborhood aggregation — for
32 collaborative filtering. Specifically, LightGCN learns user and item embeddings by linearly
33 propagating them on the user-item interaction graph, and uses the weighted sum of the embeddings
34 learned at all layers as the final embedding.
36 We implement the model following the original author with a pairwise training mode.
37 """
39 input_type = InputType.PAIRWISE
41 def __init__(self, config, dataset):
42 super().__init__(config, dataset)
44 # load parameters info
45 self.latent_dim = config["embedding_size"] # int type:the embedding size of lightGCN
46 self.n_layers = config["n_layers"] # int type:the layer num of lightGCN
47 self.reg_weight = config["reg_weight"] # float32 type: the weight decay for l2 normalization
48 self.require_pow = config["require_pow"]
50 # define layers and loss
51 self.user_embedding = torch.nn.Embedding(num_embeddings=self.n_users, embedding_dim=self.latent_dim)
52 self.item_embedding = torch.nn.Embedding(num_embeddings=self.n_items, embedding_dim=self.latent_dim)
53 self.mf_loss = BPRLoss()
54 self.reg_loss = EmbLoss()
56 # storage variables for full sort evaluation acceleration
57 self.restore_user_e = None
58 self.restore_item_e = None
60 # generate intermediate data
61 self.norm_adj_matrix = dataset.norm_adjacency_matrix(form="torch.sparse").to(self.device)
63 # parameters initialization
64 self.apply(xavier_uniform_initialization)
65 self.other_parameter_name = ["restore_user_e", "restore_item_e"]
67 def get_ego_embeddings(self):
68 r"""Get the embedding of users and items and combine to an embedding matrix.
70 Returns:
71 Tensor of the embedding matrix. Shape of [n_items+n_users, embedding_dim]
72 """
73 user_embeddings = self.user_embedding.weight
74 item_embeddings = self.item_embedding.weight
75 ego_embeddings = torch.cat([user_embeddings, item_embeddings], dim=0)
76 return ego_embeddings
78 def forward(self):
79 all_embeddings = self.get_ego_embeddings()
80 embeddings_list = [all_embeddings]
82 for layer_idx in range(self.n_layers):
83 all_embeddings = torch.sparse.mm(self.norm_adj_matrix, all_embeddings)
84 embeddings_list.append(all_embeddings)
85 lightgcn_all_embeddings = torch.stack(embeddings_list, dim=1)
86 lightgcn_all_embeddings = torch.mean(lightgcn_all_embeddings, dim=1)
88 user_all_embeddings, item_all_embeddings = torch.split(lightgcn_all_embeddings, [self.n_users, self.n_items])
89 return user_all_embeddings, item_all_embeddings
91 def calculate_loss(self, interaction):
92 # clear the storage variable when training
93 if self.restore_user_e is not None or self.restore_item_e is not None:
94 self.restore_user_e, self.restore_item_e = None, None
96 user = interaction[self.USER_ID]
97 pos_item = interaction[self.ITEM_ID]
98 neg_item = interaction[self.NEG_ITEM_ID]
100 user_all_embeddings, item_all_embeddings = self.forward()
101 u_embeddings = user_all_embeddings[user]
102 pos_embeddings = item_all_embeddings[pos_item]
103 neg_embeddings = item_all_embeddings[neg_item]
105 # calculate BPR Loss
106 pos_scores = torch.mul(u_embeddings, pos_embeddings).sum(dim=1)
107 neg_scores = torch.mul(u_embeddings, neg_embeddings).sum(dim=1)
108 mf_loss = self.mf_loss(pos_scores, neg_scores)
110 # calculate regularization Loss
111 u_ego_embeddings = self.user_embedding(user)
112 pos_ego_embeddings = self.item_embedding(pos_item)
113 neg_ego_embeddings = self.item_embedding(neg_item)
115 reg_loss = self.reg_loss(
116 u_ego_embeddings,
117 pos_ego_embeddings,
118 neg_ego_embeddings,
119 require_pow=self.require_pow,
120 )
122 loss = mf_loss + self.reg_weight * reg_loss
124 return loss
126 def predict(self, interaction):
127 user = interaction[self.USER_ID]
128 item = interaction[self.ITEM_ID]
130 user_all_embeddings, item_all_embeddings = self.forward()
132 u_embeddings = user_all_embeddings[user]
133 i_embeddings = item_all_embeddings[item]
134 scores = torch.mul(u_embeddings, i_embeddings).sum(dim=1)
135 return scores
137 def full_sort_predict(self, interaction):
138 user = interaction[self.USER_ID]
139 if self.restore_user_e is None or self.restore_item_e is None:
140 self.restore_user_e, self.restore_item_e = self.forward()
141 # get user embedding from storage variable
142 u_embeddings = self.restore_user_e[user]
144 # dot with all item embedding to accelerate
145 scores = torch.matmul(u_embeddings, self.restore_item_e.transpose(0, 1))
147 return scores.view(-1)