Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
165 changes: 165 additions & 0 deletions docs/svg.md
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.

![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.

<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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

In ltx2_video.yml and ltx2_3_video.yml, ulysses_shards defaults to -1. However, pyconfig.py raises a ValueError if attention: ulysses_ring_custom_fixed_m is used without setting ulysses_shards > 0.

Could we either add ulysses_shards: 2 (or 8) to this YAML snippet, or change attention here to ulysses_custom (like in the PR description's VABench command) so that copying this example works out of the box?

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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LTX-2 and LTX-2.3 have 48 transformer layers. Setting svg_active_end_layer: 28 means layers 28 through 47 will still run dense attention.

Was 28 intentional here, or should this be 48 (similar to how the Wan2.2 example above uses 40 to cover all layers after layer 0)?

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.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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?

20 changes: 20 additions & 0 deletions src/maxdiffusion/configs/ltx2_3_video.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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: []
Expand Down Expand Up @@ -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.
Expand Down
20 changes: 20 additions & 0 deletions src/maxdiffusion/configs/ltx2_video.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Two small config things here (and in ltx2_3_video.yml):

  1. pyconfig.py only allows passing a flag on the command line if that key already exists in the YAML file. In ltx2_pipeline.py (lines 356–371), we read several SVG keys that aren't in this YAML yet:

    • 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
    • svg_implementation: 'official_svg'

    Could you add those defaults here so users can override them from the CLI if needed?

  2. svg_high_noise_density and svg_low_noise_density are for Wan 2.2's two-expert setup. Since LTX-2 only has a single transformer and ltx2_pipeline.py only uses svg_spatial_density, setting these two flags on LTX-2 won't actually change model density. Should we remove them from the LTX-2 configs (and ltx2_pipeline.py) to avoid confusion?

# 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
Expand Down Expand Up @@ -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.
Expand Down
20 changes: 20 additions & 0 deletions src/maxdiffusion/generate_ltx2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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)?

]
config_dict = {k: getattr(config, k, None) for k in config_keys}

Expand Down
55 changes: 51 additions & 4 deletions src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)),

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LTX-2 and LTX-2.3 have 48 layers by default (matching num_layers: int = 48 in LTX2VideoTransformer3DModel and svg_num_layers: int = 48 in LTX2Attention). Let's change the fallback default here from 28 to 48:

Suggested change
"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,
}

# 2. eval_shape
p_model_factory = partial(create_model, ltx2_config=ltx2_config)
transformer = nnx.eval_shape(p_model_factory, rngs=rngs)
Expand Down Expand Up @@ -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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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?

# 2. Encode inputs (Text)
t0_encode = time.perf_counter()
(
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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)

Expand Down
Loading
Loading