Coverage for hopwise/model/knowledge_aware_recommender/userkgat.py: 14%
180 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 : 2020/9/15
2# @Author : Shanlei Mu
3# @Email : slmu@ruc.edu.cn
5r"""KGAT
6##################################################
7Reference:
8 Xiang Wang et al. "KGAT: Knowledge Graph Attention Network for Recommendation." in SIGKDD 2019.
10Reference code:
11 https://github.com/xiangwang1223/knowledge_graph_attention_network
12"""
14import numpy as np
15import scipy.sparse as sp
16import torch
17import torch.nn.functional as F
18from torch import nn
20from hopwise.model.abstract_recommender import KnowledgeRecommender
21from hopwise.model.init import xavier_normal_initialization
22from hopwise.model.loss import BPRLoss, EmbLoss
23from hopwise.utils import InputType
26class Aggregator(nn.Module):
27 """GNN Aggregator layer"""
29 def __init__(self, input_dim, output_dim, dropout, aggregator_type):
30 super().__init__()
31 self.input_dim = input_dim
32 self.output_dim = output_dim
33 self.dropout = dropout
34 self.aggregator_type = aggregator_type
36 self.message_dropout = nn.Dropout(dropout)
38 if self.aggregator_type == "gcn":
39 self.W = nn.Linear(self.input_dim, self.output_dim)
40 elif self.aggregator_type == "graphsage":
41 self.W = nn.Linear(self.input_dim * 2, self.output_dim)
42 elif self.aggregator_type == "bi":
43 self.W1 = nn.Linear(self.input_dim, self.output_dim)
44 self.W2 = nn.Linear(self.input_dim, self.output_dim)
45 else:
46 raise NotImplementedError
48 self.activation = nn.LeakyReLU()
50 def forward(self, norm_matrix, ego_embeddings):
51 side_embeddings = torch.sparse.mm(norm_matrix, ego_embeddings)
53 if self.aggregator_type == "gcn":
54 ego_embeddings = self.activation(self.W(ego_embeddings + side_embeddings))
55 elif self.aggregator_type == "graphsage":
56 ego_embeddings = self.activation(self.W(torch.cat([ego_embeddings, side_embeddings], dim=1)))
57 elif self.aggregator_type == "bi":
58 add_embeddings = ego_embeddings + side_embeddings
59 sum_embeddings = self.activation(self.W1(add_embeddings))
60 bi_embeddings = torch.mul(ego_embeddings, side_embeddings)
61 bi_embeddings = self.activation(self.W2(bi_embeddings))
62 ego_embeddings = bi_embeddings + sum_embeddings
63 else:
64 raise NotImplementedError
66 ego_embeddings = self.message_dropout(ego_embeddings)
68 return ego_embeddings
71class UserKGAT(KnowledgeRecommender):
72 r"""UserKGAT is a KGAT adapatation to learn from a KG with both user and item realtions, so users be KG entities.
73 KGAT is a knowledge-based recommendation model. It combines knowledge graph and the user-item interaction
74 graph to a new graph called collaborative knowledge graph (CKG). This model learns the representations of users and
75 items by exploiting the structure of CKG. It adopts a GNN-based architecture and define the attention on the CKG.
76 """
78 input_type = InputType.PAIRWISE
80 def __init__(self, config, dataset):
81 super().__init__(config, dataset)
83 # load dataset info
84 ckg_coo = dataset.ckg_graph(form="coo", value_field="relation_id")
85 self.all_hs = torch.LongTensor(ckg_coo.row).to(self.device)
86 self.all_ts = torch.LongTensor(ckg_coo.col).to(self.device)
87 self.all_rs = torch.LongTensor(ckg_coo.data).to(self.device)
88 self.matrix_size = torch.Size([self.n_entities, self.n_entities])
90 # load parameters info
91 self.embedding_size = config["embedding_size"]
92 self.kg_embedding_size = config["kg_embedding_size"]
93 self.layers = [self.embedding_size] + config["layers"]
94 self.aggregator_type = config["aggregator_type"]
95 self.mess_dropout = config["mess_dropout"]
96 self.reg_weight = config["reg_weight"]
98 # generate intermediate data
99 self.A_in = self.init_graph(ckg_coo) # init the attention matrix by the structure of ckg
101 if config["preload_weight"] is not None and config["preload_weight"]["recipeemb_id"] is not None:
102 self.item_entity_embedding_matrix = dataset.get_preload_weight("recipeemb_id")
104 # define layers and loss
105 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size)
106 self.relation_embedding = nn.Embedding(self.n_relations, self.kg_embedding_size)
107 self.trans_w = nn.Embedding(self.n_relations, self.embedding_size * self.kg_embedding_size)
108 self.aggregator_layers = nn.ModuleList()
109 for idx, (input_dim, output_dim) in enumerate(zip(self.layers[:-1], self.layers[1:])):
110 self.aggregator_layers.append(Aggregator(input_dim, output_dim, self.mess_dropout, self.aggregator_type))
111 self.tanh = nn.Tanh()
112 self.mf_loss = BPRLoss()
113 self.reg_loss = EmbLoss()
114 self.restore_entity_e = None
116 # parameters initialization
117 self.apply(xavier_normal_initialization)
118 self.other_parameter_name = ["restore_entity_e"]
120 def init_graph(self, ckg_coo):
121 r"""Get the initial attention matrix through the collaborative knowledge graph
123 Args:
124 ckg_coo (scipy.sparse.coo_matrix): COO adjacency of the CKG whose ``data`` holds
125 the relation id of each edge.
127 Returns:
128 torch.sparse.FloatTensor: Sparse tensor of the attention matrix
129 """
130 node_num = ckg_coo.shape[0]
132 adj_list = []
133 for rel_type in range(1, self.n_relations, 1):
134 rel_mask = ckg_coo.data == rel_type
135 sub_graph = sp.coo_matrix(
136 (np.ones(rel_mask.sum()), (ckg_coo.row[rel_mask], ckg_coo.col[rel_mask])),
137 shape=(node_num, node_num),
138 ).astype("float")
139 rowsum = np.array(sub_graph.sum(1))
140 d_inv = np.power(rowsum, -1).flatten()
141 d_inv[np.isinf(d_inv)] = 0.0
142 d_mat_inv = sp.diags(d_inv)
143 norm_adj = d_mat_inv.dot(sub_graph).tocoo()
144 adj_list.append(norm_adj)
146 final_adj_matrix = sum(adj_list).tocoo()
147 indices = torch.LongTensor([final_adj_matrix.row, final_adj_matrix.col])
148 values = torch.FloatTensor(final_adj_matrix.data)
149 adj_matrix_tensor = torch.sparse.FloatTensor(indices, values, self.matrix_size)
150 return adj_matrix_tensor.to(self.device)
152 def _get_ego_embeddings(self):
153 return self.entity_embedding.weight
155 def forward(self):
156 ego_embeddings = self._get_ego_embeddings()
157 embeddings_list = [ego_embeddings]
158 for aggregator in self.aggregator_layers:
159 ego_embeddings = aggregator(self.A_in, ego_embeddings)
160 norm_embeddings = F.normalize(ego_embeddings, p=2, dim=1)
161 embeddings_list.append(norm_embeddings)
162 kgat_all_embeddings = torch.cat(embeddings_list, dim=1)
163 return kgat_all_embeddings
165 def _get_kg_embedding(self, h, r, pos_t, neg_t):
166 h_e = self.entity_embedding(h).unsqueeze(1)
167 pos_t_e = self.entity_embedding(pos_t).unsqueeze(1)
168 neg_t_e = self.entity_embedding(neg_t).unsqueeze(1)
169 r_e = self.relation_embedding(r)
170 r_trans_w = self.trans_w(r).view(r.size(0), self.embedding_size, self.kg_embedding_size)
172 h_e = torch.bmm(h_e, r_trans_w).squeeze(1)
173 pos_t_e = torch.bmm(pos_t_e, r_trans_w).squeeze(1)
174 neg_t_e = torch.bmm(neg_t_e, r_trans_w).squeeze(1)
176 return h_e, r_e, pos_t_e, neg_t_e
178 def calculate_loss(self, interaction):
179 if self.restore_entity_e is not None:
180 self.restore_entity_e = None
182 # get loss for training rs
183 user = interaction[self.USER_ID]
184 pos_item = interaction[self.ITEM_ID]
185 neg_item = interaction[self.NEG_ITEM_ID]
187 entity_all_embeddings = self.forward()
188 u_embeddings = entity_all_embeddings[user]
189 pos_embeddings = entity_all_embeddings[
190 self.n_users + pos_item
191 ] # reindex since first entity_embeddings are users
192 neg_embeddings = entity_all_embeddings[self.n_users + neg_item]
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 mf_loss = self.mf_loss(pos_scores, neg_scores)
197 reg_loss = self.reg_loss(u_embeddings, pos_embeddings, neg_embeddings)
198 loss = mf_loss + self.reg_weight * reg_loss
200 return loss
202 def calculate_kg_loss(self, interaction):
203 r"""Calculate the training loss for a batch data of KG.
205 Args:
206 interaction (Interaction): Interaction class of the batch.
208 Returns:
209 torch.Tensor: Training loss, shape: []
210 """
212 if self.restore_entity_e is not None:
213 self.restore_entity_e = None
215 # get loss for training kg
216 h = interaction[self.HEAD_ENTITY_ID]
217 r = interaction[self.RELATION_ID]
218 pos_t = interaction[self.TAIL_ENTITY_ID]
219 neg_t = interaction[self.NEG_TAIL_ENTITY_ID]
221 h_e, r_e, pos_t_e, neg_t_e = self._get_kg_embedding(h, r, pos_t, neg_t)
222 pos_tail_score = ((h_e + r_e - pos_t_e) ** 2).sum(dim=1)
223 neg_tail_score = ((h_e + r_e - neg_t_e) ** 2).sum(dim=1)
224 kg_loss = F.softplus(pos_tail_score - neg_tail_score).mean()
225 kg_reg_loss = self.reg_loss(h_e, r_e, pos_t_e, neg_t_e)
226 loss = kg_loss + self.reg_weight * kg_reg_loss
228 return loss
230 def generate_transE_score(self, hs, ts, r):
231 r"""Calculating scores for triples in KG.
233 Args:
234 hs (torch.Tensor): head entities
235 ts (torch.Tensor): tail entities
236 r (int): the relation id between hs and ts
238 Returns:
239 torch.Tensor: the scores of (hs, r, ts)
240 """
242 all_embeddings = self._get_ego_embeddings()
243 h_e = all_embeddings[hs]
244 t_e = all_embeddings[ts]
245 r_e = self.relation_embedding.weight[r]
246 r_trans_w = self.trans_w.weight[r].view(self.embedding_size, self.kg_embedding_size)
248 h_e = torch.matmul(h_e, r_trans_w)
249 t_e = torch.matmul(t_e, r_trans_w)
251 kg_score = torch.mul(t_e, self.tanh(h_e + r_e)).sum(dim=1)
253 return kg_score
255 def update_attentive_A(self):
256 r"""Update the attention matrix using the updated embedding matrix"""
257 kg_score_list, row_list, col_list = [], [], []
258 # To reduce the GPU memory consumption, we calculate the scores of KG triples according to the type of relation
259 for rel_idx in range(1, self.n_relations, 1):
260 triple_index = torch.where(self.all_rs == rel_idx)
261 kg_score = self.generate_transE_score(self.all_hs[triple_index], self.all_ts[triple_index], rel_idx)
262 row_list.append(self.all_hs[triple_index])
263 col_list.append(self.all_ts[triple_index])
264 kg_score_list.append(kg_score)
265 kg_score = torch.cat(kg_score_list, dim=0)
266 row = torch.cat(row_list, dim=0)
267 col = torch.cat(col_list, dim=0)
268 indices = torch.cat([row, col], dim=0).view(2, -1)
269 # Current PyTorch version does not support softmax on SparseCUDA, temporarily move to CPU to calculate softmax
270 A_in = torch.sparse.FloatTensor(indices, kg_score, self.matrix_size).cpu()
271 A_in = torch.sparse.softmax(A_in, dim=1).to(self.device)
272 self.A_in = A_in
274 def predict(self, interaction):
275 user = interaction[self.USER_ID]
276 item = interaction[self.ITEM_ID]
278 entity_all_embeddings = self.forward()
280 u_embeddings = entity_all_embeddings[user]
281 i_embeddings = entity_all_embeddings[self.n_users + item]
282 scores = torch.mul(u_embeddings, i_embeddings).sum(dim=1)
283 return scores
285 def full_sort_predict(self, interaction):
286 user = interaction[self.USER_ID]
287 if self.restore_entity_e is None:
288 self.restore_entity_e = self.forward()
289 u_embeddings = self.restore_entity_e[user]
290 i_embeddings = self.restore_entity_e[self.n_users : self.n_users + self.n_items]
292 scores = torch.matmul(u_embeddings, i_embeddings.transpose(0, 1))
294 return scores.view(-1)