Coverage for hopwise/model/knowledge_aware_recommender/kgrec.py: 93%
305 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"""KGREC
2##################################################
3Reference:
4 Yuhao Yang et al. "Knowledge Graph Self-Supervised Rationalization for Recommendation" in WWW 2021.
5Reference code:
6 https://github.com/HKUDS/KGRec
7"""
9import math
11import numpy as np
12import torch
13import torch.nn.functional as F
14from torch import nn
16from hopwise.model.abstract_recommender import KnowledgeRecommender
17from hopwise.model.init import xavier_uniform_initialization
18from hopwise.model.layers import SparseDropout
19from hopwise.model.loss import BPRLoss, EmbLoss
20from hopwise.utils import InputType
23class Contrast(torch.nn.Module):
24 def __init__(self, num_hidden: int, tau: float = 0.7):
25 super().__init__()
26 self.tau: float = tau
28 self.mlp1 = torch.nn.Sequential(
29 torch.nn.Linear(num_hidden, num_hidden, bias=True),
30 torch.nn.ReLU(),
31 torch.nn.Linear(num_hidden, num_hidden, bias=True),
32 )
33 self.mlp2 = torch.nn.Sequential(
34 torch.nn.Linear(num_hidden, num_hidden, bias=True),
35 torch.nn.ReLU(),
36 torch.nn.Linear(num_hidden, num_hidden, bias=True),
37 )
39 def sim(self, z1: torch.Tensor, z2: torch.Tensor):
40 z1 = F.normalize(z1)
41 z2 = F.normalize(z2)
42 return torch.mm(z1, z2.t())
44 def self_sim(self, z1, z2):
45 z1 = F.normalize(z1)
46 z2 = F.normalize(z2)
47 return (z1 * z2).sum(1)
49 def loss(self, z1: torch.Tensor, z2: torch.Tensor):
50 def f(x):
51 return torch.exp(x / self.tau)
53 between_sim = f(self.self_sim(z1, z2))
54 rand_item = torch.randperm(z1.shape[0])
55 neg_sim = f(self.self_sim(z1, z2[rand_item])) + f(self.self_sim(z2, z1[rand_item]))
57 return -torch.log(between_sim / (between_sim + between_sim + neg_sim))
59 def forward(self, z1: torch.Tensor, z2: torch.Tensor):
60 h1 = self.mlp1(z1)
61 h2 = self.mlp2(z2)
62 loss = self.loss(h1, h2).mean()
63 return loss
66class AttnHGCN(nn.Module):
67 """
68 Heterogeneous Graph Convolutional Network
69 """
71 def __init__(
72 self,
73 embedding_size,
74 n_hops,
75 n_users,
76 n_relations,
77 mess_dropout_rate=0.1,
78 ):
79 super().__init__()
81 self.no_attn_convs = nn.ModuleList()
83 self.embedding_size = embedding_size
84 self.n_hops = n_hops
85 self.n_relations = n_relations
86 self.n_users = n_users
87 self.mess_dropout_rate = mess_dropout_rate
89 # interact relation is ignored
90 self.relation_embedding = nn.Embedding(self.n_relations - 1, self.embedding_size)
91 self.W_Q = nn.Parameter(torch.Tensor(self.embedding_size, self.embedding_size))
93 self.n_heads = 2
94 self.d_k = self.embedding_size // self.n_heads
96 nn.init.xavier_uniform_(self.W_Q)
97 self.mess_dropout = nn.Dropout(p=self.mess_dropout_rate) # mess dropout
99 # parameters initialization
100 self.apply(xavier_uniform_initialization)
102 def shared_layer_agg(self, user_emb, entity_emb, edge_index, edge_type, inter_edge, inter_edge_w):
103 from torch_geometric.utils import scatter
104 from torch_geometric.utils import softmax as scatter_softmax
106 n_entities = entity_emb.shape[0]
107 head, tail = edge_index
109 query = (entity_emb[head] @ self.W_Q).view(-1, self.n_heads, self.d_k)
110 key = (entity_emb[tail] @ self.W_Q).view(-1, self.n_heads, self.d_k)
112 key = key * self.relation_embedding(edge_type).view(-1, self.n_heads, self.d_k)
114 edge_attn_score = (query * key).sum(dim=-1) / math.sqrt(self.d_k)
115 edge_attn_score = scatter_softmax(edge_attn_score, head)
117 neigh_relation_emb = entity_emb[tail] * self.relation_embedding(edge_type) # [-1, embedding_size]
118 value = neigh_relation_emb.view(-1, self.n_heads, self.d_k)
120 entity_agg = value * edge_attn_score.view(-1, self.n_heads, 1)
121 entity_agg = entity_agg.view(-1, self.n_heads * self.d_k)
122 # attn weight makes mean to sum
123 entity_agg = scatter(src=entity_agg, index=head, dim_size=n_entities, dim=0, reduce="sum")
125 item_agg = inter_edge_w.unsqueeze(-1) * entity_emb[inter_edge[1, :]]
126 # w_attn = self.ui_weighting(user_emb, entity_emb, inter_edge)
127 # item_agg += w_attn.unsqueeze(-1) * entity_emb[inter_edge[1, :]]
128 user_agg = scatter(src=item_agg, index=inter_edge[0, :], dim_size=user_emb.shape[0], dim=0, reduce="sum")
129 return entity_agg, user_agg
131 def forward(self, user_emb, entity_emb, edge_index, edge_type, inter_edge, inter_edge_w, item_attn=None):
132 from torch_geometric.utils import scatter
133 from torch_geometric.utils import softmax as scatter_softmax
135 if item_attn is not None:
136 item_attn = item_attn[inter_edge[1, :]]
137 item_attn = scatter_softmax(item_attn, inter_edge[0, :])
138 norm = scatter(
139 torch.ones_like(inter_edge[0, :]), inter_edge[0, :], dim=0, dim_size=user_emb.shape[0], reduce="sum"
140 )
141 norm = torch.index_select(norm, 0, inter_edge[0, :])
142 item_attn = item_attn * norm
143 inter_edge_w = inter_edge_w * item_attn
145 entity_res_emb = entity_emb # [n_entity, embedding_size]
146 user_res_emb = user_emb # [n_users, embedding_size]
147 for i in range(self.n_hops):
148 entity_emb, user_emb = self.shared_layer_agg(
149 user_emb, entity_emb, edge_index, edge_type, inter_edge, inter_edge_w
150 )
152 """message dropout"""
153 if self.mess_dropout_rate > 0.0:
154 entity_emb = self.mess_dropout(entity_emb)
155 user_emb = self.mess_dropout(user_emb)
156 entity_emb = F.normalize(entity_emb)
157 user_emb = F.normalize(user_emb)
159 """result emb"""
160 user_res_emb = torch.add(user_res_emb, user_emb)
161 entity_res_emb = torch.add(entity_res_emb, entity_emb)
163 return user_res_emb, entity_res_emb
165 def forward_ui(self, user_emb, item_emb, inter_edge, inter_edge_w):
166 item_res_emb = item_emb # [n_entity, channel]
167 for i in range(self.n_hops):
168 user_emb, item_emb = self.ui_agg(user_emb, item_emb, inter_edge, inter_edge_w)
169 """message dropout"""
170 if self.mess_dropout_rate > 0.0:
171 item_emb = self.mess_dropout(item_emb)
172 user_emb = self.mess_dropout(user_emb)
173 item_emb = F.normalize(item_emb)
174 user_emb = F.normalize(user_emb)
176 """result emb"""
177 item_res_emb = torch.add(item_res_emb, item_emb)
178 return item_res_emb
180 def forward_kg(self, entity_emb, edge_index, edge_type):
181 entity_res_emb = entity_emb
182 for i in range(self.n_hops):
183 entity_emb = self.kg_agg(entity_emb, edge_index, edge_type)
184 """message dropout"""
185 if self.mess_dropout_rate > 0.0:
186 entity_emb = self.mess_dropout(entity_emb)
187 entity_emb = F.normalize(entity_emb)
189 """result emb"""
190 entity_res_emb = torch.add(entity_res_emb, entity_emb)
191 return entity_res_emb
193 def ui_agg(self, user_emb, item_emb, inter_edge, inter_edge_w):
194 from torch_geometric.utils import scatter
196 num_items = item_emb.shape[0]
197 item_emb = inter_edge_w.unsqueeze(-1) * item_emb[inter_edge[1, :]]
198 user_agg = scatter(src=item_emb, index=inter_edge[0, :], dim_size=user_emb.shape[0], dim=0, reduce="sum")
199 user_emb = inter_edge_w.unsqueeze(-1) * user_emb[inter_edge[0, :]]
200 item_agg = scatter(src=user_emb, index=inter_edge[1, :], dim_size=num_items, dim=0, reduce="sum")
201 return user_agg, item_agg
203 def kg_agg(self, entity_emb, edge_index, edge_type):
204 from torch_geometric.utils import scatter
206 n_entities = entity_emb.shape[0]
207 head, tail = edge_index
208 edge_relation_emb = self.relation_embedding(edge_type)
209 neigh_relation_emb = entity_emb[tail] * edge_relation_emb # [-1, embedding_size]
210 entity_agg = scatter(src=neigh_relation_emb, index=head, dim_size=n_entities, dim=0, reduce="mean")
211 return entity_agg
213 @torch.no_grad()
214 def norm_attn_computer(self, entity_emb, edge_index, edge_type=None, return_logits=False):
215 from torch_geometric.utils import scatter
216 from torch_geometric.utils import softmax as scatter_softmax
218 head, tail = edge_index
220 query = (entity_emb[head] @ self.W_Q).view(-1, self.n_heads, self.d_k)
221 key = (entity_emb[tail] @ self.W_Q).view(-1, self.n_heads, self.d_k)
223 if edge_type is not None:
224 key = key * self.relation_embedding(edge_type).view(-1, self.n_heads, self.d_k)
226 edge_attn = (query * key).sum(dim=-1) / math.sqrt(self.d_k)
227 edge_attn_logits = edge_attn.mean(-1).detach()
228 # softmax by head_node
229 edge_attn_score = scatter_softmax(edge_attn_logits, head)
230 # normalization by head_node degree
231 norm = scatter(torch.ones_like(head), head, dim=0, dim_size=entity_emb.shape[0], reduce="sum")
232 norm = torch.index_select(norm, 0, head)
233 edge_attn_score = edge_attn_score * norm
235 if return_logits:
236 return edge_attn_score, edge_attn_logits
237 return edge_attn_score
240class KGRec(KnowledgeRecommender):
241 r"""KGRec is a self-supervised knowledge-aware recommender that identifies and focuses on informative knowledge
242 graph connections through an attentive rationalization mechanism. It combines generative masking reconstruction
243 and contrastive learning tasks to highlight and align meaningful knowledge and interaction signals. By masking
244 and rebuilding high-rationale edges while filtering noisy ones, KGRec learns more interpretable and noise-resistant
245 recommendations.
246 """
248 input_type = InputType.PAIRWISE
250 def __init__(self, config, dataset):
251 super().__init__(config, dataset)
253 # load parameters info
254 self.embedding_size = config["embedding_size"]
255 self.reg_weight = config["reg_weight"]
256 self.context_hops = config["context_hops"]
257 self.node_dropout_rate = config["node_dropout_rate"]
258 self.mess_dropout_rate = config["mess_dropout_rate"]
260 self.mae_coef = config["mae_coef"]
261 self.mae_msize = config["mae_msize"]
262 self.cl_coef = config["cl_coef"]
263 self.cl_tau = config["cl_tau"]
264 self.cl_drop = config["cl_drop"]
265 self.samp_func = config["samp_func"]
267 self.inter_edge, _ = dataset._create_norm_ckg_adjacency_matrix(symmetric=False)
268 self.inter_edge = self.inter_edge.to(self.device)
269 self.kg_graph = dataset.kg_graph(form="coo", value_field="relation_id") # [n_entities, n_entities]
270 # edge_index: [2, -1]; edge_type: [-1,]
271 self.edge_index, self.edge_type = self.get_edges(self.kg_graph)
273 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
274 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size)
275 self.mf_loss = BPRLoss()
276 self.reg_loss = EmbLoss()
277 self.restore_user_e = None
278 self.restore_entity_e = None
280 self.gcn = AttnHGCN(
281 embedding_size=self.embedding_size,
282 n_hops=self.context_hops,
283 n_users=self.n_users,
284 n_relations=self.n_relations,
285 mess_dropout_rate=self.mess_dropout_rate,
286 )
288 self.contrast_fn = Contrast(self.embedding_size, tau=self.cl_tau)
289 self.node_dropout = SparseDropout(p=self.node_dropout_rate)
291 # parameters initialization
292 self.apply(xavier_uniform_initialization)
294 def get_edges(self, graph):
295 index = torch.LongTensor(np.array([graph.row, graph.col]))
296 type = torch.LongTensor(np.array(graph.data))
297 return index.to(self.device), type.to(self.device)
299 def forward(self):
300 from torch_geometric.utils import scatter
302 user_emb = self.user_embedding.weight
303 entity_emb = self.entity_embedding.weight
305 """node dropout"""
306 # 1. graph sparsification;
307 if self.node_dropout_rate > 0.0:
308 edge_index, edge_type = self.relation_aware_edge_sampling(sampling_rate=self.node_dropout_rate)
309 inter_edge = self.node_dropout(self.inter_edge)
310 else:
311 edge_index, edge_type = self.edge_index, self.edge_type
312 inter_edge = self.inter_edge
313 inter_edge, inter_edge_w = inter_edge._indices(), inter_edge._values()
315 # 2. compute rationale scores;
316 edge_attn_score, _ = self.gcn.norm_attn_computer(entity_emb, edge_index, edge_type, return_logits=True)
318 # for adaptive UI MAE
319 item_attn_mean_1 = scatter(edge_attn_score, edge_index[0], dim=0, dim_size=self.n_entities, reduce="mean")
320 item_attn_mean_1[item_attn_mean_1 == 0.0] = 1.0
321 item_attn_mean_2 = scatter(edge_attn_score, edge_index[1], dim=0, dim_size=self.n_entities, reduce="mean")
322 item_attn_mean_2[item_attn_mean_2 == 0.0] = 1.0
323 item_attn_mean = (0.5 * item_attn_mean_1 + 0.5 * item_attn_mean_2)[: self.n_items]
325 # for adaptive MAE training
326 noise = -torch.log(-torch.log(torch.rand_like(edge_attn_score)))
327 edge_attn_score = edge_attn_score + noise
328 _, topk_attn_edge_id = torch.topk(edge_attn_score, self.mae_msize, sorted=False)
330 enc_edge_index, enc_edge_type, masked_edge_index, masked_edge_type, _ = self.mae_edge_mask_adapt_mixed(
331 edge_index, edge_type, topk_attn_edge_id
332 )
334 # rec task
335 user_gcn_emb, entity_gcn_emb = self.gcn(
336 user_emb, entity_emb, enc_edge_index, enc_edge_type, inter_edge, inter_edge_w
337 )
339 # MAE task with dot-product decoder
340 node_pair_emb = entity_gcn_emb[masked_edge_index.t()]
341 masked_edge_emb = self.gcn.relation_embedding(masked_edge_type)
342 mae_loss = self.create_mae_loss(node_pair_emb, masked_edge_emb)
344 # CL task
345 """adaptive sampling"""
346 cl_kg_edge, cl_kg_type = self.adaptive_kg_drop_cl(edge_index, edge_type, edge_attn_score)
347 cl_ui_edge, cl_ui_w = self.adaptive_ui_drop_cl(item_attn_mean, inter_edge, inter_edge_w)
348 item_agg_ui = self.gcn.forward_ui(user_emb, entity_emb[: self.n_items], cl_ui_edge, cl_ui_w)
349 item_agg_kg = self.gcn.forward_kg(entity_emb, cl_kg_edge, cl_kg_type)[: self.n_items]
350 cl_loss = self.contrast_fn(item_agg_ui, item_agg_kg)
352 # return user embeddings, entity/item embeddings, and edge-level rationale scores
353 return user_gcn_emb, entity_gcn_emb, mae_loss, cl_loss
355 def calculate_loss(self, interaction):
356 r"""Calculate the training loss for a batch data of KG.
358 Args:
359 interaction (Interaction): Interaction class of the batch.
361 Returns:
362 torch.Tensor: Training loss, shape: []
363 """
364 if self.restore_user_e is not None or self.restore_entity_e is not None:
365 self.restore_user_e, self.restore_entity_e = None, None
367 user = interaction[self.USER_ID]
368 pos_item = interaction[self.ITEM_ID]
369 neg_item = interaction[self.NEG_ITEM_ID]
371 user_all_embeddings, entity_all_embeddings, mae_loss, cl_loss = self.forward()
373 u_embeddings = user_all_embeddings[user]
374 pos_embeddings = entity_all_embeddings[pos_item]
375 neg_embeddings = entity_all_embeddings[neg_item]
377 pos_scores = torch.mul(u_embeddings, pos_embeddings).sum(dim=1)
378 neg_scores = torch.mul(u_embeddings, neg_embeddings).sum(dim=1)
380 # the three losses
381 mf_loss = self.mf_loss(pos_scores, neg_scores)
382 reg_loss = self.reg_loss(u_embeddings, pos_embeddings, neg_embeddings, require_pow=True)
383 bpr_loss = mf_loss + self.reg_weight * reg_loss
384 mae_loss = self.mae_coef * mae_loss
385 cl_loss = self.cl_coef * cl_loss
387 total_loss = bpr_loss + mae_loss + cl_loss
388 return total_loss
390 def relation_aware_edge_sampling(self, sampling_rate=0.5):
391 # exclude interaction
392 for i in range(self.n_relations - 1):
393 edge_index_i, edge_type_i = self.edge_sampling(
394 self.edge_index[:, self.edge_type == i],
395 self.edge_type[self.edge_type == i],
396 sampling_rate=sampling_rate,
397 )
398 if i == 0:
399 edge_index_sampled = edge_index_i
400 edge_type_sampled = edge_type_i
401 else:
402 edge_index_sampled = torch.cat([edge_index_sampled, edge_index_i], dim=1)
403 edge_type_sampled = torch.cat([edge_type_sampled, edge_type_i], dim=0)
404 return edge_index_sampled, edge_type_sampled
406 def edge_sampling(self, edge_index, edge_type, sampling_rate=0.5):
407 # edge_index: [2, -1]
408 # edge_type: [-1]
409 n_edges = edge_index.shape[1]
410 random_indices = np.random.choice(n_edges, size=int(n_edges * sampling_rate), replace=False)
411 return edge_index[:, random_indices], edge_type[random_indices]
413 def mae_edge_mask_adapt_mixed(self, edge_index, edge_type, topk_egde_id):
414 # edge_index: [2, -1]
415 # edge_type: [-1]
416 n_edges = edge_index.shape[1]
417 topk_egde_id = topk_egde_id.cpu().numpy()
418 topk_mask = np.zeros(n_edges, dtype=bool)
419 topk_mask[topk_egde_id] = True
420 # add another group of random mask
421 random_indices = np.random.choice(n_edges, size=topk_egde_id.shape[0], replace=False)
422 random_mask = np.zeros(n_edges, dtype=bool)
423 random_mask[random_indices] = True
424 # combine two masks
425 mask = topk_mask | random_mask
427 remain_edge_index = edge_index[:, ~mask]
428 remain_edge_type = edge_type[~mask]
429 masked_edge_index = edge_index[:, mask]
430 masked_edge_type = edge_type[mask]
432 return remain_edge_index, remain_edge_type, masked_edge_index, masked_edge_type, mask
434 def adaptive_kg_drop_cl(self, edge_index, edge_type, edge_attn_score):
435 keep_rate = 1 - self.cl_drop
436 _, least_attn_edge_id = torch.topk(
437 -edge_attn_score, int((1 - keep_rate) * edge_attn_score.shape[0]), sorted=False
438 )
439 cl_kg_mask = torch.ones_like(edge_attn_score).bool()
440 cl_kg_mask[least_attn_edge_id] = False
441 cl_kg_edge = edge_index[:, cl_kg_mask]
442 cl_kg_type = edge_type[cl_kg_mask]
443 return cl_kg_edge, cl_kg_type
445 def adaptive_ui_drop_cl(self, item_attn_mean, inter_edge, inter_edge_w):
446 keep_rate = 1 - self.cl_drop
447 inter_attn_prob = item_attn_mean[inter_edge[1]]
448 # add gumbel noise
449 noise = -torch.log(-torch.log(torch.rand_like(inter_attn_prob)))
450 """ prob based drop """
451 inter_attn_prob = inter_attn_prob + noise
452 inter_attn_prob = F.softmax(inter_attn_prob, dim=0)
454 if self.samp_func == "np":
455 # we observed abnormal behavior of torch.multinomial on mind
456 sampled_edge_idx = np.random.choice(
457 np.arange(inter_edge_w.shape[0]),
458 size=int(keep_rate * inter_edge_w.shape[0]),
459 replace=False,
460 p=inter_attn_prob.cpu().numpy(),
461 )
462 else:
463 sampled_edge_idx = torch.multinomial(
464 inter_attn_prob, int(keep_rate * inter_edge_w.shape[0]), replacement=False
465 )
467 return inter_edge[:, sampled_edge_idx], inter_edge_w[sampled_edge_idx] / keep_rate
469 def create_mae_loss(self, node_pair_emb, masked_edge_emb=None):
470 head_embs, tail_embs = node_pair_emb[:, 0, :], node_pair_emb[:, 1, :]
471 if masked_edge_emb is not None:
472 pos1 = tail_embs * masked_edge_emb
473 else:
474 pos1 = tail_embs
475 # scores = (pos1 - head_embs).sum(dim=1).abs().mean(dim=0)
476 scores = -torch.log(torch.sigmoid(torch.mul(pos1, head_embs).sum(1))).mean()
477 return scores
479 def predict(self, interaction):
480 user = interaction[self.USER_ID]
481 item = interaction[self.ITEM_ID]
483 user_all_embeddings, entity_all_embeddings = self.gcn(
484 self.user_embedding.weight,
485 self.entity_embedding.weight,
486 self.edge_index,
487 self.edge_type,
488 self.inter_edge._indices(),
489 self.inter_edge._values(),
490 )
492 u_embeddings = user_all_embeddings[user]
493 i_embeddings = entity_all_embeddings[item]
494 scores = torch.mul(u_embeddings, i_embeddings).sum(dim=1)
495 return scores
497 def full_sort_predict(self, interaction):
498 user = interaction[self.USER_ID]
499 if self.restore_user_e is None or self.restore_entity_e is None:
500 self.restore_user_e, self.restore_entity_e = self.gcn(
501 self.user_embedding.weight,
502 self.entity_embedding.weight,
503 self.edge_index,
504 self.edge_type,
505 self.inter_edge._indices(),
506 self.inter_edge._values(),
507 )
509 u_embeddings = self.restore_user_e[user]
510 i_embeddings = self.restore_entity_e[: self.n_items]
512 scores = torch.matmul(u_embeddings, i_embeddings.transpose(0, 1))
514 return scores.view(-1)