Coverage for hopwise/model/knowledge_aware_recommender/kglrr.py: 81%

380 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1import logging 

2import os 

3 

4import numpy as np 

5import torch 

6import torch.nn.functional as F 

7from torch import nn 

8 

9from hopwise.model.abstract_recommender import KnowledgeRecommender 

10from hopwise.utils import InputType 

11 

12 

13class GraphAttentionLayer(nn.Module): 

14 def __init__(self, in_features, out_features, dropout, alpha, concat=True): 

15 super().__init__() 

16 self.dropout = dropout 

17 self.in_features = in_features 

18 self.out_features = out_features 

19 self.alpha = alpha 

20 self.concat = concat 

21 

22 self.W = nn.Parameter(torch.empty(size=(in_features, out_features))) 

23 nn.init.xavier_uniform_(self.W.data, gain=1.414) 

24 self.a = nn.Parameter(torch.empty(size=(2 * out_features, 1))) 

25 nn.init.xavier_uniform_(self.a.data, gain=1.414) 

26 self.fc = nn.Linear(2 * out_features, out_features) 

27 

28 self.leakyrelu = nn.LeakyReLU(self.alpha) 

29 

30 def forward_relation(self, item_embs, entity_embs, relations, adj): 

31 # item_embs: N, dim 

32 # entity_embs: N, e_num, dim 

33 # relations: N, e_num, r_dim 

34 # adj: N, e_num 

35 

36 # N, e_num, dim 

37 Wh = item_embs.unsqueeze(1).expand(entity_embs.size()) 

38 # N, e_num, dim 

39 We = entity_embs 

40 a_input = torch.cat((Wh, We), dim=-1) # (N, e_num, 2*dim) 

41 # N,e,2dim -> N,e,dim 

42 e_input = torch.multiply(self.fc(a_input), relations).sum(-1) # N,e 

43 e = self.leakyrelu(e_input) # (N, e_num) 

44 

45 zero_vec = -9e15 * torch.ones_like(e) 

46 attention = torch.where(adj > 0, e, zero_vec) 

47 attention = F.softmax(attention, dim=1) 

48 attention = F.dropout(attention, self.dropout, training=self.training) # N, e_num 

49 # (N, 1, e_num) * (N, e_num, out_features) -> N, out_features 

50 entity_emb_weighted = torch.bmm(attention.unsqueeze(1), entity_embs).squeeze() 

51 h_prime = entity_emb_weighted + item_embs 

52 

53 if self.concat: 

54 return F.elu(h_prime) 

55 else: 

56 return h_prime 

57 

58 def forward(self, item_embs, entity_embs, adj): 

59 Wh = torch.mm(item_embs, self.W) # h.shape: (N, in_features), Wh.shape: (N, out_features) 

60 We = torch.matmul( 

61 entity_embs, self.W 

62 ) # entity_embs: (N, e_num, in_features), We.shape: (N, e_num, out_features) 

63 a_input = self._prepare_cat(Wh, We) # (N, e_num, 2*out_features) 

64 e = self.leakyrelu(torch.matmul(a_input, self.a).squeeze(2)) # (N, e_num) 

65 

66 zero_vec = -9e15 * torch.ones_like(e) 

67 attention = torch.where(adj > 0, e, zero_vec) 

68 attention = F.softmax(attention, dim=1) 

69 attention = F.dropout(attention, self.dropout, training=self.training) # N, e_num 

70 # (N, 1, e_num) * (N, e_num, out_features) -> N, out_features 

71 entity_emb_weighted = torch.bmm(attention.unsqueeze(1), entity_embs).squeeze() 

72 h_prime = entity_emb_weighted + item_embs 

73 

74 if self.concat: 

75 return F.elu(h_prime) 

76 else: 

77 return h_prime 

78 

79 def _prepare_cat(self, Wh, We): 

80 Wh = Wh.unsqueeze(1).expand(We.size()) # (N, e_num, out_features) 

81 return torch.cat((Wh, We), dim=-1) # (N, e_num, 2*out_features) 

82 

83 

84class KGEncoder(nn.Module): 

85 def __init__(self, config, dataset, kg_dataset): 

