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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
55 changes: 55 additions & 0 deletions docs/source/en/modular_diffusers/modular_pipeline.md
Original file line number Diff line number Diff line change
Expand Up @@ -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)
```
5 changes: 3 additions & 2 deletions src/diffusers/modular_pipelines/anima/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -167,6 +167,8 @@ def __call__(


class AnimaDenoiseLoopWrapper(LoopSequentialPipelineBlocks):
_callback_tensor_inputs = ("latents", "prompt_embeds", "negative_prompt_embeds")

model_name = "anima"

@property
Expand All @@ -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
):
Expand Down
6 changes: 6 additions & 0 deletions src/diffusers/modular_pipelines/cosmos/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -483,6 +483,8 @@ def __call__(


class Cosmos3DenoiseLoopWrapper(LoopSequentialPipelineBlocks):
_callback_tensor_inputs = ("latents",)

model_name = "cosmos3-omni"

@property
Expand Down Expand Up @@ -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,
Expand All @@ -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
):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
6 changes: 6 additions & 0 deletions src/diffusers/modular_pipelines/echo/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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)
Expand Down
5 changes: 3 additions & 2 deletions src/diffusers/modular_pipelines/ernie_image/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -176,6 +176,8 @@ def __call__(


class ErnieImageDenoiseLoopWrapper(LoopSequentialPipelineBlocks):
_callback_tensor_inputs = ("latents", "text_bth", "negative_text_bth")

model_name = "ernie-image"

@property
Expand Down Expand Up @@ -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
Expand Down
5 changes: 3 additions & 2 deletions src/diffusers/modular_pipelines/flux/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -238,6 +238,8 @@ def __call__(


class FluxDenoiseLoopWrapper(LoopSequentialPipelineBlocks):
_callback_tensor_inputs = ("latents", "prompt_embeds", "pooled_prompt_embeds")

model_name = "flux"

@property
Expand Down Expand Up @@ -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
):
Expand Down
6 changes: 3 additions & 3 deletions src/diffusers/modular_pipelines/flux2/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -398,6 +398,8 @@ def __call__(


class Flux2DenoiseLoopWrapper(LoopSequentialPipelineBlocks):
_callback_tensor_inputs = ("latents", "prompt_embeds", "negative_prompt_embeds")

model_name = "flux2"

@property
Expand Down Expand Up @@ -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
):
Expand Down
42 changes: 42 additions & 0 deletions src/diffusers/modular_pipelines/helios/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
# ========================================
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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
Expand All @@ -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

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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

Expand Down
5 changes: 3 additions & 2 deletions src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
):
Expand Down
5 changes: 3 additions & 2 deletions src/diffusers/modular_pipelines/ideogram4/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Expand Down Expand Up @@ -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)
Expand Down
Loading
Loading