Coverage for hopwise/model/knowledge_aware_recommender/ripplenet.py: 95%
210 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/9/28
2# @Author : gaole he
3# @Email : hegaole@ruc.edu.cn
5r"""RippleNet
6#####################################################
7Reference:
8 Hongwei Wang et al. "RippleNet: Propagating User Preferences on the Knowledge Graph for Recommender Systems."
9 in CIKM 2018.
10"""
12import collections
14import numpy as np
15import torch
16from torch import nn
18from hopwise.model.abstract_recommender import KnowledgeRecommender
19from hopwise.model.init import xavier_normal_initialization
20from hopwise.model.loss import BPRLoss, EmbLoss
21from hopwise.utils import InputType
24class RippleNet(KnowledgeRecommender):
25 r"""RippleNet is an knowledge enhanced matrix factorization model.
26 The original interaction matrix of :math:`n_{users} \times n_{items}`
27 and related knowledge graph is set as model input,
28 we carefully design the data interface and use ripple set to train and test efficiently.
29 We just implement the model following the original author with a pointwise training mode.
30 """
32 input_type = InputType.POINTWISE
34 def __init__(self, config, dataset):
35 super().__init__(config, dataset)
37 # load dataset info
38 self.LABEL = config["LABEL_FIELD"]
40 # load parameters info
41 self.embedding_size = config["embedding_size"]
42 self.kg_weight = config["kg_weight"]
43 self.reg_weight = config["reg_weight"]
44 self.n_hop = config["n_hop"]
45 self.n_memory = config["n_memory"]
46 self.interaction_matrix = dataset.inter_matrix(form="coo").astype(np.float32)
47 head_entities = dataset.head_entities.tolist()
48 tail_entities = dataset.tail_entities.tolist()
49 relations = dataset.relations.tolist()
50 kg = {}
51 for i in range(len(head_entities)):
52 head_ent = head_entities[i]
53 tail_ent = tail_entities[i]
54 relation = relations[i]
55 kg.setdefault(head_ent, [])
56 kg[head_ent].append((tail_ent, relation))
57 self.kg = kg
58 users = self.interaction_matrix.row.tolist()
59 items = self.interaction_matrix.col.tolist()
60 user_dict = {}
61 for i in range(len(users)):
62 user = users[i]
63 item = items[i]
64 user_dict.setdefault(user, [])
65 user_dict[user].append(item)
66 self.user_dict = user_dict
67 self.ripple_set = self._build_ripple_set()
69 # define layers and loss
70 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size)
71 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size * self.embedding_size)
72 self.transform_matrix = nn.Linear(self.embedding_size, self.embedding_size, bias=False)
73 self.softmax = torch.nn.Softmax(dim=1)
74 self.sigmoid = torch.nn.Sigmoid()
75 self.rec_loss = BPRLoss()
76 self.l2_loss = EmbLoss()
77 self.loss = nn.BCEWithLogitsLoss()
79 # parameters initialization
80 self.apply(xavier_normal_initialization)
81 self.other_parameter_name = ["ripple_set"]
83 def _build_ripple_set(self):
84 r"""Get the normalized interaction matrix of users and items according to A_values.
85 Get the ripple hop-wise ripple set for every user, w.r.t. their interaction history
87 Returns:
88 ripple_set (dict)
89 """
90 ripple_set = collections.defaultdict(list)
91 n_padding = 0
92 for user in self.user_dict:
93 for h in range(self.n_hop):
94 memories_h = []
95 memories_r = []
96 memories_t = []
98 if h == 0:
99 tails_of_last_hop = self.user_dict[user]
100 else:
101 tails_of_last_hop = ripple_set[user][-1][2]
103 for entity in tails_of_last_hop:
104 if entity not in self.kg:
105 continue
106 for tail_and_relation in self.kg[entity]:
107 memories_h.append(entity)
108 memories_r.append(tail_and_relation[1])
109 memories_t.append(tail_and_relation[0])
111 # if the current ripple set of the given user is empty,
112 # we simply copy the ripple set of the last hop here
113 if len(memories_h) == 0:
114 if h == 0:
115 # self.logger.info("user {} without 1-hop kg facts, fill with padding".format(user))
116 # raise AssertionError("User without facts in 1st hop")
117 n_padding += 1
118 memories_h = [0 for _ in range(self.n_memory)]
119 memories_r = [0 for _ in range(self.n_memory)]
120 memories_t = [0 for _ in range(self.n_memory)]
121 memories_h = torch.LongTensor(memories_h).to(self.device)
122 memories_r = torch.LongTensor(memories_r).to(self.device)
123 memories_t = torch.LongTensor(memories_t).to(self.device)
124 ripple_set[user].append((memories_h, memories_r, memories_t))
125 else:
126 ripple_set[user].append(ripple_set[user][-1])
127 else:
128 # sample a fixed-size 1-hop memory for each user
129 replace = len(memories_h) < self.n_memory
130 indices = np.random.choice(len(memories_h), size=self.n_memory, replace=replace)
131 memories_h = [memories_h[i] for i in indices]
132 memories_r = [memories_r[i] for i in indices]
133 memories_t = [memories_t[i] for i in indices]
134 memories_h = torch.LongTensor(memories_h).to(self.device)
135 memories_r = torch.LongTensor(memories_r).to(self.device)
136 memories_t = torch.LongTensor(memories_t).to(self.device)
137 ripple_set[user].append((memories_h, memories_r, memories_t))
138 self.logger.info(f"{n_padding} among {len(self.user_dict)} users are padded")
139 return ripple_set
141 def forward(self, interaction):
142 users = interaction[self.USER_ID].cpu().numpy()
143 memories_h, memories_r, memories_t = {}, {}, {}
144 for hop in range(self.n_hop):
145 memories_h[hop] = []
146 memories_r[hop] = []
147 memories_t[hop] = []
148 for user in users:
149 memories_h[hop].append(self.ripple_set[user][hop][0])
150 memories_r[hop].append(self.ripple_set[user][hop][1])
151 memories_t[hop].append(self.ripple_set[user][hop][2])
152 # memories_h, memories_r, memories_t = self.ripple_set[user]
153 item = interaction[self.ITEM_ID]
154 self.item_embeddings = self.entity_embedding(item)
156 self.h_emb_list = []
157 self.r_emb_list = []
158 self.t_emb_list = []
159 for i in range(self.n_hop):
160 # [batch size * n_memory]
161 head_ent = torch.cat(memories_h[i], dim=0)
162 relation = torch.cat(memories_r[i], dim=0)
163 tail_ent = torch.cat(memories_t[i], dim=0)
164 # self.logger.info("Hop {}, size {}".format(i, head_ent.size(), relation.size(), tail_ent.size()))
166 # [batch size * n_memory, dim]
167 self.h_emb_list.append(self.entity_embedding(head_ent))
169 # [batch size * n_memory, dim * dim]
170 self.r_emb_list.append(self.relation_embedding(relation))
172 # [batch size * n_memory, dim]
173 self.t_emb_list.append(self.entity_embedding(tail_ent))
175 o_list = self._key_addressing()
176 y = o_list[-1]
177 for i in range(self.n_hop - 1):
178 y = y + o_list[i]
179 scores = torch.sum(self.item_embeddings * y, dim=1)
180 return scores
182 def _key_addressing(self):
183 r"""Conduct reasoning for specific item and user ripple set
185 Returns:
186 o_list (dict -> torch.cuda.FloatTensor): list of torch.cuda.FloatTensor n_hop * [batch_size, embedding_size]
187 """ # noqa: E501
188 o_list = []
189 for hop in range(self.n_hop):
190 # [batch_size * n_memory, dim, 1]
191 h_emb = self.h_emb_list[hop].unsqueeze(2)
193 # [batch_size * n_memory, dim, dim]
194 r_mat = self.r_emb_list[hop].view(-1, self.embedding_size, self.embedding_size)
195 # [batch_size, n_memory, dim]
196 Rh = torch.bmm(r_mat, h_emb).view(-1, self.n_memory, self.embedding_size)
198 # [batch_size, dim, 1]
199 v = self.item_embeddings.unsqueeze(2)
201 # [batch_size, n_memory]
202 probs = torch.bmm(Rh, v).squeeze(2)
204 # [batch_size, n_memory]
205 probs_normalized = self.softmax(probs)
207 # [batch_size, n_memory, 1]
208 probs_expanded = probs_normalized.unsqueeze(2)
210 tail_emb = self.t_emb_list[hop].view(-1, self.n_memory, self.embedding_size)
212 # [batch_size, dim]
213 o = torch.sum(tail_emb * probs_expanded, dim=1)
215 self.item_embeddings = self.transform_matrix(self.item_embeddings + o)
216 # item embedding update
217 o_list.append(o)
218 return o_list
220 def calculate_loss(self, interaction):
221 label = interaction[self.LABEL]
222 output = self.forward(interaction)
223 rec_loss = self.loss(output, label)
225 kge_loss = None
226 for hop in range(self.n_hop):
227 # (batch_size * n_memory, 1, dim)
228 h_expanded = self.h_emb_list[hop].unsqueeze(1)
229 # (batch_size * n_memory, dim)
230 t_expanded = self.t_emb_list[hop]
231 # (batch_size * n_memory, dim, dim)
232 r_mat = self.r_emb_list[hop].view(-1, self.embedding_size, self.embedding_size)
233 # (N, 1, dim) (N, dim, dim) -> (N, 1, dim)
234 hR = torch.bmm(h_expanded, r_mat).squeeze(1)
235 # (N, dim) (N, dim)
236 hRt = torch.sum(hR * t_expanded, dim=1)
237 if kge_loss is None:
238 kge_loss = torch.mean(self.sigmoid(hRt))
239 else:
240 kge_loss = kge_loss + torch.mean(self.sigmoid(hRt))
242 reg_loss = None
243 for hop in range(self.n_hop):
244 tp_loss = self.l2_loss(self.h_emb_list[hop], self.t_emb_list[hop], self.r_emb_list[hop])
245 if reg_loss is None:
246 reg_loss = tp_loss
247 else:
248 reg_loss = reg_loss + tp_loss
249 reg_loss = reg_loss + self.l2_loss(self.transform_matrix.weight)
250 loss = rec_loss - self.kg_weight * kge_loss + self.reg_weight * reg_loss
252 return loss
254 def predict(self, interaction):
255 scores = self.forward(interaction)
256 return scores
258 def _key_addressing_full(self):
259 r"""Conduct reasoning for specific item and user ripple set
261 Returns:
262 o_list (dict -> torch.cuda.FloatTensor): list of torch.cuda.FloatTensor
263 n_hop * [batch_size, n_item, embedding_size]
264 """
265 o_list = []
266 for hop in range(self.n_hop):
267 # [batch_size * n_memory, dim, 1]
268 h_emb = self.h_emb_list[hop].unsqueeze(2)
270 # [batch_size * n_memory, dim, dim]
271 r_mat = self.r_emb_list[hop].view(-1, self.embedding_size, self.embedding_size)
272 # [batch_size, n_memory, dim]
273 Rh = torch.bmm(r_mat, h_emb).view(-1, self.n_memory, self.embedding_size)
275 batch_size = Rh.size(0)
277 if len(self.item_embeddings.size()) == 2: # noqa: PLR2004
278 # [1, n_item, dim]
279 self.item_embeddings = self.item_embeddings.unsqueeze(0)
280 # [batch_size, n_item, dim]
281 self.item_embeddings = self.item_embeddings.expand(batch_size, -1, -1)
282 # [batch_size, dim, n_item]
283 v = self.item_embeddings.transpose(1, 2)
284 # [batch_size, dim, n_item]
285 v = v.expand(batch_size, -1, -1)
286 else:
287 assert len(self.item_embeddings.size()) == 3 # noqa: PLR2004
288 # [batch_size, dim, n_item]
289 v = self.item_embeddings.transpose(1, 2)
291 # [batch_size, n_memory, n_item]
292 probs = torch.bmm(Rh, v)
294 # [batch_size, n_memory, n_item]
295 probs_normalized = self.softmax(probs)
297 # [batch_size, n_item, n_memory]
298 probs_transposed = probs_normalized.transpose(1, 2)
300 # [batch_size, n_memory, dim]
301 tail_emb = self.t_emb_list[hop].view(-1, self.n_memory, self.embedding_size)
303 # [batch_size, n_item, dim]
304 o = torch.bmm(probs_transposed, tail_emb)
306 # [batch_size, n_item, dim] [batch_size, n_item, dim] -> [batch_size, n_item, dim]
307 self.item_embeddings = self.transform_matrix(self.item_embeddings + o)
308 # item embedding update
309 o_list.append(o)
310 return o_list
312 def full_sort_predict(self, interaction):
313 users = interaction[self.USER_ID].cpu().numpy()
314 memories_h, memories_r, memories_t = {}, {}, {}
315 for hop in range(self.n_hop):
316 memories_h[hop] = []
317 memories_r[hop] = []
318 memories_t[hop] = []
319 for user in users:
320 memories_h[hop].append(self.ripple_set[user][hop][0])
321 memories_r[hop].append(self.ripple_set[user][hop][1])
322 memories_t[hop].append(self.ripple_set[user][hop][2])
323 # memories_h, memories_r, memories_t = self.ripple_set[user]
324 # item = interaction[self.ITEM_ID]
325 self.item_embeddings = self.entity_embedding.weight[: self.n_items]
326 # self.item_embeddings = self.entity_embedding(item)
328 self.h_emb_list = []
329 self.r_emb_list = []
330 self.t_emb_list = []
331 for i in range(self.n_hop):
332 # [batch size * n_memory]
333 head_ent = torch.cat(memories_h[i], dim=0)
334 relation = torch.cat(memories_r[i], dim=0)
335 tail_ent = torch.cat(memories_t[i], dim=0)
336 # self.logger.info("Hop {}, size {}".format(i, head_ent.size(), relation.size(), tail_ent.size()))
338 # [batch size * n_memory, dim]
339 self.h_emb_list.append(self.entity_embedding(head_ent))
341 # [batch size * n_memory, dim * dim]
342 self.r_emb_list.append(self.relation_embedding(relation))
344 # [batch size * n_memory, dim]
345 self.t_emb_list.append(self.entity_embedding(tail_ent))
347 o_list = self._key_addressing_full()
348 y = o_list[-1]
349 for i in range(self.n_hop - 1):
350 y = y + o_list[i]
351 # [batch_size, n_item, dim] [batch_size, n_item, dim]
352 scores = torch.sum(self.item_embeddings * y, dim=-1)
353 return scores.view(-1)