Coverage for hopwise/model/knowledge_graph_embedding_recommender/transr.py: 78%
97 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 TransR(KnowledgeRecommender):
23 r"""TransR Rather than introducing relation-specific hyperplanes, it introduces relation-specific spaces.
24 The scoring functions is the same as TransH but h and t are projected into the space specific to relation
26 Note:
27 In this version, we sample recommender data and knowledge data separately, and put them together for training.
28 """
30 input_type = InputType.PAIRWISE
32 def __init__(self, config, dataset):
33 super().__init__(config, dataset)
35 # Load parameters info
36 self.embedding_size = config["embedding_size"]
37 self.margin = config["margin"]
38 self.device = config["device"]
39 self.ui_relation = dataset.field2token_id["relation_id"][dataset.ui_relation]
41 # Embeddings
42 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
43 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size)
44 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size)
45 self.proj_mat_e = nn.Embedding(self.n_relations, self.embedding_size * self.embedding_size)
47 # Loss
48 self.loss = nn.TripletMarginLoss(margin=self.margin, p=2, reduction="mean")
50 # Parameters initialization
51 self.apply(xavier_normal_initialization)
53 def _get_rec_embedding(self, user, pos_item, neg_item):
54 user_e = self.user_embedding(user)
55 pos_item_e = self.entity_embedding(pos_item)
56 neg_item_e = self.entity_embedding(neg_item)
57 rec_r_e = self.relation_embedding.weight[-1]
58 rec_r_e = rec_r_e.expand_as(user_e)
59 return user_e, pos_item_e, neg_item_e, rec_r_e
61 def _get_kg_embedding(self, head, relation, pos_tail, neg_tail):
62 head_e = self.entity_embedding(head)
63 pos_tail_e = self.entity_embedding(pos_tail)
64 neg_tail_e = self.entity_embedding(neg_tail)
65 relation_e = self.relation_embedding(relation)
66 return head_e, pos_tail_e, neg_tail_e, relation_e
68 def forward(self, ent, proj_mat):
69 proj_e = torch.matmul(proj_mat, ent.unsqueeze(2))
70 proj_e = proj_e.squeeze(-1)
71 return proj_e
73 def calculate_loss(self, interaction):
74 user = interaction[self.USER_ID]
76 pos_item = interaction[self.ITEM_ID]
77 neg_item = interaction[self.NEG_ITEM_ID]
79 head = interaction[self.HEAD_ENTITY_ID]
81 relation = interaction[self.RELATION_ID]
83 pos_tail = interaction[self.TAIL_ENTITY_ID]
84 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID]
86 user_e, pos_item_e, neg_item_e, rec_r_e = self._get_rec_embedding(user, pos_item, neg_item)
87 head_e, pos_tail_e, neg_tail_e, relation_e = self._get_kg_embedding(head, relation, pos_tail, neg_tail)
89 h_e = torch.cat([user_e, head_e])
90 r_e = torch.cat([rec_r_e, relation_e])
91 pos_t_e = torch.cat([pos_item_e, pos_tail_e])
92 neg_t_e = torch.cat([neg_item_e, neg_tail_e])
94 rec_rel = torch.tensor([self.ui_relation] * user.shape[0], device=self.device)
95 relation = torch.cat([rec_rel, relation])
97 proj_mat = self.proj_mat_e(relation).view(h_e.shape[0], self.embedding_size, self.embedding_size)
99 h_e_proj = self.forward(h_e, proj_mat)
100 pos_t_e_proj = self.forward(pos_t_e, proj_mat)
101 neg_t_e_proj = self.forward(neg_t_e, proj_mat)
103 loss = self.loss(h_e_proj + r_e, pos_t_e_proj, neg_t_e_proj)
105 return loss
107 def predict(self, interaction):
108 user = interaction[self.USER_ID]
109 item = interaction[self.ITEM_ID]
111 user_e = self.user_embedding(user)
112 item_e = self.entity_embedding(item)
114 rec_r_e = self.relation_embedding.weight[-1]
115 rec_r_e = rec_r_e.expand_as(user_e)
117 relation = torch.tensor([self.ui_relation] * user.shape[0], device=self.device)
118 proj_mat = self.proj_mat_e(relation).view(user_e.shape[0], self.embedding_size, self.embedding_size)
120 user_e_proj = self.forward(user_e, proj_mat)
121 item_e_proj = self.forward(item_e, proj_mat)
123 return -torch.norm(user_e_proj + rec_r_e - item_e_proj, p=2, dim=1)
125 def predict_kg(self, interaction):
126 head = interaction[self.HEAD_ENTITY_ID]
127 relation = interaction[self.RELATION_ID]
128 tail = interaction[self.TAIL_ENTITY_ID]
130 head_e = self.entity_embedding(head)
131 tail_e = self.entity_embedding(tail)
133 rec_r_e = self.relation_embedding(relation)
135 proj_mat = self.proj_mat_e(relation).view(head_e.shape[0], self.embedding_size, self.embedding_size)
137 head_e_proj = self.forward(head_e, proj_mat)
138 tail_e_proj = self.forward(tail_e, proj_mat)
140 return -torch.norm(head_e_proj + rec_r_e - tail_e_proj, p=2, dim=1)
142 def full_sort_predict(self, interaction):
143 user = interaction[self.USER_ID]
144 user_e = self.user_embedding(user)
146 rec_r_e = self.relation_embedding.weight[-1]
148 item_indices = torch.tensor(range(self.n_items)).to(self.device)
149 all_item_e = self.entity_embedding.weight[item_indices]
151 relation_users = torch.tensor([self.ui_relation] * user.shape[0], device=self.device)
152 relation_items = torch.tensor([self.ui_relation] * all_item_e.shape[0], device=self.device)
154 proj_mat_user = self.proj_mat_e(relation_users).view(user.shape[0], self.embedding_size, self.embedding_size)
155 proj_mat_items = self.proj_mat_e(relation_items).view(
156 all_item_e.shape[0], self.embedding_size, self.embedding_size
157 )
159 user_e_proj = self.forward(user_e, proj_mat_user)
160 item_e_proj = self.forward(all_item_e, proj_mat_items)
162 user_e_proj = user_e_proj.unsqueeze(1).expand(-1, item_e_proj.shape[0], -1)
163 rec_r_e = rec_r_e.unsqueeze(0).expand(1, item_e_proj.shape[0], -1)
164 item_e_proj = item_e_proj.unsqueeze(0)
166 return -torch.norm(user_e_proj + rec_r_e - item_e_proj, p=2, dim=2)