Repository navigation
[PR 2/2] LTX2: SVG config and pipeline - #498
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 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.
…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.
1879bc4 to
a33c8b2
Compare
…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.
6858a71 to
93c4518
Compare
…AOT key) Add svg_attention_enabled() so model construction, the CFG cache/MagCache check and the AOT cache key agree that SVG is off at svg_spatial_density >= 1.0. Key the LTX-2 AOT cache on svg_* settings, including those recorded in the transformer attention_config, only when SVG is active. Note in docs/svg.md that the default 5-step profile never reaches an SVG window that starts at step 10.
Expose Sparse VideoGen (SVG) attention for LTX-2 end to end. SVG is off
by default and only active when use_svg_attention is set and
svg_spatial_density < 1.0 (svg_attention_enabled()).
- ltx2_pipeline.py: build the SVG attention_config from the pyconfig and
pass the denoising step index (svg_step_index) through both the scanned
and unscanned diffusion loops. Reject SVG combined with CFG cache or
MagCache only when SVG is active. svg_num_layers defaults to 48.
- ltx2_video.yml, ltx2_3_video.yml: add all SVG options (disabled by
default) and svg_flash_block_sizes for the sparse kernel tiling. The
Wan-only svg_{high,low}_noise_density settings are not used by LTX-2.
- generate_ltx2.py: key the AOT cache on SVG settings, including those in
the transformer attention_config, only when SVG is active, so unused
svg_* values do not invalidate compiled artifacts.
- docs/svg.md: SVG-on-TPU documentation for Wan and LTX-2, including a
copyable LTX-2 example and a note on profiler_steps vs the SVG window.
- Tests: config propagation, density-1.0 fallback, svg_step_index
forwarding, svg_attention_enabled(), and AOT cache identity.
a793373 to
2a52ee9
Compare
|
|
||
| [Sparse VideoGen (SVG)](https://arxiv.org/abs/2502.01776) observes that attention heads often favor different patterns. Spatial heads concentrate attention within a frame or nearby frames. Temporal heads concentrate attention around corresponding spatial positions across frames. These patterns let us approximate dense attention while computing fewer interactions. | ||
|
|
||
|  |
There was a problem hiding this comment.
Blocker: docs/svg.md references three relative image files that are not committed in docs/ or anywhere in the repository:
- Line 11:
images/svg/head-patterns.png - Line 31:
images/svg/attention-tiles.png - Line 140:
images/svg/dense-vs-svg.png
Please add docs/images/svg/{head-patterns,attention-tiles,dense-vs-svg}.png to the PR (or remove the broken image tags if the assets are not meant to be checked in).
| ltx2_config["attention_config"] = { | ||
| "use_base2_exp": getattr(config, "use_base2_exp", False), | ||
| "use_experimental_scheduler": getattr(config, "use_experimental_scheduler", False), | ||
| "ulysses_shards": getattr(config, "ulysses_shards", -1), | ||
| "ulysses_attention_chunks": getattr(config, "ulysses_attention_chunks", 1), | ||
| "use_svg_attention": svg_attention_enabled(config), | ||
| "svg_implementation": _cfg("svg_implementation", "official_svg"), | ||
| "svg_spatial_density": float(_cfg("svg_spatial_density", 0.25)), | ||
| "svg_sample_max_row": _cfg("svg_sample_max_row", 10000), | ||
| "svg_profile_query_count": _cfg("svg_profile_query_count", 64), | ||
| "svg_profile_seed": _cfg("svg_profile_seed", 0), | ||
| "svg_dense_layer_fraction": _cfg("svg_dense_layer_fraction", 0.0), | ||
| "svg_dense_timestep_fraction": _cfg("svg_dense_timestep_fraction", 0.0), | ||
| "svg_active_start_step": _cfg("svg_active_start_step", -1), | ||
| "svg_active_end_step": _cfg("svg_active_end_step", -1), | ||
| "svg_active_start_layer": _cfg("svg_active_start_layer", -1), | ||
| "svg_active_end_layer": _cfg("svg_active_end_layer", -1), | ||
| "svg_num_train_timesteps": _cfg("svg_num_train_timesteps", 1000), | ||
| "svg_num_layers": _cfg("svg_num_layers", ltx2_config.get("num_layers", 48)), | ||
| "svg_include_first_frame": _cfg("svg_include_first_frame", True), | ||
| "svg_global_stride": _cfg("svg_global_stride", 0), | ||
| "svg_global_offset": _cfg("svg_global_offset", 0), | ||
| "svg_flash_block_sizes": getattr(config, "svg_flash_block_sizes", None) or None, | ||
| } |
There was a problem hiding this comment.
Blocker: ltx2_aot_metadata (generate_ltx2.py:L206-L210) strips svg_* keys from transformer_config["attention_config"] when svg_attention_enabled(config) is False. However, in real execution aot_cache.cached_jit (run_diffusion_loop and transformer_forward_pass) also hashes _dynamic_signature((graphdef, ...)) via _graphdef_desc(graphdef).
Because create_sharded_logical_transformer unconditionally populates all svg_* keys into ltx2_config["attention_config"] even when svg_attention_enabled(config) is False, those values are stored as Static attributes on LTX2VideoTransformer3DModel.attention_config and LTX2Attention.svg_* inside graphdef. As a result, aot_cache._dynamic_signature((graphdef,), {}) still changes when unused svg_* settings are modified (d382ae2efd22 vs. c92170d2e500 vs. 21a6fcb2b699 at svg_spatial_density=1.0), causing an AOT cache miss.
Could we only populate the svg_* keys in ltx2_config["attention_config"] when svg_attention_enabled(config) is True (and otherwise just set {"use_svg_attention": False})?
| "use_base2_exp": getattr(config, "use_base2_exp", False), | ||
| "use_experimental_scheduler": getattr(config, "use_experimental_scheduler", False), | ||
| "ulysses_shards": getattr(config, "ulysses_shards", -1), | ||
| "ulysses_attention_chunks": getattr(config, "ulysses_attention_chunks", 1), |
There was a problem hiding this comment.
Nit: "use_base2_exp", "use_experimental_scheduler", "ulysses_shards", and "ulysses_attention_chunks" are already set directly on ltx2_config at lines 336–339 and are ignored by svg_attention.init_svg_config. Passing them again inside ltx2_config["attention_config"] (and in test_generate_ltx2_aot.py:L94) is dead config and can be removed.
| def _set_svg(self, **svg): | ||
| for name, value in svg.items(): | ||
| setattr(self.config, name, value) | ||
| attention_config = {"use_base2_exp": False} | ||
| attention_config.update({k: v for k, v in svg.items() if k.startswith("svg_")}) | ||
| # Like create_sharded_logical_transformer, record the effective SVG flag. | ||
| attention_config["use_svg_attention"] = generate_ltx2.svg_attention_enabled(self.config) | ||
| self.pipeline.transformer.config = {"attention_config": attention_config} | ||
|
|
||
| def test_unused_svg_settings_do_not_change_cache_fingerprint(self): | ||
| self._set_svg(use_svg_attention=False, svg_spatial_density=0.25, svg_active_start_step=10) | ||
| _, fingerprint_a = self._fingerprint("same-revision") | ||
| self._set_svg(use_svg_attention=False, svg_spatial_density=0.5, svg_active_start_step=20) | ||
| _, fingerprint_b = self._fingerprint("same-revision") | ||
|
|
||
| self.assertEqual(fingerprint_a, fingerprint_b) | ||
|
|
||
| def test_svg_settings_change_cache_fingerprint_when_svg_is_active(self): | ||
| self._set_svg(use_svg_attention=True, svg_spatial_density=0.25, svg_active_start_step=10) | ||
| _, fingerprint_a = self._fingerprint("same-revision") | ||
| self._set_svg(use_svg_attention=True, svg_spatial_density=0.5, svg_active_start_step=10) | ||
| _, fingerprint_b = self._fingerprint("same-revision") | ||
| self._set_svg(use_svg_attention=True, svg_spatial_density=0.5, svg_active_start_step=20) | ||
| _, fingerprint_c = self._fingerprint("same-revision") | ||
|
|
||
| self.assertNotEqual(fingerprint_a, fingerprint_b) | ||
| self.assertNotEqual(fingerprint_b, fingerprint_c) | ||
|
|
||
| def test_svg_at_density_one_shares_the_dense_cache_fingerprint(self): | ||
| self._set_svg(use_svg_attention=False, svg_spatial_density=0.25) | ||
| _, dense = self._fingerprint("same-revision") | ||
| self._set_svg(use_svg_attention=True, svg_spatial_density=1.0) | ||
| _, density_one = self._fingerprint("same-revision") | ||
| self._set_svg(use_svg_attention=True, svg_spatial_density=0.25) | ||
| _, sparse = self._fingerprint("same-revision") | ||
|
|
||
| self.assertEqual(dense, density_one) | ||
| self.assertNotEqual(dense, sparse) |
There was a problem hiding this comment.
Important: These tests check ltx2_aot_metadata using a mocked self.pipeline.transformer, which missed that nnx.split(transformer)[0] (graphdef) still baked inactive svg_* attributes into aot_cache._dynamic_signature. In addition to checking _fingerprint, could we also assert that aot_cache._dynamic_signature((nnx.split(transformer)[0],), {}) on a real create_sharded_logical_transformer model stays identical across unused svg_* overrides and at svg_spatial_density=1.0?
| enable_vae_tiling: False | ||
| enable_vae_slicing: False |
There was a problem hiding this comment.
Question / Nit: enable_vae_tiling: False and enable_vae_slicing: False were added here and in ltx2_3_video.yml:L151-L152, which looks unrelated to SVG attention. Was this intentional (to expose existing generate_ltx2.py flags for CLI override), or left over from local benchmarking?
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.