Coverage for hopwise/model/knowledge_aware_recommender/mcclk.py: 90%

300 statements  

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

1# @Time : 2022/8/22 

2# @Author : Bowen Zheng 

3# @Email : 18735382001@163.com 

4 

5r"""MCCLK 

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

7Reference: 

8 Ding Zou et al. "Multi-level Cross-view Contrastive Learning for Knowledge-aware Recommender System." in SIGIR 2022. 

9 

10Reference code: 

11 https://github.com/CCIIPLab/MCCLK 

12""" # noqa: E501 

13 

14import numpy as np 

15import torch 

16import torch.nn.functional as F 

17from torch import nn 

18 

19from hopwise.model.abstract_recommender import KnowledgeRecommender 

20from hopwise.model.init import xavier_normal_initialization 

21from hopwise.model.layers import SparseDropout 

22from hopwise.model.loss import BPRLoss, EmbLoss 

23from hopwise.utils import InputType 

24 

25 

26class Aggregator(nn.Module): 

27 def __init__(self, item_only=False, attention=True): 

28 super().__init__() 

29 

30 # Only aggregate item embedding 

31 self.item_only = item_only 

32 # Whether use attention mechanism 

33 self.attention = attention 

34 

35 def forward(self, entity_emb, user_emb, relation_emb, edge_index, edge_type, inter_matrix): 

36 from torch_geometric.utils import scatter 

37 from torch_geometric.utils import softmax as scatter_softmax 

38 

39 n_entities = entity_emb.shape[0] 

40 

41 # KG aggregate 

42 head, tail = edge_index 

43 edge_relation_emb = relation_emb[edge_type] 

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

45 

46 if self.attention: 

47 # Calculate attention weights 

48 neigh_relation_emb_weight = self.calculate_sim_hrt(entity_emb[head], entity_emb[tail], edge_relation_emb) 

49 # [-1, 1] -> [-1, embedding_size] 

50 neigh_relation_emb_weight = neigh_relation_emb_weight.expand( 

51 neigh_relation_emb.shape[0], neigh_relation_emb.shape[1] 

52 ) 

53 neigh_relation_emb_weight = scatter_softmax( 

54 neigh_relation_emb_weight, index=head, dim=0 

55 ) # [-1, embedding_size] 

56 neigh_relation_emb = torch.mul(neigh_relation_emb_weight, neigh_relation_emb) 

57 

58 entity_agg = scatter( 

59 src=neigh_relation_emb, index=head, dim_size=n_entities, dim=0, reduce="mean" 

60 ) # [n_entities, embedding_size] 

61 

62 # Only aggregate item embedding 

63 if self.item_only: 

64 return entity_agg 

65 

66 user_agg = torch.sparse.mm(inter_matrix, entity_emb) # [n_users, embedding_size] 

67 # The importance of relation to user 

68 score = torch.mm(user_emb, relation_emb.t()) # [n_users, n_relations] 

69 score = torch.softmax(score, dim=-1) 

70 user_agg = user_agg + (torch.mm(score, relation_emb)) * user_agg 

71 

72 return entity_agg, user_agg 

73 

74 def calculate_sim_hrt(self, entity_emb_head, entity_emb_tail, relation_emb): 

75 r"""The calculation method of attention weight here follows the code implementation of the author, which is 

76 slightly different from that described in the paper. 

77 """ 

78 tail_relation_emb = entity_emb_tail * relation_emb 

79 tail_relation_emb = tail_relation_emb.norm(dim=1, p=2, keepdim=True) 

80 head_relation_emb = entity_emb_head * relation_emb 

81 head_relation_emb = head_relation_emb.norm(dim=1, p=2, keepdim=True) 

82 # [-1, 1, embedding_size] * [-1, embedding_size, 1] -> [-1, 1] 

83 att_weights = torch.matmul(head_relation_emb.unsqueeze(dim=1), tail_relation_emb.unsqueeze(dim=2)).squeeze( 

84 dim=-1 

85 ) 

86 att_weights = att_weights**2 

87 return att_weights 

88 

89 

90class GraphConv(nn.Module): 

91 """Graph Convolutional Network""" 

92 

93 def __init__( 

94 self, 

95 config, 

96 embedding_size, 

97 n_relations, 

98 edge_index, 

99 edge_type, 

100 inter_matrix, 

101 device, 

102 ): 

103 super().__init__() 

