Coverage for hopwise/model/general_recommender/ldiffrec.py: 60%

179 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2023/10/6 

2# @Author : Enze Liu 

3# @Email : enzeeliu@foxmail.com 

4 

5r"""DiffRec 

6################################################ 

7Reference: 

8 Wenjie Wang et al. "Diffusion Recommender Model." in SIGIR 2023. 

9 

10Reference code: 

11 https://github.com/YiyanXu/DiffRec 

12""" 

13 

14import os 

15 

16import numpy as np 

17import torch 

18import torch.nn.functional as F 

19from torch import nn 

20 

21from hopwise.model.general_recommender.diffrec import ( 

22 DNN, 

23 DiffRec, 

24 ModelMeanType, 

25 mean_flat, 

26) 

27from hopwise.model.init import xavier_normal_initialization 

28from hopwise.model.layers import MLPLayers 

29 

30 

31class AutoEncoder(nn.Module): 

32 r"""Guassian Diffusion for large-scale recommendation.""" 

33 

34 def __init__( 

35 self, 

36 item_emb, 

37 n_cate, 

38 in_dims, 

39 out_dims, 

40 device, 

41 act_func, 

42 reparam=True, 

43 dropout=0.1, 

44 ): 

45 super().__init__() 

46 

47 self.item_emb = item_emb 

48 self.n_cate = n_cate 

49 self.in_dims = in_dims 

50 self.out_dims = out_dims 

51 self.act_func = act_func 

52 self.n_item = len(item_emb) 

53 self.reparam = reparam 

54 self.dropout = nn.Dropout(dropout) 

55 

56 if n_cate == 1: # no clustering 

57 in_dims_temp = [self.n_item + 1] + self.in_dims[:-1] + [self.in_dims[-1] * 2] 

58 out_dims_temp = [self.in_dims[-1]] + self.out_dims + [self.n_item + 1] 

59 

60 self.encoder = MLPLayers(in_dims_temp, activation=self.act_func) 

61 self.decoder = MLPLayers(out_dims_temp, activation=self.act_func, last_activation=False) 

62 

63 else: 

64 from kmeans_pytorch import kmeans 

65 

66 self.cluster_ids, _ = kmeans(X=item_emb, num_clusters=n_cate, distance="euclidean", device=device) 

67 # cluster_ids(labels): [0, 1, 2, 2, 1, 0, 0, ...] 

68 category_idx = [] 

69 for i in range(n_cate): 

70 idx = np.argwhere(self.cluster_ids.numpy() == i).flatten().tolist() 

71 category_idx.append(torch.tensor(idx, dtype=int) + 1) 

72 self.category_idx = ( 

73 category_idx # [cate1: [iid1, iid2, ...], cate2: [iid3, iid4, ...], cate3: [iid5, iid6, ...]] 

74 ) 

75 self.category_map = torch.cat(tuple(category_idx), dim=-1) # map 

76 self.category_len = [len(self.category_idx[i]) for i in range(n_cate)] # item num in each category 

77 print("category length: ", self.category_len) 

78 assert sum(self.category_len) == self.n_item 

79 

80 ##### Build the Encoder and Decoder ##### 

81 encoders = [] 

82 decode_dim = [] 

83 for i in range(n_cate): 

84 if i == n_cate - 1: 

85 latent_dims = list(self.in_dims - np.array(decode_dim).sum(axis=0)) 

86 else: 

87 latent_dims = [ 

88 int(self.category_len[i] / self.n_item * self.in_dims[j]) for j in range(len(self.in_dims)) 

89 ] 

90 latent_dims = [latent_dims[j] if latent_dims[j] != 0 else 1 for j in range(len(self.in_dims))] 

91 in_dims_temp = [self.category_len[i]] + latent_dims[:-1] + [latent_dims[-1] * 2] 

92 encoders.append(MLPLayers(in_dims_temp, activation=self.act_func)) 

93 decode_dim.append(latent_dims) 

94 

95 self.encoder = nn.ModuleList(encoders) 

96 print("Latent dims of each category: ", decode_dim) 

97 

98 self.decode_dim = [decode_dim[i][::-1] for i in range(len(decode_dim))] 

99 

100 if len(out_dims) == 0: # one-layer decoder: [encoder_dim_sum, n_item] 

101 out_dim = self.in_dims[-1] 

102 self.decoder = MLPLayers([out_dim, self.n_item], activation=None) 

103 else: # multi-layer decoder: [encoder_dim, hidden_size, cate_num] 

