Repository navigation
[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.
…sts)
- Move svg_attention.py from models/wan/transformers/ to models/ so LTX-2
no longer imports from the Wan package.
- Add init_svg_config() and apply_svg_or_dense() to svg_attention and use
them from both LTX2Attention and FlaxWanAttention.
- Stop copying use_base2_exp/use_experimental_scheduler/ulysses_* into
attention_config (NNXAttentionOp takes them as args); drop unused
svg_implementation/svg_global_offset/svg_{high,low}_noise_density
attributes and raise on non-default values instead.
- Drop the unused deterministic arg and redundant spatiotemporal_shape check.
- LTX2StaticContext.spatiotemporal_shape is a static field; audio_attn1
inherits attention_config with SVG disabled; svg_num_layers defaults to
the model depth.
- Tests: assert sparse config contents, block-level video/audio dispatch,
and TPU parity of SVG at density 1.0 vs dense.
|
Please squash the commits! @jitendra-jalwaniya |
…former
Add an opt-in SVG sparse attention path to LTX-2 video self-attention
(attn1). Audio self-attention and cross-modal attention stay dense.
- Move svg_attention.py from models/wan/transformers/ to models/ so LTX-2
does not import from the Wan package, and add init_svg_config() and
apply_svg_or_dense() shared by LTX2Attention and FlaxWanAttention.
- Do not copy use_base2_exp/use_experimental_scheduler/ulysses_* into
attention_config (NNXAttentionOp takes them as args). Raise on
non-default svg_implementation/svg_global_offset instead of silently
ignoring them. The Wan-only svg_{high,low}_noise_density settings are
not used by LTX-2, which has a single transformer and reads
svg_spatial_density.
- LTX2StaticContext.spatiotemporal_shape is a static field; audio_attn1
inherits attention_config with SVG disabled; svg_num_layers defaults to
the model depth.
- Tests: sparse config contents, block-level video/audio dispatch, and
TPU parity of SVG at density 1.0 vs dense.
1879bc4 to
a33c8b2
Compare
Done |
| def init_svg_config(attention_config: Optional[Mapping[str, Any]], default_num_layers: int) -> dict[str, Any]: | ||
| """Resolves the attention-level SVG settings from `attention_config`. | ||
|
|
||
| Returns a dict keyed by the names in `SVG_ATTENTION_DEFAULTS` plus | ||
| `svg_num_layers`; other keys are ignored. Settings that configs accept but | ||
| the head-local SVG kernel does not implement fail loudly instead of being | ||
| silently dropped. | ||
| """ | ||
| attention_config = attention_config or {} | ||
| implementation = attention_config.get("svg_implementation", "official_svg") | ||
| if implementation != "official_svg": | ||
| raise ValueError(f"Unsupported svg_implementation={implementation!r}; only 'official_svg' is implemented.") | ||
| global_offset = attention_config.get("svg_global_offset", 0) | ||
| if global_offset: | ||
| raise ValueError(f"svg_global_offset={global_offset} is not supported by head-local SVG.") | ||
| resolved = {**SVG_ATTENTION_DEFAULTS, "svg_num_layers": default_num_layers} | ||
| for name in resolved: | ||
| if name in attention_config: | ||
| resolved[name] = attention_config[name] | ||
| return resolved |
There was a problem hiding this comment.
Important: When use_svg_attention is False (or forced False for cross-attention in LTX2Attention at attention_ltx2.py:L362), init_svg_config still copies every svg_* override (svg_spatial_density, svg_active_start_step, etc.) onto self, and LTX2VideoTransformer3DModel.__init__ (transformer_ltx2.py:L905-L908) stores the full self.attention_config dict on the module.
- Because
aot_cache.cached_jithashes_dynamic_signature((graphdef, ...))via_graphdef_desc(graphdef)(which serializes everyStaticattribute innnx.GraphDef.attributes), changing an unusedsvg_*flag whenuse_svg_attention=False(orsvg_spatial_density=1.0) still changesgraphdefand misses the AOT cache:use_svg_attention=False, svg_spatial_density=0.25, svg_active_start_step=10->_dynamic_signature = d382ae2efd22use_svg_attention=False, svg_spatial_density=0.50, svg_active_start_step=20->_dynamic_signature = c92170d2e500use_svg_attention=True, svg_spatial_density=1.0, svg_active_start_step=10->_dynamic_signature = 21a6fcb2b699
- In addition,
init_svg_configvalidatessvg_implementationandsvg_global_offsetbefore checking whetheruse_svg_attentionis even enabled.
If not bool(attention_config.get("use_svg_attention", False)) (or when context_dim is not None in LTX2Attention), could we return the canonical SVG_ATTENTION_DEFAULTS (with "svg_num_layers": default_num_layers) without validating or storing inactive svg_* overrides, and similarly omit inactive svg_* keys from self.attention_config in transformer_ltx2.py:L905?
|
|
||
| def _head_local_svg_attention(query, key, value, context): | ||
| from .wan.transformers import svg_attention, svg_head_local | ||
| from .wan.transformers import svg_head_local |
There was a problem hiding this comment.
Nit: Now that svg_attention.py has been moved out of models/wan/transformers/ into models/svg_attention.py so LTX-2 does not depend on the wan package, _head_local_svg_attention still has an in-line import from .wan.transformers import svg_head_local (which is a 56-line model-agnostic helper containing only inference_only and exchange_local). Consider moving svg_head_local.py alongside models/svg_attention.py (or merging it into svg_attention.py) and importing it at module top-level.
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.