diff --git a/docs/source/en/optimization/cache.md b/docs/source/en/optimization/cache.md index 3a771336e4b7..89095c6e4be3 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..e18e72980fce 100644 --- a/src/diffusers/hooks/faster_cache.py +++ b/src/diffusers/hooks/faster_cache.py @@ -20,9 +20,9 @@ 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 +from .hooks import CacheContext, HookRegistry, ModelHook, StateManager logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -160,11 +160,23 @@ class FasterCacheConfig: tensor_format: str = "BCFHW" is_guidance_distilled: bool = False - current_timestep_callback: Callable[[], int] = None + 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. " + "Please use `cache_context(name, timestep=...)` to pass the current timestep instead. " + "See the caching documentation for guidance: https://huggingface.co/docs/diffusers/main/en/optimization/cache." + ) + deprecate("current_timestep_callback", "0.45.0", depr_message) + def __repr__(self) -> str: + callback_repr = "" + if self.current_timestep_callback is not None: + callback_repr = f" current_timestep_callback={self.current_timestep_callback},\n" return ( f"FasterCacheConfig(\n" f" spatial_attention_block_skip_range={self.spatial_attention_block_skip_range},\n" @@ -180,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")" ) @@ -227,9 +240,9 @@ 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], + current_timestep_callback: Callable[[], int] | None = None, ) -> None: super().__init__() @@ -248,7 +261,7 @@ def __init__( 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 @@ -258,7 +271,19 @@ def _get_cond_input(input: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: _, cond = input.chunk(2, dim=0) return cond + 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. # We skip the unconditional branch only if the following conditions are met: @@ -270,13 +295,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 +318,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 +329,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 +340,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 +368,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 +382,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 +395,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], + current_timestep_callback: Callable[[], int] | None = None, ) -> None: super().__init__() @@ -386,7 +407,7 @@ def __init__( 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 +426,21 @@ 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() 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 +448,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 +473,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 +487,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 +495,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 +568,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)._get_timestep() < config.low_frequency_weight_update_timestep_range[1] ) return config.alpha_low_frequency if is_within_range else 1.0 @@ -554,7 +583,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)._get_timestep() < config.high_frequency_weight_update_timestep_range[1] ) return config.alpha_high_frequency if is_within_range else 1.0 @@ -581,9 +610,9 @@ 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, + current_timestep_callback=config.current_timestep_callback, ) registry = HookRegistry.check_if_exists_or_initialize(module) registry.register_hook(hook, _FASTER_CACHE_DENOISER_HOOK) @@ -627,7 +656,7 @@ 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, + 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 e7ed26b28778..aed20a7ce409 100644 --- a/src/diffusers/hooks/pyramid_attention_broadcast.py +++ b/src/diffusers/hooks/pyramid_attention_broadcast.py @@ -20,14 +20,14 @@ from ..models.attention import AttentionModuleMixin from ..models.attention_processor import Attention, MochiAttention -from ..utils import logging +from ..utils import deprecate, logging from ._common import ( _ATTENTION_CLASSES, _CROSS_TRANSFORMER_BLOCK_IDENTIFIERS, _SPATIAL_TRANSFORMER_BLOCK_IDENTIFIERS, _TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS, ) -from .hooks import HookRegistry, ModelHook +from .hooks import CacheContext, HookRegistry, ModelHook, StateManager logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -83,12 +83,24 @@ class PyramidAttentionBroadcastConfig: temporal_attention_block_identifiers: tuple[str, ...] = _TEMPORAL_TRANSFORMER_BLOCK_IDENTIFIERS cross_attention_block_identifiers: tuple[str, ...] = _CROSS_TRANSFORMER_BLOCK_IDENTIFIERS - current_timestep_callback: Callable[[], int] = None + current_timestep_callback: Callable[[], int] | None = None # TODO(aryan): add PAB for MLP layers (very limited speedup from testing with original codebase # so not added for now) + def __post_init__(self): + if self.current_timestep_callback is not None: + depr_message = ( + "Passing `current_timestep_callback` to `PyramidAttentionBroadcastConfig` is deprecated. " + "Please use `cache_context(name, timestep=...)` to pass the current timestep instead. " + "See the caching documentation for guidance: https://huggingface.co/docs/diffusers/main/en/optimization/cache." + ) + deprecate("current_timestep_callback", "0.45.0", depr_message) + def __repr__(self) -> str: + callback_repr = "" + if self.current_timestep_callback is not None: + callback_repr = f" current_timestep_callback={self.current_timestep_callback},\n" return ( f"PyramidAttentionBroadcastConfig(\n" f" spatial_attention_block_skip_range={self.spatial_attention_block_skip_range},\n" @@ -100,7 +112,7 @@ def __repr__(self) -> str: f" spatial_attention_block_identifiers={self.spatial_attention_block_identifiers},\n" f" temporal_attention_block_identifiers={self.temporal_attention_block_identifiers},\n" f" cross_attention_block_identifiers={self.cross_attention_block_identifiers},\n" - f" current_timestep_callback={self.current_timestep_callback}\n" + f"{callback_repr}" ")" ) @@ -141,7 +153,10 @@ class PyramidAttentionBroadcastHook(ModelHook): _is_stateful = True def __init__( - self, timestep_skip_range: tuple[int, int], block_skip_range: int, current_timestep_callback: Callable[[], int] + self, + timestep_skip_range: tuple[int, int], + block_skip_range: int, + current_timestep_callback: Callable[[], int] | None = None, ) -> None: super().__init__() @@ -150,31 +165,37 @@ def __init__( self.current_timestep_callback = current_timestep_callback def initialize_hook(self, module): - self.state = PyramidAttentionBroadcastState() + self.state_manager = StateManager(PyramidAttentionBroadcastState) return module def new_forward(self, module: torch.nn.Module, *args, **kwargs) -> Any: - is_within_timestep_range = ( - self.timestep_skip_range[0] < self.current_timestep_callback() < self.timestep_skip_range[1] - ) + if self.state_manager._context is None and self.current_timestep_callback is not None: + self.state_manager.set_context(CacheContext(name="inference")) + timestep = self.state_manager.context.timestep + if timestep is None and self.current_timestep_callback is not None: + timestep = self.current_timestep_callback() + if timestep is None: + raise ValueError("Pyramid Attention Broadcast requires `cache_context(name, timestep=...)`.") + state = self.state_manager.get_state() + is_within_timestep_range = self.timestep_skip_range[0] < timestep < self.timestep_skip_range[1] should_compute_attention = ( - self.state.cache is None - or self.state.iteration == 0 + state.cache is None + or state.iteration == 0 or not is_within_timestep_range - or self.state.iteration % self.block_skip_range == 0 + or state.iteration % self.block_skip_range == 0 ) if should_compute_attention: output = self.fn_ref.original_forward(*args, **kwargs) else: - output = self.state.cache + output = state.cache - self.state.cache = output - self.state.iteration += 1 + state.cache = output + state.iteration += 1 return output def reset_state(self, module: torch.nn.Module) -> None: - self.state.reset() + self.state_manager.reset() return module @@ -207,16 +228,10 @@ def apply_pyramid_attention_broadcast(module: torch.nn.Module, config: PyramidAt >>> config = PyramidAttentionBroadcastConfig( ... spatial_attention_block_skip_range=2, ... spatial_attention_timestep_skip_range=(100, 800), - ... current_timestep_callback=lambda: pipe.current_timestep, ... ) >>> apply_pyramid_attention_broadcast(pipe.transformer, config) ``` """ - if config.current_timestep_callback is None: - raise ValueError( - "The `current_timestep_callback` function must be provided in the configuration to apply Pyramid Attention Broadcast." - ) - if ( config.spatial_attention_block_skip_range is None and config.temporal_attention_block_skip_range is None @@ -291,7 +306,7 @@ def _apply_pyramid_attention_broadcast_hook( module: Attention | MochiAttention, timestep_skip_range: tuple[int, int], block_skip_range: int, - current_timestep_callback: Callable[[], int], + current_timestep_callback: Callable[[], int] | None = None, ): r""" Apply [Pyramid Attention Broadcast](https://huggingface.co/papers/2408.12588) to a given torch.nn.Module. @@ -306,8 +321,6 @@ def _apply_pyramid_attention_broadcast_hook( The number of times a specific attention broadcast is skipped before computing the attention states to re-use. If this is set to the value `N`, the attention computation will be skipped `N - 1` times (i.e., old attention states will be reused) before computing the new attention states again. - current_timestep_callback (`Callable[[], int]`): - A callback function that returns the current inference timestep. """ registry = HookRegistry.check_if_exists_or_initialize(module) hook = PyramidAttentionBroadcastHook(timestep_skip_range, block_skip_range, current_timestep_callback) 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 294213f48f4b..81457b1a9830 100644 --- a/src/diffusers/modular_pipelines/cosmos/denoise.py +++ b/src/diffusers/modular_pipelines/cosmos/denoise.py @@ -233,6 +233,7 @@ def __call__( } with components.transformer.cache_context( pass_name, + timestep=t, step_index=i, sigma=float(components.scheduler.sigmas[i]), num_inference_steps=components.scheduler.num_inference_steps, @@ -797,9 +798,11 @@ def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam("velocity", type_hint=torch.Tensor, description="Predicted (masked) transfer velocity.")] @staticmethod - def _forward(components, static, vision_tokens, vision_timesteps, context_name, step, sigma, num_inference_steps): + def _forward( + components, static, vision_tokens, vision_timesteps, context_name, step, timestep, sigma, num_inference_steps + ): with components.transformer.cache_context( - context_name, step_index=step, sigma=sigma, num_inference_steps=num_inference_steps + context_name, step_index=step, timestep=timestep, sigma=sigma, num_inference_steps=num_inference_steps ): preds_vision, _, _ = components.transformer( input_ids=static["input_ids"], @@ -847,6 +850,7 @@ def __call__( block_state.vision_timesteps, "cond", step=i, + timestep=t, sigma=float(components.scheduler.sigmas[i]), num_inference_steps=components.scheduler.num_inference_steps, ) @@ -860,6 +864,7 @@ def __call__( block_state.vision_timesteps, "cond_no_control", step=i, + timestep=t, sigma=float(components.scheduler.sigmas[i]), num_inference_steps=components.scheduler.num_inference_steps, ) @@ -873,6 +878,7 @@ def __call__( block_state.vision_timesteps, "uncond", step=i, + timestep=t, sigma=float(components.scheduler.sigmas[i]), num_inference_steps=components.scheduler.num_inference_steps, ) 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/helios/denoise.py b/src/diffusers/modular_pipelines/helios/denoise.py index 32a6f0b11b79..2e5a3f274610 100644 --- a/src/diffusers/modular_pipelines/helios/denoise.py +++ b/src/diffusers/modular_pipelines/helios/denoise.py @@ -443,7 +443,7 @@ def __call__( cond_kwargs = {k: getattr(guider_state_batch, k) for k in guider_inputs.keys()} context_name = getattr(guider_state_batch, components.guider._identifier_key) - with components.transformer.cache_context(context_name): + with components.transformer.cache_context(context_name, timestep=t): guider_state_batch.noise_pred = components.transformer( hidden_states=latent_model_input, timestep=timestep, @@ -635,7 +635,7 @@ def __call__( cond_kwargs = {kk: getattr(guider_state_batch, kk) for kk in guider_inputs.keys()} context_name = getattr(guider_state_batch, components.guider._identifier_key) - with components.transformer.cache_context(context_name): + with components.transformer.cache_context(context_name, timestep=t): guider_state_batch.noise_pred = components.transformer( hidden_states=latent_model_input, timestep=timestep, @@ -976,7 +976,7 @@ def __call__( cond_kwargs = {k: getattr(guider_state_batch, k) for k in guider_inputs.keys()} context_name = getattr(guider_state_batch, components.guider._identifier_key) - with components.transformer.cache_context(context_name): + with components.transformer.cache_context(context_name, timestep=t): guider_state_batch.noise_pred = components.transformer( hidden_states=latent_model_input, timestep=timestep, diff --git a/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py b/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py index 223d5e2eee9b..523b08a39c82 100644 --- a/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py +++ b/src/diffusers/modular_pipelines/hunyuan_video1_5/denoise.py @@ -157,7 +157,7 @@ def __call__( cond_kwargs = {input_name: getattr(guider_state_batch, input_name) for input_name in guider_inputs.keys()} context_name = getattr(guider_state_batch, components.guider._identifier_key) - with components.transformer.cache_context(context_name): + with components.transformer.cache_context(context_name, timestep=t): guider_state_batch.noise_pred = components.transformer( hidden_states=block_state.latent_model_input, image_embeds=block_state.image_embeds, @@ -370,7 +370,7 @@ def __call__( cond_kwargs = {input_name: getattr(guider_state_batch, input_name) for input_name in guider_inputs.keys()} context_name = getattr(guider_state_batch, components.guider._identifier_key) - with components.transformer.cache_context(context_name): + with components.transformer.cache_context(context_name, timestep=t): guider_state_batch.noise_pred = components.transformer( hidden_states=block_state.latent_model_input, image_embeds=block_state.image_embeds, diff --git a/src/diffusers/modular_pipelines/ltx/denoise.py b/src/diffusers/modular_pipelines/ltx/denoise.py index 44a3c91d471b..d04ed785c26c 100644 --- a/src/diffusers/modular_pipelines/ltx/denoise.py +++ b/src/diffusers/modular_pipelines/ltx/denoise.py @@ -136,7 +136,7 @@ def __call__( } context_name = getattr(guider_state_batch, components.guider._identifier_key, None) - with components.transformer.cache_context(context_name): + with components.transformer.cache_context(context_name, timestep=t): guider_state_batch.noise_pred = components.transformer( hidden_states=block_state.latent_model_input, timestep=t.expand(block_state.latent_model_input.shape[0]).to(block_state.dtype), @@ -369,7 +369,7 @@ def __call__( } context_name = getattr(guider_state_batch, components.guider._identifier_key, None) - with components.transformer.cache_context(context_name): + with components.transformer.cache_context(context_name, timestep=t): guider_state_batch.noise_pred = components.transformer( hidden_states=block_state.latent_model_input, timestep=block_state.timestep_adjusted, diff --git a/src/diffusers/modular_pipelines/ltx2/denoise.py b/src/diffusers/modular_pipelines/ltx2/denoise.py index 2fb4951377a1..4346cb72148d 100644 --- a/src/diffusers/modular_pipelines/ltx2/denoise.py +++ b/src/diffusers/modular_pipelines/ltx2/denoise.py @@ -413,7 +413,7 @@ def __call__( cond_kwargs = {name: getattr(batch, name) for name in self._guider_input_fields} cond_kwargs["spatio_temporal_guidance_blocks"] = batch.spatio_temporal_guidance_blocks cond_kwargs["isolate_modalities"] = batch.isolate_modalities - with components.transformer.cache_context(getattr(batch, identifier_key)): + with components.transformer.cache_context(getattr(batch, identifier_key), timestep=t): noise_pred_video, noise_pred_audio = components.transformer( hidden_states=block_state.latent_model_input, audio_hidden_states=block_state.audio_latent_model_input, 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/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/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/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/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/cosmos/pipeline_cosmos3_omni.py b/src/diffusers/pipelines/cosmos/pipeline_cosmos3_omni.py index edd8f44409c2..5891e0b6651f 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 2e580388874c..29824a00dd3d 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_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/flux/pipeline_flux_kontext.py b/src/diffusers/pipelines/flux/pipeline_flux_kontext.py index 3579c3abd8d1..0706a4a9b297 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 05bf8200f588..ee2ac6c0d794 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.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.py b/src/diffusers/pipelines/flux2/pipeline_flux2_klein.py index 50fd51e283f6..2a8c9075cc91 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 3150c1144f67..d56dab1678d2 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/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/helios/pipeline_helios.py b/src/diffusers/pipelines/helios/pipeline_helios.py index c80158ce8625..e5fda2e4480c 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 c32108d613ed..4051ba4d6f4c 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 4faa73d40247..25dc37b927a4 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 4b926b82c474..3729df8e8893 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 22e8346b8a14..84e4c113e177 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 1edd4f6db9f6..6d3810a518ad 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 0208559bb56e..6f8550c848ca 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 929118cc0af8..0dc53158a940 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 cba5f0aa5c68..4017fcb8268e 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 b171f83c3668..5362b7497316 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/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/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 eed7b40c48fc..81f351d978cf 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 8f107bd456c4..f4405ecfb158 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 e66f8dc18cf4..6712debaadb5 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 dc2dde5aad93..24327ba90838 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 f057f8d7908e..b75cf0aede47 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 11e65270cfb6..a92cc8f4d604 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 4e5ced0b4ec8..836e50461b77 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 e1bff845302c..c16ef736d451 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_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/ltx2/pipeline_ltx2_hdr_lora.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py index d9d95cdc0a0b..22ddd01ded21 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 2924a086c721..f174b3aefa4e 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 36b7effdfd6c..4b9d938a62eb 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 b0d2f91736ea..031583618484 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 cbb8f3b2c6c7..c9c3be3a708f 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 8bd2eb5bb5c9..65a3657c5199 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 d30aebc610b8..bf447bc93dce 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/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/ovis_image/pipeline_ovis_image.py b/src/diffusers/pipelines/ovis_image/pipeline_ovis_image.py index 841f4e1471e2..cb96c124fda1 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 34bfb8637f68..58c56f4727e8 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 ff619a781585..aa966bc9d5db 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet.py @@ -872,18 +872,19 @@ 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.transformer.cache_context("cond"): + 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( hidden_states=latents, timestep=timestep / 1000, @@ -896,7 +897,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 8c09df5574c1..2aec58dec7b2 100644 --- a/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py +++ b/src/diffusers/pipelines/qwenimage/pipeline_qwenimage_controlnet_inpaint.py @@ -843,18 +843,19 @@ 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.transformer.cache_context("cond"): + 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( hidden_states=latents, timestep=timestep / 1000, @@ -867,7 +868,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 1861141810f6..4e11ab8394aa 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 6413205424eb..2866e0bb6969 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 8b17f17c0e3e..c334099e3e6c 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 9beaa769fae2..f1b172f8ca75 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 a7b0e4d9912b..363a7a259b6c 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 e7053794e9b6..70cdbe0ec32f 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 7cb519cde4a2..b151980622e8 100644 --- a/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py +++ b/src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py @@ -780,7 +780,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, @@ -796,7 +796,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 b4a4242b0036..24b99b6c3db8 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 73b1541f9969..b740c898383e 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 c19997155d66..e5c2352a46c7 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 c18a80fdae0e..b3ff509fdc39 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 977073a60a1e..2e291d3009ec 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/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.py b/src/diffusers/pipelines/wan/pipeline_wan.py index 452911cf899d..e6c26a79ef2b 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 e96729826e70..0ec8211842c2 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 2d1f7b94750e..fc54abc70f28 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 f7689496c968..24f0f0f6f103 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/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 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]) 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/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} 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