-
Notifications
You must be signed in to change notification settings - Fork 95
[PR 2/2] LTX2: SVG config and pipeline #498
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: ltx2_svg_model
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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. | ||
|
|
||
|  | ||
|
|
||
| *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. | ||
|
|
||
| <img src="images/svg/attention-tiles.png" alt="A local attention band over a query–key tile grid, highlighting full, boundary, and skipped tiles." width="520"> | ||
|
|
||
| *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 | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. LTX-2 and LTX-2.3 have 48 transformer layers. Setting Was |
||
| 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 | ||
|
|
||
|  | ||
|
|
||
| *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. | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Small tip worth adding here: in |
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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 | ||
|
Comment on lines
+31
to
+32
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Two small config things here (and in
|
||
| # 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. | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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", | ||
|
Comment on lines
+226
to
+245
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. If In |
||
| ] | ||
| config_dict = {k: getattr(config, k, None) for k in config_keys} | ||
|
|
||
|
|
||
| Original file line number | Diff line number | Diff line change | ||||
|---|---|---|---|---|---|---|
|
|
@@ -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.<name>, 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)), | ||||||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. LTX-2 and LTX-2.3 have
Suggested change
|
||||||
| "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.") | ||||||
|
|
||||||
|
Comment on lines
+1873
to
+1876
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Up on line 348, we disable SVG when use_svg = bool(getattr(config, "use_svg_attention", False)) and (expert_density < 1.0)Here in |
||||||
| # 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) | ||||||
|
|
||||||
|
|
||||||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
In
ltx2_video.ymlandltx2_3_video.yml,ulysses_shardsdefaults to-1. However,pyconfig.pyraises aValueErrorifattention: ulysses_ring_custom_fixed_mis used without settingulysses_shards > 0.Could we either add
ulysses_shards: 2(or8) to this YAML snippet, or changeattentionhere toulysses_custom(like in the PR description's VABench command) so that copying this example works out of the box?