Skip to content

Fix batched DiffusionGemma adaptive stopping - #14386

Merged
kashif merged 7 commits into
mainfrom
fix/diffusion-gemma-adaptive-stopping
Aug 18, 2026
Merged

kashif merged 7 commits into
mainfrom
fix/diffusion-gemma-adaptive-stopping

Conversation

@kashif

@kashif kashif commented Aug 4, 2026

Copy link
Copy Markdown
Contributor

DiffusionGemma adaptive stopping currently waits for the whole batch to converge, so rows that finish early keep changing while slower rows continue. It also measures confidence from the raw model logits instead of the scheduler-shaped logits used for sampling.

This freezes each row once it is stable and confident, and uses the scheduler logits for the stopping decision. A regression test covers both cases.

Checks:

  • make quality
  • make fix-copies
  • DiffusionGemma pipeline tests (11 passed)
  • compiled decoder with static cache smoke test

@github-actions github-actions Bot added size/M PR with diff < 200 LOC tests pipelines labels Aug 4, 2026
@kashif
kashif requested a review from dg845 August 4, 2026 17:35
@kashif
kashif marked this pull request as ready for review August 4, 2026 17:36
@HuggingFaceDocBuilderDev

Copy link
Copy Markdown

The docs for this PR live here. All of your documentation changes will be reflected on that endpoint. The docs are available until 30 days after the last update.

Comment thread src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py Outdated
Comment thread tests/pipelines/diffusion_gemma/test_diffusion_gemma.py Outdated
Comment thread tests/pipelines/diffusion_gemma/test_diffusion_gemma.py Outdated
Comment thread src/diffusers/pipelines/diffusion_gemma/pipeline_diffusion_gemma.py Outdated

@dg845 dg845 left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for the PR! Are the changes (committing finished rows immediately, using scheduler-shaped logits instead of raw denoising model logits) intended to match the DiffusionGemma reference behavior? It's not obvious to me that they are bugfixes rather than modeling changes.

@kashif

kashif commented Aug 15, 2026

Copy link
Copy Markdown
Contributor Author

thanks @dg845 fixing!

Comment thread tests/pipelines/diffusion_gemma/test_diffusion_gemma.py Outdated

@dg845 dg845 left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for the changes! Can we also update the DiffusionGemma docs to reflect the new per-batch-example behavior? Specifically, I think these lines need to be updated:

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`:

(For my question about the whether the changes are a bugfix, I believe the answer is that this PR's changes update the diffusers implementation to match the transformers implementation's batched adaptive stopping behavior.)

@github-actions github-actions Bot added the documentation Improvements or additions to documentation label Aug 18, 2026
@kashif
kashif merged commit d6bfaa7 into main Aug 18, 2026
15 of 16 checks passed
@kashif
kashif deleted the fix/diffusion-gemma-adaptive-stopping branch August 18, 2026 09:12
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

documentation Improvements or additions to documentation pipelines size/M PR with diff < 200 LOC tests

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants