Coverage for hopwise/model/knowledge_aware_recommender/mcclk.py: 90%
300 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 : 2022/8/22
2# @Author : Bowen Zheng
3# @Email : 18735382001@163.com
5r"""MCCLK
6##################################################
7Reference:
8 Ding Zou et al. "Multi-level Cross-view Contrastive Learning for Knowledge-aware Recommender System." in SIGIR 2022.
10Reference code:
11 https://github.com/CCIIPLab/MCCLK
12""" # noqa: E501
14import numpy as np
15import torch
16import torch.nn.functional as F
17from torch import nn
19from hopwise.model.abstract_recommender import KnowledgeRecommender
20from hopwise.model.init import xavier_normal_initialization
21from hopwise.model.layers import SparseDropout
22from hopwise.model.loss import BPRLoss, EmbLoss
23from hopwise.utils import InputType
26class Aggregator(nn.Module):
27 def __init__(self, item_only=False, attention=True):
28 super().__init__()
30 # Only aggregate item embedding
31 self.item_only = item_only
32 # Whether use attention mechanism
33 self.attention = attention
35 def forward(self, entity_emb, user_emb, relation_emb, edge_index, edge_type, inter_matrix):
36 from torch_geometric.utils import scatter
37 from torch_geometric.utils import softmax as scatter_softmax
39 n_entities = entity_emb.shape[0]
41 # KG aggregate
42 head, tail = edge_index
43 edge_relation_emb = relation_emb[edge_type]
44 neigh_relation_emb = entity_emb[tail] * edge_relation_emb # [-1, embedding_size]
46 if self.attention:
47 # Calculate attention weights
48 neigh_relation_emb_weight = self.calculate_sim_hrt(entity_emb[head], entity_emb[tail], edge_relation_emb)
49 # [-1, 1] -> [-1, embedding_size]
50 neigh_relation_emb_weight = neigh_relation_emb_weight.expand(
51 neigh_relation_emb.shape[0], neigh_relation_emb.shape[1]
52 )
53 neigh_relation_emb_weight = scatter_softmax(
54 neigh_relation_emb_weight, index=head, dim=0
55 ) # [-1, embedding_size]
56 neigh_relation_emb = torch.mul(neigh_relation_emb_weight, neigh_relation_emb)
58 entity_agg = scatter(
59 src=neigh_relation_emb, index=head, dim_size=n_entities, dim=0, reduce="mean"
60 ) # [n_entities, embedding_size]
62 # Only aggregate item embedding
63 if self.item_only:
64 return entity_agg
66 user_agg = torch.sparse.mm(inter_matrix, entity_emb) # [n_users, embedding_size]
67 # The importance of relation to user
68 score = torch.mm(user_emb, relation_emb.t()) # [n_users, n_relations]
69 score = torch.softmax(score, dim=-1)
70 user_agg = user_agg + (torch.mm(score, relation_emb)) * user_agg
72 return entity_agg, user_agg
74 def calculate_sim_hrt(self, entity_emb_head, entity_emb_tail, relation_emb):
75 r"""The calculation method of attention weight here follows the code implementation of the author, which is
76 slightly different from that described in the paper.
77 """
78 tail_relation_emb = entity_emb_tail * relation_emb
79 tail_relation_emb = tail_relation_emb.norm(dim=1, p=2, keepdim=True)
80 head_relation_emb = entity_emb_head * relation_emb
81 head_relation_emb = head_relation_emb.norm(dim=1, p=2, keepdim=True)
82 # [-1, 1, embedding_size] * [-1, embedding_size, 1] -> [-1, 1]
83 att_weights = torch.matmul(head_relation_emb.unsqueeze(dim=1), tail_relation_emb.unsqueeze(dim=2)).squeeze(
84 dim=-1
85 )
86 att_weights = att_weights**2
87 return att_weights
90class GraphConv(nn.Module):
91 """Graph Convolutional Network"""
93 def __init__(
94 self,
95 config,
96 embedding_size,
97 n_relations,
98 edge_index,
99 edge_type,
100 inter_matrix,
101 device,
102 ):
103 super().__init__()
105 # load parameters info
106 self.n_relations = n_relations
107 self.edge_index = edge_index
108 self.edge_type = edge_type
109 self.inter_matrix = inter_matrix
110 self.embedding_size = embedding_size
111 self.n_hops = config["n_hops"]
112 self.node_dropout_rate = config["node_dropout_rate"]
113 self.mess_dropout_rate = config["mess_dropout_rate"]
114 self.topk = config["k"]
115 self.lambda_coeff = config["lambda_coeff"]
116 self.build_graph_separately = config["build_graph_separately"]
117 self.device = device
119 # define layers
120 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size)
122 # User a separate GCN to build item-item graph
123 if self.build_graph_separately:
124 r"""
125 In the original author's implementation(https://github.com/CCIIPLab/MCCLK), the process of constructing
126 k-Nearest-Neighbor item-item semantic graph(section 4.1 in paper) and encoding structural view(section 4.3.1 in paper)
127 are combined. This implementation improves the computational efficiency, but is slightly different from the
128 model structure described in the paper. We use the parameter `build_graph_separately` to control whether to
129 use a separate GCN to build a item-item semantic graph. If `build_graph_separately` is set to true, the model
130 structure will be the same as that described in the paper. Otherwise, the author's code implementation will be followed.
131 """ # noqa: E501
132 self.bg_convs = nn.ModuleList()
133 for i in range(self.n_hops):
134 self.bg_convs.append(Aggregator(item_only=True, attention=False))
136 self.convs = nn.ModuleList()
137 for i in range(self.n_hops):
138 self.convs.append(Aggregator())
140 self.node_dropout = SparseDropout(p=self.mess_dropout_rate) # node dropout
141 self.mess_dropout = nn.Dropout(p=self.mess_dropout_rate) # mess dropout
143 # parameters initialization
144 self.apply(xavier_normal_initialization)
146 def edge_sampling(self, edge_index, edge_type, rate=0.5):
147 # edge_index: [2, -1]
148 # edge_type: [-1]
149 n_edges = edge_index.shape[1]
150 random_indices = np.random.choice(n_edges, size=int(n_edges * rate), replace=False)
151 return edge_index[:, random_indices], edge_type[random_indices]
153 def forward(self, user_emb, entity_emb):
154 # node dropout
155 if self.node_dropout_rate > 0.0:
156 edge_index, edge_type = self.edge_sampling(self.edge_index, self.edge_type, self.node_dropout_rate)
157 inter_matrix = self.node_dropout(self.inter_matrix)
158 else:
159 edge_index, edge_type = self.edge_index, self.edge_type
160 inter_matrix = self.inter_matrix
162 origin_entity_emb = entity_emb
164 entity_res_emb = [entity_emb] # [n_entities, embedding_size]
165 user_res_emb = [user_emb] # [n_users, embedding_size]
166 relation_emb = self.relation_embedding.weight # [n_relations, embedding_size]
167 for i in range(len(self.convs)):
168 entity_emb, user_emb = self.convs[i](
169 entity_emb, user_emb, relation_emb, edge_index, edge_type, inter_matrix
170 )
171 # message dropout
172 if self.mess_dropout_rate > 0.0:
173 entity_emb = self.mess_dropout(entity_emb)
174 user_emb = self.mess_dropout(user_emb)
175 entity_emb = F.normalize(entity_emb)
176 user_emb = F.normalize(user_emb)
177 # result embedding
178 entity_res_emb.append(entity_emb)
179 user_res_emb.append(user_emb)
181 entity_res_emb = torch.stack(entity_res_emb, dim=1)
182 entity_res_emb = entity_res_emb.mean(dim=1, keepdim=False)
183 user_res_emb = torch.stack(user_res_emb, dim=1)
184 user_res_emb = user_res_emb.mean(dim=1, keepdim=False)
186 # build item-item graph
187 if self.build_graph_separately:
188 item_adj = self._build_graph_separately(origin_entity_emb)
189 else:
190 # build origin item-item graph
191 origin_item_adj = self.build_adj(origin_entity_emb, self.topk)
192 # update item-item graph
193 item_adj = (1 - self.lambda_coeff) * self.build_adj(
194 entity_res_emb, self.topk
195 ) + self.lambda_coeff * origin_item_adj
197 return entity_res_emb, user_res_emb, item_adj
199 def build_adj(self, context, topk):
200 r"""Construct a k-Nearest-Neighbor item-item semantic graph.
202 Returns:
203 Sparse tensor of the normalized item-item matrix.
204 """
205 # construct similarity adj matrix
206 n_entities = context.shape[0]
207 context_norm = context.div(torch.norm(context, p=2, dim=-1, keepdim=True)).cpu()
208 sim = torch.mm(context_norm, context_norm.transpose(1, 0))
209 # knn_val: [n_entities, topk] knn_index: [n_entities, topk]
210 knn_val, knn_index = torch.topk(sim, topk, dim=-1)
211 knn_val, knn_index = knn_val.to(self.device), knn_index.to(self.device)
213 y = knn_index.reshape(-1)
214 x = torch.arange(0, n_entities).unsqueeze(dim=-1).to(self.device) # [n_entities, 1]
215 x = x.expand(n_entities, topk).reshape(-1)
216 indice = torch.cat((x.unsqueeze(dim=0), y.unsqueeze(dim=0)), dim=0) # [2, n_entities * topk]
217 value = knn_val.reshape(-1)
218 adj_sparsity = torch.sparse.FloatTensor(indice.data, value.data, torch.Size([n_entities, n_entities])).to(
219 self.device
220 )
222 # normalized laplacian adj
223 rowsum = torch.sparse.sum(adj_sparsity, dim=1)
224 d_inv_sqrt = torch.pow(rowsum, -0.5)
225 d_mat_inv_sqrt_value = d_inv_sqrt._values()
226 x = torch.arange(0, n_entities).unsqueeze(dim=0).to(self.device)
227 x = x.expand(2, n_entities)
228 d_mat_inv_sqrt_indice = x
229 d_mat_inv_sqrt = torch.sparse.FloatTensor(
230 d_mat_inv_sqrt_indice,
231 d_mat_inv_sqrt_value,
232 torch.Size([n_entities, n_entities]),
233 )
234 L_norm = torch.sparse.mm(torch.sparse.mm(d_mat_inv_sqrt, adj_sparsity), d_mat_inv_sqrt)
235 return L_norm
237 def _build_graph_separately(self, entity_emb):
238 # node dropout
239 if self.node_dropout_rate > 0.0:
240 edge_index, edge_type = self.edge_sampling(self.edge_index, self.edge_type, self.node_dropout_rate)
241 inter_matrix = self.node_dropout(self.inter_matrix)
242 else:
243 edge_index, edge_type = self.edge_index, self.edge_type
244 inter_matrix = self.inter_matrix
246 origin_item_adj = self.build_adj(entity_emb, self.topk)
248 entity_res_emb = [entity_emb] # [n_entities, embedding_size]
249 relation_emb = self.relation_embedding.weight # [n_relations, embedding_size]
250 for i in range(len(self.bg_convs)):
251 entity_emb = self.bg_convs[i](entity_emb, None, relation_emb, edge_index, edge_type, inter_matrix)
252 # message dropout
253 if self.mess_dropout_rate > 0.0:
254 entity_emb = self.mess_dropout(entity_emb)
255 entity_emb = F.normalize(entity_emb)
256 # result embedding
257 entity_res_emb.append(entity_emb)
259 entity_res_emb = torch.stack(entity_res_emb, dim=1)
260 entity_res_emb = entity_res_emb.mean(dim=1, keepdim=False)
262 item_adj = (1 - self.lambda_coeff) * self.build_adj(
263 entity_res_emb, self.topk
264 ) + self.lambda_coeff * origin_item_adj
266 return item_adj
269class MCCLK(KnowledgeRecommender):
270 r"""MCCLK is a knowledge-based recommendation model.
271 It focuses on the contrastive learning in KG-aware recommendation and proposes a novel multi-level cross-view
272 contrastive learning mechanism. This model comprehensively considers three different graph views for KG-aware
273 recommendation, including global-level structural view, local-level collaborative and semantic views. It hence
274 performs contrastive learning across three views on both local and global levels, mining comprehensive graph
275 feature and structure information in a self-supervised manner.
276 """
278 input_type = InputType.PAIRWISE
280 def __init__(self, config, dataset):
281 super().__init__(config, dataset)
283 # load parameters info
284 self.embedding_size = config["embedding_size"]
285 self.reg_weight = config["reg_weight"]
286 self.lightgcn_layer = config["lightgcn_layer"]
287 self.item_agg_layer = config["item_agg_layer"]
288 self.temperature = config["temperature"]
289 self.alpha = config["alpha"]
290 self.beta = config["beta"]
291 self.loss_type = config["loss_type"]
293 # load dataset info
294 # inter_matrix: [n_users, n_entities]; inter_graph: [n_users + n_entities, n_users + n_entities]
295 self.inter_matrix, self.inter_graph = dataset._create_norm_ckg_adjacency_matrix(symmetric=False)
296 self.inter_matrix = self.inter_matrix.to(self.device)
297 self.inter_graph = self.inter_graph.to(self.device)
298 self.kg_graph = dataset.kg_graph(form="coo", value_field="relation_id") # [n_entities, n_entities]
299 # edge_index: [2, -1]; edge_type: [-1,]
300 self.edge_index, self.edge_type = self.get_edges(self.kg_graph)
302 # define layers
303 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
304 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size)
305 self.gcn = GraphConv(
306 config=config,
307 embedding_size=self.embedding_size,
308 n_relations=self.n_relations,
309 edge_index=self.edge_index,
310 edge_type=self.edge_type,
311 inter_matrix=self.inter_matrix,
312 device=self.device,
313 )
314 self.fc1 = nn.Sequential(
315 nn.Linear(self.embedding_size, self.embedding_size, bias=True),
316 nn.ReLU(),
317 nn.Linear(self.embedding_size, self.embedding_size, bias=True),
318 )
319 self.fc2 = nn.Sequential(
320 nn.Linear(self.embedding_size, self.embedding_size, bias=True),
321 nn.ReLU(),
322 nn.Linear(self.embedding_size, self.embedding_size, bias=True),
323 )
324 self.fc3 = nn.Sequential(
325 nn.Linear(self.embedding_size, self.embedding_size, bias=True),
326 nn.ReLU(),
327 nn.Linear(self.embedding_size, self.embedding_size, bias=True),
328 )
329 # define loss
330 if self.loss_type.lower() == "bpr":
331 self.rec_loss = BPRLoss()
332 elif self.loss_type.lower() == "bce":
333 self.sigmoid = nn.Sigmoid()
334 self.rec_loss = nn.BCEWithLogitsLoss()
335 else:
336 raise NotImplementedError(f"The loss type [{self.loss_type}] has not been supported.")
337 self.reg_loss = EmbLoss()
339 # storage variables for full sort evaluation acceleration
340 self.restore_user_e = None
341 self.restore_item_e = None
343 # parameters initialization
344 self.apply(xavier_normal_initialization)
346 def get_edges(self, graph):
347 index = torch.LongTensor(np.array([graph.row, graph.col]))
348 type = torch.LongTensor(np.array(graph.data))
349 return index.to(self.device), type.to(self.device)
351 def forward(self):
352 user_emb = self.user_embedding.weight
353 entity_emb = self.entity_embedding.weight
354 # Construct a k-Nearest-Neighbor item-item semantic graph and Structural View Encoder
355 entity_gcn_emb, user_gcn_emb, item_adj = self.gcn(user_emb, entity_emb)
356 # Semantic View Encoder
357 item_semantic_emb = [entity_emb]
358 item_agg_emb = entity_emb
359 for i in range(self.item_agg_layer):
360 item_agg_emb = torch.sparse.mm(item_adj, item_agg_emb)
361 item_semantic_emb.append(item_agg_emb)
362 item_semantic_emb = torch.stack(item_semantic_emb, dim=1)
363 item_semantic_emb = item_semantic_emb.mean(dim=1, keepdim=False)
364 # item_semantic_emb = F.normalize(item_semantic_emb, p=2, dim=1)
366 # Collaborative View Encoder
367 user_lightgcn_emb, item_lightgcn_emb = self.light_gcn(user_emb, entity_emb, self.inter_graph)
369 return (
370 item_semantic_emb,
371 user_lightgcn_emb,
372 item_lightgcn_emb,
373 user_gcn_emb,
374 entity_gcn_emb,
375 )
377 def light_gcn(self, user_embedding, item_embedding, adj):
378 ego_embeddings = torch.cat((user_embedding, item_embedding), dim=0)
379 all_embeddings = [ego_embeddings]
380 for i in range(self.lightgcn_layer):
381 side_embeddings = torch.sparse.mm(adj, ego_embeddings)
382 ego_embeddings = side_embeddings
383 all_embeddings += [ego_embeddings]
384 all_embeddings = torch.stack(all_embeddings, dim=1)
385 all_embeddings = all_embeddings.mean(dim=1, keepdim=False)
386 u_g_embeddings, i_g_embeddings = torch.split(all_embeddings, [self.n_users, self.n_entities], dim=0)
387 return u_g_embeddings, i_g_embeddings
389 def sim(self, z1: torch.Tensor, z2: torch.Tensor):
390 z1 = F.normalize(z1)
391 z2 = F.normalize(z2)
392 return torch.mm(z1, z2.t())
394 def calculate_loss(self, interaction):
395 if self.restore_user_e is not None or self.restore_item_e is not None:
396 self.restore_user_e, self.restore_item_e = None, None
398 # get loss for training rs
399 user = interaction[self.USER_ID]
400 pos_item = interaction[self.ITEM_ID]
401 neg_item = interaction[self.NEG_ITEM_ID]
402 all_item = torch.cat((pos_item, neg_item), dim=0)
404 (
405 item_semantic_emb,
406 user_lightgcn_emb,
407 item_lightgcn_emb,
408 user_gcn_emb,
409 item_gcn_emb,
410 ) = self.forward()
411 item_emb_1 = item_semantic_emb[all_item]
412 user_emb_1 = user_lightgcn_emb[user]
413 item_emb_2 = item_lightgcn_emb[all_item]
414 user_emb_2 = user_gcn_emb[user]
415 item_emb_3 = item_gcn_emb[all_item]
417 local_loss = self.local_level_loss(item_emb_1, item_emb_2)
418 global_loss = self.global_level_loss_1(user_emb_2, user_emb_1) + self.global_level_loss_2(
419 item_emb_3, item_emb_1 + item_emb_2
420 )
422 user_embedding = torch.cat((user_emb_2, user_emb_1), dim=-1)
423 pos_item_embedding = torch.cat(
424 (
425 item_gcn_emb[pos_item],
426 item_semantic_emb[pos_item] + item_lightgcn_emb[pos_item],
427 ),
428 dim=-1,
429 )
430 neg_item_embedding = torch.cat(
431 (
432 item_gcn_emb[neg_item],
433 item_semantic_emb[neg_item] + item_lightgcn_emb[neg_item],
434 ),
435 dim=-1,
436 )
438 pos_scores = torch.mul(user_embedding, pos_item_embedding).sum(dim=1)
439 neg_scores = torch.mul(user_embedding, neg_item_embedding).sum(dim=1)
440 if self.loss_type.lower() == "bpr":
441 rec_loss = self.rec_loss(pos_scores, neg_scores)
442 else:
443 predict = torch.cat((pos_scores, neg_scores))
444 target = torch.zeros(len(pos_item) + len(neg_item), dtype=torch.float32).to(self.device)
445 target[: len(pos_item)] = 1
446 rec_loss = self.rec_loss(predict, target)
448 reg_loss = self.reg_loss(user_embedding, pos_item_embedding, neg_item_embedding)
449 loss = (
450 rec_loss
451 + self.reg_weight * reg_loss
452 + self.beta * (self.alpha * local_loss + (1 - self.alpha) * global_loss)
453 )
455 return loss
457 def local_level_loss(self, A_embedding, B_embedding):
458 # The loss of local-level contrastive learning
459 def exp_temp(x):
460 return torch.exp(x / self.temperature)
462 A_embedding = self.fc1(A_embedding)
463 B_embedding = self.fc1(B_embedding)
464 refl_sim = exp_temp(self.sim(A_embedding, A_embedding))
465 between_sim = exp_temp(self.sim(A_embedding, B_embedding))
466 local_loss = -torch.log(between_sim.diag() / (refl_sim.sum(1) + between_sim.sum(1) - refl_sim.diag()))
467 local_loss = local_loss.mean()
468 return local_loss
470 def global_level_loss_1(self, A_embedding, B_embedding):
471 # The user embedding loss of global-level contrastive learning
472 def exp_temp(x):
473 return torch.exp(x / self.temperature)
475 A_embedding = self.fc2(A_embedding)
476 B_embedding = self.fc2(B_embedding)
478 refl_sim_1 = exp_temp(self.sim(A_embedding, A_embedding))
479 between_sim_1 = exp_temp(self.sim(A_embedding, B_embedding))
480 loss_1 = -torch.log(between_sim_1.diag() / (refl_sim_1.sum(1) + between_sim_1.sum(1) - refl_sim_1.diag()))
482 refl_sim_2 = exp_temp(self.sim(B_embedding, B_embedding))
483 between_sim_2 = exp_temp(self.sim(B_embedding, A_embedding))
484 loss_2 = -torch.log(between_sim_2.diag() / (refl_sim_2.sum(1) + between_sim_2.sum(1) - refl_sim_2.diag()))
486 global_user_loss = (loss_1 + loss_2) * 0.5
487 global_user_loss = global_user_loss.mean()
488 return global_user_loss
490 def global_level_loss_2(self, A_embedding, B_embedding):
491 # The item embedding loss of global-level contrastive learning
492 def exp_temp(x):
493 return torch.exp(x / self.temperature)
495 A_embedding = self.fc3(A_embedding)
496 B_embedding = self.fc3(B_embedding)
498 refl_sim_1 = exp_temp(self.sim(A_embedding, A_embedding))
499 between_sim_1 = exp_temp(self.sim(A_embedding, B_embedding))
500 loss_1 = -torch.log(between_sim_1.diag() / (refl_sim_1.sum(1) + between_sim_1.sum(1) - refl_sim_1.diag()))
502 refl_sim_2 = exp_temp(self.sim(B_embedding, B_embedding))
503 between_sim_2 = exp_temp(self.sim(B_embedding, A_embedding))
504 loss_2 = -torch.log(between_sim_2.diag() / (refl_sim_2.sum(1) + between_sim_2.sum(1) - refl_sim_2.diag()))
506 global_item_loss = (loss_1 + loss_2) * 0.5
507 global_item_loss = global_item_loss.mean()
508 return global_item_loss
510 def predict(self, interaction):
511 user = interaction[self.USER_ID]
512 item = interaction[self.ITEM_ID]
514 (
515 item_semantic_emb,
516 user_lightgcn_emb,
517 item_lightgcn_emb,
518 user_gcn_emb,
519 item_gcn_emb,
520 ) = self.forward()
521 item_emb_1 = item_semantic_emb[item]
522 user_emb_1 = user_lightgcn_emb[user]
523 item_emb_2 = item_lightgcn_emb[item]
524 user_emb_2 = user_gcn_emb[user]
525 item_emb_3 = item_gcn_emb[item]
527 user_embedding = torch.cat((user_emb_2, user_emb_1), dim=-1)
528 item_embedding = torch.cat((item_emb_3, item_emb_1 + item_emb_2), dim=-1)
530 scores = torch.mul(user_embedding, item_embedding).sum(dim=1)
531 if self.loss_type.lower() == "bce":
532 scores = self.sigmoid(scores)
533 return scores
535 def full_sort_predict(self, interaction):
536 user = interaction[self.USER_ID]
537 if self.restore_user_e is None or self.restore_entity_e is None:
538 (
539 item_semantic_emb,
540 user_lightgcn_emb,
541 item_lightgcn_emb,
542 user_gcn_emb,
543 entity_gcn_emb,
544 ) = self.forward()
545 self.restore_user_e = torch.cat((user_gcn_emb, user_lightgcn_emb), dim=-1)
546 self.restore_entity_e = torch.cat((entity_gcn_emb, item_semantic_emb + item_lightgcn_emb), dim=-1)
548 u_embeddings = self.restore_user_e[user]
549 i_embeddings = self.restore_entity_e[: self.n_items]
551 scores = torch.matmul(u_embeddings, i_embeddings.transpose(0, 1))
552 if self.loss_type.lower() == "bce":
553 scores = self.sigmoid(scores)
555 return scores.view(-1)