Compare commits

...
Author SHA1 Message Date
SolitaryThinker c2becf42c2 [misc]: add dated resume handoff for the BF16 MFU campaign 2026-07-28 02:50:37 -07:00
SolitaryThinker 9c86d28e11 [misc]: kill FA4 tail-tile hypothesis; anomaly reproduced standalone as forward-schedule intrinsic 2026-07-22 12:41:15 -07:00
SolitaryThinker a7317d2217 [misc]: record traced lt-pinned variant neutral on unstable tray; healthy-allocation gate required 2026-07-22 12:33:23 -07:00
SolitaryThinker 8bf6eff18f [perf]: bank first kernel-level MFU gain via pinned cuBLASLt backward algos 2026-07-22 10:58:25 -07:00
SolitaryThinker acad815f5c [perf]: cuBLASLt algo sweep finds 8-9% GEMM band headroom with bit-exact parity 2026-07-22 09:21:15 -07:00
SolitaryThinker 971827ad95 [misc]: add prioritized BF16 kernel research plan toward 50% MFU 2026-07-22 09:13:32 -07:00
SolitaryThinker e82addb927 [misc]: quantify the 50% MFU feasibility budget; reject FA4 split tuning 2026-07-22 09:10:26 -07:00
SolitaryThinker 26b455b7b6 [misc]: record 8x env-stack gate on degraded rack-059 pair; close BF16 campaign handoff 2026-07-22 09:00:15 -07:00
SolitaryThinker c459a18978 [misc]: root-cause rack-057 8x MNNVL/IMEX hang; record working multi-node preamble 2026-07-22 08:41:44 -07:00
SolitaryThinker 9435b3f29f [perf]: accept inductor coordinate-descent env stack; record BF16 campaign gates 2026-07-22 07:54:06 -07:00
SolitaryThinker 3083c59ef4 [misc]: audit and confirm the LTX2 MFU numerator; document the 2450 TFLOP/s convention 2026-07-22 06:06:33 -07:00
SolitaryThinker 6199cbee66 [misc]: re-gate tracker head as timing-neutral on 4x GB200 2026-07-22 05:27:12 -07:00
SolitaryThinker 52f1114dd9 [misc]: add resumable LTX2 MFU tracker 2026-07-22 04:13:27 -07:00
SolitaryThinker 3f3f06541c [bugfix]: make LTX video validation memory-safe 2026-07-22 03:15:20 -07:00
SolitaryThinker 0e60a0e9cc [bugfix]: isolate compiled validation inference 2026-07-22 03:14:47 -07:00
SolitaryThinker 20c36acefc [perf]: group adjacent FSDP2 modules 2026-07-21 23:26:39 -07:00
SolitaryThinker fa47ce1ab5 [perf]: enable packed LTX projections in overfit recipe 2026-07-21 21:05:59 -07:00
SolitaryThinker 7afd751915 [test]: isolate LTX SP subprocess imports 2026-07-21 21:01:52 -07:00
SolitaryThinker cc5913e5c4 [test]: cover fused parameter export 2026-07-21 21:00:42 -07:00
SolitaryThinker 26f909c520 [test]: cover packed LTX sequence parallelism 2026-07-21 20:56:27 -07:00
SolitaryThinker d016237096 [perf]: pack LTX attention projections 2026-07-21 20:51:04 -07:00
SolitaryThinker 002ec0771b [perf]: retain FSDP parameters across accumulation 2026-07-21 19:44:44 -07:00
SolitaryThinker 49508050b7 [perf]: embed uniform LTX2 timestep once 2026-07-21 18:42:06 -07:00
SolitaryThinker 7f6c290c93 [perf]: defer gradient norm materialization 2026-07-21 18:16:47 -07:00
SolitaryThinker 7c58950a92 [perf]: keep LTX2 overfit input pipeline warm 2026-07-21 17:21:41 -07:00
SolitaryThinker 0b324d0a40 [perf]: skip redundant FSDP accumulation sync 2026-07-21 16:54:40 -07:00
SolitaryThinker e42cfa5e5b [perf]: add FSDP symmetric-memory training 2026-07-21 15:49:14 -07:00
SolitaryThinker 7f139e2b28 [perf]: make FSDP reduction precision configurable 2026-07-21 07:19:23 -07:00
SolitaryThinker acdbc0a614 [perf]: batch FSDP gradient clipping 2026-07-21 07:08:38 -07:00
SolitaryThinker 421227ecd4 [perf]: omit unused LTX2 audio weights in training 2026-07-21 06:20:36 -07:00
SolitaryThinker 0431d4e8a9 [perf]: enable regional compile for training 2026-07-21 06:07:35 -07:00
SolitaryThinker cfc110cb6b [perf]: make FSDP forward reshard configurable 2026-07-21 05:18:52 -07:00
SolitaryThinker 9582f012df [perf]: add opt-in fused AdamW training
Expose PyTorch's CUDA fused AdamW through modular training config while preserving its current automatic default. Seed optimizer step state on-device when fused or capturable so distributed-checkpoint resume remains valid.
2026-07-21 04:03:24 -07:00
155 changed files with 19539 additions and 334 deletions
+8
View File
@@ -119,6 +119,7 @@ training:
tp_size: 1 # tensor parallelism
hsdp_replicate_dim: 1 # HSDP replication dimension
hsdp_shard_dim: 8 # HSDP sharding dimension
fsdp_modules_per_group: 1 # consecutive selected modules per FSDP communication group
data:
data_path: data/my_dataset
@@ -510,12 +511,19 @@ training:
tp_size: 1 # tensor parallelism group size
hsdp_replicate_dim: 1 # number of HSDP replicas
hsdp_shard_dim: 8 # number of HSDP shards
fsdp_modules_per_group: 1 # default: one selected module per communication group
```
**HSDP** shards model parameters across `hsdp_shard_dim` GPUs and replicates
across `hsdp_replicate_dim` groups. The product
`hsdp_replicate_dim * hsdp_shard_dim` should equal `num_gpus`.
`fsdp_modules_per_group` groups consecutive FSDP modules, in model traversal
order, into each communication group. Values above the default of `1` reduce
collective launches, but use larger communication buffers and can reduce
communication/compute overlap. Benchmark the target model, batch size, and GPU
topology rather than assuming a speedup.
**Sequence parallelism** splits the sequence (video frames) across `sp_size`
GPUs within each data-parallel group. Useful for long videos that don't fit on a
single GPU.
+6
View File
@@ -86,6 +86,10 @@ training:
hsdp_replicate_dim: 1 # default: 1
hsdp_shard_dim: 8 # default: -1 (defaults to num_gpus in loader)
pin_cpu_memory: false # default: false
reshard_after_forward: true # default: true; false trades memory for fewer FSDP all-gathers
fsdp_symmetric_memory: false # default: false; native NCCL symmetric-memory FSDP collectives
fsdp_modules_per_group: 1 # default: 1; consecutive selected modules per FSDP communication group
reduce_dtype: fp32 # default: "fp32"; "bf16" reduces FSDP communication and memory
# --- training.data [TYPED] -> DataConfig ---
data:
@@ -111,6 +115,7 @@ training:
lr_num_cycles: 0 # default: 0
lr_power: 0.0 # default: 0.0
min_lr_ratio: 0.5 # default: 0.5
fused: null # null=PyTorch default; true=CUDA fused AdamW
# --- training.loop [TYPED] -> TrainingLoopConfig ---
loop:
@@ -148,6 +153,7 @@ training:
precondition_outputs: false # default: false
moba_config: {} # default: {}
enable_gradient_checkpointing_type: full # default: null ("full" or null)
enable_torch_compile: false # default: false; regionally compile matched DiT blocks before FSDP
# --- training top-level [TYPED] ---
dit_precision: fp32 # default: "fp32" (master weight precision)
+15 -3
View File
@@ -27,10 +27,16 @@ training:
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
# Grouping two adjacent LTX blocks with public FSDP2 was the best
# latency/memory tradeoff measured on one four-GPU GB200 tray.
fsdp_modules_per_group: 2
data:
data_path: data/ltx2_overfit_preprocessed
dataloader_num_workers: 0
# Repeat paths virtually so the four-row fixture spans the full run
# without rebuilding the iterator or duplicating data on disk.
data_path:
data/ltx2_overfit_preprocessed: 300
dataloader_num_workers: 1
train_batch_size: 1
# LTX2Model requires 0.0: CFG dropout would zero post-connector
# embeddings, which is not the model's unconditional input.
@@ -83,8 +89,14 @@ callbacks:
sampling_steps: [8]
guidance_scale: 1.0
num_frames: 81
# Release validation-only VAE/text/audio modules before training resumes.
unload_pipeline_after_validation: true
# Required so the LTX2T2VConfig pipeline config is resolved from
# init_from (without a `pipeline:` key the loader falls back to a
# generic PipelineConfig and the LTX-2 DiT cannot be constructed).
pipeline: {}
pipeline:
dit_config:
# Persistently pack video self-attention QKV and text cross-attention KV
# projections. Checkpoints are still loaded and exported with split keys.
pack_attention_projections: true
+149 -24
View File
@@ -9,6 +9,12 @@ from fastvideo.platforms import current_platform
logger = init_logger(__name__)
def _empty_like_fa4_backward_input(x: torch.Tensor) -> torch.Tensor:
"""Mirror FA4's last-dimension contiguity normalization in fake kernels."""
return torch.empty_like(x.contiguous() if x.stride(-1) != 1 else x)
if torch.cuda.is_available():
try:
from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd
@@ -108,6 +114,7 @@ def _flash_attn_cute_forward(
softcap=0.0,
num_splits=1,
pack_gqa=None,
return_lse=True,
)[:2]
return out, lse
@@ -132,38 +139,92 @@ def _flash_attn_cute_setup_context(ctx: torch.autograd.function.FunctionCtx, inp
q, k, v, softmax_scale, causal, deterministic = inputs
out, lse = output
ctx.save_for_backward(q, k, v, out, lse)
ctx.mark_non_differentiable(lse)
ctx.softmax_scale = softmax_scale
ctx.causal = causal
ctx.deterministic = deterministic
def _flash_attn_cute_backward(
ctx: torch.autograd.function.FunctionCtx,
# ``register_autograd`` backward callbacks must be AOT-traceable. CuTe's
# Python backward skips kernel launches under fake mode, so expose the real
# dense and varlen backward kernels as opaque custom ops instead of tracing
# their Python wrappers into allocation-only graphs.
@torch.library.custom_op(
"fastvideo::_flash_attn_cute_backward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_cute_backward_op(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
out: torch.Tensor,
grad_out: torch.Tensor,
grad_lse: torch.Tensor | None,
):
del grad_lse
q, k, v, out, lse = ctx.saved_tensors
dq, dk, dv = _flash_attn_bwd(
lse: torch.Tensor,
softmax_scale: float | None,
causal: bool,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
return _flash_attn_bwd(
q,
k,
v,
out,
grad_out,
lse,
softmax_scale=ctx.softmax_scale,
causal=ctx.causal,
softmax_scale=softmax_scale,
causal=causal,
softcap=0.0,
window_size_left=None,
window_size_right=None,
deterministic=ctx.deterministic,
deterministic=deterministic,
)
@torch.library.register_fake("fastvideo::_flash_attn_cute_backward")
def _flash_attn_cute_backward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
out: torch.Tensor,
grad_out: torch.Tensor,
lse: torch.Tensor,
softmax_scale: float | None,
causal: bool,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
del out, grad_out, lse, softmax_scale, causal, deterministic
return (
_empty_like_fa4_backward_input(q),
_empty_like_fa4_backward_input(k),
_empty_like_fa4_backward_input(v),
)
def _flash_attn_cute_autograd_backward(
ctx: torch.autograd.function.FunctionCtx,
grad_out: torch.Tensor,
grad_lse: torch.Tensor | None,
):
del grad_lse
q, k, v, out, lse = ctx.saved_tensors
dq, dk, dv = torch.ops.fastvideo._flash_attn_cute_backward(
q,
k,
v,
out,
grad_out,
lse,
ctx.softmax_scale,
ctx.causal,
ctx.deterministic,
)
return dq, dk, dv, None, None, None
torch.library.register_autograd(
"fastvideo::_flash_attn_cute_forward",
_flash_attn_cute_backward,
_flash_attn_cute_autograd_backward,
setup_context=_flash_attn_cute_setup_context,
)
@@ -200,6 +261,7 @@ def _flash_attn_cute_varlen_forward(
softcap=0.0,
num_splits=1,
pack_gqa=None,
return_lse=True,
)[:2]
return out, lse
@@ -241,6 +303,7 @@ def _flash_attn_cute_varlen_setup_context(ctx: torch.autograd.function.FunctionC
) = inputs
out, lse = output
ctx.save_for_backward(q, k, v, out, lse, cu_seqlens_q, cu_seqlens_k)
ctx.mark_non_differentiable(lse)
ctx.max_seqlen_q = max_seqlen_q
ctx.max_seqlen_k = max_seqlen_k
ctx.softmax_scale = softmax_scale
@@ -248,37 +311,99 @@ def _flash_attn_cute_varlen_setup_context(ctx: torch.autograd.function.FunctionC
ctx.deterministic = deterministic
def _flash_attn_cute_varlen_backward(
ctx: torch.autograd.function.FunctionCtx,
@torch.library.custom_op(
"fastvideo::_flash_attn_cute_varlen_backward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_cute_varlen_backward_op(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
out: torch.Tensor,
grad_out: torch.Tensor,
grad_lse: torch.Tensor | None,
):
del grad_lse
q, k, v, out, lse, cu_seqlens_q, cu_seqlens_k = ctx.saved_tensors
dq, dk, dv = _flash_attn_bwd(
lse: torch.Tensor,
cu_seqlens_q: torch.Tensor,
cu_seqlens_k: torch.Tensor,
max_seqlen_q: int,
max_seqlen_k: int,
softmax_scale: float | None,
causal: bool,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
return _flash_attn_bwd(
q,
k,
v,
out,
grad_out,
lse,
softmax_scale=ctx.softmax_scale,
causal=ctx.causal,
softmax_scale=softmax_scale,
causal=causal,
softcap=0.0,
window_size_left=None,
window_size_right=None,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=ctx.max_seqlen_q,
max_seqlen_k=ctx.max_seqlen_k,
deterministic=ctx.deterministic,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k,
deterministic=deterministic,
)
@torch.library.register_fake("fastvideo::_flash_attn_cute_varlen_backward")
def _flash_attn_cute_varlen_backward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
out: torch.Tensor,
grad_out: torch.Tensor,
lse: torch.Tensor,
cu_seqlens_q: torch.Tensor,
cu_seqlens_k: torch.Tensor,
max_seqlen_q: int,
max_seqlen_k: int,
softmax_scale: float | None,
causal: bool,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
del out, grad_out, lse, cu_seqlens_q, cu_seqlens_k
del max_seqlen_q, max_seqlen_k, softmax_scale, causal, deterministic
return (
_empty_like_fa4_backward_input(q),
_empty_like_fa4_backward_input(k),
_empty_like_fa4_backward_input(v),
)
def _flash_attn_cute_varlen_autograd_backward(
ctx: torch.autograd.function.FunctionCtx,
grad_out: torch.Tensor,
grad_lse: torch.Tensor | None,
):
del grad_lse
q, k, v, out, lse, cu_seqlens_q, cu_seqlens_k = ctx.saved_tensors
dq, dk, dv = torch.ops.fastvideo._flash_attn_cute_varlen_backward(
q,
k,
v,
out,
grad_out,
lse,
cu_seqlens_q,
cu_seqlens_k,
ctx.max_seqlen_q,
ctx.max_seqlen_k,
ctx.softmax_scale,
ctx.causal,
ctx.deterministic,
)
return dq, dk, dv, None, None, None, None, None, None, None
torch.library.register_autograd(
"fastvideo::_flash_attn_cute_varlen_forward",
_flash_attn_cute_varlen_backward,
_flash_attn_cute_varlen_autograd_backward,
setup_context=_flash_attn_cute_varlen_setup_context,
)
+33 -12
View File
@@ -3,11 +3,30 @@
LTX-2 Transformer configuration for native FastVideo integration.
"""
import re
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
import re
_PACKED_PROJECTION_MAPPING: dict[str, str | tuple[str, int, int]] = {
r"^(?:model\.diffusion_model\.|diffusion_model\.|model\.)?(transformer_blocks\.\d+\.attn1)\.to_q\.(weight|bias)$":
(r"model.\1.to_qkv.\2", 0, 3),
r"^(?:model\.diffusion_model\.|diffusion_model\.|model\.)?(transformer_blocks\.\d+\.attn1)\.to_k\.(weight|bias)$":
(r"model.\1.to_qkv.\2", 1, 3),
r"^(?:model\.diffusion_model\.|diffusion_model\.|model\.)?(transformer_blocks\.\d+\.attn1)\.to_v\.(weight|bias)$":
(r"model.\1.to_qkv.\2", 2, 3),
r"^(?:model\.diffusion_model\.|diffusion_model\.|model\.)?(transformer_blocks\.\d+\.attn2)\.to_k\.(weight|bias)$":
(r"model.\1.to_kv.\2", 0, 2),
r"^(?:model\.diffusion_model\.|diffusion_model\.|model\.)?(transformer_blocks\.\d+\.attn2)\.to_v\.(weight|bias)$":
(r"model.\1.to_kv.\2", 1, 2),
}
_GATED_ATTENTION_MAPPING: dict[str, str] = {
r"^model\.diffusion_model\.(.*)\.to_gate_compress\.(.*)$": r"model.\1.to_gate_logits.\2",
r"^diffusion_model\.(.*)\.to_gate_compress\.(.*)$": r"model.\1.to_gate_logits.\2",
r"^model\.(.*)\.to_gate_compress\.(.*)$": r"model.\1.to_gate_logits.\2",
r"^(.*)\.to_gate_compress\.(.*)$": r"model.\1.to_gate_logits.\2",
}
def is_ltx2_blocks(name: str, _module) -> bool:
@@ -52,6 +71,9 @@ class LTX2VideoArchConfig(DiTArchConfig):
attention_type: str = "default"
rope_type: str = "split"
double_precision_rope: bool = True
# Opt-in persistent packing for video self-attention QKV and text
# cross-attention KV projections. The checkpoint remains split externally.
pack_attention_projections: bool = False
# LTX-2.3 gated extensions. All default OFF == LTX-2.0 behavior.
cross_attention_adaln: bool = False
caption_proj_before_connector: bool = False
@@ -116,17 +138,16 @@ class LTX2VideoArchConfig(DiTArchConfig):
# an unconditional rename would silently retarget them. Inserted at
# the front so first-match-wins matching fires the rename before the
# generic prefix-strip rules.
if self.apply_gated_attention:
gate_rules = {
r"^model\.diffusion_model\.(.*)\.to_gate_compress\.(.*)$": r"model.\1.to_gate_logits.\2",
r"^diffusion_model\.(.*)\.to_gate_compress\.(.*)$": r"model.\1.to_gate_logits.\2",
r"^model\.(.*)\.to_gate_compress\.(.*)$": r"model.\1.to_gate_logits.\2",
r"^(.*)\.to_gate_compress\.(.*)$": r"model.\1.to_gate_logits.\2",
}
self.param_names_mapping = {
**gate_rules,
**self.param_names_mapping,
}
dynamic_patterns = set(_PACKED_PROJECTION_MAPPING) | set(_GATED_ATTENTION_MAPPING)
base_mapping = {
pattern: replacement
for pattern, replacement in self.param_names_mapping.items() if pattern not in dynamic_patterns
}
self.param_names_mapping = {
**(_PACKED_PROJECTION_MAPPING if self.pack_attention_projections else {}),
**(_GATED_ATTENTION_MAPPING if self.apply_gated_attention else {}),
**base_mapping,
}
@dataclass
+2
View File
@@ -908,6 +908,8 @@ class TrainingArgs(FastVideoArgs):
mixed_precision: str = ""
train_sp_batch_size: int = 0
fsdp_sharding_startegy: str = ""
fsdp_reduce_dtype: str = "fp32"
fsdp_modules_per_group: int = 1
weighting_scheme: str = ""
logit_mean: float = 0.0
+195 -77
View File
@@ -1447,6 +1447,103 @@ class LTXLocalAttention(LocalAttention):
return output
def _init_attention_projections(
module: Any,
*,
query_dim: int,
context_dim: int | None,
inner_dim: int,
quant_config: QuantizationConfig | None,
prefix: str,
pack_attention_projections: bool,
) -> None:
"""Create split projections or persistent self-QKV/cross-KV packs."""
is_self_attention = context_dim is None
context_dim = query_dim if context_dim is None else context_dim
if pack_attention_projections and quant_config is not None:
raise ValueError(
"LTX-2 packed attention projections do not yet support linear quantization"
)
if pack_attention_projections and is_self_attention:
module.to_qkv = ReplicatedLinear(
query_dim,
3 * inner_dim,
bias=True,
prefix=f"{prefix}.to_qkv",
)
return
module.to_q = ReplicatedLinear(
query_dim,
inner_dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.to_q",
)
if pack_attention_projections:
module.to_kv = ReplicatedLinear(
context_dim,
2 * inner_dim,
bias=True,
prefix=f"{prefix}.to_kv",
)
return
module.to_k = ReplicatedLinear(
context_dim,
inner_dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.to_k",
)
module.to_v = ReplicatedLinear(
context_dim,
inner_dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.to_v",
)
def _project_attention_inputs(
module: Any,
x: torch.Tensor,
context: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Project Q/K/V while preserving the existing quantized split path."""
to_qkv = getattr(module, "to_qkv", None)
if to_qkv is not None:
if context is not x:
raise ValueError(
"Packed self-attention QKV requires query and context to be the same tensor"
)
return to_qkv(x)[0].chunk(3, dim=-1)
to_kv = getattr(module, "to_kv", None)
if to_kv is not None:
q = module.to_q(x)[0]
k, v = to_kv(context)[0].chunk(2, dim=-1)
return q, k, v
q_pre_quantized = (
module.to_q.quant_method.quantize_input(x) # type: ignore[union-attr]
if _supports_prequantized_input(module.to_q) else None)
kv_pre_quantized = None
if (_supports_prequantized_input(module.to_k)
and _supports_prequantized_input(module.to_v)):
if context is x and q_pre_quantized is not None:
kv_pre_quantized = q_pre_quantized
else:
kv_pre_quantized = module.to_k.quant_method.quantize_input( # type: ignore[union-attr]
context)
q = _linear_project_with_optional_prequant(module.to_q, x,
q_pre_quantized)
k = _linear_project_with_optional_prequant(module.to_k, context,
kv_pre_quantized)
v = _linear_project_with_optional_prequant(module.to_v, context,
kv_pre_quantized)
return q, k, v
class LTXSelfAttention(nn.Module):
"""LTX-2 attention block with RMSNorm + FastVideo LocalAttention."""
def __init__(
@@ -1460,10 +1557,12 @@ class LTXSelfAttention(nn.Module):
supported_attention_backends: tuple[AttentionBackendEnum, ...],
apply_gated_attention: bool = False,
quant_config: QuantizationConfig | None = None,
pack_attention_projections: bool = False,
prefix: str = "",
) -> None:
super().__init__()
inner_dim = dim_head * heads
is_self_attention = context_dim is None
context_dim = query_dim if context_dim is None else context_dim
self.heads = heads
@@ -1474,26 +1573,14 @@ class LTXSelfAttention(nn.Module):
# LTX-2.3 refine stage, so default-off behavior matches LTX-2.0.
self.q_norm = StageAwareRMSNorm(inner_dim, eps=norm_eps)
self.k_norm = StageAwareRMSNorm(inner_dim, eps=norm_eps)
self.to_q = ReplicatedLinear(
query_dim,
inner_dim,
bias=True,
_init_attention_projections(
self,
query_dim=query_dim,
context_dim=None if is_self_attention else context_dim,
inner_dim=inner_dim,
quant_config=quant_config,
prefix=f"{prefix}.to_q",
)
self.to_k = ReplicatedLinear(
context_dim,
inner_dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.to_k",
)
self.to_v = ReplicatedLinear(
context_dim,
inner_dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.to_v",
prefix=prefix,
pack_attention_projections=pack_attention_projections,
)
# LTX-2.3 gated attention. DISTINCT from the VSA-QAT to_gate_compress
# gate below: this gate multiplies the *self-attention output* by
@@ -1556,21 +1643,7 @@ class LTXSelfAttention(nn.Module):
) -> torch.Tensor:
context = x if context is None else context
q_pre_quantized = (
self.to_q.quant_method.quantize_input(x) # type: ignore[union-attr]
if _supports_prequantized_input(self.to_q) else None)
kv_pre_quantized = None
if (_supports_prequantized_input(self.to_k)
and _supports_prequantized_input(self.to_v)):
if context is x and q_pre_quantized is not None:
kv_pre_quantized = q_pre_quantized
else:
kv_pre_quantized = self.to_k.quant_method.quantize_input( # type: ignore[union-attr]
context)
q = _linear_project_with_optional_prequant(self.to_q, x, q_pre_quantized)
k = _linear_project_with_optional_prequant(self.to_k, context, kv_pre_quantized)
v = _linear_project_with_optional_prequant(self.to_v, context, kv_pre_quantized)
q, k, v = _project_attention_inputs(self, x, context)
gate_logits = (self.to_gate_logits(x)[0]
if self.to_gate_logits is not None else None)
gate_compress = (self.to_gate_compress(context)[0]
@@ -1647,10 +1720,12 @@ class LTXDistributedSelfAttention(nn.Module):
supported_attention_backends: tuple[AttentionBackendEnum, ...],
apply_gated_attention: bool = False,
quant_config: QuantizationConfig | None = None,
pack_attention_projections: bool = False,
prefix: str = "",
) -> None:
super().__init__()
inner_dim = dim_head * heads
is_self_attention = context_dim is None
context_dim = query_dim if context_dim is None else context_dim
self.heads = heads
@@ -1661,26 +1736,14 @@ class LTXDistributedSelfAttention(nn.Module):
# stage, preserving LTX-2.0 numerics by default.
self.q_norm = StageAwareRMSNorm(inner_dim, eps=norm_eps)
self.k_norm = StageAwareRMSNorm(inner_dim, eps=norm_eps)
self.to_q = ReplicatedLinear(
query_dim,
inner_dim,
bias=True,
_init_attention_projections(
self,
query_dim=query_dim,
context_dim=None if is_self_attention else context_dim,
inner_dim=inner_dim,
quant_config=quant_config,
prefix=f"{prefix}.to_q",
)
self.to_k = ReplicatedLinear(
context_dim,
inner_dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.to_k",
)
self.to_v = ReplicatedLinear(
context_dim,
inner_dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.to_v",
prefix=prefix,
pack_attention_projections=pack_attention_projections,
)
# LTX-2.3 gated attention (distinct from VSA-QAT to_gate_compress).
self.to_gate_logits: ReplicatedLinear | None = None
@@ -1741,21 +1804,7 @@ class LTXDistributedSelfAttention(nn.Module):
"""
context = x if context is None else context
q_pre_quantized = (
self.to_q.quant_method.quantize_input(x) # type: ignore[union-attr]
if _supports_prequantized_input(self.to_q) else None)
kv_pre_quantized = None
if (_supports_prequantized_input(self.to_k)
and _supports_prequantized_input(self.to_v)):
if context is x and q_pre_quantized is not None:
kv_pre_quantized = q_pre_quantized
else:
kv_pre_quantized = self.to_k.quant_method.quantize_input( # type: ignore[union-attr]
context)
q = _linear_project_with_optional_prequant(self.to_q, x, q_pre_quantized)
k = _linear_project_with_optional_prequant(self.to_k, context, kv_pre_quantized)
v = _linear_project_with_optional_prequant(self.to_v, context, kv_pre_quantized)
q, k, v = _project_attention_inputs(self, x, context)
gate_logits = (self.to_gate_logits(x)[0]
if self.to_gate_logits is not None else None)
gate_compress = (self.to_gate_compress(context)[0]
@@ -1812,6 +1861,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
cross_attention_adaln: bool = False,
stg_block_idx: int = 29,
quant_config: QuantizationConfig | None = None,
pack_attention_projections: bool = False,
prefix: str = "",
):
super().__init__()
@@ -1848,6 +1898,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
supported_attention_backends=video_self_attn_backends,
apply_gated_attention=video.apply_gated_attention,
quant_config=quant_config,
pack_attention_projections=pack_attention_projections,
prefix=f"{prefix}.blocks.{idx}.attn1",
)
# Text cross-attention - always local (text is replicated)
@@ -1861,6 +1912,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
supported_attention_backends=dense_attn_backends,
apply_gated_attention=video.apply_gated_attention,
quant_config=quant_config,
pack_attention_projections=pack_attention_projections,
prefix=f"{prefix}.blocks.{idx}.attn2",
)
self.ff = FeedForward(
@@ -1955,6 +2007,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
self.norm_eps = norm_eps
@torch.compiler.disable
def _register_fsdp_backward_hooks_on_output(self, vx, ax):
"""Register backward hooks on output tensors to trigger FSDP2 unshard.
@@ -2410,6 +2463,7 @@ class LTXModel(torch.nn.Module):
stg_block_idx: int = 29,
use_distributed_attention: bool = False,
quant_config: QuantizationConfig | None = None,
pack_attention_projections: bool = False,
prefix: str = "",
):
super().__init__()
@@ -2471,6 +2525,7 @@ class LTXModel(torch.nn.Module):
norm_eps=norm_eps,
use_distributed_attention=use_distributed_attention,
quant_config=quant_config,
pack_attention_projections=pack_attention_projections,
prefix=prefix,
)
@@ -2631,6 +2686,7 @@ class LTXModel(torch.nn.Module):
norm_eps: float,
use_distributed_attention: bool = False,
quant_config: QuantizationConfig | None = None,
pack_attention_projections: bool = False,
prefix: str = "",
) -> None:
video_config = (
@@ -2670,6 +2726,7 @@ class LTXModel(torch.nn.Module):
cross_attention_adaln=self.cross_attention_adaln,
stg_block_idx=self.stg_block_idx,
quant_config=quant_config,
pack_attention_projections=pack_attention_projections,
prefix=prefix,
)
for idx in range(num_layers)
@@ -2749,13 +2806,16 @@ class LTXModel(torch.nn.Module):
self-attention is skipped (STG perturbed pass).
"""
if os.getenv("LTX2_PIPELINE_DEBUG_LOG", "0") == "1":
_debug_block_log_line(
weight_log = (
"fastvideo:patchify_proj"
f":video_w_sum={self.patchify_proj.weight.float().sum().item():.6f} "
f"video_b_sum={self.patchify_proj.bias.float().sum().item():.6f} "
f"audio_w_sum={self.audio_patchify_proj.weight.float().sum().item():.6f} "
f"audio_b_sum={self.audio_patchify_proj.bias.float().sum().item():.6f}"
)
f"video_b_sum={self.patchify_proj.bias.float().sum().item():.6f}")
if self.model_type.is_audio_enabled():
weight_log += (
f" audio_w_sum={self.audio_patchify_proj.weight.float().sum().item():.6f} "
f"audio_b_sum={self.audio_patchify_proj.bias.float().sum().item():.6f}"
)
_debug_block_log_line(weight_log)
if not self.model_type.is_video_enabled() and video is not None:
raise ValueError("Video is not enabled for this model")
if not self.model_type.is_audio_enabled() and audio is not None:
@@ -2806,11 +2866,19 @@ class LTX2Transformer3DModel(BaseDiT):
lora_param_names_mapping = LTX2VideoConfig().lora_param_names_mapping
_fsdp_shard_conditions = LTX2VideoConfig()._fsdp_shard_conditions
_compile_conditions = LTX2VideoConfig()._compile_conditions
_model_type = LTXModelType.AudioVideo
def __init__(self, config: LTX2VideoConfig, hf_config: dict[str, Any]):
super().__init__(config=config, hf_config=hf_config)
arch = config.arch_config
model_type = self._model_type
self.param_names_mapping = dict(arch.param_names_mapping)
if arch.pack_attention_projections and config.quant_config is not None:
raise ValueError(
"LTX-2 packed attention projections do not yet support "
"linear quantization"
)
# Get SP world size for distributed attention
sp_world_size = get_sp_world_size()
@@ -2820,11 +2888,13 @@ class LTX2Transformer3DModel(BaseDiT):
# Validate that attention heads are divisible by SP world size
if sp_world_size > 1:
assert arch.num_attention_heads % sp_world_size == 0, (
assert (not model_type.is_video_enabled()
or arch.num_attention_heads % sp_world_size == 0), (
f"The number of video attention heads ({arch.num_attention_heads}) "
f"must be divisible by the sequence parallel size ({sp_world_size})"
)
assert arch.audio_num_attention_heads % sp_world_size == 0, (
assert (not model_type.is_audio_enabled()
or arch.audio_num_attention_heads % sp_world_size == 0), (
f"The number of audio attention heads ({arch.audio_num_attention_heads}) "
f"must be divisible by the sequence parallel size ({sp_world_size})"
)
@@ -2836,7 +2906,6 @@ class LTX2Transformer3DModel(BaseDiT):
"LTX2 VSA enabled with SP world size 1; using distributed "
"attention path for VSA compatibility")
model_type = LTXModelType.AudioVideo
self.model = LTXModel(
model_type=model_type,
num_attention_heads=arch.num_attention_heads,
@@ -2866,6 +2935,7 @@ class LTX2Transformer3DModel(BaseDiT):
stg_block_idx=arch.stg_block_idx,
use_distributed_attention=use_distributed_attention,
quant_config=config.quant_config,
pack_attention_projections=arch.pack_attention_projections,
prefix=config.prefix,
)
@@ -3138,5 +3208,53 @@ class LTX2Transformer3DModel(BaseDiT):
audio_out, output_shape=audio_shape)
return video_out, audio_out
class LTX2VideoOnlyTransformer3DModel(LTX2Transformer3DModel):
"""LTX-2 video branch used by the modular video-only trainer."""
_model_type = LTXModelType.VideoOnly
_audio_root_modules = frozenset({
"audio_adaln_single",
"audio_caption_projection",
"audio_patchify_proj",
"audio_proj_out",
"audio_prompt_adaln_single",
"av_ca_a2v_gate_adaln_single",
"av_ca_audio_scale_shift_adaln_single",
"av_ca_v2a_gate_adaln_single",
"av_ca_video_scale_shift_adaln_single",
})
_audio_root_parameters = frozenset({"audio_scale_shift_table"})
_audio_block_modules = frozenset({
"audio_attn1",
"audio_attn2",
"audio_ff",
"audio_to_video_attn",
"video_to_audio_attn",
})
_audio_block_parameters = frozenset({
"audio_prompt_scale_shift_table",
"audio_scale_shift_table",
"scale_shift_table_a2v_ca_audio",
"scale_shift_table_a2v_ca_video",
})
@classmethod
def _is_ignored_checkpoint_key(cls, key: str) -> bool:
"""Return whether an AV checkpoint key belongs to the omitted branch."""
parts = key.split(".")
if len(parts) < 2 or parts[0] != "model":
return False
if len(parts) == 2:
return parts[1] in cls._audio_root_parameters
if parts[1] in cls._audio_root_modules:
return True
if (len(parts) < 4 or parts[1] != "transformer_blocks"
or not parts[2].isdigit()):
return False
if len(parts) == 4:
return parts[3] in cls._audio_block_parameters
return parts[3] in cls._audio_block_modules
# Entry point for model registry
EntryClass = LTX2Transformer3DModel
EntryClass = [LTX2Transformer3DModel, LTX2VideoOnlyTransformer3DModel]
+4 -2
View File
@@ -1104,10 +1104,12 @@ class TransformerLoader(ComponentLoader):
cpu_offload=fastvideo_args.dit_cpu_offload,
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
fsdp_inference=fastvideo_args.use_fsdp_inference,
# TODO(will): make these configurable
# FSDP parameters stay BF16; reduction precision is an explicit
# modular-training policy with a conservative FP32 default.
default_dtype=default_dtype,
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
reduce_dtype=PRECISION_TO_TYPE[getattr(fastvideo_args, "fsdp_reduce_dtype", "fp32")],
fsdp_modules_per_group=getattr(fastvideo_args, "fsdp_modules_per_group", 1),
output_dtype=None,
training_mode=fastvideo_args.training_mode,
enable_torch_compile=fastvideo_args.enable_torch_compile,
+94 -9
View File
@@ -26,6 +26,34 @@ from fastvideo.utils import set_mixed_precision_policy, is_pin_memory_available
logger = init_logger(__name__)
def _compile_matched_submodule_forwards(
model: nn.Module,
torch_compile_kwargs: dict[str, Any] | None = None,
) -> int:
"""Compile only forwards selected by ``model._compile_conditions``."""
compile_conditions = getattr(model, "_compile_conditions", None)
if not compile_conditions:
raise ValueError("Regional torch.compile requires model._compile_conditions; "
"refusing to compile the whole FSDP model.")
compile_kwargs = dict(torch_compile_kwargs or {})
if compile_kwargs.get("fullgraph") is True:
raise ValueError("Regional FSDP training compile requires "
"fullgraph=False so runtime-only hooks can remain graph breaks.")
compile_kwargs["fullgraph"] = False
compiled_count = 0
for name, submodule in model.named_modules():
if name and any(condition(name, submodule) for condition in compile_conditions):
submodule.forward = torch.compile(submodule.forward, **compile_kwargs)
compiled_count += 1
if compiled_count == 0:
raise ValueError("model._compile_conditions matched no submodules; "
"refusing to compile the whole FSDP model.")
return compiled_count
def _maybe_quantize_model(model: nn.Module) -> None:
"""Quantize NVFP4- or FP8-tagged linear layers in-place after weights are loaded.
@@ -116,6 +144,7 @@ def maybe_load_fsdp_model(
pin_cpu_memory: bool = True,
enable_torch_compile: bool = False,
torch_compile_kwargs: dict[str, Any] | None = None,
fsdp_modules_per_group: int = 1,
) -> torch.nn.Module:
"""
Load the model with FSDP if is training, else load the model without FSDP.
@@ -151,6 +180,10 @@ def maybe_load_fsdp_model(
use_fsdp = False
logger.info("Disabling FSDP for MPS platform as it's not compatible")
if enable_torch_compile and training_mode:
compiled_count = _compile_matched_submodule_forwards(model, torch_compile_kwargs)
logger.info("Enabled regional torch.compile for %d training submodules before FSDP sharding", compiled_count)
if use_fsdp:
pin_cpu_memory = pin_cpu_memory and is_pin_memory_available()
world_size = hsdp_replicate_dim * hsdp_shard_dim
@@ -179,7 +212,8 @@ def maybe_load_fsdp_model(
mp_policy=mp_policy,
mesh=device_mesh,
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=pin_cpu_memory)
pin_cpu_memory=pin_cpu_memory,
fsdp_modules_per_group=fsdp_modules_per_group)
weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=True)
param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
@@ -209,12 +243,6 @@ def maybe_load_fsdp_model(
# are present (lazy imports inside the helper).
_maybe_quantize_model(model)
compile_in_loader = enable_torch_compile and training_mode
if compile_in_loader:
compile_kwargs = torch_compile_kwargs or {}
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s", compile_kwargs)
model = torch.compile(model, **compile_kwargs)
logger.info("torch.compile enabled for %s", type(model).__name__)
return model
@@ -227,6 +255,7 @@ def shard_model(
mesh: DeviceMesh | None = None,
fsdp_shard_conditions: list[Callable[[str, nn.Module], bool]] = [], # noqa
pin_cpu_memory: bool = True,
fsdp_modules_per_group: int = 1,
) -> None:
"""
Utility to shard a model with FSDP using the PyTorch Distributed fully_shard API.
@@ -250,12 +279,20 @@ def shard_model(
fsdp_shard_conditions (List[Callable[[str, nn.Module], bool]]): A list of functions to determine
which modules to shard with FSDP.
pin_cpu_memory (bool): If set to True, FSDP will pin the CPU memory of the offloaded parameters.
fsdp_modules_per_group (int): Number of consecutive matching modules to place in one FSDP
communication group. The default of 1 preserves per-module sharding.
Raises:
ValueError: If no layer modules were sharded, indicating that no shard_condition was triggered.
ValueError: If the grouping is invalid or no layer modules were sharded, indicating that no
shard_condition was triggered.
"""
if fsdp_modules_per_group < 1:
raise ValueError("fsdp_modules_per_group must be at least 1")
# Check if we should use size-based filtering
use_size_filtering = os.environ.get("FASTVIDEO_FSDP2_AUTOWRAP", "0") == "1"
if use_size_filtering and fsdp_modules_per_group != 1:
raise ValueError("fsdp_modules_per_group > 1 is incompatible with FASTVIDEO_FSDP2_AUTOWRAP=1")
if not fsdp_shard_conditions:
logger.warning("No FSDP shard conditions provided; nothing will be sharded.")
@@ -312,7 +349,7 @@ def shard_model(
module_kwargs = {**fsdp_kwargs, "ignored_params": local_ignored_params}
fully_shard(m, **module_kwargs)
num_layers_sharded += 1
else:
elif fsdp_modules_per_group == 1:
# Shard all modules matching conditions
for n, m in reversed(named_modules):
if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]):
@@ -325,6 +362,33 @@ def shard_model(
if num_layers_sharded == 0:
raise ValueError("No layer modules were sharded. Please check if shard conditions are working as expected.")
else:
matched_modules = [
module for name, module in named_modules
if any(shard_condition(name, module) for shard_condition in fsdp_shard_conditions)
]
if not matched_modules:
raise ValueError("No layer modules were sharded. Please check if shard conditions are working as expected.")
seen_module_params: set[nn.Parameter] = set()
for module in matched_modules:
module_params = set(module.parameters()) - ignored_params
if seen_module_params.intersection(module_params):
raise ValueError("Matched FSDP modules must not have overlapping parameter sets")
seen_module_params.update(module_params)
module_groups = [
matched_modules[start:start + fsdp_modules_per_group]
for start in range(0, len(matched_modules), fsdp_modules_per_group)
]
for group in module_groups:
module_kwargs = fsdp_kwargs
local_ignored_params = set().union(*(ignored_params_by_module[id(module)] for module in group))
if local_ignored_params:
module_kwargs = {**fsdp_kwargs, "ignored_params": local_ignored_params}
fully_shard(group, **module_kwargs)
num_layers_sharded += len(group)
# Finally shard the entire model to account for any stragglers
root_kwargs = fsdp_kwargs
@@ -368,8 +432,29 @@ def load_model_from_full_model_state_dict(
named_parameters = dict(model.named_parameters())
named_buffers = dict(model.named_buffers())
sharded_sd = {}
ignore_checkpoint_key = getattr(model, "_is_ignored_checkpoint_key", None)
ignored_checkpoint_keys = 0
if callable(ignore_checkpoint_key):
unfiltered_sd_iterator = full_sd_iterator
def _filtered_sd_iterator():
nonlocal ignored_checkpoint_keys
for source_param_name, full_tensor in unfiltered_sd_iterator:
target_param_name, _, _ = param_names_mapping( # type: ignore[misc]
source_param_name)
if (target_param_name not in meta_sd
and ignore_checkpoint_key(target_param_name)):
ignored_checkpoint_keys += 1
continue
yield source_param_name, full_tensor
full_sd_iterator = _filtered_sd_iterator()
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(full_sd_iterator,
param_names_mapping) # type: ignore
if ignored_checkpoint_keys:
logger.info("Ignored %d model-declared checkpoint keys", ignored_checkpoint_keys)
for target_param_name, full_tensor in custom_param_sd.items():
meta_sharded_param = meta_sd.get(target_param_name)
if meta_sharded_param is None:
+192 -16
View File
@@ -3,8 +3,8 @@
import contextlib
import re
from collections import defaultdict
from collections.abc import Callable, Iterator
from typing import Any
from collections.abc import Callable, Iterable, Mapping
from typing import Any, TypeAlias
import torch
@@ -12,6 +12,15 @@ from fastvideo.logger import init_logger
logger = init_logger(__name__)
ReverseParamMappingEntry: TypeAlias = (
tuple[str, int | None, int | None]
| tuple[str, int | None, int | None, int]
)
ReverseParamNamesMapping: TypeAlias = dict[
str,
ReverseParamMappingEntry | list[ReverseParamMappingEntry],
]
@contextlib.contextmanager
def set_default_torch_dtype(dtype: torch.dtype):
@@ -25,7 +34,7 @@ def set_default_torch_dtype(dtype: torch.dtype):
def get_param_names_mapping(
mapping_dict: dict[str, str]) -> Callable[[str], tuple[str, Any, Any]]:
mapping_dict: dict[str, Any]) -> Callable[[str], tuple[str, Any, Any]]:
"""
Creates a mapping function that transforms parameter names using regex patterns.
@@ -58,9 +67,9 @@ def get_param_names_mapping(
def hf_to_custom_state_dict(
hf_param_sd: dict[str, torch.Tensor] | Iterator[tuple[str, torch.Tensor]],
param_names_mapping: Callable[[str], tuple[str, Any, Any]]
) -> tuple[dict[str, torch.Tensor], dict[str, tuple[str, Any, Any]]]:
hf_param_sd: Mapping[str, torch.Tensor] | Iterable[tuple[str, torch.Tensor]],
param_names_mapping: Callable[[str], tuple[str, Any, Any]],
) -> tuple[dict[str, torch.Tensor], ReverseParamNamesMapping]:
"""
Converts a Hugging Face parameter state dictionary to a custom parameter state dictionary.
@@ -70,22 +79,67 @@ def hf_to_custom_state_dict(
Returns:
custom_param_sd (Dict[str, torch.Tensor]): The custom formatted parameter state dict
reverse_param_names_mapping (Dict[str, Tuple[str, Any, Any]]): Maps back from custom to hf
reverse_param_names_mapping: Maps direct targets to a source-name
3-tuple and merged targets to an ordered list of source-name,
merge-index, merge-total, and split-size 4-tuples.
"""
custom_param_sd = {}
to_merge_params = defaultdict(dict) # type: ignore
reverse_param_names_mapping = {}
if isinstance(hf_param_sd, dict):
hf_param_sd = hf_param_sd.items() # type: ignore
for source_param_name, full_tensor in hf_param_sd: # type: ignore
custom_param_sd: dict[str, torch.Tensor] = {}
to_merge_params: defaultdict[str, dict[int, torch.Tensor]] = defaultdict(dict)
merge_totals: dict[str, int] = {}
reverse_param_names_mapping: ReverseParamNamesMapping = {}
hf_param_items = (
hf_param_sd.items()
if isinstance(hf_param_sd, Mapping)
else hf_param_sd
)
for source_param_name, full_tensor in hf_param_items:
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
source_param_name)
reverse_param_names_mapping[target_param_name] = (source_param_name,
merge_index,
num_params_to_merge)
if merge_index is not None:
if num_params_to_merge is None:
raise ValueError(f"Missing merge total for {source_param_name!r}")
merge_index = int(merge_index)
num_params_to_merge = int(num_params_to_merge)
if full_tensor.ndim == 0:
raise ValueError(f"Cannot merge scalar parameter {source_param_name!r}")
if num_params_to_merge <= 0 or not 0 <= merge_index < num_params_to_merge:
raise ValueError(
f"Invalid merge metadata for {source_param_name!r}: "
f"index={merge_index}, total={num_params_to_merge}"
)
previous_total = merge_totals.setdefault(target_param_name, num_params_to_merge)
if previous_total != num_params_to_merge:
raise ValueError(
f"Inconsistent merge totals for {target_param_name!r}: "
f"{previous_total} and {num_params_to_merge}"
)
if merge_index in to_merge_params[target_param_name]:
raise ValueError(
f"Duplicate merge index {merge_index} for {target_param_name!r}"
)
reverse_entry = (
source_param_name,
merge_index,
num_params_to_merge,
int(full_tensor.shape[0]),
)
reverse_entries = reverse_param_names_mapping.setdefault(
target_param_name,
[],
)
if not isinstance(reverse_entries, list):
raise ValueError(f"Mixed direct and merged mappings for {target_param_name!r}")
reverse_entries.append(reverse_entry)
to_merge_params[target_param_name][merge_index] = full_tensor
if len(to_merge_params[target_param_name]) == num_params_to_merge:
reverse_entries.sort(key=lambda entry: int(entry[1]))
expected_indices = set(range(num_params_to_merge))
if set(to_merge_params[target_param_name]) != expected_indices:
raise ValueError(
f"Incomplete merge indices for {target_param_name!r}: "
f"got {sorted(to_merge_params[target_param_name])}, "
f"expected {sorted(expected_indices)}"
)
# cat at output dim according to the merge_index order
sorted_tensors = [
to_merge_params[target_param_name][i]
@@ -93,7 +147,129 @@ def hf_to_custom_state_dict(
]
full_tensor = torch.cat(sorted_tensors, dim=0)
del to_merge_params[target_param_name]
del merge_totals[target_param_name]
else:
continue
else:
if target_param_name in reverse_param_names_mapping:
raise ValueError(f"Duplicate direct mapping for {target_param_name!r}")
reverse_param_names_mapping[target_param_name] = (
source_param_name,
None,
None,
)
if target_param_name in custom_param_sd:
raise ValueError(f"Duplicate target parameter {target_param_name!r}")
custom_param_sd[target_param_name] = full_tensor
if to_merge_params:
incomplete = {
name: sorted(parts)
for name, parts in sorted(to_merge_params.items())
}
raise ValueError(f"Incomplete merged parameters: {incomplete}")
return custom_param_sd, reverse_param_names_mapping
def custom_to_hf_state_dict(
state_dict: Mapping[str, Any] | Iterable[tuple[str, Any]],
reverse_param_names_mapping: ReverseParamNamesMapping,
) -> dict[str, Any]:
"""Convert FastVideo parameter names and fused tensors back to HF format."""
if not reverse_param_names_mapping:
raise ValueError("reverse_param_names_mapping is empty")
state = dict(state_dict)
def _entries(
raw: ReverseParamMappingEntry | list[ReverseParamMappingEntry],
) -> list[ReverseParamMappingEntry]:
entries = raw if isinstance(raw, list) else [raw]
if not entries or not all(isinstance(entry, tuple) for entry in entries):
raise ValueError(f"Invalid reverse parameter mapping: {raw!r}")
return entries
def _unpack(
entry: ReverseParamMappingEntry,
) -> tuple[str, int | None, int | None, int | None]:
if len(entry) == 3:
source_key, merge_index, total = entry
return source_key, merge_index, total, None
if len(entry) == 4:
source_key, merge_index, total, split_size = entry
return source_key, merge_index, total, split_size
raise ValueError(f"Invalid reverse parameter mapping entry: {entry!r}")
merge_groups: dict[str, list[tuple[str, int, int, int | None]]] = {}
for training_key, raw_mapping in reverse_param_names_mapping.items():
for entry in _entries(raw_mapping):
source_key, merge_index, merge_total, split_size = _unpack(entry)
if merge_index is None:
continue
if merge_total is None:
raise ValueError(f"Missing merge total for {training_key!r}")
merge_groups.setdefault(training_key, []).append((
source_key,
int(merge_index),
int(merge_total),
None if split_size is None else int(split_size),
))
converted: dict[str, Any] = {}
used_keys: set[str] = set()
for training_key, splits in merge_groups.items():
if training_key not in state:
continue
tensor = state[training_key]
splits.sort(key=lambda entry: entry[1])
total = splits[0][2]
if any(split[2] != total for split in splits):
raise ValueError(f"Inconsistent merge totals for {training_key!r}")
indices = [split[1] for split in splits]
if len(splits) != total or indices != list(range(total)):
raise ValueError(
f"Incomplete reverse merge mapping for {training_key!r}: "
f"indices={indices}, total={total}"
)
recorded_sizes = [split[3] for split in splits]
if all(size is None for size in recorded_sizes):
if tensor.shape[0] % total:
raise ValueError(
f"Cannot evenly split legacy merged parameter {training_key!r} "
f"with output size {tensor.shape[0]} into {total} parts"
)
split_sizes = [tensor.shape[0] // total] * total
elif any(size is None for size in recorded_sizes):
raise ValueError(f"Partially specified split sizes for {training_key!r}")
else:
split_sizes = [int(size) for size in recorded_sizes if size is not None]
if sum(split_sizes) != tensor.shape[0]:
raise ValueError(
f"Recorded split sizes for {training_key!r} sum to "
f"{sum(split_sizes)}, expected {tensor.shape[0]}"
)
split_tensors = torch.split(tensor, split_sizes, dim=0)
for (source_key, _, _, _), split_tensor in zip(
splits,
split_tensors,
strict=True,
):
converted[source_key] = split_tensor
used_keys.add(training_key)
for training_key, value in state.items():
if training_key in used_keys:
continue
if training_key not in reverse_param_names_mapping:
converted[training_key] = value
continue
entries = _entries(reverse_param_names_mapping[training_key])
if len(entries) != 1:
raise ValueError(
f"Invalid direct reverse mapping for {training_key!r}: {entries!r}"
)
source_key, merge_index, _, _ = _unpack(entries[0])
if merge_index is None:
converted[source_key] = value
return converted
@@ -30,7 +30,7 @@ from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.logger import init_logger
from fastvideo.models.dits.ltx2 import (AudioLatentShape, DEFAULT_LTX2_AUDIO_CHANNELS, DEFAULT_LTX2_AUDIO_DOWNSAMPLE,
DEFAULT_LTX2_AUDIO_HOP_LENGTH, DEFAULT_LTX2_AUDIO_MEL_BINS,
DEFAULT_LTX2_AUDIO_SAMPLE_RATE, VideoLatentShape)
DEFAULT_LTX2_AUDIO_SAMPLE_RATE, LTXModelType, VideoLatentShape)
from fastvideo.utils import is_vsa_available
LTX2_AUDIO_CLEAN_LATENT_KEY = "ltx2_audio_clean_latent"
@@ -293,16 +293,22 @@ class LTX2DenoisingStage(PipelineStage):
raise ValueError("LTX-2 i2v timestep mask token count mismatch: "
f"expected {token_count}, got {flat_mask.shape[1]}")
timestep_template = flat_mask
audio_prompt_embeds = batch.extra.get("ltx2_audio_prompt_embeds")
audio_neg_embeds = batch.extra.get("ltx2_audio_negative_embeds")
model_type = getattr(
self.transformer,
"_model_type",
getattr(getattr(self.transformer, "model", None), "model_type", LTXModelType.AudioVideo),
)
audio_enabled = model_type.is_audio_enabled()
audio_prompt_embeds = batch.extra.get("ltx2_audio_prompt_embeds") if audio_enabled else None
audio_neg_embeds = batch.extra.get("ltx2_audio_negative_embeds") if audio_enabled else None
audio_context_p = audio_prompt_embeds[0] if audio_prompt_embeds else None
audio_context_n = audio_neg_embeds[0] if audio_neg_embeds else None
audio_latents = batch.extra.get(self.initial_audio_latents_key)
audio_latents = batch.extra.get(self.initial_audio_latents_key) if audio_enabled else None
if isinstance(audio_latents, torch.Tensor):
audio_latents = audio_latents.to(device=latents.device, dtype=latents.dtype)
# Audio conditioning: mirror video i2v mask approach.
audio_clean_latent = batch.extra.get(LTX2_AUDIO_CLEAN_LATENT_KEY)
audio_denoise_mask = batch.extra.get(LTX2_AUDIO_DENOISE_MASK_KEY)
audio_clean_latent = batch.extra.get(LTX2_AUDIO_CLEAN_LATENT_KEY) if audio_enabled else None
audio_denoise_mask = batch.extra.get(LTX2_AUDIO_DENOISE_MASK_KEY) if audio_enabled else None
if isinstance(audio_clean_latent, torch.Tensor):
audio_clean_latent = audio_clean_latent.to(device=latents.device, dtype=latents.dtype)
if isinstance(audio_denoise_mask, torch.Tensor):
@@ -411,7 +417,7 @@ class LTX2DenoisingStage(PipelineStage):
# conditioning, shift video RoPE positions forward so the
# audio prefix sits at t>=0 and video aligns with the later
# portion of audio.
video_position_offset_sec = float(batch.extra.get("video_position_offset_sec", 0.0))
video_position_offset_sec = (float(batch.extra.get("video_position_offset_sec", 0.0)) if audio_enabled else 0.0)
# Multi-modal CFG parameters (per-stream scales).
modality_scale_video = batch.ltx2_modality_scale_video
@@ -423,12 +429,12 @@ class LTX2DenoisingStage(PipelineStage):
stg_blocks_video = batch.ltx2_stg_blocks_video
stg_blocks_audio = batch.ltx2_stg_blocks_audio
do_stg_video = not math.isclose(float(stg_scale_video), 0.0)
do_stg_audio = not math.isclose(float(stg_scale_audio), 0.0)
do_stg_audio = audio_enabled and not math.isclose(float(stg_scale_audio), 0.0)
do_stg = do_stg_video or do_stg_audio
do_cfg_text = use_cfg and (cfg_scale_video != 1.0 or cfg_scale_audio != 1.0)
do_cfg_text = use_cfg and (cfg_scale_video != 1.0 or (audio_enabled and cfg_scale_audio != 1.0))
do_modality_video = not math.isclose(float(modality_scale_video), 1.0)
do_modality_audio = not math.isclose(float(modality_scale_audio), 1.0)
do_mod = do_modality_video or do_modality_audio
do_mod = audio_enabled and (do_modality_video or do_modality_audio)
do_guidance = do_cfg_text or do_mod or do_stg
if do_cfg_text and neg_prompt_embeds is None:
@@ -478,6 +484,7 @@ class LTX2DenoisingStage(PipelineStage):
# Per-sample sigma for LTX-2.3 cross-attention AdaLN prompt
# timestep. Ignored by LTX-2.0 (prompt_adaln is None).
sigma_batch = sigma.reshape(1).expand(latents.shape[0])
audio_sigma = sigma_batch if audio_enabled else None
timestep = timestep_template * sigma
audio_timestep = (audio_timestep_template * sigma if audio_timestep_template is not None else None)
latent_model_input = latents.to(target_dtype)
@@ -511,7 +518,7 @@ class LTX2DenoisingStage(PipelineStage):
audio_encoder_hidden_states=audio_context_p,
audio_timestep=audio_timestep,
video_sigma=sigma_batch,
audio_sigma=sigma_batch,
audio_sigma=audio_sigma,
video_position_offset_sec=video_position_offset_sec,
)
if isinstance(pos_outputs, tuple):
@@ -540,7 +547,7 @@ class LTX2DenoisingStage(PipelineStage):
audio_encoder_hidden_states=audio_context_n,
audio_timestep=audio_timestep,
video_sigma=sigma_batch,
audio_sigma=sigma_batch,
audio_sigma=audio_sigma,
video_position_offset_sec=video_position_offset_sec,
)
if isinstance(neg_outputs, tuple):
@@ -561,7 +568,7 @@ class LTX2DenoisingStage(PipelineStage):
audio_encoder_hidden_states=audio_context_p,
audio_timestep=audio_timestep,
video_sigma=sigma_batch,
audio_sigma=sigma_batch,
audio_sigma=audio_sigma,
skip_cross_modal_attn=True,
video_position_offset_sec=video_position_offset_sec,
)
@@ -583,7 +590,7 @@ class LTX2DenoisingStage(PipelineStage):
audio_encoder_hidden_states=audio_context_p,
audio_timestep=audio_timestep,
video_sigma=sigma_batch,
audio_sigma=sigma_batch,
audio_sigma=audio_sigma,
skip_video_self_attn_blocks=(stg_blocks_video if do_stg_video else None),
skip_audio_self_attn_blocks=(stg_blocks_audio if do_stg_audio else None),
video_position_offset_sec=video_position_offset_sec,
+123 -4
View File
@@ -21,12 +21,26 @@ from fastvideo.configs.models.dits.ltx2 import LTX2VideoArchConfig
from fastvideo.models.loader.utils import get_param_names_mapping
def _map(name: str, *, apply_gated_attention: bool) -> str:
def _map_with_metadata(
name: str,
*,
apply_gated_attention: bool = False,
pack_attention_projections: bool = False,
) -> tuple[str, int | None, int | None]:
"""Run ``name`` through a fresh config's mapping function."""
cfg = LTX2VideoArchConfig(apply_gated_attention=apply_gated_attention)
cfg = LTX2VideoArchConfig(
apply_gated_attention=apply_gated_attention,
pack_attention_projections=pack_attention_projections,
)
mapper = get_param_names_mapping(cfg.param_names_mapping)
target, _, _ = mapper(name)
return target
return mapper(name)
def _map(name: str, *, apply_gated_attention: bool) -> str:
return _map_with_metadata(
name,
apply_gated_attention=apply_gated_attention,
)[0]
class TestLTX20ParamMappingDefault:
@@ -95,6 +109,111 @@ class TestLTX23ParamMappingGated:
"model.transformer_blocks.0.attn1.to_q.weight")
class TestPackedAttentionProjectionMapping:
@pytest.mark.parametrize("prefix", [
"",
"model.",
"diffusion_model.",
"model.diffusion_model.",
])
@pytest.mark.parametrize("suffix", ["weight", "bias"])
@pytest.mark.parametrize(
"attention,projection,packed,index,total",
[
("attn1", "to_q", "to_qkv", 0, 3),
("attn1", "to_k", "to_qkv", 1, 3),
("attn1", "to_v", "to_qkv", 2, 3),
("attn2", "to_k", "to_kv", 0, 2),
("attn2", "to_v", "to_kv", 1, 2),
],
)
def test_opt_in_maps_split_checkpoint_projections_to_packed_parameters(
self,
prefix,
suffix,
attention,
projection,
packed,
index,
total,
):
source = f"{prefix}transformer_blocks.7.{attention}.{projection}.{suffix}"
assert _map_with_metadata(
source,
pack_attention_projections=True,
) == (
f"model.transformer_blocks.7.{attention}.{packed}.{suffix}",
index,
total,
)
@pytest.mark.parametrize("suffix", ["weight", "bias"])
@pytest.mark.parametrize(
"attention,projection",
[
("attn1", "to_q"),
("attn1", "to_k"),
("attn1", "to_v"),
("attn2", "to_k"),
("attn2", "to_v"),
],
)
def test_default_keeps_checkpoint_projections_split(
self,
suffix,
attention,
projection,
):
source = f"transformer_blocks.7.{attention}.{projection}.{suffix}"
assert _map_with_metadata(source) == (
f"model.{source}",
None,
None,
)
@pytest.mark.parametrize("prefix", [
"",
"model.",
"diffusion_model.",
"model.diffusion_model.",
])
def test_cross_attention_query_stays_separate(self, prefix):
assert _map_with_metadata(
f"{prefix}transformer_blocks.3.attn2.to_q.weight",
pack_attention_projections=True,
) == (
"model.transformer_blocks.3.attn2.to_q.weight",
None,
None,
)
def test_post_init_is_idempotent_and_removes_disabled_pack_rules(self):
cfg = LTX2VideoArchConfig(pack_attention_projections=True)
initial_mapping = tuple(cfg.param_names_mapping.items())
packed_patterns = {
pattern
for pattern, replacement in initial_mapping
if isinstance(replacement, tuple)
}
cfg.__post_init__()
assert tuple(cfg.param_names_mapping.items()) == initial_mapping
cfg.pack_attention_projections = False
cfg.__post_init__()
assert packed_patterns.isdisjoint(cfg.param_names_mapping)
assert get_param_names_mapping(cfg.param_names_mapping)(
"transformer_blocks.0.attn1.to_q.weight"
) == (
"model.transformer_blocks.0.attn1.to_q.weight",
None,
None,
)
class TestParamMappingRuleOrdering:
"""The gate rules must be inserted *before* the generic prefix-strip rules
so first-match-wins matching fires the rename first."""
@@ -56,6 +56,14 @@ def _extract_out(x):
return x
def _strided_input_pair(shape: tuple[int, ...], *, dtype: torch.dtype) -> tuple[torch.Tensor, torch.Tensor]:
source = torch.randn(*shape, 2, device="cuda", dtype=dtype)
return (
source[..., 0].detach().requires_grad_(True),
source.clone()[..., 0].detach().requires_grad_(True),
)
def test_capability_gate_is_lazy_and_device_keyed(monkeypatch) -> None:
if not torch.cuda.is_available():
pytest.skip("CUDA is required to import flash_attn.cute")
@@ -185,6 +193,37 @@ def test_flash_attn_func_parity_forward_backward(
_assert_close(dv_test, dv_ref, dtype=dtype, is_grad=True)
def test_flash_attn_func_torch_compile_forward_backward_parity(flash_attn_impls):
custom_flash_attn_func, _, _, _ = flash_attn_impls
dtype = torch.float16
shape = (1, 32, 2, 64)
torch.manual_seed(2)
q_eager, q_compiled = _strided_input_pair(shape, dtype=dtype)
k_eager, k_compiled = _strided_input_pair(shape, dtype=dtype)
v_eager, v_compiled = _strided_input_pair(shape, dtype=dtype)
def attention(q, k, v):
return custom_flash_attn_func(
q,
k,
v,
dropout_p=0.0,
softmax_scale=None,
causal=False,
deterministic=False,
)
out_eager = attention(q_eager, k_eager, v_eager)
out_compiled = torch.compile(attention, fullgraph=True)(q_compiled, k_compiled, v_compiled)
_assert_close(out_compiled, out_eager, dtype=dtype)
dout = torch.randn_like(out_eager)
eager_grads = torch.autograd.grad((out_eager * dout).sum(), (q_eager, k_eager, v_eager))
compiled_grads = torch.autograd.grad((out_compiled * dout).sum(), (q_compiled, k_compiled, v_compiled))
for compiled_grad, eager_grad in zip(compiled_grads, eager_grads):
_assert_close(compiled_grad, eager_grad, dtype=dtype, is_grad=True)
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
@pytest.mark.parametrize("causal", [False, True])
def test_flash_attn_varlen_func_parity_forward_backward(
@@ -260,3 +299,42 @@ def test_flash_attn_varlen_func_parity_forward_backward(
_assert_close(dq_test, dq_ref, dtype=dtype, is_grad=True)
_assert_close(dk_test, dk_ref, dtype=dtype, is_grad=True)
_assert_close(dv_test, dv_ref, dtype=dtype, is_grad=True)
def test_flash_attn_varlen_func_torch_compile_forward_backward_parity(flash_attn_impls):
_, custom_flash_attn_varlen_func, _, _ = flash_attn_impls
dtype = torch.float16
seqlens = (16, 8)
total = sum(seqlens)
max_seqlen = max(seqlens)
cu_seqlens = torch.tensor((0, seqlens[0], total), device="cuda", dtype=torch.int32)
shape = (total, 2, 64)
torch.manual_seed(3)
q_eager, q_compiled = _strided_input_pair(shape, dtype=dtype)
k_eager, k_compiled = _strided_input_pair(shape, dtype=dtype)
v_eager, v_compiled = _strided_input_pair(shape, dtype=dtype)
def attention(q, k, v, cu):
return custom_flash_attn_varlen_func(
q,
k,
v,
cu,
cu,
max_seqlen,
max_seqlen,
dropout_p=0.0,
softmax_scale=None,
causal=False,
deterministic=False,
)
out_eager = attention(q_eager, k_eager, v_eager, cu_seqlens)
out_compiled = torch.compile(attention, fullgraph=True)(q_compiled, k_compiled, v_compiled, cu_seqlens)
_assert_close(out_compiled, out_eager, dtype=dtype)
dout = torch.randn_like(out_eager)
eager_grads = torch.autograd.grad((out_eager * dout).sum(), (q_eager, k_eager, v_eager))
compiled_grads = torch.autograd.grad((out_compiled * dout).sum(), (q_compiled, k_compiled, v_compiled))
for compiled_grad, eager_grad in zip(compiled_grads, eager_grads):
_assert_close(compiled_grad, eager_grad, dtype=dtype, is_grad=True)
@@ -36,6 +36,12 @@ def test_is_ltx2_blocks_matches_only_top_level_blocks():
assert not is_ltx2_blocks("audio_transformer_blocks.0", None)
def test_fsdp_output_hook_registration_stays_outside_compiled_graph():
register_hooks = ltx2.BasicAVTransformerBlock._register_fsdp_backward_hooks_on_output
assert getattr(register_hooks, "_torchdynamo_disable", False) is True
def test_freq_grid_cache_is_bypassed_while_compiling():
args = (10000.0, 3, 96)
for generator, cached_generator in (
@@ -0,0 +1,139 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import pytest
import torch
from torch import nn
from fastvideo.models.dits.ltx2 import (
_init_attention_projections,
_project_attention_inputs,
)
@pytest.mark.parametrize("context_dim", [None, 7], ids=["self_qkv", "cross_kv"])
def test_packed_attention_projections_match_split_forward_and_backward(
context_dim: int | None) -> None:
query_dim, inner_dim = 5, 3
split, packed = nn.Module(), nn.Module()
for module, pack in ((split, False), (packed, True)):
_init_attention_projections(
module,
query_dim=query_dim,
context_dim=context_dim,
inner_dim=inner_dim,
quant_config=None,
prefix="attn",
pack_attention_projections=pack,
)
module.double()
torch.manual_seed(0)
with torch.no_grad():
for parameter in split.parameters():
parameter.copy_(torch.randn_like(parameter))
if context_dim is None:
packed.to_qkv.weight.copy_(torch.cat(
[split.to_q.weight, split.to_k.weight, split.to_v.weight]))
packed.to_qkv.bias.copy_(torch.cat(
[split.to_q.bias, split.to_k.bias, split.to_v.bias]))
else:
packed.to_q.load_state_dict(split.to_q.state_dict())
packed.to_kv.weight.copy_(torch.cat(
[split.to_k.weight, split.to_v.weight]))
packed.to_kv.bias.copy_(torch.cat(
[split.to_k.bias, split.to_v.bias]))
for module in (split, packed):
for parameter in module.parameters():
parameter.requires_grad_(True)
split_x = torch.randn(2, 4, query_dim, dtype=torch.float64,
requires_grad=True)
packed_x = split_x.detach().clone().requires_grad_(True)
if context_dim is None:
split_context, packed_context = split_x, packed_x
else:
split_context = torch.randn(2,
3,
context_dim,
dtype=torch.float64,
requires_grad=True)
packed_context = split_context.detach().clone().requires_grad_(True)
split_outputs = _project_attention_inputs(split, split_x, split_context)
packed_outputs = _project_attention_inputs(packed, packed_x,
packed_context)
for actual, expected in zip(packed_outputs, split_outputs):
torch.testing.assert_close(actual, expected)
coefficients = (0.5, 1.5, -2.0)
sum(coefficient * output.square().sum()
for coefficient, output in zip(coefficients,
split_outputs)).backward()
sum(coefficient * output.square().sum()
for coefficient, output in zip(coefficients,
packed_outputs)).backward()
torch.testing.assert_close(packed_x.grad, split_x.grad)
if context_dim is not None:
torch.testing.assert_close(packed_context.grad, split_context.grad)
split_weight_grads = [
split.to_q.weight.grad,
split.to_k.weight.grad,
split.to_v.weight.grad,
]
split_bias_grads = [
split.to_q.bias.grad,
split.to_k.bias.grad,
split.to_v.bias.grad,
]
assert set(split.state_dict()) == {
"to_q.weight",
"to_q.bias",
"to_k.weight",
"to_k.bias",
"to_v.weight",
"to_v.bias",
}
if context_dim is None:
packed_weight_grads = packed.to_qkv.weight.grad.chunk(3)
packed_bias_grads = packed.to_qkv.bias.grad.chunk(3)
assert set(packed.state_dict()) == {"to_qkv.weight", "to_qkv.bias"}
else:
torch.testing.assert_close(packed.to_q.weight.grad,
split.to_q.weight.grad)
torch.testing.assert_close(packed.to_q.bias.grad, split.to_q.bias.grad)
packed_weight_grads = (packed.to_q.weight.grad,
*packed.to_kv.weight.grad.chunk(2))
packed_bias_grads = (packed.to_q.bias.grad,
*packed.to_kv.bias.grad.chunk(2))
assert set(packed.state_dict()) == {
"to_q.weight",
"to_q.bias",
"to_kv.weight",
"to_kv.bias",
}
for actual, expected in zip(packed_weight_grads, split_weight_grads):
torch.testing.assert_close(actual, expected)
for actual, expected in zip(packed_bias_grads, split_bias_grads):
torch.testing.assert_close(actual, expected)
assert sum(parameter.numel() for parameter in packed.parameters()) == sum(
parameter.numel() for parameter in split.parameters())
@pytest.mark.parametrize("context_dim", [None, 7])
def test_packed_attention_projections_reject_linear_quantization(
context_dim: int | None) -> None:
module = nn.Module()
with pytest.raises(ValueError, match="do not yet support linear quantization"):
_init_attention_projections(
module,
query_dim=5,
context_dim=context_dim,
inner_dim=3,
quant_config=object(), # type: ignore[arg-type]
prefix="attn",
pack_attention_projections=True,
)
assert not tuple(module.parameters())
@@ -0,0 +1,89 @@
# SPDX-License-Identifier: Apache-2.0
import pytest
import torch
from fastvideo.models.dits.ltx2 import (
EntryClass,
LTXLocalAttention,
LTX2Transformer3DModel,
LTX2VideoOnlyTransformer3DModel,
LTXModel,
LTXModelType,
)
from fastvideo.models.registry import ModelRegistry
from fastvideo.platforms import AttentionBackendEnum
def _state_keys(model_type: LTXModelType, *, ltx2_3: bool) -> set[str]:
with torch.device("meta"):
model = LTXModel(
model_type=model_type,
num_attention_heads=2,
attention_head_dim=8,
in_channels=8,
out_channels=8,
num_layers=2,
cross_attention_dim=16,
caption_channels=16,
audio_num_attention_heads=2,
audio_attention_head_dim=4,
audio_in_channels=8,
audio_out_channels=8,
audio_cross_attention_dim=16,
cross_attention_adaln=ltx2_3,
caption_proj_before_connector=ltx2_3,
apply_gated_attention=ltx2_3,
)
return {f"model.{key}" for key in model.state_dict()}
def test_video_only_class_registration_keeps_av_default() -> None:
assert LTX2Transformer3DModel._model_type is LTXModelType.AudioVideo
assert LTX2VideoOnlyTransformer3DModel._model_type is LTXModelType.VideoOnly
assert EntryClass == [
LTX2Transformer3DModel,
LTX2VideoOnlyTransformer3DModel,
]
model_cls, _ = ModelRegistry.resolve_model_cls(
"LTX2VideoOnlyTransformer3DModel")
assert model_cls is LTX2VideoOnlyTransformer3DModel
@pytest.mark.parametrize("ltx2_3", [False, True])
def test_video_only_checkpoint_filter_matches_exact_av_state_difference(
ltx2_3: bool, monkeypatch: pytest.MonkeyPatch) -> None:
def init_parameterless_attention(self, *args, **kwargs):
del args, kwargs
torch.nn.Module.__init__(self)
self.backend = AttentionBackendEnum.TORCH_SDPA
monkeypatch.setattr(LTXLocalAttention, "__init__",
init_parameterless_attention)
av_keys = _state_keys(LTXModelType.AudioVideo, ltx2_3=ltx2_3)
video_keys = _state_keys(LTXModelType.VideoOnly, ltx2_3=ltx2_3)
removed_keys = av_keys - video_keys
assert removed_keys
assert all(LTX2VideoOnlyTransformer3DModel._is_ignored_checkpoint_key(key)
for key in removed_keys)
assert not any(
LTX2VideoOnlyTransformer3DModel._is_ignored_checkpoint_key(key)
for key in video_keys)
@pytest.mark.parametrize(
"key",
[
"model",
"model.audio_typo.weight",
"model.patchify_proj.weight",
"model.transformer_blocks.0.attn1.to_q.weight",
"model.transformer_blocks.0.audio_attn3.to_q.weight",
"model.transformer_blocks.bad.audio_attn1.to_q.weight",
"unrelated.weight",
],
)
def test_video_only_checkpoint_filter_rejects_near_misses(key: str) -> None:
assert not LTX2VideoOnlyTransformer3DModel._is_ignored_checkpoint_key(key)
+41 -7
View File
@@ -41,7 +41,10 @@ def _seed_everything(seed: int) -> None:
torch.backends.cudnn.allow_tf32 = False
def _build_tiny_ltx2_config() -> LTX2VideoConfig:
def _build_tiny_ltx2_config(
*,
pack_attention_projections: bool = False,
) -> LTX2VideoConfig:
arch_config = LTX2VideoArchConfig(
num_attention_heads=4,
attention_head_dim=8,
@@ -69,6 +72,7 @@ def _build_tiny_ltx2_config() -> LTX2VideoConfig:
audio_cross_attention_dim=16,
audio_positional_embedding_max_pos=[8],
av_ca_timestep_scale_multiplier=1,
pack_attention_projections=pack_attention_projections,
)
return LTX2VideoConfig(arch_config=arch_config)
@@ -150,7 +154,12 @@ def _assert_finite(name: str, tensor: torch.Tensor) -> None:
)
def _run_worker(mode: str, output_path: Path) -> None:
def _run_worker(
mode: str,
output_path: Path,
*,
pack_attention_projections: bool,
) -> None:
if mode not in {"single", "sp"}:
raise ValueError(f"Unsupported mode: {mode}")
@@ -165,7 +174,9 @@ def _run_worker(mode: str, output_path: Path) -> None:
try:
maybe_init_distributed_environment_and_model_parallel(1, sp_size)
config = _build_tiny_ltx2_config()
config = _build_tiny_ltx2_config(
pack_attention_projections=pack_attention_projections,
)
model = LTX2Transformer3DModel(config=config, hf_config={})
model = model.to(device=device, dtype=torch.float32)
_initialize_model_parameters(model)
@@ -215,6 +226,8 @@ def _run_torchrun(
mode: str,
nproc_per_node: int,
output_path: Path,
*,
pack_attention_projections: bool,
) -> None:
cmd = [
"torchrun",
@@ -231,7 +244,16 @@ def _run_torchrun(
"--output",
str(output_path),
]
if pack_attention_projections:
cmd.append("--pack-attention-projections")
env = os.environ.copy()
repo_root = str(script_path.parents[3])
inherited_pythonpath = env.get("PYTHONPATH")
env["PYTHONPATH"] = (
repo_root
if not inherited_pythonpath
else os.pathsep.join((repo_root, inherited_pythonpath))
)
env["FASTVIDEO_ATTENTION_BACKEND"] = "TORCH_SDPA"
process = subprocess.run(cmd, capture_output=True, text=True, env=env)
if process.returncode != 0:
@@ -242,7 +264,11 @@ def _run_torchrun(
)
def test_sp_gradient_matches_single_rank(tmp_path: Path) -> None:
@pytest.mark.parametrize("pack_attention_projections", [False, True])
def test_sp_gradient_matches_single_rank(
tmp_path: Path,
pack_attention_projections: bool,
) -> None:
if not torch.cuda.is_available():
pytest.skip("This test requires CUDA.")
if torch.cuda.device_count() < SP_WORLD_SIZE:
@@ -251,20 +277,23 @@ def test_sp_gradient_matches_single_rank(tmp_path: Path) -> None:
)
script_path = Path(__file__).resolve()
single_path = tmp_path / "single_rank_grads.pt"
sp_path = tmp_path / f"sp{SP_WORLD_SIZE}_grads.pt"
layout = "packed" if pack_attention_projections else "split"
single_path = tmp_path / f"single_rank_{layout}_grads.pt"
sp_path = tmp_path / f"sp{SP_WORLD_SIZE}_{layout}_grads.pt"
_run_torchrun(
script_path=script_path,
mode="single",
nproc_per_node=1,
output_path=single_path,
pack_attention_projections=pack_attention_projections,
)
_run_torchrun(
script_path=script_path,
mode="sp",
nproc_per_node=SP_WORLD_SIZE,
output_path=sp_path,
pack_attention_projections=pack_attention_projections,
)
single_grads: dict[str, torch.Tensor] = torch.load(
@@ -311,6 +340,7 @@ def _parse_args() -> argparse.Namespace:
parser.add_argument("--sp-grad-worker", action="store_true")
parser.add_argument("--mode", choices=["single", "sp"], default=None)
parser.add_argument("--output", type=str, default=None)
parser.add_argument("--pack-attention-projections", action="store_true")
return parser.parse_args()
@@ -321,4 +351,8 @@ if __name__ == "__main__":
raise SystemExit("This module is intended to be run by pytest.")
if args.mode is None or args.output is None:
raise SystemExit("--mode and --output are required in worker mode.")
_run_worker(mode=args.mode, output_path=Path(args.output))
_run_worker(
mode=args.mode,
output_path=Path(args.output),
pack_attention_projections=args.pack_attention_projections,
)
@@ -89,7 +89,6 @@ def run_training(case: dict):
"-m", "fastvideo.train.entrypoint.train",
"--config", case["config"],
"--training.checkpoint.output_dir", str(case["out_dir"]),
"--training.data.data_path", str(case["prep_dir"]),
"--callbacks.validation.dataset_file",
str(case["prep_dir"] / "validation_prompts.json"),
]
@@ -0,0 +1,104 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.models.dits.ltx2 import (
LTX2Transformer3DModel,
LTX2VideoOnlyTransformer3DModel,
)
from fastvideo.pipelines.basic.ltx2.stages.ltx2_audio_decoding import (
LTX2AudioDecodingStage,
)
from fastvideo.pipelines.basic.ltx2.stages.ltx2_denoising import (
LTX2DenoisingStage,
)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
class _RecordingVideoOnlyTransformer(LTX2VideoOnlyTransformer3DModel):
def __init__(self) -> None:
torch.nn.Module.__init__(self)
self.calls = []
def forward(self, hidden_states, **kwargs):
self.calls.append(kwargs)
return torch.zeros_like(hidden_states)
class _RecordingAudioVideoTransformer(LTX2Transformer3DModel):
def __init__(self) -> None:
torch.nn.Module.__init__(self)
self.calls = []
def forward(self, hidden_states, **kwargs):
self.calls.append(kwargs)
return torch.zeros_like(hidden_states), torch.zeros_like(kwargs["audio_hidden_states"])
def _batch() -> ForwardBatch:
text = torch.ones(1, 1, 4)
return ForwardBatch(
data_type="video",
latents=torch.ones(1, 1, 1, 1, 1),
prompt_embeds=[text],
negative_prompt_embeds=[-text],
num_inference_steps=1,
num_frames=1,
fps=24,
extra={
"ltx2_audio_prompt_embeds": [text],
"ltx2_audio_negative_embeds": [-text],
"ltx2_audio_latents": torch.ones(1, 2, 1, 1),
"video_position_offset_sec": 2.0,
},
)
def test_video_only_denoising_ignores_audio_conditioning_and_decoding() -> None:
transformer = _RecordingVideoOnlyTransformer()
batch = _batch()
batch.ltx2_cfg_scale_audio = 7.0
batch.ltx2_modality_scale_audio = 3.0
batch.ltx2_stg_scale_audio = 1.0
batch.do_classifier_free_guidance = True
args = FastVideoArgs(model_path="", disable_autocast=True)
result = LTX2DenoisingStage(
transformer,
sigmas_override=[1.0, 0.0],
).forward(batch, args)
assert len(transformer.calls) == 1
call = transformer.calls[0]
assert call["audio_hidden_states"] is None
assert call["audio_encoder_hidden_states"] is None
assert call["audio_timestep"] is None
assert call["audio_sigma"] is None
assert call["video_position_offset_sec"] == 0.0
assert result.extra["ltx2_audio_latents"] is None
decoded = LTX2AudioDecodingStage(object(), object()).forward(result, args)
assert decoded is result
assert "audio" not in decoded.extra
def test_audio_video_denoising_keeps_audio_conditioning() -> None:
transformer = _RecordingAudioVideoTransformer()
batch = _batch()
result = LTX2DenoisingStage(
transformer,
sigmas_override=[1.0, 0.0],
).forward(batch, FastVideoArgs(model_path="", disable_autocast=True))
assert len(transformer.calls) == 1
call = transformer.calls[0]
assert call["audio_hidden_states"] is not None
assert call["audio_encoder_hidden_states"] is not None
assert call["audio_timestep"] is not None
assert call["audio_sigma"] is not None
assert call["video_position_offset_sec"] == 2.0
assert result.extra["ltx2_audio_latents"] is not None
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU-only unit tests for :mod:`fastvideo.train.callbacks.grad_clip`.
Exercises ``GradNormClipCallback.on_before_optimizer_step`` against
Exercises deferred gradient-norm logging after the optimizer boundary using
synthetic ``nn.Module`` targets with manually populated gradients.
"""
from __future__ import annotations
@@ -118,10 +118,10 @@ class TestGradNormClipCallback:
cb = GradNormClipCallback(
max_grad_norm=1.0, log_grad_norms=True
)
cb.on_before_optimizer_step(
method=_Method(targets={"layer": m}, tracker=tracker),
iteration=7,
)
method = _Method(targets={"layer": m}, tracker=tracker)
cb.on_before_optimizer_step(method=method, iteration=7)
assert tracker.entries == []
cb.on_training_step_end(method=method, loss_dict={}, iteration=7)
assert len(tracker.entries) == 1
payload, step = tracker.entries[0]
assert step == 7
@@ -138,6 +138,11 @@ class TestGradNormClipCallback:
method=_Method(targets={"m": m}, tracker=tracker),
iteration=0,
)
cb.on_training_step_end(
method=_Method(targets={"m": m}, tracker=tracker),
loss_dict={},
iteration=0,
)
assert tracker.entries == []
def test_no_tracker_does_not_raise(self) -> None:
@@ -152,7 +157,9 @@ class TestGradNormClipCallback:
) -> dict[str, torch.nn.Module]:
return {"m": m}
cb.on_before_optimizer_step(method=_BareMethod(), iteration=0)
method = _BareMethod()
cb.on_before_optimizer_step(method=method, iteration=0)
cb.on_training_step_end(method=method, loss_dict={}, iteration=0)
# No assertion — must simply not raise.
def test_multiple_targets_each_logged(self) -> None:
@@ -162,9 +169,48 @@ class TestGradNormClipCallback:
}
tracker = _RecordingTracker()
cb = GradNormClipCallback(max_grad_norm=1.0)
cb.on_before_optimizer_step(
method=_Method(targets=targets, tracker=tracker),
iteration=1,
)
method = _Method(targets=targets, tracker=tracker)
cb.on_before_optimizer_step(method=method, iteration=1)
assert tracker.entries == []
cb.on_training_step_end(method=method, loss_dict={}, iteration=1)
keys = {next(iter(p)) for p, _ in tracker.entries}
assert keys == {"grad_norm/head", "grad_norm/tail"}
def test_norm_materialized_only_after_optimizer_boundary(
self, monkeypatch
) -> None:
events: list[str] = []
class _DeferredNorm:
def item(self) -> float:
events.append("item")
return 2.0
class _EventTracker:
def log(self, payload: dict[str, Any], step: int) -> None:
assert payload == {"grad_norm/layer": 2.0}
assert step == 3
events.append("log")
def _fake_clip(*_args, **_kwargs):
events.append("clip")
return _DeferredNorm()
monkeypatch.setattr(
"fastvideo.train.callbacks.grad_clip.clip_grad_norm_if_needed",
_fake_clip,
)
method = _Method(
targets={"layer": _make_module(grad_value=1.0)},
tracker=_EventTracker(),
)
cb = GradNormClipCallback(max_grad_norm=1.0)
cb.on_before_optimizer_step(method=method, iteration=3)
events.append("optimizer")
events.append("zero_grad")
cb.on_training_step_end(method=method, loss_dict={}, iteration=3)
assert events == ["clip", "optimizer", "zero_grad", "item", "log"]
@@ -14,6 +14,7 @@ distributed init and is exercised by Phase 2/3 tests.
"""
from __future__ import annotations
import contextlib
from types import SimpleNamespace
import numpy as np
@@ -29,6 +30,10 @@ from fastvideo.train.callbacks.validation import (
ValidationCallback,
_ValidationMetricStats,
)
from fastvideo.train.utils.training_config import (
ModelTrainingConfig,
TrainingConfig,
)
# ---------------------------------------------------------------------------
@@ -219,6 +224,51 @@ class TestOnValidationBegin:
assert cb.run_calls == [0]
@pytest.mark.parametrize("compile_enabled", [False, True])
def test_validation_forces_compiled_transformer_eager(
monkeypatch: pytest.MonkeyPatch,
compile_enabled: bool,
) -> None:
cb = _make_callback()
cb.training_config = TrainingConfig(
model=ModelTrainingConfig(enable_torch_compile=compile_enabled),
)
transformer = torch.nn.Linear(1, 1)
method = SimpleNamespace(student=SimpleNamespace(transformer=transformer))
events: list[tuple[str, object]] = []
@contextlib.contextmanager
def recording_ema(_transformer: torch.nn.Module):
events.append(("ema_enter", _transformer))
yield _transformer
events.append(("ema_exit", _transformer))
@contextlib.contextmanager
def recording_stance(stance: str):
events.append(("enter", stance))
yield
events.append(("exit", stance))
monkeypatch.setattr(torch.compiler, "set_stance", recording_stance)
monkeypatch.setattr(
cb,
"_find_ema_callback",
lambda: SimpleNamespace(ema_context=recording_ema),
)
monkeypatch.setattr(
cb,
"_run_validation_inner",
lambda _method, _step, model: events.append(("inner", model)),
)
cb._run_validation(method, 0)
expected = [("ema_enter", transformer), ("inner", transformer), ("ema_exit", transformer)]
if compile_enabled:
expected = [("enter", "force_eager"), *expected, ("exit", "force_eager")]
assert events == expected
# ---------------------------------------------------------------------------
# C. _find_ema_callback
# ---------------------------------------------------------------------------
@@ -0,0 +1,78 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import json
from types import SimpleNamespace
import torch
from safetensors.torch import load_file, save_file
from fastvideo.train.entrypoint.dcp_to_diffusers import _save_role_pretrained
def test_save_role_pretrained_splits_merged_parameters(tmp_path) -> None:
base = tmp_path / "base"
transformer_dir = base / "transformer"
transformer_dir.mkdir(parents=True)
(base / "model_index.json").write_text("{}", encoding="utf-8")
(transformer_dir / "config.json").write_text(
json.dumps({"_class_name": "FakeTransformer"}),
encoding="utf-8",
)
save_file(
{"stale.weight": torch.zeros(1)},
str(transformer_dir / "stale.safetensors"),
)
transformer = torch.nn.Module()
transformer.to_qkv = torch.nn.Linear(4, 6, bias=True)
with torch.no_grad():
transformer.to_qkv.weight.copy_(torch.arange(24).reshape(6, 4))
transformer.to_qkv.bias.copy_(torch.arange(6))
transformer.reverse_param_names_mapping = {
"to_qkv.weight": [
("to_q.weight", 0, 3, 1),
("to_k.weight", 1, 3, 2),
("to_v.weight", 2, 3, 3),
],
"to_qkv.bias": [
("to_q.bias", 0, 3, 1),
("to_k.bias", 1, 3, 2),
("to_v.bias", 2, 3, 3),
],
}
output = tmp_path / "export"
_save_role_pretrained(
role="student",
base_model_path=str(base),
output_dir=str(output),
model=SimpleNamespace(transformer=transformer),
)
exported = load_file(str(output / "transformer" / "model.safetensors"))
assert set(exported) == {
"to_q.weight",
"to_q.bias",
"to_k.weight",
"to_k.bias",
"to_v.weight",
"to_v.bias",
}
torch.testing.assert_close(
exported["to_q.weight"],
transformer.to_qkv.weight[:1],
)
torch.testing.assert_close(
exported["to_k.weight"],
transformer.to_qkv.weight[1:3],
)
torch.testing.assert_close(
exported["to_v.weight"],
transformer.to_qkv.weight[3:],
)
torch.testing.assert_close(exported["to_q.bias"], transformer.to_qkv.bias[:1])
torch.testing.assert_close(exported["to_k.bias"], transformer.to_qkv.bias[1:3])
torch.testing.assert_close(exported["to_v.bias"], transformer.to_qkv.bias[3:])
assert not (output / "transformer" / "stale.safetensors").exists()
@@ -5,7 +5,7 @@ Mirrors ``test_wan_finetune.py`` for the LTX-2 plugin, parametrized
over LTX-2.0 and LTX-2.3 checkpoints. LTX2-specific differences: the
synthetic ``raw_batch`` carries a single post-connector Gemma
embedding (3840-d for 2.0, 4096-d for 2.3 which has no in-DiT caption
projection) and 128-channel VAE latents; the DiT is 18.9B params, so
projection) and 128-channel VAE latents; the video-only DiT is 13B params, so
the test skips on GPUs with less than 60GB memory (e.g. the L40S CI
runner).
"""
@@ -24,6 +24,7 @@ import torch
from fastvideo.train.methods.fine_tuning.finetune import (
FineTuneMethod, )
from fastvideo.models.dits.ltx2 import LTX2VideoOnlyTransformer3DModel
from fastvideo.train.models.ltx2 import LTX2Model
from fastvideo.train.utils.config import load_run_config
@@ -144,15 +145,13 @@ def test_ltx2_finetune_single_train_step(
"all layer-0 grads are exactly zero; backward did not "
"reach the first transformer block")
# Audio / cross-modal parameters must be frozen by default.
audio_trainable = [
name for name, param in model.transformer.named_parameters()
if param.requires_grad and any(
pattern in name for pattern in ("audio", "a2v", "v2a", "av_ca"))
# Audio / cross-modal parameters are not instantiated for training.
audio_parameters = [
name for name, _ in model.transformer.named_parameters()
if LTX2VideoOnlyTransformer3DModel._is_ignored_checkpoint_key(name)
]
assert not audio_trainable, (
f"audio/cross-modal params unexpectedly trainable: "
f"{audio_trainable[:5]}")
assert not audio_parameters, (
f"audio/cross-modal params unexpectedly present: {audio_parameters[:5]}")
# Device-keyed grad-norm regression on top of the same harness.
# Skips when the current GPU has no seeded reference.
+102 -23
View File
@@ -1,18 +1,18 @@
# SPDX-License-Identifier: Apache-2.0
"""GPU loading + forward smoke test for ``LTX2Model``.
Loads the real LTX-2.0 / LTX-2.3 distilled checkpoints (18.9B at bf16,
~38GB — skips on GPUs with less than 60GB memory) via
Loads the real LTX-2.0 / LTX-2.3 distilled checkpoints into the 13B
video-only training transformer (skips on GPUs with less than 60GB memory) via
``LTX2Model.__init__`` and runs one transformer forward pass on
synthetic inputs. Catches loader or forward-signature regressions in
``fastvideo.train.models.ltx2.LTX2Model`` and the underlying
``LTX2Transformer3DModel``.
``LTX2VideoOnlyTransformer3DModel``.
LTX-2's transformer takes per-token sigma timesteps in [0, 1] shaped
[B, tokens] and a post-connector Gemma text embedding (3840-d for 2.0;
4096-d for 2.3, which has no in-DiT caption projection); it returns
the denoised x0 prediction. This mirrors the kwargs in
``LTX2Model._build_distill_input_kwargs``.
LTX-2's transformer takes sigma timesteps in [0, 1]; uniform T2V sigmas use
[B, 1] at SP=1 and expand to [B, tokens] for sequence sharding. It also takes
a post-connector Gemma text embedding (3840-d for 2.0; 4096-d for 2.3, which
has no in-DiT caption projection) and returns the denoised x0 prediction. This
mirrors the kwargs in ``LTX2Model._build_distill_input_kwargs``.
"""
from __future__ import annotations
@@ -25,6 +25,7 @@ os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29519")
from pathlib import Path
from types import SimpleNamespace
import pytest
import torch
@@ -33,8 +34,11 @@ from fastvideo.forward_context import (
get_forward_context,
set_forward_context,
)
from fastvideo.pipelines import ForwardBatch
from fastvideo.models.dits.ltx2 import LTXModelType
from fastvideo.pipelines import ForwardBatch, TrainingBatch
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.train.models.ltx2 import LTX2Model
from fastvideo.train.models.ltx2 import ltx2 as ltx2_module
from fastvideo.train.utils.config import load_run_config
_FIXTURE_DIR = Path(__file__).resolve().parent.parent / "fixtures"
@@ -56,10 +60,6 @@ class _TinyLTX2Transformer(torch.nn.Module):
def __init__(self, *, is_ltx2_3: bool) -> None:
super().__init__()
self.video_weight = torch.nn.Parameter(torch.tensor(1.0))
self.audio_weight = torch.nn.Parameter(torch.tensor(1.0))
self.a2v_weight = torch.nn.Parameter(torch.tensor(1.0))
self.v2a_weight = torch.nn.Parameter(torch.tensor(1.0))
self.av_ca_weight = torch.nn.Parameter(torch.tensor(1.0))
self.config = type("Config", (), {
"arch_config": type("Arch", (), {
"caption_proj_before_connector": is_ltx2_3,
@@ -109,6 +109,62 @@ def test_ltx2_rejects_audio_training_before_loading(
)
def test_ltx2_selects_video_only_transformer_and_forwards_backend(
monkeypatch: pytest.MonkeyPatch) -> None:
cfg = load_run_config(str(_FIXTURE_DIR / _CASES["ltx2"][0]))
transformer = _TinyLTX2Transformer(is_ltx2_3=False)
captured: dict[str, object] = {}
def fake_load_module_from_path(**kwargs):
captured.update(kwargs)
return transformer
monkeypatch.setattr(ltx2_module, "load_module_from_path",
fake_load_module_from_path)
model = LTX2Model(
init_from=cfg.models["student"]["init_from"],
training_config=cfg.training,
trainable=False,
attention_backend="TORCH_SDPA",
)
assert model.transformer is transformer
assert captured["override_transformer_cls_name"] == (
"LTX2VideoOnlyTransformer3DModel")
assert captured["attention_backend"] is AttentionBackendEnum.TORCH_SDPA
def test_ltx2_projection_packing_config_rejects_lora_before_loading(
monkeypatch: pytest.MonkeyPatch) -> None:
cfg = load_run_config(
str(_FIXTURE_DIR / _CASES["ltx2"][0]),
overrides=[
"--pipeline.dit_config.pack_attention_projections",
"true",
],
)
assert cfg.training.pipeline_config.dit_config.arch_config.pack_attention_projections
def fail_if_loaded(**_kwargs):
raise AssertionError("transformer loading must not start")
monkeypatch.setattr(ltx2_module, "load_module_from_path", fail_if_loaded)
with pytest.raises(ValueError, match="do not yet support LoRA training"):
LTX2Model(
init_from=cfg.models["student"]["init_from"],
training_config=cfg.training,
lora={"enable": True, "rank": 4},
)
def test_ltx2_rejects_role_local_sparse_attention_backend() -> None:
model = object.__new__(LTX2Model)
model.attention_backend = AttentionBackendEnum.VIDEO_SPARSE_ATTN
with pytest.raises(NotImplementedError, match="does not support VSA/VMOBA"):
model._build_attention_metadata(TrainingBatch())
@pytest.mark.parametrize("case", _CASES.keys())
def test_ltx2_wrapper_contract_runs_on_cpu(
case: str, monkeypatch: pytest.MonkeyPatch) -> None:
@@ -133,9 +189,6 @@ def test_ltx2_wrapper_contract_runs_on_cpu(
)
assert transformer.video_weight.requires_grad
for name, param in transformer.named_parameters():
if any(pattern in name for pattern in ("audio", "a2v", "v2a", "av_ca")):
assert not param.requires_grad, f"{name} was not frozen"
model._check_text_embedding_dim(torch.empty(1, 1, text_dim))
with pytest.raises(ValueError, match="text_embedding width"):
@@ -182,7 +235,7 @@ def test_ltx2_wrapper_contract_runs_on_cpu(
assert transformer.forward_fps == 12.0
assert transformer.forward_text_dim == text_dim
assert transformer.forward_timestep is not None
assert transformer.forward_timestep.shape == (1, 8)
assert transformer.forward_timestep.shape == (1, 1)
assert torch.allclose(
transformer.forward_timestep[:, 0],
batch.timesteps / 1000.0,
@@ -207,12 +260,38 @@ def test_ltx2_wrapper_contract_runs_on_cpu(
assert transformer.video_weight.grad.item() == pytest.approx(1.0)
@pytest.mark.parametrize(("sp_size", "expected_tokens"), [(1, 1), (2, 8)])
def test_ltx2_uniform_timestep_shape_follows_sequence_parallelism(
sp_size: int,
expected_tokens: int,
) -> None:
model = object.__new__(LTX2Model)
model.training_config = SimpleNamespace(
distributed=SimpleNamespace(sp_size=sp_size),
)
model._token_count = 8
timesteps = torch.tensor([250.0, 750.0])
kwargs = model._build_distill_input_kwargs(
torch.empty(2, 128, 2, 2, 2),
timesteps,
{"encoder_hidden_states": torch.empty(2, 4, 3840)},
)
assert kwargs["timestep"].shape == (2, expected_tokens)
assert torch.equal(
kwargs["timestep"],
torch.tensor([[0.25], [0.75]]).expand(2, expected_tokens),
)
assert kwargs["timestep"].is_contiguous()
@pytest.mark.usefixtures("distributed_setup")
@pytest.mark.parametrize("case", _CASES.keys())
def test_ltx2_model_loads_and_forwards(case: str):
if _gpu_too_small():
pytest.skip(f"requires a CUDA GPU with >= {_MIN_GPU_MEMORY_GB}GB "
"memory (LTX-2 DiT is 18.9B params)")
"memory (LTX-2 video DiT is 13B params)")
fixture_name, text_dim = _CASES[case]
cfg = load_run_config(str(_FIXTURE_DIR / fixture_name))
@@ -225,24 +304,24 @@ def test_ltx2_model_loads_and_forwards(case: str):
transformer = model.transformer
assert isinstance(transformer, torch.nn.Module)
assert sum(p.numel() for p in transformer.parameters()) > 0
assert transformer.model.model_type is LTXModelType.VideoOnly
assert not hasattr(transformer.model, "audio_patchify_proj")
device = torch.device("cuda:0")
dtype = torch.bfloat16
transformer = transformer.to(device=device, dtype=dtype).eval()
# LTX-2 transformer takes [B, 128, T, H, W] latents, a post-connector
# Gemma embedding, and PER-TOKEN sigmas in [0, 1] ([B, T*H*W] with
# patch size 1x1x1). Small spatial + few frames so this fits next to
# the 18.9B model.
# Gemma embedding, and a uniform per-sample T2V sigma in [0, 1]. Small
# spatial + few frames keep the activation footprint below the 13B model.
b, c, t, h, w = 1, 128, 3, 8, 8
tokens = t * h * w
hidden_states = torch.randn(b, c, t, h, w, device=device, dtype=dtype)
encoder_hidden_states = torch.randn(b,
_LTX2_TEXT_LEN,
text_dim,
device=device,
dtype=dtype)
timestep = torch.full((b, tokens), 0.5, device=device, dtype=torch.float32)
timestep = torch.full((b, 1), 0.5, device=device, dtype=torch.float32)
with torch.no_grad(), torch.autocast(device.type, dtype=dtype), \
set_forward_context(
@@ -9,6 +9,7 @@ from typing import Any
import torch
from fastvideo.train.callbacks.validation import ValidationCallback
from fastvideo.train.methods import base as method_base
from fastvideo.train.trainer import Trainer
from fastvideo.train.utils.training_config import TrainingConfig
@@ -45,6 +46,7 @@ class _DummyMethod:
self.zero_grad_steps: list[int] = []
self.optimizer_steps: list[int] = []
self.backward_calls = 0
self.gradient_sync_calls: list[bool] = []
self.tracker = None
def set_tracker(self, tracker: Any) -> None:
@@ -76,6 +78,9 @@ class _DummyMethod:
self.backward_calls += 1
(loss_map["total_loss"] / grad_accum_rounds).backward()
def set_requires_gradient_sync(self, requires_gradient_sync: bool) -> None:
self.gradient_sync_calls.append(requires_gradient_sync)
def optimizers_schedulers_step(self, iteration: int) -> None:
self.optimizer_steps.append(iteration)
@@ -133,7 +138,95 @@ def test_trainer_runs_validation_callback_during_training(
assert validation.run_calls == [0, 2]
assert method.train_start_calls == 1
assert method.backward_calls == 3
assert method.gradient_sync_calls == []
assert method.zero_grad_steps == [0, 1, 2, 3]
assert method.optimizer_steps == [1, 2, 3]
assert [step for _, step in tracker.logs] == [1, 2, 3]
assert tracker.finished is True
def test_trainer_syncs_fsdp_gradients_only_on_final_accumulation(
monkeypatch,
) -> None:
tracker = _DummyTracker()
group = SimpleNamespace(rank=0, local_rank=0, rank_in_group=0, world_size=1)
monkeypatch.setattr("fastvideo.train.trainer.get_world_group", lambda: group)
monkeypatch.setattr("fastvideo.train.trainer.get_sp_group", lambda: group)
monkeypatch.setattr(
"fastvideo.train.trainer.build_tracker",
lambda *args, **kwargs: tracker,
)
cfg = TrainingConfig()
cfg.tracker.project_name = ""
cfg.loop.gradient_accumulation_steps = 3
trainer = Trainer(cfg)
method = _DummyMethod()
trainer.run(
method,
dataloader=[{"sample": "x"}],
max_steps=1,
)
assert method.backward_calls == 3
assert method.gradient_sync_calls == [False, False, True]
assert method.optimizer_steps == [1]
def test_training_method_forwards_gradient_sync_to_trainable_fsdp_roots(
monkeypatch,
) -> None:
class _FakeFSDP(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.calls: list[tuple[str, bool, bool | None]] = []
def set_requires_gradient_sync(self, enabled: bool, *, recurse: bool) -> None:
self.calls.append(("sync", enabled, recurse))
def set_reshard_after_backward(self, enabled: bool, *, recurse: bool) -> None:
self.calls.append(("reshard", enabled, recurse))
def set_is_last_backward(self, enabled: bool) -> None:
self.calls.append(("last", enabled, None))
monkeypatch.setattr(method_base, "FSDPModule", _FakeFSDP)
trainable = _FakeFSDP()
frozen = _FakeFSDP()
owner = SimpleNamespace(
training_config=SimpleNamespace(distributed=SimpleNamespace(reshard_after_forward=False)),
_role_models={
"student": SimpleNamespace(transformer=trainable, _trainable=True),
"teacher": SimpleNamespace(transformer=frozen, _trainable=False),
}
)
method_base.TrainingMethod.set_requires_gradient_sync(owner, False)
method_base.TrainingMethod.set_requires_gradient_sync(owner, True)
assert trainable.calls == [
("sync", False, True),
("reshard", False, True),
("last", False, None),
("sync", True, True),
("reshard", True, True),
("last", True, None),
]
assert frozen.calls == []
trainable.calls.clear()
owner.training_config.distributed.reshard_after_forward = True
method_base.TrainingMethod.set_requires_gradient_sync(owner, False)
method_base.TrainingMethod.set_requires_gradient_sync(owner, True)
assert trainable.calls == [
("sync", False, True),
("last", False, None),
("sync", True, True),
("last", True, None),
]
assert frozen.calls == []
@@ -61,6 +61,10 @@ def test_minimal_yaml_applies_all_defaults(tmp_path: Path) -> None:
assert t.distributed.sp_size == 1
assert t.distributed.hsdp_replicate_dim == 1
assert t.distributed.pin_cpu_memory is False
assert t.distributed.reshard_after_forward is True
assert t.distributed.fsdp_symmetric_memory is False
assert t.distributed.fsdp_modules_per_group == 1
assert t.distributed.reduce_dtype == "fp32"
assert t.data.train_batch_size == 1
assert t.data.dataloader_num_workers == 0
@@ -72,6 +76,7 @@ def test_minimal_yaml_applies_all_defaults(tmp_path: Path) -> None:
assert t.optimizer.weight_decay == 0.0
assert t.optimizer.lr_scheduler == "constant"
assert t.optimizer.min_lr_ratio == 0.5
assert t.optimizer.fused is None
assert t.loop.max_train_steps == 0
assert t.loop.gradient_accumulation_steps == 1
@@ -86,6 +91,7 @@ def test_minimal_yaml_applies_all_defaults(tmp_path: Path) -> None:
assert t.model.weighting_scheme == "uniform"
assert t.model.precondition_outputs is False
assert t.model.moba_config == {}
assert t.model.enable_torch_compile is False
assert t.dit_precision == "fp32"
assert t.vsa_sparsity == 0.0
@@ -102,6 +108,10 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None:
"hsdp_replicate_dim": 2,
"hsdp_shard_dim": 2,
"pin_cpu_memory": True,
"reshard_after_forward": False,
"fsdp_symmetric_memory": True,
"fsdp_modules_per_group": 2,
"reduce_dtype": "bf16",
},
"data": {
"data_path": "/some/path",
@@ -121,6 +131,7 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None:
"lr_scheduler": "cosine",
"lr_warmup_steps": 100,
"min_lr_ratio": 0.1,
"fused": True,
},
"loop": {
"max_train_steps": 1000,
@@ -144,6 +155,7 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None:
"logit_mean": 0.5,
"logit_std": 1.5,
"precondition_outputs": True,
"enable_torch_compile": True,
},
"dit_precision": "bf16",
}
@@ -153,6 +165,10 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None:
assert t.distributed.num_gpus == 4
assert t.distributed.tp_size == 2
assert t.distributed.pin_cpu_memory is True
assert t.distributed.reshard_after_forward is False
assert t.distributed.fsdp_symmetric_memory is True
assert t.distributed.fsdp_modules_per_group == 2
assert t.distributed.reduce_dtype == "bf16"
assert t.data.train_batch_size == 2
assert t.data.data_path == "/some/path"
@@ -164,6 +180,7 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None:
assert t.optimizer.lr_scheduler == "cosine"
assert t.optimizer.lr_warmup_steps == 100
assert t.optimizer.min_lr_ratio == pytest.approx(0.1)
assert t.optimizer.fused is True
assert t.loop.max_train_steps == 1000
assert t.loop.gradient_accumulation_steps == 4
@@ -177,6 +194,7 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None:
assert t.vsa_sparsity == pytest.approx(0.5)
assert t.model.weighting_scheme == "logit_normal"
assert t.model.precondition_outputs is True
assert t.model.enable_torch_compile is True
assert t.dit_precision == "bf16"
@@ -292,14 +310,26 @@ def test_dotted_overrides_apply_with_type_coercion(tmp_path: Path) -> None:
overrides = [
"--training.distributed.num_gpus=4",
"--training.optimizer.learning_rate=1e-3",
"--training.optimizer.fused=true",
"--training.distributed.pin_cpu_memory=true",
"--training.distributed.reshard_after_forward=false",
"--training.distributed.fsdp_symmetric_memory=true",
"--training.distributed.fsdp_modules_per_group=2",
"--training.distributed.reduce_dtype=bf16",
"--training.model.enable_torch_compile=true",
"--training.tracker.project_name=overridden",
]
cfg = load_run_config(path, overrides=overrides)
assert cfg.training.distributed.num_gpus == 4
assert cfg.training.optimizer.learning_rate == pytest.approx(1e-3)
assert cfg.training.optimizer.fused is True
assert cfg.training.distributed.pin_cpu_memory is True
assert cfg.training.distributed.reshard_after_forward is False
assert cfg.training.distributed.fsdp_symmetric_memory is True
assert cfg.training.distributed.fsdp_modules_per_group == 2
assert cfg.training.distributed.reduce_dtype == "bf16"
assert cfg.training.model.enable_torch_compile is True
assert cfg.training.tracker.project_name == "overridden"
@@ -312,6 +342,69 @@ def test_dotted_overrides_accept_separate_value_token(tmp_path: Path) -> None:
assert cfg.training.distributed.num_gpus == 8
def test_optimizer_fused_rejects_non_bool(tmp_path: Path) -> None:
data = _minimal_yaml()
data["training"] = {"optimizer": {"fused": "true"}}
with pytest.raises(ValueError, match="training.optimizer.fused must be a bool"):
load_run_config(_write_yaml(tmp_path, data))
def test_reshard_after_forward_rejects_non_bool(tmp_path: Path) -> None:
data = _minimal_yaml()
data["training"] = {
"distributed": {
"reshard_after_forward": "false"
}
}
with pytest.raises(
ValueError,
match="training.distributed.reshard_after_forward must be a bool",
):
load_run_config(_write_yaml(tmp_path, data))
def test_fsdp_symmetric_memory_rejects_non_bool(tmp_path: Path) -> None:
data = _minimal_yaml()
data["training"] = {
"distributed": {
"fsdp_symmetric_memory": "true"
}
}
with pytest.raises(
ValueError,
match="training.distributed.fsdp_symmetric_memory must be a bool",
):
load_run_config(_write_yaml(tmp_path, data))
@pytest.mark.parametrize("value", [0, -1, True])
def test_fsdp_modules_per_group_rejects_non_positive_int(tmp_path: Path, value: object) -> None:
data = _minimal_yaml()
data["training"] = {"distributed": {"fsdp_modules_per_group": value}}
with pytest.raises(ValueError, match="training.distributed.fsdp_modules_per_group"):
load_run_config(_write_yaml(tmp_path, data))
def test_model_torch_compile_rejects_non_bool(tmp_path: Path) -> None:
data = _minimal_yaml()
data["training"] = {"model": {"enable_torch_compile": "true"}}
with pytest.raises(ValueError, match="training.model.enable_torch_compile must be a bool"):
load_run_config(_write_yaml(tmp_path, data))
def test_reduce_dtype_rejects_unknown_precision(tmp_path: Path) -> None:
data = _minimal_yaml()
data["training"] = {"distributed": {"reduce_dtype": "fp16"}}
with pytest.raises(ValueError, match="training.distributed.reduce_dtype"):
load_run_config(_write_yaml(tmp_path, data))
def test_overrides_create_intermediate_keys(tmp_path: Path) -> None:
"""Overrides into a nested key absent from YAML should still apply."""
data = _minimal_yaml()
@@ -0,0 +1,132 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from typing import Any
import pytest
import torch
from torch.distributed.fsdp import MixedPrecisionPolicy
from fastvideo.models.loader import fsdp_load
class _Block(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.weight = torch.nn.Parameter(torch.ones(2))
self.sensitive = torch.nn.Parameter(torch.ones(1))
class _Model(torch.nn.Module):
def __init__(self, block_count: int = 4) -> None:
super().__init__()
self.blocks = torch.nn.ModuleList([_Block() for _ in range(block_count)])
self.root_weight = torch.nn.Parameter(torch.ones(1))
self.root_sensitive = torch.nn.Parameter(torch.ones(1))
class _MixedDtypeModel(_Model):
def _get_parameter_dtype(self, name: str, default_dtype: torch.dtype) -> torch.dtype:
return torch.float32 if name.endswith("sensitive") else default_dtype
def _is_block(name: str, module: torch.nn.Module) -> bool:
return isinstance(module, _Block) and name.startswith("blocks.")
def _record_fully_shard(monkeypatch) -> list[tuple[Any, dict[str, Any]]]:
calls: list[tuple[Any, dict[str, Any]]] = []
monkeypatch.delenv("FASTVIDEO_FSDP2_AUTOWRAP", raising=False)
monkeypatch.setattr(fsdp_load, "fully_shard", lambda module, **kwargs: calls.append((module, kwargs)))
return calls
def _shard(model: torch.nn.Module, *, modules_per_group: int = 1) -> None:
fsdp_load.shard_model(
model,
cpu_offload=False,
mp_policy=MixedPrecisionPolicy(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
),
mesh=object(),
fsdp_shard_conditions=[_is_block],
fsdp_modules_per_group=modules_per_group,
)
def test_default_shards_each_matching_module_in_existing_reverse_order(monkeypatch) -> None:
model = _Model()
calls = _record_fully_shard(monkeypatch)
_shard(model)
assert [module for module, _ in calls] == [*reversed(model.blocks), model]
assert all(not isinstance(module, list) for module, _ in calls)
def test_groups_adjacent_matching_modules_and_shards_root_last(monkeypatch) -> None:
model = _Model()
calls = _record_fully_shard(monkeypatch)
_shard(model, modules_per_group=2)
assert [module for module, _ in calls] == [
[model.blocks[0], model.blocks[1]],
[model.blocks[2], model.blocks[3]],
model,
]
def test_grouped_modules_receive_their_combined_ignored_parameters(monkeypatch) -> None:
model = _MixedDtypeModel(block_count=3)
calls = _record_fully_shard(monkeypatch)
_shard(model, modules_per_group=2)
assert calls[0][1]["ignored_params"] == {
model.blocks[0].sensitive,
model.blocks[1].sensitive,
}
assert calls[1][1]["ignored_params"] == {model.blocks[2].sensitive}
assert calls[2][0] is model
assert calls[2][1]["ignored_params"] == {
*(block.sensitive for block in model.blocks),
model.root_sensitive,
}
@pytest.mark.parametrize("modules_per_group", [0, -1])
def test_rejects_invalid_group_size_before_sharding(monkeypatch, modules_per_group: int) -> None:
calls = _record_fully_shard(monkeypatch)
with pytest.raises(ValueError, match="must be at least 1"):
_shard(_Model(), modules_per_group=modules_per_group)
assert calls == []
def test_rejects_grouping_with_size_based_autowrap(monkeypatch) -> None:
calls = _record_fully_shard(monkeypatch)
monkeypatch.setenv("FASTVIDEO_FSDP2_AUTOWRAP", "1")
with pytest.raises(ValueError, match="incompatible with FASTVIDEO_FSDP2_AUTOWRAP"):
_shard(_Model(), modules_per_group=2)
assert calls == []
@pytest.mark.parametrize("shared_with", [1, 2])
def test_rejects_overlapping_parameter_sets_before_sharding(monkeypatch, shared_with: int) -> None:
model = _Model()
model.blocks[shared_with].weight = model.blocks[0].weight
calls = _record_fully_shard(monkeypatch)
with pytest.raises(ValueError, match="overlapping parameter sets"):
_shard(model, modules_per_group=2)
assert calls == []
@@ -85,3 +85,102 @@ def test_load_transformer_restores_backend_when_loading_fails(
attention_backend="ATTN_QAT_TRAIN",
)
assert get_global_forced_attn_backend() is None
def test_load_transformer_configures_recursive_fsdp_reshard(
monkeypatch,
tmp_path,
) -> None:
training_config = TrainingConfig(
distributed=DistributedConfig(
hsdp_shard_dim=1,
reshard_after_forward=False,
),
pipeline_config=PipelineConfig(),
)
calls: list[tuple[bool, bool]] = []
class FakeTransformer(torch.nn.Linear):
def set_reshard_after_forward(
self,
value: bool,
*,
recurse: bool,
) -> None:
calls.append((value, recurse))
monkeypatch.setattr(moduleloader, "maybe_download_model", lambda path: str(tmp_path))
monkeypatch.setattr(
moduleloader,
"verify_model_config_and_directory",
lambda path: {"transformer": ("diffusers", "FakeTransformer")},
)
monkeypatch.setattr(
moduleloader.PipelineComponentLoader,
"load_module",
lambda **kwargs: FakeTransformer(1, 1),
)
moduleloader.load_module_from_path(
model_path="fake/model",
module_type="transformer",
training_config=training_config,
)
assert calls == [(False, True)]
def test_load_transformer_configures_all_fsdp_symmetric_memory_modules(
monkeypatch,
tmp_path,
) -> None:
training_config = TrainingConfig(
distributed=DistributedConfig(
hsdp_shard_dim=1,
fsdp_symmetric_memory=True,
),
pipeline_config=PipelineConfig(),
)
calls: list[tuple[str, str, object]] = []
class FakeFSDPModule(torch.nn.Module):
def __init__(self, name: str, child: torch.nn.Module | None = None) -> None:
super().__init__()
self.name = name
if child is not None:
self.child = child
def set_force_sum_reduction_for_comms(self, value: bool) -> None:
calls.append((self.name, "force_sum", value))
def set_symm_mem_for_comm(self, backend: str) -> None:
calls.append((self.name, "symmetric_memory", backend))
transformer = FakeFSDPModule("root", FakeFSDPModule("block"))
monkeypatch.setattr(moduleloader, "FSDPModule", FakeFSDPModule)
monkeypatch.setattr(moduleloader, "maybe_download_model", lambda path: str(tmp_path))
monkeypatch.setattr(
moduleloader,
"verify_model_config_and_directory",
lambda path: {"transformer": ("diffusers", "FakeTransformer")},
)
monkeypatch.setattr(
moduleloader.PipelineComponentLoader,
"load_module",
lambda **kwargs: transformer,
)
moduleloader.load_module_from_path(
model_path="fake/model",
module_type="transformer",
training_config=training_config,
)
assert calls == [
("root", "force_sum", True),
("root", "symmetric_memory", "NCCL"),
("block", "force_sum", True),
("block", "symmetric_memory", "NCCL"),
]
@@ -0,0 +1,127 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import torch
import pytest
from fastvideo.models.loader import fsdp_load
from fastvideo.train.utils import moduleloader
from fastvideo.train.utils.training_config import (
DistributedConfig,
ModelTrainingConfig,
TrainingConfig,
)
class _RepeatedModel(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.blocks = torch.nn.ModuleList([torch.nn.Linear(1, 1), torch.nn.Linear(1, 1)])
self.output = torch.nn.Linear(1, 1)
self._compile_conditions = [lambda name, _module: name in {"blocks.0", "blocks.1"}]
self._fsdp_shard_conditions = self._compile_conditions
self.param_names_mapping = {}
def test_make_training_args_propagates_training_fsdp_options() -> None:
config = TrainingConfig(
distributed=DistributedConfig(
hsdp_shard_dim=1,
reduce_dtype="bf16",
fsdp_modules_per_group=2,
),
model=ModelTrainingConfig(enable_torch_compile=True),
)
args = moduleloader._make_training_args(config, model_path="fake/model")
assert args.enable_torch_compile is True
assert args.fsdp_reduce_dtype == "bf16"
assert args.fsdp_modules_per_group == 2
def test_compile_matched_submodule_forwards_is_regional_and_allows_graph_breaks(monkeypatch) -> None:
model = _RepeatedModel()
output_forward = model.output.forward
calls: list[tuple[object, dict[str, object]]] = []
def fake_compile(forward, **kwargs):
calls.append((forward.__self__, kwargs))
return forward
monkeypatch.setattr(fsdp_load.torch, "compile", fake_compile)
count = fsdp_load._compile_matched_submodule_forwards(model, {"backend": "eager"})
assert count == 2
assert [module for module, _ in calls] == list(model.blocks)
assert [kwargs for _, kwargs in calls] == [
{
"backend": "eager",
"fullgraph": False
},
{
"backend": "eager",
"fullgraph": False
},
]
assert model.output.forward == output_forward
@pytest.mark.parametrize(
"conditions",
[[], [lambda _name, _module: False]],
)
def test_compile_matched_submodule_forwards_rejects_missing_matches(monkeypatch, conditions) -> None:
model = _RepeatedModel()
model._compile_conditions = conditions
monkeypatch.setattr(fsdp_load.torch, "compile", lambda forward, **kwargs: forward)
with pytest.raises(ValueError, match="refusing to compile the whole FSDP model"):
fsdp_load._compile_matched_submodule_forwards(model)
def test_compile_matched_submodule_forwards_rejects_fullgraph() -> None:
with pytest.raises(ValueError, match="requires fullgraph=False"):
fsdp_load._compile_matched_submodule_forwards(_RepeatedModel(), {"fullgraph": True})
def test_fsdp_loader_compiles_regions_before_sharding(monkeypatch) -> None:
events: list[str] = []
def fake_compile(forward, **_kwargs):
events.append("compile")
return forward
def fake_shard_model(model, **_kwargs):
del model
events.append("shard")
def fake_load_state(model, _iterator, device, _dtype, **_kwargs):
model.to_empty(device=device)
monkeypatch.setattr(fsdp_load.torch, "compile", fake_compile)
monkeypatch.setattr(fsdp_load, "init_device_mesh", lambda *_args, **_kwargs: object())
monkeypatch.setattr(fsdp_load, "shard_model", fake_shard_model)
monkeypatch.setattr(fsdp_load, "safetensors_weights_iterator", lambda *_args, **_kwargs: iter(()))
monkeypatch.setattr(fsdp_load, "get_param_names_mapping", lambda _mapping: None)
monkeypatch.setattr(fsdp_load, "load_model_from_full_model_state_dict", fake_load_state)
monkeypatch.setattr(fsdp_load, "_maybe_quantize_model", lambda _model: None)
monkeypatch.setattr("fastvideo.platforms.current_platform.is_mps", lambda: False)
fsdp_load.maybe_load_fsdp_model(
model_cls=_RepeatedModel,
init_params={},
weight_dir_list=[],
device=torch.device("cpu"),
hsdp_replicate_dim=1,
hsdp_shard_dim=1,
default_dtype=torch.float32,
param_dtype=torch.float32,
reduce_dtype=torch.float32,
enable_torch_compile=True,
)
assert events == ["compile", "compile", "shard"]
@@ -0,0 +1,48 @@
# SPDX-License-Identifier: Apache-2.0
from types import SimpleNamespace
import torch
from fastvideo.train.methods.base import TrainingMethod
from fastvideo.train.utils.optimizer import build_optimizer_and_scheduler
from fastvideo.train.utils.training_config import (
OptimizerConfig,
TrainingLoopConfig,
)
def test_build_optimizer_forwards_fused_flag() -> None:
parameter = torch.nn.Parameter(torch.ones(1))
optimizer, _ = build_optimizer_and_scheduler(
params=[parameter],
optimizer_config=OptimizerConfig(fused=True),
loop_config=TrainingLoopConfig(max_train_steps=1),
learning_rate=1e-3,
betas=(0.9, 0.999),
scheduler_name="constant",
)
assert optimizer.defaults["fused"] is True
def test_resume_seed_places_fused_step_with_parameter() -> None:
parameter = torch.nn.Parameter(torch.empty(1, device="meta"))
optimizer = torch.optim.AdamW([parameter], fused=True)
method = SimpleNamespace(get_optimizers=lambda _: [optimizer])
TrainingMethod.seed_optimizer_state_for_resume(method)
assert optimizer.state[parameter]["step"].device == parameter.device
assert optimizer.state[parameter]["step"].dtype == torch.float32
def test_resume_seed_keeps_default_step_on_cpu() -> None:
parameter = torch.nn.Parameter(torch.empty(1, device="meta"))
optimizer = torch.optim.AdamW([parameter])
method = SimpleNamespace(get_optimizers=lambda _: [optimizer])
TrainingMethod.seed_optimizer_state_for_resume(method)
assert optimizer.state[parameter]["step"].device.type == "cpu"
@@ -0,0 +1,123 @@
# SPDX-License-Identifier: Apache-2.0
from types import SimpleNamespace
import pytest
import torch
from fastvideo.models.loader.utils import get_param_names_mapping, hf_to_custom_state_dict
from fastvideo.training import training_utils
def _packed_projection_mapper():
return get_param_names_mapping({
r"^q\.weight$": ("attention.to_qkv.weight", 0, 3),
r"^k\.weight$": ("attention.to_qkv.weight", 1, 3),
r"^v\.weight$": ("attention.to_qkv.weight", 2, 3),
r"^norm\.weight$": "model.norm.weight",
})
def test_merged_parameter_round_trip_preserves_unequal_split_sizes() -> None:
hf_state = {
"q.weight": torch.arange(6).reshape(2, 3),
"k.weight": torch.arange(3).reshape(1, 3) + 10,
"v.weight": torch.arange(9).reshape(3, 3) + 20,
"norm.weight": torch.arange(3),
}
custom_state, reverse_mapping = hf_to_custom_state_dict(
iter([
("v.weight", hf_state["v.weight"]),
("norm.weight", hf_state["norm.weight"]),
("q.weight", hf_state["q.weight"]),
("k.weight", hf_state["k.weight"]),
]),
_packed_projection_mapper(),
)
torch.testing.assert_close(
custom_state["attention.to_qkv.weight"],
torch.cat([
hf_state["q.weight"],
hf_state["k.weight"],
hf_state["v.weight"],
]),
)
assert reverse_mapping["attention.to_qkv.weight"] == [
("q.weight", 0, 3, 2),
("k.weight", 1, 3, 1),
("v.weight", 2, 3, 3),
]
round_trip = training_utils.custom_to_hf_state_dict(
custom_state.items(),
reverse_mapping,
)
assert round_trip.keys() == hf_state.keys()
for name, tensor in hf_state.items():
torch.testing.assert_close(round_trip[name], tensor)
def test_hf_to_custom_rejects_incomplete_merge_group() -> None:
with pytest.raises(ValueError, match="Incomplete merged parameters"):
hf_to_custom_state_dict(
{
"q.weight": torch.ones(2, 3),
"k.weight": torch.ones(1, 3),
},
_packed_projection_mapper(),
)
def test_hf_to_custom_rejects_duplicate_merge_index() -> None:
mapper = get_param_names_mapping({
r"^q\.weight$": ("attention.to_qk.weight", 0, 2),
r"^k\.weight$": ("attention.to_qk.weight", 0, 2),
})
with pytest.raises(ValueError, match="Duplicate merge index 0"):
hf_to_custom_state_dict(
{
"q.weight": torch.ones(2, 3),
"k.weight": torch.ones(1, 3),
},
mapper,
)
def test_custom_to_hf_rejects_incomplete_reverse_merge_group() -> None:
with pytest.raises(ValueError, match="Incomplete reverse merge mapping"):
training_utils.custom_to_hf_state_dict(
{"attention.to_qkv.weight": torch.ones(5, 3)},
{
"attention.to_qkv.weight": [
("q.weight", 0, 3, 2),
("v.weight", 2, 3, 3),
]
},
)
def test_clip_grad_norm_uses_local_dtensor_shards_for_foreach(monkeypatch) -> None:
class FakeDTensor:
def __init__(self, local: torch.Tensor) -> None:
self.local = local
def to_local(self) -> torch.Tensor:
return self.local
monkeypatch.setattr(torch.distributed.tensor, "DTensor", FakeDTensor)
local_grad = torch.tensor([3.0, 4.0])
parameter = SimpleNamespace(grad=FakeDTensor(local_grad))
training_utils._clip_grads_with_norm_(
[parameter],
max_norm=1.0,
total_norm=torch.tensor(5.0),
)
torch.testing.assert_close(local_grad, torch.tensor([0.6, 0.8]))
+28 -3
View File
@@ -8,7 +8,9 @@ Optionally logs per-module grad norms to the tracker.
from __future__ import annotations
from typing import TYPE_CHECKING
from typing import Any, TYPE_CHECKING
import torch
from fastvideo.logger import init_logger
from fastvideo.train.callbacks.callback import Callback
@@ -36,12 +38,16 @@ class GradNormClipCallback(Callback):
) -> None:
self._max_grad_norm = float(max_grad_norm)
self._log_grad_norms = bool(log_grad_norms)
self._pending_grad_norms: list[tuple[Any, str, torch.Tensor, int]] = []
def on_before_optimizer_step(
self,
method: TrainingMethod,
iteration: int = 0,
) -> None:
if self._pending_grad_norms:
raise RuntimeError("Previous gradient norms were not logged after the optimizer step")
max_norm = self._max_grad_norm
if max_norm <= 0.0:
return
@@ -54,8 +60,27 @@ class GradNormClipCallback(Callback):
module,
max_norm,
)
if (self._log_grad_norms and tracker is not None and grad_norm > 0.0):
if self._log_grad_norms and tracker is not None and grad_norm is not None:
self._pending_grad_norms.append((tracker, name, grad_norm, iteration))
def on_training_step_end(
self,
method: TrainingMethod,
loss_dict: dict[str, Any],
iteration: int = 0,
) -> None:
del method, loss_dict
if not self._pending_grad_norms:
return
pending = self._pending_grad_norms
self._pending_grad_norms = []
for tracker, name, grad_norm, pending_iteration in pending:
if pending_iteration != iteration:
raise RuntimeError(f"Gradient norms from step {pending_iteration} reached optimizer step {iteration}")
value = float(grad_norm.item())
if value > 0.0:
tracker.log(
{f"grad_norm/{name}": grad_norm},
{f"grad_norm/{name}": value},
iteration,
)
+5 -2
View File
@@ -272,8 +272,11 @@ class ValidationCallback(Callback):
# Look for an EMA callback to temporarily swap
# EMA weights during validation.
ema_cb = self._find_ema_callback()
ctx = ema_cb.ema_context(transformer) if ema_cb is not None else contextlib.nullcontext(transformer)
with ctx as t:
ema_ctx = ema_cb.ema_context(transformer) if ema_cb is not None else contextlib.nullcontext(transformer)
# Keep inference-only graphs out of the training transformer's compile cache.
compile_ctx = (torch.compiler.set_stance("force_eager")
if self.training_config.model.enable_torch_compile else contextlib.nullcontext())
with compile_ctx, ema_ctx as t:
self._run_validation_inner(
method,
step,
+6 -11
View File
@@ -77,6 +77,7 @@ def _save_role_pretrained(
get_model_state_dict,
)
from fastvideo.models.loader.utils import custom_to_hf_state_dict
from fastvideo.utils import maybe_download_model
def _rank() -> int:
@@ -166,18 +167,12 @@ def _save_role_pretrained(
raise TypeError(f"Expected tensor in state_dict "
f"for {module_name}.{key}, "
f"got {type(value).__name__}")
if key in reverse_mapping:
hf_key, merge_index, _ = reverse_mapping[key]
if merge_index is not None:
logger.warning(
"Skipping reverse-mapping for merged param %s "
"(merge_index=%s); saving under internal key.",
key,
merge_index,
)
hf_key = key
key = hf_key
tensor_state[key] = value.detach().cpu()
if reverse_mapping:
tensor_state = custom_to_hf_state_dict(
tensor_state,
reverse_mapping,
)
from safetensors.torch import save_file
+6
View File
@@ -69,6 +69,12 @@ def run_training_from_config(
"SLA_ATTN",
)
if tc.distributed.fsdp_symmetric_memory:
nccl_cta_policy = os.environ.setdefault("NCCL_CTA_POLICY", "2")
if nccl_cta_policy != "2":
raise ValueError("training.distributed.fsdp_symmetric_memory requires "
"NCCL_CTA_POLICY=2 before distributed initialization")
maybe_init_distributed_environment_and_model_parallel(
tc.distributed.tp_size,
tc.distributed.sp_size,
+18 -1
View File
@@ -7,6 +7,7 @@ from collections.abc import Sequence
from typing import Any, Literal, TypeAlias
import torch
from torch.distributed.fsdp import FSDPModule
from fastvideo import envs
from fastvideo.logger import init_logger
@@ -155,6 +156,17 @@ class TrainingMethod(torch.nn.Module, ABC):
grad_accum_rounds = max(1, int(grad_accum_rounds))
(loss_map["total_loss"] / grad_accum_rounds).backward()
def set_requires_gradient_sync(self, requires_gradient_sync: bool) -> None:
"""Control FSDP communication between accumulation microbatches."""
retain_parameters = not self.training_config.distributed.reshard_after_forward
for model in self._role_models.values():
transformer = getattr(model, "transformer", None)
if getattr(model, "_trainable", False) and isinstance(transformer, FSDPModule):
transformer.set_requires_gradient_sync(requires_gradient_sync, recurse=True)
if retain_parameters:
transformer.set_reshard_after_backward(requires_gradient_sync, recurse=True)
transformer.set_is_last_backward(requires_gradient_sync)
def optimizers_schedulers_step(
self,
iteration: int,
@@ -184,13 +196,18 @@ class TrainingMethod(torch.nn.Module, ABC):
"""
for opt in self.get_optimizers(0):
for group in opt.param_groups:
fused = bool(group.get("fused"))
capturable = bool(group.get("capturable"))
for p in group["params"]:
if not p.requires_grad:
continue
if len(opt.state.get(p, {})) > 0:
continue
step_device = p.device if fused or capturable else torch.device("cpu")
step_dtype = (torch.float64
if not fused and torch.get_default_dtype() == torch.float64 else torch.float32)
opt.state[p] = {
"step": torch.tensor(0.0),
"step": torch.zeros((), dtype=step_dtype, device=step_device),
"exp_avg": torch.zeros_like(p),
"exp_avg_sq": torch.zeros_like(p),
}
+25 -39
View File
@@ -2,7 +2,7 @@
"""LTX-2 model plugin (per-role instance).
Subclasses WanModel but replaces the pieces where LTX-2 differs:
- transformer class name: LTX2Transformer3DModel (blocks nested
- transformer class name: LTX2VideoOnlyTransformer3DModel (blocks nested
under ``transformer.model``, so activation checkpointing must
target the inner module)
- latents come pre-normalized from the LTX-2 VAE encoder
@@ -12,19 +12,17 @@ Subclasses WanModel but replaces the pieces where LTX-2 differs:
shifted-logit-normal with a token-count-dependent shift
(0.95 @ 1024 tokens -> 2.05 @ 4096 tokens) and a 10% uniform
mixture, instead of scheduler-index sampling
- the DiT consumes PER-TOKEN sigmas in [0, 1] shaped [B, tokens]
(not integer 0-1000 timesteps) and returns the DENOISED x0
prediction; ``predict_noise`` converts it back to the
- the DiT consumes sigmas in [0, 1] (one per sample for uniform T2V,
expanded per token when sequence parallelism needs to shard them)
and returns the DENOISED x0 prediction; ``predict_noise`` converts it back to the
framework's velocity convention v = (x_t - x0) / sigma so the
default FineTune target ``noise - clean`` applies unchanged
- temporal RoPE coordinates are divided by fps read from
``get_forward_context().forward_batch.fps``; training must set
it to the data fps or validation (which uses the preset fps)
would see different temporal frequencies than training
- the ~5.8B audio / cross-modal (a2v, v2a, av_ca) parameters are
frozen: video-only forwards never give them
gradients, and freezing keeps them out of DCP checkpoints and
optimizer state
- the ~5.8B audio / cross-modal parameters are omitted entirely; the
modular trainer only supports video batches
"""
from __future__ import annotations
@@ -35,9 +33,9 @@ import torch
import fastvideo.envs as envs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.dits.ltx2 import VideoLatentShape
from fastvideo.pipelines import ForwardBatch, TrainingBatch
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing, )
@@ -52,13 +50,6 @@ if TYPE_CHECKING:
TrainingConfig, )
from fastvideo.train.utils.lora import LoraConfig
logger = init_logger(__name__)
# Parameters whose names match any of these substrings belong to the
# audio branch or the audio<->video cross-modal machinery. They are
# unused (no gradients) in video-only forwards.
_AUDIO_PARAM_PATTERNS = ("audio", "a2v", "v2a", "av_ca")
# Official LTX-2 trainer timestep sampler constants
# (ltx_trainer/timestep_samplers.py: ShiftedLogitNormalTimestepSampler).
_SIGMA_MIN_TOKENS = 1024.0
@@ -80,7 +71,7 @@ _DEFAULT_ROPE_FPS = 24.0
class LTX2Model(WanModel):
"""LTX-2 per-role model for the modular trainer."""
_transformer_cls_name: str = "LTX2Transformer3DModel"
_transformer_cls_name: str = "LTX2VideoOnlyTransformer3DModel"
def __init__(
self,
@@ -94,6 +85,7 @@ class LTX2Model(WanModel):
transformer_override_safetensor: str
| None = None,
lora: LoraConfig | dict[str, Any] | None = None,
attention_backend: AttentionBackendEnum | str | None = None,
train_audio: bool = False,
timestep_uniform_prob: float = 0.1,
) -> None:
@@ -126,11 +118,9 @@ class LTX2Model(WanModel):
enable_gradient_checkpointing_type=(enable_gradient_checkpointing_type),
transformer_override_safetensor=(transformer_override_safetensor),
lora=lora,
attention_backend=attention_backend,
)
if trainable:
self._freeze_audio_parameters()
# No negative-prompt cache: cfg_rate is forced to 0 above and
# loading Gemma (~23GB) on every rank just for an unused
# negative embedding is wasteful.
@@ -149,7 +139,12 @@ class LTX2Model(WanModel):
enable_gradient_checkpointing_type: str | None,
training_config: TrainingConfig,
transformer_override_safetensor: str | None = None,
attention_backend: AttentionBackendEnum | str | None = None,
) -> torch.nn.Module:
arch_config = training_config.pipeline_config.dit_config.arch_config
if (getattr(arch_config, "pack_attention_projections", False) and self._lora_config is not None
and self._lora_config.enable):
raise ValueError("LTX-2 packed attention projections do not yet support LoRA training")
transformer = load_module_from_path(
model_path=init_from,
module_type="transformer",
@@ -157,6 +152,7 @@ class LTX2Model(WanModel):
disable_custom_init_weights=(disable_custom_init_weights),
override_transformer_cls_name=(self._transformer_cls_name),
transformer_override_safetensor=(transformer_override_safetensor),
attention_backend=attention_backend,
)
ckpt_type = (enable_gradient_checkpointing_type or getattr(
getattr(training_config, "model", None),
@@ -176,18 +172,6 @@ class LTX2Model(WanModel):
transformer = apply_trainable(transformer, trainable=trainable)
return transformer
def _freeze_audio_parameters(self) -> None:
frozen_params = 0
for name, param in self.transformer.named_parameters():
if any(pattern in name for pattern in _AUDIO_PARAM_PATTERNS):
param.requires_grad_(False)
frozen_params += param.numel()
logger.info(
"LTX2Model: froze %.2fB audio/cross-modal parameters "
"(video-only training)",
frozen_params / 1e9,
)
# ------------------------------------------------------------------
# Lifecycle
# ------------------------------------------------------------------
@@ -520,7 +504,8 @@ class LTX2Model(WanModel):
return training_batch
def _build_attention_metadata(self, training_batch: TrainingBatch) -> TrainingBatch:
if envs.FASTVIDEO_ATTENTION_BACKEND in ("VIDEO_SPARSE_ATTN", "VMOBA_ATTN"):
attention_backend = (self.attention_backend_name or envs.FASTVIDEO_ATTENTION_BACKEND)
if attention_backend in ("VIDEO_SPARSE_ATTN", "VMOBA_ATTN"):
raise NotImplementedError("LTX2Model does not support VSA/VMOBA attention backends")
training_batch.attn_metadata = None
return training_batch
@@ -544,12 +529,13 @@ class LTX2Model(WanModel):
if token_count is None:
token_count = int(noise_input.shape[2] * noise_input.shape[3] * noise_input.shape[4])
# The DiT wants per-token sigmas in [0, 1] shaped [B, tokens]
# (i2v-style conditioning would zero conditioned tokens; plain
# T2V uses the same sigma everywhere). ``timestep`` follows the
# framework's sigma*1000 convention.
# Plain T2V uses one sigma per sample. At SP=1 the AdaLN output
# broadcasts over tokens, so embedding that sigma once avoids a large
# redundant activation. SP sharding still requires an explicit token
# dimension so each rank receives its local timestep slice.
sigma = (timestep.to(torch.float32) / 1000.0).view(batch_size, 1)
per_token_timestep = sigma.expand(batch_size, token_count).contiguous()
if int(self.training_config.distributed.sp_size or 1) > 1:
sigma = sigma.expand(batch_size, token_count).contiguous()
return {
"hidden_states": noise_input,
@@ -557,6 +543,6 @@ class LTX2Model(WanModel):
# Post-connector embeddings are all-valid; the connector
# replaced pad positions with learnable registers.
"encoder_attention_mask": None,
"timestep": per_token_timestep,
"timestep": sigma,
"return_dict": False,
}
+2
View File
@@ -183,6 +183,8 @@ class Trainer:
step,
))
if grad_accum > 1:
method.set_requires_gradient_sync(accum_iter == grad_accum - 1)
method.backward(
loss_map,
outputs,
+40 -1
View File
@@ -24,7 +24,11 @@ from fastvideo.train.utils.training_config import (
logger = init_logger(__name__)
_TRAINING_DIT_ARCH_OVERRIDE_KEYS = ("local_attn_size", "sink_size")
_TRAINING_DIT_ARCH_OVERRIDE_KEYS = (
"local_attn_size",
"sink_size",
"pack_attention_projections",
)
@dataclass(slots=True)
@@ -334,6 +338,9 @@ def _build_training_config(
betas_raw = o.get("betas", "0.9,0.999")
betas = parse_betas(betas_raw, where="training.optimizer.betas")
fused = o.get("fused")
if fused is not None:
fused = require_bool(o, "fused", where="training.optimizer.fused")
model_path = str(t.get("model_path", "") or "")
if not model_path:
@@ -358,6 +365,12 @@ def _build_training_config(
"{'t2v', 'text_only'}, got "
f"{preprocessed_data_type!r}")
reduce_dtype = str(d.get("reduce_dtype", "fp32") or "fp32").strip().lower()
if reduce_dtype not in {"fp32", "bf16"}:
raise ValueError("training.distributed.reduce_dtype must be one of "
"{'fp32', 'bf16'}, got "
f"{reduce_dtype!r}")
return TrainingConfig(
distributed=DistributedConfig(
num_gpus=num_gpus,
@@ -366,6 +379,25 @@ def _build_training_config(
hsdp_replicate_dim=int(d.get("hsdp_replicate_dim", 1) or 1),
hsdp_shard_dim=int(d.get("hsdp_shard_dim", num_gpus) or num_gpus),
pin_cpu_memory=bool(d.get("pin_cpu_memory", False)),
reshard_after_forward=require_bool(
d,
"reshard_after_forward",
default=True,
where="training.distributed.reshard_after_forward",
),
fsdp_symmetric_memory=require_bool(
d,
"fsdp_symmetric_memory",
default=False,
where="training.distributed.fsdp_symmetric_memory",
),
fsdp_modules_per_group=require_positive_int(
d,
"fsdp_modules_per_group",
default=1,
where="training.distributed.fsdp_modules_per_group",
),
reduce_dtype=reduce_dtype,
),
data=DataConfig(
data_path=data_path,
@@ -388,6 +420,7 @@ def _build_training_config(
lr_num_cycles=int(o.get("lr_num_cycles", 0) or 0),
lr_power=float(o.get("lr_power", 0.0) or 0.0),
min_lr_ratio=float(o.get("min_lr_ratio", 0.5) or 0.5),
fused=fused,
),
loop=TrainingLoopConfig(
max_train_steps=int(lo.get("max_train_steps", 0) or 0),
@@ -415,6 +448,12 @@ def _build_training_config(
precondition_outputs=bool(m.get("precondition_outputs", False)),
moba_config=dict(m.get("moba_config", {}) or {}),
enable_gradient_checkpointing_type=(m.get("enable_gradient_checkpointing_type")),
enable_torch_compile=require_bool(
m,
"enable_torch_compile",
default=False,
where="training.model.enable_torch_compile",
),
),
pipeline_config=pipeline_config,
model_path=model_path,
+32 -1
View File
@@ -7,6 +7,7 @@ from contextlib import nullcontext
from typing import Any, TYPE_CHECKING
import torch
from torch.distributed.fsdp import FSDPModule
from fastvideo.attention.selector import (
coerce_attn_backend,
@@ -14,6 +15,7 @@ from fastvideo.attention.selector import (
)
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.fastvideo_args import ExecutionMode, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import (
PipelineComponentLoader, )
from fastvideo.utils import (
@@ -26,6 +28,8 @@ if TYPE_CHECKING:
from fastvideo.train.utils.training_config import (
TrainingConfig, )
logger = init_logger(__name__)
# ------------------------------------------------------------------
# TrainingArgs builders (only place that creates FastVideoArgs)
# ------------------------------------------------------------------
@@ -60,7 +64,9 @@ def _make_training_args(
text_encoder_cpu_offload=False,
image_encoder_cpu_offload=False,
use_fsdp_inference=False,
enable_torch_compile=False,
enable_torch_compile=tc.model.enable_torch_compile,
fsdp_reduce_dtype=tc.distributed.reduce_dtype,
fsdp_modules_per_group=tc.distributed.fsdp_modules_per_group,
)
@@ -163,4 +169,29 @@ def load_module_from_path(
if not isinstance(module, torch.nn.Module):
raise TypeError(f"Loaded {module_type!r} is not a "
f"torch.nn.Module: {type(module)}")
if module_type == "transformer":
reshard_after_forward = training_config.distributed.reshard_after_forward
set_reshard_after_forward = getattr(
module,
"set_reshard_after_forward",
None,
)
if callable(set_reshard_after_forward):
set_reshard_after_forward(
reshard_after_forward,
recurse=True,
)
elif not reshard_after_forward:
raise RuntimeError("training.distributed.reshard_after_forward=false requires "
"an FSDP-wrapped transformer")
if training_config.distributed.fsdp_symmetric_memory:
fsdp_modules = [submodule for submodule in module.modules() if isinstance(submodule, FSDPModule)]
if not fsdp_modules:
raise RuntimeError("training.distributed.fsdp_symmetric_memory=true requires "
"an FSDP-wrapped transformer")
for fsdp_module in fsdp_modules:
fsdp_module.set_force_sum_reduction_for_comms(True)
fsdp_module.set_symm_mem_for_comm("NCCL")
logger.info("Enabled FSDP symmetric-memory communication on %d modules", len(fsdp_modules))
return module
+4 -4
View File
@@ -42,6 +42,7 @@ def build_optimizer_and_scheduler(
betas=betas,
weight_decay=float(optimizer_config.weight_decay),
eps=1e-8,
fused=optimizer_config.fused,
)
scheduler = get_scheduler(
@@ -61,12 +62,11 @@ def build_optimizer_and_scheduler(
def clip_grad_norm_if_needed(
module: torch.nn.Module,
max_grad_norm: float,
) -> float:
) -> torch.Tensor | None:
if max_grad_norm <= 0.0:
return 0.0
grad_norm = (clip_grad_norm_while_handling_failing_dtensor_cases(
return None
return (clip_grad_norm_while_handling_failing_dtensor_cases(
[p for p in module.parameters()],
max_grad_norm,
foreach=None,
))
return (float(grad_norm.item()) if grad_norm is not None else 0.0)
+6
View File
@@ -18,6 +18,10 @@ class DistributedConfig:
hsdp_replicate_dim: int = 1
hsdp_shard_dim: int = -1
pin_cpu_memory: bool = False
reshard_after_forward: bool = True
fsdp_symmetric_memory: bool = False
fsdp_modules_per_group: int = 1
reduce_dtype: str = "fp32"
@dataclass(slots=True)
@@ -44,6 +48,7 @@ class OptimizerConfig:
lr_num_cycles: int = 0
lr_power: float = 0.0
min_lr_ratio: float = 0.5
fused: bool | None = None
@dataclass(slots=True)
@@ -77,6 +82,7 @@ class ModelTrainingConfig:
precondition_outputs: bool = False
moba_config: dict = field(default_factory=dict)
enable_gradient_checkpointing_type: str | None = None
enable_torch_compile: bool = False
@dataclass(slots=True)
+11 -61
View File
@@ -3,7 +3,7 @@ import json
import math
import os
import time
from collections.abc import Callable, Iterator
from collections.abc import Callable
from enum import Enum
from typing import Any
@@ -15,6 +15,10 @@ from torch.optim import Optimizer
from torch.optim.lr_scheduler import LambdaLR
from fastvideo.logger import init_logger
from fastvideo.models.loader.utils import (
ReverseParamNamesMapping,
custom_to_hf_state_dict,
)
from fastvideo.training.checkpointing_utils import (ModelWrapper, OptimizerWrapper, RandomStateWrapper,
SchedulerWrapper)
@@ -936,6 +940,11 @@ def _clip_grads_with_norm_(
max_norm = float(max_norm)
if len(grads) == 0:
return
# FSDP2 exposes sharded gradients as DTensors. Foreach only recognizes
# ordinary tensors, but mutating a DTensor's local shard updates the
# gradient in place. Localize before grouping so one foreach kernel can
# replace a Python loop and one scalar kernel per parameter.
grads = [g.to_local() if isinstance(g, torch.distributed.tensor.DTensor) else g for g in grads]
grouped_grads: dict[tuple[torch.device, torch.dtype],
tuple[list[list[torch.Tensor]],
list[int]]] = (_group_tensors_by_device_and_dtype([grads])) # type: ignore[assignment]
@@ -1016,68 +1025,9 @@ def _has_foreach_support(tensors: list[torch.Tensor], device: torch.device) -> b
return _device_has_foreach_support(device) and all(t is None or type(t) in [torch.Tensor] for t in tensors)
def custom_to_hf_state_dict(state_dict: dict[str, Any] | Iterator[tuple[str, torch.Tensor]],
reverse_param_names_mapping: dict[str, tuple[str, int, int]]) -> dict[str, Any]:
"""
Convert fastvideo's custom model format to diffusers format using reverse_param_names_mapping.
Args:
state_dict: State dict in fastvideo's custom format
reverse_param_names_mapping: Reverse mapping from fastvideo's custom format to diffusers format
Returns:
State dict in diffusers format
"""
assert len(reverse_param_names_mapping) > 0, "reverse_param_names_mapping is empty"
if isinstance(state_dict, Iterator):
state_dict = dict(state_dict)
new_state_dict = {}
# Group parameters that need to be split (merged parameters)
merge_groups: dict[str, list[tuple[str, int, int]]] = {}
# First pass: collect all merge groups
for training_key, (diffusers_key, merge_index, num_params_to_merge) in reverse_param_names_mapping.items():
if merge_index is not None:
# This is a merged parameter that needs to be split
if training_key not in merge_groups:
merge_groups[training_key] = []
merge_groups[training_key].append((diffusers_key, merge_index, num_params_to_merge))
# Second pass: handle merged parameters by splitting them
used_keys = set()
for training_key, splits in merge_groups.items():
if training_key in state_dict:
v = state_dict[training_key]
# Sort by merge_index to ensure correct order
splits.sort(key=lambda x: x[1])
total = splits[0][2]
split_size = v.shape[0] // total
split_tensors = torch.split(v, split_size, dim=0)
for diffusers_key, split_index, _ in splits:
new_state_dict[diffusers_key] = split_tensors[split_index]
used_keys.add(training_key)
# Third pass: handle regular parameters (direct mappings)
for training_key, v in state_dict.items():
if training_key in used_keys:
continue
if training_key in reverse_param_names_mapping:
diffusers_key, merge_index, _ = reverse_param_names_mapping[training_key]
if merge_index is None:
# Direct mapping
new_state_dict[diffusers_key] = v
else:
# No mapping found, keep as is
new_state_dict[training_key] = v
return new_state_dict
def _save_full_ema_safetensors_from_state(
state_dict: dict[str, Any],
reverse_param_names_mapping: dict[str, tuple[str, int, int]],
reverse_param_names_mapping: ReverseParamNamesMapping,
output_path: str,
) -> None:
"""
+9
View File
@@ -0,0 +1,9 @@
__pycache__/
*.log
*.sqlite
*.trace.json*
*.pt
*.safetensors
*.mp4
wandb/
outputs/
+11
View File
@@ -0,0 +1,11 @@
# LTX-2 MFU tracker
This directory is the resumable lab notebook for PR #1630. Preserve historical scripts and measured results; add new attempts rather than rewriting old evidence.
- Read `README.md` and `REPORT.md` before running a gate.
- Use the benchmark contract in `README.md`. Label any precision, batch, topology, or source change as a different baseline.
- Compare performance only within one healthy allocation, preferably A/B/A. Record slowest-rank step time after warmup.
- Never commit credentials, checkpoints, videos, raw logs, profiler traces, copied third-party source, or generated binaries. Record W&B/job IDs, source versions, and hashes instead.
- Put end-to-end harnesses in `harness/`, launch snapshots in `runners/`, focused experiments in `probes/`, and compact conclusions in `REPORT.md` or `reports/`.
- Update `README.md` with the current stopping point and `REPORT.md` with accepted and rejected results before handing off.
- Treat files below as frozen experiment snapshots unless a rerun explicitly supersedes one; production changes belong in the normal package directories and should later be extracted into focused PRs.
+69
View File
@@ -0,0 +1,69 @@
# LTX-2 training MFU tracker
This directory makes PR #1630 an intentionally cumulative, resumable scratch branch for LTX-2 training-performance work. It preserves enough source, configuration, and compact evidence to reproduce accepted results, understand rejected ideas, and continue toward 50% model FLOP utilization (MFU). Focused production changes can be extracted into smaller PRs later.
Start with this file, then read [`REPORT.md`](REPORT.md). The production recipe is [`examples/train/configs/overfit_ltx2_t2v.yaml`](../../../examples/train/configs/overfit_ltx2_t2v.yaml). Runner flags are the experiment-specific configuration record; there were no additional temporary YAML configs worth preserving.
## Current stopping point
The accepted measurement contract is full-parameter LTX-2 training with FP32 registered/master parameters and FP32 Adam moments, BF16 working compute and reductions, 81x480x832 video, dense FA4 attention, regional compile, no activation checkpointing, slowest-rank inclusive step time, 10 warmup steps, and the median of 20 measured steps. Do not compare a run that changes this contract as if it were the same baseline.
The harness numerator is `14.444115 percentage-seconds/sample`, derived from `353.8808175 TFLOP/sample / 2,450 TFLOP/s * 100`. Both constants are audited (see the REPORT's MFU formula audit): the FLOP figure is the exact strict no-recompute, blocks-only model-FLOP count `353,880,819,892,224` at 4,290 video/1,024 text tokens, reconciled integer-exactly against a `FlopCounterMode` measurement by `probes/audit_train_flops_per_sample.py`; the denominator is a house convention 2% below the 2,500 TFLOP/s GB200 vendor dense-BF16 peak, so reported MFU times 0.98 gives vendor-peak MFU. A different model, resolution, frame count, or FLOP convention requires a new numerator before its MFU is reportable.
| Topology | Local/global batch | Median step | Throughput | MFU | Status |
|---|---:|---:|---:|---:|---|
| 1x GB200 | 1/1 | no valid step | - | - | OOM when FP32 Adam state is materialized, even with full activation checkpointing |
| 4x GB200 | 2/8 | 0.707865914 s | 11.301575 samples/s | 40.810314% | final equal-global-batch result |
| 8x GB200, healthy two-tray allocation | 3/24 | 0.993335918 s | 24.161011 samples/s | 43.623053% | fastest accepted result |
| 8x GB200, degraded allocation | 3/24 | 1.421743279 s | 16.880685 samples/s | 30.478319% | same recipe; clocks/power were degraded, so baseline-ineligible |
The measured optimization source is `20c36acef`. Later commits `0e60a0e9c` and `3f3f06541` isolate compiled validation and unload video-only validation state; a 4x run completed steps 0-50, validations at 0 and 50, and resumed at step 51 without a compile reset or OOM ([W&B run `iopu4dwm`](https://wandb.ai/wlsaidhi/fastvideo_ltx2/runs/iopu4dwm)). A 4x/B2 A/X/B on allocation `1627918` re-gated tracker head `52f1114dd` as timing-neutral (+1.326 ms / +0.184% against a 0.044%-drift control midpoint, memory unchanged), so the table remains valid at current head; re-gate again after any future source change. The 2026-07-22 BF16 systems campaign then accepted an env-only stack for 4x work — `TORCHINDUCTOR_COORDINATE_DESCENT_TUNING=1` plus `CUBLASLT_WORKSPACE_SIZE=76800`, -1.083% / +0.432 MFU points at 4x/B2 — and measured it neutral (+0.016%) on a degraded-class 1,965 MHz 8x pair, so it stays out of `run_current.sh` until a healthy-class 8x pair re-gates it. See the REPORT campaign section for the accepted/rejected ledger (cuDNN attention, TunableOp, and 4x/B3 block-skip capacity are measured rejects), the rack-057 MNNVL/IMEX hang root cause, and the required two-tray launch preamble (socket ifnames plus raw-IP `MASTER_ADDR`).
Historical 1x W&B runs `5h2cy7m9`, `zj6wqtpg`, `nzh2k20t`, and `gbjiw7og` reported 11.71-13.46% MFU but used `training.dit_precision=bf16`. Their persistent parameters and Adam state were BF16, so they are non-comparable capacity diagnostics, not standard-training MFU results.
## Resume here
Dated handoff (2026-07-22 session, head `9c86d28e1`) — the queued gates, in priority order. All of them are blocked on the same prerequisite: a **healthy 2,062 MHz-class allocation** (check `nvidia-smi --query-gpu=clocks.max.sm` first; that day the gb-nvl-118 rack was the healthy bin while twelve rolls across racks 053/057/059 all drew 1,965 MHz trays, which are valid for ratio microgates but baseline-ineligible and too tuning-unstable for +-1% trainer gates).
- **Env stack -> recipe.** A/X/B `run_current.sh` at 4x/B2 with `TORCHINDUCTOR_COORDINATE_DESCENT_TUNING=1 CUBLASLT_WORKSPACE_SIZE=76800` (accepted -1.083% / +0.432 pp on allocation `1631376`; neutral on a degraded 8x pair). If it reproduces, fold the two envs into `run_current.sh` and refresh the stopping-point rows (also 8x/B3 on a healthy pair).
- **Pinned cuBLASLt capture verdict.** Gate [`harness/benchmark_fastvideo_train_lt_pinned.py`](harness/benchmark_fastvideo_train_lt_pinned.py) (traced custom-op variant at head; the graph-breaking variant that banked +0.234 pp is at `8bf6eff18`). The sweep margin is -8/-9.4% band, bit-exact; stable tuning on healthy hardware should decide between the variants and size the real capture (+1.9-2.2 pp potential). Production path if it wins: Inductor-level lowering or a `ReplicatedLinear` quant-method hook that keeps dbias fused.
- **Two-tray launches need the preamble** recorded in the REPORT campaign section: `GLOO_SOCKET_IFNAME=enP5p9s0 NCCL_SOCKET_IFNAME=enP5p9s0`, a fresh `MASTER_PORT`, `MASTER_ADDR` resolved to a raw IP. Rack gb-nvl-057 pairs hang NCCL MNNVL/IMEX at first collective (`NCCL_MNNVL_ENABLE=0` is diagnosis-only); rack-059 pairs passed health at 458.9 GB/s.
- **Kernel plan** in [`reports/bf16_kernel_research_plan.md`](reports/bf16_kernel_research_plan.md) carries per-item statuses: FA4 forward anomaly isolated standalone (tail-tile padding killed; remaining path is upstream CuTe forward schedule work), launch-gap graphs untouched. The feasibility budget fixes the target: 50% needs combined GEMM+FA4 throughput 58.9-64.7% vs 56.7% measured.
The general procedure:
1. Use a healthy GB200 allocation and record source SHA, GPU count/topology, clocks, power limit, throttle state, CUDA/PyTorch/FA4 versions, and W&B/job IDs. Run [`probes/pr1630_realloc1629519_native_health.py`](probes/pr1630_realloc1629519_native_health.py) and the topology/NCCL probes before accepting an 8x result.
2. Run [`runners/run_current.sh`](runners/run_current.sh). For the final configurations set `LOCAL_BATCH_SIZE=2` on 4x or `LOCAL_BATCH_SIZE=3` on 8x. This timing harness intentionally uses a dummy tracker; its structured `BF16_*` records, not W&B, are the MFU source of truth. Use [`runners/run_observed.sh`](runners/run_observed.sh) for an observable 4x training/validation run with the same production contract.
3. Attribute changes only with an A/candidate/B sandwich on one allocation. Use the slowest rank per step, discard 10 warmup steps, and report the 20-step median plus peak allocated/reserved memory. Never attribute a cross-allocation delta.
4. Check loss/gradient finiteness, FP32 parameter and Adam-state coverage, batch/accumulation semantics, and relevant export/resume or validation behavior before accepting a speedup.
5. Append the result—including rejections—to `REPORT.md`, add the exact driver or runner if it is new, update the stopping-point table only when superseded, and commit the handoff to this branch.
For profiling, use the existing FastVideo profiler/NVTX regions, PyTorch Profiler/Kineto, TorchInductor logs, CUDA-event microbenchmarks, and NCCL diagnostics. `nsys-ai` can read an exported Nsight Systems SQLite database offline but is not a collector and uses a different default GB200 FLOP convention; keep this harness authoritative.
## Artifact map
- `harness/`: end-to-end training instrumentation and Kineto drivers. `benchmark_fastvideo_train_pack_d016.py` is the detailed batch-aware driver behind the accepted rows; `benchmark_fastvideo_train_dynamic_world.py` is a lighter current-world diagnostic.
- `runners/`: `run_current.sh` is the maintained MFU entrypoint and `run_observed.sh` is the maintained W&B/validation entrypoint. Other files are frozen launch snapshots for historical A/B/A, topology, kernel, validation, and profiler gates. They deliberately retain source hashes and cluster assumptions; do not assume they run unchanged at current head.
- `probes/`: semantic/parity checks, distributed diagnostics, fixed-arena and graph prototypes, and exact-shape kernel gates.
- `reports/`: focused subreports and designs. `REPORT.md` is the authoritative chronological decision record.
Historical runners may reference `/mnt/FastVideo`, `/mnt/fv-pr1630-*`, `/mnt/fa4-cache`, `enP5p9s0`, `/mnt/te216*`, `/mnt/cutlass`, or old temporary harness aliases. Adapt checkout/data/cache paths for a new allocation. Most referenced drivers are preserved here; runners that expect the missing `bf0861ff...` harness are archival configuration records, not exact rerun entrypoints. `benchmark_fastvideo_train_pack_b3.py` was an alias of `harness/benchmark_fastvideo_train_ltx2_singleton_timestep.py`; `benchmark_fastvideo_train_pack_ga_ec4c.py` was an alias of `harness/benchmark_fastvideo_train_pack_d016.py`. The final validation runners are historical and intentionally assert the pre-fix source/diff; use the committed validation callbacks at current head for a new gate.
## Decision index
| Family | Decision | Where to continue |
|---|---|---|
| no activation checkpointing, FA4, regional compile | accepted; largest practical wins | current recipe and `run_current.sh` |
| fused AdamW, BF16 reductions, batched/deferred grad norm | accepted | production commits and full-step harnesses |
| warm repeated input, singleton timestep, packed LTX projections | accepted | recipe, packing harness, export/parity probes |
| symmetric-memory FSDP2, accumulation no-sync/retention, two-module FSDP groups | accepted, topology/policy sensitive | grouped and distributed runners |
| validation compile isolation and video-only unload | accepted correctness/capacity fixes | current source; repeat 51-step validation gate after related changes |
| raw velocity, max-autotune, attention compile flag, RMSNorm autocast removal, prefetch, 1-D mesh | rejected or timing-neutral | do not repeat without a new mechanism |
| CUTLASS/cuBLASLt fused GELU, QuACK/NVFP4 complete projection, current TP layouts | rejected by speed, safety, or quality gates | focused reports and exact-shape probes |
| fixed-arena ZeRO-2-style runtime and whole-step CUDA graphs | research only; not lifecycle-complete | fixed-arena and graph reports/probes |
| BF16 kernel research toward 50% (FA4 fwd anomaly, cuBLASLt algo pinning, launch-gap graphs, CuTe ffn GEMMs) | planned; systems ceiling is 48.34% with today's kernels | [`reports/bf16_kernel_research_plan.md`](reports/bf16_kernel_research_plan.md) |
| 1x standard full-parameter training | capacity blocked | add an explicit optimizer/parameter offload mode or use a clearly labeled nonstandard state precision |
## Artifact and credential policy
Commit source-like artifacts and compact reports only. Do not commit W&B keys, checkpoints, generated videos, raw logs, profiler traces, telemetry dumps, compiled binaries, generated CUDA, or copied third-party source. `run_observed.sh` inherits W&B authentication from the job environment or the user's W&B login. Record W&B/job IDs, package or upstream commit versions, command lines, and SHA-256 hashes in the report. Output files from historical runs are intentionally excluded.
+727
View File
@@ -0,0 +1,727 @@
# LTX-2 50%-MFU optimization gate (4x/8x GB200)
> Tracker status (2026-07-22): PR #1630 is intentionally an ongoing cumulative experiment branch. The authoritative stopping-point MFU rows are 40.810314% on 4x GB200 at local batch 2, 43.623053% on a healthy 8x two-tray allocation at local batch 3, and 30.478319% for the same 8x recipe on a degraded allocation. Standard 1x full-parameter training has no valid MFU result because FP32 master weights and FP32 Adam state OOM at the first optimizer step.
> Precision correction: historical 1x W&B runs `5h2cy7m9`, `zj6wqtpg`, `nzh2k20t`, and `gbjiw7og` used `training.dit_precision=bf16`. Their 11.71-13.46% MFU figures used persistent BF16 parameters and BF16 Adam state, so they are capacity diagnostics only and are not comparable to the standard FP32-master contract below.
Date: 2026-07-21
FastVideo source: `7f139e2b28610063d2f30526ba8f0ccae5d88944` for the original research gates; measured optimization follow-ups through `20c36acef`, followed by validation fixes `0e60a0e9c` and `3f3f06541` (PR #1630)
Fresh allocations: Slurm `1622676` and `1624214`, 4x NVIDIA GB200; Slurm `1623561`, 8x NVIDIA GB200 across two collocated trays
Workload: `FastVideo/LTX2-Distilled-Diffusers`, standard mixed precision (FP32 master weights/optimizer state, BF16 working compute and reductions), 81x480x832, dense FA4, regional compile, no activation checkpointing, repeated/prefetched input, singleton uniform-timestep embedding at SP=1, and persistent projection packing where labeled. Local batch and gradient accumulation are 1 unless the batch-capacity or accumulation sections explicitly say otherwise.
## Precision contract
The native FSDP source runs intentionally load FP32 original parameters and use FSDP mixed precision to gather/cast BF16 working parameters for compute. Those FP32 originals are the optimizer masters; `default_dtype: torch.float32` is therefore expected and is not a benchmark mismatch. The fixed-arena prototype implements the same convention explicitly with a BF16 working parameter/gradient arena and sharded FP32 masters. A separately labeled Transformer Engine gate retains FP32 masters/moments but uses BF16 registered shards and gradients; it is a measured alternative precision layout, not part of the production baseline.
## Measurement and profiler stack
- End-to-end MFU is authoritative: FastVideo's inclusive `step_time_sec` is gathered across ranks, each step uses its slowest-rank time, and the reported result is the median after 10 warmup steps.
- The deferred-gradient-norm gate also records independent batch-fetch-to-batch-fetch wall intervals and closes the final interval with a CUDA synchronization. Its true-wall result is authoritative over the trainer's internal timer when checking for timer artifacts.
- PyTorch Profiler/Kineto traces provide operator and kernel attribution; exact-shape CUDA-event microbenchmarks gate individual attention, GEMM, quantization, and collective hypotheses. TorchInductor logs and controlled NCCL A/B/A runs cover compile and communication changes.
- W&B records training/validation observability but is not the source of truth for the slowest-rank MFU rows.
- `nsys-ai` was audited at upstream commit `d51180069d29439d56e7be2f2102f44886e00bb05` and is useful as an isolated offline reader for Nsight Systems SQLite exports. It is not a collector, the current GB200 host/container has no `nsys` binary, and its 2.250 PFLOP/s GB200 default conflicts with this report's accepted 2.450 PFLOP/s convention. Do not use its MFU output or add it as a FastVideo dependency; once a collector is available, analyze each rank separately and retain this harness as authoritative.
## One-GPU capacity boundary
The last one-rank audit, at `7f139e2b28610063d2f30526ba8f0ccae5d88944`, cannot complete the first optimizer step on one 184.31 GiB GB200 under the standard precision contract. Full activation checkpointing, `reshard_after_forward=true`, expandable allocator segments, explicit lazy CUDA module loading, and regional compile on/off all still OOM during FP32 Adam moment creation. A final dense-FA4 retry also unloaded the validation-only VAE before training and failed on a 64 MiB allocation with 182.83 GiB allocated and 3.50 MiB free. The measured optimization head was not re-gated on one GPU; optimizer offload or reduced-precision state would define a different baseline.
## Verdict
A whole-transformer mega-kernel is still **not** justified, but the remaining 50%-MFU gap is now kernel research rather than systems tuning.
- For unchanged vanilla BF16, a corrected measured-head fixed-arena prototype reaches 400.616610 ms / 36.054708% MFU and leaves 111.734 ms to the 288.882 ms 50%-MFU target. That is real architecture evidence, but it is benchmark-only: the 4x/B1/accumulation-1/LTX-specific scratch runtime adds 35.829 GiB peak allocation and omits DCP/export/EMA and generic lifecycle integration. Native FSDP2 symmetric-memory collectives, repeated/prefetched input, deferred gradient-norm materialization, non-final accumulation no-sync, opt-in between-microbatch parameter retention, singleton timestep embedding, persistent projection packing, and two-modules-per-group FSDP2 sharding are the production stack. At exact head `20c36acef`, grouping saves 7.969633 ms / 1.930% without checkpointing and 23.796931 ms / 3.063% with the committed full-checkpoint recipe. BF16 registered shards with TE FP32 masters save only 5.908 ms while adding memory and lifecycle complexity. Whole-step CUDA-graph capture remains blocked before replay.
- CUTLASS is closed for the current B2 shapes. The original B1 exact-shape gate found a 72.184321 ms arithmetic opportunity, but the unrestricted trainer gate was unsafe, an exact-name B2 trainer candidate regressed 2.143% while emitting 7,330 illegal-access warnings, and the final singleton-swizzle B2 A/X/B microgate was parity/safety clean but regressed 30.752773 ms / 7.8520%. No CUTLASS source or configuration is retained.
- For BF16-equivalent low-precision training, QuACK 0.5 clears only the prequantized arithmetic gate: exact-shape dispatch reaches 3.263x paired BF16 / 60.208 ms normalized. The complete sequential operator is a hard reject at 2,072.035 ms versus 250.557 ms paired BF16, or 1,624.429 ms normalized against the required <=63.463 ms.
- Quantization is not a small wrapper cost. After amortizing immutable text context, even an ideal fused dual-orientation quantizer has a 144.0 GiB / 19.3 ms traffic floor at an ideal 8 TB/s, already larger than the 3.255 ms prequantized margin unless production is overlapped under QMM execution or the arithmetic kernel gets faster.
- The complete smoke is also below a plausible quality gate: video/self-QKV/FFN forward and dgrad cosine is about 0.814-0.817 with 58.3-58.8% relative RMS; wgrad is about 0.991 / 13.4% relative RMS. Finite output alone is insufficient.
A fresh config-7 all-pass rerun independently confirms there is no hidden QuACK integration margin. Packed NVFP4 arithmetic was exact on the all-ones smoke and finite on every production shape, but sequential fprop+dgrad+wgrad measured 80.814 ms and the optimistic deferred/batched-wgrad lower bound measured 76.032 ms. They miss the absolute 63.463 ms break-even ceiling by 17.351 and 12.569 ms respectively before quantization, scale, cache, bias, or distributed costs.
The directly deployable prequantized result had only 3.255 ms of normalized arithmetic margin and missed the <=59.524 ms / 3.3x integration target by 0.684 ms. The complete measurement closes that candidate. Deferred batched wgrad remains an optimistic arithmetic lower bound because it changes activation lifetime and reduce-scatter overlap; it cannot rescue the measured quantization and quality failures by itself.
## Fresh-allocation control and phase decomposition
The fresh allocations differ materially, so only paired same-node deltas and ratios are used for attribution. The historical external fixed-arena result was 421.850613 ms / 34.239882%. A corrected measured-head same-allocation gate supersedes that headline: 400.616610 ms / 36.054708% versus a 413.472777 ms source-control midpoint, saving 12.856167 ms / 3.109314% and adding 1.119589 MFU points with 1.295174% control drift. The best matched pushed-source B1 gate remains packed commit `d01623709`: 414.823197 ms / 34.819931% versus a 427.432493 ms split-projection control midpoint. Measured head `20c36acef` enables packing and two-module FSDP2 groups in the overfit recipe. Batch scaling reaches 43.623053% MFU on 8x/B3 with packing; larger-batch rows change global batch and remain separate from B1 attribution.
- Uninstrumented fixed-arena control: 547.683096 ms wall, 547.350 ms GPU, 26.373125% MFU.
- Event-instrumented control: 545.251277 ms; the 2.432 ms difference validates that the phase instrumentation does not materially perturb the step.
- Median event decomposition: data 0.221 ms; exposed all-gather wait 0.202 ms; forward compute 160.339 ms; backward plus overlapped reduce-scatter 345.321 ms; exposed reduce-scatter wait 3.637 ms; grad-shard copy 3.578 ms; norm/clip 6.746 ms; AdamW 14.678 ms; BF16 copy/all-gather launch 3.171 ms; GPU other 7.272 ms; host other 0.087 ms.
Only 3.838 ms is exposed collective waiting. The large backward interval is occupied by compute plus overlapped communication; it is not a 133 ms idle bubble that a graph or launch fusion can erase.
## Input-pipeline productionization
The original four-row overfit fixture exhausted each rank's epoch every step, rebuilding the stateful iterator and leaving input work on the critical path. The dataset already supports virtual path repetition, so no source abstraction or physical data copy is needed.
| 4x exact-source path | Median slowest-rank wall | MFU | Delta |
|---|---:|---:|---:|
| original four-row fixture, worker 0 | 493.569924 ms | 29.264577% | control |
| virtual repeat `:32`, worker 1 | 466.474781 ms | 30.964407% | -27.095143 ms / -5.489% |
Production commit `7c58950a9` configures the 300-step LTX-2 overfit recipe with a virtual repeat count of 300 and one worker, giving each of four ranks exactly 300 samples without expanding the 12 MB fixture. It also removes the nightly launcher's redundant scalar `data_path` override so the mapping survives. The GB200 environment parsed the exact committed YAML as `{'data/ltx2_overfit_preprocessed': 300}` with one worker, and the existing structured-path config test passed. The performance result is from exact source `7f139e2b2`; later commits do not touch the dataset path.
## Deferred gradient-norm materialization
Gradient clipping already computes the norm and enqueues the clipping work before the optimizer, but converting that CUDA scalar to a Python float forced the host to wait before fused AdamW could launch. A scratch-only A/B/A kept clipping and per-step logging intact while retaining the norm tensor until after the optimizer. Both the original trainer timer and an independent synchronized wall timer were recorded:
| 4x FA4 run | True-wall slowest-rank median | MFU | Matched result |
|---|---:|---:|---:|
| control A | 460.133193 ms | 31.391161% | control |
| scratch deferred candidate | 435.467763 ms | 33.169195% | candidate |
| control B | 459.621117 ms | 31.426134% | control |
| control midpoint | 459.877155 ms | 31.408638% | **-24.409393 ms / -5.3078%** candidate saving |
The candidate won all 20 measured true-wall steps, and 99.65% of its internal-timer delta survived the independent timer, ruling out a bookkeeping artifact. Production commit `7f6c290c9` implements the same idea in only three files: the clipping helper returns the device tensor, `GradNormClipCallback` queues it, and the existing `on_training_step_end` hook clears then materializes/logs it after the optimizer and zero-grad path. No new trainer lifecycle or configuration surface is added.
The exact committed-source confirmation selected FA4 and reached **431.193091 ms / 33.498021% MFU / 9.276586 samples/s**, saving **28.684064 ms / 6.2373%** and adding **2.08938 percentage points MFU** against the same two-control midpoint. Every rank recorded exactly 30 clips, 30 optimizer calls, and 30 gradient-norm logs; every norm was materialized after the optimizer. The three-file change passed 49 focused tests in the GB200 environment and finished review-clean. Exact-source log SHA-256: `87347668ad28ab200bf3540da01ff1c7744854c2ce2a3a0a7d1c0cb1d1788d01`.
## Singleton timestep memory gate
Plain LTX-2 T2V uses one uniform sigma per sample. At SP=1, AdaLN broadcasts one timestep embedding across all 4,290 video tokens, so commit `49508050b` passes `[B, 1]` instead of eagerly materializing `[B, 4290]`; SP>1 still expands before sequence sharding. The proof recorded 4,290 semantic tokens and one model timestep token on every rank, with finite AdaLN gradient probes. Its proof script SHA-256 is `05836dda4426309f23e1c797281f0e69cc343c522006679fef6b34e6c7d3f865`, and proof log SHA-256 is `efe7e529860771bc3564adf82b2cfa2a712dc60cf83691b174769aff11760317`.
The 4x GB200 A/X/B timing is deliberately classified as neutral. Controls measured 435.652742 and 440.625450 ms, a 438.139096 ms midpoint with 1.14% control drift; the candidate measured 434.923236 ms. The corresponding true-mean and measured-window changes were only about 0.1%. Peak allocated memory fell **13.0057 GiB / 11.416%**, and peak reserved memory fell **13.6719 GiB / 10.043%**, so the change is retained for capacity rather than timing credit. Raw log SHA-256 values are `983b52f96ceace4ddd1fa01c36f6b851c9bb74fd6c4033e8932bc99abf9f1c3a` (A), `189f0ffc65ff4ae2a449990c026a385a79222e45e95ade94977bbb981c23f70a` (X), and `505d0062845b2c1662384ab7f5e60acdebd4e6971ac66bcb521b7ab0b2bb41f5` (B).
The singleton path passed seven lightweight tests plus real LTX-2.0 and LTX-2.3 model tests in the GB200 environment. Current head passed its 80-test focused suite in 3.92 s, parsed the committed packed recipe with `PACKED_RECIPE_OK`, and remained pre-commit/review clean.
## Current-stack batch capacity
The memory reduction makes larger local batches profitable on the cumulative stack. These split-projection accumulation-1 runs were frozen at `49508050b` and inherited through `002ec0771`; the packing gate below supersedes their current timing. The harness's printed throughput/MFU numerator was fixed at B1, so the B2/B3 values below are recomputed from the measured true-wall median and the actual samples per step.
| Topology | Local batch/GPU | Median step | Aggregate throughput | MFU | Peak allocated / reserved per rank |
|---|---:|---:|---:|---:|---:|
| 4x B1 midpoint | 1 | 0.4246335145 s | 9.419888 samples/s | 34.015485% | - |
| 4x B2 | 2 | 0.7105401774 s | 11.259040 samples/s | 40.656716% | 140.496 / 163.723 GiB |
| 8x B1 midpoint | 1 | 0.4165137762 s | 19.207048 samples/s | 34.678601% | - |
| 8x B2 | 2 | 0.6963156550 s | 22.978085 samples/s | 41.487262% | 122.319 / 135.307 GiB |
| 8x B2 second-gate midpoint | 2 | 0.6996195166 s | 22.869574 samples/s | 41.291344% | - |
| 8x B3 | 3 | 0.9994570174 s | 24.013039 samples/s | 43.355887% | 160.862 / 175.172 GiB |
| 8x B4, forward reshard enabled | 4 | OOM before step 1 | - | - | 177.86 GiB allocated / 184.17 GiB total in use |
B2 gains 19.524% throughput / 6.641 MFU points on 4x and 19.634% / 6.809 points on 8x. B3 adds 5.000% throughput / 2.065 points over its 8x B2 midpoint. B4 does not fit even after enabling forward resharding, closing the capacity ladder at B3. These configurations increase global batch, so convergence and LR scaling remain separate gates.
An exact-final-head equal-global-batch gate at `20c36acef` removes that confound by comparing B1/gradient-accumulation-2 with B2/gradient-accumulation-1; both process eight samples per optimizer step. The B1 controls measured 0.7426303995 and 0.7441010120 seconds, a 0.7433657057-second midpoint / 10.761874 samples/s / 38.861435% MFU. B2 measured **0.7078659144 seconds / 11.301575 samples/s / 40.810314% MFU**, saving **35.499791 ms / 4.775549%**, raising throughput **5.014942%**, and adding **1.948878 MFU points** with only 0.198028% control drift. B2 peak allocation/reservation was 140.865/166.404 GiB, 16.109/15.293 GiB above the B1 midpoint. All runs covered 927 FP32 parameters / 13,041,520,768 elements, retained FP32 Adam moments, recorded zero master-writeback mismatches, and had finite sampled AdaLN gradients on all ranks. This is an efficiency result at unchanged optimizer batch; stochastic-run gradient hashes are not a parity claim.
The exact B2/group-2 Kineto trace is diagnostic rather than MFU evidence, but it closes the remaining systems hypothesis. Its two profiled steps were 97.2415% GPU-busy; 120.8658 of 131.3196 ms/step communication was overlapped (92.0394%), leaving only 10.4538 ms exposed. The per-step self-CUDA bands were 389.300 ms for GEMMs, 133.329 ms for FA4, and 14.134 ms for fused AdamW. The compute-event union alone was 688.602 ms, already 110.838 ms above the 577.765 ms 50%-MFU target. Further large gains must accelerate kernels, principally GEMMs and then attention; host/collective tuning cannot close this gap.
## Persistent LTX projection packing
Commit `d01623709` adds opt-in persistent packing for the video path: self-attention Q/K/V use one `to_qkv`, and text cross-attention K/V use one `to_kv`. Across 48 blocks this removes 432 projection GEMM launches over forward, dgrad, and wgrad and reduces parameter objects from 1,215 to 927 without changing the 13,041,520,768 total parameter elements. The loader records every source key and split size so DCP/Diffusers export restores strict split HF keys. Audio/cross-modal paths are unchanged; linear quantization and enabled LoRA are rejected. Current head `20c36acef` enables the otherwise-default-false option in the overfit recipe.
| Gate | Split control A | Packed | Split control B | Control midpoint | Packed delta |
|---|---:|---:|---:|---:|---:|
| 4x/B1 step | 0.4284051351 s / 33.716017% | 0.4148231971 s / 34.819931% | 0.4264598510 s / 33.869812% | 0.4274324930 s | -12.609296 ms / -2.950009%; +1.027191 pp |
| 8x/two-tray B3 step | 0.998659810983 s / 43.390496% | 0.993335918058 s / 43.623053% | 0.999399903463 s / 43.358364% | 0.999029857223 s | -5.693939 ms / -0.569947%; +0.248622 pp |
The 4x candidate reaches 9.642662 samples/s (+3.039680%) with 0.455109% control drift; the 8x candidate reaches 24.161011 samples/s with 0.074081% drift. Both are memory-neutral and completed 30 forward/backward steps with finite loss/norms, strict split-HF loading, FP32 registered masters/moments, and exact optimizer coverage. The final-head focused suite passed 80 tests in 3.92 s; a separate two-tray run passed 3 split/packed SP-gradient and fused-export tests in 150.97 s.
A real FP32 production export/reload gate printed `REAL_PACK_EXPORT_RELOAD_OK`. It wrote 52,166,223,624 bytes and 1,215 HF keys; all 192 merged internal tensors became 480 split projection tensors bit-exact against their sources. Strict split reload restored 1,215 parameter objects from the packed model's 927 while preserving 13,041,520,768 total parameter elements. Final log/script SHA-256: `c3ec531c9cf00c569be30ce1460e224225c617e17e8224af7385b9a82955f08c` / `2a4ba1280710c026fefdcd188ce0207f99a5828c277dcaf265758a8e46c06367`.
### Rejected regional max-autotune mode
A packed 4x/B1 A/B/A tested `max-autotune-no-cudagraphs` without changing source. Default controls measured 0.415032553021 s / 34.802366% MFU and 0.412693266000 s / 34.999638%; their midpoint was 0.413862909511 s / 34.901002% with 0.565232343% drift. Max-autotune measured 0.416518588027 s / 34.678200%, a **2.655679 ms / 0.641680724% regression and -0.222802 MFU points**. Triton repeatedly discarded invalid choices requiring 262,160 registers above the SM100 hardware limit of 232,448, then fell back. That regional gate did **not** enable TorchInductor's CUTLASS backend, so it does not contradict the separate CUTLASS measurements below. Keep the regional compile default; both alternatives are now rejected.
### Attention-compile environment diagnostic
`FASTVIDEO_DISABLE_ATTENTION_COMPILE=0` does not change LTX execution: its `LocalAttention.forward` and `DistributedAttention.forward` overrides bypass the decorated base forward. The packed 4x A/candidate/B medians were 0.415846481 / 0.413138084 / 0.413302450 s. The apparent -1.436381 ms / -0.346471% candidate delta is inside 0.613649% control drift and cannot be attributed to the environment variable. This is a confirmed no-op; make no source or configuration change.
## Raw velocity benchmark-only rejection
A scratch path returning raw velocity avoids the BF16 x0-to-velocity round trip and numerically improves reconstruction error at sigma=1e-3, but it has no performance or capacity case. Its true-wall A/X/B medians were 540.702049 / 542.334403 / 542.752936 ms: the candidate is 0.607 ms / 0.112% slower than the 541.727492 ms control midpoint, with no memory reduction. It remains benchmark-only and is not in the source stack. The round-trip proof script SHA-256 is `f6924a9277723fe51c2d210f5ea3fb66aafbb845d142ee50ef92effd1cb5fa4d`; raw log SHA-256 values are `83ebfafb10c4c9f10011b5cf746b7c6fdde6414d8247774ade8013f4a57a2f78` (A), `b421b2e0bcb1c9b9ac8f4871f3ad41587249a1f520faed6795fc0b10319c0c45` (X), and `6a1875640a2528a8bd126a9037138f03d34a06bf188e1fc10e06005c26cab514` (B).
## NCCL occupancy gate
`NCCL_MAX_CTAS=16` was first tested with the identical 10-warmup/20-measure workload on the same allocation:
| Run | Median slowest-rank wall | MFU | Delta |
|---|---:|---:|---:|
| paired control | 547.683096 ms | 26.373125% | control |
| `NCCL_MAX_CTAS=16` | 546.032818 ms | 26.452833% | -1.650278 ms / 0.301% |
That historical 0.301% single pair predated the final symmetric-memory stack and lacked a closing control. A fresh A/B/A on clean, frozen then-head `7c58950a9`, with FA4, regional compile, no forward reshard, symmetric memory, and BF16 reductions, reversed the sign:
| Frozen-head 4x run | Median slowest-rank wall | MFU | Delta |
|---|---:|---:|---:|
| control A, `NCCL_MAX_CTAS` unset | 558.356668 ms | 25.868976% | control |
| `NCCL_MAX_CTAS=16` | 559.553032 ms | 25.813666% | candidate |
| control B, `NCCL_MAX_CTAS` unset | 559.054499 ms | 25.836685% | control |
| control midpoint | 558.705584 ms | 25.852820% | +0.847448 ms / +0.151681% candidate regression |
All 30 losses and gradient norms were finite in every row, and control drift was only 0.124980%. Do not add the setting: it is withheld because the repeatable frozen-head result is negative, not because the effect is small. Raw log hashes are `00cc6b6632185f1f378996fdda9e62bd2ce285e57042b1c746a178a2c105b1e9` (A), `cf7351572e7ff7d585cf8ac3e4a3ee0463c03dd7ae0f9051f1423669bc85d4fd` (candidate), and `5dd6d20cbeafb70b9f62c4f37ae4391ca5937cb570f96d58a369610ac98d2ecf` (B).
## Native FSDP2 symmetric-memory gate
PyTorch 2.12 symmetric-memory communication was enabled after sharding on all 49 FSDP2 groups (root plus 48 blocks), with forced SUM reductions and `NCCL_CTA_POLICY=2` set before process-group initialization. This enables the native symmetric-memory all-gather and reduce-scatter paths and makes Copy Engine all-gather eligible on the single-node NVLink domain.
The 10-warmup/20-measure A/B/A sandwich used the same source, input, and slowest-rank timing contract:
| Run | Median slowest-rank wall | MFU | Delta |
|---|---:|---:|---:|
| baseline A | 573.191594 ms | 25.199454% | control |
| symmetric-memory AG/RS | 554.638125 ms | 26.042413% | candidate |
| baseline B | 571.166225 ms | 25.288812% | control |
| baseline midpoint | 572.178910 ms | 25.244053% | -17.540784 ms / 3.066% |
The candidate completed normally and printed proof that all 49 groups were configured. Its 17.541 ms saving established the first repeatable signal. The 30 ms threshold is retained only as a research-priority boundary, not as a rule against sound incremental optimizations.
Production commit `e42cfa5e5` adds strict `training.distributed.fsdp_symmetric_memory`, installs `NCCL_CTA_POLICY=2` before process-group initialization, forces SUM reductions, and enables native symmetric-memory communication on every transformer FSDP wrapper. Candidates deliberately launched without the environment variable and configured all 49 wrappers.
| Topology / policy | Control midpoint | Production candidate | MFU | Delta |
|---|---:|---:|---:|---:|
| 4x FULL_SHARD | 594.420553 ms | 571.917264 ms | 25.255602% | -22.503288 ms / -3.786% |
| 4x SHARD_GRAD_OP | 571.403084 ms | 553.293423 ms | 26.105705% | -18.109661 ms / -3.169% |
| 8x/two-tray FULL_SHARD | 471.689844 ms | 460.362467 ms | 31.375527% | -11.327377 ms / -2.401% |
| 8x/two-tray SHARD_GRAD_OP | 448.385441 ms | 445.343413 ms | 32.433656% | -3.042028 ms / -0.678% |
The two-tray allocation was verified as one eight-GPU MNNVL clique. A direct two-rank NCCL all-reduce test found equal large-message performance inside and across trays: 514.923 versus 513.070 GB/s at 256 MiB, and 577.143 versus 577.151 GB/s at 1 GiB. Cross-tray overhead was visible only at small sizes (1 MiB: 82.97 versus 72.73 us; 16 MiB: 76.68 versus 70.83 us).
## Conventional FSDP follow-up gates
The research-priority threshold is not an acceptance threshold. Sound, composable optimizations are kept even when their individual savings are smaller. Four additional public FSDP2 paths were tested on the exact production stack.
The accumulation-1 4x control sandwich measured 555.859681 and 551.755290 ms, or a 553.807485 ms midpoint:
| Path | Median slowest-rank wall | Delta from midpoint | Decision |
|---|---:|---:|---|
| explicit one-module forward/backward prefetch | 554.107704 ms | +0.300219 ms / +0.054% | neutral; do not add |
| explicit two-module forward/backward prefetch | 552.395538 ms | -1.411947 ms / -0.255% | drift-sensitive; do not add |
| process-group communication allocator, replacing symmetric memory | 569.133412 ms | +15.325927 ms / +2.767% | slower and mutually exclusive; reject |
On the 8x/two-tray allocation, existing HSDP `replicate=2, shard=4` maps each four-way shard group within a tray and the replicate pair across trays. It measured 465.559134 ms /31.025307% MFU versus 442.891026 ms /32.613248% for the eight-way shard, a 22.668108 ms /5.118% regression. The extra replicate all-reduce outweighs the smaller intra-tray AG/RS groups, so the existing 1x8 layout remains preferred.
A Kineto trace also exposed 49 size-one `mesh_replicate` all-reduces because a `(replicate=1, shard=4)` device mesh selects PyTorch's HSDP code path. The FP32-era trace attributed 12.146 GiB and 13.660-16.465 ms summed per rank to those events. A standards-aligned 1-D FSDP mesh removed the degenerate topology and passed 8 focused unit tests plus an exact old-2-D-to-new-1-D four-rank DCP model/AdamW resume smoke, but the current BF16/symmetric-memory A/B/A timing was neutral:
| 4x mesh path | Median slowest-rank wall | MFU | Matched result |
|---|---:|---:|---:|
| historical 2-D control A | 454.177486 ms | 31.802798% | control |
| 1-D FSDP candidate | 457.564300 ms | 31.567399% | +0.158553 ms / +0.034664% vs control midpoint |
| historical 2-D control B | 460.634008 ms | 31.357031% | control |
| historical control midpoint | 457.405747 ms | 31.578342% | midpoint |
The source candidate is not retained. Removing a trace-visible no-op is not credited as a speed optimization when the current production sandwich is indistinguishable from drift.
Gradient accumulation exposed one accepted general optimization. Native FSDP2 was previously synchronizing every microstep. Production commit `0b324d0a4` now uses `set_requires_gradient_sync(False)` and `set_is_last_backward(False)` on non-final microsteps, restoring synchronization only for the final backward. The code path is skipped entirely when accumulation is 1, so headline MFU is unchanged.
The matched 4x accumulation-2 sandwich used 30 optimizer steps, 10 warmup + 20 measured, and the same standard FP32-master/BF16-working stack:
| Path | Median slowest-rank wall | Effective MFU | Global samples/s |
|---|---:|---:|---:|
| no-sync candidate A | 1.069807470 s | 27.003205% | 7.477981 |
| forced synchronization every microstep | 1.090327118 s | 26.495012% | 7.337248 |
| no-sync candidate B | 1.074732777 s | 26.879454% | 7.443711 |
| no-sync midpoint | 1.072270124 s | 26.941187% | 7.460807 |
Suppressing the redundant first reduction saves **18.056994 ms / 1.656%**, raises effective MFU by 0.446175 points, and recorded exactly 30 non-sync plus 30 final-sync backwards in each candidate. Focused trainer tests pass in the GB200 environment (3 passed); accumulation-1 performs no additional dispatch.
Commit `002ec0771` then uses the public `set_reshard_after_backward` control to retain already-unsharded parameters across non-final backwards. This path is deliberately gated by the existing `reshard_after_forward: false` opt-in; default-policy jobs never touch backward reshard state, and the final backward restores resharding before optimization.
| 4x accumulation-2, no-sync in both paths | Median optimizer step | Effective MFU | Global samples/s |
|---|---:|---:|---:|
| force backward reshard, control A/B midpoint | 1.0214430510 s | 28.281782% | 7.83206 |
| retain params between microsteps | 1.0055416964 s | 28.729023% | 7.95591 |
This saves **15.901355 ms / 1.55675%**, adds **0.447240 MFU points**, and raises throughput **1.58137%**. Control drift was 0.12197%. Every run recorded 30 non-final plus 30 final sync calls; controls forced 30 backward reshards and the candidate forced none. Peak allocation did not increase. The two-file change passed pre-commit, focused GB200 tests (3 passed), and independent review.
## BF16 registered shards with FP32 TE masters
A source-built Transformer Engine 2.16 scratch path tested whether persistent BF16 FSDP shards can remove enough FP32-to-BF16 pack work while retaining sharded FP32 master weights and FP32 Adam moments. The stock/TE/stock 8x B1 sandwich used the same frozen `49508050b` source plus one optimizer-only scratch patch, dense FA4, singleton timesteps, BF16 reductions, symmetric-memory collectives, and 10-warmup/20-measure true-wall timing.
| 8x B1 path | Median slowest-rank wall | MFU | Global samples/s | Peak allocated / reserved |
|---|---:|---:|---:|---:|
| stock FP32 registered/master control A | 0.415380163 s | 34.773242% | 19.259466 | 82.646 / 96.264 GiB |
| TE BF16 registered + FP32 master candidate | 0.409804175 s | 35.246383% | 19.521519 | 85.666 / 93.465 GiB |
| stock FP32 registered/master control B | 0.416044141 s | 34.717746% | 19.228729 | 82.651 / 96.184 GiB |
| stock control midpoint | 0.415712152 s | 34.745472% | 19.244085 | 82.648 / 96.224 GiB |
The candidate saves **5.907977 ms / 1.421%**, raises throughput **1.442%**, and adds **0.500911 MFU points** with only 0.160% control drift. The strengthened post-timing probe checked all 1,215 trainable parameters on every rank: registered parameters were BF16 DTensors; `master_param`, `exp_avg`, and `exp_avg_sq` were FP32 DTensors with matching mesh, placements, shapes, and strides; optimizer coverage was exact; and every registered shard exactly equaled its FP32 master rounded to BF16. Losses and sampled gradients were finite.
This is not a good production trade despite the real small speedup. Peak allocation rises **3.018 GiB/rank**, registered gradients become BF16 instead of FP32, and the checkpoint has 194 FP32 tensors / 3,256,320 elements that the scratch BF16 load rounds before TE creates its masters. TE 2.16 also has no compatible prebuilt binding for this Torch 2.12/CUDA 13 environment, so it required a source build; production support would additionally need TE-specific resume initialization, FP32-master Diffusers export, checkpoint-backend compatibility checks, and guards or support for EMA and validation optimizer-state offload. The stock FP32 registered/master path remains the production convention. No TE dependency or source path is added.
## Whole-step CUDA-graph localization
Moving the complete static batch to CUDA before capture and replacing the sigma sampler's device-created constants with Python scalars advanced capture to `VideoLatentPatchifier.get_patch_grid_bounds`, where `torch.tensor(self._patch_size, device=patch_starts.device)` attempted a CPU-to-CUDA copy during capture.
A scratch-only patchifier override pre-staged that constant and passed bit-exact eager equivalence on the real shape:
```json
{"equal": true, "max_error": 0, "shape": [1, 3, 4290, 2]}
```
Capture then advanced to `_get_pixel_coords`, where another `torch.tensor(scale_factors, device=latent_coords.device)` caused the same error on all four ranks. The bounded retry stopped at this third distinct graph blocker. No graph was created or replayed, so graph latency and MFU credit remain unavailable. The repeated pattern suggests source graphability can be fixed systematically by pre-registering immutable device constants, but it does not establish a material speedup.
## Rejected FlashAttention fused-dense GELU gate
The fused-dense extension from pinned FlashAttention `82d6441` built and loaded on SM100 with CUDA 13, and executable heuristic h0 passed output and gradient parity for cuBLASLt `GELU_AUX_BIAS` plus `DGELU_BGRAD`. It was substantially slower than the unfused reference:
| Local batch | Reference | h0 fused | Reference - fused/block | Projected over 48 blocks | Fused regression |
|---|---:|---:|---:|---:|---:|
| B1 | 0.895942 ms | 1.615837 ms | -0.719894 ms | -34.554927 ms | 80.35% |
| B3 | 2.637349 ms | 4.774322 ms | -2.136973 ms | -102.574704 ms | 81.03% |
h0 also added 32 MiB peak allocation. Heuristic h1 took 15.566678 ms/block at B1 and 46.050705 ms/block at B3; h2-h4 fail `bias_act_linear_dgrad_bgrad`. The exact build and parity gate therefore reject this path before trainer integration. No source integration or FlashAttention dependency change is made.
## Existing SM100 all-pass projection gate
The node's installed packages are structurally skewed: QuACK 0.6.1 pins released `nvidia-cutlass-dsl==4.6.0`, while the FA4 environment contains `4.6.0.dev0`. Runtime shims reached incompatible shared-memory and pipeline APIs and were abandoned. A PYTHONPATH-only QuACK 0.5 shadow is compatible with the installed DSL after two moved-type aliases (`cute.core.ThrMma` and `cute.core.ThrCopy`); the environment was not installed into or modified.
The correctness smoke used packed NVFP4 inputs, BF16 output, batch `L=2`, and exact all-ones arithmetic. It passed with `max_abs=0.0`. Every exact production shape produced finite output.
The harness covers the seven valid packed projection GEMMs per block, all 48 blocks, and forward + dgrad + wgrad: 336 occurrences per phase / 1,008 GEMMs total / 300.096 TFLOP. Unique-shape latency is multiplied by its exact 48 or 144 occurrence count.
| QuACK tile / cluster | Paired BF16 | NVFP4 sequential | Speedup | Deferred-wgrad lower bound | Speedup |
|---|---:|---:|---:|---:|---:|
| config 5, 256x128 / 2x1 | 256.934 ms | 93.723 ms | 2.741x | 90.439 ms | 2.841x |
| config 6, 256x192 / 2x1 | 257.331 ms | 83.220 ms | 3.092x | 79.049 ms | 3.255x |
| config 7, 256x256 / 2x1 | 256.380 ms | 81.288 ms | 3.154x | 76.422 ms | 3.355x |
| best exact-shape dispatch, configs 5/6/7 | 257.025 ms | 78.781 ms | 3.263x | 73.916 ms | 3.477x |
Config 7 phase totals are 23.894 ms forward, 23.721 ms dgrad, and 33.673 ms sequential wgrad. The slow shapes are wgrad: text KV is 2.349x, self-QKV is 2.560x, video DD is 2.595x, and the two FFN wgrads are about 2.68x. Cross-block batched wgrad improves the aggregate, but even its self-QKV/FFN cases remain about 2.93x.
Config 6 improves most forward/dgrad shapes while config 7 remains best for every wgrad shape; config 5 wins only text dgrad. Exact-shape dispatch therefore lowers the sequential total to 78.781 ms without changing the operator boundary or requiring a new kernel. No 1CTA sweep was run.
Normalizing the paired ratios to the historical exact packed BF16 pass:
- sequential best-of: `196.431 / 3.262522 = 60.208 ms`, passing the 63.463 ms bare break-even by 3.255 ms but missing the 59.524 ms integration target by 0.684 ms;
- deferred best-of: `196.431 / 3.477274 = 56.490 ms`, passing the integration target by 3.034 ms, but not deployable without proving memory and communication overlap.
These are prequantized arithmetic lower bounds. They exclude quantization, global scale, bias/postscale epilogues, weight-cache refresh, and integration overhead. They justify an operator prototype, not a 50% end-to-end claim.
The later all-pass gate reran config 7 without ratio normalization and included exact single-shape plus batched-wgrad correctness checks. BF16 measured 256.166 ms. Sequential NVFP4 measured 80.814 ms (23.711 forward, 23.518 dgrad, 33.585 wgrad; 3.170x), while deferred batched wgrad reduced the total to 76.032 ms (3.369x). Both fail the absolute 63.463 ms break-even bound; the deferred result is still only an optimistic lower bound because it excludes every packing/epilogue cost and changes gradient lifetime. This closes the all-pass library-kernel route independently of the much slower complete-operator result below.
### Complete projection operator result
The complete scratch gate preserved the standard training contract (resident FP32 masters/moments, BF16 working tensors), refreshed both weight orientations, quantized row/transposed activation and output-gradient layouts, ran forward/dgrad/wgrad, and counted postscale, BF16 output, bias, and dbias. Immutable text-context packs were charged once across 48 blocks.
| Metric | Result | Gate |
|---|---:|---:|
| Paired bare BF16 GEMMs | 250.556930 ms | reference |
| Complete BF16 GEMMs + bias/dbias | 292.985852 ms | reference |
| Complete QuACK NVFP4 | 2,072.035073 ms | <=80.950 ms same-run / <=63.463 ms normalized |
| Ratio-normalized complete time | 1,624.428915 ms | <=63.463 ms |
| Speedup vs bare BF16 | 0.120923x | >=3.095x |
This is a hard reject for the current separate quantize/transpose/pack/epilogue path. It is about 8.27x slower than paired BF16, before autograd dispatch or distributed overlap.
## Current-head follow-up gates
### Repeated no-autocast RMSNorm
The repeated 4x A/X/B gate measured 418.399114 / 415.219851 / 413.452594 ms. The 415.925854 ms control midpoint leaves a 0.706003 ms / 0.1697% candidate saving inside 1.189% control drift. Native unweighted and weighted semantics were exact, BF16 output was preserved, the refine path still delegated, and the learned weight remained effective. Reject the source change as neutral.
### Grouped BF16 weight gradients
Exact packed LTX-2 shapes across all 48 blocks were bit-identical (`max_abs=0`), but grouping adjacent wgrad GEMMs cannot clear the 10 ms integration gate. Groups of 2/4 saved only 2.512/2.073 ms under the optimistic prepacked arithmetic ceiling; staging separate inputs regressed 11.199/11.436 ms and staging plus gradient scatter regressed 19.859/20.566 ms. This already excludes autograd hooks, buffer lifetime, and likely lost FSDP reduce-scatter overlap. Reject.
### Short-text attention dispatch
FA4 remains faster for LTX-2's short text attention: at B1 it measured 0.416 ms versus 0.809 ms FA2 and 0.852 ms SDPA; at B3 it measured 0.825 ms versus 2.341 ms FA2 and 2.468 ms SDPA. All paths were finite and passed the declared parity checks. Reject a short-text FA2/SDPA hybrid dispatch.
### QuACK 0.6.1 fused norms
The QuACK fused AdaLN and joint QK-normalization experiment regressed the projected 48-block path by 30.274514 ms on tray 0 and 31.686422 ms on tray 1. It also provided only loose BF16 parity (AdaLN output 64.3% exact, max absolute 0.03125; scale-gradient 49.7% exact, max absolute 1.0). Reject the dependency and integration.
### Fused clipping through AdamW
Passing the reciprocal clip coefficient to fused AdamW preserved the existing distributed norm and deferred logging, supported role-mapped DMD2 optimizers, and was bit-exact for parameters, moments, and step tensors on both trays. Timing did not replicate: tray 0 saved 3.263339 ms, tray 1 regressed 0.299465 ms, and the equal-tray aggregate saved only 1.481937 ms / 0.360336% (+0.127012 MFU points). The signs conflict, so reject it as neutral; fused AdamW's scaled-gradient writeback plausibly absorbs the removed foreach pass.
### Corrected current-head fixed arena
The 4x A/X/B source controls measured 416.150371 and 410.795182 ms, for a 413.472777 ms midpoint with 1.295174% drift. The fixed-arena candidate measured **400.616610 ms**, saving **12.856167 ms / 3.109314%** and moving MFU **34.935119% -> 36.054708% (+1.119589 points)**. It keeps FP32 master weights and Adam moments with a BF16 working parameter/gradient arena, covers all 927 packed parameter objects, reports zero sampled replica error, and measures a 0.015617 maximum master-precision delta. Peak allocation rises **100.832486 -> 136.661506 GiB/rank (+35.829020 GiB)**.
This is a real architecture result, not production source. The scratch runtime is limited to world size 4, accumulation 1, and LTX-specific bucket/layer ordering, and it lacks DCP/export, EMA, and generic optimizer/checkpoint lifecycle support. Keep it as the next systems architecture reference; do not fold it into PR #1630 without those gates.
### Purpose-built Triton norms
The synthetic wall projection appeared to save 2.244263 ms, but GPU kernel sums show the candidate is slower than the current compiled path: AdaLN 0.192063 versus 0.081024 ms and joint QK norm 0.353639 versus 0.157087 ms. AdaLN parity was also loose (57.4% exact output, max absolute 0.0625; scale-gradient max absolute 1.0). The wall signal is a host/autograd artifact, so reject the kernels and claim no MFU credit.
### Public FSDP2 boundary grouping
A root-only public `fully_shard` boundary is a hard reject. Its 4x A/X/B controls measured 415.065887 and 410.919773 ms, or a 412.992830 ms midpoint, while the root-only candidate measured 466.765127 ms. That is a **53.772297 ms / 13.02% regression**, moving MFU from about **34.975% to 30.945%**. Peak allocation rose by about 23.8 GiB/rank and peak reservation by about 56 GiB/rank because the entire transformer was materialized as one unit and communication overlap was lost.
Grouping consecutive block modules through PyTorch's public FSDP2 list API retains per-block decoration while reducing collective launches. The exploratory A/G2/G4/B gate chose **two modules per group** because it was faster and used less memory than four. Both final four-GB200 A/X/B gates used exact clean head `20c36acef`:
| Exact-head gate | Control A | Group-2 candidate | Control B | Control midpoint | Matched result |
|---|---:|---:|---:|---:|---:|
| optimized, no checkpointing | 414.894377 ms | **404.962008 ms / 35.667827% MFU** | 410.968906 ms | 412.931642 ms | **-7.969633 ms / -1.930%; about +0.688 pp** |
| committed recipe, full checkpointing | 777.312431 ms | **753.134414 ms / 19.178668% MFU** | 776.550259 ms | 776.931345 ms | **-23.796931 ms / -3.063%; about +0.587 pp** |
The optimized candidate added 0.500 GiB peak allocation and 2.063 GiB reservation. The full-checkpoint recipe added no peak allocation and 3.719 GiB reservation; its controls differed by only 0.098%. Both gates completed with finite losses and gradients while proving FP32 registered weights and Adam moments. The proof covers all 13,041,520,768 parameter elements, retains 48 decorated blocks in 24 block groups plus the 152,096,896-element root group, and uses BF16 working parameters, gradients, and reductions.
An exact 8x/two-tray B3 A/X/B on replacement allocation `1629519` measured 1.418004981/1.421743279/1.421487014 seconds for control A/group 2/control B. Against the 1.419745998-second / 30.521241%-MFU control midpoint, group 2 is neutral at +1.997281 ms / +0.140679% and -0.042922 MFU points, smaller than the 0.245257% control drift. The native eight-rank health gate sustained 463.195 GB/s derived bus bandwidth with MNNVL/NVLS enabled and no NCCL errors. All paths completed 30 steps with finite gradients and FP32 registered weights/moments. Absolute performance on this replacement pair was about 30.5% MFU, so it is not compared across allocations with the earlier 43.623053% pair; only the same-pair bracket is attribution evidence.
Current source exposes the generic positive `training.distributed.fsdp_modules_per_group` setting and selects `2` in the LTX-2 overfit recipe. Before the full timing gate, that exact recipe also passed a two-step smoke with full activation checkpointing. Smoke log SHA-256: `66cb9fe7d6b95d22319ea772dafa54cde434095fa6d02c184a46b474a69056ae`; runner SHA-256: `4ccb2cc9c163357180c673daf46213be02172a4a01fb6048489d8dd937d1b593`.
### CUTLASS exact-shape BF16 GEMMs
The packed LTX-2 projection mix contains five unique shapes and 15 forward/dgrad/wgrad cases, weighted to 1,008 GEMM calls and 300.095807 TFLOP per optimizer step. A cold-cache A/X/B microgate measured 280.932863 ms for current kernels, 212.677120 ms for the ATen+CUTLASS candidate, and 288.790019 ms for the closing current-kernel control. Against the **284.861441 ms** midpoint, the candidate saves **72.184321 ms / 25.3402%** and raises effective projection throughput from 1.053480 to 1.411039 PFLOP/s.
All 15 parity checks passed. Runtime proof found actual CUTLASS kernels in 8 of the 15 shape/phase cases, with the remaining winners using current NVJet kernels. The one-time cold compilation cost was 948.13 seconds; it is excluded from steady-state timing and requires persistent compiler caching in production.
The unrestricted end-to-end attempt is a hard safety reject, not a timing result. Four independent rank-local caches emitted 8,757 illegal-memory-access warnings, 105 CUTLASS errors, and 343 failed autotune choices, then produced no `BF16_RESULT`. One directly inspected FFN-up candidate (`128x128x64_0x0x1_0_tnt_align8_stream_k_2sm_epi_tma`, M/N/K 4290/16384/4096) returned CUTLASS `Error Internal` during initialization; an inspected FFN-down candidate (`128x256x64_2x2x1_0_tnt_align8_2sm_epi_tma`, 4290/4096/16384) was rejected after illegal-memory-access warnings. The failed log SHA-256 is `d6a7ae86e66fb36a3b9df1d191ab82d12c8d78fe7853f8ae7fb4c20e2660f99d`; no generic CUTLASS config or source staging is retained.
The constrained trainer gate also rejects the apparent opportunity. At B2, the current control completed at 720.135417 ms, while the exact-name candidate completed at 735.566620 ms: **+15.431203 ms / +2.142820%**, throughput -2.097866%, and corrected MFU 40.114997% -> 39.273438%. It emitted 7,330 illegal-access warnings, 18 failed-choice warnings, and four compiled-cache move tracebacks. The allowlist did not constrain generated backward layouts: logs selected unrequested `ntt` and `ttt` variants. This makes a regex-only production solution invalid even before the regression. Candidate log SHA-256: `0a12bc6ef750fab9b2c2fb43aad003f3be5c9a384a59246d3d2e44771087a8ce`.
A final direct B2-shape A/X/B microgate set TorchInductor's profiling swizzles to the singleton `[4]`, removed the optional epilogue suffix, and covered all 15 forward/dgrad/wgrad cases. It was clean: all parity checks passed and the candidate log contained zero illegal-access, failed-choice, traceback, RuntimeError, or CUDA-error strings. It was also decisively slower. Controls measured 393.934851 and 389.372927 ms, a 391.653889 ms midpoint; the candidate measured **422.406662 ms**, regressing **30.752773 ms / 7.852028%**. Only the small text projections won; the dominant video/FFN cases lost. Comparison/candidate log SHA-256 values are `7424fe05b7978d341c39e3f5a18f06a6c93f6a572a221f6b7479426cb66e7daf` and `23dae03bd6ae6405b40c3b55e07497ab355f1b464682ecf730fb616dc9f993b0`. This closes TorchInductor CUTLASS for B2; no source/config staging remains.
### Post-measurement validation-fix re-gate
Tracker head `52f1114dd` (validation compile isolation `0e60a0e9c`, video-only validation memory safety `3f3f06541`, and the tracker snapshot itself) was re-gated against measured head `20c36acef` with a 4x/B2 A/X/B sandwich on fresh allocation `1627918` (gb-nvl-118-compute02, 4x GB200, 1200 W power limit, 2062 MHz max SM clock, torch 2.12.0+cu130, CUDA 13.0, flash-attn-4 `4.0.0b20.dev2+g82d6441`). All three runs used the committed `runners/run_current.sh` and `harness/benchmark_fastvideo_train_pack_d016.py` from one container session with shared FA4/Inductor caches; only the `/mnt/FastVideo` source SHA changed between runs.
| 4x/B2 run | True-wall slowest-rank median | MFU | Matched result |
|---|---:|---:|---:|
| control A `20c36acef` | 0.721040 s | 40.064684% | control |
| candidate `52f1114dd` | 0.722525 s | 39.982318% | +1.326267 ms / +0.183898% vs midpoint |
| control B `20c36acef` | 0.721358 s | 40.047009% | control |
| control midpoint | 0.721199 s | 40.055847% | 0.044125% drift |
The candidate is within +0.184% of the control midpoint with byte-identical peak memory to control B (140.874/166.588 GiB allocated/reserved; control A measured 140.865/166.564 GiB on its cold-compile run) and no training-path mechanism: both post-measurement commits touch only validation code, and this harness removes the validation callback before training. All three runs passed the full embedded proofs — 927 FP32 DTensor parameters / 13,041,520,768 elements with FP32 Adam moments and exact optimizer coverage, finite AdaLN gradient probes and losses on every rank, and 4,290 semantic tokens with singleton model timesteps on all 30 steps. Load clocks/power were healthy (about 1.4-1.8 GHz SM, 0.91-1.17 kW against the 1.2 kW cap, no throttle reasons). Decision: current head inherits the stopping-point table as timing-neutral. The allocation-local absolute result (about 40.0-40.1% MFU at B2) is consistent with a healthy 4x tray and is not compared across allocations.
### MFU formula audit
A head-of-branch audit reverse-engineered and empirically verified the published MFU chain. The harness formula `MFU% = 14.444115 * local_batch * grad_accum / median_slowest_rank_step_sec` reproduces every published row from its recorded medians, and the README's `353.8808175 TFLOP/sample` is exactly `14.444115 * 24.5`: the harness constant is the primary value, rounded at the sixth decimal from the exact count below (relative rounding error 4e-9).
The numerator is the strict no-recompute model-FLOP count of the 48 transformer blocks at 4,290 video tokens, 1,024 text tokens, and hidden width 4,096:
| Component | Exact FLOPs/sample |
|---|---:|
| video-token block linears, `6*4290*234,881,024*48` | 290,200,202,772,480 |
| text KV projections, `6*1024*33,554,432*48` | 9,895,604,649,984 |
| self-attention, `12*4290^2*4096*48` | 43,420,719,513,600 |
| cross-attention, `12*4290*1024*4096*48` | 10,364,292,956,160 |
| total (353.880820 TFLOP) | 353,880,819,892,224 |
The first two rows sum to 300,095,807,422,464, the projection-gate figure above. The convention charges attention backward at twice forward (no flash recompute) and excludes the caption projection, patchifier, output head, and AdaLN MLP.
`probes/audit_train_flops_per_sample.py` then measured executed FLOPs with `FlopCounterMode` on allocation `1627918` at head `52f1114dd`, running the production recipe eagerly (TORCH_SDPA, compile off — compiled regions and the FA4 custom op bypass dispatch-mode counting; collectives are uncounted either way). At B1 it recorded 118,036,079,050,752 forward + 244,999,614,103,552 backward = 363,035,693,154,304 FLOPs/sample, identical on all four ranks and on both counted steps, and reconciled with the numerator integer-exactly in forward and backward separately: the counter additionally charges the flash-attention backward QK recompute, `(2*4290^2 + 2*4290*1024)*4096*48 = 8,964,168,744,960`, plus the excluded non-block layers, `190,704,517,120` (caption MLP 3840->4096->4096 and AdaLN/patchifier/output head, with no-grad inputs skipping first-layer dgrads). The numerator is therefore confirmed and is 0.054% conservative. A B2 executed-FLOP repeat was capacity-blocked: eager mode allocated 177.80 GiB and OOMed (the first attempt also exposed a transient `NCCLSymmetricMemory.cu:455` init failure), but batch linearity is structural — every counted operator's FLOPs are linear in the batch dimension — and the B2 timing gates already verify per-sample token semantics.
The denominator is a house convention. The device reports 152 SMs (SM100) with a 2,062 MHz max SM clock at a 1,200 W limit; NVIDIA's GB200 NVL72 materials give 360 PFLOPS sparse FP16/BF16 across 72 GPUs, i.e. 2,500 TFLOP/s dense per GPU, and no vendor source quotes 2,450. Published MFU values are therefore 1.020408x their vendor-peak equivalents: 43.623053% reads 42.750592%, 40.810314% reads 39.994108%, and the head re-gate 39.982318% reads 39.182672% against 2,500 TFLOP/s. `nsys-ai`'s 2,250 TFLOP/s default matches the 1,000 W HGX B200 part, not these 1,200 W GB200s. Keep 2,450 for continuity with every published row — multiply reported MFU by 0.98 for the vendor-peak value — and relabel the baseline per the README rule if the convention is ever changed.
## BF16 systems campaign toward 50% (2026-07-22)
Scope: standard training conventions only (FP32 registered masters and Adam moments, BF16 working compute and reductions); quantization deferred. All 4x gates ran the committed packed harness at B2 on allocation `1631376` (gb-nvl-118-compute04, healthy class, 2,062 MHz max SM clock). Four interleaved controls measured 0.732045/0.732073/0.731982/0.729492 s (39.462/39.461/39.466/39.600% allocation-local MFU); each candidate below is judged against its bracketing control midpoint.
A fresh Kineto category rollup (`harness/profile_fastvideo_train_pack_head.py`, B2, rank 0) decomposed the step: about 391 ms of nvjet GEMMs (band efficiency about 62% of the 2,450 convention), 133 ms FA4, 100.5 ms compiled pointwise/reduction kernels, 33.4 ms optimizer band (14.2 fused AdamW, 8.3 `_foreach_copy_`, 4.1 clip `_foreach_mul_`, 2.2 `_foreach_norm`), 4.6 ms of 203 eager `aten::sum` calls at compile-region boundaries (broadcast AdaLN/timestep gradient reductions), and mostly-overlapped FSDP staging (17.1 chunk_cat + 7.9 split_with_sizes + 25.8 PtoP). Eliminating the entire non-GEMM/attention main-stream band would land near the 50% target, which framed the gates below.
| 4x/B2 gate | Candidate median | vs bracket midpoint | Verdict |
|---|---:|---:|---|
| `TORCHINDUCTOR_COORDINATE_DESCENT_TUNING=1` | 0.727078 s / 39.732% | -4.981 ms / -0.680%; +0.270 pp (0.0039% drift) | **accept** |
| `CUBLASLT_WORKSPACE_SIZE=76800` alone | 0.732398 s / 39.443% | +0.370 ms / +0.051% | neutral alone |
| workspace + coordinate descent | 0.724103 s / 39.895% | -7.925 ms / -1.083%; +0.432 pp | **best stack; carry both envs** |
| cuDNN video self-attention (scratch swap) | 0.730204 s / 39.562% | -0.533 ms / -0.073% (0.341% drift) | neutral; reject source change |
| cuDNN + workspace + coordinate descent | 0.724667 s / 39.864% | -6.070 ms / -0.831% | no additive value over env stack |
The coordinate-descent candidate re-tunes Inductor pointwise/reduction configs only; the previously rejected `max-autotune-no-cudagraphs` gate additionally autotuned GEMMs into SM100 register-pressure fallbacks, so the two results do not conflict. The 2.98 ms spread between the two accepted-stack rows is within compile-lottery variance across coordinate-descent searches. Memory was unchanged in all rows.
Supporting microgates, all committed: `probes/benchmark_ltx2_video_attention_backends.py` measured cuDNN SDPA at the exact (B, 4290, 32, 128) fwd+bwd shape 2.31-3.45% faster per call than FA4 (about -3.2 ms/step projected) with FA2 and SDPA-flash 2-3x slower, but the trainer gate above shows the per-call win does not survive end to end; reject without source change. `probes/bench_ltx2_tunableop_gemm.py` reproduced the GEMM band at 416.5 ms single-GPU (validating the profile) and found `PYTORCH_TUNABLEOP` non-functional on this CUDA 13 build (tune pass writes no CSV; replay delta inside noise): reject. The same probe measured the 75 MB cuBLASLt workspace at -3.2% band in isolation, which motivated the workspace rows above; the trainer shows it is only additive under coordinate descent.
Batch capacity at 4x is rejected decisively. `harness/benchmark_fastvideo_train_blockskip.py` forces `block_skip` checkpointing with a stride (the LTX-2 wrapper does not plumb `n_layer`, so plain `block_skip` degenerates to full). B3 with every-6th-block checkpointing fits memory (160.9/179.9 GiB allocated/reserved) and passes all embedded proofs, but on the same tray whose B2 re-gate measured 0.7212 s / 40.06%, it ran 1.511623 s / 28.666% with the native allocator and 5.261937 s / 8.235% with `expandable_segments:True` — an allocator-pathology collapse at the reservation ceiling, not recompute cost (which models at about +2%). Do not pursue 4x/B3 or near-ceiling batch capacity without first freeing tens of GiB of real activation memory.
The 8x confirmation is blocked by rack `gb-nvl-057` infrastructure, diagnosed to root cause across two allocation pairs and recorded for the next attempt. Pair `1631584` (compute01/08, heterogeneous 4-vs-8 NIC trays) and pair `1631863` (compute01/04) both draw the rack's 1,965 MHz max-clock bin (baseline-ineligible against the 2,062 MHz rows even when working). Findings, in dependency order:
1. `/etc/hosts` on these hosts maps each host's own name to `127.0.1.1`, so Gloo full-mesh advertises loopback and peers get connection-refused. Fix: `GLOO_SOCKET_IFNAME=enP5p9s0`. Real, but not the main blocker.
2. Cross-tray TCP is healthy: torchrun static rendezvous, arbitrary fixed ports, ten sequential connections, and the c10d TCPStore phase (server on node0, both clients connected) all pass.
3. The actual hang: a `faulthandler` dump shows every rank blocked in its first `all_reduce` inside lazy NCCL comm creation, with zero NCCL output even at `NCCL_DEBUG=INFO` — a silent driver-level wait in the MNNVL/IMEX registration path. `nvidia-smi` reports fabric state Completed/Success but with sentinel `CliqueId 32766 (0x7ffe)` on both trays; NCCL's own detection line reads `MNNVL ... cliqueId 0x7ffe state 3 healthMask 0x11a9`.
4. Discriminator: `NCCL_MNNVL_ENABLE=0` makes the identical cross-tray all-reduce complete immediately over `Using network Socket`. Socket transport is diagnosis-only — a B3 step would be comm-bound and meaningless as an MFU row — so rack-057 pairs are unusable for the confirmation until IMEX is repaired.
Working preamble for two-tray pairs: `GLOO_SOCKET_IFNAME=enP5p9s0 NCCL_SOCKET_IFNAME=enP5p9s0`, a fresh `MASTER_PORT`, and `MASTER_ADDR` resolved to a raw IP (rank0's own `enP5p9s0` address on node 0; DNS resolution of the master hostname elsewhere) to stay immune to the hosts-file mapping.
A follow-up FA4 forward `num_splits` sweep on the video self-attention shape (same committed probe, extended with the text gate's split wrapper; run on a 1,965 MHz rack-059 tray) rejects split tuning: `num_splits=1` reproduces the current wrapper within its noisy 3-8% control drift, `2` is neutral-to-worse, and `4`/`8` regress 6.2-12.7% monotonically at both B2 and B3. The FA4 default schedule stands.
### 50% feasibility budget
Today's measurements make the 50% gap quantitative at 4x/B2 (2,450 TFLOP/s convention). Executed kernel work per step is 600.192 TFLOP of GEMMs (389.300 ms, 1,542 TF/s, 62.9% of peak) plus 125.498 TFLOP of FA4 including its backward recompute (133.329 ms, 941 TF/s, 38.4%), i.e. 725.690 TFLOP in 522.629 ms — 56.7% combined. The remaining step is about 166 ms of other compute (pointwise, norms, optimizer band, boundary sums) plus about 32 ms of non-busy wall. The 50% target step is 577.765 ms, so:
| Systems assumption for the non-kernel band | Kernel budget | Required combined kernel throughput | vs today |
|---|---:|---:|---:|
| today (198 ms) | 379.9 ms | 1,910 TF/s (78.0%) | x1.376 |
| realistic best systems (120 ms) | 457.8 ms | 1,585 TF/s (64.7%) | x1.142 |
| theoretical-max systems (75 ms) | 502.8 ms | 1,443 TF/s (58.9%) | x1.040 |
Equivalently, with today's kernels a perfect 75 ms systems endgame caps at **48.34% MFU**: 50% is strictly unreachable in BF16 by systems work alone. The admissible next steps are kernel research — raise the nvjet-bound GEMM band from 62.9% toward 70%+ and/or FA4 from 38.4% toward 50% on SM100 — or the separately-gated NVFP4 projection track. Torch-dispatch-level levers (workspace, TorchInductor CUTLASS, TunableOp, attention backends and FA4 split schedules, batch capacity) are measured and closed above; native cuBLASLt algo selection is NOT closed — see the sweep below.
### cuBLASLt heuristic algo pinning: open, quantified opportunity
`probes/bench_ltx2_cublaslt_algo_sweep.py` (research-plan item 2) enumerates up to 48 `cublasLtMatmulAlgoGetHeuristic` candidates per exact packed shape and training orientation via ctypes, validates each winner bit-exactly against torch in-process, and times bare GEMMs on both sides. On a 1,965 MHz tray (ratios valid; absolute ms are bin-scaled), the occurrence-weighted band improves **-7.98% at both B2 and B3** versus torch's own dispatch, **-9.4%** with per-case best-of, and **-7.7%** restricted to dgrad/wgrad only (the conservative subset that needs no bias-epilogue re-validation; production dgrad is epilogue-free and dbias is a separate compiled reduction). Every winning algo matched torch bit-exactly (`max_abs 0.0`). Case highlights at B2: `video_dd:dgrad` -23.6%, `text_kv:wgrad` -23.2%, `self_qkv:dgrad` -17.9%, `ffn_up:wgrad` -16.5%; the heuristic best loses on `ffn_up:fwd` (+15.2%), so pinning must be per-case best-of. Projected onto the healthy-tray 389.300 ms band this is **-30 to -37 ms/step at 4x/B2, about +1.9 to +2.2 MFU points**, lifting the GEMM band to roughly 68-70% of peak — inside the realistic-systems kernel requirement for 50% above.
A first end-to-end trainer gate already confirms the direction. `probes/lt_pinned_ops.cpp` (a pinned-algo `cublasLtMatmul` registry built via `cpp_extension`) plus `harness/benchmark_fastvideo_train_lt_pinned.py` (rank-0 tuning broadcast to all ranks; `ReplicatedLinear` forwards wrapped so only dgrad/wgrad dispatch to pinned algos, forward stays stock nvjet with fused bias) patched all 336 block projections and ran the standard 4x/B2 A/X/B on the 1,965 MHz tray `1632324`: controls 0.978211/0.979068 s (0.088% drift, about 29.52% allocation-local MFU), candidate **0.970947 s / -7.692 ms / -0.786%, +0.234 MFU points**, peak memory +0.13 GiB (the 128 MB Lt workspace). The scratch wrapper pays for an un-fused eager dbias, autograd-Function graph breaks in the compiled blocks, and contiguity copies, so it captures only part of the sweep's margin. A second variant re-expressed the wrapper as a fully-traced `torch.library` custom op with `register_autograd` and opaque inner dgrad/wgrad ops (no graph breaks; dbias traceable) — the first attempt crashed with the classic FakeTensor `data_ptr` error until the inner pybind calls were wrapped as opaque ops — and gated **neutral**: 0.975376/0.976246/0.978566 s (A/X/B, 0.326% drift, -0.724 ms / -0.074%), with only 288 of 336 projections pinned that round. The two trainer datapoints (+0.234 pp and +0.022 pp) bracket a deeper instability: per-case tuning margins on this 1,965 MHz tray swing several-fold between rounds (`self_qkv:dgrad` measured 3.3% and 17.9% wins in successive tunings; `ffn_down:wgrad` flipped sign), so the pinned set itself changes per run. The sweep-level opportunity stands on its two-batch-size, bit-exact evidence; the end-to-end capture verdict — and any recipe change — requires a healthy-class 2,062 MHz allocation where tuning is stable, plus the Inductor-fused integration. Both wrapper variants are committed for that session.
A third allocation roll drew rack-059 pair `1632122` (gb-nvl-059-compute05/06), which passed the native health gate with that preamble: all-rank scalar correctness and 458.914 GB/s slowest-rank sustained all-reduce bus bandwidth, matching the healthy 463 GB/s reference, so the MNNVL failure is rack-057-specific. The pair is nevertheless the 1,965 MHz clock bin and its absolute B3 result sits in the degraded-allocation class alongside the `1629519` row: the exact-head A/X/B measured 1.442031/1.443601/1.444706 s for control A / env stack / control B — a 1.443369 s / 30.021699% MFU control midpoint with 0.185% drift, and an env-stack delta of **+0.232 ms / +0.016%, i.e. neutral** (memory unchanged, 161.38 GiB peak allocated in all rows). On a degraded pair the pointwise-tuning win measured at 4x is inside noise, so the environment stack remains accepted for 4x work, is harmless on 8x, and stays **out of `run_current.sh`** until a healthy-class (about 43% baseline) 8x pair re-gates it. This row is baseline-ineligible and does not supersede the stopping-point table.
## Decision boundary
1. Do not write a whole-transformer mega-kernel or integrate the current QuACK wrapper.
2. Keep production FSDP2 symmetric memory, repeated/prefetched input, deferred gradient-norm materialization, both scoped accumulation optimizations, singleton timestep embedding, persistent projection packing, and two-module public-FSDP2 grouping. The singleton path is a capacity optimization; `d01623709` remains the matched B1 source attribution. Keep the corrected fixed arena only as architecture evidence until it has generic lifecycle, DCP/export, EMA, accumulation, and topology coverage. Reject root-only sharding, raw velocity, TE BF16 registered shards, fused dense GELU, no-autocast RMSNorm, grouped BF16 wgrad, short-text hybrid attention, QuACK/Triton norms, AdamW-integrated clipping, and current TorchInductor CUTLASS schedules. No currently confirmed production-ready stack reaches 50%; treat 50% as research, not a PR #1630 completion gate.
3. If BF16-equivalent all-pass NVFP4 remains wanted, the next admissible experiment is a purpose-built projection pipeline that fuses dual-orientation quantization with scale packing and QMM alpha/bias epilogues, and overlaps cache production with GEMM execution. This is projection-band kernel work, not a monolithic transformer kernel.
4. Require <=63.463 ms for the complete 48-block equivalent and acceptable sampled fprop/dgrad/wgrad cosine/norm before autograd or distributed integration. The current result fails both gates.
5. Keep CUDA graph work separate and secondary until it produces a measured >=32 ms bare saving.
## Artifact provenance
Hashes were captured immediately after each run. Scratch paths explicitly marked `recorded` were later reused or cleaned, so those entries are immutable provenance records rather than claims that the old bytes still exist at that path.
```text
b427c267440ae30b99e32588e155b81bbd4c67f9ade67065d0d3eea24fc26abc /tmp/pr1630_fsdp_root_control_a.log
8bf84e651bad50f559da45d2ed7c1f4df22c86c4e559b9e3569495ffda2ec65d /tmp/pr1630_fsdp_root_candidate.log
b1e45a3dce91565a151762a3897203dad7eac240be5f30924734d4d2cfaae8b4 /tmp/pr1630_fsdp_root_control_b.log
6f7a4e5138e51a7bb1a0e2f319cf30eb714d5affd7da8c1470acbed925e738a8 /tmp/run_pr1630_fsdp_root_only_4x_aba.sh
276c57e3120a8cfbaa54bcb4951ea14892d6d41da39d813eb51000e882b679de /tmp/pr1630_fsdp_buckets_control_a_fa47ce1.log
4b1e0246ccfcd9c88ccf1270dbcc6b0092fe3875069158a1ecbd193f27044182 /tmp/pr1630_fsdp_buckets_group2_fa47ce1.log
6f4cfb5c8a5cef467707d5a400ab049494eae7c12ba9931517d3e9cf735b7875 /tmp/pr1630_fsdp_buckets_group4_fa47ce1.log
2086997a99ed42f6235150437cdcfedf787f15d58fae13e8d37a2b0ff2511e09 /tmp/pr1630_fsdp_buckets_control_b_fa47ce1.log
0392370dbc0aec3bedc4b600184cd36cdd5e7c60043c19433f262e085a2d08ad /tmp/benchmark_fastvideo_train_fsdp_buckets_fa47ce1.py
9babec2096727dd05fcc261c4d1589c80bc6ba0491d35cfa9fdbf6982ce92d04 /tmp/run_pr1630_fsdp_buckets_4x.sh
66cb9fe7d6b95d22319ea772dafa54cde434095fa6d02c184a46b474a69056ae /mnt/pr1630_group2_recipe_smoke_20c36.log
4ccb2cc9c163357180c673daf46213be02172a4a01fb6048489d8dd937d1b593 /tmp/run_pr1630_group2_recipe_smoke_20c36.sh
93f0147b706ba9ea21887d4f6d8ebf684b81ec368628247cb35aa374ad1fc739 [exact-source optimized A/X/B runner] /tmp/run_pr1630_fsdp_group_source_4x_20c36.sh
9c0db1bc60feea788a8cadb069d674c56c58c588c3ea268c035751a55efcf064 /tmp/ltx2_bf16_cutlass_fa47ce1_valid/current_a.log
836da870f8356e0d126028906913988e3e612f204d238163da90dec1477c94fe /tmp/ltx2_bf16_cutlass_fa47ce1_valid/cutlass.log
44578c630d98f339269e3138d77bca6f5612e3ba38feb0194b4f2914b8d8402f /tmp/ltx2_bf16_cutlass_fa47ce1_valid/current_b.log
02c563717bdbea7249fc92a0141b6d755fd718627e3fd470bc7dd0df9014a192 /tmp/ltx2_bf16_cutlass_fa47ce1_valid/comparison.log
1864f0776c5929c75b95262a7273c81d07765a2dc64a1193eaad74bc058227ff /tmp/benchmark_ltx2_bf16_cutlass_gemm_fa47ce1.py
1bdb44c617162728bb33c33019d4514c6b9598e7ca75fd6d21ffe24173c129b7 /tmp/run_ltx2_bf16_cutlass_gemm_gate_fa47ce1.sh
6a419a1e0c7cb35b64bd622852d3847935fd9728a4d14d8fd45cf952ec55fa99 /tmp/pr1630_cutlass_allowlist_b2_current_a_20c36.log
0a12bc6ef750fab9b2c2fb43aad003f3be5c9a384a59246d3d2e44771087a8ce /tmp/pr1630_cutlass_allowlist_b2_allowlist_20c36.log
0f2ca616e34f40030df669a426fd0c207e70abb772c08ebee524a9b5307254b0 /tmp/pr1630_ltx2_b2_cutlass_swizzle4_20c36/current_a.log
23dae03bd6ae6405b40c3b55e07497ab355f1b464682ecf730fb616dc9f993b0 /tmp/pr1630_ltx2_b2_cutlass_swizzle4_20c36/cutlass.log
19f0b55e97316dfa7661e6bee17aade8afc87da3cc97f749f06d8162f02f1a7b /tmp/pr1630_ltx2_b2_cutlass_swizzle4_20c36/current_b.log
7424fe05b7978d341c39e3f5a18f06a6c93f6a572a221f6b7479426cb66e7daf /tmp/pr1630_ltx2_b2_cutlass_swizzle4_20c36/comparison.log
4f937bc1ebc80dd2c07dadbecec84c7c7eb2d7da59aaedf52aad36f95e025add /tmp/benchmark_ltx2_bf16_cutlass_gemm_fa47ce1.py
4fe47986aad94b5cf042a82c037ad854f69ec5d2951207ba80cd7bf83fbd59ef /tmp/run_ltx2_bf16_cutlass_gemm_gate_fa47ce1.sh
ff56f6ddf8e322e4f15d242c0367982b8745118470310f6bf4c9e54bb7693156 /tmp/pr1630_rmsnorm_repeat2_control_a.log
9640fd71cf69a022a7f4e32b00b6a74203d9567747291340a776c0a9a31bc2af /tmp/pr1630_rmsnorm_repeat2_candidate.log
8ec42a9a064a01377cde9e22dbd7a44712729ee13676fcc486ccfd21de16b2c0 /tmp/pr1630_rmsnorm_repeat2_control_b.log
0f76a190a2496aac31bd26efba14957fef41910a361c166ee7c8ddb37db6ed90 [recorded before scratch-path reuse] /tmp/benchmark_fastvideo_train.py
3a1db6f243e72af1b7bd276f8b571351a073d8cdd5ddaaf8476e2332a0061eba /tmp/run_pr1630_rmsnorm_no_autocast_4x_aba.sh
1eea3f150c18b8eddc9019edd957cd1e522af0d41c8678a537d1959a625289a2 /private/tmp/bench_ltx2_bf16_grouped_wgrad.py
874c17cdfd805ca1f2e1c2e3f9df80c00aa433a215269135fb4d7814d0abeed7 /private/tmp/run_ltx2_bf16_grouped_wgrad.sh
731a0f72b031d8768b199a6cd9990192d7e7dd5ae07a5238b8be5e3da21ea631 /private/tmp/bench_ltx2_bf16_grouped_wgrad_job1623561.jsonl
8f2615727cc6e6d73921bb8e0891e9564ef0e5d259e1136123da4339c90ae249 /tmp/pr1630_ltx2_text_attention_gate.log
40489af84ec0b8232d9b97705078243a9f36aee791595ffe3c02023f205e1a9e /tmp/benchmark_ltx2_text_attention_backends.py
74e2d30e1ec505545c5ae4192c6a8e572a1426d99d8c6aaf0fe39f2fb77a23ac /tmp/run_ltx2_text_attention_gate.sh
d5660696c77170c0fbcc61e5a9539aab1bcbb95d4ae32ad5b30bd342d81bdf7d /tmp/pr1630_ltx2_quack_norm_gate_node0.log
e5af3d5f9e89c566a740b0538d53afde99c95c826b63df0705fa497b7b4a0681 /tmp/pr1630_ltx2_quack_norm_gate_node1.log
fc43cbb03c67f87e20c216d1f3dfb5fcfcea0a8b0adea3ba70dc7abb095eb4f6 /tmp/benchmark_ltx2_quack_norm_gate_fa47ce1.py
f5da0421aa8c9c81adbb36c6639926f919046444f905ecd672703ba0ac404ec9 /tmp/run_ltx2_quack_norm_gate.sh
12ccb47c662016184e9959dc9f20d7163bc06f4e3382b9f4d3d32acacc7930b0 /tmp/benchmark_fastvideo_train_fused_clip_fa47ce1.py
2dda641fb5f97fd62c695c678a7b608665c09c036de3cd36ade4a363f90644c6 /tmp/run_pr1630_fused_clip_dual_4x_aba.sh
6bae2ddb903abb13816ef79ef1819c2788d31073063832c0f745cf5c2d8a7347 /tmp/pr1630_fused_clip_parity_node0.latest.log
8187717e6f30b0d1e931175543d62818bfeab3b571b2a5de927a5ae77d2abb54 /tmp/pr1630_fused_clip_control_a_node0.partial.log
1075beb1bbc33a127a915be1067c14fda8603c2ff81b4c22a123e9c29ddc5d3b /tmp/pr1630_fused_clip_candidate_node0.latest.log
0b7c61091024cb62dfee3b4cc1acd6c7c98aab6cd39c4253d28b2fd23f4bd295 /tmp/pr1630_fused_clip_control_b_node0.latest.log
e85b49fa8121766ef1debc57f9455b967e9cd6f6efecd8accd5fc7f195882443 /tmp/pr1630_fused_clip_parity_node1.latest.log
e15a993d043fb90367e6ebc577da961fbf6fd90c3c9296a51a78859fbac9905e /tmp/pr1630_fused_clip_control_a_node1.partial.log
d30eb902f93b01f4868dd45ae5449bcd4c7ebbd6ce0d4e7f6ad434664be4ff90 /tmp/pr1630_fused_clip_candidate_node1.latest.log
4becc78eff103866f5affdb467cbfe7bf9e4f786adba031dc4e62d8c18eb75c7 /tmp/pr1630_fused_clip_control_b_node1.latest.log
7b5852a45482e58df9088015d173f7d3fc049ec80fcaa65c89480acfa8859f10 /tmp/pr1630_fa47_fixed_arena_control_a.log
c27069a696d57a46cc3374574bf0d563f8f9387161498858935265104c1f3ce2 /tmp/pr1630_fa47_fixed_arena_candidate.log
3caa215e0ad0884965da38fc4c8a71e313d19aaf02f62c0e94a07126c4a62659 /tmp/pr1630_fa47_fixed_arena_control_b.log
46e6a123ad2980fe0ef70f178015b6c9a75abc2eab6d9fab0831606554c46127 /tmp/pr1630_ltx2_triton_norm_gate.log
08e113c7aa52ed13a5bfba07a9de3394d4a7abee1f4255ad3781ea3ec23ddfb5 /tmp/benchmark_ltx2_triton_norm_gate_fa47ce1.py
7c76497e07a0c2a39acf2bd223acf74885d20e0cd40b227b5289746181d8dcdc /tmp/run_ltx2_triton_norm_gate_fa47ce1.sh
4c69503260a9e2351037161e32138fc6d6d7e133993ea8f1edbfb62866bfee3d /tmp/zero2_ltx2_nccl_base_a.log
879d43933d958ffae3f3e1e5b94ba40a0f4bdf9448c217e0af0c4e00f869e8ae /tmp/zero2_ltx2_phase_repeat32.log
76e4eb26f4bbdbd173b884b379b78d10ae89a93f54119cc2db9e21edc654a49c /tmp/zero2_ltx2_nccl_cta16.log
375eac7aa12965034ca1629415b7e605bdda100b10aabb954ee985b3dfd80ab4 /tmp/bench_quack_sm100_nvfp4_v050_config5.jsonl
8b6b843d32476cbb6afc6150ebfbd2bc39dc7b8beaf9ffb71216035cdd55e963 /tmp/bench_quack_sm100_nvfp4_v050_config7.jsonl
3cdd55ebd363438f0bc5af8650eab7f12f05da6baba65aa5af1a7ee27cf0b961 /tmp/bench_quack_sm100_nvfp4_v050_config6.jsonl
2f901b87a1e227bf9b1a71b5b9945aaebb15420a1ec4f0412f1b22f28c1e7efb /tmp/bench_quack_sm100_nvfp4_v050_smoke.log
61577ec757260cb54d6df8260c8e633837d7c141e450966ab559573ca762c492 /tmp/bench_quack_sm100_nvfp4_quack050.py
a89a44b0fa17c2741817d43bb77bac2fcb43c221eea5b1e5591afe1b03b54f6a /tmp/zero2_ltx2_graph_staged_minimal.log
9037cfd2050f78a53967fa6f530ede2f0e1ebb4ed5d09cfa1e7d781740186d8c /tmp/zero2_ltx2_graph_patchifier_minimal.log
0aa302e5cdef20f543897ba767077e7ffbb97fccbaaa77eb77e0749cd8bd61fe /tmp/bf16_fsdp_symm_base_a.log
907d9526fb4b6c8e48aac3ac9e42920d5dffc06e00656aa32ffc855a52399fdc /tmp/bf16_fsdp_symm_candidate.log
7883956935bec2494372b4fd044042c32da901e21e3e35349bf7fb495c887eed /tmp/bf16_fsdp_symm_base_b.log
86209d328970e05a9da7849a6e8e9080015d7c20488d1cc0614e0b16b96ec3f1 /tmp/benchmark_fastvideo_train.py
3499dbf7080426d33bb13a7121fd45e83326eadddbc545769418200890b65f6e /tmp/benchmark_fastvideo_train_symm_mem.py
5fcd10ba65ebfdf7c2d37a2b4411a99c464ad9e39fe09444f44fd840696d1784 /tmp/bench_quack_sm100_nvfp4_complete_projection.jsonl
61e5fa6ed8aa26ed15c825cf9c7ad5aa4f34151d0803766ada7fa87866f2c258 /tmp/bench_quack_sm100_nvfp4_complete_projection_smoke.jsonl
56bae046d674a7c1f4396f117f6faca83241377c850b4e409b8fe01cb15425d5 /tmp/bench_quack_sm100_nvfp4_complete_projection.py
cfd2f8621a2cc04120caf78434087b1ac345e8357b979d3df34443fcc88554d9 /tmp/bench_quack_sm100_nvfp4_complete_projection_DESIGN.md
aa260691ce2847aa1e5cf5c852fec8c2af67de306b4bf393557bdbc3a4477d54 /tmp/pr1630_mfu_1x_fa4_no_vae.log
b8a6444936ae98641dd9bf947dc2fd513faa5650d4eedc4137b78bdbda8eb43a /tmp/pr1630_mfu_4x_full_shard_symm_production.log
0c4850071806aaaec637aee0a3201cfc21cb8bf44d195da775ff8c17c2ab9a52 /tmp/pr1630_mfu_4x_no_reshard_symm_production.log
5be182649bdbd78c6e49e862422fa11968fec7ddc2112de91e24834f80cc35fa /tmp/pr1630_mfu_8x_full_shard_node0.log
fc7c9c6e6783281aa372d17f94fb48bff005d1bcffbceed79ae5288657cf752d /tmp/pr1630_mfu_8x_full_shard_symm_production_node0.log
715fcb9385bbcc81ba605508828e4114bc159d82d31c3001afcb9f633353b31e /tmp/pr1630_mfu_8x_full_shard_cta2_b_node0.log
218902dc7568fd690bfd1cc43e5f5c02667d5fde37cfba72d815736d26420d5a /tmp/pr1630_mfu_8x_no_reshard_cta2_a_node0.log
a81f1899a00d02ee5861322649e4702981f0592ed555bc17866a1e7f9b7aef61 /tmp/pr1630_mfu_8x_no_reshard_symm_production_node0.log
f3c368eef644dfd0d806dfb8e4293c8ab3d3bb139939e03539343be1b9300592 /tmp/pr1630_mfu_8x_no_reshard_cta2_b_node0.log
f217ea1756a281b11c4be8b1b8d17a37aa6e3d43160ef8c94b54fb4f0d7261c0 /tmp/pr1630_nvlink_intra_inter_node0.log
a11d544802989fa716f2c44ffa288af3255e5b0ae0467f044dd745ab85b21b60 /tmp/bench_nvlink_intra_inter.py
2b249d4ac775494a08aa704b51bbad7b41f3fd586a6bdf774143188b7bc7ef23 /tmp/pr1630_fsdp_knobs_base_a.log
07b5e0c13c8800de85c978d1395e3ce2e52a818747cf9187d9f35cf98978da36 /tmp/pr1630_fsdp_knobs_base_b.log
c93dda310617d44464dd3aa5efa9792a81d7bb679b3fc4edc25d0f7ae84bc7ab /tmp/pr1630_fsdp_prefetch1.log
5844e6838734032785c7df4dadfacf92fcc78d884f0aff146850dd23420713d2 /tmp/pr1630_fsdp_prefetch2.log
14535b726719c1a4a8d393237c260fbc33d98eb3aaa5e8996d265b3bc2192224 /tmp/pr1630_fsdp_pg_alloc_alt.log
8eead593a61f9f74e09592aea7cef90564b7cd9596d41839e74e529709f3408b /tmp/pr1630_hsdp_base_a_node0.log
21d43573ad3224f9d3426368e677af50fbda8d6573ac355fdb6165ea18077aeb /tmp/pr1630_hsdp_2x4_node0.log
dc1cd934099e7ee72bed564213a1bcf415343d8183aa0ce3c5a4851ee9a964cb /tmp/pr1630_accum2_no_sync.log
caad7d10cbb211ca534d4120621ad2e39b85cb95935ab5bdb903a5473f8eca98 /tmp/pr1630_accum2_force_sync.log
4d86bb7ec6600c9925d9223129f2c15955e579a817e1c97624a9ae5e5aad4b4b /tmp/pr1630_accum2_no_sync_b.log
abde9be2305a70f9787f7bceec18fd2b197ad0f778bed1634ceb3e310aca8a76 [recorded before scratch-path reuse] /tmp/benchmark_fastvideo_train_fsdp_knobs.py
5e1ffa7fbbeeff10f20814f2315e4175df79d76f2a97665ca62553af235ec8e0 /tmp/bf16_input_control_maxrank2_7f139e2b.log
05f753036f10f6475c21ac7ffb2deb6c8374e0028ddc2f9c3413f85c63368405 /tmp/bf16_input_repeat32_maxrank_7f139e2b.log
b254de8912d9556c2b28040593216ae3cf5b6e496b8d2602482090b6a35df303 /tmp/pr1630_mesh_base_a.log
3d747e320d05995b0a30fd49b7ac89b1a611e3b074cc6531c26af6b75a20c41c /tmp/pr1630_mesh_1d.log
da372dae90eb840709853e7601bc8a4fe04291872a470b370cb3f5d6f1b2f5af /tmp/pr1630_mesh_base_b.log
31baeaec642fe690c564a337a87e145339a3c3cfd7a8053a26d6c7fefa59cde4 /tmp/run_fsdp_mesh_benchmark.py
230e0a749be30c7c23cac5fd8c76782c32ed3bcff282b2647d3dedaa6c193bb5 /tmp/run_fsdp_2d_control_benchmark.py
37b17e37d7ae2491ffa36ee382957664fbd51806b44a34cb35f553425c4a6360 /tmp/test_fsdp_2d_to_1d_dcp.py
cfaa897a6dcce741a72ffe826992c773323e791d83ca7d40c0f5f7b3db13649d /tmp/pr1630_mesh_dcp_smoke.log
5917084f5745e8c3d94ff7c3449ab4000c59657f347eaa3ad7bd6245427113e6 /tmp/bench_quack_sm100_nvfp4_allpass_job1622676_config7.jsonl
3286866342ef63334dc26cabfa57ea393e3888da272ba4fed7268674a56441b1 /tmp/bench_quack_sm100_nvfp4_allpass.py
8a523326e703310f36219f5fedb108d009c2887b149e135a332645e6e06daad0 /tmp/pr1630_gradnorm4_fa4_base_a.log
63ad8b88c8e2e434087e95e126217e91333fe047ae65bbd31230034fe371f0a4 /tmp/pr1630_gradnorm4_fa4_deferred.log
49fc52400bcb65485888b5dc159644ba18fc6e8eb5b785a92c070ab226aaff28 /tmp/pr1630_gradnorm4_fa4_base_b.log
87347668ad28ab200bf3540da01ff1c7744854c2ce2a3a0a7d1c0cb1d1788d01 /tmp/pr1630_gradnorm4_source_deferred.log
683f6716b8912e26c974286228404f1d4a9d7e1e458654ee10068697c1e0b681 [recorded before scratch-path reuse] /tmp/benchmark_fastvideo_train_grad_norm_sync.py
05836dda4426309f23e1c797281f0e69cc343c522006679fef6b34e6c7d3f865 /tmp/check_ltx2_singleton_timestep_parity.py
efe7e529860771bc3564adf82b2cfa2a712dc60cf83691b174769aff11760317 /tmp/singleton_gate_7f6.0.log
b3a317fc18813c34aca7b5d0998f8e5a5b5e85d94fe707720da8c47b0a614752 /tmp/benchmark_fastvideo_train_ltx2_singleton_timestep.py
983b52f96ceace4ddd1fa01c36f6b851c9bb74fd6c4033e8932bc99abf9f1c3a /tmp/pr1630_singleton_control_a.log
189f0ffc65ff4ae2a449990c026a385a79222e45e95ade94977bbb981c23f70a /tmp/pr1630_singleton_candidate.log
505d0062845b2c1662384ab7f5e60acdebd4e6971ac66bcb521b7ab0b2bb41f5 /tmp/pr1630_singleton_control_b.log
f6924a9277723fe51c2d210f5ea3fb66aafbb845d142ee50ef92effd1cb5fa4d /tmp/check_ltx2_raw_velocity_roundtrip.py
ee3b1ae78578e0226082230f566e244c6aa7c9063439218a47af0542d1211d22 /tmp/benchmark_fastvideo_train_ltx2_raw_velocity.py
83ebfafb10c4c9f10011b5cf746b7c6fdde6414d8247774ade8013f4a57a2f78 /tmp/pr1630_raw_velocity_control_a.log
b421b2e0bcb1c9b9ac8f4871f3ad41587249a1f520faed6795fc0b10319c0c45 /tmp/pr1630_raw_velocity_candidate.log
6a1875640a2528a8bd126a9037138f03d34a06bf188e1fc10e06005c26cab514 /tmp/pr1630_raw_velocity_control_b.log
c14ff07a8eb9785ae59f789fef1da9cfeb830ec29afc4e41936dc932f5637dbd /tmp/run_pr1630_batch2_aba.sh
521853c83e20eb99fe2ad459d4ee8a180e34075e49c387918203a952844c55c8 /tmp/pr1630_batch2_control_a.log
bed7081a443163326ec6401d8c8a085af4069f6a47bad34ac66882856aec9e54 /tmp/pr1630_batch2_candidate.log
d867ca9be8ed9f1a72fd2c55dee8e48d60a2a0a71dff346f27bbb456930219d3 /tmp/pr1630_batch2_control_b.log
e8c710ecc8bc247051395cd10b69f895c0ab589755d82ea616122e50a1e12962 [recorded; scratch runner cleaned] /tmp/run_pr1630_batch2_8x_aba.sh
e5b52d00895f097e88bf5f71fc890f6709a935b8b50abfeeccf7c41f3c0cf285 /tmp/pr1630_batch2_8x_control_a_node0.log
5c825d2a9dad33566669ec3324b4499af54be6e275c3abd30692ebbcec61632e /tmp/pr1630_batch2_8x_control_a_node1.log
3cc50b482a9a1ec1fd5a9ed405e6abab8c34608b90958f4166ef4886a67e824c /tmp/pr1630_batch2_8x_candidate_node0.log
073550a2f6109acdaa94c72a1f9a53bcbaee4468957d65cd36de539d45baed9b /tmp/pr1630_batch2_8x_candidate_node1.log
602656ed804d6c453d2e3a16f30d613f5569534d4dd984e2354f039b18962050 /tmp/pr1630_batch2_8x_control_b_node0.log
6f619a490cd0046c14a37a14d83f29161947df35dfea9bbb9471f7cfa7e7668a /tmp/pr1630_batch2_8x_control_b_node1.log
15dc2f617c139c05fc8e8c45dfd6d06f3f3fecc6d6999b61ced78dddb9c1f44c [recorded; scratch runner cleaned] /tmp/run_pr1630_batch3_8x_aba.sh
604804ad3139a6fbc345ba31b58c575ed82ad1d87ef55bfdeac0b90027103af3 /tmp/pr1630_batch3_8x_control_a_node0.log
48b8c8b7a65c906e7f043b9e0dbaebd829e7abdaaf85b67956145f687d0c9db6 /tmp/pr1630_batch3_8x_control_a_node1.log
117264799fba09c8e267e5f30bb8dc2a0f53e3618109e1d13c91e0347bcbb770 /tmp/pr1630_batch3_8x_candidate_node0.log
1d86886b79c59e309ea452b6e9c323a2b2a134468a8adb3962c71a323ae4a45b /tmp/pr1630_batch3_8x_candidate_node1.log
049c822078ee57013bd1be3020970ea3bfdaf084be4e91afed8331fcd904a92d /tmp/pr1630_batch3_8x_control_b_node0.log
c295dca59a641d5ea628f9f697d2db1cb3d55cc494356e2c2bda4d51ac0eeda2 /tmp/pr1630_batch3_8x_control_b_node1.log
cbb3fd5a95ab0ba2081c9fffd49767a9004a74770d1ccc0b5d83c57a12cbd1d1 /tmp/run_pr1630_batch4_full_shard_8x_aba.sh
4e21524538adc0fa7fde4cd117d6e09380cd2714b3a2e9e3fcda5ac81bc75a7e /tmp/pr1630_batch4_full_shard_8x_control_a_node0.log
9af9a84e58b446723a585cf95c7433894b876d54450ff4b05b970c87cf97362d /tmp/pr1630_batch4_full_shard_8x_control_a_node1.log
cfcf400c687750ef45bd84c24b8ef68b67b6e75d48aaa3024a58b794549d513d /tmp/pr1630_batch4_full_shard_8x_candidate_node0.log
57ef2efb41b85b09dedf8801d705d42429f98309b3850b04392810aba814216d /tmp/pr1630_batch4_full_shard_8x_candidate_node1.log
38f064416dfd68699fd58ee694dfe8d8b01a93c76d5d75b018b64ca9f6fd29cf /tmp/benchmark_fastvideo_train_fsdp_knobs.py
ab97235175edfd90a8576920b3456414bedab07eb31dde9279cf37cb4b19fcf8 /tmp/run_pr1630_accum_reshard_aba.sh
a36c89911353f21398f0f927f0490d4dfbd4bda1adac48b38bb72f2a4cfd1e6b /tmp/pr1630_accum_reshard_control_a.log
e5156b7d066912b9ad1489dcf0b0cb3146861121dcf39103b6a3ff6516a946c2 /tmp/pr1630_accum_reshard_candidate.log
e254c0ec799944f870e329c18313fb1eca8d2e1af85df3df86115eef7882a87d /tmp/pr1630_accum_reshard_control_b.log
adacec37acc19f436f289bb2ea7ff184263a39255b0412cc7ed58a0769204c97 /tmp/pr1630_te_master_8x_control_a_node0.log
d819ab40b55a546d7649d430914950f1d738a3d5a50ddda2f8a266d27fdc6722 /tmp/pr1630_te_master_8x_control_a_node1.log
6c9dc52db52ef7ae7571ea1a0f02072aa44b9f283e0b7ee7a53fbda8e8b42923 /tmp/pr1630_te_master_8x_candidate_node0.log
cc2a138f8dfdfa84233e3e5e32f4855ad3a57ea72c7413a61bf50c7744c3ff82 /tmp/pr1630_te_master_8x_candidate_node1.log
4c8e91889e9eef21f938382b21e230b928e71c44d71da8685927bfd14bb1c35b /tmp/pr1630_te_master_8x_control_b_node0.log
115a5d0e18d21683da40d3a1dbc87f11842cda23822999039067285c3507c980 /tmp/pr1630_te_master_8x_control_b_node1.log
bf0861ff481c499fd65350d2f0f16c86487db88cb03c332e4be5c76540fcc7fd /private/tmp/benchmark_fastvideo_train_ltx2_singleton_timestep.py
63e702e2b258184ffa7ae668e4e53a8b4e7f9779a0a27b93a4ad045e27212b65 /private/tmp/run_pr1630_te_master_8x_aba.sh
eac2355c930ce7aae541b97d9bedb636e0b76250b30d2348f83bca6eabc3a2e1 /private/tmp/optimizer_te_master_scratch.py
949a31c8c78b05e7529d0cd81d312740edecc5686783bc8d986fbe19c7e45d35 /mnt/pr1630_pack_4x_control_a_d016.log
ad29304eaffbd56ef2053790d310f56c3acdfd75b89113136192fb8f2b62a93e /mnt/pr1630_pack_4x_candidate_d016.log
f03b4578e1ac8931e002adcb5c00f9c1aa656cd4256138697935a1e795ead818 /mnt/pr1630_pack_4x_control_b_d016.log
bf0861ff481c499fd65350d2f0f16c86487db88cb03c332e4be5c76540fcc7fd /mnt/benchmark_fastvideo_train_pack_d016.py
78d41f61e46710ff6d4c10e2758c37ff9aff08768ed211cdb2b4525ee303a295 /mnt/run_pr1630_pack_4x_aba.sh
d4857775e66717b9ec43434d4147f01bf16c0bbfb8f9e721aff3a30755376568 /mnt/pr1630_pack_8x_b3_control_a_d016_node0.log
0265f372cb18d07812fc0ff5dc08fefff5ba720fc89b159385981ea0843c0834 /mnt/pr1630_pack_8x_b3_control_a_d016_node1.log
5e1f03976d834465cd60072459f960f40b4d0ac4dfba7d24e7b4856f4e615bcc /mnt/pr1630_pack_8x_b3_candidate_d016_node0.log
65ea51b8b98c21da82eb0a93d64dbc341b748dc5ecefdf6b4d2f033de56a4324 /mnt/pr1630_pack_8x_b3_candidate_d016_node1.log
8a8727d2777625dbabe9cbd8fd862f12aa0996e761138b1518668269166419c8 /mnt/pr1630_pack_8x_b3_control_b_d016_node0.log
22e7626f076d5e07ab70b57abbfe9df2d39e8d44e713897802ce6abc58659fc2 /mnt/pr1630_pack_8x_b3_control_b_d016_node1.log
534f201aadfc3cb02df2ad130a1193c01ebfa97ddcfd38fd4ee05a5a1e4965a5 /mnt/benchmark_fastvideo_train_pack_b3.py
ea5a9fa8898fba7904c9a649ee7fde28b1821a07747b8ba89480d877a1876592 /mnt/run_pr1630_pack_8x_b3_aba.sh
808d90a77c1d52d234a3eee30fa145b4c8ee85f4e064f0b051bd3c6255d231f9 /mnt/pr1630_compile_mode_4x_control_a_fa47.log
4da9de0d042c20b06ec465059ed6ecbcae0b019094d473f32c0c3d82c6fc2fdb /mnt/pr1630_compile_mode_4x_candidate_fa47.log
8d6573c2b13341064a96efb081cf292e87d43fd2e155d2cfb24c96eff207c327 /mnt/pr1630_compile_mode_4x_control_b_fa47.log
d7e571ac7043e4252f44260617e0de8f3467fd97ca3393393a04c64d1af482f3 /mnt/benchmark_fastvideo_train_compile_mode.py
ebc6a2493ee8a109b02561a8f6c8d9a5f5b9c09bd84241c9285b449d4f3b70b3 /mnt/run_pr1630_compile_mode_4x_aba.sh
c3ec531c9cf00c569be30ce1460e224225c617e17e8224af7385b9a82955f08c /mnt/pr1630_real_pack_export_keep_fa47.log
2a4ba1280710c026fefdcd188ce0207f99a5828c277dcaf265758a8e46c06367 /mnt/validate_ltx_pack_export.py
46b31dd97b14a28fd5da06efcd8893e8cb404908dafbc1d46e6108bc6628c62c /mnt/pr1630_attention_compile_4x_control_a_fa47.log
62aedbef5b90f9ebd162e8b45a06dbc8c74650ddd12f4a3539ef8e23fbba4580 /mnt/pr1630_attention_compile_4x_candidate_fa47.log
e81bc6f1640ad0c23b74df1fb22eb38f5c37f402d3d294baff691bb9291ad1e6 /mnt/pr1630_attention_compile_4x_control_b_fa47.log
b2fabf6ffa2d174f449c44568373a758c93c39ce2e2da01151412a6e8531dabb /mnt/run_pr1630_attention_compile_4x_aba.sh
ebb28a2930d4f4d3372b10717d54b7b60d18e006ed720e19458af6ae299da9f9 /mnt/flash-attention-82d6441e.tar.gz
578486b918fe205f17f449f254975207c3b9443b4e12d8e77e461d443d63d672 /mnt/flash-attention-82d6441eec5d4dfec120153db2c0145ae855a083/csrc/fused_dense_lib/fused_dense_lib.cpython-312-aarch64-linux-gnu.so
589a2f030bb42ea48aec9c981f638fef05c43dc4146275f83b0c0476240f6299 /mnt/bench_ltx2_fused_dense_gelu.py
c2ba5a88e215bd81199122aa5298dfc36f3d05770410ad1c491bed71eff115f6 /mnt/build_and_run_ltx2_fused_dense_gelu.sh
4fb7528407e75718c86a2596544643d4a9b2e942a6f078e511d59d0b104d1080 /mnt/pr1630_fused_dense_gelu_82d6441.log
ae32a4fb16cb4c735ba055c6ef6556ed71560d77cfb0770ea3559c78fb2b0d3b /mnt/pr1630_fused_dense_gelu_h1.log
17ac709ec74445a4ad3bb440e3bab6c77f0fca80783545644dc0f71fa94ab5a5 /mnt/pr1630_fused_dense_gelu_h2.log
17ac709ec74445a4ad3bb440e3bab6c77f0fca80783545644dc0f71fa94ab5a5 /mnt/pr1630_fused_dense_gelu_h3.log
17ac709ec74445a4ad3bb440e3bab6c77f0fca80783545644dc0f71fa94ab5a5 /mnt/pr1630_fused_dense_gelu_h4.log
1edd404b1fb886ca92439761a255861a9119878becef74f2e7e3455ccbf1b34e /mnt/pr1630_regate_52f1114/control_a.log
0e2bb3834933ab1b0aabbe65b3905e8615fc3f8d7a8ae76ac1b07f448cbff296 /mnt/pr1630_regate_52f1114/candidate.log
98e17c84a4dda1c43e5726b827bfc34e91a3721460ca999fe01d1ebb5abe1deb /mnt/pr1630_regate_52f1114/control_b.log
098b29f24cc1858ae4237f3a9dce5830ac1484f6c77c23615b9d428bfb966a41 /mnt/pr1630_regate_52f1114/health_before.log
7154bfe19e0b8b1971fae454f85b76fe4228b04b8a8a3eb34e955b4edc0488e5 /mnt/pr1630_regate_52f1114/health_during.csv
9b3f76ba4b0b6e79e04e461638d660f08c9a63880d4e433848d9af9932583384 /mnt/pr1630_regate_52f1114/health_after.log
c21dc0fff404a7b7950610a854bb14316b23899a120eca3068f9960a804f1a1c [committed runners/run_current.sh as staged for the re-gate] /mnt/pr1630_tracker_52f1114/runners/run_current.sh
ec4cd5092a691de0f0c630b5ed349f98e7d94d5795dfd11de9e242765bdcb790 [committed harness/benchmark_fastvideo_train_pack_d016.py as staged for the re-gate] /mnt/pr1630_tracker_52f1114/harness/benchmark_fastvideo_train_pack_d016.py
f39045adab530f07cbd080abc08e1b3b45bb8b80aa1ce16f605366625f492d32 /mnt/pr1630_flop_audit/b1.log
9a353a98d3697ed84b3b485cb20d52fad8acfabcb6d11cc808db88e20a48afec [failed base-patch first attempt, then NCCL symmetric-memory init failure at B2] /mnt/pr1630_flop_audit/b2.log
3b8537b9dbe4d2a707a16bef65165b2fbba1645d67ca2628a36b45b0c96a7241 [eager B2 OOM, capacity-blocked] /mnt/pr1630_flop_audit/b2_nosymm.log
64feb60e3deb87bb76564a30326922470c2113bd729b2ce85fdacf40743ddd25 [committed probes/audit_train_flops_per_sample.py] /mnt/audit_train_flops_per_sample.py
1a2d98385c77b491aed2009f9bf23e004459c97c0b6f8dcf3ab17dbc245eba52 /mnt/pr1630_head_profile_b2.log
89826925326da5da851a5766eff4b7e98fc66139851d470bd6a5e58f264cf75d /mnt/pr1630_head_profile_b2.summary.json
50cc112f6bc88990154db11880b0788933467655c4aa4ac6829e2385c87e1147 /mnt/pr1630_cdt_gate/control_a.log
b031e3152d022816eabf450f0f619a15496291aadad92207e88579e19df6c3b5 /mnt/pr1630_cdt_gate/candidate.log
2b268a684fc265810703898480799e12c88d21807543104cff6e37d899dc6c8e /mnt/pr1630_cdt_gate/control_b.log
1accd7c721551ab800ddbfe82e1a99723d99aaa14899fe7e065a2238ec3f5d16 /mnt/pr1630_ws_gate/ws75.log
00e998a4498652c4416c19b35b9b68db046b1e24256bb0491f4a3e2a0667ba79 /mnt/pr1630_ws_gate/ws75_cdt.log
132b7939e026cec3d7194ac2d10c19f56ed37b30d7bbca22b3fe28622dcf1db5 /mnt/pr1630_ws_gate/control_c.log
d5ee5c7e85c4f31583676c8af795b5882dc637f7d11bd32256642b95f75ad888 /mnt/pr1630_cudnn_gate/cudnn_alone.log
c6cbdd0e447afa14f0eb5ed37183151ce2f4705fbcffe812ca1182e2897a60e7 /mnt/pr1630_cudnn_gate/cudnn_cdt_ws.log
cf116038960d0edbe2a94982d42d5a18feb34629fb85cfdd260d889985304ed5 /mnt/pr1630_cudnn_gate/control_d.log
35a882dbb4972476e3dfb0207970a6140c1dad31e8f8a864ebfb5da0a273ee0f /mnt/pr1630_video_attn_gate_b2.log
0cc840410699cb6b28840134007eff8f28d2a2a45f9385863124aa3277e90793 /mnt/pr1630_video_attn_gate_b3.log
6fc72a3d3c1af53712662575c012df386cbbb77d4cec0a573ebe6df1901587eb [committed probes/benchmark_ltx2_video_attention_backends.py] /mnt/benchmark_ltx2_video_attention_backends.py
286afc02d989a72567c573532a45db61782c90c6b9f8b9bb8ba85d12ed26c718 /mnt/pr1630_gemm_band_default.log
653824bad44bb39142ed74c228cc338a4ef9f990e5c5338224ed6f5cb34dfee1 /mnt/pr1630_gemm_band_ws75.log
37f5c465f534b1c32a0080c0585ab7c5c705214ad8f33dea912ce63179e256cc /mnt/pr1630_gemm_band_tune.log
b973790131bf7a6f04799c482ad919b03094489ef8ac56478f271febf93fbbc0 /mnt/pr1630_gemm_band_replay.log
47cd9bf4c7a7833ad3dfcf4e86e0eb10627499748992da89e15831bf6634d56e [committed probes/bench_ltx2_tunableop_gemm.py] /mnt/bench_ltx2_tunableop_gemm.py
4ed684c2858008e4e02ca70943974b21856172dd730d18cd8ea8d84334b9eaf0 /mnt/pr1630_b3skip6_smoke.log
8a8ea4c8eabd20188f75d2f5aee98f0b4f6065d7e09f70df4f6fe052c2d8f87b /mnt/pr1630_b3skip6_expandable.log
fc2c97a86e1d0ae082b6d8bb3340e554063919943eda52f6fce1864f19d88f2e /mnt/pr1630_8x_health_1631863_node0.log
1f75874659db76179a0e3e79e7e95259999ed2989035b6f88cbe8efc00bcb205 /mnt/pr1630_8x_health_1631863_node1.log
13b309424061d50611a5e9da5ffc6f0cf7c7d10ab73f40ba655f9fd76235c6ac /mnt/pr1630_8x_stack_node0.log
f0e16c6bb91d9889d450ff072673c65076d2d0d826b6e8ea8ebeb4c07b4984fd /mnt/pr1630_8x_stack_node1.log
9e824e8988f3a63d0d53a27ce76da6f99d1eefa6aafa3dd0fefc3954e6198ca2 /mnt/pr1630_8x_nomnnvl_node0.log
0a0093e71b52f804963710606f1f77ece99d37d90b6924b7ce1c82c59c948541 /mnt/pr1630_8x_nomnnvl_node1.log
69787402d008711ea68f3b3289aa2dd05e7ee12034d4e7337d8a792e051f29f8 /mnt/pr1630_8x_crossdebug_node0.log
f3203e22560a7cc1e041b36d406eee1faef580ea5d402470dffae8a1726c259c /mnt/pr1630_8x_crossdebug_node1.log
f6bdae8f1e6aef72285ddcb15a0639a28549e67d12e0c2b9db18536255efa8a9 /mnt/pr1630_8x_health_1632122_node0.log
7ef78badee4d7077955e129d9df3ff29bc95353ca2f28dab74d08ebf6616b6d0 /mnt/pr1630_8x_health_1632122_node1.log
77f6911fa369b850477a6026e86b026a10236d990bc38b8942aa948f85f99a18 /mnt/pr1630_8x_envstack_control_a_node0.log
6df94ec852ef76334d8f45e991b8cde5bd2e38dc85fb9f3a342d8b5d1f16c16f /mnt/pr1630_8x_envstack_control_a_node1.log
15a99f49f50925df8e53c8ce628328c1cfee819bd14fa4276b630c775ef5bd67 /mnt/pr1630_8x_envstack_envstack_node0.log
4704d335a688c2eab66ad6eb4e033c12b709938689abac82398e52326105cca6 /mnt/pr1630_8x_envstack_envstack_node1.log
10eebf0d0eb9c71e5aca4eee51c66e68454767b0838af61516cce0d4e3a231f4 /mnt/pr1630_8x_envstack_control_b_node0.log
f681ac1b60beb31aa3dd868e9005684810e05ddb26966bd78885141acd7864bb /mnt/pr1630_8x_envstack_control_b_node1.log
697537479ba7f945983b0fe3e632bab39fdbb9406cf39b584587c4ad16387bea /mnt/pr1630_lt_sweep_b2.log
a02fe5338276d8dc3ce31b0af903d80eb71a9980a797e078732893f15807c463 /mnt/pr1630_lt_sweep_b3.log
41fc1e8970cbe39431a9cd56773910d3ca6cd0950ef843b5882d6c15a5d66f47 [committed probes/bench_ltx2_cublaslt_algo_sweep.py] /mnt/bench_ltx2_cublaslt_algo_sweep.py
afb1e48236ae5930ce8c12d730b935b6ad05595da9690821a808c75809f7546f /mnt/pr1630_ltpin_gate2/control_a.log
29785fe9e0eae9aa19bc6bc937a3e2a469c78f68c64402a9e044caf477b6e820 /mnt/pr1630_ltpin_gate2/ltpin.log
aae60722634e8c5a9d1f678ffe17f4c01a45ff131987e4252e7326bfb04fc0df /mnt/pr1630_ltpin_gate2/control_b.log
20f6ebe53a31768ee6da82285fef94f2a383f48c7008c732c691f2d4a3c0faaf /mnt/pr1630_fa4_tail_tile.log
b2190af81ca34e0cb8b9487ba81bc58d34fb8b2af9836a633a536105f89f96ef [committed probes/bench_fa4_tail_tile.py] /mnt/bench_fa4_tail_tile.py
3cb59dae1cdf2c1b860f17c01d7c030ac55cccde4aea772ad130aa3e557a6168 /mnt/pr1630_ltpin_gate4/control_a.log
31b2ff5e8cca54fa519a3ced4243c3266cf474b328fcc03c06224617e245d349 /mnt/pr1630_ltpin_gate4/ltpin_traced.log
faf28571bb0d48d6f2c91f967ebf426323839e886dd7384bb35d5dac09190d03 /mnt/pr1630_ltpin_gate4/control_b.log
7994f41c82ece7a3ab8c324dc3543f99cf4437c4dce3db4d0dbf36bd43b12b74 /mnt/pr1630_video_attn_splits_b2.log
bd33740e9159d6ea5a4f02e06e5fc15721fb10d8a3dadd91e20d1608175602f7 /mnt/pr1630_video_attn_splits_b3.log
647c2ffb55613bf3ded4f464fc1db1cfbe1ca5bfc50df7ab30c1aeecee3909d6 [committed probes/benchmark_ltx2_video_attention_backends.py with split sweep] /mnt/benchmark_ltx2_video_attention_backends.py
```
PR #1630's optimization stack was review-clean through measured head `20c36acef`; validation fixes followed at `0e60a0e9c` and `3f3f06541`. No dependency was added or installed. The operational GB200 launcher now forwards allocated IMEX character devices into multi-node containers so the existing MNNVL fabric is reachable.
@@ -0,0 +1,67 @@
#!/usr/bin/env python3
"""Run FastVideo training and print the same per-step metric sent to trackers."""
import argparse
import json
import statistics
import torch.distributed as dist
from fastvideo.training.trackers import DummyTracker
_step_times: list[float] = []
_original_log = DummyTracker.log
def _log(self, metrics, step):
_original_log(self, metrics, step)
if "step_time_sec" in metrics:
value = float(metrics["step_time_sec"])
if not dist.is_initialized() or dist.get_rank() == 0:
print("BF16_STEP " + json.dumps({"step": step, "step_time_sec": value}), flush=True)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True)
args, overrides = parser.parse_known_args()
DummyTracker.log = _log
import fastvideo.train.trainer as trainer_module
from fastvideo.train.entrypoint.train import main as train_main
from fastvideo.distributed import get_world_group
trainer_module.build_tracker = lambda *_args, **_kwargs: DummyTracker()
original_trainer_init = trainer_module.Trainer.__init__
class _Recorder:
def on_training_step_end(self, _method, metrics, iteration=0):
_step_times.append(float(metrics["step_time_sec"]))
def _trainer_init(self, *init_args, **init_kwargs):
original_trainer_init(self, *init_args, **init_kwargs)
self.callbacks._callbacks.pop("validation", None)
self.callbacks._callbacks["_benchmark_recorder"] = _Recorder()
trainer_module.Trainer.__init__ = _trainer_init
train_main(argparse.Namespace(config=args.config, dry_run=False), overrides=overrides or None)
world = get_world_group()
times_by_rank = [None] * world.world_size
dist.all_gather_object(times_by_rank, _step_times, group=world.cpu_group)
if world.rank == 0:
per_step_max = [max(values) for values in zip(*times_by_rank, strict=True)]
measured = per_step_max[10:30]
if len(measured) != 20:
raise RuntimeError(f"expected 30 steps, got {len(per_step_max)}")
median = statistics.median(measured)
print("BF16_RESULT " + json.dumps({
"median_step_sec": median,
"samples_per_second": 4.0 / median,
"model_mfu_percent": 14.444115 / median,
}, sort_keys=True), flush=True)
if __name__ == "__main__":
main()
@@ -0,0 +1,13 @@
#!/usr/bin/env python3
"""Run the shared timing harness and report single-GPU peak allocation."""
import runpy
import torch
import torch.distributed as dist
runpy.run_path("/mnt/benchmark_fastvideo_train.py", run_name="__main__")
if dist.get_world_size() != 1:
raise RuntimeError(f"expected one rank, got {dist.get_world_size()}")
print(f"BF16_PEAK_GIB {torch.cuda.max_memory_allocated() / 2**30:.6f}", flush=True)
@@ -0,0 +1,36 @@
#!/usr/bin/env python3
"""Scratch capacity gate: block_skip activation checkpointing with a stride.
The LTX-2 wrapper calls ``apply_activation_checkpointing`` without plumbing
``n_layer``, so ``block_skip`` degenerates to full checkpointing. This driver
forces ``block_skip`` with ``FV_BLOCK_SKIP_N`` (checkpoint every Nth block)
around the frozen packed benchmark harness. Launch with
``--models.student.enable_gradient_checkpointing_type block_skip`` so the
wrapper takes the checkpointing path at all; the patch supplies the stride.
"""
from __future__ import annotations
import os
import benchmark_fastvideo_train_pack_d016 as benchmark
import fastvideo.train.models.ltx2.ltx2 as ltx2_module
from fastvideo.training.activation_checkpoint import apply_activation_checkpointing
BLOCK_SKIP_N = int(os.environ["FV_BLOCK_SKIP_N"])
def _forced_block_skip(module, checkpointing_type="full", n_layer=1):
del checkpointing_type, n_layer
wrapped = apply_activation_checkpointing(module, checkpointing_type="block_skip", n_layer=BLOCK_SKIP_N)
wrapped_blocks = sum(
1 for _, child in module.transformer_blocks.named_children()
if type(child).__name__ == "CheckpointWrapper")
print(f"BLOCK_SKIP_APPLIED n={BLOCK_SKIP_N} wrapped_blocks={wrapped_blocks}", flush=True)
return wrapped
ltx2_module.apply_activation_checkpointing = _forced_block_skip
if __name__ == "__main__":
benchmark.main()
@@ -0,0 +1,479 @@
#!/usr/bin/env python3
"""Scratch A/B for regional torch.compile mode on packed LTX-2 training."""
from __future__ import annotations
import argparse
from collections import Counter
import gc
import hashlib
import json
import os
import statistics
import time
from typing import Any
import torch
import torch.distributed as dist
EXPECTED_STEPS = 30
WARMUP_STEPS = 10
MFU_NUMERATOR = 14.444115
GRAD_PROBE_STEP = WARMUP_STEPS
_step_times: list[float] = []
_step_starts: list[float] = []
_synced_end: float | None = None
_metrics_by_step: dict[int, dict[str, float]] = {}
_semantic_counts: Counter[str] = Counter()
_ada_grad_probe: dict[str, Any] | None = None
_peak_memory: dict[str, int] = {}
_optimizer_probe: dict[str, Any] = {}
_original_log: Any = None
def _rank_zero() -> bool:
return not dist.is_initialized() or dist.get_rank() == 0
def _log(self: Any, metrics: dict[str, Any], step: int) -> None:
if _original_log is None:
raise RuntimeError("tracker log hook was not initialized")
_original_log(self, metrics, step)
record = _metrics_by_step.setdefault(int(step), {})
for key, value in metrics.items():
if key in {"total_loss", "finetune_loss", "step_time_sec"} or key.startswith("grad_norm/"):
record[key] = float(value)
if "step_time_sec" in metrics and _rank_zero():
print(
"BF16_STEP " + json.dumps({
"step": int(step),
"step_time_sec": float(metrics["step_time_sec"]),
}),
flush=True,
)
def _wall_intervals(starts: list[float], synced_end: float) -> list[float]:
if not starts:
raise RuntimeError("no training-step wall starts were recorded")
return [next_start - start for start, next_start in zip(starts, starts[1:])] + [synced_end - starts[-1]]
def _tensor_digest(tensor: torch.Tensor) -> str:
raw = tensor.detach().contiguous().view(torch.uint8).cpu().numpy().tobytes()
return hashlib.sha256(raw).hexdigest()
def _capture_ada_grad_probe(method: Any, iteration: int) -> dict[str, Any]:
parameters: dict[str, Any] = {}
for name, parameter in method.student.transformer.named_parameters():
if "adaln_single" not in name:
continue
gradient = parameter.grad
if gradient is None:
parameters[name] = {"present": False}
continue
if isinstance(gradient, torch.distributed.tensor.DTensor):
gradient = gradient.to_local()
local = gradient.detach().contiguous().cpu()
local_float = local.float()
parameters[name] = {
"present": True,
"shape": list(local.shape),
"dtype": str(local.dtype),
"numel": local.numel(),
"finite": bool(torch.isfinite(local).all()),
"l2_norm": float(torch.linalg.vector_norm(local_float)),
"max_abs": float(local_float.abs().max()),
"mean": float(local_float.mean()),
"sha256": _tensor_digest(local),
}
if not parameters:
raise RuntimeError("found no adaln_single parameters for the gradient probe")
if not all(item.get("finite", False) for item in parameters.values() if item.get("present")):
raise RuntimeError("non-finite Ada gradient in probe")
return {"step": int(iteration), "parameters": parameters}
def _capture_optimizer_probe(method: Any) -> dict[str, Any]:
optimizers = list(method.get_optimizers(0))
if len(optimizers) != 1:
raise RuntimeError(f"expected one optimizer, got {len(optimizers)}")
optimizer = optimizers[0]
use_te_master = os.environ.get("FASTVIDEO_TE_FP32_MASTER") == "1"
optimizer_class = f"{type(optimizer).__module__}.{type(optimizer).__name__}"
if use_te_master:
if type(optimizer).__name__ != "FusedAdam" or not type(optimizer).__module__.startswith(
"transformer_engine."
):
raise RuntimeError(f"TE master mode constructed the wrong optimizer: {optimizer_class}")
elif not isinstance(optimizer, torch.optim.AdamW):
raise RuntimeError(f"stock control constructed the wrong optimizer: {optimizer_class}")
trainable_parameters = [
parameter for parameter in method.student.transformer.parameters() if parameter.requires_grad
]
optimizer_parameters = [
parameter for group in optimizer.param_groups for parameter in group["params"]
]
trainable_parameter_ids = [id(parameter) for parameter in trainable_parameters]
optimizer_parameter_ids = [id(parameter) for parameter in optimizer_parameters]
if len(optimizer_parameter_ids) != len(set(optimizer_parameter_ids)):
raise RuntimeError("optimizer contains duplicate parameter objects")
if set(trainable_parameter_ids) != set(optimizer_parameter_ids):
raise RuntimeError(
"optimizer parameter coverage differs from trainable transformer parameters: "
f"trainable={len(trainable_parameter_ids)} optimizer={len(optimizer_parameter_ids)}"
)
param_dtypes: set[str] = set()
state_dtypes: dict[str, set[str]] = {}
missing_state: dict[str, int] = {}
writeback_mismatches = torch.zeros((), dtype=torch.int64, device="cuda")
parameter_count = 0
parameter_numel = 0
for parameter in trainable_parameters:
if not isinstance(parameter, torch.distributed.tensor.DTensor):
raise RuntimeError(f"registered parameter is not a DTensor shard: {type(parameter).__name__}")
parameter_count += 1
parameter_numel += parameter.numel()
param_dtypes.add(str(parameter.dtype))
state = optimizer.state.get(parameter)
if state is None:
raise RuntimeError("trainable optimizer parameter has no state")
for name in ("master_param", "exp_avg", "exp_avg_sq"):
state_tensor = state.get(name)
if state_tensor is None:
missing_state[name] = missing_state.get(name, 0) + 1
continue
if not isinstance(state_tensor, torch.distributed.tensor.DTensor):
raise RuntimeError(f"optimizer state {name} is not a DTensor shard: {type(state_tensor).__name__}")
if (
state_tensor.shape != parameter.shape
or state_tensor.stride() != parameter.stride()
or state_tensor.placements != parameter.placements
or state_tensor.device_mesh.device_type != parameter.device_mesh.device_type
or not torch.equal(state_tensor.device_mesh.mesh, parameter.device_mesh.mesh)
or state_tensor.to_local().shape != parameter.to_local().shape
or state_tensor.to_local().stride() != parameter.to_local().stride()
):
raise RuntimeError(
f"optimizer state {name} layout does not match its registered parameter"
)
state_dtypes.setdefault(name, set()).add(str(state_tensor.dtype))
if use_te_master:
master = state.get("master_param")
if master is None:
raise RuntimeError("TE optimizer parameter is missing master_param")
local_parameter = parameter.to_local().detach()
local_master = master.to_local().detach()
if not torch.equal(local_parameter, local_master.to(local_parameter.dtype)):
writeback_mismatches += 1
expected_param_dtype = {"torch.bfloat16"} if use_te_master else {"torch.float32"}
if param_dtypes != expected_param_dtype:
raise RuntimeError(f"unexpected registered parameter dtypes: {param_dtypes}")
if state_dtypes.get("exp_avg") != {"torch.float32"} or state_dtypes.get("exp_avg_sq") != {"torch.float32"}:
raise RuntimeError(f"optimizer moments are not FP32: {state_dtypes}")
if missing_state.get("exp_avg", 0) or missing_state.get("exp_avg_sq", 0):
raise RuntimeError(f"missing optimizer moment states: {missing_state}")
if use_te_master:
if state_dtypes.get("master_param") != {"torch.float32"}:
raise RuntimeError(f"optimizer masters are not FP32: {state_dtypes}")
if missing_state.get("master_param", 0):
raise RuntimeError(f"missing TE master states: {missing_state}")
elif state_dtypes.get("master_param") or missing_state.get("master_param", 0) != parameter_count:
raise RuntimeError(f"stock optimizer unexpectedly contains master states: {state_dtypes}, {missing_state}")
if dist.is_initialized():
dist.all_reduce(writeback_mismatches, op=dist.ReduceOp.SUM)
total_writeback_mismatches = int(writeback_mismatches.item())
if total_writeback_mismatches:
raise RuntimeError(
"registered BF16 parameters differ from rounded FP32 masters: "
f"mismatches_across_ranks={total_writeback_mismatches}"
)
return {
"class": optimizer_class,
"parameter_count": parameter_count,
"parameter_numel": parameter_numel,
"parameter_dtypes": sorted(param_dtypes),
"state_dtypes": {name: sorted(dtypes) for name, dtypes in sorted(state_dtypes.items())},
"missing_state": missing_state,
"master_writeback_mismatches_across_ranks": total_writeback_mismatches,
"te_fp32_master": use_te_master,
}
def _series_digest(metrics_by_step: dict[int, dict[str, float]]) -> str:
series = [(step, sorted(metrics.items())) for step, metrics in sorted(metrics_by_step.items())]
return hashlib.sha256(json.dumps(series, separators=(",", ":")).encode()).hexdigest()
def _self_test() -> None:
assert _wall_intervals([1.0, 2.0, 4.0], 7.0) == [1.0, 2.0, 3.0]
assert len(hashlib.sha256(b"test").hexdigest()) == 64
print("SELF_TEST_OK")
def main() -> None:
global _ada_grad_probe, _optimizer_probe, _original_log, _peak_memory, _synced_end
parser = argparse.ArgumentParser()
parser.add_argument("--config")
parser.add_argument("--singleton-timestep", action="store_true")
parser.add_argument("--compile-mode", default=None)
parser.add_argument("--self-test", action="store_true")
args, overrides = parser.parse_known_args()
if args.self_test:
_self_test()
return
if not args.config:
parser.error("--config is required unless --self-test is used")
from fastvideo.training.trackers import DummyTracker
_original_log = DummyTracker.log
DummyTracker.log = _log
import fastvideo.train.trainer as trainer_module
import fastvideo.train.utils.moduleloader as moduleloader_module
from fastvideo.distributed import get_world_group
from fastvideo.train.callbacks.grad_clip import GradNormClipCallback
from fastvideo.train.entrypoint.train import main as train_main
from fastvideo.train.methods.base import TrainingMethod
from fastvideo.train.models.ltx2 import LTX2Model
original_build_kwargs = LTX2Model._build_distill_input_kwargs
original_method_optimizer_step = TrainingMethod.optimizers_schedulers_step
original_trainer_init = trainer_module.Trainer.__init__
original_trainer_iter = trainer_module.Trainer._iter_dataloader
original_trainer_run = trainer_module.Trainer.run
original_make_training_args = moduleloader_module._make_training_args
def _make_training_args(*make_args: Any, **make_kwargs: Any) -> Any:
training_args = original_make_training_args(*make_args, **make_kwargs)
if args.compile_mode:
training_args.torch_compile_kwargs = {"mode": args.compile_mode}
return training_args
moduleloader_module._make_training_args = _make_training_args
def _build_distill_input_kwargs(self: LTX2Model, *build_args: Any, **build_kwargs: Any) -> dict[str, Any]:
original_token_count = getattr(self, "_token_count", None)
if original_token_count is None:
raise RuntimeError("LTX-2 token count was not initialized before building transformer inputs")
_semantic_counts[f"semantic_token_count_{int(original_token_count)}"] += 1
if not args.singleton_timestep:
result = original_build_kwargs(self, *build_args, **build_kwargs)
else:
if int(self.training_config.distributed.sp_size or 1) != 1:
raise RuntimeError("singleton-timestep scratch path is restricted to SP=1")
# The source method uses _token_count only to expand the uniform
# per-sample sigma. Temporarily setting it to one reproduces the
# proposed source path without editing the checkout.
self._token_count = 1
try:
result = original_build_kwargs(self, *build_args, **build_kwargs)
finally:
self._token_count = original_token_count
timestep = result.get("timestep")
if not isinstance(timestep, torch.Tensor) or timestep.ndim != 2:
raise RuntimeError(f"unexpected LTX-2 timestep: {type(timestep).__name__}")
_semantic_counts[f"model_timestep_tokens_{int(timestep.shape[1])}"] += 1
return result
LTX2Model._build_distill_input_kwargs = _build_distill_input_kwargs
def _optimizer_step(self: TrainingMethod, iteration: int) -> None:
global _ada_grad_probe
if int(iteration) == GRAD_PROBE_STEP:
_ada_grad_probe = _capture_ada_grad_probe(self, iteration)
_semantic_counts["ada_grad_probes"] += 1
original_method_optimizer_step(self, iteration)
TrainingMethod.optimizers_schedulers_step = _optimizer_step
trainer_module.build_tracker = lambda *_args, **_kwargs: DummyTracker()
class _Recorder:
def on_training_step_end(self, _method: Any, metrics: dict[str, Any], iteration: int = 0) -> None:
del iteration
_step_times.append(float(metrics["step_time_sec"]))
def _trainer_init(self: Any, *init_args: Any, **init_kwargs: Any) -> None:
original_trainer_init(self, *init_args, **init_kwargs)
if int(self.training_config.loop.gradient_accumulation_steps or 1) != 1:
raise RuntimeError("scratch singleton benchmark requires gradient accumulation 1")
if int(self.training_config.distributed.sp_size or 1) != 1:
raise RuntimeError("scratch singleton benchmark requires SP=1")
grad_callbacks = [
callback for callback in self.callbacks._callbacks.values()
if isinstance(callback, GradNormClipCallback)
]
if len(grad_callbacks) != 1 or not grad_callbacks[0]._log_grad_norms:
raise RuntimeError("scratch singleton benchmark requires one logging GradNormClipCallback")
self.callbacks._callbacks.pop("validation", None)
self.callbacks._callbacks["_benchmark_recorder"] = _Recorder()
def _timed_iter(self: Any, dataloader: Any) -> Any:
iterator = original_trainer_iter(self, dataloader)
while True:
_step_starts.append(time.perf_counter())
yield next(iterator)
def _trainer_run(self: Any, method: TrainingMethod, **kwargs: Any) -> Any:
global _optimizer_probe, _peak_memory, _synced_end
vae = getattr(getattr(method, "student", None), "vae", None)
if vae is not None:
method.student.vae = None
del vae
gc.collect()
torch.cuda.empty_cache()
if _rank_zero():
print("BF16_SETUP " + json.dumps({"unused_vae": "unloaded"}), flush=True)
torch.cuda.reset_peak_memory_stats()
result = original_trainer_run(self, method, **kwargs)
torch.cuda.synchronize()
_synced_end = time.perf_counter()
_peak_memory = {
"allocated_bytes": int(torch.cuda.max_memory_allocated()),
"reserved_bytes": int(torch.cuda.max_memory_reserved()),
}
_optimizer_probe = _capture_optimizer_probe(method)
return result
trainer_module.Trainer.__init__ = _trainer_init
trainer_module.Trainer._iter_dataloader = _timed_iter
trainer_module.Trainer.run = _trainer_run
train_main(argparse.Namespace(config=args.config, dry_run=False), overrides=overrides or None)
if _synced_end is None:
raise RuntimeError("final synchronized wall endpoint was not recorded")
wall_intervals = _wall_intervals(_step_starts, _synced_end)
if len(_step_times) != EXPECTED_STEPS or len(wall_intervals) != EXPECTED_STEPS:
raise RuntimeError(
f"expected {EXPECTED_STEPS} steps, got internal={len(_step_times)} wall={len(wall_intervals)}"
)
if _ada_grad_probe is None or _semantic_counts["ada_grad_probes"] != 1:
raise RuntimeError("Ada gradient probe did not run exactly once")
payload = {
"compile_mode": args.compile_mode,
"internal_step_times": _step_times,
"wall_intervals": wall_intervals,
"end_to_end_wall_sec": _synced_end - _step_starts[0],
"measured_window_wall_sec": _synced_end - _step_starts[WARMUP_STEPS],
"metrics_by_step": _metrics_by_step,
"metric_series_digest": _series_digest(_metrics_by_step),
"semantics": dict(_semantic_counts),
"ada_grad_probe": _ada_grad_probe,
"optimizer_probe": _optimizer_probe,
"peak_memory": _peak_memory,
}
world = get_world_group()
payloads: list[dict[str, Any] | None] = [None] * world.world_size
dist.all_gather_object(payloads, payload, group=world.cpu_group)
if world.rank != 0:
return
rank_payloads = [rank_payload for rank_payload in payloads if rank_payload is not None]
if len(rank_payloads) != world.world_size:
raise RuntimeError("missing rank payload")
internal_slowest = [
max(rank_payload["internal_step_times"][index] for rank_payload in rank_payloads)
for index in range(EXPECTED_STEPS)
]
wall_slowest = [
max(rank_payload["wall_intervals"][index] for rank_payload in rank_payloads)
for index in range(EXPECTED_STEPS)
]
internal_measured = internal_slowest[WARMUP_STEPS:]
wall_measured = wall_slowest[WARMUP_STEPS:]
internal_median = statistics.median(internal_measured)
wall_median = statistics.median(wall_measured)
expected_model_tokens = 1 if args.singleton_timestep else None
semantics_by_rank = []
for rank, rank_payload in enumerate(rank_payloads):
semantics = rank_payload["semantics"]
semantic_token_counts = {
int(key.rsplit("_", 1)[1]): count
for key, count in semantics.items()
if key.startswith("semantic_token_count_")
}
model_token_counts = {
int(key.rsplit("_", 1)[1]): count
for key, count in semantics.items()
if key.startswith("model_timestep_tokens_")
}
if sum(semantic_token_counts.values()) != EXPECTED_STEPS or len(semantic_token_counts) != 1:
raise RuntimeError(f"rank {rank} recorded invalid semantic token counts: {semantic_token_counts}")
if sum(model_token_counts.values()) != EXPECTED_STEPS:
raise RuntimeError(f"rank {rank} recorded invalid timestep counts: {model_token_counts}")
if expected_model_tokens is not None and model_token_counts != {expected_model_tokens: EXPECTED_STEPS}:
raise RuntimeError(f"rank {rank} did not use singleton timesteps: {model_token_counts}")
if expected_model_tokens is None and model_token_counts != semantic_token_counts:
raise RuntimeError(
f"rank {rank} control changed timestep shape: semantic={semantic_token_counts}, model={model_token_counts}"
)
semantics_by_rank.append({
"rank": rank,
"counts": semantics,
"metric_series_digest": rank_payload["metric_series_digest"],
"peak_memory": rank_payload["peak_memory"],
"ada_grad_probe": rank_payload["ada_grad_probe"],
"optimizer_probe": rank_payload["optimizer_probe"],
"first_metrics": rank_payload["metrics_by_step"].get(1, {}),
"last_metrics": rank_payload["metrics_by_step"].get(EXPECTED_STEPS, {}),
})
print(
"BF16_TIMING " + json.dumps({
"internal_slowest_rank_sec": internal_slowest,
"true_wall_slowest_rank_sec": wall_slowest,
"final_cuda_sync": True,
}, sort_keys=True),
flush=True,
)
print(
"BF16_SEMANTICS " + json.dumps({
"singleton_timestep": args.singleton_timestep,
"gradient_probe_step": GRAD_PROBE_STEP,
"by_rank": semantics_by_rank,
"rank0_metric_series": rank_payloads[0]["metrics_by_step"],
}, sort_keys=True),
flush=True,
)
print(
"BF16_RESULT " + json.dumps({
"compile_mode": args.compile_mode,
"singleton_timestep": args.singleton_timestep,
"world_size": world.world_size,
"internal_median_step_sec": internal_median,
"internal_model_mfu_percent": MFU_NUMERATOR / internal_median,
"true_wall_median_step_sec": wall_median,
"true_wall_model_mfu_percent": MFU_NUMERATOR / wall_median,
"true_wall_sum_slowest_intervals_sec": sum(wall_measured),
"true_wall_mean_slowest_interval_sec": statistics.mean(wall_measured),
"true_measured_window_max_rank_sec": max(
rank_payload["measured_window_wall_sec"] for rank_payload in rank_payloads
),
"true_end_to_end_max_rank_sec": max(
rank_payload["end_to_end_wall_sec"] for rank_payload in rank_payloads
),
"samples_per_second_from_true_wall_median": world.world_size / wall_median,
"peak_allocated_max_rank_bytes": max(
rank_payload["peak_memory"]["allocated_bytes"] for rank_payload in rank_payloads
),
"peak_reserved_max_rank_bytes": max(
rank_payload["peak_memory"]["reserved_bytes"] for rank_payload in rank_payloads
),
}, sort_keys=True),
flush=True,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,39 @@
#!/usr/bin/env python3
"""Scratch gate: route LTX-2 video self-attention to cuDNN SDPA.
Wraps the FLASH_ATTN backend's compilable entrypoint so equal-length
query/key calls (video self-attention) run ``F.scaled_dot_product_attention``
under the cuDNN backend while unequal-length calls (text cross-attention)
stay on FA4, then runs the frozen packed benchmark harness unchanged. The
microbench precedent is a 2.3-3.4% per-call win for cuDNN at (B, 4290, 32,
128) fwd+bwd; this measures the end-to-end step.
"""
from __future__ import annotations
import torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel
import benchmark_fastvideo_train_pack_d016 as benchmark
import fastvideo.attention.backends.flash_attn as fa_module
_original = fa_module.flash_attn_func_compilable
def _hybrid_attention(q, k, v, softmax_scale=None, causal=False):
if q.shape[1] == k.shape[1]:
with sdpa_kernel(SDPBackend.CUDNN_ATTENTION):
return F.scaled_dot_product_attention(
q.transpose(1, 2),
k.transpose(1, 2),
v.transpose(1, 2),
scale=softmax_scale,
is_causal=causal,
).transpose(1, 2)
return _original(q, k, v, softmax_scale=softmax_scale, causal=causal)
fa_module.flash_attn_func_compilable = _hybrid_attention
if __name__ == "__main__":
benchmark.main()
@@ -0,0 +1,90 @@
#!/usr/bin/env python3
"""Run FastVideo training and report slowest-rank, dynamic-world metrics."""
import argparse
import gc
import json
import statistics
import torch
import torch.distributed as dist
from fastvideo.training.trackers import DummyTracker
_step_times: list[float] = []
_original_log = DummyTracker.log
def _log(self, metrics, step):
_original_log(self, metrics, step)
if "step_time_sec" in metrics:
value = float(metrics["step_time_sec"])
if not dist.is_initialized() or dist.get_rank() == 0:
print("BF16_STEP " + json.dumps({"step": step, "step_time_sec": value}), flush=True)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True)
args, overrides = parser.parse_known_args()
DummyTracker.log = _log
import fastvideo.train.trainer as trainer_module
from fastvideo.distributed import get_world_group
from fastvideo.train.entrypoint.train import main as train_main
trainer_module.build_tracker = lambda *_args, **_kwargs: DummyTracker()
original_trainer_init = trainer_module.Trainer.__init__
original_trainer_run = trainer_module.Trainer.run
class _Recorder:
def on_training_step_end(self, _method, metrics, iteration=0):
_step_times.append(float(metrics["step_time_sec"]))
def _trainer_init(self, *init_args, **init_kwargs):
original_trainer_init(self, *init_args, **init_kwargs)
self.callbacks._callbacks.pop("validation", None)
self.callbacks._callbacks["_benchmark_recorder"] = _Recorder()
def _trainer_run(self, method, **kwargs):
vae = getattr(getattr(method, "student", None), "vae", None)
if vae is not None:
method.student.vae = None
del vae
gc.collect()
torch.cuda.empty_cache()
if not dist.is_initialized() or dist.get_rank() == 0:
print("BF16_SETUP " + json.dumps({"unused_vae": "unloaded"}), flush=True)
return original_trainer_run(self, method, **kwargs)
trainer_module.Trainer.__init__ = _trainer_init
trainer_module.Trainer.run = _trainer_run
train_main(argparse.Namespace(config=args.config, dry_run=False), overrides=overrides or None)
world = get_world_group()
times_by_rank = [None] * world.world_size
dist.all_gather_object(times_by_rank, _step_times, group=world.cpu_group)
if world.rank == 0:
per_step_max = [max(values) for values in zip(*times_by_rank, strict=True)]
measured = per_step_max[10:30]
if len(measured) != 20:
raise RuntimeError(f"expected 30 steps, got {len(per_step_max)}")
median = statistics.median(measured)
print(
"BF16_RESULT " + json.dumps(
{
"median_step_sec": median,
"model_mfu_percent": 14.444115 / median,
"samples_per_second": world.world_size / median,
"world_size": world.world_size,
},
sort_keys=True,
),
flush=True,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,208 @@
#!/usr/bin/env python3
"""Scratch-only LTX-2 FSDP2 communication-bucket benchmark."""
from __future__ import annotations
import argparse
import json
import os
import sys
from collections.abc import Sequence
from typing import Any
import torch
import torch.distributed as dist
from torch import nn
from torch.distributed.fsdp import CPUOffloadPolicy, MixedPrecisionPolicy, fully_shard
BucketMode = int | str
def _chunks(items: Sequence[nn.Module], size: int) -> list[list[nn.Module]]:
if size <= 0:
raise ValueError("bucket size must be positive")
return [list(items[start:start + size]) for start in range(0, len(items), size)]
def _unique_parameters(module: nn.Module) -> list[nn.Parameter]:
return list({id(parameter): parameter for parameter in module.parameters()}.values())
def _numel(parameters: Sequence[nn.Parameter]) -> int:
return sum(parameter.numel() for parameter in parameters)
def _rank_zero() -> bool:
return not dist.is_initialized() or dist.get_rank() == 0
def _install_bucket_policy(mode: BucketMode) -> None:
import fastvideo.models.loader.fsdp_load as fsdp_load
original_shard_model = fsdp_load.shard_model
def _shard_model(
model: nn.Module,
*,
cpu_offload: bool,
reshard_after_forward: bool = True,
mp_policy: MixedPrecisionPolicy | None = MixedPrecisionPolicy(),
mesh: Any = None,
fsdp_shard_conditions: list[Any] = [], # noqa: B006
pin_cpu_memory: bool = True,
) -> None:
if mode == 1:
original_shard_model(
model,
cpu_offload=cpu_offload,
reshard_after_forward=reshard_after_forward,
mp_policy=mp_policy,
mesh=mesh,
fsdp_shard_conditions=fsdp_shard_conditions,
pin_cpu_memory=pin_cpu_memory,
)
if _rank_zero():
print(
"BF16_BUCKET_PLAN " + json.dumps({
"mode": "current",
"block_count": 48,
"blocks_per_group": 1,
"communication_group_count": 49,
"public_fully_shard_list_api": False,
}, sort_keys=True),
flush=True,
)
return
if os.environ.get("FASTVIDEO_FSDP2_AUTOWRAP", "0") == "1":
raise RuntimeError("bucket benchmark is incompatible with FASTVIDEO_FSDP2_AUTOWRAP=1")
if mp_policy is None:
raise RuntimeError("bucket benchmark requires an explicit FSDP2 mixed-precision policy")
if not fsdp_shard_conditions:
raise RuntimeError("bucket benchmark requires the model's FSDP shard condition")
named_modules = list(model.named_modules())
matched = [
(name, module)
for name, module in named_modules
if any(condition(name, module) for condition in fsdp_shard_conditions)
]
expected_names = [f"model.transformer_blocks.{index}" for index in range(48)]
names = [name for name, _ in matched]
if names != expected_names:
raise RuntimeError(f"expected the 48 ordered LTX-2 blocks, got {names}")
blocks = [module for _, module in matched]
all_parameters = _unique_parameters(model)
block_parameter_ids: set[int] = set()
block_numels: list[int] = []
for block in blocks:
parameters = _unique_parameters(block)
parameter_ids = {id(parameter) for parameter in parameters}
overlap = block_parameter_ids.intersection(parameter_ids)
if overlap:
raise RuntimeError(f"LTX-2 block parameters overlap across buckets: {len(overlap)}")
block_parameter_ids.update(parameter_ids)
block_numels.append(_numel(parameters))
root_parameters = [parameter for parameter in all_parameters if id(parameter) not in block_parameter_ids]
if not root_parameters or len(block_parameter_ids) + len(root_parameters) != len(all_parameters):
raise RuntimeError("block and root buckets do not cover model parameters exactly once")
default_param_dtype = getattr(mp_policy, "param_dtype", None)
dtype_selector = getattr(model, "_get_parameter_dtype", None)
ignored_params = set()
if callable(dtype_selector) and default_param_dtype is not None:
ignored_params = {
parameter
for name, parameter in model.named_parameters()
if dtype_selector(name, default_param_dtype) != default_param_dtype
}
if ignored_params:
raise RuntimeError("scratch LTX-2 bucket benchmark does not support ignored mixed-dtype parameters")
fsdp_kwargs: dict[str, Any] = {
"reshard_after_forward": reshard_after_forward,
"mesh": mesh,
"mp_policy": mp_policy,
}
if cpu_offload:
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(pin_memory=pin_cpu_memory)
if mode == "root":
block_groups: list[list[nn.Module]] = []
grouped_numels = [_numel(all_parameters)]
fully_shard(model, **fsdp_kwargs)
communication_group_count = 1
else:
if mode not in (2, 4):
raise ValueError(f"unsupported bucket mode: {mode!r}")
block_groups = _chunks(blocks, mode)
grouped_numels = [
sum(block_numels[index:index + mode])
for index in range(0, len(block_numels), mode)
]
for group in block_groups:
# Public PyTorch FSDP2 supports a module list as one
# communication group while retaining hooks on each module.
fully_shard(group, **fsdp_kwargs)
fully_shard(model, **fsdp_kwargs)
communication_group_count = len(block_groups) + 1
if _rank_zero():
print(
"BF16_BUCKET_PLAN " + json.dumps({
"mode": mode,
"block_count": len(blocks),
"blocks_per_group": None if mode == "root" else mode,
"block_communication_groups": len(block_groups),
"communication_group_count": communication_group_count,
"decorated_block_modules": 0 if mode == "root" else len(blocks),
"public_fully_shard_list_api": mode != "root",
"total_parameter_numel": _numel(all_parameters),
"root_straggler_numel": _numel(root_parameters),
"grouped_block_numel_min": min(grouped_numels),
"grouped_block_numel_max": max(grouped_numels),
"original_parameter_dtypes": sorted({str(parameter.dtype) for parameter in all_parameters}),
"working_parameter_dtype": str(mp_policy.param_dtype),
"reduction_dtype": str(mp_policy.reduce_dtype),
"initial_reshard_after_forward": reshard_after_forward,
"cpu_offload": cpu_offload,
}, sort_keys=True),
flush=True,
)
fsdp_load.shard_model = _shard_model
def _self_test() -> None:
modules = [nn.Identity() for _ in range(8)]
assert [len(group) for group in _chunks(modules, 1)] == [1] * 8
assert [len(group) for group in _chunks(modules, 2)] == [2] * 4
assert [len(group) for group in _chunks(modules, 4)] == [4] * 2
print("FSDP_BUCKET_SELF_TEST_OK")
def main() -> None:
parser = argparse.ArgumentParser(add_help=False)
parser.add_argument("--fsdp-blocks-per-group", choices=("1", "2", "4", "root"), required=False)
parser.add_argument("--bucket-self-test", action="store_true")
args, remaining = parser.parse_known_args()
if args.bucket_self_test:
_self_test()
return
if args.fsdp_blocks_per_group is None:
parser.error("--fsdp-blocks-per-group is required")
mode: BucketMode = args.fsdp_blocks_per_group if args.fsdp_blocks_per_group == "root" else int(
args.fsdp_blocks_per_group)
sys.argv = [sys.argv[0], *remaining]
_install_bucket_policy(mode)
import benchmark_fastvideo_train_pack_d016 as benchmark
benchmark.main()
if __name__ == "__main__":
main()
@@ -0,0 +1,167 @@
#!/usr/bin/env python3
"""Run FastVideo training with scratch-only public FSDP2 performance knobs."""
from __future__ import annotations
import argparse
import gc
import json
import statistics
import torch
import torch.distributed as dist
from torch.distributed.fsdp import FSDPModule
from fastvideo.training.trackers import DummyTracker
_step_times: list[float] = []
_gradient_sync_calls: list[bool] = []
_forced_reshard_calls = 0
_original_log = DummyTracker.log
def _log(self, metrics, step):
_original_log(self, metrics, step)
if "step_time_sec" in metrics:
value = float(metrics["step_time_sec"])
if not dist.is_initialized() or dist.get_rank() == 0:
print("BF16_STEP " + json.dumps({"step": step, "step_time_sec": value}), flush=True)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True)
parser.add_argument("--pg-alloc", action="store_true")
parser.add_argument("--prefetch-depth", type=int, default=0)
parser.add_argument("--force-sync", action="store_true")
parser.add_argument("--force-reshard-after-backward", action="store_true")
args, overrides = parser.parse_known_args()
if args.prefetch_depth < 0:
raise ValueError("--prefetch-depth must be nonnegative")
DummyTracker.log = _log
import fastvideo.train.trainer as trainer_module
import fastvideo.train.utils.moduleloader as moduleloader
from fastvideo.train.methods.base import TrainingMethod
from fastvideo.distributed import get_world_group
from fastvideo.train.entrypoint.train import main as train_main
original_load_module = moduleloader.load_module_from_path
original_set_requires_gradient_sync = TrainingMethod.set_requires_gradient_sync
def _set_requires_gradient_sync(self, enabled: bool) -> None:
global _forced_reshard_calls
enabled = True if args.force_sync else enabled
if not dist.is_initialized() or dist.get_rank() == 0:
_gradient_sync_calls.append(enabled)
original_set_requires_gradient_sync(self, enabled)
if args.force_reshard_after_backward and not enabled:
for model in self._role_models.values():
transformer = getattr(model, "transformer", None)
if getattr(model, "_trainable", False) and isinstance(transformer, FSDPModule):
transformer.set_reshard_after_backward(True, recurse=True)
if not dist.is_initialized() or dist.get_rank() == 0:
_forced_reshard_calls += 1
TrainingMethod.set_requires_gradient_sync = _set_requires_gradient_sync
def _load_module_from_path(*load_args, **load_kwargs):
module = original_load_module(*load_args, **load_kwargs)
if load_kwargs.get("module_type") != "transformer":
return module
fsdp_modules = [submodule for submodule in module.modules() if isinstance(submodule, FSDPModule)]
if not fsdp_modules:
raise RuntimeError("scratch FSDP knob benchmark requires an FSDP-wrapped transformer")
if args.pg_alloc:
for fsdp_module in fsdp_modules:
fsdp_module.set_allocate_memory_from_process_group_for_comm(True)
if args.prefetch_depth:
depth = args.prefetch_depth
for index, fsdp_module in enumerate(fsdp_modules):
forward = fsdp_modules[index + 1:index + 1 + depth]
backward = list(reversed(fsdp_modules[max(0, index - depth):index]))
if forward:
fsdp_module.set_modules_to_forward_prefetch(forward)
if backward:
fsdp_module.set_modules_to_backward_prefetch(backward)
if not dist.is_initialized() or dist.get_rank() == 0:
print(
"BF16_FSDP_KNOBS " + json.dumps({
"module_count": len(fsdp_modules),
"pg_alloc": args.pg_alloc,
"prefetch_depth": args.prefetch_depth,
}, sort_keys=True),
flush=True,
)
return module
moduleloader.load_module_from_path = _load_module_from_path
trainer_module.build_tracker = lambda *_args, **_kwargs: DummyTracker()
original_trainer_init = trainer_module.Trainer.__init__
original_trainer_run = trainer_module.Trainer.run
class _Recorder:
def on_training_step_end(self, _method, metrics, iteration=0):
del iteration
_step_times.append(float(metrics["step_time_sec"]))
def _trainer_init(self, *init_args, **init_kwargs):
original_trainer_init(self, *init_args, **init_kwargs)
self.callbacks._callbacks.pop("validation", None)
self.callbacks._callbacks["_benchmark_recorder"] = _Recorder()
def _trainer_run(self, method, **kwargs):
vae = getattr(getattr(method, "student", None), "vae", None)
if vae is not None:
method.student.vae = None
del vae
gc.collect()
torch.cuda.empty_cache()
if not dist.is_initialized() or dist.get_rank() == 0:
print("BF16_SETUP " + json.dumps({"unused_vae": "unloaded"}), flush=True)
torch.cuda.reset_peak_memory_stats()
return original_trainer_run(self, method, **kwargs)
trainer_module.Trainer.__init__ = _trainer_init
trainer_module.Trainer.run = _trainer_run
train_main(argparse.Namespace(config=args.config, dry_run=False), overrides=overrides or None)
world = get_world_group()
times_by_rank = [None] * world.world_size
dist.all_gather_object(times_by_rank, _step_times, group=world.cpu_group)
peak_memory = torch.tensor(
[torch.cuda.max_memory_allocated(), torch.cuda.max_memory_reserved()],
device="cuda",
dtype=torch.int64,
)
dist.all_reduce(peak_memory, op=dist.ReduceOp.MAX)
if world.rank == 0:
per_step_max = [max(values) for values in zip(*times_by_rank, strict=True)]
measured = per_step_max[10:30]
if len(measured) != 20:
raise RuntimeError(f"expected 30 steps, got {len(per_step_max)}")
median = statistics.median(measured)
print(
"BF16_RESULT " + json.dumps({
"median_step_sec": median,
"model_mfu_percent": 14.444115 / median,
"samples_per_second": world.world_size / median,
"world_size": world.world_size,
"pg_alloc": args.pg_alloc,
"prefetch_depth": args.prefetch_depth,
"force_sync": args.force_sync,
"force_reshard_after_backward": args.force_reshard_after_backward,
"forced_reshard_calls": _forced_reshard_calls,
"gradient_sync_false_calls": _gradient_sync_calls.count(False),
"gradient_sync_true_calls": _gradient_sync_calls.count(True),
"peak_allocated_max_rank_bytes": int(peak_memory[0]),
"peak_reserved_max_rank_bytes": int(peak_memory[1]),
}, sort_keys=True),
flush=True,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,237 @@
#!/usr/bin/env python3
"""Scratch gate: let fused AdamW apply the already-computed clip scale."""
from __future__ import annotations
import json
import sys
from typing import Any
import torch
import torch.distributed as dist
import benchmark_fastvideo_train_pack_d016 as benchmark
import fastvideo.train.callbacks.grad_clip as grad_clip_module
from fastvideo.train.callbacks.grad_clip import GradNormClipCallback
from fastvideo.train.methods.base import TrainingMethod
import fastvideo.training.training_utils as training_utils
_COUNTS = {
"fused_clips": 0,
"fallback_clips": 0,
"scaled_optimizer_steps": 0,
"cleared_scales": 0,
"ordinary_cuda_scales": 0,
}
_SCALED_OPTIMIZERS: dict[int, torch.optim.Optimizer] = {}
_VALIDATED_PAIRS: set[tuple[int, int]] = set()
_ACTIVE_OPTIMIZER: torch.optim.Optimizer | None = None
def _uses_fused_adamw(optimizer: torch.optim.Optimizer | None) -> bool:
return (
isinstance(optimizer, torch.optim.AdamW)
and bool(optimizer.param_groups)
and all(group.get("fused") is True for group in optimizer.param_groups)
and getattr(optimizer, "grad_scale", None) is None
)
def _validate_pair(module: torch.nn.Module, optimizer: torch.optim.Optimizer) -> None:
key = (id(module), id(optimizer))
if key in _VALIDATED_PAIRS:
return
module_params = {id(parameter) for parameter in module.parameters() if parameter.requires_grad}
optimizer_params = {
id(parameter)
for group in optimizer.param_groups
for parameter in group["params"]
if parameter.requires_grad
}
if module_params != optimizer_params:
raise RuntimeError(
"fused clip target and optimizer parameter sets differ: "
f"target={len(module_params)} optimizer={len(optimizer_params)}"
)
_VALIDATED_PAIRS.add(key)
def _install_fused_clip() -> None:
original_step = TrainingMethod.optimizers_schedulers_step
original_scale_grads = training_utils._clip_grads_with_norm_
def _defer_scale(
parameters: torch.Tensor | list[torch.Tensor],
max_norm: float,
total_norm: torch.Tensor,
foreach: bool | None = None,
) -> None:
del parameters, foreach
if _ACTIVE_OPTIMIZER is None:
raise RuntimeError("missing active optimizer while computing a target norm")
if isinstance(total_norm, torch.distributed.tensor.DTensor):
raise RuntimeError("clip norm remained a DTensor instead of becoming a full tensor")
if total_norm.ndim != 0 or total_norm.device.type != "cuda":
raise RuntimeError(f"expected an ordinary CUDA scalar norm, got {total_norm.shape} on {total_norm.device}")
clip_coef = float(max_norm) / (total_norm + 1e-6)
grad_scale = torch.clamp(clip_coef, max=1.0).reciprocal()
if isinstance(grad_scale, torch.distributed.tensor.DTensor) or grad_scale.ndim != 0 or grad_scale.device.type != "cuda":
raise RuntimeError(f"expected an ordinary CUDA scalar grad_scale, got {type(grad_scale)}")
_ACTIVE_OPTIMIZER.grad_scale = grad_scale
_COUNTS["fused_clips"] += 1
_COUNTS["ordinary_cuda_scales"] += 1
def _before_step(self: GradNormClipCallback, method: TrainingMethod, iteration: int = 0) -> None:
global _ACTIVE_OPTIMIZER
if self._pending_grad_norms:
raise RuntimeError("Previous gradient norms were not logged after the optimizer step")
if self._max_grad_norm <= 0.0:
return
active = {id(optimizer): optimizer for optimizer in method.get_optimizers(iteration)}
optimizer_by_role = method._optimizer_dict
tracker = getattr(method, "tracker", None)
for name, module in method.get_grad_clip_targets(iteration).items():
optimizer = optimizer_by_role.get(name)
if optimizer is not None and id(optimizer) in active and _uses_fused_adamw(optimizer):
_validate_pair(module, optimizer)
_SCALED_OPTIMIZERS[id(optimizer)] = optimizer
_ACTIVE_OPTIMIZER = optimizer
training_utils._clip_grads_with_norm_ = _defer_scale
try:
grad_norm = grad_clip_module.clip_grad_norm_if_needed(module, self._max_grad_norm)
finally:
training_utils._clip_grads_with_norm_ = original_scale_grads
_ACTIVE_OPTIMIZER = None
if getattr(optimizer, "grad_scale", None) is None:
raise RuntimeError("fused clip did not install optimizer.grad_scale")
else:
_COUNTS["fallback_clips"] += 1
grad_norm = grad_clip_module.clip_grad_norm_if_needed(module, self._max_grad_norm)
if self._log_grad_norms and tracker is not None and grad_norm is not None:
self._pending_grad_norms.append((tracker, name, grad_norm, iteration))
def _step(self: TrainingMethod, iteration: int) -> None:
scaled = dict(_SCALED_OPTIMIZERS)
try:
original_step(self, iteration)
_COUNTS["scaled_optimizer_steps"] += len(scaled)
finally:
for optimizer in scaled.values():
if hasattr(optimizer, "grad_scale"):
del optimizer.grad_scale
_COUNTS["cleared_scales"] += 1
_SCALED_OPTIMIZERS.clear()
GradNormClipCallback.on_before_optimizer_step = _before_step
TrainingMethod.optimizers_schedulers_step = _step
def _assert_tiny_exact_parity() -> None:
if not torch.cuda.is_available():
raise RuntimeError("tiny fused AdamW parity requires CUDA")
device = torch.device("cuda")
control = [torch.nn.Parameter(torch.tensor([1.25], device=device)),
torch.nn.Parameter(torch.tensor([-0.75], device=device))]
candidate = [torch.nn.Parameter(parameter.detach().clone()) for parameter in control]
kwargs = dict(lr=3e-4, betas=(0.8, 0.95), eps=1e-8, weight_decay=0.01, fused=True)
control_optimizer = torch.optim.AdamW(control, **kwargs)
candidate_optimizer = torch.optim.AdamW(candidate, **kwargs)
gradients = ((3.0, 4.0), (-4.0, 3.0), (3.0, -4.0), (-3.0, -4.0))
candidate_grad_contract: str | None = None
for step, values in enumerate(gradients, 1):
for parameter, value in zip(control, values, strict=True):
parameter.grad = torch.tensor([value], device=device)
for parameter, value in zip(candidate, values, strict=True):
parameter.grad = torch.tensor([value], device=device)
control_norm = training_utils._get_total_norm([parameter.grad for parameter in control])
candidate_norm = training_utils._get_total_norm([parameter.grad for parameter in candidate])
torch.testing.assert_close(control_norm, candidate_norm, rtol=0, atol=0)
# Norm is exactly five. Exercise the recipe's max_norm=1 clipping and
# the unclipped path on alternating steps.
max_norm = 1.0 if step % 2 else 10.0
training_utils._clip_grads_with_norm_(control, max_norm, control_norm)
expected_effective_grads = [parameter.grad.detach().clone() for parameter in control]
raw_candidate_grads = [parameter.grad.detach().clone() for parameter in candidate]
clip_coef = float(max_norm) / (candidate_norm + 1e-6)
candidate_optimizer.grad_scale = torch.clamp(clip_coef, max=1.0).reciprocal()
control_optimizer.step()
candidate_optimizer.step()
del candidate_optimizer.grad_scale
for index, (control_parameter, candidate_parameter) in enumerate(zip(control, candidate, strict=True)):
torch.testing.assert_close(control_parameter.grad, expected_effective_grads[index], rtol=0, atol=0)
torch.testing.assert_close(control_parameter, candidate_parameter, rtol=0, atol=0)
control_state = control_optimizer.state[control_parameter]
candidate_state = candidate_optimizer.state[candidate_parameter]
for key in ("step", "exp_avg", "exp_avg_sq"):
torch.testing.assert_close(control_state[key], candidate_state[key], rtol=0, atol=0)
assert control_parameter.dtype == torch.float32
assert control_state["exp_avg"].dtype == torch.float32
assert control_state["exp_avg_sq"].dtype == torch.float32
if step % 2:
wrote_scaled = all(torch.equal(parameter.grad, expected)
for parameter, expected in zip(candidate, expected_effective_grads, strict=True))
preserved_raw = all(torch.equal(parameter.grad, raw)
for parameter, raw in zip(candidate, raw_candidate_grads, strict=True))
if wrote_scaled == preserved_raw:
raise RuntimeError("candidate grads were neither uniquely scaled nor uniquely preserved")
observed = "scaled_writeback" if wrote_scaled else "raw_preserved"
if candidate_grad_contract is not None and candidate_grad_contract != observed:
raise RuntimeError(f"candidate gradient contract changed: {candidate_grad_contract} -> {observed}")
candidate_grad_contract = observed
torch.cuda.synchronize()
print("FUSED_CLIP_PARITY " + json.dumps({
"device": torch.cuda.get_device_name(),
"steps": len(gradients),
"params_moments_steps_bit_exact": True,
"registered_parameter_dtype": "torch.float32",
"moment_dtype": "torch.float32",
"candidate_grad_contract": candidate_grad_contract,
}, sort_keys=True), flush=True)
def _report_counts() -> None:
if not dist.is_initialized():
return
keys = list(_COUNTS)
counts = torch.tensor([_COUNTS[key] for key in keys], device="cuda", dtype=torch.int64)
dist.all_reduce(counts)
world_size = dist.get_world_size()
totals = dict(zip(keys, counts.cpu().tolist(), strict=True))
expected = benchmark.EXPECTED_STEPS * world_size
if totals != {
"fused_clips": expected,
"fallback_clips": 0,
"scaled_optimizer_steps": expected,
"cleared_scales": expected,
"ordinary_cuda_scales": expected,
}:
raise RuntimeError(f"unexpected fused clip counts: {totals}")
if dist.get_rank() == 0:
print("FUSED_CLIP_COUNTS " + json.dumps({
"world_size": world_size,
"per_rank_steps": benchmark.EXPECTED_STEPS,
"totals": totals,
}, sort_keys=True), flush=True)
def main() -> None:
if "--parity-only" in sys.argv:
sys.argv.remove("--parity-only")
_assert_tiny_exact_parity()
return
_install_fused_clip()
benchmark.main()
_report_counts()
if __name__ == "__main__":
main()
@@ -0,0 +1,353 @@
#!/usr/bin/env python3
"""Scratch A/B for deferring grad-norm scalar materialization past AdamW launch."""
from __future__ import annotations
import argparse
from collections import Counter
import gc
import hashlib
import json
import statistics
import time
from typing import Any
import torch
import torch.distributed as dist
EXPECTED_STEPS = 30
WARMUP_STEPS = 10
MFU_NUMERATOR = 14.444115
_step_times: list[float] = []
_step_starts: list[float] = []
_synced_end: float | None = None
_grad_norm_records: list[dict[str, Any]] = []
_pending_norm_logs: list[tuple[Any, str, torch.Tensor, int]] = []
_semantic_counts: Counter[str] = Counter()
_optimizer_returned_iteration = 0
_original_log: Any = None
def _rank_zero() -> bool:
return not dist.is_initialized() or dist.get_rank() == 0
def _log(self: Any, metrics: dict[str, Any], step: int) -> None:
if _original_log is None:
raise RuntimeError("tracker log hook was not initialized")
_original_log(self, metrics, step)
for key, raw_value in metrics.items():
if not key.startswith("grad_norm/"):
continue
value = float(raw_value)
after_optimizer = _optimizer_returned_iteration >= int(step)
_semantic_counts["grad_norm_log_calls"] += 1
_semantic_counts[
"grad_norm_logs_after_optimizer" if after_optimizer else "grad_norm_logs_before_optimizer"
] += 1
_grad_norm_records.append({
"step": int(step),
"key": key,
"value": value,
"after_optimizer": after_optimizer,
})
if "step_time_sec" in metrics and _rank_zero():
print(
"BF16_STEP " + json.dumps({
"step": step,
"step_time_sec": float(metrics["step_time_sec"]),
}),
flush=True,
)
def _flush_deferred_norm_logs(iteration: int) -> None:
global _pending_norm_logs
pending = _pending_norm_logs
_pending_norm_logs = []
for tracker, key, norm, queued_iteration in pending:
if queued_iteration != iteration:
raise RuntimeError(
f"deferred grad norm from step {queued_iteration} reached optimizer step {iteration}"
)
value = float(norm.item())
_semantic_counts["deferred_norm_materializations"] += 1
if value > 0.0:
tracker.log({key: value}, iteration)
else:
_semantic_counts["nonpositive_norm_log_skips"] += 1
def _wall_intervals(starts: list[float], synced_end: float) -> list[float]:
if not starts:
raise RuntimeError("no training-step wall starts were recorded")
return [next_start - start for start, next_start in zip(starts, starts[1:])] + [synced_end - starts[-1]]
def _series_digest(records: list[dict[str, Any]]) -> str:
semantic_series = [(record["step"], record["key"], record["value"]) for record in records]
return hashlib.sha256(json.dumps(semantic_series, separators=(",", ":")).encode()).hexdigest()
def _self_test() -> None:
global _optimizer_returned_iteration, _pending_norm_logs
events: list[Any] = ["optimizer"]
class _Norm:
def item(self) -> float:
events.append("item")
return 1.25
class _Tracker:
def log(self, metrics: dict[str, float], step: int) -> None:
events.append(("log", metrics, step))
_optimizer_returned_iteration = 3
_pending_norm_logs = [(_Tracker(), "grad_norm/student", _Norm(), 3)] # type: ignore[list-item]
_flush_deferred_norm_logs(3)
assert events == ["optimizer", "item", ("log", {"grad_norm/student": 1.25}, 3)]
assert _wall_intervals([1.0, 2.0, 4.0], 7.0) == [1.0, 2.0, 3.0]
assert not _pending_norm_logs
print("SELF_TEST_OK")
def main() -> None:
global _optimizer_returned_iteration, _original_log, _synced_end
parser = argparse.ArgumentParser()
parser.add_argument("--config")
parser.add_argument("--defer-grad-norm-materialization", action="store_true")
parser.add_argument("--source-deferred-grad-norm", action="store_true")
parser.add_argument("--self-test", action="store_true")
args, overrides = parser.parse_known_args()
if args.self_test:
_self_test()
return
if not args.config:
parser.error("--config is required unless --self-test is used")
if args.defer_grad_norm_materialization and args.source_deferred_grad_norm:
parser.error("scratch and source-deferred modes are mutually exclusive")
from fastvideo.training.trackers import DummyTracker
_original_log = DummyTracker.log
DummyTracker.log = _log
import fastvideo.train.callbacks.grad_clip as grad_clip_module
import fastvideo.train.trainer as trainer_module
from fastvideo.distributed import get_world_group
from fastvideo.train.callbacks.grad_clip import GradNormClipCallback
from fastvideo.train.entrypoint.train import main as train_main
from fastvideo.train.methods.base import TrainingMethod
from fastvideo.training.training_utils import clip_grad_norm_while_handling_failing_dtensor_cases
original_clip_grad_norm = grad_clip_module.clip_grad_norm_if_needed
original_method_optimizer_step = TrainingMethod.optimizers_schedulers_step
original_trainer_init = trainer_module.Trainer.__init__
original_trainer_iter = trainer_module.Trainer._iter_dataloader
original_trainer_run = trainer_module.Trainer.run
def _counted_clip_grad_norm(module: torch.nn.Module, max_grad_norm: float) -> float:
_semantic_counts["clip_calls"] += 1
return original_clip_grad_norm(module, max_grad_norm)
def _deferred_before_optimizer(
self: GradNormClipCallback,
method: TrainingMethod,
iteration: int = 0,
) -> None:
if self._max_grad_norm <= 0.0:
return
if _pending_norm_logs:
raise RuntimeError("previous deferred grad norm was not flushed")
tracker = getattr(method, "tracker", None)
for name, module in method.get_grad_clip_targets(iteration).items():
_semantic_counts["clip_calls"] += 1
norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[parameter for parameter in module.parameters()],
self._max_grad_norm,
foreach=None,
)
if norm is None:
_semantic_counts["missing_norms"] += 1
elif self._log_grad_norms and tracker is not None:
_pending_norm_logs.append((tracker, f"grad_norm/{name}", norm, iteration))
_semantic_counts["deferred_norm_queues"] += 1
_semantic_counts["pending_peak"] = max(
_semantic_counts["pending_peak"],
len(_pending_norm_logs),
)
if args.defer_grad_norm_materialization:
GradNormClipCallback.on_before_optimizer_step = _deferred_before_optimizer
else:
grad_clip_module.clip_grad_norm_if_needed = _counted_clip_grad_norm
def _optimizer_step(self: TrainingMethod, iteration: int) -> None:
global _optimizer_returned_iteration
_semantic_counts["optimizer_method_calls"] += 1
original_method_optimizer_step(self, iteration)
_optimizer_returned_iteration = iteration
if args.defer_grad_norm_materialization:
_flush_deferred_norm_logs(iteration)
TrainingMethod.optimizers_schedulers_step = _optimizer_step
trainer_module.build_tracker = lambda *_args, **_kwargs: DummyTracker()
class _Recorder:
def on_training_step_end(self, _method: Any, metrics: dict[str, Any], iteration: int = 0) -> None:
del iteration
_step_times.append(float(metrics["step_time_sec"]))
def _trainer_init(self: Any, *init_args: Any, **init_kwargs: Any) -> None:
original_trainer_init(self, *init_args, **init_kwargs)
if int(self.training_config.loop.gradient_accumulation_steps or 1) != 1:
raise RuntimeError("scratch grad-norm benchmark requires gradient accumulation 1")
grad_callbacks = [
callback for callback in self.callbacks._callbacks.values()
if isinstance(callback, GradNormClipCallback)
]
if len(grad_callbacks) != 1 or not grad_callbacks[0]._log_grad_norms:
raise RuntimeError("scratch grad-norm benchmark requires one logging GradNormClipCallback")
self.callbacks._callbacks.pop("validation", None)
self.callbacks._callbacks["_benchmark_recorder"] = _Recorder()
def _timed_iter(self: Any, dataloader: Any) -> Any:
iterator = original_trainer_iter(self, dataloader)
while True:
_step_starts.append(time.perf_counter())
yield next(iterator)
def _trainer_run(self: Any, method: TrainingMethod, **kwargs: Any) -> Any:
global _synced_end
vae = getattr(getattr(method, "student", None), "vae", None)
if vae is not None:
method.student.vae = None
del vae
gc.collect()
torch.cuda.empty_cache()
if _rank_zero():
print("BF16_SETUP " + json.dumps({"unused_vae": "unloaded"}), flush=True)
result = original_trainer_run(self, method, **kwargs)
torch.cuda.synchronize()
_semantic_counts["final_cuda_sync_calls"] += 1
_synced_end = time.perf_counter()
return result
trainer_module.Trainer.__init__ = _trainer_init
trainer_module.Trainer._iter_dataloader = _timed_iter
trainer_module.Trainer.run = _trainer_run
train_main(argparse.Namespace(config=args.config, dry_run=False), overrides=overrides or None)
if _synced_end is None:
raise RuntimeError("final synchronized wall endpoint was not recorded")
wall_intervals = _wall_intervals(_step_starts, _synced_end)
if len(_step_times) != EXPECTED_STEPS or len(wall_intervals) != EXPECTED_STEPS:
raise RuntimeError(
f"expected {EXPECTED_STEPS} steps, got internal={len(_step_times)} wall={len(wall_intervals)}"
)
if _pending_norm_logs:
raise RuntimeError(f"{len(_pending_norm_logs)} deferred grad norms remain after training")
payload = {
"internal_step_times": _step_times,
"wall_intervals": wall_intervals,
"end_to_end_wall_sec": _synced_end - _step_starts[0],
"measured_window_wall_sec": _synced_end - _step_starts[WARMUP_STEPS],
"semantics": dict(_semantic_counts),
"grad_norm_records": _grad_norm_records,
}
world = get_world_group()
payloads: list[dict[str, Any] | None] = [None] * world.world_size
dist.all_gather_object(payloads, payload, group=world.cpu_group)
if world.rank != 0:
return
rank_payloads = [rank_payload for rank_payload in payloads if rank_payload is not None]
if len(rank_payloads) != world.world_size:
raise RuntimeError("missing rank payload")
internal_slowest = [
max(rank_payload["internal_step_times"][index] for rank_payload in rank_payloads)
for index in range(EXPECTED_STEPS)
]
wall_slowest = [
max(rank_payload["wall_intervals"][index] for rank_payload in rank_payloads)
for index in range(EXPECTED_STEPS)
]
internal_measured = internal_slowest[WARMUP_STEPS:]
wall_measured = wall_slowest[WARMUP_STEPS:]
internal_median = statistics.median(internal_measured)
wall_median = statistics.median(wall_measured)
semantics_by_rank = []
expected_relation = (
"grad_norm_logs_after_optimizer"
if args.defer_grad_norm_materialization or args.source_deferred_grad_norm
else "grad_norm_logs_before_optimizer"
)
for rank, rank_payload in enumerate(rank_payloads):
semantics = rank_payload["semantics"]
records = rank_payload["grad_norm_records"]
if semantics.get("clip_calls") != EXPECTED_STEPS:
raise RuntimeError(f"rank {rank} recorded {semantics.get('clip_calls')} clip calls")
if semantics.get("optimizer_method_calls") != EXPECTED_STEPS:
raise RuntimeError(f"rank {rank} recorded {semantics.get('optimizer_method_calls')} optimizer calls")
if semantics.get("grad_norm_log_calls") != EXPECTED_STEPS:
raise RuntimeError(f"rank {rank} recorded {semantics.get('grad_norm_log_calls')} grad-norm logs")
if semantics.get(expected_relation) != EXPECTED_STEPS:
raise RuntimeError(f"rank {rank} failed expected log ordering: {semantics}")
if args.defer_grad_norm_materialization and semantics.get("deferred_norm_materializations") != EXPECTED_STEPS:
raise RuntimeError(f"rank {rank} failed to materialize every deferred norm")
semantics_by_rank.append({
"rank": rank,
"counts": semantics,
"grad_norm_digest": _series_digest(records),
"first_grad_norm": records[0],
"last_grad_norm": records[-1],
})
print(
"BF16_TIMING " + json.dumps({
"internal_slowest_rank_sec": internal_slowest,
"true_wall_slowest_rank_sec": wall_slowest,
"final_cuda_sync": True,
}, sort_keys=True),
flush=True,
)
print(
"BF16_SEMANTICS " + json.dumps({
"defer_grad_norm_materialization": args.defer_grad_norm_materialization,
"source_deferred_grad_norm": args.source_deferred_grad_norm,
"by_rank": semantics_by_rank,
"rank0_grad_norm_series": rank_payloads[0]["grad_norm_records"],
}, sort_keys=True),
flush=True,
)
print(
"BF16_RESULT " + json.dumps({
"defer_grad_norm_materialization": args.defer_grad_norm_materialization,
"source_deferred_grad_norm": args.source_deferred_grad_norm,
"world_size": world.world_size,
"internal_median_step_sec": internal_median,
"internal_model_mfu_percent": MFU_NUMERATOR / internal_median,
"true_wall_median_step_sec": wall_median,
"true_wall_model_mfu_percent": MFU_NUMERATOR / wall_median,
"true_wall_sum_slowest_intervals_sec": sum(wall_measured),
"true_wall_mean_slowest_interval_sec": statistics.mean(wall_measured),
"true_measured_window_max_rank_sec": max(
rank_payload["measured_window_wall_sec"] for rank_payload in rank_payloads
),
"true_end_to_end_max_rank_sec": max(
rank_payload["end_to_end_wall_sec"] for rank_payload in rank_payloads
),
"samples_per_second_from_true_wall_median": world.world_size / wall_median,
}, sort_keys=True),
flush=True,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,306 @@
#!/usr/bin/env python3
"""Scratch gate: pinned cuBLASLt dgrad/wgrad algos for the packed projections.
Follows research-plan item 2 after the algo sweep found an 8-9% weighted GEMM
band win with bit-exact parity. This driver builds the ``lt_pinned_ops`` C++
extension, tunes the eight conservative backward cases (dgrad/wgrad where the
sweep beat torch by >2%; forward stays stock nvjet with fused bias), wraps the
matching block-level ``nn.Linear`` modules in an autograd function whose
backward dispatches to the pinned algos, and runs the frozen packed benchmark
harness. Bias gradients fall back to an eager ``sum(0)`` (integration would
fuse BGRADB later).
"""
from __future__ import annotations
import ctypes
import os
import statistics
import torch
import torch.nn.functional as F
from torch.utils.cpp_extension import load
import benchmark_fastvideo_train_pack_d016 as benchmark
import fastvideo.train.trainer as trainer_module
CUDA_R_16BF = 14
CUDA_R_32F = 0
CUBLAS_COMPUTE_32F = 68
OP_N, OP_T = 0, 1
DESC_TRANSA, DESC_TRANSB = 3, 4
PREF_MAX_WORKSPACE = 1
WORKSPACE_BYTES = 128 * 1024 * 1024
MAX_ALGOS = 48
HIDDEN, FFN = 4096, 16384
VIDEO_TOKENS, TEXT_TOKENS = 11 * 15 * 26, 1024
LOCAL_BATCH = int(os.environ.get("FASTVIDEO_BENCH_LOCAL_BATCH_SIZE", "1"))
MIN_WIN_PCT = 2.0
_ext = load(
name="lt_pinned_ops",
sources=[os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "probes", "lt_pinned_ops.cpp")],
extra_ldflags=["-lcublasLt"],
with_cuda=True,
verbose=False,
)
class HeuristicResult(ctypes.Structure):
_fields_ = [
("algo", ctypes.c_uint64 * 8),
("workspaceSize", ctypes.c_size_t),
("state", ctypes.c_int),
("wavesCount", ctypes.c_float),
("reserved", ctypes.c_int * 4),
]
def _tune_case(lt, handle, workspace, name, m, n, k, transa, transb, a_shape, b_shape, d_shape, torch_fn):
device = torch.device("cuda")
a = torch.randn(a_shape, device=device, dtype=torch.bfloat16)
b = torch.randn(b_shape, device=device, dtype=torch.bfloat16)
d = torch.empty(d_shape, device=device, dtype=torch.bfloat16)
reference = torch_fn(a, b)
def check(status, what):
if status != 0:
raise RuntimeError(f"{what} failed: {status}")
op_desc = ctypes.c_void_p()
check(lt.cublasLtMatmulDescCreate(ctypes.byref(op_desc), CUBLAS_COMPUTE_32F, CUDA_R_32F), "descCreate")
for attr, value in ((DESC_TRANSA, transa), (DESC_TRANSB, transb)):
v = ctypes.c_int32(value)
check(lt.cublasLtMatmulDescSetAttribute(op_desc, attr, ctypes.byref(v), 4), "descSet")
lda, ldb = a_shape[1], b_shape[1]
def layout(rows, cols, ld):
h = ctypes.c_void_p()
check(lt.cublasLtMatrixLayoutCreate(ctypes.byref(h), CUDA_R_16BF, rows, cols, ctypes.c_int64(ld)), "layout")
return h
layout_a = layout(lda if transa == OP_T else m, m if transa == OP_T else k, lda)
layout_b = layout(ldb if transb == OP_T else k, k if transb == OP_T else n, ldb)
layout_d = layout(m, n, m)
pref = ctypes.c_void_p()
check(lt.cublasLtMatmulPreferenceCreate(ctypes.byref(pref)), "prefCreate")
ws = ctypes.c_size_t(WORKSPACE_BYTES)
check(lt.cublasLtMatmulPreferenceSetAttribute(pref, PREF_MAX_WORKSPACE, ctypes.byref(ws), 8), "prefSet")
results = (HeuristicResult * MAX_ALGOS)()
found = ctypes.c_int(0)
check(lt.cublasLtMatmulAlgoGetHeuristic(handle, op_desc, layout_a, layout_b, layout_d, layout_d,
pref, MAX_ALGOS, results, ctypes.byref(found)), "heuristic")
alpha, beta = ctypes.c_float(1.0), ctypes.c_float(0.0)
stream = ctypes.c_void_p(torch.cuda.current_stream().cuda_stream)
def run_lt(index):
check(lt.cublasLtMatmul(handle, op_desc, ctypes.byref(alpha),
ctypes.c_void_p(a.data_ptr()), layout_a,
ctypes.c_void_p(b.data_ptr()), layout_b,
ctypes.byref(beta),
ctypes.c_void_p(d.data_ptr()), layout_d,
ctypes.c_void_p(d.data_ptr()), layout_d,
ctypes.byref(results[index], 0),
ctypes.c_void_p(workspace.data_ptr()), ctypes.c_size_t(WORKSPACE_BYTES),
stream), "ltMatmul")
def time_fn(fn):
for _ in range(6):
fn()
torch.cuda.synchronize()
events = [(torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True)) for _ in range(15)]
for start, end in events:
start.record()
fn()
end.record()
torch.cuda.synchronize()
return statistics.median(start.elapsed_time(end) for start, end in events)
torch_ms = time_fn(lambda: torch_fn(a, b))
best = None
for index in range(found.value):
if results[index].state != 0:
continue
try:
run_lt(index)
except RuntimeError:
continue
torch.cuda.synchronize()
if (d.float() - reference.float()).abs().max().item() > 0.0:
continue
lt_ms = time_fn(lambda: run_lt(index))
if best is None or lt_ms < best[0]:
best = (lt_ms, index)
win_pct = 100.0 * (torch_ms - best[0]) / torch_ms if best else -1.0
spec = None
if best and win_pct >= MIN_WIN_PCT:
spec = {
"m": m, "n": n, "k": k, "transa": transa == OP_T, "transb": transb == OP_T,
"lda": lda, "ldb": ldb, "d_rows": d_shape[0], "d_cols": d_shape[1],
"algo": bytes(results[best[1]].algo),
}
print(f"LT_PIN_TUNE {name} torch_ms={torch_ms:.4f} best_lt_ms={best[0] if best else -1:.4f} "
f"win_pct={win_pct:.2f} pinned={spec is not None}", flush=True)
return spec
def _tune_all_blobs() -> dict[tuple[int, int], dict]:
"""Rank-0 tuning: returns per-shape dgrad/wgrad win metadata + algo blobs."""
lt = None
for so in ("libcublasLt.so.13", "libcublasLt.so.12", "libcublasLt.so"):
try:
lt = ctypes.CDLL(so)
break
except OSError:
continue
if lt is None:
raise OSError("libcublasLt not found")
handle = ctypes.c_void_p()
if lt.cublasLtCreate(ctypes.byref(handle)) != 0:
raise RuntimeError("cublasLtCreate failed")
workspace = torch.empty(WORKSPACE_BYTES, dtype=torch.uint8, device="cuda")
mapping: dict[tuple[int, int], dict] = {}
shapes = [
("self_qkv", VIDEO_TOKENS * LOCAL_BATCH, HIDDEN, 3 * HIDDEN),
("video_dd", VIDEO_TOKENS * LOCAL_BATCH, HIDDEN, HIDDEN),
("text_kv", TEXT_TOKENS * LOCAL_BATCH, HIDDEN, 2 * HIDDEN),
("ffn_up", VIDEO_TOKENS * LOCAL_BATCH, HIDDEN, FFN),
("ffn_down", VIDEO_TOKENS * LOCAL_BATCH, FFN, HIDDEN),
]
for name, rows, in_features, out_features in shapes:
dgrad_spec = _tune_case(
lt, handle, workspace, f"{name}:dgrad",
in_features, rows, out_features, OP_N, OP_N,
(out_features, in_features), (rows, out_features), (rows, in_features),
lambda a, b: b @ a)
wgrad_spec = _tune_case(
lt, handle, workspace, f"{name}:wgrad",
in_features, out_features, rows, OP_N, OP_T,
(rows, in_features), (rows, out_features), (out_features, in_features),
lambda a, b: b.t() @ a)
mapping[(in_features, out_features)] = {"rows": rows, "dgrad": dgrad_spec, "wgrad": wgrad_spec}
return mapping
def _register_from_spec(spec: dict | None) -> int:
if spec is None:
return -1
blob = torch.frombuffer(bytearray(spec["algo"]), dtype=torch.uint8).clone()
return _ext.register_case(spec["m"], spec["n"], spec["k"], spec["transa"], spec["transb"],
spec["lda"], spec["ldb"], spec["d_rows"], spec["d_cols"], blob)
def _tune_and_broadcast() -> dict[tuple[int, int], tuple[int, int, int]]:
import torch.distributed as dist
payload = [None]
if not dist.is_initialized() or dist.get_rank() == 0:
payload = [_tune_all_blobs()]
if dist.is_initialized():
dist.broadcast_object_list(payload, src=0)
mapping: dict[tuple[int, int], tuple[int, int, int]] = {}
for key, entry in payload[0].items():
mapping[key] = (entry["rows"], _register_from_spec(entry["dgrad"]), _register_from_spec(entry["wgrad"]))
return mapping
@torch.library.custom_op("fastvideo::lt_pinned_dgrad", mutates_args=())
def lt_pinned_dgrad(grad_out: torch.Tensor, weight: torch.Tensor, case_id: int) -> torch.Tensor:
return _ext.lt_mm(weight.contiguous(), grad_out.contiguous(), case_id)
@lt_pinned_dgrad.register_fake
def _(grad_out, weight, case_id):
return grad_out.new_empty((grad_out.shape[0], weight.shape[1]))
@torch.library.custom_op("fastvideo::lt_pinned_wgrad", mutates_args=())
def lt_pinned_wgrad(grad_out: torch.Tensor, saved_input: torch.Tensor, case_id: int) -> torch.Tensor:
return _ext.lt_mm(saved_input.contiguous(), grad_out.contiguous(), case_id)
@lt_pinned_wgrad.register_fake
def _(grad_out, saved_input, case_id):
return grad_out.new_empty((grad_out.shape[1], saved_input.shape[1]))
@torch.library.custom_op("fastvideo::lt_pinned_linear", mutates_args=())
def lt_pinned_linear(x2d: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor | None,
rows: int, dgrad_id: int, wgrad_id: int) -> torch.Tensor:
return F.linear(x2d, weight, bias)
@lt_pinned_linear.register_fake
def _(x2d, weight, bias, rows, dgrad_id, wgrad_id):
return x2d.new_empty((x2d.shape[0], weight.shape[0]))
def _lt_linear_setup(ctx, inputs, output):
x2d, weight, bias, rows, dgrad_id, wgrad_id = inputs
ctx.save_for_backward(x2d, weight)
ctx.meta = (rows, dgrad_id, wgrad_id, bias is not None)
def _lt_linear_backward(ctx, grad_out):
x2d, weight = ctx.saved_tensors
rows, dgrad_id, wgrad_id, has_bias = ctx.meta
grad_out = grad_out.contiguous()
if grad_out.shape[0] == rows and dgrad_id >= 0:
grad_input = torch.ops.fastvideo.lt_pinned_dgrad(grad_out, weight, dgrad_id)
else:
grad_input = grad_out @ weight
if grad_out.shape[0] == rows and wgrad_id >= 0:
grad_weight = torch.ops.fastvideo.lt_pinned_wgrad(grad_out, x2d, wgrad_id)
else:
grad_weight = grad_out.t() @ x2d
grad_bias = grad_out.sum(0) if has_bias else None
return grad_input, grad_weight, grad_bias, None, None, None
lt_pinned_linear.register_autograd(_lt_linear_backward, setup_context=_lt_linear_setup)
def _patched_forward(module, rows, dgrad_id, wgrad_id):
def forward(x):
bias = module.bias if not module.skip_bias_add else None
shape = x.shape
x2d = x.reshape(-1, shape[-1])
out = torch.ops.fastvideo.lt_pinned_linear(x2d, module.weight, bias, rows, dgrad_id, wgrad_id)
out = out.reshape(*shape[:-1], out.shape[-1])
output_bias = module.bias if module.skip_bias_add else None
return out, output_bias
return forward
def _install(method) -> None:
from fastvideo.layers.linear import ReplicatedLinear
mapping = _tune_and_broadcast()
patched = 0
for name, module in method.student.transformer.named_modules():
if not isinstance(module, ReplicatedLinear) or "transformer_blocks" not in name:
continue
key = (module.input_size, module.output_size)
if key not in mapping:
continue
rows, dgrad_id, wgrad_id = mapping[key]
if dgrad_id < 0 and wgrad_id < 0:
continue
module.forward = _patched_forward(module, rows, dgrad_id, wgrad_id)
patched += 1
print(f"LT_PIN_INSTALLED modules={patched}", flush=True)
_original_run = trainer_module.Trainer.run
def _run_with_pins(self, method, **kwargs):
_install(method)
return _original_run(self, method, **kwargs)
trainer_module.Trainer.run = _run_with_pins
if __name__ == "__main__":
benchmark.main()
@@ -0,0 +1,161 @@
#!/usr/bin/env python3
"""Scratch A/B for returning raw LTX-2 velocity during training.
This wraps ``benchmark_fastvideo_train_ltx2_singleton_timestep.py`` so the
two candidates remain independently selectable:
* no flags: exact source control
* ``--raw-velocity``: bypass x0 conversion and velocity reconstruction
* ``--singleton-timestep``: singleton uniform T2V timestep only
* both flags: stacked candidates
The checkout is never edited. For the raw-velocity candidate the shared
transformer's module-level ``_to_denoised`` helper is temporarily replaced by
an identity and the modular LTX-2 adapter returns that raw output directly.
Validation is disabled by the wrapped benchmark, so inference semantics are
not part of this process-local monkeypatch.
"""
from __future__ import annotations
import json
from pathlib import Path
import runpy
import sys
from typing import Any
BASE_DRIVER = Path("/mnt/benchmark_fastvideo_train_ltx2_singleton_timestep.py")
EXPECTED_STEPS = 30
def _remove_flag(flag: str) -> bool:
found = False
kept = [sys.argv[0]]
for argument in sys.argv[1:]:
if argument == flag:
found = True
else:
kept.append(argument)
sys.argv[:] = kept
return found
def main() -> None:
raw_velocity = _remove_flag("--raw-velocity")
singleton_timestep = "--singleton-timestep" in sys.argv
self_test = "--self-test" in sys.argv
if not BASE_DRIVER.is_file():
raise RuntimeError(f"missing wrapped benchmark driver: {BASE_DRIVER}")
bypass_calls = 0
if raw_velocity and not self_test:
import torch
from fastvideo.forward_context import set_forward_context
import fastvideo.models.dits.ltx2 as ltx2_dit_module
from fastvideo.train.models.ltx2 import LTX2Model
def _return_velocity(
sample: torch.Tensor,
velocity: torch.Tensor,
sigma: torch.Tensor,
calc_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
del sample, sigma, calc_dtype
nonlocal bypass_calls
bypass_calls += 1
return velocity
def _predict_raw_velocity(
self: LTX2Model,
noisy_latents: torch.Tensor,
timestep: torch.Tensor,
batch: Any,
*,
conditional: bool,
cfg_uncond: dict[str, Any] | None = None,
attn_kind: str = "dense",
clean_x: torch.Tensor | None = None,
aug_t: torch.Tensor | None = None,
) -> torch.Tensor:
if clean_x is not None or aug_t is not None:
raise NotImplementedError("LTX2Model does not support teacher forcing inputs")
device_type = self.device.type
dtype = self._get_training_dtype()
if conditional:
text_dict = batch.conditional_dict
if text_dict is None:
raise RuntimeError("Missing conditional_dict in TrainingBatch")
else:
text_dict = self._get_uncond_text_dict(batch, cfg_uncond=cfg_uncond)
if attn_kind not in ("dense", "vsa"):
raise ValueError(f"Unknown attn_kind: {attn_kind!r}")
noisy_bcthw = noisy_latents.permute(0, 2, 1, 3, 4).to(dtype)
with torch.autocast(device_type, dtype=dtype), set_forward_context(
current_timestep=batch.timesteps,
attn_metadata=batch.attn_metadata,
forward_batch=self._make_rope_forward_batch(),
):
input_kwargs = self._build_distill_input_kwargs(
noisy_bcthw,
timestep,
text_dict,
)
transformer = self._get_transformer(timestep)
velocity = transformer(**input_kwargs)
if isinstance(velocity, tuple):
velocity = velocity[0]
return velocity.permute(0, 2, 1, 3, 4)
ltx2_dit_module._to_denoised = _return_velocity
LTX2Model.predict_noise = _predict_raw_velocity
runpy.run_path(str(BASE_DRIVER), run_name="__main__")
if self_test:
print(
"RAW_VELOCITY_SELF_TEST " + json.dumps({
"raw_velocity": raw_velocity,
"singleton_timestep": singleton_timestep,
}, sort_keys=True),
flush=True,
)
return
import torch.distributed as dist
from fastvideo.distributed import get_world_group
world = get_world_group() if dist.is_initialized() else None
rank = world.rank if world is not None else 0
world_size = world.world_size if world is not None else 1
counts: list[int | None] = [None] * world_size
if world is not None:
dist.all_gather_object(
counts,
bypass_calls,
group=world.cpu_group,
)
else:
counts[0] = bypass_calls
expected = EXPECTED_STEPS if raw_velocity else 0
if any(count != expected for count in counts):
raise RuntimeError(
f"raw-velocity bypass count mismatch: expected {expected} per rank, got {counts}"
)
if rank == 0:
print(
"BF16_VARIANT " + json.dumps({
"raw_velocity": raw_velocity,
"singleton_timestep": singleton_timestep,
"to_denoised_bypass_calls_by_rank": counts,
}, sort_keys=True),
flush=True,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,472 @@
#!/usr/bin/env python3
"""Scratch A/B for embedding one uniform LTX-2 T2V timestep per sample."""
from __future__ import annotations
import argparse
from collections import Counter
import gc
import hashlib
import json
import os
import statistics
import time
from typing import Any
import torch
import torch.distributed as dist
EXPECTED_STEPS = 30
WARMUP_STEPS = 10
MFU_NUMERATOR = 14.444115
LOCAL_BATCH_SIZE = int(os.environ.get("FASTVIDEO_BENCH_LOCAL_BATCH_SIZE", "1"))
if LOCAL_BATCH_SIZE <= 0:
raise ValueError("FASTVIDEO_BENCH_LOCAL_BATCH_SIZE must be positive")
GRAD_PROBE_STEP = WARMUP_STEPS
_step_times: list[float] = []
_step_starts: list[float] = []
_synced_end: float | None = None
_metrics_by_step: dict[int, dict[str, float]] = {}
_semantic_counts: Counter[str] = Counter()
_ada_grad_probe: dict[str, Any] | None = None
_peak_memory: dict[str, int] = {}
_optimizer_probe: dict[str, Any] = {}
_original_log: Any = None
def _rank_zero() -> bool:
return not dist.is_initialized() or dist.get_rank() == 0
def _log(self: Any, metrics: dict[str, Any], step: int) -> None:
if _original_log is None:
raise RuntimeError("tracker log hook was not initialized")
_original_log(self, metrics, step)
record = _metrics_by_step.setdefault(int(step), {})
for key, value in metrics.items():
if key in {"total_loss", "finetune_loss", "step_time_sec"} or key.startswith("grad_norm/"):
record[key] = float(value)
if "step_time_sec" in metrics and _rank_zero():
print(
"BF16_STEP " + json.dumps({
"step": int(step),
"step_time_sec": float(metrics["step_time_sec"]),
}),
flush=True,
)
def _wall_intervals(starts: list[float], synced_end: float) -> list[float]:
if not starts:
raise RuntimeError("no training-step wall starts were recorded")
return [next_start - start for start, next_start in zip(starts, starts[1:])] + [synced_end - starts[-1]]
def _tensor_digest(tensor: torch.Tensor) -> str:
raw = tensor.detach().contiguous().view(torch.uint8).cpu().numpy().tobytes()
return hashlib.sha256(raw).hexdigest()
def _capture_ada_grad_probe(method: Any, iteration: int) -> dict[str, Any]:
parameters: dict[str, Any] = {}
for name, parameter in method.student.transformer.named_parameters():
if "adaln_single" not in name:
continue
gradient = parameter.grad
if gradient is None:
parameters[name] = {"present": False}
continue
if isinstance(gradient, torch.distributed.tensor.DTensor):
gradient = gradient.to_local()
local = gradient.detach().contiguous().cpu()
local_float = local.float()
parameters[name] = {
"present": True,
"shape": list(local.shape),
"dtype": str(local.dtype),
"numel": local.numel(),
"finite": bool(torch.isfinite(local).all()),
"l2_norm": float(torch.linalg.vector_norm(local_float)),
"max_abs": float(local_float.abs().max()),
"mean": float(local_float.mean()),
"sha256": _tensor_digest(local),
}
if not parameters:
raise RuntimeError("found no adaln_single parameters for the gradient probe")
if not all(item.get("finite", False) for item in parameters.values() if item.get("present")):
raise RuntimeError("non-finite Ada gradient in probe")
return {"step": int(iteration), "parameters": parameters}
def _capture_optimizer_probe(method: Any) -> dict[str, Any]:
optimizers = list(method.get_optimizers(0))
if len(optimizers) != 1:
raise RuntimeError(f"expected one optimizer, got {len(optimizers)}")
optimizer = optimizers[0]
use_te_master = os.environ.get("FASTVIDEO_TE_FP32_MASTER") == "1"
optimizer_class = f"{type(optimizer).__module__}.{type(optimizer).__name__}"
if use_te_master:
if type(optimizer).__name__ != "FusedAdam" or not type(optimizer).__module__.startswith(
"transformer_engine."
):
raise RuntimeError(f"TE master mode constructed the wrong optimizer: {optimizer_class}")
elif not isinstance(optimizer, torch.optim.AdamW):
raise RuntimeError(f"stock control constructed the wrong optimizer: {optimizer_class}")
trainable_parameters = [
parameter for parameter in method.student.transformer.parameters() if parameter.requires_grad
]
optimizer_parameters = [
parameter for group in optimizer.param_groups for parameter in group["params"]
]
trainable_parameter_ids = [id(parameter) for parameter in trainable_parameters]
optimizer_parameter_ids = [id(parameter) for parameter in optimizer_parameters]
if len(optimizer_parameter_ids) != len(set(optimizer_parameter_ids)):
raise RuntimeError("optimizer contains duplicate parameter objects")
if set(trainable_parameter_ids) != set(optimizer_parameter_ids):
raise RuntimeError(
"optimizer parameter coverage differs from trainable transformer parameters: "
f"trainable={len(trainable_parameter_ids)} optimizer={len(optimizer_parameter_ids)}"
)
param_dtypes: set[str] = set()
state_dtypes: dict[str, set[str]] = {}
missing_state: dict[str, int] = {}
writeback_mismatches = torch.zeros((), dtype=torch.int64, device="cuda")
parameter_count = 0
parameter_numel = 0
for parameter in trainable_parameters:
if not isinstance(parameter, torch.distributed.tensor.DTensor):
raise RuntimeError(f"registered parameter is not a DTensor shard: {type(parameter).__name__}")
parameter_count += 1
parameter_numel += parameter.numel()
param_dtypes.add(str(parameter.dtype))
state = optimizer.state.get(parameter)
if state is None:
raise RuntimeError("trainable optimizer parameter has no state")
for name in ("master_param", "exp_avg", "exp_avg_sq"):
state_tensor = state.get(name)
if state_tensor is None:
missing_state[name] = missing_state.get(name, 0) + 1
continue
if not isinstance(state_tensor, torch.distributed.tensor.DTensor):
raise RuntimeError(f"optimizer state {name} is not a DTensor shard: {type(state_tensor).__name__}")
if (
state_tensor.shape != parameter.shape
or state_tensor.stride() != parameter.stride()
or state_tensor.placements != parameter.placements
or state_tensor.device_mesh.device_type != parameter.device_mesh.device_type
or not torch.equal(state_tensor.device_mesh.mesh, parameter.device_mesh.mesh)
or state_tensor.to_local().shape != parameter.to_local().shape
or state_tensor.to_local().stride() != parameter.to_local().stride()
):
raise RuntimeError(
f"optimizer state {name} layout does not match its registered parameter"
)
state_dtypes.setdefault(name, set()).add(str(state_tensor.dtype))
if use_te_master:
master = state.get("master_param")
if master is None:
raise RuntimeError("TE optimizer parameter is missing master_param")
local_parameter = parameter.to_local().detach()
local_master = master.to_local().detach()
if not torch.equal(local_parameter, local_master.to(local_parameter.dtype)):
writeback_mismatches += 1
expected_param_dtype = {"torch.bfloat16"} if use_te_master else {"torch.float32"}
if param_dtypes != expected_param_dtype:
raise RuntimeError(f"unexpected registered parameter dtypes: {param_dtypes}")
if state_dtypes.get("exp_avg") != {"torch.float32"} or state_dtypes.get("exp_avg_sq") != {"torch.float32"}:
raise RuntimeError(f"optimizer moments are not FP32: {state_dtypes}")
if missing_state.get("exp_avg", 0) or missing_state.get("exp_avg_sq", 0):
raise RuntimeError(f"missing optimizer moment states: {missing_state}")
if use_te_master:
if state_dtypes.get("master_param") != {"torch.float32"}:
raise RuntimeError(f"optimizer masters are not FP32: {state_dtypes}")
if missing_state.get("master_param", 0):
raise RuntimeError(f"missing TE master states: {missing_state}")
elif state_dtypes.get("master_param") or missing_state.get("master_param", 0) != parameter_count:
raise RuntimeError(f"stock optimizer unexpectedly contains master states: {state_dtypes}, {missing_state}")
if dist.is_initialized():
dist.all_reduce(writeback_mismatches, op=dist.ReduceOp.SUM)
total_writeback_mismatches = int(writeback_mismatches.item())
if total_writeback_mismatches:
raise RuntimeError(
"registered BF16 parameters differ from rounded FP32 masters: "
f"mismatches_across_ranks={total_writeback_mismatches}"
)
return {
"class": optimizer_class,
"parameter_count": parameter_count,
"parameter_numel": parameter_numel,
"parameter_dtypes": sorted(param_dtypes),
"state_dtypes": {name: sorted(dtypes) for name, dtypes in sorted(state_dtypes.items())},
"missing_state": missing_state,
"master_writeback_mismatches_across_ranks": total_writeback_mismatches,
"te_fp32_master": use_te_master,
}
def _series_digest(metrics_by_step: dict[int, dict[str, float]]) -> str:
series = [(step, sorted(metrics.items())) for step, metrics in sorted(metrics_by_step.items())]
return hashlib.sha256(json.dumps(series, separators=(",", ":")).encode()).hexdigest()
def _self_test() -> None:
assert _wall_intervals([1.0, 2.0, 4.0], 7.0) == [1.0, 2.0, 3.0]
assert len(hashlib.sha256(b"test").hexdigest()) == 64
print("SELF_TEST_OK")
def main() -> None:
global _ada_grad_probe, _optimizer_probe, _original_log, _peak_memory, _synced_end
parser = argparse.ArgumentParser()
parser.add_argument("--config")
parser.add_argument("--singleton-timestep", action="store_true")
parser.add_argument("--self-test", action="store_true")
args, overrides = parser.parse_known_args()
if args.self_test:
_self_test()
return
if not args.config:
parser.error("--config is required unless --self-test is used")
from fastvideo.training.trackers import DummyTracker
_original_log = DummyTracker.log
DummyTracker.log = _log
import fastvideo.train.trainer as trainer_module
from fastvideo.distributed import get_world_group
from fastvideo.train.callbacks.grad_clip import GradNormClipCallback
from fastvideo.train.entrypoint.train import main as train_main
from fastvideo.train.methods.base import TrainingMethod
from fastvideo.train.models.ltx2 import LTX2Model
original_build_kwargs = LTX2Model._build_distill_input_kwargs
original_method_optimizer_step = TrainingMethod.optimizers_schedulers_step
original_trainer_init = trainer_module.Trainer.__init__
original_trainer_iter = trainer_module.Trainer._iter_dataloader
original_trainer_run = trainer_module.Trainer.run
def _build_distill_input_kwargs(self: LTX2Model, *build_args: Any, **build_kwargs: Any) -> dict[str, Any]:
original_token_count = getattr(self, "_token_count", None)
if original_token_count is None:
raise RuntimeError("LTX-2 token count was not initialized before building transformer inputs")
_semantic_counts[f"semantic_token_count_{int(original_token_count)}"] += 1
if not args.singleton_timestep:
result = original_build_kwargs(self, *build_args, **build_kwargs)
else:
if int(self.training_config.distributed.sp_size or 1) != 1:
raise RuntimeError("singleton-timestep scratch path is restricted to SP=1")
# The source method uses _token_count only to expand the uniform
# per-sample sigma. Temporarily setting it to one reproduces the
# proposed source path without editing the checkout.
self._token_count = 1
try:
result = original_build_kwargs(self, *build_args, **build_kwargs)
finally:
self._token_count = original_token_count
timestep = result.get("timestep")
if not isinstance(timestep, torch.Tensor) or timestep.ndim != 2:
raise RuntimeError(f"unexpected LTX-2 timestep: {type(timestep).__name__}")
_semantic_counts[f"model_timestep_tokens_{int(timestep.shape[1])}"] += 1
return result
LTX2Model._build_distill_input_kwargs = _build_distill_input_kwargs
def _optimizer_step(self: TrainingMethod, iteration: int) -> None:
global _ada_grad_probe
if int(iteration) == GRAD_PROBE_STEP:
_ada_grad_probe = _capture_ada_grad_probe(self, iteration)
_semantic_counts["ada_grad_probes"] += 1
original_method_optimizer_step(self, iteration)
TrainingMethod.optimizers_schedulers_step = _optimizer_step
trainer_module.build_tracker = lambda *_args, **_kwargs: DummyTracker()
class _Recorder:
def on_training_step_end(self, _method: Any, metrics: dict[str, Any], iteration: int = 0) -> None:
del iteration
_step_times.append(float(metrics["step_time_sec"]))
def _trainer_init(self: Any, *init_args: Any, **init_kwargs: Any) -> None:
original_trainer_init(self, *init_args, **init_kwargs)
if int(self.training_config.loop.gradient_accumulation_steps or 1) != 1:
raise RuntimeError("scratch singleton benchmark requires gradient accumulation 1")
if int(self.training_config.distributed.sp_size or 1) != 1:
raise RuntimeError("scratch singleton benchmark requires SP=1")
grad_callbacks = [
callback for callback in self.callbacks._callbacks.values()
if isinstance(callback, GradNormClipCallback)
]
if len(grad_callbacks) != 1 or not grad_callbacks[0]._log_grad_norms:
raise RuntimeError("scratch singleton benchmark requires one logging GradNormClipCallback")
self.callbacks._callbacks.pop("validation", None)
self.callbacks._callbacks["_benchmark_recorder"] = _Recorder()
def _timed_iter(self: Any, dataloader: Any) -> Any:
iterator = original_trainer_iter(self, dataloader)
while True:
_step_starts.append(time.perf_counter())
yield next(iterator)
def _trainer_run(self: Any, method: TrainingMethod, **kwargs: Any) -> Any:
global _optimizer_probe, _peak_memory, _synced_end
vae = getattr(getattr(method, "student", None), "vae", None)
if vae is not None:
method.student.vae = None
del vae
gc.collect()
torch.cuda.empty_cache()
if _rank_zero():
print("BF16_SETUP " + json.dumps({"unused_vae": "unloaded"}), flush=True)
torch.cuda.reset_peak_memory_stats()
result = original_trainer_run(self, method, **kwargs)
torch.cuda.synchronize()
_synced_end = time.perf_counter()
_peak_memory = {
"allocated_bytes": int(torch.cuda.max_memory_allocated()),
"reserved_bytes": int(torch.cuda.max_memory_reserved()),
}
_optimizer_probe = _capture_optimizer_probe(method)
return result
trainer_module.Trainer.__init__ = _trainer_init
trainer_module.Trainer._iter_dataloader = _timed_iter
trainer_module.Trainer.run = _trainer_run
train_main(argparse.Namespace(config=args.config, dry_run=False), overrides=overrides or None)
if _synced_end is None:
raise RuntimeError("final synchronized wall endpoint was not recorded")
wall_intervals = _wall_intervals(_step_starts, _synced_end)
if len(_step_times) != EXPECTED_STEPS or len(wall_intervals) != EXPECTED_STEPS:
raise RuntimeError(
f"expected {EXPECTED_STEPS} steps, got internal={len(_step_times)} wall={len(wall_intervals)}"
)
if _ada_grad_probe is None or _semantic_counts["ada_grad_probes"] != 1:
raise RuntimeError("Ada gradient probe did not run exactly once")
payload = {
"internal_step_times": _step_times,
"wall_intervals": wall_intervals,
"end_to_end_wall_sec": _synced_end - _step_starts[0],
"measured_window_wall_sec": _synced_end - _step_starts[WARMUP_STEPS],
"metrics_by_step": _metrics_by_step,
"metric_series_digest": _series_digest(_metrics_by_step),
"semantics": dict(_semantic_counts),
"ada_grad_probe": _ada_grad_probe,
"optimizer_probe": _optimizer_probe,
"peak_memory": _peak_memory,
}
world = get_world_group()
payloads: list[dict[str, Any] | None] = [None] * world.world_size
dist.all_gather_object(payloads, payload, group=world.cpu_group)
if world.rank != 0:
return
rank_payloads = [rank_payload for rank_payload in payloads if rank_payload is not None]
if len(rank_payloads) != world.world_size:
raise RuntimeError("missing rank payload")
internal_slowest = [
max(rank_payload["internal_step_times"][index] for rank_payload in rank_payloads)
for index in range(EXPECTED_STEPS)
]
wall_slowest = [
max(rank_payload["wall_intervals"][index] for rank_payload in rank_payloads)
for index in range(EXPECTED_STEPS)
]
internal_measured = internal_slowest[WARMUP_STEPS:]
wall_measured = wall_slowest[WARMUP_STEPS:]
internal_median = statistics.median(internal_measured)
wall_median = statistics.median(wall_measured)
expected_model_tokens = 1 if args.singleton_timestep else None
semantics_by_rank = []
for rank, rank_payload in enumerate(rank_payloads):
semantics = rank_payload["semantics"]
semantic_token_counts = {
int(key.rsplit("_", 1)[1]): count
for key, count in semantics.items()
if key.startswith("semantic_token_count_")
}
model_token_counts = {
int(key.rsplit("_", 1)[1]): count
for key, count in semantics.items()
if key.startswith("model_timestep_tokens_")
}
if sum(semantic_token_counts.values()) != EXPECTED_STEPS or len(semantic_token_counts) != 1:
raise RuntimeError(f"rank {rank} recorded invalid semantic token counts: {semantic_token_counts}")
if sum(model_token_counts.values()) != EXPECTED_STEPS:
raise RuntimeError(f"rank {rank} recorded invalid timestep counts: {model_token_counts}")
if expected_model_tokens is not None and model_token_counts != {expected_model_tokens: EXPECTED_STEPS}:
raise RuntimeError(f"rank {rank} did not use singleton timesteps: {model_token_counts}")
if expected_model_tokens is None and model_token_counts != semantic_token_counts:
raise RuntimeError(
f"rank {rank} control changed timestep shape: semantic={semantic_token_counts}, model={model_token_counts}"
)
semantics_by_rank.append({
"rank": rank,
"counts": semantics,
"metric_series_digest": rank_payload["metric_series_digest"],
"peak_memory": rank_payload["peak_memory"],
"ada_grad_probe": rank_payload["ada_grad_probe"],
"optimizer_probe": rank_payload["optimizer_probe"],
"first_metrics": rank_payload["metrics_by_step"].get(1, {}),
"last_metrics": rank_payload["metrics_by_step"].get(EXPECTED_STEPS, {}),
})
print(
"BF16_TIMING " + json.dumps({
"internal_slowest_rank_sec": internal_slowest,
"true_wall_slowest_rank_sec": wall_slowest,
"final_cuda_sync": True,
}, sort_keys=True),
flush=True,
)
print(
"BF16_SEMANTICS " + json.dumps({
"singleton_timestep": args.singleton_timestep,
"gradient_probe_step": GRAD_PROBE_STEP,
"by_rank": semantics_by_rank,
"rank0_metric_series": rank_payloads[0]["metrics_by_step"],
}, sort_keys=True),
flush=True,
)
print(
"BF16_RESULT " + json.dumps({
"singleton_timestep": args.singleton_timestep,
"local_batch_size": LOCAL_BATCH_SIZE,
"world_size": world.world_size,
"internal_median_step_sec": internal_median,
"internal_model_mfu_percent": MFU_NUMERATOR * LOCAL_BATCH_SIZE / internal_median,
"true_wall_median_step_sec": wall_median,
"true_wall_model_mfu_percent": MFU_NUMERATOR * LOCAL_BATCH_SIZE / wall_median,
"true_wall_sum_slowest_intervals_sec": sum(wall_measured),
"true_wall_mean_slowest_interval_sec": statistics.mean(wall_measured),
"true_measured_window_max_rank_sec": max(
rank_payload["measured_window_wall_sec"] for rank_payload in rank_payloads
),
"true_end_to_end_max_rank_sec": max(
rank_payload["end_to_end_wall_sec"] for rank_payload in rank_payloads
),
"samples_per_second_from_true_wall_median": (
world.world_size * LOCAL_BATCH_SIZE / wall_median
),
"peak_allocated_max_rank_bytes": max(
rank_payload["peak_memory"]["allocated_bytes"] for rank_payload in rank_payloads
),
"peak_reserved_max_rank_bytes": max(
rank_payload["peak_memory"]["reserved_bytes"] for rank_payload in rank_payloads
),
}, sort_keys=True),
flush=True,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,87 @@
#!/usr/bin/env python3
"""Benchmark FastVideo training with regional max-autotune and no CUDA graphs."""
import argparse
import json
import statistics
import torch
import torch.distributed as dist
from fastvideo.training.trackers import DummyTracker
_step_times: list[float] = []
_compile_calls = 0
_original_log = DummyTracker.log
_original_compile = torch.compile
def _log(self, metrics, step):
_original_log(self, metrics, step)
if "step_time_sec" in metrics and (not dist.is_initialized() or dist.get_rank() == 0):
print("BF16_STEP " + json.dumps({"step": step, "step_time_sec": float(metrics["step_time_sec"])}),
flush=True)
def _compile(*args, **kwargs):
global _compile_calls
if "mode" in kwargs or "options" in kwargs:
raise RuntimeError(f"unexpected compile configuration: {kwargs}")
kwargs["mode"] = "max-autotune-no-cudagraphs"
_compile_calls += 1
return _original_compile(*args, **kwargs)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True)
args, overrides = parser.parse_known_args()
torch.compile = _compile
DummyTracker.log = _log
import fastvideo.train.trainer as trainer_module
from fastvideo.distributed import get_world_group
from fastvideo.train.entrypoint.train import main as train_main
trainer_module.build_tracker = lambda *_args, **_kwargs: DummyTracker()
original_trainer_init = trainer_module.Trainer.__init__
class _Recorder:
def on_training_step_end(self, _method, metrics, iteration=0):
_step_times.append(float(metrics["step_time_sec"]))
def _trainer_init(self, *init_args, **init_kwargs):
original_trainer_init(self, *init_args, **init_kwargs)
self.callbacks._callbacks.pop("validation", None)
self.callbacks._callbacks["_benchmark_recorder"] = _Recorder()
trainer_module.Trainer.__init__ = _trainer_init
train_main(argparse.Namespace(config=args.config, dry_run=False), overrides=overrides or None)
if _compile_calls != 48:
raise RuntimeError(f"expected 48 regional compile calls, got {_compile_calls}")
world = get_world_group()
times_by_rank = [None] * world.world_size
dist.all_gather_object(times_by_rank, _step_times, group=world.cpu_group)
if world.rank == 0:
per_step_max = [max(values) for values in zip(*times_by_rank, strict=True)]
measured = per_step_max[10:30]
if len(measured) != 20:
raise RuntimeError(f"expected 30 steps, got {len(per_step_max)}")
median = statistics.median(measured)
print("BF16_COMPILE_MODE " + json.dumps({
"calls": _compile_calls,
"mode": "max-autotune-no-cudagraphs",
}, sort_keys=True), flush=True)
print("BF16_RESULT " + json.dumps({
"kind": "regional_max_autotune_no_cudagraphs",
"median_step_sec": median,
"samples_per_second": 4.0 / median,
"model_mfu_percent": 14.444115 / median,
}, sort_keys=True), flush=True)
if __name__ == "__main__":
main()
@@ -0,0 +1,488 @@
#!/usr/bin/env python3
"""Scratch A/B for embedding one uniform LTX-2 T2V timestep per sample."""
from __future__ import annotations
import argparse
from collections import Counter
import gc
import hashlib
import json
import os
import statistics
import time
from typing import Any
import torch
import torch.distributed as dist
EXPECTED_STEPS = 30
WARMUP_STEPS = 10
MFU_NUMERATOR = 14.444115
LOCAL_BATCH_SIZE = int(os.environ.get("FASTVIDEO_BENCH_LOCAL_BATCH_SIZE", "1"))
GRAD_ACCUM_STEPS = int(os.environ.get("FASTVIDEO_BENCH_GRAD_ACCUM_STEPS", "1"))
if LOCAL_BATCH_SIZE <= 0:
raise ValueError("FASTVIDEO_BENCH_LOCAL_BATCH_SIZE must be positive")
if GRAD_ACCUM_STEPS <= 0:
raise ValueError("FASTVIDEO_BENCH_GRAD_ACCUM_STEPS must be positive")
GRAD_PROBE_STEP = WARMUP_STEPS
_step_times: list[float] = []
_step_starts: list[float] = []
_synced_end: float | None = None
_metrics_by_step: dict[int, dict[str, float]] = {}
_semantic_counts: Counter[str] = Counter()
_ada_grad_probe: dict[str, Any] | None = None
_peak_memory: dict[str, int] = {}
_optimizer_probe: dict[str, Any] = {}
_original_log: Any = None
def _rank_zero() -> bool:
return not dist.is_initialized() or dist.get_rank() == 0
def _log(self: Any, metrics: dict[str, Any], step: int) -> None:
if _original_log is None:
raise RuntimeError("tracker log hook was not initialized")
_original_log(self, metrics, step)
record = _metrics_by_step.setdefault(int(step), {})
for key, value in metrics.items():
if key in {"total_loss", "finetune_loss", "step_time_sec"} or key.startswith("grad_norm/"):
record[key] = float(value)
if "step_time_sec" in metrics and _rank_zero():
print(
"BF16_STEP " + json.dumps({
"step": int(step),
"step_time_sec": float(metrics["step_time_sec"]),
}),
flush=True,
)
def _wall_intervals(starts: list[float], synced_end: float) -> list[float]:
if not starts:
raise RuntimeError("no training-step wall starts were recorded")
return [next_start - start for start, next_start in zip(starts, starts[1:])] + [synced_end - starts[-1]]
def _tensor_digest(tensor: torch.Tensor) -> str:
raw = tensor.detach().contiguous().view(torch.uint8).cpu().numpy().tobytes()
return hashlib.sha256(raw).hexdigest()
def _capture_ada_grad_probe(method: Any, iteration: int) -> dict[str, Any]:
parameters: dict[str, Any] = {}
for name, parameter in method.student.transformer.named_parameters():
if "adaln_single" not in name:
continue
gradient = parameter.grad
if gradient is None:
parameters[name] = {"present": False}
continue
if isinstance(gradient, torch.distributed.tensor.DTensor):
gradient = gradient.to_local()
local = gradient.detach().contiguous().cpu()
local_float = local.float()
parameters[name] = {
"present": True,
"shape": list(local.shape),
"dtype": str(local.dtype),
"numel": local.numel(),
"finite": bool(torch.isfinite(local).all()),
"l2_norm": float(torch.linalg.vector_norm(local_float)),
"max_abs": float(local_float.abs().max()),
"mean": float(local_float.mean()),
"sha256": _tensor_digest(local),
}
if not parameters:
raise RuntimeError("found no adaln_single parameters for the gradient probe")
if not all(item.get("finite", False) for item in parameters.values() if item.get("present")):
raise RuntimeError("non-finite Ada gradient in probe")
return {"step": int(iteration), "parameters": parameters}
def _capture_optimizer_probe(method: Any) -> dict[str, Any]:
optimizers = list(method.get_optimizers(0))
if len(optimizers) != 1:
raise RuntimeError(f"expected one optimizer, got {len(optimizers)}")
optimizer = optimizers[0]
use_te_master = os.environ.get("FASTVIDEO_TE_FP32_MASTER") == "1"
optimizer_class = f"{type(optimizer).__module__}.{type(optimizer).__name__}"
if use_te_master:
if type(optimizer).__name__ != "FusedAdam" or not type(optimizer).__module__.startswith(
"transformer_engine."
):
raise RuntimeError(f"TE master mode constructed the wrong optimizer: {optimizer_class}")
elif not isinstance(optimizer, torch.optim.AdamW):
raise RuntimeError(f"stock control constructed the wrong optimizer: {optimizer_class}")
trainable_parameters = [
parameter for parameter in method.student.transformer.parameters() if parameter.requires_grad
]
optimizer_parameters = [
parameter for group in optimizer.param_groups for parameter in group["params"]
]
trainable_parameter_ids = [id(parameter) for parameter in trainable_parameters]
optimizer_parameter_ids = [id(parameter) for parameter in optimizer_parameters]
if len(optimizer_parameter_ids) != len(set(optimizer_parameter_ids)):
raise RuntimeError("optimizer contains duplicate parameter objects")
if set(trainable_parameter_ids) != set(optimizer_parameter_ids):
raise RuntimeError(
"optimizer parameter coverage differs from trainable transformer parameters: "
f"trainable={len(trainable_parameter_ids)} optimizer={len(optimizer_parameter_ids)}"
)
param_dtypes: set[str] = set()
state_dtypes: dict[str, set[str]] = {}
missing_state: dict[str, int] = {}
writeback_mismatches = torch.zeros((), dtype=torch.int64, device="cuda")
parameter_count = 0
parameter_numel = 0
for parameter in trainable_parameters:
if not isinstance(parameter, torch.distributed.tensor.DTensor):
raise RuntimeError(f"registered parameter is not a DTensor shard: {type(parameter).__name__}")
parameter_count += 1
parameter_numel += parameter.numel()
param_dtypes.add(str(parameter.dtype))
state = optimizer.state.get(parameter)
if state is None:
raise RuntimeError("trainable optimizer parameter has no state")
for name in ("master_param", "exp_avg", "exp_avg_sq"):
state_tensor = state.get(name)
if state_tensor is None:
missing_state[name] = missing_state.get(name, 0) + 1
continue
if not isinstance(state_tensor, torch.distributed.tensor.DTensor):
raise RuntimeError(f"optimizer state {name} is not a DTensor shard: {type(state_tensor).__name__}")
if (
state_tensor.shape != parameter.shape
or state_tensor.stride() != parameter.stride()
or state_tensor.placements != parameter.placements
or state_tensor.device_mesh.device_type != parameter.device_mesh.device_type
or not torch.equal(state_tensor.device_mesh.mesh, parameter.device_mesh.mesh)
or state_tensor.to_local().shape != parameter.to_local().shape
or state_tensor.to_local().stride() != parameter.to_local().stride()
):
raise RuntimeError(
f"optimizer state {name} layout does not match its registered parameter"
)
state_dtypes.setdefault(name, set()).add(str(state_tensor.dtype))
if use_te_master:
master = state.get("master_param")
if master is None:
raise RuntimeError("TE optimizer parameter is missing master_param")
local_parameter = parameter.to_local().detach()
local_master = master.to_local().detach()
if not torch.equal(local_parameter, local_master.to(local_parameter.dtype)):
writeback_mismatches += 1
expected_param_dtype = {"torch.bfloat16"} if use_te_master else {"torch.float32"}
if param_dtypes != expected_param_dtype:
raise RuntimeError(f"unexpected registered parameter dtypes: {param_dtypes}")
if state_dtypes.get("exp_avg") != {"torch.float32"} or state_dtypes.get("exp_avg_sq") != {"torch.float32"}:
raise RuntimeError(f"optimizer moments are not FP32: {state_dtypes}")
if missing_state.get("exp_avg", 0) or missing_state.get("exp_avg_sq", 0):
raise RuntimeError(f"missing optimizer moment states: {missing_state}")
if use_te_master:
if state_dtypes.get("master_param") != {"torch.float32"}:
raise RuntimeError(f"optimizer masters are not FP32: {state_dtypes}")
if missing_state.get("master_param", 0):
raise RuntimeError(f"missing TE master states: {missing_state}")
elif state_dtypes.get("master_param") or missing_state.get("master_param", 0) != parameter_count:
raise RuntimeError(f"stock optimizer unexpectedly contains master states: {state_dtypes}, {missing_state}")
if dist.is_initialized():
dist.all_reduce(writeback_mismatches, op=dist.ReduceOp.SUM)
total_writeback_mismatches = int(writeback_mismatches.item())
if total_writeback_mismatches:
raise RuntimeError(
"registered BF16 parameters differ from rounded FP32 masters: "
f"mismatches_across_ranks={total_writeback_mismatches}"
)
return {
"class": optimizer_class,
"parameter_count": parameter_count,
"parameter_numel": parameter_numel,
"parameter_dtypes": sorted(param_dtypes),
"state_dtypes": {name: sorted(dtypes) for name, dtypes in sorted(state_dtypes.items())},
"missing_state": missing_state,
"master_writeback_mismatches_across_ranks": total_writeback_mismatches,
"te_fp32_master": use_te_master,
}
def _series_digest(metrics_by_step: dict[int, dict[str, float]]) -> str:
series = [(step, sorted(metrics.items())) for step, metrics in sorted(metrics_by_step.items())]
return hashlib.sha256(json.dumps(series, separators=(",", ":")).encode()).hexdigest()
def _self_test() -> None:
assert _wall_intervals([1.0, 2.0, 4.0], 7.0) == [1.0, 2.0, 3.0]
assert len(hashlib.sha256(b"test").hexdigest()) == 64
print("SELF_TEST_OK")
def main() -> None:
global _ada_grad_probe, _optimizer_probe, _original_log, _peak_memory, _synced_end
parser = argparse.ArgumentParser()
parser.add_argument("--config")
parser.add_argument("--singleton-timestep", action="store_true")
parser.add_argument("--self-test", action="store_true")
args, overrides = parser.parse_known_args()
if args.self_test:
_self_test()
return
if not args.config:
parser.error("--config is required unless --self-test is used")
from fastvideo.training.trackers import DummyTracker
_original_log = DummyTracker.log
DummyTracker.log = _log
import fastvideo.train.trainer as trainer_module
from fastvideo.distributed import get_world_group
from fastvideo.train.callbacks.grad_clip import GradNormClipCallback
from fastvideo.train.entrypoint.train import main as train_main
from fastvideo.train.methods.base import TrainingMethod
from fastvideo.train.models.ltx2 import LTX2Model
original_build_kwargs = LTX2Model._build_distill_input_kwargs
original_method_optimizer_step = TrainingMethod.optimizers_schedulers_step
original_trainer_init = trainer_module.Trainer.__init__
original_trainer_iter = trainer_module.Trainer._iter_dataloader
original_trainer_run = trainer_module.Trainer.run
def _build_distill_input_kwargs(self: LTX2Model, *build_args: Any, **build_kwargs: Any) -> dict[str, Any]:
original_token_count = getattr(self, "_token_count", None)
if original_token_count is None:
raise RuntimeError("LTX-2 token count was not initialized before building transformer inputs")
_semantic_counts[f"semantic_token_count_{int(original_token_count)}"] += 1
if not args.singleton_timestep:
result = original_build_kwargs(self, *build_args, **build_kwargs)
else:
if int(self.training_config.distributed.sp_size or 1) != 1:
raise RuntimeError("singleton-timestep scratch path is restricted to SP=1")
# The source method uses _token_count only to expand the uniform
# per-sample sigma. Temporarily setting it to one reproduces the
# proposed source path without editing the checkout.
self._token_count = 1
try:
result = original_build_kwargs(self, *build_args, **build_kwargs)
finally:
self._token_count = original_token_count
timestep = result.get("timestep")
if not isinstance(timestep, torch.Tensor) or timestep.ndim != 2:
raise RuntimeError(f"unexpected LTX-2 timestep: {type(timestep).__name__}")
_semantic_counts[f"model_timestep_tokens_{int(timestep.shape[1])}"] += 1
return result
LTX2Model._build_distill_input_kwargs = _build_distill_input_kwargs
def _optimizer_step(self: TrainingMethod, iteration: int) -> None:
global _ada_grad_probe
if int(iteration) == GRAD_PROBE_STEP:
_ada_grad_probe = _capture_ada_grad_probe(self, iteration)
_semantic_counts["ada_grad_probes"] += 1
original_method_optimizer_step(self, iteration)
TrainingMethod.optimizers_schedulers_step = _optimizer_step
trainer_module.build_tracker = lambda *_args, **_kwargs: DummyTracker()
class _Recorder:
def on_training_step_end(self, _method: Any, metrics: dict[str, Any], iteration: int = 0) -> None:
del iteration
_step_times.append(float(metrics["step_time_sec"]))
def _trainer_init(self: Any, *init_args: Any, **init_kwargs: Any) -> None:
original_trainer_init(self, *init_args, **init_kwargs)
actual_grad_accum = int(self.training_config.loop.gradient_accumulation_steps or 1)
if actual_grad_accum != GRAD_ACCUM_STEPS:
raise RuntimeError(
"scratch singleton benchmark gradient accumulation mismatch: "
f"config={actual_grad_accum} expected={GRAD_ACCUM_STEPS}"
)
if int(self.training_config.distributed.sp_size or 1) != 1:
raise RuntimeError("scratch singleton benchmark requires SP=1")
grad_callbacks = [
callback for callback in self.callbacks._callbacks.values()
if isinstance(callback, GradNormClipCallback)
]
if len(grad_callbacks) != 1 or not grad_callbacks[0]._log_grad_norms:
raise RuntimeError("scratch singleton benchmark requires one logging GradNormClipCallback")
self.callbacks._callbacks.pop("validation", None)
self.callbacks._callbacks["_benchmark_recorder"] = _Recorder()
def _timed_iter(self: Any, dataloader: Any) -> Any:
iterator = original_trainer_iter(self, dataloader)
microstep = 0
while True:
if microstep % GRAD_ACCUM_STEPS == 0:
_step_starts.append(time.perf_counter())
microstep += 1
yield next(iterator)
def _trainer_run(self: Any, method: TrainingMethod, **kwargs: Any) -> Any:
global _optimizer_probe, _peak_memory, _synced_end
vae = getattr(getattr(method, "student", None), "vae", None)
if vae is not None:
method.student.vae = None
del vae
gc.collect()
torch.cuda.empty_cache()
if _rank_zero():
print("BF16_SETUP " + json.dumps({"unused_vae": "unloaded"}), flush=True)
torch.cuda.reset_peak_memory_stats()
result = original_trainer_run(self, method, **kwargs)
torch.cuda.synchronize()
_synced_end = time.perf_counter()
_peak_memory = {
"allocated_bytes": int(torch.cuda.max_memory_allocated()),
"reserved_bytes": int(torch.cuda.max_memory_reserved()),
}
_optimizer_probe = _capture_optimizer_probe(method)
return result
trainer_module.Trainer.__init__ = _trainer_init
trainer_module.Trainer._iter_dataloader = _timed_iter
trainer_module.Trainer.run = _trainer_run
train_main(argparse.Namespace(config=args.config, dry_run=False), overrides=overrides or None)
if _synced_end is None:
raise RuntimeError("final synchronized wall endpoint was not recorded")
wall_intervals = _wall_intervals(_step_starts, _synced_end)
if len(_step_times) != EXPECTED_STEPS or len(wall_intervals) != EXPECTED_STEPS:
raise RuntimeError(
f"expected {EXPECTED_STEPS} steps, got internal={len(_step_times)} wall={len(wall_intervals)}"
)
if _ada_grad_probe is None or _semantic_counts["ada_grad_probes"] != 1:
raise RuntimeError("Ada gradient probe did not run exactly once")
payload = {
"internal_step_times": _step_times,
"wall_intervals": wall_intervals,
"end_to_end_wall_sec": _synced_end - _step_starts[0],
"measured_window_wall_sec": _synced_end - _step_starts[WARMUP_STEPS],
"metrics_by_step": _metrics_by_step,
"metric_series_digest": _series_digest(_metrics_by_step),
"semantics": dict(_semantic_counts),
"ada_grad_probe": _ada_grad_probe,
"optimizer_probe": _optimizer_probe,
"peak_memory": _peak_memory,
}
world = get_world_group()
payloads: list[dict[str, Any] | None] = [None] * world.world_size
dist.all_gather_object(payloads, payload, group=world.cpu_group)
if world.rank != 0:
return
rank_payloads = [rank_payload for rank_payload in payloads if rank_payload is not None]
if len(rank_payloads) != world.world_size:
raise RuntimeError("missing rank payload")
internal_slowest = [
max(rank_payload["internal_step_times"][index] for rank_payload in rank_payloads)
for index in range(EXPECTED_STEPS)
]
wall_slowest = [
max(rank_payload["wall_intervals"][index] for rank_payload in rank_payloads)
for index in range(EXPECTED_STEPS)
]
internal_measured = internal_slowest[WARMUP_STEPS:]
wall_measured = wall_slowest[WARMUP_STEPS:]
internal_median = statistics.median(internal_measured)
wall_median = statistics.median(wall_measured)
expected_model_tokens = 1 if args.singleton_timestep else None
semantics_by_rank = []
for rank, rank_payload in enumerate(rank_payloads):
semantics = rank_payload["semantics"]
semantic_token_counts = {
int(key.rsplit("_", 1)[1]): count
for key, count in semantics.items()
if key.startswith("semantic_token_count_")
}
model_token_counts = {
int(key.rsplit("_", 1)[1]): count
for key, count in semantics.items()
if key.startswith("model_timestep_tokens_")
}
expected_microsteps = EXPECTED_STEPS * GRAD_ACCUM_STEPS
if sum(semantic_token_counts.values()) != expected_microsteps or len(semantic_token_counts) != 1:
raise RuntimeError(f"rank {rank} recorded invalid semantic token counts: {semantic_token_counts}")
if sum(model_token_counts.values()) != expected_microsteps:
raise RuntimeError(f"rank {rank} recorded invalid timestep counts: {model_token_counts}")
if expected_model_tokens is not None and model_token_counts != {expected_model_tokens: expected_microsteps}:
raise RuntimeError(f"rank {rank} did not use singleton timesteps: {model_token_counts}")
if expected_model_tokens is None and model_token_counts != semantic_token_counts:
raise RuntimeError(
f"rank {rank} control changed timestep shape: semantic={semantic_token_counts}, model={model_token_counts}"
)
semantics_by_rank.append({
"rank": rank,
"counts": semantics,
"metric_series_digest": rank_payload["metric_series_digest"],
"peak_memory": rank_payload["peak_memory"],
"ada_grad_probe": rank_payload["ada_grad_probe"],
"optimizer_probe": rank_payload["optimizer_probe"],
"first_metrics": rank_payload["metrics_by_step"].get(1, {}),
"last_metrics": rank_payload["metrics_by_step"].get(EXPECTED_STEPS, {}),
})
print(
"BF16_TIMING " + json.dumps({
"internal_slowest_rank_sec": internal_slowest,
"true_wall_slowest_rank_sec": wall_slowest,
"final_cuda_sync": True,
}, sort_keys=True),
flush=True,
)
print(
"BF16_SEMANTICS " + json.dumps({
"singleton_timestep": args.singleton_timestep,
"gradient_probe_step": GRAD_PROBE_STEP,
"by_rank": semantics_by_rank,
"rank0_metric_series": rank_payloads[0]["metrics_by_step"],
}, sort_keys=True),
flush=True,
)
print(
"BF16_RESULT " + json.dumps({
"singleton_timestep": args.singleton_timestep,
"local_batch_size": LOCAL_BATCH_SIZE,
"gradient_accumulation_steps": GRAD_ACCUM_STEPS,
"world_size": world.world_size,
"internal_median_step_sec": internal_median,
"internal_model_mfu_percent": (
MFU_NUMERATOR * LOCAL_BATCH_SIZE * GRAD_ACCUM_STEPS / internal_median
),
"true_wall_median_step_sec": wall_median,
"true_wall_model_mfu_percent": (
MFU_NUMERATOR * LOCAL_BATCH_SIZE * GRAD_ACCUM_STEPS / wall_median
),
"true_wall_sum_slowest_intervals_sec": sum(wall_measured),
"true_wall_mean_slowest_interval_sec": statistics.mean(wall_measured),
"true_measured_window_max_rank_sec": max(
rank_payload["measured_window_wall_sec"] for rank_payload in rank_payloads
),
"true_end_to_end_max_rank_sec": max(
rank_payload["end_to_end_wall_sec"] for rank_payload in rank_payloads
),
"samples_per_second_from_true_wall_median": (
world.world_size * LOCAL_BATCH_SIZE * GRAD_ACCUM_STEPS / wall_median
),
"peak_allocated_max_rank_bytes": max(
rank_payload["peak_memory"]["allocated_bytes"] for rank_payload in rank_payloads
),
"peak_reserved_max_rank_bytes": max(
rank_payload["peak_memory"]["reserved_bytes"] for rank_payload in rank_payloads
),
}, sort_keys=True),
flush=True,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,134 @@
#!/usr/bin/env python3
"""Scratch candidate: run base-stage LTX RMSNorm outside CUDA autocast."""
from __future__ import annotations
import json
import os
import torch
from benchmark_fastvideo_train_pack_d016 import main
def _install_candidate() -> None:
import fastvideo.models.dits.ltx2 as ltx2
original_dispatch = ltx2._rms_norm_dispatch
original_stage_aware_forward = ltx2.StageAwareRMSNorm.forward
def rms_norm_no_autocast(
x: torch.Tensor,
eps: float,
weight: torch.Tensor | None = None,
) -> torch.Tensor:
with torch.autocast(device_type="cuda", enabled=False):
return torch.nn.functional.rms_norm(
x,
(x.shape[-1], ),
weight=weight,
eps=eps,
).to(x.dtype)
def dispatch(
x: torch.Tensor,
eps: float,
weight: torch.Tensor | None = None,
) -> torch.Tensor:
if ltx2._is_ltx2_refine_stage():
return original_dispatch(x, eps=eps, weight=weight)
return rms_norm_no_autocast(x, eps=eps, weight=weight)
def stage_aware_forward(
self: ltx2.StageAwareRMSNorm,
x: torch.Tensor,
) -> torch.Tensor:
if ltx2._is_ltx2_refine_stage():
return original_stage_aware_forward(self, x)
return rms_norm_no_autocast(x, eps=self.eps, weight=self.weight)
if not torch.cuda.is_available():
raise RuntimeError("RMSNorm candidate requires CUDA")
device = torch.device("cuda", int(os.environ.get("LOCAL_RANK", "0")))
generator = torch.Generator(device=device).manual_seed(20260721)
x = torch.randn(
(2, 3, 32),
device=device,
dtype=torch.bfloat16,
generator=generator,
)
weight = torch.linspace(
0.5,
1.5,
x.shape[-1],
device=x.device,
dtype=x.dtype,
)
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
old_unweighted = original_dispatch(x, eps=1e-6)
old_norm = ltx2.StageAwareRMSNorm(x.shape[-1], eps=1e-6).to(
device=x.device,
dtype=x.dtype,
)
with torch.no_grad():
old_norm.weight.copy_(weight)
old_weighted = original_stage_aware_forward(old_norm, x)
ltx2._rms_norm_dispatch = dispatch
ltx2.StageAwareRMSNorm.forward = stage_aware_forward
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
new_unweighted = ltx2._rms_norm_dispatch(x, eps=1e-6)
new_norm = ltx2.StageAwareRMSNorm(x.shape[-1], eps=1e-6).to(
device=x.device,
dtype=x.dtype,
)
with torch.no_grad():
new_norm.weight.copy_(weight)
new_weighted = new_norm(x)
with torch.autocast(device_type="cuda", enabled=False):
native_unweighted = torch.nn.functional.rms_norm(
x,
(x.shape[-1], ),
eps=1e-6,
).to(x.dtype)
native_weighted = torch.nn.functional.rms_norm(
x,
(x.shape[-1], ),
weight=weight,
eps=1e-6,
).to(x.dtype)
if new_unweighted.dtype != x.dtype or new_weighted.dtype != x.dtype:
raise RuntimeError(
f"RMSNorm candidate changed dtype: input={x.dtype} "
f"unweighted={new_unweighted.dtype} weighted={new_weighted.dtype}"
)
torch.testing.assert_close(new_unweighted, native_unweighted, rtol=0, atol=0)
torch.testing.assert_close(new_weighted, native_weighted, rtol=0, atol=0)
torch.testing.assert_close(new_unweighted, old_unweighted, rtol=0.02, atol=0.02)
torch.testing.assert_close(new_weighted, old_weighted, rtol=0.02, atol=0.02)
if torch.equal(new_unweighted, new_weighted):
raise RuntimeError("weighted StageAwareRMSNorm ignored its learned weight")
print(
"RMSNORM_SEMANTICS "
+ json.dumps({
"candidate_installed": True,
"input_dtype": str(x.dtype),
"output_dtype": str(new_weighted.dtype),
"native_unweighted_exact": torch.equal(new_unweighted, native_unweighted),
"native_weighted_exact": torch.equal(new_weighted, native_weighted),
"old_unweighted_max_abs_diff": float((new_unweighted - old_unweighted).abs().max()),
"old_weighted_max_abs_diff": float((new_weighted - old_weighted).abs().max()),
"stage_aware_weight_effective": not torch.equal(new_unweighted, new_weighted),
"refine_path": "delegates_to_original",
}, sort_keys=True),
flush=True,
)
if __name__ == "__main__":
_install_candidate()
main()
@@ -0,0 +1,92 @@
#!/usr/bin/env python3
"""Benchmark FastVideo training with FSDP2 symmetric-memory all-gather."""
import argparse
import json
import os
import statistics
import torch.distributed as dist
from fastvideo.training.trackers import DummyTracker
_step_times: list[float] = []
_original_log = DummyTracker.log
def _log(self, metrics, step):
_original_log(self, metrics, step)
if "step_time_sec" in metrics and (not dist.is_initialized() or dist.get_rank() == 0):
print("BF16_STEP " + json.dumps({"step": step, "step_time_sec": float(metrics["step_time_sec"])}),
flush=True)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True)
args, overrides = parser.parse_known_args()
if os.environ.get("NCCL_CTA_POLICY") != "2":
raise RuntimeError("symmetric-memory all-gather requires NCCL_CTA_POLICY=2")
import fastvideo.train.models.ltx2.ltx2 as ltx2_module
import fastvideo.train.trainer as trainer_module
from fastvideo.distributed import get_world_group
from fastvideo.train.entrypoint.train import main as train_main
from torch.distributed.fsdp import FSDPModule
original_load = ltx2_module.load_module_from_path
def load_with_symm_mem(**kwargs):
module = original_load(**kwargs)
if kwargs.get("module_type") == "transformer":
if not isinstance(module, FSDPModule):
raise RuntimeError("expected an FSDP2 transformer")
configured = []
for name, submodule in module.named_modules():
if isinstance(submodule, FSDPModule):
submodule.set_force_sum_reduction_for_comms(True)
submodule.set_symm_mem_for_comm("NCCL")
configured.append(name or "<root>")
if len(configured) != 49 or configured[0] != "<root>":
raise RuntimeError(f"expected root plus 48 FSDP2 blocks, got {configured}")
if dist.get_rank() == 0:
print(f"BF16_SYMM_MEM enabled on {len(configured)} modules", flush=True)
return module
ltx2_module.load_module_from_path = load_with_symm_mem
DummyTracker.log = _log
trainer_module.build_tracker = lambda *_args, **_kwargs: DummyTracker()
original_trainer_init = trainer_module.Trainer.__init__
class _Recorder:
def on_training_step_end(self, _method, metrics, iteration=0):
_step_times.append(float(metrics["step_time_sec"]))
def _trainer_init(self, *init_args, **init_kwargs):
original_trainer_init(self, *init_args, **init_kwargs)
self.callbacks._callbacks.pop("validation", None)
self.callbacks._callbacks["_benchmark_recorder"] = _Recorder()
trainer_module.Trainer.__init__ = _trainer_init
train_main(argparse.Namespace(config=args.config, dry_run=False), overrides=overrides or None)
world = get_world_group()
times_by_rank = [None] * world.world_size
dist.all_gather_object(times_by_rank, _step_times, group=world.cpu_group)
if world.rank == 0:
per_step_max = [max(values) for values in zip(*times_by_rank, strict=True)]
measured = per_step_max[10:30]
if len(measured) != 20:
raise RuntimeError(f"expected 30 steps, got {len(per_step_max)}")
median = statistics.median(measured)
print("BF16_RESULT " + json.dumps({
"kind": "fsdp2_symm_mem",
"median_step_sec": median,
"samples_per_second": 4.0 / median,
"model_mfu_percent": 14.444115 / median,
}, sort_keys=True), flush=True)
if __name__ == "__main__":
main()
@@ -0,0 +1,48 @@
import argparse
import sys
import torch
from fastvideo.train.entrypoint.train import main as train_main
from fastvideo.train.trainer import Trainer
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True)
parser.add_argument("--trace-dir", required=True)
args, overrides = parser.parse_known_args()
profiler = torch.profiler.profile(
activities=(torch.profiler.ProfilerActivity.CPU,
torch.profiler.ProfilerActivity.CUDA),
schedule=torch.profiler.schedule(wait=10, warmup=1, active=1, repeat=1),
on_trace_ready=torch.profiler.tensorboard_trace_handler(
args.trace_dir, use_gzip=True),
record_shapes=False,
profile_memory=False,
with_stack=False,
with_flops=False,
)
original_iter = Trainer._iter_dataloader
original_run = Trainer.run
def profiled_iter(self, dataloader):
iterator = original_iter(self, dataloader)
while True:
batch = next(iterator)
yield batch
profiler.step()
def profiled_run(self, *run_args, **run_kwargs):
with profiler:
return original_run(self, *run_args, **run_kwargs)
Trainer._iter_dataloader = profiled_iter
Trainer.run = profiled_run
train_args = argparse.Namespace(config=args.config, dry_run=False)
train_main(train_args, overrides=overrides or None)
if __name__ == "__main__":
sys.exit(main())
@@ -0,0 +1,138 @@
#!/usr/bin/env python3
"""Rank-zero Kineto trace for the packed LTX-2 training benchmark.
This wraps the already-gated benchmark harness so its semantic and optimizer
checks stay identical. Profiler timings are diagnostics, never MFU evidence.
"""
from __future__ import annotations
from collections import defaultdict
import builtins
import json
import os
from pathlib import Path
from typing import Any
import torch
import torch.distributed as dist
import benchmark_fastvideo_train_pack_d016 as benchmark
import fastvideo.train.trainer as trainer_module
WAIT_STEPS = 10
WARMUP_STEPS = 1
ACTIVE_STEPS = 2
TOTAL_STEPS = 14
TOP_ROWS = 20
OUTPUT_PREFIX = Path(os.environ.get("FASTVIDEO_KINETO_PREFIX", "/mnt/pr1630_pack_fa47ce1_rank0"))
def _cuda_time_us(event: Any, *, self_time: bool) -> float:
prefixes = ("self_", "") if self_time else ("", "self_")
for prefix in prefixes:
for suffix in ("device_time_total", "cuda_time_total"):
value = getattr(event, prefix + suffix, None)
if value is not None:
return float(value)
return 0.0
def _top_cuda_rows(profiler: torch.profiler.profile) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
operators = sorted(
profiler.key_averages(),
key=lambda event: _cuda_time_us(event, self_time=True),
reverse=True,
)[:TOP_ROWS]
operator_rows = [{
"name": event.key,
"calls": int(event.count),
"self_cuda_ms": round(_cuda_time_us(event, self_time=True) / 1000.0, 3),
"total_cuda_ms": round(_cuda_time_us(event, self_time=False) / 1000.0, 3),
} for event in operators if _cuda_time_us(event, self_time=True) > 0]
kernels: dict[str, list[float]] = defaultdict(lambda: [0.0, 0.0])
for event in profiler.events():
if not str(getattr(event, "device_type", "")).endswith("CUDA"):
continue
row = kernels[event.name]
row[0] += 1
row[1] += _cuda_time_us(event, self_time=True)
kernel_rows = [{
"name": name,
"calls": int(count),
"cuda_ms": round(cuda_us / 1000.0, 3),
} for name, (count, cuda_us) in sorted(kernels.items(), key=lambda item: item[1][1], reverse=True)[:TOP_ROWS]]
return operator_rows, kernel_rows
def _trace_ready(profiler: torch.profiler.profile) -> None:
trace_path = OUTPUT_PREFIX.with_suffix(".trace.json.gz")
summary_path = OUTPUT_PREFIX.with_suffix(".summary.json")
trace_path.parent.mkdir(parents=True, exist_ok=True)
profiler.export_chrome_trace(str(trace_path))
operators, kernels = _top_cuda_rows(profiler)
if not operators or not kernels:
raise RuntimeError("Kineto captured no CUDA operator or kernel events")
summary = {
"diagnostic_only_not_mfu_evidence": True,
"rank": 0,
"schedule": {
"wait_steps": WAIT_STEPS,
"warmup_steps": WARMUP_STEPS,
"active_steps": ACTIVE_STEPS,
},
"trace": str(trace_path),
"trace_bytes": trace_path.stat().st_size,
"top_cuda_operators": operators,
"top_cuda_kernels": kernels,
}
summary_path.write_text(json.dumps(summary, indent=2) + "\n")
builtins.print("KINETO_SUMMARY " + json.dumps(summary, separators=(",", ":")), flush=True)
class _ProfilerStep:
def __init__(self, profiler: torch.profiler.profile) -> None:
self.profiler = profiler
def on_training_step_end(self, _method: Any, _metrics: dict[str, Any], iteration: int = 0) -> None:
del iteration
self.profiler.step()
def _quiet_benchmark_print(*values: Any, **kwargs: Any) -> None:
if values and str(values[0]).startswith("BF16_"):
return
builtins.print(*values, **kwargs)
def main() -> None:
benchmark.EXPECTED_STEPS = TOTAL_STEPS
benchmark.print = _quiet_benchmark_print
original_run = trainer_module.Trainer.run
def _profiled_run(self: Any, method: Any, **kwargs: Any) -> Any:
if not dist.is_initialized() or dist.get_rank() != 0:
return original_run(self, method, **kwargs)
with torch.profiler.profile(
activities=[torch.profiler.ProfilerActivity.CPU, torch.profiler.ProfilerActivity.CUDA],
schedule=torch.profiler.schedule(wait=WAIT_STEPS, warmup=WARMUP_STEPS, active=ACTIVE_STEPS, repeat=1),
on_trace_ready=_trace_ready,
record_shapes=False,
profile_memory=False,
with_stack=False,
with_flops=False,
) as profiler:
self.callbacks._callbacks["_kineto_step"] = _ProfilerStep(profiler)
result = original_run(self, method, **kwargs)
torch.cuda.synchronize()
return result
trainer_module.Trainer.run = _profiled_run
benchmark.main()
if __name__ == "__main__":
main()
@@ -0,0 +1,155 @@
#!/usr/bin/env python3
"""Rank-zero Kineto trace with kernel-category rollup for the packed harness.
Same wrapping approach as ``profile_fastvideo_train_pack_fa47ce1.py`` but adds
a per-category CUDA-time decomposition (gemm / attention / optimizer / comm /
memcpy / compiled-fused / eager elementwise / reduce+norm / other) so the
non-GEMM, non-attention band can be attributed by name. Diagnostics only,
never MFU evidence.
"""
from __future__ import annotations
from collections import defaultdict
import builtins
import json
import os
from pathlib import Path
from typing import Any
import torch
import torch.distributed as dist
import benchmark_fastvideo_train_pack_d016 as benchmark
import fastvideo.train.trainer as trainer_module
WAIT_STEPS = 10
WARMUP_STEPS = 1
ACTIVE_STEPS = 2
TOTAL_STEPS = 14
TOP_ROWS = 40
OUTPUT_PREFIX = Path(os.environ.get("FASTVIDEO_KINETO_PREFIX", "/mnt/pr1630_pack_head_rank0"))
def _self_cuda_us(event: Any) -> float:
for name in ("self_device_time_total", "self_cuda_time_total"):
value = getattr(event, name, None)
if value is not None:
return float(value)
return 0.0
def _category(name: str) -> str:
lowered = name.lower()
if "nccl" in lowered:
return "comm"
if "memcpy" in lowered or "memset" in lowered:
return "memcpy"
if any(key in lowered for key in ("nvjet", "cutlass", "gemm", "cublas", "matmul", "splitk")):
return "gemm"
if any(key in lowered for key in ("flash", "fmha", "cute", "attention", "_attn")):
return "attention"
if "adam" in lowered or "multi_tensor" in lowered:
return "optimizer"
if lowered.startswith("triton_"):
return "compiled_fused"
if "elementwise" in lowered or "vectorized" in lowered:
return "eager_elementwise"
if any(key in lowered for key in ("reduce", "norm", "welford", "softmax")):
return "reduce_norm"
return "other"
def _trace_ready(profiler: torch.profiler.profile) -> None:
trace_path = OUTPUT_PREFIX.with_suffix(".trace.json.gz")
summary_path = OUTPUT_PREFIX.with_suffix(".summary.json")
trace_path.parent.mkdir(parents=True, exist_ok=True)
profiler.export_chrome_trace(str(trace_path))
kernels: dict[str, list[float]] = defaultdict(lambda: [0.0, 0.0])
for event in profiler.events():
if not str(getattr(event, "device_type", "")).endswith("CUDA"):
continue
row = kernels[event.name]
row[0] += 1
row[1] += _self_cuda_us(event)
if not kernels:
raise RuntimeError("Kineto captured no CUDA kernel events")
categories: dict[str, list[float]] = defaultdict(lambda: [0.0, 0.0])
for name, (count, cuda_us) in kernels.items():
row = categories[_category(name)]
row[0] += count
row[1] += cuda_us
category_rows = [{
"category": category,
"calls_per_step": round(count / ACTIVE_STEPS, 1),
"cuda_ms_per_step": round(cuda_us / 1000.0 / ACTIVE_STEPS, 3),
} for category, (count, cuda_us) in sorted(categories.items(), key=lambda item: item[1][1], reverse=True)]
kernel_rows = [{
"name": name[:160],
"calls_per_step": round(count / ACTIVE_STEPS, 1),
"cuda_ms_per_step": round(cuda_us / 1000.0 / ACTIVE_STEPS, 3),
"category": _category(name),
} for name, (count, cuda_us) in sorted(kernels.items(), key=lambda item: item[1][1], reverse=True)[:TOP_ROWS]]
summary = {
"diagnostic_only_not_mfu_evidence": True,
"rank": 0,
"active_steps": ACTIVE_STEPS,
"trace": str(trace_path),
"trace_bytes": trace_path.stat().st_size,
"total_cuda_ms_per_step": round(
sum(row[1] for row in kernels.values()) / 1000.0 / ACTIVE_STEPS, 3),
"category_rollup": category_rows,
"top_cuda_kernels": kernel_rows,
}
summary_path.write_text(json.dumps(summary, indent=2) + "\n")
builtins.print("KINETO_CATEGORIES " + json.dumps(summary, separators=(",", ":")), flush=True)
class _ProfilerStep:
def __init__(self, profiler: torch.profiler.profile) -> None:
self.profiler = profiler
def on_training_step_end(self, _method: Any, _metrics: dict[str, Any], iteration: int = 0) -> None:
del iteration
self.profiler.step()
def _quiet_benchmark_print(*values: Any, **kwargs: Any) -> None:
if values and str(values[0]).startswith("BF16_"):
return
builtins.print(*values, **kwargs)
def main() -> None:
benchmark.EXPECTED_STEPS = TOTAL_STEPS
benchmark.print = _quiet_benchmark_print
original_run = trainer_module.Trainer.run
def _profiled_run(self: Any, method: Any, **kwargs: Any) -> Any:
if not dist.is_initialized() or dist.get_rank() != 0:
return original_run(self, method, **kwargs)
with torch.profiler.profile(
activities=[torch.profiler.ProfilerActivity.CPU, torch.profiler.ProfilerActivity.CUDA],
schedule=torch.profiler.schedule(wait=WAIT_STEPS, warmup=WARMUP_STEPS, active=ACTIVE_STEPS, repeat=1),
on_trace_ready=_trace_ready,
record_shapes=False,
profile_memory=False,
with_stack=False,
with_flops=False,
) as profiler:
self.callbacks._callbacks["_kineto_step"] = _ProfilerStep(profiler)
result = original_run(self, method, **kwargs)
torch.cuda.synchronize()
return result
trainer_module.Trainer.run = _profiled_run
benchmark.main()
if __name__ == "__main__":
main()
@@ -0,0 +1,158 @@
#!/usr/bin/env python3
"""Audit the PR #1630 MFU numerator by counting executed per-sample training FLOPs.
Wraps each counted microstep (forward ``single_train_step`` plus ``backward``)
in ``torch.utils.flop_counter.FlopCounterMode`` and reports per-rank forward,
backward, and total FLOPs. Run with ``--models.student.attention_backend
TORCH_SDPA`` and ``--training.model.enable_torch_compile false``: compiled
regions and the FA4 custom op bypass the dispatch-mode counter, so FLASH_ATTN
or compiled runs undercount attention. The counter uses PyTorch's SDPA flop
formulas, which charge the flash backward recompute (about 2.5x the attention
forward) rather than the 2x pure-gradient convention.
"""
from __future__ import annotations
import argparse
import json
import os
import torch
import torch.distributed as dist
from torch.utils.flop_counter import FlopCounterMode
COUNT_STEPS = (2, 3)
_records: list[dict[str, int]] = []
_input_shapes: dict[str, list] = {}
_run_config: dict[str, int] = {}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True)
args, overrides = parser.parse_known_args()
from fastvideo.training.trackers import DummyTracker
import fastvideo.train.trainer as trainer_module
from fastvideo.distributed import get_world_group
from fastvideo.train.entrypoint.train import main as train_main
from fastvideo.train.methods.fine_tuning.finetune import FineTuneMethod
from fastvideo.train.models.ltx2 import LTX2Model
trainer_module.build_tracker = lambda *_a, **_k: DummyTracker()
original_trainer_init = trainer_module.Trainer.__init__
def _trainer_init(self, *init_args, **init_kwargs):
original_trainer_init(self, *init_args, **init_kwargs)
self.callbacks._callbacks.pop("validation", None)
_run_config["local_batch_size"] = int(self.training_config.data.train_batch_size)
_run_config["gradient_accumulation_steps"] = int(
self.training_config.loop.gradient_accumulation_steps or 1)
if bool(self.training_config.model.enable_torch_compile):
raise RuntimeError("FLOP audit requires enable_torch_compile=false")
trainer_module.Trainer.__init__ = _trainer_init
original_run = trainer_module.Trainer.run
def _trainer_run(self, method, **kwargs):
vae = getattr(getattr(method, "student", None), "vae", None)
if vae is not None:
method.student.vae = None
return original_run(self, method, **kwargs)
trainer_module.Trainer.run = _trainer_run
original_build = LTX2Model._build_distill_input_kwargs
def _build(self, *build_args, **build_kwargs):
result = original_build(self, *build_args, **build_kwargs)
if not _input_shapes:
for key, value in result.items():
if isinstance(value, torch.Tensor):
_input_shapes[key] = [list(value.shape), str(value.dtype)]
return result
LTX2Model._build_distill_input_kwargs = _build
# FineTuneMethod overrides both single_train_step and backward, so the
# subclass attributes must be patched; base-class patches never fire.
original_step = FineTuneMethod.single_train_step
original_backward = FineTuneMethod.backward
state: dict[str, object] = {"mode": None, "fwd": 0, "step": None}
def _counted_step(self, batch, step):
if int(step) in COUNT_STEPS:
mode = FlopCounterMode(display=False)
mode.__enter__()
state["mode"] = mode
state["step"] = int(step)
outputs = original_step(self, batch, step)
if state["mode"] is not None:
state["fwd"] = state["mode"].get_total_flops()
return outputs
def _counted_backward(self, loss_map, outputs, **kwargs):
result = original_backward(self, loss_map, outputs, **kwargs)
mode = state["mode"]
if mode is not None:
mode.__exit__(None, None, None)
total = int(mode.get_total_flops())
fwd = int(state["fwd"])
_records.append({
"step": state["step"],
"forward_flops": fwd,
"backward_flops": total - fwd,
"total_flops": total,
})
state["mode"] = None
return result
FineTuneMethod.single_train_step = _counted_step
FineTuneMethod.backward = _counted_backward
train_main(argparse.Namespace(config=args.config, dry_run=False), overrides=overrides or None)
if len(_records) != len(COUNT_STEPS):
raise RuntimeError(f"expected {len(COUNT_STEPS)} counted microsteps, got {len(_records)}")
properties = torch.cuda.get_device_properties(0)
payload = {
"records": _records,
"input_shapes": _input_shapes,
"run_config": _run_config,
"device": {
"name": properties.name,
"multi_processor_count": properties.multi_processor_count,
"capability": list(torch.cuda.get_device_capability(0)),
},
"torch": torch.__version__,
"attention_backend_env": os.environ.get("FASTVIDEO_ATTENTION_BACKEND"),
}
world = get_world_group()
payloads: list[dict | None] = [None] * world.world_size
dist.all_gather_object(payloads, payload, group=world.cpu_group)
if world.rank != 0:
return
totals = {tuple(sorted(r["total_flops"] for r in p["records"])) for p in payloads}
per_sample = None
batch = payloads[0]["run_config"]["local_batch_size"]
steady = payloads[0]["records"][-1]["total_flops"]
per_sample = steady / batch
print(
"FLOP_AUDIT " + json.dumps({
"by_rank": payloads,
"rank_total_sets_identical": len(totals) == 1,
"local_batch_size": batch,
"per_sample_total_flops_last_counted_step": per_sample,
"per_sample_tflops": per_sample / 1e12,
}, sort_keys=True),
flush=True,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,104 @@
#!/usr/bin/env python3
"""FA4 forward-efficiency tail-tile discriminator (plan item 1).
The trainer profile shows FA4 forward at 34.9% of peak while backward runs
46.2% — inverted vs every flash-attention generation. One candidate mechanism
is the ragged sequence: S=4,290 is 33.5 tiles of 128. This probe measures
per-FLOP throughput of the production FA4 path at S in {4,224, 4,290, 4,352}
(clean / ragged / clean tile counts), forward-only and forward+backward, at
B2 and B3. If the clean neighbors run >2-3% faster per FLOP than 4,290 in
forward, the tail hypothesis is confirmed and a seqused/padded integration is
worth pursuing; if per-FLOP throughput is flat, the anomaly lives in schedule
or occupancy instead. Ratios only; degraded-bin trays are acceptable.
"""
from __future__ import annotations
import argparse
import json
import statistics
import torch
HEADS, HEAD_DIM = 32, 128
def time_fn(fn, warmup: int, repeats: int) -> float:
for _ in range(warmup):
fn()
torch.cuda.synchronize()
events = [(torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True)) for _ in range(repeats)]
for start, end in events:
start.record()
fn()
end.record()
torch.cuda.synchronize()
return statistics.median(start.elapsed_time(end) for start, end in events)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--warmup", type=int, default=10)
parser.add_argument("--repeats", type=int, default=31)
args = parser.parse_args()
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func as fa4
torch.manual_seed(20260722)
device = torch.device("cuda:0")
rows = []
for batch in (2, 3):
for tokens in (4224, 4290, 4352):
shape = (batch, tokens, HEADS, HEAD_DIM)
q = torch.randn(shape, device=device, dtype=torch.bfloat16, requires_grad=True)
k = torch.randn(shape, device=device, dtype=torch.bfloat16, requires_grad=True)
v = torch.randn(shape, device=device, dtype=torch.bfloat16, requires_grad=True)
grad_out = torch.randn(shape, device=device, dtype=torch.bfloat16)
def fwd_only():
with torch.no_grad():
fa4(q, k, v)
def fwd_bwd():
q.grad = k.grad = v.grad = None
fa4(q, k, v).backward(grad_out)
fwd_ms = time_fn(fwd_only, args.warmup, args.repeats)
total_ms = time_fn(fwd_bwd, args.warmup, args.repeats)
d_total = HEADS * HEAD_DIM
fwd_tf = 4.0 * batch * tokens * tokens * d_total / 1e12
bwd_tf = 10.0 * batch * tokens * tokens * d_total / 1e12
row = {
"batch": batch,
"tokens": tokens,
"tiles_128": tokens / 128.0,
"fwd_ms": round(fwd_ms, 4),
"fwd_tflops_per_s": round(fwd_tf / (fwd_ms / 1000.0), 1),
"fwd_bwd_ms": round(total_ms, 4),
"bwd_ms_est": round(total_ms - fwd_ms, 4),
"bwd_tflops_per_s_est": round(bwd_tf / ((total_ms - fwd_ms) / 1000.0), 1),
}
rows.append(row)
print("FA4_TAIL_CASE " + json.dumps(row, sort_keys=True), flush=True)
del q, k, v, grad_out
torch.cuda.empty_cache()
for batch in (2, 3):
sub = {r["tokens"]: r for r in rows if r["batch"] == batch}
ragged = sub[4290]["fwd_tflops_per_s"]
clean_low, clean_high = sub[4224]["fwd_tflops_per_s"], sub[4352]["fwd_tflops_per_s"]
print(
"FA4_TAIL_VERDICT " + json.dumps({
"batch": batch,
"fwd_ragged_tflops": ragged,
"fwd_clean_low_tflops": clean_low,
"fwd_clean_high_tflops": clean_high,
"clean_over_ragged_pct": round(
100.0 * (max(clean_low, clean_high) - ragged) / ragged, 2),
}, sort_keys=True),
flush=True,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,128 @@
#!/usr/bin/env python3
"""Exact seven-pack LTX-2 all-pass lower bound with installed SM100 CuTe kernels."""
import gc
import json
import statistics
import torch
import torch.nn.functional as F
from flashinfer import fp4_quantize, mm_fp4
LAYERS = (
("self_qkv", 48, 4290, 4096, 12288),
("video_dd", 144, 4290, 4096, 4096),
("text_kv", 48, 1024, 4096, 8192),
("ffn_up", 48, 4290, 4096, 16384),
("ffn_down", 48, 4290, 16384, 4096),
)
def timed(fn, warmup=3, samples=5, inner=3):
for _ in range(warmup):
fn()
torch.cuda.synchronize()
values = []
for _ in range(samples):
begin, end = torch.cuda.Event(True), torch.cuda.Event(True)
begin.record()
for _ in range(inner):
fn()
end.record()
end.synchronize()
values.append(begin.elapsed_time(end) / inner)
return statistics.median(values)
def global_scale(x):
return (448.0 * 6.0) / x.float().abs().nan_to_num().amax().clamp_min(1e-12)
def quant(x, scale):
return fp4_quantize(x, scale, backend="cute-dsl", enable_pdl=True)
def operands(m, k, n, phase):
if phase == "fwd":
return (torch.randn(m, k, device="cuda", dtype=torch.bfloat16),
torch.randn(n, k, device="cuda", dtype=torch.bfloat16) * 0.02, True)
if phase == "dgrad":
return (torch.randn(m, n, device="cuda", dtype=torch.bfloat16) * 0.01,
(torch.randn(n, k, device="cuda", dtype=torch.bfloat16) * 0.02).T.contiguous(), True)
dy = torch.randn(m, n, device="cuda", dtype=torch.bfloat16) * 0.01
x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16)
padded = (m + 31) // 32 * 32
if padded != m:
dy, x = F.pad(dy, (0, 0, 0, padded - m)), F.pad(x, (0, 0, 0, padded - m))
return dy.T.contiguous(), x.T.contiguous(), False
def bench(name, count, m, k, n, phase):
lhs, rhs_rows, rhs_is_weight = operands(m, k, n, phase)
ls, rs = global_scale(lhs), global_scale(rhs_rows)
lq, lsf = quant(lhs, ls)
rq, rsf = quant(rhs_rows, rs)
alpha = (ls * rs).reciprocal()
out = torch.empty(lhs.shape[0], rhs_rows.shape[0], device="cuda", dtype=torch.bfloat16)
def mm():
return mm_fp4(lq, rq.T, lsf, rsf.T, alpha, out=out, backend="cute-dsl", enable_pdl=True)
def delayed():
aq, asf = quant(lhs, ls)
if rhs_is_weight:
bq, bsf = rq, rsf
else:
bq, bsf = quant(rhs_rows, rs)
return mm_fp4(aq, bq.T, asf, bsf.T, alpha, out=out, backend="cute-dsl", enable_pdl=True)
def exact():
als, brs = global_scale(lhs), global_scale(rhs_rows)
aq, asf = quant(lhs, als)
if rhs_is_weight:
bq, bsf = rq, rsf
brs = rs
else:
bq, bsf = quant(rhs_rows, brs)
return mm_fp4(aq, bq.T, asf, bsf.T, (als * brs).reciprocal(), out=out,
backend="cute-dsl", enable_pdl=True)
reference = lhs @ rhs_rows.T
actual = mm().clone()
torch.cuda.synchronize()
diff = actual.float() - reference.float()
row = {
"name": name, "phase": phase, "count": count,
"logical_shape": [m, k, n], "qmm_shape": [lhs.shape[0], lhs.shape[1], rhs_rows.shape[0]],
"bf16_ms": timed(lambda: torch.mm(lhs, rhs_rows.T, out=out)),
"prequant_ms": timed(mm), "delayed_ms": timed(delayed), "exact_ms": timed(exact),
"rhs_refresh_delayed_ms": timed(lambda: quant(rhs_rows, rs), inner=1) if rhs_is_weight else 0.0,
"rhs_refresh_exact_ms": timed(lambda: quant(rhs_rows, global_scale(rhs_rows)), inner=1) if rhs_is_weight else 0.0,
"relative_rms": float(diff.square().mean().sqrt() / reference.float().square().mean().sqrt()),
}
print(json.dumps({"kind": "case", **row}, sort_keys=True), flush=True)
return row
if __name__ == "__main__":
assert torch.cuda.get_device_capability() == (10, 0)
torch.manual_seed(20260721)
rows = []
for layer in LAYERS:
for phase in ("fwd", "dgrad", "wgrad"):
rows.append(bench(*layer, phase))
gc.collect()
torch.cuda.empty_cache()
totals = {}
for tier in ("bf16", "prequant", "delayed", "exact"):
value = sum(r[f"{tier}_ms"] * r["count"] for r in rows)
if tier == "delayed":
value += sum(r["rhs_refresh_delayed_ms"] * r["count"] for r in rows)
elif tier == "exact":
value += sum(r["rhs_refresh_exact_ms"] * r["count"] for r in rows)
totals[tier] = {"weighted_ms_48_blocks": value, "ms_per_block": value / 48}
baseline = 198.4054238319397
print(json.dumps({"kind": "aggregate", "baseline_separate_bf16_ms": baseline,
"target_ms": baseline / 2.5, "robust_target_ms": baseline / 2.7,
"totals": totals}, sort_keys=True))
@@ -0,0 +1,326 @@
#!/usr/bin/env python3
"""Reject-fast GB200 gate for deferred/grouped LTX-2 BF16 weight gradients."""
from __future__ import annotations
import argparse
import gc
import hashlib
import json
import math
import statistics
from dataclasses import dataclass
from typing import Callable
import torch
DTYPE = torch.bfloat16
MASTER_DTYPE = torch.float32
BLOCKS = 48
VIDEO_TOKENS = 4290
TEXT_TOKENS = 1024
HIDDEN = 4096
FFN = 16384
RTOL = 1.6e-2
ATOL = 1.0e-2
@dataclass(frozen=True)
class Case:
name: str
tokens: int
in_features: int
out_features: int
roles_per_block: int = 1
CASES = (
Case("self_qkv", VIDEO_TOKENS, HIDDEN, 3 * HIDDEN),
# Packed LTX has three separate 4096 -> 4096 video-token roles:
# self-attention output, text-cross query, and text-cross output.
Case("video_4096", VIDEO_TOKENS, HIDDEN, HIDDEN, roles_per_block=3),
Case("text_kv", TEXT_TOKENS, HIDDEN, 2 * HIDDEN),
Case("ffn_up", VIDEO_TOKENS, HIDDEN, FFN),
Case("ffn_down", VIDEO_TOKENS, FFN, HIDDEN),
)
def emit(kind: str, **payload: object) -> None:
print(json.dumps({"kind": kind, **payload}, sort_keys=True), flush=True)
def elapsed_ms(fn: Callable[[], None], inner: int) -> float:
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(inner):
fn()
end.record()
end.synchronize()
return start.elapsed_time(end) / inner
def paired_times(
serial: Callable[[], None],
prepacked: Callable[[], None],
staged: Callable[[], None],
staged_scatter: Callable[[], None],
*,
warmup: int,
samples: int,
inner: int,
) -> dict[str, object]:
methods = (serial, prepacked, staged, staged_scatter)
for _ in range(warmup):
for method in methods:
method()
torch.cuda.synchronize()
series = {name: [] for name in ("serial_a", "prepacked", "staged", "staged_scatter", "serial_b")}
for _ in range(samples):
series["serial_a"].append(elapsed_ms(serial, inner))
series["prepacked"].append(elapsed_ms(prepacked, inner))
series["staged"].append(elapsed_ms(staged, inner))
series["staged_scatter"].append(elapsed_ms(staged_scatter, inner))
series["serial_b"].append(elapsed_ms(serial, inner))
medians = {name: statistics.median(values) for name, values in series.items()}
paired_serial_midpoint = [
(serial_a + serial_b) / 2
for serial_a, serial_b in zip(series["serial_a"], series["serial_b"], strict=True)
]
paired_savings = {
mode: [midpoint - candidate for midpoint, candidate in zip(paired_serial_midpoint, series[mode], strict=True)]
for mode in ("prepacked", "staged", "staged_scatter")
}
medians["serial_midpoint"] = statistics.median(paired_serial_midpoint)
medians["serial_control_drift_percent"] = (
abs(medians["serial_b"] - medians["serial_a"]) / medians["serial_midpoint"] * 100
)
return {
"medians_ms": medians,
"samples_ms": series,
"paired_serial_midpoint_samples_ms": paired_serial_midpoint,
"paired_savings_median_ms": {
mode: statistics.median(values) for mode, values in paired_savings.items()
},
"paired_savings_samples_ms": paired_savings,
}
def parity_metrics(
x_sources: list[torch.Tensor],
dy_sources: list[torch.Tensor],
packed_dw: torch.Tensor,
serial_dw: torch.Tensor,
) -> dict[str, object]:
max_abs = 0.0
error_sq = 0.0
reference_sq = 0.0
close = True
sample_sha256 = hashlib.sha256()
for batch_index, (x, dy) in enumerate(zip(x_sources, dy_sources, strict=True)):
torch.mm(dy.transpose(0, 1), x, out=serial_dw)
# Chunking keeps the parity probe well below the FFN gradient's size.
for row in range(0, serial_dw.shape[0], 256):
ref = serial_dw[row:row + 256].float()
candidate = packed_dw[batch_index, row:row + 256].float()
error = candidate - ref
max_abs = max(max_abs, float(error.abs().max()))
error_sq += float(torch.sum(error * error, dtype=torch.float64))
reference_sq += float(torch.sum(ref * ref, dtype=torch.float64))
close = close and bool(torch.all(error.abs() <= ATOL + RTOL * ref.abs()))
sample = packed_dw[batch_index, ::max(1, packed_dw.shape[1] // 7),
::max(1, packed_dw.shape[2] // 7)]
sample_sha256.update(sample.contiguous().view(torch.uint8).cpu().numpy().tobytes())
relative_l2 = math.sqrt(error_sq / reference_sq) if reference_sq else 0.0
if not close:
raise RuntimeError(
f"grouped wgrad parity failed: max_abs={max_abs} relative_l2={relative_l2} "
f"rtol={RTOL} atol={ATOL}"
)
return {
"allclose": close,
"rtol": RTOL,
"atol": ATOL,
"max_abs": max_abs,
"relative_l2": relative_l2,
"sample_sha256": sample_sha256.hexdigest(),
}
def phase_context(case: Case, group: int, warmup: int, samples: int, inner: int) -> dict[str, float]:
x = torch.empty((case.tokens, case.in_features), device="cuda", dtype=DTYPE).normal_(std=0.02)
dy = torch.empty((case.tokens, case.out_features), device="cuda", dtype=DTYPE).normal_(std=0.02)
weight = torch.empty((case.out_features, case.in_features), device="cuda", dtype=DTYPE).normal_(std=0.02)
fwd_out = torch.empty((case.tokens, case.out_features), device="cuda", dtype=DTYPE)
dgrad_out = torch.empty((case.tokens, case.in_features), device="cuda", dtype=DTYPE)
def fwd() -> None:
for _ in range(group):
torch.mm(x, weight.transpose(0, 1), out=fwd_out)
def dgrad() -> None:
for _ in range(group):
torch.mm(dy, weight, out=dgrad_out)
for _ in range(warmup):
fwd()
dgrad()
torch.cuda.synchronize()
fwd_ms = statistics.median(elapsed_ms(fwd, inner) for _ in range(samples))
dgrad_ms = statistics.median(elapsed_ms(dgrad, inner) for _ in range(samples))
del x, dy, weight, fwd_out, dgrad_out
return {"forward_ms": fwd_ms, "recompute_ms": fwd_ms, "dgrad_ms": dgrad_ms}
def bench_case(case: Case, group: int, args: argparse.Namespace) -> dict[str, object]:
torch.manual_seed(20260721 + group)
x_sources = [
torch.empty((case.tokens, case.in_features), device="cuda", dtype=DTYPE).normal_(std=0.02)
for _ in range(group)
]
dy_sources = [
torch.empty((case.tokens, case.out_features), device="cuda", dtype=DTYPE).normal_(std=0.02)
for _ in range(group)
]
x_packed = torch.empty((group, case.tokens, case.in_features), device="cuda", dtype=DTYPE)
dy_packed = torch.empty((group, case.tokens, case.out_features), device="cuda", dtype=DTYPE)
packed_dw = torch.empty((group, case.out_features, case.in_features), device="cuda", dtype=DTYPE)
serial_dw = torch.empty((case.out_features, case.in_features), device="cuda", dtype=DTYPE)
scattered_dw = [torch.empty_like(serial_dw) for _ in range(group)]
def pack() -> None:
for index in range(group):
x_packed[index].copy_(x_sources[index])
dy_packed[index].copy_(dy_sources[index])
def serial() -> None:
for x, dy in zip(x_sources, dy_sources, strict=True):
torch.mm(dy.transpose(0, 1), x, out=serial_dw)
def prepacked() -> None:
torch.bmm(dy_packed.transpose(1, 2), x_packed, out=packed_dw)
def staged() -> None:
pack()
prepacked()
def staged_scatter() -> None:
staged()
for index in range(group):
scattered_dw[index].copy_(packed_dw[index])
pack()
prepacked()
torch.cuda.synchronize()
parity = parity_metrics(x_sources, dy_sources, packed_dw, serial_dw)
timings = paired_times(
serial,
prepacked,
staged,
staged_scatter,
warmup=args.warmup,
samples=args.samples,
inner=args.inner,
)
context = phase_context(case, group, args.warmup, args.samples, args.inner)
groups_per_step = BLOCKS // group
multiplier = groups_per_step * case.roles_per_block
projections = {
mode: timings["paired_savings_median_ms"][mode] * multiplier
for mode in ("prepacked", "staged", "staged_scatter")
}
result = {
"case": case.name,
"logical_shape_mkn": [case.tokens, case.in_features, case.out_features],
"group": group,
"blocks": BLOCKS,
"roles_per_block": case.roles_per_block,
"groups_per_step": groups_per_step,
"working_dtype": str(DTYPE),
"gradient_dtype": str(packed_dw.dtype),
"master_weight_dtype": str(MASTER_DTYPE),
"parity": parity,
"phase_context_ms_per_group": context,
**timings,
"projected_kernel_local_step_savings_ms": projections,
}
del x_sources, dy_sources, x_packed, dy_packed, packed_dw, serial_dw, scattered_dw
gc.collect()
torch.cuda.empty_cache()
return result
def self_test() -> None:
assert BLOCKS % 2 == 0 and BLOCKS % 4 == 0
assert sum(case.roles_per_block for case in CASES) == 7
assert DTYPE == torch.bfloat16 and MASTER_DTYPE == torch.float32
print("SELF_TEST_OK")
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--groups", type=int, nargs="+", default=[2, 4])
parser.add_argument("--warmup", type=int, default=3)
parser.add_argument("--samples", type=int, default=9)
parser.add_argument("--inner", type=int, default=2)
parser.add_argument("--self-test", action="store_true")
args = parser.parse_args()
if args.self_test:
self_test()
return
if any(group not in (2, 4) or BLOCKS % group for group in args.groups):
parser.error("groups must be 2 or 4 and divide 48")
if min(args.warmup, args.samples, args.inner) < 1:
parser.error("warmup, samples, and inner must be positive")
torch.cuda.set_device(0)
device_name = torch.cuda.get_device_name(0)
emit(
"environment",
gpu=device_name,
torch=torch.__version__,
cuda=torch.version.cuda,
blocks=BLOCKS,
checkpoint_forward_multiplicity=2,
optimizer_contract={
"working_weights": str(DTYPE),
"weight_gradients": str(DTYPE),
"resident_master_weights": str(MASTER_DTYPE),
"resident_moments": str(MASTER_DTYPE),
},
caveat=(
"kernel-local gate only; deferred wgrad can reduce FSDP reduce-scatter overlap, "
"so a passing result still requires an end-to-end trainer A/B"
),
)
results = []
for group in args.groups:
for case in CASES:
result = bench_case(case, group, args)
results.append(result)
emit("case", **result)
group_results = [result for result in results if result["group"] == group]
aggregate = {
mode: sum(result["projected_kernel_local_step_savings_ms"][mode] for result in group_results)
for mode in ("prepacked", "staged", "staged_scatter")
}
emit(
"aggregate",
group=group,
projected_kernel_local_step_savings_ms=aggregate,
conservative_gate_mode="staged_scatter",
threshold_ms=10.0,
passes_10ms_gate=aggregate["staged_scatter"] >= 10.0,
interpretation=(
"prepacked is the arithmetic ceiling; staged includes copies into contiguous bmm inputs; "
"staged_scatter additionally copies each result to a separate parameter-grad buffer"
),
)
if __name__ == "__main__":
main()
@@ -0,0 +1,241 @@
#!/usr/bin/env python3
"""cuBLASLt heuristic-algo sweep for the packed LTX-2 GEMM band (plan item 2).
Enumerates ``cublasLtMatmulAlgoGetHeuristic`` candidates per exact shape and
training orientation (fwd, dgrad, wgrad), times each against torch's own
dispatch in the same process, and reports the best-algo margin. Kill
criterion from ``reports/bf16_kernel_research_plan.md``: if the best
heuristic algo is within 2% of torch/nvjet per weighted band, native algo
pinning is closed. Bare GEMMs (no bias epilogue) on both sides; BF16 inputs,
FP32 compute/accumulate. Parity-checked against torch per case.
"""
from __future__ import annotations
import argparse
import ctypes
import json
import statistics
import torch
CUDA_R_16BF = 14
CUDA_R_32F = 0
CUBLAS_COMPUTE_32F = 68
OP_N, OP_T = 0, 1
DESC_TRANSA, DESC_TRANSB = 3, 4
PREF_MAX_WORKSPACE = 1
HIDDEN = 4096
FFN = 16384
VIDEO_TOKENS = 11 * 15 * 26
TEXT_TOKENS = 1024
WORKSPACE_BYTES = 128 * 1024 * 1024
MAX_ALGOS = 48
class HeuristicResult(ctypes.Structure):
_fields_ = [
("algo", ctypes.c_uint64 * 8),
("workspaceSize", ctypes.c_size_t),
("state", ctypes.c_int),
("wavesCount", ctypes.c_float),
("reserved", ctypes.c_int * 4),
]
def load_lt() -> ctypes.CDLL:
torch.cuda.init()
for name in ("libcublasLt.so.13", "libcublasLt.so.12", "libcublasLt.so"):
try:
return ctypes.CDLL(name)
except OSError:
continue
raise OSError("libcublasLt not found")
def check(status: int, what: str) -> None:
if status != 0:
raise RuntimeError(f"{what} failed with cublas status {status}")
def build_cases(batch: int) -> list[dict]:
m_video = VIDEO_TOKENS * batch
m_text = TEXT_TOKENS * batch
shapes = [
("self_qkv", m_video, HIDDEN, 3 * HIDDEN, 48),
("video_dd", m_video, HIDDEN, HIDDEN, 144),
("text_kv", m_text, HIDDEN, 2 * HIDDEN, 48),
("ffn_up", m_video, HIDDEN, FFN, 48),
("ffn_down", m_video, FFN, HIDDEN, 48),
]
cases = []
for name, rows, in_features, out_features, occurrences in shapes:
# Row-major training GEMMs expressed as column-major cuBLASLt calls
# (documented mapping: fwd y=xW^T, dgrad dx=dyW, wgrad dW=dy^T x).
cases.append({
"case": f"{name}:fwd", "occurrences": occurrences,
"m": out_features, "n": rows, "k": in_features,
"transa": OP_T, "transb": OP_N,
"a_shape": (out_features, in_features), "b_shape": (rows, in_features),
"d_shape": (rows, out_features),
"torch_fn": lambda a, b: b @ a.t(),
})
cases.append({
"case": f"{name}:dgrad", "occurrences": occurrences,
"m": in_features, "n": rows, "k": out_features,
"transa": OP_N, "transb": OP_N,
"a_shape": (out_features, in_features), "b_shape": (rows, out_features),
"d_shape": (rows, in_features),
"torch_fn": lambda a, b: b @ a,
})
cases.append({
"case": f"{name}:wgrad", "occurrences": occurrences,
"m": in_features, "n": out_features, "k": rows,
"transa": OP_N, "transb": OP_T,
"a_shape": (rows, in_features), "b_shape": (rows, out_features),
"d_shape": (out_features, in_features),
"torch_fn": lambda a, b: b.t() @ a,
})
return cases
def time_fn(fn, warmup: int, repeats: int) -> float:
for _ in range(warmup):
fn()
torch.cuda.synchronize()
events = [(torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True)) for _ in range(repeats)]
for start, end in events:
start.record()
fn()
end.record()
torch.cuda.synchronize()
return statistics.median(start.elapsed_time(end) for start, end in events)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--batch", type=int, default=2)
parser.add_argument("--warmup", type=int, default=8)
parser.add_argument("--repeats", type=int, default=21)
args = parser.parse_args()
torch.manual_seed(20260722)
device = torch.device("cuda:0")
lt = load_lt()
handle = ctypes.c_void_p()
check(lt.cublasLtCreate(ctypes.byref(handle)), "ltCreate")
workspace = torch.empty(WORKSPACE_BYTES, dtype=torch.uint8, device=device)
stream = ctypes.c_void_p(torch.cuda.current_stream().cuda_stream)
alpha = ctypes.c_float(1.0)
beta = ctypes.c_float(0.0)
summary = []
for case in build_cases(args.batch):
a_rm = torch.randn(case["a_shape"], device=device, dtype=torch.bfloat16)
b_rm = torch.randn(case["b_shape"], device=device, dtype=torch.bfloat16)
d = torch.empty(case["d_shape"], device=device, dtype=torch.bfloat16)
reference = case["torch_fn"](a_rm, b_rm)
# Column-major operand descriptors. Buffers: A is the weight-like
# tensor, B the activation-like tensor, both row-major contiguous.
m, n, k = case["m"], case["n"], case["k"]
op_desc = ctypes.c_void_p()
check(lt.cublasLtMatmulDescCreate(ctypes.byref(op_desc), CUBLAS_COMPUTE_32F, CUDA_R_32F), "descCreate")
for attr, value in ((DESC_TRANSA, case["transa"]), (DESC_TRANSB, case["transb"])):
v = ctypes.c_int32(value)
check(lt.cublasLtMatmulDescSetAttribute(op_desc, attr, ctypes.byref(v), 4), "descSet")
def layout(rows_cm: int, cols_cm: int, ld: int) -> ctypes.c_void_p:
handle_ = ctypes.c_void_p()
check(lt.cublasLtMatrixLayoutCreate(ctypes.byref(handle_), CUDA_R_16BF, rows_cm, cols_cm,
ctypes.c_int64(ld)), "layoutCreate")
return handle_
lda = case["a_shape"][1]
ldb = case["b_shape"][1]
a_rows_cm = lda if case["transa"] == OP_T else m
a_cols_cm = m if case["transa"] == OP_T else k
b_rows_cm = ldb if case["transb"] == OP_T else k
b_cols_cm = k if case["transb"] == OP_T else n
layout_a = layout(a_rows_cm, a_cols_cm, lda)
layout_b = layout(b_rows_cm, b_cols_cm, ldb)
layout_d = layout(m, n, m)
pref = ctypes.c_void_p()
check(lt.cublasLtMatmulPreferenceCreate(ctypes.byref(pref)), "prefCreate")
ws = ctypes.c_size_t(WORKSPACE_BYTES)
check(lt.cublasLtMatmulPreferenceSetAttribute(pref, PREF_MAX_WORKSPACE, ctypes.byref(ws), 8), "prefSet")
results = (HeuristicResult * MAX_ALGOS)()
found = ctypes.c_int(0)
status = lt.cublasLtMatmulAlgoGetHeuristic(handle, op_desc, layout_a, layout_b, layout_d, layout_d,
pref, MAX_ALGOS, results, ctypes.byref(found))
check(status, "algoGetHeuristic")
def run_lt(algo_ref) -> None:
check(
lt.cublasLtMatmul(handle, op_desc, ctypes.byref(alpha),
ctypes.c_void_p(a_rm.data_ptr()), layout_a,
ctypes.c_void_p(b_rm.data_ptr()), layout_b,
ctypes.byref(beta),
ctypes.c_void_p(d.data_ptr()), layout_d,
ctypes.c_void_p(d.data_ptr()), layout_d,
algo_ref,
ctypes.c_void_p(workspace.data_ptr()), ctypes.c_size_t(WORKSPACE_BYTES),
stream), "ltMatmul")
torch_ms = time_fn(lambda: case["torch_fn"](a_rm, b_rm), args.warmup, args.repeats)
best = None
parity_max_abs = None
for index in range(found.value):
if results[index].state != 0:
continue
algo_ref = ctypes.byref(results[index], 0)
try:
run_lt(algo_ref)
except RuntimeError:
continue
torch.cuda.synchronize()
diff = (d.float() - reference.float()).abs().max().item()
if diff > 0.25:
continue
lt_ms = time_fn(lambda: run_lt(algo_ref), args.warmup, args.repeats)
if best is None or lt_ms < best[0]:
best = (lt_ms, index)
parity_max_abs = diff
row = {
"case": case["case"],
"mnk": [m, n, k],
"occurrences": case["occurrences"],
"algos_returned": found.value,
"torch_ms": round(torch_ms, 4),
"best_lt_ms": round(best[0], 4) if best else None,
"best_algo_index": best[1] if best else None,
"parity_max_abs": parity_max_abs,
"lt_vs_torch_pct": round(100.0 * (best[0] - torch_ms) / torch_ms, 2) if best else None,
}
summary.append(row)
print("LT_SWEEP_CASE " + json.dumps(row, sort_keys=True), flush=True)
for obj, destroy in ((pref, lt.cublasLtMatmulPreferenceDestroy), (layout_a, lt.cublasLtMatrixLayoutDestroy),
(layout_b, lt.cublasLtMatrixLayoutDestroy), (layout_d, lt.cublasLtMatrixLayoutDestroy),
(op_desc, lt.cublasLtMatmulDescDestroy)):
destroy(obj)
weighted_torch = sum(r["torch_ms"] * r["occurrences"] for r in summary)
weighted_lt = sum((r["best_lt_ms"] or r["torch_ms"]) * r["occurrences"] for r in summary)
print(
"LT_SWEEP_RESULT " + json.dumps({
"batch": args.batch,
"weighted_torch_band_ms": round(weighted_torch, 2),
"weighted_best_lt_band_ms": round(weighted_lt, 2),
"band_delta_pct": round(100.0 * (weighted_lt - weighted_torch) / weighted_torch, 3),
"cases": summary,
}, sort_keys=True),
flush=True,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,224 @@
#!/usr/bin/env python3
"""Exact-shape LTX-2 BF16 GELU-epilogue microbenchmark for GB200."""
from __future__ import annotations
import argparse
import gc
import json
import statistics
from collections.abc import Callable
from typing import Any
import torch
import torch.nn.functional as F
import fused_dense_lib
BLOCKS = 48
DTYPE = torch.bfloat16
def _reference_forward(
x: torch.Tensor,
weight1: torch.Tensor,
bias1: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
pre_activation = F.linear(x, weight1, bias1)
return F.gelu(pre_activation, approximate="tanh"), pre_activation
def _reference_backward(
grad_output: torch.Tensor,
weight2: torch.Tensor,
pre_activation: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
grad_activation = F.linear(grad_output, weight2.t())
grad_pre_activation = torch.ops.aten.gelu_backward.default(
grad_activation,
pre_activation,
approximate="tanh",
)
return grad_pre_activation, grad_pre_activation.sum(dim=0)
reference_forward = torch.compile(_reference_forward, fullgraph=True, dynamic=False)
reference_backward = torch.compile(_reference_backward, fullgraph=True, dynamic=False)
def _time_ms(fn: Callable[[], Any], iterations: int) -> float:
torch.cuda.synchronize()
begin = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
begin.record()
for _ in range(iterations):
fn()
end.record()
end.synchronize()
return begin.elapsed_time(end) / iterations
def _peak_memory(fn: Callable[[], Any]) -> dict[str, float]:
gc.collect()
torch.cuda.empty_cache()
torch.cuda.synchronize()
base_allocated = torch.cuda.memory_allocated()
base_reserved = torch.cuda.memory_reserved()
torch.cuda.reset_peak_memory_stats()
outputs = fn()
torch.cuda.synchronize()
result = {
"allocated_delta_mib": (torch.cuda.max_memory_allocated() - base_allocated) / 2**20,
"reserved_delta_mib": (torch.cuda.max_memory_reserved() - base_reserved) / 2**20,
"live_output_delta_mib": (torch.cuda.memory_allocated() - base_allocated) / 2**20,
}
del outputs
gc.collect()
torch.cuda.empty_cache()
return result
def _assert_close(
name: str,
actual: torch.Tensor,
expected: torch.Tensor,
*,
atol: float,
) -> float:
if actual.dtype != DTYPE or expected.dtype != DTYPE:
raise AssertionError(f"{name}: expected BF16 tensors, got {actual.dtype} and {expected.dtype}")
torch.testing.assert_close(actual, expected, rtol=3e-3, atol=atol)
return float((actual - expected).abs().max())
def _run_shape(
m: int,
*,
warmup: int,
iterations: int,
rounds: int,
heuristic: int,
seed: int,
) -> dict[str, Any]:
d, h = 4096, 16384
torch.manual_seed(seed + m)
x = torch.empty((m, d), device="cuda", dtype=DTYPE).normal_(std=1.0)
weight1 = torch.empty((h, d), device="cuda", dtype=DTYPE).normal_(std=0.02)
bias1 = torch.empty((h,), device="cuda", dtype=DTYPE).normal_(std=0.01)
weight2 = torch.empty((d, h), device="cuda", dtype=DTYPE).normal_(std=0.02)
grad_output = torch.empty((m, d), device="cuda", dtype=DTYPE).normal_(std=1 / 32)
def reference() -> tuple[torch.Tensor, ...]:
activation, pre_activation = reference_forward(x, weight1, bias1)
grad_pre_activation, grad_bias1 = reference_backward(grad_output, weight2, pre_activation)
return activation, pre_activation, grad_pre_activation, grad_bias1
def fused() -> tuple[torch.Tensor, ...]:
activation, pre_activation = fused_dense_lib.linear_act_forward(
x,
weight1,
bias1,
True,
True,
heuristic,
)
grad_pre_activation, grad_bias1 = fused_dense_lib.bias_act_linear_dgrad_bgrad(
weight2,
grad_output,
pre_activation,
True,
heuristic,
)
return activation, pre_activation, grad_pre_activation, grad_bias1
reference_outputs = reference()
fused_outputs = fused()
names = ("activation", "pre_activation", "grad_pre_activation", "grad_bias1")
atols = (3e-2, 3e-2, 3e-2, 1.5e-1)
parity_max_abs = {
name: _assert_close(name, actual, expected, atol=atol)
for name, actual, expected, atol in zip(names, fused_outputs, reference_outputs, atols)
}
del reference_outputs, fused_outputs
gc.collect()
torch.cuda.empty_cache()
for _ in range(warmup):
reference()
fused()
torch.cuda.synchronize()
samples: dict[str, list[float]] = {"reference": [], "fused": []}
functions = {"reference": reference, "fused": fused}
for round_index in range(rounds):
order = ("reference", "fused") if round_index % 2 == 0 else ("fused", "reference")
for name in order:
samples[name].append(_time_ms(functions[name], iterations))
medians = {name: statistics.median(values) for name, values in samples.items()}
saving_ms = medians["reference"] - medians["fused"]
peak_memory = {name: _peak_memory(functions[name]) for name in ("reference", "fused")}
return {
"m": m,
"d": d,
"h": h,
"dtype": str(DTYPE),
"heuristic": heuristic,
"reference_median_ms": medians["reference"],
"fused_median_ms": medians["fused"],
"per_block_saving_ms": saving_ms,
"projected_48_block_saving_ms": BLOCKS * saving_ms,
"isolated_speedup_percent": 100 * saving_ms / medians["reference"],
"samples_ms": samples,
"parity_max_abs": parity_max_abs,
"peak_memory": peak_memory,
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--m", nargs="+", type=int, default=[4290, 12870])
parser.add_argument("--warmup", type=int, default=3)
parser.add_argument("--iterations", type=int, default=10)
parser.add_argument("--rounds", type=int, default=5)
parser.add_argument("--heuristic", type=int, default=0, choices=range(5))
parser.add_argument("--seed", type=int, default=0)
args = parser.parse_args()
if min(args.m) <= 0 or min(args.warmup, args.iterations, args.rounds) <= 0:
parser.error("shapes and measurement counts must be positive")
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required")
capability = torch.cuda.get_device_capability()
if capability[0] != 10:
raise RuntimeError(f"expected GB200/SM100, got compute capability {capability}")
torch.set_grad_enabled(False)
print(
"LTX2_FUSED_DENSE_GELU_ENV "
+ json.dumps(
{
"torch": torch.__version__,
"cuda": torch.version.cuda,
"device": torch.cuda.get_device_name(),
"capability": capability,
},
sort_keys=True,
),
flush=True,
)
for m in args.m:
result = _run_shape(
m,
warmup=args.warmup,
iterations=args.iterations,
rounds=args.rounds,
heuristic=args.heuristic,
seed=args.seed,
)
print("LTX2_FUSED_DENSE_GELU_RESULT " + json.dumps(result, sort_keys=True), flush=True)
gc.collect()
torch.cuda.empty_cache()
if __name__ == "__main__":
main()
@@ -0,0 +1,91 @@
import statistics
import time
import torch
import torch.nn.functional as F
DEVICE = torch.device("cuda:0")
DTYPE = torch.bfloat16
TOKENS = 4290
HIDDEN = 4096
FFN = 16384
class FFNBlock(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.fc1 = torch.nn.Linear(HIDDEN, FFN, bias=True, device=DEVICE, dtype=DTYPE)
self.fc2 = torch.nn.Linear(FFN, HIDDEN, bias=True, device=DEVICE, dtype=DTYPE)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.fc2(F.gelu(self.fc1(x), approximate="tanh"))
class Projections(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.layers = torch.nn.ModuleList([
torch.nn.Linear(HIDDEN, HIDDEN, bias=True, device=DEVICE, dtype=DTYPE)
for _ in range(4)
])
def forward(self, x: torch.Tensor) -> torch.Tensor:
return sum(layer(x) for layer in self.layers)
def measure(module: torch.nn.Module, batch: int, matmul_flops: float) -> tuple[float, float]:
x = torch.randn(batch, TOKENS, HIDDEN, device=DEVICE, dtype=DTYPE, requires_grad=True)
def step() -> None:
out = module(x)
out.backward(torch.ones_like(out))
x.grad = None
for parameter in module.parameters():
parameter.grad = None
for _ in range(5):
step()
torch.cuda.synchronize()
samples = []
for _ in range(12):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
step()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end))
median_ms = statistics.median(samples)
return median_ms, matmul_flops * batch / (median_ms * 1e9)
def main() -> None:
torch.manual_seed(0)
torch.cuda.set_device(DEVICE)
cases = (
("ffn", FFNBlock, 12.0 * TOKENS * HIDDEN * FFN),
("four_4096_projections", Projections, 24.0 * TOKENS * HIDDEN * HIDDEN),
)
for name, factory, flops in cases:
for compiled in (False, True):
results = []
for batch in (1, 2):
module = factory()
if compiled:
module = torch.compile(module, fullgraph=True)
median_ms, tflops = measure(module, batch, flops)
results.append((batch, median_ms, tflops))
del module
torch.cuda.empty_cache()
ratio = results[1][1] / results[0][1]
print(name, "compiled=" + str(compiled), results, "b2_over_b1=" + f"{ratio:.4f}")
if __name__ == "__main__":
started = time.time()
main()
print("wall_sec", time.time() - started)
@@ -0,0 +1,247 @@
#!/usr/bin/env python3
"""Four-GB200 LTX-2 packed BF16 DP/TP timing gate."""
import gc
import json
import os
import statistics
import torch
import torch.distributed as dist
LAYERS = (
("self_qkv", "column", "video", 4290, 4096, 12288),
("self_out", "row", "video", 4290, 4096, 4096),
("cross_q", "column", "video", 4290, 4096, 4096),
("cross_kv", "column", "text", 1024, 4096, 8192),
("cross_out", "row", "video", 4290, 4096, 4096),
("ffn_up", "column", "video", 4290, 4096, 16384),
("ffn_down", "row", "video", 4290, 16384, 4096),
)
BLOCK_COUNT = 48
CURRENT_STEP_MS = 403.725
TARGET_STEP_MS = 288.882
GB200_CAPACITY_GIB = 189471 / 1024
FIXED_ARENA_PEAK_GIB = 149.792064
FIXED_ARENA_STEADY_GIB = 97.234253
def emit(kind, **fields):
if dist.get_rank() == 0:
print(json.dumps({"kind": kind, **fields}, sort_keys=True), flush=True)
def timed(fn, sync_group, warmup=3, samples=7, inner=3):
for _ in range(warmup):
fn()
torch.cuda.synchronize()
dist.barrier(group=sync_group, device_ids=[torch.cuda.current_device()])
values = []
for _ in range(samples):
begin = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
begin.record()
for _ in range(inner):
fn()
end.record()
end.synchronize()
values.append(begin.elapsed_time(end) / inner)
return statistics.median(values)
def sanity(tp, group, tp_rank):
if tp == 1:
return 0.0
torch.manual_seed(7)
x = torch.randn(8, 16, device="cuda", dtype=torch.bfloat16)
w = torch.randn(12, 16, device="cuda", dtype=torch.bfloat16)
dy = torch.randn(8, 12, device="cuda", dtype=torch.bfloat16)
k_width, n_width = 16 // tp, 12 // tp
row_fwd = x[:, tp_rank * k_width:(tp_rank + 1) * k_width] @ w[:, tp_rank * k_width:(tp_rank + 1) * k_width].T
dist.all_reduce(row_fwd, group=group)
col_dx = dy[:, tp_rank * n_width:(tp_rank + 1) * n_width] @ w[tp_rank * n_width:(tp_rank + 1) * n_width]
dist.all_reduce(col_dx, group=group)
return max(float((row_fwd - x @ w.T).abs().max()), float((col_dx - dy @ w).abs().max()))
def case(topology, tp, batch, group, tp_rank, layer):
name, parallel, modality, base_m, k, n = layer
m = base_m * batch
local_k = k // tp if parallel == "row" else k
local_n = n // tp if parallel == "column" else n
gc.collect()
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
torch.manual_seed(20260721 + LAYERS.index(layer))
x = torch.empty(m, local_k, device="cuda", dtype=torch.bfloat16).normal_()
w = torch.empty(local_n, local_k, device="cuda", dtype=torch.bfloat16).normal_(0, 0.02)
dy = torch.empty(m, local_n, device="cuda", dtype=torch.bfloat16).normal_(0, 0.01)
y = torch.empty(m, local_n, device="cuda", dtype=torch.bfloat16)
dx = torch.empty(m, local_k, device="cuda", dtype=torch.bfloat16)
dw = torch.empty(local_n, local_k, device="cuda", dtype=torch.bfloat16)
collective_tensor = dx if parallel == "column" else y
zeros = torch.zeros_like(collective_tensor) if tp > 1 else None
def fwd():
return torch.mm(x, w.T, out=y)
def dgrad():
return torch.mm(dy, w, out=dx)
def wgrad():
return torch.mm(dy.T, x, out=dw)
def allreduce():
assert zeros is not None
dist.all_reduce(zeros, group=group)
def fwd_collective():
torch.mm(x, w.T, out=y)
if tp > 1 and parallel == "row":
dist.all_reduce(y, group=group)
return y
def dgrad_collective():
torch.mm(dy, w, out=dx)
if tp > 1 and parallel == "column":
dist.all_reduce(dx, group=group)
return dx
local = {
"fwd_ms": timed(fwd, group),
"dgrad_ms": timed(dgrad, group),
"wgrad_ms": timed(wgrad, group),
"allreduce_ms": timed(allreduce, group) if tp > 1 else 0.0,
"fwd_collective_ms": timed(fwd_collective, group) if tp > 1 else 0.0,
"dgrad_collective_ms": timed(dgrad_collective, group) if tp > 1 else 0.0,
"peak_allocated_gib": torch.cuda.max_memory_allocated() / 2**30,
}
gathered = [None] * dist.get_world_size()
dist.all_gather_object(gathered, local)
row = None
if dist.get_rank() == 0:
row = {
"topology": topology,
"name": name,
"count": BLOCK_COUNT,
"parallel": parallel,
"modality": modality,
"logical_shape": [m, k, n],
"local_gemm_shape": [m, local_k, local_n],
"allreduce_shape": list(collective_tensor.shape) if tp > 1 else None,
"allreduce_bytes": collective_tensor.numel() * 2 if tp > 1 else 0,
**{key: max(rank_row[key] for rank_row in gathered) for key in local},
}
if tp == 1:
row["fwd_collective_ms"] = row["fwd_ms"]
row["dgrad_collective_ms"] = row["dgrad_ms"]
print(json.dumps({"kind": "case", **row}, sort_keys=True), flush=True)
del x, w, dy, y, dx, dw, zeros
return row
def main():
rank = int(os.environ["RANK"])
local_rank = int(os.environ["LOCAL_RANK"])
world = int(os.environ["WORLD_SIZE"])
assert world == 4
torch.cuda.set_device(local_rank)
dist.init_process_group("nccl", device_id=torch.device("cuda", local_rank))
torch.set_grad_enabled(False)
torch.backends.cuda.matmul.allow_tf32 = False
tp2_groups = [dist.new_group([0, 1]), dist.new_group([2, 3])]
topologies = (
("dp4_b1", 1, 1, dist.group.WORLD, 0),
("tp2_dp2_b2", 2, 2, tp2_groups[rank // 2], rank % 2),
("tp4_b4", 4, 4, dist.group.WORLD, rank),
)
emit(
"environment",
torch=torch.__version__,
cuda=torch.version.cuda,
gpu=torch.cuda.get_device_name(local_rank),
capability=torch.cuda.get_device_capability(local_rank),
world_size=world,
warmup=3,
samples=7,
inner=3,
matmul_tf32=torch.backends.cuda.matmul.allow_tf32,
)
summaries = {}
for topology, tp, batch, group, tp_rank in topologies:
errors = [None] * world
dist.all_gather_object(errors, sanity(tp, group, tp_rank))
emit("sanity", topology=topology, max_abs=max(errors))
rows = [case(topology, tp, batch, group, tp_rank, layer) for layer in LAYERS]
if rank != 0:
continue
compute_ms = sum((r["fwd_ms"] + r["dgrad_ms"] + r["wgrad_ms"]) * r["count"] for r in rows)
allreduce_ms = sum(r["allreduce_ms"] * r["count"] for r in rows)
end_to_end_ms = sum(
((r["fwd_ms"] if r["parallel"] == "column" else r["fwd_collective_ms"])
+ (r["dgrad_collective_ms"] if r["parallel"] == "column" else r["dgrad_ms"])
+ r["wgrad_ms"]) * r["count"] for r in rows)
payload_bytes = sum(r["allreduce_bytes"] * r["count"] for r in rows)
ring_wire_bytes = payload_bytes * (2 * (tp - 1) / tp) if tp > 1 else 0
flops = sum(3 * 2 * (r["logical_shape"][0] * r["logical_shape"][1] * r["logical_shape"][2] // tp) * r["count"] for r in rows)
summaries[topology] = {
"tp": tp,
"dp": world // tp,
"local_batch": batch,
"compute_ms": compute_ms,
"allreduce_only_ms": allreduce_ms,
"ideal_overlap_lower_bound_ms": max(compute_ms, allreduce_ms),
"sequential_compute_collective_ms": end_to_end_ms,
"logical_flops_per_rank": flops,
"effective_compute_tflops": flops / compute_ms / 1e9,
"allreduce_payload_gib": payload_bytes / 2**30,
"ring_wire_lower_bound_gib_per_rank": ring_wire_bytes / 2**30,
"max_microbench_allocated_gib": max(r["peak_allocated_gib"] for r in rows),
}
emit("topology", topology=topology, **summaries[topology])
if rank == 0:
baseline = summaries["dp4_b1"]["compute_ms"]
required_segment = baseline - (CURRENT_STEP_MS - TARGET_STEP_MS)
for topology in ("tp2_dp2_b2", "tp4_b4"):
row = summaries[topology]
row["required_projection_segment_ms"] = required_segment
row["projected_step_compute_only_ms"] = CURRENT_STEP_MS - baseline + row["compute_ms"]
row["projected_step_ideal_overlap_ms"] = CURRENT_STEP_MS - baseline + row["ideal_overlap_lower_bound_ms"]
row["projected_step_sequential_ms"] = CURRENT_STEP_MS - baseline + row["sequential_compute_collective_ms"]
row["can_clear_target_compute_only"] = row["projected_step_compute_only_ms"] <= TARGET_STEP_MS
row["can_clear_target_ideal_overlap"] = row["projected_step_ideal_overlap_ms"] <= TARGET_STEP_MS
state_floor = {
"dp4_b1": 24.292 * 2 + 36.438,
"tp2_dp2_b2": (24.292 / 2) * 2 + 36.438,
"tp4_b4": (24.292 / 4) * 2 + 36.438,
}
emit(
"gate",
current_step_ms=CURRENT_STEP_MS,
target_step_ms=TARGET_STEP_MS,
saving_required_ms=CURRENT_STEP_MS - TARGET_STEP_MS,
dp4_projection_baseline_ms=baseline,
allreduces_per_block={"video": 6, "text": 1},
gb200_capacity_gib=GB200_CAPACITY_GIB,
fixed_arena_peak_gib=FIXED_ARENA_PEAK_GIB,
fixed_arena_steady_gib=FIXED_ARENA_STEADY_GIB,
state_floor_gib_per_rank=state_floor,
optimistic_peak_if_non_state_unchanged_gib={
name: FIXED_ARENA_PEAK_GIB - state_floor["dp4_b1"] + floor
for name, floor in state_floor.items()
},
conservative_peak_if_all_non_state_scales_with_batch_gib={
name: floor + (FIXED_ARENA_PEAK_GIB - state_floor["dp4_b1"]) * summaries[name]["local_batch"]
for name, floor in state_floor.items()
},
summaries=summaries,
)
dist.destroy_process_group()
if __name__ == "__main__":
main()
@@ -0,0 +1,105 @@
#!/usr/bin/env python3
"""Occurrence-weighted packed LTX-2 GEMM band probe for TunableOp / env gates.
Runs the five packed projection cases as real ``torch.nn.Linear`` fwd+bwd so
cuBLAS sees the production orientations (fwd addmm, dgrad NN, wgrad TN with
bias grad), weights each case by its per-step occurrence count (48 blocks,
video_dd three per block), and prints a per-step GEMM band estimate. Compare
one process without TunableOp against a tune-then-replay pair with
``PYTORCH_TUNABLEOP_ENABLED=1``. Never MFU evidence.
"""
from __future__ import annotations
import argparse
import json
import statistics
import torch
import torch.nn.functional as F
HIDDEN = 4096
FFN = 16384
VIDEO_TOKENS = 11 * 15 * 26
TEXT_TOKENS = 1024
def build_cases(batch: int) -> list[dict]:
video_rows = VIDEO_TOKENS * batch
text_rows = TEXT_TOKENS * batch
return [
{"name": "self_qkv", "rows": video_rows, "in": HIDDEN, "out": 3 * HIDDEN, "occurrences": 48},
{"name": "video_dd", "rows": video_rows, "in": HIDDEN, "out": HIDDEN, "occurrences": 144},
{"name": "text_kv", "rows": text_rows, "in": HIDDEN, "out": 2 * HIDDEN, "occurrences": 48},
{"name": "ffn_up", "rows": video_rows, "in": HIDDEN, "out": FFN, "occurrences": 48},
{"name": "ffn_down", "rows": video_rows, "in": FFN, "out": HIDDEN, "occurrences": 48},
]
def time_case(case: dict, warmup: int, repeats: int) -> dict:
device = torch.device("cuda:0")
dtype = torch.bfloat16
layer = torch.nn.Linear(case["in"], case["out"], bias=True, device=device, dtype=dtype)
x = torch.randn(case["rows"], case["in"], device=device, dtype=dtype, requires_grad=True)
grad_out = torch.randn(case["rows"], case["out"], device=device, dtype=dtype)
def step() -> None:
x.grad = None
layer.weight.grad = None
layer.bias.grad = None
F.linear(x, layer.weight, layer.bias).backward(grad_out)
for _ in range(warmup):
step()
torch.cuda.synchronize()
starts = [torch.cuda.Event(enable_timing=True) for _ in range(repeats)]
ends = [torch.cuda.Event(enable_timing=True) for _ in range(repeats)]
for start, end in zip(starts, ends, strict=True):
start.record()
step()
end.record()
torch.cuda.synchronize()
values = [float(start.elapsed_time(end)) for start, end in zip(starts, ends, strict=True)]
median_ms = statistics.median(values)
flops = 6.0 * case["rows"] * case["in"] * case["out"]
return {
"case": case["name"],
"rows": case["rows"],
"median_ms": median_ms,
"min_ms": min(values),
"occurrences": case["occurrences"],
"weighted_ms": median_ms * case["occurrences"],
"tflops": flops / (median_ms * 1e9),
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--batch", type=int, default=2)
parser.add_argument("--warmup", type=int, default=10)
parser.add_argument("--repeats", type=int, default=31)
parser.add_argument("--label", default="default")
args = parser.parse_args()
torch.manual_seed(20260722)
rows = [time_case(case, args.warmup, args.repeats) for case in build_cases(args.batch)]
total_weighted = sum(row["weighted_ms"] for row in rows)
total_flops = sum(
6.0 * case["rows"] * case["in"] * case["out"] * case["occurrences"]
for case in build_cases(args.batch))
print(
"GEMM_BAND " + json.dumps({
"label": args.label,
"batch": args.batch,
"tunableop_enabled": torch.cuda.tunable.is_enabled(),
"tunableop_tuning": torch.cuda.tunable.tuning_is_enabled(),
"per_step_gemm_band_ms": round(total_weighted, 3),
"band_tflops": round(total_flops / (total_weighted * 1e9), 1),
"cases": rows,
}, sort_keys=True),
flush=True,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,89 @@
#!/usr/bin/env python3
"""Compare two-rank NCCL all-reduce within and across GB200 trays."""
from __future__ import annotations
import json
import os
import socket
import time
import torch
import torch.distributed as dist
def benchmark(label: str, ranks: list[int], group: dist.ProcessGroup) -> None:
rank = dist.get_rank()
if rank in ranks:
sizes_and_iterations = [
(1 << 20, 200),
(16 << 20, 100),
(256 << 20, 40),
(1 << 30, 20),
]
for size_bytes, iterations in sizes_and_iterations:
tensor = torch.ones(size_bytes // 4, device="cuda", dtype=torch.float32)
for _ in range(20):
dist.all_reduce(tensor, group=group)
torch.cuda.synchronize()
dist.barrier(group=group)
start = time.perf_counter()
for _ in range(iterations):
dist.all_reduce(tensor, group=group)
torch.cuda.synchronize()
elapsed = (time.perf_counter() - start) / iterations
slowest = torch.tensor(elapsed, device="cuda", dtype=torch.float64)
dist.all_reduce(slowest, op=dist.ReduceOp.MAX, group=group)
if rank == ranks[0]:
seconds = float(slowest.item())
print(
"NVLINK_RESULT "
+ json.dumps(
{
"algorithmic_gbps": size_bytes / seconds / 1e9,
"label": label,
"latency_us": seconds * 1e6,
"ranks": ranks,
"size_bytes": size_bytes,
},
sort_keys=True,
),
flush=True,
)
del tensor, slowest
dist.barrier()
def main() -> None:
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
dist.init_process_group("nccl", device_id=torch.device("cuda", local_rank))
if dist.get_world_size() != 8:
raise RuntimeError("expected exactly eight ranks across two four-GPU trays")
intra = dist.new_group([0, 1], backend="nccl")
inter = dist.new_group([0, 4], backend="nccl")
if dist.get_rank() in (0, 4):
print(
"NVLINK_TOPOLOGY "
+ json.dumps(
{
"hostname": socket.gethostname(),
"local_rank": local_rank,
"rank": dist.get_rank(),
},
sort_keys=True,
),
flush=True,
)
dist.barrier()
benchmark("intra_tray", [0, 1], intra)
benchmark("inter_tray", [0, 4], inter)
dist.destroy_process_group()
if __name__ == "__main__":
main()
@@ -0,0 +1,365 @@
#!/usr/bin/env python3
"""Gate QuACK's existing SM100 NVFP4 kernel on the exact LTX-2 packs.
The optimized total keeps forward and dgrad sequential, but batches each
same-shaped wgrad across all 48 blocks. It is an arithmetic lower bound:
operands are prequantized and bias/global postscale/cache work is excluded.
"""
import argparse
import gc
import importlib.metadata
import json
import statistics
# name, occurrences in 48 blocks, logical M, K, N
PACKS = (
("self_qkv", 48, 4290, 4096, 12288),
("video_dd", 144, 4290, 4096, 4096),
("text_kv", 48, 1024, 4096, 8192),
("ffn_up", 48, 4290, 4096, 16384),
("ffn_down", 48, 4290, 16384, 4096),
)
CONFIGS = (
((128, 64), (1, 1)),
((128, 128), (1, 1)),
((128, 192), (1, 1)),
((128, 256), (1, 1)),
((256, 64), (2, 1)),
((256, 128), (2, 1)),
((256, 192), (2, 1)),
((256, 256), (2, 1)),
)
GATE_MS = 63.463
MARGIN_MS = 59.524
BREAK_EVEN_SPEEDUP = 3.095
MARGIN_SPEEDUP = 3.3
def emit(kind, **fields):
print(json.dumps({"kind": kind, **fields}, sort_keys=True), flush=True)
def phase_shape(pack, phase):
_, _, m, k, n = pack
if phase == "fwd":
return m, k, n
if phase == "dgrad":
return m, n, k
if phase == "wgrad":
return n, (m + 31) // 32 * 32, k
raise ValueError(phase)
def self_check():
assert sum(pack[1] for pack in PACKS) == 7 * 48
one_pass_flops = sum(2 * m * k * n * count for _, count, m, k, n in PACKS)
assert one_pass_flops == 100031935807488
assert len({phase_shape(pack, phase) for pack in PACKS for phase in ("fwd", "dgrad", "wgrad")}) == 12
def timed(torch, fn, warmup, samples, inner):
for _ in range(warmup):
fn()
torch.cuda.synchronize()
values = []
for _ in range(samples):
begin, end = torch.cuda.Event(True), torch.cuda.Event(True)
begin.record()
for _ in range(inner):
fn()
end.record()
end.synchronize()
values.append(begin.elapsed_time(end) / inner)
return statistics.median(values)
def tensors(torch, m, k, n, batch):
# Physical FP4 storage is K/2; permuting contiguous (L,M,K/2) makes K
# unit-stride while retaining the logical batch mode expected by CuTe.
a_store = torch.empty((batch, m, k // 2), device="cuda", dtype=torch.float4_e2m1fn_x2)
b_store = torch.empty((batch, n, k // 2), device="cuda", dtype=torch.float4_e2m1fn_x2)
a_store.view(torch.uint8).random_(0, 256)
b_store.view(torch.uint8).random_(0, 256)
a = a_store.permute(1, 2, 0)
b = b_store.permute(1, 2, 0)
out_store = torch.empty((batch, m, n), device="cuda", dtype=torch.bfloat16)
out = out_store.permute(1, 2, 0)
sf_k = k // 16
sfa = torch.ones((batch, (m + 127) // 128, (sf_k + 3) // 4, 512),
device="cuda", dtype=torch.float8_e4m3fn)
sfb = torch.ones((batch, (n + 127) // 128, (sf_k + 3) // 4, 512),
device="cuda", dtype=torch.float8_e4m3fn)
return a, b, out, out_store, sfa, sfb
def bench_bf16_shape(torch, shape, args):
m, k, n = shape
a = torch.randn((m, k), device="cuda", dtype=torch.bfloat16)
b = torch.randn((n, k), device="cuda", dtype=torch.bfloat16)
out = torch.empty((m, n), device="cuda", dtype=torch.bfloat16)
latency = timed(torch, lambda: torch.mm(a, b.T, out=out), args.warmup,
args.samples, args.inner)
emit("bf16", shape=shape, median_ms=latency)
del a, b, out
gc.collect()
torch.cuda.empty_cache()
return latency
def bench_shape(torch, cutlass, compile_gemm, shape, batch, configs, args):
m, k, n = shape
a, b, out, out_store, sfa, sfb = tensors(torch, m, k, n, batch)
rows = []
for index in configs:
tile, cluster = CONFIGS[index]
try:
run = compile_gemm(
cutlass.Float4E2M1FN,
cutlass.Float8E4M3FN,
16,
cutlass.BFloat16,
tile,
cluster,
a,
b,
out,
sfa,
sfb,
)
latency = timed(torch, lambda: run(a, b, out, sfa, sfb), args.warmup,
args.samples, args.inner if batch == 1 else 1)
finite = bool(torch.isfinite(out_store.flatten()[:1 << 20]).all())
if not finite:
raise RuntimeError("non-finite output sample")
rows.append((latency, index))
emit("case", shape=shape, batch=batch, config=index, tile=tile,
cluster=cluster, median_ms=latency, finite_sample=finite)
except Exception as exc:
torch.cuda.synchronize()
emit("failure", shape=shape, batch=batch, config=index, exception=repr(exc))
if not rows:
raise RuntimeError(f"all configurations failed for shape={shape}, batch={batch}")
latency, index = min(rows)
emit("best", shape=shape, batch=batch, config=index, median_ms=latency,
finite_sample=True)
del a, b, out, out_store, sfa, sfb
gc.collect()
torch.cuda.empty_cache()
return latency
def correctness_smoke(torch, cutlass, compile_gemm):
"""Prove packed FP4, BF16 output, and L>1 agree on a tiny exact case."""
batch = 2
m = n = k = 256
a_store = torch.empty((batch, m, k // 2), device="cuda",
dtype=torch.float4_e2m1fn_x2)
b_store = torch.empty((batch, n, k // 2), device="cuda",
dtype=torch.float4_e2m1fn_x2)
# E2M1 code 0b0010 is exactly 1.0; each byte stores two values.
a_store.view(torch.uint8).fill_(0x22)
b_store.view(torch.uint8).fill_(0x22)
a = a_store.permute(1, 2, 0)
b = b_store.permute(1, 2, 0)
out_store = torch.empty((batch, m, n), device="cuda", dtype=torch.bfloat16)
out = out_store.permute(1, 2, 0)
scales = torch.ones((batch, 2, 4, 512), device="cuda",
dtype=torch.float8_e4m3fn)
run = compile_gemm(
cutlass.Float4E2M1FN,
cutlass.Float8E4M3FN,
16,
cutlass.BFloat16,
(256, 128),
(2, 1),
a,
b,
out,
scales,
scales,
)
run(a, b, out, scales, scales)
torch.cuda.synchronize()
max_abs = float((out_store.float() - k).abs().max())
if max_abs != 0.0:
raise RuntimeError(f"packed NVFP4 batched smoke failed: max_abs={max_abs}")
emit("correctness", shape=(m, k, n), batch=batch, expected=float(k),
max_abs=max_abs)
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--configs", type=int, nargs="+", default=list(range(len(CONFIGS))))
parser.add_argument("--warmup", type=int, default=3)
parser.add_argument("--samples", type=int, default=5)
parser.add_argument("--inner", type=int, default=3)
parser.add_argument("--dry-run", action="store_true")
args = parser.parse_args()
self_check()
if args.dry_run:
emit("contract", packs=PACKS, configs=[CONFIGS[i] for i in args.configs],
gate_ms=GATE_MS, margin_ms=MARGIN_MS)
return
import cutlass
import cutlass.cute as cute
import torch
# QuACK 0.5 uses pre-4.6 annotation aliases only; the runtime classes are
# unchanged. Restore those names before importing its Python modules.
if not hasattr(cute.core, "ThrMma"):
cute.core.ThrMma = cute.ThrMma
if not hasattr(cute.core, "ThrCopy"):
cute.core.ThrCopy = cute.ThrCopy
try:
import quack.blockscaled_gemm_utils as blockscaled_gemm
compile_blockscaled_gemm_tvm_ffi = blockscaled_gemm.compile_blockscaled_gemm_tvm_ffi
quack_source = blockscaled_gemm.__file__
quack_api = "blockscaled_gemm_utils"
except ModuleNotFoundError:
from functools import partial
import cutlass.cute as cute
from quack.compile_utils import make_fake_tensor as fake_tensor
from quack.cute_dsl_utils import get_device_capacity, get_max_active_clusters
from quack.gemm_sm100 import GemmSm100
from quack.gemm_tvm_ffi_utils import div_for_dtype, make_scheduler_args
from quack.rounding import RoundingMode
from quack.varlen_utils import VarlenArguments
if args.configs != [7]:
raise RuntimeError(
"QuACK >=0.6 large-shape default is config 7; "
"rerun with --configs 7")
def leading_dim(tensor):
return next(i for i, stride in enumerate(tensor.stride()) if stride == 1)
def fake_compact(tensor, dtype):
logical_shape = list(tensor.shape)
ld = leading_dim(tensor)
if dtype == cutlass.Float4E2M1FN:
logical_shape[ld] *= 2
return fake_tensor(
dtype,
tuple(logical_shape),
leading_dim=ld,
divisibility=div_for_dtype(dtype),
)
def compile_tensor_like(tensor, dtype):
compile_tensor = cute.runtime.from_dlpack(tensor)
compile_tensor.element_type = dtype
marked = compile_tensor.mark_layout_dynamic(leading_dim=leading_dim(tensor))
return compile_tensor if marked is None else marked
def compile_blockscaled_gemm_tvm_ffi(
ab_dtype, sf_dtype, sf_vec_size, d_dtype, tile, cluster,
a, b, out, sfa, sfb):
if get_device_capacity(a.device)[0] != 10:
raise RuntimeError("SM100 NVFP4 benchmark requires SM100")
# QuACK 0.6.1's wheel omits blockscaled/utils.py, while importing
# its public interface pulls in an SM90 epilogue that is
# incompatible with the installed CUTLASS. Compile the same
# native GemmSm100 kernel directly with the base identity
# epilogue; no package or environment mutation is required.
gemm = partial(
GemmSm100,
sf_vec_size=sf_vec_size,
use_clc_persistence=True,
)(cutlass.Float32, ab_dtype, tile, (*cluster, 1))
gemm.rounding_mode = RoundingMode.RN
compile_epi_args = gemm.EpilogueArguments()
scheduler_args = make_scheduler_args(
get_max_active_clusters(cluster[0] * cluster[1]),
max_swizzle_size=8,
tile_count_semaphore=None,
batch_idx_permute=None,
)
stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True)
@cute.jit
def runner(a_, b_, out_, sfa_, sfb_, varlen_args_, stream_):
gemm(
a_, b_, out_, None, compile_epi_args, scheduler_args,
varlen_args_, stream_, sfa_, sfb_)
compiled = cute.compile(
runner,
fake_compact(a, ab_dtype),
fake_compact(b, ab_dtype),
fake_compact(out, d_dtype),
compile_tensor_like(sfa, sf_dtype),
compile_tensor_like(sfb, sf_dtype),
VarlenArguments(),
stream,
options="--enable-tvm-ffi",
)
def run(a_, b_, out_, sfa_, sfb_):
compiled(a_, b_, out_, sfa_, sfb_, VarlenArguments())
return run
quack_source = importlib.import_module("quack.gemm_sm100").__file__
quack_api = "gemm_sm100_direct_identity_epilogue"
assert torch.cuda.get_device_capability() == (10, 0)
assert all(0 <= index < len(CONFIGS) for index in args.configs)
torch.manual_seed(20260721)
emit("environment", gpu=torch.cuda.get_device_name(), torch=torch.__version__,
torch_source=torch.__file__, cutlass_source=cutlass.__file__,
quack=importlib.metadata.version("quack-kernels"),
quack_source=quack_source, quack_api=quack_api, configs=args.configs)
correctness_smoke(torch, cutlass, compile_blockscaled_gemm_tvm_ffi)
single = {}
bf16 = {}
for pack in PACKS:
for phase in ("fwd", "dgrad", "wgrad"):
shape = phase_shape(pack, phase)
if shape not in single:
bf16[shape] = bench_bf16_shape(torch, shape, args)
single[shape] = bench_shape(torch, cutlass, compile_blockscaled_gemm_tvm_ffi,
shape, 1, args.configs, args)
batched_wgrad = {}
for pack in PACKS:
name, count, *_ = pack
batched_wgrad[name] = bench_shape(
torch, cutlass, compile_blockscaled_gemm_tvm_ffi,
phase_shape(pack, "wgrad"), count, args.configs, args)
phases = {phase: sum(single[phase_shape(pack, phase)] * pack[1] for pack in PACKS)
for phase in ("fwd", "dgrad", "wgrad")}
bf16_phases = {phase: sum(bf16[phase_shape(pack, phase)] * pack[1] for pack in PACKS)
for phase in ("fwd", "dgrad", "wgrad")}
bf16_ms = sum(bf16_phases.values())
unbatched_ms = sum(phases.values())
deferred_wgrad_ms = sum(batched_wgrad.values())
deferred_optimized_ms = phases["fwd"] + phases["dgrad"] + deferred_wgrad_ms
emit("aggregate", bf16_ms=bf16_ms, bf16_phases_ms=bf16_phases,
unbatched_ms=unbatched_ms, phases_ms=phases,
deferred_batched_wgrad_ms=deferred_wgrad_ms,
deferred_optimized_ms=deferred_optimized_ms,
gate_ms=GATE_MS, margin_ms=MARGIN_MS,
break_even_speedup=BREAK_EVEN_SPEEDUP,
margin_speedup=MARGIN_SPEEDUP,
unbatched_speedup=bf16_ms / unbatched_ms,
deferred_speedup=bf16_ms / deferred_optimized_ms,
unbatched_ratio_break_even=bf16_ms / unbatched_ms >= BREAK_EVEN_SPEEDUP,
unbatched_ratio_margin=bf16_ms / unbatched_ms >= MARGIN_SPEEDUP,
deferred_ratio_break_even=bf16_ms / deferred_optimized_ms >= BREAK_EVEN_SPEEDUP,
deferred_ratio_margin=bf16_ms / deferred_optimized_ms >= MARGIN_SPEEDUP,
unbatched_passes_break_even=unbatched_ms <= GATE_MS,
unbatched_passes_margin=unbatched_ms <= MARGIN_MS,
deferred_passes_break_even=deferred_optimized_ms <= GATE_MS,
deferred_passes_margin=deferred_optimized_ms <= MARGIN_MS,
unbatched_effective_tflops=300095807422464 / unbatched_ms / 1e9,
deferred_effective_tflops=300095807422464 / deferred_optimized_ms / 1e9)
if __name__ == "__main__":
main()
@@ -0,0 +1,384 @@
#!/usr/bin/env python3
"""Count a complete QuACK NVFP4 projection step on the exact LTX-2 packs.
This is a scratch gate, not a FastVideo integration. It keeps FP32 master
weights resident, derives BF16 working weights once outside the timed region,
then counts exact-current quantization, both weight-cache orientations,
forward/dgrad/wgrad, output postscales, forward bias, and bias gradient.
"""
import argparse
import gc
import importlib.metadata
import json
import statistics
# name, occurrences in 48 blocks, logical M, K, N
PACKS = (
("self_qkv", 48, 4290, 4096, 12288),
("video_dd", 144, 4290, 4096, 4096),
("text_kv", 48, 1024, 4096, 8192),
("ffn_up", 48, 4290, 4096, 16384),
("ffn_down", 48, 4290, 16384, 4096),
)
CONFIGS = (
((128, 64), (1, 1)),
((128, 128), (1, 1)),
((128, 192), (1, 1)),
((128, 256), (1, 1)),
((256, 64), (2, 1)),
((256, 128), (2, 1)),
((256, 192), (2, 1)),
((256, 256), (2, 1)),
)
# Best config for each exact shape in the paired QuACK 0.5 config 5/6/7 gate.
BEST_CONFIG = {
(1024, 4096, 8192): 7,
(1024, 8192, 4096): 5,
(4096, 4320, 4096): 7,
(4096, 4320, 16384): 7,
(4290, 4096, 4096): 6,
(4290, 4096, 12288): 6,
(4290, 4096, 16384): 6,
(4290, 12288, 4096): 6,
(4290, 16384, 4096): 6,
(8192, 1024, 4096): 7,
(12288, 4320, 4096): 7,
(16384, 4320, 4096): 7,
}
HISTORICAL_BF16_MS = 196.431
GATE_MS = 63.463
MARGIN_MS = 59.524
def emit(kind, **fields):
print(json.dumps({"kind": kind, **fields}, sort_keys=True), flush=True)
def phase_shape(pack, phase):
_, _, m, k, n = pack
if phase == "fwd":
return m, k, n
if phase == "dgrad":
return m, n, k
if phase == "wgrad":
return n, (m + 31) // 32 * 32, k
raise ValueError(phase)
def self_check():
assert sum(pack[1] for pack in PACKS) == 7 * 48
assert sum(2 * m * k * n * count for _, count, m, k, n in PACKS) == 100031935807488
shapes = {phase_shape(pack, phase) for pack in PACKS for phase in ("fwd", "dgrad", "wgrad")}
assert shapes == set(BEST_CONFIG)
assert all(shape[1] % 16 == 0 for shape in shapes)
def timed(torch, fn, warmup, samples, inner):
for _ in range(warmup):
fn()
torch.cuda.synchronize()
values = []
for _ in range(samples):
begin, end = torch.cuda.Event(True), torch.cuda.Event(True)
begin.record()
for _ in range(inner):
fn()
end.record()
end.synchronize()
values.append(begin.elapsed_time(end) / inner)
return statistics.median(values)
def sampled_metrics(torch, actual, reference, limit=1 << 20):
actual = actual.flatten()[:limit].float()
reference = reference.flatten()[:limit].float()
denom = reference.square().mean().sqrt().clamp_min(1e-12)
return {
"finite": bool(torch.isfinite(actual).all()),
"relative_rms": float((actual - reference).square().mean().sqrt() / denom),
"cosine": float(torch.nn.functional.cosine_similarity(actual, reference, dim=0)),
"norm_ratio": float(actual.norm() / reference.norm().clamp_min(1e-12)),
"sampled_values": actual.numel(),
}
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--pack", choices=[pack[0] for pack in PACKS], action="append")
parser.add_argument("--warmup", type=int, default=3)
parser.add_argument("--samples", type=int, default=5)
parser.add_argument("--inner", type=int, default=1)
parser.add_argument("--smoke-only", action="store_true")
parser.add_argument("--dry-run", action="store_true")
args = parser.parse_args()
self_check()
selected = [pack for pack in PACKS if not args.pack or pack[0] in args.pack]
if args.dry_run:
emit(
"contract",
packs=selected,
best_config=[{"shape": shape, "config": config} for shape, config in BEST_CONFIG.items()],
precision={
"master_weights": "fp32",
"working_weights_activations_gradients": "bf16",
"optimizer_moments": "fp32 (outside projection timing)",
"gemm_accumulator": "fp32",
"outputs": "bf16",
},
counted=(
"x/dY exact amax and row/transpose NVFP4 quantization; BF16 transpose materialization; "
"one weight amax plus row/column cache quantization; fwd+dgrad+wgrad; "
"global postscale; forward bias; bias gradient"
),
)
return
import cutlass
import cutlass.cute as cute
import torch
# QuACK 0.5 targets the older location of these two CUTLASS DSL types.
cute.core.ThrMma = cute.ThrMma
cute.core.ThrCopy = cute.ThrCopy
import quack.blockscaled_gemm_utils as blockscaled
assert torch.cuda.get_device_capability() == (10, 0)
torch.manual_seed(20260721)
def exact_per_tensor_scale(x):
return x.float().abs().nan_to_num().amax().clamp_min(1e-12) / (448.0 * 6.0)
def quantize_and_pack(x, per_tensor_scale):
q_u8, scale_2d, _ = blockscaled.to_nvfp4_compiled(x, 16, per_tensor_scale)
return q_u8, blockscaled.pack_scale_2d_to_blocked_contig(scale_2d)
def forward_epilogue(out, alpha, bias):
out.mul_(alpha)
return out.add_(bias)
def scale_epilogue(out, alpha):
return out.mul_(alpha)
exact_per_tensor_scale = torch.compile(exact_per_tensor_scale, dynamic=True)
forward_epilogue = torch.compile(forward_epilogue, dynamic=True)
scale_epilogue = torch.compile(scale_epilogue, dynamic=True)
def quantized_operand(x, per_tensor_scale):
q_u8, scales = quantize_and_pack(x, per_tensor_scale)
rows, packed_k = q_u8.shape
operand = q_u8.view(1, rows, packed_k).permute(1, 2, 0).view(torch.float4_e2m1fn_x2)
return operand, scales
def compile_gemm(shape, a, b, out, sfa, sfb):
config = BEST_CONFIG[shape]
tile, cluster = CONFIGS[config]
run = blockscaled.compile_blockscaled_gemm_tvm_ffi(
cutlass.Float4E2M1FN,
cutlass.Float8E4M3FN,
16,
cutlass.BFloat16,
tile,
cluster,
a,
b,
out,
sfa,
sfb,
)
return run, config
emit(
"environment",
gpu=torch.cuda.get_device_name(),
torch=torch.__version__,
torch_source=torch.__file__,
cutlass_source=cutlass.__file__,
quack=importlib.metadata.version("quack-kernels"),
quack_source=blockscaled.__file__,
source_commit="7f139e2b28610063d2f30526ba8f0ccae5d88944",
)
rows = []
for pack in selected:
name, count, m, k, n = pack
padded_m = (m + 31) // 32 * 32
# The masters remain resident and untouched. Optimizer moments are
# deliberately outside this projection-only gate.
weight_master = torch.randn((n, k), device="cuda", dtype=torch.float32) * 0.02
bias_master = torch.randn((n,), device="cuda", dtype=torch.float32) * 0.02
weight = weight_master.to(torch.bfloat16)
bias = bias_master.to(torch.bfloat16)
x = torch.randn((m, k), device="cuda", dtype=torch.bfloat16)
dy = torch.randn((m, n), device="cuda", dtype=torch.bfloat16) * 0.01
fp4_fwd_store = torch.empty((1, m, n), device="cuda", dtype=torch.bfloat16)
fp4_dx_store = torch.empty((1, m, k), device="cuda", dtype=torch.bfloat16)
fp4_dw_store = torch.empty((1, n, k), device="cuda", dtype=torch.bfloat16)
fp4_dbias = torch.empty((n,), device="cuda", dtype=torch.bfloat16)
bf16_fwd = torch.empty((m, n), device="cuda", dtype=torch.bfloat16)
bf16_dx = torch.empty((m, k), device="cuda", dtype=torch.bfloat16)
bf16_dw = torch.empty((n, k), device="cuda", dtype=torch.bfloat16)
bf16_dbias = torch.empty((n,), device="cuda", dtype=torch.bfloat16)
# Materialize one set only to compile the three exact-shape QMMs.
weight_scale = exact_per_tensor_scale(weight)
qw_row, sw_row = quantized_operand(weight, weight_scale)
qw_col, sw_col = quantized_operand(weight.T.contiguous(), weight_scale)
x_scale = exact_per_tensor_scale(x)
qx, sx = quantized_operand(x, x_scale)
dy_scale = exact_per_tensor_scale(dy)
qdy, sdy = quantized_operand(dy, dy_scale)
x_t = torch.nn.functional.pad(x.T, (0, padded_m - m)).contiguous()
dy_t = torch.nn.functional.pad(dy.T, (0, padded_m - m)).contiguous()
qx_t, sx_t = quantized_operand(x_t, x_scale)
qdy_t, sdy_t = quantized_operand(dy_t, dy_scale)
run_fwd, fwd_config = compile_gemm(
phase_shape(pack, "fwd"), qx, qw_row, fp4_fwd_store.permute(1, 2, 0), sx, sw_row)
run_dgrad, dgrad_config = compile_gemm(
phase_shape(pack, "dgrad"), qdy, qw_col, fp4_dx_store.permute(1, 2, 0), sdy, sw_col)
run_wgrad, wgrad_config = compile_gemm(
phase_shape(pack, "wgrad"), qdy_t, qx_t, fp4_dw_store.permute(1, 2, 0), sdy_t, sx_t)
del qw_row, sw_row, qw_col, sw_col, qx, sx, qdy, sdy, qx_t, sx_t, qdy_t, sdy_t, x_t, dy_t
def bf16_gemms():
torch.mm(x, weight.T, out=bf16_fwd)
torch.mm(dy, weight, out=bf16_dx)
torch.mm(dy.T, x, out=bf16_dw)
def bf16_complete():
bf16_gemms()
bf16_fwd.add_(bias)
torch.sum(dy, dim=0, out=bf16_dbias)
def fp4_complete():
# One exact scale reduction is reused for both orientations of
# each underlying BF16 tensor; the quantized layouts differ.
w_scale = exact_per_tensor_scale(weight)
w_row, w_row_sf = quantized_operand(weight, w_scale)
w_col, w_col_sf = quantized_operand(weight.T.contiguous(), w_scale)
current_x_scale = exact_per_tensor_scale(x)
x_row, x_row_sf = quantized_operand(x, current_x_scale)
run_fwd(x_row, w_row, fp4_fwd_store.permute(1, 2, 0), x_row_sf, w_row_sf)
forward_epilogue(fp4_fwd_store, current_x_scale * w_scale, bias)
current_dy_scale = exact_per_tensor_scale(dy)
dy_row, dy_row_sf = quantized_operand(dy, current_dy_scale)
run_dgrad(dy_row, w_col, fp4_dx_store.permute(1, 2, 0), dy_row_sf, w_col_sf)
scale_epilogue(fp4_dx_store, current_dy_scale * w_scale)
x_transposed = torch.nn.functional.pad(x.T, (0, padded_m - m)).contiguous()
dy_transposed = torch.nn.functional.pad(dy.T, (0, padded_m - m)).contiguous()
x_col, x_col_sf = quantized_operand(x_transposed, current_x_scale)
dy_col, dy_col_sf = quantized_operand(dy_transposed, current_dy_scale)
run_wgrad(dy_col, x_col, fp4_dw_store.permute(1, 2, 0), dy_col_sf, x_col_sf)
scale_epilogue(fp4_dw_store, current_dy_scale * current_x_scale)
torch.sum(dy, dim=0, out=fp4_dbias)
def x_both_orientations():
shared_scale = exact_per_tensor_scale(x)
quantized_operand(x, shared_scale)
x_transposed = torch.nn.functional.pad(x.T, (0, padded_m - m)).contiguous()
quantized_operand(x_transposed, shared_scale)
if args.smoke_only:
fp4_complete()
bf16_complete()
torch.cuda.synchronize()
metrics = {
"fwd": sampled_metrics(torch, fp4_fwd_store[0], bf16_fwd),
"dgrad": sampled_metrics(torch, fp4_dx_store[0], bf16_dx),
"wgrad": sampled_metrics(torch, fp4_dw_store[0], bf16_dw),
"dbias": sampled_metrics(torch, fp4_dbias, bf16_dbias),
}
if not all(metric["finite"] for metric in metrics.values()):
raise RuntimeError(f"non-finite complete projection output: {metrics}")
emit(
"smoke",
name=name,
configs={"fwd": fwd_config, "dgrad": dgrad_config, "wgrad": wgrad_config},
metrics=metrics,
)
break
bf16_a_ms = timed(torch, bf16_gemms, args.warmup, args.samples, args.inner)
bf16_complete_ms = timed(torch, bf16_complete, args.warmup, args.samples, args.inner)
torch.cuda.reset_peak_memory_stats()
fp4_ms = timed(torch, fp4_complete, args.warmup, args.samples, args.inner)
peak_bytes = torch.cuda.max_memory_allocated()
# Text context is identical across all 48 blocks, so its row/transposed
# packs can live across the step and are charged once, not 48 times.
shared_x_ms = (timed(torch, x_both_orientations, args.warmup, args.samples, args.inner)
if name == "text_kv" else 0.0)
bf16_b_ms = timed(torch, bf16_gemms, args.warmup, args.samples, args.inner)
bf16_ms = (bf16_a_ms + bf16_b_ms) / 2
fp4_complete()
bf16_complete()
torch.cuda.synchronize()
metrics = {
"fwd": sampled_metrics(torch, fp4_fwd_store[0], bf16_fwd),
"dgrad": sampled_metrics(torch, fp4_dx_store[0], bf16_dx),
"wgrad": sampled_metrics(torch, fp4_dw_store[0], bf16_dw),
"dbias": sampled_metrics(torch, fp4_dbias, bf16_dbias),
}
if not all(metric["finite"] for metric in metrics.values()):
raise RuntimeError(f"non-finite complete projection output: {metrics}")
row = {
"name": name,
"count": count,
"logical_shape": [m, k, n],
"configs": {"fwd": fwd_config, "dgrad": dgrad_config, "wgrad": wgrad_config},
"bf16_gemm_a_ms": bf16_a_ms,
"bf16_gemm_b_ms": bf16_b_ms,
"bf16_gemm_midpoint_ms": bf16_ms,
"bf16_complete_ms": bf16_complete_ms,
"fp4_complete_ms": fp4_ms,
"shareable_x_quant_ms": shared_x_ms,
"speedup_vs_bf16_gemms": bf16_ms / fp4_ms,
"speedup_vs_bf16_complete": bf16_complete_ms / fp4_ms,
"peak_allocated_gib": peak_bytes / 2**30,
"metrics": metrics,
}
rows.append(row)
emit("case", **row)
del weight_master, bias_master, weight, bias, x, dy
del fp4_fwd_store, fp4_dx_store, fp4_dw_store, fp4_dbias
del bf16_fwd, bf16_dx, bf16_dw, bf16_dbias
del run_fwd, run_dgrad, run_wgrad
gc.collect()
torch.cuda.empty_cache()
if rows:
bf16_gemm_ms = sum(row["bf16_gemm_midpoint_ms"] * row["count"] for row in rows)
bf16_complete_ms = sum(row["bf16_complete_ms"] * row["count"] for row in rows)
fp4_unamortized_ms = sum(row["fp4_complete_ms"] * row["count"] for row in rows)
fp4_ms = sum(row["fp4_complete_ms"] * row["count"]
- row["shareable_x_quant_ms"] * (row["count"] - 1) for row in rows)
speedup = bf16_gemm_ms / fp4_ms
normalized_ms = HISTORICAL_BF16_MS / speedup
emit(
"aggregate",
packs=[row["name"] for row in rows],
bf16_gemm_ms=bf16_gemm_ms,
bf16_complete_ms=bf16_complete_ms,
fp4_unamortized_ms=fp4_unamortized_ms,
fp4_complete_ms=fp4_ms,
speedup_vs_bf16_gemms=speedup,
speedup_vs_bf16_complete=bf16_complete_ms / fp4_ms,
historical_bf16_ms=HISTORICAL_BF16_MS,
ratio_normalized_complete_ms=normalized_ms,
gate_ms=GATE_MS,
margin_ms=MARGIN_MS,
passes_break_even=normalized_ms <= GATE_MS,
passes_margin=normalized_ms <= MARGIN_MS,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,165 @@
#!/usr/bin/env python3
"""Tiny exact-shape gate for PyTorch's vendored SM100 NVFP4 CuTe GEMM."""
import argparse
import importlib
import json
import statistics
import time
import cutlass
import cutlass.cute as cute
import flashinfer
import torch
from cutlass.cute.runtime import from_dlpack
CONFIGS = (
((128, 64), (1, 1)),
((128, 128), (1, 1)),
((128, 192), (1, 1)),
((128, 256), (1, 1)),
((256, 64), (2, 1)),
((256, 128), (2, 1)),
((256, 192), (2, 1)),
((256, 256), (2, 1)),
)
# PyTorch 2.12's vendored template uses the pre-4.6 enum spelling; the
# installed 4.6 DSL accepts the same values as strings.
if not hasattr(cute.arch, "ProxyKind"):
cute.arch.ProxyKind = type("ProxyKind", (), {"async_shared": "async.shared"})
cute.arch.SharedSpace = type("SharedSpace", (), {"shared_cta": "cta"})
def emit(kind, **fields):
print(json.dumps({"kind": kind, **fields}, sort_keys=True), flush=True)
def quantize(x):
one = torch.ones((), device=x.device, dtype=torch.float32)
q, s = flashinfer.nvfp4_quantize(
x,
one,
sfLayout=flashinfer.SfLayout.layout_128x4,
do_shuffle=False,
)
return q.view(torch.float4_e2m1fn_x2), s.view(torch.float8_e4m3fn)
def as_cute(x):
return from_dlpack(x.detach(), assumed_align=16, enable_tvm_ffi=True)
def compile_gemm(kernel_cls, a, b, sfa, sfb, out, tile, cluster):
kernel = kernel_cls(16, tile, cluster)
max_clusters = cutlass.utils.HardwareInfo().get_max_active_clusters(cluster[0] * cluster[1])
stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True)
return cute.compile(
kernel,
as_cute(a),
as_cute(b),
as_cute(sfa),
as_cute(sfb),
as_cute(out),
max_clusters,
stream,
options="--enable-tvm-ffi",
)
def measure(fn, warmup, samples, inner):
for _ in range(warmup):
fn()
torch.cuda.synchronize()
timings = []
for _ in range(samples):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(inner):
fn()
end.record()
end.synchronize()
timings.append(start.elapsed_time(end) / inner)
return statistics.median(timings), min(timings), max(timings)
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--m", type=int, default=4290)
parser.add_argument("--k", type=int, default=4096)
parser.add_argument("--n", type=int, default=4096)
parser.add_argument("--warmup", type=int, default=3)
parser.add_argument("--samples", type=int, default=7)
parser.add_argument("--inner", type=int, default=5)
parser.add_argument("--config", type=int, nargs="*")
args = parser.parse_args()
torch.cuda.set_device(0)
torch.manual_seed(20260721)
module = importlib.import_module(
"torch._inductor.kernel.vendored_templates.cutedsl.dense_blockscaled_gemm_persistent"
)
kernel_cls = module.Sm100BlockScaledPersistentDenseGemmKernel
x = torch.randn(args.m, args.k, device="cuda", dtype=torch.bfloat16)
w = torch.randn(args.n, args.k, device="cuda", dtype=torch.bfloat16) * 0.02
aq, asf = quantize(x)
bq, bsf = quantize(w)
bq_t = bq.T
out = torch.empty(args.m, args.n, device="cuda", dtype=torch.bfloat16)
ref = x @ w.T
emit(
"environment",
torch=torch.__version__,
cutlass=cutlass.__version__,
flashinfer=flashinfer.__version__,
gpu=torch.cuda.get_device_name(),
shape=(args.m, args.k, args.n),
a=(tuple(aq.shape), tuple(aq.stride()), str(aq.dtype)),
b=(tuple(bq_t.shape), tuple(bq_t.stride()), str(bq_t.dtype)),
sfa=(tuple(asf.shape), asf.numel()),
sfb=(tuple(bsf.shape), bsf.numel()),
)
selected = range(len(CONFIGS)) if args.config is None else args.config
for index in selected:
tile, cluster = CONFIGS[index]
started = time.monotonic()
try:
compiled = compile_gemm(kernel_cls, aq, bq_t, asf, bsf, out, tile, cluster)
compile_s = time.monotonic() - started
call = lambda: compiled(aq, bq_t, asf, bsf, out)
call()
torch.cuda.synchronize()
diff = out.float() - ref.float()
rel_rms = (diff.square().mean().sqrt() / ref.float().square().mean().sqrt()).item()
median, minimum, maximum = measure(call, args.warmup, args.samples, args.inner)
emit(
"result",
index=index,
tile=tile,
cluster=cluster,
compile_seconds=compile_s,
median_ms=median,
min_ms=minimum,
max_ms=maximum,
relative_rms=rel_rms,
finite=bool(torch.isfinite(out).all()),
effective_tflops=2 * args.m * args.k * args.n / median / 1e9,
)
except Exception as exc:
torch.cuda.synchronize()
emit(
"failure",
index=index,
tile=tile,
cluster=cluster,
elapsed_seconds=time.monotonic() - started,
exception=repr(exc),
)
if __name__ == "__main__":
main()
@@ -0,0 +1,402 @@
#!/usr/bin/env python3
"""Exact B=1 LTX-2 VideoOnly raw FlashInfer NVFP4 linear benchmark.
This intentionally imports only torch and flashinfer. It compares:
* native BF16 GEMMs in their training orientations;
* prequantized NVFP4 mm_fp4 (primitive ceiling);
* static-calibrated NVFP4 qmm, including the transpose/pad work needed by
dgrad and wgrad with FlashInfer's row-major A / column-major B contract.
The four layer classes and multiplicities cover the 48 transformer blocks.
"""
import argparse
import gc
import json
import statistics
import time
from dataclasses import dataclass
from typing import Callable
import flashinfer
import torch
import torch.nn.functional as F
@dataclass(frozen=True)
class LayerShape:
name: str
count: int
m: int
k: int
n: int
LAYERS = (
LayerShape("video_dd", 288, 4290, 4096, 4096),
LayerShape("text_dd", 96, 1024, 4096, 4096),
LayerShape("ffn_up", 48, 4290, 4096, 16384),
LayerShape("ffn_down", 48, 4290, 16384, 4096),
)
PHASES = ("fwd", "dgrad", "wgrad")
BACKENDS = ("auto", "cutlass", "cudnn")
FP4_MAX_TIMES_FP8_MAX = 6.0 * 448.0
def emit(kind: str, **values: object) -> None:
print(json.dumps({"kind": kind, **values}, sort_keys=True), flush=True)
def calibrated_static_sf(tensor: torch.Tensor) -> torch.Tensor:
# Calibration is outside all timed regions. Production would retain one
# such scalar per tensor role/layer and refresh it only deliberately.
amax = tensor.abs().amax().float().clamp_min_(1e-12)
return torch.as_tensor(FP4_MAX_TIMES_FP8_MAX, device=tensor.device) / amax
def quantize(tensor: torch.Tensor, global_sf: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
return flashinfer.nvfp4_quantize(
tensor,
global_sf,
sfLayout=flashinfer.SfLayout.layout_128x4,
do_shuffle=False,
)
def mm_fp4(
lhs_fp4: torch.Tensor,
rhs_fp4: torch.Tensor,
lhs_scale: torch.Tensor,
rhs_scale: torch.Tensor,
alpha: torch.Tensor,
out: torch.Tensor,
backend: str,
) -> torch.Tensor:
return flashinfer.mm_fp4(
lhs_fp4,
rhs_fp4.T,
lhs_scale,
rhs_scale.T,
alpha,
torch.bfloat16,
out,
block_size=16,
use_8x4_sf_layout=False,
backend=backend,
use_nvfp4=True,
)
def qmm(
lhs: torch.Tensor,
rhs_rows: torch.Tensor,
lhs_sf: torch.Tensor,
rhs_sf: torch.Tensor,
alpha: torch.Tensor,
out: torch.Tensor,
backend: str,
) -> torch.Tensor:
lhs_fp4, lhs_scale = quantize(lhs, lhs_sf)
rhs_fp4, rhs_scale = quantize(rhs_rows, rhs_sf)
return mm_fp4(lhs_fp4, rhs_fp4, lhs_scale, rhs_scale, alpha, out, backend)
def benchmark(fn: Callable[[], torch.Tensor], warmup: int, samples: int, inner: int) -> dict[str, float]:
with torch.inference_mode():
for _ in range(warmup):
fn()
torch.cuda.synchronize()
values = []
for _ in range(samples):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(inner):
fn()
end.record()
end.synchronize()
values.append(start.elapsed_time(end) / inner)
return {
"median_ms": statistics.median(values),
"min_ms": min(values),
"max_ms": max(values),
}
def numerical_error(actual: torch.Tensor, reference: torch.Tensor) -> dict[str, float | bool]:
actual_f = actual.float()
reference_f = reference.float()
diff = actual_f - reference_f
reference_rms = reference_f.square().mean().sqrt()
return {
"relative_rms": (diff.square().mean().sqrt() / reference_rms.clamp_min(1e-30)).item(),
"mean_abs": diff.abs().mean().item(),
"max_abs": diff.abs().amax().item(),
"finite": bool(torch.isfinite(actual).all().item()),
}
def make_case(
layer: LayerShape,
phase: str,
) -> tuple[
Callable[[torch.Tensor], torch.Tensor],
Callable[[], tuple[torch.Tensor, torch.Tensor]],
torch.Tensor,
torch.Tensor,
tuple[int, int, int],
int,
]:
"""Return native BF16 op, FP4-orientation builder, bases, (m,k,n), FLOPs."""
m, k, n = layer.m, layer.k, layer.n
if phase == "fwd":
x = torch.empty((m, k), device="cuda", dtype=torch.bfloat16).normal_(0.0, 1.0)
weight = torch.empty((n, k), device="cuda", dtype=torch.bfloat16).normal_(0.0, 0.02)
def bf16(out: torch.Tensor) -> torch.Tensor:
return torch.mm(x, weight.T, out=out)
def orient() -> tuple[torch.Tensor, torch.Tensor]:
return x, weight
return bf16, orient, x, weight, (m, k, n), 2 * m * k * n
if phase == "dgrad":
dy = torch.empty((m, n), device="cuda", dtype=torch.bfloat16).normal_(0.0, 0.01)
weight = torch.empty((n, k), device="cuda", dtype=torch.bfloat16).normal_(0.0, 0.02)
def bf16(out: torch.Tensor) -> torch.Tensor:
return torch.mm(dy, weight, out=out)
def orient() -> tuple[torch.Tensor, torch.Tensor]:
return dy, weight.T.contiguous()
return bf16, orient, dy, weight, (m, n, k), 2 * m * n * k
if phase == "wgrad":
dy = torch.empty((m, n), device="cuda", dtype=torch.bfloat16).normal_(0.0, 0.01)
x = torch.empty((m, k), device="cuda", dtype=torch.bfloat16).normal_(0.0, 1.0)
# FlashInfer's quantizer accepts K % 16 == 0, but its CUTLASS mm_fp4
# path additionally requires logical K % 32 == 0.
padded_m = (m + 31) // 32 * 32
def bf16(out: torch.Tensor) -> torch.Tensor:
return torch.mm(dy.T, x, out=out)
def orient() -> tuple[torch.Tensor, torch.Tensor]:
pad_rows = padded_m - m
if pad_rows:
dy_for_fp4 = F.pad(dy, (0, 0, 0, pad_rows))
x_for_fp4 = F.pad(x, (0, 0, 0, pad_rows))
else:
dy_for_fp4 = dy
x_for_fp4 = x
return dy_for_fp4.T.contiguous(), x_for_fp4.T.contiguous()
return bf16, orient, dy, x, (n, padded_m, k), 2 * m * n * k
raise ValueError(phase)
def main() -> None:
started = time.monotonic()
parser = argparse.ArgumentParser()
parser.add_argument("--warmup", type=int, default=4)
parser.add_argument("--samples", type=int, default=7)
parser.add_argument("--inner", type=int, default=5)
parser.add_argument("--backends", nargs="+", default=list(BACKENDS))
args = parser.parse_args()
if not torch.cuda.is_available():
raise SystemExit("CUDA is required")
torch.cuda.set_device(0)
torch.manual_seed(20260721)
torch.backends.cuda.matmul.allow_tf32 = False
emit(
"environment",
torch=torch.__version__,
flashinfer=getattr(flashinfer, "__version__", "unknown"),
gpu=torch.cuda.get_device_name(0),
capability=torch.cuda.get_device_capability(0),
warmup=args.warmup,
samples=args.samples,
inner=args.inner,
backends=args.backends,
)
aggregate: dict[str, dict[str, dict[str, float]]] = {
"bf16": {phase: {"latency_ms": 0.0, "logical_flops": 0.0, "cases": 0.0} for phase in PHASES}
}
for backend in args.backends:
aggregate[f"fp4_mm_{backend}"] = {
phase: {"latency_ms": 0.0, "logical_flops": 0.0, "cases": 0.0} for phase in PHASES
}
aggregate[f"fp4_qmm_{backend}"] = {
phase: {"latency_ms": 0.0, "logical_flops": 0.0, "cases": 0.0} for phase in PHASES
}
for phase in PHASES:
for layer_index, layer in enumerate(LAYERS):
gc.collect()
torch.cuda.empty_cache()
torch.manual_seed(20260721 + 100 * layer_index + PHASES.index(phase))
bf16_op, orient, base_a, base_b, qshape, logical_flops = make_case(layer, phase)
qm, qk, qn = qshape
out_bf16 = torch.empty((qm, qn), device="cuda", dtype=torch.bfloat16)
bf16_timing = benchmark(
lambda: bf16_op(out_bf16),
args.warmup,
args.samples,
args.inner,
)
aggregate["bf16"][phase]["latency_ms"] += bf16_timing["median_ms"] * layer.count
aggregate["bf16"][phase]["logical_flops"] += logical_flops * layer.count
aggregate["bf16"][phase]["cases"] += layer.count
emit(
"case",
tier="bf16",
phase=phase,
layer=layer.name,
count=layer.count,
qmm_shape=qshape,
logical_shape=(layer.m, layer.k, layer.n),
logical_flops=logical_flops,
**bf16_timing,
)
# Build each operand once for static calibration and primitive-only
# timing. Calibration and this orientation work are outside the
# primitive tier, but the latter is repeated inside deployable qmm.
lhs, rhs_rows = orient()
lhs_sf = calibrated_static_sf(lhs)
rhs_sf = calibrated_static_sf(rhs_rows)
alpha = (1.0 / (lhs_sf * rhs_sf)).float()
lhs_fp4, lhs_scale = quantize(lhs, lhs_sf)
rhs_fp4, rhs_scale = quantize(rhs_rows, rhs_sf)
torch.cuda.synchronize()
for backend in args.backends:
out_fp4 = torch.empty((qm, qn), device="cuda", dtype=torch.bfloat16)
try:
mm_timing = benchmark(
lambda backend=backend: mm_fp4(
lhs_fp4,
rhs_fp4,
lhs_scale,
rhs_scale,
alpha,
out_fp4,
backend,
),
args.warmup,
args.samples,
args.inner,
)
mm_fp4(lhs_fp4, rhs_fp4, lhs_scale, rhs_scale, alpha, out_fp4, backend)
torch.cuda.synchronize()
mm_error = numerical_error(out_fp4, out_bf16)
key = f"fp4_mm_{backend}"
aggregate[key][phase]["latency_ms"] += mm_timing["median_ms"] * layer.count
aggregate[key][phase]["logical_flops"] += logical_flops * layer.count
aggregate[key][phase]["cases"] += layer.count
emit(
"case",
tier="fp4_mm",
backend=backend,
phase=phase,
layer=layer.name,
count=layer.count,
qmm_shape=qshape,
logical_shape=(layer.m, layer.k, layer.n),
logical_flops=logical_flops,
error=mm_error,
**mm_timing,
)
except Exception as exc:
torch.cuda.synchronize()
emit(
"failure",
tier="fp4_mm",
backend=backend,
phase=phase,
layer=layer.name,
qmm_shape=qshape,
exception=repr(exc),
)
continue
try:
def deployable(backend: str = backend) -> torch.Tensor:
deploy_lhs, deploy_rhs = orient()
return qmm(deploy_lhs, deploy_rhs, lhs_sf, rhs_sf, alpha, out_fp4, backend)
qmm_timing = benchmark(
deployable,
args.warmup,
args.samples,
args.inner,
)
deployable()
torch.cuda.synchronize()
qmm_error = numerical_error(out_fp4, out_bf16)
key = f"fp4_qmm_{backend}"
aggregate[key][phase]["latency_ms"] += qmm_timing["median_ms"] * layer.count
aggregate[key][phase]["logical_flops"] += logical_flops * layer.count
aggregate[key][phase]["cases"] += layer.count
emit(
"case",
tier="fp4_qmm",
backend=backend,
phase=phase,
layer=layer.name,
count=layer.count,
qmm_shape=qshape,
logical_shape=(layer.m, layer.k, layer.n),
logical_flops=logical_flops,
error=qmm_error,
**qmm_timing,
)
except Exception as exc:
torch.cuda.synchronize()
emit(
"failure",
tier="fp4_qmm",
backend=backend,
phase=phase,
layer=layer.name,
qmm_shape=qshape,
exception=repr(exc),
)
del lhs, rhs_rows, lhs_fp4, lhs_scale, rhs_fp4, rhs_scale
del base_a, base_b, out_bf16
gc.collect()
torch.cuda.empty_cache()
for tier, phases in aggregate.items():
complete = all(values["cases"] == sum(layer.count for layer in LAYERS) for values in phases.values())
total_ms = sum(values["latency_ms"] for values in phases.values())
total_flops = sum(values["logical_flops"] for values in phases.values())
phase_results = {}
for phase, values in phases.items():
latency_ms = values["latency_ms"]
phase_results[phase] = {
**values,
"effective_tflops": values["logical_flops"] / latency_ms / 1e9 if latency_ms else None,
}
emit(
"aggregate",
tier=tier,
complete=complete,
phases=phase_results,
total_latency_ms=total_ms,
total_logical_flops=total_flops,
effective_tflops=total_flops / total_ms / 1e9 if total_ms else None,
)
emit("done", elapsed_wall_seconds=time.monotonic() - started)
if __name__ == "__main__":
main()
@@ -0,0 +1,446 @@
#!/usr/bin/env python3
"""One-GPU exact-shape BF16 GEMM gate for packed LTX-2 training.
Run each backend in a fresh process. The first invocation of every compiled
shape/phase is reported as cold compile time and is never included in timing.
"""
from __future__ import annotations
import argparse
import collections
import json
import math
import os
import statistics
import subprocess
import sys
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable
import torch
import torch.nn.functional as F
BLOCKS = 48
DTYPE = torch.bfloat16
@dataclass(frozen=True)
class Shape:
key: str
m: int
k: int
n: int
roles: tuple[str, ...]
# The seven packed projections reduce to five unique shapes. Each role occurs
# once per transformer block, once in fprop, dgrad, and wgrad.
SHAPES = (
Shape("self_qkv", 4290, 4096, 12288, ("self_qkv", )),
Shape("video_4096", 4290, 4096, 4096, ("self_out", "cross_q", "cross_out")),
Shape("text_kv", 1024, 4096, 8192, ("text_kv", )),
Shape("ffn_up", 4290, 4096, 16384, ("ffn_up", )),
Shape("ffn_down", 4290, 16384, 4096, ("ffn_down", )),
)
PHASES = ("fprop", "dgrad", "wgrad")
EXPECTED_TFLOP_PER_STEP = 300.095807422464
EXPECTED_LOGICAL_GEMMS_PER_STEP = 7 * BLOCKS * len(PHASES)
def _fprop(x: torch.Tensor, w: torch.Tensor, bias: torch.Tensor) -> torch.Tensor:
return F.linear(x, w, bias)
def _dgrad(dy: torch.Tensor, w: torch.Tensor) -> torch.Tensor:
return torch.mm(dy, w)
def _wgrad(dy: torch.Tensor, x: torch.Tensor) -> torch.Tensor:
return torch.mm(dy.t(), x)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--variant", choices=("current", "cutlass"))
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--warmup", type=int, default=5)
parser.add_argument("--samples", type=int, default=15)
parser.add_argument("--inner", type=int, default=3)
parser.add_argument("--seed", type=int, default=1630)
parser.add_argument("--batch-factor", type=int, default=1)
parser.add_argument(
"--compare",
nargs=3,
type=Path,
metavar=("CURRENT_A", "CUTLASS", "CURRENT_B"),
help="compare three completed JSONL files instead of running CUDA",
)
args = parser.parse_args()
if (args.variant is None) == (args.compare is None):
parser.error("specify exactly one of --variant or --compare")
if min(args.warmup, args.samples, args.inner, args.batch_factor) < 1:
parser.error("--warmup, --samples, --inner, and --batch-factor must be positive")
return args
class Jsonl:
def __init__(self, path: Path):
path.parent.mkdir(parents=True, exist_ok=True)
self._file = path.open("w", encoding="utf-8")
def emit(self, kind: str, **values: Any) -> None:
line = json.dumps({"kind": kind, **values}, sort_keys=True)
print(line, flush=True)
self._file.write(line + "\n")
self._file.flush()
def close(self) -> None:
self._file.close()
def git_head(path: str | None) -> str | None:
if not path:
return None
try:
return subprocess.check_output(
["git", "-C", path, "rev-parse", "HEAD"],
text=True,
stderr=subprocess.DEVNULL,
).strip()
except (OSError, subprocess.CalledProcessError):
return None
def config_proof(variant: str) -> dict[str, Any]:
from torch._inductor import config
from torch._inductor.codegen.cutlass.utils import try_import_cutlass
swizzles_env = os.environ.get("CUTLASS_SWIZZLES")
if swizzles_env:
config.cutlass.cutlass_max_profiling_swizzle_options = [int(value) for value in swizzles_env.split(",")]
cutlass_available = try_import_cutlass()
proof = {
"variant": variant,
"torch": torch.__version__,
"torch_cuda": torch.version.cuda,
"device": torch.cuda.get_device_name(),
"capability": list(torch.cuda.get_device_capability()),
"fastvideo_commit": os.environ.get("FASTVIDEO_COMMIT"),
"max_autotune_gemm": config.max_autotune_gemm,
"max_autotune_gemm_backends": config.max_autotune_gemm_backends,
"max_autotune_gemm_search_space": config.max_autotune_gemm_search_space,
"inductor_cache_dir": os.environ.get("TORCHINDUCTOR_CACHE_DIR"),
"cutlass_dir": config.cutlass.cutlass_dir,
"cutlass_commit": git_head(config.cutlass.cutlass_dir),
"cutlass_importable": cutlass_available,
"cutlass_enabled_ops": config.cutlass.cutlass_enabled_ops,
"cutlass_instantiation_level": config.cutlass.cutlass_instantiation_level,
"cutlass_allowlist": config.cutlass.cutlass_op_allowlist_regex,
"cutlass_denylist": config.cutlass.cutlass_op_denylist_regex,
"cutlass_swizzles": list(config.cutlass.cutlass_max_profiling_swizzle_options),
"cublas_preferred_backend": str(torch.backends.cuda.preferred_blas_library()),
"env": {
key: os.environ.get(key)
for key in (
"CUDA_VISIBLE_DEVICES",
"TORCHINDUCTOR_MAX_AUTOTUNE_GEMM",
"TORCHINDUCTOR_MAX_AUTOTUNE_GEMM_BACKENDS",
"TORCHINDUCTOR_MAX_AUTOTUNE_GEMM_SEARCH_SPACE",
"TORCHINDUCTOR_CUTLASS_DIR",
"TORCHINDUCTOR_CUTLASS_INSTANTIATION_LEVEL",
"CUTLASS_EPILOGUE_FUSION",
"CUTLASS_SWIZZLES",
)
},
}
if variant == "cutlass":
if not config.max_autotune_gemm:
raise RuntimeError("cutlass variant requires TORCHINDUCTOR_MAX_AUTOTUNE_GEMM=1")
backends = {item.strip() for item in config.max_autotune_gemm_backends.split(",")}
if backends != {"ATEN", "CUTLASS"}:
raise RuntimeError(f"expected candidate backends ATEN,CUTLASS; got {sorted(backends)}")
if not cutlass_available:
raise RuntimeError(f"PyTorch cannot import CUTLASS from {config.cutlass.cutlass_dir}")
if list(config.cutlass.cutlass_max_profiling_swizzle_options) != [4]:
raise RuntimeError("this safety gate requires exactly CUTLASS_SWIZZLES=4")
elif config.max_autotune_gemm:
raise RuntimeError("current variant requires TORCHINDUCTOR_MAX_AUTOTUNE_GEMM=0")
return proof
def make_inputs(shape: Shape, seed: int) -> dict[str, torch.Tensor]:
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
x = torch.empty((shape.m, shape.k), device="cuda", dtype=DTYPE).normal_()
w = torch.empty((shape.n, shape.k), device="cuda", dtype=DTYPE).normal_()
w.mul_(1.0 / math.sqrt(shape.k))
bias = torch.empty((shape.n, ), device="cuda", dtype=DTYPE).normal_(std=0.01)
dy = torch.empty((shape.m, shape.n), device="cuda", dtype=DTYPE).normal_()
return {"x": x, "w": w, "bias": bias, "dy": dy}
def phase_call(
phase: str,
tensors: dict[str, torch.Tensor],
) -> tuple[Callable[..., torch.Tensor], tuple[torch.Tensor, ...]]:
if phase == "fprop":
return _fprop, (tensors["x"], tensors["w"], tensors["bias"])
if phase == "dgrad":
return _dgrad, (tensors["dy"], tensors["w"])
if phase == "wgrad":
return _wgrad, (tensors["dy"], tensors["x"])
raise AssertionError(phase)
@torch.no_grad()
def parity(candidate: torch.Tensor, reference: torch.Tensor) -> dict[str, Any]:
candidate_f = candidate.float()
reference_f = reference.float()
delta = candidate_f - reference_f
max_abs = delta.abs().max().item()
max_ref = reference_f.abs().max().item()
relative_l2 = (delta.norm() / reference_f.norm().clamp_min(1.0e-12)).item()
passed = relative_l2 <= 0.02 and max_abs <= max(1.0, 0.03 * max_ref)
return {
"passed": passed,
"max_abs": max_abs,
"max_ref": max_ref,
"relative_l2": relative_l2,
}
@torch.no_grad()
def time_cuda(
fn: Callable[..., torch.Tensor],
inputs: tuple[torch.Tensor, ...],
warmup: int,
samples: int,
inner: int,
) -> list[float]:
for _ in range(warmup):
fn(*inputs)
torch.cuda.synchronize()
timings = []
for _ in range(samples):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(inner):
output = fn(*inputs)
end.record()
end.synchronize()
timings.append(start.elapsed_time(end) / inner)
del output
return timings
@torch.no_grad()
def kernel_names(
fn: Callable[..., torch.Tensor],
inputs: tuple[torch.Tensor, ...],
) -> list[dict[str, Any]]:
"""Capture selected runtime kernel names after timing, outside the gate."""
try:
from torch.profiler import ProfilerActivity, profile
with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA]) as prof:
fn(*inputs)
torch.cuda.synchronize()
by_name: collections.defaultdict[str, float] = collections.defaultdict(float)
for event in prof.events():
if "cuda" not in str(getattr(event, "device_type", "")).lower():
continue
value = getattr(event, "self_device_time_total", None)
if value is None:
value = getattr(event, "self_cuda_time_total", 0.0)
by_name[event.name] += float(value)
return [
{"name": name, "device_time_us": value}
for name, value in sorted(by_name.items(), key=lambda item: item[1], reverse=True)[:5]
]
except Exception as exc: # Profiler proof is useful, but timing must survive profiler drift.
return [{"profiler_error": repr(exc)}]
def run_benchmark(args: argparse.Namespace, output: Jsonl) -> int:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required")
if torch.cuda.device_count() != 1:
raise RuntimeError(f"expected exactly one visible GPU, got {torch.cuda.device_count()}")
proof = config_proof(args.variant)
shapes = tuple(Shape(shape.key, shape.m * args.batch_factor, shape.k, shape.n, shape.roles) for shape in SHAPES)
expected_tflop_per_step = EXPECTED_TFLOP_PER_STEP * args.batch_factor
output.emit(
"config",
**proof,
blocks=BLOCKS,
logical_roles=[role for shape in SHAPES for role in shape.roles],
logical_shapes=[
{"role": role, "m": shape.m, "k": shape.k, "n": shape.n}
for shape in shapes
for role in shape.roles
],
unique_compile_shapes=[
{"shape": shape.key, "m": shape.m, "k": shape.k, "n": shape.n, "roles": list(shape.roles)}
for shape in shapes
],
expected_logical_gemms_per_step=EXPECTED_LOGICAL_GEMMS_PER_STEP,
expected_tflop_per_step=expected_tflop_per_step,
batch_factor=args.batch_factor,
warmup=args.warmup,
samples=args.samples,
inner=args.inner,
seed=args.seed,
)
total_weighted_ms = 0.0
total_compile_s = 0.0
all_parity_passed = True
results: list[dict[str, Any]] = []
for shape_index, shape in enumerate(shapes):
tensors = make_inputs(shape, args.seed + shape_index)
for phase in PHASES:
eager_fn, inputs = phase_call(phase, tensors)
with torch.no_grad():
reference = eager_fn(*inputs)
compiled_fn = torch.compile(eager_fn, fullgraph=True, dynamic=False)
compile_started = time.monotonic()
with torch.no_grad():
candidate = compiled_fn(*inputs)
torch.cuda.synchronize()
compile_s = time.monotonic() - compile_started
parity_stats = parity(candidate, reference)
del candidate, reference
timings = time_cuda(compiled_fn, inputs, args.warmup, args.samples, args.inner)
median_ms = statistics.median(timings)
weighted_calls = len(shape.roles) * BLOCKS
weighted_ms = median_ms * weighted_calls
operation_tflop = 2.0 * shape.m * shape.k * shape.n / 1.0e12
result = {
"variant": args.variant,
"shape": shape.key,
"roles": list(shape.roles),
"m": shape.m,
"k": shape.k,
"n": shape.n,
"phase": phase,
"input_shapes": [list(tensor.shape) for tensor in inputs],
"input_strides": [list(tensor.stride()) for tensor in inputs],
"compile_s": compile_s,
"median_ms": median_ms,
"min_ms": min(timings),
"max_ms": max(timings),
"samples_ms": timings,
"operation_tflop": operation_tflop,
"effective_petaflop_s": operation_tflop / median_ms,
"weighted_calls_per_step": weighted_calls,
"weighted_step_ms": weighted_ms,
"parity": parity_stats,
"selected_runtime_kernels": kernel_names(compiled_fn, inputs),
}
output.emit("shape_phase", **result)
results.append(result)
total_weighted_ms += weighted_ms
total_compile_s += compile_s
all_parity_passed &= parity_stats["passed"]
del tensors
torch.cuda.empty_cache()
output.emit(
"aggregate",
variant=args.variant,
unique_shapes=len(shapes),
compiled_shape_phases=len(results),
logical_roles=7,
logical_gemms_per_step=EXPECTED_LOGICAL_GEMMS_PER_STEP,
tflop_per_step=expected_tflop_per_step,
weighted_step_ms=total_weighted_ms,
effective_petaflop_s=expected_tflop_per_step / total_weighted_ms,
cold_compile_s=total_compile_s,
parity_passed=all_parity_passed,
)
return 0 if all_parity_passed else 2
def read_records(path: Path) -> list[dict[str, Any]]:
return [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines() if line.strip()]
def one_record(records: list[dict[str, Any]], kind: str) -> dict[str, Any]:
matches = [record for record in records if record.get("kind") == kind]
if len(matches) != 1:
raise RuntimeError(f"expected one {kind!r} record, found {len(matches)}")
return matches[0]
def run_compare(args: argparse.Namespace, output: Jsonl) -> int:
assert args.compare is not None
paths = args.compare
record_sets = [read_records(path) for path in paths]
aggregates = [one_record(records, "aggregate") for records in record_sets]
labels = ("current_a", "cutlass", "current_b")
for expected, aggregate in zip(("current", "cutlass", "current"), aggregates):
if aggregate["variant"] != expected:
raise RuntimeError(f"expected {expected}, got {aggregate['variant']}")
tflop_per_step = aggregates[0]["tflop_per_step"]
if any(aggregate["tflop_per_step"] != tflop_per_step for aggregate in aggregates[1:]):
raise RuntimeError("A/X/B FLOP counts differ")
baseline_ms = statistics.mean((aggregates[0]["weighted_step_ms"], aggregates[2]["weighted_step_ms"]))
candidate_ms = aggregates[1]["weighted_step_ms"]
saving_ms = baseline_ms - candidate_ms
output.emit(
"comparison",
inputs={label: str(path) for label, path in zip(labels, paths)},
weighted_step_ms={label: aggregate["weighted_step_ms"] for label, aggregate in zip(labels, aggregates)},
baseline_midpoint_ms=baseline_ms,
candidate_ms=candidate_ms,
saving_ms=saving_ms,
saving_percent=100.0 * saving_ms / baseline_ms,
control_spread_ms=abs(aggregates[0]["weighted_step_ms"] - aggregates[2]["weighted_step_ms"]),
baseline_effective_petaflop_s=tflop_per_step / baseline_ms,
candidate_effective_petaflop_s=tflop_per_step / candidate_ms,
all_parity_passed=all(aggregate["parity_passed"] for aggregate in aggregates),
pass_two_ms_gate=saving_ms >= 2.0,
)
keyed = []
for records in record_sets:
keyed.append({(r["shape"], r["phase"]): r for r in records if r.get("kind") == "shape_phase"})
if keyed[0].keys() != keyed[1].keys() or keyed[0].keys() != keyed[2].keys():
raise RuntimeError("shape/phase sets differ")
for key in sorted(keyed[0]):
a, candidate, b = (records[key] for records in keyed)
base = statistics.mean((a["weighted_step_ms"], b["weighted_step_ms"]))
output.emit(
"comparison_shape_phase",
shape=key[0],
phase=key[1],
baseline_midpoint_ms=base,
candidate_ms=candidate["weighted_step_ms"],
saving_ms=base - candidate["weighted_step_ms"],
saving_percent=100.0 * (base - candidate["weighted_step_ms"]) / base,
candidate_selected_runtime_kernels=candidate["selected_runtime_kernels"],
)
return 0
def main() -> int:
args = parse_args()
output = Jsonl(args.output)
try:
if args.compare is not None:
return run_compare(args, output)
return run_benchmark(args, output)
finally:
output.close()
if __name__ == "__main__":
sys.exit(main())
@@ -0,0 +1,412 @@
#!/usr/bin/env python3
"""Disposable GB200 gate for fused LTX-2 AdaLN and joint Q/K RMSNorm."""
from __future__ import annotations
import argparse
from collections.abc import Callable, Sequence
import json
import math
import statistics
from typing import Any
import torch
import torch.nn.functional as F
DEFAULT_BATCH = 1
DEFAULT_TOKENS = 11 * (480 // 32) * (832 // 32) # 4290
DEFAULT_DIM = 32 * 128 # 4096
DEFAULT_BLOCKS = 48
DEFAULT_TRACE_NORM_MS = 27.723
DEFAULT_STEP_MS = 427.432493
EPS = 1e-6
_quack_rmsnorm: Callable[..., torch.Tensor] | None = None
def torch_adaln(
x: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
) -> torch.Tensor:
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
normalized = F.rms_norm(x, (x.shape[-1], ), eps=EPS).to(x.dtype)
return normalized * (1 + scale) + shift
def quack_adaln(
x: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
) -> torch.Tensor:
if _quack_rmsnorm is None:
raise RuntimeError("QuACK was not initialized")
# QuACK's per-head mode supplies one affine row per sample. Moving B next
# to D is a view; the kernel accepts dynamic leading strides, so this is
# general for B>1 without copying the [B, T, D] activation.
x_tbd = x.transpose(0, 1)
weight = 1 + scale.squeeze(1)
bias = shift.squeeze(1)
return _quack_rmsnorm(x_tbd, weight=weight, bias=bias, eps=EPS).transpose(0, 1)
def torch_joint_qk(
packed_qkv: torch.Tensor,
packed_weight: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
q, k = packed_qkv[:, :, :2, :].unbind(dim=2)
q_weight, k_weight = packed_weight.unbind(dim=0)
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
q = F.rms_norm(q, (q.shape[-1], ), weight=q_weight, eps=EPS).to(q.dtype)
k = F.rms_norm(k, (k.shape[-1], ), weight=k_weight, eps=EPS).to(k.dtype)
return q, k
def quack_joint_qk(
packed_qkv: torch.Tensor,
packed_weight: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
if _quack_rmsnorm is None:
raise RuntimeError("QuACK was not initialized")
# Preserve the packed projection's stride: [B, T, 3, D] -> [B, T, 2, D]
# is a view. QuACK treats the two slots as its per-head dimension.
qk = _quack_rmsnorm(
packed_qkv[:, :, :2, :],
weight=packed_weight,
eps=EPS,
)
return qk.unbind(dim=2)
def _outputs(value: torch.Tensor | Sequence[torch.Tensor]) -> tuple[torch.Tensor, ...]:
return (value, ) if isinstance(value, torch.Tensor) else tuple(value)
def _clone_leaves(values: Sequence[torch.Tensor]) -> tuple[torch.Tensor, ...]:
return tuple(value.detach().clone().requires_grad_(True) for value in values)
def _evaluate(
function: Callable[..., Any],
values: Sequence[torch.Tensor],
grad_outputs: Sequence[torch.Tensor],
) -> tuple[tuple[torch.Tensor, ...], tuple[torch.Tensor, ...]]:
leaves = _clone_leaves(values)
outputs = _outputs(function(*leaves))
torch.autograd.backward(outputs, tuple(grad_outputs))
gradients = tuple(value.grad.detach().clone() for value in leaves)
return tuple(output.detach().clone() for output in outputs), gradients
def _comparison(
name: str,
actual: torch.Tensor,
expected: torch.Tensor,
*,
reduction_gradient: bool,
) -> dict[str, Any]:
if reduction_gradient:
atol = max(
0.1,
float(2 * torch.finfo(torch.bfloat16).eps * expected.float().abs().max()),
)
else:
atol = 0.1
rtol = 1e-3
torch.testing.assert_close(actual, expected, atol=atol, rtol=rtol)
absolute = (actual.float() - expected.float()).abs()
relative = absolute / expected.float().abs().clamp_min(1e-6)
return {
"name": name,
"shape": list(actual.shape),
"dtype": str(actual.dtype),
"atol": atol,
"rtol": rtol,
"max_abs": float(absolute.max()),
"max_rel": float(relative.max()),
"exact_fraction": float((actual == expected).float().mean()),
}
def _check_parity(
name: str,
reference: Callable[..., Any],
candidate: Callable[..., Any],
values: Sequence[torch.Tensor],
grad_outputs: Sequence[torch.Tensor],
) -> list[dict[str, Any]]:
reference_outputs, reference_gradients = _evaluate(reference, values, grad_outputs)
candidate_outputs, candidate_gradients = _evaluate(candidate, values, grad_outputs)
if len(reference_outputs) != len(candidate_outputs):
raise RuntimeError(f"{name}: output arity changed")
rows = [
_comparison(
f"{name}.output.{index}",
candidate_output,
reference_output,
reduction_gradient=False,
)
for index, (candidate_output, reference_output) in enumerate(
zip(candidate_outputs, reference_outputs, strict=True)
)
]
for index, (candidate_gradient, reference_gradient) in enumerate(
zip(candidate_gradients, reference_gradients, strict=True)
):
rows.append(
_comparison(
f"{name}.gradient.{index}",
candidate_gradient,
reference_gradient,
reduction_gradient=index > 0,
)
)
return rows
def _iteration(
function: Callable[..., Any],
values: Sequence[torch.Tensor],
grad_outputs: Sequence[torch.Tensor],
) -> None:
for value in values:
value.grad = None
torch.autograd.backward(_outputs(function(*values)), tuple(grad_outputs))
def _time_ms(
function: Callable[..., Any],
values: Sequence[torch.Tensor],
grad_outputs: Sequence[torch.Tensor],
*,
warmup: int,
repeats: int,
) -> float:
for _ in range(warmup):
_iteration(function, values, grad_outputs)
torch.cuda.synchronize()
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(repeats):
_iteration(function, values, grad_outputs)
end.record()
end.synchronize()
return float(start.elapsed_time(end) / repeats)
def _profile_once(
function: Callable[..., Any],
values: Sequence[torch.Tensor],
grad_outputs: Sequence[torch.Tensor],
) -> dict[str, Any]:
with torch.profiler.profile(
activities=[torch.profiler.ProfilerActivity.CPU, torch.profiler.ProfilerActivity.CUDA],
record_shapes=False,
profile_memory=False,
with_stack=False,
) as profiler:
_iteration(function, values, grad_outputs)
torch.cuda.synchronize()
kernels: dict[str, list[float]] = {}
for event in profiler.events():
if not str(getattr(event, "device_type", "")).endswith("CUDA"):
continue
row = kernels.setdefault(event.name, [0.0, 0.0])
row[0] += 1
row[1] += float(
getattr(event, "self_device_time_total", None)
or getattr(event, "self_cuda_time_total", None)
or 0.0
)
top = sorted(kernels.items(), key=lambda item: item[1][1], reverse=True)[:8]
return {
"launches": int(sum(row[0] for row in kernels.values())),
"summed_cuda_ms": sum(row[1] for row in kernels.values()) / 1000.0,
"top_kernels": [{
"name": kernel_name,
"calls": int(row[0]),
"cuda_ms": row[1] / 1000.0,
} for kernel_name, row in top],
}
def _benchmark(
name: str,
reference: Callable[..., Any],
candidate: Callable[..., Any],
values: Sequence[torch.Tensor],
grad_outputs: Sequence[torch.Tensor],
*,
calls_per_step: int,
warmup: int,
repeats: int,
) -> dict[str, Any]:
reference_values = _clone_leaves(values)
candidate_values = _clone_leaves(values)
reference_a = _time_ms(
reference, reference_values, grad_outputs, warmup=warmup, repeats=repeats)
candidate_ms = _time_ms(
candidate, candidate_values, grad_outputs, warmup=warmup, repeats=repeats)
reference_b = _time_ms(
reference, reference_values, grad_outputs, warmup=0, repeats=repeats)
midpoint = statistics.mean((reference_a, reference_b))
delta = midpoint - candidate_ms
return {
"name": name,
"calls_per_training_step": calls_per_step,
"control_a_ms_per_call": reference_a,
"candidate_ms_per_call": candidate_ms,
"control_b_ms_per_call": reference_b,
"control_midpoint_ms_per_call": midpoint,
"control_drift_percent": abs(reference_b - reference_a) / midpoint * 100,
"candidate_speedup_percent": delta / midpoint * 100,
"projected_step_saving_ms": delta * calls_per_step,
"profiles": {
"control": _profile_once(reference, reference_values, grad_outputs),
"candidate": _profile_once(candidate, candidate_values, grad_outputs),
},
}
def main() -> None:
global _quack_rmsnorm
parser = argparse.ArgumentParser()
parser.add_argument("--batch-size", type=int, default=DEFAULT_BATCH)
parser.add_argument("--tokens", type=int, default=DEFAULT_TOKENS)
parser.add_argument("--dim", type=int, default=DEFAULT_DIM)
parser.add_argument("--blocks", type=int, default=DEFAULT_BLOCKS)
parser.add_argument("--warmup", type=int, default=10)
parser.add_argument("--repeats", type=int, default=50)
parser.add_argument("--trace-norm-ms", type=float, default=DEFAULT_TRACE_NORM_MS)
parser.add_argument("--step-ms", type=float, default=DEFAULT_STEP_MS)
args = parser.parse_args()
if min(args.batch_size, args.tokens, args.dim, args.blocks, args.warmup, args.repeats) < 1:
parser.error("shape, block, warmup, and repeat values must be positive")
if not torch.cuda.is_available():
raise RuntimeError("this gate requires CUDA")
import quack
from quack import rmsnorm
_quack_rmsnorm = rmsnorm
device = torch.device("cuda", int(torch.cuda.current_device()))
generator = torch.Generator(device=device).manual_seed(20260722)
shape = (args.batch_size, args.tokens, args.dim)
x = torch.randn(shape, device=device, dtype=torch.bfloat16, generator=generator)
scale = torch.randn(
(args.batch_size, 1, args.dim),
device=device,
dtype=torch.bfloat16,
generator=generator,
) * 0.1
shift = torch.randn(
(args.batch_size, 1, args.dim),
device=device,
dtype=torch.bfloat16,
generator=generator,
) * 0.1
adaln_grad = torch.randn(shape, device=device, dtype=torch.bfloat16, generator=generator)
packed_qkv = torch.randn(
(args.batch_size, args.tokens, 3, args.dim),
device=device,
dtype=torch.bfloat16,
generator=generator,
)
packed_weight = torch.randn(
(2, args.dim),
device=device,
dtype=torch.bfloat16,
generator=generator,
)
q_grad = torch.randn(shape, device=device, dtype=torch.bfloat16, generator=generator)
k_grad = torch.randn(shape, device=device, dtype=torch.bfloat16, generator=generator)
parity = []
parity.extend(
_check_parity(
"adaln",
torch_adaln,
quack_adaln,
(x, scale, shift),
(adaln_grad, ),
)
)
parity.extend(
_check_parity(
"joint_qk",
torch_joint_qk,
quack_joint_qk,
(packed_qkv, packed_weight),
(q_grad, k_grad),
)
)
compiled_torch_adaln = torch.compile(torch_adaln, fullgraph=True, dynamic=False)
compiled_quack_adaln = torch.compile(quack_adaln, fullgraph=True, dynamic=False)
compiled_torch_joint_qk = torch.compile(torch_joint_qk, fullgraph=True, dynamic=False)
compiled_quack_joint_qk = torch.compile(quack_joint_qk, fullgraph=True, dynamic=False)
benchmarks = [
_benchmark(
"adaln",
compiled_torch_adaln,
compiled_quack_adaln,
(x, scale, shift),
(adaln_grad, ),
calls_per_step=2 * args.blocks,
warmup=args.warmup,
repeats=args.repeats,
),
_benchmark(
"joint_qk",
compiled_torch_joint_qk,
compiled_quack_joint_qk,
(packed_qkv, packed_weight),
(q_grad, k_grad),
calls_per_step=args.blocks,
warmup=args.warmup,
repeats=args.repeats,
),
]
projected_raw = sum(row["projected_step_saving_ms"] for row in benchmarks)
projected_capped = min(max(projected_raw, 0.0), args.trace_norm_ms)
payload = {
"environment": {
"device": torch.cuda.get_device_name(device),
"capability": list(torch.cuda.get_device_capability(device)),
"torch": torch.__version__,
"quack": getattr(quack, "__version__", "unknown"),
},
"shape_contract": {
"video_latent": [args.batch_size, 128, 11, 15, 26],
"video_tokens": list(shape),
"adaln_scale_shift": [args.batch_size, 1, args.dim],
"packed_qkv": [args.batch_size, args.tokens, 3 * args.dim],
"joint_qk_view": [args.batch_size, args.tokens, 2, args.dim],
"joint_qk_weight": [2, args.dim],
"blocks": args.blocks,
},
"parity": parity,
"benchmarks": benchmarks,
"projection": {
"trace_norm_adaln_hard_ceiling_ms": args.trace_norm_ms,
"raw_candidate_saving_ms": projected_raw,
"capped_candidate_saving_ms": projected_capped,
"baseline_step_ms": args.step_ms,
"projected_step_ms": args.step_ms - projected_capped,
"projected_step_speedup_percent": projected_capped / args.step_ms * 100,
"material_gate_ms": 2.0,
"passes_material_gate": projected_capped >= 2.0,
},
}
if not math.isfinite(projected_raw):
raise RuntimeError("non-finite timing projection")
print("LTX2_QUACK_NORM_GATE " + json.dumps(payload, sort_keys=True), flush=True)
if __name__ == "__main__":
main()
@@ -0,0 +1,234 @@
#!/usr/bin/env python3
"""Scratch GB200 gate for the exact LTX-2 overfit text-attention shape."""
from __future__ import annotations
import argparse
import json
import math
import statistics
from typing import Callable
import torch
import torch.nn.functional as F
TensorFn = Callable[[torch.Tensor, torch.Tensor, torch.Tensor], torch.Tensor]
def sdpa(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
return F.scaled_dot_product_attention(
q.transpose(1, 2),
k.transpose(1, 2),
v.transpose(1, 2),
attn_mask=None,
dropout_p=0.0,
is_causal=False,
scale=None,
).transpose(1, 2)
def metrics(actual: torch.Tensor, expected: torch.Tensor) -> dict[str, float]:
actual_f = actual.float()
expected_f = expected.float()
diff = actual_f - expected_f
expected_norm = torch.linalg.vector_norm(expected_f)
return {
"max_abs": float(diff.abs().max()),
"rmse": float(torch.sqrt(torch.mean(diff.square()))),
"relative_l2": float(torch.linalg.vector_norm(diff) / expected_norm.clamp_min(1e-20)),
"cosine": float(F.cosine_similarity(actual_f.flatten(), expected_f.flatten(), dim=0)),
}
def run_once(
fn: TensorFn,
bases: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
grad_out: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
q, k, v = (tensor.detach().clone().requires_grad_(True) for tensor in bases)
out = fn(q, k, v)
out.backward(grad_out)
assert q.grad is not None and k.grad is not None and v.grad is not None
return out.detach(), q.grad.detach(), k.grad.detach(), v.grad.detach()
def time_segment(
fn: TensorFn,
bases: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
grad_out: torch.Tensor,
warmup: int,
repeats: int,
) -> dict[str, float | list[float]]:
q, k, v = (tensor.detach().clone().requires_grad_(True) for tensor in bases)
for _ in range(warmup):
q.grad = k.grad = v.grad = None
fn(q, k, v).backward(grad_out)
torch.cuda.synchronize()
starts = [torch.cuda.Event(enable_timing=True) for _ in range(repeats)]
ends = [torch.cuda.Event(enable_timing=True) for _ in range(repeats)]
for start, end in zip(starts, ends, strict=True):
q.grad = k.grad = v.grad = None
start.record()
fn(q, k, v).backward(grad_out)
end.record()
torch.cuda.synchronize()
values = [float(start.elapsed_time(end)) for start, end in zip(starts, ends, strict=True)]
return {
"median_ms": statistics.median(values),
"min_ms": min(values),
"max_ms": max(values),
"samples_ms": values,
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--batch", type=int, choices=(1, 3), required=True)
parser.add_argument("--warmup", type=int, default=8)
parser.add_argument("--repeats", type=int, default=31)
args = parser.parse_args()
torch.manual_seed(20260722 + args.batch)
torch.cuda.manual_seed_all(20260722 + args.batch)
device = torch.device("cuda:0")
dtype = torch.bfloat16
q_shape = (args.batch, 11 * 15 * 26, 32, 128)
kv_shape = (args.batch, 1024, 32, 128)
bases = (
torch.randn(q_shape, device=device, dtype=dtype),
torch.randn(kv_shape, device=device, dtype=dtype),
torch.randn(kv_shape, device=device, dtype=dtype),
)
grad_out = torch.randn(q_shape, device=device, dtype=dtype)
from flash_attn import flash_attn_func as fa2_func
from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func as current_fa4
def fa2(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
return fa2_func(q, k, v, dropout_p=0.0, softmax_scale=None, causal=False, deterministic=False)
def make_fa4_split(num_splits: int) -> TensorFn:
class _FA4Split(torch.autograd.Function):
@staticmethod
def forward(ctx, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
out, lse = _flash_attn_fwd(
q,
k,
v,
softmax_scale=None,
causal=False,
window_size_left=None,
window_size_right=None,
softcap=0.0,
num_splits=num_splits,
pack_gqa=None,
return_lse=True,
)[:2]
ctx.save_for_backward(q, k, v, out, lse)
return out
@staticmethod
def backward(ctx, grad_out: torch.Tensor):
q, k, v, out, lse = ctx.saved_tensors
return _flash_attn_bwd(
q,
k,
v,
out,
grad_out,
lse,
softmax_scale=None,
causal=False,
softcap=0.0,
window_size_left=None,
window_size_right=None,
deterministic=False,
)
return _FA4Split.apply
functions: dict[str, TensorFn] = {
"fa4_current_splits1": current_fa4,
"fa2": fa2,
"sdpa": sdpa,
**{f"fa4_splits{splits}": make_fa4_split(splits) for splits in (1, 2, 4, 8)},
}
header = {
"batch": args.batch,
"q_shape": q_shape,
"k_shape": kv_shape,
"v_shape": kv_shape,
"dtype": str(dtype),
"causal": False,
"mask": None,
"dropout_p": 0.0,
"scale": 1.0 / math.sqrt(q_shape[-1]),
"device": torch.cuda.get_device_name(device),
"capability": torch.cuda.get_device_capability(device),
"warmup": args.warmup,
"repeats": args.repeats,
}
print("GATE_HEADER " + json.dumps(header, sort_keys=True), flush=True)
reference = run_once(sdpa, bases, grad_out)
parity: dict[str, object] = {}
for name, fn in functions.items():
try:
result = run_once(fn, bases, grad_out)
values = {
label: metrics(actual, expected)
for label, actual, expected in zip(("out", "dq", "dk", "dv"), result, reference, strict=True)
}
finite = all(torch.isfinite(tensor).all().item() for tensor in result)
parity[name] = {"finite": bool(finite), "vs_sdpa": values}
print("GATE_PARITY " + json.dumps({"batch": args.batch, "backend": name, **parity[name]}, sort_keys=True),
flush=True)
del result
except Exception as exc:
torch.cuda.synchronize()
parity[name] = {"error": repr(exc)}
print("GATE_PARITY " + json.dumps({"batch": args.batch, "backend": name, **parity[name]}, sort_keys=True),
flush=True)
candidates = ("fa2", "sdpa", "fa4_splits1", "fa4_splits2", "fa4_splits4", "fa4_splits8")
timing: dict[str, object] = {}
for candidate in candidates:
if "error" in parity.get(candidate, {}):
continue
try:
a = time_segment(current_fa4, bases, grad_out, args.warmup, args.repeats)
x = time_segment(functions[candidate], bases, grad_out, args.warmup, args.repeats)
b = time_segment(current_fa4, bases, grad_out, args.warmup, args.repeats)
midpoint_ms = (float(a["median_ms"]) + float(b["median_ms"])) / 2.0
candidate_ms = float(x["median_ms"])
row = {
"batch": args.batch,
"candidate": candidate,
"a_current": a,
"x_candidate": x,
"b_current": b,
"control_midpoint_ms": midpoint_ms,
"control_drift_pct": 100.0 * abs(float(b["median_ms"]) - float(a["median_ms"])) / midpoint_ms,
"delta_ms_per_call": candidate_ms - midpoint_ms,
"delta_pct": 100.0 * (candidate_ms - midpoint_ms) / midpoint_ms,
"projected_48_block_delta_ms": 48.0 * (candidate_ms - midpoint_ms),
}
timing[candidate] = row
print("GATE_TIMING " + json.dumps(row, sort_keys=True), flush=True)
except Exception as exc:
torch.cuda.synchronize()
timing[candidate] = {"error": repr(exc)}
print("GATE_TIMING " + json.dumps({"batch": args.batch, "candidate": candidate, "error": repr(exc)},
sort_keys=True), flush=True)
print("GATE_RESULT " + json.dumps({"header": header, "parity": parity, "timing": timing}, sort_keys=True),
flush=True)
if __name__ == "__main__":
main()
@@ -0,0 +1,568 @@
#!/usr/bin/env python3
"""Direct Triton training gate for LTX-2 RMSNorm/AdaLN hot paths."""
from __future__ import annotations
import argparse
from collections.abc import Callable, Sequence
import json
import math
import statistics
from typing import Any
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
BATCH = 1
TOKENS = 11 * (480 // 32) * (832 // 32) # 4290
DIM = 32 * 128 # 4096
BLOCKS = 48
EPS = 1e-6
TRACE_NORM_MS = 27.723
STEP_MS = 427.432493
@triton.jit
def _rms_affine_fwd(
x_ptr,
weight_ptr,
bias_ptr,
out_ptr,
rstd_ptr,
stride_x_m,
stride_x_h,
stride_x_n,
eps,
H: tl.constexpr,
N: tl.constexpr,
BLOCK_N: tl.constexpr,
HAS_BIAS: tl.constexpr,
ADD_ONE: tl.constexpr,
):
row = tl.program_id(0)
m = row // H
h = row - m * H
cols = tl.arange(0, BLOCK_N)
mask = cols < N
x = tl.load(
x_ptr + m * stride_x_m + h * stride_x_h + cols * stride_x_n,
mask=mask,
other=0.0,
).to(tl.float32)
rstd = tl.rsqrt(tl.sum(x * x, axis=0) / N + eps)
weight = tl.load(weight_ptr + h * N + cols, mask=mask, other=0.0).to(tl.float32)
if ADD_ONE:
weight += 1.0
out = x * rstd * weight
if HAS_BIAS:
bias = tl.load(bias_ptr + h * N + cols, mask=mask, other=0.0).to(tl.float32)
out += bias
tl.store(out_ptr + row * N + cols, out, mask=mask)
tl.store(rstd_ptr + row, rstd)
@triton.jit
def _rms_affine_dx(
x_ptr,
weight_ptr,
dout_ptr,
rstd_ptr,
dx_ptr,
stride_x_m,
stride_x_h,
stride_x_n,
H: tl.constexpr,
N: tl.constexpr,
BLOCK_N: tl.constexpr,
ADD_ONE: tl.constexpr,
):
row = tl.program_id(0)
m = row // H
h = row - m * H
cols = tl.arange(0, BLOCK_N)
mask = cols < N
x = tl.load(
x_ptr + m * stride_x_m + h * stride_x_h + cols * stride_x_n,
mask=mask,
other=0.0,
).to(tl.float32)
dout = tl.load(dout_ptr + row * N + cols, mask=mask, other=0.0).to(tl.float32)
weight = tl.load(weight_ptr + h * N + cols, mask=mask, other=0.0).to(tl.float32)
if ADD_ONE:
weight += 1.0
rstd = tl.load(rstd_ptr + row).to(tl.float32)
x_hat = x * rstd
weight_dout = weight * dout
correction = tl.sum(x_hat * weight_dout, axis=0) / N
dx = (weight_dout - x_hat * correction) * rstd
tl.store(dx_ptr + row * N + cols, dx, mask=mask)
@triton.jit
def _rms_affine_param_grad(
x_ptr,
dout_ptr,
rstd_ptr,
dweight_ptr,
dbias_ptr,
M,
stride_x_m,
stride_x_h,
stride_x_n,
H: tl.constexpr,
N: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
HAS_BIAS: tl.constexpr,
):
h = tl.program_id(0)
col_block = tl.program_id(1)
cols = col_block * BLOCK_N + tl.arange(0, BLOCK_N)
col_mask = cols < N
dweight = tl.zeros((BLOCK_N, ), dtype=tl.float32)
dbias = tl.zeros((BLOCK_N, ), dtype=tl.float32)
for row_start in tl.range(0, M, BLOCK_M):
rows = row_start + tl.arange(0, BLOCK_M)
row_mask = rows < M
mask = row_mask[:, None] & col_mask[None, :]
x = tl.load(
x_ptr
+ rows[:, None] * stride_x_m
+ h * stride_x_h
+ cols[None, :] * stride_x_n,
mask=mask,
other=0.0,
).to(tl.float32)
dout = tl.load(
dout_ptr + (rows[:, None] * H + h) * N + cols[None, :],
mask=mask,
other=0.0,
).to(tl.float32)
rstd = tl.load(
rstd_ptr + rows * H + h,
mask=row_mask,
other=0.0,
).to(tl.float32)
dweight += tl.sum(dout * x * rstd[:, None], axis=0)
if HAS_BIAS:
dbias += tl.sum(dout, axis=0)
tl.store(dweight_ptr + h * N + cols, dweight, mask=col_mask)
if HAS_BIAS:
tl.store(dbias_ptr + h * N + cols, dbias, mask=col_mask)
class _FusedRMSAffine(torch.autograd.Function):
@staticmethod
def forward(
ctx: Any,
x: torch.Tensor,
weight: torch.Tensor,
bias: torch.Tensor | None,
add_one: bool,
) -> torch.Tensor:
if x.ndim != 3 or weight.ndim != 2:
raise ValueError(f"expected x[M,H,N], weight[H,N], got {x.shape}, {weight.shape}")
m, h, n = x.shape
if weight.shape != (h, n):
raise ValueError(f"weight {weight.shape} does not match x {x.shape}")
if bias is not None and bias.shape != weight.shape:
raise ValueError(f"bias {bias.shape} does not match weight {weight.shape}")
if x.dtype != torch.bfloat16 or weight.dtype != torch.bfloat16:
raise TypeError("scratch kernel is specialized to BF16 activations and affine tensors")
if bias is not None and bias.dtype != torch.bfloat16:
raise TypeError("scratch kernel is specialized to BF16 bias")
if x.stride(-1) != 1 or not weight.is_contiguous() or (bias is not None and not bias.is_contiguous()):
raise ValueError("last input dimension and affine tensors must be contiguous")
block_n = triton.next_power_of_2(n)
if block_n > 65536:
raise ValueError(f"unsupported hidden dimension: {n}")
out = torch.empty((m, h, n), device=x.device, dtype=x.dtype)
rstd = torch.empty((m, h), device=x.device, dtype=torch.float32)
_rms_affine_fwd[(m * h, )](
x,
weight,
weight if bias is None else bias,
out,
rstd,
x.stride(0),
x.stride(1),
x.stride(2),
EPS,
H=h,
N=n,
BLOCK_N=block_n,
HAS_BIAS=bias is not None,
ADD_ONE=bool(add_one),
num_warps=8,
)
ctx.save_for_backward(x, weight, rstd)
ctx.has_bias = bias is not None
ctx.add_one = bool(add_one)
return out
@staticmethod
def backward(
ctx: Any,
dout: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None, None]:
x, weight, rstd = ctx.saved_tensors
m, h, n = x.shape
dout = dout.contiguous()
dx = torch.empty((m, h, n), device=x.device, dtype=x.dtype)
dweight = torch.empty_like(weight)
dbias = torch.empty_like(weight) if ctx.has_bias else None
block_n = triton.next_power_of_2(n)
_rms_affine_dx[(m * h, )](
x,
weight,
dout,
rstd,
dx,
x.stride(0),
x.stride(1),
x.stride(2),
H=h,
N=n,
BLOCK_N=block_n,
ADD_ONE=ctx.add_one,
num_warps=8,
)
grad_block_n = 32
_rms_affine_param_grad[(h, triton.cdiv(n, grad_block_n))](
x,
dout,
rstd,
dweight,
dweight if dbias is None else dbias,
m,
x.stride(0),
x.stride(1),
x.stride(2),
H=h,
N=n,
BLOCK_M=32,
BLOCK_N=grad_block_n,
HAS_BIAS=ctx.has_bias,
num_warps=4,
)
return dx, dweight, dbias, None
def triton_adaln(
x: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
) -> torch.Tensor:
x_tbd = x.transpose(0, 1)
return _FusedRMSAffine.apply(
x_tbd,
scale.squeeze(1),
shift.squeeze(1),
True,
).transpose(0, 1)
def torch_adaln(
x: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
) -> torch.Tensor:
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
normalized = F.rms_norm(x, (x.shape[-1], ), eps=EPS).to(x.dtype)
return normalized * (1 + scale) + shift
def triton_joint_qk(
packed_qkv: torch.Tensor,
packed_weight: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
batch, tokens, _, dim = packed_qkv.shape
qk = packed_qkv[:, :, :2, :].reshape(batch * tokens, 2, dim)
qk = _FusedRMSAffine.apply(qk, packed_weight, None, False)
return qk.reshape(batch, tokens, 2, dim).unbind(dim=2)
def torch_joint_qk(
packed_qkv: torch.Tensor,
packed_weight: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
q, k = packed_qkv[:, :, :2, :].unbind(dim=2)
q_weight, k_weight = packed_weight.unbind(dim=0)
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
q = F.rms_norm(q, (q.shape[-1], ), weight=q_weight, eps=EPS).to(q.dtype)
k = F.rms_norm(k, (k.shape[-1], ), weight=k_weight, eps=EPS).to(k.dtype)
return q, k
def _as_tuple(value: torch.Tensor | Sequence[torch.Tensor]) -> tuple[torch.Tensor, ...]:
return (value, ) if isinstance(value, torch.Tensor) else tuple(value)
def _clone_leaves(values: Sequence[torch.Tensor]) -> tuple[torch.Tensor, ...]:
return tuple(value.detach().clone().requires_grad_(True) for value in values)
def _run(
function: Callable[..., Any],
values: Sequence[torch.Tensor],
grad_outputs: Sequence[torch.Tensor],
) -> tuple[tuple[torch.Tensor, ...], tuple[torch.Tensor, ...]]:
leaves = _clone_leaves(values)
outputs = _as_tuple(function(*leaves))
torch.autograd.backward(outputs, tuple(grad_outputs))
return (
tuple(output.detach().clone() for output in outputs),
tuple(value.grad.detach().clone() for value in leaves),
)
def _compare(
name: str,
actual: torch.Tensor,
expected: torch.Tensor,
*,
reduction_gradient: bool,
) -> dict[str, Any]:
atol = 0.1
if reduction_gradient:
atol = max(
atol,
float(2 * torch.finfo(torch.bfloat16).eps * expected.float().abs().max()),
)
rtol = 1e-3
torch.testing.assert_close(actual, expected, atol=atol, rtol=rtol)
absolute = (actual.float() - expected.float()).abs()
return {
"name": name,
"shape": list(actual.shape),
"dtype": str(actual.dtype),
"atol": atol,
"rtol": rtol,
"max_abs": float(absolute.max()),
"exact_fraction": float((actual == expected).float().mean()),
}
def _parity(
name: str,
reference: Callable[..., Any],
candidate: Callable[..., Any],
values: Sequence[torch.Tensor],
grad_outputs: Sequence[torch.Tensor],
) -> list[dict[str, Any]]:
expected_outputs, expected_gradients = _run(reference, values, grad_outputs)
actual_outputs, actual_gradients = _run(candidate, values, grad_outputs)
rows = []
for index, (actual, expected) in enumerate(zip(actual_outputs, expected_outputs, strict=True)):
rows.append(_compare(f"{name}.output.{index}", actual, expected, reduction_gradient=False))
for index, (actual, expected) in enumerate(zip(actual_gradients, expected_gradients, strict=True)):
rows.append(_compare(f"{name}.gradient.{index}", actual, expected, reduction_gradient=index > 0))
return rows
def _iteration(
function: Callable[..., Any],
values: Sequence[torch.Tensor],
grad_outputs: Sequence[torch.Tensor],
) -> None:
for value in values:
value.grad = None
torch.autograd.backward(_as_tuple(function(*values)), tuple(grad_outputs))
def _time_ms(
function: Callable[..., Any],
values: Sequence[torch.Tensor],
grad_outputs: Sequence[torch.Tensor],
*,
warmup: int,
repeats: int,
) -> float:
for _ in range(warmup):
_iteration(function, values, grad_outputs)
torch.cuda.synchronize()
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(repeats):
_iteration(function, values, grad_outputs)
end.record()
end.synchronize()
return float(start.elapsed_time(end) / repeats)
def _profile(
function: Callable[..., Any],
values: Sequence[torch.Tensor],
grad_outputs: Sequence[torch.Tensor],
) -> dict[str, Any]:
with torch.profiler.profile(
activities=[torch.profiler.ProfilerActivity.CPU, torch.profiler.ProfilerActivity.CUDA],
record_shapes=False,
profile_memory=False,
with_stack=False,
) as profiler:
_iteration(function, values, grad_outputs)
torch.cuda.synchronize()
kernels: dict[str, list[float]] = {}
for event in profiler.events():
if not str(getattr(event, "device_type", "")).endswith("CUDA"):
continue
row = kernels.setdefault(event.name, [0.0, 0.0])
row[0] += 1
row[1] += float(
getattr(event, "self_device_time_total", None)
or getattr(event, "self_cuda_time_total", None)
or 0.0
)
return {
"launches": int(sum(row[0] for row in kernels.values())),
"summed_cuda_ms": sum(row[1] for row in kernels.values()) / 1000.0,
"top_kernels": [{
"name": name,
"calls": int(row[0]),
"cuda_ms": row[1] / 1000.0,
} for name, row in sorted(kernels.items(), key=lambda item: item[1][1], reverse=True)[:8]],
}
def _benchmark(
name: str,
reference: Callable[..., Any],
candidate: Callable[..., Any],
values: Sequence[torch.Tensor],
grad_outputs: Sequence[torch.Tensor],
*,
calls_per_step: int,
warmup: int,
repeats: int,
) -> dict[str, Any]:
reference_values = _clone_leaves(values)
candidate_values = _clone_leaves(values)
control_a = _time_ms(reference, reference_values, grad_outputs, warmup=warmup, repeats=repeats)
candidate_ms = _time_ms(candidate, candidate_values, grad_outputs, warmup=warmup, repeats=repeats)
control_b = _time_ms(reference, reference_values, grad_outputs, warmup=0, repeats=repeats)
midpoint = statistics.mean((control_a, control_b))
delta = midpoint - candidate_ms
return {
"name": name,
"calls_per_step": calls_per_step,
"control_a_ms_per_call": control_a,
"candidate_ms_per_call": candidate_ms,
"control_b_ms_per_call": control_b,
"control_midpoint_ms_per_call": midpoint,
"control_drift_percent": abs(control_a - control_b) / midpoint * 100,
"candidate_speedup_percent": delta / midpoint * 100,
"projected_step_saving_ms": delta * calls_per_step,
"profiles": {
"control": _profile(reference, reference_values, grad_outputs),
"candidate": _profile(candidate, candidate_values, grad_outputs),
},
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--warmup", type=int, default=10)
parser.add_argument("--repeats", type=int, default=50)
args = parser.parse_args()
if args.warmup < 1 or args.repeats < 1:
parser.error("--warmup and --repeats must be positive")
if not torch.cuda.is_available():
raise RuntimeError("this benchmark requires CUDA")
device = torch.device("cuda", torch.cuda.current_device())
generator = torch.Generator(device=device).manual_seed(20260722)
shape = (BATCH, TOKENS, DIM)
x = torch.randn(shape, device=device, dtype=torch.bfloat16, generator=generator)
scale = torch.randn((BATCH, 1, DIM), device=device, dtype=torch.bfloat16, generator=generator) * 0.1
shift = torch.randn((BATCH, 1, DIM), device=device, dtype=torch.bfloat16, generator=generator) * 0.1
adaln_grad = torch.randn(shape, device=device, dtype=torch.bfloat16, generator=generator)
packed_qkv = torch.randn(
(BATCH, TOKENS, 3, DIM),
device=device,
dtype=torch.bfloat16,
generator=generator,
)
packed_weight = torch.randn((2, DIM), device=device, dtype=torch.bfloat16, generator=generator)
q_grad = torch.randn(shape, device=device, dtype=torch.bfloat16, generator=generator)
k_grad = torch.randn(shape, device=device, dtype=torch.bfloat16, generator=generator)
parity = []
parity.extend(_parity("adaln", torch_adaln, triton_adaln, (x, scale, shift), (adaln_grad, )))
parity.extend(
_parity(
"joint_qk",
torch_joint_qk,
triton_joint_qk,
(packed_qkv, packed_weight),
(q_grad, k_grad),
)
)
compiled_adaln = torch.compile(torch_adaln, fullgraph=True, dynamic=False)
compiled_joint_qk = torch.compile(torch_joint_qk, fullgraph=True, dynamic=False)
benchmarks = [
_benchmark(
"adaln",
compiled_adaln,
triton_adaln,
(x, scale, shift),
(adaln_grad, ),
calls_per_step=2 * BLOCKS,
warmup=args.warmup,
repeats=args.repeats,
),
_benchmark(
"joint_qk",
compiled_joint_qk,
triton_joint_qk,
(packed_qkv, packed_weight),
(q_grad, k_grad),
calls_per_step=BLOCKS,
warmup=args.warmup,
repeats=args.repeats,
),
]
raw_saving = sum(row["projected_step_saving_ms"] for row in benchmarks)
capped_saving = min(max(raw_saving, 0.0), TRACE_NORM_MS)
if not math.isfinite(raw_saving):
raise RuntimeError("non-finite timing projection")
payload = {
"environment": {
"device": torch.cuda.get_device_name(device),
"capability": list(torch.cuda.get_device_capability(device)),
"torch": torch.__version__,
"triton": triton.__version__,
},
"shape_contract": {
"latent": [BATCH, 128, 11, 15, 26],
"hidden": list(shape),
"adaln_scale_shift": [BATCH, 1, DIM],
"packed_qkv_storage": [BATCH, TOKENS, 3 * DIM],
"joint_qk_view": [BATCH, TOKENS, 2, DIM],
"blocks": BLOCKS,
},
"accumulation": "FP32 reductions and arithmetic; BF16 outputs and returned gradients",
"parity": parity,
"benchmarks": benchmarks,
"projection": {
"raw_saving_ms_per_step": raw_saving,
"trace_capped_saving_ms_per_step": capped_saving,
"trace_family_ceiling_ms_per_step": TRACE_NORM_MS,
"baseline_step_ms": STEP_MS,
"projected_step_ms": STEP_MS - capped_saving,
"projected_latency_reduction_percent": capped_saving / STEP_MS * 100,
"material_threshold_ms": 2.0,
"passes_material_gate": capped_saving >= 2.0,
},
}
print("LTX2_TRITON_NORM_GATE " + json.dumps(payload, sort_keys=True), flush=True)
if __name__ == "__main__":
main()
@@ -0,0 +1,236 @@
#!/usr/bin/env python3
"""Scratch GB200 gate for the exact LTX-2 video self-attention shape.
Compares the production FA4 CuTe path against FA2 and forced SDPA
flash/cuDNN backends at (B, 4290, 32, 128) fwd+bwd, mirroring the text
attention gate's A/X/B-per-candidate protocol. Diagnostics feed a trainer
gate; these numbers are never MFU evidence.
"""
from __future__ import annotations
import argparse
import json
import math
import statistics
from typing import Callable
import torch
import torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel
TensorFn = Callable[[torch.Tensor, torch.Tensor, torch.Tensor], torch.Tensor]
def make_sdpa(backend: SDPBackend | None) -> TensorFn:
def _sdpa(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
qt, kt, vt = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)
if backend is None:
out = F.scaled_dot_product_attention(qt, kt, vt, attn_mask=None, dropout_p=0.0, is_causal=False)
else:
with sdpa_kernel(backend):
out = F.scaled_dot_product_attention(qt, kt, vt, attn_mask=None, dropout_p=0.0, is_causal=False)
return out.transpose(1, 2)
return _sdpa
def metrics(actual: torch.Tensor, expected: torch.Tensor) -> dict[str, float]:
actual_f = actual.float()
expected_f = expected.float()
diff = actual_f - expected_f
expected_norm = torch.linalg.vector_norm(expected_f)
return {
"max_abs": float(diff.abs().max()),
"rmse": float(torch.sqrt(torch.mean(diff.square()))),
"relative_l2": float(torch.linalg.vector_norm(diff) / expected_norm.clamp_min(1e-20)),
"cosine": float(F.cosine_similarity(actual_f.flatten(), expected_f.flatten(), dim=0)),
}
def run_once(
fn: TensorFn,
bases: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
grad_out: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
q, k, v = (tensor.detach().clone().requires_grad_(True) for tensor in bases)
out = fn(q, k, v)
out.backward(grad_out)
assert q.grad is not None and k.grad is not None and v.grad is not None
return out.detach(), q.grad.detach(), k.grad.detach(), v.grad.detach()
def time_segment(
fn: TensorFn,
bases: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
grad_out: torch.Tensor,
warmup: int,
repeats: int,
) -> dict[str, float | list[float]]:
q, k, v = (tensor.detach().clone().requires_grad_(True) for tensor in bases)
for _ in range(warmup):
q.grad = k.grad = v.grad = None
fn(q, k, v).backward(grad_out)
torch.cuda.synchronize()
starts = [torch.cuda.Event(enable_timing=True) for _ in range(repeats)]
ends = [torch.cuda.Event(enable_timing=True) for _ in range(repeats)]
for start, end in zip(starts, ends, strict=True):
q.grad = k.grad = v.grad = None
start.record()
fn(q, k, v).backward(grad_out)
end.record()
torch.cuda.synchronize()
values = [float(start.elapsed_time(end)) for start, end in zip(starts, ends, strict=True)]
return {
"median_ms": statistics.median(values),
"min_ms": min(values),
"max_ms": max(values),
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--batch", type=int, choices=(1, 2, 3), required=True)
parser.add_argument("--warmup", type=int, default=8)
parser.add_argument("--repeats", type=int, default=31)
args = parser.parse_args()
torch.manual_seed(20260722 + args.batch)
torch.cuda.manual_seed_all(20260722 + args.batch)
device = torch.device("cuda:0")
dtype = torch.bfloat16
shape = (args.batch, 11 * 15 * 26, 32, 128)
bases = (
torch.randn(shape, device=device, dtype=dtype),
torch.randn(shape, device=device, dtype=dtype),
torch.randn(shape, device=device, dtype=dtype),
)
grad_out = torch.randn(shape, device=device, dtype=dtype)
from flash_attn import flash_attn_func as fa2_func
from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func as current_fa4
def fa2(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
return fa2_func(q, k, v, dropout_p=0.0, softmax_scale=None, causal=False, deterministic=False)
def make_fa4_split(num_splits: int) -> TensorFn:
class _FA4Split(torch.autograd.Function):
@staticmethod
def forward(ctx, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
out, lse = _flash_attn_fwd(
q,
k,
v,
softmax_scale=None,
causal=False,
window_size_left=None,
window_size_right=None,
softcap=0.0,
num_splits=num_splits,
pack_gqa=None,
return_lse=True,
)[:2]
ctx.save_for_backward(q, k, v, out, lse)
return out
@staticmethod
def backward(ctx, grad_out: torch.Tensor):
q, k, v, out, lse = ctx.saved_tensors
return _flash_attn_bwd(
q,
k,
v,
out,
grad_out,
lse,
softmax_scale=None,
causal=False,
softcap=0.0,
window_size_left=None,
window_size_right=None,
deterministic=False,
)
return _FA4Split.apply
functions: dict[str, TensorFn] = {
"fa4_current": current_fa4,
"fa2": fa2,
"sdpa_flash": make_sdpa(SDPBackend.FLASH_ATTENTION),
"sdpa_cudnn": make_sdpa(SDPBackend.CUDNN_ATTENTION),
**{f"fa4_splits{splits}": make_fa4_split(splits) for splits in (1, 2, 4, 8)},
}
header = {
"batch": args.batch,
"shape": shape,
"dtype": str(dtype),
"causal": False,
"scale": 1.0 / math.sqrt(shape[-1]),
"device": torch.cuda.get_device_name(device),
"capability": torch.cuda.get_device_capability(device),
"warmup": args.warmup,
"repeats": args.repeats,
"torch": torch.__version__,
"cudnn": torch.backends.cudnn.version(),
}
print("GATE_HEADER " + json.dumps(header, sort_keys=True), flush=True)
reference = run_once(make_sdpa(None), bases, grad_out)
parity: dict[str, object] = {}
for name, fn in functions.items():
try:
result = run_once(fn, bases, grad_out)
values = {
label: metrics(actual, expected)
for label, actual, expected in zip(("out", "dq", "dk", "dv"), result, reference, strict=True)
}
finite = all(torch.isfinite(tensor).all().item() for tensor in result)
parity[name] = {"finite": bool(finite), "vs_sdpa_default": values}
del result
except Exception as exc:
torch.cuda.synchronize()
parity[name] = {"error": repr(exc)}
print("GATE_PARITY " + json.dumps({"batch": args.batch, "backend": name, **parity[name]}, sort_keys=True),
flush=True)
timing: dict[str, object] = {}
for candidate in ("fa2", "sdpa_flash", "sdpa_cudnn", "fa4_splits1", "fa4_splits2", "fa4_splits4", "fa4_splits8"):
if "error" in parity.get(candidate, {}):
continue
try:
a = time_segment(current_fa4, bases, grad_out, args.warmup, args.repeats)
x = time_segment(functions[candidate], bases, grad_out, args.warmup, args.repeats)
b = time_segment(current_fa4, bases, grad_out, args.warmup, args.repeats)
midpoint_ms = (float(a["median_ms"]) + float(b["median_ms"])) / 2.0
candidate_ms = float(x["median_ms"])
row = {
"batch": args.batch,
"candidate": candidate,
"a_current_median_ms": a["median_ms"],
"x_candidate_median_ms": candidate_ms,
"b_current_median_ms": b["median_ms"],
"control_midpoint_ms": midpoint_ms,
"control_drift_pct": 100.0 * abs(float(b["median_ms"]) - float(a["median_ms"])) / midpoint_ms,
"delta_ms_per_call": candidate_ms - midpoint_ms,
"delta_pct": 100.0 * (candidate_ms - midpoint_ms) / midpoint_ms,
"projected_48_block_delta_ms": 48.0 * (candidate_ms - midpoint_ms),
}
timing[candidate] = row
except Exception as exc:
torch.cuda.synchronize()
timing[candidate] = {"error": repr(exc)}
print("GATE_TIMING " + json.dumps({"batch": args.batch, "candidate": candidate, **timing[candidate]},
sort_keys=True), flush=True)
print("GATE_RESULT " + json.dumps({"header": header, "parity": parity, "timing": timing}, sort_keys=True),
flush=True)
if __name__ == "__main__":
main()
@@ -0,0 +1,111 @@
#!/usr/bin/env python3
"""Scratch exact-shape prequantized FP8 ceiling for LTX-2 packed linears."""
import gc
import json
import statistics
import torch
import torch.nn.functional as F
LAYERS = (
("self_qkv", 48, 4290, 4096, 12288),
("video_dd", 144, 4290, 4096, 4096),
("text_kv", 48, 1024, 4096, 8192),
("ffn_up", 48, 4290, 4096, 16384),
("ffn_down", 48, 4290, 16384, 4096),
)
def timed(fn, warmup=5, samples=9, inner=5):
for _ in range(warmup):
fn()
torch.cuda.synchronize()
values = []
for _ in range(samples):
begin = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
begin.record()
for _ in range(inner):
fn()
end.record()
end.synchronize()
values.append(begin.elapsed_time(end) / inner)
return statistics.median(values)
def operands(m, k, n, phase):
if phase == "fwd":
return (torch.randn(m, k, device="cuda", dtype=torch.bfloat16),
torch.randn(n, k, device="cuda", dtype=torch.bfloat16) * 0.02)
if phase == "dgrad":
return (torch.randn(m, n, device="cuda", dtype=torch.bfloat16) * 0.01,
(torch.randn(n, k, device="cuda", dtype=torch.bfloat16) * 0.02).T.contiguous())
dy = torch.randn(m, n, device="cuda", dtype=torch.bfloat16) * 0.01
x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16)
padded = (m + 31) // 32 * 32
if padded != m:
dy = F.pad(dy, (0, 0, 0, padded - m))
x = F.pad(x, (0, 0, 0, padded - m))
return dy.T.contiguous(), x.T.contiguous()
def bench(name, count, m, k, n, phase):
lhs, rhs_rows = operands(m, k, n, phase)
lhs8 = lhs.to(torch.float8_e4m3fn)
rhs8 = rhs_rows.to(torch.float8_e4m3fn)
scale = torch.ones((), device="cuda", dtype=torch.float32)
def fp8(fast):
return torch._scaled_mm(
lhs8,
rhs8.T,
scale_a=scale,
scale_b=scale,
out_dtype=torch.bfloat16,
use_fast_accum=fast,
)
reference = lhs @ rhs_rows.T
actual = fp8(True)
if isinstance(actual, tuple):
actual = actual[0]
torch.cuda.synchronize()
diff = actual.float() - reference.float()
row = {
"kind": "case",
"name": name,
"phase": phase,
"count": count,
"logical_shape": [m, k, n],
"qmm_shape": [lhs.shape[0], lhs.shape[1], rhs_rows.shape[0]],
"bf16_ms": timed(lambda: lhs @ rhs_rows.T),
"fp8_fast_ms": timed(lambda: fp8(True)),
"fp8_accurate_ms": timed(lambda: fp8(False)),
"relative_rms": float(diff.square().mean().sqrt() /
reference.float().square().mean().sqrt()),
}
print(json.dumps(row, sort_keys=True), flush=True)
return row
if __name__ == "__main__":
assert torch.cuda.get_device_capability() == (10, 0)
torch.manual_seed(20260721)
rows = []
for layer in LAYERS:
for phase in ("fwd", "dgrad", "wgrad"):
rows.append(bench(*layer, phase))
gc.collect()
torch.cuda.empty_cache()
totals = {
tier: sum(row[f"{tier}_ms"] * row["count"] for row in rows)
for tier in ("bf16", "fp8_fast", "fp8_accurate")
}
print(json.dumps({
"kind": "aggregate",
"totals_ms_48_blocks": totals,
"speedup_fast": totals["bf16"] / totals["fp8_fast"],
"speedup_accurate": totals["bf16"] / totals["fp8_accurate"],
}, sort_keys=True))
@@ -0,0 +1,212 @@
import argparse
import gc
import json
import statistics
import torch
import torch.nn.functional as F
import transformer_engine
import transformer_engine.pytorch as te
from transformer_engine.common.recipe import (
DelayedScaling,
Format,
MXFP8BlockScaling,
NVFP4BlockScaling,
)
LAYOUTS = {
"separate": {
"h_h": (4290, 4096, 4096, 6),
"text_kv": (1024, 4096, 4096, 2),
"ffn_up": (4290, 4096, 16384, 1),
"ffn_down": (4290, 16384, 4096, 1),
},
"packed": {
"self_qkv": (4290, 4096, 12288, 1),
"h_h": (4290, 4096, 4096, 3),
"text_kv": (1024, 4096, 8192, 1),
"ffn_up": (4290, 4096, 16384, 1),
"ffn_down": (4290, 16384, 4096, 1),
},
}
BASELINE_STEP_S = 0.479309913
BASELINE_BLOCK_GEMM_S = 0.216001
MFU_PERCENT_SECONDS = 14.444115
def _recipe(name):
if name == "bf16":
return None
if name == "fp8_delayed":
return DelayedScaling(
fp8_format=Format.HYBRID,
amax_history_len=16,
amax_compute_algo="max",
)
if name == "mxfp8":
return MXFP8BlockScaling(fp8_format=Format.E4M3)
if name in ("nvfp4", "nvfp4_primary"):
return NVFP4BlockScaling()
if name == "nvfp4_1d_weight":
return NVFP4BlockScaling(disable_2d_quantization=True)
raise ValueError(name)
def _measure(name, layout_name, shape_name, warmup, repeat, bias, alignment, m_multiplier):
base_m, k, n, count = LAYOUTS[layout_name][shape_name]
m = base_m * m_multiplier
recipe = _recipe(name)
torch.manual_seed(0)
x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16, requires_grad=True)
dy = torch.randn(m, n, device="cuda", dtype=torch.bfloat16)
if recipe is None:
layer = torch.nn.Linear(k, n, bias=bias, device="cuda", dtype=torch.bfloat16)
padded_m = m
else:
if name == "nvfp4_primary":
with te.quantized_model_init(
enabled=True,
recipe=recipe,
preserve_high_precision_init_val=False,
):
layer = te.Linear(
k,
n,
bias=bias,
device="cuda",
params_dtype=torch.bfloat16,
)
else:
layer = te.Linear(k, n, bias=bias, device="cuda", params_dtype=torch.bfloat16)
# TE block recipes require the flattened row count to be aligned.
padded_m = ((m + alignment - 1) // alignment) * alignment
def step():
layer.zero_grad(set_to_none=True)
x.grad = None
if recipe is None:
y = layer(x)
grad = dy
else:
x_in = F.pad(x, (0, 0, 0, padded_m - m)) if padded_m != m else x
with te.autocast(enabled=True, recipe=recipe):
y_padded = layer(x_in, is_first_microbatch=None)
y = y_padded[:m]
grad = dy
y.backward(grad)
for _ in range(warmup):
step()
torch.cuda.synchronize()
samples = []
for _ in range(repeat):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
step()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end))
result = {
"precision": name,
"layout": layout_name,
"shape": shape_name,
"base_m": base_m,
"m": m,
"m_multiplier": m_multiplier,
"padded_m": padded_m,
"k": k,
"n": n,
"count_per_block": count,
"median_ms": statistics.median(samples),
"min_ms": min(samples),
}
del layer, x, dy
gc.collect()
torch.cuda.empty_cache()
return result
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--warmup", type=int, default=5)
parser.add_argument("--repeat", type=int, default=10)
parser.add_argument("--no-bias", action="store_true")
parser.add_argument("--alignment", type=int, default=16)
parser.add_argument("--m-multiplier", type=int, default=1)
parser.add_argument(
"--precisions",
nargs="+",
default=["bf16", "fp8_delayed", "mxfp8", "nvfp4", "nvfp4_1d_weight"],
)
args = parser.parse_args()
print(json.dumps({
"torch": torch.__version__,
"transformer_engine": transformer_engine.__version__,
"device": torch.cuda.get_device_name(),
"nvfp4": te.is_nvfp4_available(return_reason=True),
"mxfp8": te.is_mxfp8_available(return_reason=True),
}, default=str), flush=True)
results = []
for layout_name, shapes in LAYOUTS.items():
for precision in args.precisions:
for shape_name in shapes:
result = _measure(
precision,
layout_name,
shape_name,
args.warmup,
args.repeat,
not args.no_bias,
args.alignment,
args.m_multiplier,
)
results.append(result)
print(json.dumps(result), flush=True)
totals = {
layout_name: {
precision: sum(
row["median_ms"] * row["count_per_block"]
for row in results
if row["layout"] == layout_name and row["precision"] == precision
)
for precision in args.precisions
}
for layout_name in LAYOUTS
}
current_bf16 = totals["separate"]["bf16"]
summary = {}
for layout_name, layout_totals in totals.items():
layout_bf16 = layout_totals["bf16"]
summary[layout_name] = {}
for precision, weighted_ms in layout_totals.items():
linear_ratio = weighted_ms / current_bf16
projected_step_s = (
BASELINE_STEP_S - BASELINE_BLOCK_GEMM_S + BASELINE_BLOCK_GEMM_S * linear_ratio
)
summary[layout_name][precision] = {
"weighted_ms_per_block": weighted_ms,
"weighted_speedup_vs_layout_bf16": layout_bf16 / weighted_ms,
"weighted_speedup_vs_current_bf16": current_bf16 / weighted_ms,
"projected_step_s": projected_step_s,
"projected_step_speedup": BASELINE_STEP_S / projected_step_s,
"projected_mfu_percent": MFU_PERCENT_SECONDS / projected_step_s,
}
print(json.dumps({
"projection_baseline": {
"layout": "separate",
"precision": "bf16",
"step_s": BASELINE_STEP_S,
"block_gemm_s": BASELINE_BLOCK_GEMM_S,
"m_multiplier": args.m_multiplier,
},
"summary": summary,
}), flush=True)
if __name__ == "__main__":
main()
@@ -0,0 +1,126 @@
import argparse
import gc
import json
import statistics
import torch
from torchao.float8 import Float8LinearConfig, convert_to_float8_training
LAYOUTS = {
"separate": {
"h_h": (4290, 4096, 4096, 6),
"text_kv": (1024, 4096, 4096, 2),
"ffn_up": (4290, 4096, 16384, 1),
"ffn_down": (4290, 16384, 4096, 1),
},
"packed": {
"self_qkv": (4290, 4096, 12288, 1),
"h_h": (4290, 4096, 4096, 3),
"text_kv": (1024, 4096, 8192, 1),
"ffn_up": (4290, 4096, 16384, 1),
"ffn_down": (4290, 16384, 4096, 1),
},
}
def measure(precision, layout, shape, warmup, repeat, compile_module):
m, k, n, count = LAYOUTS[layout][shape]
torch.manual_seed(0)
x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16, requires_grad=True)
dy = torch.randn(m, n, device="cuda", dtype=torch.bfloat16)
layer = torch.nn.Linear(k, n, bias=True, device="cuda", dtype=torch.bfloat16)
if precision == "fp8":
layer = convert_to_float8_training(
layer,
config=Float8LinearConfig(pad_inner_dim=True),
)
elif precision != "bf16":
raise ValueError(precision)
if compile_module:
layer = torch.compile(layer, dynamic=False)
def step():
layer.zero_grad(set_to_none=True)
x.grad = None
layer(x).backward(dy)
for _ in range(warmup):
step()
torch.cuda.synchronize()
samples = []
for _ in range(repeat):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
step()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end))
result = {
"precision": precision,
"layout": layout,
"shape": shape,
"m": m,
"k": k,
"n": n,
"count_per_block": count,
"compiled": compile_module,
"median_ms": statistics.median(samples),
"min_ms": min(samples),
}
del layer, x, dy
gc.collect()
torch.cuda.empty_cache()
return result
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--warmup", type=int, default=3)
parser.add_argument("--repeat", type=int, default=10)
parser.add_argument("--compile", action="store_true")
parser.add_argument("--layouts", nargs="+", choices=tuple(LAYOUTS), default=list(LAYOUTS))
parser.add_argument("--shapes", nargs="+")
parser.add_argument("--precisions", nargs="+", choices=("bf16", "fp8"), default=["bf16", "fp8"])
args = parser.parse_args()
print(json.dumps({
"torch": torch.__version__,
"device": torch.cuda.get_device_name(),
"compiled": args.compile,
}), flush=True)
results = []
for layout in args.layouts:
shapes = LAYOUTS[layout]
for precision in args.precisions:
for shape in shapes:
if args.shapes is not None and shape not in args.shapes:
continue
row = measure(
precision,
layout,
shape,
args.warmup,
args.repeat,
args.compile,
)
results.append(row)
print(json.dumps(row), flush=True)
totals = {
layout: {
precision: sum(
row["median_ms"] * row["count_per_block"]
for row in results
if row["layout"] == layout and row["precision"] == precision
)
for precision in args.precisions
}
for layout in args.layouts
}
print(json.dumps({"totals": totals}), flush=True)
if __name__ == "__main__":
main()
@@ -0,0 +1,284 @@
#!/usr/bin/env python3
"""Exact LTX-2 B=1 fprop/dgrad/wgrad benchmark for TorchAO NVFP4."""
import argparse
import gc
import json
import statistics
import time
from dataclasses import dataclass
from typing import Callable
import torch
import torch.nn.functional as F
from torchao.prototype.mx_formats.nvfp4_tensor import (
NVFP4Tensor,
_addmm_nvfp4_dispatch,
per_tensor_amax_to_scale,
)
@dataclass(frozen=True)
class Layer:
name: str
count: int
m: int
k: int
n: int
LAYOUTS = {
"separate": (
Layer("video_dd", 288, 4290, 4096, 4096),
Layer("text_dd", 96, 1024, 4096, 4096),
Layer("ffn_up", 48, 4290, 4096, 16384),
Layer("ffn_down", 48, 4290, 16384, 4096),
),
"packed": (
Layer("self_qkv", 48, 4290, 4096, 12288),
Layer("video_dd", 144, 4290, 4096, 4096),
Layer("text_kv", 48, 1024, 4096, 8192),
Layer("ffn_up", 48, 4290, 4096, 16384),
Layer("ffn_down", 48, 4290, 16384, 4096),
),
}
PHASES = ("fwd", "dgrad", "wgrad")
TIERS = (
"bf16",
"fp4_mm",
"fp4_mm_2level",
"fp4_eager",
"fp4_eager_dynamic",
"fp4_triton_static",
"fp4_triton_dynamic",
)
def emit(kind: str, **values: object) -> None:
print(json.dumps({"kind": kind, **values}, sort_keys=True), flush=True)
def timing(fn: Callable[[], torch.Tensor], warmup: int, samples: int, inner: int) -> dict[str, float]:
with torch.no_grad():
for _ in range(warmup):
fn()
torch.cuda.synchronize()
values = []
for _ in range(samples):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(inner):
fn()
end.record()
end.synchronize()
values.append(start.elapsed_time(end) / inner)
return {"median_ms": statistics.median(values), "min_ms": min(values), "max_ms": max(values)}
def error(actual: torch.Tensor, expected: torch.Tensor) -> dict[str, float | bool]:
actual = actual.float()
expected = expected.float()
diff = actual - expected
return {
"finite": bool(torch.isfinite(actual).all().item()),
"relative_rms": (diff.square().mean().sqrt() / expected.square().mean().sqrt().clamp_min(1e-30)).item(),
"mean_abs": diff.abs().mean().item(),
"max_abs": diff.abs().amax().item(),
}
def make_case(layer: Layer, phase: str):
m, k, n = layer.m, layer.k, layer.n
if phase == "fwd":
x = torch.empty((m, k), device="cuda", dtype=torch.bfloat16).normal_()
w = torch.empty((n, k), device="cuda", dtype=torch.bfloat16).normal_(0, 0.02)
return lambda: torch.mm(x, w.T), lambda: (x, w), (m, k, n), 2 * m * k * n
if phase == "dgrad":
dy = torch.empty((m, n), device="cuda", dtype=torch.bfloat16).normal_(0, 0.01)
w = torch.empty((n, k), device="cuda", dtype=torch.bfloat16).normal_(0, 0.02)
return lambda: torch.mm(dy, w), lambda: (dy, w.T.contiguous()), (m, n, k), 2 * m * n * k
if phase == "wgrad":
dy = torch.empty((m, n), device="cuda", dtype=torch.bfloat16).normal_(0, 0.01)
x = torch.empty((m, k), device="cuda", dtype=torch.bfloat16).normal_()
padded_m = (m + 31) // 32 * 32
def orient():
pad = padded_m - m
dy_padded = F.pad(dy, (0, 0, 0, pad)) if pad else dy
x_padded = F.pad(x, (0, 0, 0, pad)) if pad else x
return dy_padded.T.contiguous(), x_padded.T.contiguous()
return lambda: torch.mm(dy.T, x), orient, (n, padded_m, k), 2 * m * n * k
raise ValueError(phase)
def quantize(tensor: torch.Tensor, triton: bool, dynamic: bool, static_scale: torch.Tensor | None):
scale = per_tensor_amax_to_scale(tensor.abs().amax()) if dynamic else static_scale
return NVFP4Tensor.to_nvfp4(
tensor,
per_tensor_scale=scale,
is_swizzled_scales=True,
use_triton_kernel=triton,
)
def mm(lhs: NVFP4Tensor, rhs_rows: NVFP4Tensor) -> torch.Tensor:
return _addmm_nvfp4_dispatch(lhs, rhs_rows.t(), None)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--warmup", type=int, default=3)
parser.add_argument("--samples", type=int, default=7)
parser.add_argument("--inner", type=int, default=3)
parser.add_argument("--layouts", nargs="+", default=list(LAYOUTS))
parser.add_argument("--phases", nargs="+", default=list(PHASES))
parser.add_argument("--tiers", nargs="+", default=list(TIERS))
parser.add_argument("--concat", action="store_true")
args = parser.parse_args()
torch.cuda.set_device(0)
torch.manual_seed(20260721)
torch.backends.cuda.matmul.allow_tf32 = False
started = time.monotonic()
emit(
"environment",
torch=torch.__version__,
torchao=__import__("torchao").__version__,
gpu=torch.cuda.get_device_name(0),
capability=torch.cuda.get_device_capability(0),
warmup=args.warmup,
samples=args.samples,
inner=args.inner,
)
aggregates = {
layout: {tier: {phase: {"latency_ms": 0.0, "logical_flops": 0.0} for phase in PHASES} for tier in args.tiers}
for layout in args.layouts
}
for layout in args.layouts:
for phase in args.phases:
for index, layer in enumerate(LAYOUTS[layout]):
gc.collect()
torch.cuda.empty_cache()
torch.manual_seed(20260721 + 100 * index + PHASES.index(phase))
bf16, orient, qshape, logical_flops = make_case(layer, phase)
reference = bf16()
lhs, rhs = orient()
static_lhs_scale = per_tensor_amax_to_scale(lhs.abs().amax())
static_rhs_scale = per_tensor_amax_to_scale(rhs.abs().amax())
lhs_q = quantize(lhs, False, False, None)
rhs_q = quantize(rhs, False, False, None)
lhs_q_2level = quantize(lhs, False, False, static_lhs_scale)
rhs_q_2level = quantize(rhs, False, False, static_rhs_scale)
torch.cuda.synchronize()
calls = {
"bf16": bf16,
"fp4_mm": lambda: mm(lhs_q, rhs_q),
"fp4_mm_2level": lambda: mm(lhs_q_2level, rhs_q_2level),
}
def eager():
a, b = orient()
return mm(quantize(a, False, False, None), quantize(b, False, False, None))
def triton_static():
a, b = orient()
return mm(
quantize(a, True, False, static_lhs_scale),
quantize(b, True, False, static_rhs_scale),
)
def triton_dynamic():
a, b = orient()
return mm(quantize(a, True, True, None), quantize(b, True, True, None))
def eager_dynamic():
a, b = orient()
return mm(quantize(a, False, True, None), quantize(b, False, True, None))
calls.update(
fp4_eager=eager,
fp4_eager_dynamic=eager_dynamic,
fp4_triton_static=triton_static,
fp4_triton_dynamic=triton_dynamic,
)
for tier in args.tiers:
try:
result = calls[tier]()
torch.cuda.synchronize()
metrics = timing(calls[tier], args.warmup, args.samples, args.inner)
numerical = None if tier == "bf16" else error(result, reference)
aggregates[layout][tier][phase]["latency_ms"] += metrics["median_ms"] * layer.count
aggregates[layout][tier][phase]["logical_flops"] += logical_flops * layer.count
emit(
"case",
layout=layout,
tier=tier,
phase=phase,
layer=layer.name,
count=layer.count,
logical_shape=(layer.m, layer.k, layer.n),
qmm_shape=qshape,
logical_flops=logical_flops,
error=numerical,
**metrics,
)
except Exception as exc:
torch.cuda.synchronize()
emit("failure", tier=tier, phase=phase, layer=layer.name, qmm_shape=qshape, exception=repr(exc))
del reference, lhs, rhs, lhs_q, rhs_q, lhs_q_2level, rhs_q_2level
totals = {}
for layout, tiers in aggregates.items():
totals[layout] = {}
for tier, phases in tiers.items():
latency = sum(row["latency_ms"] for row in phases.values())
flops = sum(row["logical_flops"] for row in phases.values())
complete = all(row["logical_flops"] > 0 for row in phases.values())
totals[layout][tier] = {
"complete": complete,
"total_latency_ms": latency,
"total_logical_flops": flops,
"effective_tflops": flops / latency / 1e9 if latency else 0.0,
"phases": phases,
}
emit("aggregate", layout=layout, tier=tier, **totals[layout][tier])
if "separate" in totals and "bf16" in totals["separate"]:
bf16_ms = totals["separate"]["bf16"]["total_latency_ms"]
candidates = {
f"{layout}/{tier}": row
for layout, tiers in totals.items()
for tier, row in tiers.items()
if not (layout == "separate" and tier == "bf16") and row["total_latency_ms"]
}
emit(
"gate",
baseline="separate/bf16",
required_speedup=2.5,
results={name: {"speedup_vs_separate_bf16": bf16_ms / row["total_latency_ms"], "passes": row["complete"] and bf16_ms / row["total_latency_ms"] > 2.5} for name, row in candidates.items()},
)
if args.concat:
concat_total = 0.0
for name, shape, copies, dim in (
("self_qkv_weight", (4096, 4096), 3, 0),
("text_kv_weight", (4096, 4096), 2, 0),
("self_qkv_grad_output", (4290, 4096), 3, 1),
("text_kv_grad_output", (1024, 4096), 2, 1),
):
inputs = [torch.empty(shape, device="cuda", dtype=torch.bfloat16).normal_() for _ in range(copies)]
metrics = timing(lambda: torch.cat(inputs, dim=dim), args.warmup, args.samples, args.inner)
weighted_ms = metrics["median_ms"] * 48
concat_total += weighted_ms
emit("concat", name=name, count=48, weighted_ms=weighted_ms, **metrics)
del inputs
torch.cuda.empty_cache()
emit("concat_aggregate", weighted_ms=concat_total)
emit("done", elapsed_wall_seconds=time.monotonic() - started)
if __name__ == "__main__":
main()
@@ -0,0 +1,185 @@
#!/usr/bin/env python3
"""Quantify LTX-2's avoidable BF16 velocity -> x0 -> velocity error.
The modular fine-tuning target is raw flow velocity ``noise - clean``. This
script feeds that exact target through FastVideo's real ``_to_denoised``
helper and the current modular adapter's inverse, then compares the result
with returning the raw transformer velocity directly.
"""
from __future__ import annotations
import argparse
import json
import math
from typing import Any
import torch
from fastvideo.models.dits.ltx2 import _to_denoised
DEFAULT_SIGMAS = (1.0, 0.5, 0.1, 0.01, 0.001)
def _metrics(actual: torch.Tensor, expected: torch.Tensor) -> dict[str, float]:
actual = actual.float().flatten()
expected = expected.float().flatten()
error = actual - expected
expected_rms = torch.sqrt(torch.mean(expected.square()))
error_rms = torch.sqrt(torch.mean(error.square()))
denominator = torch.linalg.vector_norm(actual) * torch.linalg.vector_norm(expected)
cosine = torch.dot(actual, expected) / denominator
return {
"mse": float(torch.mean(error.square())),
"rmse": float(error_rms),
"relative_rmse": float(error_rms / expected_rms),
"max_abs": float(error.abs().max()),
"cosine": float(cosine),
"exact_fraction": float(torch.mean((actual == expected).float())),
"zero_fraction": float(torch.mean((actual == 0).float())),
}
def _run_sigma(
clean_bf16: torch.Tensor,
noise_bf16: torch.Tensor,
sigma_value: float,
) -> dict[str, Any]:
sigma = torch.tensor(
sigma_value,
device=clean_bf16.device,
dtype=torch.float32,
).view(1, 1, 1, 1, 1)
# These expressions match LTX2Model._prepare_dit_inputs and FineTuneMethod:
# x_t is accumulated in fp32 and stored in bf16, while the target is the
# bf16 subtraction noise - clean.
noisy_bf16 = (
(1.0 - sigma) * clean_bf16.float() + sigma * noise_bf16.float()
).to(torch.bfloat16)
target_bf16 = noise_bf16 - clean_bf16
# Current path: the raw DiT output is converted to bf16 x0, only for the
# modular adapter to reconstruct fp32 velocity by subtracting and dividing.
denoised_bf16 = _to_denoised(
noisy_bf16,
target_bf16,
sigma,
)
reconstructed_velocity = (
noisy_bf16.float() - denoised_bf16.float()
) / sigma
# Proposed training path: use the transformer's raw velocity as-is.
direct_velocity = target_bf16.float()
# FP32 reference separates unavoidable floating-point cancellation from
# the extra bf16 x0 materialization in the current path.
clean_fp32 = clean_bf16.float()
noise_fp32 = noise_bf16.float()
target_fp32 = noise_fp32 - clean_fp32
noisy_fp32 = (1.0 - sigma) * clean_fp32 + sigma * noise_fp32
denoised_fp32 = _to_denoised(
noisy_fp32,
target_fp32,
sigma,
)
reconstructed_fp32 = (noisy_fp32 - denoised_fp32) / sigma
return {
"sigma": sigma_value,
"direct_raw_velocity": _metrics(direct_velocity, target_bf16),
"current_bf16_roundtrip": _metrics(
reconstructed_velocity,
target_bf16,
),
"fp32_roundtrip_control": _metrics(
reconstructed_fp32,
target_fp32,
),
"denoised_x0_vs_clean": _metrics(denoised_bf16, clean_bf16),
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument(
"--device",
default="auto",
choices=("auto", "cpu", "cuda"),
)
parser.add_argument("--elements", type=int, default=262_144)
parser.add_argument("--seed", type=int, default=20260721)
args = parser.parse_args()
if args.elements <= 0:
parser.error("--elements must be positive")
device = (
"cuda"
if args.device == "auto" and torch.cuda.is_available()
else "cpu"
if args.device == "auto"
else args.device
)
if device == "cuda" and not torch.cuda.is_available():
parser.error("CUDA was requested but is unavailable")
generator = torch.Generator(device=device).manual_seed(args.seed)
shape = (1, 1, 1, 1, args.elements)
clean_bf16 = torch.randn(
shape,
generator=generator,
device=device,
dtype=torch.float32,
).to(torch.bfloat16)
noise_bf16 = torch.randn(
shape,
generator=generator,
device=device,
dtype=torch.float32,
).to(torch.bfloat16)
results = [
_run_sigma(clean_bf16, noise_bf16, sigma)
for sigma in DEFAULT_SIGMAS
]
direct_mses = [item["direct_raw_velocity"]["mse"] for item in results]
current_mses = [item["current_bf16_roundtrip"]["mse"] for item in results]
small_sigma_relative_rmse = results[-1]["current_bf16_roundtrip"]["relative_rmse"]
if any(value != 0.0 for value in direct_mses):
raise AssertionError(f"direct raw velocity was not exact: {direct_mses}")
if not all(math.isfinite(value) for value in current_mses):
raise AssertionError(f"non-finite current-path MSE: {current_mses}")
if current_mses[-1] <= current_mses[2] * 100.0:
raise AssertionError(
"expected sigma=1e-3 BF16 round-trip MSE to exceed sigma=0.1 "
f"by >100x, got {current_mses[-1]} vs {current_mses[2]}"
)
if small_sigma_relative_rmse <= 0.1:
raise AssertionError(
"expected material small-sigma error, got relative RMSE "
f"{small_sigma_relative_rmse}"
)
print(
"LTX2_RAW_VELOCITY_PROOF " + json.dumps({
"device": device,
"dtype": "torch.bfloat16",
"elements": args.elements,
"seed": args.seed,
"training_target": "noise - clean (raw flow velocity)",
"source_evidence": [
"fastvideo/train/methods/fine_tuning/finetune.py:99-100",
"fastvideo/training/ltx2_training_pipeline.py:435-438",
"fastvideo/models/schedulers/scheduling_self_forcing_flow_match.py:140-141",
],
"results": results,
}, sort_keys=True),
flush=True,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,249 @@
#!/usr/bin/env python3
"""CPU proof that uniform LTX-2 token timesteps may be embedded once."""
from __future__ import annotations
import argparse
import copy
import json
import torch
from fastvideo.models.dits.ltx2 import AdaLayerNormSingle, _to_denoised
# These are norm-based numerical gates, not bitwise-parity claims. Changing
# the number of rows presented to a GEMM may change its reduction schedule.
_FLOAT32_FORWARD_MAX_ABS = 1e-6
_FLOAT32_FORWARD_RELATIVE_L2 = 2e-6
_FLOAT32_GRADIENT_RELATIVE_L2 = 2e-6
_SIGMA_GRADIENT_RELATIVE_L2 = 5e-6
_GRADIENT_COSINE = 0.99999999
_FLOAT64_FORWARD_MAX_ABS = 1e-12
_FLOAT64_FORWARD_RELATIVE_L2 = 1e-13
_FLOAT64_GRADIENT_RELATIVE_L2 = 1e-12
def _cosine(left: torch.Tensor, right: torch.Tensor) -> float:
left = left.detach().double().flatten()
right = right.detach().double().flatten()
denominator = left.norm() * right.norm()
return float(torch.dot(left, right) / denominator) if denominator else 1.0
def _comparison(left: torch.Tensor, right: torch.Tensor) -> dict[str, float]:
left = left.detach().double()
right = right.detach().double()
difference = left - right
difference_norm = torch.linalg.vector_norm(difference)
reference_norm = torch.linalg.vector_norm(left)
return {
"max_abs": float(difference.abs().max()),
"rmse": float(torch.sqrt(torch.mean(difference.square()))),
"relative_l2": float(difference_norm / max(reference_norm, torch.finfo(torch.float64).eps)),
"relative_to_max": float(
difference.abs().max() / max(left.abs().max(), torch.finfo(torch.float64).eps)
),
"cosine": _cosine(left, right),
"exact_fraction": float(torch.mean((left == right).double())),
}
def _forward_stages(
adaln: AdaLayerNormSingle,
sigma: torch.Tensor,
hidden_dtype: torch.dtype,
) -> dict[str, torch.Tensor]:
time_projection = adaln.emb.time_proj(sigma)
timestep_linear_1 = adaln.emb.timestep_embedder.linear_1(
time_projection.to(dtype=hidden_dtype)
)
timestep_silu = adaln.emb.timestep_embedder.act(timestep_linear_1)
embedding = adaln.emb.timestep_embedder.linear_2(timestep_silu)
modulation_silu = adaln.silu(embedding)
modulation = adaln.linear(modulation_silu)
return {
"time_projection": time_projection,
"timestep_linear_1": timestep_linear_1,
"timestep_silu": timestep_silu,
"embedding": embedding,
"modulation_silu": modulation_silu,
"modulation": modulation,
}
def _run_case(
coefficient: int,
dtype: torch.dtype,
token_count: int,
) -> dict[str, object]:
torch.manual_seed(20260721 + coefficient)
torch.set_num_threads(1)
batch_size, hidden_size = 2, 32
expanded_adaln = AdaLayerNormSingle(hidden_size, embedding_coefficient=coefficient).float().to(dtype)
singleton_adaln = copy.deepcopy(expanded_adaln)
sigma_values = torch.rand(batch_size, 1, dtype=torch.float32).to(dtype)
expanded_sigma = sigma_values.detach().clone().requires_grad_(True)
singleton_sigma = expanded_sigma.detach().clone().requires_grad_(True)
expanded_input = expanded_sigma.expand(batch_size, token_count).contiguous().flatten()
singleton_input = singleton_sigma.flatten()
expanded_modulation, expanded_embedding = expanded_adaln(expanded_input, hidden_dtype=dtype)
singleton_modulation, singleton_embedding = singleton_adaln(singleton_input, hidden_dtype=dtype)
expanded_modulation = expanded_modulation.view(batch_size, token_count, coefficient, hidden_size)
singleton_modulation = singleton_modulation.view(batch_size, 1, coefficient, hidden_size)
expanded_embedding = expanded_embedding.view(batch_size, token_count, hidden_size)
singleton_embedding = singleton_embedding.view(batch_size, 1, hidden_size)
singleton_modulation_broadcast = singleton_modulation.expand_as(expanded_modulation)
singleton_embedding_broadcast = singleton_embedding.expand_as(expanded_embedding)
modulation_comparison = _comparison(expanded_modulation, singleton_modulation_broadcast)
embedding_comparison = _comparison(expanded_embedding, singleton_embedding_broadcast)
with torch.no_grad():
expanded_stages = _forward_stages(expanded_adaln, expanded_input, dtype)
singleton_stages = _forward_stages(singleton_adaln, singleton_input, dtype)
stage_comparisons: dict[str, dict[str, float]] = {}
expanded_repeat_consistency: dict[str, dict[str, float]] = {}
for name, expanded_stage in expanded_stages.items():
singleton_stage = singleton_stages[name]
expanded_stage = expanded_stage.view(batch_size, token_count, -1)
singleton_stage = singleton_stage.view(batch_size, 1, -1)
stage_comparisons[name] = _comparison(
expanded_stage,
singleton_stage.expand_as(expanded_stage),
)
expanded_repeat_consistency[name] = _comparison(
expanded_stage,
expanded_stage[:, :1].expand_as(expanded_stage),
)
# A distinct upstream gradient per token covers every downstream broadcast
# consumer: block Ada values and the final output modulation.
modulation_grad = torch.randn(
expanded_modulation.shape,
dtype=torch.float32,
).to(dtype)
embedding_grad = torch.randn(
expanded_embedding.shape,
dtype=torch.float32,
).to(dtype)
expanded_loss = (expanded_modulation * modulation_grad).sum() + (expanded_embedding * embedding_grad).sum()
singleton_loss = ((singleton_modulation * modulation_grad).sum()
+ (singleton_embedding * embedding_grad).sum())
expanded_loss.backward()
singleton_loss.backward()
parameter_checks: dict[str, dict[str, float]] = {}
worst_gradient_max_abs = 0.0
worst_gradient_relative_l2 = 0.0
worst_gradient_relative_to_max = 0.0
worst_gradient_cosine = 1.0
for (expanded_name, expanded_parameter), (singleton_name, singleton_parameter) in zip(
expanded_adaln.named_parameters(), singleton_adaln.named_parameters(), strict=True):
assert expanded_name == singleton_name
assert expanded_parameter.grad is not None and singleton_parameter.grad is not None
comparison = _comparison(expanded_parameter.grad, singleton_parameter.grad)
parameter_checks[expanded_name] = comparison
worst_gradient_max_abs = max(worst_gradient_max_abs, comparison["max_abs"])
worst_gradient_relative_l2 = max(worst_gradient_relative_l2, comparison["relative_l2"])
worst_gradient_relative_to_max = max(
worst_gradient_relative_to_max,
comparison["relative_to_max"],
)
worst_gradient_cosine = min(worst_gradient_cosine, comparison["cosine"])
assert expanded_sigma.grad is not None and singleton_sigma.grad is not None
sigma_gradient_comparison = _comparison(expanded_sigma.grad, singleton_sigma.grad)
# The wrapper's final denoising conversion also broadcasts [B, 1] over
# [B, tokens, channels].
sample = torch.randn(batch_size, token_count, hidden_size, dtype=torch.float32).to(dtype)
velocity = torch.randn_like(sample)
expanded_denoised = _to_denoised(
sample,
velocity,
expanded_sigma.detach().expand(batch_size, token_count),
calc_dtype=dtype,
)
singleton_denoised = _to_denoised(
sample,
velocity,
singleton_sigma.detach(),
calc_dtype=dtype,
)
denoised_comparison = _comparison(expanded_denoised, singleton_denoised)
return {
"coefficient": coefficient,
"dtype": str(dtype),
"token_count": token_count,
"modulation": modulation_comparison,
"embedding": embedding_comparison,
"denoised": denoised_comparison,
"forward_stages": stage_comparisons,
"expanded_repeat_consistency": expanded_repeat_consistency,
"worst_gradient_max_abs": worst_gradient_max_abs,
"worst_gradient_relative_l2": worst_gradient_relative_l2,
"worst_gradient_relative_to_max": worst_gradient_relative_to_max,
"worst_gradient_cosine": worst_gradient_cosine,
"sigma_gradient": sigma_gradient_comparison,
"parameter_checks": parameter_checks,
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--diagnostic-only", action="store_true")
parser.add_argument("--token-count", type=int, default=4290)
args = parser.parse_args()
if args.token_count <= 0:
parser.error("--token-count must be positive")
results = [
_run_case(coefficient, dtype, args.token_count)
for dtype in (torch.float32, torch.float64)
for coefficient in (6, 9)
]
print(
"LTX2_SINGLETON_PARITY " + json.dumps(results, sort_keys=True),
flush=True,
)
if args.diagnostic_only:
return
for result in results:
forward_max_abs = (
_FLOAT32_FORWARD_MAX_ABS
if result["dtype"] == "torch.float32"
else _FLOAT64_FORWARD_MAX_ABS
)
forward_relative_l2 = (
_FLOAT32_FORWARD_RELATIVE_L2
if result["dtype"] == "torch.float32"
else _FLOAT64_FORWARD_RELATIVE_L2
)
for comparison in result["forward_stages"].values():
assert comparison["max_abs"] < forward_max_abs
assert comparison["relative_l2"] < forward_relative_l2
for comparison in result["expanded_repeat_consistency"].values():
assert comparison["max_abs"] < forward_max_abs
assert comparison["relative_l2"] < forward_relative_l2
if result["dtype"] == "torch.float32":
assert result["worst_gradient_relative_l2"] < _FLOAT32_GRADIENT_RELATIVE_L2
else:
assert result["worst_gradient_relative_l2"] < _FLOAT64_GRADIENT_RELATIVE_L2
assert result["worst_gradient_cosine"] > _GRADIENT_COSINE
# get_timestep_embedding() explicitly computes in float32 even when
# the surrounding AdaLN is float64. Its sigma-gradient reduction is
# therefore expected to use the float32 tolerance in both cases.
assert result["sigma_gradient"]["relative_l2"] < _SIGMA_GRADIENT_RELATIVE_L2
assert result["sigma_gradient"]["cosine"] > _GRADIENT_COSINE
assert result["denoised"]["max_abs"] == 0.0
if __name__ == "__main__":
main()
@@ -0,0 +1,108 @@
import argparse
import statistics
import torch
import torch.distributed as dist
def timed(fn, stream: torch.cuda.Stream, warmup: int = 5, iterations: int = 20) -> tuple[float, float]:
for _ in range(warmup):
fn()
torch.cuda.synchronize()
samples = []
for _ in range(iterations):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
with torch.cuda.stream(stream):
start.record()
fn()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end))
return statistics.median(samples), statistics.mean(samples)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--symm", action="store_true")
parser.add_argument("--pg-alloc", action="store_true")
parser.add_argument("--elements", type=int, default=67_500_000)
args = parser.parse_args()
rank = int(__import__("os").environ["LOCAL_RANK"])
torch.cuda.set_device(rank)
dist.init_process_group("nccl", device_id=torch.device("cuda", rank))
group_name = dist.group.WORLD.group_name
world_size = dist.get_world_size()
stream = torch.cuda.Stream(priority=-1)
if args.symm and args.pg_alloc:
raise ValueError("choose only one allocator")
if args.symm:
import torch.distributed._symmetric_memory as symm_mem
symm_mem.set_backend("NCCL")
pool = symm_mem.get_mem_pool(torch.device("cuda", rank))
with torch.cuda.use_mem_pool(pool):
ag_out = torch.empty(args.elements * world_size, dtype=torch.bfloat16, device=rank)
rs_in = torch.empty(args.elements * world_size, dtype=torch.bfloat16, device=rank)
rs_out = torch.empty(args.elements, dtype=torch.bfloat16, device=rank)
symm_mem.rendezvous(ag_out, group=group_name)
symm_mem.rendezvous(rs_in, group=group_name)
symm_mem.rendezvous(rs_out, group=group_name)
elif args.pg_alloc:
backend = dist.group.WORLD._get_backend(torch.device("cuda", rank))
if not backend.supports_tensor_alloc(torch.device("cuda", rank)):
raise RuntimeError("ProcessGroupNCCL tensor allocator is unavailable")
ag_out = backend.allocate_tensor(
args.elements * world_size, dtype=torch.bfloat16, device=torch.device("cuda", rank)
)
rs_in = backend.allocate_tensor(
args.elements * world_size, dtype=torch.bfloat16, device=torch.device("cuda", rank)
)
rs_out = backend.allocate_tensor(
args.elements, dtype=torch.bfloat16, device=torch.device("cuda", rank)
)
else:
ag_out = torch.empty(args.elements * world_size, dtype=torch.bfloat16, device=rank)
rs_in = torch.empty(args.elements * world_size, dtype=torch.bfloat16, device=rank)
rs_out = torch.empty(args.elements, dtype=torch.bfloat16, device=rank)
ag_in = ag_out.narrow(0, rank * args.elements, args.elements)
ag_in.fill_(rank + 1)
rs_in.fill_(1)
def ag() -> None:
dist.all_gather_into_tensor(ag_out, ag_in)
def rs_bf16() -> None:
dist.reduce_scatter_tensor(rs_out, rs_in, op=dist.ReduceOp.SUM)
ag_median, ag_mean = timed(ag, stream)
rs_median, rs_mean = timed(rs_bf16, stream)
if not args.symm:
rs32_in = torch.empty(args.elements * world_size, dtype=torch.float32, device=rank)
rs32_out = torch.empty(args.elements, dtype=torch.float32, device=rank)
rs32_in.fill_(1)
def rs_fp32() -> None:
dist.reduce_scatter_tensor(rs32_out, rs32_in, op=dist.ReduceOp.SUM)
rs32_median, rs32_mean = timed(rs_fp32, stream)
else:
rs32_median = rs32_mean = float("nan")
if rank == 0:
mode = "symm_ce" if args.symm else "pg_alloc" if args.pg_alloc else "default"
print(f"mode={mode} elements_per_rank={args.elements} world={world_size}")
print(f"ag_bf16_ms median={ag_median:.3f} mean={ag_mean:.3f}")
print(f"rs_bf16_ms median={rs_median:.3f} mean={rs_mean:.3f}")
if not args.symm:
print(f"rs_fp32_ms median={rs32_median:.3f} mean={rs32_mean:.3f}")
dist.destroy_process_group()
if __name__ == "__main__":
main()
@@ -0,0 +1,102 @@
// Minimal cuBLASLt pinned-algo GEMM registry for the LTX-2 packed shapes.
// Python owns tuning (probes/bench_ltx2_cublaslt_algo_sweep.py logic) and
// registers one case per (m, n, k, transa, transb) with the winning
// 64-byte cublasLtMatmulAlgo_t blob; lt_mm dispatches by case id with zero
// Python in the hot path. BF16 in/out, FP32 compute, no epilogue.
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cublasLt.h>
#include <vector>
#include <cstring>
#include <stdexcept>
namespace {
struct PinnedCase {
cublasLtMatmulDesc_t op_desc{};
cublasLtMatrixLayout_t layout_a{};
cublasLtMatrixLayout_t layout_b{};
cublasLtMatrixLayout_t layout_d{};
cublasLtMatmulAlgo_t algo{};
int64_t m{}, n{}, k{};
int64_t d_rows{}, d_cols{};
};
cublasLtHandle_t lt_handle() {
static cublasLtHandle_t handle = [] {
cublasLtHandle_t h;
TORCH_CHECK(cublasLtCreate(&h) == CUBLAS_STATUS_SUCCESS, "cublasLtCreate failed");
return h;
}();
return handle;
}
std::vector<PinnedCase>& cases() {
static std::vector<PinnedCase> registry;
return registry;
}
torch::Tensor& workspace() {
static torch::Tensor ws = torch::empty(
{128 * 1024 * 1024},
torch::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA));
return ws;
}
void check_lt(cublasStatus_t status, const char* what) {
TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, what, " failed with status ", static_cast<int>(status));
}
int64_t register_case(int64_t m, int64_t n, int64_t k, bool transa, bool transb,
int64_t lda, int64_t ldb, int64_t d_rows, int64_t d_cols,
torch::Tensor algo_bytes) {
TORCH_CHECK(algo_bytes.dtype() == torch::kUInt8 && algo_bytes.numel() == sizeof(cublasLtMatmulAlgo_t),
"algo blob must be 64 uint8 bytes");
PinnedCase entry;
entry.m = m; entry.n = n; entry.k = k;
entry.d_rows = d_rows; entry.d_cols = d_cols;
check_lt(cublasLtMatmulDescCreate(&entry.op_desc, CUBLAS_COMPUTE_32F, CUDA_R_32F), "descCreate");
const cublasOperation_t opa = transa ? CUBLAS_OP_T : CUBLAS_OP_N;
const cublasOperation_t opb = transb ? CUBLAS_OP_T : CUBLAS_OP_N;
check_lt(cublasLtMatmulDescSetAttribute(entry.op_desc, CUBLASLT_MATMUL_DESC_TRANSA, &opa, sizeof(opa)), "setA");
check_lt(cublasLtMatmulDescSetAttribute(entry.op_desc, CUBLASLT_MATMUL_DESC_TRANSB, &opb, sizeof(opb)), "setB");
const int64_t a_rows = transa ? lda : m;
const int64_t a_cols = transa ? m : k;
const int64_t b_rows = transb ? ldb : k;
const int64_t b_cols = transb ? k : n;
check_lt(cublasLtMatrixLayoutCreate(&entry.layout_a, CUDA_R_16BF, a_rows, a_cols, lda), "layoutA");
check_lt(cublasLtMatrixLayoutCreate(&entry.layout_b, CUDA_R_16BF, b_rows, b_cols, ldb), "layoutB");
check_lt(cublasLtMatrixLayoutCreate(&entry.layout_d, CUDA_R_16BF, m, n, m), "layoutD");
std::memcpy(&entry.algo, algo_bytes.data_ptr<uint8_t>(), sizeof(cublasLtMatmulAlgo_t));
cases().push_back(entry);
return static_cast<int64_t>(cases().size()) - 1;
}
torch::Tensor lt_mm(torch::Tensor a, torch::Tensor b, int64_t case_id) {
TORCH_CHECK(case_id >= 0 && case_id < static_cast<int64_t>(cases().size()), "unknown lt case");
TORCH_CHECK(a.is_cuda() && b.is_cuda() && a.dtype() == torch::kBFloat16 && b.dtype() == torch::kBFloat16,
"lt_mm expects CUDA bf16 tensors");
TORCH_CHECK(a.is_contiguous() && b.is_contiguous(), "lt_mm expects contiguous operands");
const PinnedCase& entry = cases()[case_id];
auto d = torch::empty({entry.d_rows, entry.d_cols}, a.options());
const float alpha = 1.0f;
const float beta = 0.0f;
check_lt(cublasLtMatmul(lt_handle(), entry.op_desc, &alpha,
a.data_ptr(), entry.layout_a,
b.data_ptr(), entry.layout_b,
&beta,
d.data_ptr(), entry.layout_d,
d.data_ptr(), entry.layout_d,
&entry.algo,
workspace().data_ptr(), workspace().numel(),
at::cuda::getCurrentCUDAStream()), "ltMatmul");
return d;
}
} // namespace
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("register_case", &register_case, "register pinned cuBLASLt case");
module.def("lt_mm", &lt_mm, "run pinned cuBLASLt matmul");
}
@@ -0,0 +1,85 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import os
from typing import TYPE_CHECKING
import torch
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases,
get_scheduler,
)
if TYPE_CHECKING:
from fastvideo.train.utils.training_config import (
OptimizerConfig,
TrainingLoopConfig,
)
def build_optimizer_and_scheduler(
*,
params: list[torch.nn.Parameter],
optimizer_config: OptimizerConfig,
loop_config: TrainingLoopConfig,
learning_rate: float,
betas: tuple[float, float],
scheduler_name: str,
) -> tuple[torch.optim.Optimizer, object]:
"""Build an optimizer and LR scheduler for the scratch TE master-weight gate."""
if not params:
raise ValueError("No trainable parameters passed to build_optimizer_and_scheduler")
if os.environ.get("FASTVIDEO_TE_FP32_MASTER") == "1":
import transformer_engine.pytorch as te
optimizer = te.optimizers.FusedAdam(
params,
lr=float(learning_rate),
betas=betas,
weight_decay=float(optimizer_config.weight_decay),
eps=1e-8,
adam_w_mode=True,
master_weights=True,
master_weight_dtype=torch.float32,
exp_avg_dtype=torch.float32,
exp_avg_sq_dtype=torch.float32,
store_param_remainders=False,
)
else:
optimizer = torch.optim.AdamW(
params,
lr=float(learning_rate),
betas=betas,
weight_decay=float(optimizer_config.weight_decay),
eps=1e-8,
fused=optimizer_config.fused,
)
scheduler = get_scheduler(
str(scheduler_name),
optimizer=optimizer,
num_warmup_steps=int(optimizer_config.lr_warmup_steps),
num_training_steps=int(loop_config.max_train_steps),
num_cycles=int(optimizer_config.lr_num_cycles),
power=float(optimizer_config.lr_power),
min_lr_ratio=float(optimizer_config.min_lr_ratio),
last_epoch=-1,
)
return optimizer, scheduler
def clip_grad_norm_if_needed(
module: torch.nn.Module,
max_grad_norm: float,
) -> torch.Tensor | None:
if max_grad_norm <= 0.0:
return None
return clip_grad_norm_while_handling_failing_dtensor_cases(
[p for p in module.parameters()],
max_grad_norm,
foreach=None,
)
@@ -0,0 +1,71 @@
#!/usr/bin/env python3
"""Current-head packed LTX-2 fixed-arena gate."""
from __future__ import annotations
import importlib.util
import json
from pathlib import Path
import sys
from typing import Any
import torch
import torch.distributed as dist
EXPECTED_HEAD = "fa47ce1ab570d33bb245a49f4cd63267282b2a54"
BASE_PATH = Path(__file__).with_name("zero2_ltx2_input_probe.py")
spec = importlib.util.spec_from_file_location("zero2_ltx2_input_probe", BASE_PATH)
if spec is None or spec.loader is None:
raise RuntimeError(f"could not load fixed-arena base harness: {BASE_PATH}")
base = importlib.util.module_from_spec(spec)
sys.modules[spec.name] = base
spec.loader.exec_module(base)
base.EXPECTED_SHA = EXPECTED_HEAD
_original_init = base.LTX2FixedArenaZero2.__init__
def _checked_init(self: Any, transformer: torch.nn.Module, **kwargs: Any) -> None:
names = [name for name, parameter in transformer.named_parameters() if parameter.requires_grad]
packed_qkv = [name for name in names if ".attn1.to_qkv." in name]
packed_kv = [name for name in names if ".attn2.to_kv." in name]
split = [
name for name in names
if any(token in name for token in (".attn1.to_q.", ".attn1.to_k.", ".attn1.to_v.",
".attn2.to_k.", ".attn2.to_v."))
]
if len(names) != 927 or len(packed_qkv) != 96 or len(packed_kv) != 96 or split:
raise RuntimeError(
"packed LTX-2 projection layout was not active: "
f"params={len(names)} qkv={len(packed_qkv)} kv={len(packed_kv)} split={len(split)}")
_original_init(self, transformer, **kwargs)
fp32_fields = ("master", "master_grad", "exp_avg", "exp_avg_sq", "step")
if self.param_arena.dtype != torch.bfloat16 or self.grad_arena.dtype != torch.bfloat16:
raise RuntimeError("fixed working parameter and gradient arenas must be BF16")
if any(getattr(bucket, field).dtype != torch.float32 for bucket in self.buckets for field in fp32_fields):
raise RuntimeError("fixed-arena masters, gradients, moments, and steps must be FP32")
if not dist.is_initialized() or dist.get_rank() == 0:
print(
"FIXED_ARENA_LAYOUT " + json.dumps({
"trainable_parameter_objects": len(names),
"packed_qkv_parameters": len(packed_qkv),
"packed_kv_parameters": len(packed_kv),
"split_projection_parameters": len(split),
"arena_numel": self.param_arena.numel(),
"working_dtype": str(self.param_arena.dtype),
"master_dtype": str(self.buckets[0].master.dtype),
"master_precision_delta_max": self.master_precision_delta_max,
"moment_dtype": str(self.buckets[0].exp_avg.dtype),
}, sort_keys=True),
flush=True,
)
base.LTX2FixedArenaZero2.__init__ = _checked_init
if __name__ == "__main__":
base.main()
@@ -0,0 +1,36 @@
#!/usr/bin/env python3
import json
import os
import torch
import torch.distributed as dist
from fastvideo.distributed.parallel_state import (
get_world_group,
maybe_init_distributed_environment_and_model_parallel,
)
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
torch.cuda.set_device(local_rank)
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=1)
value = torch.tensor(float(rank + 1), device=f"cuda:{local_rank}")
dist.all_reduce(value, group=get_world_group().device_group)
torch.cuda.synchronize(local_rank)
print(
"NCCL_SMOKE "
+ json.dumps(
{
"rank": rank,
"local_rank": local_rank,
"world_size": dist.get_world_size(),
"sum": value.item(),
"host": os.uname().nodename,
},
sort_keys=True,
),
flush=True,
)
dist.barrier()
@@ -0,0 +1,85 @@
#!/usr/bin/env python3
import json
import os
import statistics
from datetime import timedelta
import torch
import torch.distributed as dist
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
torch.cuda.set_device(local_rank)
dist.init_process_group("nccl", timeout=timedelta(seconds=180))
scalar = torch.tensor(float(rank + 1), device=f"cuda:{local_rank}")
dist.all_reduce(scalar)
torch.cuda.synchronize(local_rank)
assert scalar.item() == 36.0
print(
"HEALTH_GATE "
+ json.dumps(
{
"phase": "scalar",
"rank": rank,
"local_rank": local_rank,
"world_size": dist.get_world_size(),
"sum": scalar.item(),
"host": os.uname().nodename,
},
sort_keys=True,
),
flush=True,
)
world_size = dist.get_world_size()
payload = torch.ones(128 * 1024 * 1024, dtype=torch.bfloat16, device=f"cuda:{local_rank}")
for _ in range(8):
dist.all_reduce(payload)
payload.div_(world_size)
torch.cuda.synchronize(local_rank)
dist.barrier()
latencies_ms = []
for _ in range(128):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
dist.all_reduce(payload)
end.record()
end.synchronize()
latencies_ms.append(start.elapsed_time(end))
payload.div_(world_size)
torch.cuda.synchronize(local_rank)
assert payload[0].item() == 1.0
sorted_ms = sorted(latencies_ms)
mean_ms = statistics.fmean(latencies_ms)
max_mean = torch.tensor(mean_ms, device=f"cuda:{local_rank}")
dist.all_reduce(max_mean, op=dist.ReduceOp.MAX)
payload_gb = payload.numel() * payload.element_size() / 1e9
print(
"HEALTH_GATE "
+ json.dumps(
{
"phase": "sustained_all_reduce",
"rank": rank,
"local_rank": local_rank,
"host": os.uname().nodename,
"iterations": len(latencies_ms),
"payload_bytes": payload.numel() * payload.element_size(),
"mean_ms": mean_ms,
"p50_ms": statistics.median(sorted_ms),
"p95_ms": sorted_ms[int(0.95 * (len(sorted_ms) - 1))],
"max_ms": max(sorted_ms),
"slowest_rank_mean_ms": max_mean.item(),
"slowest_rank_bus_gbps": payload_gb * (2 * (world_size - 1) / world_size) / (max_mean.item() / 1000),
"value": payload[0].item(),
},
sort_keys=True,
),
flush=True,
)
dist.barrier()
dist.destroy_process_group()
@@ -0,0 +1,38 @@
#!/usr/bin/env python3
"""Restore the historical 2-D mesh, then execute a benchmark driver."""
from __future__ import annotations
import json
import os
import runpy
import sys
from fastvideo.models.loader import fsdp_load
def _init_2d_mesh(device_type: str, replicate_dim: int, shard_dim: int):
mesh = fsdp_load.init_device_mesh(
device_type,
mesh_shape=(replicate_dim, shard_dim),
mesh_dim_names=("replicate", "shard"),
)
if int(os.environ.get("LOCAL_RANK", "0")) == 0:
print(
"BF16_FSDP_MESH " + json.dumps({
"control": "historical_2d",
"device_type": device_type,
"mesh_shape": list(mesh.mesh.shape),
"mesh_dim_names": list(mesh.mesh_dim_names or ()),
"replicate_dim": replicate_dim,
"shard_dim": shard_dim,
}, sort_keys=True),
flush=True,
)
return mesh
fsdp_load._init_fsdp_device_mesh = _init_2d_mesh
target = sys.argv.pop(1)
sys.argv[0] = target
runpy.run_path(target, run_name="__main__")
@@ -0,0 +1,36 @@
#!/usr/bin/env python3
"""Print the selected FSDP mesh, then execute a benchmark driver."""
from __future__ import annotations
import json
import os
import runpy
import sys
from fastvideo.models.loader import fsdp_load
_original_init = fsdp_load._init_fsdp_device_mesh
def _logged_init(device_type: str, replicate_dim: int, shard_dim: int):
mesh = _original_init(device_type, replicate_dim, shard_dim)
if int(os.environ.get("LOCAL_RANK", "0")) == 0:
print(
"BF16_FSDP_MESH " + json.dumps({
"device_type": device_type,
"mesh_shape": list(mesh.mesh.shape),
"mesh_dim_names": list(mesh.mesh_dim_names or ()),
"replicate_dim": replicate_dim,
"shard_dim": shard_dim,
}, sort_keys=True),
flush=True,
)
return mesh
fsdp_load._init_fsdp_device_mesh = _logged_init
target = sys.argv.pop(1)
sys.argv[0] = target
runpy.run_path(target, run_name="__main__")

Some files were not shown because too many files have changed in this diff Show More