Coverage for hopwise/model/knowledge_aware_recommender/kgat.py: 94%
184 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 KGAT(KnowledgeRecommender):
72 r"""KGAT is a knowledge-based recommendation model. It combines knowledge graph and the user-item interaction
73 graph to a new graph called collaborative knowledge graph (CKG). This model learns the representations of users and
74 items by exploiting the structure of CKG. It adopts a GNN-based architecture and define the attention on the CKG.
75 """
77 input_type = InputType.PAIRWISE
79 def __init__(self, config, dataset):
80 super().__init__(config, dataset)
82 # load dataset info
83 ckg_coo = dataset.ckg_graph(form="coo", value_field="relation_id")
84 self.all_hs = torch.LongTensor(ckg_coo.row).to(self.device)
85 self.all_ts = torch.LongTensor(ckg_coo.col).to(self.device)
86 self.all_rs = torch.LongTensor(ckg_coo.data).to(self.device)
87 self.matrix_size = torch.Size([self.n_users + self.n_entities, self.n_users + self.n_entities])
89 # load parameters info
90 self.embedding_size = config["embedding_size"]
91 self.kg_embedding_size = config["kg_embedding_size"]
92 self.layers = [self.embedding_size] + config["layers"]
93 self.aggregator_type = config["aggregator_type"]
94 self.mess_dropout = config["mess_dropout"]
95 self.reg_weight = config["reg_weight"]
97 # generate intermediate data
98 self.A_in = self.init_graph(ckg_coo) # init the attention matrix by the structure of ckg
100 # define layers and loss
101 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
102 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size)
103 self.relation_embedding = nn.Embedding(self.n_relations, self.kg_embedding_size)
104 self.trans_w = nn.Embedding(self.n_relations, self.embedding_size * self.kg_embedding_size)
105 self.aggregator_layers = nn.ModuleList()
106 for idx, (input_dim, output_dim) in enumerate(zip(self.layers[:-1], self.layers[1:])):
107 self.aggregator_layers.append(Aggregator(input_dim, output_dim, self.mess_dropout, self.aggregator_type))
108 self.tanh = nn.Tanh()
109 self.mf_loss = BPRLoss()
110 self.reg_loss = EmbLoss()
111 self.restore_user_e = None
112 self.restore_entity_e = None
114 # parameters initialization
115 self.apply(xavier_normal_initialization)
116 self.other_parameter_name = ["restore_user_e", "restore_entity_e"]
118 def init_graph(self, ckg_coo):
119 r"""Get the initial attention matrix through the collaborative knowledge graph
121 Args:
122 ckg_coo (scipy.sparse.coo_matrix): COO adjacency of the CKG whose ``data`` holds
123 the relation id of each edge.
125 Returns:
126 torch.sparse.FloatTensor: Sparse tensor of the attention matrix
127 """
128 node_num = ckg_coo.shape[0]
130 adj_list = []
131 for rel_type in range(1, self.n_relations, 1):
132 rel_mask = ckg_coo.data == rel_type
133 sub_graph = sp.coo_matrix(
134 (np.ones(rel_mask.sum()), (ckg_coo.row[rel_mask], ckg_coo.col[rel_mask])),
135 shape=(node_num, node_num),
136 ).astype("float")
137 rowsum = np.array(sub_graph.sum(1))
138 d_inv = np.power(rowsum, -1).flatten()
139 d_inv[np.isinf(d_inv)] = 0.0
140 d_mat_inv = sp.diags(d_inv)
141 norm_adj = d_mat_inv.dot(sub_graph).tocoo()
142 adj_list.append(norm_adj)
144 final_adj_matrix = sum(adj_list).tocoo()
145 indices = torch.LongTensor([final_adj_matrix.row, final_adj_matrix.col])
146 values = torch.FloatTensor(final_adj_matrix.data)
147 adj_matrix_tensor = torch.sparse.FloatTensor(indices, values, self.matrix_size)
148 return adj_matrix_tensor.to(self.device)
150 def _get_ego_embeddings(self):
151 user_embeddings = self.user_embedding.weight
152 entity_embeddings = self.entity_embedding.weight
153 ego_embeddings = torch.cat([user_embeddings, entity_embeddings], dim=0)
154 return ego_embeddings
156 def forward(self):
157 ego_embeddings = self._get_ego_embeddings()
158 embeddings_list = [ego_embeddings]
159 for aggregator in self.aggregator_layers:
160 ego_embeddings = aggregator(self.A_in, ego_embeddings)
161 norm_embeddings = F.normalize(ego_embeddings, p=2, dim=1)
162 embeddings_list.append(norm_embeddings)
163 kgat_all_embeddings = torch.cat(embeddings_list, dim=1)
164 user_all_embeddings, entity_all_embeddings = torch.split(kgat_all_embeddings, [self.n_users, self.n_entities])
165 return user_all_embeddings, entity_all_embeddings
167 def _get_kg_embedding(self, h, r, pos_t, neg_t):
168 h_e = self.entity_embedding(h).unsqueeze(1)
169 pos_t_e = self.entity_embedding(pos_t).unsqueeze(1)
170 neg_t_e = self.entity_embedding(neg_t).unsqueeze(1)
171 r_e = self.relation_embedding(r)
172 r_trans_w = self.trans_w(r).view(r.size(0), self.embedding_size, self.kg_embedding_size)
174 h_e = torch.bmm(h_e, r_trans_w).squeeze(1)
175 pos_t_e = torch.bmm(pos_t_e, r_trans_w).squeeze(1)
176 neg_t_e = torch.bmm(neg_t_e, r_trans_w).squeeze(1)
178 return h_e, r_e, pos_t_e, neg_t_e
180 def calculate_loss(self, interaction):
181 if self.restore_user_e is not None or self.restore_entity_e is not None:
182 self.restore_user_e, self.restore_entity_e = None, None
184 # get loss for training rs
185 user = interaction[self.USER_ID]
186 pos_item = interaction[self.ITEM_ID]
187 neg_item = interaction[self.NEG_ITEM_ID]
189 user_all_embeddings, entity_all_embeddings = self.forward()
190 u_embeddings = user_all_embeddings[user]
191 pos_embeddings = entity_all_embeddings[pos_item]
192 neg_embeddings = entity_all_embeddings[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 """
211 if self.restore_user_e is not None or self.restore_entity_e is not None:
212 self.restore_user_e, self.restore_entity_e = None, None
214 # get loss for training kg
215 h = interaction[self.HEAD_ENTITY_ID]
216 r = interaction[self.RELATION_ID]
217 pos_t = interaction[self.TAIL_ENTITY_ID]
218 neg_t = interaction[self.NEG_TAIL_ENTITY_ID]
220 h_e, r_e, pos_t_e, neg_t_e = self._get_kg_embedding(h, r, pos_t, neg_t)
221 pos_tail_score = ((h_e + r_e - pos_t_e) ** 2).sum(dim=1)
222 neg_tail_score = ((h_e + r_e - neg_t_e) ** 2).sum(dim=1)
223 kg_loss = F.softplus(pos_tail_score - neg_tail_score).mean()
224 kg_reg_loss = self.reg_loss(h_e, r_e, pos_t_e, neg_t_e)
225 loss = kg_loss + self.reg_weight * kg_reg_loss
227 return loss
229 def generate_transE_score(self, hs, ts, r):
230 r"""Calculating scores for triples in KG.
232 Args:
233 hs (torch.Tensor): head entities
234 ts (torch.Tensor): tail entities
235 r (int): the relation id between hs and ts
237 Returns:
238 torch.Tensor: the scores of (hs, r, ts)
239 """
240 all_embeddings = self._get_ego_embeddings()
241 h_e = all_embeddings[hs]
242 t_e = all_embeddings[ts]
243 r_e = self.relation_embedding.weight[r]
244 r_trans_w = self.trans_w.weight[r].view(self.embedding_size, self.kg_embedding_size)
246 h_e = torch.matmul(h_e, r_trans_w)
247 t_e = torch.matmul(t_e, r_trans_w)
249 kg_score = torch.mul(t_e, self.tanh(h_e + r_e)).sum(dim=1)
251 return kg_score
253 def update_attentive_A(self):
254 r"""Update the attention matrix using the updated embedding matrix"""
255 kg_score_list, row_list, col_list = [], [], []
256 # To reduce the GPU memory consumption, we calculate the scores of KG triples according to the type of relation
257 for rel_idx in range(1, self.n_relations, 1):
258 triple_index = torch.where(self.all_rs == rel_idx)
259 kg_score = self.generate_transE_score(self.all_hs[triple_index], self.all_ts[triple_index], rel_idx)
260 row_list.append(self.all_hs[triple_index])
261 col_list.append(self.all_ts[triple_index])
262 kg_score_list.append(kg_score)
263 kg_score = torch.cat(kg_score_list, dim=0)
264 row = torch.cat(row_list, dim=0)
265 col = torch.cat(col_list, dim=0)
266 indices = torch.cat([row, col], dim=0).view(2, -1)
267 # Current PyTorch version does not support softmax on SparseCUDA, temporarily move to CPU to calculate softmax
268 A_in = torch.sparse.FloatTensor(indices, kg_score, self.matrix_size).cpu()
269 A_in = torch.sparse.softmax(A_in, dim=1).to(self.device)
270 self.A_in = A_in
272 def predict(self, interaction):
273 user = interaction[self.USER_ID]
274 item = interaction[self.ITEM_ID]
276 user_all_embeddings, entity_all_embeddings = self.forward()
278 u_embeddings = user_all_embeddings[user]
279 i_embeddings = entity_all_embeddings[item]
280 scores = torch.mul(u_embeddings, i_embeddings).sum(dim=1)
281 return scores
283 def full_sort_predict(self, interaction):
284 user = interaction[self.USER_ID]
285 if self.restore_user_e is None or self.restore_entity_e is None:
286 self.restore_user_e, self.restore_entity_e = self.forward()
287 u_embeddings = self.restore_user_e[user]
288 i_embeddings = self.restore_entity_e[: self.n_items]
290 scores = torch.matmul(u_embeddings, i_embeddings.transpose(0, 1))
292 return scores.view(-1)