Coverage for hopwise/model/knowledge_aware_recommender/kgrec.py: 93%

305 statements  

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

1r"""KGREC 

2################################################## 

3Reference: 

4 Yuhao Yang et al. "Knowledge Graph Self-Supervised Rationalization for Recommendation" in WWW 2021. 

5Reference code: 

6 https://github.com/HKUDS/KGRec 

7""" 

8 

9import math 

10 

11import numpy as np 

12import torch 

13import torch.nn.functional as F 

14from torch import nn 

15 

16from hopwise.model.abstract_recommender import KnowledgeRecommender 

17from hopwise.model.init import xavier_uniform_initialization 

18from hopwise.model.layers import SparseDropout 

19from hopwise.model.loss import BPRLoss, EmbLoss 

20from hopwise.utils import InputType 

21 

22 

23class Contrast(torch.nn.Module): 

24 def __init__(self, num_hidden: int, tau: float = 0.7): 

25 super().__init__() 

26 self.tau: float = tau 

27 

28 self.mlp1 = torch.nn.Sequential( 

29 torch.nn.Linear(num_hidden, num_hidden, bias=True), 

30 torch.nn.ReLU(), 

31 torch.nn.Linear(num_hidden, num_hidden, bias=True), 

32 ) 

33 self.mlp2 = torch.nn.Sequential( 

34 torch.nn.Linear(num_hidden, num_hidden, bias=True), 

35 torch.nn.ReLU(), 

36 torch.nn.Linear(num_hidden, num_hidden, bias=True), 

37 ) 

38 

39 def sim(self, z1: torch.Tensor, z2: torch.Tensor): 

40 z1 = F.normalize(z1) 

41 z2 = F.normalize(z2) 

42 return torch.mm(z1, z2.t()) 

43 

44 def self_sim(self, z1, z2): 

45 z1 = F.normalize(z1) 

46 z2 = F.normalize(z2) 

47 return (z1 * z2).sum(1) 

48 

49 def loss(self, z1: torch.Tensor, z2: torch.Tensor): 

50 def f(x): 

51 return torch.exp(x / self.tau) 

52 

53 between_sim = f(self.self_sim(z1, z2)) 

54 rand_item = torch.randperm(z1.shape[0]) 

55 neg_sim = f(self.self_sim(z1, z2[rand_item])) + f(self.self_sim(z2, z1[rand_item])) 

56 

57 return -torch.log(between_sim / (between_sim + between_sim + neg_sim)) 

58 

59 def forward(self, z1: torch.Tensor, z2: torch.Tensor): 

60 h1 = self.mlp1(z1) 

61 h2 = self.mlp2(z2) 

62 loss = self.loss(h1, h2).mean() 

63 return loss 

64 

65 

66class AttnHGCN(nn.Module): 

67 """ 

68 Heterogeneous Graph Convolutional Network 

69 """ 

70 

71 def __init__( 

72 self, 

73 embedding_size, 

74 n_hops, 

75 n_users, 

76 n_relations, 

77 mess_dropout_rate=0.1, 

78 ): 

79 super().__init__() 

80 

81 self.no_attn_convs = nn.ModuleList() 

82 

83 self.embedding_size = embedding_size 

84 self.n_hops = n_hops 

85 self.n_relations = n_relations 

86 self.n_users = n_users 

87 self.mess_dropout_rate = mess_dropout_rate 

88 

89 # interact relation is ignored 

90 self.relation_embedding = nn.Embedding(self.n_relations - 1, self.embedding_size) 

91 self.W_Q = nn.Parameter(torch.Tensor(self.embedding_size, self.embedding_size)) 

92 

93 self.n_heads = 2 

94 self.d_k = self.embedding_size // self.n_heads 

95 

96 nn.init.xavier_uniform_(self.W_Q) 

97 self.mess_dropout = nn.Dropout(p=self.mess_dropout_rate) # mess dropout 

98 

99 # parameters initialization 

100 self.apply(xavier_uniform_initialization) 

101 

102 def shared_layer_agg(self, user_emb, entity_emb, edge_index, edge_type, inter_edge, inter_edge_w): 

103 from torch_geometric.utils import scatter 

104 from torch_geometric.utils import softmax as scatter_softmax 

105 

