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

1# @Time : 2020/9/15 

2# @Author : Shanlei Mu 

3# @Email : slmu@ruc.edu.cn 

4 

5r"""KGAT 

6################################################## 

7Reference: 

8 Xiang Wang et al. "KGAT: Knowledge Graph Attention Network for Recommendation." in SIGKDD 2019. 

9 

10Reference code: 

11 https://github.com/xiangwang1223/knowledge_graph_attention_network 

12""" 

13 

14import numpy as np 

15import scipy.sparse as sp 

16import torch 

17import torch.nn.functional as F 

18from torch import nn 

19 

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 

24 

25 

26class Aggregator(nn.Module): 

27 """GNN Aggregator layer""" 

28 

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 

35 

36 self.message_dropout = nn.Dropout(dropout) 

37 

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 

47 

48 self.activation = nn.LeakyReLU() 

49 

50 def forward(self, norm_matrix, ego_embeddings): 

51 side_embeddings = torch.sparse.mm(norm_matrix, ego_embeddings) 

52 

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 

65 

66 ego_embeddings = self.message_dropout(ego_embeddings) 

67 

68 return ego_embeddings 

69 

70 

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 """ 

76 

77 input_type = InputType.PAIRWISE 

78 

79 def __init__(self, config, dataset): 

80 super().__init__(config, dataset) 

81 

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]) 

88 

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"] 

96 

97 # generate intermediate data 

98 self.A_in = self.init_graph(ckg_coo) # init the attention matrix by the structure of ckg 

99 

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 

113 

114 # parameters initialization 

115 self.apply(xavier_normal_initialization) 

116 self.other_parameter_name = ["restore_user_e", "restore_entity_e"] 

117 

118 def init_graph(self, ckg_coo): 

119 r"""Get the initial attention matrix through the collaborative knowledge graph 

120 

121 Args: 

122 ckg_coo (scipy.sparse.coo_matrix): COO adjacency of the CKG whose ``data`` holds 

123 the relation id of each edge. 

124 

125 Returns: 

126 torch.sparse.FloatTensor: Sparse tensor of the attention matrix 

127 """ 

128 node_num = ckg_coo.shape[0] 

129 

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) 

143 

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) 

149 

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 

155 

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 

166 

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) 

173 

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) 

177 

178 return h_e, r_e, pos_t_e, neg_t_e 

179 

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 

183 

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] 

188 

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] 

193 

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 

199 

200 return loss 

201 

202 def calculate_kg_loss(self, interaction): 

203 r"""Calculate the training loss for a batch data of KG. 

204 

205 Args: 

206 interaction (Interaction): Interaction class of the batch. 

207 

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 

213 

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] 

219 

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 

226 

227 return loss 

228 

229 def generate_transE_score(self, hs, ts, r): 

230 r"""Calculating scores for triples in KG. 

231 

232 Args: 

233 hs (torch.Tensor): head entities 

234 ts (torch.Tensor): tail entities 

235 r (int): the relation id between hs and ts 

236 

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) 

245 

246 h_e = torch.matmul(h_e, r_trans_w) 

247 t_e = torch.matmul(t_e, r_trans_w) 

248 

249 kg_score = torch.mul(t_e, self.tanh(h_e + r_e)).sum(dim=1) 

250 

251 return kg_score 

252 

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 

271 

272 def predict(self, interaction): 

273 user = interaction[self.USER_ID] 

274 item = interaction[self.ITEM_ID] 

275 

276 user_all_embeddings, entity_all_embeddings = self.forward() 

277 

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 

282 

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] 

289 

290 scores = torch.matmul(u_embeddings, i_embeddings.transpose(0, 1)) 

291 

292 return scores.view(-1)