From bdb8e27f2d7b79ccb22d867fb4c54e5d04e0b7ea Mon Sep 17 00:00:00 2001 From: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com> Date: Mon, 5 Oct 2026 12:11:47 +0000 Subject: [PATCH] Fix missing cache contexts in Wan video-to-video --- .../pipelines/wan/pipeline_wan_video2video.py | 24 +++---- tests/pipelines/testing_utils/__init__.py | 2 + tests/pipelines/testing_utils/cache.py | 62 +++++++++++++++++++ .../pipelines/wan/test_wan_video_to_video.py | 10 ++- 4 files changed, 86 insertions(+), 12 deletions(-) diff --git a/src/diffusers/pipelines/wan/pipeline_wan_video2video.py b/src/diffusers/pipelines/wan/pipeline_wan_video2video.py index b192147acb64..36e1b77c0a1b 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"): + 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"): + 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/pipelines/testing_utils/__init__.py b/tests/pipelines/testing_utils/__init__.py index 9b756ec64693..b1d31cb96339 100644 --- a/tests/pipelines/testing_utils/__init__.py +++ b/tests/pipelines/testing_utils/__init__.py @@ -1,4 +1,5 @@ from .cache import ( + CacheContextTesterMixin, CacheTesterMixin, FasterCacheTesterMixin, FirstBlockCacheTesterMixin, @@ -37,6 +38,7 @@ "GroupOffloadTesterMixin", "LayerwiseCastingTesterMixin", "CacheTesterMixin", + "CacheContextTesterMixin", "PyramidAttentionBroadcastTesterMixin", "FasterCacheTesterMixin", "FirstBlockCacheTesterMixin", diff --git a/tests/pipelines/testing_utils/cache.py b/tests/pipelines/testing_utils/cache.py index 82980d57fb3c..ab559d2805d9 100644 --- a/tests/pipelines/testing_utils/cache.py +++ b/tests/pipelines/testing_utils/cache.py @@ -19,6 +19,7 @@ from diffusers import FasterCacheConfig, PyramidAttentionBroadcastConfig from diffusers.hooks.first_block_cache import FirstBlockCacheConfig +from diffusers.hooks.hooks import BaseState, HookRegistry, ModelHook, StateManager from diffusers.hooks.mag_cache import MagCacheConfig from diffusers.hooks.pyramid_attention_broadcast import PyramidAttentionBroadcastHook from diffusers.hooks.taylorseer_cache import TaylorSeerCacheConfig @@ -86,6 +87,67 @@ def run_forward(pipe): ) +class _CacheContextState(BaseState): + def __init__(self): + self.calls = 0 + + def reset(self): + self.calls = 0 + + +class _CacheContextRecorderHook(ModelHook): + _is_stateful = True + + def __init__(self): + super().__init__() + self.state_manager = StateManager(_CacheContextState) + self.calls = [] + + def pre_forward(self, module, *args, **kwargs): + state = self.state_manager.get_state() + state.calls += 1 + self.calls.append((self.state_manager.context.name, state, state.calls)) + return args, kwargs + + def reset_state(self, module): + self.state_manager.reset() + return module + + +@is_cache +class CacheContextTesterMixin(BasePipelineOutputMixin): + """Checks named denoiser contexts, state isolation, and cleanup across pipeline calls.""" + + cache_context_inputs = {} + expected_cache_contexts = {"cond", "uncond"} + + def _test_cache_context(self, expected_contexts, **inputs): + pipe = self.get_pipeline().to(torch_device) + expected_output = self.run_pipe(pipe, **inputs) + hook = _CacheContextRecorderHook() + HookRegistry.check_if_exists_or_initialize(pipe.transformer).register_hook(hook, "cache_context_recorder") + previous_states = {} + + for _ in range(2): + hook.calls.clear() + output = self.run_pipe(pipe, **inputs) + torch.testing.assert_close(output, expected_output, rtol=0, atol=0) + assert {name for name, _, _ in hook.calls} == expected_contexts + states = {} + for name in expected_contexts: + calls = [(state, count) for context, state, count in hook.calls if context == name] + states[name] = calls[0][0] + assert all(state is states[name] for state, _ in calls) + assert [count for _, count in calls] == list(range(1, len(calls) + 1)) + assert states[name].calls == 0 + assert states[name] is not previous_states.get(name) + assert len({id(state) for state in states.values()}) == len(expected_contexts) + previous_states = states + + def test_cache_context(self): + self._test_cache_context(self.expected_cache_contexts, **self.cache_context_inputs) + + @is_cache class PyramidAttentionBroadcastTesterMixin(CacheTesterMixin): PAB_CONFIG = { diff --git a/tests/pipelines/wan/test_wan_video_to_video.py b/tests/pipelines/wan/test_wan_video_to_video.py index 0c6a0f4dfd6f..799c786598a6 100644 --- a/tests/pipelines/wan/test_wan_video_to_video.py +++ b/tests/pipelines/wan/test_wan_video_to_video.py @@ -21,7 +21,7 @@ from diffusers import AutoencoderKLWan, UniPCMultistepScheduler, WanTransformer3DModel, WanVideoToVideoPipeline from ...testing_utils import assert_tensors_close -from ..testing_utils import BasePipelineTesterConfig, MemoryTesterMixin, PipelineTesterMixin +from ..testing_utils import BasePipelineTesterConfig, CacheContextTesterMixin, MemoryTesterMixin, PipelineTesterMixin class WanVideoToVideoPipelineTesterConfig(BasePipelineTesterConfig): @@ -128,3 +128,11 @@ def test_save_load_float16(self): class TestWanVideoToVideoPipelineMemory(WanVideoToVideoPipelineTesterConfig, MemoryTesterMixin): pass + + +class TestWanVideoToVideoPipelineCacheContext(WanVideoToVideoPipelineTesterConfig, CacheContextTesterMixin): + @pytest.mark.parametrize("guidance_scale", [1.0, 6.0]) + @pytest.mark.parametrize("strength", [0.5, 1.0]) + def test_cache_context(self, guidance_scale, strength): + expected_contexts = {"cond", "uncond"} if guidance_scale > 1.0 else {"cond"} + self._test_cache_context(expected_contexts, guidance_scale=guidance_scale, strength=strength)