Repository navigation
[Perf][Cosmos3] Add SeaCache support for Cosmos3 - #6922
linyueqian merged 21 commits into
Conversation
|
Codex usage limits have been reached for code reviews. Please check with the admins of this repo to increase the limits by adding credits. |
|
This PR appears to belong to: docs/design/module/cache_management.md. Module owners: @Isotr0py @princepride @SamitHuang @yzhautouskay, please review your own changes and leave a short self-review comment describing what you checked. PRs without author self-review may not be assigned a reviewer. Please take a look when you have a chance. If you would like an automated review, mention @vllm-omni-review-bot in a comment. |
Omni ReviewBot triage noteResolved as of |
MaciejBalaNV
left a comment
There was a problem hiding this comment.
Overall the changes look good. However, there is some improvements to be made about Cosmos3 specific comments and logs - I didn't mark all of these. Can you also update the Cosmos3 docs mentioning this is the recommended cache for these models?
Self-review / current stateThis PR adds a native SeaCache backend. Mainly targeting Cosmos3, but potentially possible to be used by other models. It is currently opt-in before broader quality validation, with the explicit goal of enabling SeaCache by default Remaining action items
|
|
Hi, thanks for the contribution! I think we need to be very careful about adding more diffusion caching backends unless there is a clear reason to do so, especially when the caching needs potential changes in the model. Adding too many approaches is also confusing to users since it's not clear what to pick, and it can also cause model support to be more sparse since contributions become diluted over the different backends. It would also be easier to tell whether or not a new caching backend should be added based on how it benchmarks (both for quality and accuracy) against the existing diffusion caching approaches, and not just the baseline. Then we can potentially provide better guidance to users on when to use what. Also cc @hsliuustc0106 @wtomin @NickCao @RuixiangMa in case any of you have thoughts as well |
|
Hey @alex-jw-brooks , in our internal tests for Cosmos3 models SeaCache is a clear winner and the only solution in a production-ready state. I understand the worry about confusing users with too many options, but at the same time we should make sure that vLLM-Omni stays on SOTA level. In the fast moving field it's inevitable that new options will show up, as is the case here for e.g. SeaCache (first paper Feb 2026) vs TeaCache (first paper Nov 2024). IMO we should take advantage of a modular design with easy to swap caches and extend the framework with new options, rather than limit ourselves because the amount of choices can be confusing. We can provide more examples and benchmarks for Cosmos3 model with different cache backends. We will also update the docs for Cosmos3 models explaining which cache backend should be used. |
Signed-off-by: Yuliya Zhautouskaya <yzhautouskay@nvidia.com>
Signed-off-by: Yuliya Zhautouskaya <yzhautouskay@nvidia.com>
Signed-off-by: Yuliya Zhautouskaya <yzhautouskay@nvidia.com>
Signed-off-by: Yuliya Zhautouskaya <yzhautouskay@nvidia.com>
ea2ac43 to
183a380
Compare
|
@alex-jw-brooks @MaciejBalaNV Here's diffusion-cache comparison to demonstrate SeaCache's optimal quality-speedup balance for Cosmos3 SettingsSingle-seed comparison on 1× GB200 per mode: SeaCache, Cache-DiT, and TeaCache via vLLM-Omni PR #4389. vLLM-Omni cache-acceleration post for Cache-DiT background.
Video: 1280×720, 189 frames, 35 steps, CFG 6.0. T2I: 1024×1024, 50 steps, CFG 7.0.
Environments
Cosmos3-NanoT2V
cosmos3_nano_t2v_comparison_pair1.mp4cosmos3_nano_t2v_comparison_pair2.mp4cosmos3_nano_t2v_comparison_pair3.mp4I2V
cosmos3_nano_i2v_comparison_pair1.mp4cosmos3_nano_i2v_comparison_pair2.mp4cosmos3_nano_i2v_comparison_pair3.mp4T2I
|
Signed-off-by: Yuliya Zhautouskaya <yzhautouskay@nvidia.com>
|
@hsliuustc0106 It includes:
I also merged the latest upstream |
Signed-off-by: Yuliya Zhautouskaya <yzhautouskay@nvidia.com>
Signed-off-by: Yuliya Zhautouskaya <yzhautouskay@nvidia.com>
linyueqian
left a comment
There was a problem hiding this comment.
The SeaCache implementation itself is careful and I did not find a defect in it. The indicator does latent[batch].movedim(0, -1) before apply_sea_filter, so the separable Wiener gain really runs over (T, H, W); extrapolate_residual is a correct in-place Newton divided-difference table with a zero-denominator guard; the gate forces full compute on the first and last step, on max_consecutive_cached, and whenever history or the indicator is missing; history is bounded to residual_order + 1 entries; and _synchronize_compute all-reduces the skip decision with MAX across the FS, SP and layerwise-offload groups so ranks cannot disagree and hang a collective. The get_extractor change to walk __mro__ is what lets FSDP-wrapped transformers resolve, and the two-key enabler map (Cosmos3OmniDiffusersPipeline, Cosmos3OmniPipeline) matches the alias cache_dit already uses in model_specific.py:885. Cosmos3EdgeVFMTransformer exists as a subclass in transformer_cosmos3_edge.py, so both extractor keys are live.
What holds this back is not the cache; it is two changes to the default path that the description presents as an opt-in feature. Both are inline. The first is that sampling_dtype = torch.float32 is unconditional: initial noise, sound latents, I2V image latents and the velocity mask are now created in float32 and every prediction is cast back to it, whether or not --cache-backend sea_cache is set. That is very likely the right numerical choice, but it changes what every Cosmos3 user gets on upgrade, including the noise draw at a fixed seed, and the baseline column in the PR's own table was measured on this branch, so the 1.8x is fp32-vs-fp32 and says nothing about main-vs-PR. The second is _dit_any_rank_failed, which on main imports a get_dit_group that does not exist, catches the ImportError, and silently returns the local flag; the PR points it at the world group and so turns on a cross-rank all-reduce that has never actually run. That is a bug fix and probably a good one, but it belongs in the description with a sentence on why every rank is guaranteed to reach it in lockstep.
Two smaller things worth noting rather than acting on. The transformer forward and the new extractor both raise TypeError on any unexpected kwarg where the old forward silently accepted **kwargs; CI is green so every in-tree caller is clean, but it is a tightening third-party wrappers could trip on. And the PR carries two unrelated hardenings in extract_qwen_context (torch.as_tensor for a possibly non-tensor timestep) and extract_flux2_klein_context (a None guard); harmless, but they are not SeaCache and a reader of the history will not find them here.
Validation: tests/diffusion/cache/test_seacache.py is core_model/cpu so it gates per-PR, and buildkite/vllm-omni is green at this head (build 15156). Reviewed statically; the head is on a fork and no PR code was executed. Verdict is comment rather than approve only because the float32 sampling change needs to be either stated and defended for the default path or gated behind the backend, and that is the author's call to make, not mine.
Signed-off-by: Yuliya Zhautouskaya <yzhautouskay@nvidia.com>
Signed-off-by: Yuliya Zhautouskaya <yzhautouskay@nvidia.com>
| noisy_frame_mask = extra_states.get("sea_cache_noisy_frame_mask") | ||
| if isinstance(noisy_frame_mask, torch.Tensor) and not bool(torch.any(noisy_frame_mask != 0).item()): | ||
| self._warn_once("SeaCache requires noisy vision; conditioning-only calls run in full.") | ||
| return self._run_uncached(ctx) |
There was a problem hiding this comment.
Move uncached execution outside this try block; model failures currently trigger a second forward.
| gain = reshaped_gain if gain is None else gain * reshaped_gain | ||
|
|
||
| assert gain is not None | ||
| mean_gain = gain.mean() |
There was a problem hiding this comment.
Normalize by mean_gain directly; validated gains are finite and positive.
There was a problem hiding this comment.
Approving at 7adab72c. Both items from round one have been taken out of this PR rather than argued inside it: the unconditional fp32 sampling state now lives in #7592 with its own motivation and no-cache numbers, and the _dit_any_rank_failed group change is reverted so that function is byte-for-byte main again, with a clean follow-up promised. What is left is exactly the opt-in cache and nothing that alters the default Cosmos3 path.
The cache itself did not change since my first read and was clean then: the separable Wiener filter runs over (T, H, W) after the movedim, the Newton residual extrapolation is a correct divided-difference table with a zero-denominator guard, first and last steps are forced, history is bounded to residual_order + 1, the skip decision is all-reduced with MAX across the FS, SP and layerwise-offload groups so ranks cannot diverge, and both extractor keys resolve to live classes. The runner delta against main is the SeaCache registration plus the mypy-driven tidy that came with it: a local pipeline alias asserted non-None instead of repeated self.pipeline reads, matched_state instead of shadowing state, an isinstance(self.cache_backend, CacheDiTBackend) check in place of the string compare, and an explicit error when a stepwise request reaches the scheduler with no latents; none of it changes control flow.
Validation: lane 15246 on this head is green on the general and AMD lanes (Intel is red on every main commit today, NPU is informational), tests/diffusion/cache/test_seacache.py is core_model/cpu and gates per-PR, and a cross-model panel on this head found nothing beyond the one statistics nit inline. Reviewed statically; the head is on a fork and no PR code was executed.
One request before this is merged rather than before it is approved: the branch is 56 commits behind main and two of them touch files in this diff (13b85c56 from #7427 in pipeline_cosmos3.py and its test, 68294066 from #4167 in serve.py), and the green lane ran on the head as-is, so it has never seen those changes. Please rebase onto current main; I will cycle ready after the push so the lane fires, and merge on green.
Signed-off-by: Yuliya Zhautouskaya <yzhautouskay@nvidia.com>
…3_diffusion_caching Signed-off-by: Yuliya Zhautouskaya <yzhautouskay@nvidia.com>
linyueqian
left a comment
There was a problem hiding this comment.
Re-approving at 2dd3f786, which is the rebase I asked for plus one small follow-up. The base is now current main (78934753), and the merge commit is byte-identical to the automatic merge of 6ba87dfb with main, so the #7427 and #4167 changes that the previous lane never saw are in this tree. The follow-up commit does two things and both check out: the finite-and-positive guard on gain.mean() in sea_filter.py is gone, which is safe because every axis gain is signal_scale * clean_power / (signal_scale**2 * clean_power + noise_scale**2 + 1e-16) with clean_power > 0 and signal_scale clamped into (1e-6, 1 - 1e-6), so the product is strictly positive and finite for any shape and the guard was dead; and the conditioning-only early return in hook.py now happens after the scheduler metadata is validated instead of before, which only changes which one-time warning fires when both conditions hold and adds a test pinning that a failing forward on a conditioning-only call runs exactly once.
Everything I approved at 7adab72c is otherwise unchanged: both round-one items are out of this PR (#7592 and a follow-up), the default Cosmos3 path is main's, and the cache itself was clean on my read and on a three-model panel; the one statistics nit from that panel is inline as a suggestion and does not hold the merge.
Validation: the general lane on this head has been re-fired and I will merge on green; tests/diffusion/cache/test_seacache.py is core_model/cpu and gates per-PR. Reviewed statically; the head is on a fork and no PR code was executed.
| if residual.device != ctx.hidden_states.device: | ||
| residual = residual.to(ctx.hidden_states.device) | ||
| state.consecutive_cached += 1 | ||
| self.skip_count += 1 |
There was a problem hiding this comment.
[suggestion] skip_count and consecutive_cached are bumped before the can_reuse check, so the defensive fallback a few lines down (shape, device or dtype mismatch) runs the full stack and _record_execution, which resets consecutive_cached but never takes the skip back out of skip_count. The step is counted as both skipped and executed, which only corrupts the summary statistics and only on a path that should never fire, so this is not holding the merge; moving the two increments under if can_reuse: keeps the numbers honest. Credit to the cross-model panel for spotting it.
Signed-off-by: Yuliya Zhautouskaya <yzhautouskay@nvidia.com> Signed-off-by: Matthieu Laneuville <matthieu.laneuville@surf.nl>
Signed-off-by: Yuliya Zhautouskaya <yzhautouskay@nvidia.com>

Summary
Output videos
I2V
vllm_omni_i2v_baseline_seacache_metrics_synced.mp4
T2V
vllm_omni_t2v_baseline_seacache_metrics_synced.mp4
Performance and fidelity
Both runs use Cosmos3-Nano with 35 steps, guidance 6, shift 10, 1280×720 output, 189 frames at 24 FPS, and regional
torch.compile. SeaCache uses threshold 0.25, residual order 1, maximum 2 consecutive cached calls, and power 3.LPIPS is computed across all 189 frames against the corresponding uncached output.
Reproduction commands
Both benchmarks used one GB200 with regional
torch.compileenabled.Baseline server
SeaCache server
CUDA_VISIBLE_DEVICES=0 vllm serve nvidia/Cosmos3-Nano \ --omni --host 0.0.0.0 --port 8000 --init-timeout 1800 \ --cache-backend sea_cache \ --cache-config '{"sea_threshold":0.25,"sea_residual_order":1,"sea_max_consecutive_cached":2,"sea_power_exp":3.0}' \ --no-guardrailsThe same request was executed once against each server.
T2V request
I2V request
Test Plan
vLLM Version:
0.28.0vLLM-Omni Commit:
3db6064e2561b19d449e84f2e5d7af24bd54f8bbtorch.compile.Test Result
python -m pytest -sv tests/diffusion/models/cosmos3/ tests/diffusion/cache/ -m "core_model and cpu"BEFORE SUBMITTING: read CONTRIBUTING.md and run the precheck-pr skill with the code agent for a self-check against project conventions.
(anything written below this line will be removed by GitHub Actions)