104 # decoder_modules = [[] for _ in range(n_cate)] 

105 decoders = [] 

106 for i in range(n_cate): 

107 out_dims_temp = self.decode_dim[i] + [self.category_len[i]] 

108 decoders.append( 

109 MLPLayers( 

110 out_dims_temp, 

111 activation=self.act_func, 

112 last_activation=False, 

113 ) 

114 ) 

115 self.decoder = nn.ModuleList(decoders) 

116 

117 self.apply(xavier_normal_initialization) 

118 

119 def Encode(self, batch): 

120 batch = self.dropout(batch) 

121 if self.n_cate == 1: 

122 hidden = self.encoder(batch) 

123 mu = hidden[:, : self.in_dims[-1]] 

124 logvar = hidden[:, self.in_dims[-1] :] 

125 

126 if self.training and self.reparam: 

127 latent = self.reparamterization(mu, logvar) 

128 else: 

129 latent = mu 

130 

131 kl_divergence = -0.5 * torch.mean(torch.sum(1 + logvar - mu.pow(2) - logvar.exp(), dim=1)) 

132 

133 return batch, latent, kl_divergence 

134 

135 else: 

136 batch_cate = [] 

137 for i in range(self.n_cate): 

138 batch_cate.append(batch[:, self.category_idx[i]]) 

139 # [batch_size, n_items] -> [[batch_size, n1_items], [batch_size, n2_items], [batch_size, n3_items]] 

140 latent_mu = [] 

141 latent_logvar = [] 

142 for i in range(self.n_cate): 

143 hidden = self.encoder[i](batch_cate[i]) 

144 latent_mu.append(hidden[:, : self.decode_dim[i][0]]) 

145 latent_logvar.append(hidden[:, self.decode_dim[i][0] :]) 

146 # latent: [[batch_size, latent_size1], [batch_size, latent_size2], [batch_size, latent_size3]] 

147 

148 mu = torch.cat(tuple(latent_mu), dim=-1) 

149 logvar = torch.cat(tuple(latent_logvar), dim=-1) 

150 if self.training and self.reparam: 

151 latent = self.reparamterization(mu, logvar) 

152 else: 

153 latent = mu 

154 

155 kl_divergence = -0.5 * torch.mean(torch.sum(1 + logvar - mu.pow(2) - logvar.exp(), dim=1)) 

156 

157 return torch.cat(tuple(batch_cate), dim=-1), latent, kl_divergence 

158 

159 def reparamterization(self, mu, logvar): 

160 std = torch.exp(0.5 * logvar) 

161 eps = torch.randn_like(std) 

162 return eps.mul(std).add_(mu) 

163 

164 def Decode(self, batch): 

165 if len(self.out_dims) == 0 or self.n_cate == 1: # one-layer decoder 

166 return self.decoder(batch) 

167 else: 

168 batch_cate = [] 

169 start = 0 

170 for i in range(self.n_cate): 

171 end = start + self.decode_dim[i][0] 

172 batch_cate.append(batch[:, start:end]) 

173 start = end 

174 pred_cate = [] 

175 for i in range(self.n_cate): 

176 pred_cate.append(self.decoder[i](batch_cate[i])) 

177 pred = torch.cat(tuple(pred_cate), dim=-1) 

178 

179 return pred 

180 

181 

182class LDiffRec(DiffRec): 

183 r"""L-DiffRec clusters items into groups, compresses the interaction vector over each group into a 

184 low-dimensional latent vector via a group-specific VAE, and conducts the forward and reverse 

185 diffusion processes in the latent space. 

186 """ 

187 

188 def __init__(self, config, dataset): 

189 super().__init__(config, dataset) 

190 self.n_cate = config["n_cate"] 

191 self.reparam = config["reparam"] 

192 self.ae_act_func = config["ae_act_func"] 

193 self.in_dims = config["in_dims"] 

194 self.out_dims = config["out_dims"] 

195 

196 # control loss in training 

197 self.update_count = 0 

198 self.update_count_vae = 0 

199 self.lamda = config["lamda"] 

200 self.anneal_cap = config["anneal_cap"] 

201 self.anneal_steps = config["anneal_steps"] 

202 self.vae_anneal_cap = config["vae_anneal_cap"] 

203 self.vae_anneal_steps = config["vae_anneal_steps"] 

204 

205 out_dims = self.out_dims 

206 in_dims = self.in_dims[::-1] 

207 emb_path = os.path.join(dataset.dataset_path, "item_emb.npy") 

