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

1# @Time : 2020/8/26 

2# @Author : Gaole He 

3# @Email : hegaole@ruc.edu.cn 

4 

5# UPDATE: 

6# @Time : 2020/9/16 

7# @Author : Shanlei Mu 

8# @Email : slmu@ruc.edu.cn 

9 

10r"""DGCF 

11################################################ 

12Reference: 

13 Wang Xiang et al. "Disentangled Graph Collaborative Filtering." in SIGIR 2020. 

14 

15Reference code: 

16 https://github.com/xiangwang1223/disentangled_graph_collaborative_filtering 

17""" 

18 

19import random as rd 

20 

21import numpy as np 

22import torch 

23import torch.nn.functional as F 

24from torch import nn 

25from torch.autograd import Variable 

26 

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 

31 

32 

33def sample_cor_samples(n_users, n_items, cor_batch_size): 

34 r"""This is a function that sample item ids and user ids. 

35 

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 

40 

41 Returns: 

42 list: cor_users, cor_items. The result sampled ids with both as cor_batch_size long. 

43 

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) 

50 

51 return cor_users, cor_items 

52 

53 

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 """ 

60 

61 input_type = InputType.PAIRWISE 

62 

63 def __init__(self, config, dataset): 

64 super().__init__(config, dataset) 

65 

66 # load dataset info 

67 self.interaction_matrix = dataset.inter_matrix(form="coo").astype(np.float32) 

68 

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 

80 

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 

101 

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 

110 

111 self.other_parameter_name = ["restore_user_e", "restore_item_e"] 

112 # parameters initialization 

113 self.apply(xavier_normal_initialization) 

114 

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) 

118 

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 

125 

126 def build_matrix(self, A_values): 

127 r"""Get the normalized interaction matrix of users and items according to A_values. 

128 

129 Construct the square matrix from the training data and normalize it 

130 using the laplace matrix. 

131 

132 Args: 

133 A_values (torch.cuda.FloatTensor): (num_edge, n_factors) 

134 

135 .. math:: 

136 A_{hat} = D^{-0.5} \times A \times D^{-0.5} 

137 

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)}") 

153 

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) 

157 

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 

162 

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 = [] 

171 

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) 

190 

191 iter_embeddings.append(factor_embeddings) 

192 

193 if t == self.n_iterations - 1: 

194 layer_embeddings = iter_embeddings 

195 

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) 

201 

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) 

207 

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 ) 

215 

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 

222 

223 # sum messages of neighbors, [n_users+n_items, embed_size] 

224 side_embeddings = torch.cat(layer_embeddings, dim=1) 

225 

226 ego_embeddings = side_embeddings 

227 # concatenate outputs of all layers 

228 all_embeddings += [ego_embeddings.unsqueeze(1)] 

229 

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) 

234 

235 u_g_embeddings = all_embeddings[: self.n_users, :] 

236 i_g_embeddings = all_embeddings[self.n_users :, :] 

237 

238 return u_g_embeddings, i_g_embeddings 

239 

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 

244 

245 user = interaction[self.USER_ID] 

246 pos_item = interaction[self.ITEM_ID] 

247 neg_item = interaction[self.NEG_ITEM_ID] 

248 

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] 

253 

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) 

257 

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) 

263 

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 

275 

276 def create_cor_loss(self, cor_u_embeddings, cor_i_embeddings): 

277 r"""Calculate the correlation loss for a sampled users and items. 

278 

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) 

282 

283 Returns: 

284 torch.Tensor : correlation loss. 

285 """ 

286 cor_loss = None 

287 

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) 

290 

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) 

300 

301 cor_loss /= (self.n_factors + 1.0) * self.n_factors / 2 

302 

303 return cor_loss 

304 

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) 

320 

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 

326 

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 

335 

336 D1 = _create_centered_distance(X1) 

337 D2 = _create_centered_distance(X2) 

338 

339 dcov_12 = _create_distance_covariance(D1, D2) 

340 dcov_11 = _create_distance_covariance(D1, D1) 

341 dcov_22 = _create_distance_covariance(D2, D2) 

342 

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 

349 

350 def predict(self, interaction): 

351 user = interaction[self.USER_ID] 

352 item = interaction[self.ITEM_ID] 

353 

354 u_embedding, i_embedding = self.forward() 

355 

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 

360 

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] 

366 

367 scores = torch.matmul(u_embeddings, self.restore_item_e.transpose(0, 1)) 

368 

369 return scores.view(-1)