Coverage for hopwise/model/knowledge_graph_embedding_recommender/distmult.py: 74%
86 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"""DistMult
6##################################################
7Reference:
8 Yang et al. "Embedding Entities and Relations for Learning and Inference in Knowledge Bases." in ICLR 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 DistMult(KnowledgeRecommender):
23 r"""DistMult simplify RESCAL by restricting Mr to diagonal matrices.
24 For each relation r, it introduce a vector embedding r and requires Mr = diag(r).
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 # define layers and loss
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)
45 self.loss = nn.MarginRankingLoss(margin=self.margin)
47 # parameters initialization
48 self.apply(xavier_normal_initialization)
50 def forward(self, head, relation, tail):
51 return (head * relation * tail).sum(dim=1)
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)
60 return user_e, pos_item_e, neg_item_e, rec_r_e
62 def _get_kg_embedding(self, head, pos_tail, neg_tail, relation):
63 head_e = self.entity_embedding(head)
64 pos_tail_e = self.entity_embedding(pos_tail)
65 neg_tail_e = self.entity_embedding(neg_tail)
66 relation_e = self.relation_embedding(relation)
67 return head_e, pos_tail_e, neg_tail_e, relation_e
69 def calculate_loss(self, interaction):
70 user = interaction[self.USER_ID]
72 pos_item = interaction[self.ITEM_ID]
73 neg_item = interaction[self.NEG_ITEM_ID]
75 head = interaction[self.HEAD_ENTITY_ID]
77 relation = interaction[self.RELATION_ID]
79 pos_tail = interaction[self.TAIL_ENTITY_ID]
80 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID]
82 user_e, pos_item_e, neg_item_e, rec_r_e = self._get_rec_embedding(user, pos_item, neg_item)
83 head_e, pos_tail_e, neg_tail_e, relation_e = self._get_kg_embedding(head, pos_tail, neg_tail, relation)
85 h_e = torch.cat([user_e, head_e])
86 r_e = torch.cat([rec_r_e, relation_e])
87 pos_t_e = torch.cat([pos_item_e, pos_tail_e])
88 neg_t_e = torch.cat([neg_item_e, neg_tail_e])
90 pos_score = self.forward(h_e, r_e, pos_t_e)
91 neg_score = self.forward(h_e, r_e, neg_t_e)
93 loss = self.loss(pos_score, neg_score, torch.ones_like(pos_score).to(self.device))
95 return loss
97 def predict(self, interaction):
98 user = interaction[self.USER_ID]
99 item = interaction[self.ITEM_ID]
101 user_e = self.user_embedding(user)
102 item_e = self.entity_embedding(item)
103 rec_r_e = self.relation_embedding.weight[-1]
104 rec_r_e = rec_r_e.expand_as(user_e)
106 return self.forward(user_e, rec_r_e, item_e)
108 def predict_kg(self, interaction):
109 head = interaction[self.HEAD_ENTITY_ID]
110 relation = interaction[self.RELATION_ID]
111 tail = interaction[self.TAIL_ENTITY_ID]
113 head_e = self.entity_embedding(head)
114 item_e = self.entity_embedding(tail)
115 rec_r_e = self.relation_embedding(relation)
117 return self.forward(head_e, rec_r_e, item_e)
119 def full_sort_predict(self, interaction):
120 user = interaction[self.USER_ID]
121 user_e = self.user_embedding(user)
123 rec_r_e = self.relation_embedding.weight[-1]
124 rec_r_e = rec_r_e.expand_as(user_e)
126 item_indices = torch.tensor(range(self.n_items)).to(self.device)
127 all_item_e = self.entity_embedding.weight[item_indices]
129 h = user_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1)
130 r = rec_r_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1)
131 t = all_item_e.unsqueeze(0)
133 return (h * r * t).sum(dim=-1)
135 def full_sort_predict_kg(self, interaction):
136 head = interaction[self.HEAD_ENTITY_ID]
137 relation = interaction[self.RELATION_ID]
139 head_e = self.entity_embedding(head)
140 rec_r_e = self.relation_embedding(relation)
142 h = head_e.unsqueeze(1).expand(-1, self.entity_embedding.weight.size(0), -1)
143 r = rec_r_e.unsqueeze(1).expand(-1, self.entity_embedding.weight.size(0), -1)
144 t = self.entity_embedding.weight.unsqueeze(0)
146 return (h * r * t).sum(dim=-1)