86 super().__init__() 

87 

88 self.user_history_matrix = dataset.history_item_matrix()[0].to(config["device"]) 

89 

90 self.maxhis = config["maxhis"] 

91 self.kgcn = config["kgcn"] 

92 self.dropout = config["dropout"] 

93 self.keep_prob = 1 - self.dropout # Added 

94 self.A_split = config["A_split"] 

95 self.device = config["device"] 

96 

97 self.latent_dim = config["latent_dim_rec"] 

98 self.n_layers = config["lightGCN_n_layers"] 

99 self.max_entities_per_user = config["max_entities_per_user"] 

100 self.kg_dataset = kg_dataset 

101 self.gat = GAT(self.latent_dim, self.latent_dim, dropout=0.4, alpha=0.2).train() 

102 

103 self.inter_feat = dataset.inter_feat 

104 self.num_users = dataset.user_num 

105 self.num_items = dataset.item_num 

106 

107 self.__init_weight(dataset) 

108 self.config = config 

109 

110 def __init_weight(self, dataset): 

111 self.entity_count = dataset.entity_num 

112 self.relation_count = dataset.relation_num 

113 

114 self.embedding_user = torch.nn.Embedding(num_embeddings=self.num_users, embedding_dim=self.latent_dim) 

115 # item and kg entity 

116 self.embedding_item = torch.nn.Embedding(num_embeddings=self.num_items, embedding_dim=self.latent_dim) 

117 self.embedding_entity = torch.nn.Embedding(num_embeddings=self.entity_count + 1, embedding_dim=self.latent_dim) 

118 self.embedding_relation = torch.nn.Embedding( 

119 num_embeddings=self.relation_count + 1, embedding_dim=self.latent_dim 

120 ) 

121 # relation weights 

122 self.W_R = nn.Parameter(torch.Tensor(self.relation_count, self.latent_dim, self.latent_dim)) 

123 nn.init.xavier_uniform_(self.W_R, gain=nn.init.calculate_gain("relu")) 

124 

125 nn.init.normal_(self.embedding_user.weight, std=0.1) 

126 nn.init.normal_(self.embedding_item.weight, std=0.1) 

127 nn.init.normal_(self.embedding_entity.weight, std=0.1) 

128 nn.init.normal_(self.embedding_relation.weight, std=0.1) 

129 

130 self.f = nn.Sigmoid() 

131 self.Graph = dataset.norm_adjacency_matrix().coalesce().to(self.device) 

132 self.kg_dict, self.item2relations = self.get_kg_dict(self.num_items) 

133 

134 def get_kg_dict(self, item_num): 

135 i2es = dict() 

136 i2rs = dict() 

137 for item in range(item_num): 

138 rts = self.kg_dataset.get(item, False) 

139 if rts: 

140 tails = list(set([ent for tail_list in rts.values() for ent in tail_list])) 

141 relations = list(rts.keys()) 

142 if len(tails) > self.max_entities_per_user: 

143 i2es[item] = torch.IntTensor(tails).to(self.device)[: self.max_entities_per_user] 

144 i2rs[item] = torch.IntTensor(relations).to(self.device)[: self.max_entities_per_user] 

145 else: 

146 # last embedding pos as padding idx 

147 tails.extend([self.dataset.entity_count] * (self.max_entities_per_user - len(tails))) 

148 relations.extend([self.dataset.relation_count] * (self.max_entities_per_user - len(relations))) 

149 i2es[item] = torch.IntTensor(tails).to(self.device) 

150 i2rs[item] = torch.IntTensor(relations).to(self.device) 

151 else: 

152 i2es[item] = torch.IntTensor([self.num_items] * self.max_entities_per_user).to(self.device) 

153 i2rs[item] = torch.IntTensor([self.relation_count] * self.max_entities_per_user).to(self.device) 

154 return i2es, i2rs 

155 

156 def computer(self): 

157 with torch.no_grad(): 

158 users_emb = self.embedding_user.weight 

159 items_emb = self.cal_item_embedding_from_kg(self.kg_dict) 

160 all_emb = torch.cat([users_emb, items_emb]) 

161 embs = [all_emb] 

