Spaces:
Paused
Paused
| # Copyright 2026 The HuggingFace Team. All rights reserved. | |
| # | |
| # Licensed under the Apache License, Version 2.0 (the "License"); | |
| # you may not use this file except in compliance with the License. | |
| # You may obtain a copy of the License at | |
| # | |
| # http://www.apache.org/licenses/LICENSE-2.0 | |
| # | |
| # Unless required by applicable law or agreed to in writing, software | |
| # distributed under the License is distributed on an "AS IS" BASIS, | |
| # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | |
| # See the License for the specific language governing permissions and | |
| # limitations under the License. | |
| from __future__ import annotations | |
| import math | |
| from dataclasses import dataclass | |
| import torch | |
| from ..configuration_utils import ConfigMixin, register_to_config | |
| from ..utils import BaseOutput | |
| from .scheduling_utils import SchedulerMixin | |
| class DiscreteDDIMSchedulerOutput(BaseOutput): | |
| """ | |
| Output class for the discrete DDIM scheduler. | |
| Args: | |
| prev_sample (`torch.LongTensor` of shape `(batch_size, block_length)`): | |
| Updated block tokens after the current denoising step. | |
| sampled_tokens (`torch.LongTensor` of shape `(batch_size, block_length)`): | |
| Token IDs sampled from the model logits, i.e. the predicted clean tokens `x0`. | |
| sampled_probs (`torch.Tensor` of shape `(batch_size, block_length)`): | |
| Probabilities of the sampled tokens. | |
| pred_logits (`torch.Tensor` of shape `(batch_size, block_length, vocab_size)`): | |
| The denoiser logits, passed through for self-conditioning the next step. | |
| """ | |
| prev_sample: torch.LongTensor | |
| sampled_tokens: torch.LongTensor | |
| sampled_probs: torch.Tensor | |
| pred_logits: torch.Tensor | |
| class DiscreteDDIMScheduler(SchedulerMixin, ConfigMixin): | |
| """ | |
| Discrete DDIM scheduler for the uniform corruption process, following "Structured Denoising Diffusion Models in | |
| Discrete State-Spaces" (D3PM, https://huggingface.co/papers/2107.03006). | |
| On the linear schedule the survival probability of a clean token at time `t` is `alpha(t) = 1 - t`. One denoising | |
| step from time `t` to `s < t` samples every block position from the exact posterior `q(x_s | x_t, x0)`, which for | |
| the uniform kernel decomposes into three routes: jump to the predicted clean token `x0`, stay on the current token, | |
| or jump to a uniformly random token. Unlike masked diffusion, there is no mask token; uncommitted positions carry | |
| random tokens. | |
| An optional predictor-corrector mode follows "Uniform Diffusion Models Revisited: Leave-One-Out Denoiser and | |
| Absorbing State Reformulation" via the leave-one-out (LOO) denoiser (https://huggingface.co/papers/2605.22765). | |
| When `corrector_steps > 0`, the pipeline runs that many Gibbs corrector sweeps after each predictor step (see | |
| [`~DiscreteDDIMScheduler.step_correct`]), resampling the least-confident positions from the one-coordinate | |
| conditional `Cat(alpha_s * x0_loo + (1 - alpha_s) / K)` while holding the rest fixed, which leaves the marginal | |
| `p_s` invariant and improves generation at no training cost. | |
| Args: | |
| num_inference_steps (`int`, defaults to 32): | |
| The number of denoising steps, defining the linear time grid the posterior is evaluated on. | |
| corrector_steps (`int`, defaults to 0): | |
| Number of Gibbs corrector sweeps run after each predictor step. `0` recovers plain ancestral DDIM sampling. | |
| corrector_k (`int`, defaults to 1): | |
| Number of positions resampled per corrector sweep. | |
| corrector_selection (`str`, defaults to `"lowest_log_margin"`): | |
| How the resampled positions are chosen: `"lowest_log_margin"`, `"lowest_maxprob"`, `"lowest_current_prob"`, | |
| or `"random"`. | |
| corrector_selection_tau (`float`, defaults to 1.0): | |
| Temperature of the Gumbel-top-k position selection (lower is greedier). | |
| """ | |
| order = 1 | |
| def __init__( | |
| self, | |
| num_inference_steps: int = 32, | |
| corrector_steps: int = 0, | |
| corrector_k: int = 1, | |
| corrector_selection: str = "lowest_log_margin", | |
| corrector_selection_tau: float = 1.0, | |
| ): | |
| self.num_inference_steps = num_inference_steps | |
| self.timesteps = torch.arange(num_inference_steps, dtype=torch.long) | |
| def set_timesteps(self, num_inference_steps: int, device: str | torch.device | None = None) -> None: | |
| if num_inference_steps <= 0: | |
| raise ValueError(f"`num_inference_steps` must be > 0, got {num_inference_steps}.") | |
| self.num_inference_steps = num_inference_steps | |
| self.timesteps = torch.arange(num_inference_steps, device=device, dtype=torch.long) | |
| def _sample_from_logits( | |
| logits: torch.Tensor, | |
| *, | |
| temperature: float, | |
| generator: torch.Generator | None, | |
| ) -> tuple[torch.LongTensor, torch.Tensor]: | |
| """Sample one token per position with optional temperature, returning tokens and their probabilities.""" | |
| if temperature < 0: | |
| raise ValueError(f"`temperature` must be >= 0, got {temperature}.") | |
| vocab_size = logits.shape[-1] | |
| flat_logits = logits.reshape(-1, vocab_size) | |
| probs = torch.softmax(flat_logits.float(), dim=-1) | |
| if temperature == 0.0: | |
| token = flat_logits.argmax(dim=-1, keepdim=True) | |
| else: | |
| scaled_probs = torch.softmax(flat_logits.float() / temperature, dim=-1) | |
| token = torch.multinomial(scaled_probs, num_samples=1, generator=generator) | |
| token_prob = torch.gather(probs, -1, token) | |
| return token.view(*logits.shape[:-1]), token_prob.view(*logits.shape[:-1]) | |
| def _alpha(self, step_index: int) -> float: | |
| """Survival probability `alpha = 1 - t` of a clean token at the time grid point `step_index`.""" | |
| return step_index / self.num_inference_steps | |
| def _to_loo_logits(logits: torch.Tensor, tokens: torch.LongTensor, alpha: float) -> torch.Tensor: | |
| """ | |
| Convert plain-denoiser logits to the leave-one-out posterior for the uniform kernel. | |
| Subtracts `log(1 + K * alpha / (1 - alpha))` from the observed token's logit (eq. 13 of | |
| https://huggingface.co/papers/2605.22765); renormalization happens in the following softmax. | |
| """ | |
| if alpha <= 0.0 or alpha >= 1.0: | |
| return logits | |
| delta = math.log1p(logits.shape[-1] * alpha / (1.0 - alpha)) | |
| shifted = logits.clone() | |
| src = torch.full((*tokens.shape, 1), -delta, dtype=shifted.dtype, device=shifted.device) | |
| shifted.scatter_add_(-1, tokens.unsqueeze(-1), src) | |
| return shifted | |
| def step( | |
| self, | |
| model_output: torch.Tensor, | |
| timestep: int | torch.Tensor, | |
| sample: torch.LongTensor, | |
| *, | |
| temperature: float = 0.0, | |
| generator: torch.Generator | None = None, | |
| return_dict: bool = True, | |
| ) -> DiscreteDDIMSchedulerOutput | tuple[torch.LongTensor, torch.LongTensor, torch.Tensor]: | |
| """ | |
| Sample the next block from the posterior `q(x_s | x_t, x0)` of the uniform corruption process. | |
| With `a = alpha_t / alpha_s` (survival probability from `s` to `t`) and `b = alpha_s`, the posterior mass of | |
| each route is | |
| clean: `b * (1 - a) / K + a * b * 1[x_t = x0]`, stay: `a * (1 - b) / K`, noise: `(1 - a) * (1 - b) / K`, | |
| so the last step (`b = 1`) deterministically commits the predicted clean tokens. | |
| Args: | |
| model_output (`torch.Tensor` of shape `(batch_size, block_length, vocab_size)`): | |
| Raw logits from the model for the current block. | |
| timestep (`int` or `torch.Tensor`): | |
| Current step index within the denoising schedule, in `[0, num_inference_steps - 1]`. | |
| sample (`torch.LongTensor` of shape `(batch_size, block_length)`): | |
| Current block token IDs `x_t`. | |
| temperature (`float`): | |
| Sampling temperature applied to the logits when drawing `x0`. | |
| generator (`torch.Generator`, *optional*): | |
| RNG for sampling. | |
| return_dict (`bool`): | |
| Whether to return a [`DiscreteDDIMSchedulerOutput`] or a plain tuple. | |
| """ | |
| if isinstance(timestep, torch.Tensor): | |
| step_index = int(timestep.item()) | |
| else: | |
| step_index = int(timestep) | |
| sampled_tokens, sampled_probs = self._sample_from_logits( | |
| model_output, temperature=temperature, generator=generator | |
| ) | |
| vocab_size = model_output.shape[-1] | |
| num_steps = self.num_inference_steps | |
| # `step_index` counts up from 0 to `num_inference_steps - 1`: alpha(t) = 1 - t increases towards the clean end, | |
| # with alpha_s = 1 on the final step so the predicted clean tokens are committed deterministically. | |
| alpha_t = step_index / num_steps | |
| alpha_s = (step_index + 1) / num_steps | |
| survival = alpha_t / alpha_s | |
| same = (sample == sampled_tokens).float() | |
| clean_mass = alpha_s * (1 - survival) / vocab_size + survival * alpha_s * same | |
| stay_mass = survival * (1 - alpha_s) / vocab_size * torch.ones_like(same) | |
| noise_mass = (1 - survival) * (1 - alpha_s) / vocab_size * torch.ones_like(same) | |
| route_probs = torch.stack([clean_mass, stay_mass, noise_mass], dim=-1) | |
| route_probs = route_probs / route_probs.sum(dim=-1, keepdim=True) | |
| routes = torch.multinomial(route_probs.view(-1, 3), num_samples=1, generator=generator).view_as(sample) | |
| random_tokens = torch.randint( | |
| low=0, high=vocab_size, size=sample.shape, device=sample.device, generator=generator | |
| ) | |
| prev_sample = torch.where(routes == 0, sampled_tokens, sample) | |
| prev_sample = torch.where(routes == 2, random_tokens, prev_sample) | |
| if not return_dict: | |
| return prev_sample, sampled_tokens, sampled_probs, model_output | |
| return DiscreteDDIMSchedulerOutput( | |
| prev_sample=prev_sample, | |
| sampled_tokens=sampled_tokens, | |
| sampled_probs=sampled_probs, | |
| pred_logits=model_output, | |
| ) | |
| def _select_positions( | |
| self, sample: torch.LongTensor, cond_log_probs: torch.Tensor, generator: torch.Generator | None | |
| ) -> torch.LongTensor: | |
| """Pick `corrector_k` positions per row to resample, least-confident first (Gumbel-top-k without replacement).""" | |
| selection = self.config.corrector_selection | |
| batch_size, seq_len = sample.shape | |
| k_eff = min(max(1, int(self.config.corrector_k)), seq_len) | |
| if selection == "random": | |
| scores = torch.rand(batch_size, seq_len, device=sample.device, generator=generator) | |
| return torch.topk(scores, k=k_eff, dim=-1).indices | |
| if selection == "lowest_maxprob": | |
| confidence = -cond_log_probs.max(dim=-1).values | |
| elif selection == "lowest_current_prob": | |
| confidence = -torch.gather(cond_log_probs, -1, sample.unsqueeze(-1)).squeeze(-1) | |
| elif selection == "lowest_log_margin": | |
| log_current = torch.gather(cond_log_probs, -1, sample.unsqueeze(-1)).squeeze(-1) | |
| alt = cond_log_probs.clone().scatter_(-1, sample.unsqueeze(-1), float("-inf")) | |
| confidence = -(log_current - alt.max(dim=-1).values) | |
| else: | |
| raise ValueError(f"Unknown `corrector_selection`: {selection!r}.") | |
| keys = confidence / float(self.config.corrector_selection_tau) | |
| u = torch.rand(keys.shape, device=keys.device, generator=generator).clamp_(1e-12, 1.0 - 1e-12) | |
| keys = keys + (-torch.log(-torch.log(u))) | |
| return torch.topk(keys, k=k_eff, dim=-1).indices | |
| def step_correct( | |
| self, | |
| model_output: torch.Tensor, | |
| timestep: int | torch.Tensor, | |
| sample: torch.LongTensor, | |
| *, | |
| generator: torch.Generator | None = None, | |
| return_dict: bool = True, | |
| ) -> DiscreteDDIMSchedulerOutput | tuple[torch.LongTensor, torch.LongTensor, torch.Tensor]: | |
| """ | |
| Run one Gibbs corrector sweep at the post-predictor time `s`, following the leave-one-out predictor-corrector | |
| of https://huggingface.co/papers/2605.22765. | |
| The model logits (recomputed on the current `sample`) are converted to the LOO denoiser, the one-coordinate | |
| conditional `p_s(x^l | x^{-l}) = Cat(alpha_s * x0_loo + (1 - alpha_s) / K)` is formed, the least-confident | |
| `corrector_k` positions are selected, and those positions are resampled while the rest are held fixed. The | |
| sweep preserves `p_s`, so it refines the sample without changing its marginal and needs no extra training. | |
| Args: | |
| model_output (`torch.Tensor` of shape `(batch_size, block_length, vocab_size)`): | |
| Raw logits from the model recomputed on the current (post-predictor) `sample`. | |
| timestep (`int` or `torch.Tensor`): | |
| The predictor step index just completed; the corrector runs at the following grid point `s`. | |
| sample (`torch.LongTensor` of shape `(batch_size, block_length)`): | |
| Current block token IDs to refine. | |
| generator (`torch.Generator`, *optional*): | |
| RNG for sampling. | |
| return_dict (`bool`): | |
| Whether to return a [`DiscreteDDIMSchedulerOutput`] or a plain tuple. | |
| """ | |
| if isinstance(timestep, torch.Tensor): | |
| step_index = int(timestep.item()) | |
| else: | |
| step_index = int(timestep) | |
| # The corrector acts at the cleaner time `s` reached by the predictor. | |
| alpha_s = self._alpha(step_index + 1) | |
| vocab_size = model_output.shape[-1] | |
| # Match the reference corrector, which forms the conditional in float64 (the LOO correction reaches ~log(K)). | |
| loo_logits = self._to_loo_logits(model_output.double(), sample, alpha_s) | |
| loo_log_probs = torch.log_softmax(loo_logits, dim=-1) | |
| log_uniform = math.log1p(-alpha_s) - math.log(vocab_size) | |
| cond_log_probs = torch.logaddexp( | |
| math.log(alpha_s) + loo_log_probs, torch.full_like(loo_log_probs, log_uniform) | |
| ) | |
| positions = self._select_positions(sample, cond_log_probs, generator) | |
| rows = torch.arange(sample.shape[0], device=sample.device).unsqueeze(-1).expand_as(positions) | |
| chosen_probs = cond_log_probs[rows, positions].exp() | |
| resampled = torch.multinomial( | |
| chosen_probs.reshape(-1, vocab_size), num_samples=1, generator=generator | |
| ).view_as(positions) | |
| prev_sample = sample.clone() | |
| prev_sample[rows, positions] = resampled | |
| sampled_probs = torch.gather(chosen_probs, -1, resampled.unsqueeze(-1)).squeeze(-1) | |
| if not return_dict: | |
| return prev_sample, resampled, sampled_probs, model_output | |
| return DiscreteDDIMSchedulerOutput( | |
| prev_sample=prev_sample, | |
| sampled_tokens=resampled, | |
| sampled_probs=sampled_probs, | |
| pred_logits=model_output, | |
| ) | |
| __all__ = ["DiscreteDDIMScheduler", "DiscreteDDIMSchedulerOutput"] | |