Coverage for hopwise/model/general_recommender/ncl.py: 95%

146 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1r"""NCL 

2################################################ 

3 

4Reference: 

5 Zihan Lin*, Changxin Tian*, Yupeng Hou*, Wayne Xin Zhao. "Improving Graph Collaborative Filtering with Neighborhood-enriched Contrastive Learning." in WWW 2022. 

6""" # noqa: E501 

7 

8import torch 

9import torch.nn.functional as F 

10 

11from hopwise.model.abstract_recommender import GeneralRecommender 

12from hopwise.model.init import xavier_uniform_initialization 

13from hopwise.model.loss import BPRLoss, EmbLoss 

14from hopwise.utils import InputType 

15 

16 

17class NCL(GeneralRecommender): 

18 r"""NCL is a neighborhood-enriched contrastive learning paradigm for graph collaborative filtering. 

19 Both structural and semantic neighbors are explicitly captured as contrastive learning objects. 

20 """ 

21 

22 input_type = InputType.PAIRWISE 

23 

24 def __init__(self, config, dataset): 

25 super().__init__(config, dataset) 

26 

27 # load parameters info 

28 self.latent_dim = config["embedding_size"] # int type: the embedding size of the base model 

29 self.n_layers = config["n_layers"] # int type: the layer num of the base model 

30 self.reg_weight = config["reg_weight"] # float32 type: the weight decay for l2 normalization 

31 

32 self.ssl_temp = config["ssl_temp"] 

33 self.ssl_reg = config["ssl_reg"] 

34 self.hyper_layers = config["hyper_layers"] 

35 

36 self.alpha = config["alpha"] 

37 

38 self.proto_reg = config["proto_reg"] 

39 self.k = config["num_clusters"] 

40 

41 # define layers and loss 

42 self.user_embedding = torch.nn.Embedding(num_embeddings=self.n_users, embedding_dim=self.latent_dim) 

43 self.item_embedding = torch.nn.Embedding(num_embeddings=self.n_items, embedding_dim=self.latent_dim) 

44 

45 self.mf_loss = BPRLoss() 

46 self.reg_loss = EmbLoss() 

47 

48 # storage variables for full sort evaluation acceleration 

49 self.restore_user_e = None 

50 self.restore_item_e = None 

51 

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

53 

54 # parameters initialization 

55 self.apply(xavier_uniform_initialization) 

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

57 

58 self.user_centroids = None 

59 self.user_2cluster = None 

60 self.item_centroids = None 

61 self.item_2cluster = None 

62 

63 def e_step(self): 

64 user_embeddings = self.user_embedding.weight.detach().cpu().numpy() 

65 item_embeddings = self.item_embedding.weight.detach().cpu().numpy() 

66 self.user_centroids, self.user_2cluster = self.run_kmeans(user_embeddings) 

67 self.item_centroids, self.item_2cluster = self.run_kmeans(item_embeddings) 

68 

69 def run_kmeans(self, x): 

70 """Run K-means algorithm to get k clusters of the input tensor x""" 

71 import faiss 

72 

73 kmeans = faiss.Kmeans(d=self.latent_dim, k=self.k, gpu=True) 

74 kmeans.train(x) 

75 cluster_cents = kmeans.centroids 

76 

77 _, I = kmeans.index.search(x, 1) # noqa: E741 

78 

79 # convert to cuda Tensors for broadcast 

80 centroids = torch.Tensor(cluster_cents).to(self.device) 

81 centroids = F.normalize(centroids, p=2, dim=1) 

82 

83 node2cluster = torch.LongTensor(I).squeeze().to(self.device) 

84 return centroids, node2cluster 

85 

86 def get_ego_embeddings(self): 

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

88 

89 Returns: 

90 Tensor of the embedding matrix. Shape of [n_items+n_users, embedding_dim] 

