From 1b937b701ac631ea963f044c8fca21c723aade3d Mon Sep 17 00:00:00 2001 From: Jitendra Jalwaniya Date: Tue, 29 Sep 2026 12:05:48 +0000 Subject: [PATCH 1/2] ltx2: wire SVG config through pipeline, configs, AOT metadata, docs --- docs/svg.md | 165 ++++++++++++++ src/maxdiffusion/configs/ltx2_3_video.yml | 20 ++ src/maxdiffusion/configs/ltx2_video.yml | 20 ++ src/maxdiffusion/generate_ltx2.py | 20 ++ .../pipelines/ltx2/ltx2_pipeline.py | 55 ++++- .../ltx2/test_svg_config_propagation_ltx2.py | 209 ++++++++++++++++++ 6 files changed, 485 insertions(+), 4 deletions(-) create mode 100644 docs/svg.md create mode 100644 src/maxdiffusion/tests/ltx2/test_svg_config_propagation_ltx2.py diff --git a/docs/svg.md b/docs/svg.md new file mode 100644 index 000000000..88b260d7e --- /dev/null +++ b/docs/svg.md @@ -0,0 +1,165 @@ +# Sparse VideoGen attention on TPUs + +Video diffusion models generate a video through a sequence of denoising steps. At each step, attention lets each video token gather information from other tokens across space and time. This becomes expensive as the resolution and number of frames grow: dense attention considers every query–key pair, even though many interactions contribute very little to the output. + +Sparse spatiotemporal attention takes advantage of this structure. Instead of attending everywhere, a query attends to a smaller set of positions chosen to capture the spatial and temporal information it needs. + +## How SVG chooses where to attend + +[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. + +![Attention from the same query in a spatial head and a temporal head, shown across six latent frames.](images/svg/head-patterns.png) + +*Observed attention weights for the same query in two heads. The spatial head concentrates 94.8% of its attention on the query frame, while the temporal head places substantial attention near corresponding positions in other frames. Cyan squares mark the query's spatial position; `m` is the attention mass in each displayed frame. Colors use a shared logarithmic scale. These are examples of observed attention, rather than the masks themselves.* + +SVG makes the choice separately for each head: + +1. **Profile a few queries.** Compute dense attention outputs for a small sample of query tokens, using all keys. +2. **Compare two patterns.** Compute the sampled outputs under spatial and temporal masks, then measure each one's error relative to dense attention. +3. **Use the better approximation.** Select the lower-error pattern and apply sparse attention to all queries in that head. + +The choice is recomputed at each active layer and denoising step. A head does not need to keep the same assignment throughout generation. Sparsity is configurable, so the same method can trade a smaller approximation error for a larger reduction in computation. + +This implementation reimplements routing, token placement, and kernel execution for Wan and LTX2 in MaxDiffusion. The [original SVG implementation](https://github.com/svg-project/Sparse-VideoGen) provides the reference method. + +## Making sparse attention efficient on TPU + +Skipping query–key interactions only helps if the hardware can skip the corresponding work efficiently. Our implementation arranges tokens so that both spatial and temporal patterns can use the same local-band attention kernel. Spatial heads keep frame-major order; temporal heads group corresponding spatial positions across frames. Outputs are restored to their original order afterward. + +TPUs compute attention in tiles. A tile can lie entirely inside the sparse pattern, entirely outside it, or cross its boundary. + +A local attention band over a query–key tile grid, highlighting full, boundary, and skipped tiles. + +*Sparse pattern before tile rounding. Query and key indices refer to the selected token layout. Blue indicates retained interactions and gray indicates skipped interactions. The prefix anchor is omitted for clarity.* + +We round boundary tiles to either keep or skip them, approximately preserving the attention-pair budget of the original pattern. This slightly changes which interactions are retained, but lets all selected interior tiles run through one kernel without per-token sparse masking. Only tiles touching sequence padding need an additional validity mask; their outputs are combined with the main result using a numerically stable merge. + +Token placement and restoration run after the Ulysses exchange, on each device's local heads. This limits the layout work to the heads that device will actually process. + +The masks also include an optional prefix anchor. The `svg_include_first_frame` option retains the first `H × W` keys in the selected layout. For temporal heads, that prefix spans spatial positions across frames, so it does not correspond to the original first video frame. + +## Configuration and usage + +SVG is disabled by default. To enable it, set `use_svg_attention=True` and choose the densities and the steps and layers where sparsity should be active. Calls outside that interval continue to use dense attention. + +For example, these overrides select the moderate Wan2.2 policy: + +```yaml +use_svg_attention: True +svg_high_noise_density: 0.50 +svg_low_noise_density: 0.20 +svg_active_start_step: 11 +svg_active_end_step: 40 +svg_active_start_layer: 1 +svg_active_end_layer: 40 +svg_profile_query_count: 64 +svg_sample_max_row: 10000 +svg_profile_seed: 0 +svg_include_first_frame: True +``` + +For LTX-2 (`src/maxdiffusion/configs/ltx2_video.yml` or `ltx2_3_video.yml`), use `svg_spatial_density`: + +```yaml +attention: ulysses_ring_custom_fixed_m +use_svg_attention: True +svg_spatial_density: 0.20 +svg_active_start_step: 10 +svg_active_end_step: 35 +svg_active_start_layer: 1 +svg_active_end_layer: 28 +svg_profile_query_count: 64 +svg_sample_max_row: 10000 +svg_profile_seed: 0 +svg_include_first_frame: True +``` + +Step and layer intervals are zero-based and half-open: `[11, 40)` includes steps 11 through 39. Steps refer to denoising iterations, not noise-timestep values. + +| Option | What it controls | +|---|---| +| `svg_spatial_density` | Density for single-expert Wan models and LTX-2; applies to either selected head pattern. | +| `svg_high_noise_density`, `svg_low_noise_density` | Separate densities for Wan2.2's two experts. | +| `svg_active_start_step`, `svg_active_end_step` | Denoising steps where SVG is enabled. | +| `svg_active_start_layer`, `svg_active_end_layer` | Transformer layers where SVG is enabled. | +| `svg_profile_query_count` | Number of queries sampled when choosing each head's pattern. | +| `svg_sample_max_row` | Limits query sampling to this prefix of the sequence. | +| `svg_profile_seed` | Seed for reproducible query sampling. | +| `svg_include_first_frame` | Enables the prefix anchor in the selected layout. | +| `svg_flash_block_sizes` | Sparse kernel tiling; an empty mapping uses `flash_block_sizes`. | + +Density controls the local-band width. The anchor and tile rounding affect the actual number of retained interactions, so density is not itself the fraction of total transformer FLOPs retained. + +The following is an example attention configuration for eight devices. Add it and the policy above to a Wan2.2 configuration used by `src/maxdiffusion/generate_wan.py`: + +```yaml +attention: ulysses_ring_custom_fixed_m +ici_data_parallelism: 2 +ici_context_parallelism: 4 +ulysses_shards: 2 +flash_block_sizes: + block_q: 6400 + block_kv: 2048 + block_kv_compute: 2048 + block_kv_compute_in: 1024 + heads_per_tile: 1 + vmem_limit_bytes: 67108864 +svg_flash_block_sizes: + block_q: 3328 + block_kv: 2816 + block_kv_compute: 256 + block_kv_compute_in: 256 + heads_per_tile: 1 + vmem_limit_bytes: 67108864 +``` + +Dense calls retain their configured ring split. SVG exchanges over the whole context axis, giving Ring2/Ulysses2 for dense calls and Ring1/Ulysses4 for sparse calls in this example. Tile sizes may need tuning for other shapes and devices. + +### Supported configurations + +SVG supports inference through the four custom Ulysses/ring attention backends. It requires matching self-attention QKV shapes, a matching video-token grid, and `heads_per_tile=1`. Heads and sequence length must divide evenly across the context shards, including any heads created by folding an unsharded batch. + +Training, Animate, external attention masks, periodic support, and chunked Ulysses are unsupported. SVG also cannot be combined with CFG cache or MagCache. + +## Performance and quality + +At 720p, SVG offers a configurable tradeoff between denoising latency and similarity to dense generation. The following Wan2.2 results use TPU v6e-8, 81 frames, and 40 denoising steps. Times are medians of three warm runs against same-node optimized fixed-M dense controls, using one prompt and seed. + +| Policy | Estimated total transformer FLOPs saved | Dense denoising | SVG denoising | Denoising speedup | PSNR (dB) | +|---|---:|---:|---:|---:|---:| +| Conservative | ≈27.2% | 153.46 s | 136.11 s | **1.13×** | **26.47** | +| Moderate | ≈32.2% | 153.37 s | 127.94 s | **1.20×** | **26.14** | +| Aggressive | ≈37.3% | 153.50 s | 119.86 s | **1.28×** | **24.80** | + +*PSNR is measured against dense outputs using FFmpeg's aggregate YUV metric.* + +In an earlier evaluation using the moderate SVG policy, Wan2.2 at 720p retained **98.1% of the dense baseline’s mean VBench dimension score** across a 31-prompt, 16-dimension screening subset. + +## Qualitative example + +![Three rows of giraffe video frames, with dense attention on the left and SVG on the right.](images/svg/dense-vs-svg.png) + +*Dense attention (left) and SVG (right), shown at three video frames. The overall scene and subject arrangement remain similar, with visible differences in ground texture, background detail, and coat patterns.* + +This illustration is separate from the benchmark measurements above. + +## Tests and profiling + +Run the SVG tests from the repository root: + +```bash +python -m pytest -q \ + src/maxdiffusion/tests/wan/svg_attention_test.py \ + src/maxdiffusion/tests/wan/svg_balanced_rounding_test.py \ + src/maxdiffusion/tests/wan/svg_config_propagation_test.py \ + src/maxdiffusion/tests/wan/svg_head_local_test.py \ + src/maxdiffusion/tests/wan/wan_pipeline_signature_test.py \ + src/maxdiffusion/tests/ltx2/test_svg_config_propagation_ltx2.py \ + src/maxdiffusion/tests/ltx2/test_svg_attention_ltx2.py +``` + +The tests cover routing and layout, configuration propagation, schedules, sharding, unsupported configurations, and numerical agreement with attention references. Production-kernel tests cover sparse and density-one support, aligned and padded sequences, and natural and base-2 exponentials. + +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. diff --git a/src/maxdiffusion/configs/ltx2_3_video.yml b/src/maxdiffusion/configs/ltx2_3_video.yml index 9c9c5432a..8c01da9d2 100644 --- a/src/maxdiffusion/configs/ltx2_3_video.yml +++ b/src/maxdiffusion/configs/ltx2_3_video.yml @@ -17,6 +17,24 @@ a2v_attention_kernel: 'flash' v2a_attention_kernel: 'dot_product' attention_sharding_uniform: True precision: 'bf16' +# Sparse VideoGen (SVG) Attention configuration +use_svg_attention: False +svg_spatial_density: 0.25 +svg_sample_max_row: 10000 +svg_profile_query_count: 64 +svg_profile_seed: 0 +svg_active_start_step: -1 +svg_active_end_step: -1 +svg_active_start_layer: -1 +svg_active_end_layer: -1 +svg_include_first_frame: True +svg_high_noise_density: -1.0 +svg_low_noise_density: -1.0 +# Tiling for the sparse SVG kernel only. Empty means "reuse flash_block_sizes", +# which is rarely what you want: the dense ring kernel is tuned for large kv +# compute blocks and the sparse kernel for small ones. Must stay a dict so the +# command line can override it with JSON. +svg_flash_block_sizes: {} scan_layers: True names_which_can_be_saved: [] names_which_can_be_offloaded: [] @@ -125,6 +143,8 @@ profiler_steps: 5 enable_jax_named_scopes: False replicate_vae: False +enable_vae_tiling: False +enable_vae_slicing: False run_text_encoder_on_tpu: False # Dynamically disables VAE slicing and distributes the batch dimension to avoid HBM OOM for larger batch sizes. diff --git a/src/maxdiffusion/configs/ltx2_video.yml b/src/maxdiffusion/configs/ltx2_video.yml index 23a4b104a..4778e285c 100644 --- a/src/maxdiffusion/configs/ltx2_video.yml +++ b/src/maxdiffusion/configs/ltx2_video.yml @@ -17,6 +17,24 @@ a2v_attention_kernel: 'dot_product' v2a_attention_kernel: 'dot_product' attention_sharding_uniform: True precision: 'bf16' +# Sparse VideoGen (SVG) Attention configuration +use_svg_attention: False +svg_spatial_density: 0.25 +svg_sample_max_row: 10000 +svg_profile_query_count: 64 +svg_profile_seed: 0 +svg_active_start_step: -1 +svg_active_end_step: -1 +svg_active_start_layer: -1 +svg_active_end_layer: -1 +svg_include_first_frame: True +svg_high_noise_density: -1.0 +svg_low_noise_density: -1.0 +# Tiling for the sparse SVG kernel only. Empty means "reuse flash_block_sizes", +# which is rarely what you want: the dense ring kernel is tuned for large kv +# compute blocks and the sparse kernel for small ones. Must stay a dict so the +# command line can override it with JSON. +svg_flash_block_sizes: {} # For scanning transformer layers scan_layers: True @@ -131,6 +149,8 @@ enable_jax_named_scopes: False replicate_vae: False use_bwe: False +enable_vae_tiling: False +enable_vae_slicing: False run_text_encoder_on_tpu: False # Dynamically disables VAE slicing and distributes the batch dimension to avoid HBM OOM for larger batch sizes. diff --git a/src/maxdiffusion/generate_ltx2.py b/src/maxdiffusion/generate_ltx2.py index b60f223b8..af278a61e 100644 --- a/src/maxdiffusion/generate_ltx2.py +++ b/src/maxdiffusion/generate_ltx2.py @@ -223,6 +223,26 @@ def ltx2_aot_metadata(config, pipeline, source_revision=None): "sharding", "weights_dtype", "activations_dtype", + "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", ] config_dict = {k: getattr(config, k, None) for k in config_keys} diff --git a/src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py b/src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py index eeaad473d..a60166f5e 100644 --- a/src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py +++ b/src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py @@ -336,6 +336,44 @@ def create_model(rngs: nnx.Rngs, ltx2_config: dict): dit_specs = get_sharding_specs(transformer_strategy, "ltx2_dit") ltx2_config["sharding_specs"] = dit_specs + def _cfg(name, default): + """Returns config., falling back to default if missing or None.""" + value = getattr(config, name, None) + return default if value is None else value + + high_density = float(_cfg("svg_high_noise_density", -1.0)) + low_density = float(_cfg("svg_low_noise_density", -1.0)) + expert_density = float(_cfg("svg_spatial_density", 0.25)) + + use_svg = bool(getattr(config, "use_svg_attention", False)) and (expert_density < 1.0) + + 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": use_svg, + "svg_implementation": _cfg("svg_implementation", "official_svg"), + "svg_spatial_density": expert_density, + "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", 28)), + "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_high_noise_density": high_density, + "svg_low_noise_density": low_density, + "svg_flash_block_sizes": getattr(config, "svg_flash_block_sizes", None) or None, + } + # 2. eval_shape p_model_factory = partial(create_model, ltx2_config=ltx2_config) transformer = nnx.eval_shape(p_model_factory, rngs=rngs) @@ -1832,6 +1870,10 @@ def __call__( if stg_scale > 0.0 and guidance_scale <= 1.0: raise ValueError("Spatio-temporal guidance requires guidance_scale > 1.0.") + 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.") + # 2. Encode inputs (Text) t0_encode = time.perf_counter() ( @@ -2136,6 +2178,7 @@ def __call__( is_cfg_stg_mode=do_cfg and do_stg, kv_cache=kv_cache, rope_cache=rope_cache, + svg_step_index=jnp.asarray(i, dtype=jnp.int32), ) latents_step, audio_latents_step = _select_guidance_latents( @@ -2437,6 +2480,7 @@ def transformer_forward_pass( kv_cache=None, rope_cache=None, time_embed_cache=None, + svg_step_index=None, ): """Forward pass for the transformer.""" # pylint: disable=too-many-positional-arguments,unused-argument @@ -2491,6 +2535,7 @@ def transformer_forward_pass( cached_kv=kv_cache, rope_cache=rope_cache, time_embed_cache=time_embed_cache, + svg_step_index=svg_step_index, ) return noise_pred, noise_pred_audio @@ -2598,9 +2643,9 @@ def run_diffusion_loop( def scan_body(carry, inputs): if use_kv_cache: - t, sigma_t, time_embed_cache_step = inputs + t, sigma_t, time_embed_cache_step, step_idx = inputs else: - t, sigma_t = inputs + t, sigma_t, step_idx = inputs time_embed_cache_step = None latents, audio_latents, s_state = carry @@ -2634,6 +2679,7 @@ def scan_body(carry, inputs): kv_cache=kv_cache, rope_cache=rope_cache, time_embed_cache=time_embed_cache_step, + svg_step_index=step_idx, ) latents_step, audio_latents_step = _select_guidance_latents( @@ -2680,10 +2726,11 @@ def scan_body(carry, inputs): initial_carry = (latents_jax, audio_latents_jax, scheduler_state) + step_indices = jnp.arange(len(timesteps_jax), dtype=jnp.int32) if use_kv_cache: - scan_inputs = (timesteps_jax, sigmas, time_embed_cache_full) + scan_inputs = (timesteps_jax, sigmas, time_embed_cache_full, step_indices) else: - scan_inputs = (timesteps_jax, sigmas) + scan_inputs = (timesteps_jax, sigmas, step_indices) final_carry, _ = jax.lax.scan(scan_body, initial_carry, scan_inputs) diff --git a/src/maxdiffusion/tests/ltx2/test_svg_config_propagation_ltx2.py b/src/maxdiffusion/tests/ltx2/test_svg_config_propagation_ltx2.py new file mode 100644 index 000000000..0a7d668df --- /dev/null +++ b/src/maxdiffusion/tests/ltx2/test_svg_config_propagation_ltx2.py @@ -0,0 +1,209 @@ +""" +Copyright 2026 Google LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + https://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. + +Configuration propagation tests for Sparse VideoGen attention in LTX2. +""" + +from types import SimpleNamespace +import unittest +from unittest.mock import MagicMock + +from flax import nnx +import jax +import numpy as np +from jax.sharding import Mesh + +from maxdiffusion.pipelines.ltx2.ltx2_pipeline import create_sharded_logical_transformer, LTX2Pipeline + + +class LTX2SVGConfigPropagationTest(unittest.TestCase): + + def _get_test_ltx2_config(self): + return { + "num_layers": 2, + "num_attention_heads": 2, + "attention_head_dim": 32, + "cross_attention_dim": 64, + "in_channels": 16, + "out_channels": 16, + "audio_in_channels": 4, + "audio_out_channels": 4, + "audio_num_attention_heads": 2, + "audio_attention_head_dim": 32, + "audio_cross_attention_dim": 64, + "caption_channels": 32, + "patch_size": 1, + "patch_size_t": 1, + "pos_embed_max_pos": 16, + "base_height": 32, + "base_width": 32, + "audio_pos_embed_max_pos": 16, + "audio_sampling_rate": 16000, + "audio_hop_length": 160, + "audio_scale_factor": 1.0, + } + + def test_svg_config_propagation_through_transformer_construction(self): + test_ltx2_config = self._get_test_ltx2_config() + devices = np.array(jax.devices()[:1]).reshape((1, 1)) + mesh = Mesh(devices, ("data", "fsdp")) + rngs = nnx.Rngs(0) + + cfg = SimpleNamespace( + use_svg_attention=True, + svg_spatial_density=0.25, + svg_active_start_step=8, + svg_active_end_step=30, + svg_active_start_layer=1, + svg_active_end_layer=28, + precision="DEFAULT", + flash_block_sizes={}, + activations_dtype="bfloat16", + weights_dtype="bfloat16", + attention="dot_product", + a2v_attention_kernel="dot_product", + v2a_attention_kernel="dot_product", + remat_policy="none", + names_which_can_be_saved=[], + names_which_can_be_offloaded=[], + flash_min_seq_length=0, + dropout=0.0, + scan_layers=False, + enable_jax_named_scopes=False, + use_base2_exp=False, + use_experimental_scheduler=False, + logical_axis_rules=(), + ) + + m = create_sharded_logical_transformer( + devices_array=devices, + mesh=mesh, + rngs=rngs, + config=cfg, + restored_checkpoint={"ltx2_config": dict(test_ltx2_config), "ltx2_state": {}}, + subfolder="", + ) + + # Video self-attention (attn1) must have SVG enabled + first_block = m.transformer_blocks[0] + self.assertTrue(first_block.attn1.use_svg_attention) + self.assertEqual(first_block.attn1.svg_spatial_density, 0.25) + self.assertEqual(first_block.attn1.svg_active_start_step, 8) + self.assertEqual(first_block.attn1.svg_active_end_step, 30) + self.assertEqual(first_block.attn1.svg_active_start_layer, 1) + self.assertEqual(first_block.attn1.svg_active_end_layer, 28) + + # Audio self-attention (audio_attn1) and cross attentions must remain dense + self.assertFalse(first_block.audio_attn1.use_svg_attention) + self.assertFalse(first_block.attn2.use_svg_attention) + self.assertFalse(first_block.audio_attn2.use_svg_attention) + self.assertFalse(first_block.audio_to_video_attn.use_svg_attention) + self.assertFalse(first_block.video_to_audio_attn.use_svg_attention) + + def test_svg_block_sizes_are_independent_of_the_dense_block_sizes(self): + test_ltx2_config = self._get_test_ltx2_config() + devices = np.array(jax.devices()[:1]).reshape((1, 1)) + mesh = Mesh(devices, ("data", "fsdp")) + sparse_blocks = {"block_q": 3328, "block_kv": 2816, "block_kv_compute": 256, "block_kv_compute_in": 256} + common = { + "use_svg_attention": True, + "svg_spatial_density": 0.25, + "precision": "DEFAULT", + "activations_dtype": "bfloat16", + "weights_dtype": "bfloat16", + "attention": "dot_product", + "a2v_attention_kernel": "dot_product", + "v2a_attention_kernel": "dot_product", + "remat_policy": "none", + "names_which_can_be_saved": [], + "names_which_can_be_offloaded": [], + "flash_min_seq_length": 0, + "dropout": 0.0, + "scan_layers": False, + "enable_jax_named_scopes": False, + "use_base2_exp": False, + "use_experimental_scheduler": False, + "logical_axis_rules": (), + } + + def build(**extra): + return create_sharded_logical_transformer( + devices_array=devices, + mesh=mesh, + rngs=nnx.Rngs(0), + config=SimpleNamespace(**common, **extra), + restored_checkpoint={"ltx2_config": dict(test_ltx2_config), "ltx2_state": {}}, + subfolder="", + ) + + tuned = build(flash_block_sizes={}, svg_flash_block_sizes=sparse_blocks) + self.assertEqual(tuned.transformer_blocks[0].attn1.svg_flash_block_sizes, sparse_blocks) + + default = build(flash_block_sizes={}, svg_flash_block_sizes={}) + self.assertIsNone(default.transformer_blocks[0].attn1.svg_flash_block_sizes) + + absent = build(flash_block_sizes={}) + self.assertIsNone(absent.transformer_blocks[0].attn1.svg_flash_block_sizes) + + def test_svg_disabled_when_use_svg_attention_is_false(self): + test_ltx2_config = self._get_test_ltx2_config() + devices = np.array(jax.devices()[:1]).reshape((1, 1)) + mesh = Mesh(devices, ("data", "fsdp")) + + cfg_dense = SimpleNamespace( + use_svg_attention=False, + precision="DEFAULT", + flash_block_sizes={}, + activations_dtype="bfloat16", + weights_dtype="bfloat16", + attention="dot_product", + a2v_attention_kernel="dot_product", + v2a_attention_kernel="dot_product", + remat_policy="none", + names_which_can_be_saved=[], + names_which_can_be_offloaded=[], + flash_min_seq_length=0, + dropout=0.0, + scan_layers=False, + enable_jax_named_scopes=False, + use_base2_exp=False, + use_experimental_scheduler=False, + logical_axis_rules=(), + ) + m_dense = create_sharded_logical_transformer( + devices_array=devices, + mesh=mesh, + rngs=nnx.Rngs(0), + config=cfg_dense, + restored_checkpoint={"ltx2_config": dict(test_ltx2_config), "ltx2_state": {}}, + subfolder="", + ) + self.assertFalse(m_dense.transformer_blocks[0].attn1.use_svg_attention) + + def test_cfg_and_magcache_incompatibility_validation(self): + pipeline = LTX2Pipeline.__new__(LTX2Pipeline) + pipeline.check_inputs = MagicMock() + pipeline.config = SimpleNamespace(use_svg_attention=True, use_cfg_cache=True, use_magcache=False) + + with self.assertRaisesRegex(ValueError, "SVG sparse attention cannot be combined with CFG cache or MagCache"): + pipeline(prompt="test prompt") + + pipeline.config = SimpleNamespace(use_svg_attention=True, use_cfg_cache=False, use_magcache=True) + with self.assertRaisesRegex(ValueError, "SVG sparse attention cannot be combined with CFG cache or MagCache"): + pipeline(prompt="test prompt") + + +if __name__ == "__main__": + unittest.main() From b8acc0f7519ab635246b6e189fdd353a04f28315 Mon Sep 17 00:00:00 2001 From: Jitendra Jalwaniya Date: Thu, 1 Oct 2026 12:03:37 +0000 Subject: [PATCH 2/2] ltx2/svg: address #498 review (docs example, config keys, 48-layer default, 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. --- docs/svg.md | 6 +- src/maxdiffusion/configs/ltx2_3_video.yml | 9 ++- src/maxdiffusion/configs/ltx2_video.yml | 9 ++- src/maxdiffusion/generate_ltx2.py | 2 - .../pipelines/ltx2/ltx2_pipeline.py | 6 +- .../ltx2/test_svg_config_propagation_ltx2.py | 73 +++++++++++++++++++ 6 files changed, 91 insertions(+), 14 deletions(-) diff --git a/docs/svg.md b/docs/svg.md index 88b260d7e..ed1dd9ec2 100644 --- a/docs/svg.md +++ b/docs/svg.md @@ -58,16 +58,16 @@ svg_profile_seed: 0 svg_include_first_frame: True ``` -For LTX-2 (`src/maxdiffusion/configs/ltx2_video.yml` or `ltx2_3_video.yml`), use `svg_spatial_density`: +For LTX-2 (`src/maxdiffusion/configs/ltx2_video.yml` or `ltx2_3_video.yml`), use `svg_spatial_density`. LTX-2 and LTX-2.3 have 48 transformer layers, so `[1, 48)` covers every layer after layer 0: ```yaml -attention: ulysses_ring_custom_fixed_m +attention: ulysses_custom use_svg_attention: True svg_spatial_density: 0.20 svg_active_start_step: 10 svg_active_end_step: 35 svg_active_start_layer: 1 -svg_active_end_layer: 28 +svg_active_end_layer: 48 svg_profile_query_count: 64 svg_sample_max_row: 10000 svg_profile_seed: 0 diff --git a/src/maxdiffusion/configs/ltx2_3_video.yml b/src/maxdiffusion/configs/ltx2_3_video.yml index 8c01da9d2..58af55690 100644 --- a/src/maxdiffusion/configs/ltx2_3_video.yml +++ b/src/maxdiffusion/configs/ltx2_3_video.yml @@ -28,8 +28,13 @@ svg_active_end_step: -1 svg_active_start_layer: -1 svg_active_end_layer: -1 svg_include_first_frame: True -svg_high_noise_density: -1.0 -svg_low_noise_density: -1.0 +svg_implementation: 'official_svg' +svg_dense_layer_fraction: 0.0 +svg_dense_timestep_fraction: 0.0 +svg_num_train_timesteps: 1000 +svg_num_layers: 48 +svg_global_stride: 0 +svg_global_offset: 0 # Tiling for the sparse SVG kernel only. Empty means "reuse flash_block_sizes", # which is rarely what you want: the dense ring kernel is tuned for large kv # compute blocks and the sparse kernel for small ones. Must stay a dict so the diff --git a/src/maxdiffusion/configs/ltx2_video.yml b/src/maxdiffusion/configs/ltx2_video.yml index 4778e285c..6e8ddaa8e 100644 --- a/src/maxdiffusion/configs/ltx2_video.yml +++ b/src/maxdiffusion/configs/ltx2_video.yml @@ -28,8 +28,13 @@ svg_active_end_step: -1 svg_active_start_layer: -1 svg_active_end_layer: -1 svg_include_first_frame: True -svg_high_noise_density: -1.0 -svg_low_noise_density: -1.0 +svg_implementation: 'official_svg' +svg_dense_layer_fraction: 0.0 +svg_dense_timestep_fraction: 0.0 +svg_num_train_timesteps: 1000 +svg_num_layers: 48 +svg_global_stride: 0 +svg_global_offset: 0 # Tiling for the sparse SVG kernel only. Empty means "reuse flash_block_sizes", # which is rarely what you want: the dense ring kernel is tuned for large kv # compute blocks and the sparse kernel for small ones. Must stay a dict so the diff --git a/src/maxdiffusion/generate_ltx2.py b/src/maxdiffusion/generate_ltx2.py index af278a61e..529d3badd 100644 --- a/src/maxdiffusion/generate_ltx2.py +++ b/src/maxdiffusion/generate_ltx2.py @@ -240,8 +240,6 @@ def ltx2_aot_metadata(config, pipeline, source_revision=None): "svg_include_first_frame", "svg_global_stride", "svg_global_offset", - "svg_high_noise_density", - "svg_low_noise_density", "svg_flash_block_sizes", ] config_dict = {k: getattr(config, k, None) for k in config_keys} diff --git a/src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py b/src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py index a60166f5e..73402b86b 100644 --- a/src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py +++ b/src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py @@ -341,8 +341,6 @@ def _cfg(name, default): value = getattr(config, name, None) return default if value is None else value - high_density = float(_cfg("svg_high_noise_density", -1.0)) - low_density = float(_cfg("svg_low_noise_density", -1.0)) expert_density = float(_cfg("svg_spatial_density", 0.25)) use_svg = bool(getattr(config, "use_svg_attention", False)) and (expert_density < 1.0) @@ -365,12 +363,10 @@ def _cfg(name, default): "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", 28)), + "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_high_noise_density": high_density, - "svg_low_noise_density": low_density, "svg_flash_block_sizes": getattr(config, "svg_flash_block_sizes", None) or None, } diff --git a/src/maxdiffusion/tests/ltx2/test_svg_config_propagation_ltx2.py b/src/maxdiffusion/tests/ltx2/test_svg_config_propagation_ltx2.py index 0a7d668df..5b33d2ce5 100644 --- a/src/maxdiffusion/tests/ltx2/test_svg_config_propagation_ltx2.py +++ b/src/maxdiffusion/tests/ltx2/test_svg_config_propagation_ltx2.py @@ -18,13 +18,16 @@ from types import SimpleNamespace import unittest +from unittest import mock from unittest.mock import MagicMock from flax import nnx import jax +import jax.numpy as jnp import numpy as np from jax.sharding import Mesh +from maxdiffusion.pipelines.ltx2 import ltx2_pipeline from maxdiffusion.pipelines.ltx2.ltx2_pipeline import create_sharded_logical_transformer, LTX2Pipeline @@ -192,6 +195,76 @@ def test_svg_disabled_when_use_svg_attention_is_false(self): ) self.assertFalse(m_dense.transformer_blocks[0].attn1.use_svg_attention) + def test_svg_disabled_when_spatial_density_is_one(self): + test_ltx2_config = self._get_test_ltx2_config() + devices = np.array(jax.devices()[:1]).reshape((1, 1)) + mesh = Mesh(devices, ("data", "fsdp")) + + cfg_full_density = SimpleNamespace( + use_svg_attention=True, + svg_spatial_density=1.0, + precision="DEFAULT", + flash_block_sizes={}, + activations_dtype="bfloat16", + weights_dtype="bfloat16", + attention="dot_product", + a2v_attention_kernel="dot_product", + v2a_attention_kernel="dot_product", + remat_policy="none", + names_which_can_be_saved=[], + names_which_can_be_offloaded=[], + flash_min_seq_length=0, + dropout=0.0, + scan_layers=False, + enable_jax_named_scopes=False, + use_base2_exp=False, + use_experimental_scheduler=False, + logical_axis_rules=(), + ) + m = create_sharded_logical_transformer( + devices_array=devices, + mesh=mesh, + rngs=nnx.Rngs(0), + config=cfg_full_density, + restored_checkpoint={"ltx2_config": dict(test_ltx2_config), "ltx2_state": {}}, + subfolder="", + ) + # A density of 1.0 is dense attention, so SVG must fall back to the dense path. + self.assertFalse(m.transformer_blocks[0].attn1.use_svg_attention) + + def test_transformer_forward_pass_forwards_svg_step_index(self): + captured = {} + + def fake_transformer(**kwargs): + captured.update(kwargs) + return kwargs["hidden_states"], kwargs["audio_hidden_states"] + + latents = jnp.zeros((1, 4, 8), dtype=jnp.float32) + audio_latents = jnp.zeros((1, 2, 8), dtype=jnp.float32) + with mock.patch.object(ltx2_pipeline.nnx, "merge", return_value=fake_transformer): + # Call the unjitted function so a cached trace cannot bypass the fake transformer. + ltx2_pipeline.transformer_forward_pass.fn( + None, + {}, + latents, + audio_latents, + jnp.asarray(0.5, dtype=jnp.float32), + jnp.zeros((1, 3, 8), dtype=jnp.float32), + jnp.zeros((1, 3, 8), dtype=jnp.float32), + None, + None, + latent_num_frames=1, + latent_height=2, + latent_width=2, + audio_num_frames=2, + fps=24, + global_batch_size=1, + svg_step_index=7, + ) + + self.assertIn("svg_step_index", captured) + self.assertEqual(captured["svg_step_index"], 7) + def test_cfg_and_magcache_incompatibility_validation(self): pipeline = LTX2Pipeline.__new__(LTX2Pipeline) pipeline.check_inputs = MagicMock()