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

1# @Time : 2020/8/31 

2# @Author : Changxin Tian 

3# @Email : cx.tian@outlook.com 

4 

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 

9 

10r"""LightGCN 

11################################################ 

12 

13Reference: 

14 Xiangnan He et al. "LightGCN: Simplifying and Powering Graph Convolution Network for Recommendation." in SIGIR 2020. 

15 

16Reference code: 

17 https://github.com/kuandeng/LightGCN 

18""" # noqa: E501 

19 

20import torch 

21 

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 

26 

27 

28class LightGCN(GeneralRecommender): 

29 r"""LightGCN is a GCN-based recommender model. 

30 

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. 

35 

36 We implement the model following the original author with a pairwise training mode. 

37 """ 

38 

39 input_type = InputType.PAIRWISE 

40 

41 def __init__(self, config, dataset): 

42 super().__init__(config, dataset) 

43 

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"] 

49 

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() 

55 

56 # storage variables for full sort evaluation acceleration 

57 self.restore_user_e = None 

58 self.restore_item_e = None 

59 

60 # generate intermediate data 

61 self.norm_adj_matrix = dataset.norm_adjacency_matrix(form="torch.sparse").to(self.device) 

62 

63 # parameters initialization 

64 self.apply(xavier_uniform_initialization) 

65 self.other_parameter_name = ["restore_user_e", "restore_item_e"] 

66 

67 def get_ego_embeddings(self): 

68 r"""Get the embedding of users and items and combine to an embedding matrix. 

69 

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 

77 

78 def forward(self): 

79 all_embeddings = self.get_ego_embeddings() 

80 embeddings_list = [all_embeddings] 

81 

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) 

87 

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 

90 

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 

95 

96 user = interaction[self.USER_ID] 

97 pos_item = interaction[self.ITEM_ID] 

98 neg_item = interaction[self.NEG_ITEM_ID] 

99 

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] 

104 

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) 

109 

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) 

114 

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 ) 

121 

122 loss = mf_loss + self.reg_weight * reg_loss 

123 

124 return loss 

125 

126 def predict(self, interaction): 

127 user = interaction[self.USER_ID] 

128 item = interaction[self.ITEM_ID] 

129 

130 user_all_embeddings, item_all_embeddings = self.forward() 

131 

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 

136 

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] 

143 

144 # dot with all item embedding to accelerate 

145 scores = torch.matmul(u_embeddings, self.restore_item_e.transpose(0, 1)) 

146 

147 return scores.view(-1)