diff --git a/docs/source/en/optimization/cache.md b/docs/source/en/optimization/cache.md index 9f775ec3b88c..3a771336e4b7 100644 --- a/docs/source/en/optimization/cache.md +++ b/docs/source/en/optimization/cache.md @@ -72,8 +72,14 @@ pipeline.transformer.enable_cache(config) [SeaCache](https://huggingface.co/papers/2602.18993) compares Spectral-Evolution-Aware (SEA) indicators between successive denoising steps. When the accumulated indicator change remains below a threshold, it skips the expensive -transformer block stack and predicts its output from cached residuals. The indicator is computed from the raw vision -latents, including clean conditioning frames for image-to-video generation. +transformer block stack and predicts its output from cached residuals. Build the indicator from the visual latents that +form the generated output. Include clean conditioning frames when they are part of that output trajectory, as in +image-to-video and video-to-video generation. Exclude separate visual hints that condition the generation but are not +part of the output. Text conditioning is excluded because it is not a visual latent. + +Cosmos 3 Transfer packs control hints as separate visual sequences, so its adapter excludes them from the indicator. +Control-CFG branches compare the same output trajectory while retaining their own cached residuals. Control hints still +condition the transformer. The implementation provides built-in adapters for the following models: @@ -87,8 +93,9 @@ Other video transformers can integrate with the generic path when they use `Cach block list, and register the block input/output layout in `TransformerBlockRegistry`. The pipeline must enter a `cache_context` for every transformer call, attach `step_index`, `sigma`, and `num_inference_steps`, and use separate context names for independent trajectories such as conditional and unconditional guidance. Pass a `raw_vision_callback` -that returns the noisy vision latents when no built-in adapter is available. Validate output quality and tune the cache -parameters for each model and scheduler; support and benchmark results do not transfer automatically from Cosmos 3. +that returns the visual latents forming the generated output when no built-in adapter is available. Validate output +quality and tune the cache parameters for each model and scheduler; support and benchmark results do not transfer +automatically from Cosmos 3. ### Cosmos 3 diff --git a/src/diffusers/hooks/sea_cache.py b/src/diffusers/hooks/sea_cache.py index 5cb78db1f6b4..d228b0772e48 100644 --- a/src/diffusers/hooks/sea_cache.py +++ b/src/diffusers/hooks/sea_cache.py @@ -64,8 +64,9 @@ class SeaCacheConfig: power_exp (`float`, defaults to `3.0`): Exponent of the SEA clean-signal power prior. SeaCache uses `3.0` for video features. raw_vision_callback (`Callable`, *optional*): - Advanced model adapter returning raw vision latents with shape `(C, T, H, W)`. When omitted, a built-in - adapter is used if one is available. + Advanced model adapter returning the visual latents forming the generated output, each with shape `(C, T, + H, W)`. Include clean conditioning frames within the output trajectory, but exclude separate visual hints + that are not part of the output. When omitted, a built-in adapter is used if one is available. Example: ```python @@ -326,7 +327,6 @@ def _prepare_cosmos3_raw_vision_metadata( return None raw_vision = [] - has_noisy_vision = False for latent, noisy_frame_indexes in zip(vision_tokens, vision_noisy_frame_indexes): if not isinstance(latent, torch.Tensor) or not isinstance(noisy_frame_indexes, torch.Tensor): return None @@ -340,10 +340,12 @@ def _prepare_cosmos3_raw_vision_metadata( noisy_frame_indexes = noisy_frame_indexes.flatten().to(device=latent.device, dtype=torch.long) if torch.any(noisy_frame_indexes < 0) or torch.any(noisy_frame_indexes >= latent.shape[1]): return None - has_noisy_vision = has_noisy_vision or noisy_frame_indexes.numel() > 0 - raw_vision.append(latent) + # A sequence with noisy frames belongs to the generated output. Keep that sequence whole so clean conditioning + # frames remain in the indicator, but exclude separate clean hints that are not part of the output. + if noisy_frame_indexes.numel() > 0: + raw_vision.append(latent) - return raw_vision if raw_vision and has_noisy_vision else None + return raw_vision or None def _prepare_wan_t2v_raw_vision_metadata( diff --git a/tests/models/transformers/test_models_transformer_cosmos3.py b/tests/models/transformers/test_models_transformer_cosmos3.py index 6b04dd77f8c6..c897d5bfdc69 100644 --- a/tests/models/transformers/test_models_transformer_cosmos3.py +++ b/tests/models/transformers/test_models_transformer_cosmos3.py @@ -108,7 +108,106 @@ def output_shape(self) -> tuple[int, ...]: return (1, 2, 1, 1, 1) -class TestCosmos3OmniTransformerModel(Cosmos3OmniTransformerTesterConfig, ModelTesterMixin): +class TestCosmos3OmniTransformerSeaCache(Cosmos3OmniTransformerTesterConfig, SeaCacheTesterMixin): + cache_input_key = "vision_tokens" + + def test_sea_cache_tracks_output_visual_trajectory(self): + model = self.model_class(**self.get_init_dict()).to(torch_device).eval() + model.enable_cache(SeaCacheConfig(threshold=2.0, cache_end_steps=0)) + target = torch.randn(1, 2, 2, 1, 1, device=torch_device) + control = torch.randn_like(target) + target_with_changed_clean_frame = target.clone() + target_with_changed_clean_frame[:, :, 0] += 10 + inputs = self.get_dummy_inputs() + inputs.update( + sequence_length=6, + position_ids=torch.zeros(3, 6, dtype=torch.long, device=torch_device), + vision_tokens=[control, target], + vision_token_shapes=[(2, 1, 1)] * 2, + vision_sequence_indexes=torch.arange(2, 6, device=torch_device), + vision_mse_loss_indexes=torch.tensor([5], device=torch_device), + vision_noisy_frame_indexes=[ + torch.tensor([], dtype=torch.long, device=torch_device), + torch.tensor([1], device=torch_device), + ], + ) + layer_calls = 0 + + def count_layer_calls(_module, _args, _output): + nonlocal layer_calls + layer_calls += 1 + + model.layers[0].register_forward_hook(count_layer_calls) + decisions = [] + for step, (current_control, current_target) in enumerate( + ( + (control, target), + (control + 100, target), + (control + 100, target_with_changed_clean_frame), + ) + ): + inputs["vision_tokens"] = [current_control, current_target] + with ( + torch.no_grad(), + model.cache_context("cond", step_index=step, sigma=0.9 - step * 0.3, num_inference_steps=3), + ): + model(**inputs) + state = model._diffusers_hook.get_hook(_SEA_CACHE_ROOT_HOOK).state_manager._state_cache["cond"] + decisions.append(state.gate_should_compute) + + assert decisions == [True, False, True] + assert layer_calls == 2 + + @pytest.mark.parametrize("residual_order", [0, 1]) + def test_sea_cache_transfer_branches_share_indicator_with_separate_histories(self, residual_order): + model = self.model_class(**self.get_init_dict()).to(torch_device).eval() + model.enable_cache(SeaCacheConfig(threshold=100.0, residual_order=residual_order, cache_end_steps=0)) + root_hook = model._diffusers_hook.get_hook(_SEA_CACHE_ROOT_HOOK) + target = torch.randn(1, 2, 2, 1, 1, device=torch_device) + control = torch.randn_like(target) + decisions = [] + + for step in range(6): + states = [] + for context, with_control in (("cond", True), ("cond_no_control", False), ("uncond", True)): + inputs = self.get_dummy_inputs() + sequence_length = 6 if with_control else 4 + inputs.update( + sequence_length=sequence_length, + position_ids=torch.zeros(3, sequence_length, dtype=torch.long, device=torch_device), + vision_tokens=[control, target] if with_control else [target], + vision_token_shapes=[(2, 1, 1)] * (2 if with_control else 1), + vision_sequence_indexes=torch.arange(2, sequence_length, device=torch_device), + vision_mse_loss_indexes=torch.tensor([sequence_length - 1], device=torch_device), + vision_noisy_frame_indexes=( + [ + torch.tensor([], dtype=torch.long, device=torch_device), + torch.tensor([1], device=torch_device), + ] + if with_control + else [torch.tensor([1], device=torch_device)] + ), + ) + with ( + torch.no_grad(), + model.cache_context(context, step_index=step, sigma=0.9 - step * 0.1, num_inference_steps=6), + ): + output = model(**inputs) + assert torch.isfinite(output.sample[-1]).all() + states.append(root_hook.state_manager._state_cache[context]) + + assert all(len(state.previous_indicator) == 1 for state in states) + for state in states[1:]: + torch.testing.assert_close(state.previous_indicator[0], states[0].previous_indicator[0]) + assert state.gate_should_compute == states[0].gate_should_compute + assert state.history is not states[0].history + assert states[0].history[-1][2].shape != states[1].history[-1][2].shape + decisions.append(states[0].gate_should_compute) + + assert any(decisions) and not all(decisions) + model._reset_stateful_cache() + assert all(not state.history and state.previous_indicator is None for state in states) + def test_cosmos3_supports_sea_cache_without_changing_state_dict_keys(self): model = self.model_class(**self.get_init_dict()).to(torch_device).eval() state_dict_keys = set(model.state_dict()) @@ -423,6 +522,8 @@ def test_cosmos3_sea_cache_regional_compile_fullgraph_without_recompile(self): assert refreshed.sample[0].shape == self.output_shape assert root_hook.state_manager._state_cache["cond"].history[-1][0] == 2 + +class TestCosmos3OmniTransformerModel(Cosmos3OmniTransformerTesterConfig, ModelTesterMixin): def test_cosmos3_decoder_layer_cache_metadata_tracks_generation_stream(self): metadata = TransformerBlockRegistry.get(Cosmos3VLTextMoTDecoderLayer) @@ -571,10 +672,6 @@ def test_cosmos3_nemotron_rms_norm_multiplies_in_float32(self): torch.testing.assert_close(norm(hidden_states), expected, rtol=0, atol=0) -class TestCosmos3OmniTransformerSeaCache(Cosmos3OmniTransformerTesterConfig, SeaCacheTesterMixin): - cache_input_key = "vision_tokens" - - class TestCosmos3OmniTransformerMemory(Cosmos3OmniTransformerTesterConfig, MemoryTesterMixin): @pytest.mark.skip("The transformer returns one tensor list per generated modality.") def test_layerwise_casting_training(self):