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