Coverage for hopwise/model/knowledge_aware_recommender/kgin.py: 80%

190 statements  

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

1# @Time : 2021/3/25 

2# @Author : Wenqi Sun 

3# @Email : wenqisun@pku.edu.cn 

4 

5# UPDATE: 

6# @Time : 2022/8/31 

7# @Author : Bowen Zheng 

8# @Email : 18735382001@163.com 

9 

10r"""KGIN 

11################################################## 

12Reference: 

13 Xiang Wang et al. "Learning Intents behind Interactions with Knowledge Graph for Recommendation." in WWW 2021. 

14Reference code: 

15 https://github.com/huangtinglin/Knowledge_Graph_based_Intent_Network 

16""" 

17 

18import numpy as np 

19import torch 

20import torch.nn.functional as F 

21from torch import nn 

22 

23from hopwise.model.abstract_recommender import KnowledgeRecommender 

24from hopwise.model.init import xavier_uniform_initialization 

25from hopwise.model.layers import SparseDropout 

26from hopwise.model.loss import BPRLoss, EmbLoss 

27from hopwise.utils import InputType 

28 

29 

30class Aggregator(nn.Module): 

31 """Relational Path-aware Convolution Network""" 

32 

33 def __init__( 

34 self, 

35 ): 

36 super().__init__() 

37 

38 def forward( 

39 self, 

40 entity_emb, 

41 user_emb, 

42 latent_emb, 

43 relation_emb, 

44 edge_index, 

45 edge_type, 

46 interact_mat, 

47 disen_weight_att, 

48 ): 

49 from torch_geometric.utils import scatter 

50 

51 n_entities = entity_emb.shape[0] 

52 

53 """KG aggregate""" 

54 head, tail = edge_index 

55 edge_relation_emb = relation_emb[edge_type] 

56 neigh_relation_emb = entity_emb[tail] * edge_relation_emb # [-1, embedding_size] 

57 entity_agg = scatter(src=neigh_relation_emb, index=head, dim_size=n_entities, dim=0, reduce="mean") 

58 

59 """cul user->latent factor attention""" 

60 score_ = torch.mm(user_emb, latent_emb.t()) 

61 score = nn.Softmax(dim=1)(score_) # [n_users, n_factors] 

62 """user aggregate""" 

63 user_agg = torch.sparse.mm(interact_mat, entity_emb) # [n_users, embedding_size] 

64 disen_weight = torch.mm(nn.Softmax(dim=-1)(disen_weight_att), relation_emb) # [n_factors, embedding_size] 

65 user_agg = (torch.mm(score, disen_weight)) * user_agg + user_agg # [n_users, embedding_size] 

66 

67 return entity_agg, user_agg 

68 

69 

70class GraphConv(nn.Module): 

71 """Graph Convolutional Network""" 

72 

73 def __init__( 

74 self, 

75 embedding_size, 

76 n_hops, 

77 n_users, 

78 n_factors, 

79 n_relations, 

80 edge_index, 

81 edge_type, 

82 interact_mat, 

83 ind, 

84 tmp, 

85 device, 

86 node_dropout_rate=0.5, 

87 mess_dropout_rate=0.1, 

88 ): 

89 super().__init__() 

90 

91 self.embedding_size = embedding_size 

92 self.n_hops = n_hops 

93 self.n_relations = n_relations 

94 self.n_users = n_users 

95 self.n_factors = n_factors 

96 self.edge_index = edge_index 

97 self.edge_type = edge_type 

98 self.interact_mat = interact_mat 

99 self.node_dropout_rate = node_dropout_rate 

100 self.mess_dropout_rate = mess_dropout_rate 

101 self.ind = ind 

102 self.temperature = tmp 

103 self.device = device 

104 

105 # define layers 

106 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size) 

107 disen_weight_att = nn.init.xavier_uniform_(torch.empty(n_factors, n_relations)) 

108 self.disen_weight_att = nn.Parameter(disen_weight_att) 

109 self.convs = nn.ModuleList() 