104 

105 # load parameters info 

106 self.n_relations = n_relations 

107 self.edge_index = edge_index 

108 self.edge_type = edge_type 

109 self.inter_matrix = inter_matrix 

110 self.embedding_size = embedding_size 

111 self.n_hops = config["n_hops"] 

112 self.node_dropout_rate = config["node_dropout_rate"] 

113 self.mess_dropout_rate = config["mess_dropout_rate"] 

114 self.topk = config["k"] 

115 self.lambda_coeff = config["lambda_coeff"] 

116 self.build_graph_separately = config["build_graph_separately"] 

117 self.device = device 

118 

119 # define layers 

120 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size) 

121 

122 # User a separate GCN to build item-item graph 

123 if self.build_graph_separately: 

124 r""" 

125 In the original author's implementation(https://github.com/CCIIPLab/MCCLK), the process of constructing 

126 k-Nearest-Neighbor item-item semantic graph(section 4.1 in paper) and encoding structural view(section 4.3.1 in paper) 

127 are combined. This implementation improves the computational efficiency, but is slightly different from the 

128 model structure described in the paper. We use the parameter `build_graph_separately` to control whether to 

129 use a separate GCN to build a item-item semantic graph. If `build_graph_separately` is set to true, the model 

130 structure will be the same as that described in the paper. Otherwise, the author's code implementation will be followed. 

131 """ # noqa: E501 

132 self.bg_convs = nn.ModuleList() 

133 for i in range(self.n_hops): 

134 self.bg_convs.append(Aggregator(item_only=True, attention=False)) 

135 

136 self.convs = nn.ModuleList() 

137 for i in range(self.n_hops): 

138 self.convs.append(Aggregator()) 

139 

140 self.node_dropout = SparseDropout(p=self.mess_dropout_rate) # node dropout 

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

142 

143 # parameters initialization 

144 self.apply(xavier_normal_initialization) 

145 

146 def edge_sampling(self, edge_index, edge_type, rate=0.5): 

147 # edge_index: [2, -1] 

148 # edge_type: [-1] 

149 n_edges = edge_index.shape[1] 

150 random_indices = np.random.choice(n_edges, size=int(n_edges * rate), replace=False) 

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

152 

153 def forward(self, user_emb, entity_emb): 

154 # node dropout 

155 if self.node_dropout_rate > 0.0: 

156 edge_index, edge_type = self.edge_sampling(self.edge_index, self.edge_type, self.node_dropout_rate) 

157 inter_matrix = self.node_dropout(self.inter_matrix) 

158 else: 

159 edge_index, edge_type = self.edge_index, self.edge_type 

160 inter_matrix = self.inter_matrix 

161 

162 origin_entity_emb = entity_emb 

163 

164 entity_res_emb = [entity_emb] # [n_entities, embedding_size] 

165 user_res_emb = [user_emb] # [n_users, embedding_size] 

166 relation_emb = self.relation_embedding.weight # [n_relations, embedding_size] 

167 for i in range(len(self.convs)): 

168 entity_emb, user_emb = self.convs[i]( 

169 entity_emb, user_emb, relation_emb, edge_index, edge_type, inter_matrix 

170 ) 

171 # message dropout 

172 if self.mess_dropout_rate > 0.0: 

173 entity_emb = self.mess_dropout(entity_emb) 

174 user_emb = self.mess_dropout(user_emb) 

175 entity_emb = F.normalize(entity_emb) 

176 user_emb = F.normalize(user_emb) 

177 # result embedding 

178 entity_res_emb.append(entity_emb) 

179 user_res_emb.append(user_emb) 

180 

181 entity_res_emb = torch.stack(entity_res_emb, dim=1) 

182 entity_res_emb = entity_res_emb.mean(dim=1, keepdim=False) 

183 user_res_emb = torch.stack(user_res_emb, dim=1) 

184 user_res_emb = user_res_emb.mean(dim=1, keepdim=False) 

185 

186 # build item-item graph 

187 if self.build_graph_separately: 

188 item_adj = self._build_graph_separately(origin_entity_emb) 

189 else: 

190 # build origin item-item graph 

191 origin_item_adj = self.build_adj(origin_entity_emb, self.topk) 

192 # update item-item graph 

193 item_adj = (1 - self.lambda_coeff) * self.build_adj( 

194 entity_res_emb, self.topk 

195 ) + self.lambda_coeff * origin_item_adj 

