From 0a1f36769ccdd9a3db6adf27d5e54cc6ebd2dee0 Mon Sep 17 00:00:00 2001 From: sayakpaul Date: Tue, 29 Sep 2026 10:33:34 +0530 Subject: [PATCH] feat: implement cache_context for pab and fastercache. --- docs/source/en/optimization/cache.md | 9 +- docs/source/zh/optimization/cache.md | 2 - src/diffusers/hooks/faster_cache.py | 88 +++++++++---------- .../hooks/pyramid_attention_broadcast.py | 51 ++++------- src/diffusers/models/cache_utils.py | 1 - .../modular_pipelines/cosmos/denoise.py | 10 ++- .../modular_pipelines/helios/denoise.py | 6 +- .../hunyuan_video1_5/denoise.py | 4 +- .../modular_pipelines/ltx/denoise.py | 4 +- .../modular_pipelines/ltx2/denoise.py | 2 +- .../pipelines/allegro/pipeline_allegro.py | 17 ++-- .../pipelines/cogvideo/pipeline_cogvideox.py | 2 +- .../pipeline_cogvideox_fun_control.py | 2 +- .../pipeline_cogvideox_image2video.py | 2 +- .../pipeline_cogvideox_video2video.py | 2 +- .../pipelines/cogview4/pipeline_cogview4.py | 4 +- .../pipelines/cosmos/pipeline_cosmos3_omni.py | 4 +- src/diffusers/pipelines/flux/pipeline_flux.py | 4 +- .../pipelines/flux/pipeline_flux_kontext.py | 42 ++++----- .../flux/pipeline_flux_kontext_inpaint.py | 42 ++++----- .../pipelines/flux2/pipeline_flux2_klein.py | 4 +- .../flux2/pipeline_flux2_klein_inpaint.py | 4 +- .../pipelines/helios/pipeline_helios.py | 4 +- .../helios/pipeline_helios_pyramid.py | 4 +- .../hunyuan_image/pipeline_hunyuanimage.py | 2 +- .../pipeline_hunyuanimage_refiner.py | 2 +- .../pipeline_hunyuan_skyreels_image2video.py | 34 +++---- .../hunyuan_video/pipeline_hunyuan_video.py | 4 +- .../pipeline_hunyuan_video_framepack.py | 50 ++++++----- .../pipeline_hunyuan_video_image2video.py | 34 +++---- .../pipeline_hunyuan_video1_5.py | 2 +- .../pipeline_hunyuan_video1_5_image2video.py | 2 +- .../pipelines/latte/pipeline_latte.py | 15 ++-- .../longcat_image/pipeline_longcat_image.py | 4 +- .../pipeline_longcat_image_edit.py | 4 +- src/diffusers/pipelines/ltx/pipeline_ltx.py | 2 +- .../pipelines/ltx/pipeline_ltx_condition.py | 2 +- .../ltx/pipeline_ltx_i2v_long_multi_prompt.py | 2 +- .../pipelines/ltx/pipeline_ltx_image2video.py | 2 +- src/diffusers/pipelines/ltx2/pipeline_ltx2.py | 6 +- .../pipelines/ltx2/pipeline_ltx2_condition.py | 6 +- .../pipelines/ltx2/pipeline_ltx2_hdr_lora.py | 6 +- .../pipelines/ltx2/pipeline_ltx2_ic_lora.py | 6 +- .../ltx2/pipeline_ltx2_image2video.py | 6 +- .../pipelines/lucy/pipeline_lucy_edit.py | 4 +- .../pipelines/mochi/pipeline_mochi.py | 2 +- .../motif_video/pipeline_motif_video.py | 2 +- .../pipeline_motif_video_image2video.py | 2 +- .../ovis_image/pipeline_ovis_image.py | 4 +- .../pipelines/qwenimage/pipeline_qwenimage.py | 4 +- .../pipeline_qwenimage_controlnet.py | 4 +- .../pipeline_qwenimage_controlnet_inpaint.py | 4 +- .../qwenimage/pipeline_qwenimage_edit.py | 4 +- .../pipeline_qwenimage_edit_inpaint.py | 4 +- .../qwenimage/pipeline_qwenimage_edit_plus.py | 4 +- .../qwenimage/pipeline_qwenimage_img2img.py | 4 +- .../qwenimage/pipeline_qwenimage_inpaint.py | 4 +- .../qwenimage/pipeline_qwenimage_layered.py | 4 +- .../qwenimage21/pipeline_qwenimage21.py | 4 +- .../skyreels_v2/pipeline_skyreels_v2.py | 4 +- .../pipeline_skyreels_v2_diffusion_forcing.py | 4 +- ...eline_skyreels_v2_diffusion_forcing_i2v.py | 4 +- ...eline_skyreels_v2_diffusion_forcing_v2v.py | 4 +- .../skyreels_v2/pipeline_skyreels_v2_i2v.py | 4 +- src/diffusers/pipelines/wan/pipeline_wan.py | 1 + .../pipelines/wan/pipeline_wan_animate.py | 4 +- .../pipelines/wan/pipeline_wan_i2v.py | 4 +- .../pipelines/wan/pipeline_wan_vace.py | 4 +- tests/models/testing_utils/cache.py | 44 ++++------ tests/pipelines/test_pipelines_common.py | 64 ++++++-------- tests/pipelines/testing_utils/cache.py | 68 ++++++-------- 71 files changed, 369 insertions(+), 399 deletions(-) diff --git a/docs/source/en/optimization/cache.md b/docs/source/en/optimization/cache.md index 9f775ec3b88c..c16c73d442c4 100644 --- a/docs/source/en/optimization/cache.md +++ b/docs/source/en/optimization/cache.md @@ -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. @@ -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) ``` @@ -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), diff --git a/docs/source/zh/optimization/cache.md b/docs/source/zh/optimization/cache.md index 7bf3b3c4286b..6deac3c36747 100644 --- a/docs/source/zh/optimization/cache.md +++ b/docs/source/zh/optimization/cache.md @@ -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) ``` @@ -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), diff --git a/src/diffusers/hooks/faster_cache.py b/src/diffusers/hooks/faster_cache.py index 01544aa4b430..48bdb979e6fe 100644 --- a/src/diffusers/hooks/faster_cache.py +++ b/src/diffusers/hooks/faster_cache.py @@ -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 @@ -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: @@ -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: @@ -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 @@ -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: @@ -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 ) @@ -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): @@ -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) @@ -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) @@ -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): @@ -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 @@ -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__() @@ -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( @@ -405,13 +399,17 @@ 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 @@ -419,21 +417,21 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any: # 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: @@ -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 ) @@ -458,7 +456,7 @@ 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] @@ -466,20 +464,20 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any: # 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 @@ -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 @@ -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 @@ -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, ) @@ -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) diff --git a/src/diffusers/hooks/pyramid_attention_broadcast.py b/src/diffusers/hooks/pyramid_attention_broadcast.py index e7ed26b28778..28b6be665152 100644 --- a/src/diffusers/hooks/pyramid_attention_broadcast.py +++ b/src/diffusers/hooks/pyramid_attention_broadcast.py @@ -14,7 +14,7 @@ import re from dataclasses import dataclass -from typing import Any, Callable +from typing import Any import torch @@ -27,7 +27,7 @@ _SPATIAL_TRANSFORMER_BLOCK_IDENTIFIERS, _TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS, ) -from .hooks import HookRegistry, ModelHook +from .hooks import HookRegistry, ModelHook, StateManager logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -83,8 +83,6 @@ 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 - # TODO(aryan): add PAB for MLP layers (very limited speedup from testing with original codebase # so not added for now) @@ -100,7 +98,6 @@ 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" ")" ) @@ -140,41 +137,40 @@ class PyramidAttentionBroadcastHook(ModelHook): _is_stateful = True - def __init__( - self, timestep_skip_range: tuple[int, int], block_skip_range: int, current_timestep_callback: Callable[[], int] - ) -> None: + def __init__(self, timestep_skip_range: tuple[int, int], block_skip_range: int) -> None: super().__init__() self.timestep_skip_range = timestep_skip_range self.block_skip_range = block_skip_range - 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] - ) + timestep = self.state_manager.context.timestep + 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 @@ -207,16 +203,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 @@ -281,9 +271,7 @@ def _apply_pyramid_attention_broadcast_on_attention_class( return False logger.debug(f"Enabling Pyramid Attention Broadcast ({block_type}) in layer: {name}") - _apply_pyramid_attention_broadcast_hook( - module, timestep_skip_range, block_skip_range, config.current_timestep_callback - ) + _apply_pyramid_attention_broadcast_hook(module, timestep_skip_range, block_skip_range) return True @@ -291,7 +279,6 @@ def _apply_pyramid_attention_broadcast_hook( module: Attention | MochiAttention, timestep_skip_range: tuple[int, int], block_skip_range: int, - current_timestep_callback: Callable[[], int], ): r""" Apply [Pyramid Attention Broadcast](https://huggingface.co/papers/2408.12588) to a given torch.nn.Module. @@ -306,9 +293,7 @@ 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) + hook = PyramidAttentionBroadcastHook(timestep_skip_range, block_skip_range) registry.register_hook(hook, _PYRAMID_ATTENTION_BROADCAST_HOOK) diff --git a/src/diffusers/models/cache_utils.py b/src/diffusers/models/cache_utils.py index 886ab6032bd4..469e111a13be 100644 --- a/src/diffusers/models/cache_utils.py +++ b/src/diffusers/models/cache_utils.py @@ -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) ``` diff --git a/src/diffusers/modular_pipelines/cosmos/denoise.py b/src/diffusers/modular_pipelines/cosmos/denoise.py index 6a369357e96f..dc95ff095f6f 100644 --- a/src/diffusers/modular_pipelines/cosmos/denoise.py +++ b/src/diffusers/modular_pipelines/cosmos/denoise.py @@ -225,6 +225,7 @@ def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockSta } 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, @@ -777,9 +778,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"], @@ -825,6 +828,7 @@ def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockSta block_state.vision_timesteps, "cond", step=i, + timestep=t, sigma=float(components.scheduler.sigmas[i]), num_inference_steps=components.scheduler.num_inference_steps, ) @@ -838,6 +842,7 @@ def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockSta 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, ) @@ -851,6 +856,7 @@ def __call__(self, components: Cosmos3OmniModularPipeline, block_state: BlockSta block_state.vision_timesteps, "uncond", step=i, + timestep=t, sigma=float(components.scheduler.sigmas[i]), num_inference_steps=components.scheduler.num_inference_steps, ) diff --git a/src/diffusers/modular_pipelines/helios/denoise.py b/src/diffusers/modular_pipelines/helios/denoise.py index 5fcf01a73ffc..8d47655811df 100644 --- a/src/diffusers/modular_pipelines/helios/denoise.py +++ b/src/diffusers/modular_pipelines/helios/denoise.py @@ -431,7 +431,7 @@ def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k 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, @@ -621,7 +621,7 @@ def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k 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, @@ -956,7 +956,7 @@ def __call__(self, components: HeliosModularPipeline, block_state: BlockState, k 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, diff --git a/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py b/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py index 293fad57c93f..b160d63e60c1 100644 --- a/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py +++ b/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py @@ -155,7 +155,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, @@ -364,7 +364,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, diff --git a/src/diffusers/modular_pipelines/ltx/denoise.py b/src/diffusers/modular_pipelines/ltx/denoise.py index b3ed86b51679..dc135d1044bd 100644 --- a/src/diffusers/modular_pipelines/ltx/denoise.py +++ b/src/diffusers/modular_pipelines/ltx/denoise.py @@ -134,7 +134,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), @@ -361,7 +361,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, diff --git a/src/diffusers/modular_pipelines/ltx2/denoise.py b/src/diffusers/modular_pipelines/ltx2/denoise.py index b1c4657d4d04..254d0b97aeea 100644 --- a/src/diffusers/modular_pipelines/ltx2/denoise.py +++ b/src/diffusers/modular_pipelines/ltx2/denoise.py @@ -404,7 +404,7 @@ def __call__(self, components, block_state: BlockState, i: int, t: torch.Tensor) 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, diff --git a/src/diffusers/pipelines/allegro/pipeline_allegro.py b/src/diffusers/pipelines/allegro/pipeline_allegro.py index 9d2d2aa8bd09..31964553d9a8 100644 --- a/src/diffusers/pipelines/allegro/pipeline_allegro.py +++ b/src/diffusers/pipelines/allegro/pipeline_allegro.py @@ -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: diff --git a/src/diffusers/pipelines/cogvideo/pipeline_cogvideox.py b/src/diffusers/pipelines/cogvideo/pipeline_cogvideox.py index 9043abcab65e..2767b3c104ef 100644 --- a/src/diffusers/pipelines/cogvideo/pipeline_cogvideox.py +++ b/src/diffusers/pipelines/cogvideo/pipeline_cogvideox.py @@ -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, diff --git a/src/diffusers/pipelines/cogvideo/pipeline_cogvideox_fun_control.py b/src/diffusers/pipelines/cogvideo/pipeline_cogvideox_fun_control.py index e2b45a08ee90..4db8e16e6d03 100644 --- a/src/diffusers/pipelines/cogvideo/pipeline_cogvideox_fun_control.py +++ b/src/diffusers/pipelines/cogvideo/pipeline_cogvideox_fun_control.py @@ -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, diff --git a/src/diffusers/pipelines/cogvideo/pipeline_cogvideox_image2video.py b/src/diffusers/pipelines/cogvideo/pipeline_cogvideox_image2video.py index 42f5109bb877..3481b082bb92 100644 --- a/src/diffusers/pipelines/cogvideo/pipeline_cogvideox_image2video.py +++ b/src/diffusers/pipelines/cogvideo/pipeline_cogvideox_image2video.py @@ -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, diff --git a/src/diffusers/pipelines/cogvideo/pipeline_cogvideox_video2video.py b/src/diffusers/pipelines/cogvideo/pipeline_cogvideox_video2video.py index 3cd72b0c2126..2e02005d48fc 100644 --- a/src/diffusers/pipelines/cogvideo/pipeline_cogvideox_video2video.py +++ b/src/diffusers/pipelines/cogvideo/pipeline_cogvideox_video2video.py @@ -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, diff --git a/src/diffusers/pipelines/cogview4/pipeline_cogview4.py b/src/diffusers/pipelines/cogview4/pipeline_cogview4.py index 329b76d11e0d..ef89bfc95f23 100644 --- a/src/diffusers/pipelines/cogview4/pipeline_cogview4.py +++ b/src/diffusers/pipelines/cogview4/pipeline_cogview4.py @@ -621,7 +621,7 @@ def __call__( # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latents.shape[0]) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred_cond = self.transformer( hidden_states=latent_model_input, encoder_hidden_states=prompt_embeds, @@ -635,7 +635,7 @@ def __call__( # perform guidance if self.do_classifier_free_guidance: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): noise_pred_uncond = self.transformer( hidden_states=latent_model_input, encoder_hidden_states=negative_prompt_embeds, diff --git a/src/diffusers/pipelines/cosmos/pipeline_cosmos3_omni.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos3_omni.py index 02dc70b29cfc..7c2354cade72 100644 --- a/src/diffusers/pipelines/cosmos/pipeline_cosmos3_omni.py +++ b/src/diffusers/pipelines/cosmos/pipeline_cosmos3_omni.py @@ -1751,7 +1751,7 @@ def __call__( # --- Conditional pass --- with self.transformer.cache_context( - "cond", step_index=i, sigma=sigma, num_inference_steps=self._num_timesteps + "cond", step_index=i, timestep=t, sigma=sigma, num_inference_steps=self._num_timesteps ): preds_vision, preds_sound, preds_action = self.transformer( input_ids=cond_packed_static["input_ids"], @@ -1794,7 +1794,7 @@ def __call__( uncond_v_vision = uncond_v_sound = uncond_v_action = None if self.do_classifier_free_guidance: with self.transformer.cache_context( - "uncond", step_index=i, sigma=sigma, num_inference_steps=self._num_timesteps + "uncond", step_index=i, timestep=t, sigma=sigma, num_inference_steps=self._num_timesteps ): preds_vision, preds_sound, preds_action = self.transformer( input_ids=uncond_packed_static["input_ids"], diff --git a/src/diffusers/pipelines/flux/pipeline_flux.py b/src/diffusers/pipelines/flux/pipeline_flux.py index eb831a7975ba..771913cbc711 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux.py +++ b/src/diffusers/pipelines/flux/pipeline_flux.py @@ -897,7 +897,7 @@ def __call__( # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latents.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, @@ -914,7 +914,7 @@ def __call__( if negative_image_embeds is not None: self._joint_attention_kwargs["ip_adapter_image_embeds"] = negative_image_embeds - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/flux/pipeline_flux_kontext.py b/src/diffusers/pipelines/flux/pipeline_flux_kontext.py index e5fc95e5a1c1..3c93b2529dd2 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_kontext.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_kontext.py @@ -1038,33 +1038,35 @@ def __call__( latent_model_input = torch.cat([latents, image_latents], dim=1) timestep = t.expand(latents.shape[0]).to(latents.dtype) - noise_pred = self.transformer( - hidden_states=latent_model_input, - timestep=timestep / 1000, - guidance=guidance, - pooled_projections=pooled_prompt_embeds, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_ids, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - )[0] - noise_pred = noise_pred[:, : latents.size(1)] - - if do_true_cfg: - if negative_image_embeds is not None: - self._joint_attention_kwargs["ip_adapter_image_embeds"] = negative_image_embeds - neg_noise_pred = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, guidance=guidance, - pooled_projections=negative_pooled_prompt_embeds, - encoder_hidden_states=negative_prompt_embeds, - txt_ids=negative_text_ids, + pooled_projections=pooled_prompt_embeds, + encoder_hidden_states=prompt_embeds, + txt_ids=text_ids, img_ids=latent_ids, joint_attention_kwargs=self.joint_attention_kwargs, return_dict=False, )[0] + noise_pred = noise_pred[:, : latents.size(1)] + + if do_true_cfg: + if negative_image_embeds is not None: + self._joint_attention_kwargs["ip_adapter_image_embeds"] = negative_image_embeds + with self.transformer.cache_context("uncond", timestep=t): + neg_noise_pred = self.transformer( + hidden_states=latent_model_input, + timestep=timestep / 1000, + guidance=guidance, + pooled_projections=negative_pooled_prompt_embeds, + encoder_hidden_states=negative_prompt_embeds, + txt_ids=negative_text_ids, + img_ids=latent_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + )[0] neg_noise_pred = neg_noise_pred[:, : latents.size(1)] noise_pred = neg_noise_pred + true_cfg_scale * (noise_pred - neg_noise_pred) diff --git a/src/diffusers/pipelines/flux/pipeline_flux_kontext_inpaint.py b/src/diffusers/pipelines/flux/pipeline_flux_kontext_inpaint.py index 020f9761e121..e57b3fc24fb9 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_kontext_inpaint.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_kontext_inpaint.py @@ -1346,33 +1346,35 @@ def __call__( latent_model_input = torch.cat([latents, image_latents], dim=1) timestep = t.expand(latents.shape[0]).to(latents.dtype) - noise_pred = self.transformer( - hidden_states=latent_model_input, - timestep=timestep / 1000, - guidance=guidance, - pooled_projections=pooled_prompt_embeds, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_ids, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - )[0] - noise_pred = noise_pred[:, : latents.size(1)] - - if do_true_cfg: - if negative_image_embeds is not None: - self._joint_attention_kwargs["ip_adapter_image_embeds"] = negative_image_embeds - neg_noise_pred = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, guidance=guidance, - pooled_projections=negative_pooled_prompt_embeds, - encoder_hidden_states=negative_prompt_embeds, - txt_ids=negative_text_ids, + pooled_projections=pooled_prompt_embeds, + encoder_hidden_states=prompt_embeds, + txt_ids=text_ids, img_ids=latent_ids, joint_attention_kwargs=self.joint_attention_kwargs, return_dict=False, )[0] + noise_pred = noise_pred[:, : latents.size(1)] + + if do_true_cfg: + if negative_image_embeds is not None: + self._joint_attention_kwargs["ip_adapter_image_embeds"] = negative_image_embeds + with self.transformer.cache_context("uncond", timestep=t): + neg_noise_pred = self.transformer( + hidden_states=latent_model_input, + timestep=timestep / 1000, + guidance=guidance, + pooled_projections=negative_pooled_prompt_embeds, + encoder_hidden_states=negative_prompt_embeds, + txt_ids=negative_text_ids, + img_ids=latent_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + )[0] neg_noise_pred = neg_noise_pred[:, : latents.size(1)] noise_pred = neg_noise_pred + true_cfg_scale * (noise_pred - neg_noise_pred) diff --git a/src/diffusers/pipelines/flux2/pipeline_flux2_klein.py b/src/diffusers/pipelines/flux2/pipeline_flux2_klein.py index 92005750e551..a1d818ce69d9 100644 --- a/src/diffusers/pipelines/flux2/pipeline_flux2_klein.py +++ b/src/diffusers/pipelines/flux2/pipeline_flux2_klein.py @@ -847,7 +847,7 @@ def __call__( latent_model_input = torch.cat([latents, image_latents], dim=1).to(self.transformer.dtype) latent_image_ids = torch.cat([latent_ids, image_latent_ids], dim=1) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, # (B, image_seq_len, C) timestep=timestep / 1000, @@ -862,7 +862,7 @@ def __call__( noise_pred = noise_pred[:, : latents.size(1) :] if self.do_classifier_free_guidance: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/flux2/pipeline_flux2_klein_inpaint.py b/src/diffusers/pipelines/flux2/pipeline_flux2_klein_inpaint.py index 0f9051a99b12..a267a0873f1b 100644 --- a/src/diffusers/pipelines/flux2/pipeline_flux2_klein_inpaint.py +++ b/src/diffusers/pipelines/flux2/pipeline_flux2_klein_inpaint.py @@ -1180,7 +1180,7 @@ def __call__( latent_model_input = latent_model_input.to(self.transformer.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, # (B, image_seq_len, C) timestep=timestep / 1000, @@ -1194,7 +1194,7 @@ def __call__( noise_pred = noise_pred[:, : latents.size(1)] if self.do_classifier_free_guidance: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/helios/pipeline_helios.py b/src/diffusers/pipelines/helios/pipeline_helios.py index 90ac654bc77c..98c0acc2ace7 100644 --- a/src/diffusers/pipelines/helios/pipeline_helios.py +++ b/src/diffusers/pipelines/helios/pipeline_helios.py @@ -853,7 +853,7 @@ def __call__( latents_history_short = latents_history_short.to(transformer_dtype) latents_history_mid = latents_history_mid.to(transformer_dtype) latents_history_long = latents_history_long.to(transformer_dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep, @@ -870,7 +870,7 @@ def __call__( )[0] if self.do_classifier_free_guidance: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): noise_uncond = self.transformer( hidden_states=latent_model_input, timestep=timestep, diff --git a/src/diffusers/pipelines/helios/pipeline_helios_pyramid.py b/src/diffusers/pipelines/helios/pipeline_helios_pyramid.py index c187e436a857..17940280a601 100644 --- a/src/diffusers/pipelines/helios/pipeline_helios_pyramid.py +++ b/src/diffusers/pipelines/helios/pipeline_helios_pyramid.py @@ -997,7 +997,7 @@ def __call__( latents_history_short = latents_history_short.to(transformer_dtype) latents_history_mid = latents_history_mid.to(transformer_dtype) latents_history_long = latents_history_long.to(transformer_dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep, @@ -1014,7 +1014,7 @@ def __call__( )[0] if self.do_classifier_free_guidance: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): noise_uncond = self.transformer( hidden_states=latent_model_input, timestep=timestep, diff --git a/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage.py b/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage.py index 50239e9afa22..7d3b729215ff 100644 --- a/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage.py +++ b/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage.py @@ -796,7 +796,7 @@ def __call__( # e.g. "pred_cond"/"pred_uncond" context_name = getattr(guider_state_batch, guider._identifier_key) - with self.transformer.cache_context(context_name): + with self.transformer.cache_context(context_name, timestep=t): # Run denoiser and store noise prediction in this batch guider_state_batch.noise_pred = self.transformer( hidden_states=latents, diff --git a/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage_refiner.py b/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage_refiner.py index efdb5505e604..d784fb9f8ecf 100644 --- a/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage_refiner.py +++ b/src/diffusers/pipelines/hunyuan_image/pipeline_hunyuanimage_refiner.py @@ -687,7 +687,7 @@ def __call__( # e.g. "pred_cond"/"pred_uncond" context_name = getattr(guider_state_batch, guider._identifier_key) - with self.transformer.cache_context(context_name): + with self.transformer.cache_context(context_name, timestep=t): # Run denoiser and store noise prediction in this batch guider_state_batch.noise_pred = self.transformer( hidden_states=latent_model_input, diff --git a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_skyreels_image2video.py b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_skyreels_image2video.py index bd54f2563b52..7cdda5e4ebfa 100644 --- a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_skyreels_image2video.py +++ b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_skyreels_image2video.py @@ -725,28 +725,30 @@ def __call__( # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latents.shape[0]).to(latents.dtype) - noise_pred = self.transformer( - hidden_states=latent_model_input, - timestep=timestep, - encoder_hidden_states=prompt_embeds, - encoder_attention_mask=prompt_attention_mask, - pooled_projections=pooled_prompt_embeds, - guidance=guidance, - attention_kwargs=attention_kwargs, - return_dict=False, - )[0] - - if do_true_cfg: - neg_noise_pred = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep, - encoder_hidden_states=negative_prompt_embeds, - encoder_attention_mask=negative_prompt_attention_mask, - pooled_projections=negative_pooled_prompt_embeds, + encoder_hidden_states=prompt_embeds, + encoder_attention_mask=prompt_attention_mask, + pooled_projections=pooled_prompt_embeds, guidance=guidance, attention_kwargs=attention_kwargs, return_dict=False, )[0] + + if do_true_cfg: + with self.transformer.cache_context("uncond", timestep=t): + neg_noise_pred = self.transformer( + hidden_states=latent_model_input, + timestep=timestep, + encoder_hidden_states=negative_prompt_embeds, + encoder_attention_mask=negative_prompt_attention_mask, + pooled_projections=negative_pooled_prompt_embeds, + guidance=guidance, + attention_kwargs=attention_kwargs, + return_dict=False, + )[0] noise_pred = neg_noise_pred + true_cfg_scale * (noise_pred - neg_noise_pred) # compute the previous noisy sample x_t -> x_t-1 diff --git a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video.py b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video.py index 9e7c198c19cc..ba09b1886526 100644 --- a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video.py +++ b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video.py @@ -675,7 +675,7 @@ def __call__( # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latents.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep, @@ -688,7 +688,7 @@ def __call__( )[0] if do_true_cfg: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep, diff --git a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_framepack.py b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_framepack.py index 349481492ac0..cf024cf9801e 100644 --- a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_framepack.py +++ b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_framepack.py @@ -953,32 +953,13 @@ def __call__( self._current_timestep = t timestep = t.expand(latents.shape[0]) - noise_pred = self.transformer( - hidden_states=latents.to(transformer_dtype), - timestep=timestep, - encoder_hidden_states=prompt_embeds, - encoder_attention_mask=prompt_attention_mask, - pooled_projections=pooled_prompt_embeds, - image_embeds=image_embeds, - indices_latents=indices_latents, - guidance=guidance, - latents_clean=latents_clean.to(transformer_dtype), - indices_latents_clean=indices_clean_latents, - latents_history_2x=latents_history_2x.to(transformer_dtype), - indices_latents_history_2x=indices_latents_history_2x, - latents_history_4x=latents_history_4x.to(transformer_dtype), - indices_latents_history_4x=indices_latents_history_4x, - attention_kwargs=attention_kwargs, - return_dict=False, - )[0] - - if do_true_cfg: - neg_noise_pred = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + noise_pred = self.transformer( hidden_states=latents.to(transformer_dtype), timestep=timestep, - encoder_hidden_states=negative_prompt_embeds, - encoder_attention_mask=negative_prompt_attention_mask, - pooled_projections=negative_pooled_prompt_embeds, + encoder_hidden_states=prompt_embeds, + encoder_attention_mask=prompt_attention_mask, + pooled_projections=pooled_prompt_embeds, image_embeds=image_embeds, indices_latents=indices_latents, guidance=guidance, @@ -991,6 +972,27 @@ def __call__( attention_kwargs=attention_kwargs, return_dict=False, )[0] + + if do_true_cfg: + with self.transformer.cache_context("uncond", timestep=t): + neg_noise_pred = self.transformer( + hidden_states=latents.to(transformer_dtype), + timestep=timestep, + encoder_hidden_states=negative_prompt_embeds, + encoder_attention_mask=negative_prompt_attention_mask, + pooled_projections=negative_pooled_prompt_embeds, + image_embeds=image_embeds, + indices_latents=indices_latents, + guidance=guidance, + latents_clean=latents_clean.to(transformer_dtype), + indices_latents_clean=indices_clean_latents, + latents_history_2x=latents_history_2x.to(transformer_dtype), + indices_latents_history_2x=indices_latents_history_2x, + latents_history_4x=latents_history_4x.to(transformer_dtype), + indices_latents_history_4x=indices_latents_history_4x, + attention_kwargs=attention_kwargs, + return_dict=False, + )[0] noise_pred = neg_noise_pred + true_cfg_scale * (noise_pred - neg_noise_pred) # compute the previous noisy sample x_t -> x_t-1 diff --git a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_image2video.py b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_image2video.py index 13eb35386001..0c6a4633580f 100644 --- a/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_image2video.py +++ b/src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_image2video.py @@ -892,28 +892,30 @@ def __call__( elif image_condition_type == "token_replace": latent_model_input = torch.cat([image_latents, latents[:, :, 1:]], dim=2).to(transformer_dtype) - noise_pred = self.transformer( - hidden_states=latent_model_input, - timestep=timestep, - encoder_hidden_states=prompt_embeds, - encoder_attention_mask=prompt_attention_mask, - pooled_projections=pooled_prompt_embeds, - guidance=guidance, - attention_kwargs=attention_kwargs, - return_dict=False, - )[0] - - if do_true_cfg: - neg_noise_pred = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep, - encoder_hidden_states=negative_prompt_embeds, - encoder_attention_mask=negative_prompt_attention_mask, - pooled_projections=negative_pooled_prompt_embeds, + encoder_hidden_states=prompt_embeds, + encoder_attention_mask=prompt_attention_mask, + pooled_projections=pooled_prompt_embeds, guidance=guidance, attention_kwargs=attention_kwargs, return_dict=False, )[0] + + if do_true_cfg: + with self.transformer.cache_context("uncond", timestep=t): + neg_noise_pred = self.transformer( + hidden_states=latent_model_input, + timestep=timestep, + encoder_hidden_states=negative_prompt_embeds, + encoder_attention_mask=negative_prompt_attention_mask, + pooled_projections=negative_pooled_prompt_embeds, + guidance=guidance, + attention_kwargs=attention_kwargs, + return_dict=False, + )[0] noise_pred = neg_noise_pred + true_cfg_scale * (noise_pred - neg_noise_pred) # compute the previous noisy sample x_t -> x_t-1 diff --git a/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5.py b/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5.py index 7232ebbee5b8..14c7c00876b9 100644 --- a/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5.py +++ b/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5.py @@ -773,7 +773,7 @@ def __call__( # e.g. "pred_cond"/"pred_uncond" context_name = getattr(guider_state_batch, self.guider._identifier_key) - with self.transformer.cache_context(context_name): + with self.transformer.cache_context(context_name, timestep=t): # Run denoiser and store noise prediction in this batch guider_state_batch.noise_pred = self.transformer( hidden_states=latent_model_input, diff --git a/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5_image2video.py b/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5_image2video.py index 71a36a1c51cd..0b6201b7a00c 100644 --- a/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5_image2video.py +++ b/src/diffusers/pipelines/hunyuan_video1_5/pipeline_hunyuan_video1_5_image2video.py @@ -896,7 +896,7 @@ def __call__( # e.g. "pred_cond"/"pred_uncond" context_name = getattr(guider_state_batch, self.guider._identifier_key) - with self.transformer.cache_context(context_name): + with self.transformer.cache_context(context_name, timestep=t): # Run denoiser and store noise prediction in this batch guider_state_batch.noise_pred = self.transformer( hidden_states=latent_model_input, diff --git a/src/diffusers/pipelines/latte/pipeline_latte.py b/src/diffusers/pipelines/latte/pipeline_latte.py index 7bc7b4aa915e..12d3dc43fe41 100644 --- a/src/diffusers/pipelines/latte/pipeline_latte.py +++ b/src/diffusers/pipelines/latte/pipeline_latte.py @@ -820,13 +820,14 @@ def __call__( current_timestep = current_timestep.expand(latent_model_input.shape[0]) # predict noise model_output - noise_pred = self.transformer( - hidden_states=latent_model_input, - encoder_hidden_states=prompt_embeds, - timestep=current_timestep, - enable_temporal_attentions=enable_temporal_attentions, - 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, + timestep=current_timestep, + enable_temporal_attentions=enable_temporal_attentions, + return_dict=False, + )[0] # perform guidance if do_classifier_free_guidance: diff --git a/src/diffusers/pipelines/longcat_image/pipeline_longcat_image.py b/src/diffusers/pipelines/longcat_image/pipeline_longcat_image.py index 41ca3eb54f83..5fa5e3ef36c2 100644 --- a/src/diffusers/pipelines/longcat_image/pipeline_longcat_image.py +++ b/src/diffusers/pipelines/longcat_image/pipeline_longcat_image.py @@ -632,7 +632,7 @@ def __call__( self._current_timestep = t timestep = t.expand(latents.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred_text = self.transformer( hidden_states=latents, timestep=timestep / 1000, @@ -643,7 +643,7 @@ def __call__( return_dict=False, )[0] if self.do_classifier_free_guidance: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): noise_pred_uncond = self.transformer( hidden_states=latents, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/longcat_image/pipeline_longcat_image_edit.py b/src/diffusers/pipelines/longcat_image/pipeline_longcat_image_edit.py index 9f35bb685d9f..95aa365369b6 100644 --- a/src/diffusers/pipelines/longcat_image/pipeline_longcat_image_edit.py +++ b/src/diffusers/pipelines/longcat_image/pipeline_longcat_image_edit.py @@ -694,7 +694,7 @@ def __call__( latent_model_input = torch.cat([latents, image_latents], dim=1) timestep = t.expand(latent_model_input.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred_text = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, @@ -706,7 +706,7 @@ def __call__( )[0] noise_pred_text = noise_pred_text[:, :image_seq_len] if self.do_classifier_free_guidance: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): noise_pred_uncond = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx.py b/src/diffusers/pipelines/ltx/pipeline_ltx.py index ce9177547c52..49cf9f3f4f2e 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx.py @@ -767,7 +767,7 @@ def __call__( # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latent_model_input.shape[0]) - 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, diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx_condition.py b/src/diffusers/pipelines/ltx/pipeline_ltx_condition.py index 28d296695998..26b5fb771d06 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx_condition.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx_condition.py @@ -1198,7 +1198,7 @@ def __call__( if is_conditioning_image_or_video: timestep = torch.min(timestep, (1 - conditioning_mask_model_input) * 1000.0) - 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, diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx_i2v_long_multi_prompt.py b/src/diffusers/pipelines/ltx/pipeline_ltx_i2v_long_multi_prompt.py index 838d5afc5c5a..cf5fd877e9ef 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx_i2v_long_multi_prompt.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx_i2v_long_multi_prompt.py @@ -1312,7 +1312,7 @@ def __call__( rope_interpolation_scale=rope_interpolation_scale, frame_rate=frame_rate, ) - 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.to(dtype=self.transformer.dtype), encoder_hidden_states=prompt_embeds, diff --git a/src/diffusers/pipelines/ltx/pipeline_ltx_image2video.py b/src/diffusers/pipelines/ltx/pipeline_ltx_image2video.py index 81ecfce50efa..d5f4649ac9e8 100644 --- a/src/diffusers/pipelines/ltx/pipeline_ltx_image2video.py +++ b/src/diffusers/pipelines/ltx/pipeline_ltx_image2video.py @@ -840,7 +840,7 @@ def __call__( timestep = t.expand(latent_model_input.shape[0]) timestep = timestep.unsqueeze(-1) * (1 - conditioning_mask) - 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, diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2.py index 22948a7ecf3a..63f3e032e432 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2.py @@ -1395,7 +1395,7 @@ def __call__( # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latent_model_input.shape[0]) - with self.transformer.cache_context("cond_uncond"): + with self.transformer.cache_context("cond_uncond", timestep=t): noise_pred_video, noise_pred_audio = self.transformer( hidden_states=latent_model_input, audio_hidden_states=audio_latent_model_input, @@ -1467,7 +1467,7 @@ def __call__( noise_pred_audio = self.convert_velocity_to_x0(audio_latents, noise_pred_audio, i, audio_scheduler) if self.do_spatio_temporal_guidance: - with self.transformer.cache_context("uncond_stg"): + with self.transformer.cache_context("uncond_stg", timestep=t): noise_pred_video_uncond_stg, noise_pred_audio_uncond_stg = self.transformer( hidden_states=latents.to(dtype=prompt_embeds.dtype), audio_hidden_states=audio_latents.to(dtype=prompt_embeds.dtype), @@ -1507,7 +1507,7 @@ def __call__( video_stg_delta = audio_stg_delta = 0 if self.do_modality_isolation_guidance: - with self.transformer.cache_context("uncond_modality"): + with self.transformer.cache_context("uncond_modality", timestep=t): noise_pred_video_uncond_modality, noise_pred_audio_uncond_modality = self.transformer( hidden_states=latents.to(dtype=prompt_embeds.dtype), audio_hidden_states=audio_latents.to(dtype=prompt_embeds.dtype), diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_condition.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_condition.py index bd2ee3ec6708..104ffbdd1611 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_condition.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_condition.py @@ -1825,7 +1825,7 @@ def __call__( t_audio = audio_timesteps[i] audio_timestep = t_audio.expand(latent_model_input.shape[0]) - with self.transformer.cache_context("cond_uncond"): + with self.transformer.cache_context("cond_uncond", timestep=t): noise_pred_video, noise_pred_audio = self.transformer( hidden_states=latent_model_input, audio_hidden_states=audio_latent_model_input, @@ -1901,7 +1901,7 @@ def __call__( noise_pred_audio = self.convert_velocity_to_x0(audio_latents, noise_pred_audio, i, audio_scheduler) if self.do_spatio_temporal_guidance: - with self.transformer.cache_context("uncond_stg"): + with self.transformer.cache_context("uncond_stg", timestep=t): noise_pred_video_uncond_stg, noise_pred_audio_uncond_stg = self.transformer( hidden_states=latents.to(dtype=prompt_embeds.dtype), audio_hidden_states=audio_latents.to(dtype=prompt_embeds.dtype), @@ -1943,7 +1943,7 @@ def __call__( video_stg_delta = audio_stg_delta = 0 if self.do_modality_isolation_guidance: - with self.transformer.cache_context("uncond_modality"): + with self.transformer.cache_context("uncond_modality", timestep=t): noise_pred_video_uncond_modality, noise_pred_audio_uncond_modality = self.transformer( hidden_states=latents.to(dtype=prompt_embeds.dtype), audio_hidden_states=audio_latents.to(dtype=prompt_embeds.dtype), diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py index 91173bc6e161..f3d5966ea170 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py @@ -1375,7 +1375,7 @@ def __call__( audio_timestep = t_audio.expand(latent_model_input.shape[0]) # --- Main forward pass (cond + uncond for CFG) --- - with self.transformer.cache_context("cond_uncond"): + with self.transformer.cache_context("cond_uncond", timestep=t): noise_pred_video, noise_pred_audio = self.transformer( hidden_states=latent_model_input, audio_hidden_states=audio_latent_model_input, @@ -1445,7 +1445,7 @@ def __call__( # --- STG forward pass (video only — audio output discarded) --- if self.do_spatio_temporal_guidance: - with self.transformer.cache_context("uncond_stg"): + with self.transformer.cache_context("uncond_stg", timestep=t): noise_pred_video_uncond_stg, noise_pred_audio_uncond_stg = self.transformer( hidden_states=latents.to(dtype=connector_prompt_embeds.dtype), audio_hidden_states=audio_latents.to(dtype=connector_prompt_embeds.dtype), @@ -1483,7 +1483,7 @@ def __call__( # --- Modality isolation guidance forward pass --- if self.do_modality_isolation_guidance: - with self.transformer.cache_context("uncond_modality"): + with self.transformer.cache_context("uncond_modality", timestep=t): noise_pred_video_uncond_mod, noise_pred_audio_uncond_mod = self.transformer( hidden_states=latents.to(dtype=connector_prompt_embeds.dtype), audio_hidden_states=audio_latents.to(dtype=connector_prompt_embeds.dtype), diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_ic_lora.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_ic_lora.py index dc92b6eb965a..f5276ecd08cd 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_ic_lora.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_ic_lora.py @@ -2242,7 +2242,7 @@ def __call__( if video_self_attention_mask is not None else None ) - with self.transformer.cache_context("cond_uncond"): + with self.transformer.cache_context("cond_uncond", timestep=t): noise_pred_video, noise_pred_audio = self.transformer( hidden_states=latent_model_input, audio_hidden_states=audio_latent_model_input, @@ -2331,7 +2331,7 @@ def __call__( if video_self_attention_mask is not None else None ) - with self.transformer.cache_context("uncond_stg"): + with self.transformer.cache_context("uncond_stg", timestep=t): noise_pred_video_uncond_stg, noise_pred_audio_uncond_stg = self.transformer( hidden_states=latents.to(dtype=prompt_embeds.dtype), audio_hidden_states=audio_latents.to(dtype=prompt_embeds.dtype), @@ -2380,7 +2380,7 @@ def __call__( if video_self_attention_mask is not None else None ) - with self.transformer.cache_context("uncond_modality"): + with self.transformer.cache_context("uncond_modality", timestep=t): noise_pred_video_uncond_mod, noise_pred_audio_uncond_mod = self.transformer( hidden_states=latents.to(dtype=prompt_embeds.dtype), audio_hidden_states=audio_latents.to(dtype=prompt_embeds.dtype), diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_image2video.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_image2video.py index c7c81d26cb45..e01bedafabb2 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_image2video.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_image2video.py @@ -1470,7 +1470,7 @@ def __call__( timestep = t.expand(latent_model_input.shape[0]) video_timestep = timestep.unsqueeze(-1) * (1 - conditioning_mask) - with self.transformer.cache_context("cond_uncond"): + with self.transformer.cache_context("cond_uncond", timestep=t): noise_pred_video, noise_pred_audio = self.transformer( hidden_states=latent_model_input, audio_hidden_states=audio_latent_model_input, @@ -1544,7 +1544,7 @@ def __call__( noise_pred_audio = self.convert_velocity_to_x0(audio_latents, noise_pred_audio, i, audio_scheduler) if self.do_spatio_temporal_guidance: - with self.transformer.cache_context("uncond_stg"): + with self.transformer.cache_context("uncond_stg", timestep=t): noise_pred_video_uncond_stg, noise_pred_audio_uncond_stg = self.transformer( hidden_states=latents.to(dtype=prompt_embeds.dtype), audio_hidden_states=audio_latents.to(dtype=prompt_embeds.dtype), @@ -1585,7 +1585,7 @@ def __call__( video_stg_delta = audio_stg_delta = 0 if self.do_modality_isolation_guidance: - with self.transformer.cache_context("uncond_modality"): + with self.transformer.cache_context("uncond_modality", timestep=t): noise_pred_video_uncond_modality, noise_pred_audio_uncond_modality = self.transformer( hidden_states=latents.to(dtype=prompt_embeds.dtype), audio_hidden_states=audio_latents.to(dtype=prompt_embeds.dtype), diff --git a/src/diffusers/pipelines/lucy/pipeline_lucy_edit.py b/src/diffusers/pipelines/lucy/pipeline_lucy_edit.py index 1bd7ab4ca675..cb6bd91f3308 100644 --- a/src/diffusers/pipelines/lucy/pipeline_lucy_edit.py +++ b/src/diffusers/pipelines/lucy/pipeline_lucy_edit.py @@ -670,7 +670,7 @@ def __call__( else: timestep = t.expand(latents.shape[0]) - with current_model.cache_context("cond"): + with current_model.cache_context("cond", timestep=t): noise_pred = current_model( hidden_states=latent_model_input, timestep=timestep, @@ -680,7 +680,7 @@ def __call__( )[0] if self.do_classifier_free_guidance: - with current_model.cache_context("uncond"): + with current_model.cache_context("uncond", timestep=t): noise_uncond = current_model( hidden_states=latent_model_input, timestep=timestep, diff --git a/src/diffusers/pipelines/mochi/pipeline_mochi.py b/src/diffusers/pipelines/mochi/pipeline_mochi.py index c146d2d1e564..ec94e48d991e 100644 --- a/src/diffusers/pipelines/mochi/pipeline_mochi.py +++ b/src/diffusers/pipelines/mochi/pipeline_mochi.py @@ -648,7 +648,7 @@ def __call__( # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latent_model_input.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond_uncond"): + with self.transformer.cache_context("cond_uncond", timestep=1000 - t): noise_pred = self.transformer( hidden_states=latent_model_input, encoder_hidden_states=prompt_embeds, diff --git a/src/diffusers/pipelines/motif_video/pipeline_motif_video.py b/src/diffusers/pipelines/motif_video/pipeline_motif_video.py index 8ad37932e970..8c984a2573d7 100644 --- a/src/diffusers/pipelines/motif_video/pipeline_motif_video.py +++ b/src/diffusers/pipelines/motif_video/pipeline_motif_video.py @@ -731,7 +731,7 @@ def __call__( } context_name = getattr(guider_state_batch, self.guider._identifier_key) - with self.transformer.cache_context(context_name): + with self.transformer.cache_context(context_name, timestep=t): noise_pred = self.transformer( hidden_states=hidden_states, timestep=timestep, diff --git a/src/diffusers/pipelines/motif_video/pipeline_motif_video_image2video.py b/src/diffusers/pipelines/motif_video/pipeline_motif_video_image2video.py index 1b32ba74f24b..7ff4fcb1328b 100644 --- a/src/diffusers/pipelines/motif_video/pipeline_motif_video_image2video.py +++ b/src/diffusers/pipelines/motif_video/pipeline_motif_video_image2video.py @@ -855,7 +855,7 @@ def __call__( } context_name = getattr(guider_state_batch, self.guider._identifier_key) - with self.transformer.cache_context(context_name): + with self.transformer.cache_context(context_name, timestep=t): noise_pred = self.transformer( hidden_states=hidden_states, timestep=timestep, diff --git a/src/diffusers/pipelines/ovis_image/pipeline_ovis_image.py b/src/diffusers/pipelines/ovis_image/pipeline_ovis_image.py index b22f2f0cec2d..c377e62198d1 100644 --- a/src/diffusers/pipelines/ovis_image/pipeline_ovis_image.py +++ b/src/diffusers/pipelines/ovis_image/pipeline_ovis_image.py @@ -644,7 +644,7 @@ def __call__( # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latents.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, @@ -656,7 +656,7 @@ def __call__( )[0] if do_classifier_free_guidance: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage.py index 1da0518a4f65..6d1ac9513b13 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage.py @@ -641,7 +641,7 @@ def __call__( self._current_timestep = t # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latents.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, @@ -654,7 +654,7 @@ def __call__( )[0] if do_true_cfg: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py index f946fdf27d00..237b44f0c33c 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py @@ -883,7 +883,7 @@ def __call__( return_dict=False, ) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, @@ -896,7 +896,7 @@ def __call__( )[0] if do_true_cfg: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py index 97f510a6dbf4..014ce6299fcd 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py @@ -854,7 +854,7 @@ def __call__( return_dict=False, ) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, @@ -867,7 +867,7 @@ def __call__( )[0] if do_true_cfg: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit.py index 85abb815cf23..bfcdcd339b7b 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit.py @@ -768,7 +768,7 @@ def __call__( # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latents.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, @@ -782,7 +782,7 @@ def __call__( noise_pred = noise_pred[:, : latents.size(1)] if do_true_cfg: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_inpaint.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_inpaint.py index 57d1fdaaf99f..7918dd6aae45 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_inpaint.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_inpaint.py @@ -982,7 +982,7 @@ def __call__( # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latents.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, @@ -996,7 +996,7 @@ def __call__( noise_pred = noise_pred[:, : latents.size(1)] if do_true_cfg: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_plus.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_plus.py index 84d1b60152b1..2af04b5ea855 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_plus.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_edit_plus.py @@ -812,7 +812,7 @@ def __call__( # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latents.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, @@ -826,7 +826,7 @@ def __call__( noise_pred = noise_pred[:, : latents.size(1)] if do_true_cfg: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_img2img.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_img2img.py index 9b9af83737e5..bc232d3670f5 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_img2img.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_img2img.py @@ -743,7 +743,7 @@ def __call__( self._current_timestep = t # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latents.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, @@ -756,7 +756,7 @@ def __call__( )[0] if do_true_cfg: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_inpaint.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_inpaint.py index 3d5f0040932a..74358dcab4a8 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_inpaint.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_inpaint.py @@ -912,7 +912,7 @@ def __call__( self._current_timestep = t # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latents.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, @@ -925,7 +925,7 @@ def __call__( )[0] if do_true_cfg: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_layered.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_layered.py index 7e06a7d36ffd..be8ff4580a0c 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_layered.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_layered.py @@ -811,7 +811,7 @@ def __call__( # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latents.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, @@ -826,7 +826,7 @@ def __call__( noise_pred = noise_pred[:, : latents.size(1)] if do_true_cfg: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py b/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py index 786b09b4e3cd..7c9927a153b7 100644 --- a/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py +++ b/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py @@ -768,7 +768,7 @@ def append_target_slots(mask): latent_model_input = torch.cat([input_images_latents, latents], dim=1) timestep = t.expand(latents.shape[0]).to(latents.dtype) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, @@ -784,7 +784,7 @@ def append_target_slots(mask): noise_pred = noise_pred[:, -latents.size(1) :] if do_true_cfg: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): neg_noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep / 1000, diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2.py index 0c9e6add9937..397df7f17799 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2.py @@ -550,7 +550,7 @@ def __call__( latent_model_input = latents.to(transformer_dtype) timestep = t.expand(latents.shape[0]) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep, @@ -560,7 +560,7 @@ def __call__( )[0] if self.do_classifier_free_guidance: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): noise_uncond = self.transformer( hidden_states=latent_model_input, timestep=timestep, diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing.py index 31b75bfb336f..f432ce5f5dc5 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing.py @@ -885,7 +885,7 @@ def __call__( ) timestep[:, valid_interval_start:prefix_video_latents_frames] = addnoise_condition - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep, @@ -897,7 +897,7 @@ def __call__( )[0] if self.do_classifier_free_guidance: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): noise_uncond = self.transformer( hidden_states=latent_model_input, timestep=timestep, diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_i2v.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_i2v.py index 576681b1b957..0d7111a72540 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_i2v.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_i2v.py @@ -964,7 +964,7 @@ def __call__( ) timestep[:, valid_interval_start:prefix_video_latents_frames] = addnoise_condition - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep, @@ -976,7 +976,7 @@ def __call__( )[0] if self.do_classifier_free_guidance: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): noise_uncond = self.transformer( hidden_states=latent_model_input, timestep=timestep, diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_v2v.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_v2v.py index df6076263238..ddeee0462bcb 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_v2v.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_diffusion_forcing_v2v.py @@ -974,7 +974,7 @@ def __call__( ) timestep[:, valid_interval_start:prefix_video_latents_frames] = addnoise_condition - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep, @@ -986,7 +986,7 @@ def __call__( )[0] if self.do_classifier_free_guidance: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): noise_uncond = self.transformer( hidden_states=latent_model_input, timestep=timestep, diff --git a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_i2v.py b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_i2v.py index b1f70b60b22a..5c20ce47507d 100644 --- a/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_i2v.py +++ b/src/diffusers/pipelines/skyreels_v2/pipeline_skyreels_v2_i2v.py @@ -681,7 +681,7 @@ def __call__( latent_model_input = torch.cat([latents, condition], dim=1).to(transformer_dtype) timestep = t.expand(latents.shape[0]) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep, @@ -692,7 +692,7 @@ def __call__( )[0] if self.do_classifier_free_guidance: - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): noise_uncond = self.transformer( hidden_states=latent_model_input, timestep=timestep, diff --git a/src/diffusers/pipelines/wan/pipeline_wan.py b/src/diffusers/pipelines/wan/pipeline_wan.py index b33a2a7af3db..ffd3027837c2 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan.py +++ b/src/diffusers/pipelines/wan/pipeline_wan.py @@ -614,6 +614,7 @@ def __call__( cache_context_kwargs = { "step_index": i, + "timestep": t, "sigma": float(self.scheduler.sigmas[i]), "num_inference_steps": self._num_timesteps, } diff --git a/src/diffusers/pipelines/wan/pipeline_wan_animate.py b/src/diffusers/pipelines/wan/pipeline_wan_animate.py index a6b340c2d19f..7ce5a5d65283 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan_animate.py +++ b/src/diffusers/pipelines/wan/pipeline_wan_animate.py @@ -1112,7 +1112,7 @@ def __call__( latent_model_input = torch.cat([latents, reference_latents], dim=1).to(transformer_dtype) timestep = t.expand(latents.shape[0]) - with self.transformer.cache_context("cond"): + with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( hidden_states=latent_model_input, timestep=timestep, @@ -1128,7 +1128,7 @@ def __call__( if self.do_classifier_free_guidance: # Blank out face for unconditional guidance (set all pixels to -1) face_pixel_values_uncond = face_video_segment * 0 - 1 - with self.transformer.cache_context("uncond"): + with self.transformer.cache_context("uncond", timestep=t): noise_uncond = self.transformer( hidden_states=latent_model_input, timestep=timestep, diff --git a/src/diffusers/pipelines/wan/pipeline_wan_i2v.py b/src/diffusers/pipelines/wan/pipeline_wan_i2v.py index a98f0324e0f0..423e506b8fd2 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan_i2v.py +++ b/src/diffusers/pipelines/wan/pipeline_wan_i2v.py @@ -772,7 +772,7 @@ def __call__( latent_model_input = torch.cat([latents, condition], dim=1).to(transformer_dtype) timestep = t.expand(latents.shape[0]) - with current_model.cache_context("cond"): + with current_model.cache_context("cond", timestep=t): noise_pred = current_model( hidden_states=latent_model_input, timestep=timestep, @@ -783,7 +783,7 @@ def __call__( )[0] if self.do_classifier_free_guidance: - with current_model.cache_context("uncond"): + with current_model.cache_context("uncond", timestep=t): noise_uncond = current_model( hidden_states=latent_model_input, timestep=timestep, diff --git a/src/diffusers/pipelines/wan/pipeline_wan_vace.py b/src/diffusers/pipelines/wan/pipeline_wan_vace.py index 9186304b5953..635d9fd01325 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan_vace.py +++ b/src/diffusers/pipelines/wan/pipeline_wan_vace.py @@ -975,7 +975,7 @@ def __call__( latent_model_input = latents.to(transformer_dtype) timestep = t.expand(latents.shape[0]) - with current_model.cache_context("cond"): + with current_model.cache_context("cond", timestep=t): noise_pred = current_model( hidden_states=latent_model_input, timestep=timestep, @@ -987,7 +987,7 @@ def __call__( )[0] if self.do_classifier_free_guidance: - with current_model.cache_context("uncond"): + with current_model.cache_context("uncond", timestep=t): noise_uncond = current_model( hidden_states=latent_model_input, timestep=timestep, diff --git a/tests/models/testing_utils/cache.py b/tests/models/testing_utils/cache.py index 5e330105a4c4..5f88f6640e20 100644 --- a/tests/models/testing_utils/cache.py +++ b/tests/models/testing_utils/cache.py @@ -182,7 +182,8 @@ def _test_cache_inference(self): model.enable_cache(config) # First pass populates the cache - _ = model(**inputs_dict, return_dict=False)[0] + with model.cache_context("test", timestep=1000): + _ = model(**inputs_dict, return_dict=False)[0] # Create modified inputs for second pass (vary input tensor to simulate denoising) inputs_dict_step2 = inputs_dict.copy() @@ -192,7 +193,8 @@ def _test_cache_inference(self): ) # Second pass uses cached attention with different inputs (produces approximated output) - output_with_cache = model(**inputs_dict_step2, return_dict=False)[0] + with model.cache_context("test", timestep=500): + output_with_cache = model(**inputs_dict_step2, return_dict=False)[0] assert output_with_cache is not None, "Model output should not be None with cache enabled." assert not torch.isnan(output_with_cache).any(), "Model output contains NaN with cache enabled." @@ -218,11 +220,11 @@ def _test_cache_context_manager(self, atol=1e-5, rtol=0): model.enable_cache(config) # Run inference in first context - with model.cache_context("context_1"): + with model.cache_context("context_1", timestep=1000): output_ctx1 = model(**inputs_dict, return_dict=False)[0] # Run same inference in second context (cache should be reset) - with model.cache_context("context_2"): + with model.cache_context("context_2", timestep=1000): output_ctx2 = model(**inputs_dict, return_dict=False)[0] # Both contexts should produce the same output (first pass in each) @@ -248,7 +250,8 @@ def _test_reset_stateful_cache(self): model.enable_cache(config) - _ = model(**inputs_dict, return_dict=False)[0] + with model.cache_context("test", timestep=1000): + _ = model(**inputs_dict, return_dict=False)[0] model._reset_stateful_cache() @@ -269,13 +272,8 @@ class PyramidAttentionBroadcastConfigMixin: "spatial_attention_block_skip_range": 2, } - # Store timestep for callback (must be within default range (100, 800) for skipping to trigger) - _current_timestep = 500 - def _get_cache_config(self): - config_kwargs = self.PAB_CONFIG.copy() - config_kwargs["current_timestep_callback"] = lambda: self._current_timestep - return PyramidAttentionBroadcastConfig(**config_kwargs) + return PyramidAttentionBroadcastConfig(**self.PAB_CONFIG) def _get_hook_names(self): return [_PYRAMID_ATTENTION_BROADCAST_HOOK] @@ -645,12 +643,8 @@ class FasterCacheConfigMixin: "tensor_format": "BCHW", } - def _get_cache_config(self, current_timestep_callback=None): - config_kwargs = self.FASTER_CACHE_CONFIG.copy() - if current_timestep_callback is None: - current_timestep_callback = lambda: 1000 # noqa: E731 - config_kwargs["current_timestep_callback"] = current_timestep_callback - return FasterCacheConfig(**config_kwargs) + def _get_cache_config(self): + return FasterCacheConfig(**self.FASTER_CACHE_CONFIG) def _get_hook_names(self): return [_FASTER_CACHE_DENOISER_HOOK, _FASTER_CACHE_BLOCK_HOOK] @@ -684,17 +678,13 @@ def _test_cache_inference(self): model = self.model_class(**init_dict).to(torch_device) model.eval() - current_timestep = [1000] - config = self._get_cache_config(current_timestep_callback=lambda: current_timestep[0]) + config = self._get_cache_config() model.enable_cache(config) # First pass with timestep outside skip range - computes and populates cache - current_timestep[0] = 1000 - _ = model(**inputs_dict, return_dict=False)[0] - - # Move timestep inside skip range so subsequent passes use cache - current_timestep[0] = 500 + with model.cache_context("test", timestep=1000): + _ = model(**inputs_dict, return_dict=False)[0] # Create modified inputs for second pass inputs_dict_step2 = inputs_dict.copy() @@ -704,7 +694,8 @@ def _test_cache_inference(self): ) # Second pass uses cached attention with different inputs - output_with_cache = model(**inputs_dict_step2, return_dict=False)[0] + with model.cache_context("test", timestep=500): + output_with_cache = model(**inputs_dict_step2, return_dict=False)[0] assert output_with_cache is not None, "Model output should not be None with cache enabled." assert not torch.isnan(output_with_cache).any(), "Model output contains NaN with cache enabled." @@ -729,7 +720,8 @@ def _test_reset_stateful_cache(self): config = self._get_cache_config() model.enable_cache(config) - _ = model(**inputs_dict, return_dict=False)[0] + with model.cache_context("test", timestep=1000): + _ = model(**inputs_dict, return_dict=False)[0] model._reset_stateful_cache() diff --git a/tests/pipelines/test_pipelines_common.py b/tests/pipelines/test_pipelines_common.py index 106ba55cf149..d3f1d851d44b 100644 --- a/tests/pipelines/test_pipelines_common.py +++ b/tests/pipelines/test_pipelines_common.py @@ -2509,7 +2509,6 @@ def test_pyramid_attention_broadcast_layers(self): pipe = self.pipeline_class(**components) pipe.set_progress_bar_config(disable=None) - self.pab_config.current_timestep_callback = lambda: pipe.current_timestep denoiser = pipe.transformer if hasattr(pipe, "transformer") else pipe.unet denoiser.enable_cache(self.pab_config) @@ -2533,7 +2532,7 @@ def test_pyramid_attention_broadcast_layers(self): isinstance(hook, PyramidAttentionBroadcastHook), "Hook should be of type PyramidAttentionBroadcastHook.", ) - self.assertTrue(hook.state.cache is None, "Cache should be None at initialization.") + self.assertTrue(not hook.state_manager._state_cache, "Cache should be None at initialization.") self.assertEqual(count, expected_hooks, "Number of hooks should match the expected number.") # Perform dummy inference step to ensure state is updated @@ -2543,12 +2542,13 @@ def pab_state_check_callback(pipe, i, t, kwargs): hook = module._diffusers_hook.get_hook("pyramid_attention_broadcast") if hook is None: continue + self.assertTrue(hook.state_manager._state_cache) self.assertTrue( - hook.state.cache is not None, + all(state.cache is not None for state in hook.state_manager._state_cache.values()), "Cache should have updated during inference.", ) self.assertTrue( - hook.state.iteration == i + 1, + all(state.iteration == i + 1 for state in hook.state_manager._state_cache.values()), "Hook iteration state should have updated during inference.", ) return {} @@ -2565,13 +2565,9 @@ def pab_state_check_callback(pipe, i, t, kwargs): if hook is None: continue self.assertTrue( - hook.state.cache is None, + not hook.state_manager._state_cache, "Cache should be reset to None after inference.", ) - self.assertTrue( - hook.state.iteration == 0, - "Iteration should be reset to 0 after inference.", - ) def test_pyramid_attention_broadcast_inference(self, expected_atol: float = 0.2): # We need to use higher tolerance because we are using a random model. With a converged/trained @@ -2595,7 +2591,6 @@ def test_pyramid_attention_broadcast_inference(self, expected_atol: float = 0.2) original_image_slice = np.concatenate((original_image_slice[:8], original_image_slice[-8:])) # Run inference with PAB enabled - self.pab_config.current_timestep_callback = lambda: pipe.current_timestep denoiser = pipe.transformer if hasattr(pipe, "transformer") else pipe.unet denoiser.enable_cache(self.pab_config) @@ -2673,7 +2668,6 @@ def run_forward(pipe): original_image_slice = np.concatenate((output[:8], output[-8:])) # Run inference with FasterCache enabled - self.faster_cache_config.current_timestep_callback = lambda: pipe.current_timestep pipe = create_pipe() pipe.transformer.enable_cache(self.faster_cache_config) output = run_forward(pipe).flatten() @@ -2710,7 +2704,6 @@ def test_faster_cache_state(self): pipe = self.pipeline_class(**components) pipe.set_progress_bar_config(disable=None) - self.faster_cache_config.current_timestep_callback = lambda: pipe.current_timestep pipe.transformer.enable_cache(self.faster_cache_config) expected_hooks = 0 @@ -2746,17 +2739,25 @@ def faster_cache_state_check_callback(pipe, i, t, kwargs): for name, module in denoiser.named_modules(): if not hasattr(module, "_diffusers_hook"): continue - if name == "": - # Root denoiser module - state = module._diffusers_hook.get_hook(_FASTER_CACHE_DENOISER_HOOK).state - if not self.faster_cache_config.is_guidance_distilled: - self.assertTrue(state.low_frequency_delta is not None, "Low frequency delta should be set.") - self.assertTrue(state.high_frequency_delta is not None, "High frequency delta should be set.") - else: - # Internal blocks - state = module._diffusers_hook.get_hook(_FASTER_CACHE_BLOCK_HOOK).state - self.assertTrue(state.cache is not None and len(state.cache) == 2, "Cache should be set.") - self.assertTrue(state.iteration == i + 1, "Hook iteration state should have updated during inference.") + hook_name = _FASTER_CACHE_DENOISER_HOOK if name == "" else _FASTER_CACHE_BLOCK_HOOK + hook = module._diffusers_hook.get_hook(hook_name) + assert hook.state_manager._state_cache + for state in hook.state_manager._state_cache.values(): + if name == "": + # Root denoiser module + if not self.faster_cache_config.is_guidance_distilled: + self.assertTrue( + state.low_frequency_delta is not None, "Low frequency delta should be set." + ) + self.assertTrue( + state.high_frequency_delta is not None, "High frequency delta should be set." + ) + else: + # Internal blocks + self.assertTrue(state.cache is not None and len(state.cache) == 2, "Cache should be set.") + self.assertTrue( + state.iteration == i + 1, "Hook iteration state should have updated during inference." + ) return {} inputs = self.get_dummy_inputs(device) @@ -2764,23 +2765,12 @@ def faster_cache_state_check_callback(pipe, i, t, kwargs): inputs["callback_on_step_end"] = faster_cache_state_check_callback _ = pipe(**inputs)[0] - # After inference, reset_stateful_hooks is called within the pipeline, which should have reset the states for name, module in denoiser.named_modules(): if not hasattr(module, "_diffusers_hook"): continue - - if name == "": - # Root denoiser module - state = module._diffusers_hook.get_hook(_FASTER_CACHE_DENOISER_HOOK).state - self.assertTrue(state.iteration == 0, "Iteration should be reset to 0.") - self.assertTrue(state.low_frequency_delta is None, "Low frequency delta should be reset to None.") - self.assertTrue(state.high_frequency_delta is None, "High frequency delta should be reset to None.") - else: - # Internal blocks - state = module._diffusers_hook.get_hook(_FASTER_CACHE_BLOCK_HOOK).state - self.assertTrue(state.iteration == 0, "Iteration should be reset to 0.") - self.assertTrue(state.batch_size is None, "Batch size should be reset to None.") - self.assertTrue(state.cache is None, "Cache should be reset to None.") + hook_name = _FASTER_CACHE_DENOISER_HOOK if name == "" else _FASTER_CACHE_BLOCK_HOOK + hook = module._diffusers_hook.get_hook(hook_name) + assert not hook.state_manager._state_cache # TODO(aryan, dhruv): the cache tester mixins should probably be rewritten so that more models can be tested out diff --git a/tests/pipelines/testing_utils/cache.py b/tests/pipelines/testing_utils/cache.py index 82980d57fb3c..b6309f419f5d 100644 --- a/tests/pipelines/testing_utils/cache.py +++ b/tests/pipelines/testing_utils/cache.py @@ -32,14 +32,14 @@ class CacheTesterMixin(BasePipelineOutputMixin): Shared machinery for cache-hook tester mixins. Each cache backend subclasses this and supplies its own config, mirroring the model-level `cache.py` layout. Backends store their config *kwargs* as a dict class attribute and build a fresh config instance per test via `_get_cache_config()`; a shared config instance would leak per-test - mutations (e.g. `current_timestep_callback`) across tests. The denoiser-level enable/disable inference comparison + mutations across tests. The denoiser-level enable/disable inference comparison is shared via `_test_cache_inference`; backend-specific state/layer checks live on the subclasses. """ def _get_cache_config(self): raise NotImplementedError("Subclass must implement `_get_cache_config`.") - def _test_cache_inference(self, cache_config, num_inference_steps, expected_atol=0.1, set_timestep_callback=False): + def _test_cache_inference(self, cache_config, num_inference_steps, expected_atol=0.1): device = "cpu" # ensure determinism for the device-dependent torch.Generator def create_pipe(): @@ -59,8 +59,6 @@ def run_forward(pipe): # Run inference with cache enabled pipe = create_pipe() - if set_timestep_callback: - cache_config.current_timestep_callback = lambda: pipe.current_timestep pipe.transformer.enable_cache(cache_config) output = run_forward(pipe).flatten() image_slice_enabled = torch.cat((output[:8], output[-8:])) @@ -112,7 +110,6 @@ def test_pyramid_attention_broadcast_layers(self): pipe = self.get_pipeline(**self.get_dummy_components(**dummy_component_kwargs)) pab_config = self._get_cache_config() - pab_config.current_timestep_callback = lambda: pipe.current_timestep denoiser = pipe.transformer if hasattr(pipe, "transformer") else pipe.unet denoiser.enable_cache(pab_config) @@ -135,7 +132,7 @@ def test_pyramid_attention_broadcast_layers(self): assert isinstance(hook, PyramidAttentionBroadcastHook), ( "Hook should be of type PyramidAttentionBroadcastHook." ) - assert hook.state.cache is None, "Cache should be None at initialization." + assert not hook.state_manager._state_cache, "Cache should be None at initialization." assert count == expected_hooks, "Number of hooks should match the expected number." # Perform dummy inference step to ensure state is updated @@ -145,8 +142,13 @@ def pab_state_check_callback(pipe, i, t, kwargs): hook = module._diffusers_hook.get_hook("pyramid_attention_broadcast") if hook is None: continue - assert hook.state.cache is not None, "Cache should have updated during inference." - assert hook.state.iteration == i + 1, "Hook iteration state should have updated during inference." + assert hook.state_manager._state_cache + assert all(state.cache is not None for state in hook.state_manager._state_cache.values()), ( + "Cache should have updated during inference." + ) + assert all(state.iteration == i + 1 for state in hook.state_manager._state_cache.values()), ( + "Hook iteration state should have updated during inference." + ) return {} inputs = self.get_dummy_inputs() @@ -160,8 +162,7 @@ def pab_state_check_callback(pipe, i, t, kwargs): hook = module._diffusers_hook.get_hook("pyramid_attention_broadcast") if hook is None: continue - assert hook.state.cache is None, "Cache should be reset to None after inference." - assert hook.state.iteration == 0, "Iteration should be reset to 0 after inference." + assert not hook.state_manager._state_cache, "Cache should be reset to None after inference." def test_pyramid_attention_broadcast_inference(self, base_pipe_output, expected_atol: float = 0.2): # We need to use higher tolerance because we are using a random model. With a converged/trained model, the @@ -177,7 +178,6 @@ def test_pyramid_attention_broadcast_inference(self, base_pipe_output, expected_ # Run inference with PAB enabled pab_config = self._get_cache_config() - pab_config.current_timestep_callback = lambda: pipe.current_timestep denoiser = pipe.transformer if hasattr(pipe, "transformer") else pipe.unet denoiser.enable_cache(pab_config) @@ -223,9 +223,7 @@ def _get_cache_config(self): return FasterCacheConfig(**self.FASTER_CACHE_CONFIG) def test_faster_cache_inference(self, expected_atol: float = 0.1): - self._test_cache_inference( - self._get_cache_config(), num_inference_steps=4, expected_atol=expected_atol, set_timestep_callback=True - ) + self._test_cache_inference(self._get_cache_config(), num_inference_steps=4, expected_atol=expected_atol) def test_faster_cache_state(self): from diffusers.hooks.faster_cache import _FASTER_CACHE_BLOCK_HOOK, _FASTER_CACHE_DENOISER_HOOK @@ -244,7 +242,6 @@ def test_faster_cache_state(self): pipe = self.get_pipeline(**self.get_dummy_components(**dummy_component_kwargs)) faster_cache_config = self._get_cache_config() - faster_cache_config.current_timestep_callback = lambda: pipe.current_timestep pipe.transformer.enable_cache(faster_cache_config) # Hook registration/removal is covered at the model level (`_test_cache_hooks_registered`). Here we only @@ -257,17 +254,19 @@ def faster_cache_state_check_callback(pipe, i, t, kwargs): for name, module in denoiser.named_modules(): if not hasattr(module, "_diffusers_hook"): continue - if name == "": - # Root denoiser module - state = module._diffusers_hook.get_hook(_FASTER_CACHE_DENOISER_HOOK).state - if not faster_cache_config.is_guidance_distilled: - assert state.low_frequency_delta is not None, "Low frequency delta should be set." - assert state.high_frequency_delta is not None, "High frequency delta should be set." - else: - # Internal blocks - state = module._diffusers_hook.get_hook(_FASTER_CACHE_BLOCK_HOOK).state - assert state.cache is not None and len(state.cache) == 2, "Cache should be set." - assert state.iteration == i + 1, "Hook iteration state should have updated during inference." + hook_name = _FASTER_CACHE_DENOISER_HOOK if name == "" else _FASTER_CACHE_BLOCK_HOOK + hook = module._diffusers_hook.get_hook(hook_name) + assert hook.state_manager._state_cache + for state in hook.state_manager._state_cache.values(): + if name == "": + # Root denoiser module + if not faster_cache_config.is_guidance_distilled: + assert state.low_frequency_delta is not None, "Low frequency delta should be set." + assert state.high_frequency_delta is not None, "High frequency delta should be set." + else: + # Internal blocks + assert state.cache is not None and len(state.cache) == 2, "Cache should be set." + assert state.iteration == i + 1, "Hook iteration state should have updated during inference." return {} inputs = self.get_dummy_inputs() @@ -275,23 +274,12 @@ def faster_cache_state_check_callback(pipe, i, t, kwargs): inputs["callback_on_step_end"] = faster_cache_state_check_callback _ = pipe(**inputs)[0] - # After inference, reset_stateful_hooks is called within the pipeline, which should have reset the states for name, module in denoiser.named_modules(): if not hasattr(module, "_diffusers_hook"): continue - - if name == "": - # Root denoiser module - state = module._diffusers_hook.get_hook(_FASTER_CACHE_DENOISER_HOOK).state - assert state.iteration == 0, "Iteration should be reset to 0." - assert state.low_frequency_delta is None, "Low frequency delta should be reset to None." - assert state.high_frequency_delta is None, "High frequency delta should be reset to None." - else: - # Internal blocks - state = module._diffusers_hook.get_hook(_FASTER_CACHE_BLOCK_HOOK).state - assert state.iteration == 0, "Iteration should be reset to 0." - assert state.batch_size is None, "Batch size should be reset to None." - assert state.cache is None, "Cache should be reset to None." + hook_name = _FASTER_CACHE_DENOISER_HOOK if name == "" else _FASTER_CACHE_BLOCK_HOOK + hook = module._diffusers_hook.get_hook(hook_name) + assert not hook.state_manager._state_cache # TODO(aryan, dhruv): the cache tester mixins should probably be rewritten so that more models can be tested out