110 for i in range(self.n_hops): 

111 self.convs.append(Aggregator()) 

112 self.node_dropout = SparseDropout(p=self.mess_dropout_rate) # node dropout 

113 self.mess_dropout = nn.Dropout(p=self.mess_dropout_rate) # mess dropout 

114 

115 # parameters initialization 

116 self.apply(xavier_uniform_initialization) 

117 

118 def edge_sampling(self, edge_index, edge_type, rate=0.5): 

119 # edge_index: [2, -1] 

120 # edge_type: [-1] 

121 n_edges = edge_index.shape[1] 

122 random_indices = np.random.choice(n_edges, size=int(n_edges * rate), replace=False) 

123 return edge_index[:, random_indices], edge_type[random_indices] 

124 

125 def forward(self, user_emb, entity_emb, latent_emb): 

126 """Node dropout""" 

127 # node dropout 

128 if self.node_dropout_rate > 0.0: 

129 edge_index, edge_type = self.edge_sampling(self.edge_index, self.edge_type, self.node_dropout_rate) 

130 interact_mat = self.node_dropout(self.interact_mat) 

131 else: 

132 edge_index, edge_type = self.edge_index, self.edge_type 

133 interact_mat = self.interact_mat 

134 

135 entity_res_emb = entity_emb # [n_entities, embedding_size] 

136 user_res_emb = user_emb # [n_users, embedding_size] 

137 relation_emb = self.relation_embedding.weight # [n_relations, embedding_size] 

138 for i in range(len(self.convs)): 

139 entity_emb, user_emb = self.convs[i]( 

140 entity_emb, 

141 user_emb, 

142 latent_emb, 

143 relation_emb, 

144 edge_index, 

145 edge_type, 

146 interact_mat, 

147 self.disen_weight_att, 

148 ) 

149 """message dropout""" 

150 if self.mess_dropout_rate > 0.0: 

151 entity_emb = self.mess_dropout(entity_emb) 

152 user_emb = self.mess_dropout(user_emb) 

153 entity_emb = F.normalize(entity_emb) 

154 user_emb = F.normalize(user_emb) 

155 """result emb""" 

156 entity_res_emb = torch.add(entity_res_emb, entity_emb) 

157 user_res_emb = torch.add(user_res_emb, user_emb) 

158 

159 return ( 

160 entity_res_emb, 

161 user_res_emb, 

162 self.calculate_cor_loss(self.disen_weight_att), 

163 ) 

164 

165 def calculate_cor_loss(self, tensors): 

166 def CosineSimilarity(tensor_1, tensor_2): 

167 # tensor_1, tensor_2: [channel] 

168 normalized_tensor_1 = F.normalize(tensor_1, dim=0) 

169 normalized_tensor_2 = F.normalize(tensor_2, dim=0) 

170 return (normalized_tensor_1 * normalized_tensor_2).sum(dim=0) ** 2 # no negative 

171 

172 def DistanceCorrelation(tensor_1, tensor_2): 

173 # tensor_1, tensor_2: [channel] 

174 # ref: https://en.wikipedia.org/wiki/Distance_correlation 

175 channel = tensor_1.shape[0] 

176 zeros = torch.zeros(channel, channel).to(tensor_1.device) 

177 zero = torch.zeros(1).to(tensor_1.device) 

178 tensor_1, tensor_2 = tensor_1.unsqueeze(-1), tensor_2.unsqueeze(-1) 

179 """cul distance matrix""" 

180 a_, b_ = ( 

181 torch.matmul(tensor_1, tensor_1.t()) * 2, 

182 torch.matmul(tensor_2, tensor_2.t()) * 2, 

183 ) # [channel, channel] 

184 tensor_1_square, tensor_2_square = tensor_1**2, tensor_2**2 

185 a, b = ( 

186 torch.sqrt(torch.max(tensor_1_square - a_ + tensor_1_square.t(), zeros) + 1e-8), 

187 torch.sqrt(torch.max(tensor_2_square - b_ + tensor_2_square.t(), zeros) + 1e-8), 

188 ) # [channel, channel] 