162 if self.dropout: 

163 if self.training: 

164 g_droped = self.__dropout(self.keep_prob) 

165 else: 

166 g_droped = self.Graph 

167 else: 

168 g_droped = self.Graph 

169 

170 for layer in range(self.n_layers): 

171 all_emb = torch.sparse.mm(g_droped, all_emb) 

172 embs.append(all_emb) 

173 

174 embs = torch.stack(embs, dim=1) 

175 light_out = torch.mean(embs, dim=1) 

176 users, items = torch.split(light_out, [self.num_users, self.num_items]) 

177 return users, items 

178 

179 def __dropout_x(self, x, keep_prob): 

180 size = x.size() 

181 index = x.indices().t() 

182 values = x.values() 

183 random_index = torch.rand(len(values)) + keep_prob 

184 random_index = random_index.int().bool() 

185 index = index[random_index] 

186 values = values[random_index] / keep_prob 

187 g = torch.sparse_coo_tensor(index.t(), values, size) 

188 return g 

189 

190 def __dropout(self, keep_prob): 

191 if self.A_split: 

192 graph = [] 

193 for g in self.Graph: 

194 graph.append(self.__dropout_x(g, keep_prob)) 

195 else: 

196 graph = self.__dropout_x(self.Graph, keep_prob) 

197 return graph 

198 

199 def cal_item_embedding_from_kg(self, kg: dict): 

200 if kg is None: 

201 kg = self.kg_dict 

202 

203 if self.kgcn == "GAT": 

204 return self.cal_item_embedding_gat(kg) 

205 elif self.kgcn == "RGAT": 

206 return self.cal_item_embedding_rgat(kg) 

207 elif self.kgcn == "MEAN": 

208 raise NotImplementedError("The 'MEAN' option for kgcn is not yet implemented.") 

209 elif self.kgcn == "NO": 

210 return self.embedding_item.weight 

211 

212 def cal_item_embedding_gat(self, kg: dict): 

213 item_embs = self.embedding_item(torch.IntTensor(list(kg.keys())).to(self.device)) # item_num, emb_dim 

214 # item_num, entity_num_each 

215 item_entities = torch.stack(list(kg.values())) 

216 # item_num, entity_num_each, emb_dim 

217 entity_embs = self.embedding_entity(item_entities) 

218 # item_num, entity_num_each 

219 padding_mask = torch.where( 

220 item_entities != self.entity_count, torch.ones_like(item_entities), torch.zeros_like(item_entities) 

221 ).float() 

222 return self.gat(item_embs, entity_embs, padding_mask) 

223 

224 def cal_item_embedding_rgat(self, kg: dict): 

225 item_embs = self.embedding_item(torch.IntTensor(list(kg.keys())).to(self.device)) # item_num, emb_dim 

226 # item_num, entity_num_each 

227 item_entities = torch.stack(list(kg.values())) 

228 item_relations = torch.stack(list(self.item2relations.values())) 

229 # item_num, entity_num_each, emb_dim 

230 entity_embs = self.embedding_entity(item_entities) 

231 relation_embs = self.embedding_relation(item_relations) # item_num, entity_num_each, emb_dim 

232 # w_r = self.W_R[relation_embs] # item_num, entity_num_each, emb_dim, emb_dim 

233 # item_num, entity_num_each 

234 padding_mask = torch.where( 

235 item_entities != self.entity_count, torch.ones_like(item_entities), torch.zeros_like(item_entities) 

236 ).float() 

237 return self.gat.forward_relation(item_embs, entity_embs, relation_embs, padding_mask) 

238 

239 

240class KGLRR(KnowledgeRecommender): 

241 """ 

242 KGLRR: Reinforced logical reasoning over KGs for interpretable recommendation system 

243 """ 

244 

245 input_type = InputType.PAIRWISE 

246 

247 def __init__(self, config, dataset) -> None: 

248 super().__init__(config, dataset) 

249 

250 self.kg_dataset = dataset.ckg_dict_graph() 

251 

252 self.encoder = KGEncoder(config, dataset, self.kg_dataset) 

253 self.latent_dim = config["latent_dim_rec"] 

254 

255 self.r_logic = config["r_logic"] 

