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
88 changes: 42 additions & 46 deletions src/diffusers/hooks/faster_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@
from ..models.modeling_outputs import Transformer2DModelOutput
from ..utils import logging
from ._common import _ATTENTION_CLASSES
from .hooks import HookRegistry, ModelHook
from .hooks import HookRegistry, ModelHook, StateManager


logger = logging.get_logger(__name__) # pylint: disable=invalid-name
Expand Down Expand Up @@ -160,8 +160,6 @@ class FasterCacheConfig:
tensor_format: str = "BCFHW"
is_guidance_distilled: bool = False

current_timestep_callback: Callable[[], int] = None

_unconditional_conditional_input_kwargs_identifiers: list[str] = _UNCOND_COND_INPUT_KWARGS_IDENTIFIERS

def __repr__(self) -> str:
Expand Down Expand Up @@ -227,7 +225,6 @@ def __init__(
tensor_format: str,
is_guidance_distilled: bool,
uncond_cond_input_kwargs_identifiers: list[str],
current_timestep_callback: Callable[[], int],
low_frequency_weight_callback: Callable[[torch.nn.Module], torch.Tensor],
high_frequency_weight_callback: Callable[[torch.nn.Module], torch.Tensor],
) -> None:
Expand All @@ -243,12 +240,11 @@ def __init__(
self.tensor_format = tensor_format
self.is_guidance_distilled = is_guidance_distilled

self.current_timestep_callback = current_timestep_callback
self.low_frequency_weight_callback = low_frequency_weight_callback
self.high_frequency_weight_callback = high_frequency_weight_callback

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

@staticmethod
Expand All @@ -259,6 +255,10 @@ def _get_cond_input(input: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
return cond

def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:
timestep = self.state_manager.context.timestep
if timestep is None:
raise ValueError("FasterCache requires `cache_context(name, timestep=...)`.")
state = self.state_manager.get_state()
# Split the unconditional and conditional inputs. We only want to infer the conditional branch if the
# requirements for skipping the unconditional branch are met as described in the paper.
# We skip the unconditional branch only if the following conditions are met:
Expand All @@ -270,13 +270,13 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:
# we compute the unconditional branch at least once every few iterations to ensure minimal quality loss.
is_within_timestep_range = (
self.unconditional_batch_timestep_skip_range[0]
< self.current_timestep_callback()
< timestep
< self.unconditional_batch_timestep_skip_range[1]
)
should_skip_uncond = (
self.state.iteration > 0
state.iteration > 0
and is_within_timestep_range
and self.state.iteration % self.unconditional_batch_skip_range != 0
and state.iteration % self.unconditional_batch_skip_range != 0
and not self.is_guidance_distilled
)

Expand All @@ -293,7 +293,7 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:
output = self.fn_ref.original_forward(*args, **kwargs)

if self.is_guidance_distilled:
self.state.iteration += 1
state.iteration += 1
return output

if torch.is_tensor(output):
Expand All @@ -304,12 +304,8 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:
batch_size = hidden_states.size(0)

if should_skip_uncond:
self.state.low_frequency_delta = self.state.low_frequency_delta * self.low_frequency_weight_callback(
module
)
self.state.high_frequency_delta = self.state.high_frequency_delta * self.high_frequency_weight_callback(
module
)
state.low_frequency_delta = state.low_frequency_delta * self.low_frequency_weight_callback(module)
state.high_frequency_delta = state.high_frequency_delta * self.high_frequency_weight_callback(module)

if self.tensor_format == "BCFHW":
hidden_states = hidden_states.permute(0, 2, 1, 3, 4)
Expand All @@ -319,8 +315,8 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:
low_freq_cond, high_freq_cond = _split_low_high_freq(hidden_states.float())

# Approximate/compute the unconditional branch outputs as described in Equation 9 and 10 of the paper
low_freq_uncond = self.state.low_frequency_delta + low_freq_cond
high_freq_uncond = self.state.high_frequency_delta + high_freq_cond
low_freq_uncond = state.low_frequency_delta + low_freq_cond
high_freq_uncond = state.high_frequency_delta + high_freq_cond
uncond_freq = low_freq_uncond + high_freq_uncond

uncond_states = torch.fft.ifftshift(uncond_freq)
Expand All @@ -347,10 +343,10 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:

low_freq_uncond, high_freq_uncond = _split_low_high_freq(uncond_states.float())
low_freq_cond, high_freq_cond = _split_low_high_freq(cond_states.float())
self.state.low_frequency_delta = low_freq_uncond - low_freq_cond
self.state.high_frequency_delta = high_freq_uncond - high_freq_cond
state.low_frequency_delta = low_freq_uncond - low_freq_cond
state.high_frequency_delta = high_freq_uncond - high_freq_cond

self.state.iteration += 1
state.iteration += 1
if torch.is_tensor(output):
output = hidden_states
elif isinstance(output, tuple):
Expand All @@ -361,7 +357,7 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:
return output

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


Expand All @@ -374,7 +370,6 @@ def __init__(
timestep_skip_range: tuple[int, int],
is_guidance_distilled: bool,
weight_callback: Callable[[torch.nn.Module], float],
current_timestep_callback: Callable[[], int],
) -> None:
super().__init__()

Expand All @@ -383,10 +378,9 @@ def __init__(
self.is_guidance_distilled = is_guidance_distilled

self.weight_callback = weight_callback
self.current_timestep_callback = current_timestep_callback

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

def _compute_approximated_attention_output(
Expand All @@ -405,35 +399,39 @@ def _compute_approximated_attention_output(
return t_output + (t_output - t_2_output) * weight

def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:
timestep = self.state_manager.context.timestep
if timestep is None:
raise ValueError("FasterCache requires `cache_context(name, timestep=...)`.")
state = self.state_manager.get_state()
batch_size = [
*[arg.size(0) for arg in args if torch.is_tensor(arg)],
*[v.size(0) for v in kwargs.values() if torch.is_tensor(v)],
][0]
if self.state.batch_size is None:
if state.batch_size is None:
# Will be updated on first forward pass through the denoiser
self.state.batch_size = batch_size
state.batch_size = batch_size

# If we have to skip due to the skip conditions, then let's skip as expected.
# But, we can't skip if the denoiser wants to infer both unconditional and conditional branches. This
# is because the expected output shapes of attention layer will not match if we only return values from
# the cache (which only caches conditional branch outputs). So, if state.batch_size (which is the true
# unconditional-conditional batch size) is same as the current batch size, we don't perform the layer
# skip. Otherwise, we conditionally skip the layer based on what state.skip_callback returns.
is_within_timestep_range = (
self.timestep_skip_range[0] < self.current_timestep_callback() < self.timestep_skip_range[1]
)
is_within_timestep_range = self.timestep_skip_range[0] < timestep < self.timestep_skip_range[1]
if not is_within_timestep_range:
should_skip_attention = False
else:
should_compute_attention = self.state.iteration > 0 and self.state.iteration % self.block_skip_range == 0
should_compute_attention = state.iteration > 0 and state.iteration % self.block_skip_range == 0
should_skip_attention = not should_compute_attention
if should_skip_attention:
should_skip_attention = self.is_guidance_distilled or self.state.batch_size != batch_size
should_skip_attention = state.cache is not None and (
self.is_guidance_distilled or state.batch_size != batch_size
)

if should_skip_attention:
logger.debug("FasterCache - Skipping attention and using approximation")
if torch.is_tensor(self.state.cache[-1]):
t_2_output, t_output = self.state.cache
if torch.is_tensor(state.cache[-1]):
t_2_output, t_output = state.cache
weight = self.weight_callback(module)
output = self._compute_approximated_attention_output(t_2_output, t_output, weight, batch_size)
else:
Expand All @@ -444,7 +442,7 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:
# The zip(*state.cache) operation will give us [(A_1, A_2, ...), (B_1, B_2, ...), (C_1, C_2, ...), ...] which
# allows us to compute the approximated attention output for each tensor in the cache.
output = ()
for t_2_output, t_output in zip(*self.state.cache):
for t_2_output, t_output in zip(*state.cache):
result = self._compute_approximated_attention_output(
t_2_output, t_output, self.weight_callback(module), batch_size
)
Expand All @@ -458,28 +456,28 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any:
# both cases.
if torch.is_tensor(output):
cache_output = output
if not self.is_guidance_distilled and cache_output.size(0) == self.state.batch_size:
if not self.is_guidance_distilled and cache_output.size(0) == state.batch_size:
# The output here can be both unconditional-conditional branch outputs or just conditional branch outputs.
# This is determined at the higher-level denoiser module. We only want to cache the conditional branch outputs.
cache_output = cache_output.chunk(2, dim=0)[1]
else:
# Cache all return values and perform the same operation as above
cache_output = ()
for out in output:
if not self.is_guidance_distilled and out.size(0) == self.state.batch_size:
if not self.is_guidance_distilled and out.size(0) == state.batch_size:
out = out.chunk(2, dim=0)[1]
cache_output += (out,)

if self.state.cache is None:
self.state.cache = [cache_output, cache_output]
if state.cache is None:
state.cache = [cache_output, cache_output]
else:
self.state.cache = [self.state.cache[-1], cache_output]
state.cache = [state.cache[-1], cache_output]

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

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


Expand Down Expand Up @@ -539,7 +537,7 @@ def apply_faster_cache(module: torch.nn.Module, config: FasterCacheConfig) -> No
def low_frequency_weight_callback(module: torch.nn.Module) -> float:
is_within_range = (
config.low_frequency_weight_update_timestep_range[0]
< config.current_timestep_callback()
< module._diffusers_hook.get_hook(_FASTER_CACHE_DENOISER_HOOK).state_manager.context.timestep
< config.low_frequency_weight_update_timestep_range[1]
)
return config.alpha_low_frequency if is_within_range else 1.0
Expand All @@ -554,7 +552,7 @@ def low_frequency_weight_callback(module: torch.nn.Module) -> float:
def high_frequency_weight_callback(module: torch.nn.Module) -> float:
is_within_range = (
config.high_frequency_weight_update_timestep_range[0]
< config.current_timestep_callback()
< module._diffusers_hook.get_hook(_FASTER_CACHE_DENOISER_HOOK).state_manager.context.timestep
< config.high_frequency_weight_update_timestep_range[1]
)
return config.alpha_high_frequency if is_within_range else 1.0
Expand All @@ -581,7 +579,6 @@ def _apply_faster_cache_on_denoiser(module: torch.nn.Module, config: FasterCache
config.tensor_format,
config.is_guidance_distilled,
config._unconditional_conditional_input_kwargs_identifiers,
config.current_timestep_callback,
config.low_frequency_weight_callback,
config.high_frequency_weight_callback,
)
Expand Down Expand Up @@ -627,7 +624,6 @@ def _apply_faster_cache_on_attention_class(name: str, module: AttentionModuleMix
timestep_skip_range,
config.is_guidance_distilled,
config.attention_weight_callback,
config.current_timestep_callback,
)
registry = HookRegistry.check_if_exists_or_initialize(module)
registry.register_hook(hook, _FASTER_CACHE_BLOCK_HOOK)
Expand Down
Loading
Loading