[PR 1/2] SVG implementation for LTX 2 - #497
jitendra-jalwaniya wants to merge 1 commit into
Conversation
There was a problem hiding this comment.
Code Review
This pull request integrates Sparse VideoGen (SVG) attention into the LTX2 model. It introduces SVG configuration parameters, updates the attention layer to route between dense and sparse SVG attention based on active steps and layers, and propagates the necessary spatiotemporal and step metadata through the transformer blocks and static/block contexts. Additionally, comprehensive unit tests are added to verify the SVG activation boundaries, dispatch routing, and full model forward passes. There are no review comments, so no additional feedback is provided.
ecc5b5e to
55e6847
Compare
6e1b301 to
ee0e722
Compare
ee0e722 to
650b1c6
Compare
55e6847 to
2196275
Compare
650b1c6 to
9b30925
Compare
Perseus14
left a comment
There was a problem hiding this comment.
Nice work wiring SVG into LTX-2 and running the full 778-prompt VABench eval! The context plumbing through both the scanned and unscanned transformer paths is clean, and keeping audio/cross-attention dense makes sense.
I left inline comments on a few things to tighten up before merging:
- Deduplicating SVG config/dispatch with Wan (
attention_ltx2.pyvsFlaxWanAttentioninattention_flax.py) so we don't maintain two copies of the same ~100 lines. - Passing
num_layersfrom the model intoattention_configinstead of hardcodingsvg_num_layers: 48. - Cleaning up unused or ignored
attention_configkeys so callers aren't surprised when a key has no effect. - Adding one numerical test on TPU (
ulysses_custom) and a check thataudio_attn1stays dense.
Also a quick note on the PR description:
- Since each prompt was generated with a single seed, the
-31.6%LatentSync and+4.9%QA deltas likely include seed-to-seed variance (sparse attention is an approximation of dense). I'd frame those as "comparable / on par with dense" unless we have multi-seed numbers. - Since #493 is already merged into
main, you can remove the "depends on #493" note.
| import jax.numpy as jnp | ||
| from ... import common_types | ||
| from ..attention_flax import NNXAttentionOp | ||
| from ..wan.transformers import svg_attention |
There was a problem hiding this comment.
Nit / architecture: importing from ..wan.transformers makes the ltx2 model package depend on wan. Since svg_attention.py has no Wan-specific logic, could we move svg_attention.py (or re-export it) under a shared location like maxdiffusion/models/svg_attention.py (or alongside attention_flax.py)? Happy for this to be a quick follow-up PR if you'd rather keep this diff small.
| enable_jax_named_scopes: bool = False, | ||
| attention_config: Optional[dict] = None, | ||
| ): | ||
| attention_config = { |
There was a problem hiding this comment.
Two things on this config block:
-
Duplication with
FlaxWanAttention: This 50-line block (and the dispatch block at lines 653–710 below) is almost identical toFlaxWanAttentioninsrc/maxdiffusion/models/attention_flax.py(lines 2583–2630 and 2937–2990). If we ever change an SVG default or add a parameter, it's easy to update one model and forget the other. Could we extract a shared helper (e.g. a smallSVGConfigdataclass orinit_svg_config(attention_config, default_num_layers=...)+apply_svg_or_dense(...)insvg_attention.py) and call it from bothFlaxWanAttentionandLTX2Attention? -
Ignored / unused keys in
attention_config:use_base2_exp,use_experimental_scheduler,ulysses_shards, andulysses_attention_chunksare placed intoattention_confighere, butself.attention_op = NNXAttentionOp(...)at lines 568–571 reads the function arguments (ulysses_shards,use_base2_exp, etc.) directly instead ofattention_config[...]. That means if someone passesattention_config={"use_base2_exp": True}, it gets silently ignored. Either read them fromattention_configwhen constructingNNXAttentionOp, or remove those 4 keys fromattention_config.svg_implementation,svg_high_noise_density,svg_low_noise_density, andsvg_global_offsetare stored onself, but never actually read by_head_local_svg_attention. Let's remove the unused ones (or raise aValueErrorif non-default values are passed) so users don't think setting them changes behavior.
| k_rotary_emb: Optional[Tuple[Array, Array]] = None, | ||
| perturbation_mask: Optional[Array] = None, | ||
| cached_kv: Optional[Tuple[Array, Array]] = None, | ||
| deterministic: bool = True, |
There was a problem hiding this comment.
deterministic: bool = True is added to LTX2Attention.__call__, and line 608 checks if not deterministic: raise ValueError(...). However, LTX2VideoTransformerBlock.__call__ (in transformer_ltx2.py line 600) never passes deterministic when calling self.attn1(...), so it will always be True in practice.
Could we either wire deterministic through from LTX2VideoTransformer3DModel.__call__ -> LTX2StaticContext -> self.attn1, or drop the deterministic argument here if LTX-2 attention is already inference-only?
| if self.use_svg_attention and is_self_attention and spatiotemporal_shape is not None: | ||
| is_active = svg_attention.is_svg_active( | ||
| step_index=svg_step_index, | ||
| layer_index=svg_layer_index, | ||
| timestep=svg_timestep, | ||
| start_step=self.svg_active_start_step, | ||
| end_step=self.svg_active_end_step, | ||
| start_layer=self.svg_active_start_layer, | ||
| end_layer=self.svg_active_end_layer, | ||
| dense_layer_fraction=self.svg_dense_layer_fraction, | ||
| dense_timestep_fraction=self.svg_dense_timestep_fraction, | ||
| num_train_timesteps=self.svg_num_train_timesteps, | ||
| num_layers=self.svg_num_layers, | ||
| ) |
There was a problem hiding this comment.
- Minor: at line 653,
and spatiotemporal_shape is not Noneis redundant because lines 607–611 already raise aValueErrorifself.use_svg_attention and is_self_attentionandspatiotemporal_shape is None. - As mentioned above, this whole
is_svg_active+run_dense/run_sparse_svg+jax.lax.condblock is identical toFlaxWanAttention.__call__(attention_flax.pylines 2937–2990). Pulling this into a helper function insvg_attention.pythat takesself.attention_op(or adense_fn/sparse_fn) would cut ~60 lines of duplicate code here.
| audio_encoder_attention_mask: Optional[jax.Array] = None | ||
| a2v_cross_attention_mask: Optional[jax.Array] = None | ||
| v2a_cross_attention_mask: Optional[jax.Array] = None | ||
| spatiotemporal_shape: Optional[Tuple[int, int, int]] = None |
There was a problem hiding this comment.
Use pytree_node=False for spatiotemporal_shape:
| spatiotemporal_shape: Optional[Tuple[int, int, int]] = None | |
| spatiotemporal_shape: Optional[Tuple[int, int, int]] = struct.field(pytree_node=False, default=None) |
| ulysses_attention_chunks=ulysses_attention_chunks, | ||
| use_base2_exp=use_base2_exp, | ||
| use_experimental_scheduler=use_experimental_scheduler, | ||
| attention_config={"use_svg_attention": False}, |
There was a problem hiding this comment.
Small suggestion: instead of replacing the entire dict with {"use_svg_attention": False}, copy attention_config and override just use_svg_attention:
| attention_config={"use_svg_attention": False}, | |
| attention_config={**(attention_config or {}), "use_svg_attention": False}, |
That way, if attention_config carries other non-SVG settings in the future, audio_attn1 still receives them while keeping SVG disabled.
| self.attention_config = { | ||
| "use_base2_exp": use_base2_exp, | ||
| "use_experimental_scheduler": use_experimental_scheduler, | ||
| "ulysses_shards": ulysses_shards, | ||
| "ulysses_attention_chunks": ulysses_attention_chunks, | ||
| **(attention_config or {}), | ||
| } |
There was a problem hiding this comment.
In attention_ltx2.py (line 376), svg_num_layers defaults to a hardcoded 48. That's used by is_svg_active when svg_dense_layer_fraction > 0 (math.ceil(dense_layer_fraction * num_layers)).
If someone runs a smaller config, a future LTX model with a different layer count, or a unit test with num_layers=2, svg_dense_layer_fraction will compute the cutoff against 48 instead of the model's actual self.num_layers.
We can fix that here by passing "svg_num_layers": self.num_layers as the default before unpacking attention_config:
| self.attention_config = { | |
| "use_base2_exp": use_base2_exp, | |
| "use_experimental_scheduler": use_experimental_scheduler, | |
| "ulysses_shards": ulysses_shards, | |
| "ulysses_attention_chunks": ulysses_attention_chunks, | |
| **(attention_config or {}), | |
| } | |
| self.attention_config = { | |
| "use_base2_exp": use_base2_exp, | |
| "use_experimental_scheduler": use_experimental_scheduler, | |
| "ulysses_shards": ulysses_shards, | |
| "ulysses_attention_chunks": ulysses_attention_chunks, | |
| "svg_num_layers": self.num_layers, | |
| **(attention_config or {}), | |
| } |
| mock_apply = MagicMock(return_value=jnp.zeros((B, seq_len, dim), dtype=jnp.float32)) | ||
| with patch.object(attn.attention_op, "apply_attention", mock_apply): | ||
| _ = attn( | ||
| hidden_states=hidden_states, | ||
| spatiotemporal_shape=(F, H, W), | ||
| svg_layer_index=0, | ||
| svg_timestep=jnp.array([100.0]), | ||
| svg_step_index=3, | ||
| ) | ||
| mock_apply.assert_called_once() | ||
| _, kwargs = mock_apply.call_args | ||
| self.assertEqual(kwargs.get("spatiotemporal_shape"), (F, H, W)) | ||
| sp_cfg = kwargs.get("sparse_config_override") | ||
| self.assertIsNotNone(sp_cfg) | ||
| self.assertTrue(sp_cfg.get("use_svg_attention")) |
There was a problem hiding this comment.
Nice use of patch.object to verify dispatch on CPU! Could we also assert the rest of the fields in sparse_config_override so we know the schedule indices and band width were computed and plumbed properly? For example:
self.assertEqual(sp_cfg.get("band_width"), 128)
self.assertEqual(sp_cfg.get("svg_layer_index"), 0)
self.assertEqual(sp_cfg.get("svg_step_index"), 3)
self.assertIsNotNone(sp_cfg.get("svg_timestep"))| self.assertIsNotNone(sp_cfg) | ||
| self.assertTrue(sp_cfg.get("use_svg_attention")) | ||
|
|
||
| def test_ltx2_transformer_block_forward_svg_dispatch(self): |
There was a problem hiding this comment.
Two test coverage suggestions here:
-
Verify
audio_attn1stays dense: Right now,block(block_ctx_active)raises inside videoself.attn1beforeself.audio_attn1is ever reached. If you patchblock.attn1.attention_op.apply_attentionandblock.audio_attn1.attention_op.apply_attention, you can run a full block forward pass on an active SVG step and assert that:block.attn1was called withsparse_config_override={"use_svg_attention": True, ...}, andblock.audio_attn1was called withsparse_config_override=None(dense).
-
Test the actual SVG kernel numerically (not just the
ValueError): Every active-step test in this file currently usesattention_kernel="dot_product"and asserts that_head_local_svg_attentionraisesValueError("Head-local SVG requires a custom Ulysses attention backend"). That tests the error guard, but never executes the sparse path end-to-end. Since CI runs on a TPUv4-8, could we add a test (can be skipped if not on TPU) with a("data", "fsdp", "context")mesh andattention_kernel="ulysses_custom"that runsLTX2Attentionwithsvg_spatial_density=1.0and checksjnp.allclose(out_svg, out_dense, atol=1e-2)?
Overview
This PR extends Sparse VideoGen (SVG) spatiotemporal attention support to LTX-2 (LTX2) video generation models on Cloud TPUs, building on the custom Ulysses/ring SVG kernel infrastructure introduced for Wan (PR #480).
Self-attention in LTX-2 transformer blocks dynamically profiles query tokens to choose between spatial and temporal attention patterns per head, skipping unneeded query–key interactions while executing through hardware-aligned local-band kernels on TPU. Sparse attention is opt-in (
use_svg_attention: True), disabled by default, and configurable across denoising steps, layers, and sparsity densities. Audio self-attention and cross-modal attention remain dense to preserve temporal and semantic grounding.This is PR 1/2 (model side). It depends on #493 (pyink formatting fix on
main). The config, pipeline, AOT metadata and docs wiring are in #498.Changes in this PR:
attention_ltx2.py:LTX2Attentionaccepts anattention_configdict with SVG settings and dispatches video self-attention to the SVG kernel (or dense viajax.lax.cond) based on the active step/layer window.transformer_ltx2.py: plumbsspatiotemporal_shape,svg_timestep,svg_step_indexand per-layerlayer_indexthroughLTX2StaticContext/LTX2BlockContext(scanned and unscanned paths). Onlyattn1(video self-attention) gets SVG;audio_attn1is forced dense.tests/ltx2/test_svg_attention_ltx2.py: new unit tests.VABench Evaluation: SVG vs. Dense Attention
The end-to-end results below require both this PR and #498.
We evaluated SVG against dense attention on the Full VABench Benchmark suite (778 prompts across all 24 Easy/Hard bundles and 7 content categories) for LTX-2 synchronized text-to-audio-video (T2AV) generation at long sequence length (768 × 1280 × 241 frames,$N = 29,760$ video tokens, 10.04s @ 24 fps video + 24 kHz PCM audio) on TPU v6e-8 (8 chips), followed by a 15-dimension VABench evaluation across 8× NVIDIA A100-80GB GPUs. Each prompt was generated once per configuration:
use_svg_attention=False,attention=ulysses_customuse_svg_attention=True,svg_spatial_density=0.25,attention=ulysses_custom,svg_active_*left at defaults (SVG active on all steps and layers)1. TPU v6e-8 Generation Performance (
778 Videos @ 768 × 1280 × 241)2. 15-Dimension VABench Quality Highlights
Enabling SVG yields faster generation with comparable overall quality: most metrics are on par or slightly higher, with a small drop in judged visual realism (-1.78%):
second_desyncsecond_lsaQwen2.5-Omni-7B):Full 15-Dimension VABench Comparison Table (
778 Prompts)first_dnsmossig_bak_ovr+p808)first_nisqafirst_audioboxsecond_viclipsecond_clapsecond_imagebindsecond_desyncsecond_lsaQwen2.5-Omni-7B)third_alignmentthird_audio_realitythird_visual_realitythird_expressivenessthird_artistryfourth_qa_audiofourth_qa_visionTesting
Run the LTX-2 SVG attention unit tests from the repository root:
All existing Wan and LTX-2 unit tests continue to pass.