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
« 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
5# UPDATE:
6# @Time : 2022/8/31
7# @Author : Bowen Zheng
8# @Email : 18735382001@163.com
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"""
18import numpy as np
19import torch
20import torch.nn.functional as F
21from torch import nn
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
30class Aggregator(nn.Module):
31 """Relational Path-aware Convolution Network"""
33 def __init__(
34 self,
35 ):
36 super().__init__()
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
51 n_entities = entity_emb.shape[0]
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")
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]
67 return entity_agg, user_agg
70class GraphConv(nn.Module):
71 """Graph Convolutional Network"""
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__()
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
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
115 # parameters initialization
116 self.apply(xavier_uniform_initialization)
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]
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
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)
159 return (
160 entity_res_emb,
161 user_res_emb,
162 self.calculate_cor_loss(self.disen_weight_att),
163 )
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
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)
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
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
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 """
230 input_type = InputType.PAIRWISE
232 def __init__(self, config, dataset):
233 super().__init__(config, dataset)
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"]
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)
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
279 # parameters initialization
280 self.apply(xavier_uniform_initialization)
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)
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)
296 return user_gcn_emb, entity_gcn_emb, cor_loss
298 def calculate_loss(self, interaction):
299 r"""Calculate the training loss for a batch data of KG.
301 Args:
302 interaction (Interaction): Interaction class of the batch.
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
310 user = interaction[self.USER_ID]
311 pos_item = interaction[self.ITEM_ID]
312 neg_item = interaction[self.NEG_ITEM_ID]
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]
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
327 def predict(self, interaction):
328 user = interaction[self.USER_ID]
329 item = interaction[self.ITEM_ID]
331 user_all_embeddings, entity_all_embeddings, _ = self.forward()
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
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]
345 scores = torch.matmul(u_embeddings, i_embeddings.transpose(0, 1))
347 return scores.view(-1)