Coverage for hopwise/model/knowledge_graph_embedding_recommender/complex.py: 71%
123 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"""ComplEx
6##################################################
7Reference:
8 Trouillon et al. "Complex embeddings for simple link prediction." in ICML'16.
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.utils import InputType
22class ComplEx(KnowledgeRecommender):
23 r"""ComplEx extends DistMult by introducing complex-valued embeddings.
25 Note:
26 In this version, we sample recommender data and knowledge data separately, and put them together for training.
27 """
29 input_type = InputType.PAIRWISE
31 def __init__(self, config, dataset):
32 super().__init__(config, dataset)
34 # load parameters info
35 self.embedding_size = config["embedding_size"]
36 self.device = config["device"]
37 self.ui_relation = dataset.field2token_id["relation_id"][dataset.ui_relation]
38 # define layers and loss
39 self.user_re_embedding = nn.Embedding(self.n_users, self.embedding_size)
40 self.user_im_embedding = nn.Embedding(self.n_users, self.embedding_size)
42 self.entity_re_embedding = nn.Embedding(self.n_entities, self.embedding_size)
43 self.entity_im_embedding = nn.Embedding(self.n_entities, self.embedding_size)
45 self.relation_re_embedding = nn.Embedding(self.n_relations, self.embedding_size)
46 self.relation_im_embedding = nn.Embedding(self.n_relations, self.embedding_size)
48 self.loss = nn.BCEWithLogitsLoss()
50 # parameters initialization
51 self.apply(xavier_normal_initialization)
53 def forward(self, head_re_e, head_im_e, rec_r_re_e, rec_r_im_e, tail_re_e, tail_im_e):
54 return (
55 self.triple_dot(head_re_e, rec_r_re_e, tail_re_e)
56 + self.triple_dot(head_im_e, rec_r_re_e, tail_im_e)
57 + self.triple_dot(head_re_e, rec_r_im_e, tail_im_e)
58 - self.triple_dot(head_im_e, rec_r_im_e, tail_im_e)
59 )
61 def triple_dot(self, x, y, z):
62 return (x * y * z).sum(dim=-1)
64 def _get_rec_embeddings(self, user, positive_items, negative_items):
65 relation_users = torch.tensor([self.ui_relation] * user.shape[0], device=self.device)
66 user_re_e = self.user_re_embedding(user)
67 user_im_e = self.user_im_embedding(user)
69 pos_item_re_e = self.entity_re_embedding(positive_items)
70 pos_item_im_e = self.entity_im_embedding(positive_items)
72 neg_item_re_e = self.entity_re_embedding(negative_items)
73 neg_item_im_e = self.entity_im_embedding(negative_items)
75 rec_r_re_e = self.relation_re_embedding(relation_users)
76 rec_r_im_e = self.relation_im_embedding(relation_users)
78 return user_re_e, user_im_e, rec_r_re_e, rec_r_im_e, pos_item_re_e, pos_item_im_e, neg_item_re_e, neg_item_im_e
80 def _get_kg_embeddings(self, head, relation, positive_tails, negative_tails):
81 head_re_e = self.entity_re_embedding(head)
82 head_im_e = self.entity_im_embedding(head)
84 pos_tail_re_e = self.entity_re_embedding(positive_tails)
85 pos_tail_im_e = self.entity_im_embedding(positive_tails)
87 neg_tail_re_e = self.entity_re_embedding(negative_tails)
88 neg_tail_im_e = self.entity_im_embedding(negative_tails)
90 kg_r_re_e = self.relation_re_embedding(relation)
91 kg_r_im_e = self.relation_im_embedding(relation)
93 return head_re_e, head_im_e, kg_r_re_e, kg_r_im_e, pos_tail_re_e, pos_tail_im_e, neg_tail_re_e, neg_tail_im_e
95 def calculate_loss(self, interaction):
96 user = interaction[self.USER_ID]
98 pos_item = interaction[self.ITEM_ID]
99 neg_item = interaction[self.NEG_ITEM_ID]
101 relation = interaction[self.RELATION_ID]
103 head = interaction[self.HEAD_ENTITY_ID]
105 pos_tail = interaction[self.TAIL_ENTITY_ID]
106 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID]
108 user_re_e, user_im_e, rec_r_re_e, rec_r_im_e, pos_item_re_e, pos_item_im_e, neg_item_re_e, neg_item_im_e = (
109 self._get_rec_embeddings(user, pos_item, neg_item)
110 )
111 head_re_e, head_im_e, kg_r_re_e, kg_r_im_e, pos_tail_re_e, pos_tail_im_e, neg_tail_re_e, neg_tail_im_e = (
112 self._get_kg_embeddings(head, relation, pos_tail, neg_tail)
113 )
115 score_pos_users = self.forward(user_re_e, user_im_e, rec_r_re_e, rec_r_im_e, pos_item_re_e, pos_item_im_e)
116 score_neg_users = self.forward(user_re_e, user_im_e, rec_r_re_e, rec_r_im_e, neg_item_re_e, neg_item_im_e)
117 score_pos_kg = self.forward(head_re_e, head_im_e, kg_r_re_e, kg_r_im_e, pos_tail_re_e, pos_tail_im_e)
118 score_neg_kg = self.forward(head_re_e, head_im_e, kg_r_re_e, kg_r_im_e, neg_tail_re_e, neg_tail_im_e)
120 scores_rec = torch.cat([score_pos_users, score_neg_users], dim=0)
121 scores_kg = torch.cat([score_pos_kg, score_neg_kg], dim=0)
122 labels_rec = torch.cat([torch.ones_like(score_pos_users), torch.zeros_like(score_neg_users)], dim=0)
123 labels_kg = torch.cat([torch.ones_like(score_pos_kg), torch.zeros_like(score_neg_kg)], dim=0)
125 rec_loss = self.loss(scores_rec, labels_rec)
126 kg_loss = self.loss(scores_kg, labels_kg)
128 return rec_loss + kg_loss
130 def predict(self, interaction):
131 user = interaction[self.USER_ID]
132 item = interaction[self.ITEM_ID]
133 relation = torch.tensor([self.ui_relation] * user.shape[0], device=self.device)
135 user_re_e = self.user_re_embedding(user)
136 user_im_e = self.user_im_embedding(user)
138 item_re_e = self.entity_re_embedding(item)
139 item_im_e = self.entity_im_embedding(item)
141 rec_r_re_e = self.relation_re_embedding(relation)
142 rec_r_im_e = self.relation_im_embedding(relation)
144 return self.forward(user_re_e, user_im_e, rec_r_re_e, rec_r_im_e, item_re_e, item_im_e)
146 def predict_kg(self, interaction):
147 head = interaction[self.HEAD_ENTITY_ID]
148 relation = interaction[self.RELATION_ID]
149 tail = interaction[self.TAIL_ENTITY_ID]
151 head_re_e = self.entity_re_embedding(head)
152 head_im_e = self.entity_im_embedding(head)
154 tail_re_e = self.entity_re_embedding(tail)
155 tail_im_e = self.entity_im_embedding(tail)
157 rec_r_re_e = self.relation_re_embedding(relation)
158 rec_r_im_e = self.relation_im_embedding(relation)
160 return self.forward(head_re_e, head_im_e, rec_r_re_e, rec_r_im_e, tail_re_e, tail_im_e)
162 def full_sort_predict(self, interaction):
163 user = interaction[self.USER_ID]
164 user_re_e = self.user_re_embedding(user)
165 user_im_e = self.user_im_embedding(user)
167 rec_r_re_e = self.relation_re_embedding.weight[-1]
168 rec_r_im_e = self.relation_im_embedding.weight[-1]
169 rec_r_re_e = rec_r_re_e.expand_as(user_re_e)
170 rec_r_im_e = rec_r_im_e.expand_as(user_re_e)
172 item_indices = torch.tensor(range(self.n_items)).to(self.device)
173 all_item_re_e = self.entity_re_embedding.weight[item_indices]
174 all_item_im_e = self.entity_im_embedding.weight[item_indices]
176 user_re_e = user_re_e.unsqueeze(1).expand(-1, all_item_re_e.shape[0], -1)
177 user_im_e = user_im_e.unsqueeze(1).expand(-1, all_item_re_e.shape[0], -1)
179 rec_r_re_e = rec_r_re_e.unsqueeze(1).expand(-1, all_item_re_e.shape[0], -1)
180 rec_r_im_e = rec_r_im_e.unsqueeze(1).expand(-1, all_item_re_e.shape[0], -1)
182 all_item_re_e = all_item_re_e.unsqueeze(0)
183 all_item_im_e = all_item_im_e.unsqueeze(0)
185 return (
186 self.triple_dot(user_re_e, rec_r_re_e, all_item_re_e)
187 + self.triple_dot(user_im_e, rec_r_re_e, all_item_im_e)
188 + self.triple_dot(user_re_e, rec_r_im_e, all_item_im_e)
189 - self.triple_dot(user_im_e, rec_r_im_e, all_item_im_e)
190 )
192 def full_sort_predict_kg(self, interaction):
193 head = interaction[self.HEAD_ENTITY_ID]
194 relation = interaction[self.RELATION_ID]
195 head_re_e = self.entity_re_embedding(head)
196 head_im_e = self.entity_im_embedding(head)
198 rec_r_re_e = self.relation_re_embedding(relation)
199 rec_r_im_e = self.relation_im_embedding(relation)
201 entity_indices = torch.tensor(range(self.n_entities)).to(self.device)
202 all_entity_re_e = self.entity_re_embedding.weight[entity_indices]
203 all_entity_im_e = self.entity_im_embedding.weight[entity_indices]
205 head_re_e = head_re_e.unsqueeze(1).expand(-1, all_entity_re_e.shape[0], -1)
206 head_im_e = head_im_e.unsqueeze(1).expand(-1, all_entity_re_e.shape[0], -1)
208 rec_r_re_e = rec_r_re_e.unsqueeze(1).expand(-1, all_entity_re_e.shape[0], -1)
209 rec_r_im_e = rec_r_im_e.unsqueeze(1).expand(-1, all_entity_re_e.shape[0], -1)
211 all_entity_re_e = all_entity_re_e.unsqueeze(0)
212 all_entity_im_e = all_entity_im_e.unsqueeze(0)
214 return (
215 self.triple_dot(head_re_e, rec_r_re_e, all_entity_re_e)
216 + self.triple_dot(head_im_e, rec_r_re_e, all_entity_im_e)
217 + self.triple_dot(head_re_e, rec_r_im_e, all_entity_im_e)
218 - self.triple_dot(head_im_e, rec_r_im_e, all_entity_im_e)
219 )