Skip to content
Open
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
9 changes: 6 additions & 3 deletions docs/source/en/optimization/cache.md
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,11 @@ Caching accelerates inference by storing and reusing intermediate outputs of dif

This guide shows you how to use the caching methods supported in Diffusers.

Pyramid Attention Broadcast and FasterCache read the current timestep from the denoiser's `cache_context`.
When writing a custom denoising loop, wrap each denoiser call with `model.cache_context("cond", timestep=t)`.
Use a separate context name for each guidance branch, or `"cond_uncond"` for a combined batch, and call
`model._reset_stateful_cache()` before starting a new generation. The pipeline examples below handle this for you.

## Pyramid Attention Broadcast

[Pyramid Attention Broadcast (PAB)](https://huggingface.co/papers/2408.12588) is based on the observation that attention outputs aren't that different between successive timesteps of the generation process. The attention differences are smallest in the cross attention layers and are generally cached over a longer timestep range. This is followed by temporal attention and spatial attention layers.
Expand All @@ -36,7 +41,6 @@ pipeline.to("cuda") # or "mps", "xpu", "cpu"
config = PyramidAttentionBroadcastConfig(
spatial_attention_block_skip_range=2,
spatial_attention_timestep_skip_range=(100, 800),
current_timestep_callback=lambda: pipe.current_timestep,
)
pipeline.transformer.enable_cache(config)
```
Expand All @@ -53,13 +57,12 @@ Set up and pass a [`FasterCacheConfig`] to a pipeline's transformer to enable it
import torch
from diffusers import CogVideoXPipeline, FasterCacheConfig

pipe line= CogVideoXPipeline.from_pretrained("THUDM/CogVideoX-5b", dtype=torch.bfloat16)
pipeline = CogVideoXPipeline.from_pretrained("THUDM/CogVideoX-5b", dtype=torch.bfloat16)
pipeline.to("cuda") # or "mps", "xpu", "cpu"

config = FasterCacheConfig(
spatial_attention_block_skip_range=2,
spatial_attention_timestep_skip_range=(-1, 681),
current_timestep_callback=lambda: pipe.current_timestep,
attention_weight_callback=lambda _: 0.3,
unconditional_batch_skip_range=5,
unconditional_batch_timestep_skip_range=(-1, 781),
Expand Down
2 changes: 0 additions & 2 deletions docs/source/zh/optimization/cache.md
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,6 @@ pipeline.to("cuda")
config = PyramidAttentionBroadcastConfig(
spatial_attention_block_skip_range=2,
spatial_attention_timestep_skip_range=(100, 800),
current_timestep_callback=lambda: pipe.current_timestep,
)
pipeline.transformer.enable_cache(config)
```
Expand All @@ -57,7 +56,6 @@ pipeline.to("cuda")
config = FasterCacheConfig(
spatial_attention_block_skip_range=2,
spatial_attention_timestep_skip_range=(-1, 681),
current_timestep_callback=lambda: pipe.current_timestep,
attention_weight_callback=lambda _: 0.3,
unconditional_batch_skip_range=5,
unconditional_batch_timestep_skip_range=(-1, 781),
Expand Down
117 changes: 73 additions & 44 deletions src/diffusers/hooks/faster_cache.py

Large diffs are not rendered by default.

63 changes: 38 additions & 25 deletions src/diffusers/hooks/pyramid_attention_broadcast.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,14 +20,14 @@

from ..models.attention import AttentionModuleMixin
from ..models.attention_processor import Attention, MochiAttention
from ..utils import logging
from ..utils import deprecate, logging
from ._common import (
_ATTENTION_CLASSES,
_CROSS_TRANSFORMER_BLOCK_IDENTIFIERS,
_SPATIAL_TRANSFORMER_BLOCK_IDENTIFIERS,
_TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS,
)
from .hooks import HookRegistry, ModelHook
from .hooks import CacheContext, HookRegistry, ModelHook, StateManager


logger = logging.get_logger(__name__) # pylint: disable=invalid-name
Expand Down Expand Up @@ -83,12 +83,24 @@ class PyramidAttentionBroadcastConfig:
temporal_attention_block_identifiers: tuple[str, ...] = _TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS
cross_attention_block_identifiers: tuple[str, ...] = _CROSS_TRANSFORMER_BLOCK_IDENTIFIERS

current_timestep_callback: Callable[[], int] = None
current_timestep_callback: Callable[[], int] | None = None

# TODO(aryan): add PAB for MLP layers (very limited speedup from testing with original codebase
# so not added for now)

def __post_init__(self):
if self.current_timestep_callback is not None:
depr_message = (
"Passing `current_timestep_callback` to `PyramidAttentionBroadcastConfig` is deprecated. "
"Please use `cache_context(name, timestep=...)` to pass the current timestep instead. "
"See the caching documentation for guidance: https://huggingface.co/docs/diffusers/main/en/optimization/cache."
)
deprecate("current_timestep_callback", "0.45.0", depr_message)

def __repr__(self) -> str:
callback_repr = ""
if self.current_timestep_callback is not None:
callback_repr = f" current_timestep_callback={self.current_timestep_callback},\n"
return (
f"PyramidAttentionBroadcastConfig(\n"
f" spatial_attention_block_skip_range={self.spatial_attention_block_skip_range},\n"
Expand All @@ -100,7 +112,7 @@ def __repr__(self) -> str:
f" spatial_attention_block_identifiers={self.spatial_attention_block_identifiers},\n"
f" temporal_attention_block_identifiers={self.temporal_attention_block_identifiers},\n"
f" cross_attention_block_identifiers={self.cross_attention_block_identifiers},\n"
f" current_timestep_callback={self.current_timestep_callback}\n"
f"{callback_repr}"
")"
)

Expand Down Expand Up @@ -141,7 +153,10 @@ class PyramidAttentionBroadcastHook(ModelHook):
_is_stateful = True

def __init__(
self, timestep_skip_range: tuple[int, int], block_skip_range: int, current_timestep_callback: Callable[[], int]
self,
timestep_skip_range: tuple[int, int],
block_skip_range: int,
current_timestep_callback: Callable[[], int] | None = None,
) -> None:
super().__init__()

Expand All @@ -150,31 +165,37 @@ def __init__(
self.current_timestep_callback = current_timestep_callback

def initialize_hook(self, module):
self.state = PyramidAttentionBroadcastState()
self.state_manager = StateManager(PyramidAttentionBroadcastState)
return module

def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:
is_within_timestep_range = (
self.timestep_skip_range[0] < self.current_timestep_callback() < self.timestep_skip_range[1]
)
if self.state_manager._context is None and self.current_timestep_callback is not None:
self.state_manager.set_context(CacheContext(name="inference"))
timestep = self.state_manager.context.timestep
if timestep is None and self.current_timestep_callback is not None:
timestep = self.current_timestep_callback()
if timestep is None:
raise ValueError("Pyramid Attention Broadcast requires `cache_context(name, timestep=...)`.")
state = self.state_manager.get_state()
is_within_timestep_range = self.timestep_skip_range[0] < timestep < self.timestep_skip_range[1]
should_compute_attention = (
self.state.cache is None
or self.state.iteration == 0
state.cache is None
or state.iteration == 0
or not is_within_timestep_range
or self.state.iteration % self.block_skip_range == 0
or state.iteration % self.block_skip_range == 0
)

if should_compute_attention:
output = self.fn_ref.original_forward(*args, **kwargs)
else:
output = self.state.cache
output = state.cache

self.state.cache = output
self.state.iteration += 1
state.cache = output
state.iteration += 1
return output

def reset_state(self, module: torch.nn.Module) -> None:
self.state.reset()
self.state_manager.reset()
return module


Expand Down Expand Up @@ -207,16 +228,10 @@ def apply_pyramid_attention_broadcast(module: torch.nn.Module, config: PyramidAt
>>> config = PyramidAttentionBroadcastConfig(
... spatial_attention_block_skip_range=2,
... spatial_attention_timestep_skip_range=(100, 800),
... current_timestep_callback=lambda: pipe.current_timestep,
... )
>>> apply_pyramid_attention_broadcast(pipe.transformer, config)
```
"""
if config.current_timestep_callback is None:
raise ValueError(
"The `current_timestep_callback` function must be provided in the configuration to apply Pyramid Attention Broadcast."
)

if (
config.spatial_attention_block_skip_range is None
and config.temporal_attention_block_skip_range is None
Expand Down Expand Up @@ -291,7 +306,7 @@ def _apply_pyramid_attention_broadcast_hook(
module: Attention | MochiAttention,
timestep_skip_range: tuple[int, int],
block_skip_range: int,
current_timestep_callback: Callable[[], int],
current_timestep_callback: Callable[[], int] | None = None,
):
r"""
Apply [Pyramid Attention Broadcast](https://huggingface.co/papers/2408.12588) to a given torch.nn.Module.
Expand All @@ -306,8 +321,6 @@ def _apply_pyramid_attention_broadcast_hook(
The number of times a specific attention broadcast is skipped before computing the attention states to
re-use. If this is set to the value `N`, the attention computation will be skipped `N - 1` times (i.e., old
attention states will be reused) before computing the new attention states again.
current_timestep_callback (`Callable[[], int]`):
A callback function that returns the current inference timestep.
"""
registry = HookRegistry.check_if_exists_or_initialize(module)
hook = PyramidAttentionBroadcastHook(timestep_skip_range, block_skip_range, current_timestep_callback)
Expand Down
1 change: 0 additions & 1 deletion src/diffusers/models/cache_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,6 @@ def enable_cache(self, config) -> None:
>>> config = PyramidAttentionBroadcastConfig(
... spatial_attention_block_skip_range=2,
... spatial_attention_timestep_skip_range=(100, 800),
... current_timestep_callback=lambda: pipe.current_timestep,
... )
>>> pipe.transformer.enable_cache(config)
```
Expand Down
10 changes: 8 additions & 2 deletions src/diffusers/modular_pipelines/cosmos/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -233,6 +233,7 @@ def __call__(
}
with components.transformer.cache_context(
pass_name,
timestep=t,
step_index=i,
sigma=float(components.scheduler.sigmas[i]),
num_inference_steps=components.scheduler.num_inference_steps,
Expand Down Expand Up @@ -797,9 +798,11 @@ def intermediate_outputs(self) -> list[OutputParam]:
return [OutputParam("velocity", type_hint=torch.Tensor, description="Predicted (masked) transfer velocity.")]

@staticmethod
def _forward(components, static, vision_tokens, vision_timesteps, context_name, step, sigma, num_inference_steps):
def _forward(
components, static, vision_tokens, vision_timesteps, context_name, step, timestep, sigma, num_inference_steps
):
with components.transformer.cache_context(
context_name, step_index=step, sigma=sigma, num_inference_steps=num_inference_steps
context_name, step_index=step, timestep=timestep, sigma=sigma, num_inference_steps=num_inference_steps
):
preds_vision, _, _ = components.transformer(
input_ids=static["input_ids"],
Expand Down Expand Up @@ -847,6 +850,7 @@ def __call__(
block_state.vision_timesteps,
"cond",
step=i,
timestep=t,
sigma=float(components.scheduler.sigmas[i]),
num_inference_steps=components.scheduler.num_inference_steps,
)
Expand All @@ -860,6 +864,7 @@ def __call__(
block_state.vision_timesteps,
"cond_no_control",
step=i,
timestep=t,
sigma=float(components.scheduler.sigmas[i]),
num_inference_steps=components.scheduler.num_inference_steps,
)
Expand All @@ -873,6 +878,7 @@ def __call__(
block_state.vision_timesteps,
"uncond",
step=i,
timestep=t,
sigma=float(components.scheduler.sigmas[i]),
num_inference_steps=components.scheduler.num_inference_steps,
)
Expand Down
6 changes: 3 additions & 3 deletions src/diffusers/modular_pipelines/helios/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -443,7 +443,7 @@ def __call__(
cond_kwargs = {k: getattr(guider_state_batch, k) for k in guider_inputs.keys()}

context_name = getattr(guider_state_batch, components.guider._identifier_key)
with components.transformer.cache_context(context_name):
with components.transformer.cache_context(context_name, timestep=t):
guider_state_batch.noise_pred = components.transformer(
hidden_states=latent_model_input,
timestep=timestep,
Expand Down Expand Up @@ -635,7 +635,7 @@ def __call__(
cond_kwargs = {kk: getattr(guider_state_batch, kk) for kk in guider_inputs.keys()}

context_name = getattr(guider_state_batch, components.guider._identifier_key)
with components.transformer.cache_context(context_name):
with components.transformer.cache_context(context_name, timestep=t):
guider_state_batch.noise_pred = components.transformer(
hidden_states=latent_model_input,
timestep=timestep,
Expand Down Expand Up @@ -976,7 +976,7 @@ def __call__(
cond_kwargs = {k: getattr(guider_state_batch, k) for k in guider_inputs.keys()}

context_name = getattr(guider_state_batch, components.guider._identifier_key)
with components.transformer.cache_context(context_name):
with components.transformer.cache_context(context_name, timestep=t):
guider_state_batch.noise_pred = components.transformer(
hidden_states=latent_model_input,
timestep=timestep,
Expand Down
4 changes: 2 additions & 2 deletions src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -157,7 +157,7 @@ def __call__(
cond_kwargs = {input_name: getattr(guider_state_batch, input_name) for input_name in guider_inputs.keys()}

context_name = getattr(guider_state_batch, components.guider._identifier_key)
with components.transformer.cache_context(context_name):
with components.transformer.cache_context(context_name, timestep=t):
guider_state_batch.noise_pred = components.transformer(
hidden_states=block_state.latent_model_input,
image_embeds=block_state.image_embeds,
Expand Down Expand Up @@ -370,7 +370,7 @@ def __call__(
cond_kwargs = {input_name: getattr(guider_state_batch, input_name) for input_name in guider_inputs.keys()}

context_name = getattr(guider_state_batch, components.guider._identifier_key)
with components.transformer.cache_context(context_name):
with components.transformer.cache_context(context_name, timestep=t):
guider_state_batch.noise_pred = components.transformer(
hidden_states=block_state.latent_model_input,
image_embeds=block_state.image_embeds,
Expand Down
4 changes: 2 additions & 2 deletions src/diffusers/modular_pipelines/ltx/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -136,7 +136,7 @@ def __call__(
}

context_name = getattr(guider_state_batch, components.guider._identifier_key, None)
with components.transformer.cache_context(context_name):
with components.transformer.cache_context(context_name, timestep=t):
guider_state_batch.noise_pred = components.transformer(
hidden_states=block_state.latent_model_input,
timestep=t.expand(block_state.latent_model_input.shape[0]).to(block_state.dtype),
Expand Down Expand Up @@ -369,7 +369,7 @@ def __call__(
}

context_name = getattr(guider_state_batch, components.guider._identifier_key, None)
with components.transformer.cache_context(context_name):
with components.transformer.cache_context(context_name, timestep=t):
guider_state_batch.noise_pred = components.transformer(
hidden_states=block_state.latent_model_input,
timestep=block_state.timestep_adjusted,
Expand Down
2 changes: 1 addition & 1 deletion src/diffusers/modular_pipelines/ltx2/denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -413,7 +413,7 @@ def __call__(
cond_kwargs = {name: getattr(batch, name) for name in self._guider_input_fields}
cond_kwargs["spatio_temporal_guidance_blocks"] = batch.spatio_temporal_guidance_blocks
cond_kwargs["isolate_modalities"] = batch.isolate_modalities
with components.transformer.cache_context(getattr(batch, identifier_key)):
with components.transformer.cache_context(getattr(batch, identifier_key), timestep=t):
noise_pred_video, noise_pred_audio = components.transformer(
hidden_states=block_state.latent_model_input,
audio_hidden_states=block_state.audio_latent_model_input,
Expand Down
17 changes: 9 additions & 8 deletions src/diffusers/pipelines/allegro/pipeline_allegro.py
Original file line number Diff line number Diff line change
Expand Up @@ -883,14 +883,15 @@ def __call__(
timestep = t.expand(latent_model_input.shape[0])

# predict noise model_output
noise_pred = self.transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds,
encoder_attention_mask=prompt_attention_mask,
timestep=timestep,
image_rotary_emb=image_rotary_emb,
return_dict=False,
)[0]
with self.transformer.cache_context("cond_uncond", timestep=t):
noise_pred = self.transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds,
encoder_attention_mask=prompt_attention_mask,
timestep=timestep,
image_rotary_emb=image_rotary_emb,
return_dict=False,
)[0]

# perform guidance
if do_classifier_free_guidance:
Expand Down
2 changes: 1 addition & 1 deletion src/diffusers/pipelines/cogvideo/pipeline_cogvideox.py
Original file line number Diff line number Diff line change
Expand Up @@ -727,7 +727,7 @@ def __call__(
timestep = t.expand(latent_model_input.shape[0])

# predict noise model_output
with self.transformer.cache_context("cond_uncond"):
with self.transformer.cache_context("cond_uncond", timestep=t):
noise_pred = self.transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -793,7 +793,7 @@ def __call__(
timestep = t.expand(latent_model_input.shape[0])

# predict noise model_output
with self.transformer.cache_context("cond_uncond"):
with self.transformer.cache_context("cond_uncond", timestep=t):
noise_pred = self.transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -836,7 +836,7 @@ def __call__(
timestep = t.expand(latent_model_input.shape[0])

# predict noise model_output
with self.transformer.cache_context("cond_uncond"):
with self.transformer.cache_context("cond_uncond", timestep=t):
noise_pred = self.transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -808,7 +808,7 @@ def __call__(
timestep = t.expand(latent_model_input.shape[0])

# predict noise model_output
with self.transformer.cache_context("cond_uncond"):
with self.transformer.cache_context("cond_uncond", timestep=t):
noise_pred = self.transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds,
Expand Down
Loading
Loading