From f2387dbe135ffd9deed0ed37ea67de0d2c63db25 Mon Sep 17 00:00:00 2001 From: Akshan Krithick Date: Tue, 6 Oct 2026 22:43:47 -0700 Subject: [PATCH 1/2] fix component library resolution when a transformers model folder shadows a pipeline dir --- .../pipelines/pipeline_loading_utils.py | 6 ++++- tests/pipelines/test_pipeline_utils.py | 26 +++++++++++++++++++ 2 files changed, 31 insertions(+), 1 deletion(-) diff --git a/src/diffusers/pipelines/pipeline_loading_utils.py b/src/diffusers/pipelines/pipeline_loading_utils.py index 6958f49c8ddd..291fe599361b 100644 --- a/src/diffusers/pipelines/pipeline_loading_utils.py +++ b/src/diffusers/pipelines/pipeline_loading_utils.py @@ -923,7 +923,11 @@ def _fetch_class_library_tuple(module): pipeline_dir = module_path_items[-2] if len(module_path_items) > 2 else None path = not_compiled_module.__module__.split(".") - is_pipeline_module = pipeline_dir in path and hasattr(pipelines, pipeline_dir) + # A same-named folder in another library (e.g. `transformers.models.diffusion_gemma` vs + # `diffusers.pipelines.diffusion_gemma`) must not count as a pipeline module. + is_pipeline_module = ( + path[0] == diffusers_module.__name__ and pipeline_dir in path and hasattr(pipelines, pipeline_dir) + ) # if library is not in LOADABLE_CLASSES, then it is a custom module. # Or if it's a pipeline module, then the module is inside the pipeline diff --git a/tests/pipelines/test_pipeline_utils.py b/tests/pipelines/test_pipeline_utils.py index c62e582597fc..0d140c13adf8 100644 --- a/tests/pipelines/test_pipeline_utils.py +++ b/tests/pipelines/test_pipeline_utils.py @@ -1100,3 +1100,29 @@ def test_push_to_hub_library_name(self): # Reset repo delete_repo(repo_id, token=TOKEN) + + +class TestFetchClassLibraryTuple: + def test_diffusers_model(self): + from diffusers import UNet2DConditionModel + from diffusers.pipelines.pipeline_loading_utils import _fetch_class_library_tuple + + assert _fetch_class_library_tuple(UNet2DConditionModel) == ("diffusers", "UNet2DConditionModel") + + def test_pipeline_module_class(self): + from diffusers.pipelines.deepfloyd_if import IFWatermarker + from diffusers.pipelines.pipeline_loading_utils import _fetch_class_library_tuple + + assert _fetch_class_library_tuple(IFWatermarker) == ("deepfloyd_if", "IFWatermarker") + + def test_other_library_class_shadowing_pipeline_dir(self): + from diffusers.pipelines.pipeline_loading_utils import _fetch_class_library_tuple + + # A transformers class whose model folder shares its name with a diffusers pipeline folder + # (e.g. `transformers.models.diffusion_gemma` vs `diffusers.pipelines.diffusion_gemma`) must + # resolve to its own library, not to the pipeline folder. + class FakeModel: + pass + + FakeModel.__module__ = "transformers.models.diffusion_gemma.modeling_diffusion_gemma" + assert _fetch_class_library_tuple(FakeModel) == ("transformers", "FakeModel") From 56ab9448d9d1fa73d650758a82d5a84f33c54daa Mon Sep 17 00:00:00 2001 From: Akshan Krithick Date: Tue, 6 Oct 2026 22:43:47 -0700 Subject: [PATCH 2/2] add modular blockset for diffusion gemma --- .../en/api/pipelines/diffusion_gemma.md | 33 ++ src/diffusers/__init__.py | 4 + src/diffusers/modular_pipelines/__init__.py | 5 + .../diffusion_gemma/__init__.py | 48 +++ .../diffusion_gemma/before_denoise.py | 218 ++++++++++ .../diffusion_gemma/decoders.py | 92 ++++ .../diffusion_gemma/denoise.py | 403 ++++++++++++++++++ .../diffusion_gemma/encoders.py | 132 ++++++ .../modular_blocks_diffusion_gemma.py | 127 ++++++ .../diffusion_gemma/modular_pipeline.py | 29 ++ .../modular_pipelines/modular_pipeline.py | 1 + .../dummy_torch_and_transformers_objects.py | 30 ++ .../diffusion_gemma/__init__.py | 0 .../test_modular_pipeline_diffusion_gemma.py | 168 ++++++++ 14 files changed, 1290 insertions(+) create mode 100644 src/diffusers/modular_pipelines/diffusion_gemma/__init__.py create mode 100644 src/diffusers/modular_pipelines/diffusion_gemma/before_denoise.py create mode 100644 src/diffusers/modular_pipelines/diffusion_gemma/decoders.py create mode 100644 src/diffusers/modular_pipelines/diffusion_gemma/denoise.py create mode 100644 src/diffusers/modular_pipelines/diffusion_gemma/encoders.py create mode 100644 src/diffusers/modular_pipelines/diffusion_gemma/modular_blocks_diffusion_gemma.py create mode 100644 src/diffusers/modular_pipelines/diffusion_gemma/modular_pipeline.py create mode 100644 tests/modular_pipelines/diffusion_gemma/__init__.py create mode 100644 tests/modular_pipelines/diffusion_gemma/test_modular_pipeline_diffusion_gemma.py diff --git a/docs/source/en/api/pipelines/diffusion_gemma.md b/docs/source/en/api/pipelines/diffusion_gemma.md index 2674fbb064df..b307abef67e3 100644 --- a/docs/source/en/api/pipelines/diffusion_gemma.md +++ b/docs/source/en/api/pipelines/diffusion_gemma.md @@ -175,6 +175,31 @@ out = pipe( ) ``` +## Modular + +DiffusionGemma is also available as a [modular pipeline](../../modular_diffusers/overview): `DiffusionGemmaBlocks` +splits the run into a text encoder, generation setup, scheduler setup, the canvas loop and a decode step. The +checkpoint keeps its weights at the repository root, so build the components as for the standard pipeline and hand +them to the blocks: + +```py +import torch +from transformers import AutoProcessor, DiffusionGemmaForBlockDiffusion +from diffusers import BlockRefinementScheduler +from diffusers.modular_pipelines import DiffusionGemmaBlocks + +model_id = "google/diffusiongemma-26B-A4B-it" +pipe = DiffusionGemmaBlocks().init_pipeline() +pipe.update_components( + model=DiffusionGemmaForBlockDiffusion.from_pretrained(model_id, dtype=torch.bfloat16, device_map="auto"), + processor=AutoProcessor.from_pretrained(model_id), + scheduler=BlockRefinementScheduler.from_pretrained(model_id, subfolder="scheduler"), +) + +texts = pipe(prompt="Why is the sky blue?", gen_length=256, num_inference_steps=48, output="texts") +print(texts[0]) +``` + ## DiffusionGemmaPipeline [[autodoc]] DiffusionGemmaPipeline - all @@ -182,3 +207,11 @@ out = pipe( ## DiffusionGemmaPipelineOutput [[autodoc]] pipelines.DiffusionGemmaPipelineOutput + +## DiffusionGemmaModularPipeline + +[[autodoc]] DiffusionGemmaModularPipeline + +## DiffusionGemmaBlocks + +[[autodoc]] DiffusionGemmaBlocks diff --git a/src/diffusers/__init__.py b/src/diffusers/__init__.py index 599237a6f2a0..d186c0692f54 100644 --- a/src/diffusers/__init__.py +++ b/src/diffusers/__init__.py @@ -528,6 +528,8 @@ "Cosmos3DistilledModularPipeline", "Cosmos3OmniBlocks", "Cosmos3OmniModularPipeline", + "DiffusionGemmaBlocks", + "DiffusionGemmaModularPipeline", "EchoBlocks", "EchoModularPipeline", "ErnieImageAutoBlocks", @@ -1400,6 +1402,8 @@ Cosmos3DistilledModularPipeline, Cosmos3OmniBlocks, Cosmos3OmniModularPipeline, + DiffusionGemmaBlocks, + DiffusionGemmaModularPipeline, EchoBlocks, EchoModularPipeline, ErnieImageAutoBlocks, diff --git a/src/diffusers/modular_pipelines/__init__.py b/src/diffusers/modular_pipelines/__init__.py index aaa37918956a..80a8c91997dc 100644 --- a/src/diffusers/modular_pipelines/__init__.py +++ b/src/diffusers/modular_pipelines/__init__.py @@ -111,6 +111,10 @@ "Cosmos3OmniBlocks", "Cosmos3OmniModularPipeline", ] + _import_structure["diffusion_gemma"] = [ + "DiffusionGemmaBlocks", + "DiffusionGemmaModularPipeline", + ] _import_structure["ernie_image"] = [ "ErnieImageAutoBlocks", "ErnieImageModularPipeline", @@ -162,6 +166,7 @@ Cosmos3OmniBlocks, Cosmos3OmniModularPipeline, ) + from .diffusion_gemma import DiffusionGemmaBlocks, DiffusionGemmaModularPipeline from .echo import EchoBlocks, EchoModularPipeline from .ernie_image import ErnieImageAutoBlocks, ErnieImageModularPipeline from .flux import FluxAutoBlocks, FluxKontextAutoBlocks, FluxKontextModularPipeline, FluxModularPipeline diff --git a/src/diffusers/modular_pipelines/diffusion_gemma/__init__.py b/src/diffusers/modular_pipelines/diffusion_gemma/__init__.py new file mode 100644 index 000000000000..9eee87b9bbbd --- /dev/null +++ b/src/diffusers/modular_pipelines/diffusion_gemma/__init__.py @@ -0,0 +1,48 @@ +from typing import TYPE_CHECKING + +from ...utils import ( + DIFFUSERS_SLOW_IMPORT, + OptionalDependencyNotAvailable, + _LazyModule, + get_objects_from_module, + is_torch_available, + is_transformers_available, +) + + +_dummy_objects = {} +_import_structure = {} + +try: + if not (is_transformers_available() and is_torch_available()): + raise OptionalDependencyNotAvailable() +except OptionalDependencyNotAvailable: + from ...utils import dummy_torch_and_transformers_objects # noqa F403 + + _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) +else: + _import_structure["modular_blocks_diffusion_gemma"] = ["DiffusionGemmaBlocks"] + _import_structure["modular_pipeline"] = ["DiffusionGemmaModularPipeline"] + +if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: + try: + if not (is_transformers_available() and is_torch_available()): + raise OptionalDependencyNotAvailable() + except OptionalDependencyNotAvailable: + from ...utils.dummy_torch_and_transformers_objects import * # noqa F403 + else: + from .modular_blocks_diffusion_gemma import DiffusionGemmaBlocks + from .modular_pipeline import DiffusionGemmaModularPipeline + +else: + import sys + + sys.modules[__name__] = _LazyModule( + __name__, + globals()["__file__"], + _import_structure, + module_spec=__spec__, + ) + + for name, value in _dummy_objects.items(): + setattr(sys.modules[__name__], name, value) diff --git a/src/diffusers/modular_pipelines/diffusion_gemma/before_denoise.py b/src/diffusers/modular_pipelines/diffusion_gemma/before_denoise.py new file mode 100644 index 000000000000..c0daee9b5bed --- /dev/null +++ b/src/diffusers/modular_pipelines/diffusion_gemma/before_denoise.py @@ -0,0 +1,218 @@ +# 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. + +import inspect + +import torch +from transformers import DynamicCache, ProcessorMixin, StaticCache + +from ...schedulers import BlockRefinementScheduler +from ...utils import logging +from ...utils.import_utils import is_transformers_version +from ..modular_pipeline import ModularPipelineBlocks, PipelineState +from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import DiffusionGemmaModularPipeline + + +if is_transformers_version("<", "5.11.0"): + raise ImportError( + "`DiffusionGemmaModularPipeline` requires `transformers>=5.11.0` for `DiffusionGemmaForBlockDiffusion`." + ) + +from transformers import DiffusionGemmaForBlockDiffusion # noqa: E402 + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +class DiffusionGemmaPrepareGenerationStep(ModularPipelineBlocks): + model_name = "diffusion-gemma" + + @property + def description(self) -> str: + return ( + "Prepare step that sizes the generation into canvases, creates the encoder KV cache, and resolves " + "the EOS token used for early stopping and trimming" + ) + + @property + def expected_components(self) -> list[ComponentSpec]: + return [ + ComponentSpec("model", DiffusionGemmaForBlockDiffusion), + ComponentSpec("processor", ProcessorMixin), + ] + + @property + def inputs(self) -> list[InputParam]: + return [ + InputParam( + "prompt_ids", + required=True, + type_hint=torch.LongTensor, + description="Tokenized prompt of shape `(batch_size, prompt_length)`.", + ), + InputParam( + "gen_length", + type_hint=int, + default=256, + description="Number of tokens to generate, rounded up to a multiple of the model's `canvas_length`.", + ), + InputParam( + "cache_implementation", + type_hint=str, + description='Set to `"static"` to use a fixed-shape `StaticCache` so the decoder can be compiled.', + ), + InputParam( + "eos_token_id", + type_hint=int, + description="EOS token ID for early stopping. Falls back to the processor's tokenizer.", + ), + ] + + @property + def intermediate_outputs(self) -> list[OutputParam]: + return [ + OutputParam( + "canvas_length", + type_hint=int, + description="The model's canvas length, i.e. the number of tokens denoised per block.", + ), + OutputParam("num_canvases", type_hint=int, description="Number of canvases to generate."), + OutputParam( + "past_key_values", + type_hint=object, + description="The encoder KV cache reused across canvases and denoising steps.", + ), + OutputParam( + "eos_token_id", + type_hint=int, + description="The resolved EOS token ID (user-provided or from the processor's tokenizer).", + ), + OutputParam( + "finished", + type_hint=torch.Tensor, + description="Per-example flags marking sequences that already emitted EOS.", + ), + ] + + @torch.no_grad() + def __call__(self, components: DiffusionGemmaModularPipeline, state: PipelineState) -> PipelineState: + block_state = self.get_block_state(state) + + if block_state.gen_length <= 0: + raise ValueError(f"`gen_length` must be > 0, got {block_state.gen_length}.") + + canvas_length = components.model.config.canvas_length + num_canvases = (block_state.gen_length + canvas_length - 1) // canvas_length + + batch_size, prompt_length = block_state.prompt_ids.shape + text_config = components.model.config.get_text_config(decoder=True) + max_cache_len = prompt_length + num_canvases * canvas_length + if block_state.cache_implementation == "static": + past_key_values = StaticCache(config=text_config, max_cache_len=max_cache_len) + else: + past_key_values = DynamicCache(config=text_config) + + eos_token_id = block_state.eos_token_id + if eos_token_id is None: + tokenizer = getattr(components.processor, "tokenizer", components.processor) + eos_token_id = getattr(tokenizer, "eos_token_id", None) + + block_state.canvas_length = canvas_length + block_state.num_canvases = num_canvases + block_state.past_key_values = past_key_values + block_state.eos_token_id = eos_token_id + block_state.finished = torch.zeros(batch_size, dtype=torch.bool, device=block_state.prompt_ids.device) + + self.set_block_state(state, block_state) + return components, state + + +class DiffusionGemmaSetTimestepsStep(ModularPipelineBlocks): + model_name = "diffusion-gemma" + + @property + def description(self) -> str: + return ( + "Step that splits the per-canvas forward budget into predictor and corrector steps and configures " + "the scheduler's refinement schedule" + ) + + @property + def expected_components(self) -> list[ComponentSpec]: + return [ + ComponentSpec("scheduler", BlockRefinementScheduler), + ] + + @property + def inputs(self) -> list[InputParam]: + return [ + InputParam( + "num_inference_steps", + type_hint=int, + default=48, + description="Number of denoising steps per canvas, i.e. the per-canvas budget of model forwards.", + ), + InputParam("canvas_length", required=True, type_hint=int), + ] + + @property + def intermediate_outputs(self) -> list[OutputParam]: + return [ + OutputParam("predictor_steps", type_hint=int, description="Predictor steps run per canvas."), + OutputParam( + "corrected_steps", + type_hint=int, + description="Number of leading predictor steps that also run corrector sweeps.", + ), + OutputParam( + "corrector_steps", + type_hint=int, + description="Corrector sweeps run after each of the first `corrected_steps` predictor steps.", + ), + ] + + @torch.no_grad() + def __call__(self, components: DiffusionGemmaModularPipeline, state: PipelineState) -> PipelineState: + block_state = self.get_block_state(state) + + num_inference_steps = block_state.num_inference_steps + if num_inference_steps <= 0: + raise ValueError(f"`num_inference_steps` must be > 0, got {num_inference_steps}.") + + # `num_inference_steps` is the per-block budget of model forwards. With a corrector, fold its sweeps into + # that budget (as in https://huggingface.co/papers/2605.22765) instead of adding them on top: the first + # `corrected_steps` predictor steps each run `corrector_steps` extra forwards, so the total stays + # `num_inference_steps` and the predictor-corrector costs the same as plain ancestral sampling. + corrector_steps = getattr(components.scheduler.config, "corrector_steps", 0) + if corrector_steps > 0: + corrected_steps = (num_inference_steps - 1) // (1 + corrector_steps) + predictor_steps = num_inference_steps - corrected_steps * corrector_steps + else: + corrected_steps = 0 + predictor_steps = num_inference_steps + + # Only `BlockRefinementScheduler` takes a per-call `block_length`; the DiscreteDDIM/EntropyBound schedulers + # do not, so we pass scheduler-specific kwargs by signature. + set_timesteps_kwargs = {"device": None} + if "block_length" in inspect.signature(components.scheduler.set_timesteps).parameters: + set_timesteps_kwargs["block_length"] = block_state.canvas_length + components.scheduler.set_timesteps(predictor_steps, **set_timesteps_kwargs) + + block_state.predictor_steps = predictor_steps + block_state.corrected_steps = corrected_steps + block_state.corrector_steps = corrector_steps + + self.set_block_state(state, block_state) + return components, state diff --git a/src/diffusers/modular_pipelines/diffusion_gemma/decoders.py b/src/diffusers/modular_pipelines/diffusion_gemma/decoders.py new file mode 100644 index 000000000000..306e8c8802a1 --- /dev/null +++ b/src/diffusers/modular_pipelines/diffusion_gemma/decoders.py @@ -0,0 +1,92 @@ +# 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. + +import torch +from transformers import ProcessorMixin + +from ...utils import logging +from ..modular_pipeline import ModularPipelineBlocks, PipelineState +from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import DiffusionGemmaModularPipeline + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +class DiffusionGemmaDecodeStep(ModularPipelineBlocks): + model_name = "diffusion-gemma" + + @property + def description(self) -> str: + return ( + "Decode step that trims each generated sequence at its first EOS token and decodes the token IDs " + "into text with the processor" + ) + + @property + def expected_components(self) -> list[ComponentSpec]: + return [ + ComponentSpec("processor", ProcessorMixin), + ] + + @property + def inputs(self) -> list[InputParam]: + return [ + InputParam( + "sequences", + required=True, + type_hint=torch.LongTensor, + description="The generated token IDs of shape `(batch_size, generated_length)`.", + ), + InputParam( + "eos_token_id", + type_hint=int, + description="EOS token ID used to trim each sequence. Falls back to the processor's tokenizer.", + ), + ] + + @property + def intermediate_outputs(self) -> list[OutputParam]: + return [ + OutputParam( + "texts", + type_hint=list, + description="The decoded generated text, one string per prompt.", + ), + ] + + @torch.no_grad() + def __call__(self, components: DiffusionGemmaModularPipeline, state: PipelineState) -> PipelineState: + block_state = self.get_block_state(state) + + eos_token_id = block_state.eos_token_id + if eos_token_id is None: + tokenizer = getattr(components.processor, "tokenizer", components.processor) + eos_token_id = getattr(tokenizer, "eos_token_id", None) + + sequences = block_state.sequences + # Trim each row at its first EOS so post-EOS canvas tokens don't leak into the decoded text. + decode_sequences = sequences + if eos_token_id is not None: + decode_sequences = [ + seq[: int((seq == eos_token_id).nonzero(as_tuple=True)[0][0]) + 1] + if (seq == eos_token_id).any() + else seq + for seq in sequences + ] + + block_state.texts = components.processor.batch_decode(decode_sequences, skip_special_tokens=True) + + self.set_block_state(state, block_state) + return components, state diff --git a/src/diffusers/modular_pipelines/diffusion_gemma/denoise.py b/src/diffusers/modular_pipelines/diffusion_gemma/denoise.py new file mode 100644 index 000000000000..3d0900768134 --- /dev/null +++ b/src/diffusers/modular_pipelines/diffusion_gemma/denoise.py @@ -0,0 +1,403 @@ +# 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. + +import inspect + +import torch +import torch.nn.functional as F + +from ...schedulers import BlockRefinementScheduler +from ...utils import logging +from ...utils.import_utils import is_transformers_version +from ..modular_pipeline import LoopSequentialPipelineBlocks, ModularPipelineBlocks, PipelineState +from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import DiffusionGemmaModularPipeline + + +if is_transformers_version("<", "5.11.0"): + raise ImportError( + "`DiffusionGemmaModularPipeline` requires `transformers>=5.11.0` for `DiffusionGemmaForBlockDiffusion`." + ) + +from transformers import DiffusionGemmaForBlockDiffusion # noqa: E402 + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +class DiffusionGemmaCanvasPrefillStep(ModularPipelineBlocks): + model_name = "diffusion-gemma" + + @property + def description(self) -> str: + return ( + "Canvas step that encodes the tokens not yet in the KV cache (the whole prompt on the first canvas, " + "the last committed canvas afterwards) and builds the decoder attention mask over the cache" + ) + + @property + def expected_components(self) -> list[ComponentSpec]: + return [ + ComponentSpec("model", DiffusionGemmaForBlockDiffusion), + ] + + @property + def inputs(self) -> list[InputParam]: + return [ + InputParam("multimodal_inputs", type_hint=dict), + ] + + @property + def intermediate_outputs(self) -> list[OutputParam]: + return [ + OutputParam( + "decoder_position_ids", + type_hint=torch.LongTensor, + description="Position IDs of the canvas tokens, continuing the running sequence.", + ), + OutputParam( + "decoder_attention_mask_mapping", + type_hint=object, + description="The decoder attention mask mapping built over the populated cache plus the canvas.", + ), + ] + + @torch.no_grad() + def __call__(self, components: DiffusionGemmaModularPipeline, block_state, k: int): + device = block_state.cur_input_ids.device + canvas_length = block_state.canvas_length + batch_size = block_state.cur_input_ids.shape[0] + cur_len = block_state.cur_input_ids.shape[1] + + block_state.decoder_position_ids = torch.arange(cur_len, cur_len + canvas_length, device=device).unsqueeze(0) + + # Encode the tokens not yet in the cache so the decoder reuses the encoder KV cache instead of + # re-encoding the full sequence. + cached_len = block_state.past_key_values.get_seq_length() + torch.compiler.cudagraph_mark_step_begin() + components.model.model.encoder( + input_ids=block_state.cur_input_ids[:, cached_len:], + attention_mask=block_state.cur_attention_mask, + past_key_values=block_state.past_key_values, + position_ids=torch.arange(cached_len, cur_len, device=device).unsqueeze(0), + # Image tensors are consumed by the prompt prefill only; later blocks encode text-only canvases. + **(block_state.multimodal_inputs if cached_len == 0 and block_state.multimodal_inputs else {}), + ) + + # Decoder attends bidirectionally over the populated cache (the live padding mask) plus the always-visible + # canvas; the mask builder sizes this to the cache internally, including the static buffer for a StaticCache. + decoder_attention_mask = F.pad(block_state.cur_attention_mask.bool(), (0, canvas_length), value=True) + block_state.decoder_attention_mask_mapping = ( + components.model.model.decoder.create_diffusion_decoder_attention_mask( + config=components.model.config, + inputs_embeds=torch.empty((batch_size, canvas_length, 0), device=device), + past_key_values=block_state.past_key_values, + decoder_attention_mask=decoder_attention_mask, + ) + ) + + return components, block_state + + +class DiffusionGemmaCanvasNoiseStep(ModularPipelineBlocks): + model_name = "diffusion-gemma" + + @property + def description(self) -> str: + return "Canvas step that initializes the canvas with uniformly random tokens (the uniform corruption prior)" + + @property + def expected_components(self) -> list[ComponentSpec]: + return [ + ComponentSpec("model", DiffusionGemmaForBlockDiffusion), + ] + + @property + def inputs(self) -> list[InputParam]: + return [ + InputParam.template("generator"), + ] + + @property + def intermediate_outputs(self) -> list[OutputParam]: + return [ + OutputParam( + "canvas", + type_hint=torch.LongTensor, + description="The noisy canvas of shape `(batch_size, canvas_length)` being denoised.", + ), + ] + + @torch.no_grad() + def __call__(self, components: DiffusionGemmaModularPipeline, block_state, k: int): + device = block_state.cur_input_ids.device + batch_size = block_state.cur_input_ids.shape[0] + vocab_size = components.model.config.get_text_config(decoder=True).vocab_size + + # Start from a fully random canvas; the scheduler resets its committed state at step 0. `torch.randint` + # requires the generator and the output device to match, so (as with `randn_tensor`) a CPU generator + # samples on CPU and the result is moved to `device` afterwards. + generator = block_state.generator + rand_device = generator.device if generator is not None else device + block_state.canvas = torch.randint( + 0, vocab_size, (batch_size, block_state.canvas_length), device=rand_device, generator=generator + ).to(device) + return components, block_state + + +class DiffusionGemmaCanvasDenoiseStep(ModularPipelineBlocks): + model_name = "diffusion-gemma" + + @property + def description(self) -> str: + return ( + "Canvas step that runs the inner refinement loop: each step samples candidate tokens from the " + "denoiser logits, commits the most confident ones via the scheduler, renoises the rest, and " + "self-conditions the next step on the previous logits. The first `corrected_steps` predictor steps " + "also run corrector sweeps, and adaptive stopping leaves the loop early once the prediction is " + "stable and confident" + ) + + @property + def expected_components(self) -> list[ComponentSpec]: + return [ + ComponentSpec("model", DiffusionGemmaForBlockDiffusion), + ComponentSpec("scheduler", BlockRefinementScheduler), + ] + + @property + def inputs(self) -> list[InputParam]: + return [ + InputParam( + "temperature", + type_hint=float, + default=0.0, + description="Sampling temperature (`0.0` is greedy). Other sampling knobs are scheduler config.", + ), + InputParam( + "stability_threshold", + type_hint=int, + default=1, + description="Consecutive steps the argmax prediction must be unchanged for a canvas to count as " + "stable. Only used when `confidence_threshold` is set.", + ), + InputParam( + "confidence_threshold", + type_hint=float, + default=0.005, + description="Freeze each example once its prediction is stable and the mean per-token entropy is " + "below this value, and leave the refinement loop once every example is frozen. Set to `None` to " + "always run all steps.", + ), + InputParam.template("generator"), + ] + + @torch.no_grad() + def __call__(self, components: DiffusionGemmaModularPipeline, block_state, k: int): + device = block_state.cur_input_ids.device + batch_size, canvas_length = block_state.canvas.shape + + step_param_names = set(inspect.signature(components.scheduler.step).parameters) + self_conditioning_logits = None + finished_denoising = torch.zeros(batch_size, dtype=torch.bool, device=device) + argmax_canvas = block_state.canvas + # Adaptive stopping history: the last `stability_threshold` argmax predictions of this canvas. + argmax_history = torch.full( + (max(block_state.stability_threshold, 1), batch_size, canvas_length), + -1, + dtype=torch.long, + device=device, + ) + + for step_idx in range(block_state.predictor_steps): + # Mark a fresh step and clone the logits so a cudagraph-compiled decoder does not overwrite the + # tensors that self-conditioning and the scheduler read next. Both are no-ops otherwise. + torch.compiler.cudagraph_mark_step_begin() + logits = components.model( + decoder_input_ids=block_state.canvas, + past_key_values=block_state.past_key_values, + self_conditioning_logits=self_conditioning_logits, + decoder_attention_mask=block_state.decoder_attention_mask_mapping, + decoder_position_ids=block_state.decoder_position_ids, + ).logits.clone() + + # Pass only the kwargs the chosen scheduler accepts, so any of the schedulers can drive the loop. + step_kwargs = { + "mask_token_id": None, + "temperature": block_state.temperature, + "generator": block_state.generator, + } + step_kwargs = {name: value for name, value in step_kwargs.items() if name in step_param_names} + scheduler_output = components.scheduler.step( + model_output=logits, timestep=step_idx, sample=block_state.canvas, return_dict=True, **step_kwargs + ) + block_state.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. + pred_logits = scheduler_output.pred_logits + self_conditioning_logits = pred_logits + + # Predictor-corrector (https://huggingface.co/papers/2605.22765): refine the canvas with extra Gibbs + # sweeps on the first `corrected_steps` predictor steps. Each sweep needs fresh logits. + if step_idx < block_state.corrected_steps: + for _ in range(block_state.corrector_steps): + torch.compiler.cudagraph_mark_step_begin() + corrector_logits = components.model( + decoder_input_ids=block_state.canvas, + past_key_values=block_state.past_key_values, + self_conditioning_logits=self_conditioning_logits, + decoder_attention_mask=block_state.decoder_attention_mask_mapping, + decoder_position_ids=block_state.decoder_position_ids, + ).logits.clone() + block_state.canvas = components.scheduler.step_correct( + model_output=corrector_logits, + timestep=step_idx, + sample=block_state.canvas, + generator=block_state.generator, + ).prev_sample + + # 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 canvas once every example is finished. + if block_state.confidence_threshold is not None: + 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] = next_argmax_canvas + confident = torch.distributions.Categorical(logits=pred_logits.float()).entropy().mean(-1) < ( + block_state.confidence_threshold + ) + finished_denoising = finished_denoising | (stable & confident) + argmax_canvas = next_argmax_canvas + # Commit each converged prediction. Ancestral schedulers 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). + block_state.canvas = torch.where(finished_denoising[:, None], argmax_canvas, block_state.canvas) + if bool(finished_denoising.all()): + break + + return components, block_state + + +class DiffusionGemmaCanvasUpdateStep(ModularPipelineBlocks): + model_name = "diffusion-gemma" + + @property + def description(self) -> str: + return "Canvas step that appends the denoised canvas to the running context and tracks EOS emission" + + @property + def inputs(self) -> list[InputParam]: + return [ + InputParam("eos_token_id", type_hint=int), + ] + + @torch.no_grad() + def __call__(self, components: DiffusionGemmaModularPipeline, block_state, k: int): + block_state.cur_input_ids = torch.cat([block_state.cur_input_ids, block_state.canvas], dim=-1) + block_state.cur_attention_mask = F.pad( + block_state.cur_attention_mask, (0, block_state.canvas.shape[1]), value=1 + ) + + if block_state.eos_token_id is not None: + block_state.finished = block_state.finished | (block_state.canvas == block_state.eos_token_id).any(dim=-1) + + return components, block_state + + +class DiffusionGemmaCanvasLoopWrapper(LoopSequentialPipelineBlocks): + model_name = "diffusion-gemma" + + @property + def description(self) -> str: + return ( + "Loop that generates the text canvas by canvas: each iteration prefills the new context into the KV " + "cache, initializes a random canvas, denoises it with the inner refinement loop, and appends it to " + "the running sequence. Generation stops early once every sequence has emitted EOS" + ) + + @property + def loop_inputs(self) -> list[InputParam]: + return [ + InputParam("prompt_ids", required=True, type_hint=torch.LongTensor), + InputParam("prompt_attention_mask", required=True, type_hint=torch.LongTensor), + InputParam("num_canvases", required=True, type_hint=int), + InputParam("canvas_length", required=True, type_hint=int), + InputParam("predictor_steps", required=True, type_hint=int), + InputParam("corrected_steps", required=True, type_hint=int), + InputParam("corrector_steps", required=True, type_hint=int), + InputParam("past_key_values", required=True, type_hint=object), + InputParam("finished", required=True, type_hint=torch.Tensor), + InputParam( + "eos_early_stop", + type_hint=bool, + default=True, + description="Whether to stop generating further canvases once every sequence has emitted EOS.", + ), + ] + + @property + def loop_intermediate_outputs(self) -> list[OutputParam]: + return [ + OutputParam( + "sequences", + type_hint=torch.LongTensor, + description="The generated token IDs of shape `(batch_size, generated_length)`.", + ), + ] + + @torch.no_grad() + def __call__(self, components: DiffusionGemmaModularPipeline, state: PipelineState) -> PipelineState: + block_state = self.get_block_state(state) + + device = components.model.device + block_state.cur_input_ids = block_state.prompt_ids.to(device=device) + block_state.cur_attention_mask = block_state.prompt_attention_mask.to(device=device) + block_state.finished = block_state.finished.to(device=device) + if getattr(block_state, "multimodal_inputs", None): + block_state.multimodal_inputs = { + name: value.to(device=device) for name, value in block_state.multimodal_inputs.items() + } + prompt_length = block_state.prompt_ids.shape[1] + + for k in range(block_state.num_canvases): + components, block_state = self.loop_step(components, block_state, k=k) + if ( + block_state.eos_early_stop + and block_state.eos_token_id is not None + and bool(block_state.finished.all()) + ): + break + + block_state.sequences = block_state.cur_input_ids[:, prompt_length:] + + self.set_block_state(state, block_state) + return components, state + + +class DiffusionGemmaDenoiseStep(DiffusionGemmaCanvasLoopWrapper): + block_classes = [ + DiffusionGemmaCanvasPrefillStep, + DiffusionGemmaCanvasNoiseStep, + DiffusionGemmaCanvasDenoiseStep, + DiffusionGemmaCanvasUpdateStep, + ] + block_names = ["prefill", "noise", "denoise", "update"] + + @property + def description(self) -> str: + return ( + "Canvas denoise step that iterates over canvases.\nAt each canvas: prefill -> noise -> denoise -> update." + ) diff --git a/src/diffusers/modular_pipelines/diffusion_gemma/encoders.py b/src/diffusers/modular_pipelines/diffusion_gemma/encoders.py new file mode 100644 index 000000000000..7c3350af47db --- /dev/null +++ b/src/diffusers/modular_pipelines/diffusion_gemma/encoders.py @@ -0,0 +1,132 @@ +# 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. + +import torch +from transformers import ProcessorMixin + +from ...image_processor import PipelineImageInput +from ...utils import logging +from ..modular_pipeline import ModularPipelineBlocks, PipelineState +from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from .modular_pipeline import DiffusionGemmaModularPipeline + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +class DiffusionGemmaTextEncoderStep(ModularPipelineBlocks): + model_name = "diffusion-gemma" + + @property + def description(self) -> str: + return ( + "Text encoder step that applies the chat template to a `prompt` or a raw `messages` conversation " + "and tokenizes it into the prompt token IDs consumed by the encoder prefill" + ) + + @property + def expected_components(self) -> list[ComponentSpec]: + return [ + ComponentSpec("processor", ProcessorMixin), + ] + + @property + def inputs(self) -> list[InputParam]: + return [ + InputParam("prompt", type_hint=str, description="Prompt text, wrapped in a chat template and tokenized"), + InputParam( + "messages", + type_hint=list, + description="A raw chat conversation to encode instead of `prompt`, e.g. " + '`[{"role": "user", "content": "Hello"}]` or a multi-turn / multimodal conversation.', + ), + InputParam( + "image", + type_hint=PipelineImageInput, + description="Image(s) to pair with `prompt` for multimodal generation. For richer layouts, put the " + "image content directly in `messages`.", + ), + InputParam( + "add_generation_prompt", + type_hint=bool, + default=True, + description="Whether to add the generation prompt when applying the chat template.", + ), + ] + + @property + def intermediate_outputs(self) -> list[OutputParam]: + return [ + OutputParam( + "prompt_ids", + type_hint=torch.LongTensor, + description="Tokenized prompt of shape `(batch_size, prompt_length)`.", + ), + OutputParam( + "prompt_attention_mask", + type_hint=torch.LongTensor, + description="Attention mask for `prompt_ids`.", + ), + OutputParam( + "multimodal_inputs", + type_hint=dict, + description="Image tensors the processor produced for the encoder prefill.", + ), + ] + + @staticmethod + def check_inputs(block_state): + if block_state.prompt is None and block_state.messages is None: + raise ValueError("Provide either `prompt` or `messages`.") + if block_state.prompt is not None and block_state.messages is not None: + raise ValueError("Provide either `prompt` or `messages`, not both.") + + @torch.no_grad() + def __call__(self, components: DiffusionGemmaModularPipeline, state: PipelineState) -> PipelineState: + block_state = self.get_block_state(state) + self.check_inputs(block_state) + + def build_content(text, img): + if img is None: + return text + return [{"type": "image", "image": img}, {"type": "text", "text": text}] + + messages = block_state.messages + if messages is None: + prompt, image = block_state.prompt, block_state.image + if isinstance(prompt, list): + images = image if isinstance(image, list) else [image] * len(prompt) + messages = [[{"role": "user", "content": build_content(p, im)}] for p, im in zip(prompt, images)] + else: + messages = [{"role": "user", "content": build_content(prompt, image)}] + + encoded = components.processor.apply_chat_template( + messages, + add_generation_prompt=block_state.add_generation_prompt, + tokenize=True, + return_tensors="pt", + return_dict=True, + ) + ids = encoded["input_ids"] + mask = encoded.get("attention_mask") + if mask is None: + mask = torch.ones_like(ids, dtype=torch.long) + multimodal_keys = ("pixel_values", "image_position_ids", "mm_token_type_ids") + + block_state.prompt_ids = ids + block_state.prompt_attention_mask = mask.to(dtype=torch.long) + block_state.multimodal_inputs = {k: encoded[k] for k in multimodal_keys if k in encoded} + + self.set_block_state(state, block_state) + return components, state diff --git a/src/diffusers/modular_pipelines/diffusion_gemma/modular_blocks_diffusion_gemma.py b/src/diffusers/modular_pipelines/diffusion_gemma/modular_blocks_diffusion_gemma.py new file mode 100644 index 000000000000..c1d193897ab9 --- /dev/null +++ b/src/diffusers/modular_pipelines/diffusion_gemma/modular_blocks_diffusion_gemma.py @@ -0,0 +1,127 @@ +# 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 ...utils import logging +from ..modular_pipeline import SequentialPipelineBlocks +from .before_denoise import DiffusionGemmaPrepareGenerationStep, DiffusionGemmaSetTimestepsStep +from .decoders import DiffusionGemmaDecodeStep +from .denoise import DiffusionGemmaDenoiseStep +from .encoders import DiffusionGemmaTextEncoderStep + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +# text_encoder -> prepare_generation -> set_timesteps -> denoise -> decode +# auto_docstring +class DiffusionGemmaBlocks(SequentialPipelineBlocks): + """ + Modular blocks for DiffusionGemma block-diffusion text generation. + - `text_encoder` applies the chat template and tokenizes the prompt + - `prepare_generation` sizes the canvases, creates the KV cache and resolves EOS + - `set_timesteps` splits the forward budget and configures the scheduler + - `denoise` generates the text canvas by canvas + - `decode` trims at EOS and decodes the token IDs into text + + Components: + processor (`ProcessorMixin`) model (`DiffusionGemmaForBlockDiffusion`) scheduler (`BlockRefinementScheduler`) + + Inputs: + prompt (`str`, *optional*): + Prompt text, wrapped in a chat template and tokenized + messages (`list`, *optional*): + A raw chat conversation to encode instead of `prompt`, e.g. `[{"role": "user", "content": "Hello"}]` or a + multi-turn / multimodal conversation. + image (`Image | ndarray | Tensor | list | list | list`, *optional*): + Image(s) to pair with `prompt` for multimodal generation. For richer layouts, put the image content + directly in `messages`. + add_generation_prompt (`bool`, *optional*, defaults to True): + Whether to add the generation prompt when applying the chat template. + gen_length (`int`, *optional*, defaults to 256): + Number of tokens to generate, rounded up to a multiple of the model's `canvas_length`. + cache_implementation (`str`, *optional*): + Set to `"static"` to use a fixed-shape `StaticCache` so the decoder can be compiled. + eos_token_id (`int`, *optional*): + EOS token ID for early stopping. Falls back to the processor's tokenizer. + num_inference_steps (`int`, *optional*, defaults to 48): + Number of denoising steps per canvas, i.e. the per-canvas budget of model forwards. + eos_early_stop (`bool`, *optional*, defaults to True): + Whether to stop generating further canvases once every sequence has emitted EOS. + generator (`Generator`, *optional*): + Torch generator for deterministic generation. + temperature (`float`, *optional*, defaults to 0.0): + Sampling temperature (`0.0` is greedy). Other sampling knobs are scheduler config. + stability_threshold (`int`, *optional*, defaults to 1): + Consecutive steps the argmax prediction must be unchanged for a canvas to count as stable. Only used when + `confidence_threshold` is set. + confidence_threshold (`float`, *optional*, defaults to 0.005): + Freeze each example once its prediction is stable and the mean per-token entropy is below this value, and + leave the refinement loop once every example is frozen. Set to `None` to always run all steps. + + Outputs: + prompt_ids (`LongTensor`): + Tokenized prompt of shape `(batch_size, prompt_length)`. + prompt_attention_mask (`LongTensor`): + Attention mask for `prompt_ids`. + multimodal_inputs (`dict`): + Image tensors the processor produced for the encoder prefill. + canvas_length (`int`): + The model's canvas length, i.e. the number of tokens denoised per block. + num_canvases (`int`): + Number of canvases to generate. + past_key_values (`object`): + The encoder KV cache reused across canvases and denoising steps. + eos_token_id (`int`): + The resolved EOS token ID (user-provided or from the processor's tokenizer). + finished (`Tensor`): + Per-example flags marking sequences that already emitted EOS. + predictor_steps (`int`): + Predictor steps run per canvas. + corrected_steps (`int`): + Number of leading predictor steps that also run corrector sweeps. + corrector_steps (`int`): + Corrector sweeps run after each of the first `corrected_steps` predictor steps. + decoder_position_ids (`LongTensor`): + Position IDs of the canvas tokens, continuing the running sequence. + decoder_attention_mask_mapping (`object`): + The decoder attention mask mapping built over the populated cache plus the canvas. + canvas (`LongTensor`): + The noisy canvas of shape `(batch_size, canvas_length)` being denoised. + sequences (`LongTensor`): + The generated token IDs of shape `(batch_size, generated_length)`. + texts (`list`): + The decoded generated text, one string per prompt. + """ + + model_name = "diffusion-gemma" + + block_classes = [ + DiffusionGemmaTextEncoderStep, + DiffusionGemmaPrepareGenerationStep, + DiffusionGemmaSetTimestepsStep, + DiffusionGemmaDenoiseStep, + DiffusionGemmaDecodeStep, + ] + block_names = ["text_encoder", "prepare_generation", "set_timesteps", "denoise", "decode"] + + @property + def description(self) -> str: + return ( + "Modular blocks for DiffusionGemma block-diffusion text generation.\n" + "- `text_encoder` applies the chat template and tokenizes the prompt\n" + "- `prepare_generation` sizes the canvases, creates the KV cache and resolves EOS\n" + "- `set_timesteps` splits the forward budget and configures the scheduler\n" + "- `denoise` generates the text canvas by canvas\n" + "- `decode` trims at EOS and decodes the token IDs into text" + ) diff --git a/src/diffusers/modular_pipelines/diffusion_gemma/modular_pipeline.py b/src/diffusers/modular_pipelines/diffusion_gemma/modular_pipeline.py new file mode 100644 index 000000000000..8485921d71e0 --- /dev/null +++ b/src/diffusers/modular_pipelines/diffusion_gemma/modular_pipeline.py @@ -0,0 +1,29 @@ +# 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 ...utils import logging +from ..modular_pipeline import ModularPipeline + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +class DiffusionGemmaModularPipeline(ModularPipeline): + """ + A ModularPipeline for DiffusionGemma block-diffusion text generation. + + """ + + default_blocks_name = "DiffusionGemmaBlocks" diff --git a/src/diffusers/modular_pipelines/modular_pipeline.py b/src/diffusers/modular_pipelines/modular_pipeline.py index c501d7b1fc42..591cbf531357 100644 --- a/src/diffusers/modular_pipelines/modular_pipeline.py +++ b/src/diffusers/modular_pipelines/modular_pipeline.py @@ -202,6 +202,7 @@ def _helios_pyramid_map_fn(config_dict=None): ("ltx2.5", _create_default_map_fn("LTX25ModularPipeline")), ("minimax-h3", _create_default_map_fn("MiniMaxH3ModularPipeline")), ("minimax-music3", _create_default_map_fn("MiniMaxMusic3ModularPipeline")), + ("diffusion-gemma", _create_default_map_fn("DiffusionGemmaModularPipeline")), ("ernie-image", _create_default_map_fn("ErnieImageModularPipeline")), ] ) diff --git a/src/diffusers/utils/dummy_torch_and_transformers_objects.py b/src/diffusers/utils/dummy_torch_and_transformers_objects.py index 9ede7a7543f6..fe49ef334287 100644 --- a/src/diffusers/utils/dummy_torch_and_transformers_objects.py +++ b/src/diffusers/utils/dummy_torch_and_transformers_objects.py @@ -92,6 +92,36 @@ def from_pretrained(cls, *args, **kwargs): requires_backends(cls, ["torch", "transformers"]) +class DiffusionGemmaBlocks(metaclass=DummyObject): + _backends = ["torch", "transformers"] + + def __init__(self, *args, **kwargs): + requires_backends(self, ["torch", "transformers"]) + + @classmethod + def from_config(cls, *args, **kwargs): + requires_backends(cls, ["torch", "transformers"]) + + @classmethod + def from_pretrained(cls, *args, **kwargs): + requires_backends(cls, ["torch", "transformers"]) + + +class DiffusionGemmaModularPipeline(metaclass=DummyObject): + _backends = ["torch", "transformers"] + + def __init__(self, *args, **kwargs): + requires_backends(self, ["torch", "transformers"]) + + @classmethod + def from_config(cls, *args, **kwargs): + requires_backends(cls, ["torch", "transformers"]) + + @classmethod + def from_pretrained(cls, *args, **kwargs): + requires_backends(cls, ["torch", "transformers"]) + + class EchoBlocks(metaclass=DummyObject): _backends = ["torch", "transformers"] diff --git a/tests/modular_pipelines/diffusion_gemma/__init__.py b/tests/modular_pipelines/diffusion_gemma/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/tests/modular_pipelines/diffusion_gemma/test_modular_pipeline_diffusion_gemma.py b/tests/modular_pipelines/diffusion_gemma/test_modular_pipeline_diffusion_gemma.py new file mode 100644 index 000000000000..0ce762585f79 --- /dev/null +++ b/tests/modular_pipelines/diffusion_gemma/test_modular_pipeline_diffusion_gemma.py @@ -0,0 +1,168 @@ +# coding=utf-8 +# Copyright 2026 HuggingFace Inc. +# +# 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 types import SimpleNamespace + +import pytest +import torch + +from diffusers import BlockRefinementScheduler, EntropyBoundScheduler +from diffusers.modular_pipelines import DiffusionGemmaBlocks, DiffusionGemmaModularPipeline, ModularPipeline + +from ...testing_utils import torch_device +from ..testing_utils import ( + BaseModularPipelineTesterConfig, + ModularLoadingTesterMixin, + ModularMemoryTesterMixin, + ModularPipelineTesterMixin, + ModularWorkflowTesterMixin, +) + + +class DiffusionGemmaModularPipelineTesterConfig(BaseModularPipelineTesterConfig): + pipeline_class = DiffusionGemmaModularPipeline + pipeline_blocks_class = DiffusionGemmaBlocks + pretrained_model_name_or_path = "akshan-main/tiny-diffusion-gemma-modular-pipe" + + params = frozenset(["prompt", "messages", "gen_length"]) + batch_params = frozenset(["prompt"]) + optional_params = frozenset(["num_inference_steps", "temperature", "eos_token_id"]) + output_name = "sequences" + + def get_dummy_inputs(self, seed=0): + return { + "prompt": "Why is the sky blue?", + "generator": self.get_generator(seed), + "gen_length": 32, + "num_inference_steps": 2, + } + + +class TestDiffusionGemmaModularPipelineFast(DiffusionGemmaModularPipelineTesterConfig, ModularPipelineTesterMixin): + @pytest.mark.skip( + reason="The canvas noise is drawn from a single generator for the whole batch (same as the standard " + "pipeline), so the per-prompt generator lists this test passes are not supported" + ) + def test_inference_batch_consistent(self): + pass + + @pytest.mark.skip( + reason="The canvas noise is drawn from a single generator for the whole batch (same as the standard " + "pipeline), so the per-prompt generator lists this test passes are not supported" + ) + def test_inference_batch_single_identical(self): + pass + + adaptive_stopping_vocab_size = 8 + + def _run_adaptive_stopping(self, pipe, prompt): + pipe.model.config.get_text_config(decoder=True).vocab_size = self.adaptive_stopping_vocab_size + return pipe( + prompt=prompt, + gen_length=32, + num_inference_steps=5, + confidence_threshold=0.005, + eos_early_stop=False, + generator=self.get_generator(), + output="sequences", + ) + + def test_adaptive_stopping_freezes_finished_rows(self): + pipe = self.get_pipeline().to(torch_device) + 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) + + pipe.model.forward = forward + pipe.update_components(scheduler=BlockRefinementScheduler()) + sequences = self._run_adaptive_stopping( + pipe, ["Short prompt.", "A somewhat longer prompt for the second batch row."] + ) + + # The first row is stable from the second step and the second row from the fourth, so the loop runs four + # of the five steps and each row keeps the prediction it was frozen on. + assert forward_calls == 4 + assert bool((sequences[0] == 1).all()) + assert bool((sequences[1] == 5).all()) + + def test_adaptive_stopping_uses_scheduler_logits(self): + pipe = self.get_pipeline().to(torch_device) + 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) + + pipe.model.forward = forward + pipe.update_components(scheduler=EntropyBoundScheduler(t_max=0.1, t_min=0.1)) + sequences = self._run_adaptive_stopping(pipe, "Name a color.") + + assert forward_calls == 2 + assert bool((sequences == 0).all()) + + def test_text_output(self): + pipe = self.get_pipeline().to("cpu") + + inputs = self.get_dummy_inputs() + state = pipe(**inputs) + sequences = state.get("sequences") + texts = state.get("texts") + + assert sequences.dtype == torch.long + assert sequences.shape == (1, 32) + vocab_size = pipe.model.config.get_text_config(decoder=True).vocab_size + assert bool((sequences >= 0).all()) and bool((sequences < vocab_size).all()) + assert isinstance(texts, list) and len(texts) == 1 and isinstance(texts[0], str) + + +class TestDiffusionGemmaModularPipelineLoading(DiffusionGemmaModularPipelineTesterConfig, ModularLoadingTesterMixin): + def test_save_from_pretrained(self, tmp_path, base_pipe_output): + # The base test compares an image slice; text output is compared token for token. + base_pipe = self.get_pipeline().to(torch_device) + base_pipe.save_pretrained(str(tmp_path)) + + pipe = ModularPipeline.from_pretrained(tmp_path) + pipe.load_components(dtype=torch.float32) + pipe.to(torch_device) + + sequences = pipe(**self.get_dummy_inputs(), output=self.output_name) + assert torch.equal(sequences, base_pipe_output) + + +class TestDiffusionGemmaModularPipelineWorkflow(DiffusionGemmaModularPipelineTesterConfig, ModularWorkflowTesterMixin): + pass + + +class TestDiffusionGemmaModularPipelineMemory(DiffusionGemmaModularPipelineTesterConfig, ModularMemoryTesterMixin): + pass