106 n_entities = entity_emb.shape[0] 

107 head, tail = edge_index 

108 

109 query = (entity_emb[head] @ self.W_Q).view(-1, self.n_heads, self.d_k) 

110 key = (entity_emb[tail] @ self.W_Q).view(-1, self.n_heads, self.d_k) 

111 

112 key = key * self.relation_embedding(edge_type).view(-1, self.n_heads, self.d_k) 

113 

114 edge_attn_score = (query * key).sum(dim=-1) / math.sqrt(self.d_k) 

115 edge_attn_score = scatter_softmax(edge_attn_score, head) 

116 

117 neigh_relation_emb = entity_emb[tail] * self.relation_embedding(edge_type) # [-1, embedding_size] 

118 value = neigh_relation_emb.view(-1, self.n_heads, self.d_k) 

119 

120 entity_agg = value * edge_attn_score.view(-1, self.n_heads, 1) 

121 entity_agg = entity_agg.view(-1, self.n_heads * self.d_k) 

122 # attn weight makes mean to sum 

123 entity_agg = scatter(src=entity_agg, index=head, dim_size=n_entities, dim=0, reduce="sum") 

124 

125 item_agg = inter_edge_w.unsqueeze(-1) * entity_emb[inter_edge[1, :]] 

126 # w_attn = self.ui_weighting(user_emb, entity_emb, inter_edge) 

127 # item_agg += w_attn.unsqueeze(-1) * entity_emb[inter_edge[1, :]] 

128 user_agg = scatter(src=item_agg, index=inter_edge[0, :], dim_size=user_emb.shape[0], dim=0, reduce="sum") 

129 return entity_agg, user_agg 

130 

131 def forward(self, user_emb, entity_emb, edge_index, edge_type, inter_edge, inter_edge_w, item_attn=None): 

132 from torch_geometric.utils import scatter 

133 from torch_geometric.utils import softmax as scatter_softmax 

134 

135 if item_attn is not None: 

136 item_attn = item_attn[inter_edge[1, :]] 

137 item_attn = scatter_softmax(item_attn, inter_edge[0, :]) 

138 norm = scatter( 

139 torch.ones_like(inter_edge[0, :]), inter_edge[0, :], dim=0, dim_size=user_emb.shape[0], reduce="sum" 

140 ) 

141 norm = torch.index_select(norm, 0, inter_edge[0, :]) 

142 item_attn = item_attn * norm 

143 inter_edge_w = inter_edge_w * item_attn 

144 

145 entity_res_emb = entity_emb # [n_entity, embedding_size] 

146 user_res_emb = user_emb # [n_users, embedding_size] 

147 for i in range(self.n_hops): 

148 entity_emb, user_emb = self.shared_layer_agg( 

149 user_emb, entity_emb, edge_index, edge_type, inter_edge, inter_edge_w 

150 ) 

151 

152 """message dropout""" 

153 if self.mess_dropout_rate > 0.0: 

154 entity_emb = self.mess_dropout(entity_emb) 

155 user_emb = self.mess_dropout(user_emb) 

156 entity_emb = F.normalize(entity_emb) 

157 user_emb = F.normalize(user_emb) 

158 

159 """result emb""" 

160 user_res_emb = torch.add(user_res_emb, user_emb) 

161 entity_res_emb = torch.add(entity_res_emb, entity_emb) 

162 

163 return user_res_emb, entity_res_emb 

164 

165 def forward_ui(self, user_emb, item_emb, inter_edge, inter_edge_w): 

166 item_res_emb = item_emb # [n_entity, channel] 

167 for i in range(self.n_hops): 

168 user_emb, item_emb = self.ui_agg(user_emb, item_emb, inter_edge, inter_edge_w) 

169 """message dropout""" 

170 if self.mess_dropout_rate > 0.0: 

171 item_emb = self.mess_dropout(item_emb) 

172 user_emb = self.mess_dropout(user_emb) 

173 item_emb = F.normalize(item_emb) 

174 user_emb = F.normalize(user_emb) 

175 

176 """result emb""" 

177 item_res_emb = torch.add(item_res_emb, item_emb) 

178 return item_res_emb 

179 

180 def forward_kg(self, entity_emb, edge_index, edge_type): 

181 entity_res_emb = entity_emb 

