Coverage for hopwise/model/knowledge_aware_recommender/kgnnls.py: 97%

213 statements  

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

1# @Time : 2020/10/3 

2# @Author : Changxin Tian 

3# @Email : cx.tian@outlook.com 

4 

5r"""KGNNLS 

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

7 

8Reference: 

9 Hongwei Wang et al. "Knowledge-aware Graph Neural Networks with Label Smoothness Regularization 

10 for Recommender Systems." in KDD 2019. 

11 

12Reference code: 

13 https://github.com/hwwang55/KGNN-LS 

14""" 

15 

16import random 

17 

18import numpy as np 

19import torch 

20from torch import nn 

21 

22from hopwise.model.abstract_recommender import KnowledgeRecommender 

23from hopwise.model.init import xavier_normal_initialization 

24from hopwise.model.loss import EmbLoss 

25from hopwise.utils import InputType 

26 

27 

28class KGNNLS(KnowledgeRecommender): 

29 r"""KGNN-LS is a knowledge-based recommendation model. 

30 KGNN-LS transforms the knowledge graph into a user-specific weighted graph and then apply a graph neural network to 

31 compute personalized item embeddings. To provide better inductive bias, KGNN-LS relies on label smoothness 

32 assumption, which posits that adjacent items in the knowledge graph are likely to have similar user relevance 

33 labels/scores. Label smoothness provides regularization over the edge weights and it is equivalent to a label 

34 propagation scheme on a graph. 

35 """ 

36 

37 input_type = InputType.PAIRWISE 

38 

39 def __init__(self, config, dataset): 

40 super().__init__(config, dataset) 

41 

42 # load parameters info 

43 self.embedding_size = config["embedding_size"] 

44 self.neighbor_sample_size = config["neighbor_sample_size"] 

45 self.aggregator_class = config["aggregator"] # which aggregator to use 

46 # number of iterations when computing entity representation 

47 self.n_iter = config["n_iter"] 

48 self.reg_weight = config["reg_weight"] # weight of l2 regularization 

49 # weight of label Smoothness regularization 

50 self.ls_weight = config["ls_weight"] 

51 

52 # define embedding 

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

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

55 self.relation_embedding = nn.Embedding(self.n_relations + 1, self.embedding_size) 

56 

57 # sample neighbors and construct interaction table 

58 kg_graph = dataset.kg_graph(form="coo", value_field="relation_id") 

59 adj_entity, adj_relation = self.construct_adj(kg_graph) 

60 self.adj_entity, self.adj_relation = ( 

61 adj_entity.to(self.device), 

62 adj_relation.to(self.device), 

63 ) 

64 

65 inter_feat = dataset.inter_feat 

66 pos_users = inter_feat[dataset.uid_field] 

67 pos_items = inter_feat[dataset.iid_field] 

68 pos_label = torch.ones(pos_items.shape) 

69 pos_interaction_table, self.offset = self.get_interaction_table(pos_users, pos_items, pos_label) 

70 self.interaction_table = self.sample_neg_interaction(pos_interaction_table, self.offset) 

71 

72 # define function 

73 self.softmax = nn.Softmax(dim=-1) 

74 self.linear_layers = torch.nn.ModuleList() 

75 for i in range(self.n_iter): 

76 self.linear_layers.append( 

77 nn.Linear( 

78 (self.embedding_size if not self.aggregator_class == "concat" else self.embedding_size * 2), 

79 self.embedding_size, 

80 ) 

81 ) 

82 self.ReLU = nn.ReLU() 

83 self.Tanh = nn.Tanh() 

84 

85 self.bce_loss = nn.BCEWithLogitsLoss() 

86 self.l2_loss = EmbLoss() 

87 

88 # parameters initialization 

89 self.apply(xavier_normal_initialization) 

90 self.other_parameter_name = ["adj_entity", "adj_relation"] 

91 

92 def get_interaction_table(self, user_id, item_id, y): 

