Coverage for hopwise/model/knowledge_graph_embedding_recommender/rescal.py: 74%
91 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/20
2# @Author : Alessandro Soccol
3# @Email : alessandro.soccol@unica.it
5"""RESCAL
6##################################################
7Reference:
8 Nickel et al. "A three-way model for collective learning on multi-relational data." in ICML 2011.
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 RESCAL(KnowledgeRecommender):
23 r"""RESCAL associates each entity with a vector to capture its latent semantics.
24 Each relation is represented as a matrix which models pairwise interactions between latent vectors
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"]
40 # Embeddings
41 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
42 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size)
43 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size**2)
45 # Loss
46 self.loss = nn.MarginRankingLoss(margin=self.margin)
48 # Parameters initialization
49 self.apply(xavier_normal_initialization)
51 def forward(self, head, relation, tail):
52 hr = torch.matmul(head.view(-1, 1, 1, self.embedding_size), relation)
53 hr = hr.view(-1, self.embedding_size)
54 return (hr * tail).sum(dim=1)
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].view(1, 1, self.embedding_size, self.embedding_size)
61 rec_r_e = rec_r_e.repeat(user_e.shape[0], 1, 1, 1)
63 return user_e, pos_item_e, neg_item_e, rec_r_e
65 def _get_kg_embedding(self, head, pos_tail, neg_tail, relation):
66 head_e = self.entity_embedding(head)
67 pos_tail_e = self.entity_embedding(pos_tail)
68 neg_tail_e = self.entity_embedding(neg_tail)
69 relation_e = self.relation_embedding(relation).view(-1, 1, self.embedding_size, self.embedding_size)
71 return head_e, pos_tail_e, neg_tail_e, relation_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, pos_tail, neg_tail, relation)
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 pos_score = self.forward(h_e, r_e, pos_t_e)
95 neg_score = self.forward(h_e, r_e, neg_t_e)
97 loss = self.loss(pos_score, neg_score, torch.ones_like(pos_score).to(self.device))
99 return loss
101 def predict(self, interaction):
102 user = interaction[self.USER_ID]
103 item = interaction[self.ITEM_ID]
105 user_e = self.user_embedding(user)
106 item_e = self.entity_embedding(item)
107 rec_r_e = self.relation_embedding.weight[-1].view(1, 1, self.embedding_size, self.embedding_size)
108 rec_r_e = rec_r_e.repeat(user_e.shape[0], 1, 1, 1)
110 return self.forward(user_e, rec_r_e, item_e)
112 def predict_kg(self, interaction):
113 head = interaction[self.HEAD_ENTITY_ID]
114 relation = interaction[self.RELATION_ID]
115 tail = interaction[self.TAIL_ENTITY_ID]
117 head_e = self.entity_embedding(head)
118 tail_e = self.entity_embedding(tail)
119 rec_r_e = self.relation_embedding(relation).view(
120 relation.shape[0], 1, self.embedding_size, self.embedding_size
121 )
123 return self.forward(head_e, rec_r_e, tail_e)
125 def full_sort_predict(self, interaction):
126 user = interaction[self.USER_ID]
127 user_e = self.user_embedding(user)
129 rec_r_e = self.relation_embedding.weight[-1].view(1, 1, self.embedding_size, self.embedding_size)
130 rec_r_e = rec_r_e.repeat(user_e.shape[0], 1, 1, 1)
132 item_indices = torch.tensor(range(self.n_items)).to(self.device)
133 all_item_e = self.entity_embedding.weight[item_indices]
135 user_e = user_e.view(-1, 1, 1, self.embedding_size)
136 hr = torch.matmul(user_e, rec_r_e)
138 scores = torch.matmul(hr.squeeze(2), all_item_e.T)
139 scores = scores.squeeze(1)
140 return scores
142 def full_sort_predict_kg(self, interaction):
143 head = interaction[self.HEAD_ENTITY_ID]
144 relation = interaction[self.RELATION_ID]
146 head_e = self.entity_embedding(head)
148 rec_r_e = self.relation_embedding(relation).view(
149 relation.shape[0], 1, self.embedding_size, self.embedding_size
150 )
152 all_tail_e = self.entity_embedding.weight
154 head_e = head_e.view(-1, 1, 1, self.embedding_size)
155 hr = torch.matmul(head_e, rec_r_e)
157 scores = torch.matmul(hr.squeeze(2), all_tail_e.T)
158 scores = scores.squeeze(1)
159 return scores