182 for i in range(self.n_hops): 

183 entity_emb = self.kg_agg(entity_emb, edge_index, edge_type) 

184 """message dropout""" 

185 if self.mess_dropout_rate > 0.0: 

186 entity_emb = self.mess_dropout(entity_emb) 

187 entity_emb = F.normalize(entity_emb) 

188 

189 """result emb""" 

190 entity_res_emb = torch.add(entity_res_emb, entity_emb) 

191 return entity_res_emb 

192 

193 def ui_agg(self, user_emb, item_emb, inter_edge, inter_edge_w): 

194 from torch_geometric.utils import scatter 

195 

196 num_items = item_emb.shape[0] 

197 item_emb = inter_edge_w.unsqueeze(-1) * item_emb[inter_edge[1, :]] 

198 user_agg = scatter(src=item_emb, index=inter_edge[0, :], dim_size=user_emb.shape[0], dim=0, reduce="sum") 

199 user_emb = inter_edge_w.unsqueeze(-1) * user_emb[inter_edge[0, :]] 

200 item_agg = scatter(src=user_emb, index=inter_edge[1, :], dim_size=num_items, dim=0, reduce="sum") 

201 return user_agg, item_agg 

202 

203 def kg_agg(self, entity_emb, edge_index, edge_type): 

204 from torch_geometric.utils import scatter 

205 

206 n_entities = entity_emb.shape[0] 

207 head, tail = edge_index 

208 edge_relation_emb = self.relation_embedding(edge_type) 

209 neigh_relation_emb = entity_emb[tail] * edge_relation_emb # [-1, embedding_size] 

210 entity_agg = scatter(src=neigh_relation_emb, index=head, dim_size=n_entities, dim=0, reduce="mean") 

211 return entity_agg 

212 

213 @torch.no_grad() 

214 def norm_attn_computer(self, entity_emb, edge_index, edge_type=None, return_logits=False): 

215 from torch_geometric.utils import scatter 

216 from torch_geometric.utils import softmax as scatter_softmax 

217 

218 head, tail = edge_index 

219 

220 query = (entity_emb[head] @ self.W_Q).view(-1, self.n_heads, self.d_k) 

221 key = (entity_emb[tail] @ self.W_Q).view(-1, self.n_heads, self.d_k) 

222 

223 if edge_type is not None: 

224 key = key * self.relation_embedding(edge_type).view(-1, self.n_heads, self.d_k) 

225 

226 edge_attn = (query * key).sum(dim=-1) / math.sqrt(self.d_k) 

227 edge_attn_logits = edge_attn.mean(-1).detach() 

228 # softmax by head_node 

229 edge_attn_score = scatter_softmax(edge_attn_logits, head) 

230 # normalization by head_node degree 

231 norm = scatter(torch.ones_like(head), head, dim=0, dim_size=entity_emb.shape[0], reduce="sum") 

232 norm = torch.index_select(norm, 0, head) 

233 edge_attn_score = edge_attn_score * norm 

234 

235 if return_logits: 

236 return edge_attn_score, edge_attn_logits 

237 return edge_attn_score 

238 

239 

240class KGRec(KnowledgeRecommender): 

241 r"""KGRec is a self-supervised knowledge-aware recommender that identifies and focuses on informative knowledge 

242 graph connections through an attentive rationalization mechanism. It combines generative masking reconstruction 

243 and contrastive learning tasks to highlight and align meaningful knowledge and interaction signals. By masking 

244 and rebuilding high-rationale edges while filtering noisy ones, KGRec learns more interpretable and noise-resistant 

245 recommendations. 

246 """ 

247 

248 input_type = InputType.PAIRWISE 

249 

250 def __init__(self, config, dataset): 

251 super().__init__(config, dataset) 

252 

253 # load parameters info 

254 self.embedding_size = config["embedding_size"] 

255 self.reg_weight = config["reg_weight"] 

256 self.context_hops = config["context_hops"] 

257 self.node_dropout_rate = config["node_dropout_rate"] 

258 self.mess_dropout_rate = config["mess_dropout_rate"] 

259 

260 self.mae_coef = config["mae_coef"] 

261 self.mae_msize = config["mae_msize"] 

262 self.cl_coef = config["cl_coef"] 