91 """ 

92 user_embeddings = self.user_embedding.weight 

93 item_embeddings = self.item_embedding.weight 

94 ego_embeddings = torch.cat([user_embeddings, item_embeddings], dim=0) 

95 return ego_embeddings 

96 

97 def forward(self): 

98 all_embeddings = self.get_ego_embeddings() 

99 embeddings_list = [all_embeddings] 

100 for layer_idx in range(max(self.n_layers, self.hyper_layers * 2)): 

101 all_embeddings = torch.sparse.mm(self.norm_adj_matrix, all_embeddings) 

102 embeddings_list.append(all_embeddings) 

103 

104 lightgcn_all_embeddings = torch.stack(embeddings_list[: self.n_layers + 1], dim=1) 

105 lightgcn_all_embeddings = torch.mean(lightgcn_all_embeddings, dim=1) 

106 

107 user_all_embeddings, item_all_embeddings = torch.split(lightgcn_all_embeddings, [self.n_users, self.n_items]) 

108 return user_all_embeddings, item_all_embeddings, embeddings_list 

109 

110 def ProtoNCE_loss(self, node_embedding, user, item): 

111 user_embeddings_all, item_embeddings_all = torch.split(node_embedding, [self.n_users, self.n_items]) 

112 

113 user_embeddings = user_embeddings_all[user] # [B, e] 

114 norm_user_embeddings = F.normalize(user_embeddings) 

115 

116 user2cluster = self.user_2cluster[user] # [B,] 

117 user2centroids = self.user_centroids[user2cluster] # [B, e] 

118 pos_score_user = torch.mul(norm_user_embeddings, user2centroids).sum(dim=1) 

119 pos_score_user = torch.exp(pos_score_user / self.ssl_temp) 

120 ttl_score_user = torch.matmul(norm_user_embeddings, self.user_centroids.transpose(0, 1)) 

121 ttl_score_user = torch.exp(ttl_score_user / self.ssl_temp).sum(dim=1) 

122 

123 proto_nce_loss_user = -torch.log(pos_score_user / ttl_score_user).sum() 

124 

125 item_embeddings = item_embeddings_all[item] 

126 norm_item_embeddings = F.normalize(item_embeddings) 

127 

128 item2cluster = self.item_2cluster[item] # [B, ] 

129 item2centroids = self.item_centroids[item2cluster] # [B, e] 

130 pos_score_item = torch.mul(norm_item_embeddings, item2centroids).sum(dim=1) 

131 pos_score_item = torch.exp(pos_score_item / self.ssl_temp) 

132 ttl_score_item = torch.matmul(norm_item_embeddings, self.item_centroids.transpose(0, 1)) 

133 ttl_score_item = torch.exp(ttl_score_item / self.ssl_temp).sum(dim=1) 

134 proto_nce_loss_item = -torch.log(pos_score_item / ttl_score_item).sum() 

135 

136 proto_nce_loss = self.proto_reg * (proto_nce_loss_user + proto_nce_loss_item) 

137 return proto_nce_loss 

138 

139 def ssl_layer_loss(self, current_embedding, previous_embedding, user, item): 

140 current_user_embeddings, current_item_embeddings = torch.split(current_embedding, [self.n_users, self.n_items]) 

141 previous_user_embeddings_all, previous_item_embeddings_all = torch.split( 

142 previous_embedding, [self.n_users, self.n_items] 

143 ) 

144 

145 current_user_embeddings = current_user_embeddings[user] 

146 previous_user_embeddings = previous_user_embeddings_all[user] 

147 norm_user_emb1 = F.normalize(current_user_embeddings) 

148 norm_user_emb2 = F.normalize(previous_user_embeddings) 

149 norm_all_user_emb = F.normalize(previous_user_embeddings_all) 

150 pos_score_user = torch.mul(norm_user_emb1, norm_user_emb2).sum(dim=1) 

151 ttl_score_user = torch.matmul(norm_user_emb1, norm_all_user_emb.transpose(0, 1)) 

152 pos_score_user = torch.exp(pos_score_user / self.ssl_temp) 

153 ttl_score_user = torch.exp(ttl_score_user / self.ssl_temp).sum(dim=1) 

154 

155 ssl_loss_user = -torch.log(pos_score_user / ttl_score_user).sum() 

156 

157 current_item_embeddings = current_item_embeddings[item] 

158 previous_item_embeddings = previous_item_embeddings_all[item] 

159 norm_item_emb1 = F.normalize(current_item_embeddings) 

160 norm_item_emb2 = F.normalize(previous_item_embeddings) 

161 norm_all_item_emb = F.normalize(previous_item_embeddings_all) 

162 pos_score_item = torch.mul(norm_item_emb1, norm_item_emb2).sum(dim=1) 

163 ttl_score_item = torch.matmul(norm_item_emb1, norm_all_item_emb.transpose(0, 1)) 

164 pos_score_item = torch.exp(pos_score_item / self.ssl_temp) 

165 ttl_score_item = torch.exp(ttl_score_item / self.ssl_temp).sum(dim=1) 

166 

167 ssl_loss_item = -torch.log(pos_score_item / ttl_score_item).sum() 

168 

169 ssl_loss = self.ssl_reg * (ssl_loss_user + self.alpha * ssl_loss_item) 

170 return ssl_loss 

171 

172 def calculate_loss(self, interaction): 

173 # clear the storage variable when training 

174 if self.restore_user_e is not None or self.restore_item_e is not None: 

175 self.restore_user_e, self.restore_item_e = None, None 

176 

177 user = interaction[self.USER_ID] 

178 pos_item = interaction[self.ITEM_ID] 

179 neg_item = interaction[self.NEG_ITEM_ID] 

180 

181 user_all_embeddings, item_all_embeddings, embeddings_list = self.forward() 

182 

183 center_embedding = embeddings_list[0] 

184 context_embedding = embeddings_list[self.hyper_layers * 2] 

185 

186 ssl_loss = self.ssl_layer_loss(context_embedding, center_embedding, user, pos_item) 

187 proto_loss = self.ProtoNCE_loss(center_embedding, user, pos_item) 

188 

189 u_embeddings = user_all_embeddings[user] 

190 pos_embeddings = item_all_embeddings[pos_item] 

191 neg_embeddings = item_all_embeddings[neg_item] 

192 

193 # calculate BPR Loss 

194 pos_scores = torch.mul(u_embeddings, pos_embeddings).sum(dim=1) 

195 neg_scores = torch.mul(u_embeddings, neg_embeddings).sum(dim=1) 

196 

197 mf_loss = self.mf_loss(pos_scores, neg_scores) 

198 

199 u_ego_embeddings = self.user_embedding(user) 

200 pos_ego_embeddings = self.item_embedding(pos_item) 

201 neg_ego_embeddings = self.item_embedding(neg_item) 

202 

203 reg_loss = self.reg_loss(u_ego_embeddings, pos_ego_embeddings, neg_ego_embeddings) 

204 

205 return mf_loss + self.reg_weight * reg_loss, ssl_loss, proto_loss 

206 

207 def predict(self, interaction): 

208 user = interaction[self.USER_ID] 

209 item = interaction[self.ITEM_ID] 

210 

211 user_all_embeddings, item_all_embeddings, embeddings_list = self.forward() 

212 

213 u_embeddings = user_all_embeddings[user] 

214 i_embeddings = item_all_embeddings[item] 

215 scores = torch.mul(u_embeddings, i_embeddings).sum(dim=1) 

216 return scores 

217 

218 def full_sort_predict(self, interaction): 

219 user = interaction[self.USER_ID] 

220 if self.restore_user_e is None or self.restore_item_e is None: 

221 self.restore_user_e, self.restore_item_e, embedding_list = self.forward() 

222 # get user embedding from storage variable 

223 u_embeddings = self.restore_user_e[user] 

224 

225 # dot with all item embedding to accelerate 

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

227 

228 return scores.view(-1)