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
« 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
5r"""KGNNLS
6################################################
8Reference:
9 Hongwei Wang et al. "Knowledge-aware Graph Neural Networks with Label Smoothness Regularization
10 for Recommender Systems." in KDD 2019.
12Reference code:
13 https://github.com/hwwang55/KGNN-LS
14"""
16import random
18import numpy as np
19import torch
20from torch import nn
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
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 """
37 input_type = InputType.PAIRWISE
39 def __init__(self, config, dataset):
40 super().__init__(config, dataset)
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"]
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)
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 )
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)
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()
85 self.bce_loss = nn.BCEWithLogitsLoss()
86 self.l2_loss = EmbLoss()
88 # parameters initialization
89 self.apply(xavier_normal_initialization)
90 self.other_parameter_name = ["adj_entity", "adj_relation"]
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.
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]
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()
111 interaction_table = dict(zip(keys, values))
112 return interaction_table, offset
114 def sample_neg_interaction(self, pos_interaction_table, offset):
115 r"""Sample neg_interaction to construct train data.
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
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
137 def construct_adj(self, kg_graph):
138 r"""Get neighbors and corresponding relations for each entity in the KG.
140 Args:
141 kg_graph(scipy.sparse.coo_matrix): an undirected graph
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))
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
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])
193 return torch.from_numpy(adj_entity), torch.from_numpy(adj_relation)
195 def get_neighbors(self, items):
196 r"""Get neighbors and corresponding relations for each entity in items from adj_entity and adj_relation.
198 Args:
199 items(torch.LongTensor): The input tensor that contains item's id, shape: [batch_size, ]
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
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.
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.
237 Returns:
238 item_embeddings(torch.FloatTensor): The embeddings of items, shape: [batch_size, embedding_size]
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]
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)
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]
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)
282 output = self.linear_layers[i](output)
283 # [batch_size, -1, dim]
284 output = output.reshape(self.batch_size, -1, self.embedding_size)
286 if i == self.n_iter - 1:
287 vector = self.Tanh(output)
288 else:
289 vector = self.ReLU(output)
291 entity_vectors_next_iter.append(vector)
292 entity_vectors = entity_vectors_next_iter
294 res = entity_vectors[0].reshape(self.batch_size, self.embedding_size)
295 return res
297 def label_smoothness_predict(self, user_embeddings, user, entities, relations):
298 r"""Predict the label of items by label smoothness.
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.
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
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]
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
329 def lookup_interaction_table(x, _):
330 x = int(x)
331 label = self.interaction_table.setdefault(x, 0.5)
332 return label
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)
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
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]
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 )
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]
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
378 entity_labels_next_iter.append(output)
379 entity_labels = entity_labels_next_iter
381 predicted_labels = entity_labels[0].squeeze(-1)
382 return predicted_labels
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)
394 return user_e, item_e
396 def calculate_ls_loss(self, user, item, target):
397 r"""Calculate label smoothness loss.
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],
404 Returns:
405 ls_loss: label smoothness loss
406 """
407 user_e = self.user_embedding(user)
408 entities, relations = self.get_neighbors(item)
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
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
421 users = torch.cat((user, user))
422 items = torch.cat((pos_item, neg_item))
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)
428 ls_loss = self.calculate_ls_loss(users, items, target)
429 l2_loss = self.l2_loss(user_e, item_e)
431 loss = rec_loss + self.ls_weight * ls_loss + self.reg_weight * l2_loss
432 return loss
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)
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)
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)
449 user_e, item_e = self.forward(user, item)
450 score = torch.mul(user_e, item_e).sum(dim=1)
452 return score.view(-1)