256 self.r_length = config["r_length"] 

257 self.layers = config["layers"] 

258 self.sim_scale = config["sim_scale"] 

259 self.loss_sum = config["loss_sum"] 

260 self.l2s_weight = config["l2_loss"] 

261 self.is_explain = config["explain"] 

262 

263 self.num_items = dataset.item_num 

264 

265 self._init_weights() 

266 self.bceloss = nn.BCEWithLogitsLoss() 

267 

268 def _init_weights(self): 

269 self.true = torch.nn.Parameter( 

270 torch.from_numpy(np.random.uniform(0, 1, size=[1, self.latent_dim]).astype(np.float32)), 

271 requires_grad=False, 

272 ) 

273 

274 self.and_layer = torch.nn.Linear(self.latent_dim * 2, self.latent_dim) 

275 for i in range(self.layers): 

276 setattr(self, "and_layer_%d" % i, torch.nn.Linear(self.latent_dim * 2, self.latent_dim * 2)) 

277 

278 self.or_layer = torch.nn.Linear(self.latent_dim * 2, self.latent_dim) 

279 for i in range(self.layers): 

280 setattr(self, "or_layer_%d" % i, torch.nn.Linear(self.latent_dim * 2, self.latent_dim * 2)) 

281 

282 def logic_or(self, vector1, vector2, train=False): 

283 vector1, vector2 = self.uniform_size(vector1, vector2, train) 

284 vector = torch.cat((vector1, vector2), dim=-1) 

285 for i in range(self.layers): 

286 vector = F.relu(getattr(self, "or_layer_%d" % i)(vector)) 

287 vector = self.or_layer(vector) 

288 return vector 

289 

290 def logic_and(self, vector1, vector2, train=False): 

291 vector1, vector2 = self.uniform_size(vector1, vector2, train) 

292 vector = torch.cat((vector1, vector2), dim=-1) 

293 for i in range(self.layers): 

294 vector = F.relu(getattr(self, "and_layer_%d" % i)(vector)) 

295 vector = self.and_layer(vector) 

296 return vector 

297 

298 def logic_regularizer(self, train: bool, check_list: list, constraint, constraint_valid): 

299 # This function calculates the gap between logical expressions and the real world 

300 

301 # length 

302 r_length = constraint.norm(dim=2).sum() 

303 check_list.append(("r_length", r_length)) 

304 

305 # and 

306 r_and_true = 1 - self.similarity(self.logic_and(constraint, self.true, train=train), constraint) 

307 r_and_true = (r_and_true * constraint_valid).sum() 

308 check_list.append(("r_and_true", r_and_true)) 

309 

310 r_and_self = 1 - self.similarity(self.logic_and(constraint, constraint, train=train), constraint) 

311 r_and_self = (r_and_self * constraint_valid).sum() 

312 check_list.append(("r_and_self", r_and_self)) 

313 

314 # or 

315 r_or_true = 1 - self.similarity(self.logic_or(constraint, self.true, train=train), self.true) 

316 r_or_true = (r_or_true * constraint_valid).sum() 

317 check_list.append(("r_or_true", r_or_true)) 

318 

319 r_or_self = 1 - self.similarity(self.logic_or(constraint, constraint, train=train), constraint) 

320 r_or_self = (r_or_self * constraint_valid).sum() 

321 check_list.append(("r_or_self", r_or_self)) 

322 

323 r_loss = r_and_true + r_and_self + r_or_true + r_or_self 

324 

325 if self.r_logic > 0: 

326 r_loss = r_loss * self.r_logic 

327 else: 

328 r_loss = torch.from_numpy(np.array(0.0, dtype=np.float32)).to(self.device) 

329 r_loss.requires_grad = True 

330 

331 r_loss += r_length * self.r_length 

332 check_list.append(("r_loss", r_loss)) 

333 return r_loss 

334 

335 def similarity(self, vector1, vector2, sigmoid=True): 

336 result = F.cosine_similarity(vector1, vector2, dim=-1) 

337 result = result * self.sim_scale 

338 if sigmoid: 

339 return result.sigmoid() 

340 return result 

341 