93 r"""Get interaction_table that is used for fetching user-item interaction label in LS regularization. 

94 

95 Args: 

96 user_id(torch.Tensor): the user id in user-item interactions, shape: [n_interactions, 1] 

97 item_id(torch.Tensor): the item id in user-item interactions, shape: [n_interactions, 1] 

98 y(torch.Tensor): the label in user-item interactions, shape: [n_interactions, 1] 

99 

100 Returns: 

101 tuple: 

102 - interaction_table(dict): key: user_id * 10^offset + item_id; value: y_{user_id, item_id} 

103 - offset(int): The offset that is used for calculating the key(index) in interaction_table 

104 """ 

105 offset = len(str(self.n_entities)) 

106 offset = 10**offset 

107 keys = user_id * offset + item_id 

108 keys = keys.int().cpu().numpy().tolist() 

109 values = y.float().cpu().numpy().tolist() 

110 

111 interaction_table = dict(zip(keys, values)) 

112 return interaction_table, offset 

113 

114 def sample_neg_interaction(self, pos_interaction_table, offset): 

115 r"""Sample neg_interaction to construct train data. 

116 

117 Args: 

118 pos_interaction_table(dict): the interaction_table that only contains pos_interaction. 

119 offset(int): The offset that is used for calculating the key(index) in interaction_table 

120 

121 Returns: 

122 interaction_table(dict): key: user_id * 10^offset + item_id; value: y_{user_id, item_id} 

123 """ 

124 pos_num = len(pos_interaction_table) 

125 neg_num = 0 

126 neg_interaction_table = {} 

127 while neg_num < pos_num: 

128 user_id = random.randint(0, self.n_users) 

129 item_id = random.randint(0, self.n_items) 

130 keys = user_id * offset + item_id 

131 if keys not in pos_interaction_table: 

132 neg_interaction_table[keys] = 0.0 

133 neg_num += 1 

134 interaction_table = {**pos_interaction_table, **neg_interaction_table} 

135 return interaction_table 

136 

137 def construct_adj(self, kg_graph): 

138 r"""Get neighbors and corresponding relations for each entity in the KG. 

139 

140 Args: 

141 kg_graph(scipy.sparse.coo_matrix): an undirected graph 

142 

143 Returns: 

144 tuple: 

145 - adj_entity (torch.LongTensor): each line stores the sampled neighbor entities for a given entity, 

146 shape: [n_entities, neighbor_sample_size] 

147 - adj_relation (torch.LongTensor): each line stores the corresponding sampled neighbor relations, 

148 shape: [n_entities, neighbor_sample_size] 

149 """ 

150 # self.logger.info('constructing knowledge graph ...') 

151 # treat the KG as an undirected graph 

152 kg_dict = dict() 

153 for triple in zip(kg_graph.row, kg_graph.data, kg_graph.col): 

154 head = triple[0] 

155 relation = triple[1] 

156 tail = triple[2] 

157 if head not in kg_dict: 

158 kg_dict[head] = [] 

159 kg_dict[head].append((tail, relation)) 

160 if tail not in kg_dict: 

161 kg_dict[tail] = [] 

162 kg_dict[tail].append((head, relation)) 

163 

164 # self.logger.info('constructing adjacency matrix ...') 

165 # each line of adj_entity stores the sampled neighbor entities for a given entity 

166 # each line of adj_relation stores the corresponding sampled neighbor relations 

167 entity_num = kg_graph.shape[0] 

168 adj_entity = np.zeros([entity_num, self.neighbor_sample_size], dtype=np.int64) 

169 adj_relation = np.zeros([entity_num, self.neighbor_sample_size], dtype=np.int64) 

170 for entity in range(entity_num): 

171 if entity not in kg_dict.keys(): 

172 adj_entity[entity] = np.array([entity] * self.neighbor_sample_size) 

173 adj_relation[entity] = np.array([0] * self.neighbor_sample_size) 

174 continue 

175 

176 neighbors = kg_dict[entity] 

177 n_neighbors = len(neighbors) 

178 if n_neighbors >= self.neighbor_sample_size: 

179 sampled_indices = np.random.choice( 

180 list(range(n_neighbors)), 

181 size=self.neighbor_sample_size, 

182 replace=False, 

183 ) 

184 else: 

185 sampled_indices = np.random.choice( 

186 list(range(n_neighbors)), 

187 size=self.neighbor_sample_size, 

188 replace=True, 

189 ) 

190 adj_entity[entity] = np.array([neighbors[i][0] for i in sampled_indices]) 

