Coverage for hopwise/model/knowledge_graph_embedding_recommender/toruse.py: 64%
106 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/14
2# @Author : Alessandro Soccol
3# @Email : alessandro.soccol@unica.it
5"""TorusE
6##################################################
7Reference:
8 Takuma Ebisu and Ryutaro Ichise. "TorusE: Knowledge Graph Embedding on a Lie Group." in AAAI 2018.
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 TorusE(KnowledgeRecommender):
23 r"""TorusE projects each point in a Torus.
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.margin = config["margin"]
37 self.device = config["device"]
39 # Embeddings
40 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
41 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size)
42 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size)
44 # Loss
45 self.loss = nn.TripletMarginLoss(margin=self.margin, p=2, reduction="mean")
47 # Parameters initialization
48 self.apply(xavier_normal_initialization)
50 def _get_rec_embedding(self, user, pos_item, neg_item):
51 user_e = self.user_embedding(user)
52 pos_item_e = self.entity_embedding(pos_item)
53 neg_item_e = self.entity_embedding(neg_item)
54 rec_r_e = self.relation_embedding.weight[-1]
55 rec_r_e = rec_r_e.expand_as(user_e)
57 return user_e, pos_item_e, neg_item_e, rec_r_e
59 def _get_kg_embedding(self, head, pos_tail, neg_tail, relation):
60 head_e = self.entity_embedding(head)
61 pos_tail_e = self.entity_embedding(pos_tail)
62 neg_tail_e = self.entity_embedding(neg_tail)
63 relation_e = self.relation_embedding(relation)
64 return head_e, pos_tail_e, neg_tail_e, relation_e
66 def forward(self, head, relation, tail):
67 h_e = head.clone()
68 r_e = relation.clone()
69 t_e = tail.clone()
71 h_e.data.frac_()
72 r_e.data.frac_()
73 t_e.data.frac_()
75 h_r = h_e + r_e
76 return -(4 * torch.min((h_r - t_e) ** 2, 1 - (h_r - t_e) ** 2).sum(dim=-1))
78 def calculate_loss(self, interaction):
79 user = interaction[self.USER_ID]
81 pos_item = interaction[self.ITEM_ID]
82 neg_item = interaction[self.NEG_ITEM_ID]
84 head = interaction[self.HEAD_ENTITY_ID]
86 relation = interaction[self.RELATION_ID]
88 pos_tail = interaction[self.TAIL_ENTITY_ID]
89 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID]
91 user_e, pos_item_e, neg_item_e, rec_r_e = self._get_rec_embedding(user, pos_item, neg_item)
92 head_e, pos_tail_e, neg_tail_e, relation_e = self._get_kg_embedding(head, pos_tail, neg_tail, relation)
94 h_e = torch.cat([user_e, head_e])
95 r_e = torch.cat([rec_r_e, relation_e])
96 pos_t_e = torch.cat([pos_item_e, pos_tail_e])
97 neg_t_e = torch.cat([neg_item_e, neg_tail_e])
99 loss = self.loss(h_e + r_e, pos_t_e, neg_t_e)
101 return loss
103 def predict_kg(self, interaction):
104 head = interaction[self.HEAD_ENTITY_ID]
105 relation = interaction[self.RELATION_ID]
106 tail = interaction[self.TAIL_ENTITY_ID]
108 head_e = self.entity_embedding(head)
109 tail_e = self.entity_embedding(tail)
110 rec_r_e = self.relation_embedding(relation)
112 return self.forward(head_e, rec_r_e, tail_e)
114 def predict(self, interaction):
115 user = interaction[self.USER_ID]
116 item = interaction[self.ITEM_ID]
118 user_e = self.user_embedding(user)
119 item_e = self.entity_embedding(item)
120 rec_r_e = self.relation_embedding.weight[-1]
121 rec_r_e = rec_r_e.expand_as(user_e)
123 return self.forward(user_e, rec_r_e, item_e)
125 def full_sort_predict(self, interaction):
126 user = interaction[self.USER_ID]
127 user_e = self.user_embedding(user)
129 rec_r_e = self.relation_embedding.weight[-1]
130 rec_r_e = rec_r_e.expand_as(user_e)
132 item_indices = torch.tensor(range(self.n_items)).to(self.device)
133 all_item_e = self.entity_embedding.weight[item_indices]
135 h_e = user_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1)
136 r_e = rec_r_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1)
137 t = all_item_e.unsqueeze(0)
139 h_e = h_e.clone()
140 r_e = r_e.clone()
141 t = t.clone()
143 h_e.data.frac_()
144 r_e.data.frac_()
145 t.data.frac_()
147 h_r = h_e + r_e
148 return -(4 * torch.min((h_r - t) ** 2, 1 - (h_r - t) ** 2).sum(dim=-1))
150 def full_sort_predict_kg(self, interaction):
151 head = interaction[self.HEAD_ENTITY_ID]
152 relation = interaction[self.RELATION_ID]
154 head_e = self.entity_embedding(head)
155 rec_r_e = self.relation_embedding(relation)
157 all_tail_e = self.entity_embedding.weight
159 h_e = head_e.unsqueeze(1).expand(-1, all_tail_e.shape[0], -1)
160 r_e = rec_r_e.unsqueeze(1).expand(-1, all_tail_e.shape[0], -1)
161 t = all_tail_e.unsqueeze(0)
163 h_e = h_e.clone()
164 r_e = r_e.clone()
165 t = t.clone()
167 h_e.data.frac_()
168 r_e.data.frac_()
169 t.data.frac_()
171 h_r = h_e + r_e
172 return -(4 * torch.min((h_r - t) ** 2, 1 - (h_r - t) ** 2).sum(dim=-1))