196 

197 return entity_res_emb, user_res_emb, item_adj 

198 

199 def build_adj(self, context, topk): 

200 r"""Construct a k-Nearest-Neighbor item-item semantic graph. 

201 

202 Returns: 

203 Sparse tensor of the normalized item-item matrix. 

204 """ 

205 # construct similarity adj matrix 

206 n_entities = context.shape[0] 

207 context_norm = context.div(torch.norm(context, p=2, dim=-1, keepdim=True)).cpu() 

208 sim = torch.mm(context_norm, context_norm.transpose(1, 0)) 

209 # knn_val: [n_entities, topk] knn_index: [n_entities, topk] 

210 knn_val, knn_index = torch.topk(sim, topk, dim=-1) 

211 knn_val, knn_index = knn_val.to(self.device), knn_index.to(self.device) 

212 

213 y = knn_index.reshape(-1) 

214 x = torch.arange(0, n_entities).unsqueeze(dim=-1).to(self.device) # [n_entities, 1] 

215 x = x.expand(n_entities, topk).reshape(-1) 

216 indice = torch.cat((x.unsqueeze(dim=0), y.unsqueeze(dim=0)), dim=0) # [2, n_entities * topk] 

217 value = knn_val.reshape(-1) 

218 adj_sparsity = torch.sparse.FloatTensor(indice.data, value.data, torch.Size([n_entities, n_entities])).to( 

219 self.device 

220 ) 

221 

222 # normalized laplacian adj 

223 rowsum = torch.sparse.sum(adj_sparsity, dim=1) 

224 d_inv_sqrt = torch.pow(rowsum, -0.5) 

225 d_mat_inv_sqrt_value = d_inv_sqrt._values() 

226 x = torch.arange(0, n_entities).unsqueeze(dim=0).to(self.device) 

227 x = x.expand(2, n_entities) 

228 d_mat_inv_sqrt_indice = x 

229 d_mat_inv_sqrt = torch.sparse.FloatTensor( 

230 d_mat_inv_sqrt_indice, 

231 d_mat_inv_sqrt_value, 

232 torch.Size([n_entities, n_entities]), 

233 ) 

234 L_norm = torch.sparse.mm(torch.sparse.mm(d_mat_inv_sqrt, adj_sparsity), d_mat_inv_sqrt) 

235 return L_norm 

236 

237 def _build_graph_separately(self, entity_emb): 

238 # node dropout 

239 if self.node_dropout_rate > 0.0: 

240 edge_index, edge_type = self.edge_sampling(self.edge_index, self.edge_type, self.node_dropout_rate) 

241 inter_matrix = self.node_dropout(self.inter_matrix) 

242 else: 

243 edge_index, edge_type = self.edge_index, self.edge_type 

244 inter_matrix = self.inter_matrix 

245 

246 origin_item_adj = self.build_adj(entity_emb, self.topk) 

247 

248 entity_res_emb = [entity_emb] # [n_entities, embedding_size] 

249 relation_emb = self.relation_embedding.weight # [n_relations, embedding_size] 

250 for i in range(len(self.bg_convs)): 

251 entity_emb = self.bg_convs[i](entity_emb, None, relation_emb, edge_index, edge_type, inter_matrix) 

252 # message dropout 

253 if self.mess_dropout_rate > 0.0: 

254 entity_emb = self.mess_dropout(entity_emb) 

255 entity_emb = F.normalize(entity_emb) 

256 # result embedding 

257 entity_res_emb.append(entity_emb) 

258 

259 entity_res_emb = torch.stack(entity_res_emb, dim=1) 

260 entity_res_emb = entity_res_emb.mean(dim=1, keepdim=False) 

261 

262 item_adj = (1 - self.lambda_coeff) * self.build_adj( 

263 entity_res_emb, self.topk 

264 ) + self.lambda_coeff * origin_item_adj 

265 

266 return item_adj 

267 

268 

269class MCCLK(KnowledgeRecommender): 

270 r"""MCCLK is a knowledge-based recommendation model. 

271 It focuses on the contrastive learning in KG-aware recommendation and proposes a novel multi-level cross-view 

272 contrastive learning mechanism. This model comprehensively considers three different graph views for KG-aware 

273 recommendation, including global-level structural view, local-level collaborative and semantic views. It hence 

274 performs contrastive learning across three views on both local and global levels, mining comprehensive graph 

275 feature and structure information in a self-supervised manner. 

276 """ 

