Coverage for hopwise/model/general_recommender/dgcf.py: 94%
196 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/8/26
2# @Author : Gaole He
3# @Email : hegaole@ruc.edu.cn
5# UPDATE:
6# @Time : 2020/9/16
7# @Author : Shanlei Mu
8# @Email : slmu@ruc.edu.cn
10r"""DGCF
11################################################
12Reference:
13 Wang Xiang et al. "Disentangled Graph Collaborative Filtering." in SIGIR 2020.
15Reference code:
16 https://github.com/xiangwang1223/disentangled_graph_collaborative_filtering
17"""
19import random as rd
21import numpy as np
22import torch
23import torch.nn.functional as F
24from torch import nn
25from torch.autograd import Variable
27from hopwise.model.abstract_recommender import GeneralRecommender
28from hopwise.model.init import xavier_normal_initialization
29from hopwise.model.loss import BPRLoss, EmbLoss
30from hopwise.utils import InputType
33def sample_cor_samples(n_users, n_items, cor_batch_size):
34 r"""This is a function that sample item ids and user ids.
36 Args:
37 n_users (int): number of users in total
38 n_items (int): number of items in total
39 cor_batch_size (int): number of id to sample
41 Returns:
42 list: cor_users, cor_items. The result sampled ids with both as cor_batch_size long.
44 Note:
45 We have to sample some embedded representations out of all nodes.
46 Because we have no way to store cor-distance for each pair.
47 """
48 cor_users = rd.sample(list(range(n_users)), cor_batch_size)
49 cor_items = rd.sample(list(range(n_items)), cor_batch_size)
51 return cor_users, cor_items
54class DGCF(GeneralRecommender):
55 r"""DGCF is a disentangled representation enhanced matrix factorization model.
56 The interaction matrix of :math:`n_{users} \times n_{items}` is decomposed to :math:`n_{factors}` intent graph,
57 we carefully design the data interface and use sparse tensor to train and test efficiently.
58 We implement the model following the original author with a pairwise training mode.
59 """
61 input_type = InputType.PAIRWISE
63 def __init__(self, config, dataset):
64 super().__init__(config, dataset)
66 # load dataset info
67 self.interaction_matrix = dataset.inter_matrix(form="coo").astype(np.float32)
69 # load parameters info
70 self.embedding_size = config["embedding_size"]
71 self.n_factors = config["n_factors"]
72 self.n_iterations = config["n_iterations"]
73 self.n_layers = config["n_layers"]
74 self.reg_weight = config["reg_weight"]
75 self.cor_weight = config["cor_weight"]
76 n_batch = dataset.inter_num // config["train_batch_size"] + 1
77 self.cor_batch_size = int(max(self.n_users / n_batch, self.n_items / n_batch))
78 # ensure embedding can be divided into <n_factors> intent
79 assert self.embedding_size % self.n_factors == 0
81 # generate intermediate data
82 row = self.interaction_matrix.row.tolist()
83 col = self.interaction_matrix.col.tolist()
84 col = [item_index + self.n_users for item_index in col]
85 all_h_list = row + col # row.extend(col)
86 all_t_list = col + row # col.extend(row)
87 num_edge = len(all_h_list)
88 edge_ids = range(num_edge)
89 self.all_h_list = torch.LongTensor(all_h_list).to(self.device)
90 self.all_t_list = torch.LongTensor(all_t_list).to(self.device)
91 self.edge2head = torch.LongTensor([all_h_list, edge_ids]).to(self.device)
92 self.head2edge = torch.LongTensor([edge_ids, all_h_list]).to(self.device)
93 self.tail2edge = torch.LongTensor([edge_ids, all_t_list]).to(self.device)
94 val_one = torch.ones_like(self.all_h_list).float().to(self.device)
95 num_node = self.n_users + self.n_items
96 self.edge2head_mat = self._build_sparse_tensor(self.edge2head, val_one, (num_node, num_edge))
97 self.head2edge_mat = self._build_sparse_tensor(self.head2edge, val_one, (num_edge, num_node))
98 self.tail2edge_mat = self._build_sparse_tensor(self.tail2edge, val_one, (num_edge, num_node))
99 self.num_edge = num_edge
100 self.num_node = num_node
102 # define layers and loss
103 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
104 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size)
105 self.softmax = torch.nn.Softmax(dim=1)
106 self.mf_loss = BPRLoss()
107 self.reg_loss = EmbLoss()
108 self.restore_user_e = None
109 self.restore_item_e = None
111 self.other_parameter_name = ["restore_user_e", "restore_item_e"]
112 # parameters initialization
113 self.apply(xavier_normal_initialization)
115 def _build_sparse_tensor(self, indices, values, size):
116 # Construct the sparse matrix with indices, values and size.
117 return torch.sparse.FloatTensor(indices, values, size).to(self.device)
119 def _get_ego_embeddings(self):
120 # concat of user embeddings and item embeddings
121 user_emb = self.user_embedding.weight
122 item_emb = self.item_embedding.weight
123 ego_embeddings = torch.cat([user_emb, item_emb], dim=0)
124 return ego_embeddings
126 def build_matrix(self, A_values):
127 r"""Get the normalized interaction matrix of users and items according to A_values.
129 Construct the square matrix from the training data and normalize it
130 using the laplace matrix.
132 Args:
133 A_values (torch.cuda.FloatTensor): (num_edge, n_factors)
135 .. math::
136 A_{hat} = D^{-0.5} \times A \times D^{-0.5}
138 Returns:
139 torch.cuda.FloatTensor: Sparse tensor of the normalized interaction matrix. shape: (num_edge, n_factors)
140 """
141 norm_A_values = self.softmax(A_values)
142 factor_edge_weight = []
143 for i in range(self.n_factors):
144 tp_values = norm_A_values[:, i].unsqueeze(1)
145 # (num_edge, 1)
146 d_values = torch.sparse.mm(self.edge2head_mat, tp_values)
147 # (num_node, num_edge) (num_edge, 1) -> (num_node, 1)
148 d_values = torch.clamp(d_values, min=1e-8)
149 try:
150 assert not torch.isnan(d_values).any()
151 except AssertionError:
152 self.logger.info(f"d_values {torch.min(d_values)} {torch.max(d_values)}")
154 d_values = 1.0 / torch.sqrt(d_values)
155 head_term = torch.sparse.mm(self.head2edge_mat, d_values)
156 # (num_edge, num_node) (num_node, 1) -> (num_edge, 1)
158 tail_term = torch.sparse.mm(self.tail2edge_mat, d_values)
159 edge_weight = tp_values * head_term * tail_term
160 factor_edge_weight.append(edge_weight)
161 return factor_edge_weight
163 def forward(self):
164 ego_embeddings = self._get_ego_embeddings()
165 all_embeddings = [ego_embeddings.unsqueeze(1)]
166 # initialize with every factor value as 1
167 A_values = torch.ones((self.num_edge, self.n_factors)).to(self.device)
168 A_values = Variable(A_values, requires_grad=True)
169 for k in range(self.n_layers):
170 layer_embeddings = []
172 # split the input embedding table
173 # .... ego_layer_embeddings is a (n_factors)-length list of embeddings
174 # [n_users+n_items, embed_size/n_factors]
175 ego_layer_embeddings = torch.chunk(ego_embeddings, self.n_factors, 1)
176 for t in range(0, self.n_iterations):
177 iter_embeddings = []
178 A_iter_values = []
179 factor_edge_weight = self.build_matrix(A_values=A_values)
180 for i in range(0, self.n_factors):
181 # update the embeddings via simplified graph convolution layer
182 edge_weight = factor_edge_weight[i]
183 # (num_edge, 1)
184 edge_val = torch.sparse.mm(self.tail2edge_mat, ego_layer_embeddings[i])
185 # (num_edge, dim / n_factors)
186 edge_val = edge_val * edge_weight
187 # (num_edge, dim / n_factors)
188 factor_embeddings = torch.sparse.mm(self.edge2head_mat, edge_val)
189 # (num_node, num_edge) (num_edge, dim) -> (num_node, dim)
191 iter_embeddings.append(factor_embeddings)
193 if t == self.n_iterations - 1:
194 layer_embeddings = iter_embeddings
196 # get the factor-wise embeddings
197 # .... head_factor_embeddings is a dense tensor with the size of [all_h_list, embed_size/n_factors]
198 # .... analogous to tail_factor_embeddings
199 head_factor_embeddings = torch.index_select(factor_embeddings, dim=0, index=self.all_h_list)
200 tail_factor_embeddings = torch.index_select(ego_layer_embeddings[i], dim=0, index=self.all_t_list)
202 # .... constrain the vector length
203 # .... make the following attentive weights within the range of (0,1)
204 # to adapt to torch version
205 head_factor_embeddings = F.normalize(head_factor_embeddings, p=2, dim=1)
206 tail_factor_embeddings = F.normalize(tail_factor_embeddings, p=2, dim=1)
208 # get the attentive weights
209 # .... A_factor_values is a dense tensor with the size of [num_edge, 1]
210 A_factor_values = torch.sum(
211 head_factor_embeddings * torch.tanh(tail_factor_embeddings),
212 dim=1,
213 keepdim=True,
214 )
216 # update the attentive weights
217 A_iter_values.append(A_factor_values)
218 A_iter_values = torch.cat(A_iter_values, dim=1)
219 # (num_edge, n_factors)
220 # add all layer-wise attentive weights up.
221 A_values = A_values + A_iter_values
223 # sum messages of neighbors, [n_users+n_items, embed_size]
224 side_embeddings = torch.cat(layer_embeddings, dim=1)
226 ego_embeddings = side_embeddings
227 # concatenate outputs of all layers
228 all_embeddings += [ego_embeddings.unsqueeze(1)]
230 all_embeddings = torch.cat(all_embeddings, dim=1)
231 # (num_node, n_layer + 1, embedding_size)
232 all_embeddings = torch.mean(all_embeddings, dim=1, keepdim=False)
233 # (num_node, embedding_size)
235 u_g_embeddings = all_embeddings[: self.n_users, :]
236 i_g_embeddings = all_embeddings[self.n_users :, :]
238 return u_g_embeddings, i_g_embeddings
240 def calculate_loss(self, interaction):
241 # clear the storage variable when training
242 if self.restore_user_e is not None or self.restore_item_e is not None:
243 self.restore_user_e, self.restore_item_e = None, None
245 user = interaction[self.USER_ID]
246 pos_item = interaction[self.ITEM_ID]
247 neg_item = interaction[self.NEG_ITEM_ID]
249 user_all_embeddings, item_all_embeddings = self.forward()
250 u_embeddings = user_all_embeddings[user]
251 pos_embeddings = item_all_embeddings[pos_item]
252 neg_embeddings = item_all_embeddings[neg_item]
254 pos_scores = torch.mul(u_embeddings, pos_embeddings).sum(dim=1)
255 neg_scores = torch.mul(u_embeddings, neg_embeddings).sum(dim=1)
256 mf_loss = self.mf_loss(pos_scores, neg_scores)
258 # cul regularized
259 u_ego_embeddings = self.user_embedding(user)
260 pos_ego_embeddings = self.item_embedding(pos_item)
261 neg_ego_embeddings = self.item_embedding(neg_item)
262 reg_loss = self.reg_loss(u_ego_embeddings, pos_ego_embeddings, neg_ego_embeddings)
264 if self.n_factors > 1 and self.cor_weight > 1e-9: # noqa: PLR2004
265 cor_users, cor_items = sample_cor_samples(self.n_users, self.n_items, self.cor_batch_size)
266 cor_users = torch.LongTensor(cor_users).to(self.device)
267 cor_items = torch.LongTensor(cor_items).to(self.device)
268 cor_u_embeddings = user_all_embeddings[cor_users]
269 cor_i_embeddings = item_all_embeddings[cor_items]
270 cor_loss = self.create_cor_loss(cor_u_embeddings, cor_i_embeddings)
271 loss = mf_loss + self.reg_weight * reg_loss + self.cor_weight * cor_loss
272 else:
273 loss = mf_loss + self.reg_weight * reg_loss
274 return loss
276 def create_cor_loss(self, cor_u_embeddings, cor_i_embeddings):
277 r"""Calculate the correlation loss for a sampled users and items.
279 Args:
280 cor_u_embeddings (torch.cuda.FloatTensor): (cor_batch_size, n_factors)
281 cor_i_embeddings (torch.cuda.FloatTensor): (cor_batch_size, n_factors)
283 Returns:
284 torch.Tensor : correlation loss.
285 """
286 cor_loss = None
288 ui_embeddings = torch.cat((cor_u_embeddings, cor_i_embeddings), dim=0)
289 ui_factor_embeddings = torch.chunk(ui_embeddings, self.n_factors, 1)
291 for i in range(0, self.n_factors - 1):
292 x = ui_factor_embeddings[i]
293 # (M + N, emb_size / n_factor)
294 y = ui_factor_embeddings[i + 1]
295 # (M + N, emb_size / n_factor)
296 if i == 0:
297 cor_loss = self._create_distance_correlation(x, y)
298 else:
299 cor_loss += self._create_distance_correlation(x, y)
301 cor_loss /= (self.n_factors + 1.0) * self.n_factors / 2
303 return cor_loss
305 def _create_distance_correlation(self, X1, X2):
306 def _create_centered_distance(X):
307 """X: (batch_size, dim)
308 return: X - E(X)
309 """
310 # calculate the pairwise distance of X
311 # .... A with the size of [batch_size, embed_size/n_factors]
312 # .... D with the size of [batch_size, batch_size]
313 r = torch.sum(X * X, dim=1, keepdim=True)
314 # (N, 1)
315 # (x^2 - 2xy + y^2) -> l2 distance between all vectors
316 value = r - 2 * torch.mm(X, X.T) + r.T
317 zero_value = torch.zeros_like(value)
318 value = torch.where(value > 0.0, value, zero_value)
319 D = torch.sqrt(value + 1e-8)
321 # # calculate the centered distance of X
322 # # .... D with the size of [batch_size, batch_size]
323 # matrix - average over row - average over col + average over matrix
324 D = D - torch.mean(D, dim=0, keepdim=True) - torch.mean(D, dim=1, keepdim=True) + torch.mean(D)
325 return D
327 def _create_distance_covariance(D1, D2):
328 # calculate distance covariance between D1 and D2
329 n_samples = float(D1.size(0))
330 value = torch.sum(D1 * D2) / (n_samples * n_samples)
331 zero_value = torch.zeros_like(value)
332 value = torch.where(value > 0.0, value, zero_value)
333 dcov = torch.sqrt(value + 1e-8)
334 return dcov
336 D1 = _create_centered_distance(X1)
337 D2 = _create_centered_distance(X2)
339 dcov_12 = _create_distance_covariance(D1, D2)
340 dcov_11 = _create_distance_covariance(D1, D1)
341 dcov_22 = _create_distance_covariance(D2, D2)
343 # calculate the distance correlation
344 value = dcov_11 * dcov_22
345 zero_value = torch.zeros_like(value)
346 value = torch.where(value > 0.0, value, zero_value)
347 dcor = dcov_12 / (torch.sqrt(value) + 1e-10)
348 return dcor
350 def predict(self, interaction):
351 user = interaction[self.USER_ID]
352 item = interaction[self.ITEM_ID]
354 u_embedding, i_embedding = self.forward()
356 u_embeddings = u_embedding[user]
357 i_embeddings = i_embedding[item]
358 scores = torch.mul(u_embeddings, i_embeddings).sum(dim=1)
359 return scores
361 def full_sort_predict(self, interaction):
362 user = interaction[self.USER_ID]
363 if self.restore_user_e is None or self.restore_item_e is None:
364 self.restore_user_e, self.restore_item_e = self.forward()
365 u_embeddings = self.restore_user_e[user]
367 scores = torch.matmul(u_embeddings, self.restore_item_e.transpose(0, 1))
369 return scores.view(-1)