Coverage for hopwise/model/knowledge_graph_embedding_recommender/transd.py: 70%
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/14
2# @Author : Alessandro Soccol
3# @Email : alessandro.soccol@unica.it
5"""TransD
6##################################################
7Reference:
8 Ji et al. "Knowledge Graph Embedding via Dynamic Mapping Matrix." in ACL/IJCNLP 2015.
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 TransD(KnowledgeRecommender):
23 r"""TransD simplifies TransR by further decomposing the projection matrix into a product of two vector.
24 Also in this case, the scoring function is the same as TransH and TransR,
25 but it introduces three additional mapping vectors along with the entity and relation representation.
27 Note:
28 In this version, we sample recommender data and knowledge data separately, and put them together for training.
29 """
31 input_type = InputType.PAIRWISE
33 def __init__(self, config, dataset):
34 super().__init__(config, dataset)
36 # Load parameters info
37 self.embedding_size = config["embedding_size"]
38 self.margin = config["margin"]
39 self.device = config["device"]
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)
46 self.user_vec_embedding = nn.Embedding(self.n_users, self.embedding_size)
47 self.entity_vec_embedding = nn.Embedding(self.n_entities, self.embedding_size)
48 self.relation_vec_embedding = nn.Embedding(self.n_relations, self.embedding_size)
50 # Loss
51 self.loss = nn.TripletMarginLoss(margin=self.margin, p=2, reduction="mean")
53 # Parameters initialization
54 self.apply(xavier_normal_initialization)
56 def _get_rec_embedding(self, user, pos_item, neg_item):
57 user_e = self.user_embedding(user)
58 pos_item_e = self.entity_embedding(pos_item)
59 neg_item_e = self.entity_embedding(neg_item)
60 rec_r_e = self.relation_embedding.weight[-1]
61 rec_r_e = rec_r_e.expand_as(user_e)
63 return user_e, pos_item_e, neg_item_e, rec_r_e
65 def _get_rec_vec_embedding(self, user, pos_item, neg_item):
66 user_e = self.user_vec_embedding(user)
67 pos_item_e = self.entity_vec_embedding(pos_item)
68 neg_item_e = self.entity_vec_embedding(neg_item)
69 rec_r_e = self.relation_vec_embedding.weight[-1]
70 rec_r_e = rec_r_e.expand_as(user_e)
72 return user_e, pos_item_e, neg_item_e, rec_r_e
74 def _get_kg_embedding(self, head, pos_tail, neg_tail, relation):
75 head_e = self.entity_embedding(head)
76 pos_tail_e = self.entity_embedding(pos_tail)
77 neg_tail_e = self.entity_embedding(neg_tail)
78 relation_e = self.relation_embedding(relation)
80 return head_e, pos_tail_e, neg_tail_e, relation_e
82 def _get_kg_vec_embedding(self, head, pos_tail, neg_tail, relation):
83 head_e = self.entity_vec_embedding(head)
84 pos_tail_e = self.entity_vec_embedding(pos_tail)
85 neg_tail_e = self.entity_vec_embedding(neg_tail)
86 relation_e = self.relation_vec_embedding(relation)
88 return head_e, pos_tail_e, neg_tail_e, relation_e
90 def forward(self, ent, ent_vect, rel_vect):
91 """We note that :math:`p_r(e)_i = e^p^Te \\times r^p_i + e_i` which is
92 more efficient to compute than the matrix formulation in the original
93 paper."""
94 proj_e = rel_vect * ((ent * ent_vect).sum(dim=1).unsqueeze(1))
95 return proj_e + ent
97 def calculate_loss(self, interaction):
98 user = interaction[self.USER_ID]
100 pos_item = interaction[self.ITEM_ID]
101 neg_item = interaction[self.NEG_ITEM_ID]
103 head = interaction[self.HEAD_ENTITY_ID]
105 relation = interaction[self.RELATION_ID]
107 pos_tail = interaction[self.TAIL_ENTITY_ID]
108 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID]
110 user_e, pos_item_e, neg_item_e, rec_r_e = self._get_rec_embedding(user, pos_item, neg_item)
111 head_e, pos_tail_e, neg_tail_e, relation_e = self._get_kg_embedding(head, pos_tail, neg_tail, relation)
113 user_e_vec, pos_item_e_vec, neg_item_e_vec, rec_r_e_vec = self._get_rec_vec_embedding(user, pos_item, neg_item)
114 head_e_vec, pos_tail_e_vec, neg_tail_e_vec, relation_e_vec = self._get_kg_vec_embedding(
115 head, pos_tail, neg_tail, relation
116 )
118 h_e = torch.cat([user_e, head_e])
119 r_e = torch.cat([rec_r_e, relation_e])
120 pos_t_e = torch.cat([pos_item_e, pos_tail_e])
121 neg_t_e = torch.cat([neg_item_e, neg_tail_e])
123 h_e_vec = torch.cat([user_e_vec, head_e_vec])
124 r_e_vec = torch.cat([rec_r_e_vec, relation_e_vec])
125 pos_t_e_vec = torch.cat([pos_item_e_vec, pos_tail_e_vec])
126 neg_t_e_vec = torch.cat([neg_item_e_vec, neg_tail_e_vec])
128 h_projection = self.forward(h_e, h_e_vec, r_e_vec)
129 pos_t_e_projection = self.forward(pos_t_e, pos_t_e_vec, r_e_vec)
130 neg_t_e_projection = self.forward(neg_t_e, neg_t_e_vec, r_e_vec)
132 loss = self.loss(h_projection + r_e, pos_t_e_projection, neg_t_e_projection)
133 return loss
135 def predict(self, interaction):
136 user = interaction[self.USER_ID]
137 item = interaction[self.ITEM_ID]
139 user_e = self.user_embedding(user)
140 user_e_vec = self.user_vec_embedding(user)
142 item_e = self.entity_embedding(item)
143 item_e_vec = self.entity_vec_embedding(item)
145 rec_r_e = self.relation_embedding.weight[-1]
146 rec_r_e = rec_r_e.expand_as(user_e)
148 rec_r_e_vec = self.relation_vec_embedding.weight[-1]
150 user_projection = self.forward(user_e, user_e_vec, rec_r_e_vec)
151 item_projection = self.forward(item_e, item_e_vec, rec_r_e_vec)
153 score = -torch.norm(user_projection + rec_r_e - item_projection, p=2, dim=1)
154 return score
156 def predict_kg(self, interaction):
157 head = interaction[self.HEAD_ENTITY_ID]
158 relation = interaction[self.RELATION_ID]
159 tail = interaction[self.TAIL_ENTITY_ID]
161 head_e = self.entity_embedding(head)
162 head_e_vec = self.entity_vec_embedding(head)
164 tail_e = self.entity_embedding(tail)
165 tail_e_vec = self.entity_vec_embedding(tail)
167 rec_r_e = self.relation_embedding(relation)
168 rec_r_e_vec = self.relation_vec_embedding(relation)
170 head_projection = self.forward(head_e, head_e_vec, rec_r_e_vec)
171 tail_projection = self.forward(tail_e, tail_e_vec, rec_r_e_vec)
173 score = -torch.norm(head_projection + rec_r_e - tail_projection, p=2, dim=1)
174 return score
176 def full_sort_predict(self, interaction):
177 user = interaction[self.USER_ID]
178 user_e = self.user_embedding(user)
179 user_e_vec = self.user_vec_embedding(user)
181 rec_r_e = self.relation_embedding.weight[-1]
182 rec_r_e_vec = self.relation_vec_embedding.weight[-1]
184 users_projection = self.forward(user_e, user_e_vec, rec_r_e_vec)
186 item_indices = torch.tensor(range(self.n_items)).to(self.device)
187 all_item_e = self.entity_embedding.weight[item_indices]
188 all_item_e_vec = self.entity_vec_embedding.weight[item_indices]
190 items_projection = self.forward(all_item_e, all_item_e_vec, rec_r_e_vec)
192 h_r = (users_projection + rec_r_e).unsqueeze(1).expand(-1, items_projection.shape[0], -1)
193 t = items_projection.unsqueeze(0)
195 return -torch.norm(h_r - t, p=2, dim=2)
197 def full_sort_predict_kg(self, interaction):
198 user = interaction[self.HEAD_ENTITY_ID]
199 relation = interaction[self.RELATION_ID]
201 head_e = self.entity_embedding(user)
202 head_e_vec = self.entity_embedding(user)
204 rec_r_e = self.relation_embedding(relation)
205 rec_r_e_vec = self.relation_vec_embedding(relation)
207 heads_projection = self.forward(head_e, head_e_vec, rec_r_e_vec)
209 all_tail_e = self.entity_embedding.weight
210 all_tail_e_vec = self.entity_vec_embedding.weight
212 rec_r_e_vec = rec_r_e_vec.unsqueeze(1)
213 tails_projection = self.forward(all_tail_e, all_tail_e_vec, rec_r_e_vec)
215 h_r = (heads_projection + rec_r_e).unsqueeze(1).expand(-1, tails_projection.shape[1], -1)
217 return -torch.norm(h_r - tails_projection, p=2, dim=2)