diff --git a/docs/source/en/api/pipelines/diffusion_gemma.md b/docs/source/en/api/pipelines/diffusion_gemma.md index bb3adfe7b514..2674fbb064df 100644 --- a/docs/source/en/api/pipelines/diffusion_gemma.md +++ b/docs/source/en/api/pipelines/diffusion_gemma.md @@ -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 diff --git a/src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py b/src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py index 5222ead8813b..5d608d7c49fb 100644 --- a/src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py +++ b/src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py @@ -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"`): @@ -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 @@ -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 @@ -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. diff --git a/tests/pipelines/diffusion_gemma/test_diffusion_gemma.py b/tests/pipelines/diffusion_gemma/test_diffusion_gemma.py index c01b7adbc81f..bfffc3923016 100644 --- a/tests/pipelines/diffusion_gemma/test_diffusion_gemma.py +++ b/tests/pipelines/diffusion_gemma/test_diffusion_gemma.py @@ -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 @@ -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, @@ -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] = []