342 def uniform_size(self, vector1, vector2, train=False): 

343 # Removed vector size normalization 

344 if len(vector1.size()) < len(vector2.size()): 

345 vector1 = vector1.expand_as(vector2) 

346 elif vector2.size() != vector1.size(): 

347 vector2 = vector2.expand_as(vector1) 

348 if train: 

349 r12 = torch.Tensor(vector1.size()[:-1]).to(self.device).uniform_(0, 1).bernoulli() 

350 r12 = r12.unsqueeze(-1) 

351 new_v1 = r12 * vector1 + (-r12 + 1) * vector2 

352 new_v2 = r12 * vector2 + (-r12 + 1) * vector1 

353 return new_v1, new_v2 

354 return vector1, vector2 

355 

356 def predict(self, interaction): 

357 users = interaction[self.USER_ID] 

358 

359 history = self.encoder.user_history_matrix[users, : self.encoder.maxhis] # B * H 

360 item_embed = self.encoder.computer()[1] # item_num * V 

361 

362 his_valid = history.ge(0).float() # B * H 

363 

364 maxlen = int(his_valid.sum(dim=1).max().item()) 

365 

366 elements = item_embed[history] * his_valid.unsqueeze(-1) # B * H * V 

367 

368 tmp_o = None 

369 for i in range(maxlen): 

370 tmp_o_valid = his_valid[:, i].unsqueeze(-1) 

371 if tmp_o is None: 

372 tmp_o = elements[:, i, :] * tmp_o_valid # B * V 

373 else: 

374 # Only perform OR operation if valid; otherwise, if the history is not that long (not valid), 

375 # keep the original content unchanged 

376 tmp_o = self.logic_or(tmp_o, elements[:, i, :]) * tmp_o_valid + tmp_o * (-tmp_o_valid + 1) # B * V 

377 or_vector = tmp_o # B * V 

378 left_valid = his_valid[:, 0].unsqueeze(-1) # B * 1 

379 

380 prediction = [] 

381 for i in range(users.size(0)): 

382 sent_vector = ( 

383 left_valid[i] * self.logic_and(or_vector[i].unsqueeze(0).repeat(self.num_items, 1), item_embed) 

384 + (-left_valid[i] + 1) * item_embed 

385 ) # item_size * V 

386 ithpred = self.similarity(sent_vector, self.true, sigmoid=True) # item_size 

387 prediction.append(ithpred) 

388 

389 prediction = torch.stack(prediction).to(self.device) # [B, item_size] 

390 

391 return prediction 

392 

393 def explain(self, users, history, items): 

394 bs = users.size(0) 

395 _, item_embed = self.encoder.computer() # user_num/item_num * V 

396 

397 his_valid = history.ge(0).float() # B * H 

398 elements = item_embed[history.abs()] * his_valid.unsqueeze(-1) # B * H * V 

399 

400 similarity_rlt = [] 

401 for i in range(bs): 

402 tmp_a_valid = his_valid[i, :].unsqueeze(-1) # H 

403 tmp_item = items[i].unsqueeze(0).expand(elements[i].size(0), -1) # [H, V] 

404 tmp_a = self.logic_and(tmp_item, elements[i]) * tmp_a_valid 

405 similarity_rlt.append(self.similarity(tmp_a, self.true)) 

406 

407 return torch.stack(similarity_rlt).to(self.device) # [H, V] 

408 

409 def full_sort_predict(self, interaction): 

410 r"""Full sort prediction function. 

411 Given users, calculate the scores between users and all candidate items. 

412 

413 Args: 

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

415 

416 Returns: 

417 torch.Tensor: Predicted scores for given users and all candidate items, 

418 shape: [n_batch_users, n_candidate_items] 

419 """ 

420 # The predict function already does what is needed (users vs all items) 

421 prediction = self.predict(interaction) 

422 return prediction 

423 

424 def predict_or_and(self, users, pos, neg, history): 

425 # Store content for checking: logic regularization 

426 # Compute L2 regularization on embeddings 

427 check_list = [] 

428 bs = users.size(0) 

429 users_embed, item_embed = self.encoder.computer() 

430 

431 # Each item in the history is marked as positive, but the latter part of the history may be -1, 

