Coverage for hopwise/model/knowledge_graph_embedding_recommender/rotate.py: 72%
130 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/20
2# @Author : Alessandro Soccol
3# @Email : alessandro.soccol@unica.it
5"""RotatE
6##################################################
7Reference:
8 Sun et al. "RotatE: Knowledge Graph Embedding by Relational Rotation in Complex Space." in ICLR 2019.
10Reference code:
11 https://github.com/torchkge-team/torchkge
12"""
14import math
16import torch
17from torch import nn
19from hopwise.model.abstract_recommender import KnowledgeRecommender
20from hopwise.model.init import xavier_normal_initialization
21from hopwise.utils import InputType
24class RotatE(KnowledgeRecommender):
25 r"""RotatE models relations as rotations in a complex latent space with h, r, t belonging
26 to the set of d-dimensional complex numbers. The embedding for r belonging to the set of d-dimensional
27 complex numbers, is a rotation vector: in all its elements, the phase conveys the rotation along that axis,
28 and the modulus is equal to 1.
30 Note:
31 In this version, we sample recommender data and knowledge data separately, and put them together for training.
32 """
34 input_type = InputType.PAIRWISE
36 def __init__(self, config, dataset):
37 super().__init__(config, dataset)
39 # Load parameters info
40 self.embedding_size = config["embedding_size"]
41 self.margin = config["margin"]
42 self.device = config["device"]
43 self.ui_relation = dataset.field2token_id["relation_id"][dataset.ui_relation]
45 # Embeddings
46 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
47 self.user_embedding_im = nn.Embedding(self.n_users, self.embedding_size)
49 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size)
50 self.entity_embedding_im = nn.Embedding(self.n_entities, self.embedding_size)
52 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size)
54 # Loss
55 self.loss = nn.BCEWithLogitsLoss()
57 # Parameters initialization
58 self.apply(xavier_normal_initialization)
59 nn.init.uniform_(self.relation_embedding.weight, 0, 2 * math.pi)
61 def forward(self, head_re, head_im, relation, tail_re, tail_im):
62 rel_re, rel_im = torch.cos(relation), torch.sin(relation)
64 re_score = (rel_re * head_re - rel_im * head_im) - tail_re
65 im_score = (rel_re * head_im + rel_im * head_re) - tail_im
66 complex_score = torch.stack([re_score, im_score], dim=2)
67 score = torch.linalg.vector_norm(complex_score, dim=(1, 2))
69 return self.margin - score
71 def _get_rec_embeddings(self, user, positive_items, negative_items):
72 user_re = self.user_embedding(user)
73 user_im = self.user_embedding_im(user)
74 pos_item_re = self.entity_embedding(positive_items)
75 pos_item_im = self.entity_embedding_im(positive_items)
77 neg_item_re = self.entity_embedding(negative_items)
78 neg_item_im = self.entity_embedding_im(negative_items)
80 relation_user = torch.tensor([self.ui_relation] * user.shape[0], device=self.device)
81 rec_r_e = self.relation_embedding(relation_user)
83 return user_re, user_im, rec_r_e, pos_item_re, pos_item_im, neg_item_re, neg_item_im
85 def _get_kg_embeddings(self, head, relation, positive_tails, negative_tails):
86 head_re = self.entity_embedding(head)
87 head_im = self.entity_embedding_im(head)
88 pos_tail_re = self.entity_embedding(positive_tails)
89 pos_tail_im = self.entity_embedding_im(positive_tails)
91 neg_tail_re = self.entity_embedding(negative_tails)
92 neg_tail_im = self.entity_embedding_im(negative_tails)
94 kg_r_e = self.relation_embedding(relation)
96 return head_re, head_im, kg_r_e, pos_tail_re, pos_tail_im, neg_tail_re, neg_tail_im
98 def calculate_loss(self, interaction):
99 user = interaction[self.USER_ID]
101 pos_item = interaction[self.ITEM_ID]
102 neg_item = interaction[self.NEG_ITEM_ID]
104 head = interaction[self.HEAD_ENTITY_ID]
106 relation = interaction[self.RELATION_ID]
108 pos_tail = interaction[self.TAIL_ENTITY_ID]
109 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID]
111 user_re, user_im, rec_r_e, pos_item_re, pos_item_im, neg_item_re, neg_item_im = self._get_rec_embeddings(
112 user, pos_item, neg_item
113 )
114 head_re, head_im, kg_r_e, pos_tail_re, pos_tail_im, neg_tail_re, neg_tail_im = self._get_kg_embeddings(
115 head, relation, pos_tail, neg_tail
116 )
118 score_pos_users = self.forward(user_re, user_im, rec_r_e, pos_item_re, pos_item_im)
119 score_neg_users = self.forward(user_re, user_im, rec_r_e, neg_item_re, neg_item_im)
120 score_pos_kg = self.forward(head_re, head_im, kg_r_e, pos_tail_re, pos_tail_im)
121 score_neg_kg = self.forward(head_re, head_im, kg_r_e, neg_tail_re, neg_tail_im)
123 scores_rec = torch.cat([score_pos_users, score_neg_users], dim=0)
124 scores_kg = torch.cat([score_pos_kg, score_neg_kg], dim=0)
125 labels_rec = torch.cat([torch.ones_like(score_pos_users), torch.zeros_like(score_neg_users)], dim=0)
126 labels_kg = torch.cat([torch.ones_like(score_pos_kg), torch.zeros_like(score_neg_kg)], dim=0)
128 rec_loss = self.loss(scores_rec, labels_rec)
129 kg_loss = self.loss(scores_kg, labels_kg)
131 return rec_loss + kg_loss
133 def predict(self, interaction):
134 user = interaction[self.USER_ID]
135 item = interaction[self.ITEM_ID]
136 relation_user = torch.tensor([self.ui_relation] * user.shape[0], device=self.device)
138 user_re = self.user_embedding(user)
139 user_im = self.user_embedding_im(user)
140 item_re = self.entity_embedding(item)
141 item_im = self.entity_embedding_im(item)
143 rec_r_e = self.relation_embedding(relation_user)
145 return self.forward(user_re, user_im, rec_r_e, item_re, item_im)
147 def predict_kg(self, interaction):
148 head = interaction[self.HEAD_ENTITY_ID]
149 relation = interaction[self.RELATION_ID]
150 tail = interaction[self.TAIL_ENTITY_ID]
152 head_re = self.entity_embedding(head)
153 head_im = self.entity_embedding_im(head)
154 tail_re = self.entity_embedding(tail)
155 tail_im = self.entity_embedding_im(tail)
157 rec_r_e = self.relation_embedding(relation)
159 return self.forward(head_re, head_im, rec_r_e, tail_re, tail_im)
161 def full_sort_predict(self, interaction):
162 user = interaction[self.USER_ID]
163 relation_user = torch.tensor([self.ui_relation] * user.shape[0], device=self.device)
164 item_indices = torch.tensor(range(self.n_items)).to(self.device)
166 user_re = self.user_embedding(user)
167 user_im = self.user_embedding_im(user)
169 item_re = self.entity_embedding(item_indices)
170 item_im = self.entity_embedding_im(item_indices)
172 rel_theta = self.relation_embedding(relation_user)
174 rel_re, rel_im = torch.cos(rel_theta), torch.sin(rel_theta)
176 user_re = user_re.unsqueeze(1).expand(-1, item_indices.shape[0], -1)
177 user_im = user_im.unsqueeze(1).expand(-1, item_indices.shape[0], -1)
179 rel_re = rel_re.unsqueeze(1).expand(-1, item_indices.shape[0], -1)
180 rel_im = rel_im.unsqueeze(1).expand(-1, item_indices.shape[0], -1)
182 item_re = item_re.unsqueeze(0)
183 item_im = item_im.unsqueeze(0)
185 re_score = (rel_re * user_re - rel_im * user_im) - item_re
186 im_score = (rel_re * user_im + rel_im * user_re) - item_im
187 complex_score = torch.stack([re_score, im_score], dim=3)
188 score = torch.linalg.vector_norm(complex_score, dim=(2, 3))
190 return self.margin - score
192 def full_sort_predict_kg(self, interaction):
193 head = interaction[self.HEAD_ENTITY_ID]
194 relation = interaction[self.RELATION_ID]
196 head_re = self.entity_embedding(head)
197 head_im = self.entity_embedding_im(head)
199 tail_re = self.entity_embedding.weight
200 tail_im = self.entity_embedding_im.weight
202 rel_theta = self.relation_embedding(relation)
204 rel_re, rel_im = torch.cos(rel_theta), torch.sin(rel_theta)
206 head_re = head_re.unsqueeze(1).expand(-1, self.entity_embedding.weight.shape[0], -1)
207 head_im = head_im.unsqueeze(1).expand(-1, self.entity_embedding.weight.shape[0], -1)
209 rel_re = rel_re.unsqueeze(1).expand(-1, self.entity_embedding.weight.shape[0], -1)
210 rel_im = rel_im.unsqueeze(1).expand(-1, self.entity_embedding.weight.shape[0], -1)
212 tail_re = tail_re.unsqueeze(0)
213 tail_im = tail_im.unsqueeze(0)
215 re_score = (rel_re * head_re - rel_im * head_im) - tail_re
216 im_score = (rel_re * head_im + rel_im * head_re) - tail_im
217 complex_score = torch.stack([re_score, im_score], dim=3)
218 score = torch.linalg.vector_norm(complex_score, dim=(2, 3))
220 return self.margin - score