diff --git a/docs/source/en/api/cache.md b/docs/source/en/api/cache.md index a5ed8751118d..7b71b36670ad 100644 --- a/docs/source/en/api/cache.md +++ b/docs/source/en/api/cache.md @@ -46,3 +46,9 @@ Cache methods speedup diffusion transformers by storing and reusing intermediate [[autodoc]] MagCacheConfig [[autodoc]] apply_mag_cache + +## DualCacheConfig + +[[autodoc]] DualCacheConfig + +[[autodoc]] apply_dual_cache diff --git a/docs/source/en/optimization/cache.md b/docs/source/en/optimization/cache.md index 07db3d84b489..f1ec619ce9b6 100644 --- a/docs/source/en/optimization/cache.md +++ b/docs/source/en/optimization/cache.md @@ -163,3 +163,36 @@ image = pipe("A cat playing chess", num_inference_steps=4).images[0] > [!TIP] > For pipelines that run Classifier-Free Guidance in a **batched** manner (like SDXL or Flux), the `hidden_states` processed by the model contain both conditional and unconditional branches concatenated together. The calibration process automatically accounts for this, producing a single array of ratios that represents the joint behavior. You can use this resulting array directly without modification. + +## DuCa (Dual Feature Caching) + +[DuCa (Dual Feature Caching)](https://huggingface.co/papers/2412.18911) is a training-free feature-caching schedule for diffusion transformers. It alternates between an *aggressive* strategy that reuses cached block residuals verbatim for maximum speedup and a *conservative* strategy that damps the reused residual to arrest the quality drop caused by reusing stale features, refreshing the cache at fixed cycle boundaries. + +Set up and pass a [`DualCacheConfig`] to a pipeline's transformer to enable it: + +- `cache_interval`: Length of each cache cycle. A full recomputation happens on the first step of every cycle; the remaining steps reuse cached features. +- `aggressive_steps`: Number of steps immediately after a compute step that reuse the cached residual verbatim. The remaining steps of the cycle use the conservative strategy. +- `conservative_scale`: Multiplier applied to the cached residual on conservative steps. Values below `1.0` damp error accumulation from aging features. +- `retention_ratio`: Fraction of initial steps during which caching is disabled for stability. +- `num_inference_steps`: Number of inference steps used by the pipeline, required to resolve `retention_ratio` into a step count. + +```python +import torch +from diffusers import FluxPipeline, DualCacheConfig + +pipe = FluxPipeline.from_pretrained( + "black-forest-labs/FLUX.1-dev", + torch_dtype=torch.bfloat16, +).to("cuda") + +config = DualCacheConfig( + cache_interval=3, + aggressive_steps=1, + conservative_scale=0.95, + retention_ratio=0.2, + num_inference_steps=28, +) +pipe.transformer.enable_cache(config) + +image = pipe("A cat playing chess", num_inference_steps=28).images[0] +``` diff --git a/src/diffusers/__init__.py b/src/diffusers/__init__.py index da77fa67df52..4f1d488d0928 100644 --- a/src/diffusers/__init__.py +++ b/src/diffusers/__init__.py @@ -175,6 +175,7 @@ ) _import_structure["hooks"].extend( [ + "DualCacheConfig", "FasterCacheConfig", "FirstBlockCacheConfig", "HookRegistry", @@ -184,6 +185,7 @@ "SmoothedEnergyGuidanceConfig", "TaylorSeerCacheConfig", "TextKVCacheConfig", + "apply_dual_cache", "apply_faster_cache", "apply_first_block_cache", "apply_layer_skip", @@ -1036,6 +1038,7 @@ TangentialClassifierFreeGuidance, ) from .hooks import ( + DualCacheConfig, FasterCacheConfig, FirstBlockCacheConfig, HookRegistry, @@ -1045,6 +1048,7 @@ SmoothedEnergyGuidanceConfig, TaylorSeerCacheConfig, TextKVCacheConfig, + apply_dual_cache, apply_faster_cache, apply_first_block_cache, apply_layer_skip, diff --git a/src/diffusers/hooks/__init__.py b/src/diffusers/hooks/__init__.py index 2a9aa81608e7..0a7d5d808ed7 100644 --- a/src/diffusers/hooks/__init__.py +++ b/src/diffusers/hooks/__init__.py @@ -17,6 +17,7 @@ if is_torch_available(): from .context_parallel import apply_context_parallel + from .dual_cache import DualCacheConfig, apply_dual_cache from .faster_cache import FasterCacheConfig, apply_faster_cache from .first_block_cache import FirstBlockCacheConfig, apply_first_block_cache from .group_offloading import apply_group_offloading diff --git a/src/diffusers/hooks/dual_cache.py b/src/diffusers/hooks/dual_cache.py new file mode 100644 index 000000000000..a481284df859 --- /dev/null +++ b/src/diffusers/hooks/dual_cache.py @@ -0,0 +1,339 @@ +# Copyright 2025 The HuggingFace Team. All rights reserved. +# +# 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 dataclasses import dataclass +from typing import Tuple, Union + +import torch + +from ..utils import get_logger +from ..utils.torch_utils import unwrap_module +from ._common import _ALL_TRANSFORMER_BLOCK_IDENTIFIERS +from ._helpers import TransformerBlockRegistry +from .hooks import BaseState, HookRegistry, ModelHook, StateManager + + +logger = get_logger(__name__) # pylint: disable=invalid-name + +_DUAL_CACHE_LEADER_BLOCK_HOOK = "dual_cache_leader_block_hook" +_DUAL_CACHE_BLOCK_HOOK = "dual_cache_block_hook" + +# Step types produced by ``DualCachePolicy``. +_COMPUTE = "compute" +_AGGRESSIVE = "aggressive" +_CONSERVATIVE = "conservative" + + +class DualCachePolicy: + """ + Training-free scheduler that classifies every denoising step into one of the two caching strategies proposed in + Dual Feature Caching (DuCa), or a fresh full computation. + + Each cache cycle spans ``cache_interval`` steps. The cycle opens with a full ``"compute"`` step that refreshes the + cached residual, followed by ``aggressive_steps`` ``"aggressive"`` steps (the cached residual is reused verbatim + while it is freshest) and finally the remaining ``"conservative"`` steps (the cached residual is reused with a + damping correction as it ages). The two-phase reuse is DuCa's namesake "dual" schedule; keeping it deterministic and + model-free makes the policy independently unit-testable. + """ + + def __init__( + self, + cache_interval: int = 3, + aggressive_steps: int = 1, + retention_ratio: float = 0.2, + num_inference_steps: int = 28, + ) -> None: + if cache_interval < 1: + raise ValueError(f"`cache_interval` must be >= 1, got {cache_interval}.") + if aggressive_steps < 0: + raise ValueError(f"`aggressive_steps` must be >= 0, got {aggressive_steps}.") + # There are at most `cache_interval - 1` reuse slots after the leading compute step. + self.aggressive_steps = min(aggressive_steps, max(cache_interval - 1, 0)) + self.cache_interval = cache_interval + self.retention_ratio = retention_ratio + self.num_inference_steps = num_inference_steps + + @property + def retention_steps(self) -> int: + return int(self.retention_ratio * self.num_inference_steps + 0.5) + + def classify(self, step_index: int) -> str: + """Return the step type (``"compute"``/``"aggressive"``/``"conservative"``) for ``step_index``.""" + if step_index < self.retention_steps: + return _COMPUTE + position = (step_index - self.retention_steps) % self.cache_interval + if position == 0: + return _COMPUTE + if position <= self.aggressive_steps: + return _AGGRESSIVE + return _CONSERVATIVE + + +@dataclass +class DualCacheConfig: + r""" + Configuration for [Dual Feature Caching (DuCa)](https://huggingface.co/papers/2412.18911). + + DuCa is a training-free feature-caching schedule for Diffusion Transformers. It alternates between an *aggressive* + strategy that reuses cached block residuals verbatim for maximum speedup and a *conservative* strategy that damps + the reused residual to arrest the quality drop caused by reusing stale features, refreshing the cache at fixed cycle + boundaries. + + Args: + cache_interval (`int`, defaults to `3`): + Length of each cache cycle (`N` in the paper). A full recomputation happens on the first step of every + cycle; the remaining `cache_interval - 1` steps reuse cached features. + aggressive_steps (`int`, defaults to `1`): + Number of steps immediately after a compute step that reuse the cached residual verbatim. The remaining + steps of the cycle use the conservative strategy. Clamped to `cache_interval - 1`. + conservative_scale (`float`, defaults to `0.95`): + Multiplier applied to the cached residual on conservative steps. Values below `1.0` damp error accumulation + from aging features. This scalar is a parameter-free proxy for DuCa's ToCa selective token recomputation, + which the block-level hook architecture cannot host. + retention_ratio (`float`, defaults to `0.2`): + Fraction of initial steps during which caching is disabled for stability, mirroring the warmup convention + used by [`MagCacheConfig`]. + num_inference_steps (`int`, defaults to `28`): + Number of inference steps used by the pipeline, required to resolve `retention_ratio` into a step count. + """ + + cache_interval: int = 3 + aggressive_steps: int = 1 + conservative_scale: float = 0.95 + retention_ratio: float = 0.2 + num_inference_steps: int = 28 + + def get_policy(self) -> DualCachePolicy: + return DualCachePolicy( + cache_interval=self.cache_interval, + aggressive_steps=self.aggressive_steps, + retention_ratio=self.retention_ratio, + num_inference_steps=self.num_inference_steps, + ) + + +def _combine_residual(hidden_states: torch.Tensor, residual: torch.Tensor, scale: float) -> torch.Tensor: + """Add a (optionally scaled) cached residual back onto ``hidden_states``, tolerating text+image concat layouts.""" + if residual.device != hidden_states.device: + residual = residual.to(hidden_states.device) + if scale != 1.0: + residual = residual * scale + + if residual.shape == hidden_states.shape: + return hidden_states + residual + # Flux/SD3-style concatenation: the image tokens sit at the tail of the sequence dimension. + if ( + hidden_states.ndim == 3 + and residual.ndim == 3 + and hidden_states.shape[0] == residual.shape[0] + and hidden_states.shape[2] == residual.shape[2] + and hidden_states.shape[1] > residual.shape[1] + ): + diff = hidden_states.shape[1] - residual.shape[1] + hidden_states = hidden_states.clone() + hidden_states[:, diff:, :] = hidden_states[:, diff:, :] + residual + return hidden_states + + logger.warning( + f"DualCache: cannot align residual {tuple(residual.shape)} with input {tuple(hidden_states.shape)}; " + "returning input unchanged for this step." + ) + return hidden_states + + +class DualCacheState(BaseState): + def __init__(self) -> None: + super().__init__() + self.previous_residual: torch.Tensor = None + self.head_block_input: Union[torch.Tensor, Tuple[torch.Tensor, ...]] = None + self.should_compute: bool = True + self.step_index: int = 0 + + def reset(self): + self.previous_residual = None + self.head_block_input = None + self.should_compute = True + self.step_index = 0 + + +class DualCacheHeadHook(ModelHook): + _is_stateful = True + + def __init__(self, state_manager: StateManager, config: DualCacheConfig): + self.state_manager = state_manager + self.config = config + self.policy = config.get_policy() + self._metadata = None + + def initialize_hook(self, module): + unwrapped_module = unwrap_module(module) + self._metadata = TransformerBlockRegistry.get(unwrapped_module.__class__) + return module + + @torch.compiler.disable + def new_forward(self, module: torch.nn.Module, *args, **kwargs): + if self.state_manager._current_context is None: + self.state_manager.set_context("inference") + + arg_name = self._metadata.hidden_states_argument_name + hidden_states = self._metadata._get_parameter_from_args_kwargs(arg_name, args, kwargs) + + state: DualCacheState = self.state_manager.get_state() + state.head_block_input = hidden_states + + step_type = self.policy.classify(state.step_index) + # A reuse step can only run once a residual has been cached. + if step_type == _COMPUTE or state.previous_residual is None: + state.should_compute = True + return self.fn_ref.original_forward(*args, **kwargs) + + state.should_compute = False + scale = 1.0 if step_type == _AGGRESSIVE else self.config.conservative_scale + logger.debug(f"DualCache: reusing cache at step {state.step_index} ({step_type}, scale={scale})") + + output = _combine_residual(hidden_states, state.previous_residual, scale) + + if self._metadata.return_encoder_hidden_states_index is not None: + original_encoder_hidden_states = self._metadata._get_parameter_from_args_kwargs( + "encoder_hidden_states", args, kwargs + ) + max_idx = max( + self._metadata.return_hidden_states_index, self._metadata.return_encoder_hidden_states_index + ) + ret_list = [None] * (max_idx + 1) + ret_list[self._metadata.return_hidden_states_index] = output + ret_list[self._metadata.return_encoder_hidden_states_index] = original_encoder_hidden_states + return tuple(ret_list) + return output + + def reset_state(self, module): + self.state_manager.reset() + return module + + +class DualCacheBlockHook(ModelHook): + def __init__(self, state_manager: StateManager, config: DualCacheConfig, is_tail: bool = False): + super().__init__() + self.state_manager = state_manager + self.config = config + self.is_tail = is_tail + self._metadata = None + + def initialize_hook(self, module): + unwrapped_module = unwrap_module(module) + self._metadata = TransformerBlockRegistry.get(unwrapped_module.__class__) + return module + + @torch.compiler.disable + def new_forward(self, module: torch.nn.Module, *args, **kwargs): + if self.state_manager._current_context is None: + self.state_manager.set_context("inference") + state: DualCacheState = self.state_manager.get_state() + + if not state.should_compute: + arg_name = self._metadata.hidden_states_argument_name + hidden_states = self._metadata._get_parameter_from_args_kwargs(arg_name, args, kwargs) + if self.is_tail: + self._advance_step(state) + if self._metadata.return_encoder_hidden_states_index is not None: + encoder_hidden_states = self._metadata._get_parameter_from_args_kwargs( + "encoder_hidden_states", args, kwargs + ) + max_idx = max( + self._metadata.return_hidden_states_index, self._metadata.return_encoder_hidden_states_index + ) + ret_list = [None] * (max_idx + 1) + ret_list[self._metadata.return_hidden_states_index] = hidden_states + ret_list[self._metadata.return_encoder_hidden_states_index] = encoder_hidden_states + return tuple(ret_list) + return hidden_states + + output = self.fn_ref.original_forward(*args, **kwargs) + + if self.is_tail: + out_hidden = output[self._metadata.return_hidden_states_index] if isinstance(output, tuple) else output + in_hidden = state.head_block_input + if in_hidden is not None and out_hidden.shape == in_hidden.shape: + state.previous_residual = out_hidden - in_hidden + self._advance_step(state) + + return output + + def _advance_step(self, state: DualCacheState): + state.step_index += 1 + if state.step_index >= self.config.num_inference_steps: + state.step_index = 0 + state.previous_residual = None + + +def apply_dual_cache(module: torch.nn.Module, config: DualCacheConfig) -> None: + """ + Applies [Dual Feature Caching (DuCa)](https://huggingface.co/papers/2412.18911) to a transformer module. + + A [`DualCacheHeadHook`] on the first transformer block decides, per step, whether to run a fresh forward pass or to + reuse the cached residual (verbatim on aggressive steps, damped on conservative steps). A tail [`DualCacheBlockHook`] + caches the full-stack residual after each fresh compute. + + Args: + module (`torch.nn.Module`): + The transformer module to apply DuCa to. + config (`DualCacheConfig`): + The configuration for Dual Feature Caching. + """ + HookRegistry.check_if_exists_or_initialize(module) + + state_manager = StateManager(DualCacheState, (), {}) + blocks = [] + for name, submodule in module.named_children(): + if name not in _ALL_TRANSFORMER_BLOCK_IDENTIFIERS or not isinstance(submodule, torch.nn.ModuleList): + continue + for index, block in enumerate(submodule): + blocks.append((f"{name}.{index}", block)) + + if not blocks: + logger.warning("DualCache: No transformer blocks found to apply hooks.") + return + + if len(blocks) == 1: + name, block = blocks[0] + logger.info(f"DualCache: Applying head+tail hooks to single block '{name}'") + _apply_dual_cache_block_hook(block, state_manager, config, is_tail=True) + _apply_dual_cache_head_hook(block, state_manager, config) + return + + head_block_name, head_block = blocks.pop(0) + tail_block_name, tail_block = blocks.pop(-1) + + logger.info(f"DualCache: Applying head hook to '{head_block_name}'") + _apply_dual_cache_head_hook(head_block, state_manager, config) + for name, block in blocks: + _apply_dual_cache_block_hook(block, state_manager, config) + logger.info(f"DualCache: Applying tail hook to '{tail_block_name}'") + _apply_dual_cache_block_hook(tail_block, state_manager, config, is_tail=True) + + +def _apply_dual_cache_head_hook(block: torch.nn.Module, state_manager: StateManager, config: DualCacheConfig) -> None: + registry = HookRegistry.check_if_exists_or_initialize(block) + if registry.get_hook(_DUAL_CACHE_LEADER_BLOCK_HOOK) is not None: + registry.remove_hook(_DUAL_CACHE_LEADER_BLOCK_HOOK) + registry.register_hook(DualCacheHeadHook(state_manager, config), _DUAL_CACHE_LEADER_BLOCK_HOOK) + + +def _apply_dual_cache_block_hook( + block: torch.nn.Module, state_manager: StateManager, config: DualCacheConfig, is_tail: bool = False +) -> None: + registry = HookRegistry.check_if_exists_or_initialize(block) + if registry.get_hook(_DUAL_CACHE_BLOCK_HOOK) is not None: + registry.remove_hook(_DUAL_CACHE_BLOCK_HOOK) + registry.register_hook(DualCacheBlockHook(state_manager, config, is_tail), _DUAL_CACHE_BLOCK_HOOK) diff --git a/src/diffusers/models/cache_utils.py b/src/diffusers/models/cache_utils.py index 161fcf426f21..40139ad3f794 100644 --- a/src/diffusers/models/cache_utils.py +++ b/src/diffusers/models/cache_utils.py @@ -28,6 +28,7 @@ class CacheMixin: - [Pyramid Attention Broadcast](https://huggingface.co/papers/2408.12588) - [FasterCache](https://huggingface.co/papers/2410.19355) - [FirstBlockCache](https://github.com/chengzeyi/ParaAttention/blob/7a266123671b55e7e5a2fe9af3121f07a36afc78/README.md#first-block-cache-our-dynamic-caching) + - [Dual Feature Caching (DuCa)](https://huggingface.co/papers/2412.18911) """ _cache_config = None @@ -67,12 +68,14 @@ def enable_cache(self, config) -> None: """ from ..hooks import ( + DualCacheConfig, FasterCacheConfig, FirstBlockCacheConfig, MagCacheConfig, PyramidAttentionBroadcastConfig, TaylorSeerCacheConfig, TextKVCacheConfig, + apply_dual_cache, apply_faster_cache, apply_first_block_cache, apply_mag_cache, @@ -92,6 +95,8 @@ def enable_cache(self, config) -> None: apply_first_block_cache(self, config) elif isinstance(config, MagCacheConfig): apply_mag_cache(self, config) + elif isinstance(config, DualCacheConfig): + apply_dual_cache(self, config) elif isinstance(config, TextKVCacheConfig): apply_text_kv_cache(self, config) elif isinstance(config, PyramidAttentionBroadcastConfig): @@ -105,6 +110,7 @@ def enable_cache(self, config) -> None: def disable_cache(self) -> None: from ..hooks import ( + DualCacheConfig, FasterCacheConfig, FirstBlockCacheConfig, HookRegistry, @@ -113,6 +119,7 @@ def disable_cache(self) -> None: TaylorSeerCacheConfig, TextKVCacheConfig, ) + from ..hooks.dual_cache import _DUAL_CACHE_BLOCK_HOOK, _DUAL_CACHE_LEADER_BLOCK_HOOK from ..hooks.faster_cache import _FASTER_CACHE_BLOCK_HOOK, _FASTER_CACHE_DENOISER_HOOK from ..hooks.first_block_cache import _FBC_BLOCK_HOOK, _FBC_LEADER_BLOCK_HOOK from ..hooks.mag_cache import _MAG_CACHE_BLOCK_HOOK, _MAG_CACHE_LEADER_BLOCK_HOOK @@ -134,6 +141,9 @@ def disable_cache(self) -> None: elif isinstance(self._cache_config, MagCacheConfig): registry.remove_hook(_MAG_CACHE_LEADER_BLOCK_HOOK, recurse=True) registry.remove_hook(_MAG_CACHE_BLOCK_HOOK, recurse=True) + elif isinstance(self._cache_config, DualCacheConfig): + registry.remove_hook(_DUAL_CACHE_LEADER_BLOCK_HOOK, recurse=True) + registry.remove_hook(_DUAL_CACHE_BLOCK_HOOK, recurse=True) elif isinstance(self._cache_config, PyramidAttentionBroadcastConfig): registry.remove_hook(_PYRAMID_ATTENTION_BROADCAST_HOOK, recurse=True) elif isinstance(self._cache_config, TextKVCacheConfig): diff --git a/src/diffusers/utils/dummy_pt_objects.py b/src/diffusers/utils/dummy_pt_objects.py index 8439a2b93371..7a2f48044cf5 100644 --- a/src/diffusers/utils/dummy_pt_objects.py +++ b/src/diffusers/utils/dummy_pt_objects.py @@ -167,6 +167,21 @@ def from_pretrained(cls, *args, **kwargs): requires_backends(cls, ["torch"]) +class DualCacheConfig(metaclass=DummyObject): + _backends = ["torch"] + + def __init__(self, *args, **kwargs): + requires_backends(self, ["torch"]) + + @classmethod + def from_config(cls, *args, **kwargs): + requires_backends(cls, ["torch"]) + + @classmethod + def from_pretrained(cls, *args, **kwargs): + requires_backends(cls, ["torch"]) + + class FasterCacheConfig(metaclass=DummyObject): _backends = ["torch"] @@ -302,6 +317,10 @@ def from_pretrained(cls, *args, **kwargs): requires_backends(cls, ["torch"]) +def apply_dual_cache(*args, **kwargs): + requires_backends(apply_dual_cache, ["torch"]) + + def apply_faster_cache(*args, **kwargs): requires_backends(apply_faster_cache, ["torch"]) diff --git a/tests/hooks/test_dual_cache.py b/tests/hooks/test_dual_cache.py new file mode 100644 index 000000000000..df77f5a83145 --- /dev/null +++ b/tests/hooks/test_dual_cache.py @@ -0,0 +1,136 @@ +# Copyright 2025 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. + +import pytest +import torch + +from diffusers import DualCacheConfig, apply_dual_cache +from diffusers.hooks._helpers import TransformerBlockMetadata, TransformerBlockRegistry +from diffusers.hooks.dual_cache import DualCachePolicy +from diffusers.models import ModelMixin +from diffusers.models.cache_utils import CacheMixin + + +class DummyBlock(torch.nn.Module): + def forward(self, hidden_states, encoder_hidden_states=None, **kwargs): + return hidden_states * 2.0 + + +class DummyTransformer(ModelMixin, CacheMixin): + def __init__(self): + super().__init__() + self.transformer_blocks = torch.nn.ModuleList([DummyBlock(), DummyBlock()]) + + def forward(self, hidden_states, encoder_hidden_states=None): + for block in self.transformer_blocks: + hidden_states = block(hidden_states, encoder_hidden_states=encoder_hidden_states) + return hidden_states + + +@pytest.fixture(autouse=True) +def register_dummy_blocks(): + TransformerBlockRegistry.register( + DummyBlock, + TransformerBlockMetadata(return_hidden_states_index=None, return_encoder_hidden_states_index=None), + ) + + +def _set_context(model, context_name): + for module in model.modules(): + if hasattr(module, "_diffusers_hook"): + module._diffusers_hook._set_context(context_name) + + +def test_policy_dual_schedule(): + """The policy must open each cycle with a compute step, then aggressive, then conservative steps.""" + policy = DualCachePolicy(cache_interval=3, aggressive_steps=1, retention_ratio=0.0, num_inference_steps=9) + types = [policy.classify(i) for i in range(9)] + assert types == [ + "compute", + "aggressive", + "conservative", + "compute", + "aggressive", + "conservative", + "compute", + "aggressive", + "conservative", + ] + + +def test_policy_retention_forces_compute(): + """Steps inside the warmup window are always full computes regardless of the cycle.""" + policy = DualCachePolicy(cache_interval=2, aggressive_steps=1, retention_ratio=0.5, num_inference_steps=8) + assert policy.retention_steps == 4 + assert [policy.classify(i) for i in range(4)] == ["compute"] * 4 + # After warmup the dual cycle resumes. + assert policy.classify(4) == "compute" + assert policy.classify(5) == "aggressive" + + +def test_policy_clamps_aggressive_steps(): + """`aggressive_steps` can never exceed the number of reuse slots in a cycle.""" + policy = DualCachePolicy(cache_interval=2, aggressive_steps=5, retention_ratio=0.0, num_inference_steps=4) + assert policy.aggressive_steps == 1 + assert [policy.classify(i) for i in range(4)] == ["compute", "aggressive", "compute", "aggressive"] + + +def test_aggressive_step_reuses_residual_verbatim(): + """A fresh compute caches the residual; the next (aggressive) step reuses it as-is.""" + model = DummyTransformer() + config = DualCacheConfig( + cache_interval=2, aggressive_steps=1, retention_ratio=0.0, num_inference_steps=2 + ) + apply_dual_cache(model, config) + _set_context(model, "test_context") + + # Step 0 (compute): input 10 -> 4x -> 40, cached residual = 30. + assert torch.allclose(model(torch.tensor([[[10.0]]])), torch.tensor([[[40.0]]])) + # Step 1 (aggressive): reuse -> 11 + 30 = 41 (not the computed 44). + assert torch.allclose(model(torch.tensor([[[11.0]]])), torch.tensor([[[41.0]]])) + + +def test_conservative_step_damps_residual(): + """Conservative steps reuse the cached residual scaled by `conservative_scale`.""" + model = DummyTransformer() + config = DualCacheConfig( + cache_interval=3, + aggressive_steps=1, + conservative_scale=0.5, + retention_ratio=0.0, + num_inference_steps=3, + ) + apply_dual_cache(model, config) + _set_context(model, "test_context") + + model(torch.tensor([[[10.0]]])) # compute -> residual 30 + model(torch.tensor([[[11.0]]])) # aggressive -> 41 + # Step 2 (conservative): 12 + 30 * 0.5 = 27. + assert torch.allclose(model(torch.tensor([[[12.0]]])), torch.tensor([[[27.0]]])) + + +def test_enable_cache_dispatches_dual_cache(): + """The CacheMixin.enable_cache dispatch must route DualCacheConfig through apply_dual_cache.""" + model = DummyTransformer() + assert not model.is_cache_enabled + + model.enable_cache(DualCacheConfig(cache_interval=2, retention_ratio=0.0, num_inference_steps=2)) + assert model.is_cache_enabled + assert isinstance(model._cache_config, DualCacheConfig) + # The head hook must be registered on the first transformer block by the dispatch. + assert model.transformer_blocks[0]._diffusers_hook.get_hook("dual_cache_leader_block_hook") is not None + + model.disable_cache() + assert not model.is_cache_enabled + assert model.transformer_blocks[0]._diffusers_hook.get_hook("dual_cache_leader_block_hook") is None