Coverage for hopwise/model/knowledge_aware_recommender/userktup.py: 0%
141 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/8/6
2# @Author : Shanlei Mu
3# @Email : slmu@ruc.edu.cn
5r"""
6KTUP
7##################################################
8Reference:
9 Yixin Cao et al. "Unifying Knowledge Graph Learning and Recommendation:Towards a Better Understanding
10 of User Preferences." in WWW 2019.
12Reference code:
13 https://github.com/TaoMiner/joint-kg-recommender
14"""
16import torch
17import torch.nn.functional as F
18from torch import nn
19from torch.autograd import Variable
21from hopwise.model.abstract_recommender import KnowledgeRecommender
22from hopwise.model.init import xavier_uniform_initialization
23from hopwise.model.loss import BPRLoss, EmbMarginLoss
24from hopwise.utils import InputType
27class UserKTUP(KnowledgeRecommender):
28 r"""KTUP is a knowledge-based recommendation model. It adopts the strategy of multi-task learning to jointly learn
29 recommendation and KG-related tasks, with the goal of understanding the reasons that a user interacts with an item.
30 This method utilizes an attention mechanism to combine all preferences into a single-vector representation.
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 self.L1_flag = config["L1_flag"]
41 self.use_st_gumbel = config["use_st_gumbel"]
42 self.kg_weight = config["kg_weight"]
43 self.align_weight = config["align_weight"]
44 self.margin = config["margin"]
46 # define layers and loss
47 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
48 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size)
49 self.pref_embedding = nn.Embedding(self.n_relations, self.embedding_size)
50 self.pref_norm_embedding = nn.Embedding(self.n_relations, self.embedding_size)
51 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size)
52 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size)
53 self.relation_norm_embedding = nn.Embedding(self.n_relations, self.embedding_size)
55 self.rec_loss = BPRLoss()
56 self.kg_loss = nn.MarginRankingLoss(margin=self.margin)
57 self.reg_loss = EmbMarginLoss()
59 # parameters initialization
60 self.apply(xavier_uniform_initialization)
61 normalize_user_emb = F.normalize(self.user_embedding.weight.data, p=2, dim=1)
62 normalize_item_emb = F.normalize(self.item_embedding.weight.data, p=2, dim=1)
63 normalize_pref_emb = F.normalize(self.pref_embedding.weight.data, p=2, dim=1)
64 normalize_pref_norm_emb = F.normalize(self.pref_norm_embedding.weight.data, p=2, dim=1)
65 normalize_entity_emb = F.normalize(self.entity_embedding.weight.data, p=2, dim=1)
66 normalize_rel_emb = F.normalize(self.relation_embedding.weight.data, p=2, dim=1)
67 normalize_rel_norm_emb = F.normalize(self.relation_norm_embedding.weight.data, p=2, dim=1)
68 self.user_embedding.weight.data = normalize_user_emb
69 self.item_embedding.weight_data = normalize_item_emb
70 self.pref_embedding.weight.data = normalize_pref_emb
71 self.pref_norm_embedding.weight.data = normalize_pref_norm_emb
72 self.entity_embedding.weight.data = normalize_entity_emb
73 self.relation_embedding.weight.data = normalize_rel_emb
74 self.relation_norm_embedding.weight.data = normalize_rel_norm_emb
76 def _masked_softmax(self, logits):
77 probs = F.softmax(logits, dim=len(logits.shape) - 1)
78 return probs
80 def convert_to_one_hot(self, indices, num_classes):
81 r"""
82 Args:
83 indices (Variable): A vector containing indices,
84 whose size is (batch_size,).
85 num_classes (Variable): The number of classes, which would be
86 the second dimension of the resulting one-hot matrix.
88 Returns:
89 torch.Tensor: The one-hot matrix of size (batch_size, num_classes).
90 """
92 old_shape = indices.shape
93 new_shape = torch.Size([i for i in old_shape] + [num_classes])
94 indices = indices.unsqueeze(len(old_shape))
96 one_hot = Variable(indices.data.new(new_shape).zero_().scatter_(len(old_shape), indices.data, 1))
97 return one_hot
99 def st_gumbel_softmax(self, logits, temperature=1.0):
100 r"""Return the result of Straight-Through Gumbel-Softmax Estimation.
101 It approximates the discrete sampling via Gumbel-Softmax trick
102 and applies the biased ST estimator.
103 In the forward propagation, it emits the discrete one-hot result,
104 and in the backward propagation it approximates the categorical
105 distribution via smooth Gumbel-Softmax distribution.
107 Args:
108 logits (Variable): A un-normalized probability values,
109 which has the size (batch_size, num_classes)
110 temperature (float): A temperature parameter. The higher
111 the value is, the smoother the distribution is.
113 Returns:
114 torch.Tensor: The sampled output, which has the property explained above.
115 """
117 eps = 1e-20
118 u = logits.data.new(*logits.size()).uniform_()
119 gumbel_noise = Variable(-torch.log(-torch.log(u + eps) + eps))
120 y = logits + gumbel_noise
121 y = self._masked_softmax(logits=y / temperature)
122 y_argmax = y.max(len(y.shape) - 1)[1]
123 y_hard = self.convert_to_one_hot(indices=y_argmax, num_classes=y.size(len(y.shape) - 1)).float()
124 y = (y_hard - y).detach() + y
125 return y
127 def _get_preferences(self, user_e, item_e, use_st_gumbel=False):
128 pref_probs = (
129 torch.matmul(
130 user_e + item_e,
131 torch.t(self.pref_embedding.weight + self.relation_embedding.weight),
132 )
133 / 2
134 )
135 if use_st_gumbel:
136 # todo: different torch versions may cause the st_gumbel_softmax to report errors, wait to be test
137 pref_probs = self.st_gumbel_softmax(pref_probs)
138 relation_e = torch.matmul(pref_probs, self.pref_embedding.weight + self.relation_embedding.weight) / 2
139 norm_e = (
140 torch.matmul(
141 pref_probs,
142 self.pref_norm_embedding.weight + self.relation_norm_embedding.weight,
143 )
144 / 2
145 )
146 return pref_probs, relation_e, norm_e
148 @staticmethod
149 def _transH_projection(original, norm):
150 return original - torch.sum(original * norm, dim=len(original.size()) - 1, keepdim=True) * norm
152 def _get_score(self, h_e, r_e, t_e):
153 if self.L1_flag:
154 score = -torch.sum(torch.abs(h_e + r_e - t_e), 1)
155 else:
156 score = -torch.sum((h_e + r_e - t_e) ** 2, 1)
157 return score
159 def forward(self, user, item):
160 user_e = self.user_embedding(user)
161 item_e = self.item_embedding(item)
162 user_entity_e = self.entity_embedding(user)
163 item_entity_e = self.entity_embedding(self.n_users + item)
164 user_e = user_e + user_entity_e
165 item_e = item_e + item_entity_e
167 _, relation_e, norm_e = self._get_preferences(user_e, item_e, use_st_gumbel=self.use_st_gumbel)
168 proj_user_e = self._transH_projection(user_e, norm_e)
169 proj_item_e = self._transH_projection(item_e, norm_e)
171 return proj_user_e, relation_e, proj_item_e
173 def calculate_loss(self, interaction):
174 user = interaction[self.USER_ID]
175 pos_item = interaction[self.ITEM_ID]
176 neg_item = interaction[self.NEG_ITEM_ID]
177 proj_pos_user_e, pos_relation_e, proj_pos_item_e = self.forward(user, pos_item)
178 proj_neg_user_e, neg_relation_e, proj_neg_item_e = self.forward(user, neg_item)
180 pos_item_score = self._get_score(proj_pos_user_e, pos_relation_e, proj_pos_item_e)
181 neg_item_score = self._get_score(proj_neg_user_e, neg_relation_e, proj_neg_item_e)
183 rec_loss = self.rec_loss(pos_item_score, neg_item_score)
184 orthogonal_loss = orthogonalLoss(self.pref_embedding.weight, self.pref_norm_embedding.weight)
185 item = torch.cat([pos_item, neg_item])
186 item_align_loss = self.align_weight * alignLoss(
187 self.item_embedding(item), self.entity_embedding(item + self.n_users), self.L1_flag
188 )
189 user_align_loss = self.align_weight * alignLoss(
190 self.user_embedding(user), self.entity_embedding(user), self.L1_flag
191 )
192 align_loss = item_align_loss + user_align_loss
194 return rec_loss, orthogonal_loss, align_loss
196 def calculate_kg_loss(self, interaction):
197 r"""Calculate the training loss for a batch data of KG.
199 Args:
200 interaction (Interaction): Interaction class of the batch.
202 Returns:
203 torch.Tensor: Training loss, shape: []
204 """
206 h = interaction[self.HEAD_ENTITY_ID]
207 r = interaction[self.RELATION_ID]
208 pos_t = interaction[self.TAIL_ENTITY_ID]
209 neg_t = interaction[self.NEG_TAIL_ENTITY_ID]
211 h_e = self.entity_embedding(h)
212 pos_t_e = self.entity_embedding(pos_t)
213 neg_t_e = self.entity_embedding(neg_t)
214 r_e = self.relation_embedding(r)
215 norm_e = self.relation_norm_embedding(r)
217 proj_h_e = self._transH_projection(h_e, norm_e)
218 proj_pos_t_e = self._transH_projection(pos_t_e, norm_e)
219 proj_neg_t_e = self._transH_projection(neg_t_e, norm_e)
221 pos_tail_score = self._get_score(proj_h_e, r_e, proj_pos_t_e)
222 neg_tail_score = self._get_score(proj_h_e, r_e, proj_neg_t_e)
224 kg_loss = self.kg_loss(pos_tail_score, neg_tail_score, torch.ones(h.size(0)).to(self.device))
225 orthogonal_loss = orthogonalLoss(r_e, norm_e)
226 reg_loss = self.reg_loss(h_e, pos_t_e, neg_t_e, r_e)
227 loss = self.kg_weight * (kg_loss + orthogonal_loss + reg_loss)
228 entity = torch.cat([h, pos_t, neg_t])
229 item_entity = entity[torch.logical_and(entity >= self.n_users, entity < self.n_users + self.n_items)]
230 user_entity = entity[entity < self.n_users]
231 item_align_loss = self.align_weight * alignLoss(
232 self.item_embedding(item_entity - self.n_users), self.entity_embedding(item_entity), self.L1_flag
233 )
234 user_align_loss = self.align_weight * alignLoss(
235 self.user_embedding(user_entity), self.entity_embedding(user_entity), self.L1_flag
236 )
237 align_loss = item_align_loss + user_align_loss
239 return loss, align_loss
241 def predict(self, interaction):
242 user = interaction[self.USER_ID]
243 item = interaction[self.ITEM_ID]
244 proj_user_e, relation_e, proj_item_e = self.forward(user, item)
245 return self._get_score(proj_user_e, relation_e, proj_item_e)
248def orthogonalLoss(rel_embeddings, norm_embeddings):
249 return torch.sum(
250 torch.sum(norm_embeddings * rel_embeddings, dim=1, keepdim=True) ** 2
251 / torch.sum(rel_embeddings**2, dim=1, keepdim=True)
252 )
255def alignLoss(emb1, emb2, L1_flag=False):
256 if L1_flag:
257 distance = torch.sum(torch.abs(emb1 - emb2), 1)
258 else:
259 distance = torch.sum((emb1 - emb2) ** 2, 1)
260 return distance.mean()