208 if self.n_cate > 1: 

209 if not os.path.exists(emb_path): 

210 self.logger.exception("The item embedding file must be given when n_cate>1.") 

211 item_emb = torch.from_numpy(np.load(emb_path, allow_pickle=True)) 

212 else: 

213 item_emb = torch.zeros((self.n_items - 1, 64)) 

214 self.autoencoder = AutoEncoder( 

215 item_emb, 

216 self.n_cate, 

217 in_dims, 

218 out_dims, 

219 self.device, 

220 self.ae_act_func, 

221 self.reparam, 

222 ).to(self.device) 

223 

224 self.latent_size = in_dims[-1] 

225 dims = [self.latent_size] + config["dims_dnn"] + [self.latent_size] 

226 self.mlp = DNN( 

227 dims=dims, 

228 emb_size=self.emb_size, 

229 time_type="cat", 

230 norm=self.norm, 

231 act_func=self.mlp_act_func, 

232 ).to(self.device) 

233 

234 def calculate_loss(self, interaction): 

235 user = interaction[self.USER_ID] 

236 batch = self.get_rating_matrix(user) 

237 

238 batch_cate, batch_latent, vae_kl = self.autoencoder.Encode(batch) 

239 

240 # calculate loss in diffusion 

241 batch_size, device = batch_latent.size(0), batch_latent.device 

242 ts, pt = self.sample_timesteps(batch_size, device, "importance") 

243 noise = torch.randn_like(batch_latent) 

244 if self.noise_scale != 0.0: 

245 x_t = self.q_sample(batch_latent, ts, noise) 

246 else: 

247 x_t = batch_latent 

248 

249 model_output = self.mlp(x_t, ts) 

250 target = { 

251 ModelMeanType.START_X: batch_latent, 

252 ModelMeanType.EPSILON: noise, 

253 }[self.mean_type] 

254 

255 assert model_output.shape == target.shape == batch_latent.shape 

256 

257 mse = mean_flat((target - model_output) ** 2) 

258 

259 reloss = self.reweight_loss(batch_latent, x_t, mse, ts, target, model_output, device) 

260 

261 if self.mean_type == ModelMeanType.START_X: 

262 batch_latent_recon = model_output 

263 else: 

264 batch_latent_recon = self._predict_xstart_from_eps(x_t, ts, model_output) 

265 

266 self.update_Lt_history(ts, reloss) 

267 

268 diff_loss = (reloss / pt).mean() 

269 

270 batch_recon = self.autoencoder.Decode(batch_latent_recon) 

271 

272 if self.anneal_steps > 0: 

273 lamda = max( 

274 (1.0 - self.update_count / self.anneal_steps) * self.lamda, 

275 self.anneal_cap, 

276 ) 

277 else: 

278 lamda = max(self.lamda, self.anneal_cap) 

279 

280 if self.vae_anneal_steps > 0: 

281 anneal = min(self.vae_anneal_cap, 1.0 * self.update_count_vae / self.vae_anneal_steps) 

282 else: 

283 anneal = self.vae_anneal_cap 

284 

285 self.update_count_vae += 1 

286 self.update_count += 1 

287 vae_loss = compute_loss(batch_recon, batch_cate) + anneal * vae_kl 

288 

289 loss = lamda * diff_loss + vae_loss 

290 

291 return loss 

292 

293 def full_sort_predict(self, interaction): 

294 user = interaction[self.USER_ID] 

295 batch = self.get_rating_matrix(user) 

296 _, batch_latent, _ = self.autoencoder.Encode(batch) 

297 batch_latent_recon = super().p_sample(batch_latent) 

298 prediction = self.autoencoder.Decode(batch_latent_recon) # [batch_size, n1_items + n2_items + n3_items] 

299 if self.n_cate > 1: 

300 transform = torch.zeros((prediction.shape[0], prediction.shape[1] + 1)).to(prediction.device) 

301 transform[:, self.autoencoder.category_map] = prediction 

302 else: 

303 transform = prediction 

304 return transform 

305 

306 def predict(self, interaction): 

307 item = interaction[self.ITEM_ID] 

308 x_t = self.full_sort_predict(interaction) 

309 scores = x_t[torch.arange(len(item)).to(self.device), item] 

310 return scores 

311 

312 

313def compute_loss(recon_x, x): 

314 return -torch.mean(torch.sum(F.log_softmax(recon_x, 1) * x, -1)) # multinomial log likelihood in MultVAE