191 adj_relation[entity] = np.array([neighbors[i][1] for i in sampled_indices]) 

192 

193 return torch.from_numpy(adj_entity), torch.from_numpy(adj_relation) 

194 

195 def get_neighbors(self, items): 

196 r"""Get neighbors and corresponding relations for each entity in items from adj_entity and adj_relation. 

197 

198 Args: 

199 items(torch.LongTensor): The input tensor that contains item's id, shape: [batch_size, ] 

200 

201 Returns: 

202 tuple: 

203 - entities(list): Entities is a list of i-iter (i = 0, 1, ..., n_iter) neighbors for the batch of items. 

204 dimensions of entities: {[batch_size, 1], 

205 [batch_size, n_neighbor], 

206 [batch_size, n_neighbor^2], 

207 ..., 

208 [batch_size, n_neighbor^n_iter]} 

209 - relations(list): Relations is a list of i-iter (i = 0, 1, ..., n_iter) corresponding relations for 

210 entities. Relations have the same shape as entities. 

211 """ # noqa: E501 

212 items = torch.unsqueeze(items, dim=1) 

213 entities = [items] 

214 relations = [] 

215 for i in range(self.n_iter): 

216 index = torch.flatten(entities[i]) 

217 neighbor_entities = torch.index_select(self.adj_entity, 0, index).reshape(self.batch_size, -1) 

218 neighbor_relations = torch.index_select(self.adj_relation, 0, index).reshape(self.batch_size, -1) 

219 entities.append(neighbor_entities) 

220 relations.append(neighbor_relations) 

221 return entities, relations 

222 

223 def aggregate(self, user_embeddings, entities, relations): 

224 r"""For each item, aggregate the entity representation and its neighborhood representation into a single vector. 

225 

226 Args: 

227 user_embeddings(torch.FloatTensor): The embeddings of users, shape: [batch_size, embedding_size] 

228 entities(list): entities is a list of i-iter (i = 0, 1, ..., n_iter) neighbors for the batch of items. 

229 dimensions of entities: {[batch_size, 1], 

230 [batch_size, n_neighbor], 

231 [batch_size, n_neighbor^2], 

232 ..., 

233 [batch_size, n_neighbor^n_iter]} 

234 relations(list): relations is a list of i-iter (i = 0, 1, ..., n_iter) corresponding relations for entities. 

235 relations have the same shape as entities. 

236 

237 Returns: 

238 item_embeddings(torch.FloatTensor): The embeddings of items, shape: [batch_size, embedding_size] 

239 

240 """ # noqa: E501 

241 entity_vectors = [self.entity_embedding(i) for i in entities] 

242 relation_vectors = [self.relation_embedding(i) for i in relations] 

243 

244 for i in range(self.n_iter): 

245 entity_vectors_next_iter = [] 

246 for hop in range(self.n_iter - i): 

247 shape = ( 

248 self.batch_size, 

249 -1, 

250 self.neighbor_sample_size, 

251 self.embedding_size, 

252 ) 

253 self_vectors = entity_vectors[hop] 

254 neighbor_vectors = entity_vectors[hop + 1].reshape(shape) 

255 neighbor_relations = relation_vectors[hop].reshape(shape) 

256 

257 # mix_neighbor_vectors 

258 user_embeddings = user_embeddings.reshape( 

259 self.batch_size, 1, 1, self.embedding_size 

260 ) # [batch_size, 1, 1, dim] 

261 user_relation_scores = torch.mean( 

262 user_embeddings * neighbor_relations, dim=-1 

263 ) # [batch_size, -1, n_neighbor] 

264 user_relation_scores_normalized = torch.unsqueeze( 

265 self.softmax(user_relation_scores), dim=-1 

266 ) # [batch_size, -1, n_neighbor, 1] 

267 neighbors_agg = torch.mean( 

268 user_relation_scores_normalized * neighbor_vectors, dim=2 

269 ) # [batch_size, -1, dim] 

270 

271 if self.aggregator_class == "sum": 

272 output = (self_vectors + neighbors_agg).reshape(-1, self.embedding_size) # [-1, dim] 

273 elif self.aggregator_class == "neighbor": 

274 output = neighbors_agg.reshape(-1, self.embedding_size) # [-1, dim] 

