Coverage for hopwise/model/knowledge_graph_embedding_recommender/hole.py: 76%
94 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/21
2# @Author : Alessandro Soccol
3# @Email : alessandro.soccol@unica.it
5"""HolE
6##################################################
7Reference:
8 Nickel et al. "Holographic embeddings of knowledge graphs." in AAAI 2016.
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 HolE(KnowledgeRecommender):
23 r"""HoLE combines the expressive power of RESCAL with the efficiency and simplicity of DistMult.
24 The entity representations are composed into h ⋆ t in the set of real numbers,
25 with the circular correlation operator.
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 # Loss
47 self.sigmoid = nn.Sigmoid()
48 self.loss = nn.MarginRankingLoss(margin=self.margin)
50 # Embeddings Initialization
51 self.apply(xavier_normal_initialization)
53 def forward(self, h, r, t):
54 r_e = self.get_rolling_matrix(r)
55 hr = torch.matmul(h.view(-1, 1, self.embedding_size), r_e)
56 return (hr.view(-1, self.embedding_size) * t).sum(dim=1)
58 def get_rolling_matrix(self, x):
59 b_size, dim = x.shape
60 x = x.view(b_size, 1, dim)
61 return torch.cat([x.roll(i, dims=2) for i in range(dim)], dim=1)
63 def _get_rec_embedding(self, user, pos_item, neg_item):
64 user_e = self.user_embedding(user)
65 pos_item_e = self.entity_embedding(pos_item)
66 neg_item_e = self.entity_embedding(neg_item)
67 rec_r_e = self.relation_embedding.weight[-1]
68 rec_r_e = rec_r_e.expand_as(user_e)
70 return user_e, pos_item_e, neg_item_e, rec_r_e
72 def _get_kg_embedding(self, head, pos_tail, neg_tail, relation):
73 head_e = self.entity_embedding(head)
74 pos_tail_e = self.entity_embedding(pos_tail)
75 neg_tail_e = self.entity_embedding(neg_tail)
76 relation_e = self.relation_embedding(relation)
77 return head_e, pos_tail_e, neg_tail_e, relation_e
79 def calculate_loss(self, interaction):
80 user = interaction[self.USER_ID]
82 pos_item = interaction[self.ITEM_ID]
83 neg_item = interaction[self.NEG_ITEM_ID]
85 head = interaction[self.HEAD_ENTITY_ID]
87 relation = interaction[self.RELATION_ID]
89 pos_tail = interaction[self.TAIL_ENTITY_ID]
90 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID]
92 user_e, pos_item_e, neg_item_e, rec_r_e = self._get_rec_embedding(user, pos_item, neg_item)
93 head_e, pos_tail_e, neg_tail_e, relation_e = self._get_kg_embedding(head, pos_tail, neg_tail, relation)
95 pos_score_users = self.forward(user_e, rec_r_e, pos_item_e)
96 neg_score_users = self.forward(user_e, rec_r_e, neg_item_e)
98 pos_score_entities = self.forward(head_e, relation_e, pos_tail_e)
99 neg_score_entities = self.forward(head_e, relation_e, neg_tail_e)
101 pos_scores = torch.cat([pos_score_users, pos_score_entities])
102 neg_scores = torch.cat([neg_score_users, neg_score_entities])
104 loss = self.loss(self.sigmoid(pos_scores), self.sigmoid(neg_scores), torch.ones_like(pos_scores))
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)
113 rec_r_e = self.relation_embedding.weight[-1]
114 rec_r_e = rec_r_e.expand_as(user_e)
116 return self.forward(user_e, rec_r_e, item_e)
118 def predict_kg(self, interaction):
119 head = interaction[self.HEAD_ENTITY_ID]
120 relation = interaction[self.RELATION_ID]
121 tail = interaction[self.TAIL_ENTITY_ID]
123 head_e = self.entity_embedding(head)
124 tail_e = self.entity_embedding(tail)
125 rec_r_e = self.relation_embedding(relation)
127 return self.forward(head_e, rec_r_e, tail_e)
129 def full_sort_predict(self, interaction):
130 user = interaction[self.USER_ID]
131 user_e = self.user_embedding(user)
133 rec_r_e = self.relation_embedding.weight[-1]
134 rec_r_e = rec_r_e.expand_as(user_e)
136 item_indices = torch.tensor(range(self.n_items)).to(self.device)
137 all_item_e = self.entity_embedding.weight[item_indices]
139 r_e = self.get_rolling_matrix(rec_r_e)
141 h_e = user_e.view(user_e.shape[0], 1, self.embedding_size)
142 hr = torch.matmul(h_e, r_e).view(user_e.shape[0], self.embedding_size, 1)
144 return torch.matmul(hr.squeeze(2), all_item_e.T)
146 def full_sort_predict_kg(self, interaction):
147 head = interaction[self.HEAD_ENTITY_ID]
148 relation = interaction[self.RELATION_ID]
150 head_e = self.entity_embedding(head)
151 rec_r_e = self.relation_embedding(relation)
153 all_item_e = self.entity_embedding.weight
155 r_e = self.get_rolling_matrix(rec_r_e)
157 h_e = head_e.view(head_e.shape[0], 1, self.embedding_size)
158 hr = torch.matmul(h_e, r_e).view(head_e.shape[0], self.embedding_size, 1)
160 return torch.matmul(hr.squeeze(2), all_item_e.T)