Compare commits
33
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c2becf42c2 | ||
|
|
9c86d28e11 | ||
|
|
a7317d2217 | ||
|
|
8bf6eff18f | ||
|
|
acad815f5c | ||
|
|
971827ad95 | ||
|
|
e82addb927 | ||
|
|
26b455b7b6 | ||
|
|
c459a18978 | ||
|
|
9435b3f29f | ||
|
|
3083c59ef4 | ||
|
|
6199cbee66 | ||
|
|
52f1114dd9 | ||
|
|
3f3f06541c | ||
|
|
0e60a0e9cc | ||
|
|
20c36acefc | ||
|
|
fa47ce1ab5 | ||
|
|
7afd751915 | ||
|
|
cc5913e5c4 | ||
|
|
26f909c520 | ||
|
|
d016237096 | ||
|
|
002ec0771b | ||
|
|
49508050b7 | ||
|
|
7f6c290c93 | ||
|
|
7c58950a92 | ||
|
|
0b324d0a40 | ||
|
|
e42cfa5e5b | ||
|
|
7f139e2b28 | ||
|
|
acdbc0a614 | ||
|
|
421227ecd4 | ||
|
|
0431d4e8a9 | ||
|
|
cfc110cb6b | ||
|
|
9582f012df |
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 +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.
|
||||
|
||||
@@ -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]))
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
__pycache__/
|
||||
*.log
|
||||
*.sqlite
|
||||
*.trace.json*
|
||||
*.pt
|
||||
*.safetensors
|
||||
*.mp4
|
||||
wandb/
|
||||
outputs/
|
||||
@@ -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.
|
||||
@@ -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.
|
||||
@@ -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()
|
||||
+134
@@ -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", ®ister_case, "register pinned cuBLASLt case");
|
||||
module.def("lt_mm", <_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
Reference in New Issue
Block a user