Skip to content

[PR 1/2] SVG implementation for LTX 2 - #497

Open
jitendra-jalwaniya wants to merge 1 commit into
mainfrom
ltx2_svg_model
Open

jitendra-jalwaniya wants to merge 1 commit into
mainfrom
ltx2_svg_model

Conversation

@jitendra-jalwaniya

@jitendra-jalwaniya jitendra-jalwaniya commented Sep 29, 2026 •

Copy link
Copy Markdown
Collaborator

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: LTX2Attention accepts an attention_config dict with SVG settings and dispatches video self-attention to the SVG kernel (or dense via jax.lax.cond) based on the active step/layer window.
  • transformer_ltx2.py: plumbs spatiotemporal_shape, svg_timestep, svg_step_index and per-layer layer_index through LTX2StaticContext/LTX2BlockContext (scanned and unscanned paths). Only attn1 (video self-attention) gets SVG; audio_attn1 is 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:

  • Dense Baseline: use_svg_attention=False, attention=ulysses_custom
  • SVG Ulysses: use_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)

Metric Dense Baseline SVG Ulysses Delta (SVG vs. Dense)
Denoising Time / Video (40 steps) 72.00 s 61.62 s -10.38 s (-14.42% / 1.17× speedup)
Per-Step Denoising Latency 1.800 s / step 1.541 s / step -0.259 s / step
Total Inference / Video (end-to-end) 102.94 s 91.23 s -11.71 s (-11.38%)
Benchmark Wall Time (778 Videos) ~22.24 hours ~19.72 hours -2.52 hours saved

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%):

  • Audio-Video Synchronization & Lip-Sync:
    • Synchformer Temporal Desynchronization (second_desync $\downarrow$): Reduced from 0.6743 s $\rightarrow$ 0.6404 s (-5.03% better sync overall), with strong gains on Animals (-18.25%), Music (-12.24%), Virtual Worlds (-10.97%), and Synchronous Physical Sounds (-6.12%).
    • LatentSync Lip-Sync Error (second_lsa $\downarrow$): Reduced by -31.60% overall (0.9813 $\rightarrow$ 0.6712), and by -39.15% on Human Sounds (1.3270 $\rightarrow$ 0.8075). We observed this with video heads routed through SVG while audio and cross-modal attention stay dense; since each prompt was generated once, we have not yet measured seed-to-seed variance for this metric.
  • Cross-Modal Alignment: Higher alignment across all three embedding models: ImageBind-Huge (+1.91%), ViCLIP-L (+1.82%), and LAION-CLAP (+1.79%).
  • Audio Aesthetic Quality: AudioBox Aesthetics improved by +1.67% (3.5645 $\rightarrow$ 3.6241), with DNSMOS (+0.55%) and NISQA (+0.39%) on par.
  • Multimodal Judge & Fine-Grained QA (Qwen2.5-Omni-7B):
    • Parity on most multimodal judge criteria (alignment: 4.47 vs. 4.45 [-0.60%], audio realism: 3.94 vs. 3.91 [-0.82%], expressiveness: 4.25 vs. 4.23 [-0.45%]); visual realism is slightly lower (4.55 vs. 4.47 [-1.78%]).
    • Notable accuracy gains on multi-turn question answering: Audio QA (+4.90%) and Visual QA (+4.64%).

Full 15-Dimension VABench Comparison Table (778 Prompts)

Module Dimension Metric / Criterion Direction Dense Baseline SVG Ulysses Absolute Delta Relative Change
M1: Audio Quality & Aesthetics first_dnsmos Microsoft DNSMOS (sig_bak_ovr + p808) $\uparrow$ 1.6538 1.6629 +0.0091 +0.55%
first_nisqa NISQA v2 Speech/Audio Naturalness MOS $\uparrow$ 1.5264 1.5324 +0.0060 +0.39%
first_audiobox Meta AudioBox Aesthetics $\uparrow$ 3.5645 3.6241 +0.0596 +1.67%
M2: Cross-Modal Sync & Alignment second_viclip ViCLIP-L Text-Video Similarity $\uparrow$ 0.1918 0.1953 +0.0035 +1.82%
second_clap LAION-CLAP Text-Audio Similarity $\uparrow$ 0.3835 0.3904 +0.0069 +1.79%
second_imagebind Meta ImageBind-Huge AV Alignment $\uparrow$ 0.2144 0.2185 +0.0041 +1.91%
second_desync Synchformer Temporal Desync Offset (s) $\downarrow$ 0.6743 s 0.6404 s -0.0339 s -5.03% (Better Sync)
second_lsa LatentSync Lip-Sync Distance $\downarrow$ 0.9813 0.6712 -0.3101 -31.60% (Better Lip-Sync)
M3: Multimodal Judge (Qwen2.5-Omni-7B) third_alignment AV Semantic & Temporal Alignment (1–5) $\uparrow$ 4.4743 4.4473 -0.0270 -0.60% (Parity)
third_audio_reality Acoustic Realism & Fidelity (1–5) $\uparrow$ 3.9383 3.9062 -0.0321 -0.82% (Parity)
third_visual_reality Visual Realism & Coherence (1–5) $\uparrow$ 4.5476 4.4666 -0.0810 -1.78%
third_expressiveness Emotional & Dynamic Expressiveness (1–5) $\uparrow$ 4.2468 4.2275 -0.0193 -0.45% (Parity)
third_artistry Audiovisual Aesthetic Quality (1–5) $\uparrow$ 3.6272 3.6478 +0.0206 +0.57%
Module 3 Mean Mean Multimodal Judge Score (1–5) $\uparrow$ 4.1668 4.1391 -0.0278 -0.67% (Parity)
M4: Multi-Turn Question Answering fourth_qa_audio Audio QA Accuracy (0–1) $\uparrow$ 0.6438 0.6754 +0.0315 +4.90%
fourth_qa_vision Visual QA Accuracy (0–1) $\uparrow$ 0.6088 0.6370 +0.0283 +4.64%
Module 4 Mean Mean Multi-Modal QA Accuracy (0–1) $\uparrow$ 0.6263 0.6562 +0.0299 +4.77%

