Coverage for hopwise/model/knowledge_aware_recommender/mkr.py: 100%
116 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 : 2020/10/08
2# @Author : Xinyan Fan
3# @Email : xinyan.fan@ruc.edu.cn
5r"""MKR
6#####################################################
7Reference:
8 Hongwei Wang et al. "Multi-Task Feature Learning for Knowledge Graph Enhanced Recommendation." in WWW 2019.
10Reference code:
11 https://github.com/hsientzucheng/MKR.PyTorch
12"""
14import torch
15from torch import nn
17from hopwise.model.abstract_recommender import KnowledgeRecommender
18from hopwise.model.init import xavier_normal_initialization
19from hopwise.model.layers import MLPLayers
20from hopwise.utils import InputType
23class MKR(KnowledgeRecommender):
24 r"""MKR is a Multi-task feature learning approach for Knowledge graph enhanced Recommendation. It is a deep
25 end-to-end framework that utilizes knowledge graph embedding task to assist recommendation task. The two
26 tasks are associated by cross&compress units, which automatically share latent features and learn high-order
27 interactions between items in recommender systems and entities in the knowledge graph.
28 """
30 input_type = InputType.POINTWISE
32 def __init__(self, config, dataset):
33 super().__init__(config, dataset)
35 # load parameters info
36 self.LABEL = config["LABEL_FIELD"]
37 self.embedding_size = config["embedding_size"]
38 self.kg_embedding_size = config["kg_embedding_size"]
39 self.L = config["low_layers_num"] # the number of low layers
40 self.H = config["high_layers_num"] # the number of high layers
41 self.reg_weight = config["reg_weight"]
42 self.use_inner_product = config["use_inner_product"]
43 self.dropout_prob = config["dropout_prob"]
45 # init embeddings
46 self.user_embeddings_lookup = nn.Embedding(self.n_users, self.embedding_size)
47 self.item_embeddings_lookup = nn.Embedding(self.n_entities, self.embedding_size)
48 self.entity_embeddings_lookup = nn.Embedding(self.n_entities, self.embedding_size)
49 self.relation_embeddings_lookup = nn.Embedding(self.n_relations, self.embedding_size)
51 # define layers
52 lower_mlp_layers = []
53 high_mlp_layers = []
54 for i in range(self.L + 1):
55 lower_mlp_layers.append(self.embedding_size)
56 for i in range(self.H):
57 high_mlp_layers.append(self.embedding_size * 2)
59 self.user_mlp = MLPLayers(lower_mlp_layers, self.dropout_prob, "sigmoid")
60 self.tail_mlp = MLPLayers(lower_mlp_layers, self.dropout_prob, "sigmoid")
61 self.cc_unit = nn.Sequential()
62 for i_cnt in range(self.L):
63 self.cc_unit.add_module(f"cc_unit{i_cnt}", CrossCompressUnit(self.embedding_size))
64 self.kge_mlp = MLPLayers(high_mlp_layers, self.dropout_prob, "sigmoid")
65 self.kge_pred_mlp = MLPLayers([self.embedding_size * 2, self.embedding_size], self.dropout_prob, "sigmoid")
66 if not self.use_inner_product:
67 self.rs_pred_mlp = MLPLayers([self.embedding_size * 2, 1], self.dropout_prob, "sigmoid")
68 self.rs_mlp = MLPLayers(high_mlp_layers, self.dropout_prob, "sigmoid")
70 # loss
71 self.sigmoid_BCE = nn.BCEWithLogitsLoss()
73 # parameters initialization
74 self.apply(xavier_normal_initialization)
76 def forward(
77 self,
78 user_indices=None,
79 item_indices=None,
80 head_indices=None,
81 relation_indices=None,
82 tail_indices=None,
83 ):
84 self.item_embeddings = self.item_embeddings_lookup(item_indices)
85 self.head_embeddings = self.entity_embeddings_lookup(head_indices)
86 self.item_embeddings, self.head_embeddings = self.cc_unit(
87 [self.item_embeddings, self.head_embeddings]
88 ) # calculate feature interactions between items and entities
90 if user_indices is not None:
91 # RS
92 self.user_embeddings = self.user_embeddings_lookup(user_indices)
93 self.user_embeddings = self.user_mlp(self.user_embeddings)
95 if self.use_inner_product: # get scores by inner product.
96 self.scores = torch.sum(self.user_embeddings * self.item_embeddings, 1) # [batch_size]
97 else: # get scores by mlp layers
98 self.user_item_concat = torch.cat(
99 [self.user_embeddings, self.item_embeddings], 1
100 ) # [batch_size, emb_dim*2]
101 self.user_item_concat = self.rs_mlp(self.user_item_concat)
103 self.scores = torch.squeeze(self.rs_pred_mlp(self.user_item_concat)) # [batch_size]
104 self.scores_normalized = torch.sigmoid(self.scores)
105 outputs = [
106 self.user_embeddings,
107 self.item_embeddings,
108 self.scores,
109 self.scores_normalized,
110 ]
112 if relation_indices is not None:
113 # KGE
114 self.tail_embeddings = self.entity_embeddings_lookup(tail_indices)
115 self.relation_embeddings = self.relation_embeddings_lookup(relation_indices)
116 self.tail_embeddings = self.tail_mlp(self.tail_embeddings)
118 self.head_relation_concat = torch.cat(
119 [self.head_embeddings, self.relation_embeddings], 1
120 ) # [batch_size, emb_dim*2]
121 self.head_relation_concat = self.kge_mlp(self.head_relation_concat)
123 self.tail_pred = self.kge_pred_mlp(self.head_relation_concat) # [batch_size, 1]
124 self.tail_pred = torch.sigmoid(self.tail_pred)
125 self.scores_kge = torch.sigmoid(torch.sum(self.tail_embeddings * self.tail_pred, 1))
126 self.rmse = torch.mean(
127 torch.sqrt(torch.sum(torch.pow(self.tail_embeddings - self.tail_pred, 2), 1) / self.embedding_size)
128 )
129 outputs = [
130 self.head_embeddings,
131 self.tail_embeddings,
132 self.scores_kge,
133 self.rmse,
134 ]
136 return outputs
138 def _l2_loss(self, inputs):
139 return torch.sum(inputs**2) / 2
141 def calculate_rs_loss(self, interaction):
142 r"""Calculate the training loss for a batch data of RS."""
143 # inputs
144 self.user_indices = interaction[self.USER_ID]
145 self.item_indices = interaction[self.ITEM_ID]
146 self.head_indices = interaction[self.ITEM_ID]
147 self.labels = interaction[self.LABEL]
148 # RS model
149 user_embeddings, item_embeddings, scores, scores_normalized = self.forward(
150 user_indices=self.user_indices,
151 item_indices=self.item_indices,
152 head_indices=self.head_indices,
153 relation_indices=None,
154 tail_indices=None,
155 )
156 # loss
157 base_loss_rs = torch.mean(self.sigmoid_BCE(scores, self.labels))
158 l2_loss_rs = self._l2_loss(user_embeddings) + self._l2_loss(item_embeddings)
159 loss_rs = base_loss_rs + l2_loss_rs * self.reg_weight
161 return loss_rs
163 def calculate_kg_loss(self, interaction):
164 r"""Calculate the training loss for a batch data of KG."""
165 # inputs
166 self.item_indices = interaction[self.HEAD_ENTITY_ID]
167 self.head_indices = interaction[self.HEAD_ENTITY_ID]
168 self.relation_indices = interaction[self.RELATION_ID]
169 self.tail_indices = interaction[self.TAIL_ENTITY_ID]
170 # KGE model
171 head_embeddings, tail_embeddings, scores_kge, rmse = self.forward(
172 user_indices=None,
173 item_indices=self.item_indices,
174 head_indices=self.head_indices,
175 relation_indices=self.relation_indices,
176 tail_indices=self.tail_indices,
177 )
178 # loss
179 base_loss_kge = -scores_kge
180 l2_loss_kge = self._l2_loss(head_embeddings) + self._l2_loss(tail_embeddings)
181 loss_kge = base_loss_kge + l2_loss_kge * self.reg_weight
183 return loss_kge.sum()
185 def predict(self, interaction):
186 user = interaction[self.USER_ID]
187 item = interaction[self.ITEM_ID]
188 head = interaction[self.ITEM_ID]
190 outputs = self.forward(user, item, head)
191 _, _, scores, _ = outputs
193 return scores
196class CrossCompressUnit(nn.Module):
197 r"""This is Cross&Compress Unit for MKR model to model feature interactions between items and entities."""
199 def __init__(self, dim):
200 super().__init__()
201 self.dim = dim
202 self.fc_vv = nn.Linear(dim, 1, bias=True)
203 self.fc_ev = nn.Linear(dim, 1, bias=True)
204 self.fc_ve = nn.Linear(dim, 1, bias=True)
205 self.fc_ee = nn.Linear(dim, 1, bias=True)
207 def forward(self, inputs):
208 v, e = inputs
209 # [batch_size, dim, 1], [batch_size, 1, dim]
210 v = torch.unsqueeze(v, 2)
211 e = torch.unsqueeze(e, 1)
212 # [batch_size, dim, dim]
213 c_matrix = torch.matmul(v, e)
214 c_matrix_transpose = c_matrix.permute(0, 2, 1)
215 # [batch_size * dim, dim]
216 c_matrix = c_matrix.view(-1, self.dim)
217 c_matrix_transpose = c_matrix_transpose.contiguous().view(-1, self.dim)
218 # [batch_size, dim]
219 v_intermediate = self.fc_vv(c_matrix) + self.fc_ev(c_matrix_transpose)
220 e_intermediate = self.fc_ve(c_matrix) + self.fc_ee(c_matrix_transpose)
221 v_output = v_intermediate.view(-1, self.dim)
222 e_output = e_intermediate.view(-1, self.dim)
224 return v_output, e_output