Coverage for hopwise/model/knowledge_aware_recommender/kgcn.py: 95%
145 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/6
2# @Author : Changxin Tian
3# @Email : cx.tian@outlook.com
5r"""KGCN
6################################################
8Reference:
9 Hongwei Wang et al. "Knowledge graph convolution networks for recommender systems." in WWW 2019.
11Reference code:
12 https://github.com/hwwang55/KGCN
13"""
15import numpy as np
16import torch
17from torch import nn
19from hopwise.model.abstract_recommender import KnowledgeRecommender
20from hopwise.model.init import xavier_normal_initialization
21from hopwise.model.loss import EmbLoss
22from hopwise.utils import InputType
25class KGCN(KnowledgeRecommender):
26 r"""KGCN is a knowledge-based recommendation model that captures inter-item relatedness effectively by mining their
27 associated attributes on the KG. To automatically discover both high-order structure information and semantic
28 information of the KG, we treat KG as an undirected graph and sample from the neighbors for each entity in the KG
29 as their receptive field, then combine neighborhood information with bias when calculating the representation of a
30 given entity.
31 """
33 input_type = InputType.PAIRWISE
35 def __init__(self, config, dataset):
36 super().__init__(config, dataset)
38 # load parameters info
39 self.embedding_size = config["embedding_size"]
40 # number of iterations when computing entity representation
41 self.n_iter = config["n_iter"]
42 self.aggregator_class = config["aggregator"] # which aggregator to use
43 self.reg_weight = config["reg_weight"] # weight of l2 regularization
44 self.neighbor_sample_size = config["neighbor_sample_size"]
46 # define embedding
47 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
48 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size)
49 self.relation_embedding = nn.Embedding(self.n_relations + 1, self.embedding_size)
51 # sample neighbors
52 kg_graph = dataset.kg_graph(form="coo", value_field="relation_id")
53 adj_entity, adj_relation = self.construct_adj(kg_graph)
54 self.adj_entity, self.adj_relation = (
55 adj_entity.to(self.device),
56 adj_relation.to(self.device),
57 )
59 # define function
60 self.softmax = nn.Softmax(dim=-1)
61 self.linear_layers = torch.nn.ModuleList()
62 for i in range(self.n_iter):
63 self.linear_layers.append(
64 nn.Linear(
65 (self.embedding_size if not self.aggregator_class == "concat" else self.embedding_size * 2),
66 self.embedding_size,
67 )
68 )
69 self.ReLU = nn.ReLU()
70 self.Tanh = nn.Tanh()
72 self.bce_loss = nn.BCEWithLogitsLoss()
73 self.l2_loss = EmbLoss()
75 # parameters initialization
76 self.apply(xavier_normal_initialization)
77 self.other_parameter_name = ["adj_entity", "adj_relation"]
79 def construct_adj(self, kg_graph):
80 r"""Get neighbors and corresponding relations for each entity in the KG.
82 Args:
83 kg_graph(scipy.sparse.coo_matrix): an undirected graph
85 Returns:
86 tuple:
87 - adj_entity(torch.LongTensor): each line stores the sampled neighbor entities for a given entity,
88 shape: [n_entities, neighbor_sample_size]
89 - adj_relation(torch.LongTensor): each line stores the corresponding sampled neighbor relations,
90 shape: [n_entities, neighbor_sample_size]
91 """
92 # self.logger.info('constructing knowledge graph ...')
93 # treat the KG as an undirected graph
94 kg_dict = dict()
95 for triple in zip(kg_graph.row, kg_graph.data, kg_graph.col):
96 head = triple[0]
97 relation = triple[1]
98 tail = triple[2]
99 if head not in kg_dict:
100 kg_dict[head] = []
101 kg_dict[head].append((tail, relation))
102 if tail not in kg_dict:
103 kg_dict[tail] = []
104 kg_dict[tail].append((head, relation))
106 # self.logger.info('constructing adjacency matrix ...')
107 # each line of adj_entity stores the sampled neighbor entities for a given entity
108 # each line of adj_relation stores the corresponding sampled neighbor relations
109 entity_num = kg_graph.shape[0]
110 adj_entity = np.zeros([entity_num, self.neighbor_sample_size], dtype=np.int64)
111 adj_relation = np.zeros([entity_num, self.neighbor_sample_size], dtype=np.int64)
112 for entity in range(entity_num):
113 if entity not in kg_dict.keys():
114 adj_entity[entity] = np.array([entity] * self.neighbor_sample_size)
115 adj_relation[entity] = np.array([0] * self.neighbor_sample_size)
116 continue
118 neighbors = kg_dict[entity]
119 n_neighbors = len(neighbors)
120 if n_neighbors >= self.neighbor_sample_size:
121 sampled_indices = np.random.choice(
122 list(range(n_neighbors)),
123 size=self.neighbor_sample_size,
124 replace=False,
125 )
126 else:
127 sampled_indices = np.random.choice(
128 list(range(n_neighbors)),
129 size=self.neighbor_sample_size,
130 replace=True,
131 )
132 adj_entity[entity] = np.array([neighbors[i][0] for i in sampled_indices])
133 adj_relation[entity] = np.array([neighbors[i][1] for i in sampled_indices])
135 return torch.from_numpy(adj_entity), torch.from_numpy(adj_relation)
137 def get_neighbors(self, items):
138 r"""Get neighbors and corresponding relations for each entity in items from adj_entity and adj_relation.
140 Args:
141 items(torch.LongTensor): The input tensor that contains item's id, shape: [batch_size, ]
143 Returns:
144 tuple:
145 - entities(list): Entities is a list of i-iter (i = 0, 1, ..., n_iter) neighbors for the batch of items.
146 dimensions of entities: {[batch_size, 1],
147 [batch_size, n_neighbor],
148 [batch_size, n_neighbor^2],
149 ...,
150 [batch_size, n_neighbor^n_iter]}
151 - relations(list): Relations is a list of i-iter (i = 0, 1, ..., n_iter) corresponding relations for
152 entities. Relations have the same shape as entities.
153 """ # noqa: E501
154 items = torch.unsqueeze(items, dim=1)
155 entities = [items]
156 relations = []
157 for i in range(self.n_iter):
158 index = torch.flatten(entities[i])
159 neighbor_entities = torch.index_select(self.adj_entity, 0, index).reshape(self.batch_size, -1)
160 neighbor_relations = torch.index_select(self.adj_relation, 0, index).reshape(self.batch_size, -1)
161 entities.append(neighbor_entities)
162 relations.append(neighbor_relations)
163 return entities, relations
165 def mix_neighbor_vectors(self, neighbor_vectors, neighbor_relations, user_embeddings):
166 r"""Mix neighbor vectors on user-specific graph.
168 Args:
169 neighbor_vectors(torch.FloatTensor): The embeddings of neighbor entities(items),
170 shape: [batch_size, -1, neighbor_sample_size, embedding_size]
171 neighbor_relations(torch.FloatTensor): The embeddings of neighbor relations,
172 shape: [batch_size, -1, neighbor_sample_size, embedding_size]
173 user_embeddings(torch.FloatTensor): The embeddings of users, shape: [batch_size, embedding_size]
175 Returns:
176 neighbors_aggregated(torch.FloatTensor): The neighbors aggregated embeddings,
177 shape: [batch_size, -1, embedding_size]
179 """
180 avg = False
181 if not avg:
182 user_embeddings = user_embeddings.reshape(
183 self.batch_size, 1, 1, self.embedding_size
184 ) # [batch_size, 1, 1, dim]
185 user_relation_scores = torch.mean(
186 user_embeddings * neighbor_relations, dim=-1
187 ) # [batch_size, -1, n_neighbor]
188 user_relation_scores_normalized = self.softmax(user_relation_scores) # [batch_size, -1, n_neighbor]
190 user_relation_scores_normalized = torch.unsqueeze(
191 user_relation_scores_normalized, dim=-1
192 ) # [batch_size, -1, n_neighbor, 1]
193 neighbors_aggregated = torch.mean(
194 user_relation_scores_normalized * neighbor_vectors, dim=2
195 ) # [batch_size, -1, dim]
196 else:
197 neighbors_aggregated = torch.mean(neighbor_vectors, dim=2) # [batch_size, -1, dim]
198 return neighbors_aggregated
200 def aggregate(self, user_embeddings, entities, relations):
201 r"""For each item, aggregate the entity representation and its neighborhood representation into a single vector.
203 Args:
204 user_embeddings(torch.FloatTensor): The embeddings of users, shape: [batch_size, embedding_size]
205 entities(list): entities is a list of i-iter (i = 0, 1, ..., n_iter) neighbors for the batch of items.
206 dimensions of entities: {[batch_size, 1],
207 [batch_size, n_neighbor],
208 [batch_size, n_neighbor^2],
209 ...,
210 [batch_size, n_neighbor^n_iter]}
211 relations(list): relations is a list of i-iter (i = 0, 1, ..., n_iter) corresponding relations for entities.
212 relations have the same shape as entities.
214 Returns:
215 item_embeddings(torch.FloatTensor): The embeddings of items, shape: [batch_size, embedding_size]
217 """ # noqa: E501
218 entity_vectors = [self.entity_embedding(i) for i in entities]
219 relation_vectors = [self.relation_embedding(i) for i in relations]
221 for i in range(self.n_iter):
222 entity_vectors_next_iter = []
223 for hop in range(self.n_iter - i):
224 shape = (
225 self.batch_size,
226 -1,
227 self.neighbor_sample_size,
228 self.embedding_size,
229 )
230 self_vectors = entity_vectors[hop]
231 neighbor_vectors = entity_vectors[hop + 1].reshape(shape)
232 neighbor_relations = relation_vectors[hop].reshape(shape)
234 neighbors_agg = self.mix_neighbor_vectors(
235 neighbor_vectors, neighbor_relations, user_embeddings
236 ) # [batch_size, -1, dim]
238 if self.aggregator_class == "sum":
239 output = (self_vectors + neighbors_agg).reshape(-1, self.embedding_size) # [-1, dim]
240 elif self.aggregator_class == "neighbor":
241 output = neighbors_agg.reshape(-1, self.embedding_size) # [-1, dim]
242 elif self.aggregator_class == "concat":
243 # [batch_size, -1, dim * 2]
244 output = torch.cat([self_vectors, neighbors_agg], dim=-1)
245 output = output.reshape(-1, self.embedding_size * 2) # [-1, dim * 2]
246 else:
247 raise Exception("Unknown aggregator: " + self.aggregator_class)
249 output = self.linear_layers[i](output)
250 # [batch_size, -1, dim]
251 output = output.reshape(self.batch_size, -1, self.embedding_size)
253 if i == self.n_iter - 1:
254 vector = self.Tanh(output)
255 else:
256 vector = self.ReLU(output)
258 entity_vectors_next_iter.append(vector)
259 entity_vectors = entity_vectors_next_iter
261 item_embeddings = entity_vectors[0].reshape(self.batch_size, self.embedding_size)
263 return item_embeddings
265 def forward(self, user, item):
266 self.batch_size = item.shape[0]
267 # [batch_size, dim]
268 user_e = self.user_embedding(user)
269 # entities is a list of i-iter (i = 0, 1, ..., n_iter) neighbors for the batch of items. dimensions of entities: # noqa: E501
270 # {[batch_size, 1], [batch_size, n_neighbor], [batch_size, n_neighbor^2], ..., [batch_size, n_neighbor^n_iter]}
271 entities, relations = self.get_neighbors(item)
272 # [batch_size, dim]
273 item_e = self.aggregate(user_e, entities, relations)
275 return user_e, item_e
277 def calculate_loss(self, interaction):
278 user = interaction[self.USER_ID]
279 pos_item = interaction[self.ITEM_ID]
280 neg_item = interaction[self.NEG_ITEM_ID]
282 user_e, pos_item_e = self.forward(user, pos_item)
283 user_e, neg_item_e = self.forward(user, neg_item)
285 pos_item_score = torch.mul(user_e, pos_item_e).sum(dim=1)
286 neg_item_score = torch.mul(user_e, neg_item_e).sum(dim=1)
288 predict = torch.cat((pos_item_score, neg_item_score))
289 target = torch.zeros(len(user) * 2, dtype=torch.float32).to(self.device)
290 target[: len(user)] = 1
291 rec_loss = self.bce_loss(predict, target)
293 l2_loss = self.l2_loss(user_e, pos_item_e, neg_item_e)
294 loss = rec_loss + self.reg_weight * l2_loss
296 return loss
298 def predict(self, interaction):
299 user = interaction[self.USER_ID]
300 item = interaction[self.ITEM_ID]
301 user_e, item_e = self.forward(user, item)
302 return torch.mul(user_e, item_e).sum(dim=1)
304 def full_sort_predict(self, interaction):
305 user_index = interaction[self.USER_ID]
306 item_index = torch.tensor(range(self.n_items)).to(self.device)
308 user = torch.unsqueeze(user_index, dim=1).repeat(1, item_index.shape[0])
309 user = torch.flatten(user)
310 item = torch.unsqueeze(item_index, dim=0).repeat(user_index.shape[0], 1)
311 item = torch.flatten(item)
313 user_e, item_e = self.forward(user, item)
314 score = torch.mul(user_e, item_e).sum(dim=1)
316 return score.view(-1)