Coverage for hopwise/model/knowledge_graph_embedding_recommender/transh.py: 86%
85 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"""TransH
6##################################################
7Reference:
8 Wang Z. et al. "Knowledge Graph Embedding by Translating on Hyperplanes." in AAAI 2014.
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 TransH(KnowledgeRecommender):
23 r"""TransH Have been invented to overcome the disadvantages of TransE,
24 allowing an entity to have distinct representations when involved in different relations.
25 It introduces relation-specific hyperplanes.
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"]
40 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.norm_vec = nn.Embedding(self.n_relations, self.embedding_size)
47 # Loss
48 self.rec_loss = nn.TripletMarginLoss(margin=self.margin, p=2, reduction="mean")
50 # Parameters initialization
51 self.apply(xavier_normal_initialization)
53 def forward(self, head, relation, tail, relation_ids):
54 head_proj = self.project(head, relation_ids)
55 tail_proj = self.project(tail, relation_ids)
56 score = -torch.norm(head_proj + relation - tail_proj, 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 project(self, ent, rel):
76 return ent - (ent * self.norm_vec(rel).sum(1).view(-1, 1)) * self.norm_vec(rel)
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 relation_user = torch.tensor([self.ui_relation] * user_e.shape[0], device=self.device)
95 # Projections
96 user_e = self.project(user_e, relation_user)
97 head_e = self.project(head_e, relation)
98 pos_item_e = self.project(pos_item_e, relation_user)
99 pos_tail_e = self.project(pos_tail_e, relation)
100 neg_item_e = self.project(neg_item_e, relation_user)
101 neg_tail_e = self.project(neg_tail_e, relation)
103 h_e = torch.cat([user_e, head_e])
104 r_e = torch.cat([rec_r_e, relation_e])
105 pos_t_e = torch.cat([pos_item_e, pos_tail_e])
106 neg_t_e = torch.cat([neg_item_e, neg_tail_e])
108 loss = self.rec_loss(h_e + r_e, pos_t_e, neg_t_e)
110 return loss
112 def predict(self, interaction):
113 user = interaction[self.USER_ID]
114 item = interaction[self.ITEM_ID]
116 user_e = self.user_embedding(user)
118 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 relation_ids = torch.tensor([self.ui_relation] * user_e.shape[0], device=self.device)
125 return self.forward(user_e, rec_r_e, item_e, relation_ids)
127 def full_sort_predict(self, interaction):
128 user = interaction[self.USER_ID]
129 user_e = self.user_embedding(user)
131 rec_r_e = self.relation_embedding.weight[-1]
132 rec_r_e = rec_r_e.expand_as(user_e)
134 item_indices = torch.tensor(range(self.n_items)).to(self.device)
135 all_item_e = self.entity_embedding.weight[item_indices]
137 relation_ids_user = torch.tensor([self.ui_relation] * user_e.shape[0], device=self.device)
138 relation_ids_item = torch.tensor([self.ui_relation] * all_item_e.shape[0], device=self.device)
140 h_r = self.project(user_e, relation_ids_user) + rec_r_e
141 h_r = h_r.unsqueeze(1).expand(-1, all_item_e.shape[0], -1)
143 t = self.project(all_item_e, relation_ids_item)
144 t = t.unsqueeze(0)
145 return -torch.norm(h_r - t, p=2, dim=2)