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

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

77 

78 input_type = InputType.PAIRWISE 

79 

80 def __init__(self, config, dataset): 

81 super().__init__(config, dataset) 

82 

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

89 

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

97 

98 # generate intermediate data 

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

100 

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

103 

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 

115 

116 # parameters initialization 

117 self.apply(xavier_normal_initialization) 

118 self.other_parameter_name = ["restore_entity_e"] 

119 

120 def init_graph(self, ckg_coo): 

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

122 

123 Args: 

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

125 the relation id of each edge. 

126 

127 Returns: 

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

129 """ 

130 node_num = ckg_coo.shape[0] 

131 

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) 

145 

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) 

151 

152 def _get_ego_embeddings(self): 

153 return self.entity_embedding.weight 

154 

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 

164 

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) 

171 

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) 

175 

176 return h_e, r_e, pos_t_e, neg_t_e 

177 

178 def calculate_loss(self, interaction): 

179 if self.restore_entity_e is not None: 

180 self.restore_entity_e = None 

181 

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] 

186 

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] 

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 

212 if self.restore_entity_e is not None: 

213 self.restore_entity_e = None 

214 

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] 

220 

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 

227 

228 return loss 

229 

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

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

232 

233 Args: 

234 hs (torch.Tensor): head entities 

235 ts (torch.Tensor): tail entities 

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

237 

238 Returns: 

239 torch.Tensor: the scores of (hs, r, ts) 

240 """ 

241 

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) 

247 

248 h_e = torch.matmul(h_e, r_trans_w) 

249 t_e = torch.matmul(t_e, r_trans_w) 

250 

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

252 

253 return kg_score 

254 

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 

273 

274 def predict(self, interaction): 

275 user = interaction[self.USER_ID] 

276 item = interaction[self.ITEM_ID] 

277 

278 entity_all_embeddings = self.forward() 

279 

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 

284 

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] 

291 

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

293 

294 return scores.view(-1)