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

1# @Time : 2020/9/28 

2# @Author : gaole he 

3# @Email : hegaole@ruc.edu.cn 

4 

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""" 

11 

12import collections 

13 

14import numpy as np 

15import torch 

16from torch import nn 

17 

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 

22 

23 

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 """ 

31 

32 input_type = InputType.POINTWISE 

33 

34 def __init__(self, config, dataset): 

35 super().__init__(config, dataset) 

36 

37 # load dataset info 

38 self.LABEL = config["LABEL_FIELD"] 

39 

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() 

68 

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() 

78 

79 # parameters initialization 

80 self.apply(xavier_normal_initialization) 

81 self.other_parameter_name = ["ripple_set"] 

82 

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 

86 

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 = [] 

97 

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] 

102 

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]) 

110 

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 

140 

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) 

155 

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())) 

165 

166 # [batch size * n_memory, dim] 

167 self.h_emb_list.append(self.entity_embedding(head_ent)) 

168 

169 # [batch size * n_memory, dim * dim] 

170 self.r_emb_list.append(self.relation_embedding(relation)) 

171 

172 # [batch size * n_memory, dim] 

173 self.t_emb_list.append(self.entity_embedding(tail_ent)) 

174 

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 

181 

182 def _key_addressing(self): 

183 r"""Conduct reasoning for specific item and user ripple set 

184 

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) 

192 

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) 

197 

198 # [batch_size, dim, 1] 

199 v = self.item_embeddings.unsqueeze(2) 

200 

201 # [batch_size, n_memory] 

202 probs = torch.bmm(Rh, v).squeeze(2) 

203 

204 # [batch_size, n_memory] 

205 probs_normalized = self.softmax(probs) 

206 

207 # [batch_size, n_memory, 1] 

208 probs_expanded = probs_normalized.unsqueeze(2) 

209 

210 tail_emb = self.t_emb_list[hop].view(-1, self.n_memory, self.embedding_size) 

211 

212 # [batch_size, dim] 

213 o = torch.sum(tail_emb * probs_expanded, dim=1) 

214 

215 self.item_embeddings = self.transform_matrix(self.item_embeddings + o) 

216 # item embedding update 

217 o_list.append(o) 

218 return o_list 

219 

220 def calculate_loss(self, interaction): 

221 label = interaction[self.LABEL] 

222 output = self.forward(interaction) 

223 rec_loss = self.loss(output, label) 

224 

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)) 

241 

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 

251 

252 return loss 

253 

254 def predict(self, interaction): 

255 scores = self.forward(interaction) 

256 return scores 

257 

258 def _key_addressing_full(self): 

259 r"""Conduct reasoning for specific item and user ripple set 

260 

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) 

269 

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) 

274 

275 batch_size = Rh.size(0) 

276 

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) 

290 

291 # [batch_size, n_memory, n_item] 

292 probs = torch.bmm(Rh, v) 

293 

294 # [batch_size, n_memory, n_item] 

295 probs_normalized = self.softmax(probs) 

296 

297 # [batch_size, n_item, n_memory] 

298 probs_transposed = probs_normalized.transpose(1, 2) 

299 

300 # [batch_size, n_memory, dim] 

301 tail_emb = self.t_emb_list[hop].view(-1, self.n_memory, self.embedding_size) 

302 

303 # [batch_size, n_item, dim] 

304 o = torch.bmm(probs_transposed, tail_emb) 

305 

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 

311 

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) 

327 

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())) 

337 

338 # [batch size * n_memory, dim] 

339 self.h_emb_list.append(self.entity_embedding(head_ent)) 

340 

341 # [batch size * n_memory, dim * dim] 

342 self.r_emb_list.append(self.relation_embedding(relation)) 

343 

344 # [batch size * n_memory, dim] 

345 self.t_emb_list.append(self.entity_embedding(tail_ent)) 

346 

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)