263 self.cl_tau = config["cl_tau"] 

264 self.cl_drop = config["cl_drop"] 

265 self.samp_func = config["samp_func"] 

266 

267 self.inter_edge, _ = dataset._create_norm_ckg_adjacency_matrix(symmetric=False) 

268 self.inter_edge = self.inter_edge.to(self.device) 

269 self.kg_graph = dataset.kg_graph(form="coo", value_field="relation_id") # [n_entities, n_entities] 

270 # edge_index: [2, -1]; edge_type: [-1,] 

271 self.edge_index, self.edge_type = self.get_edges(self.kg_graph) 

272 

273 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size) 

274 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size) 

275 self.mf_loss = BPRLoss() 

276 self.reg_loss = EmbLoss() 

277 self.restore_user_e = None 

278 self.restore_entity_e = None 

279 

280 self.gcn = AttnHGCN( 

281 embedding_size=self.embedding_size, 

282 n_hops=self.context_hops, 

283 n_users=self.n_users, 

284 n_relations=self.n_relations, 

285 mess_dropout_rate=self.mess_dropout_rate, 

286 ) 

287 

288 self.contrast_fn = Contrast(self.embedding_size, tau=self.cl_tau) 

289 self.node_dropout = SparseDropout(p=self.node_dropout_rate) 

290 

291 # parameters initialization 

292 self.apply(xavier_uniform_initialization) 

293 

294 def get_edges(self, graph): 

295 index = torch.LongTensor(np.array([graph.row, graph.col])) 

296 type = torch.LongTensor(np.array(graph.data)) 

297 return index.to(self.device), type.to(self.device) 

298 

299 def forward(self): 

300 from torch_geometric.utils import scatter 

301 

302 user_emb = self.user_embedding.weight 

303 entity_emb = self.entity_embedding.weight 

304 

305 """node dropout""" 

306 # 1. graph sparsification; 

307 if self.node_dropout_rate > 0.0: 

308 edge_index, edge_type = self.relation_aware_edge_sampling(sampling_rate=self.node_dropout_rate) 

309 inter_edge = self.node_dropout(self.inter_edge) 

310 else: 

311 edge_index, edge_type = self.edge_index, self.edge_type 

312 inter_edge = self.inter_edge 

313 inter_edge, inter_edge_w = inter_edge._indices(), inter_edge._values() 

314 

315 # 2. compute rationale scores; 

316 edge_attn_score, _ = self.gcn.norm_attn_computer(entity_emb, edge_index, edge_type, return_logits=True) 

317 

318 # for adaptive UI MAE 

319 item_attn_mean_1 = scatter(edge_attn_score, edge_index[0], dim=0, dim_size=self.n_entities, reduce="mean") 

320 item_attn_mean_1[item_attn_mean_1 == 0.0] = 1.0 

321 item_attn_mean_2 = scatter(edge_attn_score, edge_index[1], dim=0, dim_size=self.n_entities, reduce="mean") 

322 item_attn_mean_2[item_attn_mean_2 == 0.0] = 1.0 

323 item_attn_mean = (0.5 * item_attn_mean_1 + 0.5 * item_attn_mean_2)[: self.n_items] 

324 

325 # for adaptive MAE training 

326 noise = -torch.log(-torch.log(torch.rand_like(edge_attn_score))) 

327 edge_attn_score = edge_attn_score + noise 

328 _, topk_attn_edge_id = torch.topk(edge_attn_score, self.mae_msize, sorted=False) 

329 

330 enc_edge_index, enc_edge_type, masked_edge_index, masked_edge_type, _ = self.mae_edge_mask_adapt_mixed( 

331 edge_index, edge_type, topk_attn_edge_id 

332 ) 

333 

334 # rec task 

335 user_gcn_emb, entity_gcn_emb = self.gcn( 

336 user_emb, entity_emb, enc_edge_index, enc_edge_type, inter_edge, inter_edge_w 

337 ) 

338 

339 # MAE task with dot-product decoder 

340 node_pair_emb = entity_gcn_emb[masked_edge_index.t()] 

341 masked_edge_emb = self.gcn.relation_embedding(masked_edge_type) 

342 mae_loss = self.create_mae_loss(node_pair_emb, masked_edge_emb) 

343 

344 # CL task 

