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
126 changes: 124 additions & 2 deletions src/maxdiffusion/models/ltx2/attention_ltx2.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
import jax.numpy as jnp
from ... import common_types
from ..attention_flax import NNXAttentionOp
from ..wan.transformers import svg_attention

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.

Nit / architecture: importing from ..wan.transformers makes the ltx2 model package depend on wan. Since svg_attention.py has no Wan-specific logic, could we move svg_attention.py (or re-export it) under a shared location like maxdiffusion/models/svg_attention.py (or alongside attention_flax.py)? Happy for this to be a quick follow-up PR if you'd rather keep this diff small.

from .logical_sharding_ltx2 import get_sharding_specs, LTX2DiTShardingSpecs

Array = common_types.Array
Expand Down Expand Up @@ -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 = {

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 things on this config block:

  1. Duplication with FlaxWanAttention: This 50-line block (and the dispatch block at lines 653–710 below) is almost identical to FlaxWanAttention in src/maxdiffusion/models/attention_flax.py (lines 2583–2630 and 2937–2990). If we ever change an SVG default or add a parameter, it's easy to update one model and forget the other. Could we extract a shared helper (e.g. a small SVGConfig dataclass or init_svg_config(attention_config, default_num_layers=...) + apply_svg_or_dense(...) in svg_attention.py) and call it from both FlaxWanAttention and LTX2Attention?

  2. Ignored / unused keys in attention_config:

    • use_base2_exp, use_experimental_scheduler, ulysses_shards, and ulysses_attention_chunks are placed into attention_config here, but self.attention_op = NNXAttentionOp(...) at lines 568–571 reads the function arguments (ulysses_shards, use_base2_exp, etc.) directly instead of attention_config[...]. That means if someone passes attention_config={"use_base2_exp": True}, it gets silently ignored. Either read them from attention_config when constructing NNXAttentionOp, or remove those 4 keys from attention_config.
    • svg_implementation, svg_high_noise_density, svg_low_noise_density, and svg_global_offset are stored on self, but never actually read by _head_local_svg_attention. Let's remove the unused ones (or raise a ValueError if non-default values are passed) so users don't think setting them changes behavior.

"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
Expand Down Expand Up @@ -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,

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.

deterministic: bool = True is added to LTX2Attention.__call__, and line 608 checks if not deterministic: raise ValueError(...). However, LTX2VideoTransformerBlock.__call__ (in transformer_ltx2.py line 600) never passes deterministic when calling self.attn1(...), so it will always be True in practice.

Could we either wire deterministic through from LTX2VideoTransformer3DModel.__call__ -> LTX2StaticContext -> self.attn1, or drop the deterministic argument here if LTX-2 attention is already inference-only?

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)
Expand Down Expand Up @@ -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,
)
Comment on lines +653 to +666

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.

  1. Minor: at line 653, and spatiotemporal_shape is not None is redundant because lines 607–611 already raise a ValueError if self.use_svg_attention and is_self_attention and spatiotemporal_shape is None.
  2. As mentioned above, this whole is_svg_active + run_dense / run_sparse_svg + jax.lax.cond block is identical to FlaxWanAttention.__call__ (attention_flax.py lines 2937–2990). Pulling this into a helper function in svg_attention.py that takes self.attention_op (or a dense_fn / sparse_fn) would cut ~60 lines of duplicate code here.


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]
Expand Down
Loading
Loading