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
« 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
5"""Analogy
6##################################################
7Reference:
8 Liu et al. "Analogical Inference for Multi-Relational Embeddings." in ICML 2017.
10Reference code:
11 https://github.com/torchkge-team/torchkge
12"""
14import torch
15from torch import nn
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
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.
28 Note:
29 In this version, we sample recommender data and knowledge data separately, and put them together for training.
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.device = config["device"]
40 self.scalar_share = config["scalar_share"]
41 self.ui_relation = dataset.field2token_id["relation_id"][dataset.ui_relation]
43 self.scalar_dim = int(self.embedding_size * self.scalar_share)
44 self.complex_dim = int(self.embedding_size - self.scalar_dim)
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)
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)
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)
59 # Loss
60 self.loss = LogisticLoss()
62 # Parameters initialization
63 self.apply(xavier_normal_initialization)
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)
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)
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)
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)
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)
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 )
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)
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)
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)
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)
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 )
136 def calculate_loss(self, interaction):
137 user = interaction[self.USER_ID]
139 pos_item = interaction[self.ITEM_ID]
140 neg_item = interaction[self.NEG_ITEM_ID]
142 relation = interaction[self.RELATION_ID]
144 head = interaction[self.HEAD_ENTITY_ID]
146 pos_tail = interaction[self.TAIL_ENTITY_ID]
147 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID]
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)
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 )
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
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)
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)
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)
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)
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)
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)
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)
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]
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)
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)
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)
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)
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]
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)
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)
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)
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)
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)
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)
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]
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)
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)
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)
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)