345 """adaptive sampling""" 

346 cl_kg_edge, cl_kg_type = self.adaptive_kg_drop_cl(edge_index, edge_type, edge_attn_score) 

347 cl_ui_edge, cl_ui_w = self.adaptive_ui_drop_cl(item_attn_mean, inter_edge, inter_edge_w) 

348 item_agg_ui = self.gcn.forward_ui(user_emb, entity_emb[: self.n_items], cl_ui_edge, cl_ui_w) 

349 item_agg_kg = self.gcn.forward_kg(entity_emb, cl_kg_edge, cl_kg_type)[: self.n_items] 

350 cl_loss = self.contrast_fn(item_agg_ui, item_agg_kg) 

351 

352 # return user embeddings, entity/item embeddings, and edge-level rationale scores 

353 return user_gcn_emb, entity_gcn_emb, mae_loss, cl_loss 

354 

355 def calculate_loss(self, interaction): 

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

357 

358 Args: 

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

360 

361 Returns: 

362 torch.Tensor: Training loss, shape: [] 

363 """ 

364 if self.restore_user_e is not None or self.restore_entity_e is not None: 

365 self.restore_user_e, self.restore_entity_e = None, None 

366 

367 user = interaction[self.USER_ID] 

368 pos_item = interaction[self.ITEM_ID] 

369 neg_item = interaction[self.NEG_ITEM_ID] 

370 

371 user_all_embeddings, entity_all_embeddings, mae_loss, cl_loss = self.forward() 

372 

373 u_embeddings = user_all_embeddings[user] 

374 pos_embeddings = entity_all_embeddings[pos_item] 

375 neg_embeddings = entity_all_embeddings[neg_item] 

376 

377 pos_scores = torch.mul(u_embeddings, pos_embeddings).sum(dim=1) 

378 neg_scores = torch.mul(u_embeddings, neg_embeddings).sum(dim=1) 

379 

380 # the three losses 

381 mf_loss = self.mf_loss(pos_scores, neg_scores) 

382 reg_loss = self.reg_loss(u_embeddings, pos_embeddings, neg_embeddings, require_pow=True) 

383 bpr_loss = mf_loss + self.reg_weight * reg_loss 

384 mae_loss = self.mae_coef * mae_loss 

385 cl_loss = self.cl_coef * cl_loss 

386 

387 total_loss = bpr_loss + mae_loss + cl_loss 

388 return total_loss 

389 

390 def relation_aware_edge_sampling(self, sampling_rate=0.5): 

391 # exclude interaction 

392 for i in range(self.n_relations - 1): 

393 edge_index_i, edge_type_i = self.edge_sampling( 

394 self.edge_index[:, self.edge_type == i], 

395 self.edge_type[self.edge_type == i], 

396 sampling_rate=sampling_rate, 

397 ) 

398 if i == 0: 

399 edge_index_sampled = edge_index_i 

400 edge_type_sampled = edge_type_i 

401 else: 

402 edge_index_sampled = torch.cat([edge_index_sampled, edge_index_i], dim=1) 

403 edge_type_sampled = torch.cat([edge_type_sampled, edge_type_i], dim=0) 

404 return edge_index_sampled, edge_type_sampled 

405 

406 def edge_sampling(self, edge_index, edge_type, sampling_rate=0.5): 

407 # edge_index: [2, -1] 

408 # edge_type: [-1] 

409 n_edges = edge_index.shape[1] 

410 random_indices = np.random.choice(n_edges, size=int(n_edges * sampling_rate), replace=False) 

411 return edge_index[:, random_indices], edge_type[random_indices] 

412 

413 def mae_edge_mask_adapt_mixed(self, edge_index, edge_type, topk_egde_id): 

414 # edge_index: [2, -1] 

415 # edge_type: [-1] 

416 n_edges = edge_index.shape[1] 

417 topk_egde_id = topk_egde_id.cpu().numpy() 

418 topk_mask = np.zeros(n_edges, dtype=bool) 

419 topk_mask[topk_egde_id] = True 

420 # add another group of random mask 

421 random_indices = np.random.choice(n_edges, size=topk_egde_id.shape[0], replace=False) 

422 random_mask = np.zeros(n_edges, dtype=bool) 

423 random_mask[random_indices] = True 

424 # combine two masks 

425 mask = topk_mask | random_mask 

426 

427 remain_edge_index = edge_index[:, ~mask] 

428 remain_edge_type = edge_type[~mask] 

429 masked_edge_index = edge_index[:, mask] 

430 masked_edge_type = edge_type[mask] 

431 

432 return remain_edge_index, remain_edge_type, masked_edge_index, masked_edge_type, mask 

433 

434 def adaptive_kg_drop_cl(self, edge_index, edge_type, edge_attn_score): 

435 keep_rate = 1 - self.cl_drop 

436 _, least_attn_edge_id = torch.topk( 

437 -edge_attn_score, int((1 - keep_rate) * edge_attn_score.shape[0]), sorted=False 

438 ) 

439 cl_kg_mask = torch.ones_like(edge_attn_score).bool() 

440 cl_kg_mask[least_attn_edge_id] = False 

441 cl_kg_edge = edge_index[:, cl_kg_mask] 

442 cl_kg_type = edge_type[cl_kg_mask] 

443 return cl_kg_edge, cl_kg_type 

444 

445 def adaptive_ui_drop_cl(self, item_attn_mean, inter_edge, inter_edge_w): 

446 keep_rate = 1 - self.cl_drop 

447 inter_attn_prob = item_attn_mean[inter_edge[1]] 

448 # add gumbel noise 

449 noise = -torch.log(-torch.log(torch.rand_like(inter_attn_prob))) 

450 """ prob based drop """ 

451 inter_attn_prob = inter_attn_prob + noise 

452 inter_attn_prob = F.softmax(inter_attn_prob, dim=0) 

453 

454 if self.samp_func == "np": 

455 # we observed abnormal behavior of torch.multinomial on mind 

456 sampled_edge_idx = np.random.choice( 

457 np.arange(inter_edge_w.shape[0]), 

458 size=int(keep_rate * inter_edge_w.shape[0]), 

459 replace=False, 

460 p=inter_attn_prob.cpu().numpy(), 

461 ) 

462 else: 

463 sampled_edge_idx = torch.multinomial( 

464 inter_attn_prob, int(keep_rate * inter_edge_w.shape[0]), replacement=False 

465 ) 

466 

467 return inter_edge[:, sampled_edge_idx], inter_edge_w[sampled_edge_idx] / keep_rate 

468 

469 def create_mae_loss(self, node_pair_emb, masked_edge_emb=None): 

470 head_embs, tail_embs = node_pair_emb[:, 0, :], node_pair_emb[:, 1, :] 

471 if masked_edge_emb is not None: 

472 pos1 = tail_embs * masked_edge_emb 

473 else: 

474 pos1 = tail_embs 

475 # scores = (pos1 - head_embs).sum(dim=1).abs().mean(dim=0) 

476 scores = -torch.log(torch.sigmoid(torch.mul(pos1, head_embs).sum(1))).mean() 

477 return scores 

478 

479 def predict(self, interaction): 

480 user = interaction[self.USER_ID] 

481 item = interaction[self.ITEM_ID] 

482 

483 user_all_embeddings, entity_all_embeddings = self.gcn( 

484 self.user_embedding.weight, 

485 self.entity_embedding.weight, 

486 self.edge_index, 

487 self.edge_type, 

488 self.inter_edge._indices(), 

489 self.inter_edge._values(), 

490 ) 

491 

492 u_embeddings = user_all_embeddings[user] 

493 i_embeddings = entity_all_embeddings[item] 

494 scores = torch.mul(u_embeddings, i_embeddings).sum(dim=1) 

495 return scores 

496 

497 def full_sort_predict(self, interaction): 

498 user = interaction[self.USER_ID] 

499 if self.restore_user_e is None or self.restore_entity_e is None: 

500 self.restore_user_e, self.restore_entity_e = self.gcn( 

501 self.user_embedding.weight, 

502 self.entity_embedding.weight, 

503 self.edge_index, 

504 self.edge_type, 

505 self.inter_edge._indices(), 

506 self.inter_edge._values(), 

507 ) 

508 

509 u_embeddings = self.restore_user_e[user] 

510 i_embeddings = self.restore_entity_e[: self.n_items] 

511 

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

513 

514 return scores.view(-1)