Coverage for hopwise/model/general_recommender/diffrec.py: 75%
309 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 copy
15import enum
16import math
18import numpy as np
19import torch
20import torch.nn.functional as F
21from torch import nn
23from hopwise.model.abstract_recommender import AutoEncoderMixin, GeneralRecommender
24from hopwise.model.init import xavier_normal_initialization
25from hopwise.model.layers import MLPLayers
26from hopwise.utils import InputType
29class ModelMeanType(enum.Enum):
30 START_X = enum.auto() # the model predicts x_0
31 EPSILON = enum.auto() # the model predicts epsilon
34class DNN(nn.Module):
35 r"""A deep neural network for the reverse diffusion preocess."""
37 def __init__(
38 self,
39 dims: list,
40 emb_size: int,
41 time_type="cat",
42 act_func="tanh",
43 norm=False,
44 dropout=0.5,
45 ):
46 super().__init__()
47 self.dims = dims
48 self.time_type = time_type
49 self.time_emb_dim = emb_size
50 self.norm = norm
52 self.emb_layer = nn.Linear(self.time_emb_dim, self.time_emb_dim)
54 if self.time_type == "cat":
55 # Concatenate timestep embedding with input
56 self.dims[0] += self.time_emb_dim
57 else:
58 raise ValueError("Unimplemented timestep embedding type %s" % self.time_type)
60 self.mlp_layers = MLPLayers(layers=self.dims, dropout=0, activation=act_func, last_activation=False)
61 self.drop = nn.Dropout(dropout)
63 self.apply(xavier_normal_initialization)
65 def forward(self, x, timesteps):
66 time_emb = timestep_embedding(timesteps, self.time_emb_dim).to(x.device)
67 emb = self.emb_layer(time_emb)
68 if self.norm:
69 x = F.normalize(x)
70 x = self.drop(x)
71 h = torch.cat([x, emb], dim=-1)
72 h = self.mlp_layers(h)
73 return h
76class DiffRec(GeneralRecommender, AutoEncoderMixin):
77 r"""DiffRec is a generative recommender model which infers users' interaction probabilities in a denoising manner.
78 Note that DiffRec simultaneously ranks all items for each user.
79 We implement the the DiffRec model with only user dataloader.
80 """
82 input_type = InputType.USERWISE
84 def __init__(self, config, dataset):
85 super().__init__(config, dataset)
87 if config["mean_type"] == "x0":
88 self.mean_type = ModelMeanType.START_X
89 elif config["mean_type"] == "eps":
90 self.mean_type = ModelMeanType.EPSILON
91 else:
92 raise ValueError("Unimplemented mean type %s" % config["mean_type"])
93 self.time_aware = config["time-aware"]
94 self.w_max = config["w_max"]
95 self.w_min = config["w_min"]
96 self.build_histroy_items(dataset)
98 self.noise_schedule = config["noise_schedule"]
99 self.noise_scale = config["noise_scale"]
100 self.noise_min = config["noise_min"]
101 self.noise_max = config["noise_max"]
102 self.steps = config["steps"]
103 self.beta_fixed = config["beta_fixed"]
104 self.emb_size = config["embedding_size"]
105 self.norm = config["norm"] # True or False
106 self.reweight = config["reweight"] # reweight the loss for different timesteps
107 if self.noise_scale == 0.0:
108 self.reweight = False
109 self.sampling_noise = config["sampling_noise"] # whether sample noise during predict
110 self.sampling_steps = config["sampling_steps"]
111 self.mlp_act_func = config["mlp_act_func"]
112 assert self.sampling_steps <= self.steps, "Too much steps in inference."
114 self.history_num_per_term = config["history_num_per_term"]
115 self.Lt_history = torch.zeros(self.steps, self.history_num_per_term, dtype=torch.float64).to(self.device)
116 self.Lt_count = torch.zeros(self.steps, dtype=int).to(self.device)
118 dims = [self.n_items] + config["dims_dnn"] + [self.n_items]
120 self.mlp = DNN(
121 dims=dims,
122 emb_size=self.emb_size,
123 time_type="cat",
124 norm=self.norm,
125 act_func=self.mlp_act_func,
126 ).to(self.device)
128 if self.noise_scale != 0.0:
129 self.betas = torch.tensor(self.get_betas(), dtype=torch.float64).to(self.device)
130 if self.beta_fixed:
131 self.betas[0] = 0.00001 # Deep Unsupervised Learning using Noneequilibrium Thermodynamics 2.4.1
132 # The variance \beta_1 of the first step is fixed to a small constant to prevent overfitting.
133 assert len(self.betas.shape) == 1, "betas must be 1-D"
134 assert len(self.betas) == self.steps, "num of betas must equal to diffusion steps"
135 assert (self.betas > 0).all() and (self.betas <= 1).all(), "betas out of range"
137 self.calculate_for_diffusion()
139 def build_histroy_items(self, dataset):
140 r"""Add time-aware reweighting to the original user-item interaction matrix when config['time-aware'] is True.""" # noqa: E501
141 if not self.time_aware:
142 super().build_histroy_items(dataset)
143 else:
144 inter_feat = copy.deepcopy(dataset.inter_feat)
145 inter_feat.sort(dataset.time_field)
146 user_ids, item_ids = (
147 inter_feat[dataset.uid_field].numpy(),
148 inter_feat[dataset.iid_field].numpy(),
149 )
151 w_max = self.w_max
152 w_min = self.w_min
153 values = np.zeros(len(inter_feat))
155 row_num = dataset.user_num
156 row_ids, col_ids = user_ids, item_ids
158 for uid in range(1, row_num + 1):
159 uindex = np.argwhere(user_ids == uid).flatten()
160 int_num = len(uindex)
161 weight = np.linspace(w_min, w_max, int_num)
162 values[uindex] = weight
164 history_len = np.zeros(row_num, dtype=np.int64)
165 for row_id in row_ids:
166 history_len[row_id] += 1
168 max_inter_num = np.max(history_len)
169 col_num = max_inter_num
171 history_matrix = np.zeros((row_num, col_num), dtype=np.int64)
172 history_value = np.zeros((row_num, col_num))
173 history_len[:] = 0
175 for row_id, value, col_id in zip(row_ids, values, col_ids):
176 if history_len[row_id] >= col_num:
177 continue
178 history_matrix[row_id, history_len[row_id]] = col_id
179 history_value[row_id, history_len[row_id]] = value
180 history_len[row_id] += 1
182 self.history_item_id = torch.LongTensor(history_matrix)
183 self.history_item_value = torch.FloatTensor(history_value)
184 self.history_item_id = self.history_item_id.to(self.device)
185 self.history_item_value = self.history_item_value.to(self.device)
187 def get_betas(self):
188 r"""Given the schedule name, create the betas for the diffusion process."""
189 if self.noise_schedule in ("linear", "linear-var"):
190 start = self.noise_scale * self.noise_min
191 end = self.noise_scale * self.noise_max
192 if self.noise_schedule == "linear":
193 return np.linspace(start, end, self.steps, dtype=np.float64)
194 else:
195 return betas_from_linear_variance(self.steps, np.linspace(start, end, self.steps, dtype=np.float64))
196 elif self.noise_schedule == "cosine":
197 return betas_for_alpha_bar(self.steps, lambda t: math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2)
198 # Deep Unsupervised Learning using Noneequilibrium Thermodynamics 2.4.1
199 elif self.noise_schedule == "binomial":
200 ts = np.arange(self.steps)
201 betas = [1 / (self.steps - t + 1) for t in ts]
202 return betas
203 else:
204 raise NotImplementedError(f"unknown beta schedule: {self.noise_schedule}!")
206 def calculate_for_diffusion(self):
207 r"""Calculate the coefficients for the diffusion process."""
208 alphas = 1.0 - self.betas
209 # [alpha_{1}, ..., alpha_{1}*...*alpha_{T}] shape (steps,)
210 self.alphas_cumprod = torch.cumprod(alphas, axis=0).to(self.device)
211 # alpha_{t-1}
212 self.alphas_cumprod_prev = torch.cat([torch.tensor([1.0]).to(self.device), self.alphas_cumprod[:-1]]).to(
213 self.device
214 )
215 # alpha_{t+1}
216 self.alphas_cumprod_next = torch.cat([self.alphas_cumprod[1:], torch.tensor([0.0]).to(self.device)]).to(
217 self.device
218 )
219 assert self.alphas_cumprod_prev.shape == (self.steps,)
221 self.sqrt_alphas_cumprod = torch.sqrt(self.alphas_cumprod)
222 self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - self.alphas_cumprod)
223 self.log_one_minus_alphas_cumprod = torch.log(1.0 - self.alphas_cumprod)
224 self.sqrt_recip_alphas_cumprod = torch.sqrt(1.0 / self.alphas_cumprod)
225 self.sqrt_recipm1_alphas_cumprod = torch.sqrt(1.0 / self.alphas_cumprod - 1)
227 self.posterior_variance = self.betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
229 self.posterior_log_variance_clipped = torch.log(
230 torch.cat([self.posterior_variance[1].unsqueeze(0), self.posterior_variance[1:]])
231 )
232 # Eq.10 coef for x_theta
233 self.posterior_mean_coef1 = self.betas * torch.sqrt(self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
234 # Eq.10 coef for x_t
235 self.posterior_mean_coef2 = (1.0 - self.alphas_cumprod_prev) * torch.sqrt(alphas) / (1.0 - self.alphas_cumprod)
237 def p_sample(self, x_start):
238 r"""Generate users' interaction probabilities in a denoising manner.
240 Args:
241 x_start (torch.FloatTensor): the input tensor that contains user's history interaction matrix,
242 for DiffRec shape: [batch_size, n_items]
243 for LDiffRec shape: [batch_size, hidden_size]
245 Returns:
246 torch.FloatTensor: the interaction probabilities,
247 for DiffRec shape: [batch_size, n_items]
248 for LDiffRec shape: [batch_size, hidden_size]
249 """
250 steps = self.sampling_steps
251 if steps == 0:
252 x_t = x_start
253 else:
254 t = torch.tensor([steps - 1] * x_start.shape[0]).to(x_start.device)
255 x_t = self.q_sample(x_start, t)
257 indices = list(range(self.steps))[::-1]
259 if self.noise_scale == 0.0:
260 for i in indices:
261 t = torch.tensor([i] * x_t.shape[0]).to(x_start.device)
262 x_t = self.mlp(x_t, t)
263 return x_t
265 for i in indices:
266 t = torch.tensor([i] * x_t.shape[0]).to(x_start.device)
267 out = self.p_mean_variance(x_t, t)
268 if self.sampling_noise:
269 noise = torch.randn_like(x_t)
270 nonzero_mask = (t != 0).float().view(-1, *([1] * (len(x_t.shape) - 1))) # no noise when t == 0
271 x_t = out["mean"] + nonzero_mask * torch.exp(0.5 * out["log_variance"]) * noise
272 else:
273 x_t = out["mean"]
274 return x_t
276 def full_sort_predict(self, interaction):
277 user = interaction[self.USER_ID]
278 x_start = self.get_rating_matrix(user)
279 scores = self.p_sample(x_start)
280 return scores
282 def predict(self, interaction):
283 item = interaction[self.ITEM_ID]
284 x_t = self.full_sort_predict(interaction)
285 scores = x_t[torch.arange(len(item)).to(self.device), item]
286 return scores
288 def calculate_loss(self, interaction):
289 user = interaction[self.USER_ID]
290 x_start = self.get_rating_matrix(user)
292 batch_size, device = x_start.size(0), x_start.device
293 ts, pt = self.sample_timesteps(batch_size, device, "importance")
294 noise = torch.randn_like(x_start)
295 if self.noise_scale != 0.0:
296 x_t = self.q_sample(x_start, ts, noise)
297 else:
298 x_t = x_start
300 model_output = self.mlp(x_t, ts)
301 target = {
302 ModelMeanType.START_X: x_start,
303 ModelMeanType.EPSILON: noise,
304 }[self.mean_type]
306 assert model_output.shape == target.shape == x_start.shape
308 mse = mean_flat((target - model_output) ** 2)
310 reloss = self.reweight_loss(x_start, x_t, mse, ts, target, model_output, device)
311 self.update_Lt_history(ts, reloss)
313 # importance sampling
314 reloss /= pt
315 mean_loss = reloss.mean()
316 return mean_loss
318 def reweight_loss(self, x_start, x_t, mse, ts, target, model_output, device):
319 if self.reweight:
320 if self.mean_type == ModelMeanType.START_X:
321 # Eq.11
322 weight = self.SNR(ts - 1) - self.SNR(ts)
323 # Eq.12
324 weight = torch.where((ts == 0), 1.0, weight)
325 loss = mse
326 elif self.mean_type == ModelMeanType.EPSILON:
327 weight = (1 - self.alphas_cumprod[ts]) / (
328 (1 - self.alphas_cumprod_prev[ts]) ** 2 * (1 - self.betas[ts])
329 )
330 weight = torch.where((ts == 0), 1.0, weight)
331 likelihood = mean_flat((x_start - self._predict_xstart_from_eps(x_t, ts, model_output)) ** 2 / 2.0)
332 loss = torch.where((ts == 0), likelihood, mse)
333 else:
334 weight = torch.tensor([1.0] * len(target)).to(device)
335 loss = mse
336 reloss = weight * loss
337 return reloss
339 def update_Lt_history(self, ts, reloss):
340 # update Lt_history & Lt_count
341 for t, loss in zip(ts, reloss):
342 if self.Lt_count[t] == self.history_num_per_term:
343 Lt_history_old = self.Lt_history.clone()
344 self.Lt_history[t, :-1] = Lt_history_old[t, 1:]
345 self.Lt_history[t, -1] = loss.detach()
346 else:
347 try:
348 self.Lt_history[t, self.Lt_count[t]] = loss.detach()
349 self.Lt_count[t] += 1
350 except Exception:
351 print(t)
352 print(self.Lt_count[t])
353 print(loss)
354 raise ValueError
356 def sample_timesteps(self, batch_size, device, method="uniform", uniform_prob=0.001):
357 if method == "importance": # importance sampling
358 if not (self.Lt_count == self.history_num_per_term).all():
359 return self.sample_timesteps(batch_size, device, method="uniform")
361 Lt_sqrt = torch.sqrt(torch.mean(self.Lt_history**2, axis=-1))
362 pt_all = Lt_sqrt / torch.sum(Lt_sqrt)
363 pt_all *= 1 - uniform_prob
364 pt_all += uniform_prob / len(pt_all) # ensure the least prob > uniform_prob
366 assert pt_all.sum(-1) - 1.0 < 1e-5 # noqa: PLR2004
368 t = torch.multinomial(pt_all, num_samples=batch_size, replacement=True)
369 pt = pt_all.gather(dim=0, index=t) * len(pt_all)
371 return t, pt
373 elif method == "uniform": # uniform sampling
374 t = torch.randint(0, self.steps, (batch_size,), device=device).long()
375 pt = torch.ones_like(t).float()
377 return t, pt
379 else:
380 raise ValueError
382 def q_sample(self, x_start, t, noise=None):
383 if noise is None:
384 noise = torch.randn_like(x_start)
385 assert noise.shape == x_start.shape
386 return (
387 self._extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start
388 + self._extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise
389 )
391 def q_posterior_mean_variance(self, x_start, x_t, t):
392 r"""Compute the mean and variance of the diffusion posterior:
393 q(x_{t-1} | x_t, x_0)
394 """
395 assert x_start.shape == x_t.shape
396 posterior_mean = (
397 self._extract_into_tensor(self.posterior_mean_coef1, t, x_t.shape) * x_start
398 + self._extract_into_tensor(self.posterior_mean_coef2, t, x_t.shape) * x_t
399 )
400 posterior_variance = self._extract_into_tensor(self.posterior_variance, t, x_t.shape)
401 posterior_log_variance_clipped = self._extract_into_tensor(self.posterior_log_variance_clipped, t, x_t.shape)
402 assert (
403 posterior_mean.shape[0]
404 == posterior_variance.shape[0]
405 == posterior_log_variance_clipped.shape[0]
406 == x_start.shape[0]
407 )
408 return posterior_mean, posterior_variance, posterior_log_variance_clipped
410 def p_mean_variance(self, x, t):
411 r"""Apply the model to get p(x_{t-1} | x_t), as well as a prediction of
412 the initial x, x_0.
413 """
414 B, C = x.shape[:2]
415 assert t.shape == (B,)
416 model_output = self.mlp(x, t)
418 model_variance = self.posterior_variance
419 model_log_variance = self.posterior_log_variance_clipped
421 model_variance = self._extract_into_tensor(model_variance, t, x.shape)
422 model_log_variance = self._extract_into_tensor(model_log_variance, t, x.shape)
424 if self.mean_type == ModelMeanType.START_X:
425 pred_xstart = model_output
426 elif self.mean_type == ModelMeanType.EPSILON:
427 pred_xstart = self._predict_xstart_from_eps(x, t, eps=model_output)
428 else:
429 raise NotImplementedError(self.mean_type)
431 model_mean, _, _ = self.q_posterior_mean_variance(x_start=pred_xstart, x_t=x, t=t)
433 assert model_mean.shape == model_log_variance.shape == pred_xstart.shape == x.shape
435 return {
436 "mean": model_mean,
437 "variance": model_variance,
438 "log_variance": model_log_variance,
439 "pred_xstart": pred_xstart,
440 }
442 def _predict_xstart_from_eps(self, x_t, t, eps):
443 assert x_t.shape == eps.shape
444 return (
445 self._extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t
446 - self._extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * eps
447 )
449 def SNR(self, t):
450 r"""Compute the signal-to-noise ratio for a single timestep."""
451 self.alphas_cumprod = self.alphas_cumprod.to(t.device)
452 return self.alphas_cumprod[t] / (1 - self.alphas_cumprod[t])
454 def _extract_into_tensor(self, arr, timesteps, broadcast_shape):
455 r"""Extract values from a 1-D torch tensor for a batch of indices.
457 Args:
458 arr (torch.Tensor): the 1-D torch tensor.
459 timesteps (torch.Tensor): a tensor of indices into the array to extract.
460 broadcast_shape (torch.Size): a larger shape of K dimensions with the batch
461 dimension equal to the length of timesteps.
463 Returns:
464 torch.Tensor: a tensor of shape [batch_size, 1, ...] where the shape has K dims.
465 """
466 # res = torch.from_numpy(arr).to(device=timesteps.device)[timesteps].float()
467 arr = arr.to(timesteps.device)
468 res = arr[timesteps].float()
469 while len(res.shape) < len(broadcast_shape):
470 res = res[..., None]
471 return res.expand(broadcast_shape)
474def betas_from_linear_variance(steps, variance, max_beta=0.999):
475 alpha_bar = 1 - variance
476 betas = []
477 betas.append(1 - alpha_bar[0])
478 for i in range(1, steps):
479 betas.append(min(1 - alpha_bar[i] / alpha_bar[i - 1], max_beta))
480 return np.array(betas)
483def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999):
484 r"""Create a beta schedule that discretizes the given alpha_t_bar function,
485 which defines the cumulative product of (1-beta) over time from t = [0,1].
487 Args:
488 num_diffusion_timesteps (int): the number of betas to produce.
489 alpha_bar (Callable): a lambda that takes an argument t from 0 to 1 and
490 produces the cumulative product of (1-beta) up to that
491 part of the diffusion process.
492 max_beta (int): the maximum beta to use; use values lower than 1 to
493 prevent singularities.
495 Returns:
496 np.ndarray: a 1-D array of beta values.
497 """
498 betas = []
499 for i in range(num_diffusion_timesteps):
500 t1 = i / num_diffusion_timesteps
501 t2 = (i + 1) / num_diffusion_timesteps
502 betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta))
503 return np.array(betas)
506def normal_kl(mean1, logvar1, mean2, logvar2):
507 r"""Compute the KL divergence between two gaussians.
509 Shapes are automatically broadcasted, so batches can be compared to
510 scalars, among other use cases.
511 """
512 tensor = None
513 for obj in (mean1, logvar1, mean2, logvar2):
514 if isinstance(obj, torch.Tensor):
515 tensor = obj
516 break
517 assert tensor is not None, "at least one argument must be a Tensor"
519 # Force variances to be Tensors. Broadcasting helps convert scalars to
520 # Tensors, but it does not work for torch.exp().
521 logvar1, logvar2 = (x if isinstance(x, torch.Tensor) else torch.tensor(x).to(tensor) for x in (logvar1, logvar2))
523 return 0.5 * (
524 -1.0 + logvar2 - logvar1 + torch.exp(logvar1 - logvar2) + ((mean1 - mean2) ** 2) * torch.exp(-logvar2)
525 )
528def mean_flat(tensor):
529 r"""Take the mean over all non-batch dimensions."""
530 return tensor.mean(dim=list(range(1, len(tensor.shape))))
533def timestep_embedding(timesteps, dim, max_period=10000):
534 r"""Create sinusoidal timestep embeddings.
536 :param timesteps: a 1-D Tensor of N indices, one per batch element.
537 These may be fractional. (N,)
538 :param dim: the dimension of the output.
539 :param max_period: controls the minimum frequency of the embeddings.
540 :return: an [N x dim] Tensor of positional embeddings.
541 """
542 half = dim // 2
543 freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half).to(
544 timesteps.device
545 ) # shape (dim//2,)
546 args = timesteps[:, None].float() * freqs[None] # (N, dim//2)
547 embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) # (N, (dim//2)*2)
548 if dim % 2:
549 # zero pad in the last dimension to ensure shape (N, dim)
550 embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
551 return embedding