277 

278 input_type = InputType.PAIRWISE 

279 

280 def __init__(self, config, dataset): 

281 super().__init__(config, dataset) 

282 

283 # load parameters info 

284 self.embedding_size = config["embedding_size"] 

285 self.reg_weight = config["reg_weight"] 

286 self.lightgcn_layer = config["lightgcn_layer"] 

287 self.item_agg_layer = config["item_agg_layer"] 

288 self.temperature = config["temperature"] 

289 self.alpha = config["alpha"] 

290 self.beta = config["beta"] 

291 self.loss_type = config["loss_type"] 

292 

293 # load dataset info 

294 # inter_matrix: [n_users, n_entities]; inter_graph: [n_users + n_entities, n_users + n_entities] 

295 self.inter_matrix, self.inter_graph = dataset._create_norm_ckg_adjacency_matrix(symmetric=False) 

296 self.inter_matrix = self.inter_matrix.to(self.device) 

297 self.inter_graph = self.inter_graph.to(self.device) 

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

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

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

301 

302 # define layers 

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

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

305 self.gcn = GraphConv( 

306 config=config, 

307 embedding_size=self.embedding_size, 

308 n_relations=self.n_relations, 

309 edge_index=self.edge_index, 

310 edge_type=self.edge_type, 

311 inter_matrix=self.inter_matrix, 

312 device=self.device, 

313 ) 

314 self.fc1 = nn.Sequential( 

315 nn.Linear(self.embedding_size, self.embedding_size, bias=True), 

316 nn.ReLU(), 

317 nn.Linear(self.embedding_size, self.embedding_size, bias=True), 

318 ) 

319 self.fc2 = nn.Sequential( 

320 nn.Linear(self.embedding_size, self.embedding_size, bias=True), 

321 nn.ReLU(), 

322 nn.Linear(self.embedding_size, self.embedding_size, bias=True), 

323 ) 

324 self.fc3 = nn.Sequential( 

325 nn.Linear(self.embedding_size, self.embedding_size, bias=True), 

326 nn.ReLU(), 

327 nn.Linear(self.embedding_size, self.embedding_size, bias=True), 

328 ) 

329 # define loss 

330 if self.loss_type.lower() == "bpr": 

331 self.rec_loss = BPRLoss() 

332 elif self.loss_type.lower() == "bce": 

333 self.sigmoid = nn.Sigmoid() 

334 self.rec_loss = nn.BCEWithLogitsLoss() 

335 else: 

336 raise NotImplementedError(f"The loss type [{self.loss_type}] has not been supported.") 

337 self.reg_loss = EmbLoss() 

338 

339 # storage variables for full sort evaluation acceleration 

340 self.restore_user_e = None 

341 self.restore_item_e = None 

342 

343 # parameters initialization 

344 self.apply(xavier_normal_initialization) 

345 

346 def get_edges(self, graph): 

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

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

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

350 

351 def forward(self): 

352 user_emb = self.user_embedding.weight 

353 entity_emb = self.entity_embedding.weight 

354 # Construct a k-Nearest-Neighbor item-item semantic graph and Structural View Encoder 

355 entity_gcn_emb, user_gcn_emb, item_adj = self.gcn(user_emb, entity_emb) 

356 # Semantic View Encoder 

357 item_semantic_emb = [entity_emb] 

358 item_agg_emb = entity_emb 

359 for i in range(self.item_agg_layer): 

360 item_agg_emb = torch.sparse.mm(item_adj, item_agg_emb) 

361 item_semantic_emb.append(item_agg_emb) 

362 item_semantic_emb = torch.stack(item_semantic_emb, dim=1) 

363 item_semantic_emb = item_semantic_emb.mean(dim=1, keepdim=False) 

364 # item_semantic_emb = F.normalize(item_semantic_emb, p=2, dim=1) 

365 

366 # Collaborative View Encoder 

367 user_lightgcn_emb, item_lightgcn_emb = self.light_gcn(user_emb, entity_emb, self.inter_graph) 

368 

369 return ( 

370 item_semantic_emb, 

371 user_lightgcn_emb, 

372 item_lightgcn_emb, 

373 user_gcn_emb, 

374 entity_gcn_emb, 

375 ) 

376 

