minimax-h3 / diffusers /schedulers /scheduling_discrete_ddim.py
multimodalart's picture
multimodalart HF Staff
Sync the split MiniMax-H3 Spaces (part 4)
9a03d09 verified
Raw
History Blame Contribute Delete
15.2 kB
# 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
@dataclass
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
@register_to_config
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)
@staticmethod
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
@staticmethod
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"]