Testing

Run the LTX-2 SVG attention unit tests from the repository root:

python -m pytest -q \
  src/maxdiffusion/tests/ltx2/test_svg_attention_ltx2.py

All existing Wan and LTX-2 unit tests continue to pass.

@github-actions

Copy link
Copy Markdown

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

@jitendra-jalwaniya jitendra-jalwaniya changed the title ltx2: add Sparse VideoGen (SVG) attention to LTX2 attention and transformer SVG implementation for LTX 2 Sep 29, 2026
@jitendra-jalwaniya jitendra-jalwaniya changed the title SVG implementation for LTX 2 [PR 4/5] SVG implementation for LTX 2 Sep 29, 2026
@jitendra-jalwaniya
jitendra-jalwaniya requested review from Perseus14 and removed request for entrpn September 29, 2026 07:55
@jitendra-jalwaniya
jitendra-jalwaniya changed the base branch from ltx2_block_benchmark_fixes to fix/pyink-main September 29, 2026 17:42
@jitendra-jalwaniya jitendra-jalwaniya changed the title [PR 4/5] SVG implementation for LTX 2 [PR 1/2] SVG implementation for LTX 2 Sep 29, 2026
@jitendra-jalwaniya
jitendra-jalwaniya changed the base branch from fix/pyink-main to main September 30, 2026 05:06

@Perseus14 Perseus14 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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:

  1. Deduplicating SVG config/dispatch with Wan (attention_ltx2.py vs FlaxWanAttention in attention_flax.py) so we don't maintain two copies of the same ~100 lines.
  2. Passing num_layers from the model into attention_config instead of hardcoding svg_num_layers: 48.
  3. Cleaning up unused or ignored attention_config keys so callers aren't surprised when a key has no effect.
  4. Adding one numerical test on TPU (ulysses_custom) and a check that audio_attn1 stays 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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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 = {

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Two things on this config block:

  1. Duplication with FlaxWanAttention: This 50-line block (and the dispatch block at lines 653–710 below) is almost identical to FlaxWanAttention in src/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 small SVGConfig dataclass or init_svg_config(attention_config, default_num_layers=...) + apply_svg_or_dense(...) in svg_attention.py) and call it from both FlaxWanAttention and LTX2Attention?

  2. Ignored / unused keys in attention_config:

    • use_base2_exp, use_experimental_scheduler, ulysses_shards, and ulysses_attention_chunks are placed into attention_config here, but self.attention_op = NNXAttentionOp(...) at lines 568–571 reads the function arguments (ulysses_shards, use_base2_exp, etc.) directly instead of attention_config[...]. That means if someone passes attention_config={"use_base2_exp": True}, it gets silently ignored. Either read them from attention_config when constructing NNXAttentionOp, or remove those 4 keys from attention_config.
    • svg_implementation, svg_high_noise_density, svg_low_noise_density, and svg_global_offset are stored on self, but never actually read by _head_local_svg_attention. Let's remove the unused ones (or raise a ValueError if 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,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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?

Comment on lines +653 to +666
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,
)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

  1. Minor: at line 653, and spatiotemporal_shape is not None is redundant because lines 607–611 already raise a ValueError if self.use_svg_attention and is_self_attention and spatiotemporal_shape is None.
  2. As mentioned above, this whole is_svg_active + run_dense / run_sparse_svg + jax.lax.cond block is identical to FlaxWanAttention.__call__ (attention_flax.py lines 2937–2990). Pulling this into a helper function in svg_attention.py that takes self.attention_op (or a dense_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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Use pytree_node=False for spatiotemporal_shape:

Suggested change
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},

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Small suggestion: instead of replacing the entire dict with {"use_svg_attention": False}, copy attention_config and override just use_svg_attention:

Suggested change
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.

Comment on lines +905 to +911
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 {}),
}

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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:

Suggested change
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 {}),
}

Comment on lines +188 to +202
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"))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Two test coverage suggestions here:

  1. Verify audio_attn1 stays dense: Right now, block(block_ctx_active) raises inside video self.attn1 before self.audio_attn1 is ever reached. If you patch block.attn1.attention_op.apply_attention and block.audio_attn1.attention_op.apply_attention, you can run a full block forward pass on an active SVG step and assert that:

    • block.attn1 was called with sparse_config_override={"use_svg_attention": True, ...}, and
    • block.audio_attn1 was called with sparse_config_override=None (dense).
  2. Test the actual SVG kernel numerically (not just the ValueError): Every active-step test in this file currently uses attention_kernel="dot_product" and asserts that _head_local_svg_attention raises ValueError("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 TPU v4-8, could we add a test (can be skipped if not on TPU) with a ("data", "fsdp", "context") mesh and attention_kernel="ulysses_custom" that runs LTX2Attention with svg_spatial_density=1.0 and checks jnp.allclose(out_svg, out_dense, atol=1e-2)?

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants