diff --git a/src/maxdiffusion/models/attention_flax.py b/src/maxdiffusion/models/attention_flax.py index 3db542586..cb54b6a0f 100644 --- a/src/maxdiffusion/models/attention_flax.py +++ b/src/maxdiffusion/models/attention_flax.py @@ -38,6 +38,7 @@ from ..kernels import custom_svg_attention_dispatch from ..kernels import custom_svg_static_range_attention from . import quantizations +from . import svg_attention from .modeling_flax_utils import get_activation LOG2E = math.log2(math.e) @@ -2047,7 +2048,7 @@ def _apply_attention( def _head_local_svg_attention(query, key, value, context): - from .wan.transformers import svg_attention, svg_head_local + from .wan.transformers import svg_head_local cfg = context["spatiotemporal_config"] grid = context["spatiotemporal_shape"] @@ -2585,49 +2586,12 @@ def __init__( "use_experimental_scheduler": False, "ulysses_shards": -1, "ulysses_attention_chunks": 1, - "use_svg_attention": False, - "svg_implementation": "official_svg", - "svg_spatial_density": 0.25, - "svg_sample_max_row": 10000, - "svg_profile_query_count": 64, - "svg_profile_seed": 0, - "svg_dense_layer_fraction": 0.0, - "svg_dense_timestep_fraction": 0.0, - "svg_active_start_step": -1, - "svg_active_end_step": -1, - "svg_active_start_layer": -1, - "svg_active_end_layer": -1, - "svg_num_train_timesteps": 1000, - "svg_num_layers": 40, - "svg_include_first_frame": True, - "svg_global_stride": 0, - "svg_global_offset": 0, - "svg_high_noise_density": -1.0, - "svg_low_noise_density": -1.0, - "svg_flash_block_sizes": None, **(attention_config or {}), } - self.use_svg_attention = attention_config["use_svg_attention"] - self.svg_implementation = attention_config["svg_implementation"] - self.svg_spatial_density = attention_config["svg_spatial_density"] - self.svg_sample_max_row = attention_config["svg_sample_max_row"] - self.svg_profile_query_count = attention_config["svg_profile_query_count"] - self.svg_profile_seed = attention_config["svg_profile_seed"] - self.svg_dense_layer_fraction = attention_config["svg_dense_layer_fraction"] - self.svg_dense_timestep_fraction = attention_config["svg_dense_timestep_fraction"] - self.svg_active_start_step = attention_config["svg_active_start_step"] - self.svg_active_end_step = attention_config["svg_active_end_step"] - self.svg_active_start_layer = attention_config["svg_active_start_layer"] - self.svg_active_end_layer = attention_config["svg_active_end_layer"] - self.svg_num_train_timesteps = attention_config["svg_num_train_timesteps"] - self.svg_num_layers = attention_config["svg_num_layers"] - self.svg_include_first_frame = attention_config["svg_include_first_frame"] - self.svg_global_stride = attention_config["svg_global_stride"] - self.svg_global_offset = attention_config["svg_global_offset"] - self.svg_high_noise_density = attention_config["svg_high_noise_density"] - self.svg_low_noise_density = attention_config["svg_low_noise_density"] - self.svg_flash_block_sizes = attention_config["svg_flash_block_sizes"] + svg_config = svg_attention.init_svg_config(attention_config, default_num_layers=40) + for name, value in svg_config.items(): + setattr(self, name, value) self.is_self_attention = is_self_attention if attention_kernel in {"flash", "cudnn_flash_te"} and mesh is None: @@ -2934,73 +2898,27 @@ def __call__( key_proj = checkpoint_name(key_proj, "key_proj") value_proj = checkpoint_name(value_proj, "value_proj") - if self.use_svg_attention and is_self_attention and spatiotemporal_shape is not None: - from .wan.transformers import svg_attention - - is_active = svg_attention.is_svg_active( - step_index=svg_step_index, - layer_index=svg_layer_index, - timestep=svg_timestep, - start_step=self.svg_active_start_step, - end_step=self.svg_active_end_step, - start_layer=self.svg_active_start_layer, - end_layer=self.svg_active_end_layer, - dense_layer_fraction=self.svg_dense_layer_fraction, - dense_timestep_fraction=self.svg_dense_timestep_fraction, - num_train_timesteps=self.svg_num_train_timesteps, - num_layers=self.svg_num_layers, + def run_attention(**svg_kwargs): + return self.attention_op.apply_attention( + query_proj, + key_proj, + value_proj, + attention_mask=encoder_attention_mask, + **svg_kwargs, ) - def run_dense(_): - return self.attention_op.apply_attention( - query_proj, - key_proj, - value_proj, - attention_mask=encoder_attention_mask, - ) - - def run_sparse_svg(_): - execution_band_width = svg_attention.svg_execution_band_width( - spatiotemporal_shape, - self.svg_spatial_density, - ) - sparse_config = { - "use_svg_attention": True, - "mask_type": "svg_spatial", - "band_width": execution_band_width, - "include_first_frame": self.svg_include_first_frame, - "global_stride": self.svg_global_stride, - "global_offset": self.svg_global_offset, - "profile_query_count": self.svg_profile_query_count, - "profile_seed": self.svg_profile_seed, - "sample_max_row": self.svg_sample_max_row, - "custom_flash_block_sizes": self.svg_flash_block_sizes, - "svg_step_index": svg_step_index, - "svg_layer_index": svg_layer_index, - "svg_timestep": svg_timestep, - } - return self.attention_op.apply_attention( - query_proj, - key_proj, - value_proj, - attention_mask=encoder_attention_mask, - spatiotemporal_shape=spatiotemporal_shape, - sparse_config_override=sparse_config, - ) - - with jax.named_scope("apply_attention"): - if isinstance(is_active, bool): - attn_output = run_sparse_svg(None) if is_active else run_dense(None) - else: - attn_output = jax.lax.cond(is_active, run_sparse_svg, run_dense, operand=None) + if self.use_svg_attention and is_self_attention: + attn_output = svg_attention.apply_svg_or_dense( + self, + run_attention, + spatiotemporal_shape, + svg_layer_index=svg_layer_index, + svg_step_index=svg_step_index, + svg_timestep=svg_timestep, + ) else: with jax.named_scope("apply_attention"): - attn_output = self.attention_op.apply_attention( - query_proj, - key_proj, - value_proj, - attention_mask=encoder_attention_mask, - ) + attn_output = run_attention() else: # NEW PATH for I2V CROSS-ATTENTION diff --git a/src/maxdiffusion/models/ltx2/attention_ltx2.py b/src/maxdiffusion/models/ltx2/attention_ltx2.py index 3e55689eb..86d655f09 100644 --- a/src/maxdiffusion/models/ltx2/attention_ltx2.py +++ b/src/maxdiffusion/models/ltx2/attention_ltx2.py @@ -22,6 +22,7 @@ import jax.numpy as jnp from ... import common_types from ..attention_flax import NNXAttentionOp +from .. import svg_attention from .logical_sharding_ltx2 import get_sharding_specs, LTX2DiTShardingSpecs Array = common_types.Array @@ -352,7 +353,14 @@ def __init__( use_base2_exp: bool = False, use_experimental_scheduler: bool = False, enable_jax_named_scopes: bool = False, + attention_config: Optional[dict] = None, ): + svg_config = svg_attention.init_svg_config(attention_config, default_num_layers=48) + for name, value in svg_config.items(): + setattr(self, name, value) + self.is_self_attention = context_dim is None + self.use_svg_attention = bool(self.use_svg_attention) and self.is_self_attention + self.heads = heads self.rope_type = rope_type self.dim_head = dim_head @@ -542,10 +550,18 @@ def __call__( k_rotary_emb: Optional[Tuple[Array, Array]] = None, perturbation_mask: Optional[Array] = None, cached_kv: Optional[Tuple[Array, Array]] = None, + spatiotemporal_shape: Optional[Tuple[int, int, int]] = None, + svg_layer_index: Optional[int | jax.Array] = None, + svg_timestep: Optional[int | float | jax.Array] = None, + svg_step_index: Optional[int | jax.Array] = None, ) -> Array: # Determine context (Self or Cross) + is_self_attention = encoder_hidden_states is None context = encoder_hidden_states if encoder_hidden_states is not None else hidden_states + if self.use_svg_attention and is_self_attention and spatiotemporal_shape is None: + raise ValueError("SVG attention requires spatiotemporal_shape.") + # 1. Project and Norm with self.named_scope("QKV Projection"): query = self.to_q(hidden_states) @@ -586,8 +602,28 @@ def __call__( with self.named_scope("Attention and Output Project"): # 4. Attention - # NNXAttentionOp expects flattened input [B, S, InnerDim] for flash kernel - attn_output = self.attention_op.apply_attention(query=query, key=key, value=value, attention_mask=attention_mask) + def run_attention(**svg_kwargs): + return self.attention_op.apply_attention( + query=query, + key=key, + value=value, + attention_mask=attention_mask, + **svg_kwargs, + ) + + if self.use_svg_attention and is_self_attention: + attn_output = svg_attention.apply_svg_or_dense( + self, + run_attention, + spatiotemporal_shape, + svg_layer_index=svg_layer_index, + svg_step_index=svg_step_index, + svg_timestep=svg_timestep, + named_scope=self.named_scope, + ) + else: + with self.named_scope("apply_attention"): + attn_output = run_attention() if perturbation_mask is not None: # value is [B, S, InnerDim] diff --git a/src/maxdiffusion/models/ltx2/transformer_ltx2.py b/src/maxdiffusion/models/ltx2/transformer_ltx2.py index 1ac3c10c4..115476fcc 100644 --- a/src/maxdiffusion/models/ltx2/transformer_ltx2.py +++ b/src/maxdiffusion/models/ltx2/transformer_ltx2.py @@ -54,6 +54,9 @@ class LTX2StaticContext: audio_encoder_attention_mask: Optional[jax.Array] = None a2v_cross_attention_mask: Optional[jax.Array] = None v2a_cross_attention_mask: Optional[jax.Array] = None + spatiotemporal_shape: Optional[Tuple[int, int, int]] = struct.field(pytree_node=False, default=None) + svg_timestep: Optional[int | float | jax.Array] = None + svg_step_index: Optional[int | jax.Array] = None @struct.dataclass @@ -63,6 +66,7 @@ class LTX2BlockContext: static: LTX2StaticContext perturbation_mask: Optional[jax.Array] = None layer_kv_cache: Optional[Mapping[str, FrozenDict]] = None + layer_index: Optional[int | jax.Array] = None def _canonicalize_attention_mask(mask: Optional[jax.Array], batch_size: int, name: str) -> Optional[jax.Array]: @@ -186,6 +190,7 @@ def __init__( use_base2_exp: bool = False, use_experimental_scheduler: bool = False, enable_jax_named_scopes: bool = False, + attention_config: Optional[dict] = None, ): self.dim = dim self.norm_eps = norm_eps @@ -230,6 +235,7 @@ def __init__( ulysses_attention_chunks=ulysses_attention_chunks, use_base2_exp=use_base2_exp, use_experimental_scheduler=use_experimental_scheduler, + attention_config=attention_config, ) self.audio_norm1 = nnx.RMSNorm( @@ -263,6 +269,7 @@ def __init__( ulysses_attention_chunks=ulysses_attention_chunks, use_base2_exp=use_base2_exp, use_experimental_scheduler=use_experimental_scheduler, + attention_config={**(attention_config or {}), "use_svg_attention": False}, ) # 2. Prompt Cross-Attention @@ -595,6 +602,10 @@ def __call__( encoder_hidden_states=None, rotary_emb=video_rotary_emb, perturbation_mask=perturbation_mask, + spatiotemporal_shape=ctx.static.spatiotemporal_shape, + svg_layer_index=ctx.layer_index, + svg_timestep=ctx.static.svg_timestep, + svg_step_index=ctx.static.svg_step_index, ) hidden_states = hidden_states + attn_hidden_states * gate_msa @@ -831,6 +842,7 @@ def __init__( use_base2_exp: bool = False, use_experimental_scheduler: bool = False, enable_jax_named_scopes: bool = False, + attention_config: Optional[dict] = None, **kwargs, ): self.spatio_temporal_guidance_blocks = spatio_temporal_guidance_blocks @@ -890,6 +902,10 @@ def __init__( self.flash_min_seq_length = flash_min_seq_length self.use_base2_exp = use_base2_exp self.use_experimental_scheduler = use_experimental_scheduler + self.attention_config = { + "svg_num_layers": self.num_layers, + **(attention_config or {}), + } if sharding_specs is None: sharding_specs = get_sharding_specs("default", "ltx2_dit") @@ -1145,6 +1161,7 @@ def init_block(rngs): use_base2_exp=self.use_base2_exp, use_experimental_scheduler=self.use_experimental_scheduler, enable_jax_named_scopes=self.enable_jax_named_scopes, + attention_config=self.attention_config, ) if self.scan_layers: @@ -1189,6 +1206,7 @@ def init_block(rngs): use_base2_exp=self.use_base2_exp, use_experimental_scheduler=self.use_experimental_scheduler, enable_jax_named_scopes=self.enable_jax_named_scopes, + attention_config=self.attention_config, ) blocks.append(block) self.transformer_blocks = nnx.List(blocks) @@ -1399,6 +1417,7 @@ def __call__( cached_kv: Optional[Dict[str, Tuple[jax.Array, jax.Array]]] = None, rope_cache: Optional[Dict[str, Tuple[jax.Array, jax.Array]]] = None, time_embed_cache: Optional[Dict[str, jax.Array]] = None, + svg_step_index: Optional[int | jax.Array] = None, ) -> Any: """ Forward pass for the full LTX2 Video/Audio Diffusion Transformer. @@ -1600,6 +1619,11 @@ def __call__( audio_encoder_hidden_states = audio_encoder_hidden_states.reshape(batch_size, -1, audio_hidden_states.shape[-1]) # 5. Run transformer blocks with self.named_scope("Transformer Blocks"): + if num_frames is not None and height is not None and width is not None: + spatiotemporal_shape = (num_frames // self.patch_size_t, height // self.patch_size, width // self.patch_size) + else: + spatiotemporal_shape = None + static_context = LTX2StaticContext( encoder_hidden_states=encoder_hidden_states, audio_encoder_hidden_states=audio_encoder_hidden_states, @@ -1620,6 +1644,9 @@ def __call__( a2v_cross_attention_mask=a2v_cross_attention_mask, v2a_cross_attention_mask=v2a_cross_attention_mask, modality_mask=modality_mask, + spatiotemporal_shape=spatiotemporal_shape, + svg_timestep=timestep, + svg_step_index=svg_step_index, ) if cached_kv is not None: @@ -1630,13 +1657,14 @@ def __call__( else: unstacked_kv_layers = [None] * self.num_layers - def apply_block_in_scan(block, hidden_states, audio_hidden_states, mask, layer_kv_cache): + def apply_block_in_scan(block, hidden_states, audio_hidden_states, mask, layer_kv_cache, layer_index=None): context = LTX2BlockContext( hidden_states=hidden_states, audio_hidden_states=audio_hidden_states, static=static_context, perturbation_mask=mask, layer_kv_cache=layer_kv_cache, + layer_index=layer_index, ) with self.named_scope("Transformer Layer"): hidden_states_out, audio_hidden_states_out = block(context) @@ -1645,24 +1673,28 @@ def apply_block_in_scan(block, hidden_states, audio_hidden_states, mask, layer_k audio_hidden_states_out.astype(audio_hidden_states.dtype), ) + layer_indices = jnp.arange(self.num_layers, dtype=jnp.int32) if perturbation_mask is None: if cached_kv is not None: def scan_fn_ltx2(carry, block_and_kv): - block, layer_kv_cache = block_and_kv + block, layer_kv_cache, layer_index = block_and_kv if isinstance(layer_kv_cache, dict): layer_kv_cache = FrozenDict(layer_kv_cache) hidden_states, audio_hidden_states, rngs_carry = carry hidden_states, audio_hidden_states = apply_block_in_scan( - block, hidden_states, audio_hidden_states, None, layer_kv_cache + block, hidden_states, audio_hidden_states, None, layer_kv_cache, layer_index=layer_index ) return (hidden_states, audio_hidden_states, rngs_carry), None else: - def scan_fn_ltx2(carry, block): + def scan_fn_ltx2(carry, block_and_idx): + block, layer_index = block_and_idx hidden_states, audio_hidden_states, rngs_carry = carry - hidden_states, audio_hidden_states = apply_block_in_scan(block, hidden_states, audio_hidden_states, None, None) + hidden_states, audio_hidden_states = apply_block_in_scan( + block, hidden_states, audio_hidden_states, None, None, layer_index=layer_index + ) return (hidden_states, audio_hidden_states, rngs_carry), None if self.scan_layers: @@ -1674,7 +1706,11 @@ def scan_fn_ltx2(carry, block): ) carry = (hidden_states, audio_hidden_states, nnx.Rngs(0)) - scan_input = (self.transformer_blocks, cached_kv) if cached_kv is not None else self.transformer_blocks + scan_input = ( + (self.transformer_blocks, cached_kv, layer_indices) + if cached_kv is not None + else (self.transformer_blocks, layer_indices) + ) (hidden_states, audio_hidden_states, _), _ = nnx.scan( rematted_scan_fn, length=self.num_layers, @@ -1685,7 +1721,7 @@ def scan_fn_ltx2(carry, block): else: for i, block in enumerate(self.transformer_blocks): hidden_states, audio_hidden_states = apply_block_in_scan( - block, hidden_states, audio_hidden_states, None, unstacked_kv_layers[i] + block, hidden_states, audio_hidden_states, None, unstacked_kv_layers[i], layer_index=i ) else: masks = jnp.ones((self.num_layers, batch_size, 1, 1), dtype=self.dtype) @@ -1697,21 +1733,23 @@ def scan_fn_ltx2(carry, block): if cached_kv is not None: def scan_fn_ltx23(carry, block_and_mask_and_kv): - block, mask, layer_kv_cache = block_and_mask_and_kv + block, mask, layer_kv_cache, layer_index = block_and_mask_and_kv if isinstance(layer_kv_cache, dict): layer_kv_cache = FrozenDict(layer_kv_cache) hidden_states, audio_hidden_states, rngs_carry = carry hidden_states, audio_hidden_states = apply_block_in_scan( - block, hidden_states, audio_hidden_states, mask, layer_kv_cache + block, hidden_states, audio_hidden_states, mask, layer_kv_cache, layer_index=layer_index ) return (hidden_states, audio_hidden_states, rngs_carry), None else: def scan_fn_ltx23(carry, block_and_mask): - block, mask = block_and_mask + block, mask, layer_index = block_and_mask hidden_states, audio_hidden_states, rngs_carry = carry - hidden_states, audio_hidden_states = apply_block_in_scan(block, hidden_states, audio_hidden_states, mask, None) + hidden_states, audio_hidden_states = apply_block_in_scan( + block, hidden_states, audio_hidden_states, mask, None, layer_index=layer_index + ) return (hidden_states, audio_hidden_states, rngs_carry), None if self.scan_layers: @@ -1723,9 +1761,9 @@ def scan_fn_ltx23(carry, block_and_mask): ) carry = (hidden_states, audio_hidden_states, nnx.Rngs(0)) scan_input = ( - (self.transformer_blocks, perturbation_mask_per_layer, cached_kv) + (self.transformer_blocks, perturbation_mask_per_layer, cached_kv, layer_indices) if cached_kv is not None - else (self.transformer_blocks, perturbation_mask_per_layer) + else (self.transformer_blocks, perturbation_mask_per_layer, layer_indices) ) (hidden_states, audio_hidden_states, _), _ = nnx.scan( rematted_scan_fn, @@ -1737,7 +1775,12 @@ def scan_fn_ltx23(carry, block_and_mask): else: for i, block in enumerate(self.transformer_blocks): hidden_states, audio_hidden_states = apply_block_in_scan( - block, hidden_states, audio_hidden_states, perturbation_mask_per_layer[i], unstacked_kv_layers[i] + block, + hidden_states, + audio_hidden_states, + perturbation_mask_per_layer[i], + unstacked_kv_layers[i], + layer_index=i, ) # 6. Output layers diff --git a/src/maxdiffusion/models/wan/transformers/svg_attention.py b/src/maxdiffusion/models/svg_attention.py similarity index 73% rename from src/maxdiffusion/models/wan/transformers/svg_attention.py rename to src/maxdiffusion/models/svg_attention.py index 7ffd1e11a..74f8fce41 100644 --- a/src/maxdiffusion/models/wan/transformers/svg_attention.py +++ b/src/maxdiffusion/models/svg_attention.py @@ -20,11 +20,121 @@ import math from numbers import Integral -from typing import Tuple +from typing import Any, Callable, Mapping, Optional, Tuple import jax import jax.numpy as jnp +# Attention-level SVG keys and their defaults. `svg_num_layers` is model +# specific and supplied by the caller. Pipeline-level keys such as +# `svg_high_noise_density` / `svg_low_noise_density` (Wan expert selection) +# are resolved into `svg_spatial_density` before reaching attention. +SVG_ATTENTION_DEFAULTS = { + "use_svg_attention": False, + "svg_spatial_density": 0.25, + "svg_sample_max_row": 10000, + "svg_profile_query_count": 64, + "svg_profile_seed": 0, + "svg_dense_layer_fraction": 0.0, + "svg_dense_timestep_fraction": 0.0, + "svg_active_start_step": -1, + "svg_active_end_step": -1, + "svg_active_start_layer": -1, + "svg_active_end_layer": -1, + "svg_num_train_timesteps": 1000, + "svg_include_first_frame": True, + "svg_global_stride": 0, + "svg_flash_block_sizes": None, +} + + +def init_svg_config(attention_config: Optional[Mapping[str, Any]], default_num_layers: int) -> dict[str, Any]: + """Resolves the attention-level SVG settings from `attention_config`. + + Returns a dict keyed by the names in `SVG_ATTENTION_DEFAULTS` plus + `svg_num_layers`; other keys are ignored. Settings that configs accept but + the head-local SVG kernel does not implement fail loudly instead of being + silently dropped. + """ + attention_config = attention_config or {} + implementation = attention_config.get("svg_implementation", "official_svg") + if implementation != "official_svg": + raise ValueError(f"Unsupported svg_implementation={implementation!r}; only 'official_svg' is implemented.") + global_offset = attention_config.get("svg_global_offset", 0) + if global_offset: + raise ValueError(f"svg_global_offset={global_offset} is not supported by head-local SVG.") + resolved = {**SVG_ATTENTION_DEFAULTS, "svg_num_layers": default_num_layers} + for name in resolved: + if name in attention_config: + resolved[name] = attention_config[name] + return resolved + + +def apply_svg_or_dense( + svg: Any, + run_attention: Callable[..., jax.Array], + spatiotemporal_shape: Tuple[int, int, int], + svg_layer_index: int | jax.Array | None, + svg_step_index: int | jax.Array | None, + svg_timestep: int | float | jax.Array | None, + named_scope: Callable[[str], Any] = jax.named_scope, +) -> jax.Array: + """Runs SVG sparse attention on active steps/layers and dense attention otherwise. + + Args: + svg: Object exposing the attributes resolved by `init_svg_config` + (the attention module itself). + run_attention: Calls the attention op. Invoked with no arguments for the + dense path, and with `spatiotemporal_shape` and `sparse_config_override` + keyword arguments for the sparse path. + spatiotemporal_shape: Latent token grid `(frames, height, width)`. + svg_layer_index: Transformer layer index (static int or traced array). + svg_step_index: Denoising step index (static int or traced array). + svg_timestep: Diffusion timestep, used by fraction-based schedules. + named_scope: Context manager factory used to label the dispatch. + """ + is_active = is_svg_active( + step_index=svg_step_index, + layer_index=svg_layer_index, + timestep=svg_timestep, + start_step=svg.svg_active_start_step, + end_step=svg.svg_active_end_step, + start_layer=svg.svg_active_start_layer, + end_layer=svg.svg_active_end_layer, + dense_layer_fraction=svg.svg_dense_layer_fraction, + dense_timestep_fraction=svg.svg_dense_timestep_fraction, + num_train_timesteps=svg.svg_num_train_timesteps, + num_layers=svg.svg_num_layers, + ) + + def run_dense(_): + return run_attention() + + def run_sparse(_): + sparse_config = { + "use_svg_attention": True, + "mask_type": "svg_spatial", + "band_width": svg_execution_band_width(spatiotemporal_shape, svg.svg_spatial_density), + "include_first_frame": svg.svg_include_first_frame, + "global_stride": svg.svg_global_stride, + "profile_query_count": svg.svg_profile_query_count, + "profile_seed": svg.svg_profile_seed, + "sample_max_row": svg.svg_sample_max_row, + "custom_flash_block_sizes": svg.svg_flash_block_sizes, + "svg_step_index": svg_step_index, + "svg_layer_index": svg_layer_index, + "svg_timestep": svg_timestep, + } + return run_attention( + spatiotemporal_shape=spatiotemporal_shape, + sparse_config_override=sparse_config, + ) + + with named_scope("apply_attention"): + if isinstance(is_active, bool): + return run_sparse(None) if is_active else run_dense(None) + return jax.lax.cond(is_active, run_sparse, run_dense, operand=None) + def svg_execution_band_width(token_grid: Tuple[int, int, int], density: float) -> int: """Compute SVG symmetric band width from retained density with 128-token block ceiling. diff --git a/src/maxdiffusion/tests/ltx2/test_svg_attention_ltx2.py b/src/maxdiffusion/tests/ltx2/test_svg_attention_ltx2.py new file mode 100644 index 000000000..9aeccdc04 --- /dev/null +++ b/src/maxdiffusion/tests/ltx2/test_svg_attention_ltx2.py @@ -0,0 +1,502 @@ +""" +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. + +Tests for Sparse VideoGen (SVG) attention in LTX2. +""" + +import unittest +from unittest.mock import patch, MagicMock + +from flax import nnx +import jax +import jax.numpy as jnp +import numpy as np +from jax.sharding import Mesh +from flax.linen import partitioning as nn_partitioning + +from maxdiffusion.models.ltx2.attention_ltx2 import LTX2Attention +from maxdiffusion.models.ltx2.transformer_ltx2 import ( + LTX2VideoTransformerBlock, + LTX2VideoTransformer3DModel, + LTX2StaticContext, + LTX2BlockContext, +) +from maxdiffusion.models.svg_attention import is_svg_active + + +class LTX2SVGAttentionTest(unittest.TestCase): + + def setUp(self): + devices = np.array(jax.devices()[:1]).reshape((1, 1)) + self.mesh = Mesh(devices, ("data", "fsdp")) + self.rngs = nnx.Rngs(0) + self.logical_axis_rules = ( + ("activation_batch", ("data", "fsdp")), + ("activation_length", None), + ("activation_embed", None), + ) + + def test_is_svg_active_ltx2_boundaries(self): + """Verifies that SVG active predicate correctly checks step, layer, and density.""" + # Active range: step in [10, 30), layer in [1, 28) + # Outside active steps: step 5, 30, 35 -> False + self.assertFalse( + is_svg_active( + layer_index=5, + step_index=5, + timestep=500.0, + start_step=10, + end_step=30, + start_layer=1, + end_layer=28, + ) + ) + self.assertFalse( + is_svg_active( + layer_index=5, + step_index=30, + timestep=200.0, + start_step=10, + end_step=30, + start_layer=1, + end_layer=28, + ) + ) + + # Outside active layers: layer 0, layer 28 -> False + self.assertFalse( + is_svg_active( + layer_index=0, + step_index=15, + timestep=500.0, + start_step=10, + end_step=30, + start_layer=1, + end_layer=28, + ) + ) + self.assertFalse( + is_svg_active( + layer_index=28, + step_index=15, + timestep=500.0, + start_step=10, + end_step=30, + start_layer=1, + end_layer=28, + ) + ) + + # Inside active steps & layers: step 15, layer 5 -> True + self.assertTrue( + is_svg_active( + layer_index=5, + step_index=15, + timestep=500.0, + start_step=10, + end_step=30, + start_layer=1, + end_layer=28, + ) + ) + + # Dynamic execution under JIT with JAX Array inputs + @jax.jit + def check_dynamic(step, layer): + return is_svg_active( + layer_index=layer, + step_index=step, + timestep=500.0, + start_step=10, + end_step=30, + start_layer=1, + end_layer=28, + ) + + self.assertTrue(bool(check_dynamic(jnp.int32(15), jnp.int32(5)))) + self.assertFalse(bool(check_dynamic(jnp.int32(5), jnp.int32(5)))) + self.assertFalse(bool(check_dynamic(jnp.int32(15), jnp.int32(0)))) + + def test_ltx2_attention_forward_dense_and_sparse_dispatch(self): + """Tests LTX2Attention routing under inactive and active SVG steps.""" + B = 1 + F, H, W = 4, 8, 8 + seq_len = F * H * W + dim = 128 + num_heads = 4 + head_dim = dim // num_heads + + attention_config = { + "use_svg_attention": True, + "svg_spatial_density": 0.25, + "svg_active_start_step": 2, + "svg_active_end_step": 10, + "svg_active_start_layer": 0, + "svg_active_end_layer": 10, + "svg_sample_max_row": 100, + "svg_profile_query_count": 16, + } + + with self.mesh, nn_partitioning.axis_rules(self.logical_axis_rules): + attn = LTX2Attention( + rngs=self.rngs, + query_dim=dim, + context_dim=None, + heads=num_heads, + dim_head=head_dim, + attention_kernel="dot_product", + mesh=self.mesh, + attention_config=attention_config, + ) + + hidden_states = jax.random.normal(jax.random.key(1), (B, seq_len, dim), dtype=jnp.float32) + + # 1. Inactive step (step 0 < active_start_step 2) -> runs dense path successfully + out_dense = attn( + hidden_states=hidden_states, + spatiotemporal_shape=(F, H, W), + svg_layer_index=0, + svg_timestep=jnp.array([100.0]), + svg_step_index=0, + ) + self.assertEqual(out_dense.shape, (B, seq_len, dim)) + self.assertTrue(jnp.all(jnp.isfinite(out_dense))) + + # 2. Active step (step 3 in [2, 10)) -> triggers SVG branch which requires custom Ulysses backend + with self.assertRaisesRegex(ValueError, "Head-local SVG requires a custom Ulysses attention backend"): + attn( + hidden_states=hidden_states, + spatiotemporal_shape=(F, H, W), + svg_layer_index=0, + svg_timestep=jnp.array([100.0]), + svg_step_index=3, + ) + + # 3. Verify mock dispatch receives SVG sparse_config_override and spatiotemporal_shape + mock_apply = MagicMock(return_value=jnp.zeros((B, seq_len, dim), dtype=jnp.float32)) + with patch.object(attn.attention_op, "apply_attention", mock_apply): + _ = attn( + hidden_states=hidden_states, + spatiotemporal_shape=(F, H, W), + svg_layer_index=0, + svg_timestep=jnp.array([100.0]), + svg_step_index=3, + ) + mock_apply.assert_called_once() + _, kwargs = mock_apply.call_args + self.assertEqual(kwargs.get("spatiotemporal_shape"), (F, H, W)) + sp_cfg = kwargs.get("sparse_config_override") + self.assertIsNotNone(sp_cfg) + self.assertTrue(sp_cfg.get("use_svg_attention")) + # 256 tokens at density 0.25 -> 256 * (1 - sqrt(0.75)) ~= 34.3, rounded up to one 128 block. + self.assertEqual(sp_cfg["band_width"], 128) + self.assertEqual(sp_cfg["svg_layer_index"], 0) + self.assertEqual(sp_cfg["svg_step_index"], 3) + self.assertIsNotNone(sp_cfg["svg_timestep"]) + + def test_ltx2_transformer_block_forward_svg_dispatch(self): + """Tests LTX2VideoTransformerBlock forward pass routing with SVG.""" + B = 1 + F, H, W = 4, 8, 8 + seq_len = F * H * W + audio_seq_len = 16 + dim = 64 + audio_dim = 64 + num_heads = 2 + head_dim = dim // num_heads + + attention_config = { + "use_svg_attention": True, + "svg_spatial_density": 0.25, + "svg_active_start_step": 2, + "svg_active_end_step": 10, + "svg_active_start_layer": 0, + "svg_active_end_layer": 10, + "svg_sample_max_row": 100, + "svg_profile_query_count": 16, + } + + with self.mesh, nn_partitioning.axis_rules(self.logical_axis_rules): + block = LTX2VideoTransformerBlock( + rngs=self.rngs, + dim=dim, + num_attention_heads=num_heads, + attention_head_dim=head_dim, + cross_attention_dim=dim, + audio_dim=audio_dim, + audio_num_attention_heads=num_heads, + audio_attention_head_dim=head_dim, + audio_cross_attention_dim=audio_dim, + attention_kernel="dot_product", + a2v_attention_kernel="dot_product", + v2a_attention_kernel="dot_product", + mesh=self.mesh, + attention_config=attention_config, + ) + + hidden_states = jax.random.normal(jax.random.key(1), (B, seq_len, dim), dtype=jnp.float32) + audio_hidden_states = jax.random.normal(jax.random.key(2), (B, audio_seq_len, audio_dim), dtype=jnp.float32) + encoder_hidden_states = jax.random.normal(jax.random.key(3), (B, 16, dim), dtype=jnp.float32) + audio_encoder_hidden_states = jax.random.normal(jax.random.key(4), (B, 16, audio_dim), dtype=jnp.float32) + + # Inactive step 0 -> dense execution succeeds + static_ctx_inactive = LTX2StaticContext( + encoder_hidden_states=encoder_hidden_states, + audio_encoder_hidden_states=audio_encoder_hidden_states, + temb=jnp.zeros((B, 6 * dim)), + temb_audio=jnp.zeros((B, 6 * audio_dim)), + temb_ca_scale_shift=jnp.zeros((B, 4 * dim)), + temb_ca_audio_scale_shift=jnp.zeros((B, 4 * audio_dim)), + temb_ca_gate=jnp.zeros((B, 1 * dim)), + temb_ca_audio_gate=jnp.zeros((B, 1 * audio_dim)), + spatiotemporal_shape=(F, H, W), + svg_timestep=jnp.array([100.0]), + svg_step_index=0, + ) + block_ctx_inactive = LTX2BlockContext( + hidden_states=hidden_states, + audio_hidden_states=audio_hidden_states, + static=static_ctx_inactive, + layer_index=0, + ) + out_h, out_a = block(block_ctx_inactive) + self.assertEqual(out_h.shape, (B, seq_len, dim)) + self.assertEqual(out_a.shape, (B, audio_seq_len, audio_dim)) + + # Active step 3 -> triggers SVG + static_ctx_active = LTX2StaticContext( + encoder_hidden_states=encoder_hidden_states, + audio_encoder_hidden_states=audio_encoder_hidden_states, + temb=jnp.zeros((B, 6 * dim)), + temb_audio=jnp.zeros((B, 6 * audio_dim)), + temb_ca_scale_shift=jnp.zeros((B, 4 * dim)), + temb_ca_audio_scale_shift=jnp.zeros((B, 4 * audio_dim)), + temb_ca_gate=jnp.zeros((B, 1 * dim)), + temb_ca_audio_gate=jnp.zeros((B, 1 * audio_dim)), + spatiotemporal_shape=(F, H, W), + svg_timestep=jnp.array([100.0]), + svg_step_index=3, + ) + block_ctx_active = LTX2BlockContext( + hidden_states=hidden_states, + audio_hidden_states=audio_hidden_states, + static=static_ctx_active, + layer_index=0, + ) + with self.assertRaisesRegex(ValueError, "Head-local SVG requires a custom Ulysses attention backend"): + block(block_ctx_active) + + # Video self-attention must take the SVG path; audio self-attention stays dense. + mock_video = MagicMock(side_effect=lambda query, **_: jnp.zeros_like(query)) + mock_audio = MagicMock(side_effect=lambda query, **_: jnp.zeros_like(query)) + with patch.object(block.attn1.attention_op, "apply_attention", mock_video), patch.object( + block.audio_attn1.attention_op, "apply_attention", mock_audio + ): + block(block_ctx_active) + mock_video.assert_called_once() + video_cfg = mock_video.call_args.kwargs.get("sparse_config_override") + self.assertIsNotNone(video_cfg) + self.assertTrue(video_cfg["use_svg_attention"]) + mock_audio.assert_called_once() + self.assertIsNone(mock_audio.call_args.kwargs.get("sparse_config_override")) + + def test_ltx2_attention_svg_full_density_matches_dense_on_tpu(self): + """At density 1.0 the SVG band covers every pair, so SVG must match dense attention.""" + devices = jax.devices() + if devices[0].platform != "tpu": + self.skipTest("Requires TPU for the ulysses_custom kernels.") + context = len(devices) + mesh = Mesh(np.array(devices).reshape((1, 1, context)), ("data", "fsdp", "context")) + axis_rules = ( + ("activation_batch", ("data", "fsdp")), + ("activation_length", "context"), + ("activation_heads", None), + ("activation_kv", None), + ("activation_self_attn_heads", None), + ("activation_self_attn_q_length", "context"), + ("activation_self_attn_kv_length", "context"), + ("activation_embed", None), + ) + grid = (4, 16, 16) + heads, head_dim = context, 128 + dim = heads * head_dim + block_sizes = { + "block_q": 256, + "block_kv": 256, + "block_kv_compute": 256, + "block_kv_compute_in": 256, + "heads_per_tile": 1, + } + + def build(attention_config): + return LTX2Attention( + rngs=nnx.Rngs(0), + query_dim=dim, + heads=heads, + dim_head=head_dim, + attention_kernel="ulysses_custom", + flash_block_sizes=block_sizes, + flash_min_seq_length=0, + mesh=mesh, + attention_config=attention_config, + ) + + with mesh, nn_partitioning.axis_rules(axis_rules): + attn_dense = build(None) + attn_svg = build({"use_svg_attention": True, "svg_spatial_density": 1.0}) + hidden_states = jax.random.normal(jax.random.key(0), (1, int(np.prod(grid)), dim), dtype=jnp.float32) + kwargs = { + "spatiotemporal_shape": grid, + "svg_layer_index": 0, + "svg_step_index": 0, + } + out_dense = attn_dense(hidden_states, **kwargs) + out_svg = attn_svg(hidden_states, **kwargs) + self.assertTrue(jnp.allclose(out_svg, out_dense, atol=1e-2), float(jnp.max(jnp.abs(out_svg - out_dense)))) + + def test_ltx2_model_full_forward_with_svg(self): + """Tests LTX2VideoTransformer3DModel full forward pass with SVG configuration.""" + B = 1 + F, H, W = 2, 8, 8 + seq_len = F * H * W + audio_seq_len = 16 + in_channels = 8 + out_channels = 8 + audio_in_channels = 4 + num_heads = 2 + head_dim = 32 + + # Step range [5, 15) so step 0 is inactive and completes full dense pass, while step 6 triggers SVG + attention_config = { + "use_svg_attention": True, + "svg_spatial_density": 0.25, + "svg_active_start_step": 5, + "svg_active_end_step": 15, + "svg_active_start_layer": 0, + "svg_active_end_layer": 10, + "svg_sample_max_row": 100, + "svg_profile_query_count": 16, + } + + with self.mesh, nn_partitioning.axis_rules(self.logical_axis_rules): + # Non-scanned blocks: step 0 statically resolves inactive and executes dense + model_unscanned = LTX2VideoTransformer3DModel( + rngs=nnx.Rngs(0), + in_channels=in_channels, + out_channels=out_channels, + patch_size=1, + patch_size_t=1, + num_attention_heads=num_heads, + attention_head_dim=head_dim, + cross_attention_dim=num_heads * head_dim, + caption_channels=16, + audio_in_channels=audio_in_channels, + audio_out_channels=audio_in_channels, + audio_num_attention_heads=num_heads, + audio_attention_head_dim=head_dim, + audio_cross_attention_dim=num_heads * head_dim, + num_layers=2, + mesh=self.mesh, + attention_kernel="dot_product", + a2v_attention_kernel="dot_product", + v2a_attention_kernel="dot_product", + scan_layers=False, + attention_config=attention_config, + ) + + hidden_states = jax.random.normal(jax.random.key(10), (B, seq_len, in_channels), dtype=jnp.float32) + audio_hidden_states = jax.random.normal(jax.random.key(11), (B, audio_seq_len, audio_in_channels), dtype=jnp.float32) + timestep = jnp.array([1.0]) + encoder_hidden_states = jax.random.normal(jax.random.key(12), (B, 16, 16), dtype=jnp.float32) + audio_encoder_hidden_states = jax.random.normal(jax.random.key(13), (B, 16, 16), dtype=jnp.float32) + + # 1. Inactive step (svg_step_index=0) -> executes full dense pass successfully + output = model_unscanned( + hidden_states=hidden_states, + audio_hidden_states=audio_hidden_states, + encoder_hidden_states=encoder_hidden_states, + audio_encoder_hidden_states=audio_encoder_hidden_states, + timestep=timestep, + num_frames=F, + height=H, + width=W, + audio_num_frames=audio_seq_len, + svg_step_index=0, + return_dict=True, + ) + + self.assertEqual(output["sample"].shape, (B, seq_len, out_channels)) + self.assertEqual(output["audio_sample"].shape, (B, audio_seq_len, audio_in_channels)) + self.assertTrue(jnp.all(jnp.isfinite(output["sample"]))) + self.assertTrue(jnp.all(jnp.isfinite(output["audio_sample"]))) + + # 2. Active step (svg_step_index=6) -> triggers SVG routing and checks backend + with self.assertRaisesRegex(ValueError, "Head-local SVG requires a custom Ulysses attention backend"): + model_unscanned( + hidden_states=hidden_states, + audio_hidden_states=audio_hidden_states, + encoder_hidden_states=encoder_hidden_states, + audio_encoder_hidden_states=audio_encoder_hidden_states, + timestep=timestep, + num_frames=F, + height=H, + width=W, + audio_num_frames=audio_seq_len, + svg_step_index=6, + return_dict=True, + ) + + # 3. Scanned blocks: fail-closed validation when backend is not a custom Ulysses kernel + model_scanned = LTX2VideoTransformer3DModel( + rngs=nnx.Rngs(0), + in_channels=in_channels, + out_channels=out_channels, + patch_size=1, + patch_size_t=1, + num_attention_heads=num_heads, + attention_head_dim=head_dim, + cross_attention_dim=num_heads * head_dim, + caption_channels=16, + audio_in_channels=audio_in_channels, + audio_out_channels=audio_in_channels, + audio_num_attention_heads=num_heads, + audio_attention_head_dim=head_dim, + audio_cross_attention_dim=num_heads * head_dim, + num_layers=2, + mesh=self.mesh, + attention_kernel="dot_product", + a2v_attention_kernel="dot_product", + v2a_attention_kernel="dot_product", + scan_layers=True, + attention_config=attention_config, + ) + with self.assertRaisesRegex(ValueError, "Head-local SVG requires a custom Ulysses attention backend"): + model_scanned( + hidden_states=hidden_states, + audio_hidden_states=audio_hidden_states, + encoder_hidden_states=encoder_hidden_states, + audio_encoder_hidden_states=audio_encoder_hidden_states, + timestep=timestep, + num_frames=F, + height=H, + width=W, + audio_num_frames=audio_seq_len, + svg_step_index=6, + return_dict=True, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/src/maxdiffusion/tests/wan/svg_attention_test.py b/src/maxdiffusion/tests/wan/svg_attention_test.py index 1be65b03e..16047511f 100644 --- a/src/maxdiffusion/tests/wan/svg_attention_test.py +++ b/src/maxdiffusion/tests/wan/svg_attention_test.py @@ -23,7 +23,7 @@ import jax.numpy as jnp import numpy as np -from maxdiffusion.models.wan.transformers import svg_attention +from maxdiffusion.models import svg_attention from maxdiffusion.kernels import custom_svg_static_range_attention as static_kernel diff --git a/src/maxdiffusion/tests/wan/svg_head_local_test.py b/src/maxdiffusion/tests/wan/svg_head_local_test.py index 630879311..2c8487cd8 100644 --- a/src/maxdiffusion/tests/wan/svg_head_local_test.py +++ b/src/maxdiffusion/tests/wan/svg_head_local_test.py @@ -26,7 +26,7 @@ from jax.sharding import Mesh, NamedSharding, PartitionSpec as P from maxdiffusion.models.wan.transformers.svg_head_local import exchange_local, inference_only -from maxdiffusion.models.wan.transformers import svg_attention as svg +from maxdiffusion.models import svg_attention as svg @pytest.mark.parametrize("routing", ["mixed", "spatial", "temporal"])