377 def light_gcn(self, user_embedding, item_embedding, adj): 

378 ego_embeddings = torch.cat((user_embedding, item_embedding), dim=0) 

379 all_embeddings = [ego_embeddings] 

380 for i in range(self.lightgcn_layer): 

381 side_embeddings = torch.sparse.mm(adj, ego_embeddings) 

382 ego_embeddings = side_embeddings 

383 all_embeddings += [ego_embeddings] 

384 all_embeddings = torch.stack(all_embeddings, dim=1) 

385 all_embeddings = all_embeddings.mean(dim=1, keepdim=False) 

386 u_g_embeddings, i_g_embeddings = torch.split(all_embeddings, [self.n_users, self.n_entities], dim=0) 

387 return u_g_embeddings, i_g_embeddings 

388 

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

390 z1 = F.normalize(z1) 

391 z2 = F.normalize(z2) 

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

393 

394 def calculate_loss(self, interaction): 

395 if self.restore_user_e is not None or self.restore_item_e is not None: 

396 self.restore_user_e, self.restore_item_e = None, None 

397 

398 # get loss for training rs 

399 user = interaction[self.USER_ID] 

400 pos_item = interaction[self.ITEM_ID] 

401 neg_item = interaction[self.NEG_ITEM_ID] 

402 all_item = torch.cat((pos_item, neg_item), dim=0) 

403 

404 ( 

405 item_semantic_emb, 

406 user_lightgcn_emb, 

407 item_lightgcn_emb, 

408 user_gcn_emb, 

409 item_gcn_emb, 

410 ) = self.forward() 

411 item_emb_1 = item_semantic_emb[all_item] 

412 user_emb_1 = user_lightgcn_emb[user] 

413 item_emb_2 = item_lightgcn_emb[all_item] 

414 user_emb_2 = user_gcn_emb[user] 

415 item_emb_3 = item_gcn_emb[all_item] 

416 

417 local_loss = self.local_level_loss(item_emb_1, item_emb_2) 

418 global_loss = self.global_level_loss_1(user_emb_2, user_emb_1) + self.global_level_loss_2( 

419 item_emb_3, item_emb_1 + item_emb_2 

420 ) 

421 

422 user_embedding = torch.cat((user_emb_2, user_emb_1), dim=-1) 

423 pos_item_embedding = torch.cat( 

424 ( 

425 item_gcn_emb[pos_item], 

426 item_semantic_emb[pos_item] + item_lightgcn_emb[pos_item], 

427 ), 

428 dim=-1, 

429 ) 

430 neg_item_embedding = torch.cat( 

431 ( 

432 item_gcn_emb[neg_item], 

433 item_semantic_emb[neg_item] + item_lightgcn_emb[neg_item], 

434 ), 

435 dim=-1, 

436 ) 

437 

438 pos_scores = torch.mul(user_embedding, pos_item_embedding).sum(dim=1) 

439 neg_scores = torch.mul(user_embedding, neg_item_embedding).sum(dim=1) 

440 if self.loss_type.lower() == "bpr": 

441 rec_loss = self.rec_loss(pos_scores, neg_scores) 

442 else: 

443 predict = torch.cat((pos_scores, neg_scores)) 

444 target = torch.zeros(len(pos_item) + len(neg_item), dtype=torch.float32).to(self.device) 

445 target[: len(pos_item)] = 1 

446 rec_loss = self.rec_loss(predict, target) 

447 

448 reg_loss = self.reg_loss(user_embedding, pos_item_embedding, neg_item_embedding) 

449 loss = ( 

450 rec_loss 

451 + self.reg_weight * reg_loss 

452 + self.beta * (self.alpha * local_loss + (1 - self.alpha) * global_loss) 

453 ) 

454 

455 return loss 

456 

457 def local_level_loss(self, A_embedding, B_embedding): 

458 # The loss of local-level contrastive learning 

459 def exp_temp(x): 

460 return torch.exp(x / self.temperature) 

461 

462 A_embedding = self.fc1(A_embedding) 

463 B_embedding = self.fc1(B_embedding) 

464 refl_sim = exp_temp(self.sim(A_embedding, A_embedding)) 

465 between_sim = exp_temp(self.sim(A_embedding, B_embedding)) 

466 local_loss = -torch.log(between_sim.diag() / (refl_sim.sum(1) + between_sim.sum(1) - refl_sim.diag())) 

