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
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
1r"""NCL
2################################################
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
8import torch
9import torch.nn.functional as F
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
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 """
22 input_type = InputType.PAIRWISE
24 def __init__(self, config, dataset):
25 super().__init__(config, dataset)
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
32 self.ssl_temp = config["ssl_temp"]
33 self.ssl_reg = config["ssl_reg"]
34 self.hyper_layers = config["hyper_layers"]
36 self.alpha = config["alpha"]
38 self.proto_reg = config["proto_reg"]
39 self.k = config["num_clusters"]
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)
45 self.mf_loss = BPRLoss()
46 self.reg_loss = EmbLoss()
48 # storage variables for full sort evaluation acceleration
49 self.restore_user_e = None
50 self.restore_item_e = None
52 self.norm_adj_matrix = dataset.norm_adjacency_matrix(form="torch.sparse").to(self.device)
54 # parameters initialization
55 self.apply(xavier_uniform_initialization)
56 self.other_parameter_name = ["restore_user_e", "restore_item_e"]
58 self.user_centroids = None
59 self.user_2cluster = None
60 self.item_centroids = None
61 self.item_2cluster = None
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)
69 def run_kmeans(self, x):
70 """Run K-means algorithm to get k clusters of the input tensor x"""
71 import faiss
73 kmeans = faiss.Kmeans(d=self.latent_dim, k=self.k, gpu=True)
74 kmeans.train(x)
75 cluster_cents = kmeans.centroids
77 _, I = kmeans.index.search(x, 1) # noqa: E741
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)
83 node2cluster = torch.LongTensor(I).squeeze().to(self.device)
84 return centroids, node2cluster
86 def get_ego_embeddings(self):
87 r"""Get the embedding of users and items and combine to an embedding matrix.
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
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)
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)
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
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])
113 user_embeddings = user_embeddings_all[user] # [B, e]
114 norm_user_embeddings = F.normalize(user_embeddings)
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)
123 proto_nce_loss_user = -torch.log(pos_score_user / ttl_score_user).sum()
125 item_embeddings = item_embeddings_all[item]
126 norm_item_embeddings = F.normalize(item_embeddings)
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()
136 proto_nce_loss = self.proto_reg * (proto_nce_loss_user + proto_nce_loss_item)
137 return proto_nce_loss
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 )
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)
155 ssl_loss_user = -torch.log(pos_score_user / ttl_score_user).sum()
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)
167 ssl_loss_item = -torch.log(pos_score_item / ttl_score_item).sum()
169 ssl_loss = self.ssl_reg * (ssl_loss_user + self.alpha * ssl_loss_item)
170 return ssl_loss
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
177 user = interaction[self.USER_ID]
178 pos_item = interaction[self.ITEM_ID]
179 neg_item = interaction[self.NEG_ITEM_ID]
181 user_all_embeddings, item_all_embeddings, embeddings_list = self.forward()
183 center_embedding = embeddings_list[0]
184 context_embedding = embeddings_list[self.hyper_layers * 2]
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)
189 u_embeddings = user_all_embeddings[user]
190 pos_embeddings = item_all_embeddings[pos_item]
191 neg_embeddings = item_all_embeddings[neg_item]
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)
197 mf_loss = self.mf_loss(pos_scores, neg_scores)
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)
203 reg_loss = self.reg_loss(u_ego_embeddings, pos_ego_embeddings, neg_ego_embeddings)
205 return mf_loss + self.reg_weight * reg_loss, ssl_loss, proto_loss
207 def predict(self, interaction):
208 user = interaction[self.USER_ID]
209 item = interaction[self.ITEM_ID]
211 user_all_embeddings, item_all_embeddings, embeddings_list = self.forward()
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
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]
225 # dot with all item embedding to accelerate
226 scores = torch.matmul(u_embeddings, self.restore_item_e.transpose(0, 1))
228 return scores.view(-1)