From fdb592ce7b267e20816ccbeb7f52a4b2d1fe8619 Mon Sep 17 00:00:00 2001 From: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com> Date: Mon, 5 Oct 2026 01:53:40 +0000 Subject: [PATCH 1/2] Add step-end callbacks and interrupt to modular pipelines ModularPipeline now accepts callback_on_step_end and callback_on_step_end_tensor_inputs, and pipe.interrupt stops denoising. Denoising loops use LoopSequentialPipelineBlocks.loop_over_timesteps and declare their callback fields. set_progress_bar_config now reaches nested loops. Addresses #12386. --- .../loop_sequential_pipeline_blocks.md | 7 + .../en/modular_diffusers/modular_pipeline.md | 55 ++++++ .../modular_pipelines/anima/denoise.py | 5 +- .../modular_pipelines/cosmos/denoise.py | 6 + .../cosmos/modular_blocks_cosmos3.py | 2 + .../modular_pipelines/ernie_image/denoise.py | 5 +- .../modular_pipelines/flux/denoise.py | 5 +- .../modular_pipelines/flux2/denoise.py | 6 +- .../modular_pipelines/helios/denoise.py | 42 ++++ .../hunyuan_video1_5/denoise.py | 5 +- .../modular_pipelines/ideogram4/denoise.py | 5 +- .../modular_pipelines/krea2/denoise.py | 5 +- .../modular_pipelines/ltx/denoise.py | 5 +- .../modular_pipelines/ltx2/denoise.py | 5 +- .../modular_pipelines/minimax_h3/denoise.py | 5 +- .../minimax_music3/denoise.py | 11 +- .../modular_pipelines/modular_pipeline.py | 109 ++++++++++- .../modular_pipelines/qwenimage/denoise.py | 5 +- .../stable_diffusion_3/denoise.py | 11 +- .../stable_diffusion_xl/denoise.py | 13 +- .../modular_pipelines/wan/denoise.py | 5 +- .../wan_animate_2/denoise.py | 8 + .../modular_pipelines/z_image/denoise.py | 5 +- .../cosmos/test_modular_pipeline_cosmos3.py | 2 +- .../helios/test_modular_pipeline_helios.py | 43 ++++- .../ltx2/test_modular_pipeline_ltx2.py | 24 +++ .../test_modular_pipeline_minimax_music3.py | 43 +++++ ...st_modular_pipeline_stable_diffusion_xl.py | 33 ++++ .../test_modular_pipeline_callbacks.py | 182 ++++++++++++++++++ .../modular_pipelines/testing_utils/common.py | 22 +++ .../wan/test_modular_pipeline_wan.py | 19 ++ 31 files changed, 658 insertions(+), 40 deletions(-) create mode 100644 tests/modular_pipelines/test_modular_pipeline_callbacks.py diff --git a/docs/source/en/modular_diffusers/loop_sequential_pipeline_blocks.md b/docs/source/en/modular_diffusers/loop_sequential_pipeline_blocks.md index 74a868922799..555be3a2f808 100644 --- a/docs/source/en/modular_diffusers/loop_sequential_pipeline_blocks.md +++ b/docs/source/en/modular_diffusers/loop_sequential_pipeline_blocks.md @@ -49,6 +49,13 @@ class LoopWrapper(LoopSequentialPipelineBlocks): The loop wrapper can pass additional arguments, like current iteration index, to the loop blocks. +Denoising loops iterate with [`~modular_pipelines.LoopSequentialPipelineBlocks.loop_over_timesteps`] instead of calling `loop_step` directly. It also stops the loop when `pipe.interrupt` is set and runs the `callback_on_step_end` passed to the pipeline. List the fields a callback may read and replace in the wrapper's `_callback_tensor_inputs`. + +```py +for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): + progress_bar.update() +``` + ## Loop blocks A loop block is a [`~modular_pipelines.ModularPipelineBlocks`], but the `__call__` method behaves differently. diff --git a/docs/source/en/modular_diffusers/modular_pipeline.md b/docs/source/en/modular_diffusers/modular_pipeline.md index e234a8aa4aae..ea32c8ed04da 100644 --- a/docs/source/en/modular_diffusers/modular_pipeline.md +++ b/docs/source/en/modular_diffusers/modular_pipeline.md @@ -495,3 +495,58 @@ The `config.json` file contains an `auto_map` key that tells [`ModularPipeline`] ``` Load custom code repositories with `trust_remote_code=True` as shown in [from_pretrained](#frompretrained). See [Custom blocks](./custom_blocks) for how to create and share your own. + +## Step callbacks + +Pass `callback_on_step_end` to observe or modify the state after each denoising step. The signature matches standard +pipeline callbacks: `callback(pipeline, step_index, timestep, callback_kwargs)`. The callback must return a dictionary; +the returned fields replace the loop's values for the following steps. Return `callback_kwargs` unchanged to only observe. + +```python +previews = [] + + +def preview(pipe, step_index, timestep, callback_kwargs): + previews.append(callback_kwargs["latents"].detach().cpu().clone()) + return callback_kwargs + + +state = pipe(prompt="A small boat on a lake", callback_on_step_end=preview) +``` + +By default, the callback receives `latents`. `pipe.callback_tensor_inputs` lists every field the pipeline's denoising +loops support, such as `prompt_embeds` or LTX2's `audio_latents`; pass the ones you need with +`callback_on_step_end_tensor_inputs`. Unlike standard pipelines, requesting or returning a field the running loop doesn't +support raises an error instead of being ignored. Values keep the loop's native shapes, including packed latents. + +```python +def edit_conditioning(pipe, step_index, timestep, callback_kwargs): + if step_index == 2: + callback_kwargs["prompt_embeds"] = replacement_prompt_embeds + return callback_kwargs + + +state = pipe( + prompt="A small boat on a lake", + callback_on_step_end=edit_conditioning, + callback_on_step_end_tensor_inputs=["latents", "prompt_embeds"], +) +``` + +`PipelineCallback` and `MultiPipelineCallbacks` are accepted and provide their own `tensor_inputs`. Standard CFG cutoff +callbacks change attributes that only standard pipelines have; modular pipelines control guidance through +the `guider` component instead. + +Set `pipe.interrupt = True` in a callback to stop denoising early. The pipeline still decodes what was generated so far; +for chunked video or audio this can be shorter than requested. `step_index` counts every denoising step of the call, +across chunks and pyramid stages, so the example below stops after ten steps however the workflow is split. + +```python +def stop_after_ten_steps(pipe, step_index, timestep, callback_kwargs): + if step_index == 9: + pipe.interrupt = True + return callback_kwargs + + +state = pipe(prompt="A small boat on a lake", callback_on_step_end=stop_after_ten_steps) +``` diff --git a/src/diffusers/modular_pipelines/anima/denoise.py b/src/diffusers/modular_pipelines/anima/denoise.py index 86a3569648bc..5d89532c90ca 100644 --- a/src/diffusers/modular_pipelines/anima/denoise.py +++ b/src/diffusers/modular_pipelines/anima/denoise.py @@ -167,6 +167,8 @@ def __call__( class AnimaDenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ("latents", "prompt_embeds", "negative_prompt_embeds") + model_name = "anima" @property @@ -193,8 +195,7 @@ def __call__( num_warmup_steps = len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): if i == len(block_state.timesteps) - 1 or ( (i + 1) > num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): diff --git a/src/diffusers/modular_pipelines/cosmos/denoise.py b/src/diffusers/modular_pipelines/cosmos/denoise.py index 294213f48f4b..a12983ba83a5 100644 --- a/src/diffusers/modular_pipelines/cosmos/denoise.py +++ b/src/diffusers/modular_pipelines/cosmos/denoise.py @@ -483,6 +483,8 @@ def __call__( class Cosmos3DenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ("latents",) + model_name = "cosmos3-omni" @property @@ -543,9 +545,12 @@ def __call__( mixed_precision_reasoner_policy=getattr(block_state, "mixed_precision_reasoner_policy", None), ) trace = [] + components._validate_callback_inputs(self.callback_tensor_inputs) try: with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: for i, t in enumerate(block_state.timesteps): + if components.interrupt: + break apply_cosmos3_mixed_precision_step( components.transformer, mixed_precision, @@ -554,6 +559,7 @@ def __call__( trace=trace, ) components, block_state = self.loop_step(components, block_state, i=i, t=t) + components._call_callback_on_step_end(block_state, t, self.callback_tensor_inputs) if i == len(block_state.timesteps) - 1 or ( (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): diff --git a/src/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py b/src/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py index ab6ebac2c8cd..8fa9d1b9777a 100644 --- a/src/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py +++ b/src/diffusers/modular_pipelines/cosmos/modular_blocks_cosmos3.py @@ -889,6 +889,8 @@ def __call__( state.set("output_chunks", []) state.set("previous_output", None) for chunk_id in range(num_chunks): + if components.interrupt: + break if chunk_id > 0: components.transformer._reset_stateful_cache() state.set("chunk_id", chunk_id) diff --git a/src/diffusers/modular_pipelines/ernie_image/denoise.py b/src/diffusers/modular_pipelines/ernie_image/denoise.py index 150947944dd5..724ec0483066 100644 --- a/src/diffusers/modular_pipelines/ernie_image/denoise.py +++ b/src/diffusers/modular_pipelines/ernie_image/denoise.py @@ -176,6 +176,8 @@ def __call__( class ErnieImageDenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ("latents", "text_bth", "negative_text_bth") + model_name = "ernie-image" @property @@ -219,8 +221,7 @@ def __call__( ) -> tuple[ErnieImageModularPipeline, PipelineState]: block_state = self.get_block_state(state) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): progress_bar.update() self.set_block_state(state, block_state) return components, state diff --git a/src/diffusers/modular_pipelines/flux/denoise.py b/src/diffusers/modular_pipelines/flux/denoise.py index 7f2e20ffcec9..e4bbd69ffea7 100644 --- a/src/diffusers/modular_pipelines/flux/denoise.py +++ b/src/diffusers/modular_pipelines/flux/denoise.py @@ -238,6 +238,8 @@ def __call__( class FluxDenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ("latents", "prompt_embeds", "pooled_prompt_embeds") + model_name = "flux" @property @@ -281,8 +283,7 @@ def __call__( len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order, 0 ) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): if i == len(block_state.timesteps) - 1 or ( (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): diff --git a/src/diffusers/modular_pipelines/flux2/denoise.py b/src/diffusers/modular_pipelines/flux2/denoise.py index fa6877180057..f7dc3fcf4a60 100644 --- a/src/diffusers/modular_pipelines/flux2/denoise.py +++ b/src/diffusers/modular_pipelines/flux2/denoise.py @@ -398,6 +398,8 @@ def __call__( class Flux2DenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ("latents", "prompt_embeds", "negative_prompt_embeds") + model_name = "flux2" @property @@ -442,9 +444,7 @@ def __call__( ) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) - + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): if i == len(block_state.timesteps) - 1 or ( (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): diff --git a/src/diffusers/modular_pipelines/helios/denoise.py b/src/diffusers/modular_pipelines/helios/denoise.py index 32a6f0b11b79..9dc72b9dbfa3 100644 --- a/src/diffusers/modular_pipelines/helios/denoise.py +++ b/src/diffusers/modular_pipelines/helios/denoise.py @@ -82,6 +82,14 @@ def sample_block_noise( return noise +def _resize_to_latent_shape(latents: torch.Tensor, latent_shape: tuple) -> torch.Tensor: + # An interrupted pyramid chunk stops at a lower stage resolution; resize it to the chunk's full latent shape. + batch_size, channels, frames, height, width = latent_shape + latents = latents.permute(0, 2, 1, 3, 4).reshape(batch_size * frames, channels, *latents.shape[-2:]) + latents = F.interpolate(latents, size=(height, width), mode="nearest") + return latents.reshape(batch_size, frames, channels, height, width).permute(0, 2, 1, 3, 4) + + # ======================================== # Chunk Loop Leaf Blocks # ======================================== @@ -365,6 +373,8 @@ def __call__( class HeliosChunkDenoiseInner(ModularPipelineBlocks): """Inner timestep loop for denoising a single chunk, using guider for guidance.""" + _callback_tensor_inputs = ("latents",) + model_name = "helios" @property @@ -432,6 +442,8 @@ def __call__( with tqdm(total=num_inference_steps) as progress_bar: for i, t in enumerate(timesteps): + if components.interrupt: + break timestep = t.expand(latents.shape[0]).to(torch.int64) latent_model_input = latents.to(transformer_dtype) @@ -463,6 +475,9 @@ def __call__( generator=block_state.generator, return_dict=False, )[0] + block_state.latents = latents + components._call_callback_on_step_end(block_state, t, self.callback_tensor_inputs) + latents = block_state.latents if i == len(timesteps) - 1 or ( (i + 1) > num_warmup_steps and (i + 1) % components.scheduler.order == 0 @@ -482,6 +497,8 @@ class HeliosPyramidChunkDenoiseInner(ModularPipelineBlocks): 3. Run timestep denoising loop (same logic as HeliosChunkDenoiseInner) """ + _callback_tensor_inputs = ("latents",) + model_name = "helios-pyramid" @property @@ -555,6 +572,8 @@ def __call__( orig_zero_init_steps = getattr(components.guider, "zero_init_steps", None) for i_s in range(pyramid_num_stages): + if components.interrupt: + break # --- Stage setup --- # Disable zero init for stages > 0 (only stage 0 should have zero init) @@ -624,6 +643,8 @@ def __call__( with tqdm(total=num_inference_steps) as progress_bar: for i, t in enumerate(timesteps): + if components.interrupt: + break timestep = t.expand(latents.shape[0]).to(torch.int64) latent_model_input = latents.to(transformer_dtype) @@ -655,6 +676,9 @@ def __call__( generator=block_state.generator, return_dict=False, )[0] + block_state.latents = latents + components._call_callback_on_step_end(block_state, t, self.callback_tensor_inputs) + latents = block_state.latents if i == len(timesteps) - 1 or ( (i + 1) > num_warmup_steps and (i + 1) % components.scheduler.order == 0 @@ -665,6 +689,9 @@ def __call__( if orig_zero_init_steps is not None: components.guider.zero_init_steps = orig_zero_init_steps + if components.interrupt: + latents = _resize_to_latent_shape(latents, block_state.latent_shape) + block_state.latents = latents return components, block_state @@ -758,7 +785,10 @@ def __call__( if not hasattr(block_state, "image_latents"): block_state.image_latents = None + components._validate_callback_inputs(self.callback_tensor_inputs) for k in range(block_state.num_latent_chunk): + if components.interrupt: + break components, block_state = self.loop_step(components, block_state, k=k) self.set_block_state(state, block_state) @@ -820,6 +850,8 @@ class HeliosPyramidDistilledChunkDenoiseInner(ModularPipelineBlocks): - Tracks start_point_list and passes DMD-specific args to scheduler.step() """ + _callback_tensor_inputs = ("latents",) + model_name = "helios-pyramid" @property @@ -897,6 +929,8 @@ def __call__( shared_kwargs["attention_kwargs"] = block_state.attention_kwargs for i_s in range(pyramid_num_stages): + if components.interrupt: + break # --- Stage setup --- patch_size = components.transformer.config.patch_size @@ -965,6 +999,8 @@ def __call__( with tqdm(total=num_inference_steps) as progress_bar: for i, t in enumerate(timesteps): + if components.interrupt: + break timestep = t.expand(latents.shape[0]).to(torch.int64) latent_model_input = latents.to(transformer_dtype) @@ -1001,12 +1037,18 @@ def __call__( dmd_timesteps=components.scheduler.timesteps, all_timesteps=timesteps, )[0] + block_state.latents = latents + components._call_callback_on_step_end(block_state, t, self.callback_tensor_inputs) + latents = block_state.latents if i == len(timesteps) - 1 or ( (i + 1) > num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): progress_bar.update() + if components.interrupt: + latents = _resize_to_latent_shape(latents, block_state.latent_shape) + block_state.latents = latents return components, block_state diff --git a/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py b/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py index 223d5e2eee9b..c2d004317c2a 100644 --- a/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py +++ b/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py @@ -203,6 +203,8 @@ def __call__( class HunyuanVideo15DenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ("latents", "prompt_embeds", "negative_prompt_embeds") + model_name = "hunyuan-video-1.5" @property @@ -234,8 +236,7 @@ def __call__( ) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): if i == len(block_state.timesteps) - 1 or ( (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): diff --git a/src/diffusers/modular_pipelines/ideogram4/denoise.py b/src/diffusers/modular_pipelines/ideogram4/denoise.py index db1a708c4315..2057d6aa8aa7 100644 --- a/src/diffusers/modular_pipelines/ideogram4/denoise.py +++ b/src/diffusers/modular_pipelines/ideogram4/denoise.py @@ -258,6 +258,8 @@ class Ideogram4DenoiseStep(LoopSequentialPipelineBlocks): The denoised latents. """ + _callback_tensor_inputs = ("latents", "prompt_embeds", "negative_prompt_embeds") + model_name = "ideogram4" block_classes = [Ideogram4LoopBeforeDenoiser, Ideogram4LoopDenoiser, Ideogram4LoopAfterDenoiser] block_names = ["before_denoiser", "denoiser", "after_denoiser"] @@ -292,8 +294,7 @@ def __call__( block_state = self.get_block_state(state) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): progress_bar.update() self.set_block_state(state, block_state) diff --git a/src/diffusers/modular_pipelines/krea2/denoise.py b/src/diffusers/modular_pipelines/krea2/denoise.py index ccb972ca74c1..258de5fb03fb 100644 --- a/src/diffusers/modular_pipelines/krea2/denoise.py +++ b/src/diffusers/modular_pipelines/krea2/denoise.py @@ -242,6 +242,8 @@ def __call__( class Krea2DenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ("latents", "prompt_embeds", "negative_prompt_embeds") + model_name = "krea2" @property @@ -275,8 +277,7 @@ def __call__( block_state = self.get_block_state(state) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): progress_bar.update() self.set_block_state(state, block_state) diff --git a/src/diffusers/modular_pipelines/ltx/denoise.py b/src/diffusers/modular_pipelines/ltx/denoise.py index 44a3c91d471b..1bc50c8c680e 100644 --- a/src/diffusers/modular_pipelines/ltx/denoise.py +++ b/src/diffusers/modular_pipelines/ltx/denoise.py @@ -191,6 +191,8 @@ def __call__( class LTXDenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ("latents", "prompt_embeds", "negative_prompt_embeds") + model_name = "ltx" @property @@ -225,8 +227,7 @@ def __call__( ) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): if i == len(block_state.timesteps) - 1 or ( (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): diff --git a/src/diffusers/modular_pipelines/ltx2/denoise.py b/src/diffusers/modular_pipelines/ltx2/denoise.py index 2fb4951377a1..43111f726ed1 100644 --- a/src/diffusers/modular_pipelines/ltx2/denoise.py +++ b/src/diffusers/modular_pipelines/ltx2/denoise.py @@ -623,6 +623,8 @@ def __call__( class LTX2DenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ("latents", "audio_latents") + model_name = "ltx2" @property @@ -655,8 +657,7 @@ def __call__(self, components, state: PipelineState) -> tuple[LTX2ModularPipelin ) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): if i == len(block_state.timesteps) - 1 or ( (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): diff --git a/src/diffusers/modular_pipelines/minimax_h3/denoise.py b/src/diffusers/modular_pipelines/minimax_h3/denoise.py index efcf36c1130b..da3384aa521e 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/denoise.py +++ b/src/diffusers/modular_pipelines/minimax_h3/denoise.py @@ -242,6 +242,8 @@ def __call__( class MiniMaxH3DenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ("latents", "audio_latents", "prompt_embeds") + model_name = "minimax-h3" @property @@ -267,8 +269,7 @@ def __call__( ) -> tuple[MiniMaxH3ModularPipeline, PipelineState]: block_state = self.get_block_state(state) with self.progress_bar(total=len(block_state.timesteps)) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): progress_bar.update() self.set_block_state(state, block_state) return components, state diff --git a/src/diffusers/modular_pipelines/minimax_music3/denoise.py b/src/diffusers/modular_pipelines/minimax_music3/denoise.py index 0a709b8bf85b..a06691301cec 100644 --- a/src/diffusers/modular_pipelines/minimax_music3/denoise.py +++ b/src/diffusers/modular_pipelines/minimax_music3/denoise.py @@ -164,6 +164,8 @@ def __call__( class MiniMaxMusic3ChunkDenoiseInner(ModularPipelineBlocks): + _callback_tensor_inputs = ("latents",) + model_name = "minimax-music3" @property @@ -213,6 +215,8 @@ def __call__( } for i, t in enumerate(timesteps): + if components.interrupt: + break if overlap > 0: time_value = t.to(latents.dtype) latents[..., :overlap] = (1.0 - (1.0 - 1e-6) * time_value) * block_state.noise_prompt + ( @@ -236,7 +240,9 @@ def __call__( components.guider.cleanup_models(components.transformer) velocity = components.guider(guider_state)[0] - latents = components.scheduler.step(velocity, t, latents, return_dict=False)[0] + block_state.latents = components.scheduler.step(velocity, t, latents, return_dict=False)[0] + components._call_callback_on_step_end(block_state, t, self.callback_tensor_inputs) + latents = block_state.latents block_state.progress_bar.update() block_state.latents = latents @@ -314,7 +320,10 @@ def __call__( num_chunks = len(block_state.chunk_starts) with self.progress_bar(total=num_chunks * block_state.num_inference_steps) as progress_bar: block_state.progress_bar = progress_bar + components._validate_callback_inputs(self.callback_tensor_inputs) for k in range(num_chunks): + if components.interrupt: + break components, block_state = self.loop_step(components, block_state, k=k) block_state.progress_bar = None diff --git a/src/diffusers/modular_pipelines/modular_pipeline.py b/src/diffusers/modular_pipelines/modular_pipeline.py index c501d7b1fc42..d7c277ca8580 100644 --- a/src/diffusers/modular_pipelines/modular_pipeline.py +++ b/src/diffusers/modular_pipelines/modular_pipeline.py @@ -21,7 +21,7 @@ from collections import OrderedDict from copy import deepcopy from dataclasses import dataclass, field -from typing import Any +from typing import Any, Callable import torch from huggingface_hub import create_repo @@ -29,6 +29,7 @@ from tqdm.auto import tqdm from typing_extensions import Self +from ..callbacks import MultiPipelineCallbacks, PipelineCallback from ..configuration_utils import ConfigMixin, FrozenDict from ..models.auto_model import AutoModel from ..models.modeling_utils import ModelMixin @@ -382,6 +383,7 @@ class ModularPipelineBlocks(ConfigMixin, PushToHubMixin): model_name = None _requirements: dict[str, str] | None = None _workflow_map = None + _callback_tensor_inputs = () @classmethod def _get_signature_keys(cls, obj): @@ -400,6 +402,14 @@ def description(self) -> str: """Description of the block. Must be implemented by subclasses.""" return "" + @property + def callback_tensor_inputs(self) -> list[str]: + """Fields a step callback can read and update, including those declared by nested blocks.""" + names = list(self._callback_tensor_inputs) + for block in self.sub_blocks.values(): + names += [name for name in block.callback_tensor_inputs if name not in names] + return names + @property def expected_components(self) -> list[ComponentSpec]: return [] @@ -1578,6 +1588,19 @@ def loop_step(self, components, state: PipelineState, **kwargs): raise return components, state + def loop_over_timesteps(self, components, block_state: BlockState, timesteps): + """ + Runs `loop_step` for each timestep and yields `(i, t)` after the step-end callback. Stops early once + `components.interrupt` is set. Loop sub-blocks must update and return the `block_state` they receive. + """ + components._validate_callback_inputs(self.callback_tensor_inputs) + for i, t in enumerate(timesteps): + if components.interrupt: + break + components, block_state = self.loop_step(components, block_state, i=i, t=t) + components._call_callback_on_step_end(block_state, t, self.callback_tensor_inputs) + yield i, t + def __call__(self, components, state: PipelineState) -> tuple["ModularPipeline", PipelineState]: raise NotImplementedError("`__call__` method needs to be implemented by the subclass") @@ -1684,6 +1707,10 @@ class ModularPipeline(ConfigMixin, PushToHubMixin): config_name = "modular_model_index.json" hf_device_map = None default_blocks_name = None + interrupt = False + _callback_on_step_end = None + _callback_on_step_end_tensor_inputs = None + _callback_step_index = 0 # YiYi TODO: add warning for passing multiple ComponentSpec/ConfigSpec with the same name def __init__( @@ -2937,11 +2964,54 @@ def _dict_to_component_spec(name: str, spec_dict: dict[str, Any]) -> ComponentSp ) def set_progress_bar_config(self, **kwargs): - for sub_block_name, sub_block in self._blocks.sub_blocks.items(): - if hasattr(sub_block, "set_progress_bar_config"): - sub_block.set_progress_bar_config(**kwargs) + blocks = [self._blocks] + while blocks: + block = blocks.pop() + if hasattr(block, "set_progress_bar_config"): + block.set_progress_bar_config(**kwargs) + blocks.extend(block.sub_blocks.values()) - def __call__(self, state: PipelineState = None, output: str | list[str] = None, **kwargs): + @property + def callback_tensor_inputs(self) -> list[str]: + """Callback fields supported by the blocks; availability depends on the active workflow.""" + return self._blocks.callback_tensor_inputs + + def _validate_callback_inputs(self, tensor_inputs: list[str]): + if self._callback_on_step_end is None: + return + for name in self._callback_on_step_end_tensor_inputs: + if name not in tensor_inputs: + raise ValueError(f"Callback input '{name}' is unavailable in this denoising loop.") + + def _call_callback_on_step_end(self, block_state: BlockState, timestep: torch.Tensor, tensor_inputs: list[str]): + callback = self._callback_on_step_end + if callback is None: + return + values = vars(block_state) + missing = [name for name in self._callback_on_step_end_tensor_inputs if name not in values] + if missing: + raise ValueError(f"Callback inputs {missing} are unavailable in this denoising loop.") + callback_kwargs = {name: values[name] for name in self._callback_on_step_end_tensor_inputs} + updates = callback(self, self._callback_step_index, timestep, callback_kwargs) + if not isinstance(updates, dict): + raise TypeError("The step callback must return a dictionary.") + for name, value in updates.items(): + if name not in tensor_inputs or name not in values: + raise ValueError(f"Callback output '{name}' is unavailable in this denoising loop.") + setattr(block_state, name, value) + for fields in values.values(): + if isinstance(fields, dict) and name in fields: + fields[name] = value + self._callback_step_index += 1 + + def __call__( + self, + state: PipelineState = None, + output: str | list[str] = None, + callback_on_step_end: Callable | PipelineCallback | MultiPipelineCallbacks | None = None, + callback_on_step_end_tensor_inputs: list[str] | None = None, + **kwargs, + ): """ Execute the pipeline by running the pipeline blocks with the given inputs. @@ -2955,6 +3025,14 @@ def __call__(self, state: PipelineState = None, output: str | list[str] = None, - str: Returns a specific intermediate value from the state (e.g. `output="image"`) - list[str]: Returns a dictionary of specific intermediate values (e.g. `output=["image", "latents"]`) + callback_on_step_end (`Callable`, `PipelineCallback` or `MultiPipelineCallbacks`, optional): + Called at the end of each denoising step as `callback(pipeline, step_index, timestep, + callback_kwargs)`, where `step_index` counts all denoising steps of the call. It must return a dict + whose fields replace the loop's values for the following steps. Set `pipeline.interrupt = True` to stop + denoising early. + callback_on_step_end_tensor_inputs (`list[str]`, optional, defaults to `["latents"]`): + Fields passed in `callback_kwargs`. Must be in `callback_tensor_inputs` of every denoising loop that + runs. Examples: @@ -2981,6 +3059,22 @@ def __call__(self, state: PipelineState = None, output: str | list[str] = None, - If `output` is list[str]: Dictionary mapping output names to their values from the state (e.g. `output=["image", "latents"]`) """ + self.interrupt = False + self._callback_step_index = 0 + if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)): + callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs + if callback_on_step_end is not None and not callable(callback_on_step_end): + raise TypeError("callback_on_step_end must be callable.") + tensor_inputs = ( + ["latents"] if callback_on_step_end_tensor_inputs is None else callback_on_step_end_tensor_inputs + ) + if not isinstance(tensor_inputs, list) or not all(isinstance(name, str) for name in tensor_inputs): + raise TypeError("callback_on_step_end_tensor_inputs must be a list of field names.") + if callback_on_step_end is not None or callback_on_step_end_tensor_inputs is not None: + for name in tensor_inputs: + if name not in self.callback_tensor_inputs: + raise ValueError(f"Callback input '{name}' is not supported by this pipeline.") + if state is None: state = PipelineState() else: @@ -3008,6 +3102,8 @@ def __call__(self, state: PipelineState = None, output: str | list[str] = None, if len(passed_kwargs) > 0: warnings.warn(f"Unexpected input '{passed_kwargs.keys()}' provided. This input will be ignored.") # Run the pipeline + self._callback_on_step_end = callback_on_step_end + self._callback_on_step_end_tensor_inputs = tensor_inputs with torch.no_grad(): try: _, state = self._blocks(self, state) @@ -3015,6 +3111,9 @@ def __call__(self, state: PipelineState = None, output: str | list[str] = None, error_msg = f"Error in block: ({self._blocks.__class__.__name__}):\n" logger.error(error_msg) raise + finally: + self._callback_on_step_end = None + self._callback_on_step_end_tensor_inputs = None if output is None: return state diff --git a/src/diffusers/modular_pipelines/qwenimage/denoise.py b/src/diffusers/modular_pipelines/qwenimage/denoise.py index 7f271782f82b..da52f3833d40 100644 --- a/src/diffusers/modular_pipelines/qwenimage/denoise.py +++ b/src/diffusers/modular_pipelines/qwenimage/denoise.py @@ -452,6 +452,8 @@ def __call__( # 2. DENOISE LOOP WRAPPER: define the denoising loop logic # ==================== class QwenImageDenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ("latents", "prompt_embeds", "negative_prompt_embeds") + model_name = "qwenimage" @property @@ -492,8 +494,7 @@ def __call__( block_state.additional_cond_kwargs = {} with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): if i == len(block_state.timesteps) - 1 or ( (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): diff --git a/src/diffusers/modular_pipelines/stable_diffusion_3/denoise.py b/src/diffusers/modular_pipelines/stable_diffusion_3/denoise.py index cde6ace66245..21ee5ad17311 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_3/denoise.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_3/denoise.py @@ -190,6 +190,14 @@ def __call__( class StableDiffusion3DenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ( + "latents", + "prompt_embeds", + "negative_prompt_embeds", + "pooled_prompt_embeds", + "negative_pooled_prompt_embeds", + ) + model_name = "stable-diffusion-3" @property @@ -217,8 +225,7 @@ def __call__( ) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): if i == len(block_state.timesteps) - 1 or ( (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): diff --git a/src/diffusers/modular_pipelines/stable_diffusion_xl/denoise.py b/src/diffusers/modular_pipelines/stable_diffusion_xl/denoise.py index 16a8b236ce2e..534029fac984 100644 --- a/src/diffusers/modular_pipelines/stable_diffusion_xl/denoise.py +++ b/src/diffusers/modular_pipelines/stable_diffusion_xl/denoise.py @@ -654,6 +654,16 @@ def __call__( # the loop wrapper that iterates over the timesteps class StableDiffusionXLDenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ( + "latents", + "prompt_embeds", + "negative_prompt_embeds", + "pooled_prompt_embeds", + "negative_pooled_prompt_embeds", + "add_time_ids", + "negative_add_time_ids", + ) + model_name = "stable-diffusion-xl" @property @@ -710,8 +720,7 @@ def __call__( ) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): if i == len(block_state.timesteps) - 1 or ( (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): diff --git a/src/diffusers/modular_pipelines/wan/denoise.py b/src/diffusers/modular_pipelines/wan/denoise.py index 3036b33868d1..9f0fb6656c7d 100644 --- a/src/diffusers/modular_pipelines/wan/denoise.py +++ b/src/diffusers/modular_pipelines/wan/denoise.py @@ -414,6 +414,8 @@ def __call__( class WanDenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ("latents", "prompt_embeds", "negative_prompt_embeds") + model_name = "wan" @property @@ -457,8 +459,7 @@ def __call__( ) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): if i == len(block_state.timesteps) - 1 or ( (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): diff --git a/src/diffusers/modular_pipelines/wan_animate_2/denoise.py b/src/diffusers/modular_pipelines/wan_animate_2/denoise.py index 1e18b3a09eeb..a4b354fdb4c0 100644 --- a/src/diffusers/modular_pipelines/wan_animate_2/denoise.py +++ b/src/diffusers/modular_pipelines/wan_animate_2/denoise.py @@ -427,6 +427,8 @@ def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnim class WanAnimate2SegmentDenoiseInner(ModularPipelineBlocks): + _callback_tensor_inputs = ("latents",) + model_name = "wan-animate-2" @property @@ -549,6 +551,8 @@ def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnim total=len(block_state.timesteps), desc=f"Segment {k + 1}/{block_state.num_segments}" ) as progress_bar: for i, t in enumerate(block_state.timesteps): + if components.interrupt: + break timestep = torch.stack([t]) components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) @@ -584,6 +588,7 @@ def __call__(self, components, block_state: BlockState, k: int) -> tuple[WanAnim generator=block_state.generator, )[0] block_state.latents = latents.squeeze(0) + components._call_callback_on_step_end(block_state, t, self.callback_tensor_inputs) progress_bar.update() @@ -740,7 +745,10 @@ def __call__(self, components, state: PipelineState) -> tuple[WanAnimate2Modular block_state.segment_frames = [] block_state.out_frames = None + components._validate_callback_inputs(self.callback_tensor_inputs) for k in range(block_state.num_segments): + if components.interrupt: + break components, block_state = self.loop_step(components, block_state, k=k) self.set_block_state(state, block_state) diff --git a/src/diffusers/modular_pipelines/z_image/denoise.py b/src/diffusers/modular_pipelines/z_image/denoise.py index 899800a5019a..a9e88e52dba6 100644 --- a/src/diffusers/modular_pipelines/z_image/denoise.py +++ b/src/diffusers/modular_pipelines/z_image/denoise.py @@ -240,6 +240,8 @@ def __call__( class ZImageDenoiseLoopWrapper(LoopSequentialPipelineBlocks): + _callback_tensor_inputs = ("latents", "prompt_embeds", "negative_prompt_embeds") + model_name = "z-image" @property @@ -283,8 +285,7 @@ def __call__( ) with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: - for i, t in enumerate(block_state.timesteps): - components, block_state = self.loop_step(components, block_state, i=i, t=t) + for i, t in self.loop_over_timesteps(components, block_state, block_state.timesteps): if i == len(block_state.timesteps) - 1 or ( (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): diff --git a/tests/modular_pipelines/cosmos/test_modular_pipeline_cosmos3.py b/tests/modular_pipelines/cosmos/test_modular_pipeline_cosmos3.py index c3bcb267758f..9015878d79cd 100644 --- a/tests/modular_pipelines/cosmos/test_modular_pipeline_cosmos3.py +++ b/tests/modular_pipelines/cosmos/test_modular_pipeline_cosmos3.py @@ -257,7 +257,7 @@ def test_transfer_chunks_reset_stateful_cache_at_boundaries(self): block = Cosmos3TransferChunkDenoiseStep() child_block = mock.Mock(side_effect=lambda components, state: (components, state)) block.sub_blocks = {"child": child_block} - components = mock.Mock() + components = mock.Mock(interrupt=False) state = mock.Mock() state.get.return_value = 3 diff --git a/tests/modular_pipelines/helios/test_modular_pipeline_helios.py b/tests/modular_pipelines/helios/test_modular_pipeline_helios.py index 5c270c20d9de..b4596fc71c0b 100644 --- a/tests/modular_pipelines/helios/test_modular_pipeline_helios.py +++ b/tests/modular_pipelines/helios/test_modular_pipeline_helios.py @@ -14,6 +14,7 @@ # limitations under the License. import pytest +import torch from diffusers.modular_pipelines import ( HeliosAutoBlocks, @@ -92,7 +93,43 @@ def get_dummy_inputs(self, seed=0): return inputs -class TestHeliosModularPipelineFast(HeliosModularPipelineTesterConfig, ModularPipelineTesterMixin): +class HeliosChunkCallbackTesterMixin: + def test_step_callback_chunks_and_stages(self): + pipe = self.get_pipeline().to("cpu") + inputs = self.get_dummy_inputs() + inputs.update(num_frames=18, num_latent_frames_per_chunk=3) + steps_per_chunk = sum(inputs.get("pyramid_num_inference_steps_list", [inputs.get("num_inference_steps")])) + steps = [] + + def record(pipeline, step, timestep, tensors): + steps.append(step) + return tensors + + full = pipe(**inputs, callback_on_step_end=record, output=["videos", "latent_chunks"]) + assert steps == list(range(2 * steps_per_chunk)) + assert len(full["latent_chunks"]) == 2 + stop_steps = range(0, 2 * steps_per_chunk, 2) + for stop_step in stop_steps: + steps.clear() + inputs["generator"] = self.get_generator() + + def stop(pipeline, step, timestep, tensors): + steps.append(step) + if step == stop_step: + pipeline.interrupt = True + return tensors + + partial = pipe(**inputs, callback_on_step_end=stop, output=["videos", "latent_chunks"]) + assert steps == list(range(stop_step + 1)) + chunks = stop_step // steps_per_chunk + 1 + assert len(partial["latent_chunks"]) == chunks + assert partial["videos"].shape == (1, chunks * 8 + 1, 3, inputs["height"], inputs["width"]) + assert torch.isfinite(partial["videos"]).all() + + +class TestHeliosModularPipelineFast( + HeliosModularPipelineTesterConfig, HeliosChunkCallbackTesterMixin, ModularPipelineTesterMixin +): @pytest.mark.skip(reason="num_videos_per_prompt") def test_num_images_per_prompt(self): pass @@ -168,7 +205,9 @@ def get_dummy_inputs(self, seed=0): return inputs -class TestHeliosPyramidModularPipelineFast(HeliosPyramidModularPipelineTesterConfig, ModularPipelineTesterMixin): +class TestHeliosPyramidModularPipelineFast( + HeliosPyramidModularPipelineTesterConfig, HeliosChunkCallbackTesterMixin, ModularPipelineTesterMixin +): def test_inference_batch_single_identical(self): # Pyramid pipeline injects noise at each stage, so batch vs single can differ more super().test_inference_batch_single_identical(expected_max_diff=5e-1) diff --git a/tests/modular_pipelines/ltx2/test_modular_pipeline_ltx2.py b/tests/modular_pipelines/ltx2/test_modular_pipeline_ltx2.py index 2ee330ece20a..27864145e729 100644 --- a/tests/modular_pipelines/ltx2/test_modular_pipeline_ltx2.py +++ b/tests/modular_pipelines/ltx2/test_modular_pipeline_ltx2.py @@ -24,6 +24,7 @@ from diffusers.modular_pipelines import ComponentSpec, LTX2AutoBlocks, LTX2ModularPipeline from diffusers.pipelines.ltx2.pipeline_ltx2_condition import LTX2VideoCondition from diffusers.pipelines.ltx2.pipeline_ltx2_ic_lora import LTX2ReferenceCondition +from diffusers.utils.testing_utils import torch_device from ..testing_utils import ( BaseModularPipelineTesterConfig, @@ -167,6 +168,29 @@ def test_auto_duration_predicts_a_grid_valid_frame_count(self): assert (num_frames - 1) % pipe.vae_temporal_compression_ratio == 0 assert 0 < num_frames <= round(2.0 * inputs["frame_rate"]) + def test_step_callback_audio_latents_update(self): + pipe = self.get_pipeline().to(torch_device) + observed = [] + modified = False + + def capture(module, args, kwargs): + if modified: + observed.append(kwargs["audio_hidden_states"].detach().clone()) + + handle = pipe.transformer.register_forward_pre_hook(capture, with_kwargs=True) + + def replace(pipeline, step, timestep, tensors): + nonlocal modified + modified = True + return {"audio_latents": torch.zeros_like(tensors["audio_latents"])} + + try: + self.run_pipe(pipe, callback_on_step_end=replace, callback_on_step_end_tensor_inputs=["audio_latents"]) + finally: + handle.remove() + assert observed + assert all(tensor.count_nonzero() == 0 for tensor in observed) + class TestLTX2Text2VideoModularPipelineLoading(LTX2Text2VideoModularPipelineTesterConfig, ModularLoadingTesterMixin): def test_guiders_round_trip_as_pretrained_components(self, tmp_path): diff --git a/tests/modular_pipelines/minimax_music3/test_modular_pipeline_minimax_music3.py b/tests/modular_pipelines/minimax_music3/test_modular_pipeline_minimax_music3.py index 923c4704dd7b..9cf40ca421e2 100644 --- a/tests/modular_pipelines/minimax_music3/test_modular_pipeline_minimax_music3.py +++ b/tests/modular_pipelines/minimax_music3/test_modular_pipeline_minimax_music3.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import math import unittest import pytest @@ -23,6 +24,7 @@ MiniMaxMusic3ModularPipeline, ModularPipeline, ) +from diffusers.modular_pipelines.modular_pipeline import SequentialPipelineBlocks from ...testing_utils import enable_full_determinism, torch_device from ..testing_utils import ( @@ -134,6 +136,47 @@ def test_output_is_stereo_waveform(self): assert audio.shape[1] == 2 assert audio.abs().max() <= 1.0 + def test_step_callback_multiple_chunks(self, monkeypatch): + blocks = MiniMaxMusic3Blocks() + blocks = SequentialPipelineBlocks.from_blocks_dict( + {name: block for name, block in blocks.sub_blocks.items() if name != "semantic_generator"} + ) + pipe = blocks.init_pipeline(self.pretrained_model_name_or_path) + pipe.load_components() + sample_hop = math.prod(pipe.vocoder.config.upsampling_ratios) + monkeypatch.setattr(type(pipe), "latent_hop_length", property(lambda pipeline: sample_hop)) + config = pipe.condition_encoder.config + frame_hiddens = torch.randn(1, 201, config.condition_hidden_dim * config.num_condition_layers) + steps = [] + + def record(pipeline, step, timestep, tensors): + steps.append(step) + return tensors + + inputs = {"frame_hiddens": frame_hiddens, "num_inference_steps": 2, "output_type": "pt"} + full = pipe( + **inputs, generator=self.get_generator(), callback_on_step_end=record, output=["audios", "latent_chunks"] + ) + assert steps == list(range(4)) + assert len(full["latent_chunks"]) == 2 + for stop_step in (0, 2): + steps.clear() + + def stop(pipeline, step, timestep, tensors): + steps.append(step) + if step == stop_step: + pipeline.interrupt = True + return tensors + + partial = pipe( + **inputs, generator=self.get_generator(), callback_on_step_end=stop, output=["audios", "latent_chunks"] + ) + assert steps == list(range(stop_step + 1)) + assert len(partial["latent_chunks"]) == stop_step // 2 + 1 + assert torch.isfinite(partial["audios"]).all() + if stop_step == 0: + assert 0 < partial["audios"].shape[-1] < full["audios"].shape[-1] + class TestMiniMaxMusic3ModularPipelineLoading(MiniMaxMusic3ModularPipelineTesterConfig, ModularLoadingTesterMixin): def test_save_from_pretrained(self, tmp_path, base_pipe_output): diff --git a/tests/modular_pipelines/stable_diffusion_xl/test_modular_pipeline_stable_diffusion_xl.py b/tests/modular_pipelines/stable_diffusion_xl/test_modular_pipeline_stable_diffusion_xl.py index 1e1777b90790..60c8bd9f9dd8 100644 --- a/tests/modular_pipelines/stable_diffusion_xl/test_modular_pipeline_stable_diffusion_xl.py +++ b/tests/modular_pipelines/stable_diffusion_xl/test_modular_pipeline_stable_diffusion_xl.py @@ -374,6 +374,39 @@ def test_stable_diffusion_xl_euler(self): def test_inference_batch_single_identical(self): super().test_inference_batch_single_identical(expected_max_diff=3e-3) + def test_step_callback_conditioning_update(self): + pipe = self.get_pipeline().to(torch_device) + observed = [] + modified = False + + def capture(module, args, kwargs): + if modified: + observed.append(kwargs["encoder_hidden_states"]) + observed.append(kwargs["added_cond_kwargs"]["text_embeds"]) + observed.append(kwargs["added_cond_kwargs"]["time_ids"]) + + handle = pipe.unet.register_forward_pre_hook(capture, with_kwargs=True) + fields = [name for name in pipe.callback_tensor_inputs if name != "latents"] + + def replace(pipeline, step, timestep, tensors): + nonlocal modified + modified = True + return {name: torch.zeros_like(value) for name, value in tensors.items()} + + try: + self.run_pipe(pipe, callback_on_step_end=replace, callback_on_step_end_tensor_inputs=fields) + finally: + handle.remove() + assert observed + assert all(tensor.count_nonzero() == 0 for tensor in observed) + + def test_set_progress_bar_config_reaches_nested_loop(self, capfd): + pipe = self.get_pipeline().to(torch_device) + pipe.set_progress_bar_config(disable=True) + capfd.readouterr() + self.run_pipe(pipe) + assert "it/s" not in capfd.readouterr().err + class TestSDXLModularPipelineIPAdapter(SDXLModularPipelineTesterConfig, SDXLModularIPAdapterTesterMixin): pass diff --git a/tests/modular_pipelines/test_modular_pipeline_callbacks.py b/tests/modular_pipelines/test_modular_pipeline_callbacks.py new file mode 100644 index 000000000000..7995ebb1a183 --- /dev/null +++ b/tests/modular_pipelines/test_modular_pipeline_callbacks.py @@ -0,0 +1,182 @@ +# 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. + +import pytest +import torch + +from diffusers import ModularPipeline, ModularPipelineBlocks +from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback +from diffusers.modular_pipelines.modular_pipeline import SequentialPipelineBlocks +from diffusers.modular_pipelines.modular_pipeline_utils import ComponentSpec, InputParam, OutputParam +from diffusers.modular_pipelines.wan.denoise import WanDenoiseLoopWrapper +from diffusers.schedulers import FlowMatchEulerDiscreteScheduler + + +class CallbackSetup(ModularPipelineBlocks): + @property + def expected_components(self): + return [ + ComponentSpec( + "scheduler", FlowMatchEulerDiscreteScheduler, config={}, default_creation_method="from_config" + ) + ] + + @property + def inputs(self): + return [ + InputParam("latents", required=True), + InputParam("prompt_embeds", required=True), + InputParam("num_inference_steps", default=3), + ] + + @property + def intermediate_outputs(self): + return [OutputParam("timesteps"), OutputParam("prompt_embeds", kwargs_type="denoiser_input_fields")] + + def __call__(self, components, state): + block_state = self.get_block_state(state) + components.scheduler.set_timesteps(block_state.num_inference_steps) + block_state.timesteps = components.scheduler.timesteps + self.set_block_state(state, block_state) + return components, state + + +class CallbackSchedulerStep(ModularPipelineBlocks): + @property + def inputs(self): + return [ + InputParam("latents", required=True), + InputParam("prompt_embeds", required=True), + InputParam.template("denoiser_input_fields"), + ] + + @property + def intermediate_outputs(self): + return [OutputParam("latents")] + + def __call__(self, components, block_state, i, t): + assert block_state.prompt_embeds is block_state.denoiser_input_fields["prompt_embeds"] + block_state.latents = components.scheduler.step( + block_state.prompt_embeds, t, block_state.latents, return_dict=False + )[0] + return components, block_state + + +class CallbackLoop(WanDenoiseLoopWrapper): + block_classes = [CallbackSchedulerStep] + block_names = ["scheduler"] + + @property + def loop_expected_components(self): + return CallbackSetup().expected_components + + +class ZeroLatentsCallback(PipelineCallback): + @property + def tensor_inputs(self): + return ["latents"] + + def callback_fn(self, pipeline, step_index, timestep, callback_kwargs): + return {"latents": torch.zeros_like(callback_kwargs["latents"])} + + +class TestModularCallbacks: + def get_pipeline(self, loops=1): + blocks = {"setup": CallbackSetup()} + for i in range(loops): + if i: + blocks[f"setup_{i}"] = CallbackSetup() + blocks[f"denoise_{i}"] = CallbackLoop() + return ModularPipeline(blocks=SequentialPipelineBlocks.from_blocks_dict(blocks)) + + def run(self, pipe, **kwargs): + return pipe(latents=torch.zeros(1, 2), prompt_embeds=torch.ones(1, 2), output="latents", **kwargs) + + def test_read_only_callback_and_global_steps(self): + pipe = self.get_pipeline(loops=2) + baseline = self.run(pipe) + steps = [] + + def callback(pipeline, step, timestep, tensors): + steps.append(step) + return tensors + + output = self.run(pipe, callback_on_step_end=callback) + torch.testing.assert_close(output, baseline, rtol=0, atol=0) + assert steps == list(range(6)) + assert not pipe.interrupt + assert pipe._callback_on_step_end is None + + def test_latents_and_conditioning_updates(self): + pipe = self.get_pipeline() + observed = [] + + def callback(pipeline, step, timestep, tensors): + observed.append(tensors["latents"].clone()) + return { + "latents": torch.zeros_like(tensors["latents"]), + "prompt_embeds": torch.zeros_like(tensors["latents"]), + } + + output = self.run(pipe, callback_on_step_end=callback) + assert observed[0].abs().sum() > 0 + assert observed[1].count_nonzero() == 0 + assert output.count_nonzero() == 0 + + def test_callback_objects(self): + pipe = self.get_pipeline() + for callback in [ + ZeroLatentsCallback(), + MultiPipelineCallbacks([ZeroLatentsCallback(), ZeroLatentsCallback()]), + ]: + output = self.run(pipe, callback_on_step_end=callback, callback_on_step_end_tensor_inputs=["invalid"]) + assert output.count_nonzero() == 0 + + def test_interrupt_and_next_call(self): + pipe = self.get_pipeline(loops=2) + steps = [] + + def callback(pipeline, step, timestep, tensors): + steps.append(step) + pipeline.interrupt = True + return tensors + + output = self.run(pipe, callback_on_step_end=callback) + assert steps == [0] + assert torch.isfinite(output).all() + assert pipe.interrupt + assert not torch.equal(self.run(pipe), output) + assert not pipe.interrupt + + @pytest.mark.parametrize( + "callback,inputs,error", + [ + (lambda *args: None, None, TypeError), + (lambda *args: {"invalid": 1}, None, ValueError), + (lambda *args: {}, ["invalid"], ValueError), + (lambda *args: {}, ["negative_prompt_embeds"], ValueError), + (0, None, TypeError), + (None, ["invalid"], ValueError), + (lambda *args: {}, "latents", TypeError), + ], + ) + def test_invalid_callbacks_and_cleanup(self, callback, inputs, error): + pipe = self.get_pipeline() + pipe.interrupt = True + with pytest.raises(error): + self.run(pipe, callback_on_step_end=callback, callback_on_step_end_tensor_inputs=inputs) + assert not pipe.interrupt + assert getattr(pipe, "_callback_on_step_end", None) is None + assert torch.isfinite(self.run(pipe)).all() diff --git a/tests/modular_pipelines/testing_utils/common.py b/tests/modular_pipelines/testing_utils/common.py index 26d7391a286b..58f5d05007c2 100644 --- a/tests/modular_pipelines/testing_utils/common.py +++ b/tests/modular_pipelines/testing_utils/common.py @@ -190,6 +190,28 @@ class ModularPipelineTesterMixin(BaseModularPipelineOutputMixin): `pretrained_model_name_or_path`, `get_dummy_inputs()` and the shared fixtures). """ + def test_step_callback_and_interrupt(self): + pipe = self.get_pipeline().to(torch_device) + steps = [] + + def record(pipeline, step, timestep, tensors): + assert isinstance(tensors["latents"], torch.Tensor) + steps.append(step) + return tensors + + self.run_pipe(pipe, callback_on_step_end=record) + assert steps and steps == list(range(len(steps))) + steps.clear() + + def stop(pipeline, step, timestep, tensors): + steps.append(step) + pipeline.interrupt = True + return tensors + + assert self.run_pipe(pipe, callback_on_step_end=stop) is not None + assert steps == [0] + assert pipe.interrupt + def test_pipeline_call_signature(self): pipe = self.get_pipeline() input_parameters = pipe.blocks.input_names diff --git a/tests/modular_pipelines/wan/test_modular_pipeline_wan.py b/tests/modular_pipelines/wan/test_modular_pipeline_wan.py index d35c21455ba9..7ec7daf2d10b 100644 --- a/tests/modular_pipelines/wan/test_modular_pipeline_wan.py +++ b/tests/modular_pipelines/wan/test_modular_pipeline_wan.py @@ -55,6 +55,25 @@ class TestWanModularPipelineFast(WanModularPipelineTesterConfig, ModularPipeline def test_num_images_per_prompt(self): pass + def test_step_callback_guidance(self): + pipe = self.get_pipeline().to("cpu") + forwards = [] + counts = [] + handle = pipe.transformer.register_forward_pre_hook(lambda module, args: forwards.append(module)) + + def cutoff(pipeline, step, timestep, tensors): + counts.append(len(forwards)) + pipeline.guider.disable() + return tensors + + try: + self.run_pipe(pipe, callback_on_step_end=cutoff) + finally: + handle.remove() + # CFG runs two forwards in the first step and one per step after the callback disables it + assert counts[0] == 2 + assert [b - a for a, b in zip(counts, counts[1:])] == [1] * (len(counts) - 1) + class TestWanModularPipelineLoading(WanModularPipelineTesterConfig, ModularLoadingTesterMixin): pass From 234f3dee549b9e2b627e96ea1cb487fcb468e466 Mon Sep 17 00:00:00 2001 From: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com> Date: Tue, 6 Oct 2026 16:44:31 +0000 Subject: [PATCH 2/2] Wire step callbacks and interrupt into the Echo denoise loop --- src/diffusers/modular_pipelines/echo/denoise.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/src/diffusers/modular_pipelines/echo/denoise.py b/src/diffusers/modular_pipelines/echo/denoise.py index cbe80c0ab609..744b91e2b623 100644 --- a/src/diffusers/modular_pipelines/echo/denoise.py +++ b/src/diffusers/modular_pipelines/echo/denoise.py @@ -411,6 +411,8 @@ class EchoDenoiseLoopStep(LoopSequentialPipelineBlocks): Packed clean first-frame tokens. """ + _callback_tensor_inputs = ("latents", "audio_latents") + model_name = "echo" block_classes = EchoDenoiseLoopBlocks.values() block_names = EchoDenoiseLoopBlocks.keys() @@ -447,9 +449,13 @@ def __call__(self, components, state: PipelineState) -> PipelineState: raise ValueError("Echo `sigmas` must be monotonically non-increasing.") block_state.sigmas = sigmas + components._validate_callback_inputs(self.callback_tensor_inputs) with self.progress_bar(total=len(sigmas) - 1) as progress_bar: for i, sigma in enumerate(sigmas[:-1]): + if components.interrupt: + break components, block_state = self.loop_step(components, block_state, i=i, sigma=sigma) + components._call_callback_on_step_end(block_state, sigma, self.callback_tensor_inputs) progress_bar.update() self.set_block_state(state, block_state)