467 local_loss = local_loss.mean() 

468 return local_loss 

469 

470 def global_level_loss_1(self, A_embedding, B_embedding): 

471 # The user embedding loss of global-level contrastive learning 

472 def exp_temp(x): 

473 return torch.exp(x / self.temperature) 

474 

475 A_embedding = self.fc2(A_embedding) 

476 B_embedding = self.fc2(B_embedding) 

477 

478 refl_sim_1 = exp_temp(self.sim(A_embedding, A_embedding)) 

479 between_sim_1 = exp_temp(self.sim(A_embedding, B_embedding)) 

480 loss_1 = -torch.log(between_sim_1.diag() / (refl_sim_1.sum(1) + between_sim_1.sum(1) - refl_sim_1.diag())) 

481 

482 refl_sim_2 = exp_temp(self.sim(B_embedding, B_embedding)) 

483 between_sim_2 = exp_temp(self.sim(B_embedding, A_embedding)) 

484 loss_2 = -torch.log(between_sim_2.diag() / (refl_sim_2.sum(1) + between_sim_2.sum(1) - refl_sim_2.diag())) 

485 

486 global_user_loss = (loss_1 + loss_2) * 0.5 

487 global_user_loss = global_user_loss.mean() 

488 return global_user_loss 

489 

490 def global_level_loss_2(self, A_embedding, B_embedding): 

491 # The item embedding loss of global-level contrastive learning 

492 def exp_temp(x): 

493 return torch.exp(x / self.temperature) 

494 

495 A_embedding = self.fc3(A_embedding) 

496 B_embedding = self.fc3(B_embedding) 

497 

498 refl_sim_1 = exp_temp(self.sim(A_embedding, A_embedding)) 

499 between_sim_1 = exp_temp(self.sim(A_embedding, B_embedding)) 

500 loss_1 = -torch.log(between_sim_1.diag() / (refl_sim_1.sum(1) + between_sim_1.sum(1) - refl_sim_1.diag())) 

501 

502 refl_sim_2 = exp_temp(self.sim(B_embedding, B_embedding)) 

503 between_sim_2 = exp_temp(self.sim(B_embedding, A_embedding)) 

504 loss_2 = -torch.log(between_sim_2.diag() / (refl_sim_2.sum(1) + between_sim_2.sum(1) - refl_sim_2.diag())) 

505 

506 global_item_loss = (loss_1 + loss_2) * 0.5 

507 global_item_loss = global_item_loss.mean() 

508 return global_item_loss 

509 

510 def predict(self, interaction): 

511 user = interaction[self.USER_ID] 

512 item = interaction[self.ITEM_ID] 

513 

514 ( 

515 item_semantic_emb, 

516 user_lightgcn_emb, 

517 item_lightgcn_emb, 

518 user_gcn_emb, 

519 item_gcn_emb, 

520 ) = self.forward() 

521 item_emb_1 = item_semantic_emb[item] 

522 user_emb_1 = user_lightgcn_emb[user] 

523 item_emb_2 = item_lightgcn_emb[item] 

524 user_emb_2 = user_gcn_emb[user] 

525 item_emb_3 = item_gcn_emb[item] 

526 

527 user_embedding = torch.cat((user_emb_2, user_emb_1), dim=-1) 

528 item_embedding = torch.cat((item_emb_3, item_emb_1 + item_emb_2), dim=-1) 

529 

530 scores = torch.mul(user_embedding, item_embedding).sum(dim=1) 

531 if self.loss_type.lower() == "bce": 

532 scores = self.sigmoid(scores) 

533 return scores 

534 

535 def full_sort_predict(self, interaction): 

536 user = interaction[self.USER_ID] 

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

538 ( 

539 item_semantic_emb, 

540 user_lightgcn_emb, 

541 item_lightgcn_emb, 

542 user_gcn_emb, 

543 entity_gcn_emb, 

544 ) = self.forward() 

545 self.restore_user_e = torch.cat((user_gcn_emb, user_lightgcn_emb), dim=-1) 

546 self.restore_entity_e = torch.cat((entity_gcn_emb, item_semantic_emb + item_lightgcn_emb), dim=-1) 

547 

548 u_embeddings = self.restore_user_e[user] 

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

550 

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

552 if self.loss_type.lower() == "bce": 

553 scores = self.sigmoid(scores) 

554 

555 return scores.view(-1)