432 # indicating it is not that long 

433 his_valid = history.ge(0).float() # B * H 

434 maxlen = int(his_valid.sum(dim=1).max().item()) 

435 elements = item_embed[history.abs()] * his_valid.unsqueeze(-1) # B * H * V 

436 

437 # For later validation, each vector should satisfy the corresponding constraint in the logical 

438 # expression; 'valid' indicates the validity of the corresponding element in the constraint vector 

439 constraint = [elements.view([bs, -1, self.latent_dim])] # B * H * V 

440 constraint_valid = [his_valid.view([bs, -1])] # B * H 

441 

442 tmp_o = None 

443 for i in range(maxlen): 

444 tmp_o_valid = his_valid[:, i].unsqueeze(-1) 

445 if tmp_o is None: 

446 tmp_o = elements[:, i, :] * tmp_o_valid # B * V 

447 else: 

448 # Only perform OR operation if valid; otherwise, if the history is not that long (not valid), 

449 # keep the original content unchanged 

450 tmp_o = self.logic_or(tmp_o, elements[:, i, :]) * tmp_o_valid + tmp_o * (-tmp_o_valid + 1) # B * V 

451 constraint.append(tmp_o.view([bs, 1, self.latent_dim])) # B * 1 * V 

452 constraint_valid.append(tmp_o_valid) # B * 1 

453 or_vector = tmp_o # B * V 

454 left_valid = his_valid[:, 0].unsqueeze(-1) # B * 1 

455 

456 right_vector_true = item_embed[pos] # B * V 

457 right_vector_false = item_embed[neg] # B * V 

458 

459 constraint.append(right_vector_true.view([bs, 1, self.latent_dim])) # B * 1 * V 

460 constraint_valid.append( 

461 torch.ones((bs, 1)).to(self.device) 

462 ) # B * 1 # Indicates that all items to be judged are valid 

463 constraint.append(right_vector_false.view([bs, 1, self.latent_dim])) # B * 1 * V 

464 constraint_valid.append(torch.ones((bs, 1)).to(self.device)) # B * 1 

465 

466 sent_vector = ( 

467 self.logic_and(or_vector, right_vector_true) * left_valid + (-left_valid + 1) * right_vector_true 

468 ) # B * V 

469 constraint.append(sent_vector.view([bs, 1, self.latent_dim])) # B * 1 * V 

470 constraint_valid.append(left_valid) # B * 1 

471 prediction_true = self.similarity(sent_vector, self.true, sigmoid=False).view([-1]) # B 

472 check_list.append(("prediction_true", prediction_true)) 

473 

474 sent_vector = ( 

475 self.logic_and(or_vector, right_vector_false) * left_valid + (-left_valid + 1) * right_vector_false 

476 ) # B * V 

477 constraint.append(sent_vector.view([bs, 1, self.latent_dim])) # B * 1 * V 

478 constraint_valid.append(left_valid) # B * 1 

479 prediction_false = self.similarity(sent_vector, self.true, sigmoid=False).view([-1]) # B 

480 check_list.append(("prediction_false", prediction_false)) 

481 

482 constraint = torch.cat(tuple(constraint), dim=1) 

483 constraint_valid = torch.cat(tuple(constraint_valid), dim=1) 

484 

485 return prediction_true, prediction_false, check_list, constraint, constraint_valid 

486 

487 def calculate_loss(self, interaction): 

488 """ 

489 Calculates the total loss by combining: 

490 - BCE Loss (rloss) 

491 - Entropy Loss (tloss) 

492 - L2 Loss (l2loss) 

493 """ 

494 # Extraction of tensors from the interaction dictionary 

495 batch_users = interaction[self.USER_ID] 

496 batch_pos = interaction[self.ITEM_ID] 

497 batch_neg = interaction[self.NEG_ITEM_ID] 

498 

499 # Build the history in the same way as in predict 

500 batch_history = self.encoder.user_history_matrix[batch_users, : self.encoder.maxhis] 

501 

502 # Forward of the model with the 3 loss components 

503 rloss, tloss, l2loss = self.forward(False, 0, batch_users, batch_pos, batch_neg, batch_history) 

