def estimate_transformer_forward(
module: torch.nn.Module,
kwargs: dict[str, Any],
output: Any,
*,
role: str,
attention_metadata: Any | None,
cross_attention_cached: bool = False,
) -> _ForwardWork | None:
"""Estimate one Wan-style Transformer invocation from its logical input."""
hidden_states = _first_tensor(kwargs.get("hidden_states"))
if hidden_states is None or hidden_states.ndim != 5:
return None
arch = _arch_config(module)
if arch is None:
return None
required = ("hidden_size", "ffn_dim", "num_layers", "patch_size")
if any(not hasattr(arch, name) for name in required):
return None
batch_size = int(hidden_states.shape[0])
raw_frames = int(hidden_states.shape[2])
raw_height = int(hidden_states.shape[3])
raw_width = int(hidden_states.shape[4])
patch_size = arch.patch_size
patch: tuple[int,
...] = ((1, patch_size,
patch_size) if isinstance(patch_size, int) else tuple(int(value) for value in patch_size))
if len(patch) != 3 or any(value <= 0 for value in patch):
return None
frames = raw_frames // patch[0]
spatial_tokens = (raw_height // patch[1]) * (raw_width // patch[2])
seq_len = frames * spatial_tokens
if seq_len <= 0:
return None
hidden_size = int(arch.hidden_size)
ffn_dim = int(arch.ffn_dim)
num_layers = int(arch.num_layers)
context = _first_tensor(kwargs.get("encoder_hidden_states"))
context_tokens = int(context.shape[-2]) if context is not None and context.ndim >= 2 else 0
image_context = _first_tensor(kwargs.get("encoder_hidden_states_image"))
if image_context is not None and image_context.ndim >= 2:
context_tokens += int(image_context.shape[-2])
is_causal = hasattr(module, "num_frame_per_block")
teacher_forcing = is_causal and _first_tensor(kwargs.get("clean_x")) is not None
query_tokens = seq_len
dense_pairs = float(seq_len * seq_len)
vsa_tiles = 0
causal_chunks = 0
if is_causal:
frames_per_block = int(getattr(module, "num_frame_per_block", 1))
causal_chunks = math.ceil(frames / max(1, frames_per_block))
local_attn_size = int(getattr(module, "local_attn_size", -1))
kv_cache = kwargs.get("kv_cache")
if kv_cache is not None:
current_start = int(kwargs.get("current_start", 0))
# Causal Wan caps the KV window at
# GLOBAL_ATTN_COMPAT_MAX_LATENT_FRAMES (21) frames when
# local_attn_size is unset; sliding_window_num_frames only
# sizes the streaming KV cache. MatrixGame2's 15-frame
# compatibility window is not modeled here.
max_frames = local_attn_size if local_attn_size >= 0 else 21
key_tokens = min(current_start + seq_len, max_frames * spatial_tokens)
attention_pairs = float(seq_len * key_tokens)
dense_pairs = float(seq_len * (current_start + seq_len))
elif teacher_forcing:
query_tokens = 2 * seq_len
attention_pairs = float(_teacher_forcing_frame_pairs(frames, frames_per_block) * spatial_tokens**2)
dense_pairs = float(query_tokens * query_tokens)
else:
attention_pairs = float(
_blockwise_causal_frame_pairs(frames, frames_per_block, local_attn_size) * spatial_tokens**2)
else:
attention_pairs, vsa_tiles = _vsa_attention_pairs(seq_len, attention_metadata)
# Causal Wan pads text to ``text_len`` before every cross-attention block.
if is_causal:
context_tokens = max(context_tokens, int(getattr(module, "text_len", 0)))
b = batch_size
seq_length = query_tokens
dim = hidden_size
# Four self-attention projections, two MLP projections, cross-attention
# Q/K/V/out projections, and the two attention matmuls.
per_layer = (8.0 * b * seq_length * dim * dim + 4.0 * b * seq_length * dim * ffn_dim +
4.0 * b * seq_length * dim * dim +
(0.0 if cross_attention_cached else 4.0 * b * context_tokens * dim * dim) +
4.0 * b * seq_length * context_tokens * dim + 4.0 * b * attention_pairs * dim)
if vsa_tiles:
# Wan VSA adds one gate projection and dense pooled QK/AV matmuls.
per_layer += 2.0 * b * seq_length * dim * dim
per_layer += 4.0 * b * vsa_tiles * vsa_tiles * dim
forward_flops = per_layer * num_layers
backward_expected = _output_requires_grad(output)
# Standard MFU counts model-algorithm FLOPs, excluding activation
# checkpoint recompute: training is approximately F + 2B = 3F.
useful_flops = forward_flops * (3.0 if backward_expected else 1.0)
return _ForwardWork(
role=role,
forward_flops=forward_flops,
useful_flops=useful_flops,
causal_chunks=causal_chunks,
query_frames=b * (2 * frames if teacher_forcing else frames),
query_tokens=b * query_tokens,
attention_pairs=b * attention_pairs,
dense_attention_pairs=b * dense_pairs,
)