From 9b309250f28c3cfa3793b989f2043403380fe5aa Mon Sep 17 00:00:00 2001 From: Jitendra Jalwaniya Date: Tue, 29 Sep 2026 06:53:39 +0000 Subject: [PATCH] ltx2: add Sparse VideoGen (SVG) attention to LTX2 attention and transformer --- .../models/ltx2/attention_ltx2.py | 126 ++++- .../models/ltx2/transformer_ltx2.py | 74 ++- .../tests/ltx2/test_svg_attention_ltx2.py | 429 ++++++++++++++++++ 3 files changed, 613 insertions(+), 16 deletions(-) create mode 100644 src/maxdiffusion/tests/ltx2/test_svg_attention_ltx2.py diff --git a/src/maxdiffusion/models/ltx2/attention_ltx2.py b/src/maxdiffusion/models/ltx2/attention_ltx2.py index 3e55689eb..65db95e5a 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 ..wan.transformers import svg_attention from .logical_sharding_ltx2 import get_sharding_specs, LTX2DiTShardingSpecs Array = common_types.Array @@ -352,7 +353,58 @@ def __init__( use_base2_exp: bool = False, use_experimental_scheduler: bool = False, enable_jax_named_scopes: bool = False, + attention_config: Optional[dict] = None, ): + attention_config = { + "use_base2_exp": use_base2_exp, + "use_experimental_scheduler": use_experimental_scheduler, + "ulysses_shards": ulysses_shards, + "ulysses_attention_chunks": ulysses_attention_chunks, + "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": 48, + "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.is_self_attention = context_dim is None + self.use_svg_attention = bool(attention_config["use_svg_attention"]) and self.is_self_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"] + self.heads = heads self.rope_type = rope_type self.dim_head = dim_head @@ -542,10 +594,22 @@ def __call__( k_rotary_emb: Optional[Tuple[Array, Array]] = None, perturbation_mask: Optional[Array] = None, cached_kv: Optional[Tuple[Array, Array]] = None, + deterministic: bool = True, + 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: + if not deterministic: + raise ValueError("SVG attention supports deterministic inference only.") + if 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 +650,66 @@ 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) + if self.use_svg_attention and is_self_attention and spatiotemporal_shape is not None: + 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_dense(_): + return self.attention_op.apply_attention( + query=query, + key=key, + value=value, + attention_mask=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=query, + key=key, + value=value, + attention_mask=attention_mask, + spatiotemporal_shape=spatiotemporal_shape, + sparse_config_override=sparse_config, + ) + + with self.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) + else: + with self.named_scope("apply_attention"): + attn_output = self.attention_op.apply_attention(query=query, key=key, value=value, attention_mask=attention_mask) 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..bdac39ca3 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]] = 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={"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,13 @@ 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 = { + "use_base2_exp": use_base2_exp, + "use_experimental_scheduler": use_experimental_scheduler, + "ulysses_shards": ulysses_shards, + "ulysses_attention_chunks": ulysses_attention_chunks, + **(attention_config or {}), + } if sharding_specs is None: sharding_specs = get_sharding_specs("default", "ltx2_dit") @@ -1145,6 +1164,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 +1209,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 +1420,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 +1622,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 +1647,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 +1660,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 +1676,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 +1709,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 +1724,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 +1736,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 +1764,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 +1778,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/tests/ltx2/test_svg_attention_ltx2.py b/src/maxdiffusion/tests/ltx2/test_svg_attention_ltx2.py new file mode 100644 index 000000000..bec100710 --- /dev/null +++ b/src/maxdiffusion/tests/ltx2/test_svg_attention_ltx2.py @@ -0,0 +1,429 @@ +""" +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.wan.transformers.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")) + + 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) + + 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()