504 

505 # Combination of the 3 components 

506 total_loss = rloss + tloss + l2loss 

507 

508 return total_loss 

509 

510 def triple_loss(self, TItemScore, FItemScore): 

511 bce_loss = self.bceloss(TItemScore.sigmoid(), torch.ones_like(TItemScore)) + self.bceloss( 

512 FItemScore.sigmoid(), torch.zeros_like(FItemScore) 

513 ) 

514 # Input positive and negative example scores, maximizing the score difference 

515 if self.loss_sum: 

516 loss = torch.sum(F.softplus(-(TItemScore - FItemScore))) 

517 else: 

518 loss = torch.mean(F.softplus(-(TItemScore - FItemScore))) 

519 return (loss + bce_loss) * 0.5 

520 

521 def l2_loss(self, users, pos, neg, history): 

522 users_embed, item_embed = self.encoder.computer() 

523 users_emb = users_embed[users] 

524 pos_emb = item_embed[pos] 

525 neg_emb = item_embed[neg] 

526 his_valid = history.ge(0).float() # B * H 

527 elements = item_embed[history.abs()] * his_valid.unsqueeze(-1) # B * H * V 

528 # L2 regularization loss 

529 reg_loss = ( 

530 (1 / 2) 

531 * (users_emb.norm(2).pow(2) + pos_emb.norm(2).pow(2) + neg_emb.norm(2).pow(2) + elements.norm(2).pow(2)) 

532 / float(len(users)) 

533 ) 

534 if not self.loss_sum: 

535 reg_loss /= users.size(0) 

536 return reg_loss * self.l2s_weight 

537 

538 def check(self, check_list): 

539 """Logs the shape and contents of tensors in the provided check_list. 

540 

541 Each element in check_list is expected to be a tuple where the first item 

542 is a string (label) and the second item is a tensor. For each tuple, this 

543 function converts the tensor to a NumPy array after detaching it from the 

544 computation graph and moving it to CPU, then logs the label, shape and 

545 array contents with a threshold of 20 elements for display. 

546 

547 Args: 

548 check_list (list of tuple): List of (label, tensor) pairs to be logged for inspection. 

549 """ 

550 

551 logging.info(os.linesep) 

552 for t in check_list: 

553 d = np.array(t[1].detach().cpu()) 

554 logging.info(os.linesep.join([t[0] + "\t" + str(d.shape), np.array2string(d, threshold=20)]) + os.linesep) 

555 

556 def forward(self, print_check: bool, return_pred: bool, *args, **kwards): 

557 prediction1, prediction0, check_list, constraint, constraint_valid = self.predict_or_and(*args, **kwards) 

558 rloss = self.logic_regularizer(False, check_list, constraint, constraint_valid) 

559 tloss = self.triple_loss(prediction1, prediction0) 

560 l2loss = self.l2_loss(*args, **kwards) 

561 

562 if print_check: 

563 self.check(check_list) 

564 

565 if return_pred: 

566 return prediction1, rloss + tloss + l2loss 

567 return rloss, tloss, l2loss 

568 

569 

570class GAT(nn.Module): 

571 def __init__(self, nfeat, nhid, dropout, alpha): 

572 """Dense version of GAT.""" 

573 super().__init__() 

574 self.dropout = dropout 

575 

576 self.layer = GraphAttentionLayer(nfeat, nhid, dropout=dropout, alpha=alpha, concat=False) 

577 

578 def forward(self, item_embs, entity_embs, adj): 

579 x = F.dropout(item_embs, self.dropout, training=self.training) 

580 y = F.dropout(entity_embs, self.dropout, training=self.training) 

581 x = self.layer(x, y, adj) 

582 x = F.dropout(x, self.dropout, training=self.training) 

583 return x 

584 

585 def forward_relation(self, item_embs, entity_embs, w_r, adj): 

586 x = F.dropout(item_embs, self.dropout, training=self.training) 

587 y = F.dropout(entity_embs, self.dropout, training=self.training) 

588 x = self.layer.forward_relation(x, y, w_r, adj) 

589 x = F.dropout(x, self.dropout, training=self.training) 

590 return x