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
« 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
5r"""DiffRec
6################################################
7Reference:
8 Wenjie Wang et al. "Diffusion Recommender Model." in SIGIR 2023.
10Reference code:
11 https://github.com/YiyanXu/DiffRec
12"""
14import os
16import numpy as np
17import torch
18import torch.nn.functional as F
19from torch import nn
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
31class AutoEncoder(nn.Module):
32 r"""Guassian Diffusion for large-scale recommendation."""
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__()
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)
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]
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)
63 else:
64 from kmeans_pytorch import kmeans
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
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)
95 self.encoder = nn.ModuleList(encoders)
96 print("Latent dims of each category: ", decode_dim)
98 self.decode_dim = [decode_dim[i][::-1] for i in range(len(decode_dim))]
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)
117 self.apply(xavier_normal_initialization)
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] :]
126 if self.training and self.reparam:
127 latent = self.reparamterization(mu, logvar)
128 else:
129 latent = mu
131 kl_divergence = -0.5 * torch.mean(torch.sum(1 + logvar - mu.pow(2) - logvar.exp(), dim=1))
133 return batch, latent, kl_divergence
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]]
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
155 kl_divergence = -0.5 * torch.mean(torch.sum(1 + logvar - mu.pow(2) - logvar.exp(), dim=1))
157 return torch.cat(tuple(batch_cate), dim=-1), latent, kl_divergence
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)
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)
179 return pred
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 """
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"]
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"]
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)
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)
234 def calculate_loss(self, interaction):
235 user = interaction[self.USER_ID]
236 batch = self.get_rating_matrix(user)
238 batch_cate, batch_latent, vae_kl = self.autoencoder.Encode(batch)
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
249 model_output = self.mlp(x_t, ts)
250 target = {
251 ModelMeanType.START_X: batch_latent,
252 ModelMeanType.EPSILON: noise,
253 }[self.mean_type]
255 assert model_output.shape == target.shape == batch_latent.shape
257 mse = mean_flat((target - model_output) ** 2)
259 reloss = self.reweight_loss(batch_latent, x_t, mse, ts, target, model_output, device)
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)
266 self.update_Lt_history(ts, reloss)
268 diff_loss = (reloss / pt).mean()
270 batch_recon = self.autoencoder.Decode(batch_latent_recon)
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)
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
285 self.update_count_vae += 1
286 self.update_count += 1
287 vae_loss = compute_loss(batch_recon, batch_cate) + anneal * vae_kl
289 loss = lamda * diff_loss + vae_loss
291 return loss
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
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
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