From 0a1f36769ccdd9a3db6adf27d5e54cc6ebd2dee0 Mon Sep 17 00:00:00 2001 From: sayakpaul Date: Tue, 29 Sep 2026 10:33:34 +0530 Subject: [PATCH 1/5] 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 From ad8168c5c00325fcaf2a22164b96f39cd1a9cc34 Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Mon, 5 Oct 2026 03:34:31 +0000 Subject: [PATCH 2/5] deprecate current_timestep_callback instead of removing. --- src/diffusers/hooks/faster_cache.py | 13 ++++++++++++- .../hooks/pyramid_attention_broadcast.py | 15 +++++++++++++-- tests/others/test_utils.py | 10 +++++++++- 3 files changed, 34 insertions(+), 4 deletions(-) diff --git a/src/diffusers/hooks/faster_cache.py b/src/diffusers/hooks/faster_cache.py index 48bdb979e6fe..8638b0a251f0 100644 --- a/src/diffusers/hooks/faster_cache.py +++ b/src/diffusers/hooks/faster_cache.py @@ -20,7 +20,7 @@ from ..models.attention import AttentionModuleMixin from ..models.modeling_outputs import Transformer2DModelOutput -from ..utils import logging +from ..utils import deprecate, logging from ._common import _ATTENTION_CLASSES from .hooks import HookRegistry, ModelHook, StateManager @@ -160,8 +160,19 @@ class FasterCacheConfig: tensor_format: str = "BCFHW" is_guidance_distilled: bool = False + current_timestep_callback: Callable[[], int] | None = None + _unconditional_conditional_input_kwargs_identifiers: list[str] = _UNCOND_COND_INPUT_KWARGS_IDENTIFIERS + def __post_init__(self): + if self.current_timestep_callback is not None: + depr_message = ( + "Passing `current_timestep_callback` to `FasterCacheConfig` is deprecated and will be ignored. " + "Please use `cache_context(name, timestep=...)` to pass the current timestep instead. " + "See the caching documentation for guidance: https://huggingface.co/docs/diffusers/main/en/optimization/cache." + ) + deprecate("current_timestep_callback", "0.45.0", depr_message) + def __repr__(self) -> str: return ( f"FasterCacheConfig(\n" diff --git a/src/diffusers/hooks/pyramid_attention_broadcast.py b/src/diffusers/hooks/pyramid_attention_broadcast.py index 28b6be665152..8a9ee6f24448 100644 --- a/src/diffusers/hooks/pyramid_attention_broadcast.py +++ b/src/diffusers/hooks/pyramid_attention_broadcast.py @@ -14,13 +14,13 @@ import re from dataclasses import dataclass -from typing import Any +from typing import Any, Callable import torch from ..models.attention import AttentionModuleMixin from ..models.attention_processor import Attention, MochiAttention -from ..utils import logging +from ..utils import deprecate, logging from ._common import ( _ATTENTION_CLASSES, _CROSS_TRANSFORMER_BLOCK_IDENTIFIERS, @@ -83,9 +83,20 @@ 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 = None + # TODO(aryan): add PAB for MLP layers (very limited speedup from testing with original codebase # so not added for now) + def __post_init__(self): + if self.current_timestep_callback is not None: + depr_message = ( + "Passing `current_timestep_callback` to `PyramidAttentionBroadcastConfig` is deprecated and will be ignored. " + "Please use `cache_context(name, timestep=...)` to pass the current timestep instead. " + "See the caching documentation for guidance: https://huggingface.co/docs/diffusers/main/en/optimization/cache." + ) + deprecate("current_timestep_callback", "0.45.0", depr_message) + def __repr__(self) -> str: return ( f"PyramidAttentionBroadcastConfig(\n" diff --git a/tests/others/test_utils.py b/tests/others/test_utils.py index bb7e5c298926..56a180355eac 100755 --- a/tests/others/test_utils.py +++ b/tests/others/test_utils.py @@ -20,7 +20,7 @@ import pytest import torch -from diffusers import __version__ +from diffusers import FasterCacheConfig, PyramidAttentionBroadcastConfig, __version__ from diffusers.utils import deprecate, torch_utils from diffusers.utils.torch_utils import TorchDeviceBackend, empty_device_cache, get_device @@ -39,6 +39,14 @@ class TestDeprecate: higher_version = ".".join([str(int(__version__.split(".")[0]) + 1)] + __version__.split(".")[1:]) lower_version = "0.0.1" + @pytest.mark.parametrize("config_class", [FasterCacheConfig, PyramidAttentionBroadcastConfig]) + def test_cache_current_timestep_callback_deprecation(self, config_class): + with pytest.warns( + FutureWarning, + match=r"`current_timestep_callback` is deprecated and will be removed in version 0\.45\.0\.", + ): + config_class(current_timestep_callback=lambda: 500) + def test_deprecate_function_arg(self): kwargs = {"deprecated_arg": 4} From 21922860100d813c25367b298b2c3939f01e980b Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Wed, 7 Oct 2026 02:47:05 +0000 Subject: [PATCH 3/5] preserve deprecated functionality. --- src/diffusers/hooks/faster_cache.py | 28 ++++- .../hooks/pyramid_attention_broadcast.py | 23 +++- tests/hooks/test_cache_timestep.py | 111 ++++++++++++++++++ 3 files changed, 152 insertions(+), 10 deletions(-) create mode 100644 tests/hooks/test_cache_timestep.py diff --git a/src/diffusers/hooks/faster_cache.py b/src/diffusers/hooks/faster_cache.py index 8638b0a251f0..4b602d2d554a 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 deprecate, logging from ._common import _ATTENTION_CLASSES -from .hooks import HookRegistry, ModelHook, StateManager +from .hooks import CacheContext, HookRegistry, ModelHook, StateManager logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -167,7 +167,7 @@ class FasterCacheConfig: def __post_init__(self): if self.current_timestep_callback is not None: depr_message = ( - "Passing `current_timestep_callback` to `FasterCacheConfig` is deprecated and will be ignored. " + "Passing `current_timestep_callback` to `FasterCacheConfig` is deprecated. " "Please use `cache_context(name, timestep=...)` to pass the current timestep instead. " "See the caching documentation for guidance: https://huggingface.co/docs/diffusers/main/en/optimization/cache." ) @@ -238,6 +238,7 @@ def __init__( uncond_cond_input_kwargs_identifiers: list[str], low_frequency_weight_callback: Callable[[torch.nn.Module], torch.Tensor], high_frequency_weight_callback: Callable[[torch.nn.Module], torch.Tensor], + current_timestep_callback: Callable[[], int] | None = None, ) -> None: super().__init__() @@ -253,6 +254,7 @@ def __init__( self.low_frequency_weight_callback = low_frequency_weight_callback self.high_frequency_weight_callback = high_frequency_weight_callback + self.current_timestep_callback = current_timestep_callback def initialize_hook(self, module): self.state_manager = StateManager(FasterCacheDenoiserState) @@ -265,10 +267,18 @@ def _get_cond_input(input: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: _, cond = input.chunk(2, dim=0) return cond - def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any: + def _get_timestep(self): + if self.state_manager._context is None and self.current_timestep_callback is not None: + self.state_manager.set_context(CacheContext(name="inference")) timestep = self.state_manager.context.timestep + if timestep is None and self.current_timestep_callback is not None: + timestep = self.current_timestep_callback() if timestep is None: raise ValueError("FasterCache requires `cache_context(name, timestep=...)`.") + return timestep + + def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any: + timestep = self._get_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. @@ -381,6 +391,7 @@ 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 = None, ) -> None: super().__init__() @@ -389,6 +400,7 @@ 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_manager = StateManager(FasterCacheBlockState) @@ -410,7 +422,11 @@ 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: + if self.state_manager._context is None and self.current_timestep_callback is not None: + self.state_manager.set_context(CacheContext(name="inference")) timestep = self.state_manager.context.timestep + if timestep is None and self.current_timestep_callback is not None: + timestep = self.current_timestep_callback() if timestep is None: raise ValueError("FasterCache requires `cache_context(name, timestep=...)`.") state = self.state_manager.get_state() @@ -548,7 +564,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] - < module._diffusers_hook.get_hook(_FASTER_CACHE_DENOISER_HOOK).state_manager.context.timestep + < module._diffusers_hook.get_hook(_FASTER_CACHE_DENOISER_HOOK)._get_timestep() < config.low_frequency_weight_update_timestep_range[1] ) return config.alpha_low_frequency if is_within_range else 1.0 @@ -563,7 +579,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] - < module._diffusers_hook.get_hook(_FASTER_CACHE_DENOISER_HOOK).state_manager.context.timestep + < module._diffusers_hook.get_hook(_FASTER_CACHE_DENOISER_HOOK)._get_timestep() < config.high_frequency_weight_update_timestep_range[1] ) return config.alpha_high_frequency if is_within_range else 1.0 @@ -592,6 +608,7 @@ def _apply_faster_cache_on_denoiser(module: torch.nn.Module, config: FasterCache config._unconditional_conditional_input_kwargs_identifiers, config.low_frequency_weight_callback, config.high_frequency_weight_callback, + current_timestep_callback=config.current_timestep_callback, ) registry = HookRegistry.check_if_exists_or_initialize(module) registry.register_hook(hook, _FASTER_CACHE_DENOISER_HOOK) @@ -635,6 +652,7 @@ def _apply_faster_cache_on_attention_class(name: str, module: AttentionModuleMix timestep_skip_range, config.is_guidance_distilled, config.attention_weight_callback, + current_timestep_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 8a9ee6f24448..9c132d5583fe 100644 --- a/src/diffusers/hooks/pyramid_attention_broadcast.py +++ b/src/diffusers/hooks/pyramid_attention_broadcast.py @@ -27,7 +27,7 @@ _SPATIAL_TRANSFORMER_BLOCK_IDENTIFIERS, _TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS, ) -from .hooks import HookRegistry, ModelHook, StateManager +from .hooks import CacheContext, HookRegistry, ModelHook, StateManager logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -91,7 +91,7 @@ class PyramidAttentionBroadcastConfig: def __post_init__(self): if self.current_timestep_callback is not None: depr_message = ( - "Passing `current_timestep_callback` to `PyramidAttentionBroadcastConfig` is deprecated and will be ignored. " + "Passing `current_timestep_callback` to `PyramidAttentionBroadcastConfig` is deprecated. " "Please use `cache_context(name, timestep=...)` to pass the current timestep instead. " "See the caching documentation for guidance: https://huggingface.co/docs/diffusers/main/en/optimization/cache." ) @@ -148,18 +148,28 @@ class PyramidAttentionBroadcastHook(ModelHook): _is_stateful = True - def __init__(self, timestep_skip_range: tuple[int, int], block_skip_range: int) -> None: + def __init__( + self, + timestep_skip_range: tuple[int, int], + block_skip_range: int, + current_timestep_callback: Callable[[], int] | None = None, + ) -> 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_manager = StateManager(PyramidAttentionBroadcastState) return module def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any: + if self.state_manager._context is None and self.current_timestep_callback is not None: + self.state_manager.set_context(CacheContext(name="inference")) timestep = self.state_manager.context.timestep + if timestep is None and self.current_timestep_callback is not None: + timestep = self.current_timestep_callback() if timestep is None: raise ValueError("Pyramid Attention Broadcast requires `cache_context(name, timestep=...)`.") state = self.state_manager.get_state() @@ -282,7 +292,9 @@ 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) + _apply_pyramid_attention_broadcast_hook( + module, timestep_skip_range, block_skip_range, config.current_timestep_callback + ) return True @@ -290,6 +302,7 @@ def _apply_pyramid_attention_broadcast_hook( module: Attention | MochiAttention, timestep_skip_range: tuple[int, int], block_skip_range: int, + current_timestep_callback: Callable[[], int] | None = None, ): r""" Apply [Pyramid Attention Broadcast](https://huggingface.co/papers/2408.12588) to a given torch.nn.Module. @@ -306,5 +319,5 @@ def _apply_pyramid_attention_broadcast_hook( attention states will be reused) before computing the new attention states again. """ registry = HookRegistry.check_if_exists_or_initialize(module) - hook = PyramidAttentionBroadcastHook(timestep_skip_range, block_skip_range) + hook = PyramidAttentionBroadcastHook(timestep_skip_range, block_skip_range, current_timestep_callback) registry.register_hook(hook, _PYRAMID_ATTENTION_BROADCAST_HOOK) diff --git a/tests/hooks/test_cache_timestep.py b/tests/hooks/test_cache_timestep.py new file mode 100644 index 000000000000..2253b2828c16 --- /dev/null +++ b/tests/hooks/test_cache_timestep.py @@ -0,0 +1,111 @@ +# Copyright 2026 HuggingFace Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from contextlib import nullcontext +from unittest.mock import Mock + +import pytest +import torch + +from diffusers import CogVideoXTransformer3DModel, FasterCacheConfig, PyramidAttentionBroadcastConfig + + +@pytest.fixture +def model(): + torch.manual_seed(0) + return CogVideoXTransformer3DModel( + num_attention_heads=2, + attention_head_dim=8, + in_channels=4, + out_channels=4, + time_embed_dim=2, + text_embed_dim=8, + num_layers=1, + sample_width=8, + sample_height=8, + sample_frames=1, + patch_size=2, + temporal_compression_ratio=4, + max_text_seq_length=8, + ).eval() + + +@pytest.fixture(params=[FasterCacheConfig, PyramidAttentionBroadcastConfig]) +def config_kwargs(request): + kwargs = {"spatial_attention_block_skip_range": 2} + if request.param is FasterCacheConfig: + kwargs["tensor_format"] = "BFCHW" + return request.param, kwargs + + +@pytest.fixture +def inputs(): + generator = torch.Generator().manual_seed(0) + return [ + { + "hidden_states": torch.randn(2, 1, 4, 8, 8, generator=generator), + "encoder_hidden_states": torch.randn(2, 8, 8, generator=generator), + "timestep": torch.full((2,), timestep), + } + for timestep in [900, 600, 500, 400, 200, 0] + ] + + +@pytest.mark.parametrize("context_name", [None, "cond"]) +@torch.no_grad() +def test_deprecated_timestep_callback_matches_cache_context(model, config_kwargs, inputs, context_name): + config_class, kwargs = config_kwargs + model.enable_cache(config_class(**kwargs)) + expected = [] + for step_inputs in inputs: + with model.cache_context("cond", timestep=step_inputs["timestep"][0]): + expected.append(model(**step_inputs).sample) + model.disable_cache() + + timestep = None + with pytest.warns(FutureWarning, match="current_timestep_callback.*0.45.0"): + config = config_class(**kwargs, current_timestep_callback=lambda: timestep) + model.enable_cache(config) + + for _ in range(2): + for step_inputs, expected_output in zip(inputs, expected): + timestep = step_inputs["timestep"][0] + context = model.cache_context(context_name) if context_name is not None else nullcontext() + with context: + output = model(**step_inputs).sample + torch.testing.assert_close(output, expected_output) + model._reset_stateful_cache() + + +@torch.no_grad() +def test_context_timestep_takes_precedence_over_callback(model, config_kwargs, inputs): + config_class, kwargs = config_kwargs + callback = Mock(side_effect=AssertionError) + with pytest.warns(FutureWarning, match="current_timestep_callback"): + config = config_class(**kwargs, current_timestep_callback=callback) + model.enable_cache(config) + + for step_inputs in inputs: + with model.cache_context("cond", timestep=step_inputs["timestep"][0]): + model(**step_inputs) + callback.assert_not_called() + + +@pytest.mark.parametrize("context_name", [None, "cond"]) +def test_timestep_is_required_without_callback(model, config_kwargs, inputs, context_name): + config_class, kwargs = config_kwargs + model.enable_cache(config_class(**kwargs)) + context = model.cache_context(context_name) if context_name is not None else nullcontext() + with context, pytest.raises(ValueError, match="cache_context"): + model(**inputs[0]) From 0a9b872981941de67d6131d279009cc66ca5ced4 Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Wed, 7 Oct 2026 02:58:51 +0000 Subject: [PATCH 4/5] up --- src/diffusers/hooks/faster_cache.py | 6 +++++- src/diffusers/hooks/pyramid_attention_broadcast.py | 4 ++++ 2 files changed, 9 insertions(+), 1 deletion(-) diff --git a/src/diffusers/hooks/faster_cache.py b/src/diffusers/hooks/faster_cache.py index 4b602d2d554a..e18e72980fce 100644 --- a/src/diffusers/hooks/faster_cache.py +++ b/src/diffusers/hooks/faster_cache.py @@ -174,6 +174,9 @@ def __post_init__(self): deprecate("current_timestep_callback", "0.45.0", depr_message) def __repr__(self) -> str: + callback_repr = "" + if self.current_timestep_callback is not None: + callback_repr = f" current_timestep_callback={self.current_timestep_callback},\n" return ( f"FasterCacheConfig(\n" f" spatial_attention_block_skip_range={self.spatial_attention_block_skip_range},\n" @@ -189,6 +192,7 @@ def __repr__(self) -> str: f" spatial_attention_block_identifiers={self.spatial_attention_block_identifiers},\n" f" temporal_attention_block_identifiers={self.temporal_attention_block_identifiers},\n" f" tensor_format={self.tensor_format},\n" + f"{callback_repr}" f")" ) @@ -252,9 +256,9 @@ 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 - self.current_timestep_callback = current_timestep_callback def initialize_hook(self, module): self.state_manager = StateManager(FasterCacheDenoiserState) diff --git a/src/diffusers/hooks/pyramid_attention_broadcast.py b/src/diffusers/hooks/pyramid_attention_broadcast.py index 9c132d5583fe..aed20a7ce409 100644 --- a/src/diffusers/hooks/pyramid_attention_broadcast.py +++ b/src/diffusers/hooks/pyramid_attention_broadcast.py @@ -98,6 +98,9 @@ def __post_init__(self): deprecate("current_timestep_callback", "0.45.0", depr_message) def __repr__(self) -> str: + callback_repr = "" + if self.current_timestep_callback is not None: + callback_repr = f" current_timestep_callback={self.current_timestep_callback},\n" return ( f"PyramidAttentionBroadcastConfig(\n" f" spatial_attention_block_skip_range={self.spatial_attention_block_skip_range},\n" @@ -109,6 +112,7 @@ def __repr__(self) -> str: f" spatial_attention_block_identifiers={self.spatial_attention_block_identifiers},\n" f" temporal_attention_block_identifiers={self.temporal_attention_block_identifiers},\n" f" cross_attention_block_identifiers={self.cross_attention_block_identifiers},\n" + f"{callback_repr}" ")" ) From e814ff6feca65906332dafeab5c9d711c23e769a Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Thu, 8 Oct 2026 03:32:12 +0000 Subject: [PATCH 5/5] propagate to existing tests. --- .../modular_pipelines/echo/denoise.py | 59 +++++++-------- .../modular_pipelines/flux/denoise.py | 46 ++++++------ .../modular_pipelines/flux2/denoise.py | 62 ++++++++-------- .../modular_pipelines/minimax_h3/denoise.py | 21 +++--- .../modular_pipelines/qwenimage/denoise.py | 59 +++++++-------- .../modular_pipelines/wan/denoise.py | 16 +++-- .../pipelines/ace_step/pipeline_ace_step.py | 59 ++++++++------- src/diffusers/pipelines/bria/pipeline_bria.py | 19 ++--- .../pipelines/chroma/pipeline_chroma.py | 38 +++++----- .../chroma/pipeline_chroma_img2img.py | 40 ++++++----- .../chroma/pipeline_chroma_inpainting.py | 40 ++++++----- .../chronoedit/pipeline_chronoedit.py | 26 +++---- .../cogview4/pipeline_cogview4_control.py | 32 +++++---- .../pipelines/flux/pipeline_flux_control.py | 23 +++--- .../flux/pipeline_flux_control_img2img.py | 23 +++--- .../flux/pipeline_flux_control_inpaint.py | 23 +++--- .../flux/pipeline_flux_controlnet.py | 44 ++++++------ ...pipeline_flux_controlnet_image_to_image.py | 29 ++++---- .../pipeline_flux_controlnet_inpainting.py | 29 ++++---- .../pipelines/flux/pipeline_flux_fill.py | 23 +++--- .../pipelines/flux/pipeline_flux_img2img.py | 38 +++++----- .../pipelines/flux/pipeline_flux_inpaint.py | 38 +++++----- .../pipelines/flux2/pipeline_flux2.py | 21 +++--- .../flux2/pipeline_flux2_klein_kv.py | 71 ++++++++++--------- .../pipelines/glm_image/pipeline_glm_image.py | 69 +++++++++--------- .../kandinsky5/pipeline_kandinsky.py | 36 +++++----- .../kandinsky5/pipeline_kandinsky_i2i.py | 36 +++++----- .../kandinsky5/pipeline_kandinsky_i2v.py | 36 +++++----- .../kandinsky5/pipeline_kandinsky_t2i.py | 36 +++++----- .../kandinsky6/pipeline_kandinsky6_ti2va.py | 4 +- .../pipelines/ltx2/pipeline_ltx2_dfr.py | 36 +++++----- .../ltx2/pipeline_ltx2_dfr_temporal_refine.py | 36 +++++----- .../pipeline_nucleusmoe_image.py | 30 ++++---- .../pipeline_qwenimage_controlnet.py | 21 +++--- .../pipeline_qwenimage_controlnet_inpaint.py | 21 +++--- .../pipeline_visualcloze_generation.py | 23 +++--- .../pipelines/wan/pipeline_wan_video2video.py | 24 ++++--- 37 files changed, 679 insertions(+), 608 deletions(-) diff --git a/src/diffusers/modular_pipelines/echo/denoise.py b/src/diffusers/modular_pipelines/echo/denoise.py index cbe80c0ab609..ea6b4c891c82 100644 --- a/src/diffusers/modular_pipelines/echo/denoise.py +++ b/src/diffusers/modular_pipelines/echo/denoise.py @@ -233,34 +233,37 @@ def __call__( audio_context = block_state.connector_audio_prompt_embeds.to(transformer_dtype) context_mask = block_state.connector_attention_mask - velocity_video, velocity_audio = components.transformer( - hidden_states=block_state.latent_model_input.to(transformer_dtype), - audio_hidden_states=block_state.audio_latent_model_input.to(transformer_dtype), - encoder_hidden_states=video_context, - audio_encoder_hidden_states=audio_context, - timestep=block_state.model_video_timestep, - audio_timestep=block_state.model_audio_timestep, - # Echo deliberately uses the first target token as its global video sigma. When a clean first frame is - # present this is zero, while the audio branch keeps the current DMD sigma. `use_cross_timestep=True` - # exchanges these values in the cross-modal modulation blocks, matching the released wrapper. - sigma=block_state.video_timestep[:, 0], - audio_sigma=block_state.audio_timestep[:, 0], - encoder_attention_mask=context_mask, - audio_encoder_attention_mask=context_mask, - num_frames=block_state.latent_num_frames, - height=block_state.latent_height, - width=block_state.latent_width, - fps=block_state.frame_rate, - audio_num_frames=block_state.audio_num_frames, - video_coords=block_state.model_video_coords, - audio_coords=block_state.model_audio_coords, - isolate_modalities=False, - spatio_temporal_guidance_blocks=None, - perturbation_mask=None, - use_cross_timestep=True, - attention_kwargs=block_state.attention_kwargs, - return_dict=False, - ) + with components.transformer.cache_context( + "inference", timestep=sigma * components.transformer.config.timestep_scale_multiplier + ): + velocity_video, velocity_audio = components.transformer( + hidden_states=block_state.latent_model_input.to(transformer_dtype), + audio_hidden_states=block_state.audio_latent_model_input.to(transformer_dtype), + encoder_hidden_states=video_context, + audio_encoder_hidden_states=audio_context, + timestep=block_state.model_video_timestep, + audio_timestep=block_state.model_audio_timestep, + # Echo deliberately uses the first target token as its global video sigma. When a clean first frame is + # present this is zero, while the audio branch keeps the current DMD sigma. `use_cross_timestep=True` + # exchanges these values in the cross-modal modulation blocks, matching the released wrapper. + sigma=block_state.video_timestep[:, 0], + audio_sigma=block_state.audio_timestep[:, 0], + encoder_attention_mask=context_mask, + audio_encoder_attention_mask=context_mask, + num_frames=block_state.latent_num_frames, + height=block_state.latent_height, + width=block_state.latent_width, + fps=block_state.frame_rate, + audio_num_frames=block_state.audio_num_frames, + video_coords=block_state.model_video_coords, + audio_coords=block_state.model_audio_coords, + isolate_modalities=False, + spatio_temporal_guidance_blocks=None, + perturbation_mask=None, + use_cross_timestep=True, + attention_kwargs=block_state.attention_kwargs, + return_dict=False, + ) velocity_video = velocity_video[:, block_state.memory_video_token_count :] velocity_audio = velocity_audio[:, block_state.memory_audio_token_count :] timestep_scale = float(components.transformer.config.timestep_scale_multiplier) diff --git a/src/diffusers/modular_pipelines/flux/denoise.py b/src/diffusers/modular_pipelines/flux/denoise.py index 7f2e20ffcec9..7ca1c47bbfb1 100644 --- a/src/diffusers/modular_pipelines/flux/denoise.py +++ b/src/diffusers/modular_pipelines/flux/denoise.py @@ -93,17 +93,18 @@ def inputs(self) -> list[tuple[str, Any]]: def __call__( self, components: FluxModularPipeline, block_state: BlockState, i: int, t: torch.Tensor ) -> tuple[FluxModularPipeline, BlockState]: - noise_pred = components.transformer( - hidden_states=block_state.latents, - timestep=t.flatten() / 1000, - guidance=block_state.guidance, - encoder_hidden_states=block_state.prompt_embeds, - pooled_projections=block_state.pooled_prompt_embeds, - joint_attention_kwargs=block_state.joint_attention_kwargs, - txt_ids=block_state.txt_ids, - img_ids=block_state.img_ids, - return_dict=False, - )[0] + with components.transformer.cache_context("inference", timestep=t): + noise_pred = components.transformer( + hidden_states=block_state.latents, + timestep=t.flatten() / 1000, + guidance=block_state.guidance, + encoder_hidden_states=block_state.prompt_embeds, + pooled_projections=block_state.pooled_prompt_embeds, + joint_attention_kwargs=block_state.joint_attention_kwargs, + txt_ids=block_state.txt_ids, + img_ids=block_state.img_ids, + return_dict=False, + )[0] block_state.noise_pred = noise_pred return components, block_state @@ -182,17 +183,18 @@ def __call__( latent_model_input = torch.cat([latent_model_input, image_latents], dim=1) timestep = t.expand(latents.shape[0]).to(latents.dtype) - noise_pred = components.transformer( - hidden_states=latent_model_input, - timestep=timestep / 1000, - guidance=block_state.guidance, - encoder_hidden_states=block_state.prompt_embeds, - pooled_projections=block_state.pooled_prompt_embeds, - joint_attention_kwargs=block_state.joint_attention_kwargs, - txt_ids=block_state.txt_ids, - img_ids=block_state.img_ids, - return_dict=False, - )[0] + with components.transformer.cache_context("inference", timestep=t): + noise_pred = components.transformer( + hidden_states=latent_model_input, + timestep=timestep / 1000, + guidance=block_state.guidance, + encoder_hidden_states=block_state.prompt_embeds, + pooled_projections=block_state.pooled_prompt_embeds, + joint_attention_kwargs=block_state.joint_attention_kwargs, + txt_ids=block_state.txt_ids, + img_ids=block_state.img_ids, + return_dict=False, + )[0] noise_pred = noise_pred[:, : latents.size(1)] block_state.noise_pred = noise_pred diff --git a/src/diffusers/modular_pipelines/flux2/denoise.py b/src/diffusers/modular_pipelines/flux2/denoise.py index fa6877180057..10723aa2fe1e 100644 --- a/src/diffusers/modular_pipelines/flux2/denoise.py +++ b/src/diffusers/modular_pipelines/flux2/denoise.py @@ -119,16 +119,17 @@ def __call__( timestep = t.expand(latents.shape[0]).to(latents.dtype) - noise_pred = components.transformer( - hidden_states=latent_model_input, - timestep=timestep / 1000, - guidance=block_state.guidance, - encoder_hidden_states=block_state.prompt_embeds, - txt_ids=block_state.txt_ids, - img_ids=img_ids, - joint_attention_kwargs=block_state.joint_attention_kwargs, - return_dict=False, - )[0] + with components.transformer.cache_context("inference", timestep=t): + noise_pred = components.transformer( + hidden_states=latent_model_input, + timestep=timestep / 1000, + guidance=block_state.guidance, + encoder_hidden_states=block_state.prompt_embeds, + txt_ids=block_state.txt_ids, + img_ids=img_ids, + joint_attention_kwargs=block_state.joint_attention_kwargs, + return_dict=False, + )[0] noise_pred = noise_pred[:, : latents.size(1)] block_state.noise_pred = noise_pred @@ -208,16 +209,17 @@ def __call__( timestep = t.expand(latents.shape[0]).to(latents.dtype) - noise_pred = components.transformer( - hidden_states=latent_model_input, - timestep=timestep / 1000, - guidance=None, - encoder_hidden_states=block_state.prompt_embeds, - txt_ids=block_state.txt_ids, - img_ids=img_ids, - joint_attention_kwargs=block_state.joint_attention_kwargs, - return_dict=False, - )[0] + with components.transformer.cache_context("inference", timestep=t): + noise_pred = components.transformer( + hidden_states=latent_model_input, + timestep=timestep / 1000, + guidance=None, + encoder_hidden_states=block_state.prompt_embeds, + txt_ids=block_state.txt_ids, + img_ids=img_ids, + joint_attention_kwargs=block_state.joint_attention_kwargs, + return_dict=False, + )[0] noise_pred = noise_pred[:, : latents.size(1)] block_state.noise_pred = noise_pred @@ -341,15 +343,17 @@ def __call__( components.guider.prepare_models(components.transformer) cond_kwargs = {input_name: getattr(guider_state_batch, input_name) for input_name in guider_inputs.keys()} - noise_pred = components.transformer( - hidden_states=latent_model_input, - timestep=timestep / 1000, - guidance=None, - img_ids=img_ids, - joint_attention_kwargs=block_state.joint_attention_kwargs, - return_dict=False, - **cond_kwargs, - )[0] + context_name = getattr(guider_state_batch, components.guider._identifier_key) + with components.transformer.cache_context(context_name, timestep=t): + noise_pred = components.transformer( + hidden_states=latent_model_input, + timestep=timestep / 1000, + guidance=None, + img_ids=img_ids, + joint_attention_kwargs=block_state.joint_attention_kwargs, + return_dict=False, + **cond_kwargs, + )[0] guider_state_batch.noise_pred = noise_pred[:, : latents.size(1)] components.guider.cleanup_models(components.transformer) diff --git a/src/diffusers/modular_pipelines/minimax_h3/denoise.py b/src/diffusers/modular_pipelines/minimax_h3/denoise.py index efcf36c1130b..35357a4d6e89 100644 --- a/src/diffusers/modular_pipelines/minimax_h3/denoise.py +++ b/src/diffusers/modular_pipelines/minimax_h3/denoise.py @@ -120,16 +120,17 @@ def __call__( for name, value in block_state.denoiser_input_fields.items() if name in inspect.signature(transformer.forward).parameters } - block_state.noise_pred, block_state.audio_noise_pred = transformer( - hidden_states=block_state.latents[None], - audio_hidden_states=block_state.audio_latents[None], - encoder_hidden_states=block_state.prompt_embeds, - timestep=unique_timesteps, - timestep_indices=timestep_indices, - attention_kwargs=block_state.attention_kwargs, - return_dict=False, - **layout_kwargs, - ) + with transformer.cache_context("inference", timestep=t): + block_state.noise_pred, block_state.audio_noise_pred = transformer( + hidden_states=block_state.latents[None], + audio_hidden_states=block_state.audio_latents[None], + encoder_hidden_states=block_state.prompt_embeds, + timestep=unique_timesteps, + timestep_indices=timestep_indices, + attention_kwargs=block_state.attention_kwargs, + return_dict=False, + **layout_kwargs, + ) return components, block_state diff --git a/src/diffusers/modular_pipelines/qwenimage/denoise.py b/src/diffusers/modular_pipelines/qwenimage/denoise.py index 7f271782f82b..f86e825a7bfa 100644 --- a/src/diffusers/modular_pipelines/qwenimage/denoise.py +++ b/src/diffusers/modular_pipelines/qwenimage/denoise.py @@ -157,16 +157,17 @@ def __call__( block_state.cond_scale = controlnet_cond_scale * block_state.controlnet_keep[i] # run controlnet for the guidance batch - controlnet_block_samples = components.controlnet( - hidden_states=block_state.latent_model_input, - controlnet_cond=block_state.control_image_latents, - conditioning_scale=block_state.cond_scale, - timestep=block_state.timestep / 1000, - img_shapes=block_state.img_shapes, - encoder_hidden_states=block_state.prompt_embeds, - encoder_hidden_states_mask=block_state.prompt_embeds_mask, - return_dict=False, - ) + with components.controlnet.cache_context("inference", timestep=t): + controlnet_block_samples = components.controlnet( + hidden_states=block_state.latent_model_input, + controlnet_cond=block_state.control_image_latents, + conditioning_scale=block_state.cond_scale, + timestep=block_state.timestep / 1000, + img_shapes=block_state.img_shapes, + encoder_hidden_states=block_state.prompt_embeds, + encoder_hidden_states_mask=block_state.prompt_embeds_mask, + return_dict=False, + ) block_state.additional_cond_kwargs["controlnet_block_samples"] = controlnet_block_samples @@ -239,15 +240,16 @@ def __call__( components.guider.prepare_models(components.transformer) cond_kwargs = {input_name: getattr(guider_state_batch, input_name) for input_name in guider_inputs.keys()} - # YiYi TODO: add cache context - guider_state_batch.noise_pred = components.transformer( - hidden_states=block_state.latent_model_input, - timestep=block_state.timestep / 1000, - attention_kwargs=block_state.attention_kwargs, - return_dict=False, - **cond_kwargs, - **block_state.additional_cond_kwargs, - )[0] + context_name = getattr(guider_state_batch, components.guider._identifier_key) + 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 / 1000, + attention_kwargs=block_state.attention_kwargs, + return_dict=False, + **cond_kwargs, + **block_state.additional_cond_kwargs, + )[0] components.guider.cleanup_models(components.transformer) @@ -326,15 +328,16 @@ def __call__( components.guider.prepare_models(components.transformer) cond_kwargs = {input_name: getattr(guider_state_batch, input_name) for input_name in guider_inputs.keys()} - # YiYi TODO: add cache context - guider_state_batch.noise_pred = components.transformer( - hidden_states=block_state.latent_model_input, - timestep=block_state.timestep / 1000, - attention_kwargs=block_state.attention_kwargs, - return_dict=False, - **cond_kwargs, - **block_state.additional_cond_kwargs, - )[0] + context_name = getattr(guider_state_batch, components.guider._identifier_key) + 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 / 1000, + attention_kwargs=block_state.attention_kwargs, + return_dict=False, + **cond_kwargs, + **block_state.additional_cond_kwargs, + )[0] components.guider.cleanup_models(components.transformer) diff --git a/src/diffusers/modular_pipelines/wan/denoise.py b/src/diffusers/modular_pipelines/wan/denoise.py index 3036b33868d1..42155eaeb6fb 100644 --- a/src/diffusers/modular_pipelines/wan/denoise.py +++ b/src/diffusers/modular_pipelines/wan/denoise.py @@ -210,13 +210,15 @@ def __call__( # Predict the noise residual # store the noise_pred in guider_state_batch so that we can apply guidance across all batches - guider_state_batch.noise_pred = components.transformer( - hidden_states=block_state.latent_model_input.to(block_state.dtype), - timestep=t.expand(block_state.latent_model_input.shape[0]).to(block_state.dtype), - attention_kwargs=block_state.attention_kwargs, - return_dict=False, - **cond_kwargs, - )[0] + context_name = getattr(guider_state_batch, components.guider._identifier_key) + with components.transformer.cache_context(context_name, timestep=t): + guider_state_batch.noise_pred = components.transformer( + hidden_states=block_state.latent_model_input.to(block_state.dtype), + timestep=t.expand(block_state.latent_model_input.shape[0]).to(block_state.dtype), + attention_kwargs=block_state.attention_kwargs, + return_dict=False, + **cond_kwargs, + )[0] components.guider.cleanup_models(components.transformer) # Perform guidance diff --git a/src/diffusers/pipelines/ace_step/pipeline_ace_step.py b/src/diffusers/pipelines/ace_step/pipeline_ace_step.py index 8f8e2e1be95b..18fd72836b3c 100644 --- a/src/diffusers/pipelines/ace_step/pipeline_ace_step.py +++ b/src/diffusers/pipelines/ace_step/pipeline_ace_step.py @@ -1180,15 +1180,18 @@ def __call__( if apply_cfg: # Batched guidance: stack (cond, null) on batch dim and run the DiT once. # Matches `acestep/models/base/modeling_acestep_v15_base.py:1972-2022`. - model_output = self.transformer( - hidden_states=torch.cat([xt, xt], dim=0), - timestep=torch.cat([t_curr_tensor, t_curr_tensor], dim=0), - timestep_r=torch.cat([t_curr_tensor, t_curr_tensor], dim=0), - encoder_hidden_states=torch.cat([encoder_hidden_states, null_encoder_hidden_states], dim=0), - context_latents=torch.cat([context_latents, context_latents], dim=0), - attention_kwargs=self.attention_kwargs, - return_dict=False, - ) + with self.transformer.cache_context("cfg", timestep=t_sched): + model_output = self.transformer( + hidden_states=torch.cat([xt, xt], dim=0), + timestep=torch.cat([t_curr_tensor, t_curr_tensor], dim=0), + timestep_r=torch.cat([t_curr_tensor, t_curr_tensor], dim=0), + encoder_hidden_states=torch.cat( + [encoder_hidden_states, null_encoder_hidden_states], dim=0 + ), + context_latents=torch.cat([context_latents, context_latents], dim=0), + attention_kwargs=self.attention_kwargs, + return_dict=False, + ) vt_cond, vt_uncond = model_output[0].chunk(2, dim=0) # ACE-Step base / SFT use APG — not vanilla CFG. The original formulation is # `pred_cond + (guidance_scale - 1) * update` with time-only normalization. @@ -1204,28 +1207,30 @@ def __call__( ) else: # Standard forward pass (no CFG) - model_output = self.transformer( - hidden_states=xt, - timestep=t_curr_tensor, - timestep_r=t_curr_tensor, - encoder_hidden_states=encoder_hidden_states, - context_latents=context_latents, - attention_kwargs=self.attention_kwargs, - return_dict=False, - ) + with self.transformer.cache_context("cond", timestep=t_sched): + model_output = self.transformer( + hidden_states=xt, + timestep=t_curr_tensor, + timestep_r=t_curr_tensor, + encoder_hidden_states=encoder_hidden_states, + context_latents=context_latents, + attention_kwargs=self.attention_kwargs, + return_dict=False, + ) vt = model_output[0] # Audio cover strength blending for cover tasks if audio_cover_strength < 1.0 and non_cover_encoder_hidden_states is not None and task_type == "cover": - nc_output = self.transformer( - hidden_states=xt, - timestep=t_curr_tensor, - timestep_r=t_curr_tensor, - encoder_hidden_states=non_cover_encoder_hidden_states, - context_latents=context_latents, - attention_kwargs=self.attention_kwargs, - return_dict=False, - ) + with self.transformer.cache_context("non_cover", timestep=t_sched): + nc_output = self.transformer( + hidden_states=xt, + timestep=t_curr_tensor, + timestep_r=t_curr_tensor, + encoder_hidden_states=non_cover_encoder_hidden_states, + context_latents=context_latents, + attention_kwargs=self.attention_kwargs, + return_dict=False, + ) vt_nc = nc_output[0] # Blend: strength * cover_vt + (1 - strength) * text2music_vt vt = audio_cover_strength * vt + (1.0 - audio_cover_strength) * vt_nc diff --git a/src/diffusers/pipelines/bria/pipeline_bria.py b/src/diffusers/pipelines/bria/pipeline_bria.py index a3f262ce935e..3809996749c0 100644 --- a/src/diffusers/pipelines/bria/pipeline_bria.py +++ b/src/diffusers/pipelines/bria/pipeline_bria.py @@ -679,15 +679,16 @@ def __call__( timestep = t.expand(latent_model_input.shape[0]) # This is predicts "v" from flow-matching or eps from diffusion - noise_pred = self.transformer( - hidden_states=latent_model_input, - timestep=timestep, - encoder_hidden_states=prompt_embeds, - attention_kwargs=self.attention_kwargs, - return_dict=False, - txt_ids=text_ids, - img_ids=latent_image_ids, - )[0] + with self.transformer.cache_context("inference", timestep=t): + noise_pred = self.transformer( + hidden_states=latent_model_input, + timestep=timestep, + encoder_hidden_states=prompt_embeds, + attention_kwargs=self.attention_kwargs, + return_dict=False, + txt_ids=text_ids, + img_ids=latent_image_ids, + )[0] # perform guidance if self.do_classifier_free_guidance: diff --git a/src/diffusers/pipelines/chroma/pipeline_chroma.py b/src/diffusers/pipelines/chroma/pipeline_chroma.py index 37562f4b6a52..d482d750e8b5 100644 --- a/src/diffusers/pipelines/chroma/pipeline_chroma.py +++ b/src/diffusers/pipelines/chroma/pipeline_chroma.py @@ -855,30 +855,32 @@ 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=latents, - timestep=timestep / 1000, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_image_ids, - attention_mask=attention_mask, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - )[0] - - if self.do_classifier_free_guidance: - 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=latents, timestep=timestep / 1000, - encoder_hidden_states=negative_prompt_embeds, - txt_ids=negative_text_ids, + encoder_hidden_states=prompt_embeds, + txt_ids=text_ids, img_ids=latent_image_ids, - attention_mask=negative_attention_mask, + attention_mask=attention_mask, joint_attention_kwargs=self.joint_attention_kwargs, return_dict=False, )[0] + + if self.do_classifier_free_guidance: + 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=latents, + timestep=timestep / 1000, + encoder_hidden_states=negative_prompt_embeds, + txt_ids=negative_text_ids, + img_ids=latent_image_ids, + attention_mask=negative_attention_mask, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + )[0] noise_pred = neg_noise_pred + guidance_scale * (noise_pred - neg_noise_pred) # compute the previous noisy sample x_t -> x_t-1 diff --git a/src/diffusers/pipelines/chroma/pipeline_chroma_img2img.py b/src/diffusers/pipelines/chroma/pipeline_chroma_img2img.py index da6615e71c7f..622dc4d224c0 100644 --- a/src/diffusers/pipelines/chroma/pipeline_chroma_img2img.py +++ b/src/diffusers/pipelines/chroma/pipeline_chroma_img2img.py @@ -937,31 +937,33 @@ def __call__( if image_embeds is not None: self._joint_attention_kwargs["ip_adapter_image_embeds"] = image_embeds - noise_pred = self.transformer( - hidden_states=latents, - timestep=timestep / 1000, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_image_ids, - attention_mask=attention_mask, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - )[0] - - if self.do_classifier_free_guidance: - if negative_image_embeds is not None: - self._joint_attention_kwargs["ip_adapter_image_embeds"] = negative_image_embeds - - noise_pred_uncond = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, - encoder_hidden_states=negative_prompt_embeds, - txt_ids=negative_text_ids, + encoder_hidden_states=prompt_embeds, + txt_ids=text_ids, img_ids=latent_image_ids, - attention_mask=negative_attention_mask, + attention_mask=attention_mask, joint_attention_kwargs=self.joint_attention_kwargs, return_dict=False, )[0] + + if self.do_classifier_free_guidance: + 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): + noise_pred_uncond = self.transformer( + hidden_states=latents, + timestep=timestep / 1000, + encoder_hidden_states=negative_prompt_embeds, + txt_ids=negative_text_ids, + img_ids=latent_image_ids, + attention_mask=negative_attention_mask, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + )[0] noise_pred = noise_pred_uncond + guidance_scale * (noise_pred - noise_pred_uncond) # compute the previous noisy sample x_t -> x_t-1 diff --git a/src/diffusers/pipelines/chroma/pipeline_chroma_inpainting.py b/src/diffusers/pipelines/chroma/pipeline_chroma_inpainting.py index fe95a103cd4f..aaa5e2fb9402 100644 --- a/src/diffusers/pipelines/chroma/pipeline_chroma_inpainting.py +++ b/src/diffusers/pipelines/chroma/pipeline_chroma_inpainting.py @@ -1118,31 +1118,33 @@ def __call__( if image_embeds is not None: self._joint_attention_kwargs["ip_adapter_image_embeds"] = image_embeds - noise_pred = self.transformer( - hidden_states=latents, - timestep=timestep / 1000, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_image_ids, - attention_mask=attention_mask, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - )[0] - - if self.do_classifier_free_guidance: - if negative_image_embeds is not None: - self._joint_attention_kwargs["ip_adapter_image_embeds"] = negative_image_embeds - - noise_pred_uncond = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + noise_pred = self.transformer( hidden_states=latents, timestep=timestep / 1000, - encoder_hidden_states=negative_prompt_embeds, - txt_ids=negative_text_ids, + encoder_hidden_states=prompt_embeds, + txt_ids=text_ids, img_ids=latent_image_ids, - attention_mask=negative_attention_mask, + attention_mask=attention_mask, joint_attention_kwargs=self.joint_attention_kwargs, return_dict=False, )[0] + + if self.do_classifier_free_guidance: + 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): + noise_pred_uncond = self.transformer( + hidden_states=latents, + timestep=timestep / 1000, + encoder_hidden_states=negative_prompt_embeds, + txt_ids=negative_text_ids, + img_ids=latent_image_ids, + attention_mask=negative_attention_mask, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + )[0] noise_pred = noise_pred_uncond + guidance_scale * (noise_pred - noise_pred_uncond) # compute the previous noisy sample x_t -> x_t-1 diff --git a/src/diffusers/pipelines/chronoedit/pipeline_chronoedit.py b/src/diffusers/pipelines/chronoedit/pipeline_chronoedit.py index 95a3ad3c20e3..e19c9da9d36e 100644 --- a/src/diffusers/pipelines/chronoedit/pipeline_chronoedit.py +++ b/src/diffusers/pipelines/chronoedit/pipeline_chronoedit.py @@ -681,24 +681,26 @@ def __call__( latent_model_input = torch.cat([latents, condition], dim=1).to(transformer_dtype) timestep = t.expand(latents.shape[0]) - noise_pred = self.transformer( - hidden_states=latent_model_input, - timestep=timestep, - encoder_hidden_states=prompt_embeds, - encoder_hidden_states_image=image_embeds, - attention_kwargs=attention_kwargs, - return_dict=False, - )[0] - - if self.do_classifier_free_guidance: - noise_uncond = 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_hidden_states=prompt_embeds, encoder_hidden_states_image=image_embeds, attention_kwargs=attention_kwargs, return_dict=False, )[0] + + if self.do_classifier_free_guidance: + with self.transformer.cache_context("uncond", timestep=t): + noise_uncond = self.transformer( + hidden_states=latent_model_input, + timestep=timestep, + encoder_hidden_states=negative_prompt_embeds, + encoder_hidden_states_image=image_embeds, + attention_kwargs=attention_kwargs, + return_dict=False, + )[0] noise_pred = noise_uncond + guidance_scale * (noise_pred - noise_uncond) # compute the previous noisy sample x_t -> x_t-1 diff --git a/src/diffusers/pipelines/cogview4/pipeline_cogview4_control.py b/src/diffusers/pipelines/cogview4/pipeline_cogview4_control.py index ba25c0ef92e6..41fdddaa0c34 100644 --- a/src/diffusers/pipelines/cogview4/pipeline_cogview4_control.py +++ b/src/diffusers/pipelines/cogview4/pipeline_cogview4_control.py @@ -675,22 +675,10 @@ def __call__( # broadcast to batch dimension in a way that's compatible with ONNX/Core ML timestep = t.expand(latents.shape[0]) - noise_pred_cond = self.transformer( - hidden_states=latent_model_input, - encoder_hidden_states=prompt_embeds, - timestep=timestep, - original_size=original_size, - target_size=target_size, - crop_coords=crops_coords_top_left, - attention_kwargs=attention_kwargs, - return_dict=False, - )[0] - - # perform guidance - if self.do_classifier_free_guidance: - noise_pred_uncond = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + noise_pred_cond = self.transformer( hidden_states=latent_model_input, - encoder_hidden_states=negative_prompt_embeds, + encoder_hidden_states=prompt_embeds, timestep=timestep, original_size=original_size, target_size=target_size, @@ -699,6 +687,20 @@ def __call__( return_dict=False, )[0] + # perform guidance + if self.do_classifier_free_guidance: + with self.transformer.cache_context("uncond", timestep=t): + noise_pred_uncond = self.transformer( + hidden_states=latent_model_input, + encoder_hidden_states=negative_prompt_embeds, + timestep=timestep, + original_size=original_size, + target_size=target_size, + crop_coords=crops_coords_top_left, + attention_kwargs=attention_kwargs, + return_dict=False, + )[0] + noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_cond - noise_pred_uncond) else: noise_pred = noise_pred_cond diff --git a/src/diffusers/pipelines/flux/pipeline_flux_control.py b/src/diffusers/pipelines/flux/pipeline_flux_control.py index b50861cb437f..aab656d0a07a 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_control.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_control.py @@ -809,17 +809,18 @@ 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 / 1000, - guidance=guidance, - pooled_projections=pooled_prompt_embeds, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_image_ids, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - )[0] + with self.transformer.cache_context("inference", timestep=t): + 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_image_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + )[0] # compute the previous noisy sample x_t -> x_t-1 latents_dtype = latents.dtype diff --git a/src/diffusers/pipelines/flux/pipeline_flux_control_img2img.py b/src/diffusers/pipelines/flux/pipeline_flux_control_img2img.py index 15d876caf6ca..ba8643e2bead 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_control_img2img.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_control_img2img.py @@ -892,17 +892,18 @@ 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 / 1000, - guidance=guidance, - pooled_projections=pooled_prompt_embeds, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_image_ids, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - )[0] + with self.transformer.cache_context("inference", timestep=t): + 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_image_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + )[0] # compute the previous noisy sample x_t -> x_t-1 latents_dtype = latents.dtype diff --git a/src/diffusers/pipelines/flux/pipeline_flux_control_inpaint.py b/src/diffusers/pipelines/flux/pipeline_flux_control_inpaint.py index 6fa0c1dff9f6..7b1119c168e6 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_control_inpaint.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_control_inpaint.py @@ -1046,17 +1046,18 @@ 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 / 1000, - guidance=guidance, - pooled_projections=pooled_prompt_embeds, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_image_ids, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - )[0] + with self.transformer.cache_context("inference", timestep=t): + 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_image_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + )[0] # compute the previous noisy sample x_t -> x_t-1 latents_dtype = latents.dtype diff --git a/src/diffusers/pipelines/flux/pipeline_flux_controlnet.py b/src/diffusers/pipelines/flux/pipeline_flux_controlnet.py index ab2cebb1e765..a86b699e493a 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_controlnet.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_controlnet.py @@ -1124,30 +1124,13 @@ def __call__( ) guidance = guidance.expand(latents.shape[0]) if guidance is not None else None - noise_pred = self.transformer( - hidden_states=latents, - timestep=timestep / 1000, - guidance=guidance, - pooled_projections=pooled_prompt_embeds, - encoder_hidden_states=prompt_embeds, - controlnet_block_samples=controlnet_block_samples, - controlnet_single_block_samples=controlnet_single_block_samples, - txt_ids=text_ids, - img_ids=latent_image_ids, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - controlnet_blocks_repeat=controlnet_blocks_repeat, - )[0] - - 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=latents, timestep=timestep / 1000, guidance=guidance, - pooled_projections=negative_pooled_prompt_embeds, - encoder_hidden_states=negative_prompt_embeds, + pooled_projections=pooled_prompt_embeds, + encoder_hidden_states=prompt_embeds, controlnet_block_samples=controlnet_block_samples, controlnet_single_block_samples=controlnet_single_block_samples, txt_ids=text_ids, @@ -1156,6 +1139,25 @@ def __call__( return_dict=False, controlnet_blocks_repeat=controlnet_blocks_repeat, )[0] + + 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=latents, + timestep=timestep / 1000, + guidance=guidance, + pooled_projections=negative_pooled_prompt_embeds, + encoder_hidden_states=negative_prompt_embeds, + controlnet_block_samples=controlnet_block_samples, + controlnet_single_block_samples=controlnet_single_block_samples, + txt_ids=text_ids, + img_ids=latent_image_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + controlnet_blocks_repeat=controlnet_blocks_repeat, + )[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/flux/pipeline_flux_controlnet_image_to_image.py b/src/diffusers/pipelines/flux/pipeline_flux_controlnet_image_to_image.py index 482fca3b59c3..4e9968bddd28 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_controlnet_image_to_image.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_controlnet_image_to_image.py @@ -961,20 +961,21 @@ def __call__( ) guidance = guidance.expand(latents.shape[0]) if guidance is not None else None - noise_pred = self.transformer( - hidden_states=latents, - timestep=timestep / 1000, - guidance=guidance, - pooled_projections=pooled_prompt_embeds, - encoder_hidden_states=prompt_embeds, - controlnet_block_samples=controlnet_block_samples, - controlnet_single_block_samples=controlnet_single_block_samples, - txt_ids=text_ids, - img_ids=latent_image_ids, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - controlnet_blocks_repeat=controlnet_blocks_repeat, - )[0] + with self.transformer.cache_context("inference", timestep=t): + noise_pred = self.transformer( + hidden_states=latents, + timestep=timestep / 1000, + guidance=guidance, + pooled_projections=pooled_prompt_embeds, + encoder_hidden_states=prompt_embeds, + controlnet_block_samples=controlnet_block_samples, + controlnet_single_block_samples=controlnet_single_block_samples, + txt_ids=text_ids, + img_ids=latent_image_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + controlnet_blocks_repeat=controlnet_blocks_repeat, + )[0] latents_dtype = latents.dtype latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0] diff --git a/src/diffusers/pipelines/flux/pipeline_flux_controlnet_inpainting.py b/src/diffusers/pipelines/flux/pipeline_flux_controlnet_inpainting.py index 26a434807694..c6cbafa98754 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_controlnet_inpainting.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_controlnet_inpainting.py @@ -1139,20 +1139,21 @@ def __call__( else: guidance = None - noise_pred = self.transformer( - hidden_states=latents, - timestep=timestep / 1000, - guidance=guidance, - pooled_projections=pooled_prompt_embeds, - encoder_hidden_states=prompt_embeds, - controlnet_block_samples=controlnet_block_samples, - controlnet_single_block_samples=controlnet_single_block_samples, - txt_ids=text_ids, - img_ids=latent_image_ids, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - controlnet_blocks_repeat=controlnet_blocks_repeat, - )[0] + with self.transformer.cache_context("inference", timestep=t): + noise_pred = self.transformer( + hidden_states=latents, + timestep=timestep / 1000, + guidance=guidance, + pooled_projections=pooled_prompt_embeds, + encoder_hidden_states=prompt_embeds, + controlnet_block_samples=controlnet_block_samples, + controlnet_single_block_samples=controlnet_single_block_samples, + txt_ids=text_ids, + img_ids=latent_image_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + controlnet_blocks_repeat=controlnet_blocks_repeat, + )[0] # compute the previous noisy sample x_t -> x_t-1 latents_dtype = latents.dtype diff --git a/src/diffusers/pipelines/flux/pipeline_flux_fill.py b/src/diffusers/pipelines/flux/pipeline_flux_fill.py index 55d3464cfb4e..eadab72708a6 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_fill.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_fill.py @@ -958,17 +958,18 @@ 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=torch.cat((latents, masked_image_latents), dim=2), - timestep=timestep / 1000, - guidance=guidance, - pooled_projections=pooled_prompt_embeds, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_image_ids, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - )[0] + with self.transformer.cache_context("inference", timestep=t): + noise_pred = self.transformer( + hidden_states=torch.cat((latents, masked_image_latents), dim=2), + timestep=timestep / 1000, + guidance=guidance, + pooled_projections=pooled_prompt_embeds, + encoder_hidden_states=prompt_embeds, + txt_ids=text_ids, + img_ids=latent_image_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + )[0] # compute the previous noisy sample x_t -> x_t-1 latents_dtype = latents.dtype diff --git a/src/diffusers/pipelines/flux/pipeline_flux_img2img.py b/src/diffusers/pipelines/flux/pipeline_flux_img2img.py index 3288bba94772..ff5797c126f6 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_img2img.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_img2img.py @@ -990,32 +990,34 @@ def __call__( self._joint_attention_kwargs["ip_adapter_image_embeds"] = image_embeds # 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=latents, - timestep=timestep / 1000, - guidance=guidance, - pooled_projections=pooled_prompt_embeds, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_image_ids, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - )[0] - - 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=latents, timestep=timestep / 1000, guidance=guidance, - pooled_projections=negative_pooled_prompt_embeds, - encoder_hidden_states=negative_prompt_embeds, + pooled_projections=pooled_prompt_embeds, + encoder_hidden_states=prompt_embeds, txt_ids=text_ids, img_ids=latent_image_ids, joint_attention_kwargs=self.joint_attention_kwargs, return_dict=False, )[0] + + 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=latents, + timestep=timestep / 1000, + guidance=guidance, + pooled_projections=negative_pooled_prompt_embeds, + encoder_hidden_states=negative_prompt_embeds, + txt_ids=text_ids, + img_ids=latent_image_ids, + joint_attention_kwargs=self.joint_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/flux/pipeline_flux_inpaint.py b/src/diffusers/pipelines/flux/pipeline_flux_inpaint.py index 15d9bb9868fd..5220518758de 100644 --- a/src/diffusers/pipelines/flux/pipeline_flux_inpaint.py +++ b/src/diffusers/pipelines/flux/pipeline_flux_inpaint.py @@ -1146,32 +1146,34 @@ def __call__( self._joint_attention_kwargs["ip_adapter_image_embeds"] = image_embeds # 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=latents, - timestep=timestep / 1000, - guidance=guidance, - pooled_projections=pooled_prompt_embeds, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_image_ids, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - )[0] - - 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=latents, timestep=timestep / 1000, guidance=guidance, - pooled_projections=negative_pooled_prompt_embeds, - encoder_hidden_states=negative_prompt_embeds, + pooled_projections=pooled_prompt_embeds, + encoder_hidden_states=prompt_embeds, txt_ids=text_ids, img_ids=latent_image_ids, joint_attention_kwargs=self.joint_attention_kwargs, return_dict=False, )[0] + + 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=latents, + timestep=timestep / 1000, + guidance=guidance, + pooled_projections=negative_pooled_prompt_embeds, + encoder_hidden_states=negative_prompt_embeds, + txt_ids=text_ids, + img_ids=latent_image_ids, + joint_attention_kwargs=self.joint_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/flux2/pipeline_flux2.py b/src/diffusers/pipelines/flux2/pipeline_flux2.py index 7bb716772e3a..49fb62f1146a 100644 --- a/src/diffusers/pipelines/flux2/pipeline_flux2.py +++ b/src/diffusers/pipelines/flux2/pipeline_flux2.py @@ -971,16 +971,17 @@ 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) - noise_pred = self.transformer( - hidden_states=latent_model_input, # (B, image_seq_len, C) - timestep=timestep / 1000, - guidance=guidance, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, # B, text_seq_len, 4 - img_ids=latent_image_ids, # B, image_seq_len, 4 - joint_attention_kwargs=self.attention_kwargs, - return_dict=False, - )[0] + with self.transformer.cache_context("inference", timestep=t): + noise_pred = self.transformer( + hidden_states=latent_model_input, # (B, image_seq_len, C) + timestep=timestep / 1000, + guidance=guidance, + encoder_hidden_states=prompt_embeds, + txt_ids=text_ids, # B, text_seq_len, 4 + img_ids=latent_image_ids, # B, image_seq_len, 4 + joint_attention_kwargs=self.attention_kwargs, + return_dict=False, + )[0] noise_pred = noise_pred[:, : latents.size(1) :] diff --git a/src/diffusers/pipelines/flux2/pipeline_flux2_klein_kv.py b/src/diffusers/pipelines/flux2/pipeline_flux2_klein_kv.py index 711246db71d5..4671ad0ae3fd 100644 --- a/src/diffusers/pipelines/flux2/pipeline_flux2_klein_kv.py +++ b/src/diffusers/pipelines/flux2/pipeline_flux2_klein_kv.py @@ -796,46 +796,49 @@ def __call__( latent_model_input = torch.cat([image_latents, latents], dim=1).to(self.transformer.dtype) latent_image_ids = torch.cat([image_latent_ids, latent_ids], dim=1) - noise_pred, kv_cache = self.transformer( - hidden_states=latent_model_input, - timestep=timestep / 1000, - guidance=None, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_image_ids, - joint_attention_kwargs=self.attention_kwargs, - return_dict=False, - kv_cache_mode="extract", - num_ref_tokens=image_latents.shape[1], - ) + with self.transformer.cache_context("reference", timestep=t): + noise_pred, kv_cache = self.transformer( + hidden_states=latent_model_input, + timestep=timestep / 1000, + guidance=None, + encoder_hidden_states=prompt_embeds, + txt_ids=text_ids, + img_ids=latent_image_ids, + joint_attention_kwargs=self.attention_kwargs, + return_dict=False, + kv_cache_mode="extract", + num_ref_tokens=image_latents.shape[1], + ) elif kv_cache is not None: # Steps 1+: use cached ref KV, no ref tokens in input - noise_pred = self.transformer( - hidden_states=latents.to(self.transformer.dtype), - timestep=timestep / 1000, - guidance=None, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_ids, - joint_attention_kwargs=self.attention_kwargs, - return_dict=False, - kv_cache=kv_cache, - kv_cache_mode="cached", - )[0] + with self.transformer.cache_context("inference", timestep=t): + noise_pred = self.transformer( + hidden_states=latents.to(self.transformer.dtype), + timestep=timestep / 1000, + guidance=None, + encoder_hidden_states=prompt_embeds, + txt_ids=text_ids, + img_ids=latent_ids, + joint_attention_kwargs=self.attention_kwargs, + return_dict=False, + kv_cache=kv_cache, + kv_cache_mode="cached", + )[0] else: # No reference images: standard forward - noise_pred = self.transformer( - hidden_states=latents.to(self.transformer.dtype), - timestep=timestep / 1000, - guidance=None, - encoder_hidden_states=prompt_embeds, - txt_ids=text_ids, - img_ids=latent_ids, - joint_attention_kwargs=self.attention_kwargs, - return_dict=False, - )[0] + with self.transformer.cache_context("inference", timestep=t): + noise_pred = self.transformer( + hidden_states=latents.to(self.transformer.dtype), + timestep=timestep / 1000, + guidance=None, + encoder_hidden_states=prompt_embeds, + txt_ids=text_ids, + img_ids=latent_ids, + joint_attention_kwargs=self.attention_kwargs, + return_dict=False, + )[0] # compute the previous noisy sample x_t -> x_t-1 latents_dtype = latents.dtype diff --git a/src/diffusers/pipelines/glm_image/pipeline_glm_image.py b/src/diffusers/pipelines/glm_image/pipeline_glm_image.py index 8794e8195771..b559f50fb360 100644 --- a/src/diffusers/pipelines/glm_image/pipeline_glm_image.py +++ b/src/diffusers/pipelines/glm_image/pipeline_glm_image.py @@ -926,24 +926,27 @@ def __call__( split_sizes = prompt_grid_thw.prod(dim=-1).tolist() prior_ids_per_image = torch.split(prompt_prior_ids, split_sizes) # Process each condition image for this sample - for condition_image, condition_image_prior_token_id in zip(prompt_images, prior_ids_per_image): + for condition_image_idx, (condition_image, condition_image_prior_token_id) in enumerate( + zip(prompt_images, prior_ids_per_image) + ): condition_image = condition_image.to(device=device, dtype=prompt_embeds.dtype) condition_latent = retrieve_latents( self.vae.encode(condition_image), generator=generator, sample_mode="argmax" ) condition_latent = (condition_latent - latents_mean) / latents_std - _ = self.transformer( - hidden_states=condition_latent, - encoder_hidden_states=torch.zeros_like(prompt_embeds)[:1, :0, ...], - prior_token_id=condition_image_prior_token_id, - prior_token_drop=torch.full_like(condition_image_prior_token_id, False, dtype=torch.bool), - timestep=torch.zeros((1,), device=device), - target_size=torch.tensor([condition_image.shape[-2:]], device=device), - crop_coords=torch.zeros((1, 2), device=device), - attention_kwargs=attention_kwargs, - kv_caches=kv_caches, - ) + with self.transformer.cache_context(f"reference_{prompt_idx}_{condition_image_idx}", timestep=0): + _ = self.transformer( + hidden_states=condition_latent, + encoder_hidden_states=torch.zeros_like(prompt_embeds)[:1, :0, ...], + prior_token_id=condition_image_prior_token_id, + prior_token_drop=torch.full_like(condition_image_prior_token_id, False, dtype=torch.bool), + timestep=torch.zeros((1,), device=device), + target_size=torch.tensor([condition_image.shape[-2:]], device=device), + crop_coords=torch.zeros((1, 2), device=device), + attention_kwargs=attention_kwargs, + kv_caches=kv_caches, + ) # Move to next sample's cache slot kv_caches.next_sample() @@ -999,28 +1002,12 @@ def __call__( if prior_token_image_ids_per_sample is not None: kv_caches.set_mode("read") - noise_pred_cond = self.transformer( - hidden_states=latent_model_input, - encoder_hidden_states=prompt_embeds, - prior_token_id=prior_token_ids, - prior_token_drop=prior_token_drop_cond, - timestep=timestep, - target_size=target_size, - crop_coords=crops_coords_top_left, - attention_kwargs=attention_kwargs, - return_dict=False, - kv_caches=kv_caches, - )[0].float() - - # perform guidance - if self.do_classifier_free_guidance: - if prior_token_image_ids_per_sample is not None: - kv_caches.set_mode("skip") - noise_pred_uncond = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + noise_pred_cond = self.transformer( hidden_states=latent_model_input, - encoder_hidden_states=negative_prompt_embeds, + encoder_hidden_states=prompt_embeds, prior_token_id=prior_token_ids, - prior_token_drop=prior_token_drop_uncond, + prior_token_drop=prior_token_drop_cond, timestep=timestep, target_size=target_size, crop_coords=crops_coords_top_left, @@ -1029,6 +1016,24 @@ def __call__( kv_caches=kv_caches, )[0].float() + # perform guidance + if self.do_classifier_free_guidance: + if prior_token_image_ids_per_sample is not None: + kv_caches.set_mode("skip") + with self.transformer.cache_context("uncond", timestep=t): + noise_pred_uncond = self.transformer( + hidden_states=latent_model_input, + encoder_hidden_states=negative_prompt_embeds, + prior_token_id=prior_token_ids, + prior_token_drop=prior_token_drop_uncond, + timestep=timestep, + target_size=target_size, + crop_coords=crops_coords_top_left, + attention_kwargs=attention_kwargs, + return_dict=False, + kv_caches=kv_caches, + )[0].float() + noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_cond - noise_pred_uncond) else: noise_pred = noise_pred_cond diff --git a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky.py b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky.py index 1ad08729e8eb..ec5171978bad 100644 --- a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky.py +++ b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky.py @@ -882,31 +882,33 @@ def __call__( timestep = t.unsqueeze(0).repeat(batch_size * num_videos_per_prompt) # Predict noise residual - pred_velocity = self.transformer( - hidden_states=latents.to(dtype), - encoder_hidden_states=prompt_embeds_qwen.to(dtype), - pooled_projections=prompt_embeds_clip.to(dtype), - timestep=timestep.to(dtype), - visual_rope_pos=visual_rope_pos, - text_rope_pos=text_rope_pos, - scale_factor=scale_factor, - sparse_params=sparse_params, - return_dict=True, - ).sample - - if self.guidance_scale > 1.0 and negative_prompt_embeds_qwen is not None: - uncond_pred_velocity = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + pred_velocity = self.transformer( hidden_states=latents.to(dtype), - encoder_hidden_states=negative_prompt_embeds_qwen.to(dtype), - pooled_projections=negative_prompt_embeds_clip.to(dtype), + encoder_hidden_states=prompt_embeds_qwen.to(dtype), + pooled_projections=prompt_embeds_clip.to(dtype), timestep=timestep.to(dtype), visual_rope_pos=visual_rope_pos, - text_rope_pos=negative_text_rope_pos, + text_rope_pos=text_rope_pos, scale_factor=scale_factor, sparse_params=sparse_params, return_dict=True, ).sample + if self.guidance_scale > 1.0 and negative_prompt_embeds_qwen is not None: + with self.transformer.cache_context("uncond", timestep=t): + uncond_pred_velocity = self.transformer( + hidden_states=latents.to(dtype), + encoder_hidden_states=negative_prompt_embeds_qwen.to(dtype), + pooled_projections=negative_prompt_embeds_clip.to(dtype), + timestep=timestep.to(dtype), + visual_rope_pos=visual_rope_pos, + text_rope_pos=negative_text_rope_pos, + scale_factor=scale_factor, + sparse_params=sparse_params, + return_dict=True, + ).sample + pred_velocity = uncond_pred_velocity + guidance_scale * (pred_velocity - uncond_pred_velocity) # Compute previous sample using the scheduler latents[:, :, :, :, :num_channels_latents] = self.scheduler.step( diff --git a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2i.py b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2i.py index 9fa3378d0143..587e9c829778 100644 --- a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2i.py +++ b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2i.py @@ -770,31 +770,33 @@ def __call__( timestep = t.unsqueeze(0).repeat(batch_size * num_images_per_prompt) # Predict noise residual - pred_velocity = self.transformer( - hidden_states=latents.to(dtype), - encoder_hidden_states=prompt_embeds_qwen.to(dtype), - pooled_projections=prompt_embeds_clip.to(dtype), - timestep=timestep.to(dtype), - visual_rope_pos=visual_rope_pos, - text_rope_pos=text_rope_pos, - scale_factor=scale_factor, - sparse_params=sparse_params, - return_dict=True, - ).sample - - if self.guidance_scale > 1.0 and negative_prompt_embeds_qwen is not None: - uncond_pred_velocity = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + pred_velocity = self.transformer( hidden_states=latents.to(dtype), - encoder_hidden_states=negative_prompt_embeds_qwen.to(dtype), - pooled_projections=negative_prompt_embeds_clip.to(dtype), + encoder_hidden_states=prompt_embeds_qwen.to(dtype), + pooled_projections=prompt_embeds_clip.to(dtype), timestep=timestep.to(dtype), visual_rope_pos=visual_rope_pos, - text_rope_pos=negative_text_rope_pos, + text_rope_pos=text_rope_pos, scale_factor=scale_factor, sparse_params=sparse_params, return_dict=True, ).sample + if self.guidance_scale > 1.0 and negative_prompt_embeds_qwen is not None: + with self.transformer.cache_context("uncond", timestep=t): + uncond_pred_velocity = self.transformer( + hidden_states=latents.to(dtype), + encoder_hidden_states=negative_prompt_embeds_qwen.to(dtype), + pooled_projections=negative_prompt_embeds_clip.to(dtype), + timestep=timestep.to(dtype), + visual_rope_pos=visual_rope_pos, + text_rope_pos=negative_text_rope_pos, + scale_factor=scale_factor, + sparse_params=sparse_params, + return_dict=True, + ).sample + pred_velocity = uncond_pred_velocity + guidance_scale * (pred_velocity - uncond_pred_velocity) latents[:, :, :, :, :num_channels_latents] = self.scheduler.step( diff --git a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2v.py b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2v.py index 34099f191891..cf5edf7721f7 100644 --- a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2v.py +++ b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_i2v.py @@ -956,31 +956,33 @@ def __call__( timestep = t.unsqueeze(0).repeat(batch_size * num_videos_per_prompt) # Predict noise residual - pred_velocity = self.transformer( - hidden_states=latents.to(dtype), - encoder_hidden_states=prompt_embeds_qwen.to(dtype), - pooled_projections=prompt_embeds_clip.to(dtype), - timestep=timestep.to(dtype), - visual_rope_pos=visual_rope_pos, - text_rope_pos=text_rope_pos, - scale_factor=scale_factor, - sparse_params=sparse_params, - return_dict=True, - ).sample - - if self.guidance_scale > 1.0 and negative_prompt_embeds_qwen is not None: - uncond_pred_velocity = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + pred_velocity = self.transformer( hidden_states=latents.to(dtype), - encoder_hidden_states=negative_prompt_embeds_qwen.to(dtype), - pooled_projections=negative_prompt_embeds_clip.to(dtype), + encoder_hidden_states=prompt_embeds_qwen.to(dtype), + pooled_projections=prompt_embeds_clip.to(dtype), timestep=timestep.to(dtype), visual_rope_pos=visual_rope_pos, - text_rope_pos=negative_text_rope_pos, + text_rope_pos=text_rope_pos, scale_factor=scale_factor, sparse_params=sparse_params, return_dict=True, ).sample + if self.guidance_scale > 1.0 and negative_prompt_embeds_qwen is not None: + with self.transformer.cache_context("uncond", timestep=t): + uncond_pred_velocity = self.transformer( + hidden_states=latents.to(dtype), + encoder_hidden_states=negative_prompt_embeds_qwen.to(dtype), + pooled_projections=negative_prompt_embeds_clip.to(dtype), + timestep=timestep.to(dtype), + visual_rope_pos=visual_rope_pos, + text_rope_pos=negative_text_rope_pos, + scale_factor=scale_factor, + sparse_params=sparse_params, + return_dict=True, + ).sample + pred_velocity = uncond_pred_velocity + guidance_scale * (pred_velocity - uncond_pred_velocity) latents[:, 1:, :, :, :num_channels_latents] = self.scheduler.step( diff --git a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_t2i.py b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_t2i.py index 46002e086a28..0ca8102f0170 100644 --- a/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_t2i.py +++ b/src/diffusers/pipelines/kandinsky5/pipeline_kandinsky_t2i.py @@ -726,31 +726,33 @@ def __call__( timestep = t.unsqueeze(0).repeat(batch_size * num_images_per_prompt) # Predict noise residual - pred_velocity = self.transformer( - hidden_states=latents.to(dtype), - encoder_hidden_states=prompt_embeds_qwen.to(dtype), - pooled_projections=prompt_embeds_clip.to(dtype), - timestep=timestep.to(dtype), - visual_rope_pos=visual_rope_pos, - text_rope_pos=text_rope_pos, - scale_factor=scale_factor, - sparse_params=sparse_params, - return_dict=True, - ).sample - - if self.guidance_scale > 1.0 and negative_prompt_embeds_qwen is not None: - uncond_pred_velocity = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + pred_velocity = self.transformer( hidden_states=latents.to(dtype), - encoder_hidden_states=negative_prompt_embeds_qwen.to(dtype), - pooled_projections=negative_prompt_embeds_clip.to(dtype), + encoder_hidden_states=prompt_embeds_qwen.to(dtype), + pooled_projections=prompt_embeds_clip.to(dtype), timestep=timestep.to(dtype), visual_rope_pos=visual_rope_pos, - text_rope_pos=negative_text_rope_pos, + text_rope_pos=text_rope_pos, scale_factor=scale_factor, sparse_params=sparse_params, return_dict=True, ).sample + if self.guidance_scale > 1.0 and negative_prompt_embeds_qwen is not None: + with self.transformer.cache_context("uncond", timestep=t): + uncond_pred_velocity = self.transformer( + hidden_states=latents.to(dtype), + encoder_hidden_states=negative_prompt_embeds_qwen.to(dtype), + pooled_projections=negative_prompt_embeds_clip.to(dtype), + timestep=timestep.to(dtype), + visual_rope_pos=visual_rope_pos, + text_rope_pos=negative_text_rope_pos, + scale_factor=scale_factor, + sparse_params=sparse_params, + return_dict=True, + ).sample + pred_velocity = uncond_pred_velocity + guidance_scale * (pred_velocity - uncond_pred_velocity) latents = self.scheduler.step(pred_velocity[:, :], t, latents, return_dict=False)[0] diff --git a/src/diffusers/pipelines/kandinsky6/pipeline_kandinsky6_ti2va.py b/src/diffusers/pipelines/kandinsky6/pipeline_kandinsky6_ti2va.py index 656f2077c47c..6226855c31bf 100644 --- a/src/diffusers/pipelines/kandinsky6/pipeline_kandinsky6_ti2va.py +++ b/src/diffusers/pipelines/kandinsky6/pipeline_kandinsky6_ti2va.py @@ -886,7 +886,7 @@ def __call__( cond_mask[:, -1] = 1 latent_model_input = torch.cat([latents, cond_latents, cond_mask], 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, audio_hidden_states=audio_latents, @@ -899,7 +899,7 @@ def __call__( return_dict=False, ) 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, audio_hidden_states=audio_latents, diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr.py index 936af16d8805..2cf21ef58ea9 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr.py @@ -1323,29 +1323,31 @@ def denoise( } if video_tile_plan is None: - noise_pred_video, noise_pred_audio = self.transformer( - hidden_states=latents.to(prompt_embeds.dtype), - audio_hidden_states=audio_latents.to(prompt_embeds.dtype), - timestep=video_timestep, - video_keyframes_mask=keyframes_mask, - video_coords=video_coords, - **transformer_kwargs, - ) + with self.transformer.cache_context("inference", timestep=t): + noise_pred_video, noise_pred_audio = self.transformer( + hidden_states=latents.to(prompt_embeds.dtype), + audio_hidden_states=audio_latents.to(prompt_embeds.dtype), + timestep=video_timestep, + video_keyframes_mask=keyframes_mask, + video_coords=video_coords, + **transformer_kwargs, + ) else: # Blending the velocity is the same as blending x0: `x0 = x - v * sigma` is affine in `v`, every tile # reads the same `latents` and the same per-token sigma, and the weights sum to one. noise_pred_video = torch.zeros_like(latents, dtype=torch.float32) noise_pred_audio = None - for tile in video_tile_plan: + for tile_idx, tile in enumerate(video_tile_plan): keep = tile.keep - tile_pred, tile_audio_pred = self.transformer( - hidden_states=latents[:, keep].to(prompt_embeds.dtype), - audio_hidden_states=audio_latents.to(prompt_embeds.dtype), - timestep=video_timestep[:, keep], - video_keyframes_mask=keyframes_mask[:, keep], - video_coords=tile.coords, - **transformer_kwargs, - ) + with self.transformer.cache_context(f"tile_{tile_idx}", timestep=t): + tile_pred, tile_audio_pred = self.transformer( + hidden_states=latents[:, keep].to(prompt_embeds.dtype), + audio_hidden_states=audio_latents.to(prompt_embeds.dtype), + timestep=video_timestep[:, keep], + video_keyframes_mask=keyframes_mask[:, keep], + video_coords=tile.coords, + **transformer_kwargs, + ) noise_pred_video.index_add_( 1, keep, tile_pred.float() * tile.weights.to(torch.float32).view(1, -1, 1) ) diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr_temporal_refine.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr_temporal_refine.py index bdf9c8e67677..9161d50bfcde 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr_temporal_refine.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_dfr_temporal_refine.py @@ -1236,29 +1236,31 @@ def denoise( } if video_tile_plan is None: - noise_pred_video, noise_pred_audio = self.transformer( - hidden_states=latents.to(prompt_embeds.dtype), - audio_hidden_states=audio_latents.to(prompt_embeds.dtype), - timestep=video_timestep, - video_keyframes_mask=keyframes_mask, - video_coords=video_coords, - **transformer_kwargs, - ) + with self.transformer.cache_context("inference", timestep=t): + noise_pred_video, noise_pred_audio = self.transformer( + hidden_states=latents.to(prompt_embeds.dtype), + audio_hidden_states=audio_latents.to(prompt_embeds.dtype), + timestep=video_timestep, + video_keyframes_mask=keyframes_mask, + video_coords=video_coords, + **transformer_kwargs, + ) else: # Blending the velocity is the same as blending x0: `x0 = x - v * sigma` is affine in `v`, every tile # reads the same `latents` and the same per-token sigma, and the weights sum to one. noise_pred_video = torch.zeros_like(latents, dtype=torch.float32) noise_pred_audio = None - for tile in video_tile_plan: + for tile_idx, tile in enumerate(video_tile_plan): keep = tile.keep - tile_pred, tile_audio_pred = self.transformer( - hidden_states=latents[:, keep].to(prompt_embeds.dtype), - audio_hidden_states=audio_latents.to(prompt_embeds.dtype), - timestep=video_timestep[:, keep], - video_keyframes_mask=keyframes_mask[:, keep], - video_coords=tile.coords, - **transformer_kwargs, - ) + with self.transformer.cache_context(f"tile_{tile_idx}", timestep=t): + tile_pred, tile_audio_pred = self.transformer( + hidden_states=latents[:, keep].to(prompt_embeds.dtype), + audio_hidden_states=audio_latents.to(prompt_embeds.dtype), + timestep=video_timestep[:, keep], + video_keyframes_mask=keyframes_mask[:, keep], + video_coords=tile.coords, + **transformer_kwargs, + ) noise_pred_video.index_add_( 1, keep, tile_pred.float() * tile.weights.to(torch.float32).view(1, -1, 1) ) diff --git a/src/diffusers/pipelines/nucleusmoe_image/pipeline_nucleusmoe_image.py b/src/diffusers/pipelines/nucleusmoe_image/pipeline_nucleusmoe_image.py index 71b7a82de4ea..7e17964b4ef9 100644 --- a/src/diffusers/pipelines/nucleusmoe_image/pipeline_nucleusmoe_image.py +++ b/src/diffusers/pipelines/nucleusmoe_image/pipeline_nucleusmoe_image.py @@ -570,27 +570,29 @@ def __call__( self._current_timestep = t timestep = t.expand(latents.shape[0]).to(latents.dtype) - noise_pred = self.transformer( - hidden_states=latents, - timestep=timestep / self.scheduler.config.num_train_timesteps, - encoder_hidden_states=prompt_embeds, - encoder_hidden_states_mask=prompt_embeds_mask, - img_shapes=img_shapes, - attention_kwargs=self._attention_kwargs, - return_dict=False, - )[0] - - if do_cfg: - neg_noise_pred = self.transformer( + with self.transformer.cache_context("cond", timestep=t): + noise_pred = self.transformer( hidden_states=latents, timestep=timestep / self.scheduler.config.num_train_timesteps, - encoder_hidden_states=negative_prompt_embeds, - encoder_hidden_states_mask=negative_prompt_embeds_mask, + encoder_hidden_states=prompt_embeds, + encoder_hidden_states_mask=prompt_embeds_mask, img_shapes=img_shapes, attention_kwargs=self._attention_kwargs, return_dict=False, )[0] + if do_cfg: + with self.transformer.cache_context("uncond", timestep=t): + neg_noise_pred = self.transformer( + hidden_states=latents, + timestep=timestep / self.scheduler.config.num_train_timesteps, + encoder_hidden_states=negative_prompt_embeds, + encoder_hidden_states_mask=negative_prompt_embeds_mask, + img_shapes=img_shapes, + attention_kwargs=self._attention_kwargs, + return_dict=False, + )[0] + comb_pred = neg_noise_pred + guidance_scale * (noise_pred - neg_noise_pred) cond_norm = torch.norm(noise_pred, dim=-1, keepdim=True) noise_norm = torch.norm(comb_pred, dim=-1, keepdim=True) diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py index d48151b3d313..aa966bc9d5db 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py @@ -872,16 +872,17 @@ def __call__( cond_scale = controlnet_cond_scale * controlnet_keep[i] # controlnet - controlnet_block_samples = self.controlnet( - hidden_states=latents, - controlnet_cond=control_image, - conditioning_scale=cond_scale, - timestep=timestep / 1000, - encoder_hidden_states=prompt_embeds, - encoder_hidden_states_mask=prompt_embeds_mask, - img_shapes=img_shapes, - return_dict=False, - ) + with self.controlnet.cache_context("cond", timestep=t): + controlnet_block_samples = self.controlnet( + hidden_states=latents, + controlnet_cond=control_image, + conditioning_scale=cond_scale, + timestep=timestep / 1000, + encoder_hidden_states=prompt_embeds, + encoder_hidden_states_mask=prompt_embeds_mask, + img_shapes=img_shapes, + return_dict=False, + ) with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( diff --git a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py index 4b4fa7587d3b..2aec58dec7b2 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py @@ -843,16 +843,17 @@ def __call__( cond_scale = controlnet_cond_scale * controlnet_keep[i] # controlnet - controlnet_block_samples = self.controlnet( - hidden_states=latents, - controlnet_cond=control_image.to(dtype=latents.dtype, device=device), - conditioning_scale=cond_scale, - timestep=timestep / 1000, - encoder_hidden_states=prompt_embeds, - encoder_hidden_states_mask=prompt_embeds_mask, - img_shapes=img_shapes, - return_dict=False, - ) + with self.controlnet.cache_context("cond", timestep=t): + controlnet_block_samples = self.controlnet( + hidden_states=latents, + controlnet_cond=control_image.to(dtype=latents.dtype, device=device), + conditioning_scale=cond_scale, + timestep=timestep / 1000, + encoder_hidden_states=prompt_embeds, + encoder_hidden_states_mask=prompt_embeds_mask, + img_shapes=img_shapes, + return_dict=False, + ) with self.transformer.cache_context("cond", timestep=t): noise_pred = self.transformer( diff --git a/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_generation.py b/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_generation.py index cd6e61cff1df..ed4eb25d229a 100644 --- a/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_generation.py +++ b/src/diffusers/pipelines/visualcloze/pipeline_visualcloze_generation.py @@ -853,17 +853,18 @@ def __call__( timestep = t.expand(latents.shape[0]).to(latents.dtype) latent_model_input = torch.cat((latents, masked_image_latents), dim=2) - 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_image_ids, - joint_attention_kwargs=self.joint_attention_kwargs, - return_dict=False, - )[0] + with self.transformer.cache_context("inference", timestep=t): + 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_image_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + )[0] # Compute the previous noisy sample x_t -> x_t-1 latents_dtype = latents.dtype diff --git a/src/diffusers/pipelines/wan/pipeline_wan_video2video.py b/src/diffusers/pipelines/wan/pipeline_wan_video2video.py index b192147acb64..45d56cffa249 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan_video2video.py +++ b/src/diffusers/pipelines/wan/pipeline_wan_video2video.py @@ -678,22 +678,24 @@ def __call__( latent_model_input = latents.to(transformer_dtype) timestep = t.expand(latents.shape[0]) - noise_pred = self.transformer( - hidden_states=latent_model_input, - timestep=timestep, - encoder_hidden_states=prompt_embeds, - attention_kwargs=attention_kwargs, - return_dict=False, - )[0] - - if self.do_classifier_free_guidance: - noise_uncond = 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_hidden_states=prompt_embeds, attention_kwargs=attention_kwargs, return_dict=False, )[0] + + if self.do_classifier_free_guidance: + with self.transformer.cache_context("uncond", timestep=t): + noise_uncond = self.transformer( + hidden_states=latent_model_input, + timestep=timestep, + encoder_hidden_states=negative_prompt_embeds, + attention_kwargs=attention_kwargs, + return_dict=False, + )[0] noise_pred = noise_uncond + guidance_scale * (noise_pred - noise_uncond) # compute the previous noisy sample x_t -> x_t-1