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

1# @Time : 2020/8/6 

2# @Author : Shanlei Mu 

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

4 

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. 

10 

11Reference code: 

12 https://github.com/TaoMiner/joint-kg-recommender 

13""" 

14 

15import torch 

16import torch.nn.functional as F 

17from torch import nn 

18from torch.autograd import Variable 

19 

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 

24 

25 

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

31 

32 input_type = InputType.PAIRWISE 

33 

34 def __init__(self, config, dataset): 

35 super().__init__(config, dataset) 

36 

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

44 

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) 

53 

54 self.rec_loss = BPRLoss() 

55 self.kg_loss = nn.MarginRankingLoss(margin=self.margin) 

56 self.reg_loss = EmbMarginLoss() 

57 

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 

74 

75 def _masked_softmax(self, logits): 

76 probs = F.softmax(logits, dim=len(logits.shape) - 1) 

77 return probs 

78 

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. 

85 

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

92 

93 one_hot = Variable(indices.data.new(new_shape).zero_().scatter_(len(old_shape), indices.data, 1)) 

94 return one_hot 

95 

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. 

103 

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. 

109 

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 

122 

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 

143 

144 @staticmethod 

145 def _transH_projection(original, norm): 

146 return original - torch.sum(original * norm, dim=len(original.size()) - 1, keepdim=True) * norm 

147 

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 

154 

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 

160 

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) 

164 

165 return proj_user_e, relation_e, proj_item_e 

166 

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) 

173 

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) 

176 

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 ) 

183 

184 return rec_loss, orthogonal_loss, align_loss 

185 

186 def calculate_kg_loss(self, interaction): 

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

188 

189 Args: 

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

191 

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] 

199 

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) 

205 

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) 

209 

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) 

212 

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 ) 

222 

223 return loss, align_loss 

224 

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) 

230 

231 

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 ) 

237 

238 

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