275 elif self.aggregator_class == "concat": 

276 # [batch_size, -1, dim * 2] 

277 output = torch.cat([self_vectors, neighbors_agg], dim=-1) 

278 output = output.reshape(-1, self.embedding_size * 2) # [-1, dim * 2] 

279 else: 

280 raise Exception("Unknown aggregator: " + self.aggregator_class) 

281 

282 output = self.linear_layers[i](output) 

283 # [batch_size, -1, dim] 

284 output = output.reshape(self.batch_size, -1, self.embedding_size) 

285 

286 if i == self.n_iter - 1: 

287 vector = self.Tanh(output) 

288 else: 

289 vector = self.ReLU(output) 

290 

291 entity_vectors_next_iter.append(vector) 

292 entity_vectors = entity_vectors_next_iter 

293 

294 res = entity_vectors[0].reshape(self.batch_size, self.embedding_size) 

295 return res 

296 

297 def label_smoothness_predict(self, user_embeddings, user, entities, relations): 

298 r"""Predict the label of items by label smoothness. 

299 

300 Args: 

301 user_embeddings(torch.FloatTensor): The embeddings of users, shape: [batch_size*2, embedding_size], 

302 user(torch.FloatTensor): the index of users, shape: [batch_size*2] 

303 entities(list): entities is a list of i-iter (i = 0, 1, ..., n_iter) neighbors for the batch of items. 

304 dimensions of entities: {[batch_size*2, 1], 

305 [batch_size*2, n_neighbor], 

306 [batch_size*2, n_neighbor^2], 

307 ..., 

308 [batch_size*2, n_neighbor^n_iter]} 

309 relations(list): relations is a list of i-iter (i = 0, 1, ..., n_iter) corresponding relations for entities. 

310 relations have the same shape as entities. 

311 

312 Returns: 

313 predicted_labels(torch.FloatTensor): The predicted label of items, shape: [batch_size*2] 

314 """ # noqa: E501 

315 # calculate initial labels; calculate updating masks for label propagation 

316 entity_labels = [] 

317 # True means the label of this item is reset to initial value during label propagation 

318 reset_masks = [] 

319 holdout_item_for_user = None 

320 

321 for entities_per_iter in entities: 

322 users = torch.unsqueeze(user, dim=1) # [batch_size, 1] 

323 user_entity_concat = users * self.offset + entities_per_iter # [batch_size, n_neighbor^i] 

324 

325 # the first one in entities is the items to be held out 

326 if holdout_item_for_user is None: 

327 holdout_item_for_user = user_entity_concat 

328 

329 def lookup_interaction_table(x, _): 

330 x = int(x) 

331 label = self.interaction_table.setdefault(x, 0.5) 

332 return label 

333 

334 initial_label = user_entity_concat.clone().cpu().double() 

335 initial_label.map_(initial_label, lookup_interaction_table) 

336 initial_label = initial_label.float().to(self.device) 

337 

338 # False if the item is held out 

339 holdout_mask = (holdout_item_for_user - user_entity_concat).bool() 

340 # True if the entity is a labeled item 

341 reset_mask = (initial_label - 0.5).bool() 

342 reset_mask = torch.logical_and(reset_mask, holdout_mask) # remove held-out items 

343 initial_label = ( 

344 holdout_mask.float() * initial_label + torch.logical_not(holdout_mask).float() * 0.5 

345 ) # label initialization 

346 

347 reset_masks.append(reset_mask) 

348 entity_labels.append(initial_label) 

349 # we do not need the reset_mask for the last iteration 

350 reset_masks = reset_masks[:-1] 

351 

352 # label propagation 

353 relation_vectors = [self.relation_embedding(i) for i in relations] 

354 for i in range(self.n_iter): 

355 entity_labels_next_iter = [] 

356 for hop in range(self.n_iter - i): 

357 masks = reset_masks[hop] 

358 self_labels = entity_labels[hop] 

359 neighbor_labels = entity_labels[hop + 1].reshape(self.batch_size, -1, self.neighbor_sample_size) 

360 neighbor_relations = relation_vectors[hop].reshape( 

361 self.batch_size, -1, self.neighbor_sample_size, self.embedding_size 

362 ) 

363 

364 # mix_neighbor_labels 

