[PR 2/2] LTX2: SVG config and pipeline - #498
jitendra-jalwaniya wants to merge 3 commits into
Conversation
There was a problem hiding this comment.
Code Review
This pull request integrates Sparse VideoGen (SVG) attention support into the LTX2 pipeline in MaxDiffusion. Key changes include adding SVG configuration parameters to the LTX2 video config files, propagating these settings through the pipeline and model construction, updating the diffusion loop to pass step indices, validating incompatibility with CFG cache and MagCache, and adding comprehensive unit tests for configuration propagation and validation. I have no feedback to provide as there are no review comments.
6e1b301 to
ee0e722
Compare
209e2ae to
b91cd12
Compare
b91cd12 to
67b6f94
Compare
ee0e722 to
650b1c6
Compare
67b6f94 to
f3f2c81
Compare
650b1c6 to
9b30925
Compare
f3f2c81 to
1b937b7
Compare
Perseus14
left a comment
There was a problem hiding this comment.
Nice work wiring SVG attention into the LTX-2 pipeline, configs, and docs! The attention_config construction and svg_step_index propagation through both the lax.scan and Python diffusion loops look clean.
I left a few inline comments to address before merging:
- Default layer count (
28->48): LTX-2 and LTX-2.3 have 48 layers, so let's update thesvg_num_layersfallback inltx2_pipeline.py(and confirm whethersvg_active_end_layer: 28indocs/svg.mdand the PR description should also be48). - Docs config snippet (
ulysses_shards):attention: ulysses_ring_custom_fixed_mindocs/svg.md(and the PR description) requiresulysses_shards > 0, otherwisepyconfig.pyraises aValueErrorsinceltx2_video.ymldefaults toulysses_shards: -1. - YAML config keys (
ltx2_video.yml/ltx2_3_video.yml): Add the remainingsvg_*keys read byltx2_pipeline.pyso they can be overridden via CLI, and clean up the unused Wan 2.2 dual-expert keys (svg_high_noise_density/svg_low_noise_density). - AOT metadata & tests: Avoid invalidating dense AOT cache hashes when
use_svg_attention=False, and add a couple of unit tests forsvg_spatial_density=1.0andsvg_step_indexforwarding.
|
|
||
| For CPU semantics checks, set `JAX_PLATFORMS=cpu` and `XLA_FLAGS=--xla_force_host_platform_device_count=8`. Run the suite separately on an eight-device TPU host to exercise the compiled kernels. Tests restricted to one platform are skipped on the other. | ||
|
|
||
| For profiling, enable `enable_jax_named_scopes=True` and capture a short warm denoising interval. Routing, placement, main attention, padding cleanup, merging, and restoration have named scopes, including `svg_route_profile`, `svg_layout_place`, `svg_union_main`, `svg_tail_cleanup`, `svg_lse_merge`, and `svg_layout_restore`. The `svg_kernel_c_tiles…` scope reports the fraction of physical tiles executed, which differs from attention-pair density and total transformer FLOP savings. |
There was a problem hiding this comment.
Small tip worth adding here: in ltx2_video.yml, profiler_steps defaults to 5 (steps 0–4), while the LTX-2 example above starts SVG at step 10 (svg_active_start_step: 10). Could you add a short note reminding users to set profiler_steps higher than svg_active_start_step (or lower svg_active_start_step when profiling) so the profile actually captures the sparse SVG steps?
| if self.config and bool(getattr(self.config, "use_svg_attention", False)): | ||
| if getattr(self.config, "use_cfg_cache", False) or getattr(self.config, "use_magcache", False): | ||
| raise ValueError("SVG sparse attention cannot be combined with CFG cache or MagCache.") | ||
|
|
There was a problem hiding this comment.
Up on line 348, we disable SVG when svg_spatial_density >= 1.0:
use_svg = bool(getattr(config, "use_svg_attention", False)) and (expert_density < 1.0)Here in __call__, though, we only check getattr(self.config, "use_svg_attention", False) without checking if svg_spatial_density < 1.0. To keep the two checks consistent, could we also check that float(getattr(self.config, "svg_spatial_density", 0.25)) < 1.0 before raising the error?
| "use_svg_attention", | ||
| "svg_implementation", | ||
| "svg_spatial_density", | ||
| "svg_sample_max_row", | ||
| "svg_profile_query_count", | ||
| "svg_profile_seed", | ||
| "svg_dense_layer_fraction", | ||
| "svg_dense_timestep_fraction", | ||
| "svg_active_start_step", | ||
| "svg_active_end_step", | ||
| "svg_active_start_layer", | ||
| "svg_active_end_layer", | ||
| "svg_num_train_timesteps", | ||
| "svg_num_layers", | ||
| "svg_include_first_frame", | ||
| "svg_global_stride", | ||
| "svg_global_offset", | ||
| "svg_high_noise_density", | ||
| "svg_low_noise_density", | ||
| "svg_flash_block_sizes", |
There was a problem hiding this comment.
If use_svg_attention is False, changing any of the svg_* config values won't change the compiled model, so we ideally shouldn't change the AOT cache hash when SVG is turned off.
In aot_cache.py, there is already a helper extract_svg_meta(config, pipeline) that returns {"use_svg_attention": False} when SVG is disabled and only hashes all the svg_* parameters when SVG is enabled. Could we reuse that helper (or only include the svg_* keys when use_svg_attention is True)?
…fault, tests) Use ulysses_custom and cover all 48 layers in the LTX-2 docs example, expose the remaining SVG keys in the LTX-2 YAMLs, drop the Wan-only high/low noise densities from the LTX-2 pipeline and configs, default svg_num_layers to 48, and test the density=1.0 fallback and svg_step_index forwarding.
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 2/2 (config and pipeline side). It depends on #497, which adds the SVG attention path to the LTX2 model.
Changes in this PR:
ltx2_pipeline.py: builds the SVGattention_configfrom the pyconfig and passes it to the transformer; passes the denoising step index (svg_step_index) through both the scanned and unscanned diffusion loops.ltx2_video.yml,ltx2_3_video.yml: add SVG options (disabled by default) andsvg_flash_block_sizesfor the sparse kernel tiling.generate_ltx2.py: add SVG options to the AOT cache metadata keys so compiled artifacts are not reused across SVG settings.docs/svg.md: new SVG-on-TPU documentation covering Wan and LTX-2.tests/ltx2/test_svg_config_propagation_ltx2.py: new config-propagation tests.VABench Evaluation: SVG vs. Dense Attention
The end-to-end results below require both #497 and this PR.
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_visionConfiguration Example
To enable SVG on LTX-2, add the following to
ltx2_video.ymlor override via CLI flags. The VABench run above usedattention: ulysses_customandsvg_spatial_density: 0.25; the example below shows a restricted step/layer window:SVG requires one of the custom Ulysses/ring attention backends (the LTX2 default
attention: flashis not supported).Testing
Run the LTX-2 SVG unit tests from the repository root:
All existing Wan and LTX-2 unit tests continue to pass.