Coverage for hopwise/model/knowledge_graph_embedding_recommender/transe.py: 69%
88 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/12
2# @Author : Alessandro Soccol
3# @Email : alessandro.soccol@unica.it
5"""TransE
6##################################################
7Reference:
8 Bordes. A et al. "Translating Embeddings for Modeling Multi-relational Data." in NeurIPS 2013.
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 TransE(KnowledgeRecommender):
23 r"""TransE a method which models relationships by interpreting them
24 as translations operating on the low-dimensional embeddings of the entities.
25 Originally created for the knowledge completion task, was adapted to make recommendation
27 .. math::
28 f_t(r)=(h+r,t)
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"]
44 # Embeddings
45 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
46 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size)
47 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size)
49 # Loss
50 self.loss = nn.TripletMarginLoss(margin=self.margin, p=2, reduction="mean")
52 # Parameters initialization
53 self.apply(xavier_normal_initialization)
55 def forward(self, user, relation, item):
56 score = -torch.norm(user + relation - item, p=2, dim=1)
57 return score
59 def _get_rec_embedding(self, user, pos_item, neg_item):
60 user_e = self.user_embedding(user)
61 pos_item_e = self.entity_embedding(pos_item)
62 neg_item_e = self.entity_embedding(neg_item)
63 rec_r_e = self.relation_embedding.weight[-1]
64 rec_r_e = rec_r_e.expand_as(user_e)
66 return user_e, pos_item_e, neg_item_e, rec_r_e
68 def _get_kg_embedding(self, head, pos_tail, neg_tail, relation):
69 head_e = self.entity_embedding(head)
70 pos_tail_e = self.entity_embedding(pos_tail)
71 neg_tail_e = self.entity_embedding(neg_tail)
72 relation_e = self.relation_embedding(relation)
73 return head_e, pos_tail_e, neg_tail_e, relation_e
75 def calculate_loss(self, interaction):
76 user = interaction[self.USER_ID]
78 pos_item = interaction[self.ITEM_ID]
79 neg_item = interaction[self.NEG_ITEM_ID]
81 head = interaction[self.HEAD_ENTITY_ID]
83 relation = interaction[self.RELATION_ID]
85 pos_tail = interaction[self.TAIL_ENTITY_ID]
86 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID]
88 user_e, pos_item_e, neg_item_e, rec_r_e = self._get_rec_embedding(user, pos_item, neg_item)
89 head_e, pos_tail_e, neg_tail_e, relation_e = self._get_kg_embedding(head, pos_tail, neg_tail, relation)
91 h_e = torch.cat([user_e, head_e])
92 r_e = torch.cat([rec_r_e, relation_e])
93 pos_t_e = torch.cat([pos_item_e, pos_tail_e])
94 neg_t_e = torch.cat([neg_item_e, neg_tail_e])
96 loss = self.loss(h_e + r_e, pos_t_e, neg_t_e)
98 return loss
100 def predict(self, interaction):
101 user = interaction[self.USER_ID]
102 item = interaction[self.ITEM_ID]
104 user_e = self.user_embedding(user)
105 item_e = self.entity_embedding(item)
107 rec_r_e = self.relation_embedding.weight[-1]
108 rec_r_e = rec_r_e.expand_as(user_e)
110 return self.forward(user_e, rec_r_e, item_e)
112 def full_sort_predict(self, interaction):
113 user = interaction[self.USER_ID]
114 user_e = self.user_embedding(user)
116 rec_r_e = self.relation_embedding.weight[-1]
117 rec_r_e = rec_r_e.expand_as(user_e)
119 item_indices = torch.tensor(range(self.n_items)).to(self.device)
120 all_item_e = self.entity_embedding.weight[item_indices]
122 user_e = user_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1)
123 rec_r_e = rec_r_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1)
124 t = all_item_e.unsqueeze(0)
126 return -torch.norm(user_e + rec_r_e - t, p=2, dim=2)
128 def predict_kg(self, interaction):
129 head = interaction[self.HEAD_ENTITY_ID]
130 relation = interaction[self.RELATION_ID]
131 tail = interaction[self.TAIL_ENTITY_ID]
133 head_e = self.entity_embedding(head)
134 relation_e = self.relation_embedding(relation)
135 tail_e = self.entity_embedding(tail)
137 return self.forward(head_e, relation_e, tail_e)
139 def full_sort_predict_kg(self, interaction):
140 head = interaction[self.HEAD_ENTITY_ID]
141 relation = interaction[self.RELATION_ID]
143 head_e = self.entity_embedding(head)
145 rel_e = self.relation_embedding(relation)
146 rel_e = rel_e.expand_as(head_e)
148 tail_indices = torch.tensor(range(self.n_entities)).to(self.device)
149 all_tail_e = self.entity_embedding.weight[tail_indices]
151 head_e = head_e.unsqueeze(1).expand(-1, all_tail_e.shape[0], -1)
152 rel_e = rel_e.unsqueeze(1).expand(-1, all_tail_e.shape[0], -1)
153 t = all_tail_e.unsqueeze(0)
154 return -torch.norm(head_e + rel_e - t, p=2, dim=2)