Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 5 additions & 5 deletions docs/source/en/api/pipelines/diffusion_gemma.md
Original file line number Diff line number Diff line change
Expand Up @@ -145,11 +145,11 @@ is not overwritten); this is the setup shown in [Usage](#usage). Drop both the `

## Adaptive stopping

A block usually converges before all `num_inference_steps` are spent, so by default the pipeline leaves a block's
denoising loop early once every example's argmax prediction is stable for `stability_threshold` steps and the mean
per-token entropy falls below `confidence_threshold` (`0.005`, the value used by the released checkpoint). This roughly
halves the number of decoder forwards at matched quality and is the largest single throughput lever. Pass
`confidence_threshold=None` to always run the full `num_inference_steps`:
A block usually converges before all `num_inference_steps` are spent, so by default the pipeline freezes each batch
example once its argmax prediction is stable for `stability_threshold` steps and its mean per-token entropy falls below
`confidence_threshold` (`0.005`, the value used by the released checkpoint). The denoising loop ends once every example
is frozen. This roughly halves the number of decoder forwards at matched quality and is the largest single throughput
lever. Pass `confidence_threshold=None` to always run the full `num_inference_steps`:

```py
output = pipe(prompt="Why is the sky blue?", gen_length=256, confidence_threshold=None) # disable adaptive stopping
Expand Down
42 changes: 25 additions & 17 deletions src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py
Original file line number Diff line number Diff line change
Expand Up @@ -219,12 +219,13 @@ def __call__(
eos_token_id (`int`, *optional*):
EOS token ID for early stopping. Falls back to the processor's tokenizer.
stability_threshold (`int`, defaults to `1`):
Number of consecutive steps the argmax prediction must be unchanged for a block to count as stable.
Only used when `confidence_threshold` is set.
Number of consecutive steps an example's argmax prediction must remain unchanged for that example to
count as stable. Only used when `confidence_threshold` is set.
confidence_threshold (`float`, *optional*, defaults to `0.005`):
Leave a block's denoising loop early once every example is stable (see `stability_threshold`) and the
mean per-token entropy of the prediction is below this value. Speeds up generation at matched quality;
the default matches the released checkpoint. Set to `None` to always run all `num_inference_steps`.
Freeze each example once it is stable (see `stability_threshold`) and the mean per-token entropy of its
scheduler-shaped prediction logits is below this value. The block's denoising loop ends once every
example is frozen. Speeds up generation at matched quality; the default matches the released
checkpoint. Set to `None` to always run all `num_inference_steps`.
generator (`torch.Generator`, *optional*):
RNG for sampling.
output_type (`str`, defaults to `"text"`):
Expand Down Expand Up @@ -347,6 +348,8 @@ def __call__(
0, text_config.vocab_size, (batch_size, canvas_length), device=device, generator=generator
)
self_conditioning_logits = None
finished_denoising = torch.zeros(batch_size, dtype=torch.bool, device=device)
argmax_canvas = canvas
# Adaptive stopping history: the last `stability_threshold` argmax predictions of this block's canvas.
argmax_history = torch.full(
(max(stability_threshold, 1), batch_size, canvas_length), -1, dtype=torch.long, device=device
Expand Down Expand Up @@ -380,7 +383,8 @@ def __call__(
canvas = scheduler_output.prev_sample
# Self-condition on the logits the scheduler sampled from: temperature-shaped for the reference
# EntropyBound sampler, the raw denoiser logits for the others.
self_conditioning_logits = scheduler_output.pred_logits
pred_logits = scheduler_output.pred_logits
self_conditioning_logits = pred_logits

# Predictor-corrector (https://huggingface.co/papers/2605.22765): a scheduler exposing `corrector_steps`
# + `step_correct` refines the canvas with extra Gibbs sweeps on the first `corrected_steps` predictor
Expand Down Expand Up @@ -408,21 +412,25 @@ def __call__(
global_step += 1
progress_bar.update()

# Adaptive stopping: leave this block early once every example's argmax prediction is stable across
# `stability_threshold` steps and confident (mean per-token entropy below `confidence_threshold`).
# Adaptive stopping: freeze each example once its scheduler-shaped prediction is stable across
# `stability_threshold` steps and confident (mean per-token entropy below `confidence_threshold`),
# then leave the block once every example is finished.
if confidence_threshold is not None:
argmax_canvas = logits.argmax(dim=-1)
stable = (argmax_history == argmax_canvas[None]).all(dim=-1).all(dim=0)
next_argmax_canvas = pred_logits.argmax(dim=-1)
next_argmax_canvas = torch.where(finished_denoising[:, None], argmax_canvas, next_argmax_canvas)
stable = (argmax_history == next_argmax_canvas[None]).all(dim=-1).all(dim=0)
argmax_history = torch.roll(argmax_history, shifts=-1, dims=0)
argmax_history[-1] = argmax_canvas
confident = torch.distributions.Categorical(logits=logits.float()).entropy().mean(-1) < (
argmax_history[-1] = next_argmax_canvas
confident = torch.distributions.Categorical(logits=pred_logits.float()).entropy().mean(-1) < (
confidence_threshold
)
if bool((stable & confident).all()):
# Commit the converged prediction. Ancestral schedulers (e.g. DiscreteDDIM) only clean the
# canvas on their final step, so the in-progress canvas may still hold noise tokens; the
# denoiser argmax is the converged answer (and equals the canvas for commit-style schedulers).
canvas = argmax_canvas
finished_denoising = finished_denoising | (stable & confident)
argmax_canvas = next_argmax_canvas
# Commit each converged prediction. Ancestral schedulers (e.g. DiscreteDDIM) only clean the canvas
# on their final step, so the in-progress canvas may still hold noise tokens; the denoiser argmax
# is the converged answer (and equals the canvas for commit-style schedulers).
canvas = torch.where(finished_denoising[:, None], argmax_canvas, canvas)
if bool(finished_denoising.all()):
break

# Append the denoised canvas and extend the context for the next block.
Expand Down
61 changes: 60 additions & 1 deletion tests/pipelines/diffusion_gemma/test_diffusion_gemma.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,9 @@
import unittest
from types import SimpleNamespace

import torch

from diffusers import BlockRefinementScheduler, DiffusionGemmaPipeline
from diffusers import BlockRefinementScheduler, DiffusionGemmaPipeline, EntropyBoundScheduler
from diffusers.utils.testing_utils import require_peft_backend, require_peft_version_greater


Expand Down Expand Up @@ -75,10 +76,23 @@ def _load_pipeline(test):


class DiffusionGemmaPipelineTest(unittest.TestCase):
adaptive_stopping_vocab_size = 8

def setUp(self):
self.pipe, self.canvas_length = _load_pipeline(self)
self.prompt = "Name a color."

def _run_adaptive_stopping(self, prompt):
self.pipe.model.config.get_text_config(decoder=True).vocab_size = self.adaptive_stopping_vocab_size
return self.pipe(
prompt=prompt,
gen_length=self.canvas_length,
num_inference_steps=5,
confidence_threshold=0.005,
eos_early_stop=False,
output_type="seq",
)

def test_generate(self):
out = self.pipe(
prompt=self.prompt,
Expand All @@ -103,6 +117,51 @@ def test_generate(self):
self.assertEqual(sequences.shape, (1, self.canvas_length))
self.assertEqual(len(texts), 1)

def test_adaptive_stopping_freezes_finished_rows(self):
forward_calls = 0

def forward(decoder_input_ids, **kwargs):
nonlocal forward_calls
batch_size, canvas_length = decoder_input_ids.shape
token_ids = ([1, 3], [1, 4], [2, 5], [2, 5], [2, 6])[forward_calls]
tokens = torch.tensor(token_ids, device=decoder_input_ids.device)[:, None].expand_as(decoder_input_ids)
logits = torch.full(
(batch_size, canvas_length, self.adaptive_stopping_vocab_size),
-100.0,
device=decoder_input_ids.device,
)
logits.scatter_(-1, tokens[..., None], 100.0)
forward_calls += 1
return SimpleNamespace(logits=logits)

self.pipe.model.forward = forward
self.pipe.scheduler = BlockRefinementScheduler()
output = self._run_adaptive_stopping(["Short prompt.", "A somewhat longer prompt for the second batch row."])

self.assertEqual(forward_calls, 4)
self.assertTrue((output.sequences[0] == 1).all())
self.assertTrue((output.sequences[1] == 5).all())

def test_adaptive_stopping_uses_scheduler_logits(self):
forward_calls = 0

def forward(decoder_input_ids, **kwargs):
nonlocal forward_calls
forward_calls += 1
batch_size, canvas_length = decoder_input_ids.shape
logits = torch.zeros(
batch_size, canvas_length, self.adaptive_stopping_vocab_size, device=decoder_input_ids.device
)
logits[..., 0] = 2.0
return SimpleNamespace(logits=logits)

self.pipe.model.forward = forward
self.pipe.scheduler = EntropyBoundScheduler(t_max=0.1, t_min=0.1)
output = self._run_adaptive_stopping(self.prompt)

self.assertEqual(forward_calls, 2)
self.assertTrue((output.sequences == 0).all())

def test_callback_receives_advertised_keys(self):
observed: list[str] = []

Expand Down
Loading