365 user_embeddings = user_embeddings.reshape( 

366 self.batch_size, 1, 1, self.embedding_size 

367 ) # [batch_size, 1, 1, dim] 

368 user_relation_scores = torch.mean( 

369 user_embeddings * neighbor_relations, dim=-1 

370 ) # [batch_size, -1, n_neighbor] 

371 user_relation_scores_normalized = self.softmax(user_relation_scores) # [batch_size, -1, n_neighbor] 

372 

373 neighbors_aggregated_label = torch.mean( 

374 user_relation_scores_normalized * neighbor_labels, dim=2 

375 ) # [batch_size, -1, dim] # [batch_size, -1] 

376 output = masks.float() * self_labels + torch.logical_not(masks).float() * neighbors_aggregated_label 

377 

378 entity_labels_next_iter.append(output) 

379 entity_labels = entity_labels_next_iter 

380 

381 predicted_labels = entity_labels[0].squeeze(-1) 

382 return predicted_labels 

383 

384 def forward(self, user, item): 

385 self.batch_size = item.shape[0] 

386 # [batch_size, dim] 

387 user_e = self.user_embedding(user) 

388 # entities is a list of i-iter (i = 0, 1, ..., n_iter) neighbors for the batch of items. dimensions of entities: # noqa: E501 

389 # {[batch_size, 1], [batch_size, n_neighbor], [batch_size, n_neighbor^2], ..., [batch_size, n_neighbor^n_iter]} 

390 entities, relations = self.get_neighbors(item) 

391 # [batch_size, dim] 

392 item_e = self.aggregate(user_e, entities, relations) 

393 

394 return user_e, item_e 

395 

396 def calculate_ls_loss(self, user, item, target): 

397 r"""Calculate label smoothness loss. 

398 

399 Args: 

400 user(torch.FloatTensor): the index of users, shape: [batch_size*2], 

401 item(torch.FloatTensor): the index of items, shape: [batch_size*2], 

402 target(torch.FloatTensor): the label of user-item, shape: [batch_size*2], 

403 

404 Returns: 

405 ls_loss: label smoothness loss 

406 """ 

407 user_e = self.user_embedding(user) 

408 entities, relations = self.get_neighbors(item) 

409 

410 predicted_labels = self.label_smoothness_predict(user_e, user, entities, relations) 

411 ls_loss = self.bce_loss(predicted_labels, target) 

412 return ls_loss 

413 

414 def calculate_loss(self, interaction): 

415 user = interaction[self.USER_ID] 

416 pos_item = interaction[self.ITEM_ID] 

417 neg_item = interaction[self.NEG_ITEM_ID] 

418 target = torch.zeros(len(user) * 2, dtype=torch.float32).to(self.device) 

419 target[: len(user)] = 1 

420 

421 users = torch.cat((user, user)) 

422 items = torch.cat((pos_item, neg_item)) 

423 

424 user_e, item_e = self.forward(users, items) 

425 predict = torch.mul(user_e, item_e).sum(dim=1) 

426 rec_loss = self.bce_loss(predict, target) 

427 

428 ls_loss = self.calculate_ls_loss(users, items, target) 

429 l2_loss = self.l2_loss(user_e, item_e) 

430 

431 loss = rec_loss + self.ls_weight * ls_loss + self.reg_weight * l2_loss 

432 return loss 

433 

434 def predict(self, interaction): 

435 user = interaction[self.USER_ID] 

436 item = interaction[self.ITEM_ID] 

437 user_e, item_e = self.forward(user, item) 

438 return torch.mul(user_e, item_e).sum(dim=1) 

439 

440 def full_sort_predict(self, interaction): 

441 user_index = interaction[self.USER_ID] 

442 item_index = torch.tensor(range(self.n_items)).to(self.device) 

443 

444 user = torch.unsqueeze(user_index, dim=1).repeat(1, item_index.shape[0]) 

445 user = torch.flatten(user) 

446 item = torch.unsqueeze(item_index, dim=0).repeat(user_index.shape[0], 1) 

447 item = torch.flatten(item) 

448 

449 user_e, item_e = self.forward(user, item) 

450 score = torch.mul(user_e, item_e).sum(dim=1) 

451 

452 return score.view(-1)