Coverage for hopwise/model/knowledge_graph_embedding_recommender/analogy.py: 67%

154 statements  

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

1# @Time : 2024/11/19 

2# @Author : Alessandro Soccol 

3# @Email : alessandro.soccol@unica.it 

4 

5"""Analogy 

6################################################## 

7Reference: 

8 Liu et al. "Analogical Inference for Multi-Relational Embeddings." in ICML 2017. 

9 

10Reference code: 

11 https://github.com/torchkge-team/torchkge 

12""" 

13 

14import torch 

15from torch import nn 

16 

17from hopwise.model.abstract_recommender import KnowledgeRecommender 

18from hopwise.model.init import xavier_normal_initialization 

19from hopwise.model.loss import LogisticLoss 

20from hopwise.utils import InputType 

21 

22 

23class Analogy(KnowledgeRecommender): 

24 r"""Analogy extends RESCAL so as to further model the analogical properties of entities and relations e.g. 

25 Interstellar is to Fantasy as Nolan is to Oppenheimer”. 

26 It employs the same scoring function as RESCAL but with some constraints. 

27 

28 Note: 

29 In this version, we sample recommender data and knowledge data separately, and put them together for training. 

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.device = config["device"] 

40 self.scalar_share = config["scalar_share"] 

41 self.ui_relation = dataset.field2token_id["relation_id"][dataset.ui_relation] 

42 

43 self.scalar_dim = int(self.embedding_size * self.scalar_share) 

44 self.complex_dim = int(self.embedding_size - self.scalar_dim) 

45 

46 # Embeddings 

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

48 self.user_re_embedding = nn.Embedding(self.n_users, self.embedding_size) 

49 self.user_im_embedding = nn.Embedding(self.n_users, self.embedding_size) 

50 

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

52 self.entity_re_embedding = nn.Embedding(self.n_entities, self.embedding_size) 

53 self.entity_im_embedding = nn.Embedding(self.n_entities, self.embedding_size) 

54 

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

56 self.relation_re_embedding = nn.Embedding(self.n_relations, self.embedding_size) 

57 self.relation_im_embedding = nn.Embedding(self.n_relations, self.embedding_size) 

58 

59 # Loss 

60 self.loss = LogisticLoss() 

61 

62 # Parameters initialization 

63 self.apply(xavier_normal_initialization) 

64 

65 def forward(self, head_e, head_re_e, head_im_e, r_e, r_re_e, r_im_e, tail_e, tail_re_e, tail_im_e): 

66 return (head_e * r_e * tail_e).sum(dim=1) + ( 

67 head_re_e * (r_re_e * tail_re_e + r_im_e * tail_im_e) 

68 + head_im_e * (r_re_e * tail_im_e - r_im_e * tail_re_e) 

69 ).sum(dim=1) 

70 

71 def _get_rec_embeddings(self, users, pos_items, neg_items): 

72 user_e = self.user_embedding(users) 

73 user_re_e = self.user_re_embedding(users) 

74 user_im_e = self.user_im_embedding(users) 

75 

76 pos_item_e = self.entity_embedding(pos_items) 

77 pos_item_re_e = self.entity_re_embedding(pos_items) 

78 pos_item_im_e = self.entity_im_embedding(pos_items) 

79 

80 neg_item_e = self.entity_embedding(neg_items) 

81 neg_item_re_e = self.entity_re_embedding(neg_items) 

82 neg_item_im_e = self.entity_im_embedding(neg_items) 

83 

84 relations = torch.tensor([self.ui_relation] * users.shape[0], device=self.device) 

85 rec_r_e = self.relation_embedding(relations) 

86 rec_r_re_e = self.relation_re_embedding(relations) 

87 rec_r_im_e = self.relation_im_embedding(relations) 

88 

89 return ( 

90 user_e, 

91 user_re_e, 

92 user_im_e, 

93 pos_item_e, 

94 pos_item_re_e, 

95 pos_item_im_e, 

96 neg_item_e, 

97 neg_item_re_e, 

98 neg_item_im_e, 

99 rec_r_e, 

100 rec_r_re_e, 

101 rec_r_im_e, 

102 ) 

103 

104 def _get_kg_embeddings(self, heads, relations, pos_tails, neg_tails): 

105 head_e = self.entity_embedding(heads) 

106 head_re_e = self.entity_re_embedding(heads) 

107 head_im_e = self.entity_im_embedding(heads) 

108 

109 neg_tail_e = self.entity_embedding(neg_tails) 

110 neg_tail_re_e = self.entity_re_embedding(neg_tails) 

111 neg_tail_im_e = self.entity_im_embedding(neg_tails) 

112 

113 pos_tail_e = self.entity_embedding(pos_tails) 

114 pos_tail_re_e = self.entity_re_embedding(pos_tails) 

115 pos_tail_im_e = self.entity_im_embedding(pos_tails) 

116 

117 r_e = self.relation_embedding(relations) 

118 r_re_e = self.relation_re_embedding(relations) 

119 r_im_e = self.relation_im_embedding(relations) 

120 

121 return ( 

122 head_e, 

123 head_re_e, 

124 head_im_e, 

125 pos_tail_e, 

126 pos_tail_re_e, 

127 pos_tail_im_e, 

128 neg_tail_e, 

129 neg_tail_re_e, 

130 neg_tail_im_e, 

131 r_e, 

132 r_re_e, 

133 r_im_e, 

134 ) 

135 

136 def calculate_loss(self, interaction): 

137 user = interaction[self.USER_ID] 

138 

139 pos_item = interaction[self.ITEM_ID] 

140 neg_item = interaction[self.NEG_ITEM_ID] 

141 

142 relation = interaction[self.RELATION_ID] 

143 

144 head = interaction[self.HEAD_ENTITY_ID] 

145 

146 pos_tail = interaction[self.TAIL_ENTITY_ID] 

147 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID] 

148 

149 ( 

150 user_e, 

151 user_re_e, 

152 user_im_e, 

153 pos_item_e, 

154 pos_item_re_e, 

155 pos_item_im_e, 

156 neg_item_e, 

157 neg_item_re_e, 

158 neg_item_im_e, 

159 rec_r_e, 

160 rec_r_re_e, 

161 rec_r_im_e, 

162 ) = self._get_rec_embeddings(user, pos_item, neg_item) 

163 ( 

164 head_e, 

165 head_re_e, 

166 head_im_e, 

167 pos_tail_e, 

168 pos_tail_re_e, 

169 pos_tail_im_e, 

170 neg_tail_e, 

171 neg_tail_re_e, 

172 neg_tail_im_e, 

173 r_e, 

174 r_re_e, 

175 r_im_e, 

176 ) = self._get_kg_embeddings(head, relation, pos_tail, neg_tail) 

177 

178 score_pos_users = self.forward( 

179 user_e, user_re_e, user_im_e, rec_r_e, rec_r_re_e, rec_r_im_e, pos_item_e, pos_item_re_e, pos_item_im_e 

180 ) 

181 score_neg_users = self.forward( 

182 user_e, user_re_e, user_im_e, rec_r_e, rec_r_re_e, rec_r_im_e, neg_item_e, neg_item_re_e, neg_item_im_e 

183 ) 

184 score_pos_kg = self.forward( 

185 head_e, head_re_e, head_im_e, r_e, r_re_e, r_im_e, pos_tail_e, pos_tail_re_e, pos_tail_im_e 

186 ) 

187 score_neg_kg = self.forward( 

188 head_e, head_re_e, head_im_e, r_e, r_re_e, r_im_e, neg_tail_e, neg_tail_re_e, neg_tail_im_e 

189 ) 

190 

191 rec_loss = self.loss(-score_pos_users, -score_neg_users) 

192 kg_loss = self.loss(-score_pos_kg, -score_neg_kg) 

193 return rec_loss + kg_loss 

194 

195 def predict(self, interaction): 

196 user = interaction[self.USER_ID] 

197 item = interaction[self.ITEM_ID] 

198 relation = torch.tensor([self.ui_relation] * user.shape[0], device=self.device) 

199 

200 user_e = self.user_embedding(user) 

201 user_re_e = self.user_re_embedding(user) 

202 user_im_e = self.user_im_embedding(user) 

203 

204 r_e = self.relation_embedding(relation) 

205 r_re_e = self.relation_re_embedding(relation) 

206 r_im_e = self.relation_im_embedding(relation) 

207 

208 item_e = self.entity_embedding(item) 

209 item_re_e = self.entity_re_embedding(item) 

210 item_im_e = self.entity_im_embedding(item) 

211 

212 return self.forward(user_e, user_re_e, user_im_e, r_e, r_re_e, r_im_e, item_e, item_re_e, item_im_e) 

213 

214 def full_sort_predict(self, interaction): 

215 user = interaction[self.USER_ID] 

216 user_e = self.user_embedding(user) 

217 user_re_e = self.user_re_embedding(user) 

218 user_im_e = self.user_im_embedding(user) 

219 

220 rec_r_e = self.relation_embedding.weight[-1] 

221 rec_r_re_e = self.relation_re_embedding.weight[-1] 

222 rec_r_im_e = self.relation_im_embedding.weight[-1] 

223 rec_r_e = rec_r_e.expand_as(user_e) 

224 rec_r_re_e = rec_r_re_e.expand_as(user_e) 

225 rec_r_im_e = rec_r_im_e.expand_as(user_e) 

226 

227 item_indices = torch.tensor(range(self.n_items)).to(self.device) 

228 all_item_e = self.entity_embedding.weight[item_indices] 

229 all_item_re_e = self.entity_re_embedding.weight[item_indices] 

230 all_item_im_e = self.entity_im_embedding.weight[item_indices] 

231 

232 user_e = user_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1) 

233 user_re_e = user_re_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1) 

234 user_im_e = user_im_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1) 

235 

236 rec_r_e = rec_r_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1) 

237 rec_r_re_e = rec_r_re_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1) 

238 rec_r_im_e = rec_r_im_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1) 

239 

240 all_item_e = all_item_e.unsqueeze(0) 

241 all_item_re_e = all_item_re_e.unsqueeze(0) 

242 all_item_im_e = all_item_im_e.unsqueeze(0) 

243 

244 return (user_e * rec_r_e * all_item_e).sum(dim=-1) + ( 

245 user_re_e * (rec_r_re_e * all_item_re_e + rec_r_im_e * all_item_im_e) 

246 + user_im_e * (rec_r_re_e * all_item_im_e - rec_r_im_e * all_item_re_e) 

247 ).sum(dim=-1) 

248 

249 def predict_kg(self, interaction): 

250 head = interaction[self.HEAD_ENTITY_ID] 

251 relation = interaction[self.RELATION_ID] 

252 tail = interaction[self.TAIL_ENTITY_ID] 

253 

254 head_e = self.entity_embedding(head) 

255 head_re_e = self.entity_re_embedding(head) 

256 head_im_e = self.entity_im_embedding(head) 

257 

258 r_e = self.relation_embedding(relation) 

259 r_re_e = self.relation_re_embedding(relation) 

260 r_im_e = self.relation_im_embedding(relation) 

261 

262 tail_e = self.entity_embedding(tail) 

263 tail_re_e = self.entity_re_embedding(tail) 

264 tail_im_e = self.entity_im_embedding(tail) 

265 

266 return self.forward(head_e, head_re_e, head_im_e, r_e, r_re_e, r_im_e, tail_e, tail_re_e, tail_im_e) 

267 

268 def full_sort_predict_kg(self, interaction): 

269 head = interaction[self.HEAD_ENTITY_ID] 

270 relation = interaction[self.RELATION_ID] 

271 head_e = self.entity_embedding(head) 

272 head_re_e = self.entity_re_embedding(head) 

273 head_im_e = self.entity_im_embedding(head) 

274 

275 rec_r_e = self.relation_embedding(relation) 

276 rec_r_re_e = self.relation_re_embedding(relation) 

277 rec_r_im_e = self.relation_im_embedding(relation) 

278 rec_r_e = rec_r_e.expand_as(head_e) 

279 rec_r_re_e = rec_r_re_e.expand_as(head_e) 

280 rec_r_im_e = rec_r_im_e.expand_as(head_e) 

281 

282 entity_indices = torch.tensor(range(self.n_entities)).to(self.device) 

283 all_entities_e = self.entity_embedding.weight[entity_indices] 

284 all_entities_re_e = self.entity_re_embedding.weight[entity_indices] 

285 all_entities_im_e = self.entity_im_embedding.weight[entity_indices] 

286 

287 head_e = head_e.unsqueeze(1).expand(-1, all_entities_e.shape[0], -1) 

288 head_re_e = head_re_e.unsqueeze(1).expand(-1, all_entities_e.shape[0], -1) 

289 head_im_e = head_im_e.unsqueeze(1).expand(-1, all_entities_e.shape[0], -1) 

290 

291 rec_r_e = rec_r_e.unsqueeze(1).expand(-1, all_entities_e.shape[0], -1) 

292 rec_r_re_e = rec_r_re_e.unsqueeze(1).expand(-1, all_entities_e.shape[0], -1) 

293 rec_r_im_e = rec_r_im_e.unsqueeze(1).expand(-1, all_entities_e.shape[0], -1) 

294 

295 all_entities_e = all_entities_e.unsqueeze(0) 

296 all_entities_re_e = all_entities_re_e.unsqueeze(0) 

297 all_entities_im_e = all_entities_im_e.unsqueeze(0) 

298 

299 return (head_e * rec_r_e * all_entities_e).sum(dim=-1) + ( 

300 head_re_e * (rec_r_re_e * all_entities_re_e + rec_r_im_e * all_entities_im_e) 

301 + head_im_e * (rec_r_re_e * all_entities_im_e - rec_r_im_e * all_entities_re_e) 

302 ).sum(dim=-1)