189 """cul distance correlation""" 

190 A = a - a.mean(dim=0, keepdim=True) - a.mean(dim=1, keepdim=True) + a.mean() 

191 B = b - b.mean(dim=0, keepdim=True) - b.mean(dim=1, keepdim=True) + b.mean() 

192 dcov_AB = torch.sqrt(torch.max((A * B).sum() / channel**2, zero) + 1e-8) 

193 dcov_AA = torch.sqrt(torch.max((A * A).sum() / channel**2, zero) + 1e-8) 

194 dcov_BB = torch.sqrt(torch.max((B * B).sum() / channel**2, zero) + 1e-8) 

195 return dcov_AB / torch.sqrt(dcov_AA * dcov_BB + 1e-8) 

196 

197 def MutualInformation(tensors): 

198 # tensors: [n_factors, dimension] 

199 # normalized_tensors: [n_factors, dimension] 

200 normalized_tensors = F.normalize(tensors, dim=1) 

201 scores = torch.mm(normalized_tensors, normalized_tensors.t()) 

202 scores = torch.exp(scores / self.temperature) 

203 cor_loss = -torch.sum(torch.log(scores.diag() / scores.sum(1))) 

204 return cor_loss 

205 

206 """cul similarity for each latent factor weight pairs""" 

207 if self.ind == "mi": 

208 return MutualInformation(tensors) 

209 elif self.ind == "distance": 

210 cor_loss = 0.0 

211 for i in range(self.n_factors): 

212 for j in range(i + 1, self.n_factors): 

213 cor_loss += DistanceCorrelation(tensors[i], tensors[j]) 

214 elif self.ind == "cosine": 

215 cor_loss = 0.0 

216 for i in range(self.n_factors): 

217 for j in range(i + 1, self.n_factors): 

218 cor_loss += CosineSimilarity(tensors[i], tensors[j]) 

219 else: 

220 raise NotImplementedError(f"The independence loss type [{self.ind}] has not been supported.") 

221 return cor_loss 

222 

223 

224class KGIN(KnowledgeRecommender): 

225 r"""KGIN is a knowledge-aware recommendation model. It combines knowledge graph and the user-item interaction 

226 graph to a new graph called collaborative knowledge graph (CKG). This model explores intents behind a user-item 

227 interaction by using auxiliary item knowledge. 

228 """ 

229 

230 input_type = InputType.PAIRWISE 

231 

232 def __init__(self, config, dataset): 

233 super().__init__(config, dataset) 

234 

235 # load parameters info 

236 self.embedding_size = config["embedding_size"] 

237 self.n_factors = config["n_factors"] 

238 self.context_hops = config["context_hops"] 

239 self.node_dropout_rate = config["node_dropout_rate"] 

240 self.mess_dropout_rate = config["mess_dropout_rate"] 

241 self.ind = config["ind"] 

242 self.sim_decay = config["sim_regularity"] 

243 self.reg_weight = config["reg_weight"] 

244 self.temperature = config["temperature"] 

245 

246 # load dataset info 

247 # inter_matrix: [n_users, n_entities]; inter_graph: [n_users + n_entities, n_users + n_entities] 

248 self.interact_mat, _ = dataset._create_norm_ckg_adjacency_matrix(symmetric=False) 

249 self.interact_mat = self.interact_mat.to(self.device) 

250 self.kg_graph = dataset.kg_graph(form="coo", value_field="relation_id") # [n_entities, n_entities] 

251 # edge_index: [2, -1]; edge_type: [-1,] 

252 self.edge_index, self.edge_type = self.get_edges(self.kg_graph) 

253 

254 # define layers and loss 

255 self.n_nodes = self.n_users + self.n_entities 

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

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

258 self.latent_embedding = nn.Embedding(self.n_factors, self.embedding_size) 

259 self.gcn = GraphConv( 

260 embedding_size=self.embedding_size, 

261 n_hops=self.context_hops, 

262 n_users=self.n_users, 

263 n_relations=self.n_relations, 

264 n_factors=self.n_factors, 

265 edge_index=self.edge_index, 

266 edge_type=self.edge_type, 

267 interact_mat=self.interact_mat, 

268 ind=self.ind, 

269 tmp=self.temperature, 

270 device=self.device, 

271 node_dropout_rate=self.node_dropout_rate, 

272 mess_dropout_rate=self.mess_dropout_rate, 

273 ) 

274 self.mf_loss = BPRLoss() 

275 self.reg_loss = EmbLoss() 

276 self.restore_user_e = None 

277 self.restore_entity_e = None 

278 

279 # parameters initialization 

280 self.apply(xavier_uniform_initialization) 

281 

282 def get_edges(self, graph): 

283 index = torch.LongTensor(np.array([graph.row, graph.col])) 

284 type = torch.LongTensor(np.array(graph.data)) 

285 return index.to(self.device), type.to(self.device) 

286 

287 def forward(self): 

288 user_embeddings = self.user_embedding.weight 

289 entity_embeddings = self.entity_embedding.weight 

290 latent_embeddings = self.latent_embedding.weight 

291 # entity_gcn_emb: [n_entities, embedding_size] 

292 # user_gcn_emb: [n_users, embedding_size] 

293 # latent_gcn_emb: [n_factors, embedding_size] 

294 entity_gcn_emb, user_gcn_emb, cor_loss = self.gcn(user_embeddings, entity_embeddings, latent_embeddings) 

295 

296 return user_gcn_emb, entity_gcn_emb, cor_loss 

297 

298 def calculate_loss(self, interaction): 

299 r"""Calculate the training loss for a batch data of KG. 

300 

301 Args: 

302 interaction (Interaction): Interaction class of the batch. 

303 

304 Returns: 

305 torch.Tensor: Training loss, shape: [] 

306 """ 

307 if self.restore_user_e is not None or self.restore_entity_e is not None: 

308 self.restore_user_e, self.restore_entity_e = None, None 

309 

310 user = interaction[self.USER_ID] 

311 pos_item = interaction[self.ITEM_ID] 

312 neg_item = interaction[self.NEG_ITEM_ID] 

313 

314 user_all_embeddings, entity_all_embeddings, cor_loss = self.forward() 

315 u_embeddings = user_all_embeddings[user] 

316 pos_embeddings = entity_all_embeddings[pos_item] 

317 neg_embeddings = entity_all_embeddings[neg_item] 

318 

319 pos_scores = torch.mul(u_embeddings, pos_embeddings).sum(dim=1) 

320 neg_scores = torch.mul(u_embeddings, neg_embeddings).sum(dim=1) 

321 mf_loss = self.mf_loss(pos_scores, neg_scores) 

322 reg_loss = self.reg_loss(u_embeddings, pos_embeddings, neg_embeddings) 

323 cor_loss = self.sim_decay * cor_loss 

324 loss = mf_loss + self.reg_weight * reg_loss + cor_loss 

325 return loss 

326 

327 def predict(self, interaction): 

328 user = interaction[self.USER_ID] 

329 item = interaction[self.ITEM_ID] 

330 

331 user_all_embeddings, entity_all_embeddings, _ = self.forward() 

332 

333 u_embeddings = user_all_embeddings[user] 

334 i_embeddings = entity_all_embeddings[item] 

335 scores = torch.mul(u_embeddings, i_embeddings).sum(dim=1) 

336 return scores 

337 

338 def full_sort_predict(self, interaction): 

339 user = interaction[self.USER_ID] 

340 if self.restore_user_e is None or self.restore_entity_e is None: 

341 self.restore_user_e, self.restore_entity_e, _ = self.forward() 

342 u_embeddings = self.restore_user_e[user] 

343 i_embeddings = self.restore_entity_e[: self.n_items] 

344 

345 scores = torch.matmul(u_embeddings, i_embeddings.transpose(0, 1)) 

346 

347 return scores.view(-1)