diff --git a/docs/quantization/h3_int8_affine.md b/docs/quantization/h3_int8_affine.md new file mode 100644 index 0000000000..d57724db73 --- /dev/null +++ b/docs/quantization/h3_int8_affine.md @@ -0,0 +1,16 @@ +# INT8 affine (group-64) for MiniMax-H3 + +Weight-only load-time INT8 for the H3 DiT on CUDA. +Implementation: `fastvideo/layers/quantization/int8_affine_config.py`. + +```python +from fastvideo.layers.quantization.int8_affine_config import INT8AffineConfig + +fastvideo_args.transformer_quant = INT8AffineConfig.for_minimax_h3() +``` + +Quantizes attention and FFN linears. Excludes `attn.to_gate_compress`, +`adaln_basis`, fp32-pinned I/O projections, and norms. + +Inference only; not the MLX QAT callback. Tests: +`fastvideo/tests/ops/quantization/test_int8_affine_config.py`. diff --git a/docs/quantization/h3_nvfp4.md b/docs/quantization/h3_nvfp4.md new file mode 100644 index 0000000000..f9619aad57 --- /dev/null +++ b/docs/quantization/h3_nvfp4.md @@ -0,0 +1,16 @@ +# NVFP4 for MiniMax-H3 + +Load-time NVFP4 for the H3 DiT on Blackwell (sm100+). +Implementation: `fastvideo/layers/quantization/nvfp4_config.py` (FlashInfer). + +```python +from fastvideo.layers.quantization.nvfp4_config import NVFP4Config + +fastvideo_args.transformer_quant = NVFP4Config.for_minimax_h3() +``` + +Requires FlashInfer with NVFP4 support. Excludes `attn.to_gate_compress`. +Compact checkpoints use the NVFP4 sidecar helpers in the same module. + +Tests: `fastvideo/tests/ops/quantization/test_nvfp4_h3_prefixes.py`, +`test_nvfp4_sidecar.py`. diff --git a/docs/quantization/h3_w4a16.md b/docs/quantization/h3_w4a16.md new file mode 100644 index 0000000000..80328ce167 --- /dev/null +++ b/docs/quantization/h3_w4a16.md @@ -0,0 +1,15 @@ +# W4A16 for MiniMax-H3 + +Weight-only 4-bit storage with bf16/fp16 activations. +Implementation: `fastvideo/layers/quantization/w4a16_config.py`. + +```python +from fastvideo.layers.quantization.w4a16_config import W4A16Config + +fastvideo_args.transformer_quant = W4A16Config.for_minimax_h3() +``` + +This is a memory lane: weights are stored in 4-bit form and dequantized before +each dense GEMM. There is no fused W4A16 kernel in-tree yet. + +Tests: `fastvideo/tests/ops/quantization/test_w4a16_config.py`. diff --git a/docs/quantization/loader_quant_params.md b/docs/quantization/loader_quant_params.md new file mode 100644 index 0000000000..c89f3d4e1c --- /dev/null +++ b/docs/quantization/loader_quant_params.md @@ -0,0 +1,17 @@ +# Quantized models and loader allowlists + +Some quantization methods register scale tensors that are absent from dense +checkpoints. Without an allowlist entry, loading fails with: + +``` +Unsupported new parameter: ...scale_weight... +``` + +`ALLOWED_NEW_PARAM_PATTERNS` in `fastvideo/models/loader/fsdp_load.py` admits +expected new parameter leaf names via substring match. Prefer +`register_buffer(..., persistent=False)` for values recomputed at load time so +they never enter `state_dict()`. + +When adding a quant config that calls `register_parameter` for scales, add the +leaf name to `ALLOWED_NEW_PARAM_PATTERNS` and mirror it in +`fastvideo/models/loader/shard_cache.py`. diff --git a/docs/training/attn_qat.md b/docs/training/attn_qat.md index 8437da9cbd..1703e0399c 100644 --- a/docs/training/attn_qat.md +++ b/docs/training/attn_qat.md @@ -105,7 +105,7 @@ The migrated recipe preserves these behaviors: |---|---| | Student fake-quantized attention | `models.student.attention_backend: ATTN_QAT_TRAIN` | | Teacher and critic full-precision attention | Role-local `FLASH_ATTN` | -| Generator update every five critic steps | `method.generator_update_interval: 5` | +| Four critic-only steps, then one student-only step | `method.generator_update_interval: 5` | | Three-step rollout | `method.dmd_denoising_steps: [1000, 757, 522]` | | Score timestep range | `method.min_timestep_ratio: 0.02`, `max_timestep_ratio: 0.98` | | Legacy guidance `cond + 2(cond - uncond)` | Standard CFG scale `3.0` | diff --git a/examples/compacth3/README.md b/examples/compacth3/README.md new file mode 100644 index 0000000000..d10315b740 --- /dev/null +++ b/examples/compacth3/README.md @@ -0,0 +1,10 @@ +# CompactH3 + +| Area | Path | +|---|---| +| Block scoring, folding, recovery | `scripts/fasth3_sprint/` | +| Recovery configs | `examples/train/configs/fasth3_*.yaml` | +| DMD2 / QAD | `examples/train/configs/distribution_matching/minimax_h3/` | +| Eval, QAD, export | `scripts/compacth3/` | + +Checkpoints and generated media are excluded. Override cluster paths in launchers when deploying elsewhere. diff --git a/examples/distill/MiniMax-H3/distill_dmd.sh b/examples/distill/MiniMax-H3/distill_dmd.sh new file mode 100755 index 0000000000..0219ac2b0e --- /dev/null +++ b/examples/distill/MiniMax-H3/distill_dmd.sh @@ -0,0 +1,26 @@ +#!/usr/bin/env bash + +set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +REPO_ROOT="$(cd "${SCRIPT_DIR}/../../.." && pwd)" +cd "${REPO_ROOT}" + +export MASTER_PORT="${MASTER_PORT:-29513}" +export FASTVIDEO_FA4="${FASTVIDEO_FA4:-1}" + +export NUM_GPUS="${NUM_GPUS:-4}" +WORLD_SIZE="${NUM_GPUS}" +SP_SIZE="${SP_SIZE:-1}" +HSDP_REPLICATE="${HSDP_REPLICATE:-1}" +HSDP_SHARD="${HSDP_SHARD:-${WORLD_SIZE}}" +CONFIG="${CONFIG:-examples/train/configs/distribution_matching/minimax_h3/dmd2_sp1_fsdp40_vidprom_v6.yaml}" +OUTPUT_DIR="${OUTPUT_DIR:-outputs/minimax_h3_dmd2_local}" + +exec bash examples/train/run.sh "${CONFIG}" \ + --training.distributed.num_gpus "${WORLD_SIZE}" \ + --training.distributed.sp_size "${SP_SIZE}" \ + --training.distributed.hsdp_replicate_dim "${HSDP_REPLICATE}" \ + --training.distributed.hsdp_shard_dim "${HSDP_SHARD}" \ + --training.checkpoint.output_dir "${OUTPUT_DIR}" \ + "$@" diff --git a/examples/inference/minimax_h3/README.md b/examples/inference/minimax_h3/README.md new file mode 100644 index 0000000000..0994e1d9dd --- /dev/null +++ b/examples/inference/minimax_h3/README.md @@ -0,0 +1,63 @@ +# MiniMax-H3 inference examples + +Basic single-request H3 examples live in `examples/inference/basic/` +(`basic_minimax_h3_t2v.py`, `basic_minimax_h3_fl2va.py`, +`basic_minimax_h3_ref2va.py`). This directory holds H3-specific benchmark +tooling. + +## `h3_vsa_dmd.py` — VSA-H3 vs dense attention, few-step DMD inference + +Benchmarks 3-step (DMD-style) H3 T2VA inference under two attention +backends and prints a latency/speedup table: + +- `dense` — `FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN` with FA4 + (`FASTVIDEO_FA4=1`). If the flash-attn package is not installed the + FLASH_ATTN request falls back to Torch SDPA (the worker log prints + "Using Torch SDPA backend"); the baseline is then SDPA, not FA4. +- `vsa` — `FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3` at + `--sparsity` (default 0.9), applied at generator boot through + `FastVideoArgs.VSA_sparsity` (`pipeline.experimental`). + `--vsa-tile-size {64,256}` (default 256) flows the same way + (`FastVideoArgs.VSA_tile_size`). At tile 256, `--vsa-kernel triton` + (default, no optional dependencies) uses the 256-to-64 expansion path + and `cutedsl` opts into the FA4 CuTe 256-tile forward (requires the + optional FA4 CuTe build, `flash_attn.cute`); at tile 64 the forward is + always the native 64-token Triton kernel and `--vsa-kernel` is ignored. +- `microbench` — model-free per-attention-layer proxy on the exact packed + H3 sequence geometry (dense FA4/SDPA vs the full `MiniMaxH3VSAImpl` + tile/pool/top-k/kernel/untile path). Useful standalone, and as the + speedup proxy when the full VSA pipeline leg is unavailable. + +Each mode boots its own generator in a fresh subprocess (the backend env +var is resolved at boot), runs `--warmup` untimed request(s), then times +`--num-prompts` requests with fixed seeds shared across modes so the +per-mode videos can be eyeballed against each other. Model-load time is +reported separately from per-request latency. A crash in one mode is +contained: its signature is saved to `//crash_signature.txt` +and the remaining modes still report. + +```bash +FASTVIDEO_FA4=1 python examples/inference/minimax_h3/h3_vsa_dmd.py \ + --model-path /path/to/MiniMax-H3 \ + --prompts-json /path/to/validation.json \ + --num-prompts 4 \ + --output-dir outputs/h3_vsa_dmd \ + --modes dense,vsa,microbench \ + --dmd-steps 1000,667,333 \ + --num-gpus 4 +``` + +`--prompts-json` expects `{"data": [{"caption": ...}]}`; without it a +built-in prompt set is used. Results land in +`//results.json`, per-mode videos in `//`, and +an aggregate `summary.json` plus a final table on stdout. + +### Caveat: dense-trained checkpoints under VSA + +With the base (dense-trained) H3 checkpoint this benchmark measures SPEED +only. The base model was never trained under VSA top-k masks, so at 90% +sparsity output-quality parity is not expected — judge quality with a +VSA-trained (sparse-student) DMD checkpoint. The 3-step DMD ladder applied +to the base checkpoint is likewise a latency proxy for a distilled +student, not a quality reference: real few-step quality requires a DMD +student checkpoint. diff --git a/examples/inference/minimax_h3/h3_vsa_dmd.py b/examples/inference/minimax_h3/h3_vsa_dmd.py new file mode 100644 index 0000000000..69824520cf --- /dev/null +++ b/examples/inference/minimax_h3/h3_vsa_dmd.py @@ -0,0 +1,604 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Benchmark MiniMax-H3 few-step (DMD-style) inference: VSA-H3 vs dense attention. + +Runs the same prompts (same seeds) through the full T2VA pipeline once per +attention mode and reports per-request end-to-end latency plus the +denoising-stage time (``FASTVIDEO_STAGE_LOGGING=1``): + +- ``dense``: FLASH_ATTN with the FA4 CuTe kernels (``FASTVIDEO_FA4=1``). + When the flash-attn package is not installed, the FLASH_ATTN request + falls back to Torch SDPA (the worker log prints "Using Torch SDPA + backend") — the reported baseline is then SDPA, not FA4. +- ``vsa``: VIDEO_SPARSE_ATTN_H3 at ``--sparsity`` (default 0.9). The + sparsity is applied at generator boot via ``FastVideoArgs.VSA_sparsity`` + (``pipeline.experimental``); the H3 denoising stage builds per-step VSA + metadata from it. ``--vsa-tile-size`` (default 256) flows the same way + (``FastVideoArgs.VSA_tile_size``) and selects the tile geometry: at 256, + ``--vsa-kernel triton`` (default, no optional deps) uses the 256-to-64 + expansion path while ``cutedsl`` opts into the FA4 CuTe 256-tile forward + (requires the optional FA4 CuTe build, ``flash_attn.cute``); at 64 the + block map is already at kernel granularity, so the forward always runs + the native 64-token Triton kernel and ``--vsa-kernel`` does not apply. +- ``microbench``: model-free attention-layer microbenchmark on the exact + packed H3 sequence geometry of the requested video shape. Times + ``block_sparse_attn_256_bshd`` (Triton and, when importable, the FA4 CuTe + path) through the real ``MiniMaxH3VSAImpl`` tile/pool/top-k/untile path + against dense flash attention and torch SDPA. Use it as the speedup proxy + when the full VSA pipeline leg is unavailable. + +The attention backend is resolved at generator boot, so each mode runs in a +fresh subprocess (one generator boot per mode). This also isolates the modes +from each other's CUDA state: a crash in one leg still leaves the other +legs' numbers and the final table intact. + +Example (one 4-GPU node): + + FASTVIDEO_FA4=1 python examples/inference/minimax_h3/h3_vsa_dmd.py \\ + --model-path /path/to/MiniMax-H3 \\ + --prompts-json validation.json --num-prompts 4 \\ + --output-dir outputs/h3_vsa_dmd --modes dense,vsa,microbench + +Caveat — sparse-trained students: with the base (dense-trained) checkpoint +this benchmark measures SPEED only. The base model was never trained under +VSA top-k masks, so at 90% sparsity output-quality parity is not expected; +the per-mode videos are written for eyeballing, but judge quality with a +VSA-trained DMD student checkpoint. Likewise the 3-step DMD ladder applied +to the base checkpoint is a latency proxy for a distilled student, not a +quality reference. +""" + +from __future__ import annotations + +import argparse +import json +import os +import signal +import statistics +import subprocess +import sys +import threading +import time +from collections import deque +from pathlib import Path + +GENERATION_MODES = ("dense", "vsa") +ALL_MODES = GENERATION_MODES + ("microbench",) + +MODE_ENV: dict[str, dict[str, str]] = { + "dense": { + "FASTVIDEO_ATTENTION_BACKEND": "FLASH_ATTN", + "FASTVIDEO_FA4": "1", + }, + "vsa": { + "FASTVIDEO_ATTENTION_BACKEND": "VIDEO_SPARSE_ATTN_H3", + "FASTVIDEO_FA4": "1", + }, + "microbench": { + "FASTVIDEO_FA4": "1", + }, +} + +DEFAULT_PROMPTS = [ + "A cinematic drone shot over coastal cliffs at sunrise, golden light, gentle ocean waves, ultra detailed.", + "A barista pours latte art in a warm cafe, steam rising, shallow depth of field, soft morning light.", + "A red fox trots across fresh snow between pine trees, breath visible in the cold air, tracking shot.", + "Neon-lit rain-soaked city street at night, reflections on wet asphalt, pedestrians with umbrellas.", +] + +WARMUP_SEED = 999 +FIRST_SEED = 1000 + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument("--model-path", required=True, help="Full modular MiniMax-H3 pipeline directory") + parser.add_argument("--prompts-json", default=None, help='Optional {"data": [{"caption": ...}]} prompt file') + parser.add_argument("--num-prompts", type=int, default=4, help="Timed requests per mode") + parser.add_argument("--sparsity", type=float, default=0.9, help="VSA sparsity for the vsa/microbench modes") + parser.add_argument("--output-dir", default="outputs/h3_vsa_dmd") + parser.add_argument("--modes", default="dense,vsa", help=f"Comma-separated subset of {ALL_MODES}") + parser.add_argument("--dmd-steps", default="1000,667,333", help="FASTVIDEO_DMD_DENOISING_STEPS ladder") + parser.add_argument("--dense-native-steps", + type=int, + default=None, + help="Run the dense mode on the scheduler's NATIVE n-step schedule instead of the " + "DMD ladder (FASTVIDEO_DMD_DENOISING_STEPS is removed for that mode). E.g. 50 turns " + "the dense leg into the teacher-style 50-step baseline, so the table compares " + "50-step dense against the few-step DMD VSA leg. Note the H3 scheduler builds an " + "n-point sigma grid ending at 0, i.e. n-1 transformer forwards") + parser.add_argument("--num-gpus", type=int, default=4) + parser.add_argument("--height", type=int, default=768) + parser.add_argument("--width", type=int, default=1344) + parser.add_argument("--num-frames", type=int, default=124) + parser.add_argument("--warmup", type=int, default=1, help="Untimed warm-up requests per mode") + parser.add_argument("--vsa-kernel", + choices=("triton", "cutedsl"), + default="triton", + help="VSA-256 kernel path; cutedsl needs the optional FA4 CuTe build " + "(ignored at --vsa-tile-size 64, which is native-Triton only)") + parser.add_argument("--vsa-tile-size", + type=int, + choices=(64, 256), + default=256, + help="VSA-H3 tile size in tokens, plumbed like sparsity via " + "FastVideoArgs.VSA_tile_size; 64 runs the native Triton block-sparse forward") + parser.add_argument("--mode-timeout", type=int, default=5400, help="Hard per-mode timeout in seconds") + parser.add_argument("--microbench-text-tokens", type=int, default=300, help="Assumed text prefix length") + parser.add_argument("--microbench-heads", default="14,56", help="Per-GPU head counts to microbench") + parser.add_argument("--_worker", choices=ALL_MODES, default=None, help=argparse.SUPPRESS) + return parser.parse_args() + + +def load_prompts(args: argparse.Namespace) -> list[str]: + if args.prompts_json: + with open(args.prompts_json) as handle: + rows = json.load(handle)["data"] + prompts = [row["caption"] for row in rows[:args.num_prompts]] + else: + prompts = [DEFAULT_PROMPTS[i % len(DEFAULT_PROMPTS)] for i in range(args.num_prompts)] + if len(prompts) < args.num_prompts: + raise ValueError(f"Requested {args.num_prompts} prompts but only {len(prompts)} available.") + return prompts + + +def dense_native_steps(mode: str, args: argparse.Namespace) -> int | None: + return args.dense_native_steps if (mode == "dense" and args.dense_native_steps) else None + + +def apply_worker_env(mode: str, args: argparse.Namespace) -> None: + """Set the mode's environment. Must run before any fastvideo import.""" + env = dict(MODE_ENV[mode]) + if dense_native_steps(mode, args): + os.environ.pop("FASTVIDEO_DMD_DENOISING_STEPS", None) + env["FASTVIDEO_DMD_DEBUG_STATS"] = "1" + else: + env["FASTVIDEO_DMD_DENOISING_STEPS"] = args.dmd_steps + env["FASTVIDEO_STAGE_LOGGING"] = "1" # per-stage timings on the result object + env["FASTVIDEO_VSA_CUTEDSL"] = "1" if args.vsa_kernel == "cutedsl" else "0" + os.environ.update(env) + + +def denoise_seconds(result) -> float | None: + stages = getattr(getattr(result, "logging_info", None), "stages", None) + if not stages: + return None + for stage_name, metrics in stages.items(): + if "denois" in stage_name.lower(): + execution_time = metrics.get("execution_time") + if execution_time is not None: + return float(execution_time) + return None + + +def run_generation_worker(args: argparse.Namespace) -> int: + mode = args._worker + apply_worker_env(mode, args) + mode_dir = Path(args.output_dir) / mode + mode_dir.mkdir(parents=True, exist_ok=True) + prompts = load_prompts(args) + dmd_steps = [int(step) for step in args.dmd_steps.split(",") if step.strip()] + native_steps = dense_native_steps(mode, args) + num_inference_steps = native_steps or len(dmd_steps) + + from fastvideo import VideoGenerator + from fastvideo.api import ( + EngineConfig, + GenerationRequest, + GeneratorConfig, + OffloadConfig, + OutputConfig, + ParallelismConfig, + PipelineSelection, + SamplingConfig, + ) + + experimental: dict[str, float] = {} + if mode == "vsa": + experimental["VSA_sparsity"] = args.sparsity + experimental["VSA_tile_size"] = args.vsa_tile_size + + schedule = (f"native {native_steps}-step schedule" if native_steps else f"dmd_steps={dmd_steps}") + print(f"[{mode}] booting generator (backend={os.environ['FASTVIDEO_ATTENTION_BACKEND']}, " + f"sparsity={experimental.get('VSA_sparsity', 0.0)}, " + f"tile={experimental.get('VSA_tile_size', '-')}, {schedule})", + flush=True) + boot_start = time.perf_counter() + generator = VideoGenerator.from_config( + GeneratorConfig( + model_path=args.model_path, + engine=EngineConfig( + num_gpus=args.num_gpus, + use_fsdp_inference=args.num_gpus > 1, + parallelism=ParallelismConfig(tp_size=1, sp_size=args.num_gpus), + offload=OffloadConfig( + dit=False, + dit_layerwise=False, + text_encoder=True, + vae=True, + pin_cpu_memory=False, + ), + ), + pipeline=PipelineSelection(experimental=experimental), + )) + load_time = time.perf_counter() - boot_start + print(f"[{mode}] generator ready in {load_time:.1f}s (model load, excluded from timings)", flush=True) + + def build_request(prompt: str, seed: int, output_name: str) -> GenerationRequest: + return GenerationRequest( + prompt=prompt, + negative_prompt="", + sampling=SamplingConfig( + height=args.height, + width=args.width, + num_frames=args.num_frames, + fps=24, + num_inference_steps=num_inference_steps, + guidance_scale=1.0, + batch_cfg=False, + seed=seed, + ), + output=OutputConfig( + output_path=str(mode_dir / output_name), + save_video=True, + return_frames=False, + ), + ) + + records: list[dict] = [] + try: + for warmup_index in range(args.warmup): + print(f"[{mode}] warm-up {warmup_index} (untimed; absorbs kernel JIT/autotune)", flush=True) + warmup_start = time.perf_counter() + generator.generate(build_request(prompts[0], WARMUP_SEED, f"warmup{warmup_index}.mp4")) + print(f"[{mode}] warm-up {warmup_index} done in {time.perf_counter() - warmup_start:.1f}s", flush=True) + + for index, prompt in enumerate(prompts): + request = build_request(prompt, FIRST_SEED + index, f"prompt{index:02d}.mp4") + request_start = time.perf_counter() + result = generator.generate(request) + e2e_seconds = time.perf_counter() - request_start + record = { + "prompt_index": index, + "seed": FIRST_SEED + index, + "e2e_seconds": e2e_seconds, + "generation_seconds": result.generation_time, + "denoise_seconds": denoise_seconds(result), + "video_path": result.video_path, + "peak_memory_mb": result.peak_memory_mb, + } + records.append(record) + parts = [f"[{mode}] {index:02d} e2e={e2e_seconds:.1f}s"] + if record["generation_seconds"] is not None: + parts.append(f"gen={record['generation_seconds']:.1f}s") + if record["denoise_seconds"] is not None: + parts.append(f"denoise={record['denoise_seconds']:.1f}s") + parts.append(f"-> {result.video_path}") + print(" ".join(parts), flush=True) + finally: + payload = { + "mode": mode, + "sparsity": args.sparsity if mode == "vsa" else 0.0, + "vsa_kernel": args.vsa_kernel if mode == "vsa" else None, + "vsa_tile_size": args.vsa_tile_size if mode == "vsa" else None, + "dmd_steps": None if native_steps else dmd_steps, + "num_inference_steps": num_inference_steps, + "shape": [args.height, args.width, args.num_frames], + "num_gpus": args.num_gpus, + "load_seconds": load_time, + "requests": records, + } + (mode_dir / "results.json").write_text(json.dumps(payload, indent=2)) + generator.shutdown() + return 0 + + +def _time_cuda_call(fn, warmup: int = 3, iters: int = 10) -> float: + import torch + for _ in range(warmup): + fn() + torch.cuda.synchronize() + start = time.perf_counter() + for _ in range(iters): + fn() + torch.cuda.synchronize() + return (time.perf_counter() - start) / iters + + +def run_microbench_worker(args: argparse.Namespace) -> int: + """Per-attention-layer proxy on the exact packed H3 geometry, single GPU. + + Measures one self-attention layer's compute (no sequence-parallel + all-to-all, which is identical across backends): dense flash attention + and SDPA on the true packed length vs the full VSA-H3 path + (tile scatter + fp32 tile pooling + top-k mask + block-sparse kernel + + untile) on the 256-padded tile buffer. + """ + mode = args._worker + apply_worker_env(mode, args) + + import torch + + from fastvideo.attention.backends.video_sparse_attn_h3 import (MiniMaxH3VSAImpl, MiniMaxH3VSAMetadataBuilder) + from fastvideo.pipelines.basic.minimax_h3.packing import (MINIMAX_H3_AUDIO_CHANNELS, audio_latent_num_frames, + video_latent_num_frames) + + device = torch.device("cuda:0") + torch.manual_seed(0) + head_dim = 128 + patch_size = (1, 2, 2) + spatial_ratio = 16 # H3 video VAE spatial compression + latent_frames = video_latent_num_frames(args.num_frames) + latent_height, latent_width = args.height // spatial_ratio, args.width // spatial_ratio + n_text = args.microbench_text_tokens + n_cond = 0 # T2V: no keyframe conditioning rows + n_audio = audio_latent_num_frames(args.num_frames) * MINIMAX_H3_AUDIO_CHANNELS + n_video = ((latent_frames // patch_size[0]) * (latent_height // patch_size[1]) * (latent_width // patch_size[2])) + seq_len = n_text + n_cond + n_audio + n_video + print(f"[microbench] packed H3 sequence for {args.height}x{args.width}x{args.num_frames}: " + f"text={n_text} (assumed) + cond={n_cond} + audio={n_audio} + video={n_video} = {seq_len} rows, " + f"head_dim={head_dim}", + flush=True) + + builder = MiniMaxH3VSAMetadataBuilder() + metadata_by_sparsity = { + sparsity: builder.build( + current_timestep=0, + raw_latent_shape=(latent_frames, latent_height, latent_width), + patch_size=patch_size, + VSA_sparsity=sparsity, + prefix_segments=(n_text, n_cond, n_audio), + device=device, + exempt=True, + tile_size=args.vsa_tile_size, + ) + for sparsity in (args.sparsity, 0.0) + } + reference_metadata = metadata_by_sparsity[args.sparsity] + print(f"[microbench] tiles: prefix={reference_metadata.num_prefix_tiles} " + f"video={reference_metadata.num_video_tiles} tile_elems={reference_metadata.tile_elems} " + f"padded_len={int(reference_metadata.variable_block_sizes.numel()) * reference_metadata.tile_elems}", + flush=True) + impl = MiniMaxH3VSAImpl(num_heads=0, head_size=head_dim, causal=False, softmax_scale=1.0, prefix="blocks.0.attn") + + flash_attn_func = None + fa_version = None + try: + from fastvideo.attention.utils import flash_attn_default + flash_attn_func = flash_attn_default.flash_attn_func + fa_version = flash_attn_default.fa_version + except Exception as error: # noqa: BLE001 - report and continue with SDPA only + print(f"[microbench] flash attention unavailable ({error}); dense rows fall back to SDPA only", flush=True) + + rows: list[dict] = [] + head_counts = [int(h) for h in args.microbench_heads.split(",") if h.strip()] + for num_heads in head_counts: + note = "per-rank slice of the sp=4 run" if num_heads == 14 else "full model on one GPU" + print(f"[microbench] heads={num_heads} ({note})", flush=True) + qkv = torch.randn(3, seq_len, num_heads, head_dim, device=device, dtype=torch.bfloat16) + query, key, value = (t.contiguous() for t in qkv.unbind(0)) + + def record(name: str, milliseconds: float, num_heads: int = num_heads) -> None: + rows.append({"heads": num_heads, "name": name, "ms": milliseconds}) + print(f"[microbench] {name:<34} {milliseconds:9.3f} ms/layer-call", flush=True) + + if flash_attn_func is not None: + def dense_flash(q=query[None], k=key[None], v=value[None]): + out = flash_attn_func(q, k, v) + return out[0] if isinstance(out, tuple) else out + + record(f"dense flash (FA{fa_version})", _time_cuda_call(dense_flash) * 1e3) + + def dense_sdpa(q=query[None], k=key[None], v=value[None]): + return torch.nn.functional.scaled_dot_product_attention( + q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)) + + record("dense torch SDPA", _time_cuda_call(dense_sdpa) * 1e3) + + kernel_choices = ["triton"] + if args.vsa_kernel == "cutedsl" and args.vsa_tile_size != 64: + kernel_choices.insert(0, "cutedsl") + for kernel in kernel_choices: + os.environ["FASTVIDEO_VSA_CUTEDSL"] = "1" if kernel == "cutedsl" else "0" + for sparsity, metadata in metadata_by_sparsity.items(): + def vsa_layer(qkv=qkv, metadata=metadata): + tiled = impl.preprocess_qkv(qkv, metadata) + q, k, v = tiled.chunk(3, dim=0) + out = impl.forward(q, k, v, None, metadata) + return impl.postprocess_output(out, metadata) + + label = f"VSA-H3 {kernel} t{args.vsa_tile_size} sparsity={sparsity:.2f}" + try: + record(label, _time_cuda_call(vsa_layer) * 1e3) + except Exception as error: # noqa: BLE001 - a kernel path may be uninstalled + print(f"[microbench] {label:<34} FAILED: {error}", flush=True) + rows.append({"heads": num_heads, "name": label, "ms": None, "error": str(error)}) + + mode_dir = Path(args.output_dir) / mode + mode_dir.mkdir(parents=True, exist_ok=True) + payload = { + "mode": mode, + "sparsity": args.sparsity, + "vsa_tile_size": args.vsa_tile_size, + "geometry": { + "seq_len": seq_len, + "text": n_text, + "cond": n_cond, + "audio": n_audio, + "video": n_video, + "head_dim": head_dim, + }, + "note": ("per-layer self-attention compute only, single GPU, excludes the sequence-parallel " + "all-to-all (identical across backends); heads=14 matches one rank of the 4-GPU sp run"), + "rows": rows, + } + (mode_dir / "results.json").write_text(json.dumps(payload, indent=2)) + return 0 + + +def run_mode_subprocess(mode: str, args: argparse.Namespace) -> dict: + """Run one mode in a fresh interpreter; stream output and survive crashes.""" + mode_dir = Path(args.output_dir) / mode + mode_dir.mkdir(parents=True, exist_ok=True) + command = [sys.executable, os.path.abspath(__file__), "--_worker", mode] + for key, value in vars(args).items(): + if key in ("_worker",) or value is None: + continue + command.extend([f"--{key.replace('_', '-')}", str(value)]) + child_env = dict(os.environ, PYTHONUNBUFFERED="1") + + print(f"\n=== mode {mode}: launching worker ===", flush=True) + start = time.perf_counter() + process = subprocess.Popen( + command, + env=child_env, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + start_new_session=True, + ) + + def kill_group() -> None: + print(f"=== mode {mode}: timeout after {args.mode_timeout}s, killing process group ===", flush=True) + try: + os.killpg(process.pid, signal.SIGKILL) + except ProcessLookupError: + pass + + watchdog = threading.Timer(args.mode_timeout, kill_group) + watchdog.start() + tail: deque[str] = deque(maxlen=60) + log_path = mode_dir / "worker.log" + with open(log_path, "w") as log_file: + assert process.stdout is not None + for line in process.stdout: + print(line, end="", flush=True) + log_file.write(line) + tail.append(line.rstrip("\n")) + return_code = process.wait() + watchdog.cancel() + elapsed = time.perf_counter() - start + + results_path = mode_dir / "results.json" + results = json.loads(results_path.read_text()) if results_path.exists() else None + status = {"mode": mode, "return_code": return_code, "elapsed_seconds": elapsed, "results": results} + if return_code != 0 or results is None: + signature = [line for line in tail if any(token in line for token in + ("Error", "error", "Traceback", "CUDA", "NCCL", "Signal", + "Segmentation", "terminate", "Killed"))] + status["crash_signature"] = signature[-15:] or list(tail)[-15:] + (mode_dir / "crash_signature.txt").write_text("\n".join(status["crash_signature"]) + "\n") + print(f"=== mode {mode}: FAILED (rc={return_code}); signature saved to {mode_dir / 'crash_signature.txt'} ===", + flush=True) + else: + print(f"=== mode {mode}: completed in {elapsed:.0f}s ===", flush=True) + return status + + +def _mode_stats(status: dict) -> dict | None: + results = status.get("results") + if not results or not results.get("requests"): + return None + requests = results["requests"] + generation = [r["generation_seconds"] for r in requests if r.get("generation_seconds") is not None] + denoise = [r["denoise_seconds"] for r in requests if r.get("denoise_seconds") is not None] + steps = results.get("num_inference_steps") + native = steps is not None and results.get("dmd_steps") is None + forwards = (steps - 1) if (native and steps > 1) else steps + mean_denoise = statistics.mean(denoise) if denoise else None + return { + "n": len(requests), + "steps": steps, + "load": results.get("load_seconds"), + "e2e": statistics.mean(r["e2e_seconds"] for r in requests), + "gen": statistics.mean(generation) if generation else None, + "denoise": mean_denoise, + "denoise_per_step": (mean_denoise / forwards) if mean_denoise is not None and forwards else None, + } + + +def _fmt(value: float | None, width: int, decimals: int = 1) -> str: + return f"{value:>{width}.{decimals}f}" if value is not None else f"{'-':>{width}}" + + +def _speedup(dense: dict | None, row: dict, metric: str) -> str: + if dense is None or dense.get(metric) is None or row.get(metric) in (None, 0): + return "-" + return f"{dense[metric] / row[metric]:.2f}x" + + +def summarize(statuses: list[dict], args: argparse.Namespace) -> None: + stats = {status["mode"]: _mode_stats(status) for status in statuses if status["mode"] in GENERATION_MODES} + dense = stats.get("dense") + + if args.dense_native_steps: + title = "H3 inference: 50-step-style dense baseline vs few-step DMD VSA" + schedule = (f"dense: native {args.dense_native_steps}-step schedule; " + f"vsa: dmd_steps={args.dmd_steps}") + else: + title = "H3 DMD 3-step inference: attention backend benchmark" + schedule = f"dmd_steps={args.dmd_steps}" + print(f"\n================ {title} ================") + vsa_kernel = "triton" if args.vsa_tile_size == 64 else args.vsa_kernel + print(f"shape={args.height}x{args.width}x{args.num_frames} gpus={args.num_gpus} " + f"{schedule} vsa sparsity={args.sparsity} (tile {args.vsa_tile_size}, {vsa_kernel})") + header = (f"{'mode':<12} {'n':>3} {'steps':>6} {'load(s)':>9} {'mean e2e(s)':>12} {'mean gen(s)':>12} " + f"{'mean denoise(s)':>16} {'denoise/step(s)':>16} {'e2e speedup':>12} {'denoise speedup':>16}") + print(header) + print("-" * len(header)) + for mode in ("dense", "vsa"): + label = f"vsa@{args.sparsity:.2f}" if mode == "vsa" else mode + row = stats.get(mode) + if row is None: + if any(status["mode"] == mode for status in statuses): + print(f"{label:<12} {'-':>3} {'-':>6} {'-':>9} {'CRASHED':>12} {'-':>12} {'-':>16} {'-':>16} " + f"{'-':>12} {'-':>16}") + continue + steps_text = str(row["steps"]) if row.get("steps") else "-" + print(f"{label:<12} {row['n']:>3} {steps_text:>6} {_fmt(row['load'], 9)} {_fmt(row['e2e'], 12)} " + f"{_fmt(row['gen'], 12)} {_fmt(row['denoise'], 16)} {_fmt(row['denoise_per_step'], 16, 2)} " + f"{_speedup(dense, row, 'e2e'):>12} {_speedup(dense, row, 'denoise'):>16}") + + micro = next((status for status in statuses if status["mode"] == "microbench"), None) + if micro and micro.get("results"): + results = micro["results"] + geometry = results["geometry"] + print(f"\nAttention-layer microbench (seq={geometry['seq_len']} rows: text {geometry['text']} + " + f"audio {geometry['audio']} + video {geometry['video']}; {results['note']}):") + for row in results["rows"]: + timing = f"{row['ms']:.3f} ms" if row.get("ms") is not None else f"FAILED: {row.get('error')}" + print(f" heads={row['heads']:>2} {row['name']:<34} {timing}") + if args.dense_native_steps: + print(f"\nNote: the native {args.dense_native_steps}-step schedule is a " + f"{args.dense_native_steps}-point sigma grid = {args.dense_native_steps - 1} transformer " + f"forwards; denoise/step divides by forwards ({args.dense_native_steps - 1} dense, " + f"{len([s for s in args.dmd_steps.split(',') if s.strip()])} vsa).") + print("\nNote: base checkpoint is dense-trained; the vsa leg measures speed, not quality parity.") + + +def main() -> None: + args = parse_args() + if args._worker in GENERATION_MODES: + sys.exit(run_generation_worker(args)) + if args._worker == "microbench": + sys.exit(run_microbench_worker(args)) + + modes = [mode.strip() for mode in args.modes.split(",") if mode.strip()] + unknown = sorted(set(modes) - set(ALL_MODES)) + if unknown: + raise ValueError(f"Unknown modes {unknown}; choose from {ALL_MODES}.") + + out_root = Path(args.output_dir) + out_root.mkdir(parents=True, exist_ok=True) + prompts = load_prompts(args) + (out_root / "prompts.txt").write_text("\n\n".join(f"[{i:02d}] {p}" for i, p in enumerate(prompts))) + + statuses = [run_mode_subprocess(mode, args) for mode in modes] + summarize(statuses, args) + summary_path = out_root / "summary.json" + summary_path.write_text(json.dumps(statuses, indent=2)) + print(f"\nPer-mode outputs and summary under: {out_root}") + sys.exit(0 if all(status["return_code"] == 0 and status["results"] is not None for status in statuses) else 2) + + +if __name__ == "__main__": + main() diff --git a/examples/train/configs/compacth3/release14b_recovery_wandb.yaml b/examples/train/configs/compacth3/release14b_recovery_wandb.yaml new file mode 100644 index 0000000000..bf9749c1aa --- /dev/null +++ b/examples/train/configs/compacth3/release14b_recovery_wandb.yaml @@ -0,0 +1,131 @@ +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/release-candidates/14b-backbone-r16-job7048-v1 + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.knowledge_distillation.minimax_h3_base_recovery.MiniMaxH3BaseRecoveryMethod + denoising_weight: 0.0 + teacher_velocity_weight: 1.0 + video_velocity_weight: 1.0 + audio_velocity_weight: 4.0 + audio_seam_weight: 2.0 + feature_weight: 1.0 + teacher_grid_points: 50 + feature_local_block_indices: + - 4 + - 6 + - 7 + - 8 + - 9 + - 11 + - 12 + - 13 + - 29 + - 30 + - 31 + - 32 + modality_energy_floor: 0.001 + low_sigma_interval_fraction: 0.5 + low_sigma_interval_count: 12 + student_state_probability: 0.0 + modality_grad_probe_every: 0 + allow_prompt_only: true +training: + distributed: + num_gpus: 16 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 4 + hsdp_shard_dim: 4 + pin_cpu_memory: true + data: + data_path: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/data/user-prompts-20260906/index-v3-audio-stratified + preprocessed_data_type: t2va + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 20260910 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 1.0e-05 + betas: + - 0.9 + - 0.999 + weight_decay: 0.0 + lr_scheduler: cosine + lr_warmup_steps: 100 + loop: + max_train_steps: 4000 + gradient_accumulation_steps: 4 + checkpoint: + output_dir: /path/bound/by/launcher + training_state_checkpointing_steps: 100 + use_cpu_process_group: true + checkpoints_total_limit: 3 + preserve_every_steps: 0 + preserve_steps: + - 100 + - 250 + - 500 + - 650 + - 700 + - 750 + - 800 + - 850 + - 900 + - 950 + - 1000 + - 1100 + - 1200 + - 1300 + - 1400 + - 1500 + - 2000 + - 4000 + - 5000 + - 6000 + resume_from_checkpoint: '' + tracker: + trackers: + - wandb + project_name: fasth3-14b-2step-qad-sprint + run_name: folded-release-long-recovery + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + dit_precision: fp32 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 + validation: + _target_: fastvideo.train.callbacks.validation.ValidationCallback + pipeline_target: fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline.MiniMaxH3Pipeline + dataset_file: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/job-scripts/release14b_validation_five.json + every_steps: 100 + run_at_start: false + sampling_steps: [50] + guidance_scale: 1.0 + num_frames: 124 + num_videos_per_prompt: 1 + use_validation_media_conditioning: false + offload_training_state: true + text_encoder_cpu_offload: true + vae_cpu_offload: true +pipeline: {} diff --git a/examples/train/configs/compacth3/release14b_validation_five.json b/examples/train/configs/compacth3/release14b_validation_five.json new file mode 100644 index 0000000000..8f8f06046e --- /dev/null +++ b/examples/train/configs/compacth3/release14b_validation_five.json @@ -0,0 +1,44 @@ +{ + "data": [ + { + "id": "speech_exact_chef", + "caption": "(S1) In a bright home kitchen, a chef looks straight at the camera and says [English] Fold the eggs gently and taste before you salt. A pot simmers behind her with soft bubbling and no music.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "music_vinyl_orbit", + "caption": "Close-up of a vinyl record spinning on a turntable in a dim listening room; warm analog jazz with double bass and brushed drums fills the room while the camera slowly orbits the platter.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "sfx_rain_tin_roof", + "caption": "Heavy rain hammers a tin roof over a porch swing; each droplet burst is crisp and close while thunder rolls in the distance and the wooden swing creaks in stereo.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "motion_kitesurf_carve", + "caption": "A kite surfer carves hard across choppy bay water while the camera dives alongside; spray hisses off the board edge, the sail flaps and snaps in the wind, and gulls cry overhead.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "multishot_train_diner", + "caption": "Two-shot scene: (shot 1) commuter train doors slide open with a pneumatic hiss and a station chime; (shot 2) cut to a diner interior where a waitress calls out an order while a grill sizzles.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + } + ] +} diff --git a/examples/train/configs/distribution_matching/minimax_h3/README.md b/examples/train/configs/distribution_matching/minimax_h3/README.md new file mode 100644 index 0000000000..c8a5db9dcc --- /dev/null +++ b/examples/train/configs/distribution_matching/minimax_h3/README.md @@ -0,0 +1,27 @@ +# MiniMax-H3 DMD2 distillation configs + +Few-step DMD2 distillation of the joint video/audio MiniMax-H3 transformer. +SFT configs live in +[`../../fine_tuning/minimax_h3/`](../../fine_tuning/minimax_h3/). + +## Files + +| File | Purpose | +|---|---| +| `release20b_dmd2_v12_dense.yaml` | Corrected four-call DMD2 recipe for the 42-block student. | +| `qad_nvfp4_4call.yaml` | NVFP4 QAD on the selected checkpoint-1400 student. | +| `release20b_validation_five.json` | Five-prompt validation panel for release grading. | + +## Launch + +```bash +# Topology is derived from the allocation. +sbatch --nodes=8 scripts/submit_release20b_dmd2_v12_corrected_16gpu.sbatch + +# Or use the checked-in shell launcher directly: +bash scripts/run_release20b_dmd2_v12_16gpu.sh +``` + +`FASTVIDEO_FA4=1` selects FA4 inside roles configured with `FLASH_ATTN`. Do not +set `FASTVIDEO_ATTENTION_BACKEND` globally; attention backends are configured +per role. diff --git a/examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call.yaml b/examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call.yaml new file mode 100644 index 0000000000..52c72c3d06 --- /dev/null +++ b/examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call.yaml @@ -0,0 +1,120 @@ +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3/checkpoint-1400 + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + quant_config: nvfp4_qat_train + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/release-candidates/base-h3-teacher-complete-v1 + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method + rollout_mode: simulate + rollout_carry: true + rollout_carry_slots: 8 + rollout_sample_type: ode + generator_update_interval: 5 + real_score_guidance_scale: 1.0 + dmd_denoising_steps: + - 999 + - 749 + - 500 + - 250 + min_timestep_ratio: 0.001 + max_timestep_ratio: 0.999 + score_timestep_shift: 2.4 + score_timestep_warp_max: 0.999 + score_timestep_continuous: true + fake_score_loss_space: x0 + modality_loss_weights: + video: 1.0 + audio: 1.0 + dmd_denom_floor_ratio: 0.05 + dmd_grad_cap: 100.0 + cfg_uncond: + text: zero + fake_score_learning_rate: 2.0e-06 + fake_score_betas: + - 0.9 + - 0.999 + fake_score_lr_scheduler: constant +training: + distributed: + num_gpus: 32 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 32 + pin_cpu_memory: true + data: + data_path: + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_50k_720_mixed_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_5s_768p/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_10k_720_mixed_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_10k_mixed_res_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_fastgen_vidprom_150k/data + preprocessed_data_type: text_only + native_shape_bucketing: false + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 42 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 2.0e-06 + betas: + - 0.9 + - 0.999 + weight_decay: 0.01 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 200 + gradient_accumulation_steps: 8 + checkpoint: + output_dir: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-qad-nvfp4-4call-v1 + resume_from_checkpoint: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3/checkpoint-1400 + training_state_checkpointing_steps: 25 + require_complete_training_checkpoint: true + checkpointing_start_step: 1400 + checkpoints_total_limit: 12 + tracker: + trackers: + - wandb + project_name: fasth3-14b-2step-qad-sprint + run_name: release20b-dmd2-qad-nvfp4-4call-v1 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 + validation: + _target_: fastvideo.train.callbacks.validation.ValidationCallback + pipeline_target: fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline.MiniMaxH3Pipeline + dataset_file: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/code/release20b-dmd2-v12-corrected-v17/examples/train/configs/distribution_matching/minimax_h3/release20b_validation_five.json + every_steps: 200 + run_at_start: false + sampling_steps: + - 4 + guidance_scale: 1.0 + use_record_dimensions: true + max_record_num_frames: 124 + num_videos_per_prompt: 1 + use_validation_media_conditioning: false + offload_training_state: true + unload_pipeline_after_validation: true + text_encoder_cpu_offload: true + vae_cpu_offload: true +model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + enable_torch_compile: false +dit_precision: fp32 +vsa: null diff --git a/examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call_r16.yaml b/examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call_r16.yaml new file mode 100644 index 0000000000..41a5dc69c3 --- /dev/null +++ b/examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call_r16.yaml @@ -0,0 +1,120 @@ +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3/checkpoint-1400-reparam-r16 + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + quant_config: nvfp4_qat_train + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/release-candidates/base-h3-teacher-complete-v1 + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method + rollout_mode: simulate + rollout_carry: true + rollout_carry_slots: 8 + rollout_sample_type: ode + generator_update_interval: 5 + real_score_guidance_scale: 1.0 + dmd_denoising_steps: + - 999 + - 749 + - 500 + - 250 + min_timestep_ratio: 0.001 + max_timestep_ratio: 0.999 + score_timestep_shift: 2.4 + score_timestep_warp_max: 0.999 + score_timestep_continuous: true + fake_score_loss_space: x0 + modality_loss_weights: + video: 1.0 + audio: 1.0 + dmd_denom_floor_ratio: 0.05 + dmd_grad_cap: 100.0 + cfg_uncond: + text: zero + fake_score_learning_rate: 2.0e-06 + fake_score_betas: + - 0.9 + - 0.999 + fake_score_lr_scheduler: constant +training: + distributed: + num_gpus: 32 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 32 + pin_cpu_memory: true + data: + data_path: + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_50k_720_mixed_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_5s_768p/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_10k_720_mixed_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_10k_mixed_res_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_fastgen_vidprom_150k/data + preprocessed_data_type: text_only + native_shape_bucketing: false + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 42 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 2.0e-06 + betas: + - 0.9 + - 0.999 + weight_decay: 0.01 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 200 + gradient_accumulation_steps: 8 + checkpoint: + output_dir: runs/compacth3-qad-nvfp4-4call-r16 + resume_from_checkpoint: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3/checkpoint-1400-reparam-r16 + training_state_checkpointing_steps: 25 + require_complete_training_checkpoint: true + checkpointing_start_step: 1400 + checkpoints_total_limit: 12 + tracker: + trackers: + - wandb + project_name: fasth3-14b-2step-qad-sprint + run_name: release20b-dmd2-qad-nvfp4-4call-v1 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 + validation: + _target_: fastvideo.train.callbacks.validation.ValidationCallback + pipeline_target: fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline.MiniMaxH3Pipeline + dataset_file: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/code/release20b-dmd2-v12-corrected-v17/examples/train/configs/distribution_matching/minimax_h3/release20b_validation_five.json + every_steps: 200 + run_at_start: false + sampling_steps: + - 4 + guidance_scale: 1.0 + use_record_dimensions: true + max_record_num_frames: 124 + num_videos_per_prompt: 1 + use_validation_media_conditioning: false + offload_training_state: true + unload_pipeline_after_validation: true + text_encoder_cpu_offload: true + vae_cpu_offload: true +model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + enable_torch_compile: false +dit_precision: fp32 +vsa: null diff --git a/examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call_r768.yaml b/examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call_r768.yaml new file mode 100644 index 0000000000..252d460084 --- /dev/null +++ b/examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call_r768.yaml @@ -0,0 +1,121 @@ +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3/checkpoint-1400 + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + quant_config: nvfp4_qat_train + adaln_rank: 768 + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/release-candidates/base-h3-teacher-complete-v1 + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method + rollout_mode: simulate + rollout_carry: true + rollout_carry_slots: 8 + rollout_sample_type: ode + generator_update_interval: 5 + real_score_guidance_scale: 1.0 + dmd_denoising_steps: + - 999 + - 749 + - 500 + - 250 + min_timestep_ratio: 0.001 + max_timestep_ratio: 0.999 + score_timestep_shift: 2.4 + score_timestep_warp_max: 0.999 + score_timestep_continuous: true + fake_score_loss_space: x0 + modality_loss_weights: + video: 1.0 + audio: 1.0 + dmd_denom_floor_ratio: 0.05 + dmd_grad_cap: 100.0 + cfg_uncond: + text: zero + fake_score_learning_rate: 2.0e-06 + fake_score_betas: + - 0.9 + - 0.999 + fake_score_lr_scheduler: constant +training: + distributed: + num_gpus: 32 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 32 + pin_cpu_memory: true + data: + data_path: + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_50k_720_mixed_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_5s_768p/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_10k_720_mixed_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_10k_mixed_res_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_fastgen_vidprom_150k/data + preprocessed_data_type: text_only + native_shape_bucketing: false + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 42 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 2.0e-06 + betas: + - 0.9 + - 0.999 + weight_decay: 0.01 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 200 + gradient_accumulation_steps: 8 + checkpoint: + output_dir: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-qad-nvfp4-4call-v1 + resume_from_checkpoint: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3/checkpoint-1400 + training_state_checkpointing_steps: 25 + require_complete_training_checkpoint: true + checkpointing_start_step: 1400 + checkpoints_total_limit: 12 + tracker: + trackers: + - wandb + project_name: fasth3-14b-2step-qad-sprint + run_name: release20b-dmd2-qad-nvfp4-4call-v1 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 + validation: + _target_: fastvideo.train.callbacks.validation.ValidationCallback + pipeline_target: fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline.MiniMaxH3Pipeline + dataset_file: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/code/release20b-dmd2-v12-corrected-v17/examples/train/configs/distribution_matching/minimax_h3/release20b_validation_five.json + every_steps: 200 + run_at_start: false + sampling_steps: + - 4 + guidance_scale: 1.0 + use_record_dimensions: true + max_record_num_frames: 124 + num_videos_per_prompt: 1 + use_validation_media_conditioning: false + offload_training_state: true + unload_pipeline_after_validation: true + text_encoder_cpu_offload: true + vae_cpu_offload: true +model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + enable_torch_compile: false +dit_precision: fp32 +vsa: null diff --git a/examples/train/configs/distribution_matching/minimax_h3/release20b_dmd2_v12_dense.yaml b/examples/train/configs/distribution_matching/minimax_h3/release20b_dmd2_v12_dense.yaml new file mode 100644 index 0000000000..19380f83c0 --- /dev/null +++ b/examples/train/configs/distribution_matching/minimax_h3/release20b_dmd2_v12_dense.yaml @@ -0,0 +1,127 @@ +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: __SELECTED_PARENT__ + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + construction_precision: bf16 + disable_custom_init_weights: true + attention_backend: TORCH_SDPA + critic: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: true + disable_custom_init_weights: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + +method: + _target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method + rollout_mode: simulate + rollout_carry: true + rollout_carry_slots: 8 + rollout_sample_type: ode + generator_update_interval: 5 + real_score_guidance_scale: 1.0 + dmd_denoising_steps: [999, 749, 500, 250] + min_timestep_ratio: 0.001 + max_timestep_ratio: 0.999 + score_timestep_shift: 2.4 + score_timestep_warp_max: 0.999 + score_timestep_continuous: true + fake_score_loss_space: x0 + modality_loss_weights: + video: 1.0 + audio: 1.0 + cfg_uncond: + text: zero + fake_score_learning_rate: 2.0e-6 + fake_score_betas: [0.9, 0.999] + fake_score_lr_scheduler: constant + +training: + distributed: + num_gpus: 16 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 16 + pin_cpu_memory: true + data: + data_path: + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_50k_720_mixed_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_5s_768p/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_10k_720_mixed_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_10k_mixed_res_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_fastgen_vidprom_150k/data + preprocessed_data_type: text_only + native_shape_bucketing: false + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 42 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 2.0e-6 + betas: [0.9, 0.999] + weight_decay: 0.01 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 1000 + gradient_accumulation_steps: 16 + checkpoint: + output_dir: __OUTPUT_DIR__ + resume_from_checkpoint: "" + save_inference_checkpoint_on_validation: true + inference_checkpoint_role: student + inference_checkpoint_dtype: bfloat16 + training_state_checkpointing_steps: 100 + require_complete_training_checkpoint: true + checkpointing_start_step: 100 + checkpoints_total_limit: 50 + tracker: + trackers: [wandb] + project_name: fasth3-14b-2step-qad-sprint + run_name: release20b-dmd2-v12-dense + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + enable_torch_compile: false + dit_precision: fp32 + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 + validation: + _target_: fastvideo.train.callbacks.validation.ValidationCallback + pipeline_target: fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline.MiniMaxH3Pipeline + dataset_file: examples/train/configs/distribution_matching/minimax_h3/release20b_validation_five.json + every_steps: 100 + run_at_start: true + sampling_steps: [4] + guidance_scale: 1.0 + use_record_dimensions: true + max_record_num_frames: 124 + num_videos_per_prompt: 1 + use_validation_media_conditioning: false + offload_training_state: true + text_encoder_cpu_offload: true + vae_cpu_offload: true + +pipeline: + dit_config: + uniform_parameter_dtype: false diff --git a/examples/train/configs/distribution_matching/minimax_h3/release20b_validation_five.json b/examples/train/configs/distribution_matching/minimax_h3/release20b_validation_five.json new file mode 100644 index 0000000000..d3258fa757 --- /dev/null +++ b/examples/train/configs/distribution_matching/minimax_h3/release20b_validation_five.json @@ -0,0 +1,42 @@ +[ + { + "id": "speech_exact_chef", + "caption": "(S1) In a bright home kitchen, a chef looks straight at the camera and says [English] Fold the eggs gently and taste before you salt. A pot simmers behind her with soft bubbling and no music.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "music_vinyl_orbit", + "caption": "Close-up of a vinyl record spinning on a turntable in a dim listening room; warm analog jazz with double bass and brushed drums fills the room while the camera slowly orbits the platter.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "sfx_rain_tin_roof", + "caption": "Heavy rain hammers a tin roof over a porch swing; each droplet burst is crisp and close while thunder rolls in the distance and the wooden swing creaks in stereo.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "motion_kitesurf_carve", + "caption": "A kite surfer carves hard across choppy bay water while the camera dives alongside; spray hisses off the board edge, the sail flaps and snaps in the wind, and gulls cry overhead.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "multishot_train_diner", + "caption": "Two-shot scene: (shot 1) commuter train doors slide open with a pneumatic hiss and a station chime; (shot 2) cut to a diner interior where a waitress calls out an order while a grill sizzles.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + } +] diff --git a/examples/train/configs/fasth3_14b_recovery.yaml b/examples/train/configs/fasth3_14b_recovery.yaml new file mode 100644 index 0000000000..f35454fb52 --- /dev/null +++ b/examples/train/configs/fasth3_14b_recovery.yaml @@ -0,0 +1,88 @@ +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/checkpoints/h18-candidates/dense-activation + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA + +method: + _target_: fastvideo.train.methods.knowledge_distillation.minimax_h3_recovery.MiniMaxH3RecoveryMethod + modality_energy_floor: 1.0e-3 + denoising_weight: 1.0 + teacher_velocity_weight: 1.0 + feature_weight: 0.01 + feature_local_block_indices: [4, 9, 14, 19] + +training: + distributed: + num_gpus: 16 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 16 + pin_cpu_memory: true + + data: + data_path: /mnt/nfs/vlm-aryan/fasth3-33b-20260806/data/h3_corpus + preprocessed_data_type: t2va + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 20260829 + num_latent_t: 37 + num_height: 512 + num_width: 896 + num_frames: 124 + + optimizer: + learning_rate: 1.0e-6 + betas: [0.9, 0.999] + weight_decay: 0.0 + lr_scheduler: constant + lr_warmup_steps: 0 + + loop: + max_train_steps: 200 + gradient_accumulation_steps: 1 + + checkpoint: + output_dir: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/h36-recovery/dense-activation + training_state_checkpointing_steps: 100 + checkpoints_total_limit: 2 + preserve_every_steps: 100 + preserve_steps: [200, 400, 600, 800, 1000] + resume_from_checkpoint: "" + + tracker: + trackers: [wandb] + project_name: fasth3-14b-2step-qad-sprint + run_name: h36-recovery-dense-activation-step200 + + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + + dit_precision: bf16 + +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 + ema: + _target_: fastvideo.train.callbacks.ema.EMACallback + decay: 0.9999 + start_iter: 0 + +pipeline: {} diff --git a/examples/train/configs/fasth3_base42_recovery.yaml b/examples/train/configs/fasth3_base42_recovery.yaml new file mode 100644 index 0000000000..c80c374af2 --- /dev/null +++ b/examples/train/configs/fasth3_base42_recovery.yaml @@ -0,0 +1,90 @@ +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/checkpoints/base42-uniform-v1 + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.knowledge_distillation.minimax_h3_base_recovery.MiniMaxH3BaseRecoveryMethod + denoising_weight: 1.0 + teacher_velocity_weight: 1.0 + feature_weight: 1.0 + teacher_grid_points: 50 + feature_local_block_indices: + - 3 + - 8 + - 13 + - 18 + - 24 + - 29 + - 34 + - 39 + modality_energy_floor: 0.001 +training: + distributed: + num_gpus: 4 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 4 + pin_cpu_memory: true + data: + data_path: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/data/day1-mask-split/train + preprocessed_data_type: t2va + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 20260829 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 3.0e-05 + betas: + - 0.9 + - 0.999 + weight_decay: 0.0 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 2 + gradient_accumulation_steps: 1 + checkpoint: + output_dir: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/base42-preflight + training_state_checkpointing_steps: 2 + use_cpu_process_group: true + checkpoints_total_limit: 8 + preserve_every_steps: 25 + preserve_steps: + - 2 + - 25 + - 50 + - 100 + - 200 + resume_from_checkpoint: '' + tracker: + trackers: + - wandb + project_name: fasth3-14b-2step-qad-sprint + run_name: base42-full-schedule-preflight + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + dit_precision: fp32 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 +pipeline: {} diff --git a/examples/train/configs/fasth3_detail_band_recovery.yaml b/examples/train/configs/fasth3_detail_band_recovery.yaml new file mode 100644 index 0000000000..9d8e90b3da --- /dev/null +++ b/examples/train/configs/fasth3_detail_band_recovery.yaml @@ -0,0 +1,90 @@ +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/preserved/activation42-step500-job6878/export-500 + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.knowledge_distillation.minimax_h3_base_recovery.MiniMaxH3BaseRecoveryMethod + denoising_weight: 1.0 + teacher_velocity_weight: 1.0 + video_velocity_weight: 1.0 + audio_velocity_weight: 4.0 + audio_seam_weight: 2.0 + feature_weight: 1.0 + teacher_grid_points: 50 + feature_local_block_indices: + - 7 + - 11 + - 13 + - 14 + modality_energy_floor: 0.001 + low_sigma_interval_fraction: 0.5 + low_sigma_interval_count: 12 + student_state_probability: 0.0 + modality_grad_probe_every: 0 + allow_prompt_only: false +training: + distributed: + num_gpus: 4 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 4 + pin_cpu_memory: true + data: + data_path: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/data/day1-mask-split/train + preprocessed_data_type: t2va + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 20260829 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 3.0e-05 + betas: + - 0.9 + - 0.999 + weight_decay: 0.0 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 300 + gradient_accumulation_steps: 1 + checkpoint: + output_dir: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/base42-detail-band/job-placeholder + training_state_checkpointing_steps: 150 + use_cpu_process_group: true + checkpoints_total_limit: 3 + preserve_every_steps: 150 + preserve_steps: + - 300 + resume_from_checkpoint: '' + tracker: + trackers: + - wandb + project_name: fasth3-14b-2step-qad-sprint + run_name: base42-detail-band + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + dit_precision: fp32 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 +pipeline: {} diff --git a/examples/train/configs/fasth3_detail_band_recovery34.yaml b/examples/train/configs/fasth3_detail_band_recovery34.yaml new file mode 100644 index 0000000000..50a7f6b51f --- /dev/null +++ b/examples/train/configs/fasth3_detail_band_recovery34.yaml @@ -0,0 +1,94 @@ +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/activation34-prompt58k-recovery/job-6972/export-200 + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.knowledge_distillation.minimax_h3_base_recovery.MiniMaxH3BaseRecoveryMethod + denoising_weight: 1.0 + teacher_velocity_weight: 1.0 + video_velocity_weight: 1.0 + audio_velocity_weight: 4.0 + audio_seam_weight: 2.0 + feature_weight: 1.0 + teacher_grid_points: 50 + feature_local_block_indices: + - 4 + - 5 + - 6 + - 7 + - 8 + - 9 + - 10 + - 15 + modality_energy_floor: 0.001 + low_sigma_interval_fraction: 0.5 + low_sigma_interval_count: 12 + student_state_probability: 0.0 + modality_grad_probe_every: 0 + allow_prompt_only: false +training: + distributed: + num_gpus: 4 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 4 + pin_cpu_memory: true + data: + data_path: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/data/day1-mask-split/train + preprocessed_data_type: t2va + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 20260829 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 3.0e-05 + betas: + - 0.9 + - 0.999 + weight_decay: 0.0 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 300 + gradient_accumulation_steps: 1 + checkpoint: + output_dir: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/activation34-detail-band/job-placeholder + training_state_checkpointing_steps: 150 + use_cpu_process_group: true + checkpoints_total_limit: 3 + preserve_every_steps: 150 + preserve_steps: + - 300 + resume_from_checkpoint: '' + tracker: + trackers: + - wandb + project_name: fasth3-14b-2step-qad-sprint + run_name: activation34-detail-band + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + dit_precision: fp32 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 +pipeline: {} diff --git a/examples/train/configs/fasth3_release_long_recovery.yaml b/examples/train/configs/fasth3_release_long_recovery.yaml new file mode 100644 index 0000000000..08f95176cc --- /dev/null +++ b/examples/train/configs/fasth3_release_long_recovery.yaml @@ -0,0 +1,82 @@ +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /path/bound/by/launcher + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.knowledge_distillation.minimax_h3_base_recovery.MiniMaxH3BaseRecoveryMethod + denoising_weight: 0.0 + teacher_velocity_weight: 1.0 + video_velocity_weight: 1.0 + audio_velocity_weight: 4.0 + audio_seam_weight: 2.0 + feature_weight: 1.0 + teacher_grid_points: 50 + feature_local_block_indices: [] + modality_energy_floor: 0.001 + low_sigma_interval_fraction: 0.5 + low_sigma_interval_count: 12 + student_state_probability: 0.0 + modality_grad_probe_every: 0 + allow_prompt_only: true +training: + distributed: + num_gpus: 16 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 4 + hsdp_shard_dim: 4 + pin_cpu_memory: true + data: + data_path: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/data/user-prompts-20260906/index-v3-audio-stratified + preprocessed_data_type: t2va + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 20260910 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 1.0e-05 + betas: [0.9, 0.999] + weight_decay: 0.0 + lr_scheduler: cosine + lr_warmup_steps: 100 + loop: + max_train_steps: 4000 + gradient_accumulation_steps: 4 + checkpoint: + output_dir: /path/bound/by/launcher + training_state_checkpointing_steps: 50 + use_cpu_process_group: true + checkpoints_total_limit: 3 + preserve_every_steps: 0 + preserve_steps: [100, 250, 500, 1000, 2000, 4000, 5000, 6000] + resume_from_checkpoint: '' + tracker: + trackers: [wandb] + project_name: fasth3-14b-2step-qad-sprint + run_name: folded-release-long-recovery + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + dit_precision: fp32 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 +pipeline: {} diff --git a/examples/train/configs/overfit_minimax_h3_t2va.yaml b/examples/train/configs/overfit_minimax_h3_t2va.yaml index ac29b37826..d3d12448b2 100644 --- a/examples/train/configs/overfit_minimax_h3_t2va.yaml +++ b/examples/train/configs/overfit_minimax_h3_t2va.yaml @@ -1,7 +1,3 @@ -# MiniMax H3 T2VA 400-step overfit experiment on one Crush-Smol record. -# -# Submit eight H200 nodes through the fixed launch script: -# bash examples/train/launch_minimax_h3_t2va_crush_smol_validation.sh models: student: @@ -83,4 +79,6 @@ callbacks: text_encoder_cpu_offload: true vae_cpu_offload: true -pipeline: {} +pipeline: + dit_config: + uniform_parameter_dtype: true diff --git a/examples/training/fasth3_14b_2step_qad/comparison_five_prompts.json b/examples/training/fasth3_14b_2step_qad/comparison_five_prompts.json new file mode 100644 index 0000000000..e972fd9dd3 --- /dev/null +++ b/examples/training/fasth3_14b_2step_qad/comparison_five_prompts.json @@ -0,0 +1,32 @@ +[ + { + "id": "speech_exact_presenter", + "category": "exact_speech", + "prompt": "(S1) In a quiet recording studio, a presenter looks directly into the camera and says [English] Fast video and clear audio arrive together. The camera is locked off, the face is well lit, and no music plays.", + "expected_audio": "Exact intelligible English sentence with visible lip motion and no music." + }, + { + "id": "motorcycle_tracking", + "category": "large_motion", + "prompt": "A motorcycle accelerates along a wet coastal road while the camera tracks beside it at speed. Water sprays from the tires, scenery moves rapidly, and the engine pitch rises with acceleration in stereo.", + "expected_audio": "Engine pitch and tire spray synchronized to acceleration and motion." + }, + { + "id": "mechanical_press", + "category": "visible_mechanical_sound", + "prompt": "Inside a clean workshop, a metal stamping press descends once onto a small steel plate. The visible impact produces one sharp metallic clang followed by a short machine hiss. Static three-quarter camera view.", + "expected_audio": "One impact clang and a short hiss synchronized to the press." + }, + { + "id": "two_shot_transition", + "category": "two_shot_transition", + "prompt": "A two-shot sequence: first, a wide view of a steam train entering a rural station with rhythmic wheel sounds; then a clean cut to a close view of the conductor opening a carriage door as steam hisses. Audio perspective changes naturally after the cut.", + "expected_audio": "Wheel rhythm before the cut and close steam hiss after it, with no desynchronization." + }, + { + "id": "glass_water_closeup", + "category": "fine_motion_and_materials", + "prompt": "A continuous close-up in a bright kitchen: a person slowly pours clear water from a glass pitcher into an empty drinking glass on a wooden counter. Fingers grip the handle naturally. The stream splashes and bubbles, the water level rises steadily, and sunlight refracts through the glass. Locked camera, no cuts, no music.", + "expected_audio": "Natural continuous pouring and splashing synchronized with the visible water stream, fading when pouring stops." + } +] diff --git a/examples/training/fasth3_14b_2step_qad/quick_gate_prompts.json b/examples/training/fasth3_14b_2step_qad/quick_gate_prompts.json new file mode 100644 index 0000000000..601f9c2055 --- /dev/null +++ b/examples/training/fasth3_14b_2step_qad/quick_gate_prompts.json @@ -0,0 +1,74 @@ +[ + { + "id": "speech_exact_presenter", + "category": "exact_speech", + "prompt": "(S1) In a quiet recording studio, a presenter looks directly into the camera and says [English] Fast video and clear audio arrive together. The camera is locked off, the face is well lit, and no music plays.", + "expected_audio": "Exact intelligible English sentence with visible lip motion and no music." + }, + { + "id": "speech_exact_barista", + "category": "exact_speech", + "prompt": "(S1) A barista behind a cafe counter smiles and says [English] Your coffee is ready by the window. Cups clink softly in the background while the camera holds a close medium shot.", + "expected_audio": "Exact intelligible English sentence, subtle cup sounds, and synchronized lips." + }, + { + "id": "dog_bark_visible", + "category": "animal_sound", + "prompt": "A golden retriever stands beside a red garden gate, looks toward the camera, and gives two distinct barks. Its mouth and chest movement visibly match each bark; birds remain faint in the distance.", + "expected_audio": "Two dog barks synchronized to the visible dog." + }, + { + "id": "mechanical_press", + "category": "visible_mechanical_sound", + "prompt": "Inside a clean workshop, a metal stamping press descends once onto a small steel plate. The visible impact produces one sharp metallic clang followed by a short machine hiss. Static three-quarter camera view.", + "expected_audio": "One impact clang and a short hiss synchronized to the press." + }, + { + "id": "violin_duet", + "category": "music", + "prompt": "Two violinists perform a gentle chamber duet on a small wooden stage. Their bow strokes are clearly visible and the stereo violin music follows the motion, with quiet room ambience and no speech.", + "expected_audio": "Coherent stereo violin duet synchronized to bow motion." + }, + { + "id": "silent_snow", + "category": "near_silence", + "prompt": "A wide locked shot of fresh snow falling over an empty field at dawn. Nothing moves except soft snowflakes and distant tree branches. The scene is nearly silent, with only a very faint winter breeze and no music or speech.", + "expected_audio": "Near silence without hiss, line noise, speech, or music." + }, + { + "id": "closeup_woman", + "category": "human_closeup", + "prompt": "Close-up portrait of a woman in warm window light listening thoughtfully, blinking naturally, then taking a quiet breath. Fine skin and eye detail, shallow depth of field, soft indoor room tone, no speech.", + "expected_audio": "Subtle room tone and breath without synthetic speech." + }, + { + "id": "closeup_man_laugh", + "category": "human_closeup", + "prompt": "Close-up portrait of an older man outdoors who breaks into a brief natural laugh. His eyes, cheeks, mouth, and shoulders move consistently; the laugh is intelligible and synchronized, with light park ambience.", + "expected_audio": "Brief synchronized natural laugh with park ambience." + }, + { + "id": "motorcycle_tracking", + "category": "large_motion", + "prompt": "A motorcycle accelerates along a wet coastal road while the camera tracks beside it at speed. Water sprays from the tires, scenery moves rapidly, and the engine pitch rises with acceleration in stereo.", + "expected_audio": "Engine pitch and tire spray synchronized to acceleration and motion." + }, + { + "id": "basketball_pan", + "category": "large_motion", + "prompt": "An athlete sprints across an indoor basketball court, catches a fast pass, and completes a powerful dunk as the camera pans quickly. Sneakers squeak, the rim rattles at impact, and the crowd reacts once.", + "expected_audio": "Squeaks, one rim impact, and crowd reaction synchronized to the action." + }, + { + "id": "two_shot_transition", + "category": "two_shot_transition", + "prompt": "A two-shot sequence: first, a wide view of a steam train entering a rural station with rhythmic wheel sounds; then a clean cut to a close view of the conductor opening a carriage door as steam hisses. Audio perspective changes naturally after the cut.", + "expected_audio": "Wheel rhythm before the cut and close steam hiss after it, with no desynchronization." + }, + { + "id": "onscreen_text", + "category": "onscreen_text", + "prompt": "A clean product-demo shot of a small electronic sign on a desk. The display clearly reads FAST H3 in large white capital letters while a hand presses one button and a single soft confirmation beep sounds. Locked camera, neutral background.", + "expected_audio": "One confirmation beep synchronized to the button press; on-screen text should read FAST H3." + } +] diff --git a/examples/training/fasth3_14b_2step_qad/release_sentinel_24.json b/examples/training/fasth3_14b_2step_qad/release_sentinel_24.json new file mode 100644 index 0000000000..2c37b372e0 --- /dev/null +++ b/examples/training/fasth3_14b_2step_qad/release_sentinel_24.json @@ -0,0 +1,26 @@ +[ + {"id":"speech_en_exact","category":"speech","prompt":"A well-lit presenter faces the locked camera in a silent studio and says [English] Fast video and clear audio arrive together. No music or background noise.","expected_audio":"Exact English sentence, clean voice, synchronized lips."}, + {"id":"speech_es_exact","category":"speech","prompt":"A woman at a quiet kitchen table looks into the camera and says [Spanish] El vaso está lleno de agua fría. No music plays.","expected_audio":"Exact intelligible Spanish sentence with synchronized lips."}, + {"id":"speech_zh_exact","category":"speech","prompt":"A man stands in a quiet library and says [Chinese] 今天的天气非常晴朗。 Static close shot, no music.","expected_audio":"Exact intelligible Mandarin sentence with synchronized lips."}, + {"id":"speech_ja_exact","category":"speech","prompt":"A woman at a train platform faces the camera and says [Japanese] 次の電車は三時に来ます。 The distant station remains quiet, with no music.","expected_audio":"Exact intelligible Japanese sentence with synchronized lips."}, + {"id":"duet_singing","category":"music","prompt":"Two singers perform a gentle acoustic duet on a small stage, trading one line each and then harmonizing. One guitar accompanies them; the audience stays quiet.","expected_audio":"Two distinct singing voices, stable harmony, and clean acoustic guitar."}, + {"id":"solo_violin","category":"music","prompt":"Close view of a violinist performing a slow lyrical melody in a dry rehearsal room. Bow direction and note attacks remain visible. No other instruments.","expected_audio":"Natural solo violin whose attacks follow the visible bow changes."}, + {"id":"drum_pattern","category":"music","prompt":"A drummer plays four measured hits: kick, snare, kick, then cymbal, with a clear pause between each. Locked front camera, no backing track.","expected_audio":"Exactly four distinct percussion events aligned with the visible strikes."}, + {"id":"piano_scale","category":"music","prompt":"Overhead close-up of two hands playing a rising eight-note piano scale, then stopping with both hands lifted. No speech and no room chatter.","expected_audio":"Eight clean rising piano notes that stop when the hands lift."}, + {"id":"paper_foley","category":"foley","prompt":"Extreme close-up of hands slowly folding crisp paper twice and tearing it once along the crease. Quiet room, no speech, no music.","expected_audio":"Two soft folds followed by one crisp synchronized tear."}, + {"id":"vegetable_chop","category":"foley","prompt":"Close-up of a chef making six evenly spaced knife cuts through a carrot on a wooden board, then placing the knife down. No music.","expected_audio":"Six dry synchronized chopping impacts and one softer knife placement."}, + {"id":"gravel_steps","category":"foley","prompt":"Low tracking shot of boots taking five slow steps across loose gravel and stopping. Wind is very faint; no speech or music.","expected_audio":"Five granular footstep crunches aligned to heel contact, then silence."}, + {"id":"zipper_fabric","category":"foley","prompt":"Close-up of a person slowly zipping a canvas backpack, tightening one strap, and setting it on a table. Quiet indoor room.","expected_audio":"Continuous zipper texture, brief fabric pull, and one soft table thump."}, + {"id":"balloon_pop","category":"av_sync","prompt":"A red balloon floats motionless. A visible pin touches it at exactly the middle of the shot and the balloon bursts once. Static camera, no music.","expected_audio":"One sharp pop exactly at visible rupture, with silence before and after."}, + {"id":"basketball_bounces","category":"av_sync","prompt":"Side view of a basketball dropped onto a gym floor. It bounces exactly three times, each bounce lower than the last, then rolls away.","expected_audio":"Three decreasing synchronized bounces followed by a quiet rolling sound."}, + {"id":"door_latch","category":"av_sync","prompt":"A hand turns a brass handle, opens a wooden door, and closes it until the latch clicks. The camera stays on the handle; no speech or music.","expected_audio":"Handle turn, hinge movement, closing thump, and final click aligned to motion."}, + {"id":"firework_single","category":"av_sync","prompt":"Night skyline with one firework launching, bursting once into a blue circle, and fading. No crowd and no background music.","expected_audio":"Launch whistle followed by one delayed boom, with no extra explosions."}, + {"id":"quiet_portrait","category":"silence_noise","prompt":"A silent locked portrait of a sleeping cat in a sunlit room. Only subtle breathing and curtain movement; explicitly no speech, music, buzzing, or hiss.","expected_audio":"Near-silence without synthetic hiss or unexpected events."}, + {"id":"snow_field","category":"silence_noise","prompt":"Wide static view of fresh snow falling over an empty field at dawn. No people, vehicles, animals, speech, or music.","expected_audio":"Very quiet natural ambience without voices, tones, or crackle."}, + {"id":"rain_window","category":"general_audio","prompt":"Continuous close-up of rain striking a window while distant traffic lights blur outside. The rainfall remains steady and no one speaks.","expected_audio":"Stable natural rain texture with faint distant traffic and no artifacts."}, + {"id":"ocean_waves","category":"general_audio","prompt":"Wide sunset beach view as three waves reach the shore in succession and foam recedes. Slow tripod pan, no people and no music.","expected_audio":"Three broad wave surges synchronized to shore contact, then receding foam."}, + {"id":"motorcycle_tracking","category":"motion","prompt":"A motorcycle accelerates along a wet coastal road while the camera tracks beside it at speed. Water sprays from the tires and the engine pitch rises in stereo.","expected_audio":"Engine pitch and tire spray synchronized to acceleration and motion."}, + {"id":"dog_frisbee","category":"motion","prompt":"A dog sprints across grass, leaps to catch a flying frisbee, lands, and runs back toward the camera in one continuous tracking shot.","expected_audio":"Footfalls, one brief jump effort, landing impact, and natural outdoor ambience."}, + {"id":"train_two_shot","category":"multishot","prompt":"First, a wide view of a steam train entering a rural station with rhythmic wheel sounds; then a clean cut to the conductor opening a carriage door as steam hisses.","expected_audio":"Wheel rhythm before the cut and close steam hiss after it, with perspective change."}, + {"id":"cafe_three_shot","category":"multishot","prompt":"Three-shot sequence in one cafe: espresso pours into a cup, milk is steamed, then the finished cup is set beside a customer. Each cut changes audio perspective naturally.","expected_audio":"Pour, steam, and cup placement in the correct shots without carryover artifacts."} +] diff --git a/examples/training/fasth3_14b_2step_qad/rescue_gate_prompts.json b/examples/training/fasth3_14b_2step_qad/rescue_gate_prompts.json new file mode 100644 index 0000000000..b8fd589f1d --- /dev/null +++ b/examples/training/fasth3_14b_2step_qad/rescue_gate_prompts.json @@ -0,0 +1,26 @@ +[ + { + "id": "speech_exact_presenter", + "category": "exact_speech", + "prompt": "(S1) In a quiet recording studio, a presenter looks directly into the camera and says [English] Fast video and clear audio arrive together. The camera is locked off, the face is well lit, and no music plays.", + "expected_audio": "Exact intelligible English sentence with visible lip motion and no music." + }, + { + "id": "motorcycle_tracking", + "category": "large_motion", + "prompt": "A motorcycle accelerates along a wet coastal road while the camera tracks beside it at speed. Water sprays from the tires, scenery moves rapidly, and the engine pitch rises with acceleration in stereo.", + "expected_audio": "Engine pitch and tire spray synchronized to acceleration and motion." + }, + { + "id": "mechanical_press", + "category": "visible_mechanical_sound", + "prompt": "Inside a clean workshop, a metal stamping press descends once onto a small steel plate. The visible impact produces one sharp metallic clang followed by a short machine hiss. Static three-quarter camera view.", + "expected_audio": "One impact clang and a short hiss synchronized to the press." + }, + { + "id": "two_shot_transition", + "category": "two_shot_transition", + "prompt": "A two-shot sequence: first, a wide view of a steam train entering a rural station with rhythmic wheel sounds; then a clean cut to a close view of the conductor opening a carriage door as steam hisses. Audio perspective changes naturally after the cut.", + "expected_audio": "Wheel rhythm before the cut and close steam hiss after it, with no desynchronization." + } +] diff --git a/examples/training/fasth3_14b_2step_qad/showcase_prompt.json b/examples/training/fasth3_14b_2step_qad/showcase_prompt.json new file mode 100644 index 0000000000..d1c217cefc --- /dev/null +++ b/examples/training/fasth3_14b_2step_qad/showcase_prompt.json @@ -0,0 +1,10 @@ +[ + { + "id": "spellblade_stand_back", + "category": "showcase_motion_speech_sync", + "expected_audio": "The English line 'Stand back!' is intelligible and synchronized, with storm wind, stone impacts, blade motion, and magical crackle matching visible events.", + "prompt": "integrated_multimodal_description: [Shot 1] A 16:9 single-take motion-graphics design in a vivid 2D anime action-fantasy style, with crisp cel-shaded planes, layered vector-like shapes, luminous energy streaks, and no readable text, opens on one lone sky-blue-haired spellblade standing on a shattered moonstone bridge above a violet storm. The spellblade pivots, raises a glowing crescent blade, and releases a high-speed spiral of gold-and-cyan energy that tears through incoming shadow shards; (S1) shouts [English]Stand back! while bracing against the magical recoil, with a sharp gasp and strained breath audible. After the action begins, the camera makes one rapid, smooth forward push-in toward the spellblade's determined face as the energy spiral fills the frame, maintaining the same continuous take through 00:05.000. Keep the spellblade as the only principal character, with no additional characters, logos, captions, subtitles, or readable signage.\n\noverall_soundscape: Violet storm wind roars around the bridge while stone fragments clatter, magical energy crackles, and the blade slices through the air. The spellblade's sharp gasp, strained breath, and shouted line remain clearly synchronized with the action.\n\nnon_diegetic_music: N/A", + "source_id": "t2va-2026082050-000102", + "source_seed": 2026082050 + } +] diff --git a/fastvideo/attention/utils/flash_attn_cute.py b/fastvideo/attention/utils/flash_attn_cute.py index ca38539924..bbc840aede 100644 --- a/fastvideo/attention/utils/flash_attn_cute.py +++ b/fastvideo/attention/utils/flash_attn_cute.py @@ -13,36 +13,23 @@ try: from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd except ImportError: - # flash_attn.cute (FA4) is simply not installed -- expected on builds - # without it; callers handle the ImportError (the FASTVIDEO_FA4 gate in - # flash_attn.py raises, the FP4 probe treats FA4 as unavailable). raise except Exception as e: - # flash_attn.cute IS installed but failed to import -- almost always an - # nvidia-cutlass-dsl (CuTe DSL) version skew, e.g. "module - # 'cutlass.cute.core' has no attribute 'ThrMma'" (an AttributeError, not - # ImportError). This is fixable by pinning a compatible - # nvidia-cutlass-dsl, so warn loudly, then re-raise as ImportError so - # callers can handle it uniformly. logger.warning( "flash_attn.cute (FA4) is installed but failed to import (%r). " "This is usually an nvidia-cutlass-dsl version mismatch -- pin a " "compatible nvidia-cutlass-dsl to restore FA4.", e) raise ImportError(f"flash_attn.cute (FA4) import failed: {e!r}") from e else: - # This error will be caught in flash_attn.py or flash_attn_no_pad.py raise ImportError("flash_attn.cute is only available on CUDA devices; this error must be handled internally") try: - # FA2 serves the calls FA4 cute cannot on pre-sm90 GPUs (backward, GQA). - # Optional so FA4-only installs can still import this module. from flash_attn import flash_attn_func as _flash_attn_2_func from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func except ImportError: _flash_attn_2_func = None _flash_attn_2_varlen_func = None -# Dynamo unwraps functools caches, so keep this cache inside the opaque helper. _SM90_OR_NEWER_BY_DEVICE: dict[int, bool] = {} @@ -63,11 +50,6 @@ def _use_fa2(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> bool: assert device_id is not None if _sm90_or_newer(device_id): return False - # Pre-sm90 FA4 cute limitations, both served by FA2 (deterministic - # capability gate, not a runtime fallback): - # * the backward asserts sm90+ (L40S/sm_89 dies on its arch check); - # * GQA fails CuTeDSL JIT in pack_gqa ("ValueError: Operation creation - # failed", observed on sm_89 with HunyuanGameCraft/LTX2). if q.shape[-2] != k.shape[-2]: return True return torch.is_grad_enabled() and any(t.requires_grad for t in (q, k, v)) @@ -94,9 +76,6 @@ def _flash_attn_cute_forward( causal: bool, deterministic: bool, ) -> tuple[torch.Tensor, torch.Tensor]: - # _flash_attn_fwd returns (out, lse) on its empty-sequence early path but - # (out, lse, p, row_max) on the main path at the pinned FA4 cute ref; take - # the first two so both arities work. out, lse = _flash_attn_fwd( q, k, @@ -108,6 +87,7 @@ def _flash_attn_cute_forward( softcap=0.0, num_splits=1, pack_gqa=None, + return_lse=True, )[:2] return out, lse @@ -132,11 +112,60 @@ 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 +@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, + 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=softmax_scale, + causal=causal, + softcap=0.0, + window_size_left=None, + window_size_right=None, + 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 torch.empty_like(q), torch.empty_like(k), torch.empty_like(v) + + def _flash_attn_cute_backward( ctx: torch.autograd.function.FunctionCtx, grad_out: torch.Tensor, @@ -144,19 +173,16 @@ def _flash_attn_cute_backward( ): del grad_lse q, k, v, out, lse = ctx.saved_tensors - dq, dk, dv = _flash_attn_bwd( + dq, dk, dv = torch.ops.fastvideo._flash_attn_cute_backward( q, k, v, out, grad_out, lse, - softmax_scale=ctx.softmax_scale, - causal=ctx.causal, - softcap=0.0, - window_size_left=None, - window_size_right=None, - deterministic=ctx.deterministic, + ctx.softmax_scale, + ctx.causal, + ctx.deterministic, ) return dq, dk, dv, None, None, None @@ -200,6 +226,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 +268,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,6 +276,67 @@ def _flash_attn_cute_varlen_setup_context(ctx: torch.autograd.function.FunctionC ctx.deterministic = deterministic +@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, + 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=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=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 torch.empty_like(q), torch.empty_like(k), torch.empty_like(v) + + def _flash_attn_cute_varlen_backward( ctx: torch.autograd.function.FunctionCtx, grad_out: torch.Tensor, @@ -255,23 +344,20 @@ def _flash_attn_cute_varlen_backward( ): del grad_lse q, k, v, out, lse, cu_seqlens_q, cu_seqlens_k = ctx.saved_tensors - dq, dk, dv = _flash_attn_bwd( + dq, dk, dv = torch.ops.fastvideo._flash_attn_cute_varlen_backward( q, k, v, out, grad_out, lse, - softmax_scale=ctx.softmax_scale, - causal=ctx.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, + 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 @@ -308,13 +394,6 @@ def flash_attn_func( return out -# --------------------------------------------------------------------------- -# FP4 (NVFP4 block-scaled) variant -# --------------------------------------------------------------------------- -# The FP4 path needs the mSFQ/mSFK scale-factor tensors that the regular -# wrapper does not expose. We register a separate custom op so that -# torch.compile can treat the kernel as an opaque boundary (the underlying -# CuTeDSL kernel uses cuda.CUstream which dynamo cannot trace). @torch.library.custom_op( @@ -359,8 +438,6 @@ def _flash_attn_cute_fp4_forward_fake( causal: bool, ) -> torch.Tensor: del k, sfq, sfk, softmax_scale, causal - # q is FP4 packed: shape (batch, seqlen, nheads, headdim/2). Output is in - # V's dtype with full headdim. batch, seqlen_q, nheads = q.shape[:3] return v.new_empty(batch, seqlen_q, nheads, v.shape[-1]) diff --git a/fastvideo/configs/models/dits/minimax_h3.py b/fastvideo/configs/models/dits/minimax_h3.py index 464f75a699..b1367769e9 100644 --- a/fastvideo/configs/models/dits/minimax_h3.py +++ b/fastvideo/configs/models/dits/minimax_h3.py @@ -25,9 +25,6 @@ class MiniMaxH3ArchConfig(DiTArchConfig): _supported_attention_backends: tuple[AttentionBackendEnum, ...] = ( AttentionBackendEnum.TORCH_SDPA, AttentionBackendEnum.FLASH_ATTN, - # FP4-quantized QK attention (fa4_fp4 on sm_100/sm_103, cutlass on - # sm_12x). Enabled for speed experiments; output quality against the - # SSIM references is not yet validated. AttentionBackendEnum.ATTN_QAT_INFER, AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3, ) @@ -82,6 +79,4 @@ class MiniMaxH3Config(DiTConfig): arch_config: MiniMaxH3ArchConfig = field(default_factory=MiniMaxH3ArchConfig) prefix: str = "minimax_h3" - # FastVideo's Fully Sharded Data Parallel (FSDP) loading path requires one - # parameter dtype, while H3 inference keeps boundary projections in FP32. uniform_parameter_dtype: bool = False diff --git a/fastvideo/dataset/parquet_dataset_map_style.py b/fastvideo/dataset/parquet_dataset_map_style.py index f99937f75e..0e13c81340 100644 --- a/fastvideo/dataset/parquet_dataset_map_style.py +++ b/fastvideo/dataset/parquet_dataset_map_style.py @@ -2,19 +2,20 @@ import os import pickle import random +from collections import defaultdict from collections.abc import Sequence +from pathlib import Path from typing import Any import pyarrow as pa import pyarrow.parquet as pq -# Torch in general import torch import tqdm -# Dataset from torch.utils.data import Dataset, Sampler from torchdata.stateful_dataloader import StatefulDataLoader from fastvideo.platforms import current_platform +from fastvideo.dataset.shape_bucket import parse_video_shape_bucket_id from fastvideo.dataset.utils import collate_rows_from_parquet_schema from fastvideo.distributed import (get_sp_world_size, get_world_group, get_world_rank, get_world_size) from fastvideo.logger import init_logger @@ -37,6 +38,7 @@ def __init__( drop_last: bool = True, drop_first_row: bool = False, seed: int = 0, + sample_bucket_ids: Sequence[str] | None = None, ): self.batch_size = batch_size self.dataset_size = dataset_size @@ -46,38 +48,105 @@ def __init__( self.global_rank = global_rank self.sp_world_size = sp_world_size - # ── epoch-level RNG ──────────────────────────────────────────────── + if batch_size <= 0 or num_sp_groups <= 0 or sp_world_size <= 0: + raise ValueError("batch_size, num_sp_groups, and sp_world_size must be positive") + rng = torch.Generator().manual_seed(self.seed) - # Create a random permutation of all indices - global_indices = torch.randperm(self.dataset_size, generator=rng) - - if drop_first_row: - # drop 0 in global_indices - global_indices = global_indices[global_indices != 0] - self.dataset_size = self.dataset_size - 1 - - if self.drop_last: - # For drop_last=True, we: - # 1. Ensure total samples is divisible by (batch_size * num_sp_groups) - # 2. This guarantees each SP group gets same number of complete batches - # 3. Prevents uneven batch sizes across SP groups at end of epoch - num_batches = self.dataset_size // self.batch_size - num_global_batches = num_batches // self.num_sp_groups - global_indices = global_indices[:num_global_batches * self.num_sp_groups * self.batch_size] - else: - if self.dataset_size % (self.num_sp_groups * self.batch_size) != 0: - # add more indices to make it divisible by (batch_size * num_sp_groups) + if sample_bucket_ids is None: + global_indices = torch.randperm(self.dataset_size, generator=rng) + if drop_first_row: + global_indices = global_indices[global_indices != 0] + self.dataset_size -= 1 + + if self.drop_last: + num_batches = self.dataset_size // self.batch_size + num_global_batches = num_batches // self.num_sp_groups + global_indices = global_indices[:num_global_batches * self.num_sp_groups * self.batch_size] + elif self.dataset_size % (self.num_sp_groups * self.batch_size) != 0: padding_size = self.num_sp_groups * self.batch_size - (self.dataset_size % (self.num_sp_groups * self.batch_size)) logger.info("Padding the dataset from %d to %d", self.dataset_size, self.dataset_size + padding_size) global_indices = torch.cat([global_indices, global_indices[:padding_size]]) + self.bucket_schedule: tuple[str, ...] | None = None + self.bucket_padding: dict[str, int] | None = None + self.num_padded_samples = 0 + else: + if len(sample_bucket_ids) != self.dataset_size: + raise ValueError("sample_bucket_ids must contain one identifier per dataset row, got " + f"{len(sample_bucket_ids)} for dataset_size={self.dataset_size}") + global_indices, self.bucket_schedule = self._bucketed_global_schedule( + sample_bucket_ids, + rng=rng, + drop_first_row=drop_first_row, + ) + if drop_first_row: + self.dataset_size -= 1 - # shard the indices to each sp group ith_sp_group = self.global_rank // self.sp_world_size sp_group_local_indices = global_indices[ith_sp_group::self.num_sp_groups] self.sp_group_local_indices = sp_group_local_indices logger.info("Dataset size for each sp group: %d", len(sp_group_local_indices)) + def _bucketed_global_schedule( + self, + sample_bucket_ids: Sequence[str], + *, + rng: torch.Generator, + drop_first_row: bool, + ) -> tuple[torch.Tensor, tuple[str, ...]]: + """Build same-shape rounds shared by every data-parallel group. + + One round contains ``num_sp_groups * batch_size`` rows from exactly + one bucket. Strided DP sharding below gives every group a distinct + local batch while preserving the same bucket at that microstep. SP + ranks map to the same group and therefore receive identical indices. + """ + by_bucket: dict[str, list[int]] = defaultdict(list) + for index, bucket_id in enumerate(sample_bucket_ids): + if drop_first_row and index == 0: + continue + parse_video_shape_bucket_id(bucket_id) + by_bucket[bucket_id].append(index) + + samples_per_round = self.num_sp_groups * self.batch_size + rounds: list[torch.Tensor] = [] + round_bucket_ids: list[str] = [] + bucket_padding: dict[str, int] = {} + padded = 0 + for bucket_id in sorted(by_bucket): + bucket_indices = torch.tensor(by_bucket[bucket_id], dtype=torch.long) + bucket_indices = bucket_indices[torch.randperm(len(bucket_indices), generator=rng)] + remainder = len(bucket_indices) % samples_per_round + if remainder: + padding_size = samples_per_round - remainder + repeats = (padding_size + len(bucket_indices) - 1) // len(bucket_indices) + padding = bucket_indices.repeat(repeats)[:padding_size] + bucket_indices = torch.cat((bucket_indices, padding)) + padded += padding_size + bucket_padding[bucket_id] = padding_size + logger.info( + "Exact-shape bucket %s has %d row(s); repeated %d row(s) to fill global microbatches of %d", + bucket_id, + len(by_bucket[bucket_id]), + padding_size, + samples_per_round, + ) + bucket_rounds = list(bucket_indices.reshape(-1, samples_per_round).unbind(0)) + rounds.extend(bucket_rounds) + round_bucket_ids.extend([bucket_id] * len(bucket_rounds)) + + if not rounds: + raise ValueError("Exact-shape bucketing requires at least one dataset row") + round_order = torch.randperm(len(rounds), generator=rng).tolist() + self.bucket_padding = bucket_padding + self.num_padded_samples = padded + if padded: + logger.info("Exact-shape bucketing repeated %d row(s) to fill bucket-local global microbatches", padded) + return ( + torch.cat([rounds[index] for index in round_order]), + tuple(round_bucket_ids[index] for index in round_order), + ) + def __iter__(self): indices = self.sp_group_local_indices for i in range(0, len(indices), self.batch_size): @@ -88,6 +157,36 @@ def __len__(self): return len(self.sp_group_local_indices) // self.batch_size +def _shape_bucket_id_from_parquet_path(file_path: str) -> str: + """Return and validate the sole ``bucket=...`` ancestor of a parquet.""" + matches = [part for part in Path(file_path).parts if part.startswith("bucket=")] + if len(matches) != 1: + raise ValueError( + "Native-shape parquet paths must have exactly one ancestor named " + "'bucket=x-f', got " + f"{file_path!r} with bucket ancestors {matches}" + ) + bucket_id = matches[0] + parse_video_shape_bucket_id(bucket_id) + return bucket_id + + +def shape_bucket_ids_from_parquet_files( + parquet_files: Sequence[str], + lengths: Sequence[int], +) -> list[str]: + """Expand canonical path bucket IDs to one identifier per dataset row.""" + if len(parquet_files) != len(lengths): + raise ValueError("parquet_files and lengths must have matching lengths") + sample_bucket_ids: list[str] = [] + for file_path, length in zip(parquet_files, lengths, strict=True): + if int(length) < 0: + raise ValueError(f"Parquet row counts must be non-negative, got {length}") + bucket_id = _shape_bucket_id_from_parquet_path(str(file_path)) + sample_bucket_ids.extend([bucket_id] * int(length)) + return sample_bucket_ids + + def _parse_data_path_specs(path: str | Sequence[str] | dict[str, int]) -> list[tuple[str, int]]: """Parse one or more dataset roots with old-framework repeat counts.""" if isinstance(path, dict): @@ -132,17 +231,14 @@ def get_parquet_files_and_length(path: str | Sequence[str] | dict[str, int]): return file_names_sorted, lengths_sorted dataset_root = os.path.realpath(os.path.expanduser(specs[0][0])) - # Check if cached info exists cache_dir = os.path.join(dataset_root, "map_style_cache") cache_file = os.path.join(cache_dir, "file_info.pkl") - # Only rank 0 checks for cache and scans files if needed if get_world_rank() == 0: cache_loaded = False file_names_sorted = None lengths_sorted = None - # First try to load existing cache if os.path.exists(cache_file): logger.info("Loading cached file info from %s", cache_file) try: @@ -180,7 +276,6 @@ def get_parquet_files_and_length(path: str | Sequence[str] | dict[str, int]): logger.info("Falling back to scanning files") cache_loaded = False - # If cache not loaded (either doesn't exist or failed to load), scan files if not cache_loaded: logger.info("Scanning parquet files to get lengths") lengths = [] @@ -198,20 +293,16 @@ def get_parquet_files_and_length(path: str | Sequence[str] | dict[str, int]): for file_path in tqdm.tqdm(file_names, desc="Reading parquet files to get lengths"): num_rows = pq.ParquetFile(file_path).metadata.num_rows lengths.append(num_rows) - # sort according to file name to ensure all rank has the same order file_names_sorted, lengths_sorted = zip(*sorted(zip(file_names, lengths, strict=True), key=lambda x: x[0]), strict=True) - # Save the cache os.makedirs(cache_dir, exist_ok=True) with open(cache_file, "wb") as f: pickle.dump((file_names_sorted, lengths_sorted), f) logger.info("Saved file info to %s", cache_file) - # Wait for rank 0 to finish creating/loading cache world_group = get_world_group() world_group.barrier() - # Now all ranks load the cache (it should exist and be valid now) logger.info("Loading cached file info from %s after barrier", cache_file) with open(cache_file, "rb") as f: file_names_sorted, lengths_sorted = pickle.load(f) @@ -226,7 +317,12 @@ def get_parquet_files_and_length(path: str | Sequence[str] | dict[str, int]): return file_names_sorted, lengths_sorted -def read_row_from_parquet_file(parquet_files: list[str], global_row_idx: int, lengths: list[int]) -> dict[str, Any]: +def read_row_from_parquet_file( + parquet_files: list[str], + global_row_idx: int, + lengths: list[int], + columns: Sequence[str] | None = None, +) -> dict[str, Any]: ''' Read a row from a parquet file. Args: @@ -235,7 +331,6 @@ def read_row_from_parquet_file(parquet_files: list[str], global_row_idx: int, le lengths: List[int] Returns: ''' - # find the parquet file and local row index cumulative = 0 file_index = 0 local_row_idx = 0 @@ -246,13 +341,10 @@ def read_row_from_parquet_file(parquet_files: list[str], global_row_idx: int, le break cumulative += lengths[file_index] else: - # If we reach here, global_row_idx is out of bounds raise IndexError(f"global_row_idx {global_row_idx} is out of bounds for dataset") parquet_file = pq.ParquetFile(parquet_files[file_index]) - # Calculate the row group to read into memory and the local idx - # This way we can avoid reading in the entire parquet file cumulative = 0 row_group_index = 0 local_index = 0 @@ -265,19 +357,15 @@ def read_row_from_parquet_file(parquet_files: list[str], global_row_idx: int, le break cumulative += num_rows else: - # If we reach here, local_row_idx is out of bounds for this parquet file raise IndexError(f"local_row_idx {local_row_idx} is out of bounds for parquet file {parquet_files[file_index]}") - row_group = parquet_file.read_row_group(row_group_index).to_pydict() + row_group = parquet_file.read_row_group(row_group_index, columns=columns).to_pydict() row_dict = {k: v[local_index] for k, v in row_group.items()} del row_group return row_dict -# ──────────────────────────────────────────────────────────────────────────── -# 2. Dataset with batched __getitems__ -# ──────────────────────────────────────────────────────────────────────────── class LatentsParquetMapStyleDataset(Dataset): """ Return latents[B,C,T,H,W] and embeddings[B,L,D] in pinned CPU memory. @@ -295,18 +383,21 @@ def __init__( drop_last: bool = True, drop_first_row: bool = False, text_padding_length: int = 512, + native_shape_bucketing: bool = False, ): super().__init__() self.path = path self.cfg_rate = cfg_rate self.parquet_schema = parquet_schema self.seed = seed - # Create a seeded random generator for deterministic CFG self.rng = random.Random(seed) logger.info("Initializing LatentsParquetMapStyleDataset with path: %s", path) self.parquet_files, self.lengths = get_parquet_files_and_length(path) self.batch = batch_size self.text_padding_length = text_padding_length + self.sample_bucket_ids = ( + shape_bucket_ids_from_parquet_files(self.parquet_files, self.lengths) if native_shape_bucketing else None + ) self.sampler = DP_SP_BatchSampler( batch_size=batch_size, dataset_size=sum(self.lengths), @@ -316,6 +407,7 @@ def __init__( drop_last=drop_last, drop_first_row=drop_first_row, seed=seed, + sample_bucket_ids=self.sample_bucket_ids, ) logger.info("Dataset initialized with %d parquet files and %d rows", len(self.parquet_files), sum(self.lengths)) @@ -326,11 +418,14 @@ def get_validation_negative_prompt(self) -> tuple[torch.Tensor, torch.Tensor, st Returns the processed negative prompt data (latents, embeddings, masks, info). """ - # Read first row from first parquet file file_path = self.parquet_files[0] row_idx = 0 - # Read the negative prompt data - row_dict = read_row_from_parquet_file([file_path], row_idx, [self.lengths[0]]) + row_dict = read_row_from_parquet_file( + [file_path], + row_idx, + [self.lengths[0]], + columns=self.parquet_schema.names, + ) batch = collate_rows_from_parquet_schema([row_dict], self.parquet_schema, @@ -347,15 +442,19 @@ def get_validation_negative_prompt(self) -> tuple[torch.Tensor, torch.Tensor, st return negative_prompt_embedding, negative_prompt_attention_mask, negative_prompt - # PyTorch calls this ONLY because the batch_sampler yields a list def __getitems__(self, indices: list[int]) -> dict[str, Any]: """ Batch fetch using read_row_from_parquet_file for each index. """ - rows = [read_row_from_parquet_file(self.parquet_files, idx, self.lengths) for idx in indices] + rows = [ + read_row_from_parquet_file( + self.parquet_files, + idx, + self.lengths, + columns=self.parquet_schema.names, + ) for idx in indices + ] - # Inject sample indices for deterministic CFG dropout - # that is reproducible across checkpoint resume. for row, idx in zip(rows, indices): row["_sample_index"] = idx @@ -364,15 +463,17 @@ def __getitems__(self, indices: list[int]) -> dict[str, Any]: self.text_padding_length, cfg_rate=self.cfg_rate, seed=self.seed) + if self.sample_bucket_ids is not None: + bucket_ids = {self.sample_bucket_ids[index] for index in indices} + if len(bucket_ids) != 1: + raise RuntimeError(f"Exact-shape sampler emitted a mixed bucket batch: {sorted(bucket_ids)}") + batch["_shape_bucket_id"] = bucket_ids.pop() return batch def __len__(self): return sum(self.lengths) -# ──────────────────────────────────────────────────────────────────────────── -# 3. Loader helper – everything else stays just like your original trainer -# ──────────────────────────────────────────────────────────────────────────── def passthrough(batch): return batch @@ -385,7 +486,9 @@ def build_parquet_map_style_dataloader(path, drop_last=True, drop_first_row=False, text_padding_length=512, - seed=42) -> tuple[LatentsParquetMapStyleDataset, StatefulDataLoader]: + seed=42, + native_shape_bucketing=False) -> tuple[LatentsParquetMapStyleDataset, + StatefulDataLoader]: dataset = LatentsParquetMapStyleDataset(path, batch_size, cfg_rate=cfg_rate, @@ -393,7 +496,8 @@ def build_parquet_map_style_dataloader(path, drop_first_row=drop_first_row, text_padding_length=text_padding_length, parquet_schema=parquet_schema, - seed=seed) + seed=seed, + native_shape_bucketing=native_shape_bucketing) loader = StatefulDataLoader( dataset, diff --git a/fastvideo/dataset/shape_bucket.py b/fastvideo/dataset/shape_bucket.py new file mode 100644 index 0000000000..35dea8d239 --- /dev/null +++ b/fastvideo/dataset/shape_bucket.py @@ -0,0 +1,49 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Portable exact-shape bucket identifiers for preprocessed video data.""" + +from __future__ import annotations + +import re +from dataclasses import dataclass + +_VIDEO_SHAPE_BUCKET_PATTERN = re.compile( + r"^bucket=(?P[1-9][0-9]*)x(?P[1-9][0-9]*)-(?P[1-9][0-9]*)f$" +) + + +@dataclass(frozen=True, slots=True) +class ExactVideoShapeBucket: + """Pixel geometry and frame count encoded by one bucket directory.""" + + width: int + height: int + num_frames: int + + @property + def bucket_id(self) -> str: + return f"bucket={self.width}x{self.height}-{self.num_frames}f" + + +def parse_video_shape_bucket_id(bucket_id: str) -> ExactVideoShapeBucket: + """Parse ``bucket=x-f`` exactly. + + Width is deliberately first so the identifier matches common media + geometry notation. The strict spelling makes independently generated data + roots portable and prevents ranks from silently assigning the same shape + two different names. + """ + match = _VIDEO_SHAPE_BUCKET_PATTERN.fullmatch(str(bucket_id)) + if match is None: + raise ValueError( + "An exact video-shape bucket must be named " + "'bucket=x-f' with positive decimal " + f"integers (for example 'bucket=1344x768-124f'), got {bucket_id!r}" + ) + return ExactVideoShapeBucket( + width=int(match.group("width")), + height=int(match.group("height")), + num_frames=int(match.group("num_frames")), + ) + + +__all__ = ["ExactVideoShapeBucket", "parse_video_shape_bucket_id"] diff --git a/fastvideo/dataset/validation_dataset.py b/fastvideo/dataset/validation_dataset.py index 755df34f65..445263e791 100644 --- a/fastvideo/dataset/validation_dataset.py +++ b/fastvideo/dataset/validation_dataset.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -# adapted from: https://github.com/a-r-r-o-w/finetrainers/blob/main/finetrainers/data/dataset.py +import json import os import pathlib @@ -20,7 +20,6 @@ def __init__(self, filename: str): super().__init__() self.filename = pathlib.Path(filename) - # get directory of filename self.dir = os.path.abspath(self.filename.parent) if not self.filename.exists(): @@ -29,7 +28,11 @@ def __init__(self, filename: str): if self.filename.suffix == ".csv": data = datasets.load_dataset("csv", data_files=self.filename.as_posix(), split="train") elif self.filename.suffix == ".json": - data = datasets.load_dataset("json", data_files=self.filename.as_posix(), split="train", field="data") + document = json.loads(self.filename.read_text(encoding="utf-8")) + rows = document.get("data") if isinstance(document, dict) else document + if not isinstance(rows, list): + raise ValueError("Validation JSON must be a row array or an object containing a 'data' row array") + data = datasets.Dataset.from_list(rows) elif self.filename.suffix == ".parquet": data = datasets.load_dataset("parquet", data_files=self.filename.as_posix(), split="train") elif self.filename.suffix == ".arrow": @@ -40,22 +43,18 @@ def __init__(self, filename: str): f"Unsupported file format {self.filename.suffix} for validation dataset. Supported formats are: {_SUPPORTED_FILE_FORMATS}" ) - # Get distributed training info self.global_rank = get_world_rank() self.world_size = get_world_size() self.sp_world_size = get_sp_world_size() self.num_sp_groups = self.world_size // self.sp_world_size - # Convert to list to get total samples self.all_samples = list(data) self.original_total_samples = len(self.all_samples) - # Extend samples to be a multiple of DP degree (num_sp_groups) remainder = self.original_total_samples % self.num_sp_groups if remainder != 0: samples_to_add = self.num_sp_groups - remainder - # Duplicate samples cyclically to reach the target additional_samples = [] for i in range(samples_to_add): additional_samples.append(self.all_samples[i % self.original_total_samples]) @@ -64,17 +63,13 @@ def __init__(self, filename: str): self.total_samples = len(self.all_samples) - # Calculate which SP group this rank belongs to self.sp_group_id = self.global_rank // self.sp_world_size - # Now all SP groups will have equal number of samples self.samples_per_sp_group = self.total_samples // self.num_sp_groups - # Calculate start and end indices for this SP group self.start_idx = self.sp_group_id * self.samples_per_sp_group self.end_idx = self.start_idx + self.samples_per_sp_group - # Get samples for this SP group self.sp_group_samples = self.all_samples[self.start_idx:self.end_idx] logger.info( @@ -98,12 +93,8 @@ def __len__(self): def __iter__(self): for sample in self.sp_group_samples: - # For consistency reasons, we mandate that "caption" is always present in the validation dataset. - # However, since the model specifications use "prompt", we create an alias here. sample["prompt"] = sample["caption"] - # Load image or video if the path is provided - # TODO(aryan): need to handle custom columns here for control conditions sample["image"] = None sample["video"] = None diff --git a/fastvideo/fastvideo_args.py b/fastvideo/fastvideo_args.py index a648a02538..317fb4a12d 100644 --- a/fastvideo/fastvideo_args.py +++ b/fastvideo/fastvideo_args.py @@ -1,5 +1,4 @@ # SPDX-License-Identifier: Apache-2.0 -# Inspired by SGLang: https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/server_args.py """The arguments of FastVideo Inference.""" import argparse import dataclasses @@ -26,9 +25,6 @@ logger = init_logger(__name__) -# Offload flags that trade device memory for host memory. All of them are a loss -# on a device where the two are the same physical pool. Keeping the policy -# centralized lets every loader and stage share one worker-local decision. UNIFIED_MEMORY_OFFLOAD_FLAGS = ( "dit_layerwise_offload", "dit_cpu_offload", @@ -91,32 +87,24 @@ def choices(cls) -> list[str]: return [workload.value for workload in cls] -# args for fastvideo framework @dataclasses.dataclass class FastVideoArgs: - # Model and path configuration (for convenience) model_path: str - # Running mode mode: ExecutionMode = ExecutionMode.INFERENCE - # Workload type workload_type: WorkloadType = WorkloadType.T2V - # Distributed executor backend distributed_executor_backend: str = "mp" - # a few attributes for ray related ray_placement_group: PlacementGroup | None = None ray_runtime_env: RuntimeEnv | None = None inference_mode: bool = True # if False == training mode - # HuggingFace specific parameters trust_remote_code: bool = False revision: str | None = None - # Parallelism num_gpus: int = 1 tp_size: int = -1 sp_size: int = -1 @@ -127,30 +115,15 @@ class FastVideoArgs: pipeline_config: PipelineConfig = field(default_factory=PipelineConfig) preprocess_config: PreprocessConfig | None = None - # LoRA parameters - # (Wenxuan) prefer to keep it here instead of in pipeline config to not make it complicated. lora_path: str | None = None lora_nickname: str = "default" # for swapping adapters in the pipeline lora_strength: float = 1.0 - # can restrict layers to adapt, e.g. ["q_proj"] - # Will adapt only q, k, v, o by default. lora_target_modules: list[str] | None = None output_type: str = "pil" - # The attention backend requested for this run. Applied per component at - # load time (each component resolves its own decision, recorded on its - # config); a role-level request (the train stack's per-role - # attention_backend) overrides it. - # - # This field is the parse-once adapter for FASTVIDEO_ATTENTION_BACKEND: - # when left unset it takes the env var's value in __post_init__, so the - # environment is an *input* read once here rather than something the loader - # consults later. None means no request: per-layer defaults, then platform - # auto-selection. attention_backend: str | None = None - # CPU offload parameters dit_cpu_offload: bool = True use_fsdp_inference: bool = False dit_layerwise_offload: bool = True @@ -159,96 +132,44 @@ class FastVideoArgs: vae_cpu_offload: bool = True pin_cpu_memory: bool = True - # MiniMax-H3 inference load order. ``None`` (auto) defers DiT/VAE load until - # after the Qwen3-VL encoder is released, but only on unified-memory - # devices (GB10 / Spark). Discrete GPUs keep the encoder resident so a - # later ``generate()`` on the same worker can re-encode. Explicit True / - # False overrides the probe. Training never defers. h3_sequential_load: bool | None = None - # MiniMax-H3 video reconstruction. ``h3-vae`` is the full ViT decoder. - # ``taeh3`` is Ollin Boer Bohan's tiny preview decoder; it changes quality - # and is opt-in. T2VA with TAEH3 does not need the video VAE weights. video_decode_backend: str = "h3-vae" taeh3_checkpoint: str | None = None taeh3_chunk_size: int = 5 - # Load each heavy component on first use and free it once the last stage - # that holds it has run, instead of keeping every component resident from - # load time to shutdown. Peak memory becomes the largest overlapping set - # rather than the sum of all components. ``None`` (auto) turns this on for - # unified-memory devices (GB10 / Spark) after the worker binds its device, - # and leaves it off on discrete GPUs. Explicit True / False overrides the - # probe. A released component is re-read from disk on the next generation, - # so this trades per-request latency for headroom. Inference only; training - # keeps every component resident. lazy_module_load: bool | None = None - # Sequence-parallel MiniMax-H3 VAE (opt-in, default off). With SP > 1 the - # video VAE's temporal chunks (decode) and clips (reference encode) are - # round-robined across the sequence-parallel ranks and reassembled - # bit-exactly on the group's first rank instead of running serially on - # one rank while the others idle. ``__post_init__`` folds the - # FASTVIDEO_VAE_PARALLEL_DECODE / FASTVIDEO_VAE_PARALLEL_ENCODE env vars - # into these fields (parse-once, like attention_backend), and - # FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY overrides the chunk-transport - # collective ("gather" or "all_gather"). vae_parallel_decode: bool = False vae_parallel_encode: bool = False vae_parallel_decode_strategy: str | None = None - # Compilation - # ``enable_torch_compile`` covers the DiT path (transformer, - # transformer_2, and the LTX-2 stage-2 transformer_refine). - # Per-component flags below let callers compile additional submodules - # independently; ``False`` leaves the component eager. enable_torch_compile: bool = False + regional_compile: bool = False enable_torch_compile_text_encoder: bool = False enable_torch_compile_vae: bool = False enable_torch_compile_audio_vae: bool = False - # ``torch_compile_kwargs`` is the master kwargs dict (applied to every - # compiled submodule unless a per-component dict below is non-empty, - # in which case the per-component dict overrides entirely — matching - # the FastVideo-internal precedent). torch_compile_kwargs: dict[str, Any] = field(default_factory=dict) torch_compile_kwargs_dit: dict[str, Any] = field(default_factory=dict) torch_compile_kwargs_text_encoder: dict[str, Any] = field(default_factory=dict) torch_compile_kwargs_vae: dict[str, Any] = field(default_factory=dict) torch_compile_kwargs_audio_vae: dict[str, Any] = field(default_factory=dict) - # Regional (per-transformer-block) fullgraph torch.compile of the DiT at - # inference — the inference-side counterpart of the training regional - # compile ported from hao-ai-lab/FastVideo#1718. Applied by the loader - # right after the transformer loads, with fullgraph=True and inductor - # options {emulate_precision_casts: True} injected (no user kwargs - # needed). MiniMax-H3 VSA is supported only by its compile-safe sm_100a - # tile-64 inference route; other VSA routes degrade the transformer to - # eager with one warning. Dense FA2/FA3/FA4 inference uses compile-visible - # custom-op boundaries. Opt-in via FASTVIDEO_INFERENCE_TORCH_COMPILE=1 (folded in - # __post_init__) or PipelineSelection.experimental - # {"inference_torch_compile": true}. Distinct from ``enable_torch_compile``, - # which keeps the pipeline-level compile semantics. inference_torch_compile: bool = False disable_autocast: bool = False - # VSA parameters VSA_sparsity: float = 0.0 # inference/validation sparsity VSA_tile_size: int = 256 # VSA-H3 tile size (256 or 64); 64 = native Triton path - # V-MoBA parameters moba_config_path: str | None = None moba_config: dict[str, Any] = field(default_factory=dict) - # Master port for distributed training/inference master_port: int | None = None - # Stage verification enable_stage_verification: bool = True - # Prompt text file for batch processing prompt_txt: str | None = None - # LTX-2 VAE tiling overrides ltx2_vae_tiling: bool | None = None ltx2_vae_spatial_tile_size_in_pixels: int | None = None ltx2_vae_spatial_tile_overlap_in_pixels: int | None = None @@ -256,12 +177,6 @@ class FastVideoArgs: ltx2_vae_temporal_tile_overlap_in_frames: int | None = None ltx2_initial_latent_path: str | None = None ltx2_audio_latent_path: str | None = None - # Generic stage-2 refine surface (preferred user-facing API). The - # ltx2_refine_* fields below remain the runtime carriers; these - # generic ones let CLI / typed-config callers set the same values - # without binding to a specific model family. ``None`` here means - # "fall back to the model_index.json default and/or the - # ltx2_refine_* runtime carrier". refine_enabled: bool | None = None refine_upsampler_path: str | None = None refine_transformer_path: str | None = None @@ -271,11 +186,6 @@ class FastVideoArgs: refine_add_noise: bool | None = None refine_noise_path: str | None = None refine_audio_noise_path: str | None = None - # LTX-2 stage-2 spatial refinement (the SR pipeline). When enabled the - # transformer runs once at half resolution, the latents are upsampled - # by the LTX2 latent upsampler, then a short stage-2 distilled - # denoising pass refines the upsampled latents. Behaviour is opt-in - # and isolated to LTX-2 today. ltx2_refine_enabled: bool = False ltx2_refine_upsampler_path: str | None = None ltx2_refine_transformer_path: str | None = None @@ -288,7 +198,6 @@ class FastVideoArgs: ltx2_legacy_native_noise_order: bool = False ltx2_use_distilled_sigmas: bool = True - # model paths for correct deallocation model_paths: dict[str, str] = field(default_factory=dict) model_loaded: dict[str, bool] = field(default_factory=lambda: { "transformer": True, @@ -298,14 +207,6 @@ class FastVideoArgs: override_text_encoder_safetensors: str | None = None # path to safetensors file for text encoder override override_text_encoder_quant: QuantizationMethods = None - # Typed transformer quantization carrier. The typed inference API - # accepts ``engine.quantization.transformer_quant: "NVFP4"`` and the - # compat layer resolves the name to a concrete ``QuantizationConfig`` - # instance (e.g. ``NVFP4Config()``); ``__post_init__`` then pins it on - # ``pipeline_config.dit_config.quant_config`` so the loader can detect - # FP4 layers via the standard ``get_quant_method`` path. ``None`` - # leaves whatever value the caller already set on ``dit_config`` - # untouched. transformer_quant: Any | None = None override_transformer_cls_name: str | None = None @@ -314,10 +215,7 @@ class FastVideoArgs: override_pipeline_cls_name: str | None = None - # # DMD parameters - # dmd_denoising_steps: List[int] | None = field(default=None) - # MoE parameters used by Wan2.2 boundary_ratio: float | None = 0.875 @property @@ -339,22 +237,13 @@ def __post_init__(self): self._resolve_refine_args() self._apply_transformer_quant() if not self.inference_torch_compile: - # Parse-once adapter (same pattern as attention_backend below): the - # environment variable is an input read once here, so the loader - # only ever consults the typed field. import fastvideo.envs as envs if envs.FASTVIDEO_INFERENCE_TORCH_COMPILE: self.inference_torch_compile = True if self.attention_backend is not None: - # Fail fast on typos instead of silently auto-selecting later. from fastvideo.attention.selector import coerce_attn_backend coerce_attn_backend(self.attention_backend) else: - # Parse-once adapter: fold the environment variable into the typed - # request so resolution has a single input and library code never - # consults the environment on the load path. The env var keeps its - # historically permissive parse — an unknown name is ignored here - # and falls through to automatic selection rather than raising. import fastvideo.envs as envs from fastvideo.attention.selector import backend_name_to_enum env_backend = envs.FASTVIDEO_ATTENTION_BACKEND @@ -367,9 +256,6 @@ def _fold_vae_parallel_env(self) -> None: """Parse-once adapters for the sequence-parallel VAE env vars.""" import fastvideo.envs as envs - # Mirrors fastvideo.models.vaes.minimax_h3_parallel.DECODE_GATHER_STRATEGIES / - # DEFAULT_DECODE_GATHER_STRATEGY (kept literal here so constructing args - # never imports model modules; a unit test pins the two in sync). strategies = ("gather", "all_gather") if not self.vae_parallel_decode and envs.FASTVIDEO_VAE_PARALLEL_DECODE: self.vae_parallel_decode = True @@ -395,15 +281,10 @@ def _apply_transformer_quant(self) -> None: dit_config = getattr(self.pipeline_config, "dit_config", None) if dit_config is None: return - # Resolve a registry name (e.g. "nvfp4_qat_train" from the CLI) to a - # QuantizationConfig instance; a bare string has no get_quant_method. tq = self.transformer_quant if isinstance(tq, str): from fastvideo.layers.quantization import get_quantization_config tq = get_quantization_config(tq)() - # Don't overwrite if the caller already set it explicitly on - # dit_config (e.g. via ``pipeline_config.dit_config.quant_config = NVFP4Config()``); - # the explicit setter wins. if getattr(dit_config, "quant_config", None) is None: dit_config.quant_config = tq @@ -459,14 +340,12 @@ def _apply_ltx2_vae_overrides(self) -> None: @staticmethod def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: - # Model and path configuration parser.add_argument( "--model-path", type=str, help="The path of the model weights. This can be a local folder or a Hugging Face repo ID.", ) - # Running mode parser.add_argument( "--mode", type=str, @@ -475,7 +354,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="The mode to run FastVideo", ) - # Workload type parser.add_argument( "--workload-type", type=str, @@ -484,7 +362,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="The workload type", ) - # distributed_executor_backend parser.add_argument( "--distributed-executor-backend", type=str, @@ -500,7 +377,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="Whether to use inference mode", ) - # HuggingFace specific parameters parser.add_argument( "--trust-remote-code", action=StoreBoolean, @@ -514,7 +390,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="The specific model version to use (can be a branch name, tag name, or commit id)", ) - # Parallelism parser.add_argument( "--num-gpus", type=int, @@ -552,7 +427,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="Set timeout for torch.distributed initialization.", ) - # Output type parser.add_argument( "--output-type", type=str, @@ -561,7 +435,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="Output type for the generated video", ) - # Attention backend (process-wide default request) parser.add_argument( "--attention-backend", type=str, @@ -572,7 +445,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: "then per-layer defaults, then automatic selection.", ) - # Prompt text file for batch processing parser.add_argument( "--prompt-txt", type=str, @@ -580,7 +452,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="Path to a text file containing prompts (one per line) for batch processing", ) - # LTX-2 VAE tiling overrides parser.add_argument( "--ltx2-vae-tiling", action=StoreBoolean, @@ -618,7 +489,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="Path to load/save a precomputed LTX-2 initial latent.", ) - # LoRA parameters (inference-time adapter loading) parser.add_argument( "--lora-path", type=str, @@ -645,7 +515,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="Optional list of module name substrings to restrict LoRA injection (e.g. q_proj k_proj v_proj).", ) - # BSA runtime control (LongCat) parser.add_argument( "--enable-bsa", action=StoreBoolean, @@ -688,7 +557,9 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: type=str, default=None, help= - "JSON string of kwargs to pass to torch.compile. Example: '{\"backend\":\"inductor\",\"mode\":\"reduce-overhead\"}'", + "JSON string of kwargs to pass to torch.compile. Example: '{\"backend\":\"inductor\",\"mode\":\"reduce-overhead\"}'. " + "Note: the modular fastvideo/train stack uses regional fullgraph compile, which rejects 'mode' " + "(it injects inductor options); express mode effects via 'options' there.", ) parser.add_argument( "--inference-torch-compile", @@ -793,7 +664,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="Disable autocast for denoising loop and vae decoding in pipeline sampling", ) - # VSA parameters parser.add_argument( "--VSA-sparsity", type=float, @@ -807,7 +677,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="VSA-H3 tile size in tokens (256 or 64); 64 runs the native Triton block-sparse path", ) - # Master port for distributed training/inference parser.add_argument( "--master-port", type=int, @@ -815,7 +684,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="Master port for distributed training/inference", ) - # Stage verification parser.add_argument( "--enable-stage-verification", action=StoreBoolean, @@ -854,10 +722,8 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: type=str, help="Path to safetensors file for initial weight loading") - # Add pipeline configuration arguments PipelineConfig.add_cli_args(parser) - # Add preprocessing configuration arguments PreprocessConfig.add_cli_args(parser) return parser @@ -865,10 +731,8 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: @classmethod def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs": provided_args = clean_cli_args(args) - # Get all fields from the dataclass attrs = [attr.name for attr in dataclasses.fields(cls)] - # Create a dictionary of attribute values, with defaults for missing attributes kwargs: dict[str, Any] = {} for attr in attrs: if attr == 'pipeline_config': @@ -878,11 +742,9 @@ def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs": preprocess_config = PreprocessConfig.from_kwargs(provided_args) kwargs['preprocess_config'] = preprocess_config elif attr == 'mode': - # Convert string to ExecutionMode enum mode_value = getattr(args, attr, FastVideoArgs.mode.value) kwargs['mode'] = ExecutionMode.from_string(mode_value) if isinstance(mode_value, str) else mode_value elif attr == 'torch_compile_kwargs': - # Parse JSON string for torch.compile kwargs torch_compile_kwargs_str = getattr(args, 'torch_compile_kwargs', None) if torch_compile_kwargs_str: try: @@ -893,17 +755,13 @@ def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs": else: kwargs['torch_compile_kwargs'] = {} elif attr == 'workload_type': - # Convert string to WorkloadType enum workload_type_value = getattr(args, 'workload_type', FastVideoArgs.workload_type.value) kwargs['workload_type'] = WorkloadType.from_string(workload_type_value) if isinstance( workload_type_value, str) else workload_type_value - # Use getattr with default value from the dataclass for potentially missing attributes else: - # Get the field to check if it has a default_factory field = dataclasses.fields(cls)[next(i for i, f in enumerate(dataclasses.fields(cls)) if f.name == attr)] if field.default_factory is not dataclasses.MISSING: - # Use the default_factory to create the default value default_value = field.default_factory() else: default_value = getattr(cls, attr, None) @@ -914,29 +772,22 @@ def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs": @classmethod def from_kwargs(cls, **kwargs: Any) -> "FastVideoArgs": - # Convert mode string to enum if necessary if 'mode' in kwargs and isinstance(kwargs['mode'], str): kwargs['mode'] = ExecutionMode.from_string(kwargs['mode']) - # Convert workload_type string to enum if necessary if 'workload_type' in kwargs and isinstance(kwargs['workload_type'], str): kwargs['workload_type'] = WorkloadType.from_string(kwargs['workload_type']) kwargs['pipeline_config'] = PipelineConfig.from_kwargs(kwargs) kwargs['preprocess_config'] = PreprocessConfig.from_kwargs(kwargs) - # Filter to only FastVideoArgs dataclass fields — pipeline-specific CLI - # args (e.g. enable_bsa, bsa_sparsity) live in PipelineConfig and must - # not be forwarded to the FastVideoArgs constructor. valid_fields = {f.name for f in dataclasses.fields(cls)} return cls(**{k: v for k, v in kwargs.items() if k in valid_fields}) def check_fastvideo_args(self) -> None: """Validate inference arguments for consistency""" - # Validate mode and inference_mode consistency assert isinstance(self.mode, ExecutionMode), f"Mode must be an ExecutionMode enum, got {type(self.mode)}" assert self.mode in ExecutionMode.choices(), f"Invalid execution mode: {self.mode}" - # Validate workload type assert isinstance(self.workload_type, WorkloadType), f"Workload type must be a WorkloadType enum, got {type(self.workload_type)}" assert self.workload_type in WorkloadType.choices(), f"Invalid workload type: {self.workload_type}" @@ -948,11 +799,6 @@ def check_fastvideo_args(self) -> None: logger.warning("Mode is '%s' but inference_mode is False. Setting inference_mode to True.", self.mode) self.inference_mode = True - # Inference policy must wait until a worker owns and binds its device: - # a unified-memory device disables layerwise offload before conflicts - # are resolved, preserving an explicit FSDP request. Training does not - # pass through the inference worker boundary, so retain its historical - # constructor-time normalization. if not self.inference_mode: self._resolve_device_offload_conflicts() @@ -980,7 +826,6 @@ def check_fastvideo_args(self) -> None: self.pipeline_config.check_pipeline_config() - # Add preprocessing config validation if needed if self.mode == ExecutionMode.PREPROCESS: if self.preprocess_config is None: raise ValueError("preprocess_config is not set in FastVideoArgs when mode is PREPROCESS") @@ -1054,11 +899,6 @@ def disable_offload_on_unified_memory(self, device_id: int = 0, *, offload_flag: try: device_name = current_platform.get_device_name(device_id) except Exception: - # Device naming is diagnostic only. NVML can be unavailable on - # an integrated GPU (for example Jetson), and its physical- - # ordinal lookup cannot interpret CUDA_VISIBLE_DEVICES UUID/MIG - # selectors. Neither case should undo an authoritative driver - # classification. device_name = current_platform.device_name for flag in enabled_flags: @@ -1112,10 +952,6 @@ def set_current_fastvideo_args(fastvideo_args: FastVideoArgs): def get_current_fastvideo_args() -> FastVideoArgs: if _current_fastvideo_args is None: - # in ci, usually when we test custom ops/modules directly, - # we don't set the fastvideo config. In that case, we set a default - # config. - # TODO(will): may need to handle this for CI. raise ValueError("Current fastvideo args is not set.") return _current_fastvideo_args @@ -1138,20 +974,16 @@ class TrainingArgs(FastVideoArgs): group_frame: bool = False group_resolution: bool = False - # text encoder & vae & diffusion model pretrained_model_name_or_path: str = "" - # DMD model paths - separate paths for each network real_score_model_path: str = "" # path for real score (teacher) model fake_score_model_path: str = "" # path for fake score (critic) model - # diffusion setting ema_decay: float = 0.0 ema_start_step: int = 0 training_cfg_rate: float = 0.0 precondition_outputs: bool = False - # validation & logs validation_dataset_file: str = "" validation_preprocessed_path: str = "" validation_sampling_steps: str = "" @@ -1163,12 +995,10 @@ class TrainingArgs(FastVideoArgs): wandb_run_name: str = "" seed: int | None = None - # output output_dir: str = "" checkpoints_total_limit: int = 0 resume_from_checkpoint: str = "" # specify the checkpoint folder to resume from - # optimizer & scheduler num_train_epochs: int = 0 max_train_steps: int = 0 gradient_accumulation_steps: int = 0 @@ -1205,37 +1035,21 @@ class TrainingArgs(FastVideoArgs): pred_decay_type: str = "" hunyuan_teacher_disable_cfg: bool = False - # master_weight_type master_weight_type: str = "" - # VSA training decay parameters VSA_decay_rate: float = 0.01 # decay rate -> 0.02 VSA_decay_interval_steps: int = 1 # decay interval steps -> 50 - # Reuse the per-step padded VSA tile buffer across attention layers during - # training. Defaults to False: under full activation checkpointing the - # cached buffer survives into the backward recompute and inflates peak - # memory (see #1423). Enable on memory-rich setups to keep the per-step - # buffer-reuse speedup. VSA_cache_tile_buf: bool = False - # LoRA training parameters lora_rank: int | None = None lora_alpha: int | None = None lora_training: bool = False ltx2_first_frame_conditioning_p: float = 0.1 - # distillation args generator_update_interval: int = 5 dfake_gen_update_ratio: int = 5 # self-forcing: how often to train generator vs critic min_timestep_ratio: float = 0.2 max_timestep_ratio: float = 0.98 - # CFG scale applied to the real (teacher) score in the DMD loss, using the - # parameterization `x = x_cond + w * (x_cond - x_uncond)`. This differs - # from the Ho & Salimans form `x_uncond + w * (x_cond - x_uncond)` by an - # offset of 1: `w_here = w_standard - 1`. So `w=0` recovers the - # conditional output, `w=-1` recovers the unconditional output, and the - # default 3.5 corresponds to a standard CFG scale of 4.5. Matches the - # original DMD2 reference implementation. real_score_guidance_scale: float = 3.5 fake_score_learning_rate: float = 0.0 # separate learning rate for fake_score_transformer, if 0.0, use learning_rate fake_score_lr_scheduler: str = "constant" # separate lr scheduler for fake_score_transformer, if not set, use lr_scheduler @@ -1244,11 +1058,9 @@ class TrainingArgs(FastVideoArgs): weight_only_checkpointing_steps: int = 0 # for inference log_visualization: bool = False visualization_steps: int = 0 - # simulate generator forward to match inference simulate_generator_forward: bool = False warp_denoising_step: bool = False - # Self-forcing specific arguments num_frame_per_block: int = 3 independent_first_frame: bool = False enable_gradient_masking: bool = True @@ -1260,41 +1072,32 @@ class TrainingArgs(FastVideoArgs): @classmethod def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs": provided_args = clean_cli_args(args) - # Get all fields from the dataclass attrs = [attr.name for attr in dataclasses.fields(cls)] logger.info(provided_args) - # Create a dictionary of attribute values, with defaults for missing attributes kwargs: dict[str, Any] = {} for attr in attrs: if attr == 'pipeline_config': pipeline_config = PipelineConfig.from_kwargs(provided_args) kwargs[attr] = pipeline_config elif attr == 'mode': - # Convert string to ExecutionMode enum mode_value = getattr(args, attr, ExecutionMode.FINETUNING.value) kwargs[attr] = ExecutionMode.from_string(mode_value) if isinstance(mode_value, str) else mode_value elif attr == 'workload_type': - # Convert string to WorkloadType enum workload_type_value = getattr(args, 'workload_type', WorkloadType.T2V.value) kwargs[attr] = WorkloadType.from_string(workload_type_value) if isinstance(workload_type_value, str) else workload_type_value - # Use getattr with default value from the dataclass for potentially missing attributes else: - # Get the field to check its default value field = dataclasses.fields(cls)[next(i for i, f in enumerate(dataclasses.fields(cls)) if f.name == attr)] - # Check if the attribute is provided in args if hasattr(args, attr): value = getattr(args, attr) else: - # Use the field's default value if field.default_factory is not dataclasses.MISSING: value = field.default_factory() elif field.default is not dataclasses.MISSING: value = field.default else: - # No default value, use None value = None kwargs[attr] = value @@ -1318,7 +1121,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: parser.add_argument("--num-width", type=int, required=True, help="Number of widths") parser.add_argument("--num-frames", type=int, required=True, help="Number of frames") - # Training batch and model configuration parser.add_argument("--train-batch-size", type=int, required=True, help="Training batch size") parser.add_argument("--num-latent-t", type=int, required=True, help="Number of latent time steps") parser.add_argument("--group-frame", action=StoreBoolean, help="Whether to group frames during training") @@ -1326,7 +1128,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: action=StoreBoolean, help="Whether to group resolutions during training") - # Model paths parser.add_argument("--pretrained-model-name-or-path", type=str, required=True, @@ -1337,7 +1138,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="Path to DiT model or model name") parser.add_argument("--cache-dir", type=str, help="Directory to cache models") - # DMD model paths - separate paths for each network parser.add_argument("--generator-model-path", type=str, help="Path to generator (student) model for DMD distillation") @@ -1348,7 +1148,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: type=str, help="Path to fake score (critic) model for DMD distillation") - # Diffusion settings parser.add_argument("--ema-decay", type=float, default=0.999, help="EMA decay rate") parser.add_argument("--ema-start-step", type=int, default=0, help="Step to start EMA") parser.add_argument("--training-cfg-rate", type=float, help="Classifier-free guidance scale") @@ -1356,7 +1155,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: action=StoreBoolean, help="Whether to precondition the outputs of the model") - # Validation and logging parser.add_argument("--validation-dataset-file", type=str, help="Path to unprocessed validation dataset") parser.add_argument("--validation-preprocessed-path", type=str, help="Path to processed validation dataset") parser.add_argument("--validation-sampling-steps", type=str, help="Validation sampling steps") @@ -1368,7 +1166,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: parser.add_argument("--wandb-run-name", type=str, help="Run name for wandb") parser.add_argument("--seed", type=int, default=42, help="Seed for deterministic training") - # Output configuration parser.add_argument("--output-dir", type=str, required=True, help="Output directory for checkpoints and logs") parser.add_argument("--checkpoints-total-limit", type=int, help="Maximum number of checkpoints to keep") parser.add_argument("--training-state-checkpointing-steps", @@ -1380,7 +1177,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: parser.add_argument("--resume-from-checkpoint", type=str, help="Path to checkpoint to resume from") parser.add_argument("--logging-dir", type=str, help="Directory for logging") - # Training configuration parser.add_argument("--num-train-epochs", type=int, help="Number of training epochs") parser.add_argument("--max-train-steps", type=int, help="Maximum number of training steps") parser.add_argument("--gradient-accumulation-steps", type=int, help="Number of steps to accumulate gradients") @@ -1425,7 +1221,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.", ) - # Additional training parameters parser.add_argument("--num-euler-timesteps", type=int, help="Number of Euler timesteps") parser.add_argument("--lr-num-cycles", type=int, help="Number of learning rate cycles") parser.add_argument("--lr-power", type=float, help="Learning rate power") @@ -1452,7 +1247,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="Whether to disable CFG for Hunyuan teacher") parser.add_argument("--master-weight-type", type=str, help="Master weight type") - # VSA parameters for training with dense to sparse adaption parser.add_argument( "--VSA-decay-rate", # decay rate, how much sparsity you want to decay each step type=float, @@ -1480,7 +1274,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="Probability of conditioning on the first frame during LTX-2 training", ) - # V-MoBA parameters parser.add_argument( "--moba-config-path", type=str, @@ -1488,7 +1281,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: help="Path to a JSON file containing V-MoBA specific configurations.", ) - # Distillation arguments parser.add_argument("--generator-update-interval", type=int, default=TrainingArgs.generator_update_interval, @@ -1533,7 +1325,6 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: action=StoreBoolean, help="Whether to warp denoising step according to the scheduler time shift") - # Self-forcing specific arguments parser.add_argument("--num-frame-per-block", type=int, default=TrainingArgs.num_frame_per_block, diff --git a/fastvideo/layers/quantization/__init__.py b/fastvideo/layers/quantization/__init__.py index 83f5076f1c..7151b8be84 100644 --- a/fastvideo/layers/quantization/__init__.py +++ b/fastvideo/layers/quantization/__init__.py @@ -11,11 +11,12 @@ "nvfp4_qat", "nvfp4_qat_train", "fp8_qat_train", + "INT8Affine", + "W4A16", ] QUANTIZATION_METHODS: list[str] = list(get_args(QuantizationMethods)) -# The customized quantization methods which will be added to this dict. _CUSTOMIZED_METHOD_TO_QUANT_CONFIG = {} @@ -58,25 +59,25 @@ def get_quantization_config(quantization: str) -> type[QuantizationConfig]: if quantization not in QUANTIZATION_METHODS: raise ValueError(f"Invalid quantization method: {quantization}") - # lazy import to avoid triggering `torch.compile` too early from .absmax_fp8 import AbsMaxFP8Config from .fp8_config import FP8Config - from .mxfp8_config import MXFP8Config from .nvfp4_config import NVFP4Config from .nvfp4_qat_config import NVFP4QATConfig from .nvfp4_qat_train_config import NVFP4QATTrainConfig from .fp8_qat_train_config import FP8QATTrainConfig + from .int8_affine_config import INT8AffineConfig + from .w4a16_config import W4A16Config method_to_config: dict[str, type[QuantizationConfig]] = { "AbsMaxFP8": AbsMaxFP8Config, "FP8": FP8Config, - "MXFP8": MXFP8Config, "NVFP4": NVFP4Config, "nvfp4_qat": NVFP4QATConfig, "nvfp4_qat_train": NVFP4QATTrainConfig, "fp8_qat_train": FP8QATTrainConfig, + "INT8Affine": INT8AffineConfig, + "W4A16": W4A16Config, } - # Update the `method_to_config` with customized quantization methods. method_to_config.update(_CUSTOMIZED_METHOD_TO_QUANT_CONFIG) return method_to_config[quantization] diff --git a/fastvideo/layers/quantization/int8_affine_config.py b/fastvideo/layers/quantization/int8_affine_config.py new file mode 100644 index 0000000000..43eb7801e6 --- /dev/null +++ b/fastvideo/layers/quantization/int8_affine_config.py @@ -0,0 +1,841 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Weight-only affine INT8 quantization (group size 64) for CUDA inference.""" + +from __future__ import annotations + +import json +import logging +import os +from collections.abc import Iterable +from typing import Any + +import torch +import torch.nn.functional as F +from torch.nn.parameter import Parameter + +from fastvideo.layers.quantization.base_config import ( + QuantizationConfig, + QuantizeMethodBase, +) +from fastvideo.models.utils import set_weight_attrs + +logger = logging.getLogger(__name__) + +DEFAULT_GROUP_SIZE = 64 +DEFAULT_BITS = 8 +_EPS = 1e-7 +_MAX_UINT8_CODE = 255 + + + + +def _group(w: torch.Tensor, group_size: int) -> torch.Tensor: + """View the last dim as ``(num_groups, group_size)``. + + Identical to ``mlx_affine_qat._group``: MLX groups along the last + (input) axis, which for a linear weight is the contraction dimension. + """ + if w.shape[-1] % group_size != 0: + raise ValueError(f"Last dim {w.shape[-1]} is not divisible by group_size {group_size}; " + "MLX affine quantization groups along the last axis.") + return w.reshape(*w.shape[:-1], w.shape[-1] // group_size, group_size) + + +def int8_affine_quantize( + w: torch.Tensor, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Quantize like ``mx.quantize(..., mode="affine")``. + + Transcribed from ``mlx_affine_qat.mlx_affine_quantize_reference``. The + decisions reproduced exactly are: per-group min/max in the input's + arithmetic, a *negative* scale when ``|w_max| >= |w_min|`` (the + quantizer anchors at the endpoint with the larger magnitude), the anchor + re-expressed as an exact integer multiple of the scale so the extreme + value round-trips exactly, ``rint`` (round-half-to-even) rounding, and + codes clamped to ``[0, 2**bits - 1]``. + + Bit-identical to ``mlx_affine_quantize_reference`` for a given input + dtype; the only representation change is that codes are ``torch.uint8`` + rather than ``torch.int32``, because ``2**bits - 1 == 255`` does not fit + a signed byte. Callers must not cast these to ``int8`` — 255 would wrap + to -1 and the dequantized weight would be wrong. + + ``scales``/``biases`` come back in ``w.dtype``, as in the reference. + Since the codes are decided *before* that cast, calling this on an fp32 + ``w`` costs nothing in code fidelity and avoids rounding the stored + scales to bf16 — which is what ``_quantize_layer_weight`` does (a bf16 + checkpoint value converts to fp32 exactly, so this is lossless input + with a higher-precision scale store). + + Returns ``(codes, scales, biases)``. Mirroring the reference, ``codes`` + comes back in the *grouped* shape ``w.shape[:-1] + (K // group_size, group_size)`` + (not ``w.shape``) and ``scales``/``biases`` in + ``w.shape[:-1] + (K // group_size,)``; pass ``out_shape=w.shape`` to + ``int8_affine_dequantize`` to flatten the grouping back out. + """ + n_bins = float((1 << bits) - 1) + grouped = _group(w, group_size).float() + + w_min = grouped.min(dim=-1).values + w_max = grouped.max(dim=-1).values + mask = w_min.abs() > w_max.abs() + scale = ((w_max - w_min) / n_bins).clamp_min(_EPS) + scale = torch.where(mask, scale, -scale) + edge = torch.where(mask, w_min, w_max) + q0 = torch.round(edge / scale) + nonzero_q0 = q0 != 0 + scale = torch.where(nonzero_q0, edge / torch.where(nonzero_q0, q0, torch.ones_like(q0)), scale) + bias = torch.where(nonzero_q0, edge, torch.zeros_like(edge)) + + codes = torch.round((grouped - bias.unsqueeze(-1)) / scale.unsqueeze(-1)) + codes = codes.clamp(min=0.0, max=n_bins) + if codes.max().item() > _MAX_UINT8_CODE: + raise ValueError(f"bits={bits} produced a code above {_MAX_UINT8_CODE}; only bits<=8 fits uint8 storage.") + return codes.to(torch.uint8), scale.to(w.dtype), bias.to(w.dtype) + + +def int8_affine_dequantize( + codes: torch.Tensor, + scales: torch.Tensor, + biases: torch.Tensor, + *, + out_shape: torch.Size | None = None, +) -> torch.Tensor: + """``code * scale + bias`` per group — the inverse of ``int8_affine_quantize``. + + Mirrors ``mlx_affine_qat.mlx_affine_dequantize_reference`` (which itself + matches MLX's *CPU* kernel). ``codes`` may be ``uint8``; the multiply is + done in the scales' dtype, so pass fp32 scales to get the fp32 stream. + + ``codes`` is accepted in either shape the quantizer's callers use: the + grouped ``(*, K // group_size, group_size)`` the reference returns, or the + flattened ``(*, K)`` a stored weight buffer naturally has. The two are + distinguished by rank (grouped codes are one rank above ``scales``). + """ + dtype = scales.dtype + if codes.dim() == scales.dim(): + codes = codes.reshape(*scales.shape, codes.shape[-1] // scales.shape[-1]) + deq = codes.to(dtype) * scales.unsqueeze(-1) + biases.unsqueeze(-1) + if out_shape is not None: + deq = deq.reshape(out_shape) + return deq + + + +_GENERIC_LINEAR_SUFFIXES: tuple[str, ...] = ( + "attn.to_q", + "attn.to_k", + "attn.to_v", + "attn.to_out", + "ff.fc_in", + "ff.fc_out", +) + +_NEVER_QUANTIZE_SUBSTRINGS: tuple[str, ...] = ( + "to_gate_compress", + "adaln_basis", +) + +_H3_FP32_KEPT_SUBSTRINGS: tuple[str, ...] = ( + "proj_in", + "audio_proj_in", + "proj_out", + "audio_proj_out", + "time_embedder", +) + +_H3_INPUT_PROJECTION_SUBSTRINGS: tuple[str, ...] = ("context_embedder", ) + +MINIMAX_H3_PREFIX = "minimax_h3" +MINIMAX_H3_NUM_LAYERS = 50 +MINIMAX_H3_NUM_REFINER_LAYERS = 2 +MINIMAX_H3_BLOCK_SCOPES: tuple[str, ...] = ( + "transformer_blocks", + "token_refiner.refiner_blocks", +) +MINIMAX_H3_INT8_AFFINE_SUFFIXES: tuple[str, ...] = _GENERIC_LINEAR_SUFFIXES + ("adaln_proj.linear", ) +MINIMAX_H3_INT8_AFFINE_EXCLUSIONS: tuple[str, ...] = ( + _NEVER_QUANTIZE_SUBSTRINGS + _H3_FP32_KEPT_SUBSTRINGS + _H3_INPUT_PROJECTION_SUBSTRINGS) + + +def minimax_h3_int8_affine_prefixes( + *, + prefix: str = MINIMAX_H3_PREFIX, + num_layers: int = MINIMAX_H3_NUM_LAYERS, + num_refiner_layers: int = MINIMAX_H3_NUM_REFINER_LAYERS, + suffixes: Iterable[str] = MINIMAX_H3_INT8_AFFINE_SUFFIXES, +) -> frozenset[str]: + """Enumerate the exact H3 linear prefixes this config targets. + + Built from H3's real module names as constructed in + ``fastvideo/models/dits/minimax_h3.py``: + ``MiniMaxH3TransformerBlock`` builds ``{prefix}.transformer_blocks.{i}.attn``, + ``.ff``, ``.adaln_proj``; ``MiniMaxH3TokenRefiner`` builds + ``{prefix}.token_refiner.refiner_blocks.{i}.attn`` / ``.ff``. Defaults + match ``MiniMaxH3ArchConfig`` (``prefix="minimax_h3"``, ``num_layers=50``, + ``num_refiner_layers=2``). + + The enumerated set is *not* how selection runs at runtime (suffix + + deny-list matching is, so depth changes cannot silently drop layers); + it exists to be asserted against in tests and to give callers who want + a literal set one place to get it. + """ + suffixes = tuple(suffixes) + prefixes: set[str] = set() + for index in range(num_layers): + for suffix in suffixes: + prefixes.add(f"{prefix}.transformer_blocks.{index}.{suffix}") + for index in range(num_refiner_layers): + for suffix in suffixes: + if suffix.startswith("adaln_proj"): + continue + prefixes.add(f"{prefix}.token_refiner.refiner_blocks.{index}.{suffix}") + return frozenset(prefixes) + + +class INT8AffineConfig(QuantizationConfig): + """Weight-only affine INT8 (group-64) quantization for CUDA DiT inference. + + Layer selection is a constructor field, not a hardcoded model list: + ``target_layers`` (explicit full prefixes) takes precedence when given, + otherwise ``layer_suffixes`` is matched with ``str.endswith``. Both are + subject to a fail-closed deny list — see ``exclude_substrings``. + + Weight-only: there is no activation quantizer, and ``INT8AffineQuantizeMethod.apply`` + dequantizes the stored codes back to the activation dtype and runs a + normal bf16/fp32 GEMM. That is the correctness-first path; a fused INT8 + GEMM is a follow-up, not a prerequisite. + + Only the INT8 arithmetic is scheme-specific — nothing here is H3-only. + Use :meth:`for_minimax_h3` for the H3 profile. + """ + + def __init__( + self, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + target_layers: Iterable[str] | None = None, + layer_suffixes: Iterable[str] | None = None, + exclude_substrings: Iterable[str] | None = None, + include_context_embedder: bool = False, + retain_original_weight: bool = True, + ) -> None: + super().__init__() + if bits < 2 or bits > 8: + raise ValueError(f"bits must be in [2, 8] (codes are stored as uint8), got {bits}") + if group_size <= 0: + raise ValueError(f"group_size must be positive, got {group_size}") + self.group_size = group_size + self.bits = bits + self.target_layers: frozenset[str] | None = (None if target_layers is None else frozenset(target_layers)) + self.layer_suffixes: tuple[str, ...] = (tuple(_GENERIC_LINEAR_SUFFIXES) + if layer_suffixes is None else tuple(layer_suffixes)) + self.exclude_substrings: tuple[str, ...] = tuple(_NEVER_QUANTIZE_SUBSTRINGS) + tuple( + exclude_substrings or ()) + self._include_context_embedder = include_context_embedder + if not include_context_embedder: + self.exclude_substrings = self.exclude_substrings + _H3_INPUT_PROJECTION_SUBSTRINGS + self.retain_original_weight = retain_original_weight + + def get_name(self) -> str: + return "INT8Affine" + + def get_supported_act_dtypes(self) -> list[torch.dtype]: + return [torch.bfloat16, torch.float16, torch.float32] + + @classmethod + def get_min_capability(cls) -> int: + """Turing (75). + + The compute path is a plain bf16/fp32 GEMM over a dequantized weight, + so no INT8 tensor-core class is required; 75 matches ``AbsMaxFP8Config`` + and keeps the config loadable on the same hosts. + """ + return 75 + + @staticmethod + def get_config_filenames() -> list[str]: + return [] + + @classmethod + def from_config(cls, config: dict[str, Any]) -> INT8AffineConfig: + return cls( + group_size=config.get("group_size", DEFAULT_GROUP_SIZE), + bits=config.get("bits", DEFAULT_BITS), + target_layers=config.get("target_layers"), + layer_suffixes=config.get("layer_suffixes"), + exclude_substrings=config.get("exclude_substrings"), + include_context_embedder=config.get("include_context_embedder", False), + retain_original_weight=config.get("retain_original_weight", True), + ) + + @classmethod + def for_minimax_h3( + cls, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + include_context_embedder: bool = False, + retain_original_weight: bool = True, + ) -> INT8AffineConfig: + """The verified MiniMax-H3 profile: attention + FFN + AdaLN GEMMs. + + Excludes H3's fp32-pinned modules, the VSA gate, and (by default) + the text input projection. + """ + return cls( + group_size=group_size, + bits=bits, + layer_suffixes=MINIMAX_H3_INT8_AFFINE_SUFFIXES, + exclude_substrings=_H3_FP32_KEPT_SUBSTRINGS, + include_context_embedder=include_context_embedder, + retain_original_weight=retain_original_weight, + ) + + def is_target_layer(self, prefix: str) -> bool: + """Whether ``prefix`` is quantized under this config. + + Deny list first (fail-closed), then ``target_layers`` if supplied, + else suffix matching. Non-``LinearBase`` layers are filtered by + :meth:`get_quant_method`, not here, so this is safe to call on any + module name. + """ + for banned in self.exclude_substrings: + if banned in prefix: + return False + if self.target_layers is not None: + return prefix in self.target_layers + return any(prefix.endswith(suffix) for suffix in self.layer_suffixes) + + def get_quant_method(self, layer: torch.nn.Module, prefix: str): + from fastvideo.layers.linear import LinearBase + + if isinstance(layer, LinearBase) and self.is_target_layer(prefix): + return INT8AffineQuantizeMethod( + layer_prefix=prefix, + group_size=self.group_size, + bits=self.bits, + retain_original_weight=self.retain_original_weight, + ) + return None + + +class INT8AffineQuantizeMethod(QuantizeMethodBase): + """Linear method for weight-only affine INT8. + + ``create_weights`` allocates the same dense bf16 Parameter an + unquantized linear would (so the BF16 checkpoint loads unchanged), and + the INT8 codes/scales/biases arrive later as non-persistent buffers from + :func:`convert_model_to_int8_affine` — i.e. conversion happens at *load* + time, not construction time, exactly mirroring ``NVFP4QuantizeMethod``. + """ + + def __init__( + self, + layer_prefix: str = "", + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + retain_original_weight: bool = True, + ) -> None: + super().__init__() + self.layer_prefix = layer_prefix + self.group_size = group_size + self.bits = bits + self.retain_original_weight = retain_original_weight + + def create_weights( + self, + layer: torch.nn.Module, + input_size_per_partition: int, + output_partition_sizes: list[int], + input_size: int, + output_size: int, + params_dtype: torch.dtype, + **extra_weight_attrs, + ): + weight = Parameter( + torch.empty( + sum(output_partition_sizes), + input_size_per_partition, + dtype=params_dtype, + ), + requires_grad=False, + ) + set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0}) + layer.register_parameter("weight", weight) + set_weight_attrs(weight, extra_weight_attrs) + + def _ensure_quantized(self, layer: torch.nn.Module) -> bool: + """Convert on first use if the loader hook never ran. + + Returns False when the layer is intentionally left dense (grad-enabled + forward: a training step must see the master weight, not a frozen + dequantized copy). The loader path is ``_maybe_quantize_model`` -> + :func:`convert_model_to_int8_affine`; this fallback exists so the + config is still *correct* if that dispatch is missing, but it warns + because reaching it means the loader hook did not fire. + """ + if getattr(layer, "_int8_affine_codes", None) is not None: + return True + weight = getattr(layer, "weight", None) + if weight is None: + raise RuntimeError(f"INT8Affine layer {self.layer_prefix!r} has no weight and no quantized buffers.") + if torch.is_grad_enabled(): + return False + logger.warning( + "INT8Affine: layer %r reached apply() unquantized; converting lazily. The loader hook " + "(_maybe_quantize_model) did not dispatch to convert_model_to_int8_affine — check its " + "isinstance chain in fastvideo/models/loader/fsdp_load.py.", + self.layer_prefix, + ) + _quantize_layer_weight(layer, weight, group_size=self.group_size, bits=self.bits) + return True + + def apply( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + if not self._ensure_quantized(layer): + weight = layer.weight + return F.linear(x, weight.to(x.dtype) if weight.dtype != x.dtype else weight, bias) + + codes = layer._int8_affine_codes + weight = int8_affine_dequantize( + codes, + layer._int8_affine_scales, + layer._int8_affine_biases, + out_shape=codes.shape, + ).to(x.dtype) + return F.linear(x, weight, bias) + + + + +def _quantize_layer_weight( + mod: torch.nn.Module, + weight: torch.Tensor, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, +) -> None: + """Quantize one linear's weight in place into non-persistent buffers.""" + from torch.distributed.tensor import DTensor # type: ignore + + weight_local = weight.to_local() if isinstance(weight, DTensor) else weight # type: ignore[arg-type] + w32 = weight_local.detach().float().nan_to_num() + if w32.shape[-1] % group_size != 0: + raise ValueError(f"INT8Affine layer {mod!r}: input dim {w32.shape[-1]} is not divisible by " + f"group_size {group_size}.") + codes, scales, biases = int8_affine_quantize(w32, group_size=group_size, bits=bits) + mod.register_buffer("_int8_affine_codes", codes.reshape(w32.shape).contiguous(), persistent=False) + mod.register_buffer("_int8_affine_scales", scales.to(torch.float32).contiguous(), persistent=False) + mod.register_buffer("_int8_affine_biases", biases.to(torch.float32).contiguous(), persistent=False) + + +def convert_model_to_int8_affine(model: torch.nn.Module, ) -> None: + """Quantize every INT8Affine-tagged linear in-place after weights load. + + Mirrors ``convert_model_to_nvfp4`` / ``convert_model_to_fp8``: walk the + module tree once, convert each layer whose ``quant_method`` is an + :class:`INT8AffineQuantizeMethod`, and register the int8 codes plus + per-group scales/biases as non-persistent buffers (so they are not + written back into ``state_dict``/checkpoints). + + Callers: the loader hook ``_maybe_quantize_model`` in + ``fastvideo/models/loader/fsdp_load.py``. *That hook is not edited by + this module* — it dispatches on an explicit ``isinstance`` chain, so it + needs a matching branch (see the module report). Without it, + ``INT8AffineQuantizeMethod.apply`` converts lazily on first forward and + logs a warning, so inference is still correct, just later and noisier. + """ + converted = 0 + purged = 0 + shapes: set[tuple[int, int]] = set() + for mod in model.modules(): + qm = getattr(mod, "quant_method", None) + if not isinstance(qm, INT8AffineQuantizeMethod): + continue + weight = getattr(mod, "weight", None) + if weight is None: + continue + _quantize_layer_weight(mod, weight, group_size=qm.group_size, bits=qm.bits) + converted += 1 + shapes.add((qm.group_size, qm.bits)) + if not qm.retain_original_weight: + original = mod._parameters.get("weight") + if original is not None: + original.grad = None + mod.register_parameter("weight", None) + purged += 1 + + if converted: + logger.info("INT8Affine conversion receipt: quantized %d linear layers (%s); purged %d original bf16 " + "weight tensors.", converted, + ", ".join(f"group_size={g}, bits={b}" for g, b in sorted(shapes)), purged) + + + +INT8_AFFINE_SIDECAR_SUFFIX = ".int8affine.safetensors" +INT8_AFFINE_DIR_SIDECAR_NAME = "int8_affine.safetensors" +_INT8_AFFINE_SIDECAR_FORMAT = "fastvideo.int8_affine" +_INT8_AFFINE_SIDECAR_VERSION = 1 +_INT8_AFFINE_SIDECAR_METADATA_KEY = "fastvideo_int8_affine" +_INT8_AFFINE_SIDECAR_KEY_SEP = "::" +_INT8_AFFINE_SIDECAR_BUFFERS = ( + "_int8_affine_codes", + "_int8_affine_scales", + "_int8_affine_biases", +) +_INT8_AFFINE_SIDECAR_DTYPES = { + "_int8_affine_codes": torch.uint8, + "_int8_affine_scales": torch.float32, + "_int8_affine_biases": torch.float32, +} + + +def _sidecar_key(module_fqn: str, buffer_name: str) -> str: + return f"{module_fqn}{_INT8_AFFINE_SIDECAR_KEY_SEP}{buffer_name}" + + +def _int8_affine_tagged_modules(model: torch.nn.Module) -> list[tuple[str, torch.nn.Module, INT8AffineQuantizeMethod]]: + tagged = [] + for fqn, mod in model.named_modules(): + qm = getattr(mod, "quant_method", None) + if isinstance(qm, INT8AffineQuantizeMethod): + tagged.append((fqn, mod, qm)) + return tagged + + +def _is_dtensor(tensor: torch.Tensor) -> bool: + try: + from torch.distributed.tensor import DTensor # type: ignore + except ImportError: # pragma: no cover - depends on the torch build + return False + return isinstance(tensor, DTensor) + + +def int8_affine_sidecar_state_dict(model: torch.nn.Module) -> dict[str, torch.Tensor]: + """Collect the quantized tensors of every INT8 affine linear in *model*. + + Keys are ``"::"`` and values are detached CPU + copies. Modules whose buffers are missing (never converted) are skipped; + the returned mapping is what :func:`save_int8_affine_checkpoint` writes. + + FSDP note: a DTensor buffer is saved as this rank's local shard, so a + sharded save is only reloadable into an identically sharded model. + """ + state: dict[str, torch.Tensor] = {} + for fqn, mod, _ in _int8_affine_tagged_modules(model): + for name in _INT8_AFFINE_SIDECAR_BUFFERS: + tensor = getattr(mod, name, None) + if tensor is None: + continue + if _is_dtensor(tensor): + tensor = tensor.to_local() # type: ignore[attr-defined] + state[_sidecar_key(fqn, name)] = tensor.detach().to("cpu", copy=True).contiguous() + return state + + +def save_int8_affine_checkpoint( + model: torch.nn.Module, + path: str | os.PathLike[str], + *, + extra_metadata: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Write the model's INT8 affine tensors to a compact sidecar safetensors file. + + The file is roughly the INT8 size (1 byte/code plus two fp32 per group — + 8.5 bits/weight at ``group_size=64``) instead of the dense bf16 size. + Returns a receipt dict (also logged) with the module count and both sizes. + Raises ``RuntimeError`` when the model has no INT8 affine linears, which + usually means the config's layer selection did not cover the model's layer + paths, and when the tagged layers carry no buffers (never converted). + + The manifest under the ``fastvideo_int8_affine`` metadata key carries the + format name/version, the scheme (``group_size``/``bits``), the per-layer + weight and buffer shapes, and the quantized module fqns (the keys of + ``layers``), so a loader can validate a sidecar against a model without + materializing the tensors. + """ + from safetensors.torch import save_file + + state = int8_affine_sidecar_state_dict(model) + tagged = _int8_affine_tagged_modules(model) + if not tagged: + raise RuntimeError("No INT8 affine linear layers found in this model; nothing to serialize. " + "Check that the model was built with an INT8AffineConfig whose layer selection " + "covers its layer paths (e.g. INT8AffineConfig.for_minimax_h3()).") + if not state: + raise RuntimeError(f"Found {len(tagged)} INT8Affine-tagged linear layers but none carry quantized " + "buffers. Call convert_model_to_int8_affine(model) before saving a sidecar.") + + layers: dict[str, dict[str, Any]] = {} + quant_prefixes: dict[str, str] = {} + dense_bytes = 0 + group_sizes: set[int] = set() + bit_widths: set[int] = set() + for fqn, mod, qm in tagged: + codes = getattr(mod, "_int8_affine_codes", None) + weight = getattr(mod, "weight", None) + if codes is None and weight is None: + continue + weight_shape = [int(dim) for dim in (weight if weight is not None else codes).shape] + tensors = { + name: [int(dim) for dim in getattr(mod, name).shape] + for name in _INT8_AFFINE_SIDECAR_BUFFERS if getattr(mod, name, None) is not None + } + layers[fqn] = { + "weight_shape": weight_shape, + "group_size": int(qm.group_size), + "bits": int(qm.bits), + "tensors": tensors, + } + quant_prefixes[fqn] = getattr(qm, "layer_prefix", "") or "" + dense_bytes += weight_shape[0] * weight_shape[1] * 2 + group_sizes.add(int(qm.group_size)) + bit_widths.add(int(qm.bits)) + + metadata: dict[str, Any] = { + "format": _INT8_AFFINE_SIDECAR_FORMAT, + "version": _INT8_AFFINE_SIDECAR_VERSION, + "group_size": group_sizes.pop() if len(group_sizes) == 1 else None, + "bits": bit_widths.pop() if len(bit_widths) == 1 else None, + "num_layers": len(layers), + "layers": layers, + "quant_prefixes": quant_prefixes, + "model_class": type(model).__name__, + } + if extra_metadata: + metadata.update(extra_metadata) + + payload = dict(state) + serialized_bytes = sum(t.numel() * t.element_size() for t in payload.values()) + save_file(payload, os.fspath(path), metadata={_INT8_AFFINE_SIDECAR_METADATA_KEY: json.dumps(metadata)}) + + receipt = { + "path": os.fspath(path), + "num_layers": len(layers), + "num_tensors": len(payload), + "quantized_bytes": serialized_bytes, + "dense_bfloat16_bytes": dense_bytes, + "compression_ratio": (dense_bytes / serialized_bytes) if serialized_bytes else 0.0, + } + logger.info( + "INT8Affine sidecar: wrote %d quantized modules / %d tensors (%d bytes) to %s " + "(%.2f GiB quantized vs %.2f GiB dense bf16, %.2fx smaller).", + receipt["num_layers"], + receipt["num_tensors"], + serialized_bytes, + receipt["path"], + serialized_bytes / (1 << 30), + dense_bytes / (1 << 30), + receipt["compression_ratio"], + ) + return receipt + + +def read_int8_affine_sidecar_metadata(path: str | os.PathLike[str]) -> dict[str, Any]: + """Return the manifest of a sidecar file without materializing its tensors.""" + from safetensors import safe_open + + with safe_open(os.fspath(path), framework="pt", device="cpu") as handle: + raw = handle.metadata() or {} + if _INT8_AFFINE_SIDECAR_METADATA_KEY not in raw: + raise ValueError(f"{os.fspath(path)} is not a FastVideo INT8 affine sidecar " + f"(no {_INT8_AFFINE_SIDECAR_METADATA_KEY!r} metadata).") + return json.loads(raw[_INT8_AFFINE_SIDECAR_METADATA_KEY]) + + +def int8_affine_sidecar_path_for(checkpoint_path: str | os.PathLike[str]) -> str: + """Conventional sidecar path for a transformer checkpoint or directory. + + ``.../transformer.safetensors`` -> ``.../transformer.int8affine.safetensors``; + a directory -> ``/int8_affine.safetensors``. + """ + raw = os.fspath(checkpoint_path) + if os.path.isdir(raw): + return os.path.join(raw, INT8_AFFINE_DIR_SIDECAR_NAME) + if raw.endswith(".safetensors"): + return raw[:-len(".safetensors")] + INT8_AFFINE_SIDECAR_SUFFIX + return raw + INT8_AFFINE_SIDECAR_SUFFIX + + +def _sidecar_target_device(mod: torch.nn.Module, name: str) -> torch.device | None: + """Device the restored buffer should live on. + + Mirrors ``_quantize_layer_weight``, which registers the buffers on the + (local) weight's device; falls back to an existing buffer, then to the + module's parameter device so a purge-then-restore still lands on GPU. + """ + weight = getattr(mod, "weight", None) + if weight is not None and weight.device.type != "meta": + return weight.device + existing = getattr(mod, name, None) + if existing is not None and existing.device.type != "meta": + return existing.device + for param in mod.parameters(recurse=False): + if param.device.type != "meta": + return param.device + return None + + +def _expected_sidecar_shapes(name: str, weight_shape: tuple[int, int], group_size: int) -> tuple[tuple[int, ...], ...]: + """The one shape a fresh conversion would produce for *name*. + + Unlike the NVFP4 sidecar there is no padded variant to accept: the + quantizer groups along the last axis and ``_group`` refuses a K that is + not divisible by ``group_size``, so an exact divisor is the only + legitimate layout. A padded scales tensor would not merely be unusual — + ``int8_affine_dequantize`` recovers the group width as + ``codes.shape[-1] // scales.shape[-1]``, so extra groups silently regroup + every code in the row. + """ + out_dim, in_dim = weight_shape + if name == "_int8_affine_codes": + return ((out_dim, in_dim), ) + if in_dim % group_size: + raise ValueError(f"Sidecar declares a ({out_dim}, {in_dim}) weight with group_size {group_size}, " + "which does not divide the input dim; this layout cannot be dequantized.") + return ((out_dim, in_dim // group_size), ) + + +def load_int8_affine_checkpoint( + model: torch.nn.Module, + path: str | os.PathLike[str], + *, + strict: bool = True, +) -> int: + """Restore INT8 affine tensors from a sidecar, skipping ``convert_model_to_int8_affine``. + + Registers ``_int8_affine_codes`` / ``_int8_affine_scales`` / + ``_int8_affine_biases`` on every INT8Affine-tagged linear from the + sidecar, byte-for-byte as the conversion would have produced them. The + dense bf16 weights are never touched (they may be absent entirely), and + nothing here needs a GPU or a fused INT8 kernel — the dequantize-then-GEMM + reference path in :meth:`INT8AffineQuantizeMethod.apply` is pure PyTorch, + so a pre-quantized checkpoint loads on any host. + + ``strict`` raises on any layer-set mismatch (a sidecar that does not + describe this model); with ``strict=False`` those are logged and skipped, + leaving those layers unconverted. Scheme (``group_size``/``bits``) and + per-tensor shape/dtype mismatches are **never** downgraded: a mis-read + code buffer produces garbage output with no error, so those always raise. + + Returns the number of layers restored. + """ + from safetensors import safe_open + + tagged = _int8_affine_tagged_modules(model) + if not tagged: + raise RuntimeError("No INT8 affine linear layers are attached to this model, so a sidecar cannot be " + "restored. This is the silent-dense failure mode: the model's INT8AffineConfig " + "layer selection does not cover its layer paths (for MiniMax-H3 use " + "INT8AffineConfig.for_minimax_h3()).") + + manifest = read_int8_affine_sidecar_metadata(path) + if manifest.get("format") != _INT8_AFFINE_SIDECAR_FORMAT: + raise ValueError(f"Unsupported INT8 affine sidecar format {manifest.get('format')!r} in " + f"{os.fspath(path)}.") + if int(manifest.get("version", -1)) != _INT8_AFFINE_SIDECAR_VERSION: + raise ValueError(f"Unsupported INT8 affine sidecar version {manifest.get('version')!r} in " + f"{os.fspath(path)} (this build reads version {_INT8_AFFINE_SIDECAR_VERSION}).") + + saved_layers: dict[str, dict[str, Any]] = manifest.get("layers", {}) + model_fqns = {fqn for fqn, _, _ in tagged} + missing = sorted(model_fqns - set(saved_layers)) + extra = sorted(set(saved_layers) - model_fqns) + if missing or extra: + message = (f"INT8 affine sidecar {os.fspath(path)} does not match this model: " + f"{len(missing)} layers missing from the sidecar, {len(extra)} layers not in the model. " + f"First missing={missing[:3]}, first extra={extra[:3]}.") + if strict: + raise ValueError(message) + logger.warning("%s Restoring the intersection only.", message) + + restored = 0 + with safe_open(os.fspath(path), framework="pt", device="cpu") as handle: + available = set(handle.keys()) + for fqn, mod, qm in tagged: + if fqn not in saved_layers: + continue + entry = saved_layers[fqn] + try: + weight_shape = tuple(int(dim) for dim in entry["weight_shape"]) + group_size = int(entry["group_size"]) + bits = int(entry["bits"]) + except (KeyError, TypeError, ValueError) as exc: + raise ValueError(f"INT8 affine sidecar {os.fspath(path)} entry for {fqn!r} is malformed: " + f"{entry!r} does not carry an integer weight_shape/group_size/bits.") from exc + if len(weight_shape) != 2: + raise ValueError(f"INT8 affine sidecar entry {fqn!r} declares weight shape {list(weight_shape)}; " + "a linear weight is 2-D.") + if group_size != qm.group_size or bits != qm.bits: + raise ValueError(f"INT8 affine sidecar {os.fspath(path)} was written for {fqn!r} with " + f"group_size={group_size}, bits={bits}, but this model quantizes it with " + f"group_size={qm.group_size}, bits={qm.bits}.") + weight = getattr(mod, "weight", None) + if weight is not None and tuple(int(dim) for dim in weight.shape) != weight_shape: + raise ValueError(f"INT8 affine sidecar entry {fqn!r} describes a {list(weight_shape)} weight, but " + f"this model's layer has shape {list(weight.shape)}.") + tensors: dict[str, torch.Tensor] = {} + for name in _INT8_AFFINE_SIDECAR_BUFFERS: + key = _sidecar_key(fqn, name) + if key not in available: + continue + tensor = handle.get_tensor(key) + expected_dtype = _INT8_AFFINE_SIDECAR_DTYPES[name] + if tensor.dtype != expected_dtype: + raise ValueError(f"INT8 affine sidecar tensor {key} has dtype {tensor.dtype}, expected " + f"{expected_dtype}. Codes are uint8 because a bits=8 code spans [0, 255] " + "and does not fit int8; casting would silently corrupt them.") + expected = _expected_sidecar_shapes(name, weight_shape, group_size) + if tuple(tensor.shape) not in expected: + raise ValueError(f"INT8 affine sidecar tensor {key} has shape {tuple(tensor.shape)}, expected " + f"one of {list(expected)} for a {list(weight_shape)} linear with " + f"group_size={group_size}.") + device = _sidecar_target_device(mod, name) + if device is not None: + tensor = tensor.to(device=device, non_blocking=True) + tensors[name] = tensor + if set(tensors) != set(_INT8_AFFINE_SIDECAR_BUFFERS): + message = (f"INT8 affine sidecar entry for {fqn!r} is incomplete (has {sorted(tensors)}); " + f"all of {list(_INT8_AFFINE_SIDECAR_BUFFERS)} are required.") + if strict: + raise ValueError(message) + logger.warning("%s Skipping this layer.", message) + continue + for name, tensor in tensors.items(): + mod.register_buffer(name, tensor, persistent=False) + restored += 1 + + logger.info("INT8Affine sidecar: restored %d quantized modules from %s (dense weights untouched).", restored, + os.fspath(path)) + return restored + + +__all__ = [ + "DEFAULT_BITS", + "DEFAULT_GROUP_SIZE", + "INT8AffineConfig", + "INT8AffineQuantizeMethod", + "INT8_AFFINE_DIR_SIDECAR_NAME", + "INT8_AFFINE_SIDECAR_SUFFIX", + "MINIMAX_H3_BLOCK_SCOPES", + "MINIMAX_H3_INT8_AFFINE_EXCLUSIONS", + "MINIMAX_H3_INT8_AFFINE_SUFFIXES", + "MINIMAX_H3_PREFIX", + "convert_model_to_int8_affine", + "int8_affine_dequantize", + "int8_affine_quantize", + "int8_affine_sidecar_path_for", + "int8_affine_sidecar_state_dict", + "load_int8_affine_checkpoint", + "minimax_h3_int8_affine_prefixes", + "read_int8_affine_sidecar_metadata", + "save_int8_affine_checkpoint", +] diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 284622d595..fa621ce930 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -1,25 +1,10 @@ # SPDX-License-Identifier: Apache-2.0 -"""NVFP4 linear quantization for supported transformer layer sets. - -NVFP4 is NVIDIA's block-scaled FP4 format (e2m1 mantissa, fp32 alpha, -``layout_128x4`` scale layout, group size 16) — distinct from -generic FP4 / OCP-FP4 / MX-FP4. We name the public surface ``NVFP4`` -explicitly so downstream callers don't conflate it with other FP4 -variants that may land later (e.g. AMD's MX-FP4 or vendor-neutral -e3m0). - -The registered config targets the curated LTX-2 deployment set and the -main MiniMax-H3 transformer-block FFN linears. - -`flashinfer` is imported lazily inside the call paths that need it. -This keeps ``import fastvideo`` cheap on hosts where flashinfer is -not installed; only the actual NVFP4 quantize / matmul ops fail at -use time, with a clear error. -""" +"""NVFP4 quantization (FlashInfer-backed) for LTX-2 and MiniMax-H3.""" from __future__ import annotations +import json import logging -import re +import os from typing import Any import torch @@ -75,17 +60,44 @@ def _require_flashinfer() -> tuple[Any, Any, Any]: _LTX2_NVFP4_LINEAR_PREFIXES = frozenset(f"ltx2.blocks.{block_idx}.{suffix}" for block_idx in range(48) for suffix in _LTX2_NVFP4_BLOCK_LINEAR_SUFFIXES) | frozenset( ("ltx2.adaln_single.linear", )) -_MINIMAX_H3_NVFP4_FF_PREFIX = re.compile(r"(?:^|\.)transformer_blocks\.\d+\.ff\.(?:fc_in|fc_out)$") +MINIMAX_H3_DIT_PREFIX = "minimax_h3" +MINIMAX_H3_NUM_LAYERS = 50 +MINIMAX_H3_BLOCK_LINEAR_SUFFIXES = ( + "attn.to_q", + "attn.to_k", + "attn.to_v", + "attn.to_out", + "ff.fc_in", + "ff.fc_out", +) +MINIMAX_H3_NVFP4_LINEAR_PREFIXES = frozenset(f"{MINIMAX_H3_DIT_PREFIX}.transformer_blocks.{block_idx}.{suffix}" + for block_idx in range(MINIMAX_H3_NUM_LAYERS) + for suffix in MINIMAX_H3_BLOCK_LINEAR_SUFFIXES) -def is_ltx2_nvfp4_linear_prefix(prefix: str) -> bool: - """Return whether *prefix* belongs to the LTX-2 NVFP4 deployment set.""" - return prefix in _LTX2_NVFP4_LINEAR_PREFIXES +_ALWAYS_EXCLUDED_LINEAR_SUFFIXES = ("attn.to_gate_compress", ) +MINIMAX_H3_NVFP4_EXCLUDED_LINEAR_SUFFIXES = _ALWAYS_EXCLUDED_LINEAR_SUFFIXES + + +def _matches_linear_suffix(prefix: str, suffixes: frozenset[str] | tuple[str, ...]) -> bool: + """True when *prefix* is one of *suffixes*, or ends at a dot boundary. + + Entries may be a full module path or a trailing suffix + (``"attn.to_gate_compress"``), so one pattern covers every block of a + stack. The dot boundary keeps ``"ff.fc_in"`` from matching a hypothetical + ``"cross_ff.fc_in"``. + """ + return any(prefix == suffix or prefix.endswith("." + suffix) for suffix in suffixes) -def is_minimax_h3_nvfp4_linear_prefix(prefix: str) -> bool: - """Return whether *prefix* is a main MiniMax-H3 transformer-block FFN linear.""" - return _MINIMAX_H3_NVFP4_FF_PREFIX.search(prefix) is not None +def is_ltx2_nvfp4_linear_prefix(prefix: str) -> bool: + """Return whether *prefix* belongs to the LTX-2 NVFP4 deployment set. + + Kept as a module-level predicate for callers that predate the + ``NVFP4Config.layer_prefixes`` field; per-instance configs should use + :meth:`NVFP4Config.is_nvfp4_linear_prefix` instead. + """ + return prefix in _LTX2_NVFP4_LINEAR_PREFIXES def _is_ltx2_refine_only_prefix(prefix: str) -> bool: @@ -242,9 +254,6 @@ def _nvfp4_quantize( x_for_quant = x logical_rows = x.shape[0] else: - # Sequence-parallel can feed either logical rows or row-padded - # rows. Normalize to the kernel tile shape for swizzled layouts - # so both paths share a stable quantization contract. row_tile = 8 if sf_layout == SfLayout.layout_8x4.value else 128 logical_rows = x.shape[0] pad_rows = (-logical_rows) % row_tile @@ -315,10 +324,6 @@ def __init__(self, layer_prefix: str = ""): self.x_global_sf = torch.tensor(1.0, device="cuda", dtype=torch.float32) self.layer_prefix = layer_prefix self._is_refine_only_layer = _is_ltx2_refine_only_prefix(layer_prefix) - # Set from NVFP4Config.retain_original_weights in get_quant_method: - # True = retain every original bf16 weight; None/False (default) = - # purge the purgeable set. Refine-only layers are always retained -- - # the base stage profile runs them dense by deployment contract. self._retain_original_weights: bool | None = None def create_weights(self, layer: torch.nn.Module, input_size_per_partition: int, output_partition_sizes: list[int], @@ -360,16 +365,10 @@ def apply( | None = None, ) -> torch.Tensor: SfLayout, _, _ = _require_flashinfer() - # The original bf16 weight may have been purged after FP4 conversion - # (see convert_model_to_nvfp4); the packed FP4 weight keeps the - # output dim as its first dimension (only K is packed 2-per-byte). weight = getattr(layer, "weight", None) out_dim = weight.shape[0] if weight is not None else layer._nvfp4_weight.shape[0] original_shape = x.shape - # Stage-aware profile: keep refine-only FP4 layers in dense mode - # during stage-1 denoising so the base path doesn't pay the - # quantize/dequantize tax for layers it never touches. stage_profile = _get_ltx2_fp4_stage_profile(default="refine") if self._is_refine_only_layer and stage_profile == "base": if weight is None: @@ -382,9 +381,6 @@ def apply( return out.view(*original_shape[:-1], out_dim) if pre_quantized is not None: x_fp4, x_scale, x_global_sf = pre_quantized - # FlashInfer fused norm+quant APIs may return 3D tensors for - # 3D inputs. mm_fp4 only accepts 2D tensors, so flatten - # batch/sequence dims here. if x_fp4.dim() > 2: x_fp4 = x_fp4.view(-1, x_fp4.shape[-1]) if x_scale.dim() > 2: @@ -427,27 +423,46 @@ def apply( class NVFP4Config(QuantizationConfig): - """Select NVFP4 for the supported LTX-2 and MiniMax-H3 linear sets. + """NVFP4 quantization configuration, parameterized by layer paths. NVFP4 is NVIDIA's block-scaled FP4 (e2m1 mantissa, fp32 alpha, - ``layout_128x4`` scale layout, group size 16). LTX-2 uses its curated - attention and FFN deployment set. MiniMax-H3 uses only ``fc_in`` and - ``fc_out`` in each main transformer-block FFN. + ``layout_128x4`` scale layout, group size 16). + + Which linears get quantized is set by ``layer_prefixes``. The default is + the historical LTX-2 set, so ``NVFP4Config()`` behaves exactly as it did + before the field existed. Other models must pass their own set (see + :meth:`for_minimax_h3`) — a model whose layer paths are not covered + attaches no quant methods at all and silently runs dense, which is the + failure mode this field exists to remove. """ - def __init__(self, layer_profile: str = "refine", retain_original_weights: bool | None = None): + def __init__( + self, + layer_profile: str = "refine", + retain_original_weights: bool | None = None, + layer_prefixes: frozenset[str] | set[str] | list[str] | None = None, + exclude_prefixes: frozenset[str] | set[str] | list[str] | None = None, + ): super().__init__() - # ``base``: stage-1 set (no attn2.to_out, no cross-modal AV - # projections). ``refine``: full stage-2 set. self.layer_profile = layer_profile - # Original bf16 ``layer.weight`` retention after FP4 conversion. - # Default (None/False): purge the purgeable originals -- every - # always-FP4 layer. Refine-only layers (the cross-modal AV - # projections) are ALWAYS retained: the ``base`` stage profile runs - # them dense by deployment contract, including the distilled - # single-stage deploy. True: retain everything (debugging / - # pre-purge behavior). self.retain_original_weights = retain_original_weights + self.layer_prefixes: frozenset[str] = (frozenset(_LTX2_NVFP4_LINEAR_PREFIXES) + if layer_prefixes is None else frozenset(layer_prefixes)) + self.exclude_prefixes: frozenset[str] = frozenset(exclude_prefixes) if exclude_prefixes else frozenset() + + def is_nvfp4_linear_prefix(self, prefix: str) -> bool: + """Whether *prefix* is quantized under this config. + + Exclusions are checked first, and ``_ALWAYS_EXCLUDED_LINEAR_SUFFIXES`` + is unconditional: a prefix listed in ``layer_prefixes`` by mistake + (e.g. a glob that swept up ``attn.to_gate_compress``) still comes back + False here, so the gate cannot be quantized by any caller. + """ + if _matches_linear_suffix(prefix, _ALWAYS_EXCLUDED_LINEAR_SUFFIXES): + return False + if _matches_linear_suffix(prefix, self.exclude_prefixes): + return False + return prefix in self.layer_prefixes def get_name(self): return "nvfp4" @@ -463,20 +478,31 @@ def get_min_capability(cls): def get_config_filenames(): return [] + @classmethod + def for_minimax_h3(cls, **kwargs: Any) -> NVFP4Config: + """The MiniMax-H3 NVFP4 layer set (300 block linears). + + ``attn.to_gate_compress`` (the VSA sparse-attention gate) is excluded; + see ``MINIMAX_H3_NVFP4_EXCLUDED_LINEAR_SUFFIXES``. Keyword arguments + (e.g. ``retain_original_weights``) are forwarded to the constructor. + """ + kwargs.setdefault("layer_prefixes", MINIMAX_H3_NVFP4_LINEAR_PREFIXES) + kwargs.setdefault("exclude_prefixes", MINIMAX_H3_NVFP4_EXCLUDED_LINEAR_SUFFIXES) + return cls(**kwargs) + @classmethod def from_config(cls, config: dict[str, Any]) -> NVFP4Config: return cls( layer_profile=config.get("layer_profile", "refine"), retain_original_weights=config.get("retain_original_weights"), + layer_prefixes=config.get("layer_prefixes"), + exclude_prefixes=config.get("exclude_prefixes"), ) def get_quant_method(self, layer: torch.nn.Module, prefix: str): from fastvideo.layers.linear import LinearBase - # LTX-2 switches its active subset by stage at runtime. MiniMax-H3 - # uses the fixed main-transformer FFN set selected by its prefix. - if isinstance(layer, LinearBase) and (is_ltx2_nvfp4_linear_prefix(prefix) - or is_minimax_h3_nvfp4_linear_prefix(prefix)): + if isinstance(layer, LinearBase) and self.is_nvfp4_linear_prefix(prefix): method = NVFP4QuantizeMethod(layer_prefix=prefix) method._retain_original_weights = self.retain_original_weights return method @@ -487,9 +513,6 @@ def convert_model_to_nvfp4(model: torch.nn.Module) -> None: SfLayout, _, _ = _require_flashinfer() from torch.distributed.tensor import DTensor # type: ignore - purged = 0 - retained = 0 - purged_bytes = 0 for mod in model.modules(): qm = getattr(mod, "quant_method", None) if isinstance(qm, NVFP4QuantizeMethod): @@ -522,25 +545,45 @@ def convert_model_to_nvfp4(model: torch.nn.Module) -> None: persistent=False, ) - retain_flag = getattr(qm, "_retain_original_weights", None) - # Refine-only layers are NEVER purgeable: the "base" stage profile - # runs them dense by deployment contract (the distilled - # single-stage deploy included — its forward context is the base - # profile, so e.g. audio_to_video_attn routes dense every step). - # retain_original_weights therefore only widens retention - # (True = keep everything); it cannot narrow it below the - # dense-capable set. - retain = qm._is_refine_only_layer or retain_flag is True - if retain: - retained += 1 - elif isinstance(weight, DTensor): - # ponytail: purging FSDP-sharded originals needs per-shard - # resharding bookkeeping; skip until a sharded deploy needs it. - retained += 1 - else: - purged_bytes += weight.numel() * weight.element_size() - purged += 1 - mod.register_parameter("weight", None) + _apply_dense_weight_policy(model) + + +def _apply_dense_weight_policy(model: torch.nn.Module) -> None: + """Drop the bf16 ``weight`` of every NVFP4 linear that no longer needs it. + + Refine-only layers are NEVER purgeable: the "base" stage profile runs them + dense by deployment contract (the distilled single-stage deploy included — + its forward context is the base profile, so e.g. ``audio_to_video_attn`` + routes dense every step). ``retain_original_weights`` therefore only widens + retention (True = keep everything); it cannot narrow it below the + dense-capable set. + + Shared by :func:`convert_model_to_nvfp4` and + :func:`load_nvfp4_checkpoint` so a restored sidecar frees the same memory a + fresh conversion would. + """ + from torch.distributed.tensor import DTensor # type: ignore + + purged = 0 + retained = 0 + purged_bytes = 0 + for mod in model.modules(): + qm = getattr(mod, "quant_method", None) + if not isinstance(qm, NVFP4QuantizeMethod): + continue + weight = getattr(mod, "weight", None) + if weight is None: + continue + retain_flag = getattr(qm, "_retain_original_weights", None) + retain = getattr(qm, "_is_refine_only_layer", False) or retain_flag is True + if retain: + retained += 1 + elif isinstance(weight, DTensor): + retained += 1 + else: + purged_bytes += weight.numel() * weight.element_size() + purged += 1 + mod.register_parameter("weight", None) if purged or retained: logger.info( @@ -553,10 +596,328 @@ def convert_model_to_nvfp4(model: torch.nn.Module) -> None: ) + +NVFP4_SIDECAR_SUFFIX = ".nvfp4.safetensors" +NVFP4_DIR_SIDECAR_NAME = "nvfp4.safetensors" +_NVFP4_SIDECAR_FORMAT = "fastvideo.nvfp4" +_NVFP4_SIDECAR_VERSION = 1 +_NVFP4_SIDECAR_METADATA_KEY = "fastvideo_nvfp4" +_NVFP4_SIDECAR_KEY_SEP = "::" +_NVFP4_SIDECAR_SF_LAYOUT = "layout_128x4" +_NVFP4_SIDECAR_DO_SHUFFLE = False +_NVFP4_SIDECAR_BLOCK_SIZE = 16 +_NVFP4_SIDECAR_BUFFERS = ( + "_nvfp4_weight", + "_nvfp4_weight_scale", + "_weight_global_sf", + "_nvfp4_alpha", +) + + +def _sidecar_key(module_fqn: str, buffer_name: str) -> str: + return f"{module_fqn}{_NVFP4_SIDECAR_KEY_SEP}{buffer_name}" + + +def _split_sidecar_key(key: str) -> tuple[str, str]: + module_fqn, _, buffer_name = key.rpartition(_NVFP4_SIDECAR_KEY_SEP) + return module_fqn, buffer_name + + +def _nvfp4_tagged_modules(model: torch.nn.Module) -> list[tuple[str, torch.nn.Module, NVFP4QuantizeMethod]]: + tagged = [] + for fqn, mod in model.named_modules(): + qm = getattr(mod, "quant_method", None) + if isinstance(qm, NVFP4QuantizeMethod): + tagged.append((fqn, mod, qm)) + return tagged + + +def _is_dtensor(tensor: torch.Tensor) -> bool: + try: + from torch.distributed.tensor import DTensor # type: ignore + except ImportError: # pragma: no cover - depends on the torch build + return False + return isinstance(tensor, DTensor) + + +def nvfp4_sidecar_state_dict(model: torch.nn.Module) -> dict[str, torch.Tensor]: + """Collect the quantized tensors of every NVFP4 linear in *model*. + + Keys are ``"::"`` and values are detached CPU + copies. Modules whose buffers are missing (never converted) are skipped; + the returned mapping is what :func:`save_nvfp4_checkpoint` writes. + + FSDP note: a DTensor buffer is saved as this rank's local shard, so a + sharded save is only reloadable into an identically sharded model. + """ + state: dict[str, torch.Tensor] = {} + for fqn, mod, _ in _nvfp4_tagged_modules(model): + for name in _NVFP4_SIDECAR_BUFFERS: + tensor = getattr(mod, name, None) + if tensor is None: + continue + if _is_dtensor(tensor): + tensor = tensor.to_local() # type: ignore[attr-defined] + state[_sidecar_key(fqn, name)] = tensor.detach().to("cpu", copy=True).contiguous() + return state + + +def save_nvfp4_checkpoint( + model: torch.nn.Module, + path: str | os.PathLike[str], + *, + extra_metadata: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Write the model's NVFP4 tensors to a compact sidecar safetensors file. + + The file is roughly the packed-FP4 size (0.5 byte/code plus 1 byte per 16 + codes) instead of the dense bf16 size — about 3.6x smaller for a + power-of-two K. Returns a receipt dict (also logged) describing the layer + count and both sizes. Raises ``RuntimeError`` when the model has no NVFP4 + linears, which usually means the config's ``layer_prefixes`` did not cover + the model's layer paths. + """ + from safetensors.torch import save_file + + state = nvfp4_sidecar_state_dict(model) + tagged = _nvfp4_tagged_modules(model) + if not tagged: + raise RuntimeError("No NVFP4 linear layers found in this model; nothing to serialize. " + "Check that the model was built with an NVFP4Config whose " + "layer_prefixes cover its layer paths (e.g. NVFP4Config.for_minimax_h3()).") + if not state: + raise RuntimeError(f"Found {len(tagged)} NVFP4-tagged linear layers but none carry quantized " + "buffers. Call convert_model_to_nvfp4(model) before saving a sidecar.") + + layers: dict[str, list[int]] = {} + quant_prefixes: dict[str, str] = {} + dense_bytes = 0 + for fqn, mod, qm in tagged: + packed = getattr(mod, "_nvfp4_weight", None) + weight = getattr(mod, "weight", None) + if packed is None and weight is None: + continue + if weight is not None: + out_dim, in_dim = int(weight.shape[0]), int(weight.shape[1]) + else: + out_dim, in_dim = int(packed.shape[0]), int(packed.shape[1]) * 2 + layers[fqn] = [out_dim, in_dim] + quant_prefixes[fqn] = getattr(qm, "layer_prefix", "") or "" + dense_bytes += out_dim * in_dim * 2 + + metadata = { + "format": _NVFP4_SIDECAR_FORMAT, + "version": _NVFP4_SIDECAR_VERSION, + "sf_layout": _NVFP4_SIDECAR_SF_LAYOUT, + "do_shuffle": _NVFP4_SIDECAR_DO_SHUFFLE, + "block_size": _NVFP4_SIDECAR_BLOCK_SIZE, + "num_layers": len(layers), + "layers": layers, + "quant_prefixes": quant_prefixes, + "model_class": type(model).__name__, + } + if extra_metadata: + metadata.update(extra_metadata) + + payload = dict(state) + serialized_bytes = sum(t.numel() * t.element_size() for t in payload.values()) + save_file(payload, os.fspath(path), metadata={_NVFP4_SIDECAR_METADATA_KEY: json.dumps(metadata)}) + + receipt = { + "path": os.fspath(path), + "num_layers": len(layers), + "num_tensors": len(payload), + "quantized_bytes": serialized_bytes, + "dense_bfloat16_bytes": dense_bytes, + "compression_ratio": (dense_bytes / serialized_bytes) if serialized_bytes else 0.0, + } + logger.info( + "NVFP4 sidecar: wrote %d layers / %d tensors to %s (%.2f GiB quantized vs " + "%.2f GiB dense bf16, %.2fx smaller).", + receipt["num_layers"], + receipt["num_tensors"], + receipt["path"], + serialized_bytes / (1 << 30), + dense_bytes / (1 << 30), + receipt["compression_ratio"], + ) + return receipt + + +def read_nvfp4_sidecar_metadata(path: str | os.PathLike[str]) -> dict[str, Any]: + """Return the manifest of a sidecar file without materializing its tensors.""" + from safetensors import safe_open + + with safe_open(os.fspath(path), framework="pt", device="cpu") as handle: + raw = handle.metadata() or {} + if _NVFP4_SIDECAR_METADATA_KEY not in raw: + raise ValueError(f"{os.fspath(path)} is not a FastVideo NVFP4 sidecar " + f"(no {_NVFP4_SIDECAR_METADATA_KEY!r} metadata).") + return json.loads(raw[_NVFP4_SIDECAR_METADATA_KEY]) + + +def nvfp4_sidecar_path_for(checkpoint_path: str | os.PathLike[str]) -> str: + """Conventional sidecar path for a transformer checkpoint or directory. + + ``.../transformer.safetensors`` -> ``.../transformer.nvfp4.safetensors``; + a directory -> ``/nvfp4.safetensors``. + """ + raw = os.fspath(checkpoint_path) + if os.path.isdir(raw): + return os.path.join(raw, NVFP4_DIR_SIDECAR_NAME) + if raw.endswith(".safetensors"): + return raw[:-len(".safetensors")] + NVFP4_SIDECAR_SUFFIX + return raw + NVFP4_SIDECAR_SUFFIX + + +def _sidecar_target_device(mod: torch.nn.Module, name: str) -> torch.device | None: + """Device the restored buffer should live on. + + Mirrors ``convert_model_to_nvfp4``, which registers the buffers on the + (local) weight's device; falls back to an existing buffer, then to the + module's parameter device so a purge-then-restore still lands on GPU. + """ + weight = getattr(mod, "weight", None) + if weight is not None and weight.device.type != "meta": + return weight.device + existing = getattr(mod, name, None) + if existing is not None and existing.device.type != "meta": + return existing.device + for param in mod.parameters(recurse=False): + if param.device.type != "meta": + return param.device + return None + + +def load_nvfp4_checkpoint( + model: torch.nn.Module, + path: str | os.PathLike[str], + *, + strict: bool = True, + purge_dense_weights: bool = True, +) -> int: + """Restore NVFP4 tensors from a sidecar, skipping ``convert_model_to_nvfp4``. + + Registers ``_nvfp4_weight`` / ``_nvfp4_weight_scale`` / + ``_weight_global_sf`` / ``_nvfp4_alpha`` on every NVFP4-tagged linear from + the sidecar, byte-for-byte as the conversion would have produced them. The + dense bf16 weights are never touched (they may be absent entirely), and no + flashinfer call is made — this works on a host that only needs to *serve* a + pre-quantized checkpoint. + + ``strict`` raises on any layer-set or shape mismatch (a sidecar that does + not describe this model); with ``strict=False`` the mismatches are logged + and skipped, leaving those layers unconverted. ``purge_dense_weights`` + applies the same retention policy as ``convert_model_to_nvfp4``. + + Returns the number of layers restored. + """ + from safetensors import safe_open + + tagged = _nvfp4_tagged_modules(model) + if not tagged: + raise RuntimeError("No NVFP4 linear layers are attached to this model, so a sidecar cannot be " + "restored. This is the silent-dense failure mode: the model's NVFP4Config " + "layer_prefixes do not cover its layer paths (for MiniMax-H3 use " + "NVFP4Config.for_minimax_h3()).") + + manifest = read_nvfp4_sidecar_metadata(path) + if manifest.get("format") != _NVFP4_SIDECAR_FORMAT: + raise ValueError(f"Unsupported NVFP4 sidecar format {manifest.get('format')!r} in {os.fspath(path)}.") + if int(manifest.get("version", -1)) != _NVFP4_SIDECAR_VERSION: + raise ValueError(f"Unsupported NVFP4 sidecar version {manifest.get('version')!r} in " + f"{os.fspath(path)} (this build reads version {_NVFP4_SIDECAR_VERSION}).") + expected_layout = { + "sf_layout": _NVFP4_SIDECAR_SF_LAYOUT, + "do_shuffle": _NVFP4_SIDECAR_DO_SHUFFLE, + "block_size": _NVFP4_SIDECAR_BLOCK_SIZE, + } + for key, expected in expected_layout.items(): + if manifest.get(key) != expected: + raise ValueError(f"NVFP4 sidecar {os.fspath(path)} was written with {key}=" + f"{manifest.get(key)!r}, but this build quantizes with {key}={expected!r}.") + + saved_layers: dict[str, list[int]] = manifest.get("layers", {}) + model_fqns = {fqn for fqn, _, _ in tagged} + missing = sorted(model_fqns - set(saved_layers)) + extra = sorted(set(saved_layers) - model_fqns) + if missing or extra: + message = (f"NVFP4 sidecar {os.fspath(path)} does not match this model: " + f"{len(missing)} layers missing from the sidecar, {len(extra)} layers not in the model. " + f"First missing={missing[:3]}, first extra={extra[:3]}.") + if strict: + raise ValueError(message) + logger.warning("%s Restoring the intersection only.", message) + + restored = 0 + with safe_open(os.fspath(path), framework="pt", device="cpu") as handle: + available = set(handle.keys()) + for fqn, mod, _ in tagged: + if fqn not in saved_layers: + continue + out_dim, in_dim = (int(value) for value in saved_layers[fqn]) + tensors: dict[str, torch.Tensor] = {} + for name in _NVFP4_SIDECAR_BUFFERS: + key = _sidecar_key(fqn, name) + if key not in available: + continue + tensor = handle.get_tensor(key) + expected = _expected_sidecar_shapes(name, out_dim, in_dim) + if tuple(tensor.shape) not in expected: + raise ValueError(f"NVFP4 sidecar tensor {key} has shape {tuple(tensor.shape)}, expected one of " + f"{list(expected)} for a ({out_dim}, {in_dim}) linear.") + device = _sidecar_target_device(mod, name) + if device is not None: + tensor = tensor.to(device=device, non_blocking=True) + tensors[name] = tensor + if "_nvfp4_weight" not in tensors or "_nvfp4_weight_scale" not in tensors: + message = (f"NVFP4 sidecar entry for {fqn!r} is incomplete (has " + f"{sorted(tensors)}); the packed weight and its block scales are both required.") + if strict: + raise ValueError(message) + logger.warning("%s Skipping this layer.", message) + continue + for name, tensor in tensors.items(): + mod.register_buffer(name, tensor, persistent=False) + restored += 1 + + if purge_dense_weights: + _apply_dense_weight_policy(model) + + logger.info("NVFP4 sidecar: restored %d quantized layers from %s (dense weights %s).", restored, os.fspath(path), + "purged per policy" if purge_dense_weights else "left in place") + return restored + + +def _expected_sidecar_shapes(name: str, out_dim: int, in_dim: int) -> tuple[tuple[int, ...], ...]: + """The shapes a fresh conversion could produce for *name*. + + ``_nvfp4_quantize`` narrows the packed weight back to the logical row count + but returns the block scales as the kernel emitted them, i.e. still padded + to the 128-row tile: a layer whose output dim is not a multiple of 128 gets + a scale tensor with more rows than the weight. Both are accepted so a + sidecar written by a real flashinfer conversion validates. + """ + if name == "_nvfp4_weight": + return ((out_dim, (in_dim + 1) // 2), ) + if name == "_nvfp4_weight_scale": + padded_rows = ((out_dim + 127) // 128) * 128 + shapes = {(out_dim, (in_dim + 15) // 16), (padded_rows, (in_dim + 15) // 16)} + return tuple(sorted(shapes)) + return ((), ) + + __all__ = [ + "MINIMAX_H3_NVFP4_EXCLUDED_LINEAR_SUFFIXES", + "MINIMAX_H3_NVFP4_LINEAR_PREFIXES", "NVFP4Config", "NVFP4QuantizeMethod", + "NVFP4_SIDECAR_SUFFIX", "convert_model_to_nvfp4", "is_ltx2_nvfp4_linear_prefix", - "is_minimax_h3_nvfp4_linear_prefix", + "load_nvfp4_checkpoint", + "nvfp4_sidecar_path_for", + "nvfp4_sidecar_state_dict", + "read_nvfp4_sidecar_metadata", + "save_nvfp4_checkpoint", ] diff --git a/fastvideo/layers/quantization/nvfp4_qat_config.py b/fastvideo/layers/quantization/nvfp4_qat_config.py index bcd9fd04c7..4ea019040a 100644 --- a/fastvideo/layers/quantization/nvfp4_qat_config.py +++ b/fastvideo/layers/quantization/nvfp4_qat_config.py @@ -45,20 +45,15 @@ logger = logging.getLogger(__name__) -# Wan-style attention + FFN projection layers. Matched as substrings of the -# layer prefix (e.g. "blocks.0.attn1.to_q" contains "to_q"). Wan's "to_q"/ -# "to_k"/"to_v" also substring-match Kandinsky5's "to_query"/"to_key"/ -# "to_value", so only Kandinsky5's out-projection and FFN names (which don't -# share a substring with Wan's "to_out"/"ffn.fc_in"/"ffn.fc_out") need to be -# listed explicitly below. DEFAULT_FP4_LAYERS = ( + "ff.fc_in", + "ff.fc_out", "ffn.fc_in", "ffn.fc_out", "to_q", "to_k", "to_v", "to_out", - # Kandinsky5 "self_attention.out_layer", "cross_attention.out_layer", "feed_forward.mlp.fc_in", @@ -85,11 +80,6 @@ def __init__(self) -> None: super().__init__() self.weight_fp4 = None self.weight_scale = None - # Static input global scale factor. Matches the FastVideo-Quantization - # production path; recomputing it per-call via a ``.max()`` reduction - # (the previous behavior) adds a sync point, costs a kernel launch, - # and produces a data-dependent value that prevents CUDA-graph - # capture under ``torch.compile(mode='reduce-overhead')``. self.x_global_sf = torch.tensor(1.0, device="cuda", dtype=torch.float32) def create_weights(self, layer: torch.nn.Module, input_size_per_partition: int, output_partition_sizes: list[int], @@ -105,8 +95,6 @@ def create_weights(self, layer: torch.nn.Module, input_size_per_partition: int, set_weight_attrs(weight, extra_weight_attrs) def apply(self, layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor | None = None) -> torch.Tensor: - # ``_fp4_weight`` carries the (out, in/2) packed fp4 weight, so its - # row count is the output dim even after the dense weight is popped. out_dim = layer._fp4_weight.shape[0] original_shape = x.shape @@ -206,7 +194,6 @@ def convert_model_to_fp4(model: torch.nn.Module) -> None: weight_local = weight.to_local() if isinstance(weight, DTensor) else weight # type: ignore[arg-type] - # Only the reduced scalar needs fp32; avoid a full fp32 copy. weight_absmax = (weight_local.detach().abs().nan_to_num().amax().to(dtype=torch.float32)) weight_global_sf = (448 * 6) / weight_absmax fp4_w, fp4_s = _nvfp4_quantize( @@ -223,8 +210,6 @@ def convert_model_to_fp4(model: torch.nn.Module) -> None: persistent=False, ) - # Drop the dense weight as soon as the fp4 buffers are installed - # so it cannot keep occupying GPU memory. removed_weight = mod._parameters.pop("weight", None) if removed_weight is not None: removed_weight.grad = None diff --git a/fastvideo/layers/quantization/w4a16_config.py b/fastvideo/layers/quantization/w4a16_config.py new file mode 100644 index 0000000000..317c5cc289 --- /dev/null +++ b/fastvideo/layers/quantization/w4a16_config.py @@ -0,0 +1,917 @@ +# SPDX-License-Identifier: Apache-2.0 +"""W4A16 weight-only quantization for CUDA DiT inference.""" + +from __future__ import annotations + +import json +import logging +import os +from collections.abc import Iterable +from typing import Any + +import torch +import torch.nn.functional as F +from torch.nn.parameter import Parameter + +from fastvideo.layers.quantization.base_config import ( + QuantizationConfig, + QuantizeMethodBase, +) +from fastvideo.models.utils import set_weight_attrs + +logger = logging.getLogger(__name__) + +DEFAULT_GROUP_SIZE = 64 +DEFAULT_BITS = 4 +_EPS = 1e-8 +_LAZY_CONVERSION_WARNED = False + + + +def _group(w: torch.Tensor, group_size: int) -> torch.Tensor: + """View the last dim as ``(num_groups, group_size)``. + + Grouping is along the last (input/contraction) axis, matching + ``int8_affine_config._group`` and MLX's affine quantizer. + """ + if w.shape[-1] % group_size != 0: + raise ValueError(f"Last dim {w.shape[-1]} is not divisible by group_size {group_size}.") + return w.reshape(*w.shape[:-1], w.shape[-1] // group_size, group_size) + + +def w4a16_quantize( + w: torch.Tensor, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Group-wise affine quantization of ``w`` to ``bits``-bit unsigned codes. + + Per group of ``group_size`` contiguous values along the last axis, this + solves ``w ~= (code - zero) * scale`` with a **min/max affine** fit: + ``scale = (max - min) / (2**bits - 1)`` and ``zero`` the code that + reproduces ``min`` exactly, so both endpoints of the group round-trip + exactly. Codes are ``rint``-rounded (round-half-to-even) and clamped to + ``[0, 2**bits - 1]``. + + The fit is done in fp32 regardless of ``w.dtype``, so a bf16 checkpoint + value (which converts to fp32 exactly) yields a higher-precision scale + store than solving in bf16 would. + + Returns ``(codes, scales, zeros)``: + + - ``codes`` — ``torch.uint8``, shape ``w.shape[:-1] + (K // 2,)``, **two + 4-bit codes packed per byte**. For ``bits=4`` only; ``bits=8`` returns + one code per byte at shape ``w.shape``. + - ``scales`` — ``torch.float32``, shape ``w.shape[:-1] + (K // group_size,)``. + - ``zeros`` — ``torch.float32``, same shape as ``scales``. + + Packing convention (ours — no kernel consumes it yet): the **low nibble is + the lower K index**, i.e. ``packed[..., j] = codes[..., 2j] | codes[..., 2j+1] << 4``. + A future kernel has to be written against this layout; it does not match + AWQ's interleaved layout, GPTQ's ``g_idx`` layout, or bitsandbytes' order. + """ + if bits != 4 and bits != 8: + raise ValueError(f"W4A16 stores codes as uint8; bits must be 4 or 8, got {bits}") + if group_size <= 0: + raise ValueError(f"group_size must be positive, got {group_size}") + + w32 = w.detach().float().nan_to_num() + max_code = float((1 << bits) - 1) + grouped = _group(w32, group_size) + + w_min = grouped.amin(dim=-1) + w_max = grouped.amax(dim=-1) + scales = (w_max - w_min).clamp_min(_EPS) / max_code + zeros = torch.round(-w_min / scales).clamp_(0.0, max_code) + + codes = torch.round(grouped / scales.unsqueeze(-1) + zeros.unsqueeze(-1)) + codes = codes.clamp_(0.0, max_code).to(torch.uint8).reshape(w32.shape) + + if bits == 8: + return codes, scales, zeros + return _pack_4bit(codes), scales, zeros + + +def _pack_4bit(codes: torch.Tensor) -> torch.Tensor: + """Pack ``[0, 15]`` codes two-per-byte along the last axis. + + Low nibble = lower K index. The last dim must be even, which every H3 + linear input dim is (5376, 7168, 14336, 2688 are all even). + """ + if codes.shape[-1] % 2: + raise ValueError(f"4-bit packing needs an even last dim, got {codes.shape[-1]}") + codes = codes.to(torch.uint8).reshape(*codes.shape[:-1], codes.shape[-1] // 2, 2) + low = codes[..., 0] + high = codes[..., 1] & 0x0F + return (low | (high << 4)).contiguous() + + +def _unpack_4bit(packed: torch.Tensor) -> torch.Tensor: + """Inverse of :func:`_pack_4bit`; returns uint8 codes at ``2 * packed.shape[-1]``.""" + low = packed & 0x0F + high = (packed >> 4) & 0x0F + return torch.stack((low, high), dim=-1).reshape(*packed.shape[:-1], packed.shape[-1] * 2) + + +def w4a16_dequantize( + codes: torch.Tensor, + scales: torch.Tensor, + zeros: torch.Tensor, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + out_shape: tuple[int, ...] | torch.Size | None = None, + out_dtype: torch.dtype | None = None, +) -> torch.Tensor: + """Reconstruct the dense weight from 4-bit codes, group scales and zeros. + + ``codes`` is the packed tensor :func:`w4a16_quantize` returned (or, for + ``bits=8``, the unpacked one). Pass ``out_shape`` to name the logical + weight shape; otherwise the reconstruction keeps the code layout's own + last dim (``2 * packed.shape[-1]`` for 4-bit). + + The arithmetic runs in fp32 and is cast to ``out_dtype`` at the end — + ``(code - zero) * scale`` in fp32 is the more accurate side of the split, + same choice ``int8_affine_config`` makes. + """ + if bits == 4: + codes = _unpack_4bit(codes) + if out_shape is not None and tuple(codes.shape) != tuple(out_shape): + codes = codes.reshape(out_shape) + grouped = _group(codes.float(), group_size) + dense = (grouped - zeros.unsqueeze(-1)) * scales.unsqueeze(-1) + dense = dense.reshape(codes.shape) + return dense if out_dtype is None else dense.to(out_dtype) + + +MINIMAX_H3_PREFIX = "minimax_h3" +MINIMAX_H3_NUM_LAYERS = 50 +MINIMAX_H3_NUM_REFINER_LAYERS = 2 +MINIMAX_H3_BLOCK_SCOPES: tuple[str, ...] = ( + "transformer_blocks", + "token_refiner.refiner_blocks", +) +MINIMAX_H3_BLOCK_LINEAR_SUFFIXES: tuple[str, ...] = ( + "attn.to_q", + "attn.to_k", + "attn.to_v", + "attn.to_out", + "ff.fc_in", + "ff.fc_out", +) +MINIMAX_H3_MAIN_STACK_LINEAR_SUFFIXES: tuple[str, ...] = MINIMAX_H3_BLOCK_LINEAR_SUFFIXES + ("adaln_proj.linear", ) + +_GENERIC_LINEAR_SUFFIXES: tuple[str, ...] = ( + "attn.to_q", + "attn.to_k", + "attn.to_v", + "attn.to_out", + "ff.fc_in", + "ff.fc_out", +) + +_NEVER_QUANTIZE_SUBSTRINGS: tuple[str, ...] = ("attn.to_gate_compress", ) + +_H3_FP32_KEPT_SUBSTRINGS: tuple[str, ...] = ( + "proj_in", + "audio_proj_in", + "proj_out", + "audio_proj_out", + "time_embedder", +) + + +def _matches_linear_suffix(prefix: str, suffixes: Iterable[str]) -> bool: + """True when *prefix* is one of *suffixes* or ends at a dot boundary. + + The dot boundary keeps ``"ff.fc_in"`` from matching a hypothetical + ``"cross_ff.fc_in"``. + """ + return any(prefix == suffix or prefix.endswith("." + suffix) for suffix in suffixes) + + +def minimax_h3_w4a16_prefixes( + *, + prefix: str = MINIMAX_H3_PREFIX, + num_layers: int = MINIMAX_H3_NUM_LAYERS, + num_refiner_layers: int = MINIMAX_H3_NUM_REFINER_LAYERS, +) -> frozenset[str]: + """Enumerate the exact H3 linear prefixes the H3 W4A16 profile targets. + + Built from H3's real module names as constructed in + ``fastvideo/models/dits/minimax_h3.py``: ``MiniMaxH3TransformerBlock`` + builds ``{prefix}.transformer_blocks.{i}.attn`` / ``.ff`` / ``.adaln_proj``; + ``MiniMaxH3TokenRefiner`` builds ``{prefix}.token_refiner.refiner_blocks.{i}.attn`` / ``.ff``. + + 50 main blocks x 7 linears + 2 refiner blocks x 6 linears = **362 linears**. + + This is the allowlist :meth:`W4A16Config.for_minimax_h3` hands to + ``target_layers``. It is a plain function so a caller (or a test) can + regenerate it from the architecture constants rather than trusting a + literal — the H3 profile is *derived*, not hardcoded. + """ + prefixes: set[str] = set() + for index in range(num_layers): + for suffix in MINIMAX_H3_MAIN_STACK_LINEAR_SUFFIXES: + prefixes.add(f"{prefix}.transformer_blocks.{index}.{suffix}") + for index in range(num_refiner_layers): + for suffix in MINIMAX_H3_BLOCK_LINEAR_SUFFIXES: + prefixes.add(f"{prefix}.token_refiner.refiner_blocks.{index}.{suffix}") + return frozenset(prefixes) + + +class W4A16Config(QuantizationConfig): + """Weight-only 4-bit (group-wise affine) quantization with 16-bit activations. + + Layer selection is a constructor field, not a hardcoded model list: + ``target_layers`` is an explicit allowlist of full module paths and takes + precedence when given; otherwise ``layer_suffixes`` is matched with a + dot-boundary suffix rule. Both are subject to the fail-closed deny list + (``_NEVER_QUANTIZE_SUBSTRINGS``), which the constructor can only widen. + + Weight-only by construction: there is no activation quantizer and + :class:`W4A16QuantizeMethod` runs a dense 16-bit GEMM over a dequantized + weight. Use :meth:`for_minimax_h3` for the H3 profile. + """ + + def __init__( + self, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + target_layers: Iterable[str] | None = None, + layer_suffixes: Iterable[str] | None = None, + exclude_substrings: Iterable[str] | None = None, + retain_original_weight: bool = True, + ) -> None: + super().__init__() + if bits not in (4, 8): + raise ValueError(f"W4A16 stores codes as uint8; bits must be 4 or 8, got {bits}") + if group_size <= 0: + raise ValueError(f"group_size must be positive, got {group_size}") + self.group_size = group_size + self.bits = bits + self.target_layers: frozenset[str] | None = (None if target_layers is None else frozenset(target_layers)) + self.layer_suffixes: tuple[str, ...] = (tuple(_GENERIC_LINEAR_SUFFIXES) + if layer_suffixes is None else tuple(layer_suffixes)) + self.exclude_substrings: tuple[str, ...] = tuple(_NEVER_QUANTIZE_SUBSTRINGS) + tuple(exclude_substrings or ()) + self.retain_original_weight = retain_original_weight + + def get_name(self) -> str: + return "W4A16" + + def get_supported_act_dtypes(self) -> list[torch.dtype]: + return [torch.bfloat16, torch.float16, torch.float32] + + @classmethod + def get_min_capability(cls) -> int: + """Turing (75). + + The compute path is a plain 16-bit GEMM over a dequantized weight, so + no 4-bit tensor-core class is required and the config stays loadable + wherever the other reference paths are. **This is not a claim that + 4-bit runs fast there.** The deployment target for this lane is Ada + (sm_89, RTX 4090 / RTX 6000 Ada); making it a *fast* path needs a + fused W4A16 GEMM that does not exist in this repository yet. + """ + return 75 + + @staticmethod + def get_config_filenames() -> list[str]: + return [] + + @classmethod + def from_config(cls, config: dict[str, Any]) -> W4A16Config: + return cls( + group_size=config.get("group_size", DEFAULT_GROUP_SIZE), + bits=config.get("bits", DEFAULT_BITS), + target_layers=config.get("target_layers"), + layer_suffixes=config.get("layer_suffixes"), + exclude_substrings=config.get("exclude_substrings"), + retain_original_weight=config.get("retain_original_weight", True), + ) + + @classmethod + def for_minimax_h3( + cls, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + retain_original_weight: bool = True, + ) -> W4A16Config: + """The MiniMax-H3 profile: 362 attention / FFN / AdaLN linears. + + The allowlist comes from :func:`minimax_h3_w4a16_prefixes`, and the + deny list additionally carries H3's fp32-pinned modules (``proj_in``, + ``audio_proj_in``, ``time_embedder``, ``proj_out``, ``audio_proj_out``) + — those are excluded both by construction (they are not in the + allowlist) and by name, so a later widening of the allowlist cannot + silently reach them. + + ``group_size`` must divide the input dim of every targeted linear: + with the released H3 config (hidden 5376, inner 7168, ffn 14336, + adaln 2688) that holds for 32, 64 and 128. + """ + return cls( + group_size=group_size, + bits=bits, + target_layers=minimax_h3_w4a16_prefixes(), + exclude_substrings=_H3_FP32_KEPT_SUBSTRINGS, + retain_original_weight=retain_original_weight, + ) + + def is_target_layer(self, prefix: str) -> bool: + """Whether ``prefix`` is quantized under this config. + + Deny list first (fail-closed), then ``target_layers`` if supplied, + else suffix matching. Non-``LinearBase`` layers are filtered by + :meth:`get_quant_method`, not here, so this is safe to call on any + module name. + """ + for banned in self.exclude_substrings: + if banned in prefix: + return False + if self.target_layers is not None: + return prefix in self.target_layers + return _matches_linear_suffix(prefix, self.layer_suffixes) + + def get_quant_method(self, layer: torch.nn.Module, prefix: str): + from fastvideo.layers.linear import LinearBase + + if not isinstance(layer, LinearBase) or not self.is_target_layer(prefix): + return None + input_size = getattr(layer, "input_size", None) + if input_size is not None and input_size % self.group_size: + logger.warning( + "W4A16: skipping layer %r — input dim %d is not divisible by group_size %d. " + "The layer runs dense.", prefix, input_size, self.group_size) + return None + if self.bits == 4 and input_size is not None and input_size % 2: + logger.warning( + "W4A16: skipping layer %r — input dim %d is odd, so 4-bit codes cannot be " + "packed two-per-byte. The layer runs dense.", prefix, input_size) + return None + return W4A16QuantizeMethod( + layer_prefix=prefix, + group_size=self.group_size, + bits=self.bits, + retain_original_weight=self.retain_original_weight, + ) + + +class W4A16QuantizeMethod(QuantizeMethodBase): + """Linear method for weight-only 4-bit affine quantization. + + ``create_weights`` allocates the same dense Parameter an unquantized linear + would (so a BF16 checkpoint loads unchanged), and the 4-bit codes, group + scales and zeros arrive later as non-persistent buffers from + :func:`convert_model_to_w4a16` — conversion happens at *load* time, not + construction time, exactly mirroring ``NVFP4QuantizeMethod`` and + ``INT8AffineQuantizeMethod``. + + ``apply`` is the **reference path**: dequantize the whole weight to the + activation dtype, then ``F.linear``. It is bit-for-bit a dense 16-bit GEMM + over an approximation of the original weight, which is what makes it + useful as a correctness oracle — and it is why it is not a performance + path (see the module docstring's "Kernel status"). + """ + + def __init__( + self, + layer_prefix: str = "", + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + retain_original_weight: bool = True, + ) -> None: + super().__init__() + self.layer_prefix = layer_prefix + self.group_size = group_size + self.bits = bits + self.retain_original_weight = retain_original_weight + + def create_weights( + self, + layer: torch.nn.Module, + input_size_per_partition: int, + output_partition_sizes: list[int], + input_size: int, + output_size: int, + params_dtype: torch.dtype, + **extra_weight_attrs, + ): + weight = Parameter( + torch.empty( + sum(output_partition_sizes), + input_size_per_partition, + dtype=params_dtype, + ), + requires_grad=False, + ) + set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0}) + layer.register_parameter("weight", weight) + set_weight_attrs(weight, extra_weight_attrs) + + def _ensure_quantized(self, layer: torch.nn.Module) -> bool: + """Convert on first use if the loader hook never ran. + + Returns False when the layer is intentionally left dense (grad-enabled + forward: a training step must see the master weight, not a frozen + dequantized copy). The loader path is ``_maybe_quantize_model`` -> + :func:`convert_model_to_w4a16`; this fallback exists so the config is + still *correct* if that dispatch is missing, but it warns because + reaching it means the loader hook did not fire. + """ + if getattr(layer, "_w4a16_codes", None) is not None: + return True + weight = getattr(layer, "weight", None) + if weight is None: + raise RuntimeError(f"W4A16 layer {self.layer_prefix!r} has no weight and no quantized buffers.") + if torch.is_grad_enabled(): + return False + global _LAZY_CONVERSION_WARNED + if not _LAZY_CONVERSION_WARNED: + _LAZY_CONVERSION_WARNED = True + logger.warning( + "W4A16: layer %r reached apply() unconverted; converting lazily (this message is logged once " + "per process, not once per layer). The loader hook (_maybe_quantize_model) did not dispatch to " + "convert_model_to_w4a16 — check its isinstance chain in " + "fastvideo/models/loader/fsdp_load.py.", self.layer_prefix) + _quantize_layer_weight(layer, weight, group_size=self.group_size, bits=self.bits) + return True + + def apply( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + if not self._ensure_quantized(layer): + weight = layer.weight + if weight is None: + raise RuntimeError( + f"W4A16 layer {self.layer_prefix!r} is in the dense (grad-enabled) branch, but its " + "original bf16 weight was purged (W4A16Config(retain_original_weight=False)). Training " + "needs the master weight; load with retain_original_weight left at its default for any " + "run that takes gradients.") + return F.linear(x, weight.to(x.dtype) if weight.dtype != x.dtype else weight, bias) + + weight = w4a16_dequantize( + layer._w4a16_codes, + layer._w4a16_scales, + layer._w4a16_zeros, + group_size=self.group_size, + bits=self.bits, + out_shape=layer._w4a16_weight_shape, + out_dtype=x.dtype, + ) + return F.linear(x, weight, bias) + + + + +def _quantize_layer_weight( + mod: torch.nn.Module, + weight: torch.Tensor, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, +) -> None: + """Quantize one linear's weight in place into non-persistent buffers.""" + from torch.distributed.tensor import DTensor # type: ignore + + weight_local = weight.to_local() if isinstance(weight, DTensor) else weight # type: ignore[arg-type] + if weight_local.shape[-1] % group_size: + raise ValueError(f"W4A16 layer {mod!r}: input dim {weight_local.shape[-1]} is not divisible by " + f"group_size {group_size}.") + codes, scales, zeros = w4a16_quantize(weight_local, group_size=group_size, bits=bits) + mod.register_buffer("_w4a16_codes", codes.contiguous(), persistent=False) + mod.register_buffer("_w4a16_scales", scales.to(torch.float32).contiguous(), persistent=False) + mod.register_buffer("_w4a16_zeros", zeros.to(torch.float32).contiguous(), persistent=False) + mod._w4a16_weight_shape = tuple(weight_local.shape) + + +def convert_model_to_w4a16(model: torch.nn.Module) -> None: + """Quantize every W4A16-tagged linear in-place after weights load. + + Mirrors ``convert_model_to_nvfp4`` / ``convert_model_to_int8_affine``: walk + the module tree once, convert each layer whose ``quant_method`` is a + :class:`W4A16QuantizeMethod`, and register the 4-bit codes plus per-group + scales/zeros as non-persistent buffers (so they are not written back into + ``state_dict``/checkpoints). + + Callers: the loader hook ``_maybe_quantize_model`` in + ``fastvideo/models/loader/fsdp_load.py``. *That hook is not edited by this + module* — it dispatches on an explicit ``isinstance`` chain, so it needs a + matching branch (see the module report). Without it, + :meth:`W4A16QuantizeMethod.apply` converts lazily on first forward and logs + a warning, so inference is still correct, just later and noisier. + """ + converted = 0 + purged = 0 + schemes: set[tuple[int, int]] = set() + for mod in model.modules(): + qm = getattr(mod, "quant_method", None) + if not isinstance(qm, W4A16QuantizeMethod): + continue + weight = getattr(mod, "weight", None) + if weight is None: + continue + _quantize_layer_weight(mod, weight, group_size=qm.group_size, bits=qm.bits) + converted += 1 + schemes.add((qm.group_size, qm.bits)) + if not qm.retain_original_weight: + original = mod._parameters.get("weight") + if original is not None: + original.grad = None + mod.register_parameter("weight", None) + purged += 1 + + if converted: + logger.info( + "W4A16 conversion receipt: quantized %d linear layers (%s, reference dequantize-then-GEMM " + "path); purged %d original bf16 weight tensors.", converted, + ", ".join(f"group_size={g}, bits={b}" for g, b in sorted(schemes)), purged) + logger.info("W4A16: no fused 4-bit GEMM is available in this build, so forward compute is a dense " + "16-bit GEMM over a dequantized weight. Expect BF16-comparable VRAM for the transient " + "dense weight and slower-than-BF16 step times.") + + + +W4A16_SIDECAR_SUFFIX = ".w4a16.safetensors" +W4A16_DIR_SIDECAR_NAME = "w4a16.safetensors" +_W4A16_SIDECAR_FORMAT = "fastvideo.w4a16" +_W4A16_SIDECAR_VERSION = 1 +_W4A16_SIDECAR_METADATA_KEY = "fastvideo_w4a16" +_W4A16_SIDECAR_KEY_SEP = "::" +_W4A16_SIDECAR_BUFFERS = ( + "_w4a16_codes", + "_w4a16_scales", + "_w4a16_zeros", +) +_W4A16_SIDECAR_DTYPES = { + "_w4a16_codes": torch.uint8, + "_w4a16_scales": torch.float32, + "_w4a16_zeros": torch.float32, +} + + +def _sidecar_key(module_fqn: str, buffer_name: str) -> str: + return f"{module_fqn}{_W4A16_SIDECAR_KEY_SEP}{buffer_name}" + + +def _w4a16_tagged_modules(model: torch.nn.Module) -> list[tuple[str, torch.nn.Module, W4A16QuantizeMethod]]: + tagged = [] + for fqn, mod in model.named_modules(): + qm = getattr(mod, "quant_method", None) + if isinstance(qm, W4A16QuantizeMethod): + tagged.append((fqn, mod, qm)) + return tagged + + +def _is_dtensor(tensor: torch.Tensor) -> bool: + try: + from torch.distributed.tensor import DTensor # type: ignore + except ImportError: # pragma: no cover - depends on the torch build + return False + return isinstance(tensor, DTensor) + + +def w4a16_sidecar_state_dict(model: torch.nn.Module) -> dict[str, torch.Tensor]: + """Collect the quantized tensors of every W4A16 linear in *model*. + + Keys are ``"::"`` and values are detached CPU + copies. Modules whose buffers are missing (never converted) are skipped; + the returned mapping is what :func:`save_w4a16_checkpoint` writes. + + ``_w4a16_weight_shape`` is not a tensor and so is not collected here; it + travels in the manifest instead (see :func:`save_w4a16_checkpoint`). + + FSDP note: a DTensor buffer is saved as this rank's local shard, so a + sharded save is only reloadable into an identically sharded model. + """ + state: dict[str, torch.Tensor] = {} + for fqn, mod, _ in _w4a16_tagged_modules(model): + for name in _W4A16_SIDECAR_BUFFERS: + tensor = getattr(mod, name, None) + if tensor is None: + continue + if _is_dtensor(tensor): + tensor = tensor.to_local() # type: ignore[attr-defined] + state[_sidecar_key(fqn, name)] = tensor.detach().to("cpu", copy=True).contiguous() + return state + + +def save_w4a16_checkpoint( + model: torch.nn.Module, + path: str | os.PathLike[str], + *, + extra_metadata: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Write the model's W4A16 tensors to a compact sidecar safetensors file. + + The file is roughly the 4-bit size (5.0 bits/weight at ``group_size=64``: + 4 for the packed codes plus two fp32 per group) instead of the dense bf16 + size. Returns a receipt dict (also logged) with the module count and both + sizes. Raises ``RuntimeError`` when the model has no W4A16 linears, which + usually means the config's layer selection did not cover the model's layer + paths, and when the tagged layers carry no buffers (never converted). + + The manifest under the ``fastvideo_w4a16`` metadata key carries the format + name/version, the scheme (``group_size``/``bits``), the per-layer weight and + buffer shapes, and the quantized module fqns (the keys of ``layers``), so a + loader can validate a sidecar against a model without materializing the + tensors. The per-layer ``weight_shape`` is what a load uses to restore + ``_w4a16_weight_shape``, which the packed codes cannot express. + """ + from safetensors.torch import save_file + + state = w4a16_sidecar_state_dict(model) + tagged = _w4a16_tagged_modules(model) + if not tagged: + raise RuntimeError("No W4A16 linear layers found in this model; nothing to serialize. Check that the " + "model was built with a W4A16Config whose layer selection covers its layer paths " + "(e.g. W4A16Config.for_minimax_h3()).") + if not state: + raise RuntimeError(f"Found {len(tagged)} W4A16-tagged linear layers but none carry quantized buffers. " + "Call convert_model_to_w4a16(model) before saving a sidecar.") + + layers: dict[str, dict[str, Any]] = {} + quant_prefixes: dict[str, str] = {} + dense_bytes = 0 + group_sizes: set[int] = set() + bit_widths: set[int] = set() + for fqn, mod, qm in tagged: + codes = getattr(mod, "_w4a16_codes", None) + weight = getattr(mod, "weight", None) + if codes is None and weight is None: + continue + if weight is not None: + weight_shape = [int(dim) for dim in weight.shape] + elif getattr(mod, "_w4a16_weight_shape", None) is not None: + weight_shape = [int(dim) for dim in mod._w4a16_weight_shape] + else: + raise RuntimeError(f"W4A16 layer {fqn!r} has no dense weight and no recorded " + "_w4a16_weight_shape; its logical weight shape cannot be recovered from the " + "packed codes. Re-run convert_model_to_w4a16(model) before saving.") + tensors = { + name: [int(dim) for dim in getattr(mod, name).shape] + for name in _W4A16_SIDECAR_BUFFERS if getattr(mod, name, None) is not None + } + layers[fqn] = { + "weight_shape": weight_shape, + "group_size": int(qm.group_size), + "bits": int(qm.bits), + "tensors": tensors, + } + quant_prefixes[fqn] = getattr(qm, "layer_prefix", "") or "" + dense_bytes += weight_shape[0] * weight_shape[1] * 2 + group_sizes.add(int(qm.group_size)) + bit_widths.add(int(qm.bits)) + + metadata: dict[str, Any] = { + "format": _W4A16_SIDECAR_FORMAT, + "version": _W4A16_SIDECAR_VERSION, + "group_size": group_sizes.pop() if len(group_sizes) == 1 else None, + "bits": bit_widths.pop() if len(bit_widths) == 1 else None, + "num_layers": len(layers), + "layers": layers, + "quant_prefixes": quant_prefixes, + "model_class": type(model).__name__, + } + if extra_metadata: + metadata.update(extra_metadata) + + payload = dict(state) + serialized_bytes = sum(t.numel() * t.element_size() for t in payload.values()) + save_file(payload, os.fspath(path), metadata={_W4A16_SIDECAR_METADATA_KEY: json.dumps(metadata)}) + + receipt = { + "path": os.fspath(path), + "num_layers": len(layers), + "num_tensors": len(payload), + "quantized_bytes": serialized_bytes, + "dense_bfloat16_bytes": dense_bytes, + "compression_ratio": (dense_bytes / serialized_bytes) if serialized_bytes else 0.0, + } + logger.info( + "W4A16 sidecar: wrote %d quantized modules / %d tensors (%d bytes) to %s " + "(%.2f GiB quantized vs %.2f GiB dense bf16, %.2fx smaller).", + receipt["num_layers"], + receipt["num_tensors"], + serialized_bytes, + receipt["path"], + serialized_bytes / (1 << 30), + dense_bytes / (1 << 30), + receipt["compression_ratio"], + ) + return receipt + + +def read_w4a16_sidecar_metadata(path: str | os.PathLike[str]) -> dict[str, Any]: + """Return the manifest of a sidecar file without materializing its tensors.""" + from safetensors import safe_open + + with safe_open(os.fspath(path), framework="pt", device="cpu") as handle: + raw = handle.metadata() or {} + if _W4A16_SIDECAR_METADATA_KEY not in raw: + raise ValueError(f"{os.fspath(path)} is not a FastVideo W4A16 sidecar " + f"(no {_W4A16_SIDECAR_METADATA_KEY!r} metadata).") + return json.loads(raw[_W4A16_SIDECAR_METADATA_KEY]) + + +def w4a16_sidecar_path_for(checkpoint_path: str | os.PathLike[str]) -> str: + """Conventional sidecar path for a transformer checkpoint or directory. + + ``.../transformer.safetensors`` -> ``.../transformer.w4a16.safetensors``; + a directory -> ``/w4a16.safetensors``. + """ + raw = os.fspath(checkpoint_path) + if os.path.isdir(raw): + return os.path.join(raw, W4A16_DIR_SIDECAR_NAME) + if raw.endswith(".safetensors"): + return raw[:-len(".safetensors")] + W4A16_SIDECAR_SUFFIX + return raw + W4A16_SIDECAR_SUFFIX + + +def _sidecar_target_device(mod: torch.nn.Module, name: str) -> torch.device | None: + """Device the restored buffer should live on. + + Mirrors ``_quantize_layer_weight``, which registers the buffers on the + (local) weight's device; falls back to an existing buffer, then to the + module's parameter device so a purge-then-restore still lands on GPU. + """ + weight = getattr(mod, "weight", None) + if weight is not None and weight.device.type != "meta": + return weight.device + existing = getattr(mod, name, None) + if existing is not None and existing.device.type != "meta": + return existing.device + for param in mod.parameters(recurse=False): + if param.device.type != "meta": + return param.device + return None + + +def _expected_sidecar_shapes(name: str, weight_shape: tuple[int, int], group_size: int, + bits: int) -> tuple[tuple[int, ...], ...]: + """The one shape a fresh conversion would produce for *name*. + + Unlike the NVFP4 sidecar there is no padded variant to accept: the + quantizer groups along the last axis and ``_group`` refuses a K that is + not divisible by ``group_size``, and the packed code interpretation is + fixed by ``bits``. A sidecar that disagrees would not merely be unusual — + reading it back would unpack the wrong nibble order or the wrong number of + groups, which is silent corruption rather than an error. + """ + out_dim, in_dim = weight_shape + if name == "_w4a16_codes": + if bits == 8: + return ((out_dim, in_dim), ) + if in_dim % 2: + raise ValueError(f"Sidecar declares a ({out_dim}, {in_dim}) weight at bits={bits}, but 4-bit codes " + "pack two per byte and need an even input dim.") + return ((out_dim, in_dim // 2), ) + if in_dim % group_size: + raise ValueError(f"Sidecar declares a ({out_dim}, {in_dim}) weight with group_size {group_size}, " + "which does not divide the input dim; this layout cannot be dequantized.") + return ((out_dim, in_dim // group_size), ) + + +def load_w4a16_checkpoint( + model: torch.nn.Module, + path: str | os.PathLike[str], + *, + strict: bool = True, +) -> int: + """Restore W4A16 tensors from a sidecar, skipping ``convert_model_to_w4a16``. + + Registers ``_w4a16_codes`` / ``_w4a16_scales`` / ``_w4a16_zeros`` on every + W4A16-tagged linear from the sidecar, byte-for-byte as the conversion would + have produced them, and restores the ``_w4a16_weight_shape`` attribute that + :meth:`W4A16QuantizeMethod.apply` needs (it is not a buffer, so nothing else + would). The dense bf16 weights are never touched (they may be absent + entirely), and nothing here needs a GPU or a 4-bit kernel — the + dequantize-then-GEMM reference path in ``apply`` is pure PyTorch, so a + pre-quantized checkpoint loads on any host. + + ``strict`` raises on any layer-set mismatch (a sidecar that does not + describe this model); with ``strict=False`` those are logged and skipped, + leaving those layers unconverted. Scheme (``group_size``/``bits``) and + per-tensor shape/dtype mismatches are **never** downgraded: mis-read or + mis-unpacked codes produce garbage output with no error, so those always + raise. + + Returns the number of layers restored. + """ + from safetensors import safe_open + + tagged = _w4a16_tagged_modules(model) + if not tagged: + raise RuntimeError("No W4A16 linear layers are attached to this model, so a sidecar cannot be restored. " + "This is the silent-dense failure mode: the model's W4A16Config layer selection " + "does not cover its layer paths (for MiniMax-H3 use " + "W4A16Config.for_minimax_h3()).") + + manifest = read_w4a16_sidecar_metadata(path) + if manifest.get("format") != _W4A16_SIDECAR_FORMAT: + raise ValueError(f"Unsupported W4A16 sidecar format {manifest.get('format')!r} in {os.fspath(path)}.") + if int(manifest.get("version", -1)) != _W4A16_SIDECAR_VERSION: + raise ValueError(f"Unsupported W4A16 sidecar version {manifest.get('version')!r} in {os.fspath(path)} " + f"(this build reads version {_W4A16_SIDECAR_VERSION}).") + + saved_layers: dict[str, dict[str, Any]] = manifest.get("layers", {}) + model_fqns = {fqn for fqn, _, _ in tagged} + missing = sorted(model_fqns - set(saved_layers)) + extra = sorted(set(saved_layers) - model_fqns) + if missing or extra: + message = (f"W4A16 sidecar {os.fspath(path)} does not match this model: {len(missing)} layers missing " + f"from the sidecar, {len(extra)} layers not in the model. First missing={missing[:3]}, " + f"first extra={extra[:3]}.") + if strict: + raise ValueError(message) + logger.warning("%s Restoring the intersection only.", message) + + restored = 0 + with safe_open(os.fspath(path), framework="pt", device="cpu") as handle: + available = set(handle.keys()) + for fqn, mod, qm in tagged: + if fqn not in saved_layers: + continue + entry = saved_layers[fqn] + try: + weight_shape = tuple(int(dim) for dim in entry["weight_shape"]) + group_size = int(entry["group_size"]) + bits = int(entry["bits"]) + except (KeyError, TypeError, ValueError) as exc: + raise ValueError(f"W4A16 sidecar {os.fspath(path)} entry for {fqn!r} is malformed: {entry!r} " + "does not carry an integer weight_shape/group_size/bits.") from exc + if len(weight_shape) != 2: + raise ValueError(f"W4A16 sidecar entry {fqn!r} declares weight shape {list(weight_shape)}; " + "a linear weight is 2-D.") + if group_size != qm.group_size or bits != qm.bits: + raise ValueError(f"W4A16 sidecar {os.fspath(path)} was written for {fqn!r} with " + f"group_size={group_size}, bits={bits}, but this model quantizes it with " + f"group_size={qm.group_size}, bits={qm.bits}.") + weight = getattr(mod, "weight", None) + if weight is not None and tuple(int(dim) for dim in weight.shape) != weight_shape: + raise ValueError(f"W4A16 sidecar entry {fqn!r} describes a {list(weight_shape)} weight, but " + f"this model's layer has shape {list(weight.shape)}.") + tensors: dict[str, torch.Tensor] = {} + for name in _W4A16_SIDECAR_BUFFERS: + key = _sidecar_key(fqn, name) + if key not in available: + continue + tensor = handle.get_tensor(key) + expected_dtype = _W4A16_SIDECAR_DTYPES[name] + if tensor.dtype != expected_dtype: + raise ValueError(f"W4A16 sidecar tensor {key} has dtype {tensor.dtype}, expected " + f"{expected_dtype}. Codes are uint8 bit patterns; a cast would silently " + "unpack to different 4-bit values.") + expected = _expected_sidecar_shapes(name, weight_shape, group_size, bits) + if tuple(tensor.shape) not in expected: + raise ValueError(f"W4A16 sidecar tensor {key} has shape {tuple(tensor.shape)}, expected " + f"one of {list(expected)} for a {list(weight_shape)} linear with " + f"group_size={group_size}, bits={bits}.") + device = _sidecar_target_device(mod, name) + if device is not None: + tensor = tensor.to(device=device, non_blocking=True) + tensors[name] = tensor + if set(tensors) != set(_W4A16_SIDECAR_BUFFERS): + message = (f"W4A16 sidecar entry for {fqn!r} is incomplete (has {sorted(tensors)}); " + f"all of {list(_W4A16_SIDECAR_BUFFERS)} are required.") + if strict: + raise ValueError(message) + logger.warning("%s Skipping this layer.", message) + continue + for name, tensor in tensors.items(): + mod.register_buffer(name, tensor, persistent=False) + mod._w4a16_weight_shape = weight_shape + restored += 1 + + logger.info("W4A16 sidecar: restored %d quantized modules from %s (dense weights untouched).", restored, + os.fspath(path)) + return restored + + +__all__ = [ + "DEFAULT_BITS", + "DEFAULT_GROUP_SIZE", + "MINIMAX_H3_BLOCK_LINEAR_SUFFIXES", + "MINIMAX_H3_BLOCK_SCOPES", + "MINIMAX_H3_NUM_LAYERS", + "MINIMAX_H3_NUM_REFINER_LAYERS", + "MINIMAX_H3_PREFIX", + "W4A16Config", + "W4A16QuantizeMethod", + "W4A16_DIR_SIDECAR_NAME", + "W4A16_SIDECAR_SUFFIX", + "convert_model_to_w4a16", + "load_w4a16_checkpoint", + "minimax_h3_w4a16_prefixes", + "read_w4a16_sidecar_metadata", + "save_w4a16_checkpoint", + "w4a16_dequantize", + "w4a16_quantize", + "w4a16_sidecar_path_for", + "w4a16_sidecar_state_dict", +] diff --git a/fastvideo/models/dits/minimax_h3.py b/fastvideo/models/dits/minimax_h3.py index 370d11d47b..42a27286d2 100644 --- a/fastvideo/models/dits/minimax_h3.py +++ b/fastvideo/models/dits/minimax_h3.py @@ -83,8 +83,6 @@ def __init__(self, rope_freq_dim: int, rope_theta: float) -> None: def forward(self, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """Build rotary tensors on the device that owns the packed positions.""" position_ids = position_ids.to(torch.float32) - # Analytic rotary positional embedding (RoPE) state is non-persistent, - # so runtime coordinates own the device after loading or state offload. inv_freq = self.inv_freq.to(position_ids.device) freqs = position_ids.unsqueeze(-1) * inv_freq.view(1, 1, -1) freqs_t, freqs_h, freqs_w = freqs.unbind(dim=1) @@ -188,11 +186,6 @@ def __init__( prefix=f"{prefix}.to_out", ) self.fuse_qknorm_rope = fuse_qknorm_rope - # VSA carries a learned gate on its pooled-compression branch. The H3 - # checkpoint has no such weight, so the loader zero-initializes it - # (ALLOWED_NEW_PARAM_PATTERNS) and the branch is exactly disabled - # until finetuned. Built only when VSA-H3 actually resolves, keeping - # the FLASH/SDPA paths and their state_dict untouched. resolved_backend = get_attn_backend(attention_head_dim, get_compute_dtype(), supported_attention_backends=supported_attention_backends) @@ -207,8 +200,6 @@ def __init__( fa4_packed_varlen=fa4_packed_varlen, ) self.to_gate_compress: ReplicatedLinear | None = None - # None = unchecked; the first forward tests the loaded weight once and - # skips the gate branch entirely while it is structurally zero. self._gate_compress_active: bool | None = None if use_vsa: self.to_gate_compress = ReplicatedLinear( @@ -229,6 +220,7 @@ def _gate_active(self) -> bool: reach a zero gate for it to ever train). """ if torch.is_grad_enabled(): + self._gate_compress_active = None return True if self._gate_compress_active is None: if torch.compiler.is_compiling(): @@ -252,8 +244,6 @@ def _resolve_gate_compress_for_compile(self) -> None: return if self._gate_compress_active is None: weight = self.to_gate_compress.weight - # bool() on a DTensor reduction resolves collectively, so every - # rank caches the same answer. self._gate_compress_active = bool((weight != 0).any()) @staticmethod @@ -298,8 +288,6 @@ def forward( query = self._apply_rotary_emb(query, rotary_emb) key = self._apply_rotary_emb(key, rotary_emb) - # H3 rotates only 96/128 channels, which the generic `freqs_cis` - # branch cannot express. Apply it above, then pass no RoPE here. extra_attention_kwargs = {} if self.to_gate_compress is not None and self._gate_active(): gate_compress, _ = self.to_gate_compress(hidden_states) @@ -481,7 +469,6 @@ def __init__( fuse_modulate: bool = False, fuse_qknorm_rope: bool = False, fuse_swiglu: bool = False, - fa4_packed_varlen: bool = False, ) -> None: super().__init__() self.norm1 = nn.RMSNorm(hidden_size, eps=norm_eps) @@ -494,7 +481,7 @@ def __init__( quant_config, prefix=f"{prefix}.attn", fuse_qknorm_rope=fuse_qknorm_rope, - fa4_packed_varlen=fa4_packed_varlen, + fa4_packed_varlen=True, ) self.norm2 = nn.RMSNorm(hidden_size, eps=norm_eps) self.ff = MiniMaxH3FeedForward( @@ -592,16 +579,12 @@ class MiniMaxH3Transformer3DModel(BaseDiT): def _get_parameter_dtype(self, name: str, default_dtype: torch.dtype) -> torch.dtype: """Keep the released input, timestep, and output projections in FP32. - Factorized AdaLN uses FP16; BF16 is ~1.7x worse there. + Folded AdaLN parameters follow the enclosing FSDP policy. In + particular, an FP32 training load must keep them as FP32 optimizer + masters; pinning the folded weights to FP16 here quantizes every + small Adam update before the next forward. Release checkpoints are + still exported in BF16, matching the validated 42-block parent. """ - # Precedence: the factorized-AdaLN FP16 pin wins over - # uniform_parameter_dtype on purpose. Under FSDP's one-dtype rule the - # resulting mix hard-fails at load time, which beats silently training - # AdaLN in BF16. Rank-reduced checkpoints are inference artifacts -- - # train from the full-rank release. - if getattr(self, "adaln_rank", None) is not None and ( - ".adaln_proj." in name or name.startswith(("norm_out.linear.", "adaln_basis."))): - return torch.float16 if self.config.uniform_parameter_dtype: return default_dtype return torch.float32 if name.split(".", 1)[0] in self._keep_in_fp32_modules else default_dtype @@ -666,14 +649,6 @@ def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None: prefix=f"{config.prefix}.time_embedder", ) self.adaln_rank: int | None = arch.adaln_rank - if self.adaln_rank is not None and config.uniform_parameter_dtype: - raise ValueError( - "Rank-reduced AdaLN checkpoints (adaln_rank set) cannot be trained: " - "uniform_parameter_dtype needs one dtype for every trainable " - "parameter, but factorized AdaLN weights are pinned to FP16 " - "(BF16 reconstructs them ~1.7x worse). Fine-tune the full-rank " - "checkpoint instead, then re-fit the basis with " - "scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py.") adaln_dim = self.adaln_rank or arch.time_embed_dim self.adaln_basis = ReplicatedLinear( arch.time_embed_dim, @@ -684,8 +659,6 @@ def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None: ) if self.adaln_rank else None self.rope = MiniMaxH3RotaryPosEmbed(arch.rope_freq_dim, arch.rope_theta) - # per-generation caches for loop-invariant work (see _rotary_for / - # _refined_text); plain attrs, never in state_dict self._rope_cache: tuple | None = None self._text_cache: tuple | None = None self.token_refiner = MiniMaxH3TokenRefiner( @@ -697,8 +670,6 @@ def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None: arch.norm_eps, arch.qk_norm_eps, arch.final_norm_eps, - # The refiner attends over the text stream only; the packed-sequence - # VSA backend must never be selected for it. tuple(backend for backend in self.supported_attention_backends if backend != AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3), config.quant_config, @@ -720,7 +691,6 @@ def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None: fuse_modulate="modulate" in self.enabled_fusions, fuse_qknorm_rope="qknorm_rope" in self.enabled_fusions, fuse_swiglu="swiglu" in self.enabled_fusions, - fa4_packed_varlen=envs.FASTVIDEO_MINIMAX_H3_FA4_PACKED_VARLEN, ) for index in range(arch.num_layers) ]) self.norm_out = MiniMaxH3AdaLayerNormOut( @@ -911,9 +881,6 @@ def forward( rotary_emb = self._rotary_for(position_ids, text_embeds.dtype) sp_world_size = get_sp_world_size() if model_parallel_is_initialized() else 1 - # text/video/audio indices partition [0, sequence_length), so the - # uninitialized buffer is fully overwritten; in-place index_copy_ avoids - # the three full-buffer clones out-of-place index_copy would make. packed_hidden_states = text_embeds.new_empty((text_embeds.shape[0], sequence_length, text_embeds.shape[-1])) packed_hidden_states.index_copy_(1, text_indices, text_embeds) packed_hidden_states.index_copy_(1, video_indices, video_embeds.to(text_embeds.dtype)) @@ -935,8 +902,6 @@ def forward( local_timestep_indices, _ = sequence_model_parallel_shard(local_timestep_indices, dim=0) rotary_emb = (rotary_cos, rotary_sin) - # The eager driver owns profiling markers while each block's compiled - # forward owns the graph that the marker surrounds. for block_index, block in enumerate(self.transformer_blocks): with nvtx_range(f"minimax_h3.transformer_block.{block_index}"): packed_hidden_states = block( diff --git a/fastvideo/models/loader/component_loader.py b/fastvideo/models/loader/component_loader.py index 148a6d1bb0..704510c18d 100644 --- a/fastvideo/models/loader/component_loader.py +++ b/fastvideo/models/loader/component_loader.py @@ -45,12 +45,34 @@ safetensors_weights_iterator, ) from fastvideo.models.registry import ModelRegistry +from fastvideo.platforms import AttentionBackendEnum from fastvideo.utils import PRECISION_TO_TYPE, is_pin_memory_available from fastvideo.hooks.layerwise_offload import enable_layerwise_offload logger = init_logger(__name__) +def _teacher_critic_attention_context(fastvideo_args: FastVideoArgs): + """Mask only a generator-only QAT request for teacher/critic loads. + + ``_loading_teacher_critic_model`` predates role-local attention backends and + also controls custom-weight/quant-config handling. Treating the flag as a + blanket request for automatic attention used to erase an explicit + ``FLASH_ATTN`` request from DMD teacher and critic roles. Preserve every + explicit dense request; only suppress ``ATTN_QAT_TRAIN``, which is the + generator-only policy the flag was introduced to isolate. + """ + if not hasattr(fastvideo_args, "_loading_teacher_critic_model"): + return nullcontext() + + active_scope = _active_component_attention_backend_scope() + requested = (active_scope.backend if active_scope is not None else + coerce_attn_backend(getattr(fastvideo_args, "attention_backend", None))) + if requested is AttentionBackendEnum.ATTN_QAT_TRAIN: + return _component_attention_backend_scope(None, component="transformer") + return nullcontext() + + class ComponentLoader(ABC): """Base class for loading a specific type of model component.""" @@ -83,7 +105,6 @@ def for_module_type(cls, module_type: str, transformers_or_diffusers: str) -> "C Returns: A component loader for the specified module type """ - # Map of module types to their loader classes and expected library module_loaders = { "scheduler": (SchedulerLoader, "diffusers"), "audio_scheduler": (SchedulerLoader, "diffusers"), @@ -112,29 +133,18 @@ def for_module_type(cls, module_type: str, transformers_or_diffusers: str) -> "C "processor": (ProcessorLoader, "transformers"), "upsampler": (UpsamplerLoader, "diffusers"), "upsampler_2": (UpsamplerLoader, "diffusers"), - # Stable Audio's `StableAudioMultiConditioner` bundles T5 + - # NumberConditioners; not a pure text encoder, so it gets - # its own loader. "conditioner": (ConditionerLoader, "fastvideo"), - # LTX-2 spatial / temporal upsamplers — share the - # UpsamplerLoader path with the upsampler/upsampler_2 keys - # so the SR pipeline picks up real weights instead of the - # generic config-only loader. "spatial_upsampler": (UpsamplerLoader, "diffusers"), "temporal_upsampler": (UpsamplerLoader, "diffusers"), } if module_type in module_loaders: loader_cls, expected_library = module_loaders[module_type] - # Allow fastvideo.* libraries for custom implementations (e.g. Cosmos2_5Pipeline) - # that aren't available in diffusers/transformers yet is_fastvideo_module = transformers_or_diffusers.startswith("fastvideo.") if not is_fastvideo_module: - # Assert that the library matches what's expected for this module type assert transformers_or_diffusers == expected_library, f"{module_type} must be loaded from {expected_library}, got {transformers_or_diffusers}" return loader_cls() - # For unknown module types, use a generic loader logger.warning( "No specific loader found for module type: %s. Using generic loader.", module_type, @@ -173,8 +183,6 @@ def _prepare_weights( """Prepare weights for the model. If the model is not local, it will be downloaded.""" - # model_name_or_path = (self._maybe_download_from_modelscope( - # model_name_or_path, revision) or model_name_or_path) is_local = os.path.isdir(model_name_or_path) assert is_local, "Model path must be a local directory" @@ -223,7 +231,6 @@ def _get_weights_iterator(self, source: "Source", to_cpu: bool) -> Generator[tup if self.counter_before_loading_weights == 0.0: self.counter_before_loading_weights = time.perf_counter() - # Apply the prefix. return ((source.prefix + name, tensor) for (name, tensor) in weights_iterator) def _get_all_weights( @@ -249,12 +256,6 @@ def _get_all_weights( def load(self, model_path: str, fastvideo_args: FastVideoArgs): """Load the text encoders based on the model path, and inference args.""" - # model_config: PretrainedConfig = get_hf_config( - # model=model_path, - # trust_remote_code=fastvideo_args.trust_remote_code, - # revision=fastvideo_args.revision, - # model_override_args=None, - # ) model_config = get_diffusers_config(model=model_path) model_config.pop("_name_or_path", None) model_config.pop("transformers_version", None) @@ -328,7 +329,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): encoder_precision = encoder_precisions[idx] target_device = get_local_torch_device() - # TODO(will): add support for other dtypes return self.load_model( model_path, encoder_config, @@ -356,9 +356,6 @@ def load_model( offload_flag=offload_flag) if requested_cpu_offload and disable_cpu_offload: - # Direct loader callers can choose a CPU target before the worker - # applies its device-local policy. Reset both the request and the - # target so the model is never constructed on the host first. logger.info("Disabling %s on unified-memory device %d", offload_flag, device_id) cpu_offload = False target_device = runtime_device @@ -411,9 +408,6 @@ def load_model( dtype=PRECISION_TO_TYPE[dtype], device=target_device, ) - # HF passthrough encoders return before FastVideo's FSDP - # wrapping path, so the text stage needs their placement to - # put token tensors on the same device. model._fastvideo_input_device = target_device return model.eval() @@ -450,10 +444,6 @@ def load_model( self.counter_after_loading_weights - self.counter_before_loading_weights, ) - # Strict check for unquantized models and for serialized checkpoint - # schemes, which track every loaded tensor. It runs before the - # post-load hooks so a missing tensor is reported by name here - # instead of as one anonymous "not loaded" failure inside a hook. weights_not_loaded = weights_to_load - loaded_weights if weights_not_loaded and (model_config.quant_config is None or checkpoint_quant_config is not None): raise ValueError("Following weights were not initialized from " @@ -464,14 +454,12 @@ def load_model( logger.info("Validated %d serialized %s text-encoder linears", processed_linears, checkpoint_quant_config.get_name()) - # Explicitly move model to target device after loading weights model = model.to(target_device) from fastvideo.platforms import current_platform if use_cpu_offload: pin_cpu_memory = fastvideo_args.pin_cpu_memory and is_pin_memory_available() - # Disable FSDP for MPS as it's not compatible if current_platform.is_mps(): logger.info("Disabling FSDP sharding for MPS platform as it's not compatible") elif current_platform.is_npu(): @@ -509,12 +497,6 @@ class ImageEncoderLoader(TextEncoderLoader): def load(self, model_path: str, fastvideo_args: FastVideoArgs): """Load the text encoders based on the model path, and inference args.""" - # model_config: PretrainedConfig = get_hf_config( - # model=model_path, - # trust_remote_code=fastvideo_args.trust_remote_code, - # revision=fastvideo_args.revision, - # model_override_args=None, - # ) with open(os.path.join(model_path, "config.json")) as f: model_config = json.load(f) model_config.pop("_name_or_path", None) @@ -562,7 +544,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): target_device = (torch.device("mps") if current_platform.is_mps() else torch.device("cpu")) else: target_device = get_local_torch_device() - # TODO(will): add support for other dtypes return self.load_model( model_path, encoder_config, @@ -626,8 +607,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): logger.info("Loading tokenizer from %s", model_path) resolved_model_path = model_path - # LTX2 checkpoints may not ship a top-level tokenizer/ directory. - # In that case, tokenizer assets live under text_encoder/gemma/. if not os.path.isdir(resolved_model_path): ltx2_gemma_path = os.path.normpath(os.path.join(resolved_model_path, "..", "text_encoder", "gemma")) if os.path.isdir(ltx2_gemma_path): @@ -638,8 +617,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): ) resolved_model_path = ltx2_gemma_path - # Cosmos2.5 stores an AutoProcessor config in `tokenizer/config.json` (not a tokenizer - # config). Use its `_name_or_path` (e.g. Qwen/Qwen2.5-VL-7B-Instruct) as the source. tokenizer_cfg_path = os.path.join(resolved_model_path, "config.json") if os.path.exists(tokenizer_cfg_path): try: @@ -660,14 +637,8 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): ) return processor except Exception: - # If parsing fails, fall through to AutoTokenizer below. pass - # Only Flux2 full's Mistral3 (require_processor=True) must load via - # AutoProcessor. Gate the processor_config.json shortcut on that flag so - # existing encoders (e.g. HunyuanVideo 1.5 / Qwen2.5-VL) stay on the - # historical AutoTokenizer path below even if their tokenizer dir happens - # to ship a processor_config.json. require_processor = False if hasattr(fastvideo_args.pipeline_config, "text_encoder_configs"): try: @@ -693,9 +664,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): try: tokenizer = AutoTokenizer.from_pretrained( resolved_model_path, # "/tokenizer" - # in v0, this was same string as encoder_name "ClipTextModel" - # TODO(will): pass these tokenizer kwargs from inference args? Maybe - # other method of config? local_files_only=os.path.isdir(resolved_model_path), ) except (OSError, ValueError): @@ -775,8 +743,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): is_gen3c = pipeline_name.startswith("Gen3C") is_cosmos25 = pipeline_name == "Cosmos25Config" - # GEN3C: prefer tokenizer-backed VAE checkpoint when available. - # This aligns latent conditioning with the GEN3C temporal contract. if is_gen3c and class_name in ("AutoencoderKLWan", "AutoencoderKLGen3CTokenizer"): from fastvideo.models.vaes.gen3c_tokenizer_vae import (AutoencoderKLGen3CTokenizer) @@ -825,7 +791,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): class_name, ) - # Cosmos2.5 uses a Wan2.1 VAE stored as `tokenizer.safetensors` under the VAE folder. if class_name == "AutoencoderKLWan" and is_cosmos25: from fastvideo.models.vaes.cosmos25wanvae import Cosmos25WanVAE @@ -855,7 +820,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): ).to(target_device) return vae.eval() - # LTX-2 uses CausalVideoAutoencoder with nested "vae" config if class_name == "CausalVideoAutoencoder" and "vae" in config: vae_cls, _ = ModelRegistry.resolve_model_cls(class_name) vae = vae_cls(config).to(target_device) @@ -874,17 +838,13 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): vae_cls, _ = ModelRegistry.resolve_model_cls(class_name) vae = vae_cls(vae_config).to(target_device) - # Find all safetensors files safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors")) if not safetensors_list: raise ValueError(f"No safetensors files found in {model_path}") - # Common case: a single `.safetensors` checkpoint file. - # Some models may be sharded into multiple files; in that case we merge. loaded = {} for sf_file in safetensors_list: loaded.update(safetensors_load_file(sf_file)) - # LTX-2 CausalVideoAutoencoder needs per_channel_statistics remapping if class_name == "CausalVideoAutoencoder" and "vae" in config: per_channel_prefixes = ( "per_channel_statistics.", @@ -907,8 +867,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): break loaded = remapped - # Diffusers-format AutoencoderKL checkpoints should match exactly; load - # strictly so missing/unexpected keys are surfaced early. strict_load = class_name in {"AutoencoderKL", "AutoencoderKLMiniMaxH3"} vae.load_state_dict(loaded, strict=strict_load) if (class_name == "AutoencoderKLWan" and getattr(vae.config, "use_light_vae", False) @@ -950,10 +908,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): return audio_vae.eval() precision = getattr(fastvideo_args.pipeline_config, "audio_decoder_precision", "bf16") - # MMAudio normalizes its magnitude-preserving convolution weights in - # fp32 and only then casts the whole feature utility module to bf16. - # Constructing/loading directly in bf16 quantizes the unnormalized - # checkpoint first and changes the decoded mel trajectory. construction_precision = "fp32" if class_name == "MMAudioVAE" else precision construction_device = torch.device("cpu") if class_name == "MMAudioVAE" else target_device with set_default_torch_dtype(PRECISION_TO_TYPE[construction_precision]): @@ -992,8 +946,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): target_device = get_local_torch_device() precision = getattr(fastvideo_args.pipeline_config, "vocoder_precision", "bf16") - # Canonical BigVGAN likewise removes parametrized weight norm in fp32 - # before the official MMAudio feature module is cast to bf16. construction_precision = "fp32" if class_name == "BigVGANV2" else precision construction_device = torch.device("cpu") if class_name == "BigVGANV2" else target_device with set_default_torch_dtype(PRECISION_TO_TYPE[construction_precision]): @@ -1046,36 +998,24 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): fastvideo_args.model_paths["transformer"] = model_path - # Config from Diffusers supersedes fastvideo's model config dit_config = deepcopy(fastvideo_args.pipeline_config.dit_config) dit_config.update_model_arch(config) - # Generator-only QAT for DMD distillation: the teacher (real_score) and - # critic (fake_score) transformers load with this flag set and must stay - # full precision. Drop the nvfp4_qat quant from their copied config, and - # build their attention under a scope that ignores any process-wide - # ATTN_QAT_TRAIN request so it falls back to dense. The generator loads - # without the flag and keeps both. The scope is exception-safe and - # needs no env mutation or selector cache flush (the request is part - # of the resolution cache key). - _qat_generator_only = hasattr(fastvideo_args, "_loading_teacher_critic_model") - if _qat_generator_only: + _teacher_or_critic = hasattr(fastvideo_args, "_loading_teacher_critic_model") + if _teacher_or_critic: dit_config.quant_config = None model_cls, _ = ModelRegistry.resolve_model_cls(cls_name) - # Find all safetensors files safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors")) if not safetensors_list: raise ValueError(f"No safetensors files found in {model_path}") - # arch_config can infer architecture from weight keys (e.g. Flux2 layer counts) update_fn = getattr(dit_config.arch_config, "update_from_weight_keys", None) if callable(update_fn): weight_keys = _collect_safetensors_keys(safetensors_list) update_fn(weight_keys) - # Check if we should use custom initialization weights custom_weights_path = getattr(fastvideo_args, "init_weights_from_safetensors", None) use_custom_weights = (custom_weights_path and os.path.exists(custom_weights_path) and not hasattr(fastvideo_args, "_loading_teacher_critic_model")) @@ -1099,23 +1039,13 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): default_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision] - # Load the model using FSDP loader logger.info("Loading model from %s, default_dtype: %s", cls_name, default_dtype) assert fastvideo_args.hsdp_shard_dim is not None - # Cosmos2.5 checkpoints can include extra entries not present in the - # instantiated model (e.g. pos_embedder ranges / *_extra_state). Load - # non-strictly for Cosmos2.5 only; keep upstream strict behavior for others. strict_load = not (cls_name.startswith("Cosmos25") or cls_name == "Cosmos25Transformer3DModel" or getattr(fastvideo_args.pipeline_config, "prefix", "") == "Cosmos25") - attention_context = (_component_attention_backend_scope(None, component="transformer") - if _qat_generator_only else nullcontext()) + attention_context = _teacher_critic_attention_context(fastvideo_args) with attention_context: - # dit_config is what the model is handed and keeps as `self.config`, - # so recording here makes the decision readable from the loaded - # transformer — and records the narrowed one for teacher/critic. resolved = record_resolved_attention_backend(dit_config) - # Every worker records its resolved backend so distributed profile - # snapshots can prove that all ranks use the requested kernels. logger.info("Worker %s transformer attention backend: %s", os.environ.get("RANK", "0"), resolved.name if resolved else "automatic selection", @@ -1134,7 +1064,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): 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 default_dtype=default_dtype, param_dtype=torch.bfloat16, reduce_dtype=torch.float32, @@ -1142,11 +1071,10 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): training_mode=fastvideo_args.training_mode, enable_torch_compile=fastvideo_args.enable_torch_compile, torch_compile_kwargs=fastvideo_args.torch_compile_kwargs, + regional_compile=getattr(fastvideo_args, "regional_compile", False), + pre_fsdp_transform=getattr(fastvideo_args, "_pre_fsdp_transform", None), inference_regional_compile=fastvideo_args.inference_torch_compile, inference_vsa_tile_size=fastvideo_args.VSA_tile_size, - # Only the whole-parameter half of the adapter is applied here, while - # tensors are still unsharded; LoRAPipeline merges the low-rank half - # once the module tree exists. lora_path=getattr(fastvideo_args, "lora_path", None), lora_strength=getattr(fastvideo_args, "lora_strength", 1.0), ) @@ -1162,7 +1090,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): model = model.eval() if fastvideo_args.inference_mode and fastvideo_args.dit_layerwise_offload: - # Check if model has nn.ModuleList for layerwise offload compatibility has_module_list = any(isinstance(m, nn.ModuleList) for m in model.children()) if has_module_list: enable_layerwise_offload(model) @@ -1216,17 +1143,11 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): precision = getattr(fastvideo_args.pipeline_config, "precision", "fp16") target_dtype = PRECISION_TO_TYPE.get(precision, torch.float16) - # Without this merge the model falls back to its dataclass - # defaults (e.g. SA-1.0's 3-conditioner spec — wrong for SA-small). from dataclasses import fields as _fields from fastvideo.configs.models.encoders import ( StableAudioConditionerConfig, ) if model_cls.__name__ == "StableAudioMultiConditioner": cond_config = StableAudioConditionerConfig() - # `update_model_arch` is strict (raises on unknown keys); the - # converter writes a few non-arch keys (`_class_name`, - # `_diffusers_version`, `_name_or_path`) that must be filtered - # out first. valid = {f.name for f in _fields(cond_config.arch_config)} cond_config.update_model_arch({k: v for k, v in config.items() if k in valid}) with set_default_torch_dtype(target_dtype): @@ -1239,8 +1160,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): if not os.path.isfile(weights): raise FileNotFoundError(f"Conditioner weights not found: {weights}") state = safetensors_load_file(weights) - # Non-strict: T5 weights live outside this checkpoint (fetched in - # the conditioner's `__init__` from the standard HF repo). model.load_state_dict(state, strict=False) return model.to(device=target_device, dtype=target_dtype).eval() @@ -1257,12 +1176,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): raise ValueError("Model config does not contain a _class_name attribute. " "Only diffusers format is supported.") - # The base PipelineConfig declares ``upsampler_config`` as a - # single ``UpsamplerConfig`` instance, but Hunyuan15 narrows it - # to a tuple of two configs (one per SR target). We only treat - # the attribute as a multi-config when it actually is one; - # otherwise the LTX-2 branch below handles the single-class - # path that takes the diffusers config dict directly. upsampler_config_attr = getattr(fastvideo_args.pipeline_config, "upsampler_config", None) if isinstance(upsampler_config_attr, list | tuple): try: @@ -1272,9 +1185,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): upsampler_cfg = deepcopy(upsampler_config_attr[1]) upsampler_cfg.update_model_config(config_dict) elif class_name == "LTX2LatentUpsampler": - # LTX-2 pipeline_config does not declare upsampler_config; the - # `LTX2LatentUpsampler` wrapper takes the raw diffusers config - # dict directly via LatentUpsamplerConfigurator. upsampler_cfg = deepcopy(config_dict) else: raise AttributeError("pipeline_config.upsampler_config is missing; cannot build " @@ -1298,11 +1208,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): for sf_file in safetensors_list: loaded.update(safetensors_load_file(sf_file)) - # The LTX-2 latent upsampler wrapper exposes the actual conv - # stack at ``self.model``; checkpoint state_dicts may be saved - # without the ``model.`` prefix when the inner module was - # serialised directly. Strip / forward as needed so both layouts - # load cleanly. target_module = getattr(model, "model", model) if loaded and all(k.startswith("model.") for k in loaded): stripped = {k[len("model."):]: v for k, v in loaded.items()} @@ -1346,7 +1251,6 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): model_config = get_diffusers_config(model=model_path) logger.info("Diffusers Model config: %s", model_config) - # This is a placeholder - in a real implementation, you'd need to handle this properly return None else: raise ValueError(f"Unsupported library: {self.library}") @@ -1384,19 +1288,8 @@ def load_module( component_model_path, ) - # Get the appropriate loader for this module type loader = ComponentLoader.for_module_type(module_name, transformers_or_diffusers) - # Resolve this component's attention backend ONCE, here, from the - # explicit request and the environment. A per-role request already - # resolved upstream (the train stack's moduleloader) wins; otherwise - # fastvideo_args.attention_backend is the process-wide default, applied - # per component. - # - # The loaders record the decision on their own config rather than this - # function doing it after the fact, because a loader may narrow it for - # one component: the DMD teacher/critic transformers build dense inside - # a nested scope (see `record_resolved_attention_backend`). if _active_component_attention_backend_scope() is not None: attention_context = nullcontext() else: diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index e076359ee6..276ac2c2a2 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -1,8 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -# Adapted from torchtune -# Copyright 2024 The TorchTune Authors. -# Copyright 2025 The FastVideo Authors. from __future__ import annotations import os @@ -16,11 +13,17 @@ from torch import nn from torch.distributed import DeviceMesh, init_device_mesh from torch.distributed._tensor import distribute_tensor +from torch.distributed.tensor import DTensor, Replicate, Shard from torch.distributed.fsdp import (CPUOffloadPolicy, FSDPModule, MixedPrecisionPolicy, fully_shard) from torch.nn.modules.module import _IncompatibleKeys from fastvideo.logger import init_logger from fastvideo.models.loader.lora_patch import DenseLoRAPatch +from fastvideo.models.loader.shard_cache import ( + shard_cache_context, + try_load_from_shard_cache, + write_shard_cache, +) from fastvideo.models.loader.utils import (get_param_names_mapping, hf_to_custom_state_dict) from fastvideo.models.loader.weight_utils import safetensors_weights_iterator from fastvideo.utils import set_mixed_precision_policy, is_pin_memory_available @@ -28,6 +31,81 @@ logger = init_logger(__name__) +def _dtensor_from_cpu_full_tensor( + full_tensor: torch.Tensor, + meta_sharded_param: DTensor, + device: torch.device, + target_dtype: torch.dtype, +) -> DTensor | None: + """Materialize only this rank's DTensor shard on CUDA. + + The checkpoint iterator already yields CPU tensors. Moving the complete + tensor to every GPU before ``distribute_tensor`` creates an avoidable + full-tensor H2D peak (about 1 GiB for H3's largest matrices), which can + OOM a multi-role DMD2 process even though the final local shards fit. + Replicate+Shard meshes can be sliced exactly on CPU and reconstructed as + a DTensor from the local piece. Unknown placements retain the established + full-tensor fallback. + """ + coordinate = meta_sharded_param.device_mesh.get_coordinate() + if coordinate is None: + return None + local = full_tensor + for mesh_dim, placement in enumerate(meta_sharded_param.placements): + if isinstance(placement, Replicate): + continue + if not isinstance(placement, Shard): + return None + chunks = torch.tensor_split( + local, + int(meta_sharded_param.device_mesh.size(mesh_dim)), + dim=int(placement.dim), + ) + local = chunks[int(coordinate[mesh_dim])] + local = local.contiguous().to(device=device, dtype=target_dtype) + return DTensor.from_local( + local, + meta_sharded_param.device_mesh, + meta_sharded_param.placements, + run_check=False, + shape=meta_sharded_param.shape, + stride=meta_sharded_param.stride(), + ) + + +def _mixed_precision_module_groups( + model: nn.Module, + default_param_dtype: torch.dtype | None, +) -> tuple[list[tuple[str, nn.Module]], set[nn.Parameter]]: + """Resolve declared FP32 FSDP groups and uncovered mixed parameters.""" + dtype_selector = getattr(model, "_get_parameter_dtype", None) + if not callable(dtype_selector) or default_param_dtype is None: + return [], set() + + mixed_params = { + parameter + for name, parameter in model.named_parameters() + if dtype_selector(_strip_checkpoint_wrapper_prefix(name), default_param_dtype) != default_param_dtype + } + declared = set(getattr(model, "_keep_in_fp32_modules", ())) + groups: list[tuple[str, nn.Module]] = [] + covered: set[nn.Parameter] = set() + for name, module in model.named_modules(): + clean_name = _strip_checkpoint_wrapper_prefix(name) + if clean_name not in declared or any(clean_name.startswith(f"{parent}.") for parent, _ in groups): + continue + parameters = list(module.named_parameters()) + if not parameters: + continue + if not all( + dtype_selector(_strip_checkpoint_wrapper_prefix(f"{clean_name}.{child_name}"), + default_param_dtype) == torch.float32 for child_name, _ in parameters): + continue + groups.append((clean_name, module)) + covered.update(parameter for _, parameter in parameters) + return groups, mixed_params - covered + + def _summarize_param_names(names: set[str]) -> str: """Collapse per-layer parameter names into one ``blocks.*.suffix xN`` entry each.""" families: dict[str, int] = {} @@ -54,7 +132,6 @@ def _maybe_quantize_model(model: nn.Module, *, defer_weight_conversion_until_lor silent, so the same walk emits a one-line receipt with the count of linears that actually carry the QAT method. """ - # Defer imports: these modules pull in heavy symbols at module-load time. from fastvideo.layers.linear import LinearBase from fastvideo.layers.quantization.nvfp4_config import ( NVFP4QuantizeMethod, @@ -70,6 +147,14 @@ def _maybe_quantize_model(model: nn.Module, *, defer_weight_conversion_until_lor FP8QuantizeMethod, convert_model_to_fp8, ) + from fastvideo.layers.quantization.int8_affine_config import ( + INT8AffineQuantizeMethod, + convert_model_to_int8_affine, + ) + from fastvideo.layers.quantization.w4a16_config import ( + W4A16QuantizeMethod, + convert_model_to_w4a16, + ) from fastvideo.layers.quantization.mxfp8_config import ( MXFP8QuantizeMethod, convert_model_to_mxfp8, @@ -94,6 +179,14 @@ def _maybe_quantize_model(model: nn.Module, *, defer_weight_conversion_until_lor logger.info("Converting loaded model weights for FP8 linear layers") convert_model_to_fp8(model) return + if isinstance(qm, INT8AffineQuantizeMethod): + logger.info("Converting loaded model weights for INT8 affine linear layers") + convert_model_to_int8_affine(model) + return + if isinstance(qm, W4A16QuantizeMethod): + logger.info("Converting loaded model weights for W4A16 linear layers") + convert_model_to_w4a16(model) + return if isinstance(qm, MXFP8QuantizeMethod): if defer_weight_conversion_until_lora_merge: logger.info("Deferring MXFP8 weight conversion until the inference LoRA merge completes") @@ -101,9 +194,6 @@ def _maybe_quantize_model(model: nn.Module, *, defer_weight_conversion_until_lor logger.info("Converting loaded model weights for MXFP8 linear layers") convert_model_to_mxfp8(model) return - # QAT-train configs are mutually exclusive with the inference schemes - # above (one quant_config per model), so when they're active the loop - # always runs to completion and the counts below are model-wide. if isinstance(qm, NVFP4QATTrainQuantizeMethod): qat_train_attached += 1 elif isinstance(mod, LinearBase): @@ -113,7 +203,6 @@ def _maybe_quantize_model(model: nn.Module, *, defer_weight_conversion_until_lor qat_train_skipped) -# TODO(PY): move this to utils elsewhere @contextlib.contextmanager def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]: """ @@ -171,14 +260,11 @@ def _validate_fsdp_inference_quantization(init_params: dict[str, Any], fsdp_infe if isinstance(quant_config, NVFP4QATTrainConfig): return - # TODO: (David) Currently reject every FSDP inference quantization config except NVFP4QATTrainConfig. - # Support FSDP inference with precomputed quantized weights. raise NotImplementedError( "FSDP inference supports unquantized transformers and NVFP4QATTrainConfig only; " f"got {type(quant_config).__name__}.") -# Supports optional torch.compile for FSDP-wrapped models during training def maybe_load_fsdp_model( model_cls: type[nn.Module], init_params: dict[str, Any], @@ -197,6 +283,8 @@ def maybe_load_fsdp_model( pin_cpu_memory: bool = True, enable_torch_compile: bool = False, torch_compile_kwargs: dict[str, Any] | None = None, + pre_fsdp_transform: Callable[[nn.Module], nn.Module] | None = None, + regional_compile: bool = False, inference_regional_compile: bool = False, inference_vsa_tile_size: int | None = None, lora_path: str | None = None, @@ -214,8 +302,6 @@ def maybe_load_fsdp_model( """ _validate_fsdp_inference_quantization(init_params, fsdp_inference) - # NOTE(will): cast_forward_inputs=True shouldn't be needed as we are - # manually casting the inputs to the model mp_policy = MixedPrecisionPolicy(param_dtype, reduce_dtype, output_dtype, cast_forward_inputs=False) set_mixed_precision_policy( @@ -229,17 +315,25 @@ def maybe_load_fsdp_model( with set_default_dtype(default_dtype), torch.device("meta"): model = model_cls(**init_params) + if pre_fsdp_transform is not None: + model = pre_fsdp_transform(model) + dtype_selector = getattr(model, "_get_parameter_dtype", None) - has_mixed_parameter_dtypes = callable(dtype_selector) and any( - dtype_selector(name, param_dtype) != param_dtype for name, _ in model.named_parameters()) - if training_mode and has_mixed_parameter_dtypes: + parameter_dtype_overrides = [] + if callable(dtype_selector): + parameter_dtype_overrides = [ + (clean_name, str(selected_dtype)) for name, _ in model.named_parameters() + if (selected_dtype := dtype_selector(clean_name := _strip_checkpoint_wrapper_prefix(name), + default_dtype)) != default_dtype + ] + _, ungrouped_mixed_params = _mixed_precision_module_groups(model, param_dtype) + if training_mode and ungrouped_mixed_params: raise NotImplementedError("FSDP training with model-selected mixed parameter dtypes requires " - "separate gradient synchronization for replicated parameters.") + "separate gradient synchronization for replicated parameters or " + "declared FP32 module groups.") - # Check if we should use FSDP use_fsdp = training_mode or fsdp_inference - # Disable FSDP for MPS as it's not compatible from fastvideo.platforms import current_platform if current_platform.is_mps(): use_fsdp = False @@ -256,14 +350,12 @@ def maybe_load_fsdp_model( with torch.device("cpu"): device_mesh = init_device_mesh( "npu", - # (Replicate(), Shard(dim=0)) mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim), mesh_dim_names=("replicate", "shard"), ) else: device_mesh = init_device_mesh( "cuda", - # (Replicate(), Shard(dim=0)) mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim), mesh_dim_names=("replicate", "shard"), ) @@ -275,11 +367,6 @@ def maybe_load_fsdp_model( fsdp_shard_conditions=model._fsdp_shard_conditions, pin_cpu_memory=pin_cpu_memory) - # Host offload is already disabled on unified memory (GB10). Staging the - # 35B FastH3 DiT on CPU and then copying to CUDA doubled that working set - # and took minutes. Follow cpu_offload: read onto the accelerator. - weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=cpu_offload) - logger.info("Loading transformer weights with to_cpu=%s", cpu_offload) param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping) dense_lora_patch = DenseLoRAPatch.from_adapter( lora_path, @@ -287,9 +374,6 @@ def maybe_load_fsdp_model( strength=lora_strength, ) if dense_lora_patch is not None: - # H3's compression gate is created only by the VSA attention backend. Loading a - # VSA student under dense attention would otherwise warn about 50 unmatched - # replacements and continue with a silently incomplete model. model_parameter_names = {name for name, _ in model.named_parameters()} missing_vsa_gates = sorted(name for name in dense_lora_patch.replacement_parameters if "gate_compress" in name and name not in model_parameter_names) @@ -298,50 +382,68 @@ def maybe_load_fsdp_model( "This LoRA adapter provides MiniMax H3 VSA compression gates, but the selected attention backend " "did not construct them. Use attention_backend='VIDEO_SPARSE_ATTN_H3'. Missing parameters: " + ", ".join(missing_vsa_gates[:3]) + (" ..." if len(missing_vsa_gates) > 3 else "")) - load_model_from_full_model_state_dict( - model, - weight_iterator, - device, - default_dtype, - strict=strict, - cpu_offload=cpu_offload, - param_names_mapping=param_names_mapping_fn, - dense_lora_patch=dense_lora_patch, - ) + + shard_cache_ctx = None + if use_fsdp and not cpu_offload and dense_lora_patch is None: + shard_cache_ctx = shard_cache_context( + weight_dir_list=weight_dir_list, + device_mesh=device_mesh, + hsdp_replicate_dim=hsdp_replicate_dim, + hsdp_shard_dim=hsdp_shard_dim, + default_dtype=default_dtype, + param_dtype=param_dtype, + param_names_mapping=model.param_names_mapping, + parameter_dtype_overrides=parameter_dtype_overrides, + ) + cache_hit = (shard_cache_ctx is not None + and try_load_from_shard_cache(model, shard_cache_ctx, device, strict=strict)) + if not cache_hit: + weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=cpu_offload) + logger.info("Loading transformer weights with to_cpu=%s", cpu_offload) + load_model_from_full_model_state_dict( + model, + weight_iterator, + device, + default_dtype, + strict=strict, + cpu_offload=cpu_offload, + param_names_mapping=param_names_mapping_fn, + dense_lora_patch=dense_lora_patch, + ) + if shard_cache_ctx is not None: + write_shard_cache(model, shard_cache_ctx) if hasattr(model, "materialize_non_persistent_buffers"): model.materialize_non_persistent_buffers(device=device, dtype=default_dtype) for n, p in chain(model.named_parameters(), model.named_buffers()): if p.is_meta: raise RuntimeError(f"Unexpected param or buffer {n} on meta device.") - # Avoid unintended computation graph accumulation during inference if isinstance(p, torch.nn.Parameter): p.requires_grad = False - # Post-load weight quantization. We detect the active scheme by the - # ``quant_method`` attached to each linear layer at construction time - # (via ``QuantizationConfig.get_quant_method``). The loader's - # responsibility is just to materialize the quantized weight buffers - # from the freshly-loaded bf16 weights. No-op when no quantized layers - # are present (lazy imports inside the helper). _maybe_quantize_model(model, defer_weight_conversion_until_lora_merge=lora_path is not None) - compile_in_loader = enable_torch_compile and training_mode - if compile_in_loader: - unsupported = _prepare_model_for_compile(model, regional=False) - if unsupported is not None: - logger.warning("Training torch.compile requested but disabled: %s. Model stays eager.", unsupported) + if enable_torch_compile and training_mode: + if not regional_compile: + unsupported = _prepare_model_for_compile(model, regional=False) + if unsupported is not None: + logger.warning("Training torch.compile requested but disabled: %s. Model stays eager.", unsupported) + else: + compile_kwargs = torch_compile_kwargs or {} + logger.info("Enabling whole-model torch.compile with kwargs=%s", compile_kwargs) + model = torch.compile(model, **compile_kwargs) else: - 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__) + unsupported = _regional_compile_unsupported_reason(init_params, training_mode=True) + if unsupported is not None: + logger.warning( + "enable_torch_compile requested but disabled: %s. " + "Training continues in eager mode.", unsupported) + else: + _compile_model_regions(model, torch_compile_kwargs or {}) elif inference_regional_compile and not training_mode: - # Inference-side counterpart of the #1718 training regional compile: - # per-block fullgraph compile right after the transformer loads, no - # user kwargs needed (fullgraph + emulate_precision_casts injected). unsupported = _regional_compile_unsupported_reason( init_params, vsa_tile_size=inference_vsa_tile_size, + training_mode=False, ) if unsupported is None: unsupported = _prepare_model_for_compile(model, regional=True) @@ -356,41 +458,38 @@ def maybe_load_fsdp_model( return model +def _strip_checkpoint_wrapper_prefix(name: str) -> str: + """Canonicalize an FQN produced after activation-checkpoint wrapping.""" + return name.replace("._checkpoint_wrapped_module.", ".").removeprefix("_checkpoint_wrapped_module.") + + def _regional_compile_unsupported_reason( init_params: dict[str, Any], *, vsa_tile_size: int | None = None, + training_mode: bool = False, ) -> str | None: - """Return why regional fullgraph compile cannot run, or None if it can. - - Dense FA2, FA3, and FA4 inference all route through compile-visible - custom-op boundaries. FA3's raw autograd.Function carve-out applies only - to grad-enabled calls, outside this inference-only loader path. - - The legacy VSA backend remains outside the fullgraph support envelope. - MiniMax H3's VSA backend is supported only through the inference-only - sm_100a tile-64 route; its regional hook resolves loaded compression - gates and probes the kernel before block capture. - """ + """Return why regional fullgraph compile cannot run, or ``None``.""" try: - from fastvideo.attention.layer import _attention_compile_explicitly_disabled + from fastvideo.attention.layer import (_attention_compile_disabled, + _attention_compile_explicitly_disabled) except Exception: # pragma: no cover - attention stack not importable pass else: - if _attention_compile_explicitly_disabled(): - # The escape hatch wraps attention forwards in - # torch.compiler.disable, which is a hard dynamo error inside a - # fullgraph region ("Skip inlining `torch.compiler.disable()`d - # function"). Degrade to eager instead, matching the hatch's - # debugging intent. + compile_disabled = (_attention_compile_disabled() + if training_mode else _attention_compile_explicitly_disabled()) + if compile_disabled: return ("FASTVIDEO_DISABLE_ATTENTION_COMPILE=1 keeps attention " "forwards out of compiled graphs via torch.compiler." "disable, which fullgraph regional compile cannot trace; " - "this model stays eager") + "this role stays eager") config = init_params.get("config") resolved = getattr(config, "_resolved_attention_backend", None) resolved_name = getattr(resolved, "name", "") if resolved_name == "VIDEO_SPARSE_ATTN_H3": + if training_mode: + return ("VIDEO_SPARSE_ATTN_H3 training uses sparse kernels and collectives " + "that are not fullgraph-traceable; this role stays eager") if os.environ.get("FASTVIDEO_H3_VSA_PROBE"): return ("FASTVIDEO_H3_VSA_PROBE records tensors and files from the VSA-H3 attention body, which " "regional fullgraph compile cannot capture; this model stays eager") @@ -404,7 +503,18 @@ def _regional_compile_unsupported_reason( return (f"attention backend resolved to {resolved_name}, whose Triton " "kernels, sequence-parallel collectives, and sync metadata " "guard graph-break (incompatible with fullgraph regional " - "compile); this model stays eager") + "compile); this role stays eager") + if not training_mode or resolved is None or resolved_name != "FLASH_ATTN": + return None + try: + from fastvideo.attention.utils.flash_attn_default import fa_version + except Exception: # pragma: no cover - flash-attn stack not importable + return None + if fa_version == "3": + return ("attention backend resolved to FLASH_ATTN with flash-attn 3, " + "whose grad-enabled path graph-breaks (incompatible with " + "fullgraph regional compile); use FA2, FA4 (FASTVIDEO_FA4=1), " + "or TORCH_SDPA for compiled training") return None @@ -421,25 +531,20 @@ def _enable_regional_attention_compile(model: nn.Module) -> int: def _compile_model_regions(model: nn.Module, compile_kwargs: dict[str, Any]) -> int: - """Compile repeated mathematical regions of a loaded model. + """Compile repeated mathematical regions after FSDP setup. Only the selected module ``forward`` is replaced. This keeps activation - checkpoint wrappers structurally transparent while any module-level hooks - (FSDP pre/post, layerwise offload) execute outside the compiled region. + checkpoint wrappers structurally transparent while FSDP pre/post hooks + execute outside the compiled region. """ compile_conditions = getattr(model, "_compile_conditions", None) if not compile_conditions: raise ValueError(f"{type(model).__name__} does not declare _compile_conditions") if compile_kwargs.get("fullgraph", True) is not True: - raise ValueError("Regional compile requires fullgraph=True") + raise ValueError("Regional training compile requires fullgraph=True") if "mode" in compile_kwargs: - # torch.compile forbids passing both `mode` and `options`, and - # regional compile always injects options (emulate_precision_casts) - # to match the training-side regional-compile configuration. Fail here - # with an actionable message instead of letting torch raise a - # mode/options conflict about an `options` key the user never wrote. - raise ValueError("Regional compile sets inductor options " + raise ValueError("Regional training compile sets inductor options " "(emulate_precision_casts) and cannot be combined " "with torch_compile_kwargs['mode']. Remove 'mode' or " "express its effect via torch_compile_kwargs['options'].") @@ -452,9 +557,6 @@ def _compile_model_regions(model: nn.Module, compile_kwargs: dict[str, Any]) -> if not name: continue if any(condition(name, submodule) for condition in compile_conditions): - # Activation checkpoint wrappers are control-flow boundaries, not - # mathematical regions. Keep their saved-tensor/recompute logic - # eager and compile only the repeated block they own. compile_target = getattr(submodule, "_checkpoint_wrapped_module", submodule) compile_target.forward = torch.compile(compile_target.forward, **kwargs) compiled_count += 1 @@ -462,7 +564,7 @@ def _compile_model_regions(model: nn.Module, compile_kwargs: dict[str, Any]) -> if compiled_count == 0: raise ValueError(f"No submodules in {type(model).__name__} matched _compile_conditions") logger.info( - "Enabled regional torch.compile for %d submodules in %s with kwargs=%s", + "Enabled regional torch.compile for %d submodules in %s after FSDP setup with kwargs=%s", compiled_count, type(model).__name__, kwargs, @@ -506,7 +608,6 @@ def shard_model( Raises: ValueError: If no layer modules were sharded, indicating that no shard_condition was triggered. """ - # Check if we should use size-based filtering use_size_filtering = os.environ.get("FASTVIDEO_FSDP2_AUTOWRAP", "0") == "1" if not fsdp_shard_conditions: @@ -514,15 +615,12 @@ def shard_model( return default_param_dtype = getattr(mp_policy, "param_dtype", None) - dtype_selector = getattr(model, "_get_parameter_dtype", None) - ignored_params: set[nn.Parameter] = 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 - } + fp32_groups, ignored_params = _mixed_precision_module_groups(model, default_param_dtype) named_modules = list(model.named_modules()) + fp32_group_ids = {id(module) for _, module in fp32_groups} + fp32_group_params = { + parameter for _, module in fp32_groups for parameter in module.parameters() + } ignored_params_by_module = { id(module): ignored_params.intersection(set(module.parameters())) for _, module in named_modules @@ -536,27 +634,25 @@ def shard_model( if cpu_offload: fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(pin_memory=pin_cpu_memory) - # iterating in reverse to start with - # lowest-level modules first num_layers_sharded = 0 if use_size_filtering: - # Size-based filtering mode min_params = int(os.environ.get("FASTVIDEO_FSDP2_MIN_PARAMS", "10000000")) logger.info("Using size-based filtering with threshold: %.2fM", min_params / 1e6) for n, m in reversed(named_modules): if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]): - # Count all parameters + if id(m) in fp32_group_ids: + continue + if fp32_group_params.intersection(set(m.parameters())): + raise ValueError(f"FSDP shard condition for {n!r} contains a declared FP32 compute group") param_count = sum(p.numel() for p in m.parameters(recurse=True)) - # Skip small modules if param_count < min_params: logger.info("Skipping module %s (%.2fM params < %.2fM threshold)", n, param_count / 1e6, min_params / 1e6) continue - # Shard this module logger.info("Sharding module %s (%.2fM params)", n, param_count / 1e6) module_kwargs = fsdp_kwargs local_ignored_params = ignored_params_by_module[id(m)] @@ -565,9 +661,12 @@ def shard_model( fully_shard(m, **module_kwargs) num_layers_sharded += 1 else: - # 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]): + if id(m) in fp32_group_ids: + continue + if fp32_group_params.intersection(set(m.parameters())): + raise ValueError(f"FSDP shard condition for {n!r} contains a declared FP32 compute group") module_kwargs = fsdp_kwargs local_ignored_params = ignored_params_by_module[id(m)] if local_ignored_params: @@ -578,14 +677,34 @@ def shard_model( if num_layers_sharded == 0: raise ValueError("No layer modules were sharded. Please check if shard conditions are working as expected.") - # Finally shard the entire model to account for any stragglers + if fp32_groups: + fp32_kwargs = { + **fsdp_kwargs, + "mp_policy": MixedPrecisionPolicy( + param_dtype=torch.float32, + reduce_dtype=mp_policy.reduce_dtype, + output_dtype=mp_policy.output_dtype, + cast_forward_inputs=mp_policy.cast_forward_inputs, + ), + } + for _, module in fp32_groups: + fully_shard(module, **fp32_kwargs) + logger.info("Sharded FP32 compute modules: %s", [name for name, _ in fp32_groups]) + root_kwargs = fsdp_kwargs if ignored_params: root_kwargs = {**fsdp_kwargs, "ignored_params": ignored_params} fully_shard(model, **root_kwargs) -# TODO(PY): device mesh for cfg parallel +ALLOWED_NEW_PARAM_PATTERNS: tuple[str, ...] = ("gate_compress", "proj_l", "scale_weight", "scale_input") + + +def is_allowed_new_param(param_name: str) -> bool: + """Whether ``param_name`` may be zero-initialized instead of loaded.""" + return any(pattern in param_name for pattern in ALLOWED_NEW_PARAM_PATTERNS) + + def load_model_from_full_model_state_dict( model: FSDPModule | torch.nn.Module, full_sd_iterator: Generator[tuple[str, torch.Tensor], None, None], @@ -621,24 +740,15 @@ def load_model_from_full_model_state_dict( NotImplementedError: If got FSDP with more than 1D. """ meta_sd = model.state_dict() - named_parameters = dict(model.named_parameters()) - named_buffers = dict(model.named_buffers()) + named_parameters = {_strip_checkpoint_wrapper_prefix(k): v for k, v in model.named_parameters()} + named_buffers = {_strip_checkpoint_wrapper_prefix(k): v for k, v in model.named_buffers()} sharded_sd = {} custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(full_sd_iterator, param_names_mapping) # type: ignore - # Drain rather than iterate. Production safetensors values may retain - # memory-mapped shard storage, while mapped or merged parameters can own - # ordinary allocations. Keeping the dict retains all of that source - # storage until loading finishes; popping releases each reference as soon - # as its conversion completes and lowers the host/unified-memory working - # set. for target_param_name in list(custom_param_sd): full_tensor = custom_param_sd.pop(target_param_name) meta_sharded_param = meta_sd.get(target_param_name) if meta_sharded_param is None: - # Some checkpoints include extra entries that are not part of the - # instantiated model's state_dict (e.g. `_extra_state` keys from - # some FSDP checkpoint formats). These can be safely skipped. if (target_param_name.endswith("._extra_state") or target_param_name.endswith("_extra_state")): logger.warning( "Skipping non-parameter checkpoint key: %s", @@ -646,8 +756,6 @@ def load_model_from_full_model_state_dict( ) continue - # For non-strict loads, treat this as an "unexpected key" and skip it - # (mirrors torch.nn.Module.load_state_dict(strict=False)). if not strict: logger.warning( "Skipping unexpected checkpoint key (not present in model): %s", @@ -663,17 +771,11 @@ def load_model_from_full_model_state_dict( if callable(dtype_selector): target_dtype = dtype_selector(target_param_name, param_dtype) if dense_lora_patch is not None: - # Returns float32 when a delta was added, so the cast below is what lands - # the parameter in its storage dtype. full_tensor = dense_lora_patch.apply_to(target_param_name, full_tensor) if not hasattr(meta_sharded_param, "device_mesh"): full_tensor = full_tensor.to(device=device, dtype=target_dtype) target_param = named_parameters.get(target_param_name) weight_loader = getattr(target_param, "weight_loader", None) - # Gated on a shape mismatch: only fused/stacked params with a custom - # weight_loader (e.g. Qwen3's merged QKV/gate-up) take this path. - # Existing models whose unsharded params match the checkpoint shape - # fall through to the original `sharded_tensor = full_tensor` below. if target_param is not None and callable(weight_loader) and tuple(target_param.shape) != tuple( full_tensor.shape): loaded_param = nn.Parameter(torch.empty(tuple(target_param.shape), device=device, dtype=target_dtype), @@ -683,15 +785,24 @@ def load_model_from_full_model_state_dict( weight_loader(loaded_param, full_tensor) sharded_tensor = loaded_param.data else: - # In cases where parts of the model aren't sharded, some parameters will be plain tensors. sharded_tensor = full_tensor else: - full_tensor = full_tensor.to(device=device, dtype=target_dtype) - sharded_tensor = distribute_tensor( - full_tensor, - meta_sharded_param.device_mesh, - meta_sharded_param.placements, - ) + sharded_tensor = None + if full_tensor.device.type == "cpu" and isinstance(meta_sharded_param, DTensor): + sharded_tensor = _dtensor_from_cpu_full_tensor( + full_tensor, + meta_sharded_param, + device, + target_dtype, + ) + if sharded_tensor is None: + full_tensor = full_tensor.to(device=device, dtype=target_dtype) + sharded_tensor = distribute_tensor( + full_tensor, + meta_sharded_param.device_mesh, + meta_sharded_param.placements, + src_data_rank=None, + ) if cpu_offload: sharded_tensor = sharded_tensor.cpu() if target_param_name in named_buffers: @@ -702,10 +813,6 @@ def load_model_from_full_model_state_dict( model.reverse_param_names_mapping = reverse_param_names_mapping unused_keys = set(meta_sd.keys()) - set(sharded_sd.keys()) if unused_keys: - # Say which of these the adapter is about to fill in. Reporting all of them as - # "unloaded" was accurate when zero-init was the only outcome; with an adapter - # supplying real values it reads as a problem that is not one. Names are - # summarized because a 50-layer model prints 50 near-identical lines otherwise. from_adapter = ({key for key in unused_keys if dense_lora_patch.provides(key)} if dense_lora_patch is not None else set()) zero_init = unused_keys - from_adapter @@ -716,18 +823,13 @@ def load_model_from_full_model_state_dict( logger.warning("Found unloaded parameters in meta state dict, zero-initializing: %d (%s)", len(zero_init), _summarize_param_names(zero_init)) - # List of allowed parameter name patterns - ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l"] # Can be extended as needed for new_param_name in unused_keys: - # An adapter that ships the parameter outright both supplies the value and - # authorizes it: the allowlist exists to catch a checkpoint silently missing a - # weight, which is not the case when something deliberately provides one. adapter_value = (dense_lora_patch.replacement_for(new_param_name) if dense_lora_patch is not None else None) - if adapter_value is None and not any(pattern in new_param_name for pattern in ALLOWED_NEW_PARAM_PATTERNS): + if adapter_value is None and not is_allowed_new_param(new_param_name): logger.error("Unsupported new parameter: %s. Allowed patterns: %s", new_param_name, - ALLOWED_NEW_PARAM_PATTERNS) + list(ALLOWED_NEW_PARAM_PATTERNS)) raise ValueError(f"New parameter '{new_param_name}' is not supported. " - f"Currently only parameters containing {ALLOWED_NEW_PARAM_PATTERNS} are allowed.") + f"Currently only parameters containing {list(ALLOWED_NEW_PARAM_PATTERNS)} are allowed.") meta_sharded_param = meta_sd.get(new_param_name) target_dtype = param_dtype dtype_selector = getattr(model, "_get_parameter_dtype", None) @@ -749,10 +851,8 @@ def load_model_from_full_model_state_dict( if cpu_offload: sharded_tensor = sharded_tensor.cpu() elif not hasattr(meta_sharded_param, "device_mesh"): - # Initialize with zeros sharded_tensor = torch.zeros_like(meta_sharded_param, device=device, dtype=target_dtype) else: - # Initialize with zeros and distribute full_tensor = torch.zeros_like(meta_sharded_param, device=device, dtype=target_dtype) sharded_tensor = distribute_tensor( full_tensor, @@ -766,5 +866,4 @@ def load_model_from_full_model_state_dict( if dense_lora_patch is not None: dense_lora_patch.report_unapplied() - # choose `assign=True` since we cannot call `copy_` on meta tensor return model.load_state_dict(sharded_sd, strict=strict, assign=True) diff --git a/fastvideo/models/loader/shard_cache.py b/fastvideo/models/loader/shard_cache.py new file mode 100644 index 0000000000..e954af5ebd --- /dev/null +++ b/fastvideo/models/loader/shard_cache.py @@ -0,0 +1,342 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Per-rank sharded base-weight cache for fast training relaunches. + +After a full checkpoint load, each shard rank persists its local DTensor +chunks (post-rename, post-cast — exactly what ``assign=True`` installed) as +one safetensors file in a cache directory (typically tmpfs). Subsequent +launches with an identical (checkpoint, mesh layout, dtype, name-mapping) +tuple rebuild the model from those chunks via ``DTensor.from_local`` — +skipping the full-tensor reads, per-rank H2D of the whole checkpoint, and +the distribute/scatter step. + +Opt-in via ``FASTVIDEO_WEIGHT_SHARD_CACHE=`` (e.g. ``/dev/shm/fv-wcache``). +Any validation failure or exception degrades to the normal full load — the +cache can never fail a run. +""" + +from __future__ import annotations + +import hashlib +import json +import os +import shutil +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import torch +import torch.distributed as dist +import torch.nn as nn +from torch.distributed.tensor import DTensor + +from fastvideo.logger import init_logger + +logger = init_logger(__name__) + +_FORMAT_VERSION = 1 +_ENV_DIR = "FASTVIDEO_WEIGHT_SHARD_CACHE" +_ENV_MAX_GB = "FASTVIDEO_WEIGHT_SHARD_CACHE_MAX_GB" +_ENV_PER_NODE = "FASTVIDEO_WEIGHT_SHARD_CACHE_PER_NODE" +_ALLOWED_NEW_PARAM_PATTERNS = ("gate_compress", "proj_l", "scale_weight", "scale_input") +_WRITE_MARGIN_BYTES = 5 << 30 + + +@dataclass +class ShardCacheContext: + entry_dir: Path + key: str + shard_index: int + num_shards: int + is_writer: bool # replicate-coordinate 0 writes; replicas are identical + + +def _expand_weight_files(weight_dir_list: list[str]) -> list[str]: + files: list[str] = [] + for entry in weight_dir_list: + if os.path.isdir(entry): + files.extend(str(p) for p in sorted(Path(entry).glob("*.safetensors"))) + else: + files.append(entry) + return files + + +def _shard_file(entry_dir: Path, shard_index: int, num_shards: int) -> Path: + return entry_dir / f"shard{shard_index}-of-{num_shards}.safetensors" + + +def shard_cache_context( + *, + weight_dir_list: list[str], + device_mesh: Any, + hsdp_replicate_dim: int, + hsdp_shard_dim: int, + default_dtype: torch.dtype, + param_dtype: torch.dtype, + param_names_mapping: dict[str, str] | None, + parameter_dtype_overrides: list[tuple[str, str]] | None = None, +) -> ShardCacheContext | None: + root = os.environ.get(_ENV_DIR) + if not root: + return None + try: + files = _expand_weight_files(weight_dir_list) + if not files: + return None + stats = sorted((os.path.basename(p), os.stat(p).st_size, os.stat(p).st_mtime_ns) for p in files) + mapping_items = sorted((param_names_mapping or {}).items()) + key_fields = [ + _FORMAT_VERSION, + stats, + int(hsdp_replicate_dim), + int(hsdp_shard_dim), + str(default_dtype), + str(param_dtype), + mapping_items, + ] + if parameter_dtype_overrides: + key_fields.append(sorted(parameter_dtype_overrides)) + key_material = json.dumps(key_fields, sort_keys=True) + key = hashlib.sha256(key_material.encode()).hexdigest()[:16] + coordinate = device_mesh.get_coordinate() + if coordinate is None: + return None + replicate_index, shard_index = int(coordinate[0]), int(coordinate[1]) + per_node_root = os.environ.get(_ENV_PER_NODE, "0") == "1" + return ShardCacheContext( + entry_dir=Path(root) / key, + key=key, + shard_index=shard_index, + num_shards=int(hsdp_shard_dim), + is_writer=per_node_root or replicate_index == 0, + ) + except Exception as exc: # noqa: BLE001 - cache must never fail a load + logger.warning("shard cache disabled for this load (context error): %s", exc) + return None + + +def _all_ranks_agree(local_ok: bool, device: torch.device) -> bool: + if not (dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1): + return local_ok + flag = torch.tensor([1 if local_ok else 0], device=device, dtype=torch.int32) + dist.all_reduce(flag, op=dist.ReduceOp.MIN) + return bool(flag.item()) + + +def _validate_entry( + entry: dict[str, Any], + meta_param: torch.Tensor, + expected_dtype: torch.dtype, +) -> bool: + if entry["dtype"] != str(expected_dtype): + return False + if list(entry["global_shape"]) != list(meta_param.shape): + return False + is_dtensor = isinstance(meta_param, DTensor) + if entry["kind"] != ("dtensor" if is_dtensor else "tensor"): + return False + if is_dtensor: + if entry["placements"] != [str(p) for p in meta_param.placements]: + return False + if list(entry["local_shape"]) != list(meta_param.to_local().shape): + return False + return True + + +def try_load_from_shard_cache( + model: nn.Module, + ctx: ShardCacheContext, + device: torch.device, + *, + strict: bool = True, +) -> bool: + """Assemble the model's state dict from cached local shards. Returns False + (without mutating the model) on any mismatch.""" + try: + manifest_path = ctx.entry_dir / "manifest.json" + shard_path = _shard_file(ctx.entry_dir, ctx.shard_index, ctx.num_shards) + local_ok = manifest_path.is_file() and shard_path.is_file() + manifest: dict[str, Any] = {} + if local_ok: + manifest = json.loads(manifest_path.read_text()) + local_ok = (manifest.get("format_version") == _FORMAT_VERSION and manifest.get("key") == ctx.key) + meta_sd = model.state_dict() + dtype_selector = getattr(model, "_get_parameter_dtype", None) + if local_ok: + params_table = manifest["params"] + for name, meta_param in meta_sd.items(): + entry = params_table.get(name) + if entry is None: + if not any(pattern in name for pattern in _ALLOWED_NEW_PARAM_PATTERNS): + local_ok = False + break + continue + expected_dtype = meta_param.dtype + if callable(dtype_selector): + expected_dtype = dtype_selector(name, expected_dtype) + if not _validate_entry(entry, meta_param, expected_dtype): + local_ok = False + break + if not _all_ranks_agree(local_ok, device): + if local_ok: + logger.info("shard cache: another rank missed entry %s; falling back to full load", ctx.key) + return False + + from safetensors import safe_open + + from fastvideo.models.loader.fsdp_load import ( + _strip_checkpoint_wrapper_prefix, ) + + named_buffers = {_strip_checkpoint_wrapper_prefix(k): v for k, v in model.named_buffers()} + sharded_sd: dict[str, Any] = {} + with safe_open(str(shard_path), framework="pt", device=str(device)) as f: + cached_keys = set(f.keys()) + for name, meta_param in meta_sd.items(): + if name in cached_keys: + local = f.get_tensor(name) + if isinstance(meta_param, DTensor): + tensor: torch.Tensor = DTensor.from_local( + local, + meta_param.device_mesh, + meta_param.placements, + run_check=False, + shape=meta_param.shape, + stride=meta_param.stride(), + ) + else: + tensor = local + else: + target_dtype = meta_param.dtype + if callable(dtype_selector): + target_dtype = dtype_selector(name, target_dtype) + if isinstance(meta_param, DTensor): + local = torch.zeros(meta_param.to_local().shape, device=device, dtype=target_dtype) + tensor = DTensor.from_local( + local, + meta_param.device_mesh, + meta_param.placements, + run_check=False, + shape=meta_param.shape, + stride=meta_param.stride(), + ) + else: + tensor = torch.zeros(meta_param.shape, device=device, dtype=target_dtype) + sharded_sd[name] = tensor if name in named_buffers else nn.Parameter(tensor) + + reverse_map = manifest.get("reverse_param_names_mapping", {}) + model.reverse_param_names_mapping = {k: tuple(v) for k, v in reverse_map.items()} + model.load_state_dict(sharded_sd, strict=strict, assign=True) + for p in (shard_path, manifest_path): + try: + os.utime(p) + except OSError: + pass + logger.info( + "shard cache HIT %s: %d tensors from %s", + ctx.key, + len(sharded_sd), + shard_path, + ) + return True + except Exception as exc: # noqa: BLE001 - cache must never fail a load + logger.warning("shard cache load failed (%s); falling back to full load", exc) + return False + + +def write_shard_cache(model: nn.Module, ctx: ShardCacheContext) -> None: + """Persist this rank's local shards after a successful full load.""" + try: + from safetensors.torch import save_file + + tensors: dict[str, torch.Tensor] = {} + params_table: dict[str, Any] = {} + for name, value in model.state_dict().items(): + if isinstance(value, DTensor): + local = value.to_local().detach().to("cpu", copy=True).contiguous() + params_table[name] = { + "kind": "dtensor", + "dtype": str(value.dtype), + "global_shape": list(value.shape), + "placements": [str(p) for p in value.placements], + "local_shape": list(local.shape), + } + else: + local = value.detach().to("cpu", copy=True).contiguous() + params_table[name] = { + "kind": "tensor", + "dtype": str(value.dtype), + "global_shape": list(value.shape), + } + tensors[name] = local + + ctx.entry_dir.mkdir(parents=True, exist_ok=True) + needed = sum(t.numel() * t.element_size() for t in tensors.values()) + free = shutil.disk_usage(ctx.entry_dir).free + if free < needed + _WRITE_MARGIN_BYTES: + logger.warning( + "shard cache: skipping write (%.1f GiB needed, %.1f GiB free at %s)", + needed / 2**30, + free / 2**30, + ctx.entry_dir, + ) + _barrier_if_initialized() + return + + if ctx.is_writer: + shard_path = _shard_file(ctx.entry_dir, ctx.shard_index, ctx.num_shards) + tmp_path = shard_path.with_suffix(".safetensors.tmp") + save_file(tensors, str(tmp_path)) + os.replace(tmp_path, shard_path) + _barrier_if_initialized() + + rank = dist.get_rank() if (dist.is_available() and dist.is_initialized()) else 0 + if ctx.is_writer: + reverse_map = { + k: list(v) + for k, v in getattr(model, "reverse_param_names_mapping", {}).items() + } + manifest = { + "format_version": _FORMAT_VERSION, + "key": ctx.key, + "num_shards": ctx.num_shards, + "params": params_table, + "reverse_param_names_mapping": reverse_map, + } + manifest_tmp = ctx.entry_dir / f"manifest.json.tmp.{ctx.shard_index}" + manifest_tmp.write_text(json.dumps(manifest)) + os.replace(manifest_tmp, ctx.entry_dir / "manifest.json") + if rank == 0: + logger.info("shard cache WRITE %s: %d tensors -> %s", ctx.key, len(tensors), ctx.entry_dir) + _gc_cache_root(ctx.entry_dir.parent, keep=ctx.entry_dir.name) + except Exception as exc: # noqa: BLE001 - cache must never fail a run + logger.warning("shard cache write failed (non-fatal): %s", exc) + _barrier_if_initialized() + + +def _barrier_if_initialized() -> None: + if dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1: + dist.barrier() + + +def _gc_cache_root(root: Path, keep: str) -> None: + """Drop least-recently-written entries beyond the size cap.""" + try: + max_bytes = float(os.environ.get(_ENV_MAX_GB, "300")) * 2**30 + entries = [] + for child in root.iterdir(): + if not child.is_dir(): + continue + manifest = child / "manifest.json" + size = sum(f.stat().st_size for f in child.glob("*") if f.is_file()) + mtime = manifest.stat().st_mtime if manifest.is_file() else 0.0 + entries.append((mtime, size, child)) + total = sum(size for _, size, _ in entries) + for mtime, size, child in sorted(entries): + if total <= max_bytes: + break + if child.name == keep: + continue + shutil.rmtree(child, ignore_errors=True) + total -= size + logger.info("shard cache GC: evicted %s (%.1f GiB)", child, size / 2**30) + except Exception as exc: # noqa: BLE001 + logger.warning("shard cache GC failed (non-fatal): %s", exc) diff --git a/fastvideo/models/schedulers/scheduling_minimax_h3.py b/fastvideo/models/schedulers/scheduling_minimax_h3.py index 647075caa2..3cb21db195 100644 --- a/fastvideo/models/schedulers/scheduling_minimax_h3.py +++ b/fastvideo/models/schedulers/scheduling_minimax_h3.py @@ -1,5 +1,4 @@ # SPDX-License-Identifier: Apache-2.0 -# Copyright 2025 The MiniMax authors and The HuggingFace Team. from dataclasses import dataclass @@ -51,6 +50,11 @@ def set_shift(self, shift: float) -> None: raise ValueError(f"`shift` must be positive, got {shift}.") self._shift = float(shift) + def shift_sigmas(self, base_sigmas: torch.Tensor) -> torch.Tensor: + """Warp unshifted base noise amounts onto this scheduler's shifted grid.""" + base = torch.as_tensor(base_sigmas, dtype=torch.float32) + return self._shift * base / (1 + (self._shift - 1) * base) + def set_timesteps( self, num_inference_steps: int | None = None, @@ -62,7 +66,7 @@ def set_timesteps( raise ValueError("`set_timesteps` requires explicit `sigmas` or " f"`num_inference_steps` >= 2, got {num_inference_steps}.") base = torch.linspace(1.0, 0.0, int(num_inference_steps), dtype=torch.float32) - sigma_tensor = self._shift * base / (1 + (self._shift - 1) * base) + sigma_tensor = self.shift_sigmas(base) sigma_tensor = torch.unique_consecutive(sigma_tensor) else: sigma_tensor = torch.as_tensor(sigmas, dtype=torch.float32).flatten().cpu() @@ -119,8 +123,6 @@ def step( if not isinstance(timestep, torch.Tensor): timestep = torch.tensor(timestep, dtype=sample.dtype) - # H3 deliberately derives x0's sigma from the transformer timestep, while - # the Euler ratio below uses the stored grid. Keep the two float32 paths separate. sigma_from_timestep = 1 - timestep.to(device=sample.device, dtype=sample.dtype) while sigma_from_timestep.ndim < sample.ndim: sigma_from_timestep = sigma_from_timestep.unsqueeze(-1) diff --git a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_decoding.py b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_decoding.py index 1eae7005e5..deb7f2cb4b 100644 --- a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_decoding.py +++ b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_decoding.py @@ -7,7 +7,7 @@ import torch -from fastvideo.distributed import get_local_torch_device, get_sp_group, get_world_group, model_parallel_is_initialized +from fastvideo.distributed import get_local_torch_device, get_sp_group, model_parallel_is_initialized from fastvideo.fastvideo_args import FastVideoArgs from fastvideo.logger import init_logger from fastvideo.models.vaes.minimax_h3_audio import MiniMaxH3AudioVAE @@ -40,9 +40,9 @@ def _layout(batch: ForwardBatch) -> MiniMaxH3PackedLayout: def _decode_participation(fastvideo_args: FastVideoArgs, want_parallel: bool) -> tuple[Any, bool, bool]: """Resolve (sp_group, is_output_rank, parallel) for the VAE decode stages. - The existing serial path keeps its global-rank-zero output ownership. - Parallel decode assembles once per sequence-parallel group, on that - group's first rank. ``parallel`` is only true when every group rank will + Both paths produce one output per sequence-parallel group, on that + group's first rank. This retains every data-parallel validation sample + instead of silently keeping only global rank zero. ``parallel`` is only true when every group rank will run the decode body — the collectives inside require uniform participation, so no rank-dependent branch may guard them. """ @@ -51,7 +51,7 @@ def _decode_participation(fastvideo_args: FastVideoArgs, want_parallel: bool) -> sp_group = get_sp_group() if bool(want_parallel) and sp_group.world_size > 1: return sp_group, sp_group.is_first_rank, True - return sp_group, get_world_group().is_first_rank, False + return sp_group, sp_group.is_first_rank, False class MiniMaxH3VideoDecodingStage(PipelineStage): @@ -81,9 +81,6 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward placeholder = torch.empty((0, 3, 0, 0, 0), device="cpu", dtype=torch.float32) sp_group, is_output_rank, parallel = _decode_participation(fastvideo_args, fastvideo_args.vae_parallel_decode) if not is_output_rank and not parallel: - # Consumers read the output rank's ForwardBatch. Keep a - # verifier-compatible placeholder on other ranks and avoid - # duplicating the full VAE decode and CPU output buffer. batch.output = placeholder return batch @@ -127,8 +124,6 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward try: latents = self.vae.denormalize_latents(latents.to(device=device, dtype=torch.float32)) if fastvideo_args.output_type == "latent": - # No collectives on this path, so uniform participation is - # trivial: every rank returns here. batch.output = latents.detach().float().cpu() if is_output_rank else placeholder return batch @@ -140,8 +135,6 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward dtype=torch.float32, pin_memory=fastvideo_args.pin_cpu_memory and is_pin_memory_available(), ) - # Attribute the streamed decoder computation while retaining - # per-chunk device-to-host transfer and pinned-buffer reuse. with ( nvtx_range("minimax_h3.vae"), torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"), @@ -184,9 +177,7 @@ def verify_output(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> V @torch.no_grad() def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch: """Decode H3 audio latents into a stereo CPU waveform.""" - # Audio decode is sub-second, so preserve the serial path's global - # rank-zero ownership. - if model_parallel_is_initialized() and not get_world_group().is_first_rank: + if model_parallel_is_initialized() and not get_sp_group().is_first_rank: batch.extra["audio"] = torch.empty((0, 2), device="cpu", dtype=torch.float32) batch.extra["audio_sample_rate"] = self.audio_vae.sampling_rate self._clear_runtime(batch) @@ -209,8 +200,6 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward self._clear_runtime(batch) return batch - # The range isolates waveform synthesis from packing and runtime - # cleanup so the audio decoder has one stable timeline boundary. with nvtx_range("minimax_h3.audio_vae"): decoded = self.audio_vae.decode(latents).sample.float() if decoded.ndim != 3 or decoded.shape[0] != 2 or decoded.shape[1] != 1: diff --git a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py index bce4ea978e..0216bd8d5b 100644 --- a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py +++ b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py @@ -3,6 +3,8 @@ from __future__ import annotations +import os + from typing import Any import torch @@ -11,6 +13,7 @@ from fastvideo.distributed import get_local_torch_device from fastvideo.fastvideo_args import FastVideoArgs from fastvideo.forward_context import set_forward_context +from fastvideo.logger import init_logger from fastvideo.hooks.activation_trace import trace_step from fastvideo.profiler import nvtx_range, profiler_region from fastvideo.pipelines.basic.minimax_h3.packing import ( @@ -25,6 +28,8 @@ from fastvideo.pipelines.stages.validators import VerificationResult from fastvideo.utils import get_compute_dtype +logger = init_logger(__name__) + def _h3_vsa_metadata_builder(transformer: Any, fastvideo_args: FastVideoArgs) -> Any: """Builder instance when the transformer resolved to VSA-H3, else None. @@ -82,8 +87,6 @@ def _set_dmd_schedule(self, steps: list[int], grid_points: int, device: torch.de for scheduler in (self.scheduler, self.audio_scheduler): shift = float(scheduler.shift) sigmas = shift * base / (1 + (shift - 1) * base) - # Explicit sigmas are already shifted. Scheduler timesteps are H3 - # clean time (1 - sigma); passing integer rungs to step() is wrong. scheduler.set_timesteps(sigmas=sigmas, device=device) def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult: @@ -117,11 +120,21 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward device = get_local_torch_device() dmd_steps = fastvideo_args.pipeline_config.dmd_denoising_steps - if dmd_steps is None: + if not dmd_steps: + env_steps = os.environ.get("FASTVIDEO_DMD_DENOISING_STEPS", "").strip() + if env_steps: + dmd_steps = [int(s) for s in env_steps.split(",") if s.strip()] + stochastic_renoise = bool(dmd_steps) and ( + bool(getattr(fastvideo_args.pipeline_config, "dmd_stochastic_renoise", False)) + or os.environ.get("FASTVIDEO_DMD_STOCHASTIC_RENOISE", "0").strip().lower() in ("1", "true", "yes")) + if stochastic_renoise: + logger.info("MiniMax-H3 DMD denoising with stochastic fresh-noise re-noising.") + if dmd_steps: + self._set_dmd_schedule(dmd_steps, batch.num_inference_steps, device) + logger.info("MiniMax-H3 denoising with explicit DMD steps %s.", list(dmd_steps)) + else: self.scheduler.set_timesteps(batch.num_inference_steps, device=device) self.audio_scheduler.set_timesteps(batch.num_inference_steps, device=device) - else: - self._set_dmd_schedule(dmd_steps, batch.num_inference_steps, device) video_timesteps = self.scheduler.timesteps audio_timesteps = self.audio_scheduler.timesteps if video_timesteps is None or audio_timesteps is None: @@ -154,17 +167,14 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward if vsa_metadata_builder is not None: vsa_patch_size = fastvideo_args.pipeline_config.dit_config.patch_size vsa_prefix_segments = _h3_vsa_prefix_segments(layout, vsa_patch_size) - # Per-request knobs (sweeps flip these between generate_video calls - # without respawning workers); mode None defers to the env default. vsa_mode = batch.extra.get("vsa_mode", "exempt") if vsa_mode not in ("exempt", "compete"): raise ValueError(f"vsa_mode must be 'exempt' or 'compete', got {vsa_mode!r}.") vsa_exempt = vsa_mode == "exempt" vsa_dense_layers = tuple(batch.extra.get("vsa_dense_layers", ())) vsa_dense_first_n = int(batch.extra.get("vsa_dense_first_n_steps", 0)) - # Run-level tile geometry (256 default, 64 = native Triton path), - # plumbed like the run-level sparsity; the builder validates the - # value against VSA_H3_TILE_SHAPES. + vsa_sparsity_base = (float(batch.VSA_sparsity) + if float(batch.VSA_sparsity) > 0.0 else float(fastvideo_args.VSA_sparsity)) vsa_tile_size = int(fastvideo_args.VSA_tile_size) try: @@ -173,18 +183,13 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward batch.latents = batch.latents.to(device) batch.audio_latents = batch.audio_latents.to(device) - # The stage range groups the complete denoising loop while the - # indexed model ranges retain timing detail for every H3 block. with profiler_region("inference_denoising"), nvtx_range("minimax_h3.dit"): for index, (video_timestep, audio_timestep) in enumerate(zip(video_timesteps, audio_timesteps, strict=True)): unique_timesteps, timestep_indices = row_timestep_plan[index] attn_metadata = None if vsa_metadata_builder is not None: - # Optional schedule: run the first N steps dense (sparsity 0 - # selects every tile — parity-proven ≡ dense ≤2e-4); early - # steps set global structure and are the most damage-prone. - vsa_sparsity = 0.0 if index < vsa_dense_first_n else float(batch.VSA_sparsity) + vsa_sparsity = 0.0 if index < vsa_dense_first_n else vsa_sparsity_base attn_metadata = vsa_metadata_builder.build( current_timestep=index, raw_latent_shape=(layout.num_video_latent_frames, layout.latent_height, @@ -197,11 +202,6 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward dense_layers=vsa_dense_layers, tile_size=vsa_tile_size, ) - # Under torch.compile(mode="reduce-overhead") each denoising - # step must be marked, or cudagraph trees flag cross-step - # reuse of pooled outputs as "accessing tensor output of - # CUDAGraphs that has been overwritten" (surfaces at sp=1; - # sp>1 is masked by collective-induced graph breaks). torch.compiler.cudagraph_mark_step_begin() with trace_step(index), set_forward_context( current_timestep=index, @@ -223,18 +223,49 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward video_start = layout.num_condition_video_rows audio_start = layout.num_condition_audio_rows - batch.latents[video_start:] = self.scheduler.step( - video_velocity[0, video_start:].float(), - video_timestep, - batch.latents[video_start:], - return_dict=False, - )[0] - batch.audio_latents[audio_start:] = self.audio_scheduler.step( - audio_velocity[0, audio_start:].float(), - audio_timestep, - batch.audio_latents[audio_start:], - return_dict=False, - )[0] + if os.environ.get("FASTVIDEO_DMD_DEBUG_STATS", "0") == "1": + with torch.no_grad(): + for tag, lat, vel, st, sig in ( + ("video", batch.latents, video_velocity, video_start, self.scheduler.sigmas), + ("audio", batch.audio_latents, audio_velocity, audio_start, + self.audio_scheduler.sigmas), + ): + s = float(sig[index]) + xin = lat[st:].float() + x0dbg = xin + s * vel[0, st:].float() + logger.info( + "DMD_DEBUG step=%d %s sigma=%.4f in(std=%.4f,mean=%.4f) " + "x0(std=%.4f,mean=%.4f) v(std=%.4f)", index, tag, s, xin.std(), xin.mean(), + x0dbg.std(), x0dbg.mean(), vel[0, st:].float().std()) + if stochastic_renoise: + assert self.scheduler.sigmas is not None and self.audio_scheduler.sigmas is not None + for latents, velocity, start, sigmas in ( + (batch.latents, video_velocity, video_start, self.scheduler.sigmas), + (batch.audio_latents, audio_velocity, audio_start, self.audio_scheduler.sigmas), + ): + sigma = float(sigmas[index]) + sigma_next = float(sigmas[index + 1]) + sample = latents[start:].float() + pred_x0 = sample + sigma * velocity[0, start:].float() + if sigma_next > 0.0: + noise = torch.randn_like(pred_x0) + nxt = (1.0 - sigma_next) * pred_x0 + sigma_next * noise + else: + nxt = pred_x0 + latents[start:] = nxt.to(latents.dtype) + else: + batch.latents[video_start:] = self.scheduler.step( + video_velocity[0, video_start:].float(), + video_timestep, + batch.latents[video_start:], + return_dict=False, + )[0] + batch.audio_latents[audio_start:] = self.audio_scheduler.step( + audio_velocity[0, audio_start:].float(), + audio_timestep, + batch.audio_latents[audio_start:], + return_dict=False, + )[0] batch.step_index = index batch.timestep = video_timestep finally: diff --git a/fastvideo/pipelines/pipeline_batch_info.py b/fastvideo/pipelines/pipeline_batch_info.py index 2a7d2533e2..f3a41c86d3 100644 --- a/fastvideo/pipelines/pipeline_batch_info.py +++ b/fastvideo/pipelines/pipeline_batch_info.py @@ -1,5 +1,4 @@ # SPDX-License-Identifier: Apache-2.0 -# Inspired by SGLang: https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/model_executor/forward_batch_info.py """ Data structures for functional pipeline processing. @@ -29,7 +28,6 @@ class PipelineLoggingInfo: """Simple approach using OrderedDict to track stage metrics.""" def __init__(self): - # OrderedDict preserves insertion order and allows easy access self.stages: OrderedDict[str, dict[str, Any]] = OrderedDict() def add_stage_execution_time(self, stage_name: str, execution_time: float): @@ -67,39 +65,31 @@ class ForwardBatch: execution, allowing methods to update specific components without needing to manage numerous individual parameters. """ - # TODO(will): double check that args are separate from fastvideo_args - # properly. Also maybe think about providing an abstraction for pipeline - # specific arguments. data_type: str generator: torch.Generator | list[torch.Generator] | None = None - # Image inputs image_path: str | None = None image_embeds: list[torch.Tensor] = field(default_factory=list) pil_image: torch.Tensor | PIL.Image.Image | None = None last_image: torch.Tensor | PIL.Image.Image | None = None references: list[Any] | None = None preprocessed_image: torch.Tensor | None = None - # Text inputs prompt: str | list[str] | None = None negative_prompt: str | list[str] | None = None prompt_path: str | None = None output_path: str = "outputs/" output_video_name: str | None = None - # Video inputs video_path: str | None = None video_latent: torch.Tensor | None = None - # Refine inputs (LongCat) refine_from: str | None = None t_thresh: float = 0.5 spatial_refine_only: bool = False num_cond_frames: int = 0 stage1_video: list[PIL.Image.Image] | None = None # Loaded frames from refine_from - # Primary encoder embeddings prompt_embeds: list[torch.Tensor] = field(default_factory=list) negative_prompt_embeds: list[torch.Tensor] | None = None prompt_attention_mask: list[torch.Tensor] | None = None @@ -107,79 +97,63 @@ class ForwardBatch: clip_embedding_pos: list[torch.Tensor] | None = None clip_embedding_neg: list[torch.Tensor] | None = None - # Additional text-related parameters max_sequence_length: int | None = None prompt_template: dict[str, Any] | None = None do_classifier_free_guidance: bool = False - # When True, ``guidance_scale`` is passed into models that use embedded guidance (e.g. FLUX) - # and must not imply classic dual-forward CFG. Use ``true_cfg_scale > 1`` for true CFG. use_embedded_guidance: bool = False true_cfg_scale: float = 1.0 - # Batch info batch_size: int | None = None num_videos_per_prompt: int = 1 seed: int | None = None seeds: list[int] | None = None - # Tracking if embeddings are already processed is_prompt_processed: bool = False - # Latent tensors latents: torch.Tensor | None = None audio_latents: torch.Tensor | None = None lq_latents: torch.Tensor | None = None raw_latent_shape: tuple[int, ...] | None = None noise_pred: torch.Tensor | None = None image_latent: torch.Tensor | None = None - # Normalized clean first frame for Wan TI2V and causal-DMD conditioning. first_frame_latent: torch.Tensor | None = None - # Action control inputs (Matrix-Game) mouse_cond: torch.Tensor | None = None # Shape: (B, T, 2) keyboard_cond: torch.Tensor | None = None # Shape: (B, T, K) grid_sizes: torch.Tensor | None = None # Shape: (3,) [F,H,W] num_iterations: int | None = None use_base_model: bool = False - # Camera control inputs (HYWorld) pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path - # Camera/action control inputs (GameCraft) camera_states: torch.Tensor | None = None # Plücker coordinates [B, T, 6, H, W] gt_latents: torch.Tensor | None = None # Ground truth latents for conditioning [B, 16, T, H, W] conditioning_mask: torch.Tensor | None = None # Mask for conditioning [B, 1, T, H, W] camera_trajectory: str | None = None # Camera trajectory file/identifier action_list: list[str] | None = None # List of actions (e.g., ['forward', 'left']) action_speed_list: list[float] | None = None # Speed for each action - # Camera control inputs (LingBotWorld and LingBotWorld2) c2ws_plucker_emb: torch.Tensor | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat] action_path: str | None = None # Directory containing poses.npy and intrinsics.npy - # Camera control inputs (GEN3C) trajectory_type: str | None = None movement_distance: float | None = None camera_rotation: str | None = None - # Latent dimensions height_latents: list[int] | int | None = None width_latents: list[int] | int | None = None num_frames: list[int] | int = 1 # Default for image models - # Original dimensions (before VAE scaling) height: list[int] | int | None = None width: list[int] | int | None = None height_sr: list[int] | int | None = None width_sr: list[int] | int | None = None fps: list[int] | int | None = None - # Timesteps timesteps: torch.Tensor | None = None timestep: torch.Tensor | float | int | None = None step_index: int | None = None boundary_ratio: float | None = None - # Scheduler parameters num_inference_steps: int = 50 num_inference_steps_sr: int = 50 guidance_scale: float = 1.0 @@ -191,37 +165,27 @@ class ForwardBatch: eta: float = 0.0 sigmas: list[float] | None = None - # TeaCache enable_teacache: bool = False - # LTX-2 multi-modal CFG parameters ltx2_cfg_scale_video: float = 1.0 ltx2_cfg_scale_audio: float = 1.0 ltx2_modality_scale_video: float = 1.0 ltx2_modality_scale_audio: float = 1.0 ltx2_rescale_scale: float = 0.0 - # STG (Spatio-Temporal Guidance) parameters ltx2_stg_scale_video: float = 0.0 ltx2_stg_scale_audio: float = 0.0 ltx2_stg_blocks_video: list[int] = field(default_factory=list) ltx2_stg_blocks_audio: list[int] = field(default_factory=list) - # LTX-2 image / video / continuation conditioning ltx2_images: list[tuple[str, int, float]] | None = None ltx2_image_crf: float = 33.0 ltx2_conditioning_latent_stage1: torch.Tensor | None = None ltx2_conditioning_latent_stage2: torch.Tensor | None = None ltx2_video_conditions: list[tuple[list[str], int, float]] | None = None - # Stable Audio (T2A): clip start/end in seconds. Parallels the - # `SamplingParam` fields of the same name; the - # `StableAudioConditioningStage` / `DecodingStage` read them. audio_start_in_s: float | None = None audio_end_in_s: float | None = None - # Stable Audio A2A variation + inpainting payloads (parallel to - # `SamplingParam`). `Any` because we accept torch tensors or numpy - # arrays the user supplies; the latent-prep stage normalises shapes. init_audio: Any = None init_audio_strength: float | None = None init_noise_level: float | None = None @@ -230,13 +194,10 @@ class ForwardBatch: n_tokens: int | None = None - # Other parameters that may be needed by specific schedulers extra_step_kwargs: dict[str, Any] = field(default_factory=dict) - # Component modules (populated by the pipeline) modules: dict[str, Any] = field(default_factory=dict) - # Final output (after pipeline completion) output: torch.Tensor | None = None return_trajectory_latents: bool = False return_trajectory_decoded: bool = False @@ -247,26 +208,20 @@ class ForwardBatch: continuation_state: "ContinuationState | None" = None return_continuation_state: bool = False - # Extra parameters that might be needed by specific pipeline implementations extra: dict[str, Any] = field(default_factory=dict) - # Misc save_video: bool = True return_frames: bool = False is_cfg_negative: bool = False - # VSA parameters VSA_sparsity: float = 0.0 - # Logging info logging_info: PipelineLoggingInfo = field(default_factory=PipelineLoggingInfo) def __post_init__(self): """Initialize dependent fields after dataclass initialization.""" - # LTX-2 text CFG scales; FLUX uses ``use_embedded_guidance`` so ``guidance_scale > 1`` alone - # does not enable classifier-free guidance. ltx2_text_cfg_enabled = (self.ltx2_cfg_scale_video != 1.0 or self.ltx2_cfg_scale_audio != 1.0) if self.use_embedded_guidance: self.do_classifier_free_guidance = (self.true_cfg_scale > 1.0) or ltx2_text_cfg_enabled @@ -286,64 +241,51 @@ class TrainingBatch: current_timestep: int = 0 current_vsa_sparsity: float = 0.0 - # Dataloader batch outputs latents: torch.Tensor | None = None raw_latent_shape: tuple[int, ...] | None = None noise_latents: torch.Tensor | None = None encoder_hidden_states: torch.Tensor | None = None encoder_attention_mask: torch.Tensor | None = None - # LTX related audio inputs audio_latents: torch.Tensor | None = None audio_noisy_model_input: torch.Tensor | None = None audio_timesteps: torch.Tensor | None = None - # Audio follows an independent noise schedule, so multimodal supervised - # fine-tuning needs its own sigma to reconstruct the clean-audio target. audio_sigmas: torch.Tensor | None = None audio_noise: torch.Tensor | None = None audio_encoder_hidden_states: torch.Tensor | None = None audio_encoder_attention_mask: torch.Tensor | None = None conditioning_mask: torch.Tensor | None = None - # i2v preprocessed_image: torch.Tensor | None = None image_embeds: torch.Tensor | None = None image_latents: torch.Tensor | None = None infos: list[dict[str, Any]] | None = None mask_lat_size: torch.Tensor | None = None - # ODE trajectory supervision trajectory_latents: torch.Tensor | None = None trajectory_timesteps: torch.Tensor | None = None - # Transformer inputs noisy_model_input: torch.Tensor | None = None timesteps: torch.Tensor | None = None sigmas: torch.Tensor | None = None noise: torch.Tensor | None = None - # MiniMax H3 reuses the packed row boundaries from batch preparation to - # split the transformer's joint sequence back into video and audio outputs. minimax_h3_layout: Any | None = None + minimax_h3_dmd_layout: Any | None = None attn_metadata_vsa: AttentionMetadata | None = None attn_metadata: AttentionMetadata | None = None - # input kwargs input_kwargs: dict[str, Any] | None = None - # Training loss loss: torch.Tensor | None = None - # Training outputs total_loss: float | None = None grad_norm: float | None = None - # Distillation-specific attributes encoder_hidden_states_neg: torch.Tensor | None = None encoder_attention_mask_neg: torch.Tensor | None = None conditional_dict: dict[str, Any] | None = None unconditional_dict: dict[str, Any] | None = None - # Distillation losses generator_loss: float = 0.0 fake_score_loss: float = 0.0 diff --git a/fastvideo/pipelines/preprocess/preprocess_minimax_h3_overfit.py b/fastvideo/pipelines/preprocess/preprocess_minimax_h3_overfit.py index 33a26ce991..53c938fc2e 100644 --- a/fastvideo/pipelines/preprocess/preprocess_minimax_h3_overfit.py +++ b/fastvideo/pipelines/preprocess/preprocess_minimax_h3_overfit.py @@ -36,8 +36,6 @@ from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_input_preparation import MINIMAX_H3_KEYFRAMES_KEY from fastvideo.utils import verify_model_config_and_directory -# Match the H3 validation geometry and media rates, and limit both streams to -# the same maximum duration used by validation. NUM_FRAMES = 124 VIDEO_HEIGHT = 768 VIDEO_WIDTH = 1344 @@ -127,8 +125,6 @@ def _load_component( Using the registered loader keeps preprocessing aligned with the component classes and precision policy that produce H3 inference conditioning. """ - # Diffusers modular manifests can append loading metadata after the - # provider and architecture fields defined by the component contract. transformers_or_diffusers, _ = model_index[name][:2] return PipelineComponentLoader.load_module( module_name=name, @@ -153,8 +149,6 @@ def encode_video_latents( """ print("Loading MiniMax H3 video VAE") vae = _load_component("vae", model_path, model_index, fastvideo_args) - # Feed [B, C, T, H, W] pixels and retain [C, T, H, W] latents; H3 later - # flattens every video token in (C, patch_t, patch_h, patch_w) order. pixels = torch.from_numpy(frames.copy()).permute(3, 0, 1, 2)[None] pixels = pixels.to(device=torch.device("cuda:0"), dtype=torch.float32).div_(255.0) generator = torch.Generator("cpu").manual_seed(42) @@ -274,8 +268,6 @@ def write_parquet(record: dict[str, Any], output_dir: Path) -> Path: output_dir.mkdir(parents=True, exist_ok=True) for parquet_path in output_dir.glob("*.parquet"): parquet_path.unlink() - # The map-style cache stores Parquet file metadata and row counts, so a - # replacement shard set requires cache reconstruction. shutil.rmtree(output_dir / "map_style_cache", ignore_errors=True) table = pa.table( {name: [record[name]] @@ -313,21 +305,35 @@ def validate_preprocessed_training_data( print(f"Validated one Crush-Smol H3 training row in {expected_path}") -def main() -> None: - """Encode the selected Crush-Smol record with each H3 component in sequence. +def main( + video: Path | None = None, + prompt: str | None = None, + output_dir: Path = OUTPUT_DIR, + model_path: Path = MODEL_PATH, +) -> None: + """Encode one audio-video-caption record with each H3 component in sequence. + With no arguments this encodes the pinned Crush-Smol record; ``--video`` + + ``--prompt`` preprocess an arbitrary mp4 (with soundtrack) instead. Releasing each component before loading the next component keeps video, audio, and text preprocessing within one GPU's memory. """ _init_single_process_distributed() - resolved_model_path = MODEL_PATH.resolve() + resolved_model_path = model_path.resolve() if not resolved_model_path.is_dir(): raise FileNotFoundError(f"Filtered MiniMax H3 model directory is missing at {resolved_model_path}") model_index = verify_model_config_and_directory(str(resolved_model_path)) - video_path, caption = load_crush_smol_training_sample( - DATA_DIR / "videos2caption.json", - DATA_DIR / "videos", - ) + if video is None: + video_path, caption = load_crush_smol_training_sample( + DATA_DIR / "videos2caption.json", + DATA_DIR / "videos", + ) + else: + if not prompt or not prompt.strip(): + raise ValueError("--prompt is required when --video is given") + video_path, caption = video, prompt.strip() + if not video_path.is_file(): + raise FileNotFoundError(f"Training video is missing at {video_path}") frames, waveform = load_training_media(video_path) pipeline_config = MiniMaxH3PipelineConfig() fastvideo_args = FastVideoArgs( @@ -345,21 +351,34 @@ def main() -> None: audio_latents = encode_audio_latents(waveform, resolved_model_path, model_index, fastvideo_args) text_embedding = encode_text_embedding(caption, resolved_model_path, model_index, fastvideo_args) record = build_parquet_record( - file_name=TRAINING_VIDEO_NAME, + file_name=video_path.name, caption=caption, video_latents=video_latents, audio_latents=audio_latents, text_embedding=text_embedding, ) - output_path = write_parquet(record, OUTPUT_DIR) - print(f"Wrote one Crush-Smol MiniMax H3 T2VA record to {output_path}") + output_path = write_parquet(record, output_dir) + print(f"Wrote one MiniMax H3 T2VA record to {output_path}") if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("--validate-only", action="store_true") + parser.add_argument("--video", + type=Path, + default=None, + help="mp4 with a soundtrack (>=124 frames at 24 fps); " + "default: the pinned Crush-Smol record") + parser.add_argument("--prompt", type=str, default=None, help="caption for --video") + parser.add_argument("--output-dir", type=Path, default=OUTPUT_DIR) + parser.add_argument("--model-path", type=Path, default=MODEL_PATH) cli_args = parser.parse_args() if cli_args.validate_only: validate_preprocessed_training_data() else: - main() + main( + video=cli_args.video, + prompt=cli_args.prompt, + output_dir=cli_args.output_dir, + model_path=cli_args.model_path, + ) diff --git a/fastvideo/pipelines/preprocess/preprocess_minimax_h3_text_only.py b/fastvideo/pipelines/preprocess/preprocess_minimax_h3_text_only.py new file mode 100644 index 0000000000..843abf58f5 --- /dev/null +++ b/fastvideo/pipelines/preprocess/preprocess_minimax_h3_text_only.py @@ -0,0 +1,212 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Encode a prompt list into MiniMax H3 text-only conditioning rows (data-free DMD2). + +Each line of ``--prompts-file`` is one complete H3 prompt document (for the +VidProM-H3 set: ``integrated_multimodal_description: ... overall_soundscape: +... non_diegetic_music: ...`` fed to the model verbatim as one string). Every +prompt is tokenized raw (no chat template, ``add_special_tokens=False``) and +encoded through ``MiniMaxH3ConditioningStage`` — the same Qwen3-VL layer-50 +path that produced the validated t2va overfit rows — then written as +``pyarrow_schema_text_only`` records for ``rollout_mode: simulate`` training +(``training.data.preprocessed_data_type: text_only``). + +Embeddings are stored as float32 to match the training collate, which decodes +``text_embedding_bytes`` with a hard-coded ``np.float32`` +(``fastvideo/dataset/utils.py``). At ~300 tokens x 5120 dims that is ~6 MB per +prompt — budget ~380 GB for the full 63k VidProM set. + +Sharding: ``--num-shards N`` splits the prompt list round-robin by line index; +run one process per GPU with distinct ``--shard-index``/``CUDA_VISIBLE_DEVICES`` +(and a distinct ``MASTER_PORT`` — each process initializes a one-rank process +group). Each shard writes ``/shard_XX/``; the training dataloader +walks the directory tree, so pointing ``training.data.data_path`` at +``--output-dir`` picks up every shard. Restarting a shard resumes after the +rows already on disk (delete the shard directory to re-encode from scratch). +""" + +from __future__ import annotations + +import argparse +import json +import os +import time +from pathlib import Path +from typing import Any + +import pyarrow.parquet as pq +import torch + +from fastvideo.configs.pipelines.minimax_h3 import MiniMaxH3PipelineConfig +from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter, records_to_table) +from fastvideo.dataset.dataloader.record_schema import text_only_record_creator +from fastvideo.dataset.dataloader.schema import pyarrow_schema_text_only +from fastvideo.fastvideo_args import FastVideoArgs +from fastvideo.models.loader.component_loader import PipelineComponentLoader +from fastvideo.pipelines import ForwardBatch +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_conditioning import MiniMaxH3ConditioningStage +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_input_preparation import MINIMAX_H3_KEYFRAMES_KEY +from fastvideo.utils import verify_model_config_and_directory + + +def _init_single_process_distributed(shard_index: int) -> None: + """Initialize the one-rank process groups required by component loaders. + + Concurrent shards on one host must not share a rendezvous port, so the + default port is offset by the shard index (an explicit ``MASTER_PORT`` + still wins). + """ + os.environ.setdefault("MASTER_ADDR", "127.0.0.1") + os.environ.setdefault("MASTER_PORT", str(29531 + shard_index)) + os.environ.setdefault("RANK", "0") + os.environ.setdefault("WORLD_SIZE", "1") + os.environ.setdefault("LOCAL_RANK", "0") + from fastvideo.distributed import maybe_init_distributed_environment_and_model_parallel + + maybe_init_distributed_environment_and_model_parallel(1, 1) + + +def _load_component( + name: str, + model_path: Path, + model_index: dict[str, Any], + fastvideo_args: FastVideoArgs, +) -> Any: + """Load one checkpoint component through the inference component registry.""" + transformers_or_diffusers, _ = model_index[name][:2] + return PipelineComponentLoader.load_module( + module_name=name, + component_model_path=str(model_path / name), + transformers_or_diffusers=transformers_or_diffusers, + fastvideo_args=fastvideo_args, + ) + + +def _load_shard_prompts( + prompts_file: Path, + shard_index: int, + num_shards: int, + jsonl_field: str | None = None, +) -> list[tuple[int, str, str]]: + """Return this shard's ``(global_line_index, prompt, text_name)`` triples. + + Default mode treats each non-empty line as one complete prompt document. + With ``jsonl_field`` set, each line is a JSON record; the prompt is taken + verbatim from that field (embedded newlines preserved — required by the + h3-t2va-condition-v1 contract, which forbids altering ``prompt_compiled``) + and the record's ``id`` becomes the parquet row id when present. + """ + entries: list[tuple[int, str, str]] = [] + with prompts_file.open(encoding="utf-8") as handle: + for index, line in enumerate(handle): + line = line.strip() + if not line: + continue + if jsonl_field is None: + entries.append((index, line, f"vidprom_{index:05d}")) + continue + record = json.loads(line) + prompt = record.get(jsonl_field) + if not isinstance(prompt, str) or not prompt: + raise ValueError(f"{prompts_file}:{index + 1}: empty or non-string field {jsonl_field!r}") + entries.append((index, prompt, str(record.get("id") or f"jsonl_{index:05d}"))) + return entries[shard_index::num_shards] + + +def _count_existing_rows(shard_dir: Path) -> int: + """Count rows already written under a shard directory (resume offset).""" + total = 0 + for parquet_path in sorted(shard_dir.rglob("*.parquet")): + total += pq.ParquetFile(parquet_path).metadata.num_rows + return total + + +def main(args: argparse.Namespace) -> None: + _init_single_process_distributed(args.shard_index) + + model_path = args.model_path.resolve() + if not model_path.is_dir(): + raise FileNotFoundError(f"MiniMax H3 model directory is missing at {model_path}") + model_index = verify_model_config_and_directory(str(model_path)) + + shard = _load_shard_prompts(args.prompts_file, args.shard_index, args.num_shards, args.jsonl_field) + shard_dir = args.output_dir / f"shard_{args.shard_index:02d}" + already_done = _count_existing_rows(shard_dir) if shard_dir.is_dir() else 0 + if already_done >= len(shard): + print(f"Shard {args.shard_index}/{args.num_shards}: all {len(shard)} rows already encoded") + return + todo = shard[already_done:] + if args.limit is not None: + todo = todo[:args.limit] + print(f"Shard {args.shard_index}/{args.num_shards}: {len(shard)} prompts total, " + f"{already_done} already on disk, encoding {len(todo)} now -> {shard_dir}") + + fastvideo_args = FastVideoArgs( + model_path=str(model_path), + pipeline_config=MiniMaxH3PipelineConfig(), + num_gpus=1, + tp_size=1, + sp_size=1, + hsdp_shard_dim=1, + use_fsdp_inference=False, + vae_cpu_offload=False, + text_encoder_cpu_offload=False, + ) + print("Loading MiniMax H3 tokenizer, processor, and Qwen3-VL encoder") + tokenizer = _load_component("tokenizer", model_path, model_index, fastvideo_args) + processor = _load_component("processor", model_path, model_index, fastvideo_args) + conditioner = _load_component("text_encoder", model_path, model_index, fastvideo_args) + stage = MiniMaxH3ConditioningStage( + conditioner=conditioner, + tokenizer=tokenizer, + processor=processor, + ) + + writer = ParquetDatasetWriter(out_dir=str(shard_dir), samples_per_file=args.samples_per_file) + records: list[dict[str, Any]] = [] + started = time.monotonic() + with torch.inference_mode(): + for done, (global_index, prompt, text_name) in enumerate(todo, start=1): + batch = ForwardBatch(data_type="video", prompt=prompt) + batch.extra[MINIMAX_H3_KEYFRAMES_KEY] = [] + batch = stage.forward(batch, fastvideo_args) + if not batch.prompt_embeds: + raise RuntimeError(f"MiniMax H3 conditioning returned no embedding for line {global_index}") + text_embedding = batch.prompt_embeds[0].squeeze(0).float().cpu().contiguous().numpy() + records.append( + text_only_record_creator( + text_name=text_name, + text_embedding=text_embedding, + caption=prompt, + )) + + if len(records) >= args.flush_every or done == len(todo): + writer.append_table(records_to_table(records, pyarrow_schema_text_only)) + records = [] + written = writer.flush(write_remainder=done == len(todo)) + rate = done / (time.monotonic() - started) + remaining = (len(todo) - done) / rate if rate > 0 else float("inf") + print(f"[shard {args.shard_index}] {done}/{len(todo)} encoded " + f"({rate:.2f} prompts/s, ~{remaining / 60:.0f} min left, flushed {written} rows)") + + print(f"Shard {args.shard_index} complete: {already_done + len(todo)}/{len(shard)} rows in {shard_dir}") + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--prompts-file", type=Path, required=True, help="one H3 prompt document per line") + parser.add_argument("--model-path", type=Path, required=True, help="MiniMax-H3 checkpoint directory") + parser.add_argument("--output-dir", type=Path, required=True, help="dataset root; shards write shard_XX/ under it") + parser.add_argument("--shard-index", type=int, default=0) + parser.add_argument("--num-shards", type=int, default=1) + parser.add_argument("--samples-per-file", type=int, default=64) + parser.add_argument("--flush-every", type=int, default=256, help="rows buffered between parquet flushes") + parser.add_argument("--limit", type=int, default=None, help="encode at most N prompts this run (smoke tests)") + parser.add_argument("--jsonl-field", + type=str, + default=None, + help="treat --prompts-file as JSONL and take the prompt verbatim from " + "this field (record 'id' becomes the row id when present)") + cli_args = parser.parse_args() + if not 0 <= cli_args.shard_index < cli_args.num_shards: + parser.error(f"--shard-index {cli_args.shard_index} must be in [0, {cli_args.num_shards})") + main(cli_args) diff --git a/fastvideo/tests/attention/test_compile_policy.py b/fastvideo/tests/attention/test_compile_policy.py new file mode 100644 index 0000000000..f9ccb2f4a6 --- /dev/null +++ b/fastvideo/tests/attention/test_compile_policy.py @@ -0,0 +1,19 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Attention compile-boundary policy tests.""" + +from fastvideo.attention.layer import _attention_compile_disabled + + +def test_attention_compile_is_enabled_by_default(monkeypatch) -> None: + monkeypatch.delenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", raising=False) + + assert not _attention_compile_disabled() + + +def test_attention_compile_escape_hatch(monkeypatch) -> None: + monkeypatch.setenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", "1") + + assert _attention_compile_disabled() + + monkeypatch.setenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", "0") + assert not _attention_compile_disabled() diff --git a/fastvideo/tests/attention/test_flash_attn_cute_custom_op.py b/fastvideo/tests/attention/test_flash_attn_cute_custom_op.py index a1d0a5167d..210aad800c 100644 --- a/fastvideo/tests/attention/test_flash_attn_cute_custom_op.py +++ b/fastvideo/tests/attention/test_flash_attn_cute_custom_op.py @@ -173,6 +173,37 @@ def test_flash_attn_func_parity_forward_backward(flash_attn_impls, dtype: torch. _assert_close(dv_test, dv_ref, dtype=dtype, is_grad=True) +def test_flash_attn_func_fullgraph_compile_backward(flash_attn_impls): + custom_flash_attn_func, _, _, _ = flash_attn_impls + + def loss(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor: + return custom_flash_attn_func(q, k, v).square().mean() + + torch.manual_seed(2) + base_inputs = [ + torch.randn( + 1, + 64, + 4, + 64, + device="cuda", + dtype=torch.bfloat16, + requires_grad=False, + ) for _ in range(3) + ] + eager_inputs = [tensor.clone().requires_grad_(True) for tensor in base_inputs] + compiled_inputs = [tensor.clone().requires_grad_(True) for tensor in base_inputs] + + eager_loss = loss(*eager_inputs) + eager_grads = torch.autograd.grad(eager_loss, eager_inputs) + compiled_loss = torch.compile(loss, fullgraph=True)(*compiled_inputs) + compiled_grads = torch.autograd.grad(compiled_loss, compiled_inputs) + + _assert_close(compiled_loss, eager_loss, dtype=torch.bfloat16) + for compiled_grad, eager_grad in zip(compiled_grads, eager_grads): + _assert_close(compiled_grad, eager_grad, dtype=torch.bfloat16, 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(flash_attn_impls, dtype: torch.dtype, causal: bool): diff --git a/fastvideo/tests/attention/test_vsa_h3_inference_metadata_parity.py b/fastvideo/tests/attention/test_vsa_h3_inference_metadata_parity.py new file mode 100644 index 0000000000..d20812e285 --- /dev/null +++ b/fastvideo/tests/attention/test_vsa_h3_inference_metadata_parity.py @@ -0,0 +1,170 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU checks that the VSA-H3 inference (validation/denoising-stage) metadata +path builds exactly what the training path builds, stays in-bounds, and that +malformed geometry fails synchronously instead of as an async kernel fault. + +Shapes mirror the v7 DMD2 validation request that motivated this test: +768x1344, 124 frames -> video latents (37, 48, 84), 207 audio latents, +patch (1, 2, 2), 3-step DMD ladder, 90% sparsity (jobs 2307/2321).""" + +import math + +import pytest +import torch + +from fastvideo.attention.backends.video_sparse_attn_h3 import (_TILE_ELEMS, MiniMaxH3VSAMetadataBuilder, + _build_block_mask, _h3_tile_geometry, + _validate_h3_tile_geometry) +from fastvideo.pipelines.basic.minimax_h3.packing import (MINIMAX_H3_TEXT_TAG, audio_latent_num_frames, + build_packed_sequence, video_latent_num_frames) +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_denoising import _h3_vsa_prefix_segments + +_CPU = torch.device("cpu") +_PATCH = (1, 2, 2) +_NUM_FRAMES = 124 # -> 37 latent frames, 207 audio latents +_LATENT = (video_latent_num_frames(_NUM_FRAMES), 768 // 16, 1344 // 16) +_NUM_AUDIO = audio_latent_num_frames(_NUM_FRAMES) +_SPARSITY = 0.9 +_DMD_STEPS = 3 + +_TEXT_LENS = [7, 100, 255, 256, 257, 500] + + +def _layout(text_len: int, anchors: tuple[str, ...] = ()): + tags = torch.full((text_len, ), MINIMAX_H3_TEXT_TAG, dtype=torch.long) + return build_packed_sequence(tags, *_LATENT, _NUM_AUDIO, _PATCH, anchors) + + +def _inference_metadata(builder, layout, step_index: int, sparsity: float = _SPARSITY): + """Exactly the MiniMaxH3DenoisingStage.forward calling convention.""" + return builder.build( + current_timestep=step_index, + raw_latent_shape=(layout.num_video_latent_frames, layout.latent_height, layout.latent_width), + patch_size=_PATCH, + VSA_sparsity=sparsity, + prefix_segments=_h3_vsa_prefix_segments(layout, _PATCH), + device=_CPU, + exempt=True, + dense_layers=(), + ) + + +def _training_metadata(builder, layout, sparsity: float = _SPARSITY): + """Exactly the MiniMaxH3Model._maybe_build_vsa_metadata calling convention.""" + return builder.build( + current_timestep=0, + raw_latent_shape=(layout.num_video_latent_frames, layout.latent_height, layout.latent_width), + patch_size=_PATCH, + VSA_sparsity=sparsity, + prefix_segments=_h3_vsa_prefix_segments(layout, _PATCH), + device=_CPU, + ) + + +def _assert_in_bounds(meta, layout, tag: str): + n_tiles = meta.num_prefix_tiles + meta.num_video_tiles + sizes = meta.variable_block_sizes + assert sizes.numel() == n_tiles, tag + assert int(sizes.min()) >= 1 and int(sizes.max()) <= _TILE_ELEMS, tag + assert int(sizes.sum()) == meta.total_seq_length == layout.sequence_length, tag + idx = meta.untile_combined_index + assert idx.numel() == meta.total_seq_length, tag + assert int(idx.min()) >= 0 and int(idx.max()) < n_tiles * _TILE_ELEMS, tag + assert idx.unique().numel() == idx.numel(), f"{tag}: untile index must be injective" + assert bool((idx % _TILE_ELEMS < sizes[idx // _TILE_ELEMS]).all()), tag + + +@pytest.mark.parametrize("text_len", _TEXT_LENS) +def test_inference_metadata_matches_training(text_len): + layout = _layout(text_len) + assert _h3_vsa_prefix_segments(layout, _PATCH) == (text_len, 0, _NUM_AUDIO * 2) + + meta_train = _training_metadata(MiniMaxH3VSAMetadataBuilder(), layout) + _assert_in_bounds(meta_train, layout, f"train text={text_len}") + + infer_builder = MiniMaxH3VSAMetadataBuilder() # one builder per denoise loop, as the stage does + for step in range(_DMD_STEPS): + meta_inf = _inference_metadata(infer_builder, layout, step) + _assert_in_bounds(meta_inf, layout, f"infer text={text_len} step={step}") + assert meta_inf.total_seq_length == meta_train.total_seq_length + assert meta_inf.num_prefix_tiles == meta_train.num_prefix_tiles + assert meta_inf.num_video_tiles == meta_train.num_video_tiles + assert meta_inf.exempt == meta_train.exempt + assert meta_inf.dense_layers == meta_train.dense_layers + assert meta_inf.VSA_sparsity == meta_train.VSA_sparsity + assert torch.equal(meta_inf.variable_block_sizes, meta_train.variable_block_sizes) + assert torch.equal(meta_inf.untile_combined_index, meta_train.untile_combined_index) + + +def test_keyframe_conditioned_layout_in_bounds(): + """Image-conditioned validation adds condition keyframe rows to the prefix.""" + layout = _layout(100, anchors=("first", )) + rows_per_frame = (_LATENT[1] // _PATCH[1]) * (_LATENT[2] // _PATCH[2]) + assert _h3_vsa_prefix_segments(layout, _PATCH) == (100, rows_per_frame, _NUM_AUDIO * 2) + meta = _inference_metadata(MiniMaxH3VSAMetadataBuilder(), layout, 0) + _assert_in_bounds(meta, layout, "keyframe-conditioned") + + +def test_route_a_expansion_in_bounds(): + """The 256->64 route-A remap the Triton fallback consumes stays in-bounds.""" + try: + from fastvideo_kernel import block_sparse_attn_256 + except Exception as exc: # triton driver probing raises RuntimeError on GPU-less hosts + pytest.skip(f"fastvideo_kernel unavailable here: {exc}") + layout = _layout(100) + meta = _inference_metadata(MiniMaxH3VSAMetadataBuilder(), layout, 0) + n_tiles = meta.variable_block_sizes.numel() + scores = torch.randn(1, 4, n_tiles, n_tiles) + mask = _build_block_mask(scores, meta.num_prefix_tiles, meta.num_video_tiles, _SPARSITY, exempt=True) + mask64, sizes64 = block_sparse_attn_256._expand_mask_and_sizes_256_to_64(mask, meta.variable_block_sizes) + assert mask64.shape[-2:] == (4 * n_tiles, 4 * n_tiles) + assert sizes64.numel() == 4 * n_tiles + assert int(sizes64.min()) >= 0 and int(sizes64.max()) <= 64 + assert int(sizes64.sum()) == meta.total_seq_length + per_tile = sizes64.view(n_tiles, 4) + assert bool((per_tile[:, 0] > 0).all()), "every logical tile keeps at least one valid 64-block" + assert bool((per_tile[:, :-1] >= per_tile[:, 1:]).all()), "child sizes must be non-increasing" + assert torch.equal(per_tile.sum(dim=1), meta.variable_block_sizes.to(per_tile.dtype)) + + +def test_geometry_guard_rejects_corruption(): + """The synchronous guard must catch what would otherwise be an async fault.""" + prefix = (100, _NUM_AUDIO * 2) + dit_shape = tuple(d // p for d, p in zip(_LATENT, _PATCH, strict=True)) + (_, sizes, untile, _, _) = _h3_tile_geometry(prefix, dit_shape, _CPU) + + with pytest.raises(ValueError, match="tile sizes out of bounds"): + bad = sizes.clone() + bad[0] = _TILE_ELEMS + 1 + _validate_h3_tile_geometry(prefix, dit_shape, bad, untile) + with pytest.raises(ValueError, match="tile sizes out of bounds"): + bad = sizes.clone() + bad[-1] += 1 # sum mismatch + _validate_h3_tile_geometry(prefix, dit_shape, bad, untile) + with pytest.raises(ValueError, match="untile index"): + _validate_h3_tile_geometry(prefix, dit_shape, sizes, untile[:-1]) + with pytest.raises(ValueError, match="injective"): + bad = untile.clone() + bad[1] = int(bad[0]) # duplicate slot + _validate_h3_tile_geometry(prefix, dit_shape, sizes, bad) + with pytest.raises(ValueError, match="injective"): + bad = untile.clone() + bad[0] = sizes.numel() * _TILE_ELEMS # beyond the padded buffer + _validate_h3_tile_geometry(prefix, dit_shape, sizes, bad) + partial = int((sizes < _TILE_ELEMS).nonzero()[0]) + with pytest.raises(ValueError, match="injective"): + bad = untile.clone() + bad[0] = partial * _TILE_ELEMS + int(sizes[partial]) # first pad slot + _validate_h3_tile_geometry(prefix, dit_shape, sizes, bad) + + _validate_h3_tile_geometry(prefix, dit_shape, sizes, untile) + assert int(sizes.sum()) == sum(prefix) + math.prod(dit_shape) + + +if __name__ == "__main__": + for _text_len in _TEXT_LENS: + test_inference_metadata_matches_training(_text_len) + test_keyframe_conditioned_layout_in_bounds() + test_route_a_expansion_in_bounds() + test_geometry_guard_rejects_corruption() + print("all VSA-H3 inference-metadata parity checks passed") diff --git a/fastvideo/tests/attention/test_vsa_h3_metadata.py b/fastvideo/tests/attention/test_vsa_h3_metadata.py index 24292f26b4..a9879c23a3 100644 --- a/fastvideo/tests/attention/test_vsa_h3_metadata.py +++ b/fastvideo/tests/attention/test_vsa_h3_metadata.py @@ -16,12 +16,7 @@ _720P = dict(raw_latent_shape=(30, 44, 80), patch_size=(1, 2, 2), prefix_segments=(512, 1760, 400)) _TINY = dict(raw_latent_shape=(8, 8, 12), patch_size=(1, 2, 2), prefix_segments=(7, 5, 3)) -# (4,4,4) coverage: dit grid (9, 10, 13) is ragged in all three dims -# (t: 4+4+1, h: 4+4+2, w: 4+4+4+1) and every prefix segment leaves a -# partial tail tile at 64 (70 -> 64+6, 5 -> 5, 130 -> 64+64+2). _TINY64 = dict(raw_latent_shape=(9, 20, 26), patch_size=(1, 2, 2), prefix_segments=(70, 5, 130)) -# production-shape request: 768x1344, 124 frames -> latents (37, 48, 84), -# patch (1,2,2) -> token grid (37, 24, 42); text 300 + audio 414 rows. _PROD = dict(raw_latent_shape=(37, 48, 84), patch_size=(1, 2, 2), prefix_segments=(300, 0, 414)) _CPU = torch.device("cpu") @@ -69,16 +64,12 @@ def test_geometry_720p(): assert meta.num_prefix_tiles == 2 + 7 + 2 assert meta.num_video_tiles == 8 * 3 * 5 assert int(meta.variable_block_sizes.sum()) == seq - # (permutation coverage of [0, seq) is implied by the roundtrip below: - # untile_combined_index scatters seq distinct rows and recovers all of x) - # segment purity: no prefix tile straddles a segment boundary boundaries = [512, 512 + 1760, 512 + 1760 + 400] start = 0 for size in meta.variable_block_sizes[:meta.num_prefix_tiles].tolist(): end = start + size assert all(not (start < b < end) for b in boundaries), (start, end) start = end - # untile(tile(x)) == x x = torch.randn(1, seq, 2, 4) buf = _impl().tile(x, meta) assert buf.shape[1] == meta.variable_block_sizes.numel() * _TILE_ELEMS @@ -143,9 +134,28 @@ def test_prefix_queries_stay_dense_at_high_sparsity(): "video rows should actually be sparse at 75%" -# --------------------------------------------------------------------------- -# 64-token (4,4,4) tile geometry -# --------------------------------------------------------------------------- +def test_tile_under_grad_does_not_reuse_shared_buffer(): + """Training forwards need fresh tile buffers; in-place reuse of the shared + holder would trip autograd's saved-tensor version check at backward.""" + meta = _build(_TINY) + impl = _impl() + seq = meta.total_seq_length + + x1 = torch.randn(1, seq, 2, 8, requires_grad=True) + x2 = torch.randn(1, seq, 2, 8, requires_grad=True) + buf1 = impl.tile(x1, meta) + saved = (buf1 * buf1).sum() # saves buf1 for backward, like the kernel + buf2 = impl.tile(x2, meta) + assert buf1 is not buf2 + saved.backward() # raises "modified by an inplace operation" on reuse + assert x1.grad is not None and x2.grad is None + + with torch.no_grad(): + y1 = impl.tile(x1.detach(), meta) + y2 = impl.tile(x2.detach(), meta) + assert y1 is y2 + + def test_geometry_tile64_ragged_tails(): @@ -163,7 +173,6 @@ def test_geometry_tile64_ragged_tails(): assert int(meta.variable_block_sizes.max()) <= 64 assert meta.variable_block_sizes[:meta.num_prefix_tiles].tolist() == [64, 6, 5, 64, 64, 2] - # per-tile valid sizes: product of the per-dim clamped tails expected = torch.tensor([ min(4, t - 4 * tt) * min(4, h - 4 * hh) * min(4, w - 4 * ww) for tt in range(n_t) for hh in range(n_h) for ww in range(n_w) @@ -172,16 +181,13 @@ def test_geometry_tile64_ragged_tails(): assert torch.equal(meta.variable_block_sizes[meta.num_prefix_tiles:], expected) assert int(expected.min()) == 1 * 2 * 1 # the (t,h,w) ragged corner - # every packed video row lands in the 3D tile its (t,h,w) coordinate says idx = meta.untile_combined_index row = torch.arange(t * h * w) row_t, row_h, row_w = row // (h * w), (row // w) % h, row % w expected_tile = meta.num_prefix_tiles + ((row_t // 4) * n_h + row_h // 4) * n_w + row_w // 4 assert torch.equal(idx[prefix_len:] // 64, expected_tile) - # and in a non-pad slot of that tile assert bool((idx % 64 < meta.variable_block_sizes[idx // 64]).all()) - # untile(tile(x)) == x on the 64-wide padded buffer x = torch.randn(1, seq, 2, 4) buf = _impl().tile(x, meta) assert buf.shape[1] == meta.variable_block_sizes.numel() * 64 @@ -198,7 +204,6 @@ def test_geometry_tile64_production_shape(): sizes_vid = meta64.variable_block_sizes[meta64.num_prefix_tiles:] assert int(sizes_vid.max()) == 64 and int(sizes_vid.min()) == 1 * 4 * 2 # (t, w) ragged corner - # same packed sequence under the default 256 geometry, fewer tiles meta256 = _build(_PROD) assert meta256.tile_elems == _TILE_ELEMS assert meta256.num_prefix_tiles == 2 + 2 @@ -236,7 +241,6 @@ def test_geometry_guard_enforces_tile64_bound(): sizes[0] = 65 with pytest.raises(ValueError, match="tile sizes out of bounds"): _validate_h3_tile_geometry(prefix, dit_shape, sizes, meta.untile_combined_index, 64) - # the untampered tile-64 geometry passes its own bound _validate_h3_tile_geometry(prefix, dit_shape, meta.variable_block_sizes, meta.untile_combined_index, 64) @@ -251,6 +255,7 @@ def test_builder_rejects_unknown_tile_size(): test_mask_policy() test_sparsity_zero_matches_dense_sdpa() test_prefix_queries_stay_dense_at_high_sparsity() + test_tile_under_grad_does_not_reuse_shared_buffer() test_geometry_tile64_ragged_tails() test_geometry_tile64_production_shape() test_sparsity_zero_matches_dense_sdpa_tile64() diff --git a/fastvideo/tests/dataset/test_exact_shape_bucket_sampler.py b/fastvideo/tests/dataset/test_exact_shape_bucket_sampler.py new file mode 100644 index 0000000000..34997e3866 --- /dev/null +++ b/fastvideo/tests/dataset/test_exact_shape_bucket_sampler.py @@ -0,0 +1,114 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Exact-shape global microbatch scheduling contracts.""" + +from pathlib import Path + +import pytest + +from fastvideo.dataset.parquet_dataset_map_style import ( + DP_SP_BatchSampler, + shape_bucket_ids_from_parquet_files, +) +from fastvideo.dataset.shape_bucket import parse_video_shape_bucket_id + + +@pytest.mark.parametrize( + "bucket_id", + [ + "bucket=0x768-124f", + "bucket=1344x0-124f", + "bucket=1344x768-0f", + "bucket=1344X768-124f", + "bucket=1344x768-124F", + "bucket=01344x768-124f", + "1344x768-124f", + "bucket=1344x768", + "bucket=1344x768-124f-extra", + ], +) +def test_portable_bucket_id_rejects_noncanonical_spelling(bucket_id: str) -> None: + with pytest.raises(ValueError, match="bucket=x"): + parse_video_shape_bucket_id(bucket_id) + + +def test_bucket_ids_expand_from_canonical_parquet_ancestors() -> None: + files = [ + "/shared/source-a/data/bucket=1344x768-124f/part-0.parquet", + "/shared/source-b/data/bucket=768x1344-362f/part-1.parquet", + ] + assert shape_bucket_ids_from_parquet_files(files, [2, 1]) == [ + "bucket=1344x768-124f", + "bucket=1344x768-124f", + "bucket=768x1344-362f", + ] + + with pytest.raises(ValueError, match="exactly one ancestor"): + shape_bucket_ids_from_parquet_files(["/shared/data/part.parquet"], [1]) + malformed = str(Path("/shared/data/bucket=1344-768-124f/part.parquet")) + with pytest.raises(ValueError, match="positive decimal"): + shape_bucket_ids_from_parquet_files([malformed], [1]) + + +def _rank_samplers(bucket_ids: list[str], *, seed: int = 17) -> list[DP_SP_BatchSampler]: + return [ + DP_SP_BatchSampler( + batch_size=1, + dataset_size=len(bucket_ids), + num_sp_groups=32, + sp_world_size=1, + global_rank=rank, + drop_last=True, + seed=seed, + sample_bucket_ids=bucket_ids, + ) for rank in range(32) + ] + + +def test_rare_bucket_is_padded_not_dropped_and_all_ranks_share_schedule() -> None: + rare = "bucket=480x832-124f" + common = "bucket=1344x768-362f" + bucket_ids = [rare] * 7 + [common] * 64 + samplers = _rank_samplers(bucket_ids) + + assert all(sampler.bucket_schedule == samplers[0].bucket_schedule for sampler in samplers) + assert samplers[0].bucket_schedule is not None + assert samplers[0].bucket_schedule.count(rare) == 1 + assert samplers[0].bucket_schedule.count(common) == 2 + assert samplers[0].bucket_padding == {rare: 25} + assert samplers[0].num_padded_samples == 25 + + batches_by_rank = [list(sampler) for sampler in samplers] + observed_originals: set[int] = set() + for step, scheduled_bucket in enumerate(samplers[0].bucket_schedule): + step_indices = [batches_by_rank[rank][step][0] for rank in range(32)] + assert {bucket_ids[index] for index in step_indices} == {scheduled_bucket} + observed_originals.update(step_indices) + if scheduled_bucket == rare: + assert set(step_indices) == set(range(7)) + assert len(step_indices) == 32 + + assert observed_originals == set(range(len(bucket_ids))) + + +def test_bucket_schedule_is_seeded_and_sp_ranks_share_indices() -> None: + bucket_ids = (["bucket=1344x768-124f"] * 16 + ["bucket=768x1344-362f"] * 16) + + def sampler(rank: int, seed: int) -> DP_SP_BatchSampler: + return DP_SP_BatchSampler( + batch_size=1, + dataset_size=len(bucket_ids), + num_sp_groups=4, + sp_world_size=2, + global_rank=rank, + drop_last=True, + seed=seed, + sample_bucket_ids=bucket_ids, + ) + + first = [list(sampler(rank, 9)) for rank in range(8)] + replay = [list(sampler(rank, 9)) for rank in range(8)] + changed = [list(sampler(rank, 10)) for rank in range(8)] + assert first == replay + assert first != changed + for sp_leader in range(0, 8, 2): + assert first[sp_leader] == first[sp_leader + 1] diff --git a/fastvideo/tests/dataset/test_parquet_dataset_map_style.py b/fastvideo/tests/dataset/test_parquet_dataset_map_style.py index e899cd0ae6..22352c7615 100644 --- a/fastvideo/tests/dataset/test_parquet_dataset_map_style.py +++ b/fastvideo/tests/dataset/test_parquet_dataset_map_style.py @@ -5,15 +5,20 @@ import pickle +import pyarrow as pa +import pyarrow.parquet as pq import pytest from fastvideo.dataset import parquet_dataset_map_style as parquet_dataset -from fastvideo.dataset.parquet_dataset_map_style import _parse_data_path_specs +from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2va, pyarrow_schema_text_only +from fastvideo.dataset.parquet_dataset_map_style import ( + _parse_data_path_specs, + LatentsParquetMapStyleDataset, + read_row_from_parquet_file, +) def test_parse_data_path_specs_accepts_old_repeat_string() -> None: - # Dataset parsing keeps compatibility with the old "path:repeat" string - # form used by existing training configs. assert _parse_data_path_specs("data/path1:2,data/path2:1") == [ ("data/path1", 2), ("data/path2", 1), @@ -21,8 +26,6 @@ def test_parse_data_path_specs_accepts_old_repeat_string() -> None: def test_parse_data_path_specs_accepts_yaml_mapping() -> None: - # New YAML mapping form should reach the dataset layer as path -> repeat - # and parse to the same internal spec representation. assert _parse_data_path_specs({ "data/path1": 1, "data/path2": 2, @@ -51,8 +54,6 @@ def test_parse_data_path_specs_drops_non_positive_repeats() -> None: def test_parse_data_path_specs_rejects_malformed_repeats() -> None: - # Repeat counts must parse as ints; anything else is an explicit error - # rather than a silently mis-weighted dataset. with pytest.raises(ValueError): _parse_data_path_specs("data/a:abc") with pytest.raises(ValueError): @@ -78,8 +79,6 @@ def _write_root_cache(dataset_root, filename: str, length: int) -> str: def test_get_parquet_files_and_length_repeats_single_path(tmp_path, monkeypatch) -> None: - # get_parquet_files_and_length applies repeat counts after reading the - # per-root parquet cache, so a repeated root duplicates both names and rows. dataset_root = tmp_path / "dataset" parquet_file = _write_root_cache(dataset_root, "sample.parquet", 7) @@ -95,8 +94,6 @@ def test_get_parquet_files_and_length_repeats_single_path(tmp_path, monkeypatch) def test_get_parquet_files_and_length_mixes_roots_and_resorts(tmp_path, monkeypatch) -> None: - # Multiple roots are expanded per repeat count and then globally re-sorted - # by filename, so the mix order is independent of the mapping order. root_a = tmp_path / "dataset_a" root_b = tmp_path / "dataset_b" file_a = _write_root_cache(root_a, "a.parquet", 5) @@ -115,7 +112,70 @@ def test_get_parquet_files_and_length_mixes_roots_and_resorts(tmp_path, monkeypa def test_get_parquet_files_and_length_raises_when_all_repeats_dropped() -> None: - # Zero/negative repeats are dropped at parse time; if that leaves nothing - # to read, the mix branch fails loudly instead of yielding an empty dataset. with pytest.raises(FileNotFoundError): parquet_dataset.get_parquet_files_and_length({"data/a": 0}) + + +def test_read_row_projects_text_columns_from_t2va_superset(tmp_path) -> None: + parquet_path = tmp_path / "sample.parquet" + row = { + "id": ["sample-0"], + "vae_latent_bytes": [b"video-must-not-be-read"], + "vae_latent_shape": [[24, 2, 4, 4]], + "vae_latent_dtype": ["float32"], + "audio_latent_bytes": [b"audio-must-not-be-read"], + "audio_latent_shape": [[2, 32, 8]], + "audio_latent_dtype": ["float32"], + "text_embedding_bytes": [b"text"], + "text_embedding_shape": [[1, 4]], + "text_embedding_dtype": ["float32"], + "file_name": ["sample.mp4"], + "caption": ["prompt"], + "media_type": ["video"], + "width": [64], + "height": [64], + "num_frames": [5], + "duration_sec": [5.0 / 24.0], + "fps": [24.0], + "audio_sample_rate": [32_000], + } + pq.write_table(pa.Table.from_pydict(row, schema=pyarrow_schema_t2va), parquet_path) + text_columns = [ + "id", + "text_embedding_bytes", + "text_embedding_shape", + "text_embedding_dtype", + "caption", + ] + + projected = read_row_from_parquet_file([str(parquet_path)], 0, [1], columns=text_columns) + + assert projected == { + "id": "sample-0", + "text_embedding_bytes": b"text", + "text_embedding_shape": [1, 4], + "text_embedding_dtype": "float32", + "caption": "prompt", + } + + +def test_dataset_projects_its_declared_schema_columns(monkeypatch) -> None: + observed = {} + + def fake_read(parquet_files, global_row_idx, lengths, columns=None): + observed["columns"] = columns + return {"id": "sample-0"} + + monkeypatch.setattr(parquet_dataset, "read_row_from_parquet_file", fake_read) + monkeypatch.setattr(parquet_dataset, "collate_rows_from_parquet_schema", lambda rows, *args, **kwargs: rows[0]) + dataset = LatentsParquetMapStyleDataset.__new__(LatentsParquetMapStyleDataset) + dataset.parquet_files = ("unused.parquet", ) + dataset.lengths = (1, ) + dataset.parquet_schema = pyarrow_schema_text_only + dataset.text_padding_length = 512 + dataset.cfg_rate = 0.0 + dataset.seed = 42 + dataset.sample_bucket_ids = None + + assert dataset.__getitems__([0]) == {"id": "sample-0", "_sample_index": 0} + assert observed["columns"] == pyarrow_schema_text_only.names diff --git a/fastvideo/tests/dataset/test_validation_dataset.py b/fastvideo/tests/dataset/test_validation_dataset.py new file mode 100644 index 0000000000..e511164331 --- /dev/null +++ b/fastvideo/tests/dataset/test_validation_dataset.py @@ -0,0 +1,33 @@ +# SPDX-License-Identifier: Apache-2.0 +import json + +import pytest + +from fastvideo.dataset.validation_dataset import ValidationDataset + + +@pytest.mark.parametrize("wrapped", [False, True]) +def test_validation_json_accepts_array_and_legacy_data_wrapper( + tmp_path, + monkeypatch: pytest.MonkeyPatch, + wrapped: bool, +) -> None: + rows = [{ + "caption": "A held-out prompt", + "source": "source-a", + "sample_id": "sample-1", + "width": 128, + "height": 80, + "num_frames": 39, + }] + document = {"data": rows} if wrapped else rows + manifest = tmp_path / "validation.json" + manifest.write_text(json.dumps(document), encoding="utf-8") + monkeypatch.setattr("fastvideo.dataset.validation_dataset.get_world_rank", lambda: 0) + monkeypatch.setattr("fastvideo.dataset.validation_dataset.get_world_size", lambda: 1) + monkeypatch.setattr("fastvideo.dataset.validation_dataset.get_sp_world_size", lambda: 1) + + dataset = ValidationDataset(str(manifest)) + + assert dataset.all_samples == rows + assert list(dataset)[0]["prompt"] == rows[0]["caption"] diff --git a/fastvideo/tests/loader/test_shard_cache.py b/fastvideo/tests/loader/test_shard_cache.py new file mode 100644 index 0000000000..68784710fb --- /dev/null +++ b/fastvideo/tests/loader/test_shard_cache.py @@ -0,0 +1,215 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU contract tests for the sharded base-weight cache. + +Runs on a single-process gloo group with a (1, 1) CPU device mesh: DTensor +round-trip through write/load, the FQN reconciliation matrix (allowed +zero-init params, disallowed extras, shape mismatches), and the +never-fail-the-run contract. +""" + +import os + +import pytest +import torch +import torch.distributed as dist +import torch.nn as nn +from torch.distributed.device_mesh import init_device_mesh +from torch.distributed.tensor import Replicate, Shard, distribute_tensor + +from fastvideo.models.loader.shard_cache import ( + ShardCacheContext, + try_load_from_shard_cache, + write_shard_cache, +) + + +@pytest.fixture(scope="module") +def cpu_mesh(): + os.environ.setdefault("MASTER_ADDR", "127.0.0.1") + os.environ.setdefault("MASTER_PORT", "29581") + if not dist.is_initialized(): + dist.init_process_group("gloo", rank=0, world_size=1) + return init_device_mesh("cpu", (1, 1), mesh_dim_names=("replicate", "shard")) + + +def _make_model( + cpu_mesh, + *, + extra_param: str | None = None, + weight_rows: int = 8, + dtype: torch.dtype = torch.float32, +) -> nn.Module: + model = nn.Module() + placements = (Replicate(), Shard(0)) + weight = distribute_tensor(torch.randn(weight_rows, 4, dtype=dtype), cpu_mesh, placements) + bias = distribute_tensor(torch.randn(weight_rows, dtype=dtype), cpu_mesh, placements) + model.register_parameter("weight", nn.Parameter(weight)) + model.register_parameter("bias", nn.Parameter(bias)) + model.register_buffer("scale", torch.full((1, ), 2.0)) + if extra_param is not None: + extra = distribute_tensor(torch.randn(4, 4), cpu_mesh, placements) + model.register_parameter(extra_param.replace(".", "_"), nn.Parameter(extra)) + model.reverse_param_names_mapping = {"weight": ("hf.weight", None, None)} + return model + + +def _ctx(tmp_path) -> ShardCacheContext: + return ShardCacheContext(entry_dir=tmp_path / "entry", key="testkey", shard_index=0, num_shards=1, is_writer=True) + + +def test_round_trip_restores_tensors_and_reverse_mapping(cpu_mesh, tmp_path): + src = _make_model(cpu_mesh) + ctx = _ctx(tmp_path) + write_shard_cache(src, ctx) + + dst = _make_model(cpu_mesh) + with torch.no_grad(): + dst.weight.mul_(0) + dst.bias.mul_(0) + dst.reverse_param_names_mapping = {} + assert try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + assert torch.equal(dst.weight.to_local(), src.weight.to_local()) + assert torch.equal(dst.bias.to_local(), src.bias.to_local()) + assert torch.equal(dst.scale, src.scale) + assert dst.reverse_param_names_mapping == {"weight": ("hf.weight", None, None)} + + +def test_allowed_new_param_zero_inits_on_hit(cpu_mesh, tmp_path): + src = _make_model(cpu_mesh) + ctx = _ctx(tmp_path) + write_shard_cache(src, ctx) + + dst = _make_model(cpu_mesh) + gate = distribute_tensor(torch.randn(4, 4), cpu_mesh, (Replicate(), Shard(0))) + dst.register_parameter("to_gate_compress", nn.Parameter(gate)) + assert try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + assert torch.equal(dst.to_gate_compress.to_local(), torch.zeros(4, 4)) + + +def test_disallowed_missing_param_misses_without_mutation(cpu_mesh, tmp_path): + src = _make_model(cpu_mesh) + ctx = _ctx(tmp_path) + write_shard_cache(src, ctx) + + dst = _make_model(cpu_mesh) + mystery = distribute_tensor(torch.randn(4, 4), cpu_mesh, (Replicate(), Shard(0))) + dst.register_parameter("mystery", nn.Parameter(mystery)) + before = dst.weight.to_local().clone() + assert not try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + assert torch.equal(dst.weight.to_local(), before) + + +def test_shape_mismatch_misses(cpu_mesh, tmp_path): + src = _make_model(cpu_mesh) + ctx = _ctx(tmp_path) + write_shard_cache(src, ctx) + + dst = _make_model(cpu_mesh, weight_rows=9) + assert not try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + + +def test_missing_entry_misses_cleanly(cpu_mesh, tmp_path): + dst = _make_model(cpu_mesh) + ctx = ShardCacheContext(entry_dir=tmp_path / "absent", key="k2", shard_index=0, num_shards=1, is_writer=True) + assert not try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + + +def test_ac_wrapped_buffer_stays_buffer_on_cache_hit(cpu_mesh, tmp_path): + """Warm-path counterpart of the cold-path AC-prefix fix in fsdp_load. + + Under pre-FSDP activation checkpointing the model handed to + ``try_load_from_shard_cache`` is already checkpoint-wrapped: + ``state_dict()`` (and manifest) keys are clean, but raw + ``named_buffers()`` keys carry the ``_checkpoint_wrapped_module.`` + segment. The buffer-membership test must compare canonical names, or a + cached persistent buffer inside a wrapped block is reassigned as an + ``nn.Parameter`` by ``load_state_dict(assign=True)`` — on warm boots + only, silently diverging from the (already fixed) cold-boot path. + """ + from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import ( + checkpoint_wrapper, ) + + def _block_model() -> nn.Module: + model = nn.Module() + block = nn.Module() + weight = distribute_tensor(torch.randn(8, 4), cpu_mesh, (Replicate(), Shard(0))) + block.register_parameter("weight", nn.Parameter(weight)) + block.register_buffer("gain", torch.full((4, ), 3.0)) + model.block = block + model.reverse_param_names_mapping = {} + return model + + src = _block_model() + src.block = checkpoint_wrapper(src.block) + ctx = _ctx(tmp_path) + write_shard_cache(src, ctx) + assert "block.gain" in src.state_dict() + + dst = _block_model() + dst.block = checkpoint_wrapper(dst.block) + with torch.no_grad(): + dst.block.weight.mul_(0) + dst.block.gain.mul_(0) + assert try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + + buffer_names = {name for name, _ in dst.named_buffers()} + parameter_names = {name for name, _ in dst.named_parameters()} + assert "block._checkpoint_wrapped_module.gain" in buffer_names + assert not any(name.endswith("gain") for name in parameter_names) + assert torch.equal(dst.block.gain, torch.full((4, ), 3.0)) + assert torch.equal(dst.block.weight.to_local(), src.block.weight.to_local()) + + +def test_nonzero_rank_writer_emits_manifest_for_per_node_roots(cpu_mesh, tmp_path, monkeypatch): + """Every shard writer must write the manifest, not just global rank 0. + + With FASTVIDEO_WEIGHT_SHARD_CACHE_PER_NODE=1 each node keeps its own copy + of the entry; ranks on non-head nodes write their shard files into their + node's tmpfs but (pre-fix) never a manifest, so try_load_from_shard_cache + failed its `manifest.json` existence check there and the all-rank vote + turned every multi-node warm boot into a full load (observed on the + h3-compile-ab job 2592 b2_on_warm leg: shard4-7 present on the second + tray, manifest.json absent). + """ + import fastvideo.models.loader.shard_cache as sc + + src = _make_model(cpu_mesh) + ctx = _ctx(tmp_path) + monkeypatch.setattr(sc.dist, "get_rank", lambda: 4) + write_shard_cache(src, ctx) + assert (ctx.entry_dir / "manifest.json").is_file() + + dst = _make_model(cpu_mesh) + with torch.no_grad(): + dst.weight.mul_(0) + dst.reverse_param_names_mapping = {} + assert try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + assert torch.equal(dst.weight.to_local(), src.weight.to_local()) + + +def test_model_selected_dtype_rejects_stale_cache_and_hits_fresh_cache(cpu_mesh, tmp_path): + stale = _make_model(cpu_mesh, dtype=torch.bfloat16) + stale_ctx = _ctx(tmp_path / "stale") + write_shard_cache(stale, stale_ctx) + + def select_dtype(name: str, default: torch.dtype) -> torch.dtype: + return torch.float32 if name == "weight" else default + + destination = _make_model(cpu_mesh, dtype=torch.bfloat16) + destination._get_parameter_dtype = select_dtype + assert not try_load_from_shard_cache(destination, stale_ctx, torch.device("cpu")) + + fresh = _make_model(cpu_mesh, dtype=torch.bfloat16) + fresh_weight = distribute_tensor( + torch.randn(8, 4, dtype=torch.float32), + cpu_mesh, + (Replicate(), Shard(0)), + ) + fresh.weight = nn.Parameter(fresh_weight) + fresh_ctx = _ctx(tmp_path / "fresh") + write_shard_cache(fresh, fresh_ctx) + + destination = _make_model(cpu_mesh, dtype=torch.bfloat16) + destination._get_parameter_dtype = select_dtype + assert try_load_from_shard_cache(destination, fresh_ctx, torch.device("cpu")) + assert destination.weight.dtype == torch.float32 diff --git a/fastvideo/tests/ops/quantization/test_allowlist_mirror.py b/fastvideo/tests/ops/quantization/test_allowlist_mirror.py new file mode 100644 index 0000000000..6057f994aa --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_allowlist_mirror.py @@ -0,0 +1,43 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Guard against drift between the two copies of the zero-init allowlist. + +``fsdp_load`` owns the canonical ``ALLOWED_NEW_PARAM_PATTERNS`` and uses it to +decide which model parameters may be zero-initialized when a checkpoint does not +supply them. ``shard_cache`` keeps a literal mirror of that tuple because +``fsdp_load`` imports ``shard_cache`` -- importing back would be circular. + +A mismatch does not fail a run; it silently keeps the weight shard cache cold +for any quantization config that registers new parameters (for example +``AbsMaxFP8``, whose ``scale_weight`` / ``scale_input`` appear in the model but +never in the checkpoint). That is an invisible performance regression, so this +test pins the two tuples together. +""" + +import unittest + +from fastvideo.models.loader.fsdp_load import ALLOWED_NEW_PARAM_PATTERNS +from fastvideo.models.loader.shard_cache import _ALLOWED_NEW_PARAM_PATTERNS + + +class TestAllowlistMirror(unittest.TestCase): + """``shard_cache`` must mirror ``fsdp_load``'s allowlist exactly.""" + + def test_mirror_matches_canonical_source(self): + self.assertEqual( + _ALLOWED_NEW_PARAM_PATTERNS, + ALLOWED_NEW_PARAM_PATTERNS, + "shard_cache._ALLOWED_NEW_PARAM_PATTERNS has drifted from " + "fsdp_load.ALLOWED_NEW_PARAM_PATTERNS. Update the mirror in " + "fastvideo/models/loader/shard_cache.py to match.", + ) + + def test_quant_scale_params_are_shared(self): + """The specific names a quantizing loader depends on are present in both.""" + for name in ("scale_weight", "scale_input"): + with self.subTest(param=name): + self.assertIn(name, ALLOWED_NEW_PARAM_PATTERNS) + self.assertIn(name, _ALLOWED_NEW_PARAM_PATTERNS) + + +if __name__ == "__main__": + unittest.main() diff --git a/fastvideo/tests/ops/quantization/test_int8_affine_config.py b/fastvideo/tests/ops/quantization/test_int8_affine_config.py new file mode 100644 index 0000000000..4810056668 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_int8_affine_config.py @@ -0,0 +1,303 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU-only tests for the affine group-64 INT8 quantization config. + +No CUDA, no model weights, no inference stack: every test here runs on a +laptop. The three things being pinned are (1) the module imports without +flashinfer/CUDA, (2) the quantizer reproduces the validated MLX affine math +and round-trips within a stated tolerance, and (3) layer selection picks +H3's attention/FFN GEMMs and *excludes* ``attn.to_gate_compress``. + +Test 3 is the load-bearing one: ``to_gate_compress`` is H3's VSA sparse +routing gate, and quantizing it changes a discrete routing decision rather +than perturbing an output. +""" + +from __future__ import annotations + +import importlib +import math +import sys + +import pytest +import torch + +from fastvideo.layers.quantization.int8_affine_config import ( + INT8AffineConfig, + MINIMAX_H3_INT8_AFFINE_SUFFIXES, + int8_affine_dequantize, + int8_affine_quantize, + minimax_h3_int8_affine_prefixes, +) + + + +def test_config_module_imports_without_cuda_or_flashinfer(): + """The config must import on a CPU-only host with no flashinfer.""" + assert "flashinfer" not in sys.modules, "importing the config must not pull in flashinfer" + module = importlib.import_module("fastvideo.layers.quantization.int8_affine_config") + assert module is not None + assert "flashinfer" not in sys.modules + + +def test_config_metadata_is_well_formed(): + cfg = INT8AffineConfig() + assert cfg.get_name() == "INT8Affine" + assert cfg.group_size == 64 + assert cfg.bits == 8 + assert torch.bfloat16 in cfg.get_supported_act_dtypes() + assert cfg.get_config_filenames() == [] + assert INT8AffineConfig.get_min_capability() >= 70 + rebuilt = INT8AffineConfig.from_config({"group_size": 32, "bits": 4}) + assert (rebuilt.group_size, rebuilt.bits) == (32, 4) + + +def test_invalid_bits_and_group_size_are_rejected(): + with pytest.raises(ValueError): + INT8AffineConfig(bits=16) + with pytest.raises(ValueError): + INT8AffineConfig(group_size=0) + + + + +def _independent_affine_reference(w: torch.Tensor, group_size: int = 64, bits: int = 8): + """A from-scratch restatement of the MLX affine algorithm. + + Written against the description in ``mlx_affine_qat.py``'s module + docstring (per-group min/max, magnitude-anchored sign, exact-integer + anchor re-expression, rint rounding, clamp to ``[0, 2**bits-1]``) rather + than by copying the implementation, so it is a real cross-check of the + transcription and not a tautology. + """ + n_bins = float((1 << bits) - 1) + flat = w.reshape(-1, w.shape[-1] // group_size, group_size).float() + lo = flat.amin(dim=-1) + hi = flat.amax(dim=-1) + anchor = torch.where(lo.abs() > hi.abs(), lo, hi) + other = torch.where(lo.abs() > hi.abs(), hi, lo) + step = ((hi - lo) / n_bins).clamp_min(1e-7) + step = torch.where(lo.abs() > hi.abs(), step, -step) + q0 = torch.round(anchor / step) + use = q0 != 0 + step = torch.where(use, anchor / torch.where(use, q0, torch.ones_like(q0)), step) + zero = torch.where(use, anchor, torch.zeros_like(anchor)) + del other + codes = torch.round((flat - zero.unsqueeze(-1)) / step.unsqueeze(-1)).clamp(0.0, n_bins) + return codes.to(torch.int64), step, zero + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16, torch.float16]) +def test_quantize_matches_independent_reference(dtype): + torch.manual_seed(0) + w = torch.randn(37, 256, dtype=dtype) + codes, scales, biases = int8_affine_quantize(w, group_size=64, bits=8) + ref_codes, ref_scales, ref_bias = _independent_affine_reference(w, group_size=64, bits=8) + + assert codes.dtype == torch.uint8 + assert codes.shape == (*w.shape[:-1], w.shape[-1] // 64, 64) + assert scales.shape == (*w.shape[:-1], w.shape[-1] // 64) + assert torch.equal(codes.reshape(ref_codes.shape).to(torch.int64), ref_codes) + torch.testing.assert_close(scales, ref_scales.to(dtype), rtol=0, atol=0) + torch.testing.assert_close(biases, ref_bias.to(dtype), rtol=0, atol=0) + + +def test_quantize_matches_in_tree_mlx_reference_when_available(): + """Parity against the real ``mlx_affine_qat`` transcription, if present. + + That module lives in a sibling worktree today, so this is skipped unless + it has landed next to us; when it does, the two transcriptions must agree + bit-for-bit on codes and exactly on the fp32 scales/biases. + """ + mlx = pytest.importorskip("fastvideo.layers.quantization.mlx_affine_qat", + reason="mlx_affine_qat.py is not in this tree yet") + torch.manual_seed(1) + w = torch.randn(16, 128, dtype=torch.float32) + codes, scales, biases = int8_affine_quantize(w, group_size=64, bits=8) + ref_codes, ref_scales, ref_bias = mlx.mlx_affine_quantize_reference(w, group_size=64, bits=8) + assert torch.equal(codes.reshape(ref_codes.shape).to(torch.int64), ref_codes.to(torch.int64)) + torch.testing.assert_close(scales, ref_scales, rtol=0, atol=0) + torch.testing.assert_close(biases, ref_bias, rtol=0, atol=0) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +def test_roundtrip_error_is_within_documented_tolerance(dtype): + """Group-64 8-bit affine on N(0,1) weights. + + Tolerances are measured over seeds 0-5, not guessed, then given ~2x + headroom. Group-64 8-bit spends ~1/255 of the per-group range per step + (step ~= 0.016 for a 64-sample N(0,1) group), so the reconstruction error + is ~step/sqrt(12) RMS and ~step/2 at worst: + + - fp32 source: max/max|w| <= 0.0055, rms/max|w| <= 0.0013 + - bf16 source: max/max|w| <= 0.0093, rms/max|w| <= 0.0021 (the source is + itself bf16-rounded before quantizing) + """ + torch.manual_seed(2) + w = torch.randn(64, 1024, dtype=dtype) + codes, scales, biases = int8_affine_quantize(w, group_size=64, bits=8) + deq = int8_affine_dequantize(codes, scales, biases, out_shape=w.shape).float() + + assert deq.shape == w.shape + top = w.float().abs().max().item() + rel_max = (w.float() - deq).abs().max().item() / top + rel_rms = math.sqrt(((w.float() - deq)**2).mean().item()) / top + max_limit, rms_limit = (0.02, 0.005) if dtype is torch.bfloat16 else (0.01, 0.003) + assert rel_max < max_limit, f"max relative error {rel_max:.5f} exceeded {max_limit} for {dtype}" + assert rel_rms < rms_limit, f"rms relative error {rel_rms:.5f} exceeded {rms_limit} for {dtype}" + + grouped_w = w.float().reshape(-1, 64) + grouped_deq = deq.reshape(-1, 64) + extreme = grouped_w.abs().max(dim=-1).values + at_extreme = grouped_w.abs() == extreme.unsqueeze(-1) + rtol, atol = (1e-5, 1e-5) if dtype is torch.float32 else (1e-2, 1e-2) + assert torch.allclose(grouped_w[at_extreme], grouped_deq[at_extreme], rtol=rtol, atol=atol) + + +def test_group_size_must_divide_input_dim(): + with pytest.raises(ValueError, match="not divisible"): + int8_affine_quantize(torch.randn(4, 100), group_size=64) + + +def test_codes_never_exceed_uint8_range(): + torch.manual_seed(3) + for _ in range(5): + w = torch.randn(8, 256) * torch.rand(8, 1) * 10 + codes, _, _ = int8_affine_quantize(w, group_size=64, bits=8) + assert codes.max().item() <= 255 + + + +_H3_INCLUDED = [ + "minimax_h3.transformer_blocks.0.attn.to_q", + "minimax_h3.transformer_blocks.0.attn.to_k", + "minimax_h3.transformer_blocks.0.attn.to_v", + "minimax_h3.transformer_blocks.0.attn.to_out", + "minimax_h3.transformer_blocks.0.ff.fc_in", + "minimax_h3.transformer_blocks.0.ff.fc_out", + "minimax_h3.transformer_blocks.49.attn.to_q", + "minimax_h3.transformer_blocks.49.ff.fc_out", + "minimax_h3.token_refiner.refiner_blocks.0.attn.to_q", + "minimax_h3.token_refiner.refiner_blocks.0.attn.to_out", + "minimax_h3.token_refiner.refiner_blocks.1.ff.fc_in", + "minimax_h3.transformer_blocks.0.adaln_proj.linear", +] + +_H3_EXCLUDED = [ + "minimax_h3.transformer_blocks.0.attn.to_gate_compress", + "minimax_h3.transformer_blocks.49.attn.to_gate_compress", + "minimax_h3.adaln_basis", + "minimax_h3.proj_in", + "minimax_h3.audio_proj_in", + "minimax_h3.proj_out", + "minimax_h3.audio_proj_out", + "minimax_h3.time_embedder.fc_in", + "minimax_h3.time_embedder.fc_out", + "minimax_h3.context_embedder", + "minimax_h3.transformer_blocks.0.norm1", + "minimax_h3.rope", +] + + +def test_h3_layer_selection_include_and_exclude_sets(): + cfg = INT8AffineConfig.for_minimax_h3() + for prefix in _H3_INCLUDED: + assert cfg.is_target_layer(prefix), f"expected {prefix!r} to be quantized" + for prefix in _H3_EXCLUDED: + assert not cfg.is_target_layer(prefix), f"expected {prefix!r} to be EXCLUDED" + + +def test_to_gate_compress_is_excluded_even_by_a_broad_allowlist(): + """The deny list is fail-closed: no constructor argument re-enables it. + + This is the regression guard for the failure mode the recon flagged — + a name that matches no "norm"/"scale_shift_table"-style heuristic being + silently swept into a broad suffix rule. + """ + hostile = INT8AffineConfig( + layer_suffixes=("to_q", "to_k", "to_v", "to_out", "to_gate_compress"), + exclude_substrings=(), # caller tries to clear the deny list + ) + assert not hostile.is_target_layer("minimax_h3.transformer_blocks.7.attn.to_gate_compress") + assert hostile.is_target_layer("minimax_h3.transformer_blocks.7.attn.to_q") + + explicit = INT8AffineConfig( + target_layers=("minimax_h3.transformer_blocks.7.attn.to_gate_compress", ), + ) + assert not explicit.is_target_layer("minimax_h3.transformer_blocks.7.attn.to_gate_compress") + + +def test_enumerated_h3_prefixes_agree_with_selection(): + """The literal enumerated set and the suffix rule must pick the same layers.""" + cfg = INT8AffineConfig.for_minimax_h3() + enumerated = minimax_h3_int8_affine_prefixes() + assert len(enumerated) == 50 * 7 + 2 * 6 # 50 blocks x 7 suffixes, 2 refiner blocks x 6 + for prefix in enumerated: + assert cfg.is_target_layer(prefix), f"enumerated prefix {prefix!r} not selected by the suffix rule" + selected = { + f"minimax_h3.{scope}.{i}.{suffix}" + for scope, count in (("transformer_blocks", 50), ("token_refiner.refiner_blocks", 2)) + for i in range(count) + for suffix in MINIMAX_H3_INT8_AFFINE_SUFFIXES + if not (scope != "transformer_blocks" and suffix.startswith("adaln_proj")) + if cfg.is_target_layer(f"minimax_h3.{scope}.{i}.{suffix}") + } + assert selected == set(enumerated) + + +def test_non_linear_layers_get_no_quant_method(): + from fastvideo.layers.linear import ReplicatedLinear + + cfg = INT8AffineConfig.for_minimax_h3() + linear = ReplicatedLinear(64, 64, bias=False, quant_config=cfg, prefix="minimax_h3.transformer_blocks.0.attn.to_q") + assert linear.quant_method is not None + assert linear.quant_method.__class__.__name__ == "INT8AffineQuantizeMethod" + + gate = ReplicatedLinear(64, 64, bias=False, quant_config=cfg, + prefix="minimax_h3.transformer_blocks.0.attn.to_gate_compress") + assert gate.quant_method.__class__.__name__ == "UnquantizedLinearMethod" + + + + +def test_conversion_and_apply_match_dense_linear_within_tolerance(): + """End-to-end: load-time conversion then apply() dequantizes and matches. + + Uses ``retain_original_weight=False`` to exercise the purge path too. + """ + from fastvideo.layers.linear import ReplicatedLinear + from fastvideo.layers.quantization.int8_affine_config import convert_model_to_int8_affine + + torch.manual_seed(4) + cfg = INT8AffineConfig.for_minimax_h3(retain_original_weight=False) + layer = ReplicatedLinear(128, 256, bias=False, quant_config=cfg, + prefix="minimax_h3.transformer_blocks.0.attn.to_q") + with torch.no_grad(): + layer.weight.copy_(torch.randn(256, 128)) + + convert_model_to_int8_affine(layer) + assert layer._int8_affine_codes.dtype == torch.uint8 + assert layer._int8_affine_codes.shape == (256, 128) + assert layer._int8_affine_scales.shape == (256, 2) + assert layer.weight is None, "retain_original_weight=False should purge the bf16 weight" + + x = torch.randn(4, 128, dtype=torch.bfloat16) + out, _ = layer(x) + assert out.shape == (4, 256) + assert torch.isfinite(out).all() + + +def test_apply_falls_back_to_dense_under_grad(): + """A training step must see the master weight, not a frozen dequant copy.""" + from fastvideo.layers.linear import ReplicatedLinear + + torch.manual_seed(5) + cfg = INT8AffineConfig.for_minimax_h3() + layer = ReplicatedLinear(64, 64, bias=False, quant_config=cfg, + prefix="minimax_h3.transformer_blocks.0.attn.to_q") + with torch.no_grad(): + layer.weight.copy_(torch.randn(64, 64)) + x = torch.randn(2, 64, dtype=torch.bfloat16) + with torch.enable_grad(): + out, _ = layer(x) + assert out.shape == (2, 64) + assert not hasattr(layer, "_int8_affine_codes"), "grad-enabled forward must not quantize in place" diff --git a/fastvideo/tests/ops/quantization/test_int8_dispatch.py b/fastvideo/tests/ops/quantization/test_int8_dispatch.py new file mode 100644 index 0000000000..2e8f7d2c82 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_int8_dispatch.py @@ -0,0 +1,68 @@ +# SPDX-License-Identifier: Apache-2.0 +"""The loader must dispatch INT8 affine configs to their conversion hook. + +``_maybe_quantize_model`` walks the module tree and converts weights for +whichever quantization method it finds attached. It is an explicit +``isinstance`` chain, so a new quantization config is silently inert until it +gains a branch here -- the model keeps its dense weights and quietly produces +unquantized output with no error. + +This test pins the INT8Affine branch so that adding or reordering the chain +cannot silently drop it. +""" + +import unittest + +import torch +import torch.nn as nn + +from fastvideo.layers.quantization import int8_affine_config +from fastvideo.layers.quantization.int8_affine_config import ( + INT8AffineQuantizeMethod, + convert_model_to_int8_affine, +) +from fastvideo.models.loader import fsdp_load + + +class _FakeLinear(nn.Module): + """Minimal stand-in: a dense weight plus an attached quant method.""" + + def __init__(self, prefix: str = "minimax_h3.transformer_blocks.0.attn.to_q"): + super().__init__() + self.weight = nn.Parameter(torch.randn(64, 64, dtype=torch.bfloat16)) + self.quant_method = INT8AffineQuantizeMethod(layer_prefix=prefix) + + +class TestInt8Dispatch(unittest.TestCase): + + def setUp(self): + self._calls = [] + self._orig = int8_affine_config.convert_model_to_int8_affine + + def _spy(model): + self._calls.append(model) + return self._orig(model) + + int8_affine_config.convert_model_to_int8_affine = _spy + self.addCleanup(setattr, int8_affine_config, "convert_model_to_int8_affine", self._orig) + + def test_int8_layer_triggers_conversion(self): + model = _FakeLinear() + fsdp_load._maybe_quantize_model(model) + self.assertEqual(len(self._calls), 1, "INT8AffineQuantizeMethod did not reach its conversion hook") + self.assertTrue(hasattr(model, "_int8_affine_codes"), "quantized buffers were not registered") + + def test_unquantized_model_is_untouched(self): + """A plain layer must not trigger any conversion.""" + model = nn.Linear(8, 8) + fsdp_load._maybe_quantize_model(model) + self.assertEqual(self._calls, [], "conversion ran on a model with no quant method") + + def test_conversion_is_exported_by_the_config_module(self): + """The hook the loader imports must be the one the config defines.""" + self.assertTrue(callable(convert_model_to_int8_affine)) + self.assertTrue(callable(self._orig)) + + +if __name__ == "__main__": + unittest.main() diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_h3_prefixes.py b/fastvideo/tests/ops/quantization/test_nvfp4_h3_prefixes.py new file mode 100644 index 0000000000..c8ed2f2df2 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_nvfp4_h3_prefixes.py @@ -0,0 +1,150 @@ +# SPDX-License-Identifier: Apache-2.0 +"""NVFP4 layer-prefix selection: LTX-2 default, MiniMax-H3 opt-in. + +``NVFP4Config`` used to hardcode the LTX-2 layer paths, so the same config +handed to MiniMax-H3 attached no quant methods at all and the model ran dense +without an error. These CPU-only tests pin both halves of the fix: the default +still selects exactly the LTX-2 set, and an H3-configured instance selects the +300 block linears while refusing to quantize the VSA compression gate. + +No flashinfer and no CUDA are required — the tests use real +``ReplicatedLinear`` layers (which is where ``get_quant_method`` is called from) +with only ``NVFP4QuantizeMethod.__init__`` replaced, because the real one +allocates ``x_global_sf`` on ``cuda``. The same shim is used by +``test_nvfp4_purge.py``. +""" +from __future__ import annotations + +import sys + +import pytest +import torch + +import fastvideo.layers.quantization.nvfp4_config as nv +from fastvideo.layers.linear import ReplicatedLinear, UnquantizedLinearMethod + +_H3_BLOCK = "minimax_h3.transformer_blocks.{idx}.{suffix}" +_GATE_PREFIX = _H3_BLOCK.format(idx=0, suffix="attn.to_gate_compress") + + +@pytest.fixture(autouse=True) +def _cpu_quantize_method(monkeypatch): + """Build NVFP4QuantizeMethod without its cuda-allocated ``x_global_sf``.""" + + def _init(self, layer_prefix: str = ""): + self.weight_fp4 = None + self.weight_scale = None + self.x_global_sf = torch.tensor(1.0, dtype=torch.float32) + self.layer_prefix = layer_prefix + self._is_refine_only_layer = nv._is_ltx2_refine_only_prefix(layer_prefix) + self._retain_original_weights = None + + monkeypatch.setattr(nv.NVFP4QuantizeMethod, "__init__", _init) + + +def _linear(quant_config, prefix: str) -> ReplicatedLinear: + return ReplicatedLinear(8, 8, bias=False, quant_config=quant_config, prefix=prefix) + + +def test_default_config_keeps_the_ltx2_layer_set() -> None: + """No regression: the historical default is still the LTX-2 set.""" + config = nv.NVFP4Config() + assert config.layer_prefixes == nv._LTX2_NVFP4_LINEAR_PREFIXES + assert len(config.layer_prefixes) == 577 + assert config.exclude_prefixes == frozenset() + assert config.is_nvfp4_linear_prefix("ltx2.blocks.0.attn1.to_q") + assert config.is_nvfp4_linear_prefix("ltx2.adaln_single.linear") + assert not config.is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=0, suffix="attn.to_q")) + + +def test_h3_config_selects_the_300_block_linears() -> None: + config = nv.NVFP4Config.for_minimax_h3() + assert len(config.layer_prefixes) == 300 + assert config.layer_prefixes == nv.MINIMAX_H3_NVFP4_LINEAR_PREFIXES + for idx in range(nv.MINIMAX_H3_NUM_LAYERS): + for suffix in nv.MINIMAX_H3_BLOCK_LINEAR_SUFFIXES: + assert config.is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=idx, suffix=suffix)) + assert config.is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=49, suffix="ff.fc_out")) + assert not config.is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=50, suffix="ff.fc_out")) + assert not config.is_nvfp4_linear_prefix("minimax_h3.token_refiner.blocks.0.attn.to_q") + assert not config.is_nvfp4_linear_prefix("minimax_h3.proj_in") + assert not config.is_nvfp4_linear_prefix("minimax_h3.transformer_blocks.0.adaln_proj.linear") + assert not config.is_nvfp4_linear_prefix("ltx2.blocks.0.attn1.to_q") + + +def test_h3_config_excludes_the_vsa_gate() -> None: + """``attn.to_gate_compress`` must never be quantized.""" + config = nv.NVFP4Config.for_minimax_h3() + assert not config.is_nvfp4_linear_prefix(_GATE_PREFIX) + assert not config.is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=37, suffix="attn.to_gate_compress")) + assert _GATE_PREFIX not in config.layer_prefixes + assert config.exclude_prefixes == frozenset(nv.MINIMAX_H3_NVFP4_EXCLUDED_LINEAR_SUFFIXES) + + +def test_gate_is_excluded_even_if_a_caller_allowlists_it() -> None: + """The gate exclusion is unconditional, not just absent from the set. + + A caller who builds the prefix set with a glob (or hand-lists every linear + in the block) must still not get a quantized VSA gate. + """ + config = nv.NVFP4Config(layer_prefixes=frozenset({_GATE_PREFIX, _H3_BLOCK.format(idx=0, suffix="attn.to_q")})) + assert not config.is_nvfp4_linear_prefix(_GATE_PREFIX) + assert config.is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=0, suffix="attn.to_q")) + assert _GATE_PREFIX not in nv._LTX2_NVFP4_LINEAR_PREFIXES + assert not nv.NVFP4Config(layer_prefixes=[_GATE_PREFIX]).is_nvfp4_linear_prefix(_GATE_PREFIX) + + +def test_exclude_prefixes_match_a_suffix_or_a_full_path() -> None: + full_path = "custom.blocks.3.attn.to_out" + config = nv.NVFP4Config(layer_prefixes=[full_path, "custom.other"], exclude_prefixes=["attn.to_out"]) + assert config.is_nvfp4_linear_prefix("custom.other") + assert not config.is_nvfp4_linear_prefix(full_path) + nested = nv.NVFP4Config(layer_prefixes=["custom.cross_ff.fc_in"], exclude_prefixes=["ff.fc_in"]) + assert nested.is_nvfp4_linear_prefix("custom.cross_ff.fc_in") + + +def test_from_config_round_trips_layer_prefixes() -> None: + config = nv.NVFP4Config.from_config({ + "layer_profile": "base", + "layer_prefixes": ["minimax_h3.transformer_blocks.0.attn.to_q"], + "exclude_prefixes": ["attn.to_gate_compress"], + }) + assert config.layer_profile == "base" + assert config.layer_prefixes == frozenset({"minimax_h3.transformer_blocks.0.attn.to_q"}) + assert config.exclude_prefixes == frozenset({"attn.to_gate_compress"}) + assert nv.NVFP4Config.from_config({}).layer_prefixes == nv._LTX2_NVFP4_LINEAR_PREFIXES + + +def test_for_minimax_h3_forwards_kwargs() -> None: + config = nv.NVFP4Config.for_minimax_h3(retain_original_weights=True) + assert config.retain_original_weights is True + assert config.layer_prefixes == nv.MINIMAX_H3_NVFP4_LINEAR_PREFIXES + + +def test_get_quant_method_attaches_for_h3_and_skips_the_gate() -> None: + """End-to-end through ``ReplicatedLinear``, which is the real call site.""" + h3 = nv.NVFP4Config.for_minimax_h3() + selected = _linear(h3, _H3_BLOCK.format(idx=0, suffix="attn.to_q")) + assert isinstance(selected.quant_method, nv.NVFP4QuantizeMethod) + assert selected.quant_method.layer_prefix == _H3_BLOCK.format(idx=0, suffix="attn.to_q") + + gate = _linear(h3, _GATE_PREFIX) + assert type(gate.quant_method) is UnquantizedLinearMethod + + default = nv.NVFP4Config() + untagged = _linear(default, _H3_BLOCK.format(idx=0, suffix="attn.to_q")) + assert type(untagged.quant_method) is UnquantizedLinearMethod + ltx2 = _linear(default, "ltx2.blocks.0.attn1.to_q") + assert isinstance(ltx2.quant_method, nv.NVFP4QuantizeMethod) + + +def test_module_imports_without_flashinfer(monkeypatch) -> None: + """The H3 surface must import on hosts with no flashinfer (only the + kernels fail, at use time).""" + monkeypatch.setitem(sys.modules, "flashinfer", None) + monkeypatch.delitem(sys.modules, "fastvideo.layers.quantization.nvfp4_config", raising=False) + import importlib + + reloaded = importlib.import_module("fastvideo.layers.quantization.nvfp4_config") + assert len(reloaded.MINIMAX_H3_NVFP4_LINEAR_PREFIXES) == 300 + assert reloaded.NVFP4Config.for_minimax_h3().is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=0, suffix="attn.to_k")) diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_sidecar.py b/fastvideo/tests/ops/quantization/test_nvfp4_sidecar.py new file mode 100644 index 0000000000..9b912a6180 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_nvfp4_sidecar.py @@ -0,0 +1,323 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU-only round-trip tests for the compact NVFP4 checkpoint sidecar. + +``convert_model_to_nvfp4`` registers its packed tensors with +``persistent=False``, so a saved checkpoint carries dense bf16 weights only and +the FP4 payload is rebuilt at every load. ``save_nvfp4_checkpoint`` / +``load_nvfp4_checkpoint`` write and restore that payload directly. + +flashinfer, the FP4 GEMM and ``NVFP4QuantizeMethod.__init__``'s cuda allocation +are all stubbed, so this runs on any host. The conversion, the buffer +registration and the retention policy under test are the real ones. +""" +from __future__ import annotations + +import json +import logging + +import pytest +import torch +import torch.nn as nn +from safetensors import safe_open +from safetensors.torch import save_file + +import fastvideo.layers.quantization.nvfp4_config as nv +from fastvideo.layers.linear import ReplicatedLinear + +_BLOCK = "minimax_h3.transformer_blocks.{idx}" +_IN_DIM = 8 +_OUT_DIM = 4 + + +@pytest.fixture(autouse=True) +def _stub_flashinfer(monkeypatch): + """Patch the three flashinfer touch points the real path would use.""" + import types + + fake_sf_layout = types.SimpleNamespace(layout_128x4=None) + monkeypatch.setattr(nv, "_require_flashinfer", lambda: (fake_sf_layout, None, None)) + monkeypatch.setattr(nv, "_nvfp4_quantize", _fake_quantize) + + def _init(self, layer_prefix: str = ""): + self.weight_fp4 = None + self.weight_scale = None + self.x_global_sf = torch.tensor(1.0, dtype=torch.float32) + self.layer_prefix = layer_prefix + self._is_refine_only_layer = nv._is_ltx2_refine_only_prefix(layer_prefix) + self._retain_original_weights = None + + monkeypatch.setattr(nv.NVFP4QuantizeMethod, "__init__", _init) + + +def _fake_quantize(weight, global_sf, sfLayout=None, do_shuffle=False): + """Deterministic stand-in for ``nvfp4_quantize`` with the real shapes.""" + out_dim, in_dim = weight.shape[0], weight.shape[-1] + packed = torch.arange(out_dim * ((in_dim + 1) // 2), dtype=torch.uint8).view(out_dim, (in_dim + 1) // 2) + scales = torch.arange(out_dim * ((in_dim + 15) // 16), dtype=torch.uint8).view(out_dim, (in_dim + 15) // 16) + return packed, scales + + +def _module() -> nn.Module: + """An empty submodule, so FQNs can be assembled explicitly.""" + return nn.Module() + + +def _build_model(*, num_blocks: int = 2, seed: int = 0, quant_config=None) -> nn.Module: + """A miniature H3-shaped DiT: ``minimax_h3.transformer_blocks.{i}.{attn.to_q,ff.fc_in}``.""" + config = quant_config if quant_config is not None else nv.NVFP4Config.for_minimax_h3() + generator = torch.Generator().manual_seed(seed) + root = _module() + dit = _module() + blocks = nn.ModuleList() + for idx in range(num_blocks): + block = _module() + attn = _module() + attn.to_q = ReplicatedLinear(_IN_DIM, + _OUT_DIM, + bias=False, + quant_config=config, + prefix=f"{_BLOCK.format(idx=idx)}.attn.to_q") + block.attn = attn + ff = _module() + ff.fc_in = ReplicatedLinear(_IN_DIM, + _OUT_DIM, + bias=False, + quant_config=config, + prefix=f"{_BLOCK.format(idx=idx)}.ff.fc_in") + block.ff = ff + blocks.append(block) + dit.transformer_blocks = blocks + root.minimax_h3 = dit + for _, param in root.named_parameters(): + param.data.copy_(torch.randn(param.shape, generator=generator)) + return root + + +def _buffers(model: nn.Module) -> dict[str, torch.Tensor]: + return nv.nvfp4_sidecar_state_dict(model) + + +def _rewrite_sidecar(src, dst, *, metadata: dict | None = None, tensors: dict | None = None) -> str: + """Copy a sidecar, optionally replacing its manifest or one tensor.""" + with safe_open(src, framework="pt", device="cpu") as handle: + payload = {key: handle.get_tensor(key) for key in handle.keys()} + manifest = json.loads(handle.metadata()[nv._NVFP4_SIDECAR_METADATA_KEY]) + if metadata is not None: + manifest.update(metadata) + if tensors is not None: + payload.update(tensors) + save_file(payload, str(dst), metadata={nv._NVFP4_SIDECAR_METADATA_KEY: json.dumps(manifest)}) + return str(dst) + + +def test_convert_purges_weights_and_sidecar_shrinks_the_payload(tmp_path) -> None: + model = _build_model() + nv.convert_model_to_nvfp4(model) + assert model.minimax_h3.transformer_blocks[0].attn.to_q.weight is None + + path = tmp_path / "nvfp4.safetensors" + receipt = nv.save_nvfp4_checkpoint(model, path) + assert path.exists() + assert receipt["num_layers"] == 4 + assert receipt["num_tensors"] == 16 # 4 buffers x 4 layers + assert receipt["quantized_bytes"] == sum(t.numel() * t.element_size() for t in _buffers(model).values()) + assert receipt["quantized_bytes"] < receipt["dense_bfloat16_bytes"] + assert receipt["compression_ratio"] > 1.0 + + manifest = nv.read_nvfp4_sidecar_metadata(path) + assert manifest["format"] == nv._NVFP4_SIDECAR_FORMAT + assert manifest["version"] == nv._NVFP4_SIDECAR_VERSION + assert manifest["sf_layout"] == "layout_128x4" + assert manifest["do_shuffle"] is False + assert manifest["block_size"] == 16 + assert manifest["layers"][f"{_BLOCK.format(idx=1)}.ff.fc_in"] == [_OUT_DIM, _IN_DIM] + assert manifest["quant_prefixes"][f"{_BLOCK.format(idx=0)}.attn.to_q"] == f"{_BLOCK.format(idx=0)}.attn.to_q" + + +def test_sidecar_state_dict_uses_module_fqn_keys_and_cpu_tensors() -> None: + model = _build_model() + nv.convert_model_to_nvfp4(model) + state = _buffers(model) + key = f"{_BLOCK.format(idx=0)}.attn.to_q::{nv._NVFP4_SIDECAR_BUFFERS[0]}" + assert key in state + assert state[key].device.type == "cpu" + assert {name for _, name in (k.split("::") for k in state)} == set(nv._NVFP4_SIDECAR_BUFFERS) + assert state[f"{_BLOCK.format(idx=0)}.attn.to_q::_nvfp4_weight"].dtype is torch.uint8 + + +def test_load_restores_buffers_without_reconverting(tmp_path) -> None: + """The restored tensors must be identical to a fresh conversion's.""" + source = _build_model(seed=1) + nv.convert_model_to_nvfp4(source) + expected = {key: value.clone() for key, value in _buffers(source).items()} + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + target = _build_model(seed=2) + restored = nv.load_nvfp4_checkpoint(target, path) + assert restored == 4 + actual = _buffers(target) + assert set(actual) == set(expected) + for key, value in expected.items(): + assert torch.equal(actual[key], value), key + assert actual[key].dtype == value.dtype + assert target.minimax_h3.transformer_blocks[0].attn.to_q.weight is None + + +def test_load_does_not_need_flashinfer(monkeypatch, tmp_path) -> None: + """Serving a pre-quantized checkpoint must not require the FP4 kernels.""" + source = _build_model(seed=3) + nv.convert_model_to_nvfp4(source) + expected = {key: value.clone() for key, value in _buffers(source).items()} + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + def _boom(): + raise ImportError("NVFP4 quantization requires flashinfer.") + + monkeypatch.setattr(nv, "_require_flashinfer", _boom) + target = _build_model(seed=4) + assert nv.load_nvfp4_checkpoint(target, path) == 4 + assert all(torch.equal(_buffers(target)[key], value) for key, value in expected.items()) + + +def test_load_into_a_model_whose_dense_weights_were_never_loaded(tmp_path) -> None: + """The compact case: the checkpoint carries no bf16 weights at all.""" + source = _build_model(seed=5) + nv.convert_model_to_nvfp4(source) + expected = {key: value.clone() for key, value in _buffers(source).items()} + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + target = _build_model(seed=6) + for module in target.modules(): + if getattr(module, "quant_method", None) is not None and hasattr(module, "weight"): + module.register_parameter("weight", None) + assert nv.load_nvfp4_checkpoint(target, path) == 4 + assert all(torch.equal(_buffers(target)[key], value) for key, value in expected.items()) + + +def test_load_purges_dense_weights_unless_asked_not_to(tmp_path) -> None: + source = _build_model(seed=7) + nv.convert_model_to_nvfp4(source) + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + keep = _build_model(seed=8) + nv.load_nvfp4_checkpoint(keep, path, purge_dense_weights=False) + assert keep.minimax_h3.transformer_blocks[0].attn.to_q.weight is not None + + +def _sidecar_of(model: nn.Module, path) -> str: + nv.convert_model_to_nvfp4(model) + nv.save_nvfp4_checkpoint(model, path) + return str(path) + + +def test_layer_set_mismatch_strict_and_lenient(tmp_path, caplog) -> None: + source = _build_model(num_blocks=1, seed=10) + nv.convert_model_to_nvfp4(source) + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + wider = _build_model(num_blocks=3, seed=13) + bigger = _build_model(num_blocks=2, seed=11) + with pytest.raises(ValueError, match="missing from the sidecar"): + nv.load_nvfp4_checkpoint(bigger, path) + + with caplog.at_level(logging.WARNING): + assert nv.load_nvfp4_checkpoint(bigger, path, strict=False) == 2 + assert any("does not match this model" in record.message for record in caplog.records) + assert getattr(bigger.minimax_h3.transformer_blocks[1].attn.to_q, "_nvfp4_weight", None) is None + + smaller = _build_model(num_blocks=1, seed=12) + with pytest.raises(ValueError, match="not in the model"): + nv.load_nvfp4_checkpoint(smaller, _sidecar_of(wider, tmp_path / "wide.safetensors")) + + +def test_layout_and_version_mismatches_are_fatal(tmp_path) -> None: + source = _build_model(seed=13) + nv.convert_model_to_nvfp4(source) + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + target = _build_model(seed=14) + bad_layout = _rewrite_sidecar(path, tmp_path / "layout.safetensors", metadata={"sf_layout": "layout_linear"}) + with pytest.raises(ValueError, match="sf_layout"): + nv.load_nvfp4_checkpoint(target, bad_layout) + with pytest.raises(ValueError, match="sf_layout"): + nv.load_nvfp4_checkpoint(target, bad_layout, strict=False) + + bad_shuffle = _rewrite_sidecar(path, tmp_path / "shuffle.safetensors", metadata={"do_shuffle": True}) + with pytest.raises(ValueError, match="do_shuffle"): + nv.load_nvfp4_checkpoint(target, bad_shuffle) + + bad_version = _rewrite_sidecar(path, tmp_path / "version.safetensors", metadata={"version": 99}) + with pytest.raises(ValueError, match="version"): + nv.load_nvfp4_checkpoint(target, bad_version) + + bad_format = _rewrite_sidecar(path, tmp_path / "format.safetensors", metadata={"format": "something.else"}) + with pytest.raises(ValueError, match="format"): + nv.load_nvfp4_checkpoint(target, bad_format) + + +def test_tensor_shape_mismatch_is_rejected(tmp_path) -> None: + source = _build_model(seed=15) + nv.convert_model_to_nvfp4(source) + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + key = f"{_BLOCK.format(idx=0)}.attn.to_q::_nvfp4_weight" + bad = _rewrite_sidecar(path, tmp_path / "shape.safetensors", tensors={key: torch.zeros(3, 7, dtype=torch.uint8)}) + with pytest.raises(ValueError, match="expected"): + nv.load_nvfp4_checkpoint(_build_model(seed=16), bad) + + +def test_block_scales_padded_to_the_128_row_tile_are_accepted(tmp_path) -> None: + """``nvfp4_quantize`` returns block scales padded to the 128-row tile. + + ``_nvfp4_quantize`` narrows the packed weight back to the logical row count + but not the scales, so a model whose output dim is not a multiple of 128 + can legitimately hold a scale tensor with more rows than the weight. A + sidecar carrying that must load rather than fail shape validation. + """ + source = _build_model(seed=19) + nv.convert_model_to_nvfp4(source) + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + scale_key = f"{_BLOCK.format(idx=0)}.attn.to_q::_nvfp4_weight_scale" + padded = _rewrite_sidecar(path, + tmp_path / "padded.safetensors", + tensors={scale_key: torch.zeros(128, (_IN_DIM + 15) // 16, dtype=torch.uint8)}) + assert nv.load_nvfp4_checkpoint(_build_model(seed=20), padded) == 4 + + +def test_saving_an_unconverted_model_raises(tmp_path) -> None: + model = _build_model(seed=17) + with pytest.raises(RuntimeError, match="convert_model_to_nvfp4"): + nv.save_nvfp4_checkpoint(model, tmp_path / "nvfp4.safetensors") + + +def test_a_model_with_no_nvfp4_layers_raises_with_the_prefix_hint(tmp_path) -> None: + """The silent-dense failure mode must be reported, not ignored.""" + empty = _build_model(seed=18, quant_config=nv.NVFP4Config()) + with pytest.raises(RuntimeError, match="for_minimax_h3"): + nv.save_nvfp4_checkpoint(empty, tmp_path / "nvfp4.safetensors") + with pytest.raises(RuntimeError, match="for_minimax_h3"): + nv.load_nvfp4_checkpoint(empty, tmp_path / "nvfp4.safetensors") + + +def test_sidecar_path_helper(tmp_path) -> None: + assert nv.nvfp4_sidecar_path_for("/models/h3/transformer.safetensors") == "/models/h3/transformer.nvfp4.safetensors" + assert nv.nvfp4_sidecar_path_for("/models/h3") == "/models/h3.nvfp4.safetensors" + assert nv.nvfp4_sidecar_path_for(str(tmp_path)) == str(tmp_path / "nvfp4.safetensors") + + +def test_read_metadata_rejects_a_foreign_file(tmp_path) -> None: + from safetensors.torch import save_file as _save + + plain = tmp_path / "plain.safetensors" + _save({"w": torch.zeros(2, 2)}, str(plain)) + with pytest.raises(ValueError, match="not a FastVideo NVFP4 sidecar"): + nv.read_nvfp4_sidecar_metadata(plain) diff --git a/fastvideo/tests/ops/quantization/test_quant_param_allowlist.py b/fastvideo/tests/ops/quantization/test_quant_param_allowlist.py new file mode 100644 index 0000000000..fe72d39ebe --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_quant_param_allowlist.py @@ -0,0 +1,126 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU-only contract tests for the ``fsdp_load`` zero-init allowlist. + +A quantization config registers scale tensors that no checkpoint carries +(``AbsMaxFP8`` -> ``scale_weight`` / ``scale_input``; see +``fastvideo/layers/quantization/absmax_fp8.py``). ``fsdp_load`` rejects every +model parameter that is absent from the incoming state dict, so enabling +``engine.quantization.transformer_quant`` used to abort the load with:: + + Unsupported new parameter: transformer_blocks.0.attn.to_out.scale_input + +These tests pin both halves of that contract, on CPU and without any +distributed process group: + +* quant scale parameters are accepted and zero-initialized, +* a parameter that is genuinely missing from the checkpoint still raises, so + the allowlist is not a blanket exemption. + +See ``docs/quantization/loader_quant_params.md`` before adding a quantization +config that registers new parameters. +""" + +import unittest + +import torch +import torch.nn as nn + +from fastvideo.layers.quantization.absmax_fp8 import AbsMaxFP8LinearMethod +from fastvideo.models.loader.fsdp_load import ( + ALLOWED_NEW_PARAM_PATTERNS, + is_allowed_new_param, + load_model_from_full_model_state_dict, +) + +DTYPE = torch.float32 +REPORTED_FQN = "transformer_blocks.0.attn.to_out.scale_input" + + +class _QuantToOut(nn.Module): + """A module holding an AbsMaxFP8 linear's parameters under their real names.""" + + def __init__(self, in_features: int = 3, out_features: int = 2) -> None: + super().__init__() + AbsMaxFP8LinearMethod().create_weights( + self, + input_size_per_partition=in_features, + output_partition_sizes=[out_features], + input_size=in_features, + output_size=out_features, + params_dtype=DTYPE, + ) + + +class _Block(nn.Module): + + def __init__(self) -> None: + super().__init__() + self.attn = _Attn() + + +class _Attn(nn.Module): + + def __init__(self) -> None: + super().__init__() + self.to_out = _QuantToOut() + + +def _model() -> nn.Module: + model = nn.Module() + model.add_module("transformer_blocks", nn.ModuleList([_Block()])) + return model + + +def _load(model: nn.Module, checkpoint: dict[str, torch.Tensor]): + return load_model_from_full_model_state_dict( + model, + ((name, tensor) for name, tensor in checkpoint.items()), + device=torch.device("cpu"), + param_dtype=DTYPE, + strict=False, + cpu_offload=False, + param_names_mapping=lambda name: (name, None, None), + training_mode=False, + ) + + +class TestQuantParamAllowlist(unittest.TestCase): + + def test_reported_fqn_is_allowed(self): + self.assertTrue(is_allowed_new_param(REPORTED_FQN)) + + def test_quant_scale_params_are_zero_initialized(self): + model = _model() + checkpoint = {"transformer_blocks.0.attn.to_out.weight": torch.ones(2, 3, dtype=DTYPE)} + _load(model, checkpoint) + + to_out = model.transformer_blocks[0].attn.to_out + self.assertEqual(to_out.scale_weight.shape, (1, )) + self.assertEqual(to_out.scale_input.shape, (1, )) + self.assertTrue(torch.equal(to_out.scale_weight, torch.zeros(1, dtype=DTYPE))) + self.assertTrue(torch.equal(to_out.scale_input, torch.zeros(1, dtype=DTYPE))) + self.assertTrue(torch.equal(to_out.weight, torch.ones(2, 3, dtype=DTYPE))) + + def test_missing_real_weight_still_raises(self): + model = _model() + with self.assertRaisesRegex(ValueError, "is not supported"): + _load(model, {}) + + def test_bare_scale_param_is_not_admitted(self): + self.assertFalse(is_allowed_new_param("transformer_blocks.0.attn.to_out.scale")) + + def test_every_absmax_fp8_registered_param_is_allowed(self): + to_out = _QuantToOut() + new_params = [name for name, _ in to_out.named_parameters() if name != "weight"] + self.assertEqual(sorted(new_params), ["scale_input", "scale_weight"]) + for name in new_params: + fqn = f"transformer_blocks.0.attn.to_out.{name}" + self.assertTrue(is_allowed_new_param(fqn), f"{fqn} is not in {ALLOWED_NEW_PARAM_PATTERNS}") + + def test_existing_attention_patterns_still_allowed(self): + for name in ("transformer_blocks.0.attn.to_gate_compress.weight", "blocks.0.attn1.attn_impl.proj_l.weight"): + self.assertTrue(is_allowed_new_param(name)) + + +if __name__ == "__main__": + unittest.main() diff --git a/fastvideo/tests/ops/quantization/test_quant_sidecar_roundtrip.py b/fastvideo/tests/ops/quantization/test_quant_sidecar_roundtrip.py new file mode 100644 index 0000000000..491741dea2 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_quant_sidecar_roundtrip.py @@ -0,0 +1,656 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU-only round-trip tests for the INT8 affine and W4A16 checkpoint sidecars. + +``convert_model_to_int8_affine`` / ``convert_model_to_w4a16`` register their +quantized tensors with ``persistent=False``, so a saved checkpoint carries +dense bf16 weights only and the low-bit payload is rebuilt at every load — by +re-running a conversion that starts from those dense weights. On the hardware +these lanes target (24-48 GB Ada parts, a 32 GB RTX 5090) H3's bf16 DiT does +not fit, so that rebuild is impossible and a pre-quantized checkpoint is the +only thing that can be served. ``save_*_checkpoint`` / ``load_*_checkpoint`` +write and restore that payload directly. + +No CUDA, no GPU, no flashinfer: INT8 affine and W4A16 are pure PyTorch, and +every test here runs on a laptop. The quantizers, the buffer registration, the +manifest and the load-time validation under test are the real ones. + +The load-bearing assertions are the ones about *silent* corruption: +``torch.equal`` (not ``allclose``) for the round trip, an exact dtype check on +the codes, and a shape check per tensor — a mis-read or mis-unpacked code +buffer produces plausible-looking garbage with no error anywhere. +""" + +from __future__ import annotations + +import json +import logging +import sys +from collections.abc import Callable +from typing import Any, NamedTuple + +import pytest +import torch +import torch.nn as nn +from safetensors import safe_open +from safetensors.torch import save_file + +from fastvideo.layers.linear import ReplicatedLinear +from fastvideo.layers.quantization import int8_affine_config as i8 +from fastvideo.layers.quantization import w4a16_config as w4 + +IN_DIM = 128 +OUT_DIM = 32 +GROUP_SIZE = 64 +_BLOCK = "minimax_h3.transformer_blocks.{idx}" +_Q = f"{_BLOCK.format(idx=0)}.attn.to_q" +_FF = f"{_BLOCK.format(idx=1)}.ff.fc_in" + + +class _Scheme(NamedTuple): + """Everything the two lanes do not share, addressed by name.""" + + name: str + module: Any + make_config: Callable[..., Any] + buffers: tuple[str, ...] + format_name: str + metadata_key: str + suffix: str + dir_name: str + convert: Callable[[nn.Module], None] + state_dict: Callable[[nn.Module], dict[str, torch.Tensor]] + save: Callable[..., dict[str, Any]] + load: Callable[..., int] + read_metadata: Callable[[Any], dict[str, Any]] + path_for: Callable[[Any], str] + wiring_hint: str + + +INT8 = _Scheme( + name="int8_affine", + module=i8, + make_config=i8.INT8AffineConfig, + buffers=i8._INT8_AFFINE_SIDECAR_BUFFERS, + format_name=i8._INT8_AFFINE_SIDECAR_FORMAT, + metadata_key=i8._INT8_AFFINE_SIDECAR_METADATA_KEY, + suffix=i8.INT8_AFFINE_SIDECAR_SUFFIX, + dir_name=i8.INT8_AFFINE_DIR_SIDECAR_NAME, + convert=i8.convert_model_to_int8_affine, + state_dict=i8.int8_affine_sidecar_state_dict, + save=i8.save_int8_affine_checkpoint, + load=i8.load_int8_affine_checkpoint, + read_metadata=i8.read_int8_affine_sidecar_metadata, + path_for=i8.int8_affine_sidecar_path_for, + wiring_hint="for_minimax_h3", +) + +W4A16 = _Scheme( + name="w4a16", + module=w4, + make_config=w4.W4A16Config, + buffers=w4._W4A16_SIDECAR_BUFFERS, + format_name=w4._W4A16_SIDECAR_FORMAT, + metadata_key=w4._W4A16_SIDECAR_METADATA_KEY, + suffix=w4.W4A16_SIDECAR_SUFFIX, + dir_name=w4.W4A16_DIR_SIDECAR_NAME, + convert=w4.convert_model_to_w4a16, + state_dict=w4.w4a16_sidecar_state_dict, + save=w4.save_w4a16_checkpoint, + load=w4.load_w4a16_checkpoint, + read_metadata=w4.read_w4a16_sidecar_metadata, + path_for=w4.w4a16_sidecar_path_for, + wiring_hint="for_minimax_h3", +) + +_ALL = (INT8, W4A16) +_by_name = pytest.mark.parametrize("scheme", _ALL, ids=[s.name for s in _ALL]) + + + + +def _build(scheme: _Scheme, *, num_blocks: int = 2, seed: int = 0, config=None) -> nn.Module: + """A miniature H3-shaped DiT: ``minimax_h3.transformer_blocks.{i}.{attn.to_q,ff.fc_in}``.""" + cfg = config if config is not None else scheme.make_config() + generator = torch.Generator().manual_seed(seed) + root = nn.Module() + dit = nn.Module() + blocks = nn.ModuleList() + for idx in range(num_blocks): + block = nn.Module() + attn = nn.Module() + attn.to_q = ReplicatedLinear(IN_DIM, + OUT_DIM, + bias=False, + quant_config=cfg, + prefix=f"{_BLOCK.format(idx=idx)}.attn.to_q") + ff = nn.Module() + ff.fc_in = ReplicatedLinear(IN_DIM, + OUT_DIM, + bias=False, + quant_config=cfg, + prefix=f"{_BLOCK.format(idx=idx)}.ff.fc_in") + block.attn = attn + block.ff = ff + blocks.append(block) + dit.transformer_blocks = blocks + root.minimax_h3 = dit + for _, param in root.named_parameters(): + param.data.copy_(torch.randn(param.shape, generator=generator)) + return root + + +def _tagged(model: nn.Module) -> dict[str, nn.Module]: + return {fqn: mod for fqn, mod in model.named_modules() if getattr(mod, "quant_method", None) is not None} + + +def _rewrite_sidecar(scheme: _Scheme, src, dst, *, metadata: dict | None = None, tensors: dict | None = None) -> str: + """Copy a sidecar, optionally replacing its manifest or some tensors.""" + with safe_open(src, framework="pt", device="cpu") as handle: + payload = {key: handle.get_tensor(key) for key in handle.keys()} + manifest = json.loads(handle.metadata()[scheme.metadata_key]) + if metadata is not None: + manifest.update(metadata) + if tensors is not None: + payload.update(tensors) + save_file(payload, str(dst), metadata={scheme.metadata_key: json.dumps(manifest)}) + return str(dst) + + +def _sidecar_of(scheme: _Scheme, model: nn.Module, path) -> str: + scheme.convert(model) + scheme.save(model, path) + return str(path) + + + + +@_by_name +def test_state_dict_does_not_carry_the_quantized_buffers(scheme: _Scheme) -> None: + """The whole reason for the sidecar: the buffers are non-persistent. + + A plain ``state_dict()`` is the dense bf16 weights and nothing else, so a + checkpoint written the ordinary way cannot be served on a host that cannot + hold the dense weights and re-convert. + """ + model = _build(scheme) + scheme.convert(model) + + keys = list(model.state_dict()) + assert keys, "the dense weights are persistent, so this must not be empty" + for key in keys: + for buffer_name in scheme.buffers: + assert buffer_name not in key, f"{key} leaked a non-persistent buffer" + assert set(keys) == {f"{fqn}.weight" for fqn in _tagged(model)} + for fqn, mod in _tagged(model).items(): + assert getattr(mod, scheme.buffers[0]) is not None, fqn + assert len(scheme.state_dict(model)) == 4 * len(scheme.buffers) + + +@_by_name +def test_sidecar_state_dict_uses_module_fqn_keys_and_cpu_tensors(scheme: _Scheme) -> None: + model = _build(scheme) + scheme.convert(model) + state = scheme.state_dict(model) + + assert f"{_Q}::{scheme.buffers[0]}" in state + assert all(value.device.type == "cpu" for value in state.values()) + assert {key.split("::", 1)[1] for key in state} == set(scheme.buffers) + assert state[f"{_Q}::{scheme.buffers[0]}"].dtype is torch.uint8 + for name in scheme.buffers[1:]: + assert state[f"{_Q}::{name}"].dtype is torch.float32 + + +def test_int8_codes_are_uint8_because_they_exceed_the_int8_range() -> None: + """Affine bits=8 codes span [0, 255]; int8 storage would wrap 255 to -1. + + Pinned on a crafted weight rather than a random one so the assertion is + about the scheme, not about a lucky seed. + """ + model = _build(INT8) + linear = _tagged(model)[_Q] + weight = torch.full((OUT_DIM, IN_DIM), -1.0) + weight[:, 0] = 3.0 # one large positive value per row + linear.weight.data.copy_(weight) + INT8.convert(model) + + codes = linear._int8_affine_codes + assert codes.dtype is torch.uint8 + assert codes.max().item() == 255 + as_int8 = codes.to(torch.int8) + assert as_int8.min().item() == -1 + assert not torch.equal(as_int8.to(torch.int64), codes.to(torch.int64)) + + + + +@_by_name +def test_save_then_load_into_a_fresh_module_is_bit_identical(scheme: _Scheme, tmp_path) -> None: + source = _build(scheme, seed=1) + scheme.convert(source) + expected = {key: value.clone() for key, value in scheme.state_dict(source).items()} + + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + target = _build(scheme, seed=2) + assert scheme.load(target, path) == 4 + actual = scheme.state_dict(target) + + assert set(actual) == set(expected) + for key, value in expected.items(): + assert actual[key].dtype == value.dtype, key + assert torch.equal(actual[key], value), f"{key} is not bit-identical (allclose would hide this)" + assert not torch.equal(target.minimax_h3.transformer_blocks[0].attn.to_q.weight, + source.minimax_h3.transformer_blocks[0].attn.to_q.weight) + + +@_by_name +def test_loaded_buffers_reproduce_the_source_forward(scheme: _Scheme, tmp_path) -> None: + """A loaded layer must compute what the converted source layer computes. + + This is the end-to-end payoff: for W4A16 it only holds if the load also + restored ``_w4a16_weight_shape``, which is a plain attribute and not a + buffer, so nothing else would carry it across. + """ + source = _build(scheme, seed=3) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + target = _build(scheme, seed=4) + scheme.load(target, path) + + x = torch.randn(2, IN_DIM, generator=torch.Generator().manual_seed(5)) + with torch.no_grad(): + want, _ = _tagged(source)[_Q](x) + got, _ = _tagged(target)[_Q](x) + assert torch.equal(got, want) + + +def test_w4a16_load_restores_the_logical_weight_shape_attribute(tmp_path) -> None: + """``_w4a16_weight_shape`` is not a buffer, so the manifest must carry it.""" + source = _build(W4A16, seed=6) + W4A16.convert(source) + path = tmp_path / "w4a16.safetensors" + W4A16.save(source, path) + + target = _build(W4A16, seed=7) + linear = _tagged(target)[_Q] + assert getattr(linear, "_w4a16_weight_shape", None) is None + W4A16.load(target, path) + assert tuple(linear._w4a16_weight_shape) == (OUT_DIM, IN_DIM) + assert linear._w4a16_codes.shape == (OUT_DIM, IN_DIM // 2) + + +@_by_name +def test_save_and_load_with_no_dense_weights_at_all(scheme: _Scheme, tmp_path) -> None: + """The target-hardware case: the bf16 weights are never present. + + With ``retain_original_weight=False`` the converted model keeps no dense + copy, and on a card that cannot hold them the loading model has none + either (``weight`` is ``None``). Save and load must both work from the + quantized buffers alone — this is the path that exercises the manifest's + per-layer ``weight_shape`` for the shape it cannot read off a weight. + """ + config = scheme.make_config(retain_original_weight=False) + source = _build(scheme, seed=34, config=config) + scheme.convert(source) + assert _tagged(source)[_Q].weight is None + + path = tmp_path / "sidecar.safetensors" + receipt = scheme.save(source, path) + assert receipt["num_layers"] == 4 + expected = {key: value.clone() for key, value in scheme.state_dict(source).items()} + + target = _build(scheme, seed=35, config=scheme.make_config(retain_original_weight=False)) + for mod in _tagged(target).values(): + mod.register_parameter("weight", None) + assert scheme.load(target, path) == 4 + + actual = scheme.state_dict(target) + assert set(actual) == set(expected) + assert all(torch.equal(actual[key], value) for key, value in expected.items()) + x = torch.randn(2, IN_DIM, generator=torch.Generator().manual_seed(36)) + with torch.no_grad(): + assert torch.equal(_tagged(target)[_Q](x)[0], _tagged(source)[_Q](x)[0]) + + +def test_w4a16_bits_8_sidecar_uses_the_unpacked_code_shape(tmp_path) -> None: + """``bits=8`` stores one code per byte, so the code shape is the weight shape. + + The manifest and the load-time shape check both have to follow ``bits`` + rather than assume the 4-bit packed layout. + """ + config = W4A16.make_config(bits=8) + source = _build(W4A16, seed=42, config=config) + W4A16.convert(source) + assert _tagged(source)[_Q]._w4a16_codes.shape == (OUT_DIM, IN_DIM) + + path = tmp_path / "w4a16_bits8.safetensors" + W4A16.save(source, path) + manifest = W4A16.read_metadata(path) + assert manifest["bits"] == 8 + assert manifest["layers"][_Q]["tensors"][W4A16.buffers[0]] == [OUT_DIM, IN_DIM] + + target = _build(W4A16, seed=43, config=W4A16.make_config(bits=8)) + assert W4A16.load(target, path) == 4 + assert torch.equal(W4A16.state_dict(target)[f"{_Q}::{W4A16.buffers[0]}"], + W4A16.state_dict(source)[f"{_Q}::{W4A16.buffers[0]}"]) + x = torch.randn(2, IN_DIM, generator=torch.Generator().manual_seed(44)) + with torch.no_grad(): + assert torch.equal(_tagged(target)[_Q](x)[0], _tagged(source)[_Q](x)[0]) + + +def test_w4a16_codes_preserve_the_low_nibble_first_packing(tmp_path) -> None: + """The packed bytes must round-trip the packer's nibble order exactly. + + The convention is the module's own (``_pack_4bit``): the low nibble is the + lower K index. Re-deriving the expected byte from the source layer's own + quantizer output is what makes this a check of the *serialized* bytes + rather than of the packing function. + """ + source = _build(W4A16, seed=8) + W4A16.convert(source) + linear = _tagged(source)[_Q] + codes = linear._w4a16_codes + weight = linear.weight.detach().float() + unpacked = w4._unpack_4bit(codes) + expected = w4._pack_4bit(unpacked) + assert torch.equal(codes, expected) + assert codes.shape == (OUT_DIM, IN_DIM // 2) + assert unpacked.shape == (OUT_DIM, IN_DIM) + raw, _, _ = w4.w4a16_quantize(weight, group_size=GROUP_SIZE, bits=4) + assert torch.equal(unpacked, w4._unpack_4bit(raw)) + + path = tmp_path / "w4a16.safetensors" + W4A16.save(source, path) + target = _build(W4A16, seed=9) + W4A16.load(target, path) + assert torch.equal(_tagged(target)[_Q]._w4a16_codes, codes) + + + + +@_by_name +def test_manifest_round_trips_scheme_and_layer_inventory(scheme: _Scheme, tmp_path) -> None: + model = _build(scheme) + scheme.convert(model) + path = tmp_path / "sidecar.safetensors" + receipt = scheme.save(model, path) + + manifest = scheme.read_metadata(path) + assert manifest["format"] == scheme.format_name + assert manifest["version"] == 1 + assert manifest["group_size"] == GROUP_SIZE + assert manifest["bits"] == (8 if scheme is INT8 else 4) + assert manifest["num_layers"] == 4 + assert set(manifest["layers"]) == {_Q, _FF, f"{_BLOCK.format(idx=0)}.ff.fc_in", + f"{_BLOCK.format(idx=1)}.attn.to_q"} + assert manifest["quant_prefixes"][_Q] == _Q + assert manifest["model_class"] == "Module" + + entry = manifest["layers"][_Q] + assert entry["weight_shape"] == [OUT_DIM, IN_DIM] + assert entry["group_size"] == GROUP_SIZE + assert entry["bits"] == (8 if scheme is INT8 else 4) + assert set(entry["tensors"]) == set(scheme.buffers) + + with safe_open(path, framework="pt", device="cpu") as handle: + for fqn, layer in manifest["layers"].items(): + for name, shape in layer["tensors"].items(): + assert list(handle.get_tensor(f"{fqn}::{name}").shape) == shape + + assert receipt["num_layers"] == 4 + assert receipt["num_tensors"] == 4 * len(scheme.buffers) + assert receipt["quantized_bytes"] == sum(t.numel() * t.element_size() for t in scheme.state_dict(model).values()) + assert receipt["quantized_bytes"] < receipt["dense_bfloat16_bytes"] + assert receipt["compression_ratio"] > 1.0 + + +@_by_name +def test_save_logs_a_receipt_with_the_module_count_and_bytes(scheme: _Scheme, tmp_path, caplog) -> None: + model = _build(scheme) + scheme.convert(model) + with caplog.at_level(logging.INFO): + receipt = scheme.save(model, tmp_path / "sidecar.safetensors") + message = "\n".join(record.message for record in caplog.records) + assert "4 quantized modules" in message + assert str(receipt["quantized_bytes"]) in message + assert "dense bf16" in message + + +@_by_name +def test_extra_metadata_is_merged_into_the_manifest(scheme: _Scheme, tmp_path) -> None: + model = _build(scheme) + scheme.convert(model) + path = tmp_path / "sidecar.safetensors" + scheme.save(model, path, extra_metadata={"source_checkpoint": "h3-bf16"}) + assert scheme.read_metadata(path)["source_checkpoint"] == "h3-bf16" + + + + +@_by_name +def test_wrong_code_dtype_is_rejected_at_load(scheme: _Scheme, tmp_path) -> None: + """Codes cast to int8 must fail, not silently dequantize to garbage.""" + source = _build(scheme, seed=10) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + key = f"{_Q}::{scheme.buffers[0]}" + with safe_open(path, framework="pt", device="cpu") as handle: + bad_codes = handle.get_tensor(key).to(torch.int8) + bad = _rewrite_sidecar(scheme, path, tmp_path / "int8codes.safetensors", tensors={key: bad_codes}) + + with pytest.raises(ValueError, match="dtype"): + scheme.load(_build(scheme, seed=11), bad) + with pytest.raises(ValueError, match="dtype"): + scheme.load(_build(scheme, seed=11), bad, strict=False) + + +@_by_name +def test_wrong_float_dtype_is_rejected_at_load(scheme: _Scheme, tmp_path) -> None: + """A bf16 scale store is not a bit-exact restore of the fp32 constants.""" + source = _build(scheme, seed=12) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + key = f"{_Q}::{scheme.buffers[1]}" + with safe_open(path, framework="pt", device="cpu") as handle: + bad_scales = handle.get_tensor(key).to(torch.bfloat16) + bad = _rewrite_sidecar(scheme, path, tmp_path / "bf16scales.safetensors", tensors={key: bad_scales}) + with pytest.raises(ValueError, match="dtype"): + scheme.load(_build(scheme, seed=13), bad) + + +@_by_name +def test_wrong_shape_is_rejected_at_load(scheme: _Scheme, tmp_path) -> None: + source = _build(scheme, seed=14) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + key = f"{_Q}::{scheme.buffers[0]}" + bad = _rewrite_sidecar(scheme, + path, + tmp_path / "shape.safetensors", + tensors={key: torch.zeros(OUT_DIM, IN_DIM * 2, dtype=torch.uint8)}) + with pytest.raises(ValueError, match="expected"): + scheme.load(_build(scheme, seed=15), bad) + with pytest.raises(ValueError, match="expected"): + scheme.load(_build(scheme, seed=15), bad, strict=False) + + +@_by_name +def test_scale_shape_mismatch_is_rejected_at_load(scheme: _Scheme, tmp_path) -> None: + """A regrouped scales tensor would regroup every code in the row.""" + source = _build(scheme, seed=16) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + key = f"{_Q}::{scheme.buffers[1]}" + bad = _rewrite_sidecar(scheme, + path, + tmp_path / "scales.safetensors", + tensors={key: torch.zeros(OUT_DIM, IN_DIM, dtype=torch.float32)}) + with pytest.raises(ValueError, match="expected"): + scheme.load(_build(scheme, seed=17), bad) + + +@_by_name +def test_group_size_and_bits_mismatches_are_fatal(scheme: _Scheme, tmp_path) -> None: + """A different scheme means different dequantize arithmetic over the bytes.""" + source = _build(scheme, seed=18) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + other_group = scheme.make_config(group_size=GROUP_SIZE * 2) + with pytest.raises(ValueError, match="group_size"): + scheme.load(_build(scheme, seed=19, config=other_group), path) + with pytest.raises(ValueError, match="group_size"): + scheme.load(_build(scheme, seed=19, config=other_group), path, strict=False) + + other_bits = scheme.make_config(bits=4 if scheme is INT8 else 8) + with pytest.raises(ValueError, match="bits"): + scheme.load(_build(scheme, seed=19, config=other_bits), path) + + +@_by_name +def test_weight_shape_that_disagrees_with_the_model_is_fatal(scheme: _Scheme, tmp_path) -> None: + """A sidecar describing a differently-shaped layer must not load into this one.""" + source = _build(scheme, seed=20) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + layers = dict(scheme.read_metadata(path)["layers"]) + layers[_Q] = {**layers[_Q], "weight_shape": [OUT_DIM, IN_DIM * 2]} + reshaped = _rewrite_sidecar(scheme, path, tmp_path / "ws.safetensors", metadata={"layers": layers}) + with pytest.raises(ValueError, match="shape"): + scheme.load(_build(scheme, seed=21), reshaped) + + +@_by_name +def test_malformed_layer_entry_is_reported_not_a_key_error(scheme: _Scheme, tmp_path) -> None: + """A hand-edited manifest must fail with a message naming the layer.""" + source = _build(scheme, seed=37) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + layers = dict(scheme.read_metadata(path)["layers"]) + layers[_Q] = {key: value for key, value in layers[_Q].items() if key != "group_size"} + malformed = _rewrite_sidecar(scheme, path, tmp_path / "malformed.safetensors", metadata={"layers": layers}) + with pytest.raises(ValueError, match="malformed"): + scheme.load(_build(scheme, seed=38), malformed) + + +@_by_name +def test_format_and_version_mismatches_are_fatal(scheme: _Scheme, tmp_path) -> None: + source = _build(scheme, seed=22) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + bad_format = _rewrite_sidecar(scheme, path, tmp_path / "format.safetensors", metadata={"format": "other"}) + with pytest.raises(ValueError, match="format"): + scheme.load(_build(scheme, seed=23), bad_format) + + bad_version = _rewrite_sidecar(scheme, path, tmp_path / "version.safetensors", metadata={"version": 99}) + with pytest.raises(ValueError, match="version"): + scheme.load(_build(scheme, seed=23), bad_version) + + +@_by_name +def test_incomplete_layer_entry_is_rejected(scheme: _Scheme, tmp_path) -> None: + """A layer missing its code buffer must not silently become dense.""" + source = _build(scheme, seed=24) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + with safe_open(path, framework="pt", device="cpu") as handle: + payload = {key: handle.get_tensor(key) for key in handle.keys()} + manifest = json.loads(handle.metadata()[scheme.metadata_key]) + del payload[f"{_Q}::{scheme.buffers[0]}"] + stripped = tmp_path / "stripped.safetensors" + save_file(payload, str(stripped), metadata={scheme.metadata_key: json.dumps(manifest)}) + + with pytest.raises(ValueError, match="incomplete"): + scheme.load(_build(scheme, seed=25), stripped) + + +@_by_name +def test_layer_set_mismatch_strict_and_lenient(scheme: _Scheme, tmp_path, caplog) -> None: + source = _build(scheme, num_blocks=1, seed=26) + path = tmp_path / "one.safetensors" + scheme.convert(source) + scheme.save(source, path) + + bigger = _build(scheme, num_blocks=2, seed=27) + with pytest.raises(ValueError, match="missing from the sidecar"): + scheme.load(bigger, path) + + with caplog.at_level(logging.WARNING): + assert scheme.load(bigger, path, strict=False) == 2 + assert any("does not match this model" in record.message for record in caplog.records) + assert getattr(bigger.minimax_h3.transformer_blocks[1].attn.to_q, scheme.buffers[0], None) is None + + wider = _build(scheme, num_blocks=3, seed=28) + with pytest.raises(ValueError, match="not in the model"): + scheme.load(_build(scheme, num_blocks=1, seed=29), _sidecar_of(scheme, wider, + tmp_path / "wide.safetensors")) + + +@_by_name +def test_read_metadata_rejects_a_foreign_file(scheme: _Scheme, tmp_path) -> None: + plain = tmp_path / "plain.safetensors" + save_file({"w": torch.zeros(2, 2)}, str(plain)) + with pytest.raises(ValueError, match="not a FastVideo"): + scheme.read_metadata(plain) + + +@_by_name +def test_saving_an_unconverted_model_raises(scheme: _Scheme, tmp_path) -> None: + model = _build(scheme, seed=30) + with pytest.raises(RuntimeError, match="convert_model_to"): + scheme.save(model, tmp_path / "sidecar.safetensors") + + +@_by_name +def test_a_model_with_no_tagged_layers_raises_with_the_wiring_hint(scheme: _Scheme, tmp_path) -> None: + """The silent-dense failure mode must be reported, not ignored.""" + empty = _build(scheme, seed=31, config=scheme.make_config(target_layers=["nothing.matches.this"])) + with pytest.raises(RuntimeError, match=scheme.wiring_hint): + scheme.save(empty, tmp_path / "sidecar.safetensors") + with pytest.raises(RuntimeError, match=scheme.wiring_hint): + scheme.load(empty, tmp_path / "sidecar.safetensors") + + + + +@_by_name +def test_load_needs_no_gpu_and_no_flashinfer(scheme: _Scheme, tmp_path, monkeypatch) -> None: + """Serving a pre-quantized checkpoint on the target host must not import kernels.""" + source = _build(scheme, seed=32) + scheme.convert(source) + expected = {key: value.clone() for key, value in scheme.state_dict(source).items()} + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + assert "flashinfer" not in sys.modules, "neither lane may pull in flashinfer" + target = _build(scheme, seed=33) + assert scheme.load(target, path) == 4 + assert "flashinfer" not in sys.modules + assert all(torch.equal(scheme.state_dict(target)[key], value) for key, value in expected.items()) + + +def test_sidecar_path_helpers(tmp_path) -> None: + for scheme in _ALL: + assert scheme.path_for("/models/h3/transformer.safetensors") == f"/models/h3/transformer{scheme.suffix}" + assert scheme.path_for("/models/h3") == f"/models/h3{scheme.suffix}" + assert scheme.path_for(str(tmp_path)) == str(tmp_path / scheme.dir_name) diff --git a/fastvideo/tests/ops/quantization/test_w4a16_config.py b/fastvideo/tests/ops/quantization/test_w4a16_config.py new file mode 100644 index 0000000000..6520341747 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_w4a16_config.py @@ -0,0 +1,434 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU-only tests for the W4A16 (4-bit weight / 16-bit activation) config. + +Scope: the quantizer round-trip, the packed-code layout, the MiniMax-H3 layer +selection (including the ``attn.to_gate_compress`` exclusion) and the load-time +conversion path. Nothing here needs CUDA, a GPU kernel, or flashinfer — the +compute path under test is the documented ``dequantize then dense GEMM`` +reference, which is pure PyTorch. + +Tolerance note: W4A16's 4-bit codes give a per-element reconstruction error +bounded by half the group's quantizer step, +``scale = (max - min) / 15``. Tests assert against that *derived* bound rather +than a hand-picked constant, so the assertion stays meaningful if the group +size changes. +""" + +from __future__ import annotations + +import pytest +import torch +import torch.nn as nn +import torch.nn.functional as F + +from fastvideo.layers.linear import ReplicatedLinear +from fastvideo.layers.quantization.w4a16_config import ( + DEFAULT_BITS, + DEFAULT_GROUP_SIZE, + MINIMAX_H3_MAIN_STACK_LINEAR_SUFFIXES, + W4A16Config, + W4A16QuantizeMethod, + convert_model_to_w4a16, + minimax_h3_w4a16_prefixes, + w4a16_dequantize, + w4a16_quantize, +) + + +def _random_weight(out_dim: int, in_dim: int, *, generator: torch.Generator) -> torch.Tensor: + """Deterministic, roughly-linear-layer-shaped weights (non-uniform rows).""" + return torch.randn(out_dim, in_dim, generator=generator) * 0.02 + + + + +def test_config_imports_without_cuda_dependencies(): + """Importing the config must not require CUDA, a kernel, or flashinfer.""" + config = W4A16Config() + assert config.get_name() == "W4A16" + assert torch.bfloat16 in config.get_supported_act_dtypes() + assert config.get_config_filenames() == [] + assert W4A16Config.get_min_capability() >= 70 + assert DEFAULT_BITS == 4 + + +def test_config_rejects_unsupported_bits(): + with pytest.raises(ValueError): + W4A16Config(bits=3) + with pytest.raises(ValueError): + W4A16Config(group_size=0) + + +def test_from_config_round_trips_fields(): + config = W4A16Config.from_config({ + "group_size": 32, + "bits": 4, + "target_layers": ["a.b"], + "exclude_substrings": ["keep_me_dense"], + "retain_original_weight": False, + }) + assert config.group_size == 32 + assert config.target_layers == frozenset({"a.b"}) + assert "keep_me_dense" in config.exclude_substrings + assert config.retain_original_weight is False + + + + +@pytest.mark.parametrize("group_size", [32, 64, 128]) +def test_quantize_dequantize_round_trip_within_step_bound(group_size: int): + """Every reconstructed element is within half a quantizer step of the source.""" + generator = torch.Generator().manual_seed(0) + weight = _random_weight(48, 4 * group_size, generator=generator) + + codes, scales, zeros = w4a16_quantize(weight, group_size=group_size, bits=4) + restored = w4a16_dequantize(codes, scales, zeros, group_size=group_size, bits=4, out_shape=weight.shape) + + assert codes.dtype == torch.uint8 + assert codes.shape == (weight.shape[0], weight.shape[1] // 2) + assert scales.shape == (weight.shape[0], weight.shape[1] // group_size) + assert zeros.shape == scales.shape + assert restored.shape == weight.shape + + error = (restored - weight).abs() + per_element_bound = (scales.unsqueeze(-1) / 2 + 1e-6) + assert torch.all(error.reshape(weight.shape[0], -1, group_size) <= per_element_bound) + + +def test_round_trip_is_exact_for_uniform_group(): + """A group that lands exactly on the code grid reconstructs exactly. + + ``zero`` is a rounded integer, so the anchor is only exact when + ``-min / scale`` is already integral -- here the ladder spans codes 0..15 + with a step of 0.1, giving ``zero = 8`` exactly. + """ + group_size = 64 + ladder = ((torch.arange(16).float() - 8) * 0.1).repeat(group_size // 16) + weight = ladder.repeat(2, 1) + + codes, scales, zeros = w4a16_quantize(weight, group_size=group_size, bits=4) + restored = w4a16_dequantize(codes, scales, zeros, group_size=group_size, bits=4, out_shape=weight.shape) + + assert torch.allclose(restored, weight, atol=1e-6) + + +def test_round_trip_holds_across_activation_dtypes(): + """bf16/fp16 sources still reconstruct within the step bound.""" + group_size = 64 + generator = torch.Generator().manual_seed(1) + base = _random_weight(16, group_size * 2, generator=generator) + for dtype in (torch.bfloat16, torch.float16, torch.float32): + weight = base.to(dtype) + codes, scales, zeros = w4a16_quantize(weight, group_size=group_size, bits=4) + restored = w4a16_dequantize(codes, scales, zeros, group_size=group_size, bits=4, out_shape=weight.shape) + error = (restored - weight.float()).abs() + bound = (scales.unsqueeze(-1) / 2 + 1e-2) + assert torch.all(error.reshape(weight.shape[0], -1, group_size) <= bound), dtype + + +def test_quantize_rejects_indivisible_group_size(): + weight = torch.randn(4, 100) + with pytest.raises(ValueError, match="not divisible"): + w4a16_quantize(weight, group_size=64) + + +def test_packed_codes_store_four_bits_per_weight(): + """The 4-bit storage claim is real: 4 bits/weight for the codes themselves. + + (Excludes the per-group scales/zeros, which add ``2 * 32 / group_size`` + bits per weight -- 1 bit/weight at ``group_size=64``.) + """ + weight = torch.randn(32, 128) + codes, scales, zeros = w4a16_quantize(weight, group_size=64, bits=4) + assert codes.numel() == weight.numel() // 2 + assert codes.element_size() == 1 + assert codes.numel() * 8 == weight.numel() * 4 + assert scales.numel() == zeros.numel() == weight.numel() // 64 + + +def test_packed_nibble_layout_is_low_first(): + """Low nibble holds the lower K index -- the documented packing order. + + Asserted structurally rather than against hand-computed codes: a strictly + increasing row must unpack to a non-decreasing code sequence. Swapping the + nibble order would make the unpacked sequence zig-zag within every pair, + so this distinguishes the two conventions without depending on where the + zero-point rounding lands. + """ + from fastvideo.layers.quantization.w4a16_config import _unpack_4bit + + group_size = 64 + weight = torch.linspace(-1.0, 1.0, group_size).unsqueeze(0) + codes, _, _ = w4a16_quantize(weight, group_size=group_size, bits=4) + unpacked = _unpack_4bit(codes) + + assert unpacked.shape == weight.shape + assert unpacked.dtype == torch.uint8 + deltas = unpacked[0, 1:].to(torch.int16) - unpacked[0, :-1].to(torch.int16) + assert torch.all(deltas >= 0), unpacked[0] + + +def test_nan_does_not_poison_the_group(): + """nan_to_num matches the other configs: one NaN must not nuke its group.""" + group_size = 64 + weight = torch.randn(1, group_size) + weight[0, 3] = float("nan") + codes, scales, zeros = w4a16_quantize(weight, group_size=group_size, bits=4) + restored = w4a16_dequantize(codes, scales, zeros, group_size=group_size, bits=4, out_shape=weight.shape) + assert torch.isfinite(restored).all() + + + + +def test_h3_allowlist_is_362_linears(): + config = W4A16Config.for_minimax_h3() + assert len(config.target_layers) == 362 + assert len(config.target_layers) == 50 * 7 + 2 * 6 + assert len(minimax_h3_w4a16_prefixes()) == 362 + + +def test_h3_selection_includes_the_intended_linears(): + config = W4A16Config.for_minimax_h3() + included = [ + "minimax_h3.transformer_blocks.0.attn.to_q", + "minimax_h3.transformer_blocks.0.attn.to_k", + "minimax_h3.transformer_blocks.0.attn.to_v", + "minimax_h3.transformer_blocks.0.attn.to_out", + "minimax_h3.transformer_blocks.0.ff.fc_in", + "minimax_h3.transformer_blocks.0.ff.fc_out", + "minimax_h3.transformer_blocks.0.adaln_proj.linear", + "minimax_h3.transformer_blocks.49.attn.to_q", + "minimax_h3.transformer_blocks.49.ff.fc_out", + "minimax_h3.token_refiner.refiner_blocks.0.attn.to_q", + "minimax_h3.token_refiner.refiner_blocks.1.ff.fc_out", + ] + for prefix in included: + assert config.is_target_layer(prefix), prefix + assert "adaln_proj.linear" in MINIMAX_H3_MAIN_STACK_LINEAR_SUFFIXES + + +def test_h3_selection_excludes_to_gate_compress(): + """The VSA gate steers discrete sparse routing -- it must never be quantized.""" + config = W4A16Config.for_minimax_h3() + for index in (0, 17, 49): + assert not config.is_target_layer(f"minimax_h3.transformer_blocks.{index}.attn.to_gate_compress") + assert not config.is_target_layer("minimax_h3.token_refiner.refiner_blocks.0.attn.to_gate_compress") + + +def test_to_gate_compress_exclusion_cannot_be_widened_away(): + """A caller cannot opt the gate back in, even by naming it in the allowlist.""" + gate = "minimax_h3.transformer_blocks.0.attn.to_gate_compress" + config = W4A16Config(target_layers=[gate, "minimax_h3.transformer_blocks.0.attn.to_q"]) + assert not config.is_target_layer(gate) + assert config.is_target_layer("minimax_h3.transformer_blocks.0.attn.to_q") + + +def test_h3_gate_module_is_built_and_left_dense(monkeypatch): + """The strongest form of the gate check: the real H3 attention module. + + H3 only builds ``attn.to_gate_compress`` when the VSA backend resolves + (``use_vsa`` guards the construction), which does not happen on a CPU-only + host. Forcing the backend resolution is enough to get the module built and + observe which quant method it receives -- this is the wiring the exclusion + actually has to protect, not just the string predicate. + """ + from fastvideo.layers.linear import UnquantizedLinearMethod + + import fastvideo.models.dits.minimax_h3 as h3_module + from fastvideo.platforms import AttentionBackendEnum + + class _FakeVSABackend: + + def get_name(self) -> str: + return "VIDEO_SPARSE_ATTN_H3" + + monkeypatch.setattr(h3_module, "get_attn_backend", lambda *args, **kwargs: _FakeVSABackend()) + + prefix = "minimax_h3.transformer_blocks.0.attn" + attention = h3_module.MiniMaxH3Attention( + 64, + 2, + 32, + 1e-5, + (AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3, ), + W4A16Config.for_minimax_h3(), + prefix=prefix, + ) + assert attention.to_gate_compress is not None + assert isinstance(attention.to_gate_compress.quant_method, UnquantizedLinearMethod) + assert isinstance(attention.to_q.quant_method, W4A16QuantizeMethod) + assert isinstance(attention.to_out.quant_method, W4A16QuantizeMethod) + + +def test_h3_selection_excludes_fp32_pinned_modules(): + config = W4A16Config.for_minimax_h3() + for prefix in ( + "minimax_h3.proj_in", + "minimax_h3.audio_proj_in", + "minimax_h3.time_embedder.fc_in", + "minimax_h3.proj_out", + "minimax_h3.audio_proj_out", + "minimax_h3.adaln_basis", + "minimax_h3.norm_out.linear", + "minimax_h3.context_embedder", + ): + assert not config.is_target_layer(prefix), prefix + + +def test_generic_suffix_matching_respects_dot_boundaries(): + """``ff.fc_in`` must not match a hypothetical ``cross_ff.fc_in``.""" + config = W4A16Config() + assert config.is_target_layer("model.blocks.0.ff.fc_in") + assert not config.is_target_layer("model.blocks.0.cross_ff.fc_in") + + +def test_get_quant_method_skips_indivisible_and_odd_dims(): + """A group must fit inside one row and 4-bit codes need an even last dim.""" + config = W4A16Config(target_layers=["blk.a", "blk.b", "blk.c"], group_size=64) + assert isinstance(config.get_quant_method(ReplicatedLinear(64, 8, bias=False, prefix="blk.a"), "blk.a"), + W4A16QuantizeMethod) + assert config.get_quant_method(ReplicatedLinear(100, 8, bias=False, prefix="blk.b"), "blk.b") is None + odd = W4A16Config(target_layers=["blk.c"], group_size=5) + assert odd.get_quant_method(ReplicatedLinear(25, 8, bias=False, prefix="blk.c"), "blk.c") is None + + +def test_get_quant_method_ignores_non_linear_layers(): + config = W4A16Config() + assert config.get_quant_method(nn.RMSNorm(64), "minimax_h3.transformer_blocks.0.attn.norm_q") is None + + + + +def _tiny_model(prefix: str, quant_config: W4A16Config, in_dim: int = 64, out_dim: int = 32): + linear = ReplicatedLinear(in_dim, out_dim, bias=False, quant_config=quant_config, prefix=prefix) + module = nn.Module() + module.add_module("linear", linear) + return module, linear + + +def test_convert_registers_non_persistent_buffers_and_keeps_weight(): + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"]) + model, linear = _tiny_model("block.linear", config) + linear.weight.data.normal_() + + convert_model_to_w4a16(model) + + assert linear._w4a16_codes.dtype == torch.uint8 + assert tuple(linear._w4a16_weight_shape) == tuple(linear.weight.shape) + assert list(model.state_dict().keys()) == ["linear.weight"] + assert linear.weight is not None # retained by default + + +def test_apply_matches_the_documented_dequantize_then_gemm_reference(): + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"]) + model, linear = _tiny_model("block.linear", config) + linear.weight.data.normal_() + convert_model_to_w4a16(model) + + x = torch.randn(4, 64) + reference = F.linear( + x, + w4a16_dequantize(linear._w4a16_codes, + linear._w4a16_scales, + linear._w4a16_zeros, + out_shape=linear._w4a16_weight_shape, + out_dtype=x.dtype), + ) + out, _ = linear(x) + assert torch.equal(out, reference) + + +def test_apply_converts_lazily_when_the_loader_hook_never_ran(): + """No conversion hook -> still correct, just later and noisier.""" + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"]) + linear = ReplicatedLinear(64, 32, bias=False, quant_config=config, prefix="block.linear") + linear.weight.data.normal_() + assert getattr(linear, "_w4a16_codes", None) is None + + x = torch.randn(2, 64) + with torch.no_grad(): + out, _ = linear(x) + assert out.shape == (2, 32) + assert getattr(linear, "_w4a16_codes", None) is not None + + +def test_grad_enabled_forward_stays_dense(): + """Under grad the master weight must be used -- no frozen dequantized copy.""" + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"]) + linear = ReplicatedLinear(64, 32, bias=False, quant_config=config, prefix="block.linear") + linear.weight.data.normal_() + + x = torch.randn(2, 64) + out, _ = linear(x) # grad enabled by default in pytest + assert out.shape == (2, 32) + assert getattr(linear, "_w4a16_codes", None) is None + + +def test_purging_frees_the_dense_weight_when_opted_in(): + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"], retain_original_weight=False) + model, linear = _tiny_model("block.linear", config) + linear.weight.data.normal_() + + convert_model_to_w4a16(model) + + assert linear.weight is None + assert "linear.weight" not in model.state_dict() + assert linear._w4a16_codes is not None + out, _ = linear(torch.randn(2, 64)) + assert out.shape == (2, 32) + + +def test_untargeted_layer_is_untouched_by_convert(): + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"]) + model, linear = _tiny_model("block.linear", config) + dense = ReplicatedLinear(64, 32, bias=False, prefix="block.other") + model.add_module("other", dense) + linear.weight.data.normal_() + dense.weight.data.normal_() + + before = dense.weight.clone() + convert_model_to_w4a16(model) + + assert getattr(dense, "_w4a16_codes", None) is None + assert torch.equal(dense.weight, before) + + +def test_bias_is_preserved_on_the_quantized_path(): + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"]) + linear = ReplicatedLinear(64, 32, bias=True, quant_config=config, prefix="block.linear") + model = nn.Module() + model.add_module("linear", linear) + linear.weight.data.normal_() + + convert_model_to_w4a16(model) + x = torch.randn(3, 64) + out, out_bias = linear(x) + assert out_bias is None + reference = F.linear( + x, + w4a16_dequantize(linear._w4a16_codes, + linear._w4a16_scales, + linear._w4a16_zeros, + out_shape=linear._w4a16_weight_shape, + out_dtype=x.dtype), linear.bias) + assert torch.equal(out, reference) + + +def test_default_group_size_divides_every_h3_targeted_input_dim(): + """The H3 profile's group size is only valid if it divides all its K dims. + + Documented H3 dims (``MiniMaxH3ArchConfig``): hidden 5376, attention inner + 50 * 128 = 7168, ffn 14336, adaln 2688. This test pins that arithmetic so a + config change that breaks divisibility fails here rather than at load time. + """ + h3_input_dims = (5376, 7168, 14336, 2688) + for dim in h3_input_dims: + assert dim % DEFAULT_GROUP_SIZE == 0, dim + assert dim % 2 == 0, dim diff --git a/fastvideo/tests/train/callbacks/test_callback.py b/fastvideo/tests/train/callbacks/test_callback.py index 13a022a06b..e1250177a9 100644 --- a/fastvideo/tests/train/callbacks/test_callback.py +++ b/fastvideo/tests/train/callbacks/test_callback.py @@ -23,10 +23,6 @@ ) -# ``fastvideo.logger.init_logger`` sets ``propagate=False`` on its -# loggers, so the standard ``caplog`` fixture cannot observe them. -# This helper attaches a temporary handler directly to the target -# logger and yields the captured records. @contextmanager def _capture_logger(name: str, level: int = logging.WARNING) -> Iterator[list[logging.LogRecord]]: logger = logging.getLogger(name) @@ -48,9 +44,6 @@ def emit(self, record: logging.LogRecord) -> None: logger.setLevel(prev_level) -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- class _RecordingCallback(Callback): @@ -83,9 +76,6 @@ def __init__(self, **_: Any) -> None: pass -# --------------------------------------------------------------------------- -# A. Callback base class -# --------------------------------------------------------------------------- class TestCallbackBase: @@ -96,19 +86,16 @@ def test_default_hooks_return_none(self) -> None: assert (cb.on_training_step_end(method=None, loss_dict={}) is None) assert cb.on_before_optimizer_step(method=None) is None assert cb.on_validation_begin(method=None) is None + assert cb.will_run_validation() is False assert cb.on_validation_end(method=None) is None assert cb.on_train_end(method=None) is None def test_default_state_dict_round_trip(self) -> None: cb = Callback() assert cb.state_dict() == {} - # Default load_state_dict accepts arbitrary state without raising. assert cb.load_state_dict({"unrelated": 1}) is None -# --------------------------------------------------------------------------- -# B. CallbackDict construction -# --------------------------------------------------------------------------- class TestCallbackDictInit: @@ -118,7 +105,6 @@ def test_empty_config(self) -> None: assert cb_dict._callbacks == {} def test_builtin_name_resolves_without_target(self) -> None: - # ``grad_clip`` is a registered builtin. cfg = {"grad_clip": {"max_grad_norm": 0.5}} tc = object() cb_dict = CallbackDict(cfg, training_config=tc) @@ -128,7 +114,6 @@ def test_builtin_name_resolves_without_target(self) -> None: GradNormClipCallback, ) cb = cb_dict._callbacks["grad_clip"] assert isinstance(cb, GradNormClipCallback) - # CallbackDict wires up training_config + back-pointer. assert cb.training_config is tc assert cb._callback_dict is cb_dict @@ -161,7 +146,6 @@ def test_non_callback_target_raises(self) -> None: CallbackDict(cfg, training_config=object()) def test_builtin_registry_has_expected_entries(self) -> None: - # Sanity: protect the builtin registry from silent shrinkage. assert set(_BUILTIN_CALLBACKS) >= { "grad_clip", "validation", @@ -169,9 +153,6 @@ def test_builtin_registry_has_expected_entries(self) -> None: } -# --------------------------------------------------------------------------- -# C. Dispatch via __getattr__ -# --------------------------------------------------------------------------- class TestCallbackDictDispatch: @@ -196,8 +177,6 @@ def test_dispatch_to_hook_some_callbacks_skip(self) -> None: sink: list[str] = [] cb_dict = self._build(sink) - # The base Callback subclass below only implements one hook; - # dispatch should still fan out without raising. class _OnlyValidation(Callback): @@ -209,7 +188,6 @@ def on_validation_end(self, method, iteration: int = 0) -> None: assert "vend:11" in sink def test_dispatch_unknown_hook_is_noop(self) -> None: - # Methods that don't exist on any callback should not raise. cb_dict = self._build([]) cb_dict.totally_made_up_hook(method=None, iteration=0) @@ -219,9 +197,6 @@ def test_underscore_attribute_raises(self) -> None: getattr(cb_dict, "_does_not_exist") -# --------------------------------------------------------------------------- -# D. state_dict / load_state_dict -# --------------------------------------------------------------------------- class TestCallbackDictStateDict: diff --git a/fastvideo/tests/train/callbacks/test_ema.py b/fastvideo/tests/train/callbacks/test_ema.py index 3e84079aa3..786319ad13 100644 --- a/fastvideo/tests/train/callbacks/test_ema.py +++ b/fastvideo/tests/train/callbacks/test_ema.py @@ -16,9 +16,6 @@ from fastvideo.train.callbacks.ema import EMACallback -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- class _Student: @@ -47,6 +44,18 @@ def __init__( self.tracker = tracker +class _CriticOnlyMethod(_Method): + + def __init__(self, transformer: torch.nn.Module) -> None: + super().__init__(transformer) + self._student_optimizer = object() + self._critic_optimizer = object() + + def get_optimizers(self, iteration: int) -> list[object]: + del iteration + return [self._critic_optimizer] + + def _tiny_transformer(*, fill: float = 0.0) -> torch.nn.Module: m = torch.nn.Linear(4, 2, bias=False) with torch.no_grad(): @@ -54,9 +63,6 @@ def _tiny_transformer(*, fill: float = 0.0) -> torch.nn.Module: return m -# --------------------------------------------------------------------------- -# A. on_train_start -# --------------------------------------------------------------------------- class TestOnTrainStart: @@ -67,7 +73,6 @@ def test_initializes_ema_from_student(self) -> None: cb.on_train_start(_Method(transformer), iteration=0) assert cb.student_ema is not None - # Shadow shape matches transformer parameter. shadow = cb.student_ema.shadow["weight"] assert shadow.shape == transformer.weight.shape assert torch.allclose(shadow, transformer.weight.detach().cpu()) @@ -78,16 +83,12 @@ def test_missing_transformer_raises(self) -> None: cb.on_train_start(_Method(transformer=None), iteration=0) -# --------------------------------------------------------------------------- -# B. on_training_step_end (decay math + start_iter gating) -# --------------------------------------------------------------------------- class TestOnTrainingStepEnd: def test_no_op_before_train_start(self) -> None: cb = EMACallback() - # student_ema is None until on_train_start. cb.on_training_step_end(_Method(transformer=None), loss_dict={}, iteration=0) assert not cb._ema_started @@ -96,31 +97,38 @@ def test_skipped_until_start_iter(self) -> None: cb = EMACallback(decay=0.5, start_iter=10) cb.on_train_start(_Method(transformer), iteration=0) - # Mutate transformer to drift it away from initial shadow. with torch.no_grad(): transformer.weight.fill_(7.0) cb.on_training_step_end(_Method(transformer), loss_dict={}, iteration=5) - # Below start_iter: shadow is untouched, _ema_started False. assert not cb._ema_started assert torch.allclose( cb.student_ema.shadow["weight"], torch.full((2, 4), 1.0), ) + def test_skips_iterations_without_student_optimizer(self) -> None: + transformer = _tiny_transformer(fill=1.0) + cb = EMACallback(decay=0.5, start_iter=0) + method = _CriticOnlyMethod(transformer) + cb.on_train_start(method, iteration=0) + + with torch.no_grad(): + transformer.weight.fill_(7.0) + cb.on_training_step_end(method, loss_dict={}, iteration=1) + + assert not cb._ema_started + assert torch.allclose(cb.student_ema.shadow["weight"], torch.full((2, 4), 1.0)) + def test_first_active_step_reinits_then_updates(self) -> None: transformer = _tiny_transformer(fill=1.0) cb = EMACallback(decay=0.9, start_iter=10) cb.on_train_start(_Method(transformer), iteration=0) - # Drift transformer so that re-init has a visible effect. with torch.no_grad(): transformer.weight.fill_(5.0) cb.on_training_step_end(_Method(transformer), loss_dict={}, iteration=10) - # First active step: shadow is re-initialized from the - # current transformer (5.0) and *then* update() applies decay - # against the same value, so shadow stays at 5.0. assert cb._ema_started assert torch.allclose( cb.student_ema.shadow["weight"], @@ -132,10 +140,7 @@ def test_subsequent_step_applies_decay(self) -> None: cb = EMACallback(decay=0.9, start_iter=0) cb.on_train_start(_Method(transformer), iteration=0) - # Step 0: re-init at 2.0, then update against 2.0 → still 2.0. cb.on_training_step_end(_Method(transformer), loss_dict={}, iteration=0) - # Step 1: drift transformer to 12.0, expect - # shadow = 0.9 * 2.0 + 0.1 * 12.0 = 3.0. with torch.no_grad(): transformer.weight.fill_(12.0) cb.on_training_step_end(_Method(transformer), loss_dict={}, iteration=1) @@ -156,9 +161,6 @@ def test_tracker_logs_decay(self) -> None: assert any(payload.get("ema/decay") == 0.99 and step == 0 for payload, step in tracker.entries) -# --------------------------------------------------------------------------- -# C. ema_context -# --------------------------------------------------------------------------- class TestEmaContext: @@ -166,7 +168,6 @@ class TestEmaContext: def test_passthrough_when_inactive(self) -> None: transformer = _tiny_transformer(fill=3.0) cb = EMACallback() - # No on_train_start → student_ema is None. with cb.ema_context(transformer) as t: assert t is transformer assert torch.allclose(t.weight, torch.full((2, 4), 3.0)) @@ -177,10 +178,7 @@ def test_swaps_weights_then_restores(self) -> None: method = _Method(transformer) cb.on_train_start(method, iteration=0) - # decay=0 → after one step the shadow == current weights == 1.0. cb.on_training_step_end(method, loss_dict={}, iteration=0) - # Drift transformer; ema_context should swap shadow (1.0) in - # for the duration and restore the post-drift value (9.0). with torch.no_grad(): transformer.weight.fill_(9.0) @@ -190,9 +188,6 @@ def test_swaps_weights_then_restores(self) -> None: assert torch.allclose(transformer.weight, torch.full((2, 4), 9.0)) -# --------------------------------------------------------------------------- -# D. State dict round-trip -# --------------------------------------------------------------------------- class TestStateDict: @@ -212,10 +207,8 @@ def test_round_trip_preserves_shadow_and_started_flag(self) -> None: assert "student_ema" in state assert state["ema_started"] is True - # Build a fresh callback and load. fresh = EMACallback(decay=0.5, start_iter=0) fresh.on_train_start(_Method(_tiny_transformer(fill=0.0)), iteration=0) - # Sanity: fresh shadow != saved shadow before load. assert not torch.allclose( fresh.student_ema.shadow["weight"], cb.student_ema.shadow["weight"], @@ -229,7 +222,6 @@ def test_round_trip_preserves_shadow_and_started_flag(self) -> None: def test_load_without_student_ema_only_sets_flag(self) -> None: cb = EMACallback() - # student_ema is None — load must not attempt to assign shadow. cb.load_state_dict({"ema_started": True}) assert cb._ema_started is True assert cb.student_ema is None diff --git a/fastvideo/tests/train/callbacks/test_latent_vis_shape_context.py b/fastvideo/tests/train/callbacks/test_latent_vis_shape_context.py new file mode 100644 index 0000000000..2f3ec3d9ac --- /dev/null +++ b/fastvideo/tests/train/callbacks/test_latent_vis_shape_context.py @@ -0,0 +1,50 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Latent visualization keeps native packed-shape context.""" + +from types import SimpleNamespace + +import torch + +from fastvideo.train.callbacks.latent_vis import LatentVisCallback + + +def test_latent_vis_passes_batch_local_layout_to_decoder(monkeypatch) -> None: + layout = object() + latent = torch.ones(1, 9) + decoded: list[tuple[torch.Tensor, object]] = [] + + class Student: + + @staticmethod + def decode_vis_latents(value, *, layout): + decoded.append((value, layout)) + return torch.zeros(1, 1, 3, 2, 2, dtype=torch.uint8).numpy() + + class Tracker: + + def video(self, clip, *, fps, format): + assert fps == 24 + assert format == "mp4" + return clip + + def log_artifacts(self, artifacts, iteration): + assert set(artifacts) == {"latent_vis/generator_pred_video"} + assert iteration == 8 + + monkeypatch.setattr( + "fastvideo.train.callbacks.latent_vis.get_world_group", + lambda: SimpleNamespace(rank=0), + ) + callback = LatentVisCallback(every_steps=1, keys=["generator_pred_video"]) + callback.tracker = Tracker() + method = SimpleNamespace( + student=Student(), + latent_vis={ + "generator_pred_video": latent, + "_fv_latent_layout": layout, + }, + ) + + callback.on_training_step_end(method, {}, iteration=8) + + assert decoded == [(latent, layout)] diff --git a/fastvideo/tests/train/callbacks/test_validation.py b/fastvideo/tests/train/callbacks/test_validation.py index e8baf8d1df..1f0cb63589 100644 --- a/fastvideo/tests/train/callbacks/test_validation.py +++ b/fastvideo/tests/train/callbacks/test_validation.py @@ -21,6 +21,7 @@ import torch from fastvideo.api.sampling_param import SamplingParam +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_input_preparation import resolve_target_num_frames from fastvideo.train.callbacks.callback import CallbackDict from fastvideo.train.callbacks.ema import EMACallback from fastvideo.train.callbacks.validation import ( @@ -31,9 +32,6 @@ _ValidationMetricStats, ) -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- _PIPE_TARGET = "fastvideo.pipelines.basic.wan.wan_pipeline.WanPipeline" @@ -44,6 +42,8 @@ def _make_callback( sampling_steps: list[int] | None = None, guidance_scale: float | None = None, num_frames: int | None = None, + use_record_dimensions: bool = False, + max_record_num_frames: int | None = None, num_videos_per_prompt: int = 1, use_validation_media_conditioning: bool = True, sampling_timesteps: list[int] | None = None, @@ -59,6 +59,8 @@ def _make_callback( sampling_steps=sampling_steps, guidance_scale=guidance_scale, num_frames=num_frames, + use_record_dimensions=use_record_dimensions, + max_record_num_frames=max_record_num_frames, num_videos_per_prompt=num_videos_per_prompt, use_validation_media_conditioning=use_validation_media_conditioning, sampling_timesteps=sampling_timesteps, @@ -68,9 +70,6 @@ def _make_callback( ) -# --------------------------------------------------------------------------- -# A. Constructor coercions / defaults -# --------------------------------------------------------------------------- class TestConstructor: @@ -84,6 +83,8 @@ def test_defaults(self) -> None: assert cb.sampling_steps == [40] assert cb.guidance_scale is None assert cb.num_frames is None + assert cb.use_record_dimensions is False + assert cb.max_record_num_frames is None assert cb.num_videos_per_prompt == 1 assert cb.use_validation_media_conditioning is True assert cb.run_at_start is True @@ -93,7 +94,6 @@ def test_defaults(self) -> None: assert cb.offload_training_state is False assert cb.unload_pipeline_after_validation is False assert cb.attn_qat_infer is False - # Lazy fields not yet populated. assert cb._pipeline is None assert cb._sampling_param is None assert cb.validation_random_generator is None @@ -101,8 +101,6 @@ def test_defaults(self) -> None: def test_string_inputs_are_coerced(self) -> None: """Verify YAML-compatible scalar strings become callback value types.""" - # YAML often produces strings for numeric fields; the - # constructor must coerce them. cb = ValidationCallback( pipeline_target=_PIPE_TARGET, dataset_file="x.json", @@ -110,6 +108,8 @@ def test_string_inputs_are_coerced(self) -> None: sampling_steps=["20", "40"], # type: ignore[arg-type] guidance_scale="4.5", # type: ignore[arg-type] num_frames="77", # type: ignore[arg-type] + use_record_dimensions="true", # type: ignore[arg-type] + max_record_num_frames="345", # type: ignore[arg-type] num_videos_per_prompt="2", # type: ignore[arg-type] use_validation_media_conditioning="false", # type: ignore[arg-type] sampling_timesteps=["1000", "500"], @@ -123,6 +123,8 @@ def test_string_inputs_are_coerced(self) -> None: assert cb.sampling_steps == [20, 40] assert cb.guidance_scale == 4.5 assert cb.num_frames == 77 + assert cb.use_record_dimensions is True + assert cb.max_record_num_frames == 345 assert cb.num_videos_per_prompt == 2 assert cb.use_validation_media_conditioning is False assert cb.run_at_start is False @@ -137,6 +139,12 @@ def test_init_rejects_nonpositive_video_count(self) -> None: with pytest.raises(ValueError, match="num_videos_per_prompt must be positive"): _make_callback(num_videos_per_prompt=0) + @pytest.mark.parametrize("max_record_num_frames", [0, -1]) + def test_init_rejects_nonpositive_record_frame_cap(self, max_record_num_frames: int) -> None: + """A configured native-record cap must describe a usable request.""" + with pytest.raises(ValueError, match="max_record_num_frames must be positive"): + _make_callback(max_record_num_frames=max_record_num_frames) + def test_pipeline_kwargs_collected(self) -> None: cb = ValidationCallback( pipeline_target=_PIPE_TARGET, @@ -144,7 +152,6 @@ def test_pipeline_kwargs_collected(self) -> None: extra_arg=123, another="value", ) - # Unknown kwargs are stashed for the pipeline factory. assert cb.pipeline_kwargs == { "extra_arg": 123, "another": "value", @@ -190,9 +197,6 @@ def test_metrics_mapping_is_coerced(self) -> None: assert cb.metrics_config.log_prefix == "custom/validation" -# --------------------------------------------------------------------------- -# B. on_validation_begin gating -# --------------------------------------------------------------------------- class _NoRunValidation(ValidationCallback): @@ -220,6 +224,7 @@ class TestOnValidationBegin: def test_skipped_when_every_steps_zero(self) -> None: cb = _make_recording(every_steps=0) + assert cb.will_run_validation(0) is False cb.on_validation_begin(method=None, iteration=0) cb.on_validation_begin(method=None, iteration=1000) assert cb.run_calls == [] @@ -232,6 +237,8 @@ def test_skipped_on_off_iter(self) -> None: def test_runs_on_match(self) -> None: cb = _make_recording(every_steps=50) + assert cb.will_run_validation(50) is True + assert cb.will_run_validation(51) is False cb.on_validation_begin(method=None, iteration=50) cb.on_validation_begin(method=None, iteration=100) assert cb.run_calls == [50, 100] @@ -257,6 +264,8 @@ def test_on_validation_begin_skips_step_zero_when_disabled(self) -> None: cb.on_validation_begin(method=None, iteration=0) cb.on_validation_begin(method=None, iteration=20) assert cb.run_calls == [20] + assert cb.will_run_validation(0) is False + assert cb.will_run_validation(20) is True class TestH3ValidationContract: @@ -292,6 +301,122 @@ def test_prepare_validation_batch_forwards_video_count( assert batch.num_videos_per_prompt == 3 + def test_prepare_validation_batch_honors_complete_record_dimensions( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + """V10 samples each held-out prompt at its native spatial/temporal shape.""" + cb = _make_callback( + num_frames=77, + use_record_dimensions=True, + ) + cb.training_config = SimpleNamespace( + data=SimpleNamespace( + num_height=64, + num_width=96, + num_latent_t=2, + ), + pipeline_config=SimpleNamespace(vae_config=SimpleNamespace( + arch_config=SimpleNamespace(temporal_compression_ratio=4), ), ), + model_path="unused", + vsa_sparsity=0.0, + ) + monkeypatch.setattr( + "fastvideo.train.callbacks.validation.make_inference_args", + lambda *args, **kwargs: SimpleNamespace(), + ) + + batch = cb._prepare_validation_batch( + SamplingParam(), + { + "prompt": "Generate synchronized media.", + "width": 128, + "height": 80, + "num_frames": 39, + }, + num_inference_steps=4, + ) + + assert (batch.width, batch.height, batch.num_frames) == (128, 80, 39) + assert batch.n_tokens == 10 * 10 * 16 + + def test_record_dimensions_require_complete_triplet(self) -> None: + cb = _make_callback(use_record_dimensions=True) + cb.training_config = SimpleNamespace( + data=SimpleNamespace( + num_height=64, + num_width=96, + num_latent_t=2, + ), + pipeline_config=SimpleNamespace(vae_config=SimpleNamespace( + arch_config=SimpleNamespace(temporal_compression_ratio=4), ), ), + ) + + with pytest.raises(ValueError, match="must provide width, height, and num_frames together"): + cb._validation_sampling_dimensions({ + "prompt": "partial", + "width": 128, + }) + + @pytest.mark.parametrize( + ("record_num_frames", "max_record_num_frames", "expected_num_frames"), + [ + (328, 345, 328), + (362, 345, 345), + (362, None, 362), + ], + ) + def test_record_frame_cap_preserves_shorter_and_legacy_geometry( + self, + record_num_frames: int, + max_record_num_frames: int | None, + expected_num_frames: int, + ) -> None: + """The opt-in cap affects only over-limit native record lengths.""" + cb = _make_callback( + use_record_dimensions=True, + max_record_num_frames=max_record_num_frames, + ) + cb.training_config = SimpleNamespace( + data=SimpleNamespace( + num_height=64, + num_width=96, + num_latent_t=2, + ), + pipeline_config=SimpleNamespace(vae_config=SimpleNamespace( + arch_config=SimpleNamespace(temporal_compression_ratio=4), ), ), + ) + + assert cb._validation_sampling_dimensions({ + "width": 1344, + "height": 768, + "num_frames": record_num_frames, + }) == (768, 1344, expected_num_frames) + + def test_record_frame_cap_matches_h3_inference_boundary(self) -> None: + """The v10 cap is the largest released H3 geometry below 15 seconds.""" + assert resolve_target_num_frames(345) == 345 + with pytest.raises(ValueError, match="aligned num_frames=362"): + resolve_target_num_frames(362) + + def test_record_dimensions_are_ignored_without_opt_in(self) -> None: + cb = _make_callback(num_frames=77, use_record_dimensions=False) + cb.training_config = SimpleNamespace( + data=SimpleNamespace( + num_height=64, + num_width=96, + num_latent_t=2, + ), + pipeline_config=SimpleNamespace(vae_config=SimpleNamespace( + arch_config=SimpleNamespace(temporal_compression_ratio=4), ), ), + ) + + assert cb._validation_sampling_dimensions({ + "width": 128, + "height": 80, + "num_frames": 39, + }) == (64, 96, 77) + def test_prepare_validation_batch_ignores_media_for_text_only_generation( self, monkeypatch: pytest.MonkeyPatch, @@ -374,8 +499,12 @@ def forward(self, batch, inference_args): monkeypatch.setattr( cb, "_prepare_validation_batch", - lambda sampling_param, validation_batch, num_inference_steps: SimpleNamespace(prompt=validation_batch[ - "caption"], ), + lambda sampling_param, validation_batch, num_inference_steps: SimpleNamespace( + prompt=validation_batch["caption"], + width=96, + height=64, + num_frames=5, + ), ) result = cb._run_validation_for_steps(50, transformer=torch.nn.Identity()) @@ -384,6 +513,13 @@ def forward(self, batch, inference_args): assert len(result.videos) == 8 assert result.audio_sample_rates == [32_000] * 8 assert len(result.audio_waveforms) == 8 + assert result.metadata == [{ + "source": "unknown", + "sample_id": None, + "width": 96, + "height": 64, + "num_frames": 5, + }] * 8 for prompt_index, waveform in enumerate(result.audio_waveforms): assert torch.is_tensor(waveform) torch.testing.assert_close(waveform, torch.full((32, 2), float(prompt_index))) @@ -431,6 +567,64 @@ def log_artifacts(self, artifacts, step): assert artifacts["validation_videos_50_steps"] == filenames assert {key: artifacts[key] for key in scalar_metrics} == scalar_metrics + def test_log_validation_artifacts_keeps_references_and_metadata_in_one_event(self) -> None: + """Generated/reference pairs and source/shape receipts remain aligned.""" + + class FakeWandbTracker: + + def __init__(self) -> None: + self.artifact_calls = [] + + def video(self, filename, *, caption, fps): + return (filename, caption, fps) + + def log_artifacts(self, artifacts, step): + self.artifact_calls.append((artifacts, step)) + + cb = _make_callback() + cb.tracker = FakeWandbTracker() + metadata = [{ + "source": "nuva/50k", + "sample_id": "sample-1", + "width": 1344, + "height": 768, + "num_frames": 345, + "reference_num_frames": 362, + }] + caption = cb._validation_artifact_caption("A prompt", metadata[0]) + ref_caption = cb._validation_artifact_caption( + "A prompt", + metadata[0], + prefix="held-out reference", + use_reference_num_frames=True, + ) + scalar_metrics = cb._validation_metadata_scalar_metrics( + metadata, + num_inference_steps=4, + ) + + cb._log_validation_video_artifacts( + ["generated.mp4"], + [caption], + key="validation_videos_4_steps", + step=100, + fps=24, + reference_video_filenames=["reference.mp4"], + reference_captions=[ref_caption], + reference_key="validation_references_4_steps", + scalar_metrics=scalar_metrics, + ) + + artifacts, step = cb.tracker.artifact_calls[0] + assert step == 100 + assert artifacts["validation_videos_4_steps"][0][0] == "generated.mp4" + assert artifacts["validation_references_4_steps"][0][0] == "reference.mp4" + assert "source=nuva/50k" in artifacts["validation_videos_4_steps"][0][1] + assert "shape=1344x768x345f" in artifacts["validation_videos_4_steps"][0][1] + assert "shape=1344x768x362f" in artifacts["validation_references_4_steps"][0][1] + assert artifacts["validation/4_steps/source/nuva_50k_count"] == 1.0 + assert artifacts["validation/4_steps/shape/1344x768x345f_count"] == 1.0 + class TestAttnQatInferValidation: @@ -483,16 +677,12 @@ def test_context_requires_qat_training_attention(self) -> None: pass -# --------------------------------------------------------------------------- -# C. _find_ema_callback -# --------------------------------------------------------------------------- class TestFindEmaCallback: def test_returns_none_without_callback_dict(self) -> None: cb = _make_callback() - # _callback_dict is not set on bare instances. assert cb._find_ema_callback() is None def test_returns_none_when_no_ema_registered(self) -> None: @@ -513,23 +703,17 @@ def test_finds_ema_callback(self) -> None: assert found is ema -# --------------------------------------------------------------------------- -# D. state_dict / load_state_dict (rng round-trip) -# --------------------------------------------------------------------------- class TestStateDict: def test_state_dict_empty_without_generator(self) -> None: cb = _make_callback() - # validation_random_generator is None until on_train_start. assert cb.state_dict() == {} def test_round_trip_preserves_rng_state(self) -> None: cb = _make_callback() gen = torch.Generator(device="cpu").manual_seed(123) - # Advance RNG so a default-init generator on the receiving - # side is observably different. for _ in range(5): torch.randn(4, generator=gen) cb.validation_random_generator = gen @@ -537,27 +721,20 @@ def test_round_trip_preserves_rng_state(self) -> None: state = cb.state_dict() assert "validation_rng" in state - # Receiver: fresh generator with a different seed. fresh = _make_callback() fresh.validation_random_generator = (torch.Generator(device="cpu").manual_seed(999)) fresh.load_state_dict(state) - # After load, both generators draw the same next sample. a = torch.randn(8, generator=cb.validation_random_generator) b = torch.randn(8, generator=fresh.validation_random_generator) assert torch.equal(a, b) def test_load_without_generator_is_noop(self) -> None: cb = _make_callback() - # Generator is None: load must not raise even when state has - # an rng entry. cb.load_state_dict({"validation_rng": torch.tensor([1, 2, 3], dtype=torch.uint8)}) assert cb.validation_random_generator is None -# --------------------------------------------------------------------------- -# E. metric result aggregation -# --------------------------------------------------------------------------- class TestMetricAggregation: @@ -728,9 +905,6 @@ def test_available_paths_requires_all_paths(self) -> None: assert ValidationCallback._available_paths(["a.mp4", None]) is None -# --------------------------------------------------------------------------- -# F. action overlay plumbing -# --------------------------------------------------------------------------- class TestActionOverlay: @@ -962,7 +1136,6 @@ def _encoder(hidden_size: int, text_len: int = 0) -> SimpleNamespace: )) def test_copies_loaded_width(self) -> None: - # ByT5 at the generic T5 default vs the checkpoint's real width. validation = SimpleNamespace(text_encoder_configs=(self._encoder(512), )) loaded = SimpleNamespace(text_encoder_configs=(self._encoder(1472), )) @@ -971,8 +1144,6 @@ def test_copies_loaded_width(self) -> None: assert validation.text_encoder_configs[0].arch_config.hidden_size == 1472 def test_preserves_training_owned_fields(self) -> None: - # ``text_len`` is a FastVideo-only field absent from HF configs, so - # the loader never populates it and the training value must survive. validation = SimpleNamespace(text_encoder_configs=(self._encoder(512, text_len=1000), )) loaded = SimpleNamespace(text_encoder_configs=(self._encoder(1472, text_len=0), )) @@ -981,8 +1152,6 @@ def test_preserves_training_owned_fields(self) -> None: assert validation.text_encoder_configs[0].arch_config.text_len == 1000 def test_keeps_config_objects_unshared(self) -> None: - # The second call site writes into ``tc.pipeline_config`` itself, so - # the merge must not alias the loaded encoder objects into it. validation = SimpleNamespace(text_encoder_configs=(self._encoder(512), )) original = validation.text_encoder_configs loaded = SimpleNamespace(text_encoder_configs=(self._encoder(1472), )) @@ -993,8 +1162,6 @@ def test_keeps_config_objects_unshared(self) -> None: assert validation.text_encoder_configs[0] is not loaded.text_encoder_configs[0] def test_skips_unpopulated_loaded_width(self) -> None: - # ``TextEncoderArchConfig.hidden_size`` defaults to 0; a loader that - # never filled it must not clobber a real training-side width. validation = SimpleNamespace(text_encoder_configs=(self._encoder(3584), )) loaded = SimpleNamespace(text_encoder_configs=(self._encoder(0), )) @@ -1030,5 +1197,4 @@ def test_missing_or_empty_encoders_is_noop( validation: SimpleNamespace, loaded: SimpleNamespace, ) -> None: - # Pipelines without text encoders must not raise here. ValidationCallback._keep_loaded_encoder_widths(validation, loaded) diff --git a/fastvideo/tests/train/callbacks/test_validation_sampling_contract.py b/fastvideo/tests/train/callbacks/test_validation_sampling_contract.py new file mode 100644 index 0000000000..826f547da6 --- /dev/null +++ b/fastvideo/tests/train/callbacks/test_validation_sampling_contract.py @@ -0,0 +1,95 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Validation must sample at the operating point training teaches. + +Two independent knobs have to survive the training-config -> validation hop: +the few-step denoising ladder and the VSA attention contract (sparsity AND +tile geometry). v8 shipped 2400 steps of validation that missed both — three +forwards on the scheduler's native grid at tile 256, while training ran four +forwards on ``[999, 749, 500, 250]`` at tile 64 — with nothing raising. +""" + +from __future__ import annotations + +import types + +import pytest + +from fastvideo.train.callbacks.validation import ValidationCallback + +LADDER = [999, 749, 500, 250] + + +def _callback(**kwargs): + """A callback instance without touching distributed state.""" + return ValidationCallback( + pipeline_target="fastvideo.pipelines.basic.minimax_h3." + "minimax_h3_pipeline.MiniMaxH3Pipeline", + dataset_file="unused.json", + **kwargs, + ) + + +def _method(ladder=LADDER): + cfg = {} if ladder is None else {"dmd_denoising_steps": list(ladder)} + return types.SimpleNamespace(method_config=cfg) + + +def test_ladder_is_inherited_when_unset() -> None: + cb = _callback(sampling_steps=[4]) + assert cb.sampling_timesteps is None + cb._adopt_training_sampling_contract(_method()) + assert cb.sampling_timesteps == LADDER + + +def test_matching_ladder_is_accepted() -> None: + cb = _callback(sampling_steps=[4], sampling_timesteps=LADDER) + cb._adopt_training_sampling_contract(_method()) + assert cb.sampling_timesteps == LADDER + + +def test_diverging_ladder_raises() -> None: + cb = _callback(sampling_steps=[4], sampling_timesteps=[1000, 667, 333]) + with pytest.raises(ValueError, match="disagrees with the trained ladder"): + cb._adopt_training_sampling_contract(_method()) + + +@pytest.mark.parametrize("ladder", [None, []]) +def test_non_dmd_methods_are_left_alone(ladder) -> None: + cb = _callback(sampling_steps=[40]) + cb._adopt_training_sampling_contract(_method(ladder)) + assert cb.sampling_timesteps is None + + +def test_method_without_config_is_tolerated() -> None: + cb = _callback(sampling_steps=[40]) + cb._adopt_training_sampling_contract(types.SimpleNamespace()) + assert cb.sampling_timesteps is None + + +def test_attention_contract_accepts_matching_args() -> None: + tc = types.SimpleNamespace(vsa_sparsity=0.9, vsa_tile_size=64) + args = types.SimpleNamespace(VSA_sparsity=0.9, VSA_tile_size=64) + ValidationCallback._assert_attention_contract(args, tc) + + +def test_attention_contract_catches_tile_size_drift() -> None: + """The exact v8 failure: sparsity propagated, tile size left at default.""" + tc = types.SimpleNamespace(vsa_sparsity=0.9, vsa_tile_size=64) + args = types.SimpleNamespace(VSA_sparsity=0.9, VSA_tile_size=256) + with pytest.raises(ValueError, match="VSA_tile_size=256"): + ValidationCallback._assert_attention_contract(args, tc) + + +def test_attention_contract_catches_sparsity_drift() -> None: + tc = types.SimpleNamespace(vsa_sparsity=0.9, vsa_tile_size=64) + args = types.SimpleNamespace(VSA_sparsity=0.0, VSA_tile_size=64) + with pytest.raises(ValueError, match="VSA_sparsity=0.0"): + ValidationCallback._assert_attention_contract(args, tc) + + +def test_make_inference_args_propagates_both_vsa_knobs() -> None: + """The missing line that caused the drift, pinned at its source.""" + from fastvideo.train.utils.moduleloader import make_inference_args + src = __import__("inspect").getsource(make_inference_args) + assert "args.VSA_sparsity = tc.vsa_sparsity" in src + assert "args.VSA_tile_size = tc.vsa_tile_size" in src diff --git a/fastvideo/tests/train/fixtures/minimax_h3_dmd2_min.yaml b/fastvideo/tests/train/fixtures/minimax_h3_dmd2_min.yaml new file mode 100644 index 0000000000..a29e4ae288 --- /dev/null +++ b/fastvideo/tests/train/fixtures/minimax_h3_dmd2_min.yaml @@ -0,0 +1,61 @@ +# Minimal H3 DMD2 trio and method configuration for CPU contract tests. +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: MiniMaxAI/MiniMax-H3 + trainable: true + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: MiniMaxAI/MiniMax-H3 + trainable: false + critic: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: MiniMaxAI/MiniMax-H3 + trainable: true + +method: + _target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method + rollout_mode: data_latent + generator_update_interval: 1 + real_score_guidance_scale: 3.5 + dmd_denoising_steps: [1000, 757, 522] + min_timestep_ratio: 0.02 + max_timestep_ratio: 0.98 + cfg_uncond: + text: zero + fake_score_learning_rate: 1.0e-3 + fake_score_betas: [0.0, 0.999] + fake_score_lr_scheduler: constant + +training: + distributed: + num_gpus: 1 + sp_size: 1 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 1 + data: + data_path: /tmp/minimax_h3_t2va + preprocessed_data_type: t2va + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 42 + # Tiny geometry: [1, 24, 2, 4, 4] video latents, [1, 2, 32, 8] audio. + num_latent_t: 2 + num_height: 64 + num_width: 64 + num_frames: 5 + optimizer: + learning_rate: 1.0e-3 + betas: [0.0, 0.999] + weight_decay: 0.0 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 4 + model: + precondition_outputs: false + dit_precision: bf16 + +callbacks: {} +pipeline: {} diff --git a/fastvideo/tests/train/methods/test_dmd2_data_forcing.py b/fastvideo/tests/train/methods/test_dmd2_data_forcing.py new file mode 100644 index 0000000000..019e03828a --- /dev/null +++ b/fastvideo/tests/train/methods/test_dmd2_data_forcing.py @@ -0,0 +1,416 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Per-batch data forcing for carried DMD2 (v9, FastGen data-driven regime). + +CPU contract tests for ``rollout_data_forcing``: routing by latent +presence under mixed loading, forced-input noising math (uniform grid-rung +draw, per-modality shifts), walk pausing, uniform first-call seeding, and +knob validation. ``rollout_data_forcing: false`` (the default) keeps the +carried walk byte-identical — ``test_dmd2_rollout_carry.py`` covers that +path unchanged. +""" +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch + +from fastvideo.tests.train.methods.test_dmd2_rollout_carry import ( + _GRID, + _LATENT_SHAPE, + _CarryStudent, + _make_method, + _stub_losses, +) +from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method +from fastvideo.train.models.minimax_h3.minimax_h3 import shift_noise_amount + +_VIDEO_NUMEL = 6 +_AUDIO_NUMEL = 2 +assert _VIDEO_NUMEL + _AUDIO_NUMEL == _LATENT_SHAPE[1] + + +class _ForcingStudent(_CarryStudent): + """Carry fake that also packs real latents for ``latents_source='data'``.""" + + def __init__(self) -> None: + super().__init__() + self.prepare_sources: list[str] = [] + self.add_noise_calls: list[dict] = [] + + def prepare_batch(self, raw_batch, *, generator, latents_source): + self.prepare_sources.append(latents_source) + self.prepare_calls.append(raw_batch) + if latents_source == "data": + latents = torch.cat( + ( + raw_batch["vae_latent"].reshape(1, -1), + raw_batch["audio_latent"].reshape(1, -1), + ), + dim=1, + ) + else: + latents = torch.zeros(_LATENT_SHAPE) + batch = SimpleNamespace( + latents=latents, + timesteps=torch.tensor([0.0]), + attn_metadata=None, + attn_metadata_vsa="vsa-metadata", + dmd_latent_vis_dict={}, + fake_score_latent_vis_dict={}, + ) + self.last_batch = batch + return batch + + def add_noise(self, clean, noise, timestep): + noisy = super().add_noise(clean, noise, timestep) + self.add_noise_calls.append({ + "clean": clean, + "noise": noise, + "timestep": float(timestep.reshape(-1)[0]), + "noisy": noisy, + }) + return noisy + + +def _latent_batch(seed: int = 1) -> dict[str, torch.Tensor]: + generator = torch.Generator().manual_seed(seed) + return { + "vae_latent": torch.randn(1, _VIDEO_NUMEL, generator=generator), + "audio_latent": torch.randn(1, _AUDIO_NUMEL, generator=generator), + "text_embedding": torch.randn(1, 4, generator=generator), + } + + +def _text_only_batch() -> dict[str, torch.Tensor]: + """Mixed loading under the t2va schema: latent columns come through empty.""" + return { + "vae_latent": torch.zeros(1, 0), + "audio_latent": torch.zeros(1, 0), + "text_embedding": torch.ones(1, 4), + } + + + + +def test_data_forcing_defaults_off_and_latent_batches_walk() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1, student=_ForcingStudent()) + _stub_losses(method) + student = method.student + + assert method._rollout_data_forcing is False + _, _, metrics = method.single_train_step(_latent_batch(), iteration=0) + + assert student.prepare_sources == ["zeros"] + assert "data_forced" not in metrics + assert metrics["rollout_step"] == 0.0 + + + + +def test_routing_by_latent_presence_under_mixed_loading() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1, student=_ForcingStudent(), data_forcing=True) + _stub_losses(method) + student = method.student + + _, _, walk_metrics = method.single_train_step(_text_only_batch(), iteration=0) + _, _, forced_metrics = method.single_train_step(_latent_batch(), iteration=1) + + assert student.prepare_sources == ["zeros", "data"] + assert walk_metrics["data_forced"] == 0.0 + assert walk_metrics["rollout_step"] == 0.0 + assert forced_metrics["data_forced"] == 1.0 + assert "rollout_step" not in forced_metrics + + +def test_half_present_latent_pair_fails_loud() -> None: + method = _make_method(slots=1, sample_type="ode", student=_ForcingStudent(), data_forcing=True) + batch = _latent_batch() + batch["audio_latent"] = torch.zeros(1, 0) + with pytest.raises(ValueError, match="exactly one of"): + method.single_train_step(batch, iteration=0) + + + + +def test_forced_input_is_real_latents_noised_at_a_grid_rung() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1, student=_ForcingStudent(), data_forcing=True) + _stub_losses(method) + student = method.student + method._carry_slot_seeded[0] = True + + batch = _latent_batch(seed=3) + packed_real = torch.cat( + (batch["vae_latent"].reshape(1, -1), batch["audio_latent"].reshape(1, -1)), + dim=1, + ) + method.single_train_step(batch, iteration=0) + + forced = student.add_noise_calls[0] + assert forced["timestep"] in [float(t) for t in _GRID] + torch.testing.assert_close(forced["clean"], packed_real) + sigma = forced["timestep"] / 1000.0 + torch.testing.assert_close(forced["noisy"], (1.0 - sigma) * packed_real + sigma * forced["noise"]) + main = student.predict_calls[-1] + assert main["timestep"] == forced["timestep"] + assert main["grad_enabled"] is True + assert main["attn_kind"] == "vsa" + vis = student.last_batch.dmd_latent_vis_dict + torch.testing.assert_close(vis["generator_timestep"], torch.tensor([forced["timestep"]])) + + +def test_forced_rung_draw_covers_the_whole_grid_and_never_zero() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1, student=_ForcingStudent(), data_forcing=True) + _stub_losses(method) + method._carry_slot_seeded[0] = True + + for call in range(64): + method.single_train_step(_latent_batch(seed=call), iteration=call) + drawn = {call["timestep"] for call in method.student.add_noise_calls} + assert drawn == {float(t) for t in _GRID} + assert 0.0 not in drawn + + + + +def test_forced_batches_pause_the_walk_and_text_batches_resume_it() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1, student=_ForcingStudent(), data_forcing=True) + _stub_losses(method) + + method.single_train_step(_text_only_batch(), iteration=0) + carried = method._carry_slots[0] + assert carried is not None and carried["rung"] == 1 + state_before = carried["state"].clone() + + method.single_train_step(_latent_batch(seed=5), iteration=1) + method.single_train_step(_latent_batch(seed=6), iteration=2) + assert method._carry_slots[0] is carried + torch.testing.assert_close(carried["state"], state_before) + + _, _, metrics = method.single_train_step(_text_only_batch(), iteration=3) + assert metrics["rollout_step"] == 1.0 + assert method._carry_slots[0] is not None + assert method._carry_slots[0]["rung"] == 2 + + +def test_forced_batch_at_walk_boundary_leaves_the_boundary_state() -> None: + method = _make_method(slots=1, grid=[999, 500], sample_type="ode", interval=1, student=_ForcingStudent(), + data_forcing=True) + _stub_losses(method) + + method.single_train_step(_text_only_batch(), iteration=0) + method.single_train_step(_text_only_batch(), iteration=1) + assert method._carry_slots[0] is None + + _, _, forced_metrics = method.single_train_step(_latent_batch(), iteration=2) + assert forced_metrics["data_forced"] == 1.0 + assert method._carry_slots[0] is None + _, _, metrics = method.single_train_step(_text_only_batch(), iteration=3) + assert metrics["rollout_step"] == 0.0 + + + + +def test_forced_first_call_runs_stagger_prewalk_with_uniform_forward_count() -> None: + forced_student = _ForcingStudent() + forced_method = _make_method(slots=1, sample_type="ode", interval=1, student=forced_student, data_forcing=True, + rank=1, world=2) + _stub_losses(forced_method) + text_student = _ForcingStudent() + text_method = _make_method(slots=1, sample_type="ode", interval=1, student=text_student, data_forcing=True, rank=1, + world=2) + _stub_losses(text_method) + + latent = _latent_batch(seed=7) + forced_method.single_train_step(latent, iteration=0) + text_method.single_train_step(_text_only_batch(), iteration=0) + + assert len(forced_student.predict_calls) == len(_GRID) + assert len(text_student.predict_calls) == len(_GRID) + + carried = forced_method._carry_slots[0] + assert carried is not None + assert carried["rung"] == (1 * 1 + 0) % len(_GRID) + assert forced_method._carry_slot_seeded[0] is True + torch.testing.assert_close(carried["raw_batch"]["vae_latent"], latent["vae_latent"]) + + _, _, metrics = forced_method.single_train_step(_text_only_batch(), iteration=1) + assert metrics["rollout_step"] == float(carried["rung"]) + adopted = forced_student.prepare_calls[-1] + torch.testing.assert_close(adopted["text_embedding"], latent["text_embedding"]) + + + + +def test_data_forcing_requires_rollout_carry() -> None: + with pytest.raises(ValueError, match="rollout_carry: true"): + _make_method(carry=False, sample_type=None, student=_ForcingStudent(), data_forcing=True) + + +def test_data_forcing_rejects_non_bool() -> None: + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", {"rollout_data_forcing": "yes"}) + object.__setattr__(method, "_rollout_carry", True) + with pytest.raises(ValueError, match="must be a bool"): + method._parse_rollout_data_forcing() + + +def test_data_forcing_requires_explicit_legacy_mixed_regime_opt_in() -> None: + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", {"rollout_data_forcing": True}) + object.__setattr__(method, "_rollout_carry", True) + with pytest.raises(ValueError, match="not a FastGen recipe"): + method._parse_rollout_data_forcing() + + method.method_config["allow_mixed_rollout_regimes"] = True + assert method._parse_rollout_data_forcing() is True + + +def test_data_forcing_requires_t2va_schema() -> None: + method = object.__new__(DMD2Method) + object.__setattr__(method, "_rollout_mode", "simulate") + object.__setattr__(method, "_rollout_data_forcing", True) + object.__setattr__( + method, + "training_config", + SimpleNamespace(data=SimpleNamespace(preprocessed_data_type="text_only")), + ) + with pytest.raises(ValueError, match="t2va"): + method._validate_preprocessed_data_type() + object.__setattr__( + method, + "training_config", + SimpleNamespace(data=SimpleNamespace(preprocessed_data_type="t2va")), + ) + method._validate_preprocessed_data_type() + + +def test_batch_classifier_contract() -> None: + assert DMD2Method._batch_has_latents(_latent_batch()) is True + assert DMD2Method._batch_has_latents(_text_only_batch()) is False + assert DMD2Method._batch_has_latents({"text_embedding": torch.ones(1, 4)}) is False + with pytest.raises(ValueError, match="exactly one of"): + DMD2Method._batch_has_latents({ + "vae_latent": torch.ones(1, 3), + "audio_latent": torch.zeros(1, 0), + }) + + + + +def _build_forcing_trio(monkeypatch: pytest.MonkeyPatch, *, interval: int) -> DMD2Method: + from fastvideo.tests.train.methods.test_minimax_h3_dmd2 import ( + _FIXTURE, + _make_model, + ) + from fastvideo.train.utils.config import load_run_config + + config = load_run_config(str(_FIXTURE)) + config.method["rollout_mode"] = "simulate" + config.method["generator_update_interval"] = interval + config.method["rollout_carry"] = True + config.method["rollout_carry_slots"] = 1 + config.method["rollout_sample_type"] = "ode" + config.method["rollout_data_forcing"] = True + config.method["allow_mixed_rollout_regimes"] = True + student = _make_model(monkeypatch, config.training, scale=1.0) + teacher = _make_model(monkeypatch, config.training, trainable=False, scale=0.5) + critic = _make_model(monkeypatch, config.training, scale=0.25) + student.init_preprocessors = lambda training_config: None + method = DMD2Method( + cfg=config, + role_models={ + "student": student, + "teacher": teacher, + "critic": critic, + }, + ) + method.cuda_generator = torch.Generator(device="cpu").manual_seed(0) + return method + + +def test_forced_noising_per_modality_shift_on_real_h3_trio(monkeypatch: pytest.MonkeyPatch) -> None: + """The forced input mixes each modality at its own shifted sigma.""" + from fastvideo.tests.train.methods.test_minimax_h3_dmd2 import _raw_batch + + method = _build_forcing_trio(monkeypatch, interval=5) + student = method.student + method._carry_slot_seeded[0] = True + + records: list[dict] = [] + original_add_noise = student.add_noise + + def spy_add_noise(clean, noise, timestep): + noisy = original_add_noise(clean, noise, timestep) + records.append({ + "clean": clean, + "noise": noise, + "timestep": timestep, + "noisy": noisy, + }) + return noisy + + monkeypatch.setattr(student, "add_noise", spy_add_noise) + + raw = _raw_batch(seed=1) + loss_map, outputs, metrics = method.single_train_step(raw, iteration=1) + + forced = records[0] + expected_packed = student.pack_latents( + raw["vae_latent"].permute(0, 2, 1, 3, 4).to(torch.bfloat16), + raw["audio_latent"].to(torch.bfloat16), + ) + torch.testing.assert_close(forced["clean"], expected_packed) + rung = int(forced["timestep"].reshape(-1)[0]) + assert rung in method.method_config["dmd_denoising_steps"] + + slices = dict(student.modality_slices()) + base = torch.tensor([rung / 1000.0], dtype=torch.float64) + for name, shift in (("video", 12.0), ("audio", 3.0)): + sigma = shift_noise_amount(base, shift) + expected = ((1.0 - sigma) * forced["clean"][:, slices[name]].to(torch.float64) + + sigma * forced["noise"][:, slices[name]].to(torch.float64)).to(torch.bfloat16) + torch.testing.assert_close(forced["noisy"][:, slices[name]], expected) + + assert metrics["data_forced"] == 1.0 + assert metrics["update_student"] == 0.0 + assert loss_map["fake_score_loss"].item() > 0.0 + assert method._carry_slots[0] is None + + method.backward(loss_map, outputs) + assert method.critic.transformer.scale.grad is not None + assert torch.isfinite(method.critic.transformer.scale.grad) + + +def test_forced_student_step_and_walk_adoption_on_real_h3_trio(monkeypatch: pytest.MonkeyPatch) -> None: + """A forced first call seeds the walk; the walk later reuses its prompt.""" + from fastvideo.tests.train.methods.test_minimax_h3_dmd2 import _raw_batch + + method = _build_forcing_trio(monkeypatch, interval=1) + student = method.student + + latent_raw = _raw_batch(seed=1) + loss_map, outputs, metrics = method.single_train_step(latent_raw, iteration=0) + + assert metrics["data_forced"] == 1.0 + assert metrics["update_student"] == 1.0 + assert torch.isfinite(loss_map["total_loss"]) + assert loss_map["generator_loss"].item() > 0.0 + assert "generator_loss_video" in metrics and "generator_loss_audio" in metrics + carried = method._carry_slots[0] + assert carried is not None and carried["rung"] == 0 + method.backward(loss_map, outputs) + assert student.transformer.scale.grad is not None + assert torch.isfinite(student.transformer.scale.grad) + + text_raw = { + "vae_latent": torch.zeros(1, 0), + "audio_latent": torch.zeros(1, 0), + "text_embedding": _raw_batch(seed=9)["text_embedding"], + "text_attention_mask": torch.tensor([[1, 1, 0, 0]], dtype=torch.float32), + } + _, _, walk_metrics = method.single_train_step(text_raw, iteration=1) + assert walk_metrics["data_forced"] == 0.0 + assert walk_metrics["rollout_step"] == 0.0 + adopted_text = latent_raw["text_embedding"][:, :2].to(torch.bfloat16) + torch.testing.assert_close(student.transformer.last_encoder_hidden_states, adopted_text) diff --git a/fastvideo/tests/train/methods/test_dmd2_fake_score_loss_space.py b/fastvideo/tests/train/methods/test_dmd2_fake_score_loss_space.py new file mode 100644 index 0000000000..a2731b3403 --- /dev/null +++ b/fastvideo/tests/train/methods/test_dmd2_fake_score_loss_space.py @@ -0,0 +1,136 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Critic regression-space coverage for DMD2.""" +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch + +from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method + +_SPLIT = 24 +_TOTAL = 32 +_SIGMA = {"video": 0.8, "audio": 0.5} +_TIMESTEP = torch.tensor([500]) + + +class _Student: + + def modality_slices(self): + return (("video", slice(0, _SPLIT)), ("audio", slice(_SPLIT, _TOTAL))) + + def add_noise(self, clean, noise, timestep): + out = clean.clone() + for name, sl in self.modality_slices(): + sigma = _SIGMA[name] + out[:, sl] = (1.0 - sigma) * clean[:, sl] + sigma * noise[:, sl] + return out + + +class _Critic: + + def __init__(self): + self.scale = torch.nn.Parameter(torch.tensor(0.1)) + self.predict_noise_calls = 0 + self.predict_x0_calls = 0 + + def predict_noise(self, noisy, timestep, batch, *, conditional, cfg_uncond, attn_kind): + self.predict_noise_calls += 1 + return self.scale * torch.ones_like(noisy) + + def predict_x0(self, noisy, timestep, batch, *, conditional, cfg_uncond, attn_kind): + self.predict_x0_calls += 1 + return self.scale * torch.ones_like(noisy) + + +def _method(space: str | None, critic: _Critic) -> DMD2Method: + method = object.__new__(DMD2Method) + config = {} if space is None else {"fake_score_loss_space": space} + object.__setattr__(method, "method_config", config) + object.__setattr__(method, "student", _Student()) + object.__setattr__(method, "critic", critic) + object.__setattr__(method, "cuda_generator", torch.Generator().manual_seed(0)) + object.__setattr__(method, "_cfg_uncond", None) + object.__setattr__(method, "_fake_score_loss_space", method._parse_fake_score_loss_space()) + object.__setattr__(method, "_sample_score_timestep", lambda device: _TIMESTEP) + return method + + +def _run(space: str | None, critic: _Critic) -> torch.Tensor: + method = _method(space, critic) + batch = SimpleNamespace(timesteps=None, attn_metadata=None, fake_score_latent_vis_dict=None) + gen = torch.randn(1, _TOTAL, generator=torch.Generator().manual_seed(7)) + loss, _, _, metrics = method._critic_flow_matching_loss(batch, generator_pred_x0=gen) + assert set(metrics) == {"fake_score_loss_video", "fake_score_loss_audio"} + return loss + + +def _expected(space: str | dict[str, str]) -> float: + gen = torch.randn(1, _TOTAL, generator=torch.Generator().manual_seed(7)) + if space == "x0": + pred_x0 = 0.1 * torch.ones_like(gen) + return sum(torch.mean((pred_x0[:, sl] - gen[:, sl])**2).item() + for _, sl in _Student().modality_slices()) + + noise = torch.randn(gen.shape, generator=torch.Generator().manual_seed(0)) + target = noise - gen + pred = 0.1 * torch.ones_like(gen) + expected = 0.0 + for name, sl in _Student().modality_slices(): + mse = torch.mean((pred[:, sl] - target[:, sl])**2).item() + space_m = space.get(name, "velocity") if isinstance(space, dict) else space + weight = _SIGMA[name]**2 if space_m == "x0" else 1.0 + expected += weight * mse + return expected + + +def test_default_is_velocity_space() -> None: + loss = _run(None, _Critic()) + assert loss.item() == pytest.approx(_expected("velocity"), rel=1e-5) + + +def test_x0_space_uses_direct_x0_regression() -> None: + critic = _Critic() + loss = _run("x0", critic) + assert loss.item() == pytest.approx(_expected("x0"), rel=1e-4) + assert critic.predict_x0_calls == 1 + assert critic.predict_noise_calls == 0 + + +def test_x0_space_keeps_critic_gradient() -> None: + critic = _Critic() + loss = _run("x0", critic) + loss.backward() + assert critic.scale.grad is not None + assert critic.scale.grad.abs().item() > 0.0 + + +def test_invalid_loss_space_rejected() -> None: + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", {"fake_score_loss_space": "eps"}) + with pytest.raises(ValueError, match="velocity, x0"): + method._parse_fake_score_loss_space() + + +def test_per_modality_space_mapping() -> None: + """Legacy mixed spaces retain their one-forward compatibility path.""" + spec = {"video": "x0", "audio": "velocity"} + loss = _run(spec, _Critic()) + assert loss.item() == pytest.approx(_expected(spec), rel=1e-4) + + +def test_mapping_unknown_modality_falls_back_to_default() -> None: + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", {"fake_score_loss_space": {"video": "x0"}}) + mapping = method._parse_fake_score_loss_space() + object.__setattr__(method, "_fake_score_loss_space", mapping) + assert method._fake_score_space_for("video") == "x0" + assert method._fake_score_space_for("audio") == "velocity" + + +def test_mapping_invalid_value_rejected() -> None: + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", {"fake_score_loss_space": {"audio": "eps"}}) + with pytest.raises(ValueError, match="velocity, x0"): + method._parse_fake_score_loss_space() diff --git a/fastvideo/tests/train/methods/test_dmd2_fastgen_parity.py b/fastvideo/tests/train/methods/test_dmd2_fastgen_parity.py new file mode 100644 index 0000000000..9056974912 --- /dev/null +++ b/fastvideo/tests/train/methods/test_dmd2_fastgen_parity.py @@ -0,0 +1,227 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU regression tests for DMD2 behavior matched to the FastGen H3 recipe. + +The references below intentionally spell out the small pieces of FastGen math +instead of importing a second checkout. FastGen's rectified-flow schedule uses +float64 arithmetic over ``max_t=0.999`` and casts latent results back only once. +""" +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch + +from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method +from fastvideo.train.models.minimax_h3.minimax_h3_dmd import MiniMaxH3DMDModel + +_FASTGEN_MAX_T = 0.999 +_TIMESTEP_SCALE = 1000.0 + + +def _fastgen_time_shift(value: torch.Tensor, shift: float) -> torch.Tensor: + """FastGen ``time_shift`` on the H3 schedule's finite time domain.""" + value = value.to(torch.float64) + return value * shift * _FASTGEN_MAX_T / (value * (shift - 1.0) + _FASTGEN_MAX_T) + + +def _h3_adapter() -> MiniMaxH3DMDModel: + model = MiniMaxH3DMDModel.__new__(MiniMaxH3DMDModel) + model.training_config = SimpleNamespace(data=SimpleNamespace( + num_latent_t=2, + num_frames=5, + num_height=64, + num_width=64, + )) + return model + + +def _score_method() -> DMD2Method: + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", { + "min_timestep_ratio": 0.001, + "max_timestep_ratio": 0.999, + "score_timestep_shift": 2.4, + "score_timestep_warp_max": _FASTGEN_MAX_T, + "score_timestep_continuous": True, + }) + object.__setattr__(method, "student", SimpleNamespace( + num_train_timesteps=1000, + shift_and_clamp_timestep=lambda timestep: timestep, + )) + lo, hi = method._parse_score_timestep_bounds() + object.__setattr__(method, "_score_min_timestep", lo) + object.__setattr__(method, "_score_max_timestep", hi) + object.__setattr__(method, "_score_timestep_shift", 2.4) + object.__setattr__(method, "_score_timestep_warp_max", _FASTGEN_MAX_T) + object.__setattr__(method, "_score_timestep_continuous", True) + object.__setattr__(method, "cuda_generator", torch.Generator().manual_seed(0)) + return method + + +@pytest.mark.parametrize("unit_draw", [0.0, 0.37, 1.0]) +def test_shifted_score_sampler_draws_uniform_coordinate_before_inverse_warp( + monkeypatch: pytest.MonkeyPatch, + unit_draw: float, +) -> None: + """The random draw is uniform on FastGen's pre-shift ``[.001, .999]``. + + H3 represents the score time on an unshifted base clock. Applying the + inverse 2.4 warp there makes its video shift compose to 5 and its audio + shift compose to 1.25, exactly matching FastGen's video-clock schedule. + """ + + def fixed_rand(size, *, device=None, dtype=None, generator=None): + del generator + assert tuple(size) == (1, ) + return torch.full((1, ), unit_draw, device=device, dtype=dtype) + + monkeypatch.setattr(torch, "rand", fixed_rand) + method = _score_method() + sampled = method._sample_score_timestep(torch.device("cpu")) + + pre_shift = torch.tensor( + [0.001 + unit_draw * (0.999 - 0.001)], + dtype=torch.float64, + ) + expected_base = _fastgen_time_shift(pre_shift, 1.0 / 2.4) * _TIMESTEP_SCALE + assert sampled.dtype == torch.float64 + torch.testing.assert_close(sampled, expected_base, rtol=0.0, atol=1e-10) + + sigma_video, sigma_audio = _h3_adapter()._noise_amounts(sampled) + torch.testing.assert_close( + sigma_video, + _fastgen_time_shift(pre_shift, 5.0), + rtol=0.0, + atol=1e-12, + ) + torch.testing.assert_close( + sigma_audio, + _fastgen_time_shift(pre_shift, 1.25), + rtol=0.0, + atol=1e-12, + ) + + +def _packed_bfloat16_pair(model: MiniMaxH3DMDModel) -> tuple[torch.Tensor, torch.Tensor]: + slices = dict(model.modality_slices()) + total = slices["audio"].stop + clean = torch.linspace(-2.75, 3.125, total, dtype=torch.float32).reshape(1, -1).to(torch.bfloat16) + noise = torch.linspace(1.875, -3.5, total, dtype=torch.float32).reshape(1, -1).to(torch.bfloat16) + return clean, noise + + +def _reference_sigmas(timestep: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + base = (timestep.reshape(-1)[:1].to(torch.float64) / _TIMESTEP_SCALE).clamp(0.0, _FASTGEN_MAX_T) + return _fastgen_time_shift(base, 12.0), _fastgen_time_shift(base, 3.0) + + +def _reference_mix(clean: torch.Tensor, noise: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor: + return ((1.0 - sigma) * clean.to(torch.float64) + sigma * noise.to(torch.float64)).to(clean.dtype) + + +def _reference_unmix(noisy: torch.Tensor, clean: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor: + return ((noisy.to(torch.float64) - (1.0 - sigma) * clean.to(torch.float64)) / sigma.clamp_min(1e-6)).to( + noisy.dtype) + + +def test_h3_add_noise_matches_fastgen_fp64_then_cast_for_bfloat16() -> None: + model = _h3_adapter() + clean, noise = _packed_bfloat16_pair(model) + timestep = torch.tensor([413.375], dtype=torch.float64) + sigma_video, sigma_audio = _reference_sigmas(timestep) + clean_video, clean_audio = model.unpack_latents(clean) + noise_video, noise_audio = model.unpack_latents(noise) + expected = model.pack_latents( + _reference_mix(clean_video, noise_video, sigma_video), + _reference_mix(clean_audio, noise_audio, sigma_audio), + ) + + actual = model.add_noise(clean, noise, timestep) + + assert actual.dtype == torch.bfloat16 + torch.testing.assert_close(actual, expected, rtol=0.0, atol=0.0) + + +def test_h3_extract_eps_matches_fastgen_fp64_then_cast_at_low_audio_sigma() -> None: + model = _h3_adapter() + clean, noisy = _packed_bfloat16_pair(model) + timestep = _fastgen_time_shift(torch.tensor([0.001], dtype=torch.float64), 1.0 / 2.4) * _TIMESTEP_SCALE + sigma_video, sigma_audio = _reference_sigmas(timestep) + noisy_video, noisy_audio = model.unpack_latents(noisy) + clean_video, clean_audio = model.unpack_latents(clean) + expected = model.pack_latents( + _reference_unmix(noisy_video, clean_video, sigma_video), + _reference_unmix(noisy_audio, clean_audio, sigma_audio), + ) + + actual = model.extract_eps(noisy, clean, timestep) + + assert actual.dtype == torch.bfloat16 + torch.testing.assert_close(actual, expected, rtol=0.0, atol=0.0) + + +def test_h3_predict_x0_helper_matches_fastgen_fp64_then_cast() -> None: + model = _h3_adapter() + noisy, pred_noise = _packed_bfloat16_pair(model) + sigma = torch.tensor([0.3141592653589793], dtype=torch.float64) + expected = (noisy.to(torch.float64) - sigma * pred_noise.to(torch.float64)).to(torch.bfloat16) + + actual = model._to_x0(noisy, pred_noise, sigma) + + assert actual.dtype == torch.bfloat16 + torch.testing.assert_close(actual, expected, rtol=0.0, atol=0.0) + + +class _ScoreStudent: + + def modality_slices(self): + return (("video", slice(0, 3)), ("audio", slice(3, 5))) + + def add_noise(self, clean, noise, timestep): + del timestep + return 0.75 * clean + 0.25 * noise + + +class _DirectX0Critic: + + def __init__(self) -> None: + self.value = torch.nn.Parameter(torch.tensor(0.25)) + self.predict_x0_calls = 0 + + def predict_x0(self, noisy, timestep, batch, *, conditional, cfg_uncond, attn_kind): + del timestep, batch, conditional, cfg_uncond, attn_kind + self.predict_x0_calls += 1 + return self.value * torch.ones_like(noisy) + + def predict_noise(self, *args, **kwargs): + del args, kwargs + raise AssertionError("global x0 regression must call critic.predict_x0 directly") + + +def test_critic_x0_objective_is_direct_predict_x0_mse_per_modality() -> None: + method = object.__new__(DMD2Method) + critic = _DirectX0Critic() + object.__setattr__(method, "method_config", {"fake_score_loss_space": "x0"}) + object.__setattr__(method, "student", _ScoreStudent()) + object.__setattr__(method, "critic", critic) + object.__setattr__(method, "cuda_generator", torch.Generator().manual_seed(11)) + object.__setattr__(method, "_cfg_uncond", None) + object.__setattr__(method, "_fake_score_loss_space", {"__default__": "x0"}) + object.__setattr__(method, "_sample_score_timestep", lambda device: torch.tensor([500.0], device=device)) + batch = SimpleNamespace(timesteps=None, attn_metadata=None, fake_score_latent_vis_dict=None) + generated_x0 = torch.tensor([[1.0, -2.0, 4.0, 8.0, -16.0]]) + + loss, _, _, metrics = method._critic_flow_matching_loss(batch, generator_pred_x0=generated_x0) + + pred_x0 = torch.full_like(generated_x0, 0.25) + expected_video = torch.mean((pred_x0[:, :3] - generated_x0[:, :3])**2) + expected_audio = torch.mean((pred_x0[:, 3:] - generated_x0[:, 3:])**2) + assert critic.predict_x0_calls == 1 + torch.testing.assert_close(metrics["fake_score_loss_video"], expected_video) + torch.testing.assert_close(metrics["fake_score_loss_audio"], expected_audio) + torch.testing.assert_close(loss, expected_video + expected_audio) + + loss.backward() + assert critic.value.grad is not None + assert torch.isfinite(critic.value.grad) diff --git a/fastvideo/tests/train/methods/test_dmd2_rollout_carry.py b/fastvideo/tests/train/methods/test_dmd2_rollout_carry.py new file mode 100644 index 0000000000..406492a9a3 --- /dev/null +++ b/fastvideo/tests/train/methods/test_dmd2_rollout_carry.py @@ -0,0 +1,610 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Carried backward-simulation rollout for DMD2 (FastGen port). + +CPU contract tests for ``rollout_carry``: slot round-robin, rung +progression and clearing, first-ever staggered starts with a uniform +pre-walk, ODE renoise with per-modality shifts, carried conditioning, +and knob validation. ``rollout_carry: false`` (the default) keeps the +existing rollout path byte-identical; the pre-existing DMD2 suites +(``test_minimax_h3_dmd2.py`` et al.) cover that path unchanged. +""" +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch + +from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method +from fastvideo.train.models.minimax_h3 import MiniMaxH3DMDModel +from fastvideo.train.models.minimax_h3.minimax_h3 import shift_noise_amount + +_GRID = [999, 749, 500, 250] +_LATENT_SHAPE = (1, 8) + + +class _CarryStudent: + """Rectified-flow fake on one packed tensor (sigma = t / 1000).""" + + device = torch.device("cpu") + + def __init__(self) -> None: + self.prepare_calls: list[dict] = [] + self.predict_calls: list[dict] = [] + self.last_batch: SimpleNamespace | None = None + self.last_pred: torch.Tensor | None = None + + def prepare_batch(self, raw_batch, *, generator, latents_source): + assert latents_source == "zeros" + self.prepare_calls.append(raw_batch) + batch = SimpleNamespace( + latents=torch.zeros(_LATENT_SHAPE), + timesteps=torch.tensor([0.0]), + attn_metadata=None, + attn_metadata_vsa="vsa-metadata", + dmd_latent_vis_dict={}, + fake_score_latent_vis_dict={}, + ) + self.last_batch = batch + return batch + + @staticmethod + def _sigma(timestep: torch.Tensor) -> float: + return float(timestep.reshape(-1)[0]) / 1000.0 + + def predict_x0(self, noisy, timestep, batch, *, conditional, cfg_uncond, attn_kind): + self.predict_calls.append({ + "timestep": float(timestep.reshape(-1)[0]), + "grad_enabled": torch.is_grad_enabled(), + "attn_kind": attn_kind, + }) + pred = noisy * 0.5 + self.last_pred = pred + return pred + + def add_noise(self, clean, noise, timestep): + sigma = self._sigma(timestep) + return (1.0 - sigma) * clean + sigma * noise + + def extract_eps(self, noisy, clean, timestep): + sigma = self._sigma(timestep) + return (noisy - (1.0 - sigma) * clean) / sigma + + +def _make_method( + *, + slots: int = 1, + grid: list[int] | None = None, + interval: int = 1, + sample_type: str | None = "ode", + carry: bool = True, + rank: int = 0, + world: int = 1, + grad_accum: int | None = None, + rollout_mode: str = "simulate", + student: object | None = None, + data_forcing: bool | None = None, + native_shape_bucketing: bool = False, +) -> DMD2Method: + method = object.__new__(DMD2Method) + config: dict = { + "rollout_mode": rollout_mode, + "rollout_carry": carry, + "dmd_denoising_steps": list(grid or _GRID), + "generator_update_interval": interval, + } + if carry: + config["rollout_carry_slots"] = slots + if sample_type is not None: + config["rollout_sample_type"] = sample_type + if data_forcing is not None: + config["rollout_data_forcing"] = data_forcing + if data_forcing: + config["allow_mixed_rollout_regimes"] = True + object.__setattr__(method, "method_config", config) + object.__setattr__(method, "student", student if student is not None else _CarryStudent()) + object.__setattr__( + method, + "training_config", + SimpleNamespace( + loop=SimpleNamespace(gradient_accumulation_steps=(slots if grad_accum is None else grad_accum)), + distributed=SimpleNamespace(sp_size=1), + data=SimpleNamespace(native_shape_bucketing=native_shape_bucketing), + ), + ) + object.__setattr__(method, "cuda_generator", torch.Generator().manual_seed(0)) + object.__setattr__(method, "_cfg_uncond", None) + object.__setattr__(method, "_denoising_step_list", None) + object.__setattr__(method, "_rollout_mode", method._parse_rollout_mode()) + object.__setattr__(method, "_rollout_carry_rank_world", lambda: (rank, world)) + knobs = method._parse_rollout_carry() + object.__setattr__(method, "_rollout_carry", knobs[0]) + object.__setattr__(method, "_rollout_carry_slot_count", knobs[1]) + object.__setattr__(method, "_rollout_sample_type", knobs[2]) + method._init_rollout_carry_state() + object.__setattr__( + method, + "_rollout_data_forcing", + method._parse_rollout_data_forcing(), + ) + return method + + +def _stub_losses(method: DMD2Method) -> list[torch.Tensor]: + """Replace the loss paths with recorders; carry mechanics stay real.""" + critic_preds: list[torch.Tensor] = [] + object.__setattr__(method, "_dmd_loss", lambda pred, batch: (torch.zeros(()), {})) + + def _critic(batch, *, generator_pred_x0=None): + critic_preds.append(generator_pred_x0) + return torch.zeros(()), "critic-ctx", {}, {} + + object.__setattr__(method, "_critic_flow_matching_loss", _critic) + return critic_preds + + +def _replay_ode_walk(state: torch.Tensor, grid: list[int], hops: int) -> torch.Tensor: + """Reference walk under the fake student's flow: pred = 0.5 * state.""" + for rung in range(hops): + sigma = grid[rung] / 1000.0 + pred = state * 0.5 + eps = (state - (1.0 - sigma) * pred) / sigma + sigma_next = grid[rung + 1] / 1000.0 + state = (1.0 - sigma_next) * pred + sigma_next * eps + return state + + + + +def test_slot_round_robin_and_rung_progression_with_two_slots() -> None: + method = _make_method(slots=2, sample_type="sde", interval=1) + _stub_losses(method) + student = method.student + + rungs = [] + forwards_per_call = [] + for call in range(10): + before = len(student.predict_calls) + _, _, metrics = method.single_train_step({"text_embedding": torch.ones(1, 4)}, iteration=call) + rungs.append(metrics["rollout_step"]) + forwards_per_call.append(len(student.predict_calls) - before) + + assert rungs == [0.0, 1.0, 1.0, 2.0, 2.0, 3.0, 3.0, 0.0, 0.0, 1.0] + assert forwards_per_call == [4, 4, 1, 1, 1, 1, 1, 1, 1, 1] + assert method._carry_slots[0] is not None and method._carry_slots[0]["rung"] == 1 + assert method._carry_slots[1] is not None and method._carry_slots[1]["rung"] == 2 + + +def test_carry_slots_cleared_after_last_rung() -> None: + method = _make_method(slots=1, grid=[999, 500], sample_type="ode", interval=1) + _stub_losses(method) + student = method.student + + counts, rungs = [], [] + for call in range(3): + before = len(student.predict_calls) + _, _, metrics = method.single_train_step({"x": torch.ones(1)}, iteration=call) + rungs.append(metrics["rollout_step"]) + counts.append(len(student.predict_calls) - before) + + assert rungs == [0.0, 1.0, 0.0] + assert counts == [2, 1, 1] + assert method._carry_slots[0] is not None + assert method._carry_slots[0]["rung"] == 1 + + + + +@pytest.mark.parametrize( + ("rank", "slot", "expected_offset"), + [(0, 0, 0), (0, 1, 1), (1, 0, 2), (1, 1, 3), (2, 0, 0), (2, 1, 1)], +) +def test_stagger_offsets_follow_rank_slot_formula(rank: int, slot: int, expected_offset: int) -> None: + method = _make_method(slots=2, sample_type="ode", rank=rank, world=3) + student = method.student + step_list = method._get_denoising_step_list(torch.device("cpu")) + state = torch.randn(_LATENT_SHAPE, generator=torch.Generator().manual_seed(3)) + batch = SimpleNamespace(dmd_latent_vis_dict={}) + + snapshot, rung = method._staggered_start(state, batch, step_list, slot) + + assert rung == expected_offset == (rank * 2 + slot) % len(_GRID) + assert len(student.predict_calls) == len(_GRID) - 1 + assert all(not call["grad_enabled"] for call in student.predict_calls) + assert all(call["attn_kind"] == "vsa" for call in student.predict_calls) + torch.testing.assert_close(snapshot, _replay_ode_walk(state, _GRID, expected_offset)) + + +@pytest.mark.parametrize("rank", [0, 1, 7]) +@pytest.mark.parametrize("slot", [0, 1]) +def test_native_shape_stagger_is_rank_synchronous(rank: int, slot: int) -> None: + method = _make_method( + slots=2, + sample_type="ode", + rank=rank, + world=8, + native_shape_bucketing=True, + ) + step_list = method._get_denoising_step_list(torch.device("cpu")) + state = torch.randn(_LATENT_SHAPE, generator=torch.Generator().manual_seed(3)) + batch = SimpleNamespace(dmd_latent_vis_dict={}) + + _, rung = method._staggered_start(state, batch, step_list, slot) + + assert rung == slot + + + + +def _h3_adapter() -> MiniMaxH3DMDModel: + model = MiniMaxH3DMDModel.__new__(MiniMaxH3DMDModel) + model.training_config = SimpleNamespace(data=SimpleNamespace( + num_latent_t=2, + num_frames=5, + num_height=64, + num_width=64, + )) + return model + + +def test_ode_renoise_analytic_identity_per_modality() -> None: + """From x_t = (1-s_m)x0 + s_m*eps and pred = x0, the advanced state is + (1-s'_m)x0 + s'_m*eps for both the video-shift and audio-shift slices.""" + model = _h3_adapter() + slices = dict(model.modality_slices()) + total = slices["audio"].stop + generator = torch.Generator().manual_seed(11) + x0 = torch.randn(1, total, generator=generator) + eps = torch.randn(1, total, generator=generator) + timestep = torch.tensor([500], dtype=torch.long) + next_timestep = torch.tensor([250], dtype=torch.long) + + x_t = model.add_noise(x0, eps, timestep) + implied = model.extract_eps(x_t, x0, timestep) + torch.testing.assert_close(implied, eps, rtol=1e-5, atol=1e-5) + + advanced = model.add_noise(x0, implied, next_timestep) + for name, shift in (("video", 12.0), ("audio", 3.0)): + sigma_next = float(shift_noise_amount(torch.tensor([0.25]), shift)) + expected = (1.0 - sigma_next) * x0[:, slices[name]] + sigma_next * eps[:, slices[name]] + torch.testing.assert_close(advanced[:, slices[name]], expected, rtol=1e-5, atol=1e-5) + + +def test_method_renoise_ode_uses_adapter_extract_eps() -> None: + model = _h3_adapter() + method = _make_method(slots=1, grid=[500, 250], sample_type="ode", student=model) + total = dict(model.modality_slices())["audio"].stop + generator = torch.Generator().manual_seed(13) + x0 = torch.randn(1, total, generator=generator) + eps = torch.randn(1, total, generator=generator) + timestep = torch.tensor([500], dtype=torch.long) + x_t = model.add_noise(x0, eps, timestep) + + step_list = method._get_denoising_step_list(torch.device("cpu")) + advanced = method._renoise(x_t, x0, timestep, 1, step_list) + expected = model.add_noise(x0, eps, torch.tensor([250], dtype=torch.long)) + torch.testing.assert_close(advanced, expected, rtol=1e-5, atol=1e-5) + + +def test_method_renoise_sde_draws_fresh_noise() -> None: + method = _make_method(slots=1, grid=[999, 500], sample_type="sde") + step_list = method._get_denoising_step_list(torch.device("cpu")) + state = torch.randn(_LATENT_SHAPE, generator=torch.Generator().manual_seed(5)) + pred = state * 0.5 + + advanced = method._renoise(state, pred, torch.tensor([999]), 1, step_list) + noise = torch.randn(_LATENT_SHAPE, generator=torch.Generator().manual_seed(0)) + torch.testing.assert_close(advanced, 0.5 * pred + 0.5 * noise) + + + + +def test_mid_walk_reuses_carried_conditioning_and_ignores_fresh_batches() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1) + _stub_losses(method) + student = method.student + + batch_a = {"text_embedding": torch.full((1, 4), 1.0), "info_list": ["prompt-a"]} + batch_b = {"text_embedding": torch.full((1, 4), 2.0), "info_list": ["prompt-b"]} + batch_c = {"text_embedding": torch.full((1, 4), 3.0), "info_list": ["prompt-c"]} + + method.single_train_step(batch_a, iteration=0) + snapshot = student.prepare_calls[0] + assert snapshot is not batch_a + torch.testing.assert_close(snapshot["text_embedding"], batch_a["text_embedding"]) + assert snapshot["info_list"] == ["prompt-a"] + + for call in range(1, 4): + method.single_train_step(batch_b, iteration=call) + assert student.prepare_calls[call] is snapshot + torch.testing.assert_close(student.prepare_calls[call]["text_embedding"], batch_a["text_embedding"]) + + method.single_train_step(batch_c, iteration=4) + torch.testing.assert_close(student.prepare_calls[4]["text_embedding"], batch_c["text_embedding"]) + assert student.prepare_calls[4]["info_list"] == ["prompt-c"] + + + + +def test_student_phase_forward_has_grad_and_sets_ctx_and_vis() -> None: + method = _make_method(slots=1, sample_type="ode", interval=5) + _stub_losses(method) + student = method.student + + _, outputs, metrics = method.single_train_step({"x": torch.ones(1)}, iteration=5) + + assert metrics["update_student"] == 1.0 + assert metrics["rollout_step"] == 0.0 + main = student.predict_calls[-1] + assert main["grad_enabled"] is True + assert main["attn_kind"] == "vsa" + assert main["timestep"] == float(_GRID[0]) + student_ctx = outputs["_fv_backward"]["student_ctx"] + assert student_ctx[0] is student.last_batch.timesteps + assert student_ctx[1] == "vsa-metadata" + assert outputs["_fv_backward"]["critic_ctx"] is None + vis = student.last_batch.dmd_latent_vis_dict + torch.testing.assert_close(vis["generator_timestep"], torch.tensor([float(_GRID[0])])) + torch.testing.assert_close(vis["generator_pred_video"], student.last_pred) + assert "generator_timestep" in method.latent_vis + + +def test_critic_phase_forward_is_no_grad_and_feeds_carried_pred() -> None: + method = _make_method(slots=1, sample_type="ode", interval=5) + critic_preds = _stub_losses(method) + student = method.student + + _, outputs, metrics = method.single_train_step({"x": torch.ones(1)}, iteration=1) + + assert metrics["update_student"] == 0.0 + assert metrics["rollout_step"] == 0.0 + main = student.predict_calls[-1] + assert main["grad_enabled"] is False + assert main["attn_kind"] == "vsa" + assert len(critic_preds) == 1 + assert critic_preds[0] is student.last_pred + assert outputs["_fv_backward"]["critic_ctx"] == "critic-ctx" + assert outputs["_fv_backward"]["student_ctx"] is None + assert method._carry_slots[0] is not None + assert method._carry_slots[0]["rung"] == 1 + assert "generator_timestep" in student.last_batch.dmd_latent_vis_dict + + +def test_carried_advance_state_is_detached_and_matches_ode_math() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1) + _stub_losses(method) + + method.single_train_step({"x": torch.ones(1)}, iteration=0) + carried = method._carry_slots[0] + assert carried is not None + assert carried["rung"] == 1 + assert not carried["state"].requires_grad + state0 = torch.randn(_LATENT_SHAPE, generator=torch.Generator().manual_seed(0)) + torch.testing.assert_close(carried["state"], _replay_ode_walk(state0, _GRID, 1)) + + + + +def _parse_only( + config: dict, + *, + grad_accum: int = 1, + streams: tuple[int, int] = (0, 1), + native_shape_bucketing: bool = False, +): + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", dict(config)) + object.__setattr__(method, "student", _CarryStudent()) + object.__setattr__( + method, + "training_config", + SimpleNamespace( + loop=SimpleNamespace(gradient_accumulation_steps=grad_accum), + distributed=SimpleNamespace(sp_size=1), + data=SimpleNamespace(native_shape_bucketing=native_shape_bucketing), + ), + ) + object.__setattr__(method, "_rollout_mode", method._parse_rollout_mode()) + object.__setattr__(method, "_rollout_carry_rank_world", lambda: streams) + return method._parse_rollout_carry() + + +def test_rollout_carry_defaults_off() -> None: + knobs = _parse_only({ + "rollout_mode": "simulate", + "dmd_denoising_steps": _GRID, + }) + assert knobs == (False, 0, "sde") + + +def test_rollout_carry_defaults_slots_to_grad_accum() -> None: + knobs = _parse_only( + { + "rollout_mode": "simulate", + "rollout_carry": True, + "rollout_sample_type": "ode", + "dmd_denoising_steps": _GRID, + "generator_update_interval": 5, + }, + grad_accum=3, + ) + assert knobs == (True, 3, "ode") + + +def test_rollout_carry_requires_simulate_mode() -> None: + with pytest.raises(ValueError, match="rollout_mode: simulate"): + _parse_only({ + "rollout_mode": "data_latent", + "rollout_carry": True, + "dmd_denoising_steps": _GRID, + }) + + +def test_rollout_carry_slots_must_match_grad_accum() -> None: + with pytest.raises(ValueError, match="gradient_accumulation_steps"): + _parse_only( + { + "rollout_mode": "simulate", + "rollout_carry": True, + "rollout_carry_slots": 4, + "dmd_denoising_steps": _GRID, + }, + grad_accum=2, + ) + + +def test_carry_knobs_require_rollout_carry_enabled() -> None: + with pytest.raises(ValueError, match="rollout_carry: true"): + _parse_only({ + "rollout_mode": "simulate", + "rollout_carry_slots": 2, + "dmd_denoising_steps": _GRID, + }) + with pytest.raises(ValueError, match="rollout_carry: true"): + _parse_only({ + "rollout_mode": "simulate", + "rollout_sample_type": "ode", + "dmd_denoising_steps": _GRID, + }) + + +def test_rollout_sample_type_rejects_unknown_values() -> None: + with pytest.raises(ValueError, match="ode, sde"): + _parse_only({ + "rollout_mode": "simulate", + "rollout_carry": True, + "rollout_sample_type": "euler", + "dmd_denoising_steps": _GRID, + }) + + +def test_ode_requires_student_extract_eps() -> None: + method = object.__new__(DMD2Method) + object.__setattr__( + method, + "method_config", + { + "rollout_mode": "simulate", + "rollout_carry": True, + "rollout_sample_type": "ode", + "dmd_denoising_steps": _GRID, + }, + ) + object.__setattr__(method, "student", SimpleNamespace()) + object.__setattr__( + method, + "training_config", + SimpleNamespace( + loop=SimpleNamespace(gradient_accumulation_steps=1), + distributed=SimpleNamespace(sp_size=1), + ), + ) + object.__setattr__(method, "_rollout_mode", "simulate") + with pytest.raises(ValueError, match="extract_eps"): + method._parse_rollout_carry() + + +def test_coverage_guard_rejects_uncovered_rung_phases() -> None: + with pytest.raises(ValueError, match="cannot cover"): + DMD2Method._validate_rollout_carry_coverage(streams=1, grid_len=4, interval=2) + DMD2Method._validate_rollout_carry_coverage(streams=2, grid_len=4, interval=2) + DMD2Method._validate_rollout_carry_coverage(streams=1, grid_len=4, interval=5) + + +def test_native_shape_coverage_does_not_count_rank_staggering() -> None: + config = { + "rollout_mode": "simulate", + "rollout_carry": True, + "rollout_carry_slots": 1, + "rollout_sample_type": "ode", + "dmd_denoising_steps": _GRID, + "generator_update_interval": 2, + } + assert _parse_only(config, streams=(0, 2)) == (True, 1, "ode") + with pytest.raises(ValueError, match="cannot cover"): + _parse_only(config, streams=(0, 2), native_shape_bucketing=True) + + + + +def _build_carry_trio(monkeypatch: pytest.MonkeyPatch, *, interval: int) -> DMD2Method: + from fastvideo.tests.train.methods.test_minimax_h3_dmd2 import ( + _FIXTURE, + _make_model, + ) + from fastvideo.train.utils.config import load_run_config + + config = load_run_config(str(_FIXTURE)) + config.method["rollout_mode"] = "simulate" + config.method["generator_update_interval"] = interval + config.method["rollout_carry"] = True + config.method["rollout_carry_slots"] = 1 + config.method["rollout_sample_type"] = "ode" + student = _make_model(monkeypatch, config.training, scale=1.0) + teacher = _make_model(monkeypatch, config.training, trainable=False, scale=0.5) + critic = _make_model(monkeypatch, config.training, scale=0.25) + student.init_preprocessors = lambda training_config: None + method = DMD2Method( + cfg=config, + role_models={ + "student": student, + "teacher": teacher, + "critic": critic, + }, + ) + method.cuda_generator = torch.Generator(device="cpu").manual_seed(0) + return method + + +def test_full_carried_student_step_on_real_h3_trio(monkeypatch: pytest.MonkeyPatch) -> None: + """End-to-end carried step: real prepare_batch, losses, and ODE advance.""" + from fastvideo.tests.train.methods.test_minimax_h3_dmd2 import _raw_batch + + method = _build_carry_trio(monkeypatch, interval=1) + student = method.student + + loss_map, outputs, metrics = method.single_train_step(_raw_batch(seed=1), iteration=0) + + assert metrics["update_student"] == 1.0 + assert metrics["rollout_step"] == 0.0 # (rank 0 * 1 + slot 0) % 3 + assert torch.isfinite(loss_map["total_loss"]) + assert loss_map["generator_loss"].item() > 0.0 + assert "generator_loss_video" in metrics and "generator_loss_audio" in metrics + carried = method._carry_slots[0] + assert carried is not None and carried["rung"] == 1 + assert carried["state"].shape == student.prepare_batch( + carried["raw_batch"], + generator=torch.Generator().manual_seed(0), + latents_source="zeros", + ).latents.shape + assert not carried["state"].requires_grad + + method.backward(loss_map, outputs) + assert student.transformer.scale.grad is not None + assert torch.isfinite(student.transformer.scale.grad) + + adopted = _raw_batch(seed=1) + adopted_text = adopted["text_embedding"][:, :2].to(torch.bfloat16) + method.single_train_step(_raw_batch(seed=9), iteration=1) + torch.testing.assert_close(student.transformer.last_encoder_hidden_states, adopted_text) + + +def test_full_carried_critic_step_on_real_h3_trio(monkeypatch: pytest.MonkeyPatch) -> None: + from fastvideo.tests.train.methods.test_minimax_h3_dmd2 import _raw_batch + + method = _build_carry_trio(monkeypatch, interval=5) + critic = method.critic + + loss_map, outputs, metrics = method.single_train_step(_raw_batch(seed=1), iteration=1) + + assert metrics["update_student"] == 0.0 + assert metrics["rollout_step"] == 0.0 + assert loss_map["generator_loss"].item() == 0.0 + assert loss_map["fake_score_loss"].item() > 0.0 + assert "fake_score_loss_video" in metrics and "fake_score_loss_audio" in metrics + assert method._carry_slots[0] is not None and method._carry_slots[0]["rung"] == 1 + + method.backward(loss_map, outputs) + assert critic.transformer.scale.grad is not None + assert torch.isfinite(critic.transformer.scale.grad) + assert method.student.transformer.scale.grad is None diff --git a/fastvideo/tests/train/methods/test_dmd2_timestep_bounds.py b/fastvideo/tests/train/methods/test_dmd2_timestep_bounds.py index 8f100f3874..ee56b856ae 100644 --- a/fastvideo/tests/train/methods/test_dmd2_timestep_bounds.py +++ b/fastvideo/tests/train/methods/test_dmd2_timestep_bounds.py @@ -4,6 +4,7 @@ from types import SimpleNamespace import pytest +import torch from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method @@ -18,7 +19,7 @@ def _method_with_ratios(min_ratio, max_ratio) -> DMD2Method: return method -def test_dmd2_score_timestep_bounds_match_legacy_recipe() -> None: +def test_dmd2_score_timestep_bounds_apply_ratios() -> None: method = _method_with_ratios(0.02, 0.98) assert method._parse_score_timestep_bounds() == (20, 980) @@ -39,3 +40,63 @@ def test_dmd2_score_timestep_bounds_reject_invalid_ranges( method = _method_with_ratios(min_ratio, max_ratio) with pytest.raises(ValueError, match="0 <= min <= max <= 1"): method._parse_score_timestep_bounds() + + +@pytest.mark.parametrize("warp_max", [0.0, -0.1, 1.1]) +def test_dmd2_score_timestep_warp_max_rejects_invalid_endpoint(warp_max: float) -> None: + method = _method_with_ratios(0.001, 0.999) + lo, hi = method._parse_score_timestep_bounds() + object.__setattr__(method, "_score_min_timestep", lo) + object.__setattr__(method, "_score_max_timestep", hi) + method.method_config["score_timestep_warp_max"] = warp_max + with pytest.raises(ValueError, match="0 < max <= 1"): + method._parse_score_timestep_warp_max() + + +def test_dmd2_score_timestep_warp_max_must_cover_upper_bound() -> None: + method = _method_with_ratios(0.001, 1.0) + lo, hi = method._parse_score_timestep_bounds() + object.__setattr__(method, "_score_min_timestep", lo) + object.__setattr__(method, "_score_max_timestep", hi) + method.method_config["score_timestep_warp_max"] = 0.999 + with pytest.raises(ValueError, match="must not exceed"): + method._parse_score_timestep_warp_max() + + +def test_dmd2_score_timestep_continuous_requires_bool() -> None: + method = _method_with_ratios(0.001, 0.999) + method.method_config["score_timestep_continuous"] = 1 + with pytest.raises(ValueError, match="must be a bool"): + method._parse_score_timestep_continuous() + + +def _sampler(min_ratio: float, max_ratio: float, shift: float) -> DMD2Method: + method = _method_with_ratios(min_ratio, max_ratio) + lo, hi = method._parse_score_timestep_bounds() + object.__setattr__(method, "_score_min_timestep", lo) + object.__setattr__(method, "_score_max_timestep", hi) + object.__setattr__(method, "_score_timestep_shift", shift) + object.__setattr__(method, "cuda_generator", torch.Generator().manual_seed(0)) + method.student.shift_and_clamp_timestep = lambda t: t + return method + + +def test_uniform_sampler_draws_in_bounds_without_boundary_atoms() -> None: + """shift=1 samples uniformly inside the configured integer bounds.""" + method = _sampler(0.02, 0.98, shift=1.0) + device = torch.device("cpu") + draws = torch.cat([method._sample_score_timestep(device) for _ in range(4000)]) + + assert int(draws.min()) >= 20 + assert int(draws.max()) <= 980 + assert int((draws == 20).sum()) < 20 + assert int((draws == 980).sum()) < 20 + + +def test_shifted_sampler_respects_bounds() -> None: + method = _sampler(0.005, 0.98, shift=12.0) + device = torch.device("cpu") + draws = torch.cat([method._sample_score_timestep(device) for _ in range(2000)]) + + assert int(draws.min()) >= 5 + assert int(draws.max()) <= 980 diff --git a/fastvideo/tests/train/methods/test_dmd2_vsd_normalizer.py b/fastvideo/tests/train/methods/test_dmd2_vsd_normalizer.py new file mode 100644 index 0000000000..e7f0102178 --- /dev/null +++ b/fastvideo/tests/train/methods/test_dmd2_vsd_normalizer.py @@ -0,0 +1,47 @@ +# SPDX-License-Identifier: Apache-2.0 +"""The VSD normalizer must stay finite when |gen - real| degenerates to 0.""" +from __future__ import annotations + +from types import SimpleNamespace + +import torch + +from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method + + +class _EchoTeacher: + """Returns the generator's own prediction: |gen - real| == 0 exactly.""" + + def __init__(self, gen: torch.Tensor): + self._gen = gen + + def predict_x0(self, noisy, timestep, batch, *, conditional, cfg_uncond, attn_kind): + return self._gen.detach().clone() + + +class _OffsetCritic(_EchoTeacher): + + def predict_x0(self, noisy, timestep, batch, *, conditional, cfg_uncond, attn_kind): + return self._gen.detach().clone() + 1.0 + + +def test_degenerate_denominator_yields_finite_loss_and_grad() -> None: + gen = torch.randn(1, 16, generator=torch.Generator().manual_seed(3)).requires_grad_(True) + + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", {}) + object.__setattr__(method, "cuda_generator", torch.Generator().manual_seed(0)) + object.__setattr__(method, "_cfg_uncond", None) + object.__setattr__(method, "student", SimpleNamespace(add_noise=lambda clean, noise, t: clean)) + object.__setattr__(method, "teacher", _EchoTeacher(gen)) + object.__setattr__(method, "critic", _OffsetCritic(gen)) + object.__setattr__(method, "_sample_score_timestep", lambda device: torch.tensor([500])) + + batch = SimpleNamespace(dmd_latent_vis_dict={}) + loss, metrics = method._dmd_loss(gen, batch) + + assert torch.isfinite(loss) + assert metrics == {} + loss.backward() + assert gen.grad is not None + assert torch.isfinite(gen.grad).all() diff --git a/fastvideo/tests/train/methods/test_minimax_h3_dmd2.py b/fastvideo/tests/train/methods/test_minimax_h3_dmd2.py new file mode 100644 index 0000000000..c1e76ee47b --- /dev/null +++ b/fastvideo/tests/train/methods/test_minimax_h3_dmd2.py @@ -0,0 +1,1122 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU contract tests for MiniMax H3 DMD2 distillation. + +Covers the packed dual-modality adapter (MiniMaxH3DMDModel) and one full +DMD2Method.single_train_step on a tiny CPU trio: student rollout, critic +flow-matching loss, generator DMD loss, both backwards, both optimizers. +""" + +import math +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch +import yaml + +from fastvideo.attention.backends.video_sparse_attn_h3 import MiniMaxH3VSAMetadata +from fastvideo.forward_context import get_forward_context +from fastvideo.pipelines.basic.minimax_h3.packing import audio_latent_num_frames, video_latent_num_frames +from fastvideo.platforms import AttentionBackendEnum +from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method +from fastvideo.train.models.minimax_h3 import MiniMaxH3DMDModel, MiniMaxH3Model +from fastvideo.train.models.minimax_h3.minimax_h3 import shift_noise_amount +from fastvideo.train.utils.config import load_run_config + +_FIXTURE = Path(__file__).resolve().parent.parent / "fixtures" / "minimax_h3_dmd2_min.yaml" +_REPO_ROOT = Path(__file__).resolve().parents[4] +_EXPERIMENT_CONFIG = (_REPO_ROOT / "examples/train/configs/distribution_matching/minimax_h3/dmd2_sp1_fsdp40_nuva_v9_dataforce_vsa64.yaml") +_V10_EXPERIMENT_CONFIG = (_REPO_ROOT / "examples/train/configs/distribution_matching/minimax_h3/dmd2_sp1_fsdp64_v10_dataonly_mixed_vsa64.yaml") +_V10_MAXSHAPE_CONFIG = (_REPO_ROOT / "examples/train/configs/distribution_matching/minimax_h3/dmd2_sp1_fsdp64_v10_maxshape_gate_vsa64.yaml") +_V10_PREPARE_LAUNCHER = _REPO_ROOT / "examples/train/slurm/prepare_h3_dmd2_v10_slinky.sh" +_H3_SBATCH = _REPO_ROOT / "examples/train/slurm/dmd2_32xgb200.sbatch" +_V10_GATED_LAUNCHER = _REPO_ROOT / "scripts/train/run_h3_v10_gated.sh" +_V10_MAXSHAPE_RUNNER = _REPO_ROOT / "scripts/train/run_h3_v10_maxshape_gate.sh" +_V10_KERNEL_GATE = _REPO_ROOT / "scripts/train/gate_h3_v10_kernel.sh" +_V10_KERNEL_REBUILD = _REPO_ROOT / "scripts/train/rebuild_h3_v10_kernel.sh" +_V10_KERNEL_RECEIPT_HELPER = _REPO_ROOT / "scripts/train/h3_v10_kernel_receipt.py" + +_VIDEO_SHAPE = (1, 2, 24, 4, 4) +_AUDIO_SHAPE = (1, 2, 32, 8) +_PACKED_NUMEL = math.prod(_VIDEO_SHAPE) + math.prod(_AUDIO_SHAPE) + + +class _TinyJointTransformer(torch.nn.Module): + """Scale packed H3 rows with one trainable parameter.""" + + patch_size = (1, 2, 2) + + def __init__(self, scale: float = 1.0) -> None: + super().__init__() + self.scale = torch.nn.Parameter(torch.tensor(scale)) + self.last_encoder_hidden_states: torch.Tensor | None = None + self.last_attn_metadata = None + + def forward(self, **kwargs): + self.last_encoder_hidden_states = kwargs["encoder_hidden_states"] + self.last_attn_metadata = get_forward_context().attn_metadata + return ( + kwargs["hidden_states"] * self.scale, + kwargs["audio_hidden_states"] * self.scale, + ) + + +def _make_model( + monkeypatch: pytest.MonkeyPatch, + training_config, + *, + trainable: bool = True, + scale: float = 1.0, +) -> MiniMaxH3DMDModel: + monkeypatch.setattr(MiniMaxH3Model, "device", property(lambda _self: torch.device("cpu"))) + model = MiniMaxH3DMDModel.__new__(MiniMaxH3DMDModel) + model._trainable = trainable + model.transformer = _TinyJointTransformer(scale) + model.training_config = training_config + model.sp_group = None + model.attention_backend = None + return model + + +def _tiny_training_config(): + return SimpleNamespace( + data=SimpleNamespace( + num_latent_t=2, + num_frames=5, + num_height=64, + num_width=64, + ), + distributed=SimpleNamespace(sp_size=1), + vsa_sparsity=0.0, + vsa_tile_size=256, + ) + + +def _raw_batch(seed: int = 1) -> dict[str, torch.Tensor]: + generator = torch.Generator().manual_seed(seed) + return { + "vae_latent": torch.randn(1, 24, 2, 4, 4, generator=generator), + "audio_latent": torch.randn(1, 2, 32, 8, generator=generator), + "text_embedding": torch.randn(1, 4, 5120, generator=generator), + "text_attention_mask": torch.tensor([[1, 1, 0, 0]], dtype=torch.float32), + } + + +def _native_raw_batch( + width: int, + height: int, + num_frames: int, + *, + seed: int = 1, +) -> dict: + generator = torch.Generator().manual_seed(seed) + return { + "vae_latent": torch.randn( + 1, + 24, + video_latent_num_frames(num_frames), + height // 16, + width // 16, + generator=generator, + dtype=torch.bfloat16, + ), + "audio_latent": torch.randn( + 1, + 2, + 32, + audio_latent_num_frames(num_frames), + generator=generator, + dtype=torch.bfloat16, + ), + "text_embedding": torch.randn(1, 4, 5120, generator=generator), + "text_attention_mask": torch.tensor([[1, 1, 0, 0]], dtype=torch.float32), + "_shape_bucket_id": f"bucket={width}x{height}-{num_frames}f", + "info_list": [{ + "width": width, + "height": height, + "num_frames": num_frames, + "fps": 24.0, + "audio_sample_rate": 32_000, + }], + } + + +def _build_method( + monkeypatch: pytest.MonkeyPatch, + *, + rollout_mode: str, + generator_update_interval: int = 1, +) -> DMD2Method: + config = load_run_config(str(_FIXTURE)) + config.method["rollout_mode"] = rollout_mode + config.method["generator_update_interval"] = generator_update_interval + student = _make_model(monkeypatch, config.training, scale=1.0) + teacher = _make_model(monkeypatch, config.training, trainable=False, scale=0.5) + critic = _make_model(monkeypatch, config.training, scale=0.25) + student.init_preprocessors = lambda training_config: None + method = DMD2Method( + cfg=config, + role_models={ + "student": student, + "teacher": teacher, + "critic": critic, + }, + ) + method.cuda_generator = torch.Generator(device="cpu").manual_seed(0) + return method + + + + +@pytest.mark.parametrize("rollout_mode", ["data_latent", "simulate"]) +def test_dmd2_student_iteration_updates_only_student( + monkeypatch: pytest.MonkeyPatch, + rollout_mode: str, +) -> None: + """Student iterations do not train or step the critic.""" + method = _build_method(monkeypatch, rollout_mode=rollout_mode) + student = method.student + teacher = method.teacher + critic = method.critic + + loss_map, outputs, metrics = method.single_train_step(_raw_batch(), iteration=0) + + assert metrics["update_student"] == 1.0 + for key in ("total_loss", "generator_loss", "fake_score_loss"): + assert torch.isfinite(loss_map[key]), key + assert loss_map["generator_loss"].item() > 0.0 + assert loss_map["fake_score_loss"].item() == 0.0 + torch.testing.assert_close(loss_map["total_loss"], loss_map["generator_loss"]) + assert "generator_pred_video" in method.latent_vis + assert "real_score_pred_video" in method.latent_vis + assert "faker_score_pred_video" in method.latent_vis + + method.backward(loss_map, outputs) + assert student.transformer.scale.grad is not None + assert torch.isfinite(student.transformer.scale.grad) + assert critic.transformer.scale.grad is None + assert teacher.transformer.scale.grad is None + assert method.get_optimizers(0) == [method._student_optimizer] + assert method.get_lr_schedulers(0) == [method._student_lr_scheduler] + assert method.get_grad_clip_targets(0) == {"student": student.transformer} + + student_before = student.transformer.scale.detach().clone() + critic_before = critic.transformer.scale.detach().clone() + method.optimizers_schedulers_step(0) + assert student.transformer.scale.detach() != student_before + torch.testing.assert_close(critic.transformer.scale.detach(), critic_before) + + +def test_dmd2_critic_iteration_updates_only_critic(monkeypatch: pytest.MonkeyPatch) -> None: + """Off-interval iterations do not train or step the student.""" + method = _build_method( + monkeypatch, + rollout_mode="data_latent", + generator_update_interval=5, + ) + + student = method.student + critic = method.critic + loss_map, outputs, metrics = method.single_train_step(_raw_batch(), iteration=1) + + assert metrics["update_student"] == 0.0 + assert loss_map["generator_loss"].item() == 0.0 + assert loss_map["fake_score_loss"].item() > 0.0 + torch.testing.assert_close(loss_map["total_loss"], loss_map["fake_score_loss"]) + method.backward(loss_map, outputs) + assert student.transformer.scale.grad is None + assert critic.transformer.scale.grad is not None + assert torch.isfinite(critic.transformer.scale.grad) + assert method.get_optimizers(1) == [method._critic_optimizer] + assert method.get_lr_schedulers(1) == [method._critic_lr_scheduler] + assert method.get_grad_clip_targets(1) == {"critic": critic.transformer} + + student_before = student.transformer.scale.detach().clone() + critic_before = critic.transformer.scale.detach().clone() + method.optimizers_schedulers_step(1) + torch.testing.assert_close(student.transformer.scale.detach(), student_before) + assert critic.transformer.scale.detach() != critic_before + + +def test_dmd2_five_step_cadence_and_resume_state(monkeypatch: pytest.MonkeyPatch) -> None: + method = _build_method( + monkeypatch, + rollout_mode="simulate", + generator_update_interval=5, + ) + + assert [method._should_update_student(i) for i in range(1, 6)] == [ + False, + False, + False, + False, + True, + ] + method.method_config["generator_update_interval"] = 0 + with pytest.raises(ValueError, match="must be positive"): + method._should_update_student(0) + + method.seed_optimizer_state_for_resume() + for optimizer in (method._student_optimizer, method._critic_optimizer): + assert optimizer.state + assert all("exp_avg" in state for state in optimizer.state.values()) + + method.method_config.pop("generator_update_interval") + assert [method._should_update_student(i) for i in range(1, 6)] == [ + False, + False, + False, + False, + True, + ] + + + + +def test_packed_adapter_roundtrip_and_prepare_batch(monkeypatch: pytest.MonkeyPatch) -> None: + """Verify pack/unpack inversion and packed clean latents in the batch.""" + model = _make_model(monkeypatch, _tiny_training_config()) + video = torch.randn(_VIDEO_SHAPE) + audio = torch.randn(_AUDIO_SHAPE) + + packed = model.pack_latents(video, audio) + assert packed.shape == (1, _PACKED_NUMEL) + video_out, audio_out = model.unpack_latents(packed) + torch.testing.assert_close(video_out, video) + torch.testing.assert_close(audio_out, audio) + + raw_batch = _raw_batch() + batch = model.prepare_batch( + raw_batch, + generator=torch.Generator().manual_seed(7), + ) + assert batch.latents.shape == (1, _PACKED_NUMEL) + video_clean, audio_clean = model.unpack_latents(batch.latents) + torch.testing.assert_close( + video_clean, + raw_batch["vae_latent"].permute(0, 2, 1, 3, 4).to(torch.bfloat16), + ) + torch.testing.assert_close(audio_clean, batch.audio_latents) + + +def test_native_layout_is_batch_local_across_successive_shapes(monkeypatch: pytest.MonkeyPatch) -> None: + """A later shape must not mutate how an earlier packed tensor is split.""" + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + + first = model.prepare_batch( + _native_raw_batch(64, 64, 5), + generator=torch.Generator().manual_seed(3), + ) + first_packed = first.latents.clone() + first_layout = first.minimax_h3_dmd_layout + second = model.prepare_batch( + _native_raw_batch(96, 64, 22), + generator=torch.Generator().manual_seed(4), + ) + + assert first.minimax_h3_dmd_layout is first_layout + assert first_layout != second.minimax_h3_dmd_layout + first_video, first_audio = model.unpack_latents(first_packed, layout=first_layout) + assert first_video.shape == (1, 2, 24, 4, 4) + assert first_audio.shape == (1, 2, 32, 8) + second_video, second_audio = model.unpack_latents(second.latents, layout=second.minimax_h3_dmd_layout) + assert second_video.shape == (1, 7, 24, 4, 6) + assert second_audio.shape == (1, 2, 32, 37) + + noise = torch.zeros_like(second.latents) + mixed = model.add_noise_for_batch(second.latents, noise, torch.tensor([500]), second) + assert mixed.shape == second.latents.shape + slices = dict(model.modality_slices_for_batch(second)) + assert slices["video"].stop == math.prod(second_video.shape) + assert slices["audio"].stop == second.latents.shape[1] + + prediction = model.predict_noise( + second.latents, + torch.tensor([500]), + second, + conditional=True, + ) + torch.testing.assert_close(prediction, -second.latents) + + +def test_dmd_losses_follow_successive_native_modality_slices(monkeypatch: pytest.MonkeyPatch) -> None: + """Student and critic phases both use the active batch's exact split.""" + method = _build_method( + monkeypatch, + rollout_mode="data_latent", + generator_update_interval=2, + ) + method.training_config.data.native_shape_bucketing = True + method.method_config["modality_loss_weights"] = {"video": 0.25, "audio": 2.0} + + student_losses, _, student_metrics = method.single_train_step( + _native_raw_batch(64, 64, 5), + iteration=2, + ) + critic_losses, _, critic_metrics = method.single_train_step( + _native_raw_batch(96, 64, 22), + iteration=3, + ) + + assert student_metrics["update_student"] == 1.0 + assert {"generator_loss_video", "generator_loss_audio"} <= student_metrics.keys() + assert torch.isfinite(student_losses["generator_loss"]) + assert critic_metrics["update_student"] == 0.0 + assert {"fake_score_loss_video", "fake_score_loss_audio"} <= critic_metrics.keys() + assert torch.isfinite(critic_losses["fake_score_loss"]) + + +@pytest.mark.parametrize( + ("width", "height", "num_frames"), + [ + (1344, 768, 124), + (768, 1344, 362), + (832, 480, 90), + (480, 832, 124), + ], +) +def test_native_validation_accepts_min_max_portrait_and_lowres( + monkeypatch: pytest.MonkeyPatch, + width: int, + height: int, + num_frames: int, +) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + raw = _native_raw_batch(width, height, num_frames) + + video, audio = model._resolve_clean_latents(raw, "data", torch.bfloat16, torch.device("cpu")) + + assert video.shape == ( + 1, + 24, + video_latent_num_frames(num_frames), + height // 16, + width // 16, + ) + assert audio.shape == (1, 2, 32, audio_latent_num_frames(num_frames)) + + +@pytest.mark.parametrize( + ("mutation", "message"), + [ + (lambda raw: raw["info_list"][0].update(width=96), "disagrees with row metadata"), + (lambda raw: raw["info_list"][0].update(fps=30.0), "24 fps clock"), + (lambda raw: raw["info_list"][0].update(audio_sample_rate=44_100), "32000 Hz"), + (lambda raw: raw.update(audio_latent=raw["audio_latent"][..., :-1]), "audio clock"), + (lambda raw: raw.update(vae_latent=raw["vae_latent"][:, :, :-1]), "vae_latent shape"), + ], +) +def test_native_validation_rejects_bucket_metadata_and_clock_mismatches( + monkeypatch: pytest.MonkeyPatch, + mutation, + message: str, +) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + raw = _native_raw_batch(64, 64, 5) + mutation(raw) + + with pytest.raises(ValueError, match=message): + model._resolve_clean_latents(raw, "data", torch.bfloat16, torch.device("cpu")) + + +def test_native_validation_requires_production_canvas_multiple(monkeypatch: pytest.MonkeyPatch) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + raw = _native_raw_batch(64, 64, 5) + raw["_shape_bucket_id"] = "bucket=80x64-5f" + raw["info_list"][0]["width"] = 80 + raw["vae_latent"] = torch.zeros(1, 24, 2, 4, 5) + + with pytest.raises(ValueError, match="canvas multiple 32"): + model._resolve_clean_latents(raw, "data", torch.bfloat16, torch.device("cpu")) + + +def test_legacy_fixed_data_path_still_truncates_to_config(monkeypatch: pytest.MonkeyPatch) -> None: + config = _tiny_training_config() + model = _make_model(monkeypatch, config) + raw = _raw_batch() + raw["vae_latent"] = torch.cat((raw["vae_latent"], raw["vae_latent"][:, :, :1]), dim=2) + raw["audio_latent"] = torch.cat((raw["audio_latent"], raw["audio_latent"][..., :2]), dim=-1) + + video, audio = model._resolve_clean_latents(raw, "data", torch.bfloat16, torch.device("cpu")) + + assert video.shape == (1, 24, 2, 4, 4) + assert audio.shape == (1, 2, 32, 8) + + +def test_simulate_zeros_follow_native_shape_bucket(monkeypatch: pytest.MonkeyPatch) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + + raw = {"_shape_bucket_id": "bucket=96x64-22f"} + video, audio = model._resolve_clean_latents(raw, "zeros", torch.bfloat16, torch.device("cpu")) + + assert video.shape == (1, 24, 7, 4, 6) + assert audio.shape == (1, 2, 32, 37) + + +def test_native_simulate_requires_shape_bucket(monkeypatch: pytest.MonkeyPatch) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + + with pytest.raises(ValueError, match="data-free batches require.*_shape_bucket_id"): + model._resolve_clean_latents({}, "zeros", torch.bfloat16, torch.device("cpu")) + + +@pytest.mark.parametrize( + ("bucket_id", "video_shape", "audio_frames"), + [ + ("bucket=1760x768-362f", (1, 24, 107, 48, 110), audio_latent_num_frames(362)), + ("bucket=768x1344-124f", (1, 24, 37, 84, 48), audio_latent_num_frames(124)), + ], +) +def test_native_simulate_zeros_cover_production_extremes( + monkeypatch: pytest.MonkeyPatch, + bucket_id: str, + video_shape: tuple[int, ...], + audio_frames: int, +) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + + video, audio = model._resolve_clean_latents( + {"_shape_bucket_id": bucket_id}, + "zeros", + torch.bfloat16, + torch.device("cpu"), + ) + + assert video.shape == video_shape + assert audio.shape == (1, 2, 32, audio_frames) + + +def test_native_simulate_rejects_non_aligned_canvas(monkeypatch: pytest.MonkeyPatch) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + + with pytest.raises(ValueError, match="canvas multiple 32"): + model._resolve_clean_latents( + {"_shape_bucket_id": "bucket=80x64-5f"}, + "zeros", + torch.bfloat16, + torch.device("cpu"), + ) + + +def test_packed_add_noise_applies_modality_shifts(monkeypatch: pytest.MonkeyPatch) -> None: + """One shared base timestep must map to two shifted noise amounts.""" + model = _make_model(monkeypatch, _tiny_training_config()) + clean = torch.ones(1, _PACKED_NUMEL) + noise = torch.zeros(1, _PACKED_NUMEL) + + torch.testing.assert_close( + model.add_noise(clean, noise, torch.tensor([0])), + clean, + ) + torch.testing.assert_close( + model.add_noise(clean, noise, torch.tensor([1000])), + noise, + ) + + mixed = model.add_noise(clean, noise, torch.tensor([500])) + video_mixed, audio_mixed = model.unpack_latents(mixed) + base = torch.tensor([0.5]) + torch.testing.assert_close( + video_mixed, + torch.full(_VIDEO_SHAPE, float(1.0 - shift_noise_amount(base, 12.0))), + ) + torch.testing.assert_close( + audio_mixed, + torch.full(_AUDIO_SHAPE, float(1.0 - shift_noise_amount(base, 3.0))), + ) + + +def test_packed_predict_noise_plumbs_timesteps_and_tolerates_vsa(monkeypatch: pytest.MonkeyPatch, ) -> None: + """Explicit method timesteps must rewrite both modality clean-times.""" + model = _make_model(monkeypatch, _tiny_training_config()) + batch = model.prepare_batch(_raw_batch(), generator=torch.Generator().manual_seed(7)) + noisy = torch.randn(1, _PACKED_NUMEL).to(torch.bfloat16) + timestep = torch.tensor([757], dtype=torch.long) + + prediction = model.predict_noise( + noisy, + timestep, + batch, + conditional=True, + attn_kind="vsa", + ) + + base = torch.tensor([0.757]) + torch.testing.assert_close( + batch.timesteps, + 1.0 - shift_noise_amount(base.double(), 12.0), + ) + torch.testing.assert_close( + batch.audio_timesteps, + 1.0 - shift_noise_amount(base.double(), 3.0), + ) + torch.testing.assert_close(prediction, -noisy) + + x0 = model.predict_x0(noisy, timestep, batch, conditional=True) + noisy_video, noisy_audio = model.unpack_latents(noisy) + sigma_video = shift_noise_amount(base.double(), 12.0) + sigma_audio = shift_noise_amount(base.double(), 3.0) + expected = model.pack_latents( + (noisy_video.double() + sigma_video * noisy_video.double()).to(torch.bfloat16), + (noisy_audio.double() + sigma_audio * noisy_audio.double()).to(torch.bfloat16), + ) + torch.testing.assert_close(x0, expected) + + +def test_uncond_forward_zeroes_text_and_guards_policies(monkeypatch: pytest.MonkeyPatch) -> None: + """Teacher-CFG unconditional forwards zero text; other policies fail fast.""" + model = _make_model(monkeypatch, _tiny_training_config()) + batch = model.prepare_batch(_raw_batch(), generator=torch.Generator().manual_seed(7)) + noisy = torch.randn(1, _PACKED_NUMEL).to(torch.bfloat16) + timestep = torch.tensor([500], dtype=torch.long) + + model.predict_noise( + noisy, + timestep, + batch, + conditional=False, + cfg_uncond={"text": "zero"}, + ) + assert torch.all(model.transformer.last_encoder_hidden_states == 0) + + model.predict_noise( + noisy, + timestep, + batch, + conditional=True, + cfg_uncond={"text": "zero"}, + ) + assert torch.any(model.transformer.last_encoder_hidden_states != 0) + + with pytest.raises(ValueError, match="cfg_uncond"): + model.predict_noise(noisy, timestep, batch, conditional=False) + with pytest.raises(ValueError, match="negative-prompt"): + model.set_requires_negative_conditioning(True) + model.set_requires_negative_conditioning(False) + + + + +def test_prepare_batch_builds_vsa_h3_metadata(monkeypatch: pytest.MonkeyPatch) -> None: + """The VSA-H3 role gets real packed-sequence metadata; dense view stays None.""" + tc = _tiny_training_config() + tc.vsa_sparsity = 0.35 + tc.vsa_tile_size = 256 + model = _make_model(monkeypatch, tc) + model.attention_backend = AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3 + + batch = model.prepare_batch(_raw_batch(), generator=torch.Generator().manual_seed(7)) + + meta = batch.attn_metadata_vsa + assert isinstance(meta, MiniMaxH3VSAMetadata) + assert batch.attn_metadata is None + assert meta.VSA_sparsity == pytest.approx(0.35) + assert meta.total_seq_length == 26 + assert meta.num_prefix_tiles == 2 + assert meta.num_video_tiles == 1 + assert meta.variable_block_sizes.tolist() == [2, 16, 8] + assert int(meta.variable_block_sizes.sum()) == meta.total_seq_length + + +def test_predict_noise_routes_vsa_metadata_by_attn_kind(monkeypatch: pytest.MonkeyPatch) -> None: + """Student "vsa" forwards see the VSA metadata; "dense" forwards see None.""" + model = _make_model(monkeypatch, _tiny_training_config()) + model.attention_backend = AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3 + batch = model.prepare_batch(_raw_batch(), generator=torch.Generator().manual_seed(7)) + noisy = torch.randn(1, _PACKED_NUMEL).to(torch.bfloat16) + timestep = torch.tensor([757], dtype=torch.long) + + model.predict_noise(noisy, timestep, batch, conditional=True, attn_kind="vsa") + assert model.transformer.last_attn_metadata is batch.attn_metadata_vsa + assert isinstance(model.transformer.last_attn_metadata, MiniMaxH3VSAMetadata) + + model.predict_noise(noisy, timestep, batch, conditional=True, attn_kind="dense") + assert model.transformer.last_attn_metadata is None + + +def _ctor_training_config() -> SimpleNamespace: + """The minimum surface MiniMaxH3Model.__init__ reads from TrainingConfig.""" + return SimpleNamespace( + pipeline_config=SimpleNamespace(dit_config=SimpleNamespace(uniform_parameter_dtype=False)), + data=SimpleNamespace( + train_batch_size=1, + training_cfg_rate=0.0, + preprocessed_data_type="t2va", + ), + model=SimpleNamespace(enable_gradient_checkpointing_type=None), + ) + + +def test_per_role_attention_backend_override_resolves(monkeypatch: pytest.MonkeyPatch) -> None: + """Each role's backend reaches the loader; unsupported backends fail fast.""" + captured: dict[str, AttentionBackendEnum | None] = {} + + def _fake_load(**kwargs): + captured[kwargs["model_path"]] = kwargs["attention_backend"] + return _TinyJointTransformer() + + monkeypatch.setattr( + "fastvideo.train.models.minimax_h3.minimax_h3.load_module_from_path", + _fake_load, + ) + + student_config = _ctor_training_config() + student = MiniMaxH3DMDModel( + init_from="role/student", + training_config=student_config, + trainable=True, + attention_backend="VIDEO_SPARSE_ATTN_H3", + ) + teacher = MiniMaxH3DMDModel( + init_from="role/teacher", + training_config=_ctor_training_config(), + trainable=False, + attention_backend="FLASH_ATTN", + ) + + assert student.attention_backend is AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3 + assert student_config.pipeline_config.dit_config.uniform_parameter_dtype is False + assert teacher.attention_backend is AttentionBackendEnum.FLASH_ATTN + assert captured["role/student"] is AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3 + assert captured["role/teacher"] is AttentionBackendEnum.FLASH_ATTN + assert not any(p.requires_grad for p in teacher.transformer.parameters()) + + with pytest.raises(ValueError, match="supports the attention backends"): + MiniMaxH3DMDModel( + init_from="role/bad", + training_config=_ctor_training_config(), + attention_backend="VIDEO_SPARSE_ATTN", + ) + + + + +def test_h3_dmd2_fixture_resolves_trio_contract() -> None: + """The fixture must wire the H3 DMD trio through the modular builder path.""" + config = load_run_config(str(_FIXTURE)) + + for role in ("student", "teacher", "critic"): + assert config.models[role]["_target_"] == ("fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel") + assert config.models["teacher"]["trainable"] is False + assert config.method["_target_"] == ("fastvideo.train.methods.distribution_matching.dmd2.DMD2Method") + assert config.training.data.preprocessed_data_type == "t2va" + + +def test_h3_dmd2_current_config_pins_recipe() -> None: + """The production config pins the intended H3 DMD2 knobs.""" + config = yaml.safe_load(_EXPERIMENT_CONFIG.read_text()) + method = config["method"] + training = config["training"] + + for role in ("student", "teacher", "critic"): + assert config["models"][role]["_target_"] == ("fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel") + assert config["models"]["teacher"]["trainable"] is False + assert config["models"]["critic"]["trainable"] is True + assert method["_target_"] == ("fastvideo.train.methods.distribution_matching.dmd2.DMD2Method") + assert method["rollout_mode"] == "simulate" + assert method["rollout_carry"] is True + assert (method["rollout_carry_slots"] == training["loop"]["gradient_accumulation_steps"]) + assert training["loop"]["gradient_accumulation_steps"] == 4 + assert method["rollout_sample_type"] == "ode" + assert method["rollout_data_forcing"] is True + assert method["allow_mixed_rollout_regimes"] is True + assert method["generator_update_interval"] == 5 + assert method["real_score_guidance_scale"] == 1.0 + assert method["dmd_denoising_steps"] == [999, 749, 500, 250] + assert "warp_denoising_step" not in method + assert method["score_timestep_shift"] == 2.4 + assert method["score_timestep_warp_max"] == 0.999 + assert method["score_timestep_continuous"] is True + assert method["min_timestep_ratio"] == 0.001 + assert method["max_timestep_ratio"] == 0.999 + assert method["fake_score_loss_space"] == "x0" + assert method["cfg_uncond"] == {"text": "zero"} + assert method["fake_score_learning_rate"] == training["optimizer"]["learning_rate"] + assert method["fake_score_betas"] == [0.9, 0.999] + assert method["fake_score_lr_scheduler"] == "constant" + assert training["optimizer"]["betas"] == [0.9, 0.999] + assert training["dit_precision"] == "fp32" + assert training["checkpoint"]["output_dir"].endswith("v9_dataforce_vsa64") + assert training["vsa"] == {"sparsity": 0.9, "tile_size": 64} + assert (config["models"]["student"]["attention_backend"] == "VIDEO_SPARSE_ATTN_H3") + assert config["models"]["teacher"]["attention_backend"] == "FLASH_ATTN" + assert config["models"]["critic"]["attention_backend"] == "FLASH_ATTN" + assert config["pipeline"]["dit_config"]["uniform_parameter_dtype"] is False + assert training["data"]["preprocessed_data_type"] == "t2va" + data_paths = training["data"]["data_path"] + assert any("nuva_t2va" in str(path) for path in data_paths) + assert any("text_only" in str(path) for path in data_paths) + assert training["data"]["train_batch_size"] == 1 + assert training["data"]["training_cfg_rate"] == 0.0 + assert config["callbacks"]["grad_clip"]["max_grad_norm"] == 1.0 + assert training["model"]["enable_torch_compile"] is True + assert "torch_compile_kwargs" not in training["model"] + + +def test_h3_dmd2_v10_config_pins_data_only_native_shape_recipe() -> None: + """V10 is a fresh 64-GPU, global-batch-64, all-real-latent lineage.""" + config = yaml.safe_load(_V10_EXPERIMENT_CONFIG.read_text()) + method = config["method"] + training = config["training"] + distributed = training["distributed"] + data = training["data"] + + assert method["rollout_mode"] == "data_latent" + for carry_key in ( + "rollout_carry", + "rollout_carry_slots", + "rollout_sample_type", + "rollout_data_forcing", + ): + assert carry_key not in method + assert method["dmd_denoising_steps"] == [999, 749, 500, 250] + assert method["fake_score_learning_rate"] == 2.0e-6 + assert training["optimizer"]["learning_rate"] == 2.0e-6 + + assert distributed == { + "num_gpus": 64, + "sp_size": 1, + "tp_size": 1, + "hsdp_replicate_dim": 1, + "hsdp_shard_dim": 64, + } + global_batch = (distributed["num_gpus"] // distributed["sp_size"] * data["train_batch_size"] * + training["loop"]["gradient_accumulation_steps"]) + assert global_batch == 64 + assert data["preprocessed_data_type"] == "t2va" + assert data["native_shape_bucketing"] is True + assert len(data["data_path"]) == 5 + assert all(path.startswith("/mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/") + and path.endswith("/data") for path in data["data_path"]) + + checkpoint = training["checkpoint"] + assert checkpoint["output_dir"].endswith("v10_dataonly_mixed_vsa64") + assert training["loop"]["gradient_accumulation_steps"] == 1 + assert checkpoint["save_inference_checkpoint_on_validation"] is True + assert checkpoint["inference_checkpoint_role"] == "student" + assert checkpoint["inference_checkpoint_dtype"] == "bfloat16" + assert checkpoint["training_state_checkpointing_steps"] == 100 + assert checkpoint["require_complete_training_checkpoint"] is True + assert checkpoint["checkpointing_start_step"] == 100 + assert checkpoint["checkpoints_total_limit"] == 3 + assert training["tracker"]["run_name"] == "dmd2_sp1_fsdp64_v10_dataonly_mixed_vsa64" + assert training["model"]["enable_torch_compile"] is True + assert training["model"]["torch_compile_kwargs"] == { + "dynamic": True, + "recompile_limit": 32, + } + assert training["vsa"] == {"sparsity": 0.9, "tile_size": 64} + assert config["models"]["student"]["attention_backend"] == "VIDEO_SPARSE_ATTN_H3" + for role in ("teacher", "critic"): + assert config["models"][role]["attention_backend"] == "FLASH_ATTN" + + validation = config["callbacks"]["validation"] + assert validation["dataset_file"].endswith("/validation/heldout60.json") + assert validation["every_steps"] == 100 + assert validation["run_at_start"] is True + assert validation["sampling_steps"] == [4] + assert validation["use_record_dimensions"] is True + assert validation["max_record_num_frames"] == 345 + assert validation["use_validation_media_conditioning"] is False + + +def test_h3_dmd2_v10_prepare_launcher_pins_finalized_data_and_execution_clone() -> None: + """The non-submitting helper gates the dedicated clone and immutable dataset.""" + launcher = _V10_PREPARE_LAUNCHER.read_text() + + assert "/mnt/lustre/vlm-wlsaidhi/fastvideo/FastVideo-v10" in launcher + assert ('readonly CONFIG="${REPO}/examples/train/configs/distribution_matching/minimax_h3/' + 'dmd2_sp1_fsdp64_v10_dataonly_mixed_vsa64.yaml"') in launcher + assert 'readonly DATA_ROOT="/mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3"' in launcher + assert 'readonly BASE_DATA_ROOT="/mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v2"' in launcher + assert 'readonly VALIDATION_MANIFEST="${DATA_ROOT}/validation/heldout60.json"' in launcher + assert "readonly VALIDATION_MAX_RECORD_NUM_FRAMES=345" in launcher + assert ('readonly OUTPUT_DIR="/mnt/lustre/vlm-wlsaidhi/fastvideo/outputs/' + 'minimax_h3_dmd2_sp1_fsdp64_v10_dataonly_mixed_vsa64"') in launcher + assert "readonly NUM_NODES=16" in launcher + assert "readonly WORLD_SIZE=64" in launcher + assert "readonly HSDP_REPLICATE=1" in launcher + assert "readonly HSDP_SHARD=64" in launcher + assert "readonly GRADIENT_ACCUMULATION_STEPS=1" in launcher + assert "readonly GLOBAL_BATCH_SIZE=64" in launcher + assert 'PARTITION="${PARTITION:-hpc-rack-3}"' in launcher + assert "readonly EXPECTED_TRAINING_ROWS=60549" in launcher + assert "readonly EXPECTED_SHAPE_BUCKETS=87" in launcher + assert "readonly EXPECTED_SCHEDULED_ROWS=63424" in launcher + assert "readonly EXPECTED_PADDED_ROWS=2875" in launcher + assert "readonly EXPECTED_STEPS_PER_EPOCH=991" in launcher + assert "readonly EXPECTED_VALIDATION_ROWS=60" in launcher + assert "readonly EXPECTED_VALIDATION_DP_PADDING=4" in launcher + assert "readonly MIN_OUTPUT_FREE_BYTES=" in launcher + for removed_override in ("CONFIG", "DATA_ROOT", "VALIDATION_MANIFEST", "OUTPUT_DIR"): + assert f'${{{removed_override}:-' not in launcher + assert 'readonly REVIEWED_V10_COMMIT="7635a5295b027000a00f6d70789c5cb5886218c3"' in launcher + assert 'merge-base --is-ancestor "${REVIEWED_V10_COMMIT}" HEAD' in launcher + assert "actual_data_paths = training[\"data\"][\"data_path\"]" in launcher + assert 'actual_validation = validation["dataset_file"]' in launcher + assert "actual_validation_max_record_num_frames" in launcher + assert 'validation.get("sampling_steps") != [4]' in launcher + assert "actual_output = training[\"checkpoint\"][\"output_dir\"]" in launcher + assert "actual_topology != expected_topology" in launcher + assert "actual_global_batch_size != global_batch_size" in launcher + assert 'compile_kwargs.get("dynamic") is not True' in launcher + assert "kernel receipt source" in launcher + assert 'readonly MAXSHAPE_AUDIT_ROOT=' in launcher + assert 'audit_root.glob("job-*/RESULT.json")' in launcher + assert "fastvideo-h3-v10-maxshape-gate-v1" in launcher + assert "no successful final-commit 64-GPU 1760x768x362 capacity receipt" in launcher + assert "available_bytes < MIN_OUTPUT_FREE_BYTES" in launcher + assert "validated same-recipe pre-step-100 restart namespace" in launcher + assert 'checkpoint_dir_re = re.compile(r"^checkpoint-([0-9]+)$")' in launcher + assert 'dcp_metadata = checkpoint / "dcp" / ".metadata"' in launcher + assert 'marker != "complete\\n"' in launcher + assert 'expected_rng_names = {f"rng_state_rank{rank}.pt" for rank in range(world_size)}' in launcher + assert 'len(rng_paths) != world_size' in launcher + assert "incomplete training checkpoint is not newer than the latest strict fallback" in launcher + assert "required_inference_steps = set(range(0, latest_step + 1, 100))" in launcher + assert "missing unlimited-retention inference checkpoints" in launcher + assert "completed validation step {step} does not have all 64 DP-padded names" in launcher + assert 'saved_method.get("dmd_denoising_steps") != [999, 749, 500, 250]' in launcher + assert "belongs to a different data/heldout60 recipe" in launcher + assert "shard/tensor counts differ from its index" in launcher + assert "step-zero validation must preserve all 64 DP-padded four-forward names" in launcher + assert 'require_file "${DATA_ROOT}/READY.json"' in launcher + assert 'require_file "${DATA_ROOT}/DERIVATION_RECEIPT.json"' in launcher + assert 'require_file "${source_root}/READY.json"' in launcher + assert 'require_file "${source_root}/MANIFEST.json"' in launcher + assert 'require_file "${source_root}/MANIFEST_rows.jsonl"' in launcher + assert 'require_file "${source_root}/data/map_style_cache/file_info.pkl"' in launcher + assert "finalize_dataset.py" in launcher + assert "derive_filtered_dataset.py" in launcher + assert '"excluded_resolutions": ["576x576", "640x480", "832x480"]' in launcher + assert '"excluded_validation_rows": 4' in launcher + assert '"training_holdout_policy": "preserve_base_validation_conditioning_ids"' in launcher + assert '"validation_payload_path": "validation/heldout60.json"' in launcher + assert "minimax-h3-native-t2va-filtered-derivation-v2" in launcher + assert '"validation_summary_sha256": sha256(data_root / "validation" / "manifest.json")' in launcher + assert "validation DP padding does not repeat exactly the first four retained records" in launcher + assert "excluded resolution leaked into v3 cache" in launcher + assert "--verify-only" in launcher + assert 'require_file "${REPO}/scripts/train/run_h3_v10_gated.sh"' in launcher + assert "--nodes=%q" in launcher + assert "--no-requeue" in launcher + assert "--wrap" not in launcher + assert '"${REPO}/scripts/train/run_h3_v10_gated.sh" "${execution_commit}"' in launcher + assert 'git -C "${REPO}" status --porcelain' in launcher + assert "This helper never calls sbatch" in launcher + + sbatch = _H3_SBATCH.read_text() + assert "HSDP_REPLICATE * HSDP_SHARD != WORLD_SIZE" in sbatch + assert 'git -C "${REPO}" status --porcelain' in sbatch + assert "V10 SOURCE GATE FAILED: execution checkout is dirty" in sbatch + assert "export HOME=" not in sbatch + for runtime_variable in ( + "HF_HOME", + "XDG_CACHE_HOME", + "TORCH_HOME", + "TRITON_CACHE_DIR", + "TORCHINDUCTOR_CACHE_DIR", + "TORCH_EXTENSIONS_DIR", + "FLASHINFER_WORKSPACE_BASE", + "CUDA_CACHE_PATH", + "NUMBA_CACHE_DIR", + "WANDB_CONFIG_DIR", + "WANDB_CACHE_DIR", + "WANDB_DATA_DIR", + "NETRC", + ): + assert f"export {runtime_variable}=" in sbatch + + +def test_h3_dmd2_v10_maxshape_gate_matches_production_capacity_contract() -> None: + config = yaml.safe_load(_V10_MAXSHAPE_CONFIG.read_text()) + training = config["training"] + distributed = training["distributed"] + checkpoint = training["checkpoint"] + + assert distributed == { + "num_gpus": 64, + "sp_size": 1, + "tp_size": 1, + "hsdp_replicate_dim": 1, + "hsdp_shard_dim": 64, + } + assert training["data"]["data_path"].endswith("/v10_maxshape_64g/data") + assert training["data"]["native_shape_bucketing"] is True + assert training["data"]["train_batch_size"] == 1 + assert training["loop"] == {"max_train_steps": 2, "gradient_accumulation_steps": 1} + assert config["method"]["generator_update_interval"] == 2 + assert training["model"]["enable_torch_compile"] is True + assert training["model"]["torch_compile_kwargs"] == { + "dynamic": True, + "recompile_limit": 32, + } + assert checkpoint["resume_from_checkpoint"] == "latest" + assert checkpoint["save_inference_checkpoint_on_validation"] is False + assert checkpoint["training_state_checkpointing_steps"] == 0 + assert "validation" not in config.get("callbacks", {}) + + runner = _V10_MAXSHAPE_RUNNER.read_text() + assert '"${SLURM_JOB_NUM_NODES:-0}" != "16"' in runner + assert "/v10_mixed_native_v3/" in runner + assert "/v10_mixed_native_v2/" not in runner + assert '"${staged}" -ef "${source}"' in runner + assert "all_64_gpus_sampled" in runner + assert "critic_grad_finite_positive" in runner + assert "student_grad_finite_positive" in runner + assert "dense_teacher_and_critic_compiled" in runner + assert "vsa_grad_used_triton64" in runner + assert '"execution_commit_unchanged": observed_execution_commit == expected_execution_commit' in runner + assert '"execution_checkout_clean": execution_checkout_clean' in runner + assert '"execution_commit": expected_execution_commit' in runner + + sbatch = _H3_SBATCH.read_text() + assert 'TRAIN_LOG_ROOT="${H3_V10_TRAIN_LOG_ROOT:-${REPO}/examples/train/logs}"' in sbatch + + +def test_h3_dmd2_v10_gated_launcher_defers_requeue_until_all_gates_pass() -> None: + launcher = _V10_GATED_LAUNCHER.read_text() + + assert launcher.startswith("#!/bin/bash\n") + assert "#SBATCH --nodes=16" in launcher + assert "#SBATCH --gpus-per-node=4" in launcher + assert "#SBATCH --no-requeue" in launcher + assert "--wrap" not in launcher + assert '(( $# != 1 ))' in launcher + assert '[[ ! "$1" =~ ^[0-9a-f]{40}$ ]]' in launcher + assert 'readonly EXPECTED_V10_COMMIT="$1"' in launcher + assert 'requested commit ${EXPECTED_V10_COMMIT} != execution HEAD ${execution_commit}' in launcher + assert launcher.count("require_exact_execution_checkout") == 4 + assert 'MAXSHAPE_RECEIPT="${MAXSHAPE_ROOT}/audit/job-${SLURM_JOB_ID:' in launcher + assert 'if [[ ! -f "${MAXSHAPE_RECEIPT}" ]]' in launcher + assert '"job_id": job_id' in launcher + assert '"${MAXSHAPE_CONFIG}" "${EXPECTED_V10_COMMIT}" "${SLURM_JOB_ID}"' in launcher + assert 'bash "${REPO}/scripts/train/run_h3_v10_maxshape_gate.sh"' in launcher + preflight = 'bash "${REPO}/examples/train/slurm/prepare_h3_dmd2_v10_slinky.sh"' + enable_requeue = 'scontrol update JobId="${SLURM_JOB_ID}" Requeue=1' + production = 'exec bash "${REPO}/examples/train/slurm/dmd2_32xgb200.sbatch"' + assert launcher.index(preflight) < launcher.index(enable_requeue) < launcher.index(production) + assert 'scontrol update JobId="${SLURM_JOB_ID}" Requeue=0' in launcher + assert "export H3_V10_KERNEL_GATE=0" in launcher + assert "export GATE_TEST=0" in launcher + + +def test_h3_dmd2_v10_kernel_gate_pins_import_order_and_real_gpu_checks() -> None: + gated_launcher = _V10_GATED_LAUNCHER.read_text() + sbatch = _H3_SBATCH.read_text() + gate = _V10_KERNEL_GATE.read_text() + rebuild = _V10_KERNEL_REBUILD.read_text() + receipt_helper = _V10_KERNEL_RECEIPT_HELPER.read_text() + expected_pythonpath = "${KERNEL_PREFIX}:${FA4_OVERLAY}:${FA4_CUTLASS_PACKAGES}" + + assert f'export PYTHONPATH="{expected_pythonpath}"' in gated_launcher + assert "export H3_V10_KERNEL_GATE=1" in gated_launcher + assert 'if [ "${H3_V10_KERNEL_GATE}" = "1" ]' in sbatch + assert "scripts/train/gate_h3_v10_kernel.sh" in sbatch + assert "python' -m pytest" in sbatch + assert "test_vsa_triton_backward_scale.py" in gate + assert "test_forward_matches_reference[64]" in gate + assert "test_real_sm100a_no_grad_route_receipt" in gate + assert "timeout --signal=TERM --kill-after=30s 300s" in gate + assert "FASTVIDEO_KERNEL_V10_RECEIPT.json" in gate + assert 'if source_commit != execution_commit:' in gate + assert 'observed_wheel_sha256 != receipt.get("wheel_sha256")' in gate + assert 'observed_prefix_tree_sha256 != receipt.get("installed_prefix_tree_sha256")' in gate + assert '"installed_prefix_tree_sha256": installed_prefix_tree_sha256' in rebuild + assert 'UV="${UV:-${KERNEL_ROOT}/tools/uv}"' in rebuild + assert "/home/vlm-wlsaidhi/.local/bin/uv" not in rebuild + assert '"__pycache__" not in relative.parts' in receipt_helper + assert 'path.suffix != ".pyc"' in receipt_helper + assert "907f2100e" in rebuild and "56d4a6074" in rebuild + assert "TORCH_CUDA_ARCH_LIST=10.0a" in rebuild + + +def test_h3_dmd2_v10_kernel_prefix_receipt_hashes_only_stable_installed_files(tmp_path: Path) -> None: + from scripts.train.h3_v10_kernel_receipt import RECEIPT_FILENAME, installed_prefix_tree_sha256 + + prefix = tmp_path / "prefix" + package = prefix / "fastvideo_kernel" + package.mkdir(parents=True) + installed = package / "kernel.so" + installed.write_bytes(b"installed-kernel-v1") + (prefix / "metadata.txt").write_text("metadata-v1", encoding="utf-8") + + receipt = prefix / RECEIPT_FILENAME + receipt.write_text("receipt-v1", encoding="utf-8") + bytecode_dir = package / "__pycache__" + bytecode_dir.mkdir() + bytecode = bytecode_dir / "module.cpython-312.pyc" + bytecode.write_bytes(b"bytecode-v1") + stray_bytecode = package / "generated.pyc" + stray_bytecode.write_bytes(b"stray-v1") + + original = installed_prefix_tree_sha256(prefix) + receipt.write_text("receipt-v2", encoding="utf-8") + bytecode.write_bytes(b"bytecode-v2") + stray_bytecode.write_bytes(b"stray-v2") + assert installed_prefix_tree_sha256(prefix) == original + + installed.write_bytes(b"installed-kernel-v2") + assert installed_prefix_tree_sha256(prefix) != original + + +def test_validation_dmd_sigmas_match_training_noise_amounts() -> None: + """``pipeline_config.dmd_denoising_steps`` replays the trained jump points. + + The H3 denoising stage normalizes the method's integer steps to base time + and lets each scheduler apply its own shift; the resulting clean-times must + match ``1 - shift_noise_amount(base)`` — the exact noising the packed DMD + adapter applies during training rollouts — with one forward per step. + """ + from fastvideo.models.schedulers.scheduling_minimax_h3 import MiniMaxH3Scheduler + + steps = [1000, 667, 333] + base = torch.tensor([step / 1000.0 for step in steps] + [0.0], dtype=torch.float32) + video = MiniMaxH3Scheduler(shift=12.0) + audio = MiniMaxH3Scheduler(shift=3.0) + video.set_timesteps(sigmas=video.shift_sigmas(base)) + audio.set_timesteps(sigmas=audio.shift_sigmas(base)) + + assert video.num_inference_steps == len(steps) + assert audio.num_inference_steps == len(steps) + for index, step in enumerate(steps): + base_step = torch.tensor([step / 1000.0]) + assert video.timesteps[index].item() == pytest.approx(1.0 - shift_noise_amount(base_step, 12.0).item()) + assert audio.timesteps[index].item() == pytest.approx(1.0 - shift_noise_amount(base_step, 3.0).item()) + + +def test_validation_callback_injects_method_denoising_steps() -> None: + """The callback copies the trained step list onto the validation config.""" + from fastvideo.train.callbacks.validation import ValidationCallback + + callback = ValidationCallback.__new__(ValidationCallback) + callback.method = SimpleNamespace(method_config={"dmd_denoising_steps": [1000, 667, 333]}) + + config = SimpleNamespace(dmd_denoising_steps=None) + callback._inject_method_denoising_steps(config) + assert config.dmd_denoising_steps == [1000, 667, 333] + + explicit = SimpleNamespace(dmd_denoising_steps=[1000, 500]) + callback._inject_method_denoising_steps(explicit) + assert explicit.dmd_denoising_steps == [1000, 500] + + callback.method = SimpleNamespace(method_config={"dmd_denoising_steps": [1000, 757], "warp_denoising_step": True}) + warped = SimpleNamespace(dmd_denoising_steps=None) + callback._inject_method_denoising_steps(warped) + assert warped.dmd_denoising_steps is None diff --git a/fastvideo/tests/train/methods/test_minimax_h3_finetune.py b/fastvideo/tests/train/methods/test_minimax_h3_finetune.py index 0d1663e1a2..2ac6612546 100644 --- a/fastvideo/tests/train/methods/test_minimax_h3_finetune.py +++ b/fastvideo/tests/train/methods/test_minimax_h3_finetune.py @@ -47,7 +47,8 @@ class _IdentityJointTransformer: patch_size = (1, 2, 2) def __call__(self, **kwargs): - return kwargs["hidden_states"], kwargs["audio_hidden_states"] + self.autocast_enabled = torch.is_autocast_enabled("cpu") + return kwargs["hidden_states"].float(), kwargs["audio_hidden_states"].float() @pytest.mark.parametrize( @@ -144,6 +145,50 @@ def test_h3_uniform_parameter_dtype_uses_fsdp_dtype() -> None: assert parameter_dtype == torch.bfloat16 +def test_h3_default_parameter_dtype_keeps_compute_boundaries_fp32() -> None: + model = cast( + MiniMaxH3Transformer3DModel, + SimpleNamespace( + config=MiniMaxH3Config(uniform_parameter_dtype=False), + _keep_in_fp32_modules=MiniMaxH3Transformer3DModel._keep_in_fp32_modules, + ), + ) + + assert MiniMaxH3Transformer3DModel._get_parameter_dtype( + model, + "proj_in.weight", + torch.bfloat16, + ) == torch.float32 + assert MiniMaxH3Transformer3DModel._get_parameter_dtype( + model, + "transformer_blocks.0.attn.to_q.weight", + torch.bfloat16, + ) == torch.bfloat16 + + +def test_h3_folded_adaln_keeps_fp32_training_master() -> None: + """Rank-reduced AdaLN must not silently demote an FP32 training load.""" + model = cast( + MiniMaxH3Transformer3DModel, + SimpleNamespace( + config=MiniMaxH3Config(uniform_parameter_dtype=False), + adaln_rank=768, + _keep_in_fp32_modules=MiniMaxH3Transformer3DModel._keep_in_fp32_modules, + ), + ) + + assert MiniMaxH3Transformer3DModel._get_parameter_dtype( + model, + "transformer_blocks.0.adaln_proj.linear.weight", + torch.float32, + ) == torch.float32 + assert MiniMaxH3Transformer3DModel._get_parameter_dtype( + model, + "transformer_blocks.0.adaln_proj.linear.weight", + torch.bfloat16, + ) == torch.bfloat16 + + def test_h3_materializes_rotary_frequencies_on_loader_device() -> None: """Verify that checkpoint loading moves analytic rotary state to the model device.""" model = cast( @@ -209,8 +254,11 @@ def test_h3_plugin_prepares_and_restores_joint_latent_shapes(monkeypatch: pytest ) assert isinstance(prediction, tuple) + assert model.transformer.autocast_enabled is False assert prediction[0].shape == batch.latents.shape assert prediction[1].shape == batch.audio_latents.shape + assert prediction[0].dtype == batch.noisy_model_input.dtype + assert prediction[1].dtype == batch.audio_noisy_model_input.dtype torch.testing.assert_close( prediction[0], -batch.noisy_model_input.permute(0, 2, 1, 3, 4), diff --git a/fastvideo/tests/train/trainer/test_validation.py b/fastvideo/tests/train/trainer/test_validation.py index 5a08a36f8c..a3b1927f43 100644 --- a/fastvideo/tests/train/trainer/test_validation.py +++ b/fastvideo/tests/train/trainer/test_validation.py @@ -84,6 +84,27 @@ def optimizers_zero_grad(self, iteration: int) -> None: self.weight.grad = None +class _RecordingCheckpointManager: + + def __init__(self) -> None: + self.inference_events: list[tuple[int, bool]] = [] + self.training_events: list[int] = [] + self.final_events: list[int] = [] + + def maybe_resume(self, *, resume_from_checkpoint: str) -> None: + assert resume_from_checkpoint == "" + return None + + def maybe_save_inference(self, step: int, *, validation_scheduled: bool) -> None: + self.inference_events.append((step, validation_scheduled)) + + def maybe_save(self, step: int) -> None: + self.training_events.append(step) + + def save_final(self, step: int) -> None: + self.final_events.append(step) + + def test_trainer_runs_validation_callback_during_training(monkeypatch, ) -> None: tracker = _DummyTracker() group = SimpleNamespace(rank=0, local_rank=0, rank_in_group=0, world_size=1) @@ -119,6 +140,7 @@ def test_trainer_runs_validation_callback_during_training(monkeypatch, ) -> None callback_configs=callback_configs, ) method = _DummyMethod() + checkpoint_manager = _RecordingCheckpointManager() trainer.run( method, @@ -126,6 +148,7 @@ def test_trainer_runs_validation_callback_during_training(monkeypatch, ) -> None "sample": "x" }], max_steps=3, + checkpoint_manager=checkpoint_manager, ) validation = trainer.callbacks._callbacks["validation"] @@ -137,3 +160,11 @@ def test_trainer_runs_validation_callback_during_training(monkeypatch, ) -> None assert method.optimizer_steps == [1, 2, 3] assert [step for _, step in tracker.logs] == [1, 2, 3] assert tracker.finished is True + assert checkpoint_manager.inference_events == [ + (0, True), + (1, False), + (2, True), + (3, False), + ] + assert checkpoint_manager.training_events == [1, 2, 3] + assert checkpoint_manager.final_events == [3] diff --git a/fastvideo/tests/train/utils/test_checkpoint.py b/fastvideo/tests/train/utils/test_checkpoint.py index 1bbab529cb..2f8b4a5edc 100644 --- a/fastvideo/tests/train/utils/test_checkpoint.py +++ b/fastvideo/tests/train/utils/test_checkpoint.py @@ -4,29 +4,33 @@ Covers the pure-Python portions of the checkpoint manager: name parsing, resume-path resolution, metadata round-trip, rolling-delete cleanup, the ``_is_stateful`` predicate, and the ``maybe_save`` gating -logic. Code paths that touch DCP (``dcp.save`` / ``dcp.load``) and -CUDA RNG snapshots are intentionally not covered here — those need a -GPU runner and will be tested in later phases. +logic. The inference staging path is covered with mocked DCP I/O; real +distributed collectives and CUDA RNG snapshots require a GPU runner. """ from __future__ import annotations +import json from pathlib import Path +from types import SimpleNamespace from typing import Any import pytest +import torch +import fastvideo.train.utils.checkpoint as checkpoint_module +import fastvideo.train.utils.inference_checkpoint as inference_checkpoint +from fastvideo.train.methods.base import TrainingMethod from fastvideo.train.utils.checkpoint import ( CheckpointConfig, CheckpointManager, + _FullModelState, _find_latest_checkpoint, + _is_complete_training_checkpoint, _is_stateful, _parse_step_from_dir, _resolve_resume_checkpoint, ) -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- def _make_checkpoint_dir( @@ -34,35 +38,78 @@ def _make_checkpoint_dir( step: int, *, with_dcp: bool = True, + with_metadata: bool = True, ) -> Path: - """Create a fake ``checkpoint-/dcp`` directory tree.""" + """Create a fake ``checkpoint-/dcp`` directory tree. + + ``dcp/.metadata`` is dcp.save's completion marker (written last); + ``_find_latest_checkpoint`` requires it, so a complete fake checkpoint + must include it. ``with_metadata=False`` fakes a crashed mid-write save. + """ ckpt_dir = output_dir / f"checkpoint-{step}" ckpt_dir.mkdir(parents=True, exist_ok=True) if with_dcp: (ckpt_dir / "dcp").mkdir(exist_ok=True) + if with_metadata: + (ckpt_dir / "dcp" / ".metadata").touch() return ckpt_dir +def _publish_fake_training_checkpoint( + checkpoint_dir: Path, + *, + step: int, + world_size: int = 1, +) -> None: + (checkpoint_dir / "metadata.json").write_text( + json.dumps({ + "step": step, + "config": { + "training": { + "distributed": { + "num_gpus": world_size, + }, + }, + }, + }), + encoding="utf-8", + ) + for rank in range(world_size): + (checkpoint_dir / f"rng_state_rank{rank}.pt").write_bytes(b"rng") + (checkpoint_dir / ".complete").write_text("complete\n", encoding="utf-8") + + def _make_manager( tmp_path: Path, *, save_steps: int = 0, + save_inference_on_validation: bool = False, keep_last: int = 0, + start_step: int = 0, raw_config: dict[str, Any] | None = None, + require_complete_training_checkpoint: bool = False, ) -> CheckpointManager: """Build a minimal ``CheckpointManager`` for tests that don't touch DCP.""" return CheckpointManager( method=None, dataloader=None, output_dir=str(tmp_path), - config=CheckpointConfig(save_steps=save_steps, keep_last=keep_last), - raw_config=raw_config, + config=CheckpointConfig( + save_steps=save_steps, + keep_last=keep_last, + start_step=start_step, + save_inference_on_validation=save_inference_on_validation, + require_complete_training_checkpoint=require_complete_training_checkpoint, + ), + raw_config=( + raw_config + if raw_config is not None + else ({"training": {"distributed": {"num_gpus": 1}}} + if require_complete_training_checkpoint else None) + ), ) -# --------------------------------------------------------------------------- -# A. _is_stateful predicate -# --------------------------------------------------------------------------- class _Full: @@ -98,9 +145,36 @@ def test_is_stateful_false_when_missing_load_state_dict() -> None: assert _is_stateful(_MissingLoad()) is False -# --------------------------------------------------------------------------- -# B. _parse_step_from_dir -# --------------------------------------------------------------------------- +def test_full_model_state_keeps_frozen_inference_parameters() -> None: + module = torch.nn.Linear(3, 2) + module.requires_grad_(False) + + state = _FullModelState(module).state_dict() + + assert set(state) == {"weight", "bias"} + + +def test_inference_checkpoint_role_is_explicit() -> None: + student = SimpleNamespace( + transformer=torch.nn.Linear(2, 2), + _init_from="base/student", + ) + fake_method = SimpleNamespace(_role_models={"student": student}) + + modules = TrainingMethod.inference_checkpoint_modules(fake_method, "student") + base_path = TrainingMethod.inference_checkpoint_base_model_path(fake_method, "student") + + assert modules == {"transformer": student.transformer} + assert base_path == "base/student" + + +def test_unknown_inference_checkpoint_role_raises() -> None: + fake_method = SimpleNamespace(_role_models={}) + + with pytest.raises(ValueError, match="known roles"): + TrainingMethod.inference_checkpoint_modules(fake_method, "ema") + + def test_parse_step_valid(tmp_path: Path) -> None: @@ -116,9 +190,6 @@ def test_parse_step_invalid_raises(tmp_path: Path) -> None: _parse_step_from_dir(tmp_path / "not-a-checkpoint") -# --------------------------------------------------------------------------- -# C. _find_latest_checkpoint -# --------------------------------------------------------------------------- def test_find_latest_returns_none_on_nonexistent_dir(tmp_path: Path) -> None: @@ -138,8 +209,46 @@ def test_find_latest_returns_largest_step(tmp_path: Path) -> None: assert latest.name == "checkpoint-200" +def test_find_latest_default_preserves_legacy_dcp_completion_contract(tmp_path: Path) -> None: + legacy = _make_checkpoint_dir(tmp_path, 10) + + assert not (legacy / ".complete").exists() + assert _find_latest_checkpoint(tmp_path) == legacy + + +def test_find_latest_strict_skips_unpublished_newer_checkpoint(tmp_path: Path) -> None: + older = _make_checkpoint_dir(tmp_path, 5) + _publish_fake_training_checkpoint(older, step=5, world_size=2) + newer = _make_checkpoint_dir(tmp_path, 10) + _publish_fake_training_checkpoint(newer, step=10, world_size=2) + (newer / ".complete").unlink() + + latest = _find_latest_checkpoint(tmp_path, require_complete_marker=True) + + assert latest == older + + +@pytest.mark.parametrize("defect", ["marker", "metadata_step", "missing_rng", "extra_rng"]) +def test_strict_training_checkpoint_requires_complete_publication( + tmp_path: Path, + defect: str, +) -> None: + checkpoint = _make_checkpoint_dir(tmp_path, 10) + _publish_fake_training_checkpoint(checkpoint, step=10, world_size=2) + if defect == "marker": + (checkpoint / ".complete").write_text("incomplete\n", encoding="utf-8") + elif defect == "metadata_step": + _publish_fake_training_checkpoint(checkpoint, step=9, world_size=2) + elif defect == "missing_rng": + (checkpoint / "rng_state_rank1.pt").unlink() + else: + (checkpoint / "rng_state_rank2.pt").write_bytes(b"stale") + + assert not _is_complete_training_checkpoint(checkpoint, require_complete_marker=True) + assert _find_latest_checkpoint(tmp_path, require_complete_marker=True) is None + + def test_find_latest_skips_dirs_without_dcp_subdir(tmp_path: Path) -> None: - # checkpoint-10 is "corrupted" — has no dcp/ subdir, must be skipped. _make_checkpoint_dir(tmp_path, 10, with_dcp=False) _make_checkpoint_dir(tmp_path, 5, with_dcp=True) latest = _find_latest_checkpoint(tmp_path) @@ -147,6 +256,14 @@ def test_find_latest_skips_dirs_without_dcp_subdir(tmp_path: Path) -> None: assert latest.name == "checkpoint-5" +def test_find_latest_skips_incomplete_dcp_save(tmp_path: Path) -> None: + _make_checkpoint_dir(tmp_path, 10, with_metadata=False) + _make_checkpoint_dir(tmp_path, 5) + latest = _find_latest_checkpoint(tmp_path) + assert latest is not None + assert latest.name == "checkpoint-5" + + def test_find_latest_skips_non_checkpoint_dirs(tmp_path: Path) -> None: (tmp_path / "logs").mkdir() (tmp_path / "wandb").mkdir() @@ -157,9 +274,6 @@ def test_find_latest_skips_non_checkpoint_dirs(tmp_path: Path) -> None: assert latest.name == "checkpoint-7" -# --------------------------------------------------------------------------- -# D. _resolve_resume_checkpoint -# --------------------------------------------------------------------------- def test_resolve_latest_with_no_checkpoints_returns_none(tmp_path: Path) -> None: @@ -176,6 +290,28 @@ def test_resolve_latest_returns_latest_checkpoint(tmp_path: Path) -> None: assert resolved.name == "checkpoint-30" +def test_resolve_explicit_strict_checkpoint_rejects_missing_marker(tmp_path: Path) -> None: + checkpoint = _make_checkpoint_dir(tmp_path, 42) + + with pytest.raises(ValueError, match="incomplete"): + _resolve_resume_checkpoint( + str(checkpoint), + output_dir=str(tmp_path), + require_complete_marker=True, + ) + + +def test_resolve_latest_strict_refuses_fresh_start_over_incomplete_state(tmp_path: Path) -> None: + _make_checkpoint_dir(tmp_path, 42) + + with pytest.raises(ValueError, match="refusing to start from scratch"): + _resolve_resume_checkpoint( + "latest", + output_dir=str(tmp_path), + require_complete_marker=True, + ) + + def test_resolve_explicit_checkpoint_dir(tmp_path: Path) -> None: ckpt = _make_checkpoint_dir(tmp_path, 42) resolved = _resolve_resume_checkpoint(str(ckpt), output_dir=str(tmp_path)) @@ -220,9 +356,6 @@ def test_resolve_unknown_dir_raises(tmp_path: Path) -> None: _resolve_resume_checkpoint(str(bogus), output_dir=str(tmp_path)) -# --------------------------------------------------------------------------- -# E. metadata read/write -# --------------------------------------------------------------------------- def test_write_metadata_roundtrip_with_step(tmp_path: Path) -> None: @@ -254,16 +387,85 @@ def test_write_metadata_includes_raw_config(tmp_path: Path) -> None: assert loaded["config"] == raw +def test_strict_manager_requires_world_size_in_saved_config(tmp_path: Path) -> None: + with pytest.raises(ValueError, match="num_gpus"): + CheckpointManager( + method=None, + dataloader=None, + output_dir=str(tmp_path), + config=CheckpointConfig( + save_steps=1, + keep_last=1, + require_complete_training_checkpoint=True, + ), + raw_config={}, + ) + + +def test_save_publishes_training_complete_after_rng_barrier( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + raw_config = {"training": {"distributed": {"num_gpus": 1}}} + method = SimpleNamespace(checkpoint_state=lambda: {}) + manager = CheckpointManager( + method=method, + dataloader=None, + output_dir=str(tmp_path), + config=CheckpointConfig( + save_steps=1, + keep_last=0, + require_complete_training_checkpoint=True, + ), + raw_config=raw_config, + ) + checkpoint = tmp_path / "checkpoint-10" + checkpoint.mkdir() + marker = checkpoint / ".complete" + marker.write_text("complete\n", encoding="utf-8") + events: list[str] = [] + + def fake_dcp_save(states: dict[str, Any], *, checkpoint_id: str) -> None: + assert states == {} + assert not marker.exists() + events.append("dcp") + dcp_dir = Path(checkpoint_id) + dcp_dir.mkdir(parents=True, exist_ok=True) + (dcp_dir / ".metadata").touch() + + def fake_rng_save(checkpoint_dir: Path) -> None: + assert not marker.exists() + events.append("rng") + (checkpoint_dir / "rng_state_rank0.pt").write_bytes(b"rng") + + def fake_barrier() -> None: + events.append("barrier") + + publish = checkpoint_module._publish_training_checkpoint_complete + + def record_publish(checkpoint_dir: Path) -> None: + assert events[-1] == "barrier" + events.append("publish") + publish(checkpoint_dir) + + monkeypatch.setattr(checkpoint_module.dcp, "save", fake_dcp_save) + monkeypatch.setattr(checkpoint_module, "_barrier", fake_barrier) + monkeypatch.setattr(checkpoint_module, "_publish_training_checkpoint_complete", record_publish) + monkeypatch.setattr(manager, "_save_rng_snapshot", fake_rng_save) + + manager.save(10) + + assert events == ["barrier", "dcp", "barrier", "rng", "barrier", "publish", "barrier"] + assert marker.read_text(encoding="utf-8") == "complete\n" + assert _is_complete_training_checkpoint(checkpoint, require_complete_marker=True) + + def test_load_metadata_raises_on_missing_file(tmp_path: Path) -> None: ckpt_dir = _make_checkpoint_dir(tmp_path, 7) - # No metadata.json written. with pytest.raises(FileNotFoundError, match="metadata"): CheckpointManager.load_metadata(ckpt_dir) -# --------------------------------------------------------------------------- -# F. _cleanup_old_checkpoints (rolling delete) -# --------------------------------------------------------------------------- def test_cleanup_keep_last_zero_is_noop(tmp_path: Path) -> None: @@ -304,9 +506,44 @@ def test_cleanup_skips_non_checkpoint_dirs(tmp_path: Path) -> None: assert remaining == ["checkpoint-3", "logs", "wandb"] -# --------------------------------------------------------------------------- -# G. maybe_save gating logic -# --------------------------------------------------------------------------- +def test_cleanup_never_removes_inference_checkpoints(tmp_path: Path) -> None: + mgr = _make_manager(tmp_path, keep_last=1) + for step in (1, 2, 3): + _make_checkpoint_dir(tmp_path, step) + inference_dir = tmp_path / "inference" / f"checkpoint-{step}" + inference_dir.mkdir(parents=True) + (inference_dir / ".complete").touch() + + mgr._cleanup_old_checkpoints() + + assert sorted(path.name for path in tmp_path.glob("checkpoint-*")) == ["checkpoint-3"] + assert sorted(path.name for path in (tmp_path / "inference").iterdir()) == [ + "checkpoint-1", + "checkpoint-2", + "checkpoint-3", + ] + + +def test_strict_cleanup_does_not_count_incomplete_newer_directories(tmp_path: Path) -> None: + mgr = _make_manager( + tmp_path, + keep_last=2, + require_complete_training_checkpoint=True, + ) + for step in (100, 200, 300): + checkpoint = _make_checkpoint_dir(tmp_path, step) + _publish_fake_training_checkpoint(checkpoint, step=step) + _make_checkpoint_dir(tmp_path, 400) + + mgr._cleanup_old_checkpoints() + + assert sorted(path.name for path in tmp_path.glob("checkpoint-*")) == [ + "checkpoint-200", + "checkpoint-300", + "checkpoint-400", + ] + + def _record_save_calls(mgr: CheckpointManager) -> list[int]: @@ -353,3 +590,146 @@ def test_maybe_save_triggers_on_each_interval(tmp_path: Path) -> None: for step in range(1, 41): mgr.maybe_save(step=step) assert calls == [10, 20, 30, 40] + + +def _record_both_save_calls( + mgr: CheckpointManager, +) -> tuple[list[int], list[int]]: + training_calls: list[int] = [] + inference_calls: list[int] = [] + + def fake_training_save(step: int) -> None: + training_calls.append(step) + mgr._last_saved_step = step + + def fake_inference_save(step: int) -> None: + inference_calls.append(step) + mgr._last_inference_saved_step = step + + mgr.save = fake_training_save # type: ignore[method-assign] + mgr.save_inference = fake_inference_save # type: ignore[method-assign] + return training_calls, inference_calls + + +def test_inference_save_tracks_validation_events_and_ignores_start_gate(tmp_path: Path) -> None: + mgr = _make_manager( + tmp_path, + save_steps=10, + save_inference_on_validation=True, + start_step=12, + ) + training_calls, inference_calls = _record_both_save_calls(mgr) + + mgr.maybe_save_inference(0, validation_scheduled=True) + mgr.maybe_save_inference(4, validation_scheduled=False) + mgr.maybe_save_inference(10, validation_scheduled=True) + mgr.maybe_save_inference(10, validation_scheduled=True) + for step in range(1, 21): + mgr.maybe_save(step) + + assert training_calls == [20] + assert inference_calls == [0, 10] + + +def test_disabled_validation_inference_checkpointing_is_no_op(tmp_path: Path) -> None: + mgr = _make_manager( + tmp_path, + save_inference_on_validation=False, + ) + _, inference_calls = _record_both_save_calls(mgr) + mgr.maybe_save_inference(0, validation_scheduled=True) + + assert inference_calls == [] + + +def test_save_final_dedupes_training_checkpoint(tmp_path: Path) -> None: + mgr = _make_manager( + tmp_path, + save_steps=10, + save_inference_on_validation=True, + ) + training_calls, inference_calls = _record_both_save_calls(mgr) + + mgr.maybe_save_inference(20, validation_scheduled=True) + mgr.maybe_save(20) + mgr.save_final(20) + + assert training_calls == [20] + assert inference_calls == [20] + + +def test_save_final_does_not_create_off_validation_inference_product(tmp_path: Path) -> None: + mgr = _make_manager( + tmp_path, + save_steps=10, + save_inference_on_validation=True, + ) + training_calls, inference_calls = _record_both_save_calls(mgr) + + mgr.save_final(17) + + assert training_calls == [17] + assert inference_calls == [] + + +def test_save_inference_stages_full_state_and_publishes_once( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + module = torch.nn.Linear(2, 2) + module.requires_grad_(False) + base_model = tmp_path / "base" + base_model.mkdir() + + class _Method: + cuda_generator = None + + def inference_checkpoint_modules(self, role: str) -> dict[str, torch.nn.Module]: + assert role == "student" + return {"transformer": module} + + def inference_checkpoint_base_model_path(self, role: str) -> str: + assert role == "student" + return str(base_model) + + manager = CheckpointManager( + method=_Method(), + dataloader=None, + output_dir=str(tmp_path / "run"), + config=CheckpointConfig( + save_steps=0, + keep_last=0, + save_inference_on_validation=True, + ), + ) + staged_states: list[dict[str, Any]] = [] + + def fake_dcp_save(states: dict[str, Any], *, checkpoint_id: str) -> None: + staged_states.append(states) + dcp_dir = Path(checkpoint_id) + dcp_dir.mkdir(parents=True, exist_ok=True) + (dcp_dir / ".metadata").touch() + + def fake_export(**kwargs: Any) -> Path: + target = Path(kwargs["output_dir"]) + target.mkdir(parents=True) + (target / ".complete").touch() + return target + + monkeypatch.setattr("fastvideo.train.utils.checkpoint.dcp.save", fake_dcp_save) + monkeypatch.setattr(inference_checkpoint, "export_inference_checkpoint", fake_export) + monkeypatch.setattr( + inference_checkpoint, + "validate_complete_inference_checkpoint", + lambda path, *, step: path if (path / ".complete").is_file() else None, + ) + + manager.save_inference(10) + manager.save_inference(10) + + assert len(staged_states) == 1 + state = staged_states[0] + assert set(state) == {"roles.student.transformer"} + assert isinstance(state["roles.student.transformer"], _FullModelState) + assert not (tmp_path / "run" / ".inference-staging").exists() + assert (tmp_path / "run" / "inference" / "checkpoint-10" / ".complete").is_file() diff --git a/fastvideo/tests/train/utils/test_config.py b/fastvideo/tests/train/utils/test_config.py index 8471eea8c9..810c962cd2 100644 --- a/fastvideo/tests/train/utils/test_config.py +++ b/fastvideo/tests/train/utils/test_config.py @@ -33,9 +33,6 @@ def _minimal_yaml() -> dict[str, Any]: } -# --------------------------------------------------------------------------- -# Happy path -# --------------------------------------------------------------------------- def test_minimal_yaml_loads_happy_path(tmp_path: Path) -> None: @@ -46,7 +43,6 @@ def test_minimal_yaml_loads_happy_path(tmp_path: Path) -> None: assert cfg.method["_target_"] == ("fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod") assert "student" in cfg.models assert cfg.callbacks == {} - # raw retains the original YAML dict for downstream logging. assert "models" in cfg.raw and "method" in cfg.raw @@ -75,6 +71,10 @@ def test_minimal_yaml_applies_all_defaults(tmp_path: Path) -> None: assert t.loop.gradient_accumulation_steps == 1 assert t.checkpoint.output_dir == "" + assert t.checkpoint.save_inference_checkpoint_on_validation is False + assert t.checkpoint.inference_checkpoint_role == "student" + assert t.checkpoint.inference_checkpoint_dtype == "bfloat16" + assert t.checkpoint.require_complete_training_checkpoint is False assert t.checkpoint.checkpoints_total_limit == 0 assert t.tracker.trackers == [] @@ -84,9 +84,12 @@ 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.model.torch_compile_kwargs == {} assert t.dit_precision == "fp32" assert t.vsa_sparsity == 0.0 + assert t.vsa_tile_size == 256 assert t.pipeline_config is None @@ -126,7 +129,11 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None: }, "checkpoint": { "output_dir": "/out", + "save_inference_checkpoint_on_validation": True, + "inference_checkpoint_role": "student", + "inference_checkpoint_dtype": "float16", "training_state_checkpointing_steps": 50, + "require_complete_training_checkpoint": True, "checkpoints_total_limit": 3, }, "tracker": { @@ -135,13 +142,18 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None: "run_name": "myrun", }, "vsa": { - "sparsity": 0.5 + "sparsity": 0.5, + "tile_size": 64, }, "model": { "weighting_scheme": "logit_normal", "logit_mean": 0.5, "logit_std": 1.5, "precondition_outputs": True, + "enable_torch_compile": True, + "torch_compile_kwargs": { + "dynamic": False, + }, }, "dit_precision": "bf16", } @@ -167,20 +179,24 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None: assert t.loop.gradient_accumulation_steps == 4 assert t.checkpoint.output_dir == "/out" + assert t.checkpoint.save_inference_checkpoint_on_validation is True + assert t.checkpoint.inference_checkpoint_role == "student" + assert t.checkpoint.inference_checkpoint_dtype == "float16" + assert t.checkpoint.require_complete_training_checkpoint is True assert t.checkpoint.checkpoints_total_limit == 3 assert t.tracker.trackers == ["wandb"] assert t.tracker.project_name == "myproj" assert t.vsa_sparsity == pytest.approx(0.5) + assert t.vsa_tile_size == 64 assert t.model.weighting_scheme == "logit_normal" assert t.model.precondition_outputs is True + assert t.model.enable_torch_compile is True + assert t.model.torch_compile_kwargs == {"dynamic": False} assert t.dit_precision == "bf16" -# --------------------------------------------------------------------------- -# Schema validation -# --------------------------------------------------------------------------- def test_missing_models_raises(tmp_path: Path) -> None: @@ -190,6 +206,46 @@ def test_missing_models_raises(tmp_path: Path) -> None: load_run_config(_write_yaml(tmp_path, data)) +def test_invalid_vsa_tile_size_raises(tmp_path: Path) -> None: + data = _minimal_yaml() + data["training"] = {"vsa": {"tile_size": 128}} + with pytest.raises(ValueError, match="training.vsa.tile_size must be 64 or 256"): + load_run_config(_write_yaml(tmp_path, data)) + + +def test_validation_inference_checkpoint_flag_requires_bool(tmp_path: Path) -> None: + data = _minimal_yaml() + data["training"] = {"checkpoint": {"save_inference_checkpoint_on_validation": 1}} + with pytest.raises(ValueError, match="save_inference_checkpoint_on_validation"): + load_run_config(_write_yaml(tmp_path, data)) + + +def test_training_checkpoint_completion_flag_requires_bool(tmp_path: Path) -> None: + data = _minimal_yaml() + data["training"] = {"checkpoint": {"require_complete_training_checkpoint": 1}} + with pytest.raises(ValueError, match="require_complete_training_checkpoint"): + load_run_config(_write_yaml(tmp_path, data)) + + +def test_enabled_inference_checkpoint_requires_role(tmp_path: Path) -> None: + data = _minimal_yaml() + data["training"] = { + "checkpoint": { + "save_inference_checkpoint_on_validation": True, + "inference_checkpoint_role": "", + } + } + with pytest.raises(ValueError, match="inference_checkpoint_role"): + load_run_config(_write_yaml(tmp_path, data)) + + +def test_invalid_inference_checkpoint_dtype_raises(tmp_path: Path) -> None: + data = _minimal_yaml() + data["training"] = {"checkpoint": {"inference_checkpoint_dtype": "fp8"}} + with pytest.raises(ValueError, match="inference_checkpoint_dtype"): + load_run_config(_write_yaml(tmp_path, data)) + + def test_missing_method_raises(tmp_path: Path) -> None: data = _minimal_yaml() del data["method"] @@ -223,9 +279,6 @@ def test_missing_config_file_raises(tmp_path: Path) -> None: load_run_config(str(tmp_path / "does_not_exist.yaml")) -# --------------------------------------------------------------------------- -# Special parsing -# --------------------------------------------------------------------------- @pytest.mark.parametrize("betas_value, expected", [ @@ -273,8 +326,6 @@ def test_pipeline_quant_config_rejects_unknown_name(tmp_path: Path) -> None: def test_data_path_mapping_parses_repeat_counts(tmp_path: Path) -> None: - # Config loading should preserve structured multi-dataset paths so the - # dataset layer can interpret repeat counts later. data = _minimal_yaml() data["training"] = { "data": { @@ -294,8 +345,6 @@ def test_data_path_mapping_parses_repeat_counts(tmp_path: Path) -> None: def test_dotted_override_replaces_mapping_data_path(tmp_path: Path) -> None: - # A dict-valued data_path is a single leaf for overrides: a scalar - # --training.data.data_path replaces the whole mapping. data = _minimal_yaml() data["training"] = { "data": { @@ -342,7 +391,6 @@ def test_dotted_overrides_accept_separate_value_token(tmp_path: Path) -> None: def test_overrides_create_intermediate_keys(tmp_path: Path) -> None: """Overrides into a nested key absent from YAML should still apply.""" data = _minimal_yaml() - # No `training.checkpoint` block in the minimal YAML. path = _write_yaml(tmp_path, data) cfg = load_run_config( path, diff --git a/fastvideo/tests/train/utils/test_inference_checkpoint.py b/fastvideo/tests/train/utils/test_inference_checkpoint.py new file mode 100644 index 0000000000..b4514e5785 --- /dev/null +++ b/fastvideo/tests/train/utils/test_inference_checkpoint.py @@ -0,0 +1,476 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU contracts for bounded-memory modular inference checkpoint export.""" + +from __future__ import annotations + +import json +from pathlib import Path + +import pytest +import torch +import torch.distributed.checkpoint as dcp +from safetensors.torch import load_file, save_file + +import fastvideo.train.utils.inference_checkpoint as inference_checkpoint +from fastvideo.train.utils.inference_checkpoint import ( + InferenceCheckpointExportError, + UnsupportedMergedReverseMappingError, + export_inference_checkpoint, + export_inference_checkpoint_from_dcp, + validate_complete_inference_checkpoint, +) + + +class _LiveModule(torch.nn.Module): + + def __init__(self, reverse_mapping: dict | None = None) -> None: + super().__init__() + self.master = torch.nn.Parameter( + torch.tensor([1.25, -2.5], dtype=torch.float32) + ) + self.reverse_param_names_mapping = reverse_mapping or {} + + +@pytest.fixture +def base_model_dir(tmp_path: Path) -> Path: + base = tmp_path / "base" + transformer = base / "transformer" + transformer.mkdir(parents=True) + (transformer / "config.json").write_text( + json.dumps({"_class_name": "FakeTransformer"}), + encoding="utf-8", + ) + (base / "modular_model_index.json").write_text( + json.dumps({"_class_name": "FakePipeline"}), + encoding="utf-8", + ) + vae = base / "vae" + vae.mkdir() + (vae / "config.json").write_text("{}", encoding="utf-8") + (base / ".cache").mkdir() + return base + + +def _write_base_transformer_weights(base: Path, tensors: dict[str, torch.Tensor]) -> None: + """Write one indexed base shard to define the exact export contract.""" + transformer = base / "transformer" + filename = "diffusion_pytorch_model-00001-of-00001.safetensors" + save_file(tensors, transformer / filename) + index = { + "metadata": { + "total_size": sum(tensor.numel() * tensor.element_size() for tensor in tensors.values()) + }, + "weight_map": {key: filename for key in tensors}, + } + (transformer / "diffusion_pytorch_model.safetensors.index.json").write_text( + json.dumps(index, indent=2, sort_keys=True) + "\n", + encoding="utf-8", + ) + + +def _save_model_only_dcp( + root: Path, + tensors: dict[str, torch.Tensor], + *, + role: str = "student", + module_name: str = "transformer", +) -> Path: + checkpoint = root / "temporary-model-checkpoint" + dcp_dir = checkpoint / "dcp" + state = { + f"roles.{role}.{module_name}.{key}": value.clone() + for key, value in tensors.items() + } + dcp.save(state, checkpoint_id=str(dcp_dir)) + assert (dcp_dir / ".metadata").is_file() + return checkpoint + + +def _load_exported_tensors(transformer_dir: Path) -> dict[str, torch.Tensor]: + tensors: dict[str, torch.Tensor] = {} + for shard in sorted(transformer_dir.glob("*.safetensors")): + tensors.update(load_file(shard)) + return tensors + + +def test_export_casts_maps_shards_and_publishes_complete_layout( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights( + base_model_dir, + { + "disk.proj.weight": torch.empty(2, 2), + "disk.position_ids": torch.empty(2, dtype=torch.int64), + }, + ) + source = _save_model_only_dcp( + tmp_path, + { + "proj.weight": torch.arange(4, dtype=torch.float32).reshape(2, 2), + "block.attn.to_gate_compress.weight": torch.tensor( + [0.25, 0.5, 0.75, 1.0], dtype=torch.float32 + ), + "position_ids": torch.tensor([3, 7], dtype=torch.int64), + }, + ) + module = _LiveModule( + { + "proj.weight": ("disk.proj.weight", None, None), + "position_ids": ("disk.position_ids", None, None), + } + ) + master_before = module.master.detach().clone() + + result = export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=100, + module=module, + base_model_dir=base_model_dir, + dtype="bfloat16", + max_shard_size_bytes=16, + ) + + assert result == (tmp_path / "run" / "inference" / "checkpoint-100").resolve() + assert (result / ".complete").read_text(encoding="utf-8") == "complete\n" + assert (result / "metadata.json").is_file() + assert (result / "modular_model_index.json").is_symlink() + assert (result / "vae").is_symlink() + assert not (result / ".cache").exists() + assert (result / "transformer" / "config.json").is_file() + assert not (result / "transformer" / "config.json").is_symlink() + assert not (result / "transformer" / "base-weights.safetensors").exists() + + exported = _load_exported_tensors(result / "transformer") + assert set(exported) == { + "disk.proj.weight", + "disk.position_ids", + "block.attn.to_gate_compress.weight", + } + assert exported["disk.proj.weight"].dtype == torch.bfloat16 + assert exported["block.attn.to_gate_compress.weight"].dtype == torch.bfloat16 + assert exported["disk.position_ids"].dtype == torch.int64 + torch.testing.assert_close( + exported["disk.proj.weight"], + torch.arange(4, dtype=torch.bfloat16).reshape(2, 2), + ) + assert torch.equal(exported["disk.position_ids"], torch.tensor([3, 7])) + assert module.master.dtype == torch.float32 + assert torch.equal(module.master.detach(), master_before) + + metadata = json.loads((result / "metadata.json").read_text(encoding="utf-8")) + assert metadata["kind"] == "inference" + assert metadata["step"] == 100 + assert metadata["role"] == "student" + assert metadata["module"] == "transformer" + assert metadata["dtype"] == "bfloat16" + assert metadata["tensor_count"] == 3 + assert metadata["total_size"] == 32 + assert metadata["shard_count"] == 3 + assert all(size <= 16 for size in metadata["shard_sizes"]) + + index = json.loads( + ( + result + / "transformer" + / "diffusion_pytorch_model.safetensors.index.json" + ).read_text(encoding="utf-8") + ) + assert index["metadata"]["total_size"] == 32 + assert set(index["weight_map"]) == set(exported) + assert len(set(index["weight_map"].values())) == 3 + + +def test_export_is_atomic_on_failure_and_retry_succeeds( + tmp_path: Path, + base_model_dir: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _write_base_transformer_weights( + base_model_dir, + { + "a": torch.empty(4), + "b": torch.empty(4), + }, + ) + source = _save_model_only_dcp( + tmp_path, + { + "a": torch.ones(4, dtype=torch.float32), + "b": torch.ones(4, dtype=torch.float32), + }, + ) + module = _LiveModule() + real_save_file = inference_checkpoint.save_file + calls = 0 + + def fail_after_first_write(tensors, filename, metadata=None): + nonlocal calls + calls += 1 + real_save_file(tensors, filename, metadata=metadata) + if calls == 1: + raise OSError("injected shard write failure") + + monkeypatch.setattr(inference_checkpoint, "save_file", fail_after_first_write) + kwargs = { + "dcp_checkpoint": source, + "output_dir": tmp_path / "run", + "step": 7, + "module": module, + "base_model_dir": base_model_dir, + "dtype": torch.bfloat16, + "max_shard_size_bytes": 8, + } + with pytest.raises(OSError, match="injected"): + export_inference_checkpoint_from_dcp(**kwargs) + + inference_root = tmp_path / "run" / "inference" + assert not (inference_root / "checkpoint-7").exists() + assert list(inference_root.glob(".checkpoint-7.tmp-*")) == [] + + monkeypatch.setattr(inference_checkpoint, "save_file", real_save_file) + result = export_inference_checkpoint_from_dcp(**kwargs) + assert (result / ".complete").is_file() + + for child in (source / "dcp").iterdir(): + child.unlink() + (source / "dcp").rmdir() + source.rmdir() + assert export_inference_checkpoint_from_dcp(**kwargs) == result + + +def test_export_rejects_merged_reverse_mapping( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights(base_model_dir, {"q.weight": torch.empty(2, 2)}) + source = _save_model_only_dcp( + tmp_path, + {"fused_qkv.weight": torch.ones(6, 2)}, + ) + module = _LiveModule( + {"fused_qkv.weight": ("q.weight", 0, 3)} + ) + + with pytest.raises(UnsupportedMergedReverseMappingError, match="model-specific split"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=3, + module=module, + base_model_dir=base_model_dir, + max_shard_size_bytes=1024, + ) + assert not (tmp_path / "run" / "inference" / "checkpoint-3").exists() + + +def test_export_refuses_incomplete_source_or_destination( + tmp_path: Path, + base_model_dir: Path, +) -> None: + source = tmp_path / "incomplete-source" / "dcp" + source.mkdir(parents=True) + with pytest.raises(FileNotFoundError, match="missing .metadata"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=1, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + + complete_source = _save_model_only_dcp( + tmp_path / "other", + {"weight": torch.ones(1)}, + ) + incomplete_destination = tmp_path / "run" / "inference" / "checkpoint-2" + incomplete_destination.mkdir(parents=True) + (incomplete_destination / "metadata.json").write_text("{}", encoding="utf-8") + with pytest.raises(InferenceCheckpointExportError, match="Refusing to overwrite incomplete"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=complete_source, + output_dir=tmp_path / "run", + step=2, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + + +def test_export_rejects_unknown_unmapped_tensor( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights(base_model_dir, {"weight": torch.empty(2)}) + source = _save_model_only_dcp( + tmp_path, + { + "weight": torch.ones(2), + "surprise.weight": torch.ones(2), + }, + ) + + with pytest.raises(InferenceCheckpointExportError, match="unknown base key"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=4, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + + +def test_export_rejects_missing_base_tensor( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights( + base_model_dir, + { + "weight": torch.empty(2), + "bias": torch.empty(2), + }, + ) + source = _save_model_only_dcp(tmp_path, {"weight": torch.ones(2)}) + + with pytest.raises(InferenceCheckpointExportError, match="missing 1 base transformer tensors"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=5, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + + +def test_export_rejects_base_shape_mismatch( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights(base_model_dir, {"weight": torch.empty(3)}) + source = _save_model_only_dcp(tmp_path, {"weight": torch.ones(2)}) + + with pytest.raises(InferenceCheckpointExportError, match="shape .* != base transformer shape"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=6, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + + +def test_export_rejects_base_index_header_mismatch( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights(base_model_dir, {"weight": torch.empty(2)}) + index_path = base_model_dir / "transformer" / "diffusion_pytorch_model.safetensors.index.json" + index = json.loads(index_path.read_text(encoding="utf-8")) + index["weight_map"]["ghost"] = next(iter(index["weight_map"].values())) + index_path.write_text(json.dumps(index), encoding="utf-8") + source = _save_model_only_dcp(tmp_path, {"weight": torch.ones(2)}) + + with pytest.raises(InferenceCheckpointExportError, match="index/header mismatch"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=7, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + + +def test_completed_checkpoint_validator_rejects_missing_and_corrupt_shards( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights(base_model_dir, {"weight": torch.empty(2)}) + source = _save_model_only_dcp(tmp_path, {"weight": torch.ones(2)}) + + missing = export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "missing-run", + step=8, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + missing_shard = next((missing / "transformer").glob("*.safetensors")) + missing_shard.unlink() + with pytest.raises(InferenceCheckpointExportError, match="shards differ from its index"): + validate_complete_inference_checkpoint(missing, step=8) + + corrupt = export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "corrupt-run", + step=9, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + corrupt_shard = next((corrupt / "transformer").glob("*.safetensors")) + corrupt_shard.write_bytes(b"corrupt") + with pytest.raises(InferenceCheckpointExportError, match="Cannot read inference checkpoint shard"): + validate_complete_inference_checkpoint(corrupt, step=9) + + +def test_checkpoint_manager_adapter_uses_exact_target( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights(base_model_dir, {"weight": torch.empty(2)}) + source = _save_model_only_dcp(tmp_path, {"weight": torch.ones(2)}) + module = _LiveModule() + target = tmp_path / "run" / "inference" / "checkpoint-9" + + result = export_inference_checkpoint( + dcp_dir=source / "dcp", + output_dir=target, + base_model_path=base_model_dir, + role="student", + modules={"transformer": module}, + dtype="bfloat16", + step=9, + raw_config={"not": "serialized"}, + ) + + assert result == target.resolve() + metadata = json.loads((result / "metadata.json").read_text(encoding="utf-8")) + assert metadata["config"] == {"not": "serialized"} + assert "source_dcp" not in metadata + + for child in (source / "dcp").iterdir(): + child.unlink() + (source / "dcp").rmdir() + source.rmdir() + assert export_inference_checkpoint( + dcp_dir=source / "dcp", + output_dir=target, + base_model_path=base_model_dir, + role="student", + modules={"transformer": module}, + dtype="bfloat16", + step=9, + raw_config={"not": "serialized"}, + ) == result + + with pytest.raises(ValueError, match="CheckpointManager output_dir"): + export_inference_checkpoint( + dcp_dir=source / "dcp", + output_dir=tmp_path / "wrong-target", + base_model_path=base_model_dir, + role="student", + modules={"transformer": module}, + dtype="bfloat16", + step=9, + ) + + with pytest.raises(InferenceCheckpointExportError, match="exactly one module"): + export_inference_checkpoint( + dcp_dir=source / "dcp", + output_dir=tmp_path / "run" / "inference" / "checkpoint-10", + base_model_path=base_model_dir, + role="student", + modules={}, + dtype="bfloat16", + step=10, + ) diff --git a/fastvideo/tests/train/utils/test_inference_checkpoint_distributed.py b/fastvideo/tests/train/utils/test_inference_checkpoint_distributed.py new file mode 100644 index 0000000000..088e2ae32c --- /dev/null +++ b/fastvideo/tests/train/utils/test_inference_checkpoint_distributed.py @@ -0,0 +1,286 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Two-GPU FSDP2 gate for validation-triggered inference checkpoints. + +This test intentionally launches a real two-rank NCCL worker and must run on a +compute node. It covers frozen parameters, DCP staging, bounded rank-zero +export, strict reload, RNG preservation, and rank-zero export failure +propagation without leaving a long-running NCCL collective outstanding. +""" + +from __future__ import annotations + +import argparse +import json +import os +import random +import subprocess +import sys +from pathlib import Path +from typing import Any + +import numpy as np +import pytest +import torch +import torch.distributed as dist +from safetensors.torch import load_file, save_file +from torch.distributed import init_device_mesh +from torch.distributed.fsdp import fully_shard + +import fastvideo.train.utils.inference_checkpoint as inference_checkpoint +from fastvideo.train.utils.checkpoint import CheckpointConfig, CheckpointManager + +WORLD_SIZE = 2 + + +class _TinyTransformer(torch.nn.Module): + + def __init__(self, device: torch.device | str = "cpu") -> None: + super().__init__() + self.frozen = torch.nn.Parameter( + torch.tensor([1.25, -2.5, 3.75, -4.0], device=device), + requires_grad=False, + ) + self.trainable = torch.nn.Parameter( + torch.tensor([[0.5, 1.0], [1.5, 2.0]], device=device), + ) + self.reverse_param_names_mapping: dict[str, tuple[str, None, None]] = {} + + +class _Method: + + def __init__(self, transformer: torch.nn.Module, base_model_dir: Path, generator: torch.Generator) -> None: + self.transformer = transformer + self.base_model_dir = base_model_dir + self.cuda_generator = generator + + def inference_checkpoint_modules(self, role: str) -> dict[str, torch.nn.Module]: + if role != "student": + raise ValueError(role) + return {"transformer": self.transformer} + + def inference_checkpoint_base_model_path(self, role: str) -> str: + if role != "student": + raise ValueError(role) + return str(self.base_model_dir) + + +def _write_base_model(base_model_dir: Path) -> None: + transformer_dir = base_model_dir / "transformer" + transformer_dir.mkdir(parents=True) + (transformer_dir / "config.json").write_text( + json.dumps({"_class_name": "TinyTransformer"}), + encoding="utf-8", + ) + filename = "diffusion_pytorch_model-00001-of-00001.safetensors" + tensors = { + "frozen": torch.empty(4), + "trainable": torch.empty(2, 2), + } + save_file(tensors, transformer_dir / filename) + index = { + "metadata": { + "total_size": sum(tensor.numel() * tensor.element_size() for tensor in tensors.values()) + }, + "weight_map": {key: filename for key in tensors}, + } + (transformer_dir / "diffusion_pytorch_model.safetensors.index.json").write_text( + json.dumps(index, indent=2, sort_keys=True) + "\n", + encoding="utf-8", + ) + (base_model_dir / "modular_model_index.json").write_text( + json.dumps({"_class_name": "TinyPipeline"}), + encoding="utf-8", + ) + + +def _load_exported_state(checkpoint_dir: Path) -> dict[str, torch.Tensor]: + state: dict[str, torch.Tensor] = {} + for shard in sorted((checkpoint_dir / "transformer").glob("*.safetensors")): + state.update(load_file(shard)) + return state + + +def _capture_rng(generator: torch.Generator) -> dict[str, Any]: + numpy_state = np.random.get_state() + return { + "torch": torch.get_rng_state().clone(), + "python": random.getstate(), + "numpy": ( + numpy_state[0], + numpy_state[1].copy(), + numpy_state[2], + numpy_state[3], + numpy_state[4], + ), + "cuda": torch.cuda.get_rng_state().clone(), + "generator": generator.get_state().clone(), + } + + +def _rng_equal(before: dict[str, Any], after: dict[str, Any]) -> bool: + before_numpy = before["numpy"] + after_numpy = after["numpy"] + return bool( + torch.equal(before["torch"], after["torch"]) + and before["python"] == after["python"] + and before_numpy[0] == after_numpy[0] + and np.array_equal(before_numpy[1], after_numpy[1]) + and before_numpy[2:] == after_numpy[2:] + and torch.equal(before["cuda"], after["cuda"]) + and torch.equal(before["generator"], after["generator"]) + ) + + +def _run_worker(work_dir: Path, result_path: Path) -> None: + dist.init_process_group("nccl") + rank = dist.get_rank() + local_rank = int(os.environ["LOCAL_RANK"]) + device = torch.device("cuda", local_rank) + torch.cuda.set_device(device) + try: + if dist.get_world_size() != WORLD_SIZE: + raise RuntimeError(f"Expected world size {WORLD_SIZE}, got {dist.get_world_size()}") + + base_model_dir = work_dir / "base" + output_dir = work_dir / "run" + if rank == 0: + _write_base_model(base_model_dir) + dist.barrier() + + torch.manual_seed(1000 + rank) + torch.cuda.manual_seed(2000 + rank) + random.seed(3000 + rank) + np.random.seed(4000 + rank) + generator = torch.Generator(device=device) + generator.manual_seed(5000 + rank) + + transformer = _TinyTransformer(device=device) + mesh = init_device_mesh("cuda", (WORLD_SIZE, )) + fully_shard(transformer, mesh=mesh) + method = _Method(transformer, base_model_dir, generator) + manager = CheckpointManager( + method=method, + dataloader=None, + output_dir=str(output_dir), + config=CheckpointConfig( + save_steps=0, + keep_last=0, + save_inference_on_validation=True, + inference_role="student", + inference_dtype="float32", + ), + ) + + success_rng_before = _capture_rng(generator) + manager.save_inference(1) + success_rng_after = _capture_rng(generator) + + checkpoint_dir = output_dir / "inference" / "checkpoint-1" + exported = _load_exported_state(checkpoint_dir) + reloaded = _TinyTransformer() + incompatible = reloaded.load_state_dict(exported, strict=True) + expected = _TinyTransformer().state_dict() + reload_equal = (not incompatible.missing_keys and not incompatible.unexpected_keys + and all(torch.equal(reloaded.state_dict()[key], value) for key, value in expected.items())) + frozen_present = "frozen" in exported and torch.equal(exported["frozen"], expected["frozen"]) + + real_export = inference_checkpoint.export_inference_checkpoint + if rank == 0: + + def _injected_failure(**_: Any) -> Path: + raise RuntimeError("injected rank-zero export failure") + + inference_checkpoint.export_inference_checkpoint = _injected_failure + + failure_rng_before = _capture_rng(generator) + failure: str | None = None + try: + manager.save_inference(2) + except RuntimeError as error: + failure = str(error) + finally: + if rank == 0: + inference_checkpoint.export_inference_checkpoint = real_export + failure_rng_after = _capture_rng(generator) + + local_result = { + "rank": rank, + "success_rng_equal": _rng_equal(success_rng_before, success_rng_after), + "failure_rng_equal": _rng_equal(failure_rng_before, failure_rng_after), + "reload_equal": reload_equal, + "frozen_present": frozen_present, + "failure": failure, + } + gathered: list[dict[str, Any] | None] = [None] * WORLD_SIZE + dist.all_gather_object(gathered, local_result) + if rank == 0: + result_path.write_text(json.dumps(gathered, indent=2, sort_keys=True) + "\n", encoding="utf-8") + dist.barrier() + finally: + dist.destroy_process_group() + + +def test_two_gpu_fsdp2_inference_checkpoint_gate(tmp_path: Path) -> None: + if not torch.cuda.is_available(): + pytest.skip("This test requires CUDA and must run on a compute node.") + if torch.cuda.device_count() < WORLD_SIZE: + pytest.skip(f"This test requires at least {WORLD_SIZE} CUDA devices.") + + result_path = tmp_path / "results.json" + command = [ + sys.executable, + "-m", + "torch.distributed.run", + "--standalone", + "--nproc_per_node", + str(WORLD_SIZE), + str(Path(__file__).resolve()), + "--worker", + "--work-dir", + str(tmp_path / "worker"), + "--result", + str(result_path), + ] + environment = os.environ.copy() + environment.setdefault("TORCHDYNAMO_DISABLE", "1") + try: + process = subprocess.run( + command, + capture_output=True, + text=True, + env=environment, + timeout=180, + ) + except subprocess.TimeoutExpired as error: + raise RuntimeError( + "Two-GPU inference checkpoint worker timed out after 180 seconds\n" + f"STDOUT:\n{error.stdout}\nSTDERR:\n{error.stderr}") from error + if process.returncode != 0: + raise RuntimeError( + f"Two-GPU inference checkpoint worker failed with code {process.returncode}\n" + f"STDOUT:\n{process.stdout}\nSTDERR:\n{process.stderr}") + + results = json.loads(result_path.read_text(encoding="utf-8")) + assert len(results) == WORLD_SIZE + assert all(result["success_rng_equal"] for result in results) + assert all(result["failure_rng_equal"] for result in results) + assert all(result["reload_equal"] for result in results) + assert all(result["frozen_present"] for result in results) + errors = [result["failure"] for result in results] + assert all(error is not None and "injected rank-zero export failure" in error for error in errors) + assert len(set(errors)) == 1 + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--worker", action="store_true") + parser.add_argument("--work-dir", type=Path) + parser.add_argument("--result", type=Path) + return parser.parse_args() + + +if __name__ == "__main__": + args = _parse_args() + if not args.worker or args.work_dir is None or args.result is None: + raise SystemExit("--worker, --work-dir, and --result are required") + _run_worker(args.work_dir, args.result) diff --git a/fastvideo/tests/train/utils/test_moduleloader_attention_backend.py b/fastvideo/tests/train/utils/test_moduleloader_attention_backend.py index a7e957c550..580049cedc 100644 --- a/fastvideo/tests/train/utils/test_moduleloader_attention_backend.py +++ b/fastvideo/tests/train/utils/test_moduleloader_attention_backend.py @@ -3,15 +3,22 @@ from __future__ import annotations +from types import SimpleNamespace + import torch import pytest -from fastvideo.attention.selector import _active_component_attention_backend_scope +from fastvideo.attention.selector import ( + _active_component_attention_backend_scope, + _component_attention_backend_scope, +) from fastvideo.configs.pipelines.base import PipelineConfig +from fastvideo.models.loader import component_loader from fastvideo.platforms import AttentionBackendEnum from fastvideo.train.utils import moduleloader from fastvideo.train.utils.training_config import ( DistributedConfig, + ModelTrainingConfig, TrainingConfig, ) @@ -37,7 +44,10 @@ def _fake_load_module(**kwargs): del kwargs scope = _active_component_attention_backend_scope() captured.append((scope.backend, scope.component) if scope else (None, None)) - return torch.nn.Linear(1, 1) + module = torch.nn.Linear(1, 1) + module.config = SimpleNamespace( # type: ignore[attr-defined] + _resolved_attention_backend=scope.backend if scope else None, ) + return module monkeypatch.setattr( moduleloader.PipelineComponentLoader, @@ -57,6 +67,107 @@ def _fake_load_module(**kwargs): assert _active_component_attention_backend_scope() is None +@pytest.mark.parametrize("role", ["student", "teacher", "critic"]) +def test_role_attention_backend_receipt_matches_request( + monkeypatch, + tmp_path, + role: str, +) -> None: + """Every DMD role returns a construction receipt for its explicit backend.""" + training_config = TrainingConfig( + distributed=DistributedConfig(hsdp_shard_dim=1), + pipeline_config=PipelineConfig(), + ) + requested = (AttentionBackendEnum.ATTN_QAT_TRAIN + if role == "student" else AttentionBackendEnum.FLASH_ATTN) + + monkeypatch.setattr(moduleloader, "maybe_download_model", lambda path: str(tmp_path)) + monkeypatch.setattr( + moduleloader, + "verify_model_config_and_directory", + lambda path: {"transformer": ("diffusers", "FakeTransformer")}, + ) + + def _fake_load_module(**kwargs): + del kwargs + scope = _active_component_attention_backend_scope() + assert scope is not None + module = torch.nn.Linear(1, 1) + module.config = SimpleNamespace( # type: ignore[attr-defined] + _resolved_attention_backend=scope.backend, ) + return module + + monkeypatch.setattr(moduleloader.PipelineComponentLoader, "load_module", _fake_load_module) + + moduleloader.load_module_from_path( + model_path=f"fake/{role}", + module_type="transformer", + training_config=training_config, + disable_custom_init_weights=(role != "student"), + attention_backend=requested, + ) + + +def test_role_attention_backend_receipt_mismatch_fails(monkeypatch, tmp_path) -> None: + """A request that gets narrowed or lost may not silently start training.""" + training_config = TrainingConfig( + distributed=DistributedConfig(hsdp_shard_dim=1), + pipeline_config=PipelineConfig(), + ) + monkeypatch.setattr(moduleloader, "maybe_download_model", lambda path: str(tmp_path)) + monkeypatch.setattr( + moduleloader, + "verify_model_config_and_directory", + lambda path: {"transformer": ("diffusers", "FakeTransformer")}, + ) + + def _fake_load_module(**kwargs): + del kwargs + module = torch.nn.Linear(1, 1) + module.config = SimpleNamespace( # type: ignore[attr-defined] + _resolved_attention_backend=AttentionBackendEnum.TORCH_SDPA, ) + return module + + monkeypatch.setattr(moduleloader.PipelineComponentLoader, "load_module", _fake_load_module) + + with pytest.raises(RuntimeError, match="requested attention backend FLASH_ATTN.*recorded TORCH_SDPA"): + moduleloader.load_module_from_path( + model_path="fake/teacher", + module_type="transformer", + training_config=training_config, + disable_custom_init_weights=True, + attention_backend="FLASH_ATTN", + ) + + +def test_teacher_critic_preserves_explicit_flash_attention_scope() -> None: + """disable_custom_init_weights must not erase a role-local dense backend.""" + args = SimpleNamespace( + _loading_teacher_critic_model=True, + attention_backend=None, + ) + + with _component_attention_backend_scope(AttentionBackendEnum.FLASH_ATTN, component="teacher"): + with component_loader._teacher_critic_attention_context(args): + scope = _active_component_attention_backend_scope() + assert scope is not None + assert scope.backend is AttentionBackendEnum.FLASH_ATTN + + +def test_teacher_critic_masks_generator_only_qat_attention_scope() -> None: + """The historical student-only QAT policy still narrows teacher/critic.""" + args = SimpleNamespace( + _loading_teacher_critic_model=True, + attention_backend=None, + ) + + with _component_attention_backend_scope(AttentionBackendEnum.ATTN_QAT_TRAIN, component="teacher"): + with component_loader._teacher_critic_attention_context(args): + scope = _active_component_attention_backend_scope() + assert scope is not None + assert scope.backend is None + + def test_load_transformer_restores_backend_when_loading_fails( monkeypatch, tmp_path, @@ -92,3 +203,75 @@ def _raise_during_load(**kwargs): attention_backend="ATTN_QAT_TRAIN", ) assert _active_component_attention_backend_scope() is None + + +def test_training_args_propagate_compile_settings() -> None: + training_config = TrainingConfig( + distributed=DistributedConfig(hsdp_shard_dim=1), + model=ModelTrainingConfig( + enable_torch_compile=True, + torch_compile_kwargs={"dynamic": False}, + ), + pipeline_config=PipelineConfig(), + ) + + args = moduleloader._make_training_args(training_config, model_path="fake/model") + + assert args.enable_torch_compile is True + assert args.torch_compile_kwargs == {"dynamic": False} + + +def test_load_transformer_forwards_pre_fsdp_transform(monkeypatch, tmp_path) -> None: + training_config = TrainingConfig( + distributed=DistributedConfig(hsdp_shard_dim=1), + pipeline_config=PipelineConfig(), + ) + + def transform(module): + return module + + captured = None + + monkeypatch.setattr(moduleloader, "maybe_download_model", lambda path: str(tmp_path)) + monkeypatch.setattr( + moduleloader, + "verify_model_config_and_directory", + lambda path: {"transformer": ("diffusers", "FakeTransformer")}, + ) + + def _fake_load_module(**kwargs): + nonlocal captured + captured = getattr(kwargs["fastvideo_args"], "_pre_fsdp_transform", None) + return torch.nn.Linear(1, 1) + + monkeypatch.setattr(moduleloader.PipelineComponentLoader, "load_module", _fake_load_module) + + moduleloader.load_module_from_path( + model_path="fake/model", + module_type="transformer", + training_config=training_config, + pre_fsdp_transform=transform, + ) + + assert captured is transform + + +def test_pre_fsdp_transform_rejects_non_transformer(monkeypatch, tmp_path) -> None: + training_config = TrainingConfig( + distributed=DistributedConfig(hsdp_shard_dim=1), + pipeline_config=PipelineConfig(), + ) + monkeypatch.setattr(moduleloader, "maybe_download_model", lambda path: str(tmp_path)) + monkeypatch.setattr( + moduleloader, + "verify_model_config_and_directory", + lambda path: {"vae": ("diffusers", "FakeVAE")}, + ) + + with pytest.raises(ValueError, match="only be set when loading a transformer"): + moduleloader.load_module_from_path( + model_path="fake/model", + module_type="vae", + training_config=training_config, + pre_fsdp_transform=lambda module: module, + ) diff --git a/fastvideo/tests/train/utils/test_torch_compile.py b/fastvideo/tests/train/utils/test_torch_compile.py new file mode 100644 index 0000000000..fa0bd66ade --- /dev/null +++ b/fastvideo/tests/train/utils/test_torch_compile.py @@ -0,0 +1,241 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Regional training compile policy regression tests.""" + +from __future__ import annotations + +import pytest +import torch + +from fastvideo.models.loader.fsdp_load import _compile_model_regions +from fastvideo.train.utils.activation_checkpoint import apply_activation_checkpointing + + +class _RepeatedModel(torch.nn.Module): + _compile_conditions = [ + lambda name, module: name.startswith("blocks.") and name.count(".") == 1 + ] + + def __init__(self) -> None: + super().__init__() + self.blocks = torch.nn.ModuleList([ + torch.nn.Linear(4, 4), + torch.nn.Linear(4, 4), + ]) + + def forward(self, value: torch.Tensor) -> torch.Tensor: + for block in self.blocks: + value = block(value) + return value + + +def test_regional_compile_preserves_checkpoint_state_dict(monkeypatch) -> None: + model = apply_activation_checkpointing(_RepeatedModel()) + state_dict_keys = list(model.state_dict()) + calls: list[tuple[torch.nn.Module, dict]] = [] + + def _fake_compile(forward, **kwargs): + calls.append((forward.__self__, kwargs)) + return forward + + monkeypatch.setattr(torch, "compile", _fake_compile) + + assert _compile_model_regions(model, {}) == 2 + assert [target for target, _ in calls] == [ + block._checkpoint_wrapped_module for block in model.blocks + ] + assert [kwargs for _, kwargs in calls] == [ + {"fullgraph": True, "options": {"emulate_precision_casts": True}}, + {"fullgraph": True, "options": {"emulate_precision_casts": True}}, + ] + assert list(model.state_dict()) == state_dict_keys + + +def test_regional_compile_dispatches_grad_and_no_grad_calls(monkeypatch) -> None: + model = _RepeatedModel() + compiled_calls = 0 + + def _fake_compile(eager_forward, **kwargs): + del kwargs + + def _compiled(*args, **forward_kwargs): + nonlocal compiled_calls + compiled_calls += 1 + return eager_forward(*args, **forward_kwargs) + + return _compiled + + monkeypatch.setattr(torch, "compile", _fake_compile) + _compile_model_regions(model, {}) + + value = torch.randn(2, 4) + model(value) + assert compiled_calls == 2 + + with torch.no_grad(): + model(value) + assert compiled_calls == 4 + + +def test_regional_compile_forwards_supported_kwargs(monkeypatch) -> None: + model = _RepeatedModel() + calls: list[dict] = [] + global_recompile_limit = torch._dynamo.config.recompile_limit + + def _fake_compile(forward, **kwargs): + calls.append(kwargs) + return forward + + monkeypatch.setattr(torch, "compile", _fake_compile) + + _compile_model_regions(model, { + "dynamic": False, + "recompile_limit": 32, + "options": { + "emulate_precision_casts": False, + "max_autotune": True, + }, + }) + + assert calls == [ + { + "fullgraph": True, + "dynamic": False, + "recompile_limit": 32, + "options": { + "emulate_precision_casts": False, + "max_autotune": True, + }, + }, + { + "fullgraph": True, + "dynamic": False, + "recompile_limit": 32, + "options": { + "emulate_precision_casts": False, + "max_autotune": True, + }, + }, + ] + assert torch._dynamo.config.recompile_limit == global_recompile_limit + + +def test_regional_compile_rejects_partial_graph_mode() -> None: + with pytest.raises(ValueError, match="fullgraph=True"): + _compile_model_regions(_RepeatedModel(), {"fullgraph": False}) + + +def test_regional_compile_rejects_mode_kwarg() -> None: + """`mode` conflicts with the always-injected inductor options. + + torch.compile forbids mode+options together; the loader must fail with an + actionable message rather than letting torch blame an `options` key the + user never wrote (the CLI help's own example uses `mode`). + """ + model = _RepeatedModel() + with pytest.raises(ValueError, match="mode"): + _compile_model_regions(model, {"mode": "reduce-overhead"}) + + +def test_checkpoint_wrapper_prefix_normalization() -> None: + """AC-wrapped blocks must not break name-keyed weight-loader lookups. + + checkpoint_wrapper strips its prefix from state_dict() keys via hooks but + NOT from named_parameters()/named_buffers(); checkpoint keys are clean. + Pre-fix, a loaded buffer inside a wrapped block missed the named_buffers + membership test and was silently converted into a trainable nn.Parameter + by load_state_dict(assign=True). + """ + from fastvideo.models.loader.fsdp_load import _strip_checkpoint_wrapper_prefix + + class _BufferBlock(torch.nn.Module): + + def __init__(self) -> None: + super().__init__() + self.lin = torch.nn.Linear(4, 4) + self.register_buffer("freq", torch.arange(4.0)) + + def forward(self, value: torch.Tensor) -> torch.Tensor: + return self.lin(value) + self.freq + + class _BufferModel(torch.nn.Module): + + def __init__(self) -> None: + super().__init__() + self.blocks = torch.nn.ModuleList([_BufferBlock()]) + + from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import ( + checkpoint_wrapper, ) + + model = _BufferModel() + model.blocks[0] = checkpoint_wrapper(model.blocks[0]) + + assert "blocks.0.freq" in model.state_dict() + raw_buffer_names = {name for name, _ in model.named_buffers()} + assert "blocks.0.freq" not in raw_buffer_names + assert "blocks.0._checkpoint_wrapped_module.freq" in raw_buffer_names + + clean_buffers = {_strip_checkpoint_wrapper_prefix(name) for name, _ in model.named_buffers()} + clean_params = {_strip_checkpoint_wrapper_prefix(name) for name, _ in model.named_parameters()} + assert clean_buffers == {"blocks.0.freq"} + assert clean_params == {"blocks.0.lin.weight", "blocks.0.lin.bias"} + + loaded = { + "blocks.0.freq": torch.ones(4), + "blocks.0.lin.weight": torch.ones(4, 4), + "blocks.0.lin.bias": torch.ones(4), + } + sharded_sd = { + key: (value if key in clean_buffers else torch.nn.Parameter(value)) + for key, value in loaded.items() + } + model.load_state_dict(sharded_sd, assign=True) + assert any("freq" in name for name, _ in model.named_buffers()) + assert not any("freq" in name for name, _ in model.named_parameters()) + + +def test_regional_compile_unsupported_when_attention_compile_disabled(monkeypatch) -> None: + """The attention-eager escape hatch must degrade the role, not crash. + + FASTVIDEO_DISABLE_ATTENTION_COMPILE=1 wraps attention forwards in + torch.compiler.disable; under fullgraph regional compile that raises + `torch._dynamo.exc.Unsupported: Skip inlining torch.compiler.disable()'d + function` at the first training step (observed on the h3-compile-ab + b4_on_attneager leg, job 2610). The guard must reject regional compile + for the role so it falls back to eager with the standard warning. + """ + from fastvideo.models.loader.fsdp_load import _regional_compile_unsupported_reason + + monkeypatch.setenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", "1") + reason = _regional_compile_unsupported_reason({"config": None}) + assert reason is not None + assert "FASTVIDEO_DISABLE_ATTENTION_COMPILE" in reason + + monkeypatch.setenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", "0") + assert _regional_compile_unsupported_reason({"config": None}) is None + + +def test_regional_compile_unsupported_for_vsa_backends() -> None: + """A VSA-backed role must fall back to eager instead of hard-failing. + + The VSA backends (Triton block-sparse kernels behind sequence-parallel + all-to-alls plus a host-synced metadata guard) are not fullgraph-traceable. + enable_torch_compile is a per-run switch shared by every DMD role, so the + loader must skip the VSA student while dense FLASH_ATTN/SDPA roles compile. + """ + from fastvideo.models.loader.fsdp_load import _regional_compile_unsupported_reason + from fastvideo.platforms import AttentionBackendEnum + + class _Config: + pass + + for backend in (AttentionBackendEnum.VIDEO_SPARSE_ATTN, AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3): + config = _Config() + config._resolved_attention_backend = backend + reason = _regional_compile_unsupported_reason({"config": config}) + assert reason is not None + assert backend.name in reason + + sdpa_config = _Config() + sdpa_config._resolved_attention_backend = AttentionBackendEnum.TORCH_SDPA + assert _regional_compile_unsupported_reason({"config": sdpa_config}) is None + assert _regional_compile_unsupported_reason({"config": None}) is None diff --git a/fastvideo/tests/train/utils/test_tracking.py b/fastvideo/tests/train/utils/test_tracking.py new file mode 100644 index 0000000000..06e19912d5 --- /dev/null +++ b/fastvideo/tests/train/utils/test_tracking.py @@ -0,0 +1,135 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU tests for the modular trainer's wandb tracking guard. + +The Trainer logs each step's loss dict through ``build_tracker``'s tracker on +rank 0; these tests pin the graceful no-op path (wandb missing / no +credentials) and the logged loss-dict shape with a monkeypatched wandb. +""" + +import builtins +import sys +from types import SimpleNamespace + +from fastvideo.train.utils import tracking +from fastvideo.train.utils.training_config import ( + CheckpointConfig, + TrackerConfig, +) +from fastvideo.training.trackers import DummyTracker, WandbTracker + + +class _FakeRun: + + def __init__(self, **init_kwargs): + self.init_kwargs = init_kwargs + self.logged = [] + + def log(self, metrics, step=None): + self.logged.append((metrics, step)) + + def finish(self): + pass + + +def _fake_wandb(): + runs = [] + + def _init(**kwargs): + run = _FakeRun(**kwargs) + runs.append(run) + return run + + return SimpleNamespace(init=_init, runs=runs, api=SimpleNamespace(api_key="key")) + + +def test_wandb_usable_false_when_not_importable(monkeypatch) -> None: + real_import = builtins.__import__ + + def _blocked(name, *args, **kwargs): + if name == "wandb": + raise ImportError("no wandb") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", _blocked) + monkeypatch.delitem(sys.modules, "wandb", raising=False) + assert tracking._wandb_usable() is False + + +def test_wandb_usable_false_without_credentials(monkeypatch) -> None: + monkeypatch.delenv("WANDB_API_KEY", raising=False) + monkeypatch.delenv("WANDB_MODE", raising=False) + monkeypatch.setitem( + sys.modules, + "wandb", + SimpleNamespace(api=SimpleNamespace(api_key=None)), + ) + assert tracking._wandb_usable() is False + + +def test_wandb_usable_with_env_key(monkeypatch) -> None: + monkeypatch.setenv("WANDB_API_KEY", "key") + monkeypatch.setitem(sys.modules, "wandb", SimpleNamespace()) + assert tracking._wandb_usable() is True + + +def test_build_tracker_noops_cleanly_without_wandb(monkeypatch, tmp_path) -> None: + """A wandb-project config must degrade to the dummy tracker, not crash.""" + monkeypatch.setattr(tracking, "get_world_group", lambda: SimpleNamespace(rank=0)) + monkeypatch.setattr(tracking, "_wandb_usable", lambda: False) + + tracker = tracking.build_tracker( + TrackerConfig(project_name="h3-dmd2-vsa"), + CheckpointConfig(output_dir=str(tmp_path)), + config={"method": {}}, + ) + + assert isinstance(tracker, DummyTracker) + tracker.log({"total_loss": 1.0}, 1) # must not raise + tracker.finish() + + +def test_build_tracker_logs_loss_dict_with_monkeypatched_wandb(monkeypatch, tmp_path) -> None: + fake = _fake_wandb() + monkeypatch.setitem(sys.modules, "wandb", fake) + monkeypatch.setenv("WANDB_API_KEY", "key") + monkeypatch.setattr(tracking, "get_world_group", lambda: SimpleNamespace(rank=0)) + + run_config = {"method": {"dmd_denoising_steps": [1000, 757, 522]}} + tracker = tracking.build_tracker( + TrackerConfig(project_name="h3-dmd2-vsa", run_name="dmd2_vsa0_overfit"), + CheckpointConfig(output_dir=str(tmp_path)), + config=run_config, + ) + + assert isinstance(tracker, WandbTracker) + (run, ) = fake.runs + assert run.init_kwargs["project"] == "h3-dmd2-vsa" + assert run.init_kwargs["name"] == "dmd2_vsa0_overfit" + assert run.init_kwargs["config"] == run_config + + metrics = { + "total_loss": 0.5, + "generator_loss": 0.25, + "fake_score_loss": 0.25, + "update_student": 1.0, + "step_time_sec": 0.1, + "vsa_sparsity": 0.0, + } + tracker.log(metrics, 7) + assert run.logged == [(metrics, 7)] + + +def test_build_tracker_nonzero_rank_never_inits_wandb(monkeypatch, tmp_path) -> None: + fake = _fake_wandb() + monkeypatch.setitem(sys.modules, "wandb", fake) + monkeypatch.setenv("WANDB_API_KEY", "key") + monkeypatch.setattr(tracking, "get_world_group", lambda: SimpleNamespace(rank=1)) + + tracker = tracking.build_tracker( + TrackerConfig(project_name="h3-dmd2-vsa"), + CheckpointConfig(output_dir=str(tmp_path)), + config=None, + ) + + assert isinstance(tracker, DummyTracker) + assert fake.runs == [] diff --git a/fastvideo/train/attn_qat/README.md b/fastvideo/train/attn_qat/README.md index 55edd75a1f..343aae1b88 100644 --- a/fastvideo/train/attn_qat/README.md +++ b/fastvideo/train/attn_qat/README.md @@ -71,7 +71,7 @@ The migration preserves these training semantics: |---|---| | Student fake-quantized attention | `models.student.attention_backend: ATTN_QAT_TRAIN` | | Teacher/critic full-precision attention | Role-local `FLASH_ATTN` | -| Generator update every five critic steps | `method.generator_update_interval: 5` | +| Four critic-only steps, then one student-only step | `method.generator_update_interval: 5` | | Three-step rollout | `method.dmd_denoising_steps: [1000, 757, 522]` | | Score timestep range | `method.min_timestep_ratio: 0.02`, `max_timestep_ratio: 0.98` | | Legacy teacher guidance `cond + 2(cond-uncond)` | Standard CFG scale `3.0` | diff --git a/fastvideo/train/callbacks/callback.py b/fastvideo/train/callbacks/callback.py index ca057c151b..11e68fa5e7 100644 --- a/fastvideo/train/callbacks/callback.py +++ b/fastvideo/train/callbacks/callback.py @@ -1,8 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -"""Callback base class and CallbackDict manager. - -Adapted from FastGen's callback pattern to FastVideo's types. -""" +"""Callback base class and CallbackDict manager.""" from __future__ import annotations @@ -19,11 +16,11 @@ logger = init_logger(__name__) -# Well-known callback names that don't need ``_target_`` in YAML. _BUILTIN_CALLBACKS: dict[str, str] = { "grad_clip": "fastvideo.train.callbacks.grad_clip.GradNormClipCallback", "validation": "fastvideo.train.callbacks.validation.ValidationCallback", "ema": "fastvideo.train.callbacks.ema.EMACallback", + "latent_vis": "fastvideo.train.callbacks.latent_vis.LatentVisCallback", } @@ -38,10 +35,6 @@ class Callback: training_config: TrainingConfig method: TrainingMethod _callback_dict: CallbackDict | None - # Yaml dict key under which this callback was declared (e.g. - # "validation_short"). Set by ``CallbackDict`` after instantiation. - # Useful for callbacks that want to disambiguate themselves from - # sibling instances in tracker keys, log paths, etc. name: str = "" def on_train_start( @@ -73,6 +66,11 @@ def on_validation_begin( ) -> None: pass + def will_run_validation(self, iteration: int = 0) -> bool: + """Return whether this callback will validate at ``iteration``.""" + del iteration + return False + def on_validation_end( self, method: TrainingMethod, @@ -179,3 +177,7 @@ def _dispatch(*args: Any, **kwargs: Any) -> None: fn(*args, **kwargs) return _dispatch + + def will_run_validation(self, iteration: int = 0) -> bool: + """Return whether any configured callback schedules validation now.""" + return any(cb.will_run_validation(iteration) for cb in self._callbacks.values()) diff --git a/fastvideo/train/callbacks/ema.py b/fastvideo/train/callbacks/ema.py index 2bd9f01f09..1f6f4b1645 100644 --- a/fastvideo/train/callbacks/ema.py +++ b/fastvideo/train/callbacks/ema.py @@ -53,9 +53,6 @@ def __init__( self._ema_started = False self.student_ema: EMA_FSDP | None = None - # ---------------------------------------------------------- - # Hooks - # ---------------------------------------------------------- def on_train_start( self, @@ -92,6 +89,9 @@ def on_training_step_end( if self.student_ema is None: return + student_optimizer = getattr(method, "_student_optimizer", None) + if student_optimizer is not None and student_optimizer not in method.get_optimizers(iteration): + return if iteration < self._start_iter: return if not self._ema_started: @@ -113,9 +113,6 @@ def on_training_step_end( iteration, ) - # ---------------------------------------------------------- - # EMA context manager - # ---------------------------------------------------------- @contextlib.contextmanager def ema_context( @@ -132,9 +129,6 @@ def ema_context( else: yield transformer - # ---------------------------------------------------------- - # Checkpoint state - # ---------------------------------------------------------- def state_dict(self) -> dict[str, Any]: if self.student_ema is None: diff --git a/fastvideo/train/callbacks/latent_vis.py b/fastvideo/train/callbacks/latent_vis.py new file mode 100644 index 0000000000..4d72c00a80 --- /dev/null +++ b/fastvideo/train/callbacks/latent_vis.py @@ -0,0 +1,93 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Intermediate-latent visualization callback for distillation methods. + +Port of the legacy ``fastvideo/training/distillation_pipeline.py`` latent +logging to the modular trainer: every ``every_steps`` iterations, rank 0 +decodes the method's latest latent snapshots (``method.latent_vis`` — the +student's rollout prediction plus the real- and fake-score predictions on +generator-update steps) through the model's ``decode_vis_latents`` hook and +logs them to the tracker as videos. + +Both hooks are optional: methods that never populate ``latent_vis`` or models +without ``decode_vis_latents`` turn this callback into a no-op. +""" + +from __future__ import annotations + +from typing import Any, TYPE_CHECKING + +import torch + +from fastvideo.distributed import get_world_group +from fastvideo.logger import init_logger +from fastvideo.train.callbacks.callback import Callback +from fastvideo.training.trackers import DummyTracker + +if TYPE_CHECKING: + from fastvideo.train.methods.base import TrainingMethod + +logger = init_logger(__name__) + +_DEFAULT_KEYS = ( + "generator_pred_video", + "real_score_pred_video", + "faker_score_pred_video", +) + + +class LatentVisCallback(Callback): + """Decode and log intermediate training latents as tracker videos.""" + + def __init__( + self, + every_steps: int = 100, + keys: list[str] | None = None, + fps: int = 24, + ) -> None: + self.every_steps = int(every_steps) + self.keys = tuple(keys) if keys else _DEFAULT_KEYS + self.fps = int(fps) + self.tracker: Any = DummyTracker() + + def on_train_start( + self, + method: TrainingMethod, + iteration: int = 0, + ) -> None: + tracker = getattr(method, "tracker", None) + if tracker is not None: + self.tracker = tracker + + def on_training_step_end( + self, + method: TrainingMethod, + loss_dict: dict[str, Any], + iteration: int = 0, + ) -> None: + if self.every_steps <= 0 or iteration % self.every_steps != 0: + return + if get_world_group().rank != 0: + return + vis = getattr(method, "latent_vis", None) + if not vis: + return + decode = getattr(getattr(method, "student", None), "decode_vis_latents", None) + if decode is None: + return + + artifacts: dict[str, Any] = {} + latent_layout = vis.get("_fv_latent_layout") + for key in self.keys: + latent = vis.get(key) + if not isinstance(latent, torch.Tensor): + continue + try: + clip = decode(latent) if latent_layout is None else decode(latent, layout=latent_layout) + except Exception as exc: + logger.warning("Latent visualization decode failed for %r: %s", key, exc) + continue + art = self.tracker.video(clip, fps=self.fps, format="mp4") + if art is not None: + artifacts[f"latent_vis/{key}"] = art + if artifacts: + self.tracker.log_artifacts(artifacts, iteration) diff --git a/fastvideo/train/callbacks/validation.py b/fastvideo/train/callbacks/validation.py index e83775ed04..b1e8247e31 100644 --- a/fastvideo/train/callbacks/validation.py +++ b/fastvideo/train/callbacks/validation.py @@ -12,6 +12,7 @@ import gc import json import os +import re import time from copy import deepcopy from dataclasses import dataclass, field @@ -57,6 +58,7 @@ class _ValidationStepResult: overlay_videos: list[list[np.ndarray]] = field(default_factory=list) overlay_captions: list[str] = field(default_factory=list) ref_videos: list[str | None] = field(default_factory=list) + metadata: list[dict[str, Any]] = field(default_factory=list) actions: list[dict[str, Any] | None] = field(default_factory=list) mouse_pitch_signs: list[int | None] = field(default_factory=list) @@ -136,6 +138,8 @@ def __init__( sampling_steps: list[int] | None = None, guidance_scale: float | None = None, num_frames: int | None = None, + use_record_dimensions: bool = False, + max_record_num_frames: int | None = None, num_videos_per_prompt: int = 1, use_validation_media_conditioning: bool = True, output_dir: str | None = None, @@ -160,6 +164,10 @@ def __init__( self.sampling_steps = ([int(s) for s in sampling_steps] if sampling_steps else [40]) self.guidance_scale = (float(guidance_scale) if guidance_scale is not None else None) self.num_frames = (int(num_frames) if num_frames is not None else None) + self.use_record_dimensions = self._coerce_bool(use_record_dimensions) + self.max_record_num_frames = (int(max_record_num_frames) if max_record_num_frames is not None else None) + if self.max_record_num_frames is not None and self.max_record_num_frames <= 0: + raise ValueError("callbacks.validation.max_record_num_frames must be positive") self.num_videos_per_prompt = int(num_videos_per_prompt) if self.num_videos_per_prompt <= 0: raise ValueError("callbacks.validation.num_videos_per_prompt must be positive") @@ -167,7 +175,6 @@ def __init__( self.output_dir = (str(output_dir) if output_dir is not None else None) self.sampling_timesteps = ([int(s) for s in sampling_timesteps] if sampling_timesteps is not None else None) self.overlay_actions = self._coerce_bool(overlay_actions) - # Validation-only action amplification for world model; training keeps raw action values. self.keyboard_value_scale = float(keyboard_value_scale) metrics_config = pipeline_kwargs.pop("metrics", None) self.metrics_config = self._parse_metrics_config(metrics_config) @@ -176,7 +183,6 @@ def __init__( self.attn_qat_infer = self._coerce_bool(attn_qat_infer) self.pipeline_kwargs = dict(pipeline_kwargs) - # Set after on_train_start. self._pipeline: Any | None = None self._pipeline_key: tuple[Any, ...] | None = None self._sampling_param: SamplingParam | None = None @@ -236,22 +242,47 @@ def _coerce_bool(value: Any) -> bool: return value.strip().lower() in {"1", "true", "yes", "on"} return bool(value) - # ---------------------------------------------------------- - # Callback hooks - # ---------------------------------------------------------- - def _adopt_training_denoising_ladder( + @staticmethod + def _assert_attention_contract( + inference_args: Any, + tc: Any, + ) -> None: + """Fail if validation would sample off the training attention contract. + + Sparsity and tile geometry are separate knobs and both have to survive + the training-config to inference-args hop: v8 validated a + tile-64-trained student at the tile-256 default because only the + sparsity was propagated, and nothing downstream noticed. + """ + for train_attr, args_attr in ( + ("vsa_sparsity", "VSA_sparsity"), + ("vsa_tile_size", "VSA_tile_size"), + ): + trained = getattr(tc, train_attr, None) + sampled = getattr(inference_args, args_attr, None) + if trained is None or sampled is None: + continue + if sampled != trained: + raise ValueError(f"validation would sample at {args_attr}={sampled} while " + f"training runs {train_attr}={trained}. The attention " + "contract must match; fix the training-config to " + "inference-args propagation rather than the symptom.") + + def _adopt_training_sampling_contract( self, method: TrainingMethod, ) -> None: - """Keep validation on the timestep ladder the method trains against. - - ``sampling_steps`` only sets ``num_inference_steps``, which the - scheduler turns into an N-point sigma grid -- N-1 forwards on its own - spacing, not the trained ladder. Validation reaches the trained - operating point only when ``sampling_timesteps`` repeats - ``method.dmd_denoising_steps`` exactly, so inherit it by default and - refuse to run when the two disagree. + """Keep validation sampling on the operating point training teaches. + + A few-step method trains its student on an explicit timestep ladder. + Validation only reaches that ladder when ``sampling_timesteps`` is + configured: ``sampling_steps`` sets ``num_inference_steps``, and the + public scheduler turns N of those into an N-point sigma grid, i.e. + N-1 forwards on the scheduler's own spacing. Duplicating the ladder by + hand in the callback config is the footgun that silently validated v8 + at three forwards on the wrong grid for its whole run, so derive it + from the method instead, and refuse to run when the two disagree. """ method_config = getattr(method, "method_config", None) if not isinstance(method_config, dict): @@ -264,8 +295,8 @@ def _adopt_training_denoising_ladder( if self.sampling_timesteps is None: self.sampling_timesteps = trained logger.info( - "validation: inheriting the trained denoising ladder %s " - "(%d forwards)", + "validation: adopting the trained denoising ladder %s " + "(%d forwards) from the training method", trained, len(trained), ) @@ -274,8 +305,9 @@ def _adopt_training_denoising_ladder( raise ValueError("callbacks.validation.sampling_timesteps " f"{self.sampling_timesteps} disagrees with the trained ladder " f"{trained} (method.dmd_denoising_steps). Validation would " - "sample off the operating point the student is trained for. " - "Drop the override to inherit the ladder, or align the two.") + "sample off the operating point the student was distilled " + "for. Drop the callback override to inherit the ladder, or " + "align the two deliberately.") def on_train_start( self, @@ -284,7 +316,7 @@ def on_train_start( ) -> None: self.method = method tc = self.training_config - self._adopt_training_denoising_ladder(method) + self._adopt_training_sampling_contract(method) self.world_group = get_world_group() self.sp_group = get_sp_group() @@ -309,19 +341,19 @@ def on_validation_begin( iteration: int = 0, ) -> None: """Run the optional step-zero baseline and each scheduled validation event.""" - if self.every_steps <= 0: - return - # Step zero measures the checkpoint before the first optimizer update. - if iteration == 0 and not self.run_at_start: - return - if iteration % self.every_steps != 0: + if not self.will_run_validation(iteration): return self._run_validation(method, iteration) - # ---------------------------------------------------------- - # Core validation logic - # ---------------------------------------------------------- + def will_run_validation(self, iteration: int = 0) -> bool: + """Return whether this callback schedules validation at ``iteration``.""" + if self.every_steps <= 0: + return False + if iteration == 0 and not self.run_at_start: + return False + return iteration % self.every_steps == 0 + def _run_validation( self, @@ -335,8 +367,6 @@ def _run_validation( method, validation_transformer=transformer, ): - # 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, self._attn_qat_infer_context(t): @@ -420,6 +450,17 @@ def _validation_memory_context( self._restore_inactive_role_modules(module_records) self._restore_optimizer_states(optimizer_tensor_records) self._empty_cuda_cache() + off_cuda = [ + f"{name} on {getattr(p, '_local_tensor', p).device}" + for name, p in validation_transformer.named_parameters() + if getattr(p, "_local_tensor", p).device.type != "cuda" + ] + if off_cuda: + logger.warning( + "Post-validation: %d validation-transformer params off-CUDA; first: %s", + len(off_cuda), + off_cuda[:5], + ) def _offload_optimizer_states_to_cpu( self, @@ -504,6 +545,14 @@ def _offload_inactive_role_modules_to_cpu( device = self._first_cuda_tensor_device(module) if device is None: continue + if self._is_fsdp_managed(module): + logger.info( + "Keeping role %r transformer on %s during validation " + "(FSDP-managed modules do not survive .to() round-trips).", + role, + device, + ) + continue try: module.to("cpu") except Exception as exc: @@ -532,6 +581,17 @@ def _restore_inactive_role_modules( device, ) + @staticmethod + def _is_fsdp_managed(module: torch.nn.Module) -> bool: + try: + from torch.distributed.fsdp import FSDPModule + from torch.distributed.tensor import DTensor + except ImportError: + return False + if isinstance(module, FSDPModule): + return True + return any(isinstance(p, DTensor) for p in module.parameters(recurse=True)) + @staticmethod def _first_cuda_tensor_device(module: torch.nn.Module) -> torch.device | None: for tensor in list(module.parameters(recurse=True)) + list(module.buffers(recurse=True)): @@ -590,8 +650,6 @@ def _run_validation_inner( transformer=transformer, ) - # Every rank participates in sequence-parallel inference, but - # only the group leader retains decoded media for saving. if self.rank_in_sp_group != 0: continue @@ -630,6 +688,10 @@ def _run_validation_inner( result.ref_videos, local_videos.indices, ) + local_metadata = self._select_by_indices( + result.metadata, + local_videos.indices, + ) local_actions = self._select_by_indices( result.actions, local_videos.indices, @@ -655,22 +717,26 @@ def _run_validation_inner( all_overlay_video_filenames = list(local_overlay_video_filenames) all_captions = list(local_captions) all_overlay_captions = list(local_overlay_captions) + all_ref_videos = list(local_ref_videos) + all_metadata = list(local_metadata) all_audio_video_count = local_videos.audio_video_count all_metric_stats = local_metric_stats for sp_idx in range(1, num_sp_groups): - # Sequence-parallel group leaders occupy the first - # global rank in each contiguous group. src = (sp_idx * self.sp_world_size) recv_v = (self.world_group.recv_object(src=src)) recv_c = (self.world_group.recv_object(src=src)) recv_ov = (self.world_group.recv_object(src=src)) recv_oc = (self.world_group.recv_object(src=src)) + recv_ref = (self.world_group.recv_object(src=src)) + recv_metadata = (self.world_group.recv_object(src=src)) recv_m = (self.world_group.recv_object(src=src)) recv_audio_video_count = (self.world_group.recv_object(src=src)) all_video_filenames.extend(recv_v) all_overlay_video_filenames.extend(recv_ov) all_captions.extend(recv_c) all_overlay_captions.extend(recv_oc) + all_ref_videos.extend(recv_ref) + all_metadata.extend(recv_metadata) all_audio_video_count += int(recv_audio_video_count) self._merge_metric_stats( all_metric_stats, @@ -681,14 +747,37 @@ def _run_validation_inner( all_metric_stats, step=step, ) - # Media and completion counts share one tracker event so - # artifacts and verification data remain aligned. + display_captions = [ + self._validation_artifact_caption(caption, metadata) + for caption, metadata in zip(all_captions, all_metadata, strict=True) + ] + reference_filenames: list[str] = [] + reference_captions: list[str] = [] + for caption, ref_video, metadata in zip( + all_captions, + all_ref_videos, + all_metadata, + strict=True, + ): + if ref_video is None or not os.path.isfile(ref_video): + continue + reference_filenames.append(ref_video) + reference_captions.append( + self._validation_artifact_caption( + caption, + metadata, + prefix="held-out reference", + use_reference_num_frames=True, + )) self._log_validation_video_artifacts( all_video_filenames, - all_captions, + display_captions, key=f"validation_videos_{num_inference_steps}_steps", step=step, fps=sp.fps, + reference_video_filenames=reference_filenames, + reference_captions=reference_captions, + reference_key=f"validation_references_{num_inference_steps}_steps", scalar_metrics={ f"validation/{num_inference_steps}_steps_video_count": float(len(all_video_filenames)), @@ -696,6 +785,12 @@ def _run_validation_inner( (time.perf_counter() - validation_started_at), f"validation/{num_inference_steps}_steps_audio_video_count": float(all_audio_video_count), + f"validation/{num_inference_steps}_steps_reference_video_count": + float(len(reference_filenames)), + **self._validation_metadata_scalar_metrics( + all_metadata, + num_inference_steps=num_inference_steps, + ), }, ) if all_overlay_video_filenames: @@ -724,6 +819,14 @@ def _run_validation_inner( local_overlay_captions, dst=0, ) + self.world_group.send_object( + local_ref_videos, + dst=0, + ) + self.world_group.send_object( + local_metadata, + dst=0, + ) self.world_group.send_object( local_metric_stats, dst=0, @@ -784,8 +887,6 @@ def _save_validation_videos( audio_sample_rate=audio_sample_rate, ) except Exception as exc: - # Validation media is diagnostic output, so one failed write - # must not terminate training or prevent later artifact writes. logger.exception( "Failed to save validation media %s on rank %s; skipping artifact: %s", fname, @@ -798,7 +899,6 @@ def _save_validation_videos( saved.filenames.append(fname) saved.indices.append(i) if audio is not None: - # The media writer verifies requested streams before returning. saved.audio_video_count += 1 return saved @@ -818,6 +918,9 @@ def _log_validation_video_artifacts( step: int, fps: int, scalar_metrics: dict[str, float] | None = None, + reference_video_filenames: list[str] | None = None, + reference_captions: list[str] | None = None, + reference_key: str | None = None, ) -> None: """Log validation media and its scalar verification data at one step.""" video_logs = [] @@ -833,18 +936,87 @@ def _log_validation_video_artifacts( ) if art is not None: video_logs.append(art) + artifacts: dict[str, Any] = {} if video_logs: - artifacts: dict[str, Any] = {key: video_logs} - if scalar_metrics: - artifacts.update(scalar_metrics) + artifacts[key] = video_logs + if ((reference_video_filenames is None) != (reference_captions is None) + or (reference_video_filenames is not None) != (reference_key is not None)): + raise ValueError("Validation reference filenames, captions, and key must be provided together.") + if reference_video_filenames is not None: + reference_logs = [] + for fname, cap in zip( + reference_video_filenames, + reference_captions or [], + strict=True, + ): + art = self.tracker.video( + fname, + caption=cap, + fps=fps, + ) + if art is not None: + reference_logs.append(art) + if reference_logs: + assert reference_key is not None + artifacts[reference_key] = reference_logs + if scalar_metrics: + artifacts.update(scalar_metrics) + if artifacts: self.tracker.log_artifacts( artifacts, step, ) - # ---------------------------------------------------------- - # Metric evaluation - # ---------------------------------------------------------- + @staticmethod + def _validation_artifact_caption( + caption: str, + metadata: dict[str, Any], + *, + prefix: str = "generated", + use_reference_num_frames: bool = False, + ) -> str: + fields = [prefix] + source = metadata.get("source") + sample_id = metadata.get("sample_id") + if source: + fields.append(f"source={source}") + if sample_id: + fields.append(f"id={sample_id}") + width = metadata.get("width") + height = metadata.get("height") + num_frames = (metadata.get("reference_num_frames", metadata.get("num_frames")) + if use_reference_num_frames else metadata.get("num_frames")) + if width and height and num_frames: + fields.append(f"shape={width}x{height}x{num_frames}f") + return f"[{' | '.join(fields)}] {caption}" + + @staticmethod + def _validation_metric_segment(value: Any) -> str: + segment = re.sub(r"[^A-Za-z0-9_.-]+", "_", str(value).strip()) + return segment.strip("_") or "unknown" + + @classmethod + def _validation_metadata_scalar_metrics( + cls, + metadata: list[dict[str, Any]], + *, + num_inference_steps: int, + ) -> dict[str, float]: + metrics: dict[str, float] = {} + prefix = f"validation/{num_inference_steps}_steps" + for record in metadata: + source = cls._validation_metric_segment(record.get("source", "unknown")) + source_key = f"{prefix}/source/{source}_count" + metrics[source_key] = metrics.get(source_key, 0.0) + 1.0 + width = record.get("width") + height = record.get("height") + num_frames = record.get("num_frames") + if width and height and num_frames: + shape = f"{int(width)}x{int(height)}x{int(num_frames)}f" + shape_key = f"{prefix}/shape/{shape}_count" + metrics[shape_key] = metrics.get(shape_key, 0.0) + 1.0 + return metrics + def _metric_device(self) -> str: device = self.metrics_config.device @@ -1176,9 +1348,6 @@ def _validation_mouse_pitch_sign( return -1 return int(self.metrics_config.mouse_pitch_sign) - # ---------------------------------------------------------- - # Pipeline management - # ---------------------------------------------------------- def _get_sampling_param(self) -> SamplingParam: if self._sampling_param is None: @@ -1275,6 +1444,25 @@ def _sync_runtime_dit_arch_config( getattr(transformer, name), ) + def _inject_method_denoising_steps(self, validation_config: Any) -> None: + """Sample validation at the method's trained DMD jump points. + + Distillation students only ever denoise from ``dmd_denoising_steps``; + stages that honor ``pipeline_config.dmd_denoising_steps`` should visit + the same points instead of the scheduler's native grid. An explicit + value already on the pipeline config wins, and warped lists stay out + (their raw entries are grid indices, not timesteps). + """ + method_config = getattr(self.method, "method_config", None) + if not isinstance(method_config, dict): + return + steps = method_config.get("dmd_denoising_steps") + if not steps or bool(method_config.get("warp_denoising_step", False)): + return + if getattr(validation_config, "dmd_denoising_steps", "unset") is not None: + return + validation_config.dmd_denoising_steps = [int(step) for step in steps] + def _get_pipeline( self, *, @@ -1293,8 +1481,6 @@ def _get_pipeline( ) loaded_modules: dict[str, Any] = {"transformer": transformer} - # Distillation methods build the flow-match scheduler their few-step DMD - # sampler needs; inject it so the pipeline doesn't fall back to UniPC. method_scheduler = getattr(self.method, "_sf_scheduler", None) if method_scheduler is not None: loaded_modules["scheduler"] = method_scheduler @@ -1313,8 +1499,6 @@ def _get_pipeline( kwargs["flow_shift"] = float(flow_shift) kwargs.update(self.pipeline_kwargs) - # The pipeline class comes from a YAML target, so static analysis cannot - # infer the dynamically resolved ``from_pretrained`` class method. self._pipeline = PipelineCls.from_pretrained( # type: ignore[attr-defined] self._pipeline_model_path(), **kwargs, @@ -1326,6 +1510,7 @@ def _get_pipeline( validation_config, loaded_config, ) + self._inject_method_denoising_steps(validation_config) self._pipeline.fastvideo_args.pipeline_config = validation_config arch_config = self._pipeline.fastvideo_args.pipeline_config.dit_config.arch_config logger.info( @@ -1338,9 +1523,6 @@ def _get_pipeline( self._pipeline_key = key return self._pipeline - # ---------------------------------------------------------- - # Batch preparation - # ---------------------------------------------------------- def _prepare_validation_batch( self, @@ -1352,34 +1534,23 @@ def _prepare_validation_batch( tc = self.training_config sampling_param.prompt = validation_batch["prompt"] - sampling_param.height = tc.data.num_height - sampling_param.width = tc.data.num_width + height, width, num_frames = self._validation_sampling_dimensions(validation_batch) + sampling_param.height = height + sampling_param.width = width sampling_param.num_inference_steps = int(num_inference_steps) sampling_param.data_type = "video" if self.guidance_scale is not None: sampling_param.guidance_scale = float(self.guidance_scale) sampling_param.seed = self.seed - # Output multiplicity belongs in SamplingParam so pipeline stages - # allocate the same batch dimension that validation expects to log. sampling_param.num_videos_per_prompt = self.num_videos_per_prompt - # SamplingParam is cached across records; clearing the path prevents a - # prior record's image or video from conditioning a later prompt. sampling_param.image_path = None if self.use_validation_media_conditioning: - # Image-to-video pipelines use an image or the first frame of a validation video. img_path = (validation_batch.get("image_path") or validation_batch.get("video_path")) if img_path is not None and (img_path.startswith("http") or os.path.isfile(img_path)): sampling_param.image_path = img_path - temporal_compression_factor = int( - tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio # type: ignore[union-attr] - ) - default_num_frames = ((tc.data.num_latent_t - 1) * temporal_compression_factor + 1) - if self.num_frames is not None: - sampling_param.num_frames = int(self.num_frames) - else: - sampling_param.num_frames = int(default_num_frames) + sampling_param.num_frames = num_frames latents_size = [ (sampling_param.num_frames - 1) // 4 + 1, @@ -1397,6 +1568,7 @@ def _prepare_validation_batch( tc, model_path=tc.model_path, ) + self._assert_attention_contract(inference_args, tc) batch = ForwardBatch( **shallow_asdict(sampling_param), @@ -1406,18 +1578,10 @@ def _prepare_validation_batch( VSA_sparsity=tc.vsa_sparsity, timesteps=sampling_timesteps_tensor, ) - # shallow_asdict(sampling_param) copies list-typed fields by - # reference. sampling_param is cached and reused across every - # validation sample/step, so without this reset every ForwardBatch - # would share (and keep appending to) the same - # prompt_attention_mask/negative_attention_mask list forever -- - # index [0] would then hold the *first-ever* validation sample's - # mask instead of the current one, mismatching prompt_embeds. batch.prompt_attention_mask = [] batch.negative_attention_mask = [] batch._inference_args = inference_args # type: ignore[attr-defined] - # Conditionally set I2V fields. if ("image" in validation_batch and validation_batch["image"] is not None): batch.pil_image = validation_batch["image"] @@ -1429,6 +1593,50 @@ def _prepare_validation_batch( return batch + def _validation_sampling_dimensions( + self, + validation_batch: dict[str, Any], + ) -> tuple[int, int, int]: + """Resolve output geometry, optionally from one validation record. + + Native-shape validation is explicit because cached ``SamplingParam`` + instances are shared across records. A complete record triplet wins; + partial metadata fails instead of combining dimensions from unrelated + shapes. ``max_record_num_frames`` optionally caps only the temporal + member of a complete record triplet; fixed/default geometry is never + changed. With the option off (the default), legacy callback/config + behavior is unchanged. + """ + tc = self.training_config + temporal_compression_factor = int( + tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio # type: ignore[union-attr] + ) + default_num_frames = ((tc.data.num_latent_t - 1) * temporal_compression_factor + 1) + dimensions = { + "height": int(tc.data.num_height), + "width": int(tc.data.num_width), + "num_frames": (int(self.num_frames) if self.num_frames is not None else int(default_num_frames)), + } + + if self.use_record_dimensions: + present = {name: validation_batch.get(name) is not None for name in dimensions} + if any(present.values()) and not all(present.values()): + missing = [name for name, is_present in present.items() if not is_present] + raise ValueError("Native-shape validation records must provide width, height, and num_frames together; " + f"missing {missing} for prompt {validation_batch.get('prompt')!r}.") + if all(present.values()): + dimensions = {name: int(validation_batch[name]) for name in dimensions} + if self.max_record_num_frames is not None: + dimensions["num_frames"] = min(dimensions["num_frames"], self.max_record_num_frames) + + for name, value in dimensions.items(): + if value <= 0: + raise ValueError(f"Validation {name} must be positive, got {value}") + if dimensions["height"] % 8 or dimensions["width"] % 8: + raise ValueError("Validation width and height must be divisible by 8, got " + f"{dimensions['width']}x{dimensions['height']}") + return dimensions["height"], dimensions["width"], dimensions["num_frames"] + def _attach_action_conditions( self, batch: ForwardBatch, @@ -1462,9 +1670,6 @@ def _attach_action_conditions( tensor, ) - # ---------------------------------------------------------- - # Validation loop - # ---------------------------------------------------------- def _run_validation_for_steps( self, @@ -1494,6 +1699,7 @@ def _run_validation_for_steps( tc, model_path=tc.model_path, ) + self._assert_attention_contract(inference_args, tc) self._sync_runtime_dit_arch_config( inference_args.pipeline_config, transformer, @@ -1503,8 +1709,6 @@ def _run_validation_for_steps( pipeline.fastvideo_args.pipeline_config, ) - # Propagate sampling_timesteps to pipeline_config so - # causal/DMD denoising stages can read them. if (self.sampling_timesteps is not None and inference_args.pipeline_config.dmd_denoising_steps is None): inference_args.pipeline_config.dmd_denoising_steps = ([int(s) for s in self.sampling_timesteps]) @@ -1515,6 +1719,7 @@ def _run_validation_for_steps( captions: list[str] = [] overlay_captions: list[str] = [] ref_videos: list[str | None] = [] + metadata: list[dict[str, Any]] = [] actions: list[dict[str, Any] | None] = [] mouse_pitch_signs: list[int | None] = [] @@ -1526,7 +1731,7 @@ def _run_validation_for_steps( ) assert (batch.prompt is not None and isinstance(batch.prompt, str)) - ref_video = validation_batch.get("ref_video") + ref_video = (validation_batch.get("ref_video") or validation_batch.get("video_path")) action = self._validation_actions(validation_batch) with torch.no_grad(): @@ -1539,8 +1744,6 @@ def _run_validation_for_steps( if self.rank_in_sp_group != 0: continue - # Append metadata only on the group leader so every list position - # describes the same decoded video throughout save and logging. output_audio = output_batch.extra.get("audio") output_audio_sample_rate = output_batch.extra.get("audio_sample_rate") if (output_audio is None) != (output_audio_sample_rate is None): @@ -1550,8 +1753,6 @@ def _run_validation_for_steps( raise TypeError("Validation pipeline audio must be a torch.Tensor or numpy.ndarray; " f"got {type(output_audio).__name__}.") if torch.is_tensor(output_audio): - # The returned validation result stays on CPU so MP4 encoding - # does not retain a generation tensor on the GPU. output_audio = output_audio.detach().cpu() video = rearrange( @@ -1571,6 +1772,17 @@ def _run_validation_for_steps( audio_waveforms.append(output_audio) audio_sample_rates.append(int(output_audio_sample_rate) if output_audio_sample_rate is not None else None) ref_videos.append(ref_video if isinstance(ref_video, str) else None) + record_metadata: dict[str, Any] = { + "source": validation_batch.get("source", "unknown"), + "sample_id": validation_batch.get("sample_id", validation_batch.get("id")), + "width": int(batch.width), + "height": int(batch.height), + "num_frames": int(batch.num_frames), + } + reference_num_frames = validation_batch.get("num_frames") + if reference_num_frames is not None and int(reference_num_frames) != int(batch.num_frames): + record_metadata["reference_num_frames"] = int(reference_num_frames) + metadata.append(record_metadata) actions.append(action) mouse_pitch_signs.append(self._validation_mouse_pitch_sign(validation_batch)) if self.overlay_actions: @@ -1590,6 +1802,7 @@ def _run_validation_for_steps( overlay_videos=overlay_videos, overlay_captions=overlay_captions, ref_videos=ref_videos, + metadata=metadata, actions=actions, mouse_pitch_signs=mouse_pitch_signs, ) @@ -1609,9 +1822,6 @@ def _post_process_validation_frames( action=action, ) - # ---------------------------------------------------------- - # State management - # ---------------------------------------------------------- def state_dict(self) -> dict[str, Any]: state: dict[str, Any] = {} diff --git a/fastvideo/train/entrypoint/dcp_to_diffusers.py b/fastvideo/train/entrypoint/dcp_to_diffusers.py index da9d8e0b56..db00b7cb45 100644 --- a/fastvideo/train/entrypoint/dcp_to_diffusers.py +++ b/fastvideo/train/entrypoint/dcp_to_diffusers.py @@ -62,6 +62,7 @@ def _save_role_pretrained( output_dir: str, module_names: list[str] | None = None, overwrite: bool = False, + link_base: bool = False, model: Any, ) -> str: """Export a role's modules into a diffusers-style model dir. @@ -105,10 +106,6 @@ def _barrier() -> None: "Pass --overwrite to replace it.") def _copy_or_link(src: str, dest: str) -> None: - # Resolve symlinks ourselves: os.link's follow_symlinks=True - # default isn't honored on all filesystems (e.g. some - # network/overlay mounts), which can silently hard-link to - # the symlink itself instead of its target. real_src = os.path.realpath(src) try: os.link(real_src, dest) @@ -117,16 +114,40 @@ def _copy_or_link(src: str, dest: str) -> None: logger.info( "Creating pretrained export dir at %s " - "(base=%s)", + "(base=%s, link_base=%s)", dst, local_base, + link_base, ) - shutil.copytree( - local_base, - dst, - symlinks=False, - copy_function=_copy_or_link, - ) + if link_base: + rewritten = set(module_names or ["transformer"]) + dst.mkdir(parents=True, exist_ok=True) + for entry in sorted(local_base.iterdir()): + if entry.name == ".cache" or entry.name.startswith(".git"): + continue + target = dst / entry.name + if entry.is_dir() and entry.name in rewritten: + target.mkdir() + for f in sorted(entry.iterdir()): + if f.name.endswith(".safetensors") or f.name.endswith(".safetensors.index.json"): + continue + shutil.copy2(os.path.realpath(f), target / f.name) + else: + os.symlink(os.path.realpath(entry), target) + else: + try: + shutil.copytree( + local_base, + dst, + symlinks=False, + copy_function=_copy_or_link, + ignore=shutil.ignore_patterns(".cache", ".git*"), + ) + except shutil.Error as exc: + logger.warning("copytree finished with %d skipped entries (first: %s)", len(exc.args[0]), + exc.args[0][0] if exc.args[0] else "?") + if not ((dst / "modular_model_index.json").is_file() or (dst / "model_index.json").is_file()): + raise FileNotFoundError(f"Export dir {dst} is missing its model index after copy.") _barrier() @@ -160,13 +181,9 @@ def _copy_or_link(src: str, dest: str) -> None: if _rank() == 0: for path in module_dir.glob("*.safetensors"): path.unlink(missing_ok=True) + for path in module_dir.glob("*.safetensors.index.json"): + path.unlink(missing_ok=True) - # Convert internal parameter names back to HF format. - # load_model_from_full_model_state_dict builds reverse_param_names_mapping - # (internal_key → hf_key) and stores it on the module. Without this, - # the exported safetensors would have internal keys (e.g. - # "patch_embedding.proj.bias") and the next load would double-map them - # (e.g. → "patch_embedding.proj.proj.bias"). reverse_mapping: dict = getattr(modules[module_name], "reverse_param_names_mapping", {}) tensor_state: dict[str, torch.Tensor] = {} @@ -233,6 +250,8 @@ def convert( role: str = "student", overwrite: bool = False, verify: bool = False, + weights_only: bool = False, + link_base: bool = False, ) -> str: """Load a DCP checkpoint and export as a diffusers model. @@ -243,6 +262,8 @@ def convert( from fastvideo.distributed import ( maybe_init_distributed_environment_and_model_parallel, ) from fastvideo.train.utils.builder import build_from_config + from fastvideo.train.utils.instantiate import instantiate + from fastvideo.training.checkpointing_utils import ModelWrapper from fastvideo.train.utils.checkpoint import ( CheckpointManager, _resolve_resume_checkpoint, @@ -254,7 +275,6 @@ def convert( import torch.distributed.checkpoint as dcp - # -- Resolve checkpoint directory -- resolved = _resolve_resume_checkpoint( checkpoint_dir, output_dir=checkpoint_dir, @@ -263,7 +283,6 @@ def convert( if not dcp_dir.is_dir(): raise FileNotFoundError(f"Missing dcp/ under {resolved}") - # -- Obtain config -- cfg: RunConfig if config_path is not None: cfg = load_run_config(config_path) @@ -278,32 +297,34 @@ def convert( tc = cfg.training - # -- Init distributed (1 GPU is enough; DCP reshards) -- maybe_init_distributed_environment_and_model_parallel( tp_size=1, sp_size=1, ) - # Override distributed config so model loading uses 1 GPU. tc.distributed.tp_size = 1 tc.distributed.sp_size = 1 tc.distributed.num_gpus = 1 tc.distributed.hsdp_replicate_dim = 1 tc.distributed.hsdp_shard_dim = 1 - # -- Build model (loads pretrained weights + FSDP) -- - _, method, _, _ = build_from_config(cfg) - - # -- Load DCP weights into the model -- - states = method.checkpoint_state() + if weights_only: + if role not in cfg.models: + raise KeyError(f"Role {role!r} is not present in the checkpoint config") + model = instantiate(cfg.models[role], training_config=tc) + if model.transformer is None: + raise ValueError(f"Role {role!r} has no transformer to export") + states = {f"roles.{role}.transformer": ModelWrapper(model.transformer)} + else: + _, method, _, _ = build_from_config(cfg) + states = method.checkpoint_state() + model = method._role_models[role] logger.info( "Loading DCP checkpoint from %s", resolved, ) dcp.load(states, checkpoint_id=str(dcp_dir)) - # -- Export to diffusers format -- - model = method._role_models[role] base_model_path = str(tc.model_path) if not base_model_path: raise ValueError("Cannot determine base_model_path from " @@ -321,6 +342,7 @@ def convert( base_model_path=base_model_path, output_dir=output_dir, overwrite=overwrite, + link_base=link_base, model=model, ) logger.info("Export complete: %s", result) @@ -443,6 +465,22 @@ def main() -> None: "the exported directory to catch key-mapping bugs " "immediately."), ) + parser.add_argument( + "--weights-only", + action="store_true", + help=("Load only roles.* module weights from the checkpoint, " + "skipping optimizer/scheduler states (halves GPU memory " + "and allows exporting via a shim config whose optimizer " + "differs from the checkpoint's)."), + ) + parser.add_argument( + "--link-base", + action="store_true", + help=("Symlink base-model components into the export dir instead " + "of copying them (only the exported module dirs are real). " + "Saves hundreds of GB per export; the export then depends on " + "the base model dir staying in place."), + ) args = parser.parse_args(sys.argv[1:]) convert( @@ -452,6 +490,8 @@ def main() -> None: role=args.role, overwrite=args.overwrite, verify=args.verify, + weights_only=args.weights_only, + link_base=args.link_base, ) diff --git a/fastvideo/train/entrypoint/train.py b/fastvideo/train/entrypoint/train.py index 3f71e2a721..8c6e0e5946 100644 --- a/fastvideo/train/entrypoint/train.py +++ b/fastvideo/train/entrypoint/train.py @@ -47,7 +47,6 @@ def run_training_from_config( from fastvideo.train.utils.builder import build_from_config from fastvideo.train.utils.config import load_run_config - # Enable deterministic mode for reproducibility. torch.backends.cudnn.benchmark = False torch.backends.cudnn.deterministic = True @@ -56,9 +55,7 @@ def run_training_from_config( model_path_lower = str(tc.model_path).lower() - # Auto-set attention backend for model families that require a specific - # backend at load time, unless the user already overrode it explicitly. - if tc.vsa_sparsity > 0.0: + if tc.vsa_sparsity > 0.0 and "minimax" not in model_path_lower: os.environ.setdefault( "FASTVIDEO_ATTENTION_BACKEND", "VIDEO_SPARSE_ATTN", @@ -87,8 +84,6 @@ def run_training_from_config( callback_configs=cfg.callbacks, ) - # Attach the exact YAML used for this run to the - # tracker (e.g., W&B Files). trainer.tracker.log_file( os.path.abspath(os.path.expanduser(config_path)), name="run.yaml", @@ -97,6 +92,11 @@ def run_training_from_config( ckpt_config = CheckpointConfig( save_steps=int(tc.checkpoint.training_state_checkpointing_steps or 0), keep_last=int(tc.checkpoint.checkpoints_total_limit or 0), + start_step=int(tc.checkpoint.checkpointing_start_step or 0), + save_inference_on_validation=bool(tc.checkpoint.save_inference_checkpoint_on_validation), + inference_role=str(tc.checkpoint.inference_checkpoint_role or "student"), + inference_dtype=str(tc.checkpoint.inference_checkpoint_dtype or "bfloat16"), + require_complete_training_checkpoint=bool(tc.checkpoint.require_complete_training_checkpoint), ) checkpoint_manager = CheckpointManager( diff --git a/fastvideo/train/methods/base.py b/fastvideo/train/methods/base.py index 5168313c5f..8709d5dec2 100644 --- a/fastvideo/train/methods/base.py +++ b/fastvideo/train/methods/base.py @@ -40,7 +40,6 @@ class TrainingMethod(torch.nn.Module, ABC): generator instead of relying on global RNG state. """ - # Shared CUDA RNG generator (initialized in on_train_start). cuda_generator: torch.Generator | None = None def __init__( @@ -58,7 +57,6 @@ def __init__( self.method_config: dict[str, Any] = dict(cfg.method) self.validation_config: dict[str, Any] = dict(getattr(cfg, "validation", {}) or {}) - # Build nn.ModuleDict for FSDP / checkpoint visibility. self.role_modules = torch.nn.ModuleDict() for role, model in role_models.items(): mods: dict[str, torch.nn.Module] = {} @@ -68,7 +66,6 @@ def __init__( if mods: self.role_modules[role] = torch.nn.ModuleDict(mods) - # ------------------------------------------------------------------ def set_tracker(self, tracker: Any) -> None: self.tracker = tracker @@ -109,6 +106,17 @@ def _optimizer_dict(self) -> dict[str, Any]: def _lr_scheduler_dict(self) -> dict[str, Any]: ... + def apply_configured_lrs(self) -> None: + """Re-apply the config's learning rates to live optimizers/schedulers. + + Called after a checkpoint resume when + ``training.checkpoint.reset_lr_on_resume`` is set: DCP restores + param-group ``lr``/``initial_lr`` and scheduler ``base_lrs`` from the + checkpoint, which would otherwise silently override an LR change in + the YAML. Methods with per-role LRs override this. + """ + return None + def checkpoint_state(self) -> dict[str, Any]: """Return DCP-ready checkpoint state for all trainable roles. @@ -144,6 +152,41 @@ def checkpoint_state(self) -> dict[str, Any]: return states + def inference_checkpoint_modules( + self, + role: str = "student", + ) -> dict[str, torch.nn.Module]: + """Return the complete modules that form an inference checkpoint. + + This intentionally does not reuse :meth:`checkpoint_state`: resumable + state includes optimizers and every trainable role, while inference + exports select one explicit role and must retain frozen parameters too. + Methods with a different deployable role contract may override this + hook; EMA is never selected implicitly. + """ + model = self._role_models.get(role) + if model is None: + raise ValueError(f"Inference checkpoint role {role!r} is not available; " + f"known roles: {sorted(self._role_models)}") + transformer = getattr(model, "transformer", None) + if not isinstance(transformer, torch.nn.Module): + raise ValueError(f"Inference checkpoint role {role!r} has no transformer module") + return {"transformer": transformer} + + def inference_checkpoint_base_model_path( + self, + role: str = "student", + ) -> str: + """Return the immutable model directory supplying non-trained parts.""" + model = self._role_models.get(role) + if model is None: + raise ValueError(f"Inference checkpoint role {role!r} is not available; " + f"known roles: {sorted(self._role_models)}") + path = str(getattr(model, "_init_from", "") or "") + if not path: + raise ValueError(f"Inference checkpoint role {role!r} does not expose its init_from model path") + return path + def backward( self, loss_map: dict[str, torch.Tensor], @@ -182,7 +225,7 @@ def seed_optimizer_state_for_resume(self) -> None: DCP needs matching entries to load into; without them the saved optimizer state is silently dropped. """ - for opt in self.get_optimizers(0): + for opt in self._optimizer_dict.values(): for group in opt.param_groups: for p in group["params"]: if not p.requires_grad: @@ -195,7 +238,6 @@ def seed_optimizer_state_for_resume(self) -> None: "exp_avg_sq": torch.zeros_like(p), } - # -- Shared hooks (override in subclasses as needed) -- def manages_optimization(self) -> bool: """Whether the method owns backward/optimizer stepping internally. @@ -263,7 +305,6 @@ def on_train_start(self) -> None: global_rank = int(world_group.rank) sp_size = int(self.training_config.distributed.sp_size or 1) - # Ranks within the same SP group share a seed. sp_group_seed = seed + (global_rank // sp_size) if sp_size > 1 else seed + global_rank set_random_seed(seed) @@ -275,6 +316,6 @@ def on_train_start(self) -> None: def _infer_attn_kind(self) -> Literal["dense", "vsa"]: """Derive metadata mode from the student's resolved backend.""" backend = (self.student.attention_backend_name or envs.FASTVIDEO_ATTENTION_BACKEND) - if backend == "VIDEO_SPARSE_ATTN": + if backend in ("VIDEO_SPARSE_ATTN", "VIDEO_SPARSE_ATTN_H3"): return "vsa" return "dense" diff --git a/fastvideo/train/methods/distribution_matching/dmd2.py b/fastvideo/train/methods/distribution_matching/dmd2.py index ee1890bcb9..97b250c1e3 100644 --- a/fastvideo/train/methods/distribution_matching/dmd2.py +++ b/fastvideo/train/methods/distribution_matching/dmd2.py @@ -3,9 +3,13 @@ from __future__ import annotations +import json +import math +from pathlib import Path from typing import Any, Literal import torch +import torch.distributed as dist import torch.nn.functional as F from fastvideo.train.methods.base import TrainingMethod, LogScalar @@ -55,6 +59,13 @@ def __init__( raise ValueError("DMD2Method requires critic to be trainable") self._cfg_uncond = self._parse_cfg_uncond() self._rollout_mode = self._parse_rollout_mode() + ( + self._rollout_carry, + self._rollout_carry_slot_count, + self._rollout_sample_type, + ) = self._parse_rollout_carry() + self._init_rollout_carry_state() + self._rollout_data_forcing = self._parse_rollout_data_forcing() self._validate_preprocessed_data_type() self._configure_student_negative_conditioning() self._denoising_step_list: torch.Tensor | None = (None) @@ -62,8 +73,11 @@ def __init__( self._score_min_timestep, self._score_max_timestep, ) = self._parse_score_timestep_bounds() + self._score_timestep_shift = self._parse_score_timestep_shift() + self._score_timestep_warp_max = self._parse_score_timestep_warp_max() + self._score_timestep_continuous = self._parse_score_timestep_continuous() + self._fake_score_loss_space = self._parse_fake_score_loss_space() - # Initialize preprocessors on student. self.student.init_preprocessors(self.training_config) self._init_optimizers_and_schedulers() @@ -82,7 +96,6 @@ def _lr_scheduler_dict(self) -> dict[str, Any]: "critic": self._critic_lr_scheduler, } - # TrainingMethod override: single_train_step def single_train_step( self, batch: dict[str, Any], @@ -92,6 +105,9 @@ def single_train_step( dict[str, Any], dict[str, LogScalar], ]: + if self._rollout_carry: + return self._carried_train_step(batch, iteration) + latents_source: Literal["data", "zeros"] = "data" if self._rollout_mode == "simulate": latents_source = "zeros" @@ -107,22 +123,29 @@ def single_train_step( generator_loss = torch.zeros( (), device=training_batch.latents.device, - dtype=training_batch.latents.dtype, + dtype=torch.float32, ) student_ctx = None + generator_metrics: dict[str, LogScalar] = {} + fake_score_loss = torch.zeros_like(generator_loss) + critic_ctx = None + critic_outputs: dict[str, Any] = {} + critic_metrics: dict[str, LogScalar] = {} if update_student: generator_pred_x0 = self._student_rollout(training_batch, with_grad=True) student_ctx = ( training_batch.timesteps, training_batch.attn_metadata_vsa, ) - generator_loss = self._dmd_loss(generator_pred_x0, training_batch) - - ( - fake_score_loss, - critic_ctx, - critic_outputs, - ) = self._critic_flow_matching_loss(training_batch) + generator_loss, generator_metrics = self._dmd_loss(generator_pred_x0, training_batch) + training_batch.dmd_latent_vis_dict["generator_pred_video"] = generator_pred_x0.detach() + else: + ( + fake_score_loss, + critic_ctx, + critic_outputs, + critic_metrics, + ) = self._critic_flow_matching_loss(training_batch) total_loss = generator_loss + fake_score_loss loss_map = { @@ -137,10 +160,19 @@ def single_train_step( "student_ctx": student_ctx, "critic_ctx": critic_ctx, } - metrics: dict[str, LogScalar] = {"update_student": float(update_student)} + metrics: dict[str, LogScalar] = { + "update_student": float(update_student), + **generator_metrics, + **critic_metrics, + } + self.latent_vis = { + **(training_batch.fake_score_latent_vis_dict or {}), + **(training_batch.dmd_latent_vis_dict or {}), + "_fv_latent_layout": + getattr(training_batch, "minimax_h3_dmd_layout", None), + } return loss_map, outputs, metrics - # TrainingMethod override: backward def backward( self, loss_map: dict[str, torch.Tensor], @@ -168,6 +200,8 @@ def backward( student_ctx, grad_accum_rounds=grad_accum_rounds, ) + self._assert_finite_gradients("student", self.student) + return critic_ctx = backward_ctx.get("critic_ctx") if critic_ctx is None: @@ -177,39 +211,111 @@ def backward( critic_ctx, grad_accum_rounds=grad_accum_rounds, ) + self._assert_finite_gradients("critic", self.critic) + + @staticmethod + def _local_tensor(tensor: torch.Tensor) -> torch.Tensor: + return getattr(tensor, "_local_tensor", tensor) + + @classmethod + def _assert_finite_gradients(cls, role: str, model: ModelBase) -> None: + """Abort on numerical corruption instead of applying a partial update.""" + bad: list[str] = [] + for name, parameter in model.transformer.named_parameters(): + if parameter.grad is None: + continue + if not bool(torch.isfinite(cls._local_tensor(parameter.grad)).all()): + bad.append(name) + if len(bad) == 8: + break + if bad: + raise RuntimeError( + f"Nonfinite {role} gradients before clipping/Adam: {bad}" + ) - # TrainingMethod override: get_optimizers def get_optimizers( self, iteration: int, ) -> list[torch.optim.Optimizer]: - optimizers: list[torch.optim.Optimizer] = [] - optimizers.append(self._critic_optimizer) if self._should_update_student(iteration): - optimizers.append(self._student_optimizer) - return optimizers + return [self._student_optimizer] + return [self._critic_optimizer] - # TrainingMethod override: get_lr_schedulers def get_lr_schedulers( self, iteration: int, ) -> list[Any]: - schedulers: list[Any] = [] - schedulers.append(self._critic_lr_scheduler) if self._should_update_student(iteration): - schedulers.append(self._student_lr_scheduler) - return schedulers + return [self._student_lr_scheduler] + return [self._critic_lr_scheduler] - # TrainingMethod override: get_grad_clip_targets def get_grad_clip_targets( self, iteration: int, ) -> dict[str, torch.nn.Module]: - targets: dict[str, torch.nn.Module] = {} if self._should_update_student(iteration): - targets["student"] = (self.student.transformer) - targets["critic"] = self.critic.transformer - return targets + return {"student": self.student.transformer} + return {"critic": self.critic.transformer} + + def optimizers_schedulers_step(self, iteration: int) -> None: + """Prove the first critic and student Adam updates are finite FP32.""" + role = "student" if self._should_update_student(iteration) else "critic" + model = self.student if role == "student" else self.critic + optimizer = self.get_optimizers(iteration)[0] + verified = getattr(self, "_verified_optimizer_roles", set()) + first = role not in verified + probes: list[tuple[torch.Tensor, torch.Tensor]] = [] + if first: + for parameter in model.transformer.parameters(): + if not parameter.requires_grad: + continue + local = self._local_tensor(parameter).detach().reshape(-1) + if local.dtype != torch.float32: + raise RuntimeError( + f"DMD2 {role} master weights must be FP32, got {local.dtype}" + ) + if local.numel() and len(probes) < 16: + probes.append((local, local[:4096].clone())) + + super().optimizers_schedulers_step(iteration) + + if first: + changed = sum( + int(torch.count_nonzero(current[:before.numel()] != before)) + for current, before in probes + ) + moments = [ + self._local_tensor(value) + for state in optimizer.state.values() + for key, value in state.items() + if key in {"exp_avg", "exp_avg_sq"} and torch.is_tensor(value) + ] + if ( + not probes + or changed == 0 + or not moments + or any(value.dtype != torch.float32 for value in moments) + or any(not bool(torch.isfinite(value).all()) for value in moments) + or any( + not bool(torch.isfinite(current[:before.numel()]).all()) + for current, before in probes + ) + ): + raise RuntimeError(f"No finite FP32 Adam update for {role}") + rank = dist.get_rank() if dist.is_initialized() else 0 + root = Path(self.training_config.checkpoint.output_dir) + root.mkdir(parents=True, exist_ok=True) + (root / f"dmd2_update_{role}_rank{rank}.json").write_text( + json.dumps({ + "role": role, + "iteration": iteration, + "changed_probe_elements": changed, + "passed": True, + }) + "\n", + encoding="utf-8", + ) + verified.add(role) + self._verified_optimizer_roles = verified def _parse_rollout_mode(self, ) -> Literal["simulate", "data_latent"]: """Parse how DMD2 obtains the latent point used for rollout. @@ -236,6 +342,216 @@ def _parse_rollout_mode(self, ) -> Literal["simulate", "data_latent"]: "{simulate, data_latent}, got " f"{raw!r}") + def _parse_rollout_carry(self) -> tuple[bool, int, Literal["ode", "sde"]]: + """Parse the carried backward-simulation knobs. + + ``rollout_carry: true`` walks the student's own sampling grid one + rung per ``single_train_step`` call, carrying the trajectory in + memory across calls (FastGen's backward simulation): exactly one + generation forward per call instead of one full-grid walk. Off + (default) keeps the existing full-rollout behavior unchanged. + + ``rollout_carry_slots`` is the number of independent trajectory + streams per rank; it must equal + ``training.loop.gradient_accumulation_steps`` because the trainer + calls ``single_train_step`` once per accumulation round and the + slots are selected round-robin over calls. Defaults to that value. + + ``rollout_sample_type`` picks how the walk re-noises onto the next + rung: ``sde`` draws fresh noise (the existing rollout behavior), + ``ode`` reuses the noise the current state implies per modality — + the deterministic step the FastGen H3 recipe uses. + """ + raw_carry = self.method_config.get("rollout_carry", None) + if raw_carry is None: + raw_carry = False + if not isinstance(raw_carry, bool): + raise ValueError("method.rollout_carry must be a bool, got " + f"{type(raw_carry).__name__}") + carry = bool(raw_carry) + + raw_sample_type = self.method_config.get("rollout_sample_type", None) + sample_type: Literal["ode", "sde"] = "sde" + if raw_sample_type is not None: + if not isinstance(raw_sample_type, str): + raise ValueError("method.rollout_sample_type must be a " + "string, got " + f"{type(raw_sample_type).__name__}") + normalized = raw_sample_type.strip().lower() + if normalized not in ("ode", "sde"): + raise ValueError("method.rollout_sample_type must be one of " + f"{{ode, sde}}, got {raw_sample_type!r}") + if not carry: + raise ValueError("method.rollout_sample_type requires " + "method.rollout_carry: true") + sample_type = normalized # type: ignore[assignment] + + slots_raw = get_optional_int( + self.method_config, + "rollout_carry_slots", + where="method.rollout_carry_slots", + ) + if not carry: + if slots_raw is not None: + raise ValueError("method.rollout_carry_slots requires " + "method.rollout_carry: true") + return False, 0, sample_type + + if self._rollout_mode != "simulate": + raise ValueError("method.rollout_carry: true requires " + "method.rollout_mode: simulate") + + grad_accum = max( + 1, + int(self.training_config.loop.gradient_accumulation_steps or 1), + ) + slots = grad_accum if slots_raw is None else int(slots_raw) + if slots <= 0: + raise ValueError("method.rollout_carry_slots must be positive, " + f"got {slots}") + if slots != grad_accum: + raise ValueError("method.rollout_carry_slots must equal " + "training.loop.gradient_accumulation_steps, got " + f"slots={slots} vs " + f"gradient_accumulation_steps={grad_accum}") + + if sample_type == "ode" and not callable(getattr(self.student, "extract_eps", None)): + raise ValueError("method.rollout_sample_type: ode requires the " + "student model to implement " + "extract_eps(noisy_latents, clean_latents, " + "timestep)") + + _, stagger_groups = self._rollout_carry_rank_world() + if bool(getattr(getattr(self.training_config, "data", None), "native_shape_bucketing", False)): + stagger_groups = 1 + self._validate_rollout_carry_coverage( + streams=stagger_groups * slots, + grid_len=self._rollout_grid_length(), + interval=self._generator_update_interval(), + ) + return True, slots, sample_type + + @staticmethod + def _validate_rollout_carry_coverage( + *, + streams: int, + grid_len: int, + interval: int, + ) -> None: + """Reject configurations that leave student rungs untrained. + + Student updates revisit a stream's rung modulo + ``gcd(len(dmd_denoising_steps), generator_update_interval)``, so + every residue class must be represented by a (rank, slot) stream; + the consecutive stagger offsets cover all classes exactly when + there are at least ``gcd`` streams. + """ + phase_classes = math.gcd(grid_len, interval) + if streams < phase_classes: + raise ValueError("Carried backward-simulation DMD2 cannot cover every student " + f"rung with len(dmd_denoising_steps)={grid_len}, " + f"generator_update_interval={interval}, and " + f"{streams} trajectory stream(s) (stagger groups x slots). " + "Student updates preserve the rung modulo " + f"gcd(grid, interval)={phase_classes}, but only {streams} " + "stream phase(s) are present. Use at least that many " + "rank-slot streams or choose a coprime update interval.") + + def _rollout_carry_rank_world(self) -> tuple[int, int]: + """Stagger rank and stream-group count for the carried rollout. + + Mirrors ``TrainingMethod.on_train_start``'s RNG grouping: ranks + inside one sequence-parallel group shard the same document and must + walk one shared trajectory, so they share a stagger rank (with + ``sp_size=1`` this is exactly the global rank). Falls back to a + single group when the distributed world is not initialized (CPU + tests, single-process runs). + """ + try: + from fastvideo.distributed import get_world_group + world_group = get_world_group() + global_rank = int(world_group.rank) + world_size = int(world_group.world_size) + except (AssertionError, ImportError, RuntimeError): + global_rank, world_size = 0, 1 + sp_size = max( + 1, + int(getattr(self.training_config.distributed, "sp_size", 1) or 1), + ) + return global_rank // sp_size, max(1, world_size // sp_size) + + def _parse_rollout_data_forcing(self) -> bool: + """Parse per-batch data forcing for the carried walk. + + ``rollout_data_forcing: true`` routes latent-bearing batches (t2va + parquet rows in a mixed ``data_path``) onto FastGen's data-driven + student inputs — the real packed latents forward-noised at a + uniformly drawn grid rung (``sample_from_t_list`` semantics) — + while text-only batches keep walking the carried backward + simulation. Off (default) keeps every batch on the walk, + byte-identical to the carry-only behavior. + """ + raw = self.method_config.get("rollout_data_forcing", None) + if raw is None: + return False + if not isinstance(raw, bool): + raise ValueError("method.rollout_data_forcing must be a bool, " + f"got {type(raw).__name__}") + if raw and not self._rollout_carry: + raise ValueError("method.rollout_data_forcing: true requires " + "method.rollout_carry: true; without the carry, " + "always-forced inputs are " + "method.rollout_mode: data_latent") + if raw: + allow_mixed = self.method_config.get( + "allow_mixed_rollout_regimes", + False, + ) + if not isinstance(allow_mixed, bool): + raise ValueError("method.allow_mixed_rollout_regimes must be a bool, " + f"got {type(allow_mixed).__name__}") + if not allow_mixed: + raise ValueError("method.rollout_data_forcing mixes carried and data-latent " + "rollout regimes per batch, which is not a FastGen recipe. " + "Choose one global regime with rollout_mode={simulate, " + "data_latent}; set allow_mixed_rollout_regimes: true only " + "to reproduce the legacy v9 experiment.") + return raw + + @staticmethod + def _batch_has_latents(batch: dict[str, Any]) -> bool: + """Classify a mixed-loading batch as latent-bearing or text-only. + + Under the t2va parquet schema the collate emits empty (numel-0) + tensors for latent columns a text-only row does not carry, so + presence means "key exists and non-empty". A row carrying exactly + one of the pair is corrupt data, not a batch type. + """ + video = batch.get("vae_latent") + audio = batch.get("audio_latent") + has_video = isinstance(video, torch.Tensor) and video.numel() > 0 + has_audio = isinstance(audio, torch.Tensor) and audio.numel() > 0 + if has_video != has_audio: + raise ValueError("Mixed-loading batch carries exactly one of " + "vae_latent/audio_latent non-empty; a t2va row " + "must carry both and a text_only row neither " + f"(vae_latent={'present' if has_video else 'empty/missing'}, " + f"audio_latent={'present' if has_audio else 'empty/missing'})") + return has_video + + def _rollout_grid_length(self) -> int: + raw = self.method_config.get("dmd_denoising_steps", None) + if not isinstance(raw, list) or not raw: + raise ValueError("method_config.dmd_denoising_steps must " + "be set for DMD2 distillation") + return len(raw) + + def _init_rollout_carry_state(self) -> None: + slots = max(0, int(self._rollout_carry_slot_count)) + self._carry_call_count = 0 + self._carry_slots: list[dict[str, Any] | None] = [None] * slots + self._carry_slot_seeded: list[bool] = [False] * slots + def _validate_preprocessed_data_type(self) -> None: data_type = str(getattr( self.training_config.data, @@ -246,6 +562,13 @@ def _validate_preprocessed_data_type(self) -> None: raise ValueError("training.data.preprocessed_data_type='text_only' " "requires method.rollout_mode='simulate'; " "data_latent rollout requires vae_latent data.") + if self._rollout_data_forcing and data_type != "t2va": + raise ValueError("method.rollout_data_forcing: true requires " + "training.data.preprocessed_data_type='t2va': the " + "t2va parquet schema is the superset that reads " + "latent columns; text-only roots mixed into the " + "same data_path yield empty latent columns and " + "route to the carried walk.") def _uses_negative_prompt_conditioning(self) -> bool: if self._cfg_uncond is None: @@ -314,7 +637,6 @@ def _parse_cfg_uncond(self, ) -> dict[str, Any] | None: def _init_optimizers_and_schedulers(self) -> None: tc = self.training_config - # Student optimizer/scheduler. student_lr = float(tc.optimizer.learning_rate) student_betas = tc.optimizer.betas student_sched = str(tc.optimizer.lr_scheduler) @@ -331,8 +653,6 @@ def _init_optimizers_and_schedulers(self) -> None: scheduler_name=student_sched, ) - # Critic optimizer/scheduler — must be set in - # method config. critic_lr_raw = get_optional_float( self.method_config, "fake_score_learning_rate", @@ -370,20 +690,87 @@ def _init_optimizers_and_schedulers(self) -> None: scheduler_name=critic_sched, ) - def _should_update_student( - self, - iteration: int, - ) -> bool: + @staticmethod + def _add_noise_for_batch( + model: ModelBase, + clean_latents: torch.Tensor, + noise: torch.Tensor, + timestep: torch.Tensor, + batch: Any, + ) -> torch.Tensor: + """Call the batch-aware hook while retaining lightweight test doubles.""" + hook = getattr(model, "add_noise_for_batch", None) + if hook is not None and batch is not None: + return hook(clean_latents, noise, timestep, batch) + return model.add_noise(clean_latents, noise, timestep) + + @staticmethod + def _extract_eps_for_batch( + model: ModelBase, + noisy_latents: torch.Tensor, + clean_latents: torch.Tensor, + timestep: torch.Tensor, + batch: Any, + ) -> torch.Tensor: + hook = getattr(model, "extract_eps_for_batch", None) + if hook is not None and batch is not None: + return hook(noisy_latents, clean_latents, timestep, batch) + extractor = getattr(model, "extract_eps", None) + if not callable(extractor): + raise TypeError(f"{type(model).__name__} does not implement extract_eps") + return extractor(noisy_latents, clean_latents, timestep) + + def _modality_slices(self, batch: Any) -> tuple[tuple[str, slice], ...] | None: + """Return slices used to normalize packed modalities independently.""" + batch_getter = getattr(self.student, "modality_slices_for_batch", None) + if batch_getter is not None: + slices = tuple(batch_getter(batch)) + return slices or None + getter = getattr(self.student, "modality_slices", None) + if getter is None: + return None + slices = tuple(getter()) + return slices or None + + def _modality_weight(self, name: str) -> float: + raw = self.method_config.get("modality_loss_weights", None) + if not isinstance(raw, dict): + return 1.0 + value = raw.get(name, 1.0) + return 1.0 if value is None else float(value) + + def apply_configured_lrs(self) -> None: + """Force student/critic LRs back to the configured values (post-resume).""" + student_lr = float(self.training_config.optimizer.learning_rate) + critic_lr = float(self.method_config.get("fake_score_learning_rate")) + for optimizer, scheduler, lr in ( + (self._student_optimizer, self._student_lr_scheduler, student_lr), + (self._critic_optimizer, self._critic_lr_scheduler, critic_lr), + ): + for group in optimizer.param_groups: + group["lr"] = lr + if "initial_lr" in group: + group["initial_lr"] = lr + if hasattr(scheduler, "base_lrs"): + scheduler.base_lrs = [lr] * len(scheduler.base_lrs) + + def _generator_update_interval(self) -> int: interval = get_optional_int( self.method_config, "generator_update_interval", where="method.generator_update_interval", ) if interval is None: - interval = 1 + interval = 5 if interval <= 0: - return True - return iteration % interval == 0 + raise ValueError("method.generator_update_interval must be positive") + return interval + + def _should_update_student( + self, + iteration: int, + ) -> bool: + return iteration % self._generator_update_interval() == 0 def _get_denoising_step_list( self, @@ -432,7 +819,7 @@ def _sample_rollout_timestep( return step_list[index] def _parse_score_timestep_bounds(self) -> tuple[int, int]: - """Resolve the score-model timestep window used by legacy DMD. + """Resolve the score-model timestep window. The student rollout schedule is controlled separately by ``dmd_denoising_steps``. These bounds apply only to the randomly @@ -461,15 +848,127 @@ def _parse_score_timestep_bounds(self) -> tuple[int, int]: int(max_ratio * num_timesteps), ) - def _sample_score_timestep(self, device: torch.device) -> torch.Tensor: - timestep = torch.randint( - 0, - int(self.student.num_train_timesteps), - [1], - device=device, - dtype=torch.long, - generator=self.cuda_generator, + def _parse_score_timestep_shift(self) -> float: + """Resolve the rational warp used by score-time sampling. + + Legacy integer sampling draws uniformly in the warped coordinate and + inverts it. Continuous FastGen parity draws the pre-warp coordinate + directly and applies this inverse warp before the model adapter adds + its modality-specific clock. + """ + shift = get_optional_float( + self.method_config, + "score_timestep_shift", + where="method.score_timestep_shift", ) + shift = 1.0 if shift is None else float(shift) + if shift <= 0.0: + raise ValueError("method.score_timestep_shift must be > 0, " + f"got {shift}") + return shift + + def _parse_score_timestep_warp_max(self) -> float: + """Resolve the endpoint used by the continuous rational time warp.""" + warp_max = get_optional_float( + self.method_config, + "score_timestep_warp_max", + where="method.score_timestep_warp_max", + ) + warp_max = 1.0 if warp_max is None else float(warp_max) + if not 0.0 < warp_max <= 1.0: + raise ValueError("method.score_timestep_warp_max must satisfy " + f"0 < max <= 1, got {warp_max}") + max_ratio = self._score_max_timestep / float(self.student.num_train_timesteps) + if max_ratio > warp_max: + raise ValueError("method.max_timestep_ratio must not exceed " + "method.score_timestep_warp_max, got " + f"{max_ratio} > {warp_max}") + return warp_max + + def _parse_score_timestep_continuous(self) -> bool: + """Select FastGen-style continuous score times instead of integer bins.""" + raw = self.method_config.get("score_timestep_continuous", False) + if not isinstance(raw, bool): + raise ValueError("method.score_timestep_continuous must be a bool, " + f"got {type(raw).__name__}") + return raw + + def _parse_fake_score_loss_space(self) -> dict[str, str]: + """Resolve the critic regression space, globally or per modality. + + ``velocity`` is plain velocity MSE. A global ``x0`` setting calls the + critic's x0 prediction directly, matching FastGen without estimating + sigma from rounded latents. Legacy mixed mappings retain the original + single-forward sigma-squared conversion for their x0 modalities. + """ + raw = self.method_config.get("fake_score_loss_space", None) + if raw is None: + return {"__default__": "velocity"} + if isinstance(raw, str): + mapping = {"__default__": raw} + elif isinstance(raw, dict): + mapping = {str(k).strip().lower(): str(v) for k, v in raw.items()} + mapping.setdefault("__default__", "velocity") + else: + raise ValueError("method.fake_score_loss_space must be a string " + "or a {modality: space} mapping, got " + f"{type(raw).__name__}") + normalized: dict[str, str] = {} + for key, value in mapping.items(): + space = str(value).strip().lower() + if space not in ("velocity", "x0"): + raise ValueError("method.fake_score_loss_space values must be " + f"one of {{velocity, x0}}, got {value!r} " + f"for {key!r}") + normalized[key] = space + return normalized + + def _fake_score_space_for(self, modality_name: str) -> str: + return self._fake_score_loss_space.get( + modality_name, + self._fake_score_loss_space["__default__"], + ) + + def _sample_score_timestep(self, device: torch.device) -> torch.Tensor: + shift = self._score_timestep_shift + num_timesteps = float(self.student.num_train_timesteps) + t_lo = self._score_min_timestep / num_timesteps + t_hi = self._score_max_timestep / num_timesteps + + if getattr(self, "_score_timestep_continuous", False): + u = torch.rand( + [1], + device=device, + dtype=torch.float64, + generator=self.cuda_generator, + ) * (t_hi - t_lo) + t_lo + inverse_shift = 1.0 / shift + warp_max = getattr(self, "_score_timestep_warp_max", 1.0) + t = (u * inverse_shift * warp_max / (u * (inverse_shift - 1.0) + warp_max)) + timestep = t * num_timesteps + timestep = self.student.shift_and_clamp_timestep(timestep) + return timestep.clamp(0.0, warp_max * num_timesteps) + + if shift == 1.0: + timestep = torch.randint( + self._score_min_timestep, + self._score_max_timestep + 1, + [1], + device=device, + dtype=torch.long, + generator=self.cuda_generator, + ) + else: + sigma_lo = shift * t_lo / (1.0 + (shift - 1.0) * t_lo) + sigma_hi = shift * t_hi / (1.0 + (shift - 1.0) * t_hi) + u = torch.rand( + [1], + device=device, + dtype=torch.float32, + generator=self.cuda_generator, + ) * (sigma_hi - sigma_lo) + sigma_lo + t = u / (shift - (shift - 1.0) * u) + timestep = (t * num_timesteps).round().to(torch.long) timestep = self.student.shift_and_clamp_timestep(timestep) return timestep.clamp( self._score_min_timestep, @@ -495,7 +994,7 @@ def _student_rollout( dtype=dtype, generator=self.cuda_generator, ) - noisy_latents = self.student.add_noise(latents, noise, timestep) + noisy_latents = self._add_noise_for_batch(self.student, latents, noise, timestep, batch) pred_x0 = self.student.predict_x0( noisy_latents, timestep, @@ -561,11 +1060,13 @@ def _student_rollout( dtype=pred_clean.dtype, generator=self.cuda_generator, ) - current_noise_latents = (self.student.add_noise( + current_noise_latents = self._add_noise_for_batch( + self.student, pred_clean, noise, next_timestep_tensor, - )) + batch, + ) noise_latents.append(current_noise_latents.clone()) if noise_latent_index >= 0: @@ -598,12 +1099,455 @@ def _student_rollout( batch.dmd_latent_vis_dict["generator_timestep"] = target_timestep.float().detach() return pred_x0 + + def _carried_train_step( + self, + batch: dict[str, Any], + iteration: int, + ) -> tuple[ + dict[str, torch.Tensor], + dict[str, Any], + dict[str, LogScalar], + ]: + """One backward-simulation call: one generation forward, carried state. + + A multistep student is only ever correct on its own sampling + trajectory, and walking the full grid every call costs + ``len(dmd_denoising_steps)`` forwards. Instead the walk is spread + over consecutive calls: each call pays for exactly one student + forward at the carried rung, both phases (student and critic) + consume it — the critic is fit on the same simulated states the + student trains on — and both advance the trajectory. Each + grad-accum round owns an independent slot, selected round-robin + because the trainer does not pass the round index. + + An empty slot (first ever use, cleared after a finished trajectory, + or after a resume — the carry is transient and never checkpointed) + starts a fresh trajectory from noise and adopts the incoming loader + batch's conditioning; mid-walk calls ignore the fresh loader batch + and rebuild the training batch from the carried raw batch, since a + trajectory keeps the prompt it set out with. + """ + slot = self._carry_call_count % self._rollout_carry_slot_count + self._carry_call_count += 1 + + if self._rollout_data_forcing and self._batch_has_latents(batch): + return self._data_forced_train_step(batch, slot, iteration) + + carried = self._carry_slots[slot] + raw_batch = (self._carry_snapshot_raw_batch(batch) if carried is None else carried["raw_batch"]) + + training_batch = self.student.prepare_batch( + raw_batch, + generator=self.cuda_generator, + latents_source="zeros", + ) + latents = training_batch.latents + device = latents.device + step_list = self._get_denoising_step_list(device) + + if carried is None: + rung = 0 + state = torch.randn( + latents.shape, + device=device, + dtype=latents.dtype, + generator=self.cuda_generator, + ) + if not self._carry_slot_seeded[slot]: + self._carry_slot_seeded[slot] = True + state, rung = self._staggered_start( + state, + training_batch, + step_list, + slot, + ) + else: + rung = int(carried["rung"]) + state = carried["state"] + + timestep = step_list[rung] * torch.ones( + 1, + device=device, + dtype=torch.long, + ) + + update_student = self._should_update_student(iteration) + + generator_loss = torch.zeros((), device=device, dtype=torch.float32) + fake_score_loss = torch.zeros_like(generator_loss) + student_ctx = None + critic_ctx = None + critic_outputs: dict[str, Any] = {} + generator_metrics: dict[str, LogScalar] = {} + critic_metrics: dict[str, LogScalar] = {} + if update_student: + generator_pred_x0 = self.student.predict_x0( + state, + timestep, + training_batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="vsa", + ) + student_ctx = ( + training_batch.timesteps, + training_batch.attn_metadata_vsa, + ) + generator_loss, generator_metrics = self._dmd_loss(generator_pred_x0, training_batch) + training_batch.dmd_latent_vis_dict["generator_pred_video"] = generator_pred_x0.detach() + else: + with torch.no_grad(): + generator_pred_x0 = self.student.predict_x0( + state, + timestep, + training_batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="vsa", + ) + ( + fake_score_loss, + critic_ctx, + critic_outputs, + critic_metrics, + ) = self._critic_flow_matching_loss( + training_batch, + generator_pred_x0=generator_pred_x0, + ) + training_batch.dmd_latent_vis_dict["generator_timestep"] = timestep.float().detach() + + self._advance_carry( + slot, + state, + generator_pred_x0, + timestep, + rung, + step_list, + training_batch, + raw_batch, + ) + + total_loss = generator_loss + fake_score_loss + loss_map = { + "total_loss": total_loss, + "generator_loss": generator_loss, + "fake_score_loss": fake_score_loss, + } + outputs: dict[str, Any] = dict(critic_outputs) + outputs["_fv_backward"] = { + "update_student": update_student, + "student_ctx": student_ctx, + "critic_ctx": critic_ctx, + } + metrics: dict[str, LogScalar] = { + "update_student": float(update_student), + "rollout_step": float(rung), + **generator_metrics, + **critic_metrics, + } + if self._rollout_data_forcing: + metrics["data_forced"] = 0.0 + self.latent_vis = { + **(training_batch.fake_score_latent_vis_dict or {}), + **(training_batch.dmd_latent_vis_dict or {}), + "_fv_latent_layout": + getattr(training_batch, "minimax_h3_dmd_layout", None), + } + return loss_map, outputs, metrics + + def _data_forced_train_step( + self, + batch: dict[str, Any], + slot: int, + iteration: int, + ) -> tuple[ + dict[str, torch.Tensor], + dict[str, Any], + dict[str, LogScalar], + ]: + """One data-forced call: train on real latents noised at a grid rung. + + FastGen's data-driven multistep student inputs (its + ``backward_simulation: false`` regime): ``t_student`` is drawn + uniformly over the student grid's rungs — ``sample_from_t_list`` + semantics, never t=0 — and the real packed latents are + forward-noised to that rung under each modality's shift, exactly + the uncarried ``rollout_mode: data_latent`` math. The slot's + carried walk pauses untouched and resumes on this stream's next + text-only batch: FastGen picks one regime per config, so pausing + is the minimal per-batch composition of its two modes. Both + phases consume the same forced generation, mirroring the carried + step's critic passthrough. + + The slot's one-time stagger pre-walk still runs on its first-ever + call even when that call is data-forced: the pre-walk's FSDP + collective count must stay uniform across ranks, and ranks whose + first batch is text-only run theirs on this same call. The seeded + walk adopts this batch's conditioning and waits at its stagger + rung. + """ + training_batch = self.student.prepare_batch( + batch, + generator=self.cuda_generator, + latents_source="data", + ) + latents = training_batch.latents + device = latents.device + if not self._carry_slot_seeded[slot]: + self._carry_slot_seeded[slot] = True + step_list = self._get_denoising_step_list(device) + state = torch.randn( + latents.shape, + device=device, + dtype=latents.dtype, + generator=self.cuda_generator, + ) + state, rung = self._staggered_start( + state, + training_batch, + step_list, + slot, + ) + self._carry_slots[slot] = { + "state": state.detach(), + "rung": rung, + "raw_batch": self._carry_snapshot_raw_batch(batch), + } + + forced_timestep = self._sample_rollout_timestep(device) + noise = torch.randn( + latents.shape, + device=device, + dtype=latents.dtype, + generator=self.cuda_generator, + ) + noisy_latents = self._add_noise_for_batch(self.student, latents, noise, forced_timestep, training_batch) + + update_student = self._should_update_student(iteration) + + generator_loss = torch.zeros((), device=device, dtype=torch.float32) + fake_score_loss = torch.zeros_like(generator_loss) + student_ctx = None + critic_ctx = None + critic_outputs: dict[str, Any] = {} + generator_metrics: dict[str, LogScalar] = {} + critic_metrics: dict[str, LogScalar] = {} + if update_student: + generator_pred_x0 = self.student.predict_x0( + noisy_latents, + forced_timestep, + training_batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="vsa", + ) + student_ctx = ( + training_batch.timesteps, + training_batch.attn_metadata_vsa, + ) + generator_loss, generator_metrics = self._dmd_loss(generator_pred_x0, training_batch) + training_batch.dmd_latent_vis_dict["generator_pred_video"] = generator_pred_x0.detach() + else: + with torch.no_grad(): + generator_pred_x0 = self.student.predict_x0( + noisy_latents, + forced_timestep, + training_batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="vsa", + ) + ( + fake_score_loss, + critic_ctx, + critic_outputs, + critic_metrics, + ) = self._critic_flow_matching_loss( + training_batch, + generator_pred_x0=generator_pred_x0, + ) + training_batch.dmd_latent_vis_dict["generator_timestep"] = forced_timestep.float().detach() + + total_loss = generator_loss + fake_score_loss + loss_map = { + "total_loss": total_loss, + "generator_loss": generator_loss, + "fake_score_loss": fake_score_loss, + } + outputs: dict[str, Any] = dict(critic_outputs) + outputs["_fv_backward"] = { + "update_student": update_student, + "student_ctx": student_ctx, + "critic_ctx": critic_ctx, + } + metrics: dict[str, LogScalar] = { + "update_student": float(update_student), + "data_forced": 1.0, + **generator_metrics, + **critic_metrics, + } + self.latent_vis = { + **(training_batch.fake_score_latent_vis_dict or {}), + **(training_batch.dmd_latent_vis_dict or {}), + "_fv_latent_layout": + getattr(training_batch, "minimax_h3_dmd_layout", None), + } + return loss_map, outputs, metrics + + def _carry_snapshot_raw_batch( + self, + batch: dict[str, Any], + ) -> dict[str, Any]: + """Adopt the incoming loader batch as a trajectory's conditioning. + + A trajectory keeps the prompt it set out with for its whole walk, + so the raw dict is snapshotted (tensors detached and kept on the + student device) and mid-walk calls rebuild the training batch from + it via ``prepare_batch``; H3's ``prepare_batch`` reads the dict + without mutating it and rebuilds the packed layout and VSA + attention metadata deterministically on every call. + """ + device = self.student.device + snapshot: dict[str, Any] = {} + for key, value in batch.items(): + if isinstance(value, torch.Tensor): + snapshot[key] = value.detach().to(device) + else: + snapshot[key] = value + return snapshot + + def _staggered_start( + self, + state: torch.Tensor, + training_batch: Any, + step_list: torch.Tensor, + slot: int, + ) -> tuple[torch.Tensor, int]: + """Pre-walk a fresh trajectory and snapshot it at this stream's rung. + + First-ever fill of a slot only. Each (rank, slot) stream starts at + ``(stagger_rank * slots + slot) % len(grid)``, spreading the + streams evenly across the grid so student updates (which revisit + rungs modulo ``gcd(len(grid), generator_update_interval)``) see + every rung. Every rank walks the whole grid under ``no_grad`` + regardless of its offset so the FSDP forwards issue a uniform + collective count — a rank-dependent count would desynchronize the + all-gathers and hang; only the kept snapshot differs per rank. + """ + rank, _ = self._rollout_carry_rank_world() + grid_len = len(step_list) + data_config = getattr(self.training_config, "data", None) + stagger_rank = 0 if bool(getattr(data_config, "native_shape_bucketing", False)) else rank + offset = (stagger_rank * self._rollout_carry_slot_count + slot) % grid_len + device = state.device + snapshot = state + with torch.no_grad(): + for rung in range(grid_len - 1): + timestep = step_list[rung] * torch.ones( + 1, + device=device, + dtype=torch.long, + ) + pred_x0 = self.student.predict_x0( + state, + timestep, + training_batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="vsa", + ) + state = self._renoise( + state, + pred_x0.detach(), + timestep, + rung + 1, + step_list, + training_batch, + ) + if rung + 1 == offset: + snapshot = state + return snapshot, offset + + def _renoise( + self, + state: torch.Tensor, + pred_x0: torch.Tensor, + timestep: torch.Tensor, + next_rung: int, + step_list: torch.Tensor, + batch: Any | None = None, + ) -> torch.Tensor: + """Re-noise an x0 prediction made at ``timestep`` onto the next rung. + + ``sde`` draws fresh noise — the existing full-rollout hop. + ``ode`` reuses the noise the current state implies per modality + (``eps_m = (x_t - alpha_m(t) x0) / sigma_m(t)`` with each + modality's shifted sigma, via the adapter's ``extract_eps``), the + deterministic step the FastGen H3 recipe uses. + """ + device = state.device + next_timestep = step_list[next_rung] * torch.ones( + 1, + device=device, + dtype=torch.long, + ) + if self._rollout_sample_type == "ode": + eps = self._extract_eps_for_batch(self.student, state, pred_x0, timestep, batch) + else: + eps = torch.randn( + state.shape, + device=device, + dtype=pred_x0.dtype, + generator=self.cuda_generator, + ) + return self._add_noise_for_batch(self.student, pred_x0, eps, next_timestep, batch) + + def _advance_carry( + self, + slot: int, + state: torch.Tensor, + generator_pred_x0: torch.Tensor, + timestep: torch.Tensor, + rung: int, + step_list: torch.Tensor, + training_batch: Any, + raw_batch: dict[str, Any], + ) -> None: + """Hand the one paid-for step to the slot, or clear a finished walk. + + The advanced state is detached and produced under ``no_grad``: it + feeds a later call, not a gradient path. Walking past the last rung + ends the trajectory (the terminal clean sample is never trained + on), so the slot empties and the next call starts fresh at rung 0. + """ + if rung + 1 >= len(step_list): + self._carry_slots[slot] = None + return + with torch.no_grad(): + next_state = self._renoise( + state, + generator_pred_x0.detach(), + timestep, + rung + 1, + step_list, + training_batch, + ) + self._carry_slots[slot] = { + "state": next_state.detach(), + "rung": rung + 1, + "raw_batch": raw_batch, + } + def _critic_flow_matching_loss( self, batch: Any, - ) -> tuple[torch.Tensor, Any, dict[str, Any]]: - with torch.no_grad(): - generator_pred_x0 = self._student_rollout(batch, with_grad=False) + *, + generator_pred_x0: torch.Tensor | None = None, + ) -> tuple[torch.Tensor, Any, dict[str, Any], dict[str, LogScalar]]: + if generator_pred_x0 is None: + with torch.no_grad(): + generator_pred_x0 = self._student_rollout(batch, with_grad=False) device = generator_pred_x0.device fake_score_timestep = self._sample_score_timestep(device) @@ -614,18 +1558,62 @@ def _critic_flow_matching_loss( dtype=generator_pred_x0.dtype, generator=self.cuda_generator, ) - noisy_x0 = self.student.add_noise(generator_pred_x0, noise, fake_score_timestep) - - pred_noise = self.critic.predict_noise( - noisy_x0, + noisy_x0 = self._add_noise_for_batch( + self.student, + generator_pred_x0, + noise, fake_score_timestep, batch, - conditional=True, - cfg_uncond=self._cfg_uncond, - attn_kind="dense", ) - target = noise - generator_pred_x0 - flow_matching_loss = torch.mean((pred_noise - target)**2) + + slices = self._modality_slices(batch) + emit_modality_metrics = slices is not None + if slices is None: + slices = (("packed", slice(None)), ) + all_x0 = all(self._fake_score_space_for(name) == "x0" for name, _ in slices) + + pred_x0: torch.Tensor | None = None + pred_noise: torch.Tensor | None = None + target: torch.Tensor | None = None + if all_x0: + pred_x0 = self.critic.predict_x0( + noisy_x0, + fake_score_timestep, + batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="dense", + ) + else: + pred_noise = self.critic.predict_noise( + noisy_x0, + fake_score_timestep, + batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="dense", + ) + target = noise - generator_pred_x0 + + flow_matching_loss = torch.zeros((), device=device, dtype=torch.float32) + metrics: dict[str, LogScalar] = {} + for name, modality in slices: + if all_x0: + assert pred_x0 is not None + loss_m = torch.mean((pred_x0[:, modality].float() - generator_pred_x0[:, modality].float())**2) + else: + assert pred_noise is not None and target is not None + loss_m = torch.mean((pred_noise[:, modality].float() - target[:, modality].float())**2) + if not all_x0 and self._fake_score_space_for(name) == "x0": + assert target is not None + with torch.no_grad(): + num = torch.mean((noisy_x0[:, modality].float() - generator_pred_x0[:, modality].float())**2) + den = torch.mean(target[:, modality].float()**2) + sigma_sq = num / den + loss_m = sigma_sq * loss_m + flow_matching_loss = flow_matching_loss + self._modality_weight(name) * loss_m + if emit_modality_metrics: + metrics[f"fake_score_loss_{name}"] = loss_m.detach() batch.fake_score_latent_vis_dict = { "generator_pred_video": generator_pred_x0, @@ -636,13 +1624,14 @@ def _critic_flow_matching_loss( flow_matching_loss, (batch.timesteps, batch.attn_metadata), outputs, + metrics, ) def _dmd_loss( self, generator_pred_x0: torch.Tensor, batch: Any, - ) -> torch.Tensor: + ) -> tuple[torch.Tensor, dict[str, LogScalar]]: guidance_scale = get_optional_float( self.method_config, "real_score_guidance_scale", @@ -661,7 +1650,13 @@ def _dmd_loss( dtype=generator_pred_x0.dtype, generator=self.cuda_generator, ) - noisy_latents = self.student.add_noise(generator_pred_x0, noise, timestep) + noisy_latents = self._add_noise_for_batch( + self.student, + generator_pred_x0, + noise, + timestep, + batch, + ) faker_x0 = self.critic.predict_x0( noisy_latents, @@ -679,22 +1674,43 @@ def _dmd_loss( cfg_uncond=self._cfg_uncond, attn_kind="dense", ) - real_uncond_x0 = self.teacher.predict_x0( - noisy_latents, - timestep, - batch, - conditional=False, - cfg_uncond=self._cfg_uncond, - attn_kind="dense", - ) - real_cfg_x0 = real_uncond_x0 + (real_cond_x0 - real_uncond_x0) * guidance_scale - - denom = torch.abs(generator_pred_x0 - real_cfg_x0).mean() - grad = (faker_x0 - real_cfg_x0) / denom - grad = torch.nan_to_num(grad) - - loss = 0.5 * F.mse_loss( - generator_pred_x0.float(), - (generator_pred_x0.float() - grad.float()).detach(), - ) - return loss + if float(guidance_scale) == 1.0: + real_cfg_x0 = real_cond_x0 + else: + real_uncond_x0 = self.teacher.predict_x0( + noisy_latents, + timestep, + batch, + conditional=False, + cfg_uncond=self._cfg_uncond, + attn_kind="dense", + ) + real_cfg_x0 = real_uncond_x0 + (real_cond_x0 - real_uncond_x0) * guidance_scale + + batch.dmd_latent_vis_dict.update({ + "real_score_pred_video": real_cfg_x0.detach(), + "faker_score_pred_video": faker_x0.detach(), + "dmd_timestep": timestep.detach(), + }) + + slices = self._modality_slices(batch) + emit_modality_metrics = slices is not None + if slices is None: + slices = (("packed", slice(None)), ) + loss = torch.zeros((), device=device, dtype=torch.float32) + metrics: dict[str, LogScalar] = {} + for name, modality in slices: + gen_m = generator_pred_x0[:, modality].float() + with torch.no_grad(): + real_m = real_cfg_x0[:, modality].float() + denom = (gen_m - real_m).abs().mean() + 1e-6 + grad = (faker_x0[:, modality].float() - real_m) / denom + if not bool(torch.isfinite(grad).all()): + raise RuntimeError( + f"Nonfinite DMD2 distribution direction for {name}" + ) + loss_m = 0.5 * F.mse_loss(gen_m, (gen_m - grad).detach()) + loss = loss + self._modality_weight(name) * loss_m + if emit_modality_metrics: + metrics[f"generator_loss_{name}"] = loss_m.detach() + return loss, metrics diff --git a/fastvideo/train/models/base.py b/fastvideo/train/models/base.py index 4ef95a3dcd..81ce682a6a 100644 --- a/fastvideo/train/models/base.py +++ b/fastvideo/train/models/base.py @@ -18,8 +18,6 @@ from fastvideo.train.utils.lora import LoraConfig from fastvideo.pipelines import TrainingBatch -# Video models return one flow tensor. Joint video/audio models return an -# ordered pair so training methods can apply each modality's scheduler target. NoisePrediction: TypeAlias = torch.Tensor | tuple[torch.Tensor, torch.Tensor] @@ -89,9 +87,6 @@ def _enable_lora_if_configured( ) return True - # ------------------------------------------------------------------ - # Lifecycle - # ------------------------------------------------------------------ def init_preprocessors( # noqa: B027 self, @@ -117,9 +112,6 @@ def decode_latents( """ raise NotImplementedError(f"{type(self).__name__} does not implement decode_latents()") - # ------------------------------------------------------------------ - # Timestep helpers - # ------------------------------------------------------------------ @property def num_train_timesteps(self) -> int: @@ -130,9 +122,6 @@ def shift_and_clamp_timestep(self, timestep: torch.Tensor) -> torch.Tensor: """Apply model/pipeline timestep shifting and clamp.""" return timestep - # ------------------------------------------------------------------ - # Runtime primitives - # ------------------------------------------------------------------ @abstractmethod def prepare_batch( @@ -153,6 +142,22 @@ def add_noise( ) -> torch.Tensor: """Apply forward-process noise at *timestep*.""" + def add_noise_for_batch( + self, + clean_latents: torch.Tensor, + noise: torch.Tensor, + timestep: torch.Tensor, + batch: TrainingBatch, + ) -> torch.Tensor: + """Apply noise with optional immutable geometry carried by ``batch``. + + Video-only and fixed-shape models keep their existing behavior. Joint + packed models override this hook when splitting the tensor requires + batch-local shape context. + """ + del batch + return self.add_noise(clean_latents, noise, timestep) + @abstractmethod def predict_noise( self, diff --git a/fastvideo/train/models/minimax_h3/__init__.py b/fastvideo/train/models/minimax_h3/__init__.py index acbdc51ae2..006658f47c 100644 --- a/fastvideo/train/models/minimax_h3/__init__.py +++ b/fastvideo/train/models/minimax_h3/__init__.py @@ -3,3 +3,7 @@ from fastvideo.train.models.minimax_h3.minimax_h3 import ( MiniMaxH3Model as MiniMaxH3Model, ) +from fastvideo.train.models.minimax_h3.minimax_h3_dmd import ( + MiniMaxH3DMDLatentLayout as MiniMaxH3DMDLatentLayout, + MiniMaxH3DMDModel as MiniMaxH3DMDModel, +) diff --git a/fastvideo/train/models/minimax_h3/minimax_h3.py b/fastvideo/train/models/minimax_h3/minimax_h3.py index 0788880182..149df14076 100644 --- a/fastvideo/train/models/minimax_h3/minimax_h3.py +++ b/fastvideo/train/models/minimax_h3/minimax_h3.py @@ -3,16 +3,22 @@ from __future__ import annotations +import math from typing import Any, Literal, TYPE_CHECKING import torch +import fastvideo.envs as envs +from fastvideo.attention.backends.video_sparse_attn_h3 import ( + MiniMaxH3VSAMetadataBuilder, ) +from fastvideo.dataset.shape_bucket import parse_video_shape_bucket_id from fastvideo.distributed import get_sp_group from fastvideo.forward_context import set_forward_context from fastvideo.models.schedulers.scheduling_minimax_h3 import MiniMaxH3Scheduler from fastvideo.pipelines import TrainingBatch from fastvideo.pipelines.basic.minimax_h3.packing import ( MINIMAX_H3_AUDIO_CHANNELS, + MINIMAX_H3_CANVAS_MULTIPLE, MINIMAX_H3_TEXT_TAG, MiniMaxH3PackedLayout, audio_latent_num_frames, @@ -21,7 +27,10 @@ patchify_video_latents, unpack_audio_tokens, unpatchify_video_tokens, + video_latent_num_frames, ) +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_denoising import ( + _h3_vsa_prefix_segments, ) from fastvideo.platforms import AttentionBackendEnum from fastvideo.train.models.base import ModelBase, NoisePrediction @@ -32,19 +41,32 @@ if TYPE_CHECKING: from fastvideo.train.utils.training_config import TrainingConfig -# H3 maps one shared denoising stage through modality-specific scheduler -# shifts so video and audio remain synchronized at different noise amounts. _VIDEO_SCHEDULER_SHIFT = 12.0 _AUDIO_SCHEDULER_SHIFT = 3.0 _VIDEO_LATENT_CHANNELS = 24 _AUDIO_LATENT_CHANNELS = 32 +_AUDIO_SAMPLE_RATE = 32_000 + +_ALLOWED_ATTENTION_BACKENDS = ( + AttentionBackendEnum.TORCH_SDPA, + AttentionBackendEnum.FLASH_ATTN, + AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3, +) -def shift_noise_amount(base_noise_amount: torch.Tensor, shift: float) -> torch.Tensor: - """Apply the MiniMax H3 rational shift to a unit noise amount.""" +def shift_noise_amount( + base_noise_amount: torch.Tensor, + shift: float, + *, + max_noise_amount: float = 1.0, +) -> torch.Tensor: + """Apply the MiniMax H3 rational shift on ``[0, max_noise_amount]``.""" if shift <= 0: raise ValueError(f"shift must be positive, got {shift}") - return shift * base_noise_amount / (1.0 + (shift - 1.0) * base_noise_amount) + if not 0.0 < max_noise_amount <= 1.0: + raise ValueError("max_noise_amount must satisfy 0 < max <= 1, " + f"got {max_noise_amount}") + return (shift * base_noise_amount * max_noise_amount / (base_noise_amount * (shift - 1.0) + max_noise_amount)) class MiniMaxH3Model(ModelBase): @@ -62,36 +84,33 @@ def __init__( enable_gradient_checkpointing_type: str | None = None, transformer_override_safetensor: str | None = None, attention_backend: AttentionBackendEnum | str | None = AttentionBackendEnum.TORCH_SDPA, + construction_precision: str | None = None, ) -> None: """Validate the single-document T2VA contract and load the transformer.""" super().__init__( trainable=trainable, attention_backend=attention_backend, ) - # PyTorch scaled dot product attention (SDPA) provides dense attention - # without adding another attention-kernel dependency to H3 training. - if self.attention_backend != AttentionBackendEnum.TORCH_SDPA: - raise ValueError("MiniMaxH3Model requires the TORCH_SDPA attention backend") + if self.attention_backend not in _ALLOWED_ATTENTION_BACKENDS: + allowed = ", ".join(b.name for b in _ALLOWED_ATTENTION_BACKENDS) + raise ValueError("MiniMaxH3Model supports the attention backends " + f"{{{allowed}}}, got {self.attention_backend}") if training_config.pipeline_config is None: raise ValueError("MiniMaxH3Model requires a resolved MiniMax H3 pipeline config") - # Packed row indices describe one text-video-audio document without a - # batch offset, so each data-parallel replica consumes one sample. if int(training_config.data.train_batch_size) != 1: raise ValueError("MiniMaxH3Model requires training.data.train_batch_size=1") - # Classifier-free guidance (CFG) dropout replaces text embeddings with - # zeros, but H3 training does not define a zero-vector branch. if float(training_config.data.training_cfg_rate) != 0.0: raise ValueError("MiniMaxH3Model requires training.data.training_cfg_rate=0.0") - # Joint supervision requires paired video and stereo-audio latents from - # every parquet row. - if str(training_config.data.preprocessed_data_type) != "t2va": - raise ValueError("MiniMaxH3Model requires training.data.preprocessed_data_type='t2va'") - - # FastVideo's Fully Sharded Data Parallel loading path requires one BF16 - # parameter dtype, including modules that H3 inference keeps in FP32. - training_config.pipeline_config.dit_config.uniform_parameter_dtype = True # type: ignore[attr-defined] + if str(training_config.data.preprocessed_data_type) not in ("t2va", "text_only"): + raise ValueError("MiniMaxH3Model requires training.data.preprocessed_data_type " + "'t2va' or 'text_only'") + configured_precision = str(getattr(training_config, "dit_precision", "fp32")) + if trainable and construction_precision not in (None, configured_precision): + raise ValueError("A trainable MiniMaxH3 role cannot override construction_precision; " + "FP32 optimizer masters must follow training.dit_precision") self._init_from = str(init_from) + self._construction_precision = construction_precision self.training_config = training_config self.transformer = self._load_transformer( trainable=trainable, @@ -123,6 +142,7 @@ def _load_transformer( override_transformer_cls_name=self._transformer_cls_name, transformer_override_safetensor=transformer_override_safetensor, attention_backend=self.attention_backend, + construction_precision=self._construction_precision, ) checkpointing_type = (enable_gradient_checkpointing_type or self.training_config.model.enable_gradient_checkpointing_type) @@ -135,15 +155,17 @@ def _load_transformer( def init_preprocessors(self, training_config: TrainingConfig) -> None: """Load precomputed text embeddings and paired video-audio latents.""" - from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2va + from fastvideo.dataset.dataloader.schema import (pyarrow_schema_t2va, pyarrow_schema_text_only) from fastvideo.train.utils.dataloader import build_parquet_t2v_train_dataloader self.sp_group = get_sp_group() text_config = training_config.pipeline_config.text_encoder_configs[0] # type: ignore[union-attr] + parquet_schema = (pyarrow_schema_text_only + if str(training_config.data.preprocessed_data_type) == "text_only" else pyarrow_schema_t2va) self.dataloader = build_parquet_t2v_train_dataloader( training_config.data, text_len=int(text_config.arch_config.text_len), - parquet_schema=pyarrow_schema_t2va, + parquet_schema=parquet_schema, ) self.start_step = 0 @@ -154,30 +176,57 @@ def _resolve_clean_latents( dtype: torch.dtype, device: torch.device, ) -> tuple[torch.Tensor, torch.Tensor]: - """Resolve fixed visual and stereo-audio latent tensors for one sample.""" + """Resolve fixed or native visual and stereo-audio latents.""" data_config = self.training_config.data + native_shapes = bool(getattr(data_config, "native_shape_bucketing", False)) if latents_source == "data": if "vae_latent" not in raw_batch or "audio_latent" not in raw_batch: raise ValueError("A T2VA batch requires vae_latent and audio_latent tensors") video_latents = raw_batch["vae_latent"] audio_latents = raw_batch["audio_latent"] elif latents_source == "zeros": + latent_frames = int(data_config.num_latent_t) + height = int(data_config.num_height) + width = int(data_config.num_width) + num_frames = int(data_config.num_frames) + if native_shapes: + bucket_id = raw_batch.get("_shape_bucket_id") + if not isinstance(bucket_id, str): + raise ValueError( + "Native-shape data-free batches require the exact-shape sampler to set _shape_bucket_id") + bucket = parse_video_shape_bucket_id(bucket_id) + width = bucket.width + height = bucket.height + num_frames = bucket.num_frames + latent_frames = video_latent_num_frames(num_frames) + if width % MINIMAX_H3_CANVAS_MULTIPLE or height % MINIMAX_H3_CANVAS_MULTIPLE: + raise ValueError(f"Native pixel geometry {width}x{height} must use the H3 canvas multiple " + f"{MINIMAX_H3_CANVAS_MULTIPLE}") + latent_geometry = (latent_frames, height // 16, width // 16) + patch_size = tuple(int(value) for value in self.transformer.patch_size) + if any(value % patch for value, patch in zip(latent_geometry, patch_size, strict=True)): + raise ValueError( + f"Native latent geometry {latent_geometry} is not divisible by transformer patch {patch_size}") video_latents = torch.zeros( 1, _VIDEO_LATENT_CHANNELS, - data_config.num_latent_t, - data_config.num_height // 16, - data_config.num_width // 16, + latent_frames, + height // 16, + width // 16, ) audio_latents = torch.zeros( 1, MINIMAX_H3_AUDIO_CHANNELS, _AUDIO_LATENT_CHANNELS, - audio_latent_num_frames(data_config.num_frames), + audio_latent_num_frames(num_frames), ) else: raise ValueError(f"Unknown latents_source: {latents_source!r}") + if not isinstance(video_latents, torch.Tensor): + raise ValueError(f"vae_latent must be a tensor, got {type(video_latents).__name__}") + if not isinstance(audio_latents, torch.Tensor): + raise ValueError(f"audio_latent must be a tensor, got {type(audio_latents).__name__}") if video_latents.ndim != 5 or tuple(video_latents.shape[:2]) != (1, _VIDEO_LATENT_CHANNELS): raise ValueError("vae_latent must have shape [1, 24, latent_frames, latent_height, latent_width], " f"got {tuple(video_latents.shape)}") @@ -188,19 +237,100 @@ def _resolve_clean_latents( ): raise ValueError("audio_latent must have shape [1, 2, 32, audio_frames], " f"got {tuple(audio_latents.shape)}") - if data_config.num_latent_t > 0: - video_latents = video_latents[:, :, :data_config.num_latent_t] - expected_audio_frames = audio_latent_num_frames(data_config.num_frames) - audio_latents = audio_latents[:, :, :, :expected_audio_frames] - if video_latents.shape[2] != data_config.num_latent_t: - raise ValueError("vae_latent contains fewer frames than training.data.num_latent_t") - if audio_latents.shape[-1] != expected_audio_frames: - raise ValueError("audio_latent length does not match training.data.num_frames") + + if latents_source == "data" and native_shapes: + self._validate_native_latents(raw_batch, video_latents, audio_latents) + elif not native_shapes: + if data_config.num_latent_t > 0: + video_latents = video_latents[:, :, :data_config.num_latent_t] + expected_audio_frames = audio_latent_num_frames(data_config.num_frames) + audio_latents = audio_latents[:, :, :, :expected_audio_frames] + if video_latents.shape[2] != data_config.num_latent_t: + raise ValueError("vae_latent contains fewer frames than training.data.num_latent_t") + if audio_latents.shape[-1] != expected_audio_frames: + raise ValueError("audio_latent length does not match training.data.num_frames") return ( video_latents.to(device=device, dtype=dtype), audio_latents.to(device=device, dtype=dtype), ) + def _validate_native_latents( + self, + raw_batch: dict[str, Any], + video_latents: torch.Tensor, + audio_latents: torch.Tensor, + ) -> None: + """Cross-check the canonical bucket, row metadata, and latent clocks.""" + bucket_id = raw_batch.get("_shape_bucket_id") + if not isinstance(bucket_id, str): + raise ValueError("Native-shape T2VA batches require the exact-shape sampler to set _shape_bucket_id") + bucket = parse_video_shape_bucket_id(bucket_id) + + infos = raw_batch.get("info_list") + if not isinstance(infos, list) or len(infos) != 1 or not isinstance(infos[0], dict): + raise ValueError("Native-shape T2VA batches require exactly one info_list metadata record") + info = infos[0] + + def _metadata_int(name: str) -> int: + value = info.get(name) + if value is None or isinstance(value, bool): + raise ValueError(f"T2VA metadata {name!r} must be a positive integer, got {value!r}") + try: + result = int(value) + except (TypeError, ValueError) as exc: + raise ValueError(f"T2VA metadata {name!r} must be a positive integer, got {value!r}") from exc + if result <= 0 or result != value: + raise ValueError(f"T2VA metadata {name!r} must be a positive integer, got {value!r}") + return result + + width = _metadata_int("width") + height = _metadata_int("height") + num_frames = _metadata_int("num_frames") + audio_sample_rate = _metadata_int("audio_sample_rate") + if (width, height, num_frames) != (bucket.width, bucket.height, bucket.num_frames): + raise ValueError(f"Shape bucket {bucket_id!r} disagrees with row metadata " + f"width={width}, height={height}, num_frames={num_frames}") + if audio_sample_rate != _AUDIO_SAMPLE_RATE: + raise ValueError( + f"MiniMax H3 T2VA audio must be encoded at {_AUDIO_SAMPLE_RATE} Hz, got {audio_sample_rate}") + fps_value = info.get("fps") + if fps_value is None: + raise ValueError("T2VA metadata 'fps' must be numeric, got None") + try: + fps = float(fps_value) + except (TypeError, ValueError) as exc: + raise ValueError(f"T2VA metadata 'fps' must be numeric, got {info.get('fps')!r}") from exc + if not math.isfinite(fps) or not math.isclose(fps, 24.0, rel_tol=0.0, abs_tol=0.05): + raise ValueError(f"MiniMax H3 T2VA video must use the 24 fps clock, got {fps}") + + if width % MINIMAX_H3_CANVAS_MULTIPLE or height % MINIMAX_H3_CANVAS_MULTIPLE: + raise ValueError(f"Native pixel geometry {width}x{height} must use the H3 canvas multiple " + f"{MINIMAX_H3_CANVAS_MULTIPLE}") + expected_video_shape = ( + 1, + _VIDEO_LATENT_CHANNELS, + video_latent_num_frames(num_frames), + height // 16, + width // 16, + ) + expected_audio_shape = ( + 1, + MINIMAX_H3_AUDIO_CHANNELS, + _AUDIO_LATENT_CHANNELS, + audio_latent_num_frames(num_frames), + ) + if tuple(video_latents.shape) != expected_video_shape: + raise ValueError(f"vae_latent shape does not match {bucket_id!r}: expected {expected_video_shape}, " + f"got {tuple(video_latents.shape)}") + if tuple(audio_latents.shape) != expected_audio_shape: + raise ValueError(f"audio_latent shape does not match the {num_frames}-frame H3 audio clock: " + f"expected {expected_audio_shape}, got {tuple(audio_latents.shape)}") + latent_geometry = (expected_video_shape[2], expected_video_shape[3], expected_video_shape[4]) + patch_size = tuple(int(value) for value in self.transformer.patch_size) + if any(value % patch for value, patch in zip(latent_geometry, patch_size, strict=True)): + raise ValueError( + f"Native latent geometry {latent_geometry} is not divisible by transformer patch {patch_size}") + def _sample_noise_amounts( self, generator: torch.Generator, @@ -214,8 +344,6 @@ def _sample_noise_amounts( dtype=torch.float32, ) if int(self.training_config.distributed.sp_size) > 1: - # Sequence-parallel ranks shard one document and therefore require - # identical video and audio noise amounts for that document. self.sp_group.broadcast(base_noise_amount, src=0) return ( shift_noise_amount(base_noise_amount, _VIDEO_SCHEDULER_SHIFT), @@ -259,8 +387,6 @@ def prepare_batch( _, _, video_frames, latent_height, latent_width = video_latents.shape num_audio_latents = audio_latents.shape[-1] text_token_tags = torch.full((int(valid_text.sum()), ), MINIMAX_H3_TEXT_TAG, dtype=torch.long) - # H3 self-attention consumes one interleaved document, so the layout - # owns the row tags, positions, and modality output indices together. layout = build_packed_sequence( text_token_tags, video_frames, @@ -287,15 +413,44 @@ def prepare_batch( training_batch.audio_noise = audio_noise training_batch.sigmas = video_sigmas training_batch.audio_sigmas = audio_sigmas - # ModelBase exposes clean-time timesteps while the loss consumes the - # complementary noise amounts stored in the sigma fields. training_batch.timesteps = 1.0 - video_noise_amount training_batch.audio_timesteps = 1.0 - audio_noise_amount training_batch.minimax_h3_layout = layout training_batch.attn_metadata = None training_batch.attn_metadata_vsa = None + self._maybe_build_vsa_metadata(training_batch) return training_batch + def _maybe_build_vsa_metadata(self, batch: TrainingBatch) -> None: + """Populate the VSA view when this model runs the VSA-H3 backend. + + Methods route it through ``predict_noise(attn_kind="vsa")``; dense + models keep ``attn_metadata_vsa`` as ``None`` so their forwards stay + dense. One builder per model keeps the padded tile buffer reused + across steps (see ``MiniMaxH3VSAMetadata.tile_buf_holder``). + """ + backend = (getattr(self, "attention_backend_name", None) or envs.FASTVIDEO_ATTENTION_BACKEND) + if backend != "VIDEO_SPARSE_ATTN_H3": + return + layout = batch.minimax_h3_layout + patch_size = tuple(self.transformer.patch_size) + builder = getattr(self, "_vsa_metadata_builder", None) + if builder is None: + builder = self._vsa_metadata_builder = MiniMaxH3VSAMetadataBuilder() + batch.attn_metadata_vsa = builder.build( + current_timestep=0, + raw_latent_shape=( + layout.num_video_latent_frames, + layout.latent_height, + layout.latent_width, + ), + patch_size=patch_size, + VSA_sparsity=float(self.training_config.vsa_sparsity), + prefix_segments=_h3_vsa_prefix_segments(layout, patch_size), + device=self.device, + tile_size=int(self.training_config.vsa_tile_size), + ) + def add_noise( self, clean_latents: torch.Tensor, @@ -320,10 +475,9 @@ def predict_noise( ) -> NoisePrediction: """Pack modality timesteps and convert H3 outputs to noise-minus-clean.""" del timestep - if not conditional or cfg_uncond is not None: - raise ValueError("MiniMaxH3Model predicts one conditional T2VA sample") - if attn_kind != "dense": - raise ValueError("MiniMaxH3Model supports dense attention for training") + if attn_kind not in ("dense", "vsa"): + raise ValueError(f"Unknown attn_kind: {attn_kind!r}") + attn_metadata = (batch.attn_metadata_vsa if attn_kind == "vsa" else batch.attn_metadata) layout = batch.minimax_h3_layout if not isinstance(layout, MiniMaxH3PackedLayout): raise RuntimeError("prepare_batch() must set TrainingBatch.minimax_h3_layout") @@ -332,12 +486,18 @@ def predict_noise( if batch.timesteps is None or batch.audio_timesteps is None: raise RuntimeError("prepare_batch() must set video and audio timesteps") + encoder_hidden_states = batch.encoder_hidden_states + if not conditional: + if (cfg_uncond or {}).get("text") != "zero": + raise ValueError("MiniMaxH3Model unconditional forwards require " + "method.cfg_uncond={'text': 'zero'}") + encoder_hidden_states = torch.zeros_like(encoder_hidden_states) + dtype = torch.bfloat16 device = self.device + video_input_dtype = noisy_latents.dtype + audio_input_dtype = batch.audio_noisy_model_input.dtype video_bcthw = noisy_latents.permute(0, 2, 1, 3, 4).to(dtype) - # Match H3 checkpoint token order: video rows flatten - # (C, patch_t, patch_h, patch_w), while audio rows flatten stereo - # channel, time, and latent feature dimensions in that order. video_rows = patchify_video_latents(video_bcthw, self.transformer.patch_size) audio_latents = batch.audio_noisy_model_input.to(dtype) num_audio_latents = audio_latents.shape[-1] @@ -352,14 +512,14 @@ def predict_noise( unique_timesteps = unique_timesteps.to(device) timestep_indices = timestep_indices.to(device) - with torch.autocast(device.type, dtype=dtype), set_forward_context( + with set_forward_context( current_timestep=unique_timesteps, - attn_metadata=None, + attn_metadata=attn_metadata, ): video_velocity, audio_velocity = self.transformer( hidden_states=video_rows[None], audio_hidden_states=audio_rows[None], - encoder_hidden_states=batch.encoder_hidden_states, + encoder_hidden_states=encoder_hidden_states, timestep=unique_timesteps, timestep_indices=timestep_indices, token_tags=layout.token_tags.to(device), @@ -379,7 +539,10 @@ def predict_noise( self.transformer.patch_size, ).permute(0, 2, 1, 3, 4) audio_prediction = unpack_audio_tokens(audio_velocity[0], num_audio_latents)[None] - return -video_prediction, -audio_prediction + return ( + (-video_prediction).to(video_input_dtype), + (-audio_prediction).to(audio_input_dtype), + ) def backward( self, @@ -390,6 +553,19 @@ def backward( ) -> None: """Restore the forward context and average accumulated microbatch gradients.""" timesteps, attn_metadata = ctx + expected_device_type = self.device.type + offenders = [] + for name, param in self.transformer.named_parameters(): + local = getattr(param, "_local_tensor", param) + if local.device.type != expected_device_type: + offenders.append(f"param {name} on {local.device}") + if param.grad is not None: + grad_local = getattr(param.grad, "_local_tensor", param.grad) + if grad_local.device.type != expected_device_type: + offenders.append(f"grad {name} on {grad_local.device}") + if offenders: + raise RuntimeError(f"{len(offenders)} training tensors off-CUDA before backward; " + f"first offenders: {offenders[:8]}") with set_forward_context( current_timestep=timesteps, attn_metadata=attn_metadata, diff --git a/fastvideo/train/models/minimax_h3/minimax_h3_dmd.py b/fastvideo/train/models/minimax_h3/minimax_h3_dmd.py new file mode 100644 index 0000000000..26438badb1 --- /dev/null +++ b/fastvideo/train/models/minimax_h3/minimax_h3_dmd.py @@ -0,0 +1,498 @@ +# SPDX-License-Identifier: Apache-2.0 +"""MiniMax H3 distribution-matching adapter (packed dual-modality latents).""" + +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import Any, cast, Literal + +import torch + +from fastvideo.pipelines import TrainingBatch +from fastvideo.pipelines.basic.minimax_h3.packing import ( + MINIMAX_H3_AUDIO_CHANNELS, + audio_latent_num_frames, +) +from fastvideo.train.models.minimax_h3.minimax_h3 import ( + _AUDIO_LATENT_CHANNELS, + _AUDIO_SCHEDULER_SHIFT, + _VIDEO_LATENT_CHANNELS, + _VIDEO_SCHEDULER_SHIFT, + MiniMaxH3Model, + shift_noise_amount, +) + +_DMD_TIMESTEP_SCALE = 1000 +_FASTGEN_MAX_T = 0.999 + + +@dataclass(frozen=True, slots=True) +class MiniMaxH3DMDLatentLayout: + """Exact batch-local shapes behind one packed H3 DMD latent tensor.""" + + video_shape: tuple[int, int, int, int, int] + audio_shape: tuple[int, int, int, int] + + def __post_init__(self) -> None: + if self.video_shape[0] != 1 or self.video_shape[2] != _VIDEO_LATENT_CHANNELS: + raise ValueError("DMD video latents must have shape [1, T, 24, H, W], got " + f"{self.video_shape}") + if self.audio_shape[:3] != (1, MINIMAX_H3_AUDIO_CHANNELS, _AUDIO_LATENT_CHANNELS): + raise ValueError("DMD audio latents must have shape [1, 2, 32, Ta], got " + f"{self.audio_shape}") + if any(value <= 0 for value in self.video_shape + self.audio_shape): + raise ValueError("DMD latent layout dimensions must all be positive") + + @classmethod + def from_latents( + cls, + video_latents: torch.Tensor, + audio_latents: torch.Tensor, + ) -> MiniMaxH3DMDLatentLayout: + if video_latents.ndim != 5 or audio_latents.ndim != 4: + raise ValueError("H3 DMD pack expects video [1,T,24,H,W] and audio [1,2,32,Ta], " + f"got {tuple(video_latents.shape)} and {tuple(audio_latents.shape)}") + return cls( + video_shape=cast(tuple[int, int, int, int, int], tuple(int(value) for value in video_latents.shape)), + audio_shape=cast(tuple[int, int, int, int], tuple(int(value) for value in audio_latents.shape)), + ) + + @property + def video_numel(self) -> int: + return math.prod(self.video_shape) + + @property + def audio_numel(self) -> int: + return math.prod(self.audio_shape) + + @property + def packed_numel(self) -> int: + return self.video_numel + self.audio_numel + + def modality_slices(self) -> tuple[tuple[str, slice], ...]: + return ( + ("video", slice(0, self.video_numel)), + ("audio", slice(self.video_numel, self.packed_numel)), + ) + + +class MiniMaxH3DMDModel(MiniMaxH3Model): + """Present H3's dual (video, audio) streams to DMD2 as one packed tensor. + + ``DMD2Method``'s rollout and loss math assume one latent tensor per + sample. This adapter flattens both modality latents into one ``[1, N]`` + tensor (video's ``[1, T, 24, H, W]`` elements first, stereo audio's + ``[1, 2, 32, Ta]`` elements after) so ``dmd2.py`` stays model-agnostic. + Method timesteps become one shared base noise amount that is shifted per + modality (video 12.0, audio 3.0). Continuous score times use FastGen's + ``max_t=0.999`` domain; integer rollout rungs preserve the release + pipeline's unit-domain grid. + + ``modality_slices()`` exposes the packed video/audio column ranges so + DMD2 computes losses and normalizers per modality instead of one packed + mean that video's element count would dominate. + """ + + @property + def num_train_timesteps(self) -> int: + return _DMD_TIMESTEP_SCALE + + + def _modality_shapes(self) -> tuple[tuple[int, int, int, int, int], tuple[int, int, int, int]]: + """Return the ``[1, T, C, H, W]`` video and ``[1, 2, 32, Ta]`` audio shapes.""" + data = self.training_config.data + video_shape = ( + 1, + int(data.num_latent_t), + _VIDEO_LATENT_CHANNELS, + int(data.num_height) // 16, + int(data.num_width) // 16, + ) + audio_shape = ( + 1, + MINIMAX_H3_AUDIO_CHANNELS, + _AUDIO_LATENT_CHANNELS, + audio_latent_num_frames(int(data.num_frames)), + ) + return video_shape, audio_shape + + def _fixed_latent_layout(self) -> MiniMaxH3DMDLatentLayout: + video_shape, audio_shape = self._modality_shapes() + return MiniMaxH3DMDLatentLayout( + video_shape=video_shape, + audio_shape=audio_shape, + ) + + def _batch_latent_layout(self, batch: TrainingBatch) -> MiniMaxH3DMDLatentLayout: + layout = batch.minimax_h3_dmd_layout + if not isinstance(layout, MiniMaxH3DMDLatentLayout): + raise RuntimeError("prepare_batch() must set TrainingBatch.minimax_h3_dmd_layout") + return layout + + def pack_latents( + self, + video_latents: torch.Tensor, + audio_latents: torch.Tensor, + *, + layout: MiniMaxH3DMDLatentLayout | None = None, + ) -> torch.Tensor: + """Flatten both modality latents into one ``[1, N]`` tensor.""" + actual_layout = MiniMaxH3DMDLatentLayout.from_latents(video_latents, audio_latents) + if layout is not None and actual_layout != layout: + raise ValueError(f"Latent tensors do not match their batch layout: {actual_layout} != {layout}") + return torch.cat( + (video_latents.reshape(1, -1), audio_latents.reshape(1, -1)), + dim=1, + ) + + def modality_slices( + self, + *, + layout: MiniMaxH3DMDLatentLayout | None = None, + ) -> tuple[tuple[str, slice], ...]: + """Named packed-latent column slices for per-modality DMD2 losses. + + Video is ~3.6M packed elements against audio's ~15-30k, so a single + global mean would give audio <1% of the distillation signal; DMD2 + consumes these slices to normalize and weight each stream separately. + """ + return (layout or self._fixed_latent_layout()).modality_slices() + + def modality_slices_for_batch(self, batch: TrainingBatch) -> tuple[tuple[str, slice], ...]: + """Return per-modality slices for this batch's native shape.""" + return self.modality_slices(layout=self._batch_latent_layout(batch)) + + def unpack_latents( + self, + packed: torch.Tensor, + *, + layout: MiniMaxH3DMDLatentLayout | None = None, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Split one packed ``[1, N]`` tensor back into (video, audio) latents.""" + resolved = layout or self._fixed_latent_layout() + if packed.shape != (1, resolved.packed_numel): + raise ValueError(f"Packed latents must have shape [1, {resolved.packed_numel}], got {tuple(packed.shape)}") + return ( + packed[:, :resolved.video_numel].reshape(resolved.video_shape), + packed[:, resolved.video_numel:].reshape(resolved.audio_shape), + ) + + def _noise_amounts(self, timestep: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Map one method timestep to both modality noise amounts in FP64.""" + warp_max = (_FASTGEN_MAX_T if timestep.is_floating_point() else 1.0) + base = (timestep.reshape(-1)[:1].to(torch.float64) / _DMD_TIMESTEP_SCALE) + base = base.clamp(0.0, warp_max) + return ( + shift_noise_amount( + base, + _VIDEO_SCHEDULER_SHIFT, + max_noise_amount=warp_max, + ), + shift_noise_amount( + base, + _AUDIO_SCHEDULER_SHIFT, + max_noise_amount=warp_max, + ), + ) + + + def set_requires_negative_conditioning(self, requires: bool) -> None: + """Fail fast: H3 cannot encode negative prompts at training time.""" + if requires: + raise ValueError("MiniMaxH3DMDModel has no negative-prompt encoder; set " + "method.cfg_uncond={'text': 'zero'} for unconditional forwards") + + def prepare_batch( + self, + raw_batch: dict[str, Any], + *, + generator: torch.Generator, + latents_source: Literal["data", "zeros"] = "data", + ) -> TrainingBatch: + """Prepare the T2VA batch, then expose clean latents in packed form.""" + batch = super().prepare_batch( + raw_batch, + generator=generator, + latents_source=latents_source, + ) + if batch.latents is None or batch.audio_latents is None: + raise RuntimeError("MiniMax H3 batch preparation did not produce paired latents") + layout = MiniMaxH3DMDLatentLayout.from_latents(batch.latents, batch.audio_latents) + batch.minimax_h3_dmd_layout = layout + batch.latents = self.pack_latents(batch.latents, batch.audio_latents, layout=layout) + return batch + + def add_noise( + self, + clean_latents: torch.Tensor, + noise: torch.Tensor, + timestep: torch.Tensor, + ) -> torch.Tensor: + """Noise legacy fixed-shape packed latents at one shared timestep.""" + return self._add_noise_with_layout( + clean_latents, + noise, + timestep, + self._fixed_latent_layout(), + ) + + def add_noise_for_batch( + self, + clean_latents: torch.Tensor, + noise: torch.Tensor, + timestep: torch.Tensor, + batch: TrainingBatch, + ) -> torch.Tensor: + """Noise packed latents using immutable geometry from ``batch``.""" + layout = self._batch_latent_layout(batch) + if not bool(getattr(self.training_config.data, "native_shape_bucketing", False)): + return self.add_noise(clean_latents, noise, timestep) + return self._add_noise_with_layout( + clean_latents, + noise, + timestep, + layout, + ) + + def _add_noise_with_layout( + self, + clean_latents: torch.Tensor, + noise: torch.Tensor, + timestep: torch.Tensor, + layout: MiniMaxH3DMDLatentLayout, + ) -> torch.Tensor: + sigma_video, sigma_audio = self._noise_amounts(timestep) + clean_video, clean_audio = self.unpack_latents(clean_latents, layout=layout) + noise_video, noise_audio = self.unpack_latents(noise, layout=layout) + return self.pack_latents( + self._mix(clean_video, noise_video, sigma_video), + self._mix(clean_audio, noise_audio, sigma_audio), + layout=layout, + ) + + @staticmethod + def _mix( + clean: torch.Tensor, + noise: torch.Tensor, + sigma: torch.Tensor, + ) -> torch.Tensor: + original_dtype = clean.dtype + clean_fp64 = clean.to(torch.float64) + noise_fp64 = noise.to(torch.float64) + sigma_fp64 = sigma.to(device=clean.device, dtype=torch.float64) + return ((1.0 - sigma_fp64) * clean_fp64 + sigma_fp64 * noise_fp64).to(original_dtype) + + def extract_eps( + self, + noisy_latents: torch.Tensor, + clean_latents: torch.Tensor, + timestep: torch.Tensor, + ) -> torch.Tensor: + """Recover noise for legacy fixed-shape packed latents. + + The inverse of :meth:`add_noise` at the same shared timestep: with + ``x_t = (1 - sigma_m) x0 + sigma_m eps`` under each modality's + shifted sigma (video 12.0, audio 3.0), the implied noise is + ``eps_m = (x_t - (1 - sigma_m) x0) / sigma_m``. DMD2's ODE renoise + uses this to step a carried trajectory deterministically between + grid rungs. + """ + return self._extract_eps_with_layout( + noisy_latents, + clean_latents, + timestep, + self._fixed_latent_layout(), + ) + + def extract_eps_for_batch( + self, + noisy_latents: torch.Tensor, + clean_latents: torch.Tensor, + timestep: torch.Tensor, + batch: TrainingBatch, + ) -> torch.Tensor: + """Recover packed noise using immutable geometry from ``batch``.""" + layout = self._batch_latent_layout(batch) + if not bool(getattr(self.training_config.data, "native_shape_bucketing", False)): + return self.extract_eps(noisy_latents, clean_latents, timestep) + return self._extract_eps_with_layout( + noisy_latents, + clean_latents, + timestep, + layout, + ) + + def _extract_eps_with_layout( + self, + noisy_latents: torch.Tensor, + clean_latents: torch.Tensor, + timestep: torch.Tensor, + layout: MiniMaxH3DMDLatentLayout, + ) -> torch.Tensor: + sigma_video, sigma_audio = self._noise_amounts(timestep) + noisy_video, noisy_audio = self.unpack_latents(noisy_latents, layout=layout) + clean_video, clean_audio = self.unpack_latents(clean_latents, layout=layout) + return self.pack_latents( + self._unmix(noisy_video, clean_video, sigma_video), + self._unmix(noisy_audio, clean_audio, sigma_audio), + layout=layout, + ) + + @staticmethod + def _unmix( + noisy: torch.Tensor, + clean: torch.Tensor, + sigma: torch.Tensor, + ) -> torch.Tensor: + original_dtype = noisy.dtype + noisy_fp64 = noisy.to(torch.float64) + clean_fp64 = clean.to(torch.float64) + sigma_fp64 = sigma.to(device=noisy.device, dtype=torch.float64) + eps = ((noisy_fp64 - (1.0 - sigma_fp64) * clean_fp64) / sigma_fp64.clamp_min(1e-6)) + return eps.to(original_dtype) + + def predict_noise( + self, + noisy_latents: torch.Tensor, + timestep: torch.Tensor, + batch: TrainingBatch, + *, + conditional: bool, + cfg_uncond: dict[str, Any] | None = None, + attn_kind: Literal["dense", "vsa"] = "dense", + ) -> torch.Tensor: + """Run one packed joint forward at an explicit method timestep. + + Both modality clean-time fields on ``batch`` are rewritten from + ``timestep`` so the packed-row timestep plan and the backward + forward-context stay coherent with this call. + """ + sigma_video, sigma_audio = self._noise_amounts(timestep) + layout = self._batch_latent_layout(batch) + noisy_video, noisy_audio = self.unpack_latents(noisy_latents, layout=layout) + batch.timesteps = (1.0 - sigma_video).to(noisy_latents.device) + batch.audio_timesteps = (1.0 - sigma_audio).to(noisy_latents.device) + batch.audio_noisy_model_input = noisy_audio + video_pred, audio_pred = super().predict_noise( + noisy_video, + timestep, + batch, + conditional=conditional, + cfg_uncond=cfg_uncond, + attn_kind=attn_kind, + ) + return self.pack_latents(video_pred, audio_pred, layout=layout) + + def predict_x0( + self, + noisy_latents: torch.Tensor, + timestep: torch.Tensor, + batch: TrainingBatch, + *, + conditional: bool, + cfg_uncond: dict[str, Any] | None = None, + attn_kind: Literal["dense", "vsa"] = "dense", + ) -> torch.Tensor: + """Convert packed noise-minus-clean predictions to packed clean latents.""" + pred_noise = self.predict_noise( + noisy_latents, + timestep, + batch, + conditional=conditional, + cfg_uncond=cfg_uncond, + attn_kind=attn_kind, + ) + sigma_video, sigma_audio = self._noise_amounts(timestep) + layout = self._batch_latent_layout(batch) + noisy_video, noisy_audio = self.unpack_latents(noisy_latents, layout=layout) + pred_video, pred_audio = self.unpack_latents(pred_noise, layout=layout) + return self.pack_latents( + self._to_x0(noisy_video, pred_video, sigma_video), + self._to_x0(noisy_audio, pred_audio, sigma_audio), + layout=layout, + ) + + @staticmethod + def _to_x0( + noisy: torch.Tensor, + pred_noise: torch.Tensor, + sigma: torch.Tensor, + ) -> torch.Tensor: + original_dtype = noisy.dtype + noisy_fp64 = noisy.to(torch.float64) + pred_noise_fp64 = pred_noise.to(torch.float64) + sigma_fp64 = sigma.to(device=noisy.device, dtype=torch.float64) + return (noisy_fp64 - sigma_fp64 * pred_noise_fp64).to(original_dtype) + + + def _load_vis_vae(self) -> Any: + """Lazily load the H3 video VAE for visualization decodes. + + The module stays CPU-resident between decodes; ``decode_vis_latents`` + moves it to the GPU per call. Loading mirrors the H3 preprocess + scripts: the inference component registry keeps precision policy and + normalization identical to the published decode recipe. + """ + vae = getattr(self, "_vis_vae_module", None) + if vae is not None: + return vae + import os + + from fastvideo.fastvideo_args import FastVideoArgs + from fastvideo.models.loader.component_loader import PipelineComponentLoader + from fastvideo.utils import verify_model_config_and_directory + + model_index = verify_model_config_and_directory(self._init_from) + transformers_or_diffusers, _ = model_index["vae"][:2] + args = FastVideoArgs( + model_path=self._init_from, + pipeline_config=self.training_config.pipeline_config, + num_gpus=1, + tp_size=1, + sp_size=1, + hsdp_shard_dim=1, + use_fsdp_inference=False, + vae_cpu_offload=True, + text_encoder_cpu_offload=True, + ) + vae = PipelineComponentLoader.load_module( + module_name="vae", + component_model_path=os.path.join(self._init_from, "vae"), + transformers_or_diffusers=transformers_or_diffusers, + fastvideo_args=args, + ) + vae.to("cpu") + self._vis_vae_module = vae + return vae + + @torch.no_grad() + def decode_vis_latents( + self, + packed: torch.Tensor, + *, + layout: MiniMaxH3DMDLatentLayout | None = None, + ) -> Any: + """Decode the packed video stream into a uint8 ``[B, T, C, H, W]`` clip. + + Follows ``MiniMaxH3VideoDecodingStage``: denormalize latents, decode + under the published FP16-autocast-over-FP32 recipe, denormalize + pixels. The audio stream is dropped — the tracker artifact is a + silent video. + """ + video_latents, _ = self.unpack_latents(packed.detach(), layout=layout) + latents = video_latents.permute(0, 2, 1, 3, 4).to(device=self.device, dtype=torch.float32) + vae = self._load_vis_vae() + vae.to(self.device) + try: + latents = vae.denormalize_latents(latents) + with torch.autocast(self.device.type, dtype=torch.float16, enabled=self.device.type == "cuda"): + video = vae.decode(latents).sample + video = vae.denormalize_pixels(video.float()).clamp_(0.0, 1.0).cpu() + finally: + vae.to("cpu") + video = video.permute(0, 2, 1, 3, 4) + return (video * 255.0).to(torch.uint8).numpy() + + +__all__ = ["MiniMaxH3DMDLatentLayout", "MiniMaxH3DMDModel"] diff --git a/fastvideo/train/trainer.py b/fastvideo/train/trainer.py index be6054f772..549c50628e 100644 --- a/fastvideo/train/trainer.py +++ b/fastvideo/train/trainer.py @@ -11,6 +11,7 @@ from tqdm.auto import tqdm from fastvideo.distributed import get_sp_group, get_world_group +from fastvideo.logger import init_logger from fastvideo.train.callbacks.callback import CallbackDict from fastvideo.train.methods.base import LogScalar, TrainingMethod from fastvideo.train.utils.tracking import build_tracker @@ -19,6 +20,42 @@ from fastvideo.train.utils.training_config import ( TrainingConfig, ) +logger = init_logger(__name__) + + +def _verify_master_weight_precision(method: TrainingMethod, tc: TrainingConfig) -> None: + """Refuse to train on low-precision master weights unless opted in. + + Optimizer steps applied in-place to bf16/fp16 parameters round away + updates below ~half an ulp of each weight's magnitude; O(1)-magnitude + parameters (norm gains) freeze entirely at typical distillation learning + rates, and ``zeros_like``-allocated optimizer state inherits the same + starved dtype. FP32 sharded masters (``training.dit_precision: fp32``) + fix both; ordinary FSDP groups still compute in BF16, while models may + declare narrower FP32 compute boundaries. + """ + if bool(getattr(tc.model, "allow_low_precision_master_weights", False)): + return + offenders: list[str] = [] + for role, model in getattr(method, "_role_models", {}).items(): + if not getattr(model, "_trainable", False): + continue + transformer = getattr(model, "transformer", None) + if transformer is None: + continue + for name, param in transformer.named_parameters(): + if param.requires_grad and param.dtype != torch.float32: + offenders.append(f"{role}:{name} ({param.dtype})") + break + if offenders: + raise RuntimeError("Trainable master weights are not fp32: " + f"{offenders}. bf16/fp16 parameter storage silently rounds away " + "optimizer updates below ~half an ulp per weight (norm-scale " + "parameters freeze completely). Set training.dit_precision: fp32 " + "(FP32 sharded masters; ordinary groups compute in BF16), " + "or acknowledge the effect explicitly with " + "training.model.allow_low_precision_master_weights: true.") + def _coerce_log_scalar( value: Any, @@ -114,6 +151,7 @@ def run( ) method.set_tracker(self.tracker) + _verify_master_weight_precision(method, tc) method.on_train_start() self.callbacks.on_train_start( method, @@ -127,6 +165,15 @@ def run( resumed_step = (checkpoint_manager.maybe_resume(resume_from_checkpoint=(resume_from_checkpoint))) if resumed_step is not None: start_step = int(resumed_step) + if bool(getattr(tc.checkpoint, "reset_lr_on_resume", False)): + method.apply_configured_lrs() + logger.info("reset_lr_on_resume: re-applied configured learning rates at step %s", start_step) + initial_validation_scheduled = self.callbacks.will_run_validation(iteration=start_step) + if checkpoint_manager is not None: + checkpoint_manager.maybe_save_inference( + start_step, + validation_scheduled=initial_validation_scheduled, + ) self.callbacks.on_validation_begin( method, iteration=start_step, @@ -136,9 +183,6 @@ def run( data_stream = self._iter_dataloader(dataloader) - # Restore the RNG snapshot LAST — after dcp.load, - # after iter(dataloader), after everything that may - # have advanced the RNG as a side-effect. if (checkpoint_manager is not None and resume_from_checkpoint): checkpoint_manager.load_rng_snapshot(resume_from_checkpoint, ) progress = tqdm( @@ -147,13 +191,10 @@ def run( desc="Steps", disable=self.local_rank > 0, ) - # Allow method-specific optimization flow (e.g. DiffusionNFT). method_manages_optimization = bool(method.manages_optimization()) for step in progress: t0 = time.perf_counter() - # Accumulate on GPU during grad-accum; materialise - # to CPU once per step right before logging. loss_sums: dict[str, float | torch.Tensor] = {} metric_sums: dict[str, float | torch.Tensor] = {} if method_manages_optimization: @@ -214,8 +255,6 @@ def run( method.optimizers_schedulers_step(step) method.optimizers_zero_grad(step) - # Single CPU sync point: materialise GPU tensors - # to float right before logging. divisor = 1 if method_manages_optimization else grad_accum metrics = {k: float(v) / divisor for k, v in loss_sums.items()} metrics.update({k: float(v) / divisor for k, v in metric_sums.items()}) @@ -230,7 +269,12 @@ def run( iteration=step, ) + validation_scheduled = self.callbacks.will_run_validation(iteration=step) if checkpoint_manager is not None: + checkpoint_manager.maybe_save_inference( + step, + validation_scheduled=validation_scheduled, + ) checkpoint_manager.maybe_save(step) self.callbacks.on_validation_begin( diff --git a/fastvideo/train/utils/checkpoint.py b/fastvideo/train/utils/checkpoint.py index 7a3d6c1055..ec8e9d2132 100644 --- a/fastvideo/train/utils/checkpoint.py +++ b/fastvideo/train/utils/checkpoint.py @@ -2,11 +2,13 @@ from __future__ import annotations +import contextlib import json import os import random import re import shutil +import time from dataclasses import dataclass from pathlib import Path from typing import Any @@ -27,6 +29,8 @@ logger = init_logger(__name__) _CHECKPOINT_DIR_RE = re.compile(r"^checkpoint-(\d+)$") +_TRAINING_CHECKPOINT_COMPLETE_MARKER = ".complete" +_RANK_RNG_STATE_RE = re.compile(r"^rng_state_rank(\d+)\.pt$") def _is_stateful(obj: Any) -> bool: @@ -52,7 +56,62 @@ def _parse_step_from_dir(checkpoint_dir: Path) -> int: return int(match.group(1)) -def _find_latest_checkpoint(output_dir: Path) -> Path | None: +def _saved_checkpoint_world_size(metadata: dict[str, Any]) -> int | None: + try: + world_size = metadata["config"]["training"]["distributed"]["num_gpus"] + except (KeyError, TypeError): + return None + if isinstance(world_size, bool) or not isinstance(world_size, int) or world_size <= 0: + return None + return world_size + + +def _is_complete_training_checkpoint( + checkpoint_dir: Path, + *, + require_complete_marker: bool, +) -> bool: + """Return whether ``checkpoint_dir`` is safe to select for resume. + + ``dcp/.metadata`` is the historical completion contract. Strict callers + additionally require the marker published after every rank has written its + RNG snapshot. Keeping strictness opt-in preserves compatibility with + checkpoints created before the stronger marker existed. + """ + dcp_metadata = checkpoint_dir / "dcp" / ".metadata" + if not dcp_metadata.is_file(): + return False + if not require_complete_marker: + return True + + try: + step = _parse_step_from_dir(checkpoint_dir) + marker = (checkpoint_dir / _TRAINING_CHECKPOINT_COMPLETE_MARKER).read_text(encoding="utf-8") + metadata = json.loads((checkpoint_dir / "metadata.json").read_text(encoding="utf-8")) + except (OSError, TypeError, ValueError): + return False + if not isinstance(metadata, dict) or marker != "complete\n" or metadata.get("step") != step: + return False + + world_size = _saved_checkpoint_world_size(metadata) + if world_size is None: + return False + expected_rng_names = {f"rng_state_rank{rank}.pt" for rank in range(world_size)} + try: + actual_rng_paths = list(checkpoint_dir.glob("rng_state_rank*.pt")) + actual_rng_names = {path.name for path in actual_rng_paths if _RANK_RNG_STATE_RE.match(path.name)} + rng_files_complete = all(path.is_file() and path.stat().st_size > 0 for path in actual_rng_paths) + except OSError: + return False + return (actual_rng_names == expected_rng_names and len(actual_rng_paths) == len(expected_rng_names) + and rng_files_complete) + + +def _find_latest_checkpoint( + output_dir: Path, + *, + require_complete_marker: bool = False, +) -> Path | None: if not output_dir.exists(): return None @@ -62,7 +121,10 @@ def _find_latest_checkpoint(output_dir: Path) -> Path | None: continue if not _CHECKPOINT_DIR_RE.match(child.name): continue - if not (child / "dcp").is_dir(): + if not _is_complete_training_checkpoint( + child, + require_complete_marker=require_complete_marker, + ): continue try: step = _parse_step_from_dir(child) @@ -76,7 +138,27 @@ def _find_latest_checkpoint(output_dir: Path) -> Path | None: return candidates[-1][1] -def _resolve_resume_checkpoint(resume_from_checkpoint: str, *, output_dir: str) -> Path | None: +def _publish_training_checkpoint_complete(checkpoint_dir: Path) -> None: + """Atomically publish the marker that makes a training checkpoint visible.""" + marker = checkpoint_dir / _TRAINING_CHECKPOINT_COMPLETE_MARKER + temporary = checkpoint_dir / f"{_TRAINING_CHECKPOINT_COMPLETE_MARKER}.tmp-{os.getpid()}-{time.time_ns()}" + try: + with temporary.open("w", encoding="utf-8") as handle: + handle.write("complete\n") + handle.flush() + os.fsync(handle.fileno()) + os.replace(temporary, marker) + finally: + with contextlib.suppress(FileNotFoundError): + temporary.unlink() + + +def _resolve_resume_checkpoint( + resume_from_checkpoint: str, + *, + output_dir: str, + require_complete_marker: bool = False, +) -> Path | None: """Resolve a user-provided resume path to a concrete checkpoint dir. Accepted values: @@ -89,8 +171,16 @@ def _resolve_resume_checkpoint(resume_from_checkpoint: str, *, output_dir: str) if str(resume_from_checkpoint).strip().lower() == "latest": out = Path(os.path.expanduser(str(output_dir))).resolve() - latest = _find_latest_checkpoint(out) + latest = _find_latest_checkpoint( + out, + require_complete_marker=require_complete_marker, + ) if latest is None: + has_checkpoint_dirs = out.is_dir() and any(child.is_dir() and _CHECKPOINT_DIR_RE.match(child.name) + for child in out.iterdir()) + if require_complete_marker and has_checkpoint_dirs: + raise ValueError(f"No complete resumable checkpoint found under {out}; " + "refusing to start from scratch in a non-empty training namespace") logger.info( "resume_from_checkpoint='latest' but no " "checkpoints found under %s; starting from " @@ -110,14 +200,20 @@ def _resolve_resume_checkpoint(resume_from_checkpoint: str, *, output_dir: str) if path.is_dir() and _CHECKPOINT_DIR_RE.match(path.name): if not (path / "dcp").is_dir(): raise FileNotFoundError(f"Missing dcp dir under checkpoint: {path / 'dcp'}") + if not _is_complete_training_checkpoint( + path, + require_complete_marker=require_complete_marker, + ): + raise ValueError(f"Checkpoint is incomplete under the configured resume policy: {path}") return path - # Treat as output_dir -> pick latest. - latest = _find_latest_checkpoint(path) + latest = _find_latest_checkpoint( + path, + require_complete_marker=require_complete_marker, + ) if latest is not None: return latest - # Give a clearer error message. out = Path(os.path.expanduser(str(output_dir))).resolve() raise ValueError("Could not resolve resume checkpoint. Expected a checkpoint directory " f"named 'checkpoint-' (with 'dcp/' inside), or an output_dir " @@ -184,6 +280,11 @@ def load_state_dict( class CheckpointConfig: save_steps: int keep_last: int + start_step: int = 0 + save_inference_on_validation: bool = False + inference_role: str = "student" + inference_dtype: str = "bfloat16" + require_complete_training_checkpoint: bool = False class CheckpointManager: @@ -207,18 +308,28 @@ def __init__( self.dataloader = dataloader self.output_dir = str(output_dir) self.config = config + save_inference = bool(config.save_inference_on_validation) + inference_role = str(config.inference_role or "") + if save_inference and (not inference_role or "." in inference_role): + raise ValueError("inference_role must be a non-empty DCP key segment when inference saving is enabled") + if save_inference and str(config.inference_dtype) not in {"bfloat16", "float16", "float32"}: + raise ValueError("inference_dtype must be bfloat16, float16, or float32") + if config.require_complete_training_checkpoint: + metadata = {"config": raw_config} + if _saved_checkpoint_world_size(metadata) is None: + raise ValueError("require_complete_training_checkpoint needs a positive " + "training.distributed.num_gpus value in the saved raw config") self._callbacks = callbacks self._raw_config = raw_config self._last_saved_step: int | None = None + self._last_inference_saved_step: int | None = None def _build_states(self) -> dict[str, Any]: states: dict[str, Any] = self.method.checkpoint_state() - # Dataloader (optional but recommended for exact resume). if _is_stateful(self.dataloader): states["dataloader"] = self.dataloader - # Callback state (e.g. EMA shadow weights, validation RNG). if self._callbacks is not None and _is_stateful(self._callbacks): states["callbacks"] = _CallbackStateWrapper(self._callbacks, ) @@ -230,49 +341,189 @@ def _checkpoint_dir(self, step: int) -> Path: def _dcp_dir(self, step: int) -> Path: return self._checkpoint_dir(step) / "dcp" + def _inference_checkpoint_dir(self, step: int) -> Path: + return Path(self.output_dir) / "inference" / f"checkpoint-{step}" + + def _inference_staging_dir(self, step: int) -> Path: + return Path(self.output_dir) / ".inference-staging" / f"checkpoint-{step}" + def maybe_save(self, step: int) -> None: - save_steps = int(self.config.save_steps or 0) - if save_steps <= 0: + if step < int(self.config.start_step or 0): return - if step % save_steps != 0: + + save_steps = int(self.config.save_steps or 0) + if save_steps > 0 and step % save_steps == 0 and self._last_saved_step != step: + self.save(step) + + def maybe_save_inference(self, step: int, *, validation_scheduled: bool) -> None: + """Save the model evaluated by one scheduled validation event. + + This event-driven policy includes step-zero validation and deliberately + ignores the start gate and cadence used for resumable training state. + """ + if not validation_scheduled or not bool(self.config.save_inference_on_validation): return - if self._last_saved_step == step: + if self._last_inference_saved_step == step: return - self.save(step) + self.save_inference(step) def save_final(self, step: int) -> None: - save_steps = int(self.config.save_steps or 0) - if save_steps <= 0: - return - self.save(step) + if int(self.config.save_steps or 0) > 0 and self._last_saved_step != step: + self.save(step) def save(self, step: int) -> None: checkpoint_dir = self._checkpoint_dir(step) dcp_dir = self._dcp_dir(step) os.makedirs(dcp_dir, exist_ok=True) + if _rank() == 0: + with contextlib.suppress(FileNotFoundError): + (checkpoint_dir / _TRAINING_CHECKPOINT_COMPLETE_MARKER).unlink() + _barrier() + states = self._build_states() if _rank() == 0: logger.info( - "Saving checkpoint to %s", + "Saving resumable training checkpoint to %s", checkpoint_dir, ) self._write_metadata(checkpoint_dir, step) dcp.save(states, checkpoint_id=str(dcp_dir)) _barrier() - # Save RNG state AFTER dcp.save so it captures the - # exact state the continuous run continues with. - # dcp.save triggers FSDP all-gather ops that can - # advance the RNG between when DCP captures it and - # when the save completes. self._save_rng_snapshot(checkpoint_dir) _barrier() + if _rank() == 0: + _publish_training_checkpoint_complete(checkpoint_dir) + _barrier() + self._last_saved_step = step self._cleanup_old_checkpoints() + def save_inference(self, step: int) -> None: + """Save one deployable inference checkpoint for the configured role. + + DCP is used only as a temporary, distributed staging format so FSDP2 + ranks never gather the full fp32 model into one process. Rank zero then + streams bounded tensor groups into a bf16/fp16/fp32 modular model + directory and publishes it atomically. + """ + role = str(self.config.inference_role or "student") + modules = self.method.inference_checkpoint_modules(role) + base_model_path = self.method.inference_checkpoint_base_model_path(role) + checkpoint_dir = self._inference_checkpoint_dir(step) + + already_complete: bool | None = None + existing_error: str | None = None + if _rank() == 0: + try: + from fastvideo.train.utils.inference_checkpoint import ( + validate_complete_inference_checkpoint, ) + + already_complete = (validate_complete_inference_checkpoint(checkpoint_dir, step=step) is not None) + except Exception as error: + existing_error = f"{type(error).__name__}: {error}" + if dist.is_available() and dist.is_initialized(): + complete_payload: list[Any] = [already_complete, existing_error] + dist.broadcast_object_list(complete_payload, src=0) + already_complete = bool(complete_payload[0]) + existing_error = complete_payload[1] + if existing_error is not None: + raise RuntimeError(f"Existing inference checkpoint failed validation at step {step}: {existing_error}") + if already_complete: + if _rank() == 0: + logger.info("Inference checkpoint already complete at %s; skipping", checkpoint_dir) + self._last_inference_saved_step = step + return + + staging_dir = self._inference_staging_dir(step) + dcp_dir = staging_dir / "dcp" + export_status_path = staging_dir / "export-status.json" + if _rank() == 0: + shutil.rmtree(staging_dir, ignore_errors=True) + os.makedirs(dcp_dir, exist_ok=True) + _barrier() + + states = {f"roles.{role}.{module_name}": _FullModelState(module) for module_name, module in modules.items()} + if not states: + raise ValueError(f"Inference checkpoint role {role!r} exposes no modules") + + torch_rng = torch.get_rng_state() + python_rng = random.getstate() + numpy_rng = np.random.get_state() + cuda_rng = torch.cuda.get_rng_state() if torch.cuda.is_available() else None + generator = getattr(self.method, "cuda_generator", None) + generator_rng = generator.get_state() if generator is not None else None + try: + if _rank() == 0: + logger.info("Staging inference role %s with DCP at %s", role, dcp_dir) + dcp.save(states, checkpoint_id=str(dcp_dir)) + _barrier() + finally: + torch.set_rng_state(torch_rng) + random.setstate(python_rng) + np.random.set_state(numpy_rng) + if cuda_rng is not None: + torch.cuda.set_rng_state(cuda_rng) + if generator is not None and generator_rng is not None: + generator.set_state(generator_rng) + + export_error: str | None = None + if _rank() == 0: + try: + from fastvideo.train.utils.inference_checkpoint import ( + export_inference_checkpoint, ) + + export_inference_checkpoint( + dcp_dir=dcp_dir, + output_dir=checkpoint_dir, + base_model_path=base_model_path, + role=role, + modules=modules, + dtype=str(self.config.inference_dtype), + step=step, + raw_config=self._raw_config, + ) + except Exception as error: # propagate the rank-zero failure collectively + logger.exception("Inference checkpoint export failed at step %s", step) + export_error = f"{type(error).__name__}: {error}" + status_tmp = export_status_path.with_suffix(".tmp") + status_tmp.write_text( + json.dumps({ + "complete": export_error is None, + "error": export_error + }) + "\n", + encoding="utf-8", + ) + os.replace(status_tmp, export_status_path) + else: + last_log = time.monotonic() + while not export_status_path.is_file(): + time.sleep(2.0) + now = time.monotonic() + if now - last_log >= 60.0: + logger.info("Waiting for rank-zero inference export at step %s", step) + last_log = now + try: + status = json.loads(export_status_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as error: + raise RuntimeError(f"Invalid inference export status at step {step}: {export_status_path}") from error + if status.get("complete") is not True: + export_error = str(status.get("error") or "rank-zero export failed without an error message") + if export_error is not None: + raise RuntimeError(f"Inference checkpoint export failed at step {step}: {export_error}") + + _barrier() + if _rank() == 0: + shutil.rmtree(staging_dir, ignore_errors=True) + staging_root = staging_dir.parent + with contextlib.suppress(OSError): + staging_root.rmdir() + _barrier() + self._last_inference_saved_step = step + def _write_metadata( self, checkpoint_dir: Path, @@ -328,13 +579,13 @@ def load_rng_snapshot( resolved = _resolve_resume_checkpoint( checkpoint_path, output_dir=self.output_dir, + require_complete_marker=self.config.require_complete_training_checkpoint, ) if resolved is None: return rank = _rank() rng_path = resolved / f"rng_state_rank{rank}.pt" if not rng_path.is_file(): - # Fall back to legacy single-file snapshot. rng_path = resolved / "rng_state.pt" if not rng_path.is_file(): logger.warning( @@ -370,6 +621,7 @@ def maybe_resume(self, *, resume_from_checkpoint: str | None) -> int | None: resolved = _resolve_resume_checkpoint( resume_from_checkpoint, output_dir=self.output_dir, + require_complete_marker=self.config.require_complete_training_checkpoint, ) if resolved is None: return None @@ -398,6 +650,9 @@ def _cleanup_old_checkpoints(self) -> None: continue if not _CHECKPOINT_DIR_RE.match(child.name): continue + if (self.config.require_complete_training_checkpoint + and not _is_complete_training_checkpoint(child, require_complete_marker=True)): + continue try: step = _parse_step_from_dir(child) except Exception: diff --git a/fastvideo/train/utils/config.py b/fastvideo/train/utils/config.py index d396b02585..b2561bc51e 100644 --- a/fastvideo/train/utils/config.py +++ b/fastvideo/train/utils/config.py @@ -65,7 +65,6 @@ def _safe_asdict(obj: Any) -> Any: return resolved -# ---- parsing helpers (kept for use by methods) ---- def _resolve_existing_file(path: str) -> str: @@ -138,7 +137,6 @@ def parse_betas(raw: Any, *, where: str) -> tuple[float, float]: f"got {type(raw).__name__}") -# ---- config convenience helpers ---- def require_positive_int( @@ -264,8 +262,6 @@ def _parse_pipeline_config( pipeline_raw, dit_arch_overrides = _split_training_dit_arch_overrides(pipeline_raw) - # Derive model_path from models.student.init_from — - # needed by PipelineConfig.from_kwargs. model_path: str | None = None student_cfg = models.get("student") if student_cfg is not None: @@ -382,6 +378,32 @@ def _build_training_config( "{'t2v', 't2va', 'text_only'}, got " f"{preprocessed_data_type!r}") + vsa_tile_size = int(vs.get("tile_size", 256) or 256) + if vsa_tile_size not in (64, 256): + raise ValueError(f"training.vsa.tile_size must be 64 or 256, got {vsa_tile_size!r}") + + save_inference_checkpoint_on_validation = require_bool( + ck, + "save_inference_checkpoint_on_validation", + default=False, + where="training.checkpoint.save_inference_checkpoint_on_validation", + ) + inference_checkpoint_role = str(ck.get("inference_checkpoint_role", "student") or "").strip() + if save_inference_checkpoint_on_validation and not inference_checkpoint_role: + raise ValueError("training.checkpoint.inference_checkpoint_role must be non-empty " + "when validation inference checkpointing is enabled") + inference_checkpoint_dtype = str(ck.get("inference_checkpoint_dtype", "bfloat16") or "bfloat16").strip().lower() + if inference_checkpoint_dtype not in {"bfloat16", "float16", "float32"}: + raise ValueError("training.checkpoint.inference_checkpoint_dtype must be one of " + "['bfloat16', 'float16', 'float32'], got " + f"{inference_checkpoint_dtype!r}") + require_complete_training_checkpoint = require_bool( + ck, + "require_complete_training_checkpoint", + default=False, + where="training.checkpoint.require_complete_training_checkpoint", + ) + return TrainingConfig( distributed=DistributedConfig( num_gpus=num_gpus, @@ -402,6 +424,7 @@ def _build_training_config( num_width=int(da.get("num_width", 0) or 0), num_latent_t=int(da.get("num_latent_t", 0) or 0), num_frames=int(da.get("num_frames", 0) or 0), + native_shape_bucketing=bool(da.get("native_shape_bucketing", False)), ), optimizer=OptimizerConfig( learning_rate=float(o.get("learning_rate", 0.0) or 0.0), @@ -420,8 +443,14 @@ def _build_training_config( checkpoint=CheckpointConfig( output_dir=str(ck.get("output_dir", "") or ""), resume_from_checkpoint=str(ck.get("resume_from_checkpoint", "") or ""), + save_inference_checkpoint_on_validation=save_inference_checkpoint_on_validation, + inference_checkpoint_role=inference_checkpoint_role or "student", + inference_checkpoint_dtype=inference_checkpoint_dtype, training_state_checkpointing_steps=int(ck.get("training_state_checkpointing_steps", 0) or 0), + require_complete_training_checkpoint=require_complete_training_checkpoint, checkpoints_total_limit=int(ck.get("checkpoints_total_limit", 0) or 0), + checkpointing_start_step=int(ck.get("checkpointing_start_step", 0) or 0), + reset_lr_on_resume=bool(ck.get("reset_lr_on_resume", False)), ), tracker=TrackerConfig( trackers=list(tr.get("trackers", []) or []), @@ -430,6 +459,7 @@ def _build_training_config( run_name=str(tr.get("run_name", "") or ""), ), vsa_sparsity=float(vs.get("sparsity", 0.0) or 0.0), + vsa_tile_size=vsa_tile_size, vsa_cache_tile_buf=bool(vs.get("cache_tile_buf", False) or False), model=ModelTrainingConfig( weighting_scheme=str(m.get("weighting_scheme", "uniform") or "uniform"), @@ -439,6 +469,9 @@ 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")), + allow_low_precision_master_weights=bool(m.get("allow_low_precision_master_weights", False)), + enable_torch_compile=bool(m.get("enable_torch_compile", False)), + torch_compile_kwargs=dict(m.get("torch_compile_kwargs", {}) or {}), ), pipeline_config=pipeline_config, model_path=model_path, @@ -481,17 +514,14 @@ def _cast_value(raw: str) -> Any: return False if raw.lower() in ("none", "null"): return None - # Try int try: return int(raw) except ValueError: pass - # Try float try: return float(raw) except ValueError: pass - # Try YAML list literal like [1, 2] if raw.startswith("[") and raw.endswith("]"): try: return yaml.safe_load(raw) @@ -536,13 +566,11 @@ def load_run_config( raw = yaml.safe_load(f) cfg = _require_mapping(raw, where=path) - # Apply CLI overrides before building typed config. if overrides: parsed = _parse_cli_overrides(overrides) _apply_overrides(cfg, parsed) logger.info("Applied CLI overrides: %s", parsed) - # --- models --- models_raw = _require_mapping(cfg.get("models"), where="models") models: dict[str, dict[str, Any]] = {} for role, model_cfg_raw in models_raw.items(): @@ -553,23 +581,19 @@ def load_run_config( "'_target_' key") models[role_str] = dict(model_cfg) - # --- method --- method_raw = _require_mapping(cfg.get("method"), where="method") if "_target_" not in method_raw: raise ValueError("method must have a '_target_' key") method = dict(method_raw) - # --- callbacks --- callbacks_raw = cfg.get("callbacks", None) if callbacks_raw is None: callbacks: dict[str, dict[str, Any]] = {} else: callbacks = _require_mapping(callbacks_raw, where="callbacks") - # --- pipeline config --- pipeline_config = _parse_pipeline_config(cfg, models=models) - # --- training config --- training_raw = _require_mapping(cfg.get("training"), where="training") t = dict(training_raw) training = _build_training_config(t, models=models, pipeline_config=pipeline_config) diff --git a/fastvideo/train/utils/dataloader.py b/fastvideo/train/utils/dataloader.py index 44292fafd8..767f112f7d 100644 --- a/fastvideo/train/utils/dataloader.py +++ b/fastvideo/train/utils/dataloader.py @@ -29,6 +29,7 @@ def build_parquet_t2v_train_dataloader( drop_last=True, text_padding_length=int(text_len), seed=int(data_config.seed or 0), + native_shape_bucketing=bool(data_config.native_shape_bucketing), )) return dataloader @@ -53,4 +54,4 @@ def build_parquet_matrixgame2_train_dataloader( text_padding_length=512, seed=int(data_config.seed or 0), )) - return dataloader \ No newline at end of file + return dataloader diff --git a/fastvideo/train/utils/inference_checkpoint.py b/fastvideo/train/utils/inference_checkpoint.py new file mode 100644 index 0000000000..e4895a7366 --- /dev/null +++ b/fastvideo/train/utils/inference_checkpoint.py @@ -0,0 +1,619 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Bounded-memory export of model-only DCP state for inference. + +The modular trainer stores resumable state with PyTorch Distributed +Checkpoint (DCP), while FastVideo inference consumes a Diffusers-style model +directory. This module bridges those formats without gathering a full model +in memory: rank 0 reads one bounded shard at a time from an already-complete +role/module-only DCP checkpoint and publishes an immutable inference directory +with an atomic rename. + +The manager-facing entry point is :func:`export_inference_checkpoint`; the +lower-level :func:`export_inference_checkpoint_from_dcp` exposes the same +rank-0-only conversion with a run-root destination. The checkpoint manager +owns distributed coordination around the temporary DCP save and local export. +""" + +from __future__ import annotations + +import json +import os +import re +import shutil +import tempfile +from collections.abc import Mapping +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import torch +import torch.distributed as dist +import torch.distributed.checkpoint as dcp +from safetensors import safe_open +from safetensors.torch import save_file +from torch.distributed.checkpoint import FileSystemReader + +DEFAULT_MAX_SHARD_SIZE_BYTES = 5 * 1024**3 +_FORMAT_VERSION = 1 +_ALLOWED_NATIVE_EXTRA_KEYS = (re.compile(r"(?:^|\.)attn\.to_gate_compress\.weight$"), ) + + +class InferenceCheckpointExportError(RuntimeError): + """The model-only checkpoint could not be exported safely.""" + + +class UnsupportedMergedReverseMappingError(InferenceCheckpointExportError): + """A fused training parameter would need to be split for inference.""" + + +def _is_allowed_native_extra(key: str) -> bool: + return any(pattern.search(key) for pattern in _ALLOWED_NATIVE_EXTRA_KEYS) + + +@dataclass(frozen=True, slots=True) +class _TensorPlan: + checkpoint_key: str + internal_key: str + output_key: str + shape: tuple[int, ...] + source_dtype: torch.dtype + output_dtype: torch.dtype + output_nbytes: int + + +def _rank() -> int: + if dist.is_available() and dist.is_initialized(): + return int(dist.get_rank()) + return 0 + + +def _normalize_output_dtype(dtype: torch.dtype | str) -> torch.dtype: + if isinstance(dtype, str): + name = dtype.removeprefix("torch.") + resolved = getattr(torch, name, None) + if not isinstance(resolved, torch.dtype): + raise ValueError(f"Unsupported inference checkpoint dtype: {dtype!r}") + dtype = resolved + if not isinstance(dtype, torch.dtype): + raise TypeError("dtype must be a torch.dtype or torch dtype name") + if not torch.empty((), dtype=dtype).is_floating_point(): + raise ValueError(f"Inference checkpoint dtype must be floating point, got {dtype}") + return dtype + + +def _resolve_dcp_dir(checkpoint: str | os.PathLike[str]) -> Path: + path = Path(checkpoint).expanduser().resolve() + if path.name != "dcp" and (path / "dcp").is_dir(): + path = path / "dcp" + if not path.is_dir(): + raise FileNotFoundError(f"Inference checkpoint DCP directory not found: {path}") + if not (path / ".metadata").is_file(): + raise FileNotFoundError(f"Incomplete inference checkpoint DCP (missing .metadata): {path}") + return path + + +def validate_complete_inference_checkpoint(path: Path, *, step: int) -> Path | None: + """Validate a published inference checkpoint without loading its tensors.""" + if not (path.exists() or path.is_symlink()): + return None + complete_path = path / ".complete" + metadata_path = path / "metadata.json" + if not complete_path.is_file() or not metadata_path.is_file(): + raise InferenceCheckpointExportError(f"Refusing to overwrite incomplete inference checkpoint: {path}") + try: + complete_marker = complete_path.read_text(encoding="utf-8") + except OSError as exc: + raise InferenceCheckpointExportError(f"Cannot read inference completion marker: {complete_path}") from exc + if complete_marker != "complete\n": + raise InferenceCheckpointExportError(f"Invalid inference completion marker: {complete_path}") + try: + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise InferenceCheckpointExportError( + f"Invalid metadata for existing inference checkpoint: {metadata_path}") from exc + if (metadata.get("format_version") != _FORMAT_VERSION or metadata.get("kind") != "inference" + or metadata.get("step") != step): + raise InferenceCheckpointExportError( + f"Existing inference checkpoint metadata does not match step={step}: {metadata_path}") + role = metadata.get("role") + dtype = metadata.get("dtype") + if not isinstance(role, str) or not role or "." in role or dtype not in {"bfloat16", "float16", "float32"}: + raise InferenceCheckpointExportError(f"Inference checkpoint role/dtype metadata is invalid: {metadata_path}") + module_name = str(metadata.get("module") or "") + if not module_name or "/" in module_name or "\\" in module_name or module_name in {".", ".."}: + raise InferenceCheckpointExportError( + f"Inference checkpoint has an invalid module name {module_name!r}: {metadata_path}") + module_dir = path / module_name + index_path = module_dir / "diffusion_pytorch_model.safetensors.index.json" + if not index_path.is_file(): + raise InferenceCheckpointExportError(f"Inference checkpoint is missing its module index: {index_path}") + if not (module_dir / "config.json").is_file(): + raise InferenceCheckpointExportError(f"Inference checkpoint is missing its module config: {module_dir}") + if not any((path / name).is_file() for name in ("model_index.json", "modular_model_index.json")): + raise InferenceCheckpointExportError(f"Inference checkpoint is missing its model index: {path}") + try: + index = json.loads(index_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise InferenceCheckpointExportError(f"Invalid inference checkpoint index: {index_path}") from exc + weight_map = index.get("weight_map") + if not isinstance(weight_map, dict) or not weight_map: + raise InferenceCheckpointExportError(f"Inference checkpoint index has no weight map: {index_path}") + expected_by_shard: dict[str, set[str]] = {} + for key, filename in weight_map.items(): + if not isinstance(key, str) or not key or not isinstance(filename, str) or not filename: + raise InferenceCheckpointExportError(f"Invalid key/shard entry in {index_path}: {key!r} -> {filename!r}") + if Path(filename).name != filename: + raise InferenceCheckpointExportError( + f"Inference checkpoint index contains an invalid shard path: {filename!r}") + expected_by_shard.setdefault(filename, set()).add(key) + actual_shards = {shard.name for shard in module_dir.glob("*.safetensors") if shard.is_file()} + expected_shards = set(expected_by_shard) + if actual_shards != expected_shards: + raise InferenceCheckpointExportError( + f"Inference checkpoint shards differ from its index under {module_dir}: " + f"missing={sorted(expected_shards - actual_shards)} extra={sorted(actual_shards - expected_shards)}") + shard_sizes: list[int] = [] + output_shapes: dict[str, tuple[int, ...]] = {} + for filename in sorted(expected_by_shard): + expected_keys = expected_by_shard[filename] + shard_path = module_dir / filename + try: + with safe_open(str(shard_path), framework="pt", device="cpu") as handle: + actual_keys = set(handle.keys()) + for key in actual_keys: + output_shapes[key] = tuple(int(dim) for dim in handle.get_slice(key).get_shape()) + except Exception as exc: + raise InferenceCheckpointExportError(f"Cannot read inference checkpoint shard: {shard_path}") from exc + if actual_keys != expected_keys: + raise InferenceCheckpointExportError( + f"Inference checkpoint shard keys differ from its index for {shard_path}: " + f"missing={sorted(expected_keys - actual_keys)[:10]} extra={sorted(actual_keys - expected_keys)[:10]}") + shard_sizes.append(shard_path.stat().st_size) + if int(metadata.get("tensor_count", -1)) != len(weight_map): + raise InferenceCheckpointExportError( + f"Inference checkpoint tensor_count does not match its index: {metadata_path}") + if int(metadata.get("shard_count", -1)) != len(expected_shards): + raise InferenceCheckpointExportError( + f"Inference checkpoint shard_count does not match its index: {metadata_path}") + logical_shard_sizes = metadata.get("shard_sizes") + total_size = metadata.get("total_size") + max_shard_size = metadata.get("max_shard_size_bytes") + index_metadata = index.get("metadata") + index_total_size = index_metadata.get("total_size") if isinstance(index_metadata, dict) else None + if (not isinstance(logical_shard_sizes, list) or len(logical_shard_sizes) != len(expected_shards) + or any(not isinstance(size, int) or size < 0 + for size in logical_shard_sizes) or not isinstance(max_shard_size, int) or max_shard_size <= 0 + or any(size > max_shard_size for size in logical_shard_sizes) or not isinstance(total_size, int) + or total_size != sum(logical_shard_sizes) or index_total_size != total_size): + raise InferenceCheckpointExportError( + f"Inference checkpoint logical shard sizes are inconsistent: {metadata_path}") + recorded_shard_sizes = metadata.get("shard_file_sizes") + if not isinstance(recorded_shard_sizes, list) or recorded_shard_sizes != shard_sizes: + raise InferenceCheckpointExportError( + f"Inference checkpoint shard_file_sizes do not match files on disk: {metadata_path}") + + base_model_dir = metadata.get("base_model_dir") + if not isinstance(base_model_dir, str) or not base_model_dir: + raise InferenceCheckpointExportError(f"Inference checkpoint has no base_model_dir: {metadata_path}") + base_shapes = _component_tensor_shapes(Path(base_model_dir) / module_name) + for key, shape in output_shapes.items(): + expected_shape = base_shapes.get(key) + if expected_shape is None: + if not _is_allowed_native_extra(key): + raise InferenceCheckpointExportError( + f"Inference checkpoint contains unknown non-native tensor {key!r}: {path}") + elif shape != expected_shape: + raise InferenceCheckpointExportError( + f"Inference checkpoint tensor {key!r} shape {shape} != base transformer shape {expected_shape}") + missing_base_keys = set(base_shapes) - set(output_shapes) + if missing_base_keys: + raise InferenceCheckpointExportError( + f"Inference checkpoint is missing {len(missing_base_keys)} base transformer tensors; " + f"first={sorted(missing_base_keys)[:10]}") + return path + + +def _component_tensor_shapes(module_dir: Path) -> dict[str, tuple[int, ...]]: + """Read the base component's exact key/shape contract from safetensors headers.""" + index_candidates = ( + module_dir / "diffusion_pytorch_model.safetensors.index.json", + module_dir / "model.safetensors.index.json", + ) + index_path = next((path for path in index_candidates if path.is_file()), None) + expected_files: set[str] | None = None + indexed_keys: set[str] | None = None + if index_path is not None: + try: + index = json.loads(index_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise InferenceCheckpointExportError(f"Invalid base transformer index: {index_path}") from exc + weight_map = index.get("weight_map") + if not isinstance(weight_map, dict) or not weight_map: + raise InferenceCheckpointExportError(f"Base transformer index has no weight map: {index_path}") + indexed_keys = set(weight_map) + expected_files = {str(filename) for filename in weight_map.values()} + + files = sorted(module_dir.glob("*.safetensors")) + if expected_files is not None: + files = [module_dir / filename for filename in sorted(expected_files)] + if not files or any(not path.is_file() for path in files): + raise InferenceCheckpointExportError(f"Base transformer safetensors are incomplete under {module_dir}") + + shapes: dict[str, tuple[int, ...]] = {} + for path in files: + try: + with safe_open(str(path), framework="pt", device="cpu") as handle: + for key in handle.keys(): # noqa: SIM118 + if key in shapes: + raise InferenceCheckpointExportError(f"Duplicate base transformer tensor key {key!r}") + shapes[key] = tuple(int(dim) for dim in handle.get_slice(key).get_shape()) + except InferenceCheckpointExportError: + raise + except Exception as exc: + raise InferenceCheckpointExportError(f"Cannot read base transformer safetensors: {path}") from exc + if indexed_keys is not None and set(shapes) != indexed_keys: + raise InferenceCheckpointExportError( + f"Base transformer index/header mismatch under {module_dir}: " + f"missing={sorted(indexed_keys - set(shapes))[:10]} extra={sorted(set(shapes) - indexed_keys)[:10]}") + return shapes + + +def _mapping_output_key( + internal_key: str, + reverse_mapping: Mapping[str, Any], +) -> str: + entry = reverse_mapping.get(internal_key) + if entry is None: + return internal_key + if not isinstance(entry, tuple | list) or len(entry) != 3: + raise InferenceCheckpointExportError(f"Invalid reverse mapping for {internal_key!r}: expected " + "(output_key, merge_index, num_params_to_merge)") + output_key, merge_index, num_params_to_merge = entry + if merge_index is not None or num_params_to_merge not in (None, 1): + raise UnsupportedMergedReverseMappingError(f"Cannot stream merged reverse mapping for {internal_key!r}: " + f"output_key={output_key!r}, merge_index={merge_index!r}, " + f"num_params_to_merge={num_params_to_merge!r}. " + "A model-specific split exporter is required.") + if not isinstance(output_key, str) or not output_key: + raise InferenceCheckpointExportError( + f"Invalid output key in reverse mapping for {internal_key!r}: {output_key!r}") + return output_key + + +def _tensor_output_dtype(source_dtype: torch.dtype, configured_dtype: torch.dtype) -> torch.dtype: + if torch.empty((), dtype=source_dtype).is_floating_point(): + return configured_dtype + return source_dtype + + +def _build_tensor_plan( + *, + dcp_dir: Path, + state_prefix: str, + reverse_mapping: Mapping[str, Any], + base_shapes: Mapping[str, tuple[int, ...]], + output_dtype: torch.dtype, + max_shard_size_bytes: int, +) -> list[_TensorPlan]: + metadata = FileSystemReader(str(dcp_dir)).read_metadata() + plans: list[_TensorPlan] = [] + output_keys: set[str] = set() + + for checkpoint_key in sorted(metadata.state_dict_metadata): + if not checkpoint_key.startswith(state_prefix): + continue + tensor_metadata = metadata.state_dict_metadata[checkpoint_key] + properties = getattr(tensor_metadata, "properties", None) + shape = getattr(tensor_metadata, "size", None) + source_dtype = getattr(properties, "dtype", None) + if shape is None or not isinstance(source_dtype, torch.dtype): + raise InferenceCheckpointExportError(f"Inference state contains a non-tensor value at {checkpoint_key!r}; " + "safetensors exports support tensors only") + + internal_key = checkpoint_key[len(state_prefix):] + if not internal_key: + raise InferenceCheckpointExportError(f"Empty module key under DCP prefix {state_prefix!r}") + output_key = _mapping_output_key(internal_key, reverse_mapping) + tensor_shape = tuple(int(dim) for dim in shape) + expected_shape = base_shapes.get(output_key) + if expected_shape is None: + if internal_key in reverse_mapping or not _is_allowed_native_extra(output_key): + raise InferenceCheckpointExportError( + f"Inference tensor {internal_key!r} maps to unknown base key {output_key!r}; " + "only MiniMax-H3 attn.to_gate_compress.weight may be exported as a native extra") + elif tensor_shape != expected_shape: + raise InferenceCheckpointExportError( + f"Inference tensor {output_key!r} shape {tensor_shape} != base transformer shape {expected_shape}") + if output_key in output_keys: + raise InferenceCheckpointExportError(f"Reverse mapping produces duplicate inference key {output_key!r}") + output_keys.add(output_key) + + tensor_output_dtype = _tensor_output_dtype(source_dtype, output_dtype) + numel = 1 + for dim in tensor_shape: + numel *= dim + output_nbytes = numel * torch.empty((), dtype=tensor_output_dtype).element_size() + if output_nbytes > max_shard_size_bytes: + raise InferenceCheckpointExportError( + f"Tensor {checkpoint_key!r} requires {output_nbytes} bytes after casting, " + f"which exceeds max_shard_size_bytes={max_shard_size_bytes}") + plans.append( + _TensorPlan( + checkpoint_key=checkpoint_key, + internal_key=internal_key, + output_key=output_key, + shape=tensor_shape, + source_dtype=source_dtype, + output_dtype=tensor_output_dtype, + output_nbytes=output_nbytes, + )) + + if not plans: + raise InferenceCheckpointExportError(f"No tensor keys found under DCP prefix {state_prefix!r} in {dcp_dir}") + missing_base_keys = set(base_shapes) - output_keys + if missing_base_keys: + raise InferenceCheckpointExportError( + f"Inference checkpoint is missing {len(missing_base_keys)} base transformer tensors; " + f"first={sorted(missing_base_keys)[:10]}") + return plans + + +def _group_shards(plans: list[_TensorPlan], max_shard_size_bytes: int) -> list[list[_TensorPlan]]: + shards: list[list[_TensorPlan]] = [] + current: list[_TensorPlan] = [] + current_nbytes = 0 + for plan in plans: + if current and current_nbytes + plan.output_nbytes > max_shard_size_bytes: + shards.append(current) + current = [] + current_nbytes = 0 + current.append(plan) + current_nbytes += plan.output_nbytes + if current: + shards.append(current) + return shards + + +def _prepare_model_layout(temp_dir: Path, base_model_dir: Path, module_name: str) -> Path: + if not any((base_model_dir / name).is_file() for name in ("model_index.json", "modular_model_index.json")): + raise FileNotFoundError( + f"Base model directory has no model_index.json or modular_model_index.json: {base_model_dir}") + base_module_dir = base_model_dir / module_name + base_config = base_module_dir / "config.json" + if not base_config.is_file(): + raise FileNotFoundError(f"Base model component config not found: {base_config}") + + module_dir = temp_dir / module_name + module_dir.mkdir(parents=True) + shutil.copy2(base_config, module_dir / "config.json") + + reserved = {module_name, "metadata.json", ".complete"} + for entry in sorted(base_model_dir.iterdir(), key=lambda item: item.name): + if entry.name in reserved or entry.name == ".cache" or entry.name.startswith(".git"): + continue + target = temp_dir / entry.name + target.symlink_to(entry.resolve(), target_is_directory=entry.is_dir()) + return module_dir + + +def _write_tensor_shards( + *, + dcp_dir: Path, + module_dir: Path, + shards: list[list[_TensorPlan]], +) -> tuple[dict[str, str], int, list[int], list[int]]: + weight_map: dict[str, str] = {} + total_size = 0 + shard_sizes: list[int] = [] + shard_file_sizes: list[int] = [] + shard_count = len(shards) + + for shard_index, shard in enumerate(shards, start=1): + filename = (f"diffusion_pytorch_model-{shard_index:05d}-of-{shard_count:05d}.safetensors") + state = {plan.checkpoint_key: torch.empty(plan.shape, dtype=plan.source_dtype, device="cpu") for plan in shard} + dcp.load(state, checkpoint_id=str(dcp_dir), no_dist=True) + + output_tensors: dict[str, torch.Tensor] = {} + shard_size = 0 + for plan in shard: + tensor = state[plan.checkpoint_key] + if tensor.is_floating_point(): + tensor = tensor.to(dtype=plan.output_dtype) + output_tensors[plan.output_key] = tensor.detach().cpu().contiguous() + weight_map[plan.output_key] = filename + shard_size += plan.output_nbytes + save_file(output_tensors, module_dir / filename) + total_size += shard_size + shard_sizes.append(shard_size) + shard_file_sizes.append((module_dir / filename).stat().st_size) + del output_tensors, state + + index = { + "metadata": { + "total_size": total_size + }, + "weight_map": weight_map, + } + index_path = module_dir / "diffusion_pytorch_model.safetensors.index.json" + index_path.write_text(json.dumps(index, indent=2, sort_keys=True) + "\n", encoding="utf-8") + return weight_map, total_size, shard_sizes, shard_file_sizes + + +def export_inference_checkpoint_from_dcp( + *, + dcp_checkpoint: str | os.PathLike[str], + output_dir: str | os.PathLike[str], + step: int, + module: torch.nn.Module, + base_model_dir: str | os.PathLike[str], + role: str = "student", + module_name: str = "transformer", + dtype: torch.dtype | str = torch.bfloat16, + max_shard_size_bytes: int = DEFAULT_MAX_SHARD_SIZE_BYTES, + raw_config: Mapping[str, Any] | None = None, +) -> Path: + """Export one role/module-only DCP as an immutable inference checkpoint. + + ``CheckpointManager`` is expected to call this function on rank 0 after a + collective model-only DCP save has completed. ``dcp_checkpoint`` may name + that DCP directory directly or its parent containing ``dcp/``. The source + is never modified or removed. + + Floating tensors are cast into ``dtype`` in independent CPU buffers; + integer and boolean tensors retain their source dtype. Parameter names are + converted with ``module.reverse_param_names_mapping``. Unmapped names are + retained for FastVideo-native inference parameters such as MiniMax-H3 VSA + gates. Merged mappings fail rather than silently writing incompatible + weights. + + The completed model is atomically renamed to + ``/inference/checkpoint-``. A valid existing completed + export is returned unchanged, making retries idempotent. + """ + + if _rank() != 0: + raise InferenceCheckpointExportError("export_inference_checkpoint_from_dcp is rank-0-only; " + "the checkpoint manager must coordinate other ranks") + if isinstance(step, bool) or not isinstance(step, int) or step < 0: + raise ValueError(f"step must be a non-negative integer, got {step!r}") + if not role or "." in role: + raise ValueError(f"role must be a non-empty DCP key segment, got {role!r}") + if not module_name or "." in module_name: + raise ValueError(f"module_name must be a non-empty DCP key segment, got {module_name!r}") + if not 0 < max_shard_size_bytes <= DEFAULT_MAX_SHARD_SIZE_BYTES: + raise ValueError("max_shard_size_bytes must be in " + f"[1, {DEFAULT_MAX_SHARD_SIZE_BYTES}], got {max_shard_size_bytes}") + + output_dtype = _normalize_output_dtype(dtype) + run_output_dir = Path(output_dir).expanduser().resolve() + inference_root = run_output_dir / "inference" + final_dir = inference_root / f"checkpoint-{step}" + existing = validate_complete_inference_checkpoint(final_dir, step=step) + if existing is not None: + return existing + + dcp_dir = _resolve_dcp_dir(dcp_checkpoint) + base_dir = Path(base_model_dir).expanduser().resolve() + base_shapes = _component_tensor_shapes(base_dir / module_name) + reverse_mapping = getattr(module, "reverse_param_names_mapping", {}) + if reverse_mapping is None: + reverse_mapping = {} + if not isinstance(reverse_mapping, Mapping): + raise InferenceCheckpointExportError("module.reverse_param_names_mapping must be a mapping") + + state_prefix = f"roles.{role}.{module_name}." + plans = _build_tensor_plan( + dcp_dir=dcp_dir, + state_prefix=state_prefix, + reverse_mapping=reverse_mapping, + base_shapes=base_shapes, + output_dtype=output_dtype, + max_shard_size_bytes=max_shard_size_bytes, + ) + shards = _group_shards(plans, max_shard_size_bytes) + + inference_root.mkdir(parents=True, exist_ok=True) + temp_dir = Path(tempfile.mkdtemp( + prefix=f".checkpoint-{step}.tmp-", + dir=str(inference_root), + )) + try: + module_dir = _prepare_model_layout(temp_dir, base_dir, module_name) + weight_map, total_size, shard_sizes, shard_file_sizes = _write_tensor_shards( + dcp_dir=dcp_dir, + module_dir=module_dir, + shards=shards, + ) + metadata = { + "format_version": _FORMAT_VERSION, + "kind": "inference", + "step": step, + "role": role, + "module": module_name, + "dtype": str(output_dtype).removeprefix("torch."), + "base_model_dir": str(base_dir), + "tensor_count": len(weight_map), + "total_size": total_size, + "shard_count": len(shards), + "shard_sizes": shard_sizes, + "shard_file_sizes": shard_file_sizes, + "max_shard_size_bytes": max_shard_size_bytes, + } + if raw_config is not None: + metadata["config"] = raw_config + (temp_dir / "metadata.json").write_text( + json.dumps(metadata, indent=2, sort_keys=True) + "\n", + encoding="utf-8", + ) + (temp_dir / ".complete").write_text("complete\n", encoding="utf-8") + try: + temp_dir.rename(final_dir) + except FileExistsError as exc: + raise InferenceCheckpointExportError(f"Inference checkpoint appeared concurrently: {final_dir}") from exc + except BaseException: + shutil.rmtree(temp_dir, ignore_errors=True) + raise + validated = validate_complete_inference_checkpoint(final_dir, step=step) + if validated is None: + raise InferenceCheckpointExportError(f"Published inference checkpoint disappeared: {final_dir}") + return validated + + +def export_inference_checkpoint( + *, + dcp_dir: str | os.PathLike[str], + output_dir: str | os.PathLike[str], + base_model_path: str | os.PathLike[str], + role: str, + modules: Mapping[str, torch.nn.Module], + dtype: torch.dtype | str, + step: int, + raw_config: Mapping[str, Any] | None = None, +) -> Path: + """CheckpointManager adapter for one deployable role/module checkpoint. + + ``output_dir`` is the manager's fully resolved target, + ``/inference/checkpoint-``. The temporary DCP may contain only + the role/module state selected by ``modules``. Multi-component deployment + is intentionally rejected until its component layout and atomicity + contract are defined. + + ``raw_config`` is persisted in ``metadata.json`` when supplied, matching + resumable checkpoint provenance. The ephemeral DCP staging path is not + persisted because the manager removes it after a successful export. + """ + + if len(modules) != 1: + raise InferenceCheckpointExportError("Inference checkpoint export currently supports exactly one module; " + f"got {sorted(modules)}") + module_name, module = next(iter(modules.items())) + if not isinstance(module, torch.nn.Module): + raise TypeError(f"Inference checkpoint module {module_name!r} must be a torch.nn.Module") + + final_dir = Path(output_dir).expanduser().resolve() + expected_name = f"checkpoint-{step}" + if final_dir.name != expected_name or final_dir.parent.name != "inference": + raise ValueError("CheckpointManager output_dir must be " + f"/inference/{expected_name}, got {final_dir}") + run_output_dir = final_dir.parent.parent + return export_inference_checkpoint_from_dcp( + dcp_checkpoint=dcp_dir, + output_dir=run_output_dir, + step=step, + module=module, + base_model_dir=base_model_path, + role=role, + module_name=module_name, + dtype=dtype, + raw_config=raw_config, + ) + + +__all__ = [ + "DEFAULT_MAX_SHARD_SIZE_BYTES", + "InferenceCheckpointExportError", + "UnsupportedMergedReverseMappingError", + "export_inference_checkpoint", + "export_inference_checkpoint_from_dcp", + "validate_complete_inference_checkpoint", +] diff --git a/fastvideo/train/utils/moduleloader.py b/fastvideo/train/utils/moduleloader.py index 822e41c17a..d37ffc9c9a 100644 --- a/fastvideo/train/utils/moduleloader.py +++ b/fastvideo/train/utils/moduleloader.py @@ -5,12 +5,15 @@ import os from contextlib import nullcontext from typing import Any, TYPE_CHECKING +from collections.abc import Callable import torch from fastvideo.attention.selector import ( + NO_REQUEST, _component_attention_backend_scope, coerce_attn_backend, + component_attention_backend, ) from fastvideo.configs.pipelines.base import PipelineConfig from fastvideo.fastvideo_args import ExecutionMode, TrainingArgs @@ -26,9 +29,6 @@ from fastvideo.train.utils.training_config import ( TrainingConfig, ) -# ------------------------------------------------------------------ -# TrainingArgs builders (only place that creates FastVideoArgs) -# ------------------------------------------------------------------ def _make_training_args( @@ -38,9 +38,6 @@ def _make_training_args( ) -> TrainingArgs: """Build a TrainingArgs for PipelineComponentLoader.""" pipeline_config = tc.pipeline_config or PipelineConfig() - # Propagate dit_precision from TrainingConfig to PipelineConfig - # so that TransformerLoader.load() picks up the correct - # default_dtype (e.g. fp32 master weights for training). if tc.dit_precision and tc.dit_precision != pipeline_config.dit_precision: pipeline_config.dit_precision = tc.dit_precision return TrainingArgs( @@ -60,7 +57,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, + regional_compile=True, + torch_compile_kwargs=tc.model.torch_compile_kwargs, ) @@ -73,14 +72,12 @@ def make_inference_args( args = _make_training_args(tc, model_path=model_path) args.inference_mode = True args.mode = ExecutionMode.INFERENCE - args.dit_cpu_offload = True + args.dit_cpu_offload = False args.VSA_sparsity = tc.vsa_sparsity + args.VSA_tile_size = tc.vsa_tile_size return args -# ------------------------------------------------------------------ -# Module loading -# ------------------------------------------------------------------ def load_module_from_path( @@ -92,6 +89,8 @@ def load_module_from_path( override_transformer_cls_name: str | None = None, transformer_override_safetensor: str | None = None, attention_backend: AttentionBackendEnum | str | None = None, + construction_precision: str | None = None, + pre_fsdp_transform: Callable[[torch.nn.Module], torch.nn.Module] | None = None, ) -> torch.nn.Module: """Load one pipeline component with its role-scoped attention policy. @@ -104,6 +103,9 @@ def load_module_from_path( scoped to this load call. """ fastvideo_args: Any = _make_training_args(training_config, model_path=model_path) + original_dit_precision = fastvideo_args.pipeline_config.dit_precision + if construction_precision is not None: + fastvideo_args.pipeline_config.dit_precision = str(construction_precision) local_model_path = maybe_download_model(model_path) config = verify_model_config_and_directory(local_model_path) @@ -117,43 +119,53 @@ def load_module_from_path( raise ValueError(f"Module {module_type!r} has null value in " f"config at {local_model_path}") - # Trailing modular-manifest metadata does not change component dispatch; - # the provider and architecture remain the first two fields. transformers_or_diffusers, _architecture = module_info[:2] component_path = os.path.join(local_model_path, module_type) - # fastvideo_args is freshly built above and never escapes this function, - # so overrides are plain assignments — nothing to save or restore. if override_transformer_cls_name is not None: fastvideo_args.override_transformer_cls_name = str(override_transformer_cls_name) if transformer_override_safetensor: fastvideo_args.init_weights_from_safetensors = str(transformer_override_safetensor) + if pre_fsdp_transform is not None: + if module_type != "transformer": + raise ValueError("pre_fsdp_transform can only be set when loading " + f"a transformer, got module_type={module_type!r}") + fastvideo_args._pre_fsdp_transform = pre_fsdp_transform + if attention_backend is not None and module_type != "transformer": raise ValueError("attention_backend can only be set when loading " f"a transformer, got module_type={module_type!r}") resolved_attention_backend = coerce_attn_backend(attention_backend) - # Per-role request delivered as a construction scope: process-local, - # exception-safe, and part of the selector's cache key (no global - # mutation, no cache flushes between roles). attention_context = (nullcontext() if resolved_attention_backend is None else _component_attention_backend_scope( resolved_attention_backend, component=module_type)) if disable_custom_init_weights: fastvideo_args._loading_teacher_critic_model = True - # Attention implementations are bound while transformer layers are - # constructed. Scope the override to this one role so student, - # teacher, and critic can use independent backends in one process. - with attention_context: - module = PipelineComponentLoader.load_module( - module_name=module_type, - component_model_path=component_path, - transformers_or_diffusers=(transformers_or_diffusers), - fastvideo_args=fastvideo_args, - ) + try: + with attention_context: + module = PipelineComponentLoader.load_module( + module_name=module_type, + component_model_path=component_path, + transformers_or_diffusers=(transformers_or_diffusers), + fastvideo_args=fastvideo_args, + ) + finally: + fastvideo_args.pipeline_config.dit_precision = original_dit_precision if not isinstance(module, torch.nn.Module): raise TypeError(f"Loaded {module_type!r} is not a " f"torch.nn.Module: {type(module)}") + if resolved_attention_backend is not None: + receipt = component_attention_backend(module) + if receipt is NO_REQUEST: + raise RuntimeError(f"Loaded {module_type!r} from {model_path!r} did not record its " + f"requested attention backend {resolved_attention_backend.name}. " + "The component loader must stamp the construction decision on " + "module.config._resolved_attention_backend.") + if receipt is not resolved_attention_backend: + raise RuntimeError(f"Loaded {module_type!r} from {model_path!r} requested attention " + f"backend {resolved_attention_backend.name}, but recorded " + f"{receipt.name}.") return module diff --git a/fastvideo/train/utils/optimizer.py b/fastvideo/train/utils/optimizer.py index 43a79d98df..e6485d146b 100644 --- a/fastvideo/train/utils/optimizer.py +++ b/fastvideo/train/utils/optimizer.py @@ -18,6 +18,70 @@ ) +class AdamWBeta1Zero(torch.optim.Optimizer): + """AdamW specialized for ``beta1 == 0``: identical update, no ``exp_avg``. + + With ``beta1 = 0`` Adam's first moment reduces to the raw gradient + (``m_t = g_t``, bias correction 1), so the buffer only doubles optimizer + state for nothing — one full parameter-sized tensor per model. The op + sequence below mirrors ``torch.optim.AdamW``'s single-tensor path + exactly, so the parameter trajectory is bitwise-equivalent to + ``AdamW(betas=(0.0, beta2))``. + """ + + def __init__( + self, + params, + lr: float, + beta2: float, + eps: float = 1e-8, + weight_decay: float = 0.0, + ) -> None: + if not 0.0 <= beta2 < 1.0: + raise ValueError(f"Invalid beta2: {beta2}") + defaults = dict(lr=float(lr), beta2=float(beta2), eps=float(eps), weight_decay=float(weight_decay)) + super().__init__(params, defaults) + + @torch.no_grad() + def step(self, closure=None): + loss = None + if closure is not None: + with torch.enable_grad(): + loss = closure() + for group in self.param_groups: + lr = group["lr"] + beta2 = group["beta2"] + eps = group["eps"] + weight_decay = group["weight_decay"] + for p in group["params"]: + if p.grad is None: + continue + grad = p.grad + state = self.state[p] + if not state: + state["step"] = 0 + if p.dtype == torch.float32: + state["exp_avg_sq"] = torch.zeros_like(p, memory_format=torch.preserve_format) + else: + state["exp_avg_sq"] = torch.zeros_like(p, dtype=torch.float32) + state["step"] += 1 + exp_avg_sq = state["exp_avg_sq"] + p.mul_(1 - lr * weight_decay) + if p.dtype == torch.float32: + exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2) + bias_correction2_sqrt = (1 - beta2**state["step"])**0.5 + denom = (exp_avg_sq.sqrt() / bias_correction2_sqrt).add_(eps) + p.addcdiv_(grad, denom, value=-lr) + else: + grad_f = grad.float() + exp_avg_sq.mul_(beta2).addcmul_(grad_f, grad_f, value=1 - beta2) + bias_correction2_sqrt = (1 - beta2**state["step"])**0.5 + denom = (exp_avg_sq.sqrt() / bias_correction2_sqrt).add_(eps) + update = grad_f.div_(denom).mul_(-lr) + p.add_(update.to(p.dtype)) + return loss + + def build_optimizer_and_scheduler( *, params: list[torch.nn.Parameter], @@ -36,13 +100,22 @@ def build_optimizer_and_scheduler( raise ValueError("No trainable parameters passed to " "build_optimizer_and_scheduler") - optimizer = torch.optim.AdamW( - params, - lr=float(learning_rate), - betas=betas, - weight_decay=float(optimizer_config.weight_decay), - eps=1e-8, - ) + if float(betas[0]) == 0.0: + optimizer: torch.optim.Optimizer = AdamWBeta1Zero( + params, + lr=float(learning_rate), + beta2=float(betas[1]), + weight_decay=float(optimizer_config.weight_decay), + eps=1e-8, + ) + else: + optimizer = torch.optim.AdamW( + params, + lr=float(learning_rate), + betas=betas, + weight_decay=float(optimizer_config.weight_decay), + eps=1e-8, + ) scheduler = get_scheduler( str(scheduler_name), diff --git a/fastvideo/train/utils/tracking.py b/fastvideo/train/utils/tracking.py index 91e06889be..f8a3084ccb 100644 --- a/fastvideo/train/utils/tracking.py +++ b/fastvideo/train/utils/tracking.py @@ -6,6 +6,7 @@ from typing import Any, TYPE_CHECKING from fastvideo.distributed import get_world_group +from fastvideo.logger import init_logger from fastvideo.training.trackers import ( initialize_trackers, Trackers, @@ -17,6 +18,24 @@ TrackerConfig, ) +logger = init_logger(__name__) + + +def _wandb_usable() -> bool: + """True when wandb can actually start a run (importable + credentials).""" + try: + import wandb + except Exception: + return False + if os.environ.get("WANDB_API_KEY"): + return True + if os.environ.get("WANDB_MODE") in ("offline", "disabled"): + return True + try: + return wandb.api.api_key is not None + except Exception: + return False + def build_tracker( tracker_config: TrackerConfig, @@ -33,6 +52,11 @@ def build_tracker( trackers.append(Trackers.WANDB.value) if world_group.rank != 0: trackers = [] + if Trackers.WANDB.value in trackers and not _wandb_usable(): + logger.warning("wandb tracking requested but wandb is not importable " + "or no credentials are configured (WANDB_API_KEY); " + "continuing without the wandb tracker.") + trackers = [t for t in trackers if t != Trackers.WANDB.value] tracker_log_dir = (checkpoint_config.output_dir or os.getcwd()) if trackers: @@ -43,11 +67,36 @@ def build_tracker( tracker_run_name = tracker_config.run_name or None project = (tracker_config.project_name or "fastvideo") - return initialize_trackers( - trackers, - experiment_name=project, - config=tracker_config_dict, - log_dir=tracker_log_dir, - entity=tracker_entity, - run_name=tracker_run_name, - ) + try: + return initialize_trackers( + trackers, + experiment_name=project, + config=tracker_config_dict, + log_dir=tracker_log_dir, + entity=tracker_entity, + run_name=tracker_run_name, + ) + except Exception as exc: + if Trackers.WANDB.value not in trackers: + raise + logger.warning("Tracker init failed (%s); retrying wandb in offline mode.", exc) + os.environ["WANDB_MODE"] = "offline" + try: + return initialize_trackers( + trackers, + experiment_name=project, + config=tracker_config_dict, + log_dir=tracker_log_dir, + entity=tracker_entity, + run_name=tracker_run_name, + ) + except Exception as offline_exc: + logger.warning("Offline tracker init also failed (%s); continuing without trackers.", offline_exc) + return initialize_trackers( + [], + experiment_name=project, + config=None, + log_dir=tracker_log_dir, + entity=None, + run_name=None, + ) diff --git a/fastvideo/train/utils/training_config.py b/fastvideo/train/utils/training_config.py index 25a546cef0..31df20dd23 100644 --- a/fastvideo/train/utils/training_config.py +++ b/fastvideo/train/utils/training_config.py @@ -32,6 +32,7 @@ class DataConfig: num_width: int = 0 num_latent_t: int = 0 num_frames: int = 0 + native_shape_bucketing: bool = False @dataclass(slots=True) @@ -56,8 +57,14 @@ class TrainingLoopConfig: class CheckpointConfig: output_dir: str = "" resume_from_checkpoint: str = "" + save_inference_checkpoint_on_validation: bool = False + inference_checkpoint_role: str = "student" + inference_checkpoint_dtype: str = "bfloat16" training_state_checkpointing_steps: int = 0 + require_complete_training_checkpoint: bool = False checkpoints_total_limit: int = 0 + checkpointing_start_step: int = 0 + reset_lr_on_resume: bool = False @dataclass(slots=True) @@ -77,6 +84,9 @@ class ModelTrainingConfig: precondition_outputs: bool = False moba_config: dict = field(default_factory=dict) enable_gradient_checkpointing_type: str | None = None + allow_low_precision_master_weights: bool = False + enable_torch_compile: bool = False + torch_compile_kwargs: dict = field(default_factory=dict) @dataclass(slots=True) @@ -88,11 +98,7 @@ class TrainingConfig: checkpoint: CheckpointConfig = field(default_factory=CheckpointConfig) tracker: TrackerConfig = field(default_factory=TrackerConfig) vsa_sparsity: float = 0.0 - # Reuse the per-step padded VSA tile buffer across attention layers. - # Defaults to False for training: under full activation checkpointing the - # cached buffer survives into the backward recompute and inflates peak - # memory (see #1423). Enable on memory-rich setups to keep the per-step - # buffer-reuse speedup. + vsa_tile_size: int = 256 vsa_cache_tile_buf: bool = False model: ModelTrainingConfig = field(default_factory=ModelTrainingConfig) pipeline_config: PipelineConfig | None = None diff --git a/mkdocs.yml b/mkdocs.yml index 3a67cb01c4..0a6a63ce03 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -217,6 +217,11 @@ nav: - Video Sparse Attention: attention/vsa/index.md - Sliding Tile Attention (Archived): attention/sta/index.md - Backend Development: contributing/attention_backend.md + - Quantization: + - Quantized Checkpoint Loading: quantization/loader_quant_params.md + - Affine INT8 for MiniMax-H3: quantization/h3_int8_affine.md + - NVFP4 for MiniMax-H3: quantization/h3_nvfp4.md + - W4A16 (4-bit weight) for MiniMax-H3: quantization/h3_w4a16.md - Utilities: - LoRA: utilities/lora.md - Debugging: utilities/debugging.md diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py b/scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py index 4a5cdfb4cb..5e7a61f582 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py @@ -28,7 +28,6 @@ --src /path/to/MiniMax-H3/transformer \ --dst /path/to/MiniMax-H3-r16/transformer --rank 16 - # then assemble a model dir whose other components point at the original ln -s /path/to/MiniMax-H3/{text_encoder,tokenizer,processor,vae,audio_vae,\ scheduler,audio_scheduler,modular_model_index.json} /path/to/MiniMax-H3-r16/ @@ -70,6 +69,7 @@ def parse_args() -> argparse.Namespace: p.add_argument("--grid", type=int, default=4096, help="Timestep samples used to fit the basis.") p.add_argument("--freq-dim", type=int, default=256) p.add_argument("--report-only", action="store_true", help="Fit and report error without writing.") + p.add_argument("--receipt", type=Path, help="Optional JSON receipt for automated parity gates.") return p.parse_args() @@ -94,11 +94,9 @@ def load_keys(src: Path, index: dict[str, str], keys: list[str]) -> dict[str, to def fit_basis(src: Path, index: dict[str, str], rank: int, grid: int, - freq_dim: int) -> tuple[torch.Tensor, torch.Tensor]: + freq_dim: int) -> tuple[torch.Tensor, torch.Tensor, float]: """Return (V [time_embed_dim, rank], U [grid, time_embed_dim]) in float64.""" embedder = load_keys(src, index, list(TIME_EMBEDDER_KEYS)) - # t is the DiT's timestep input: scheduler.timesteps = 1 - sigmas, so t in [0, 1]. - # Condition rows pin t to 0.999 / 1.0, so the endpoint must be included. t = torch.linspace(0.0, 1.0, grid, dtype=torch.float64) h = timestep_embedding(t, freq_dim) h = h @ embedder["time_embedder.linear_1.weight"].double().T + embedder["time_embedder.linear_1.bias"].double() @@ -108,17 +106,24 @@ def fit_basis(src: Path, index: dict[str, str], rank: int, grid: int, _, s, vh = torch.linalg.svd(u, full_matrices=False) residual = ((s[rank:]**2).sum() / (s**2).sum()).sqrt() print(f"basis: U={tuple(u.shape)} rank={rank} relative residual ||U-U_r||/||U|| = {residual:.3e}") - return vh[:rank].T.contiguous(), u + return vh[:rank].T.contiguous(), u, float(residual) def main() -> None: args = parse_args() src, dst = Path(args.src), Path(args.dst) - index_map = json.loads((src / INDEX_NAME).read_text())["weight_map"] + index_path = src / INDEX_NAME + if index_path.exists(): + index_map = json.loads(index_path.read_text())["weight_map"] + else: + single = src / "model.safetensors" + if not single.exists(): + raise SystemExit(f"{src} has neither {INDEX_NAME} nor model.safetensors") + with safe_open(str(single), framework="pt") as handle: + index_map = {key: "model.safetensors" for key in handle.keys()} + + basis, u, residual = fit_basis(src, index_map, args.rank, args.grid, args.freq_dim) - basis, u = fit_basis(src, index_map, args.rank, args.grid, args.freq_dim) - - # Worst-case induced error on the actual modulation outputs. worst = 0.0 scale = 0.0 for key in sorted(k for k in index_map if k.endswith(ADALN_SUFFIX) or k == NORM_OUT_WEIGHT): @@ -129,7 +134,21 @@ def main() -> None: scale = max(scale, ref.abs().max().item()) print(f"modulation error over all projections: max|err|={worst:.3e} " f"(|Wu|max={scale:.3f}, relative={worst / scale:.3e})") + receipt = { + "schema_version": 1, + "source": str(src.resolve()), + "destination": None if args.report_only else str(dst.resolve()), + "rank": args.rank, + "grid_points": args.grid, + "basis_relative_residual": residual, + "modulation_max_abs_error": worst, + "modulation_reference_absmax": scale, + "modulation_relative_max_error": worst / scale, + } if args.report_only: + if args.receipt: + args.receipt.parent.mkdir(parents=True, exist_ok=True) + args.receipt.write_text(json.dumps(receipt, indent=2) + "\n") return dst.mkdir(parents=True, exist_ok=True) @@ -143,7 +162,6 @@ def main() -> None: tensor = f.get_tensor(key) total_before += tensor.numel() if key.endswith(ADALN_SUFFIX) or key == NORM_OUT_WEIGHT: - # [out, time_embed_dim] @ [time_embed_dim, rank] -> [out, rank] tensor = (tensor.double() @ basis).to(torch.float16) tensors[key] = tensor total_after += tensor.numel() @@ -151,7 +169,6 @@ def main() -> None: save_file(tensors, str(dst / shard), metadata={"format": "pt"}) print(f"wrote {shard} ({len(tensors)} tensors)") - # ReplicatedLinear stores [out_features, in_features], so the basis is V.T. basis_shard = "diffusion_pytorch_model-adaln-basis.safetensors" save_file({"adaln_basis.weight": basis.T.to(torch.float16).contiguous()}, str(dst / basis_shard), @@ -171,6 +188,14 @@ def main() -> None: print(f"\nparameters: {total_before / 1e9:.3f}B -> {total_after / 1e9:.3f}B " f"({100 * (1 - total_after / total_before):.1f}% removed)") print(f"bf16 footprint: {total_before * 2 / 1e9:.1f} GB -> ~{total_after * 2 / 1e9:.1f} GB") + receipt.update({ + "parameters_before": total_before, + "parameters_after": total_after, + "fraction_removed": 1 - total_after / total_before, + }) + if args.receipt: + args.receipt.parent.mkdir(parents=True, exist_ok=True) + args.receipt.write_text(json.dumps(receipt, indent=2) + "\n") if __name__ == "__main__": diff --git a/scripts/checkpoint_conversion/export_h3_dmd2_student.py b/scripts/checkpoint_conversion/export_h3_dmd2_student.py new file mode 100644 index 0000000000..45719f8f0d --- /dev/null +++ b/scripts/checkpoint_conversion/export_h3_dmd2_student.py @@ -0,0 +1,197 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Export a MiniMax-H3 DMD2 training checkpoint's student into an inference model dir. + +The modular trainer's DCP checkpoints hold every role and optimizer +(``roles.student.transformer.*`` is the piece inference needs, stored as the +fp32 master weights). This script streams just those tensors out of the DCP +shards in bounded memory (no process group, no GPU), renames them from +fastvideo layer names back to the on-disk checkpoint convention (the inverse +of ``MiniMaxH3ArchConfig.param_names_mapping``), casts to bf16, and writes a +sharded safetensors ``transformer/`` next to symlinks into the base model dir +for every other component — so the output loads through the standard +inference pipeline with ~66 GB of new bytes instead of a full copy. + +One param per block has no on-disk counterpart: ``attn.to_gate_compress``, +the VSA gate. The base checkpoint default-initializes it; a VSA-trained +student's gate is learned, so it is exported under its fastvideo name (no +mapping rule touches it, and the loader resolves it to the model param +verbatim). Dense-only consumers ignore it. + +Usage:: + + python scripts/checkpoint_conversion/export_h3_dmd2_student.py \ + --checkpoint /path/to/outputs//checkpoint-1400 \ + --output-dir /path/to/exports/-step1400 \ + [--base-model /mnt/lustre/vlm-k1kong/models/MiniMax-H3] \ + [--role student] [--dtype bfloat16] [--copy-components] + +``--checkpoint latest --run-dir `` picks the newest +``checkpoint-*`` whose ``dcp/.metadata`` exists (the strict-resume contract). +""" + +from __future__ import annotations + +import argparse +import json +import re +import shutil +from pathlib import Path + +import torch +from safetensors.torch import save_file + +INVERSE_PARAM_RULES: tuple[tuple[str, str], ...] = ( + (r"^time_embedder\.fc_in\.(.*)$", r"time_embedder.linear_1.\1"), + (r"^time_embedder\.fc_out\.(.*)$", r"time_embedder.linear_2.\1"), + (r"^(.*)\.attn\.to_out\.(weight|bias)$", r"\1.attn.to_out.0.\2"), + (r"^(.*)\.ff\.fc_in\.(.*)$", r"\1.ff.net.0.proj.\2"), + (r"^(.*)\.ff\.fc_out\.(.*)$", r"\1.ff.net.2.\2"), +) + +EXPECTED_NEW_PARAM_PATTERNS = (re.compile(r"\.attn\.to_gate_compress\."), ) + +SHARD_BUDGET_BYTES = 5 * 1024**3 # ~5 GB per safetensors shard (bf16) + + +def to_disk_name(name: str) -> str: + for pattern, repl in INVERSE_PARAM_RULES: + new, n = re.subn(pattern, repl, name) + if n: + return new + return name + + +def find_latest_checkpoint(run_dir: Path) -> Path: + candidates = sorted( + (p for p in run_dir.glob("checkpoint-*") if (p / "dcp" / ".metadata").exists()), + key=lambda p: int(p.name.rsplit("-", 1)[-1]), + ) + if not candidates: + raise FileNotFoundError(f"no complete checkpoint-*/dcp/.metadata under {run_dir}") + return candidates[-1] + + +def main(args: argparse.Namespace) -> None: + if args.checkpoint == "latest": + if args.run_dir is None: + raise SystemExit("--checkpoint latest requires --run-dir") + checkpoint = find_latest_checkpoint(args.run_dir) + else: + checkpoint = Path(args.checkpoint) + dcp_dir = checkpoint / "dcp" + if not (dcp_dir / ".metadata").exists(): + raise FileNotFoundError(f"{dcp_dir}/.metadata missing — incomplete checkpoint, refusing") + base_model = args.base_model + if not any((base_model / name).exists() for name in ("model_index.json", "modular_model_index.json")): + raise FileNotFoundError(f"{base_model} does not look like a model dir " + "(no model_index.json or modular_model_index.json)") + out_dtype = getattr(torch, args.dtype) + prefix = f"roles.{args.role}.transformer." + + from torch.distributed.checkpoint import FileSystemReader + import torch.distributed.checkpoint as dcp + + reader = FileSystemReader(str(dcp_dir)) + metadata = reader.read_metadata() + param_meta = { + key: meta + for key, meta in metadata.state_dict_metadata.items() + if key.startswith(prefix) + } + if not param_meta: + raise SystemExit(f"no keys under {prefix!r} in {dcp_dir}") + print(f"{checkpoint.name}: {len(param_meta)} tensors under {prefix!r}") + + base_transformer = base_model / "transformer" + base_index = base_transformer / "diffusion_pytorch_model.safetensors.index.json" + base_keys: set[str] = set() + if base_index.exists(): + base_keys = set(json.loads(base_index.read_text())["weight_map"]) + + def nbytes(meta) -> int: + numel = 1 + for dim in meta.size: + numel *= dim + return numel * torch.finfo(out_dtype).bits // 8 + + ordered = sorted(param_meta) + shards: list[list[str]] = [[]] + acc = 0 + for key in ordered: + size = nbytes(param_meta[key]) + if shards[-1] and acc + size > SHARD_BUDGET_BYTES: + shards.append([]) + acc = 0 + shards[-1].append(key) + acc += size + + out_transformer = args.output_dir / "transformer" + out_transformer.mkdir(parents=True, exist_ok=True) + + weight_map: dict[str, str] = {} + total_size = 0 + unexpected_new: list[str] = [] + n_shards = len(shards) + for shard_idx, keys in enumerate(shards, start=1): + fname = f"diffusion_pytorch_model-{shard_idx:05d}-of-{n_shards:05d}.safetensors" + state = { + key: torch.empty(tuple(param_meta[key].size), dtype=param_meta[key].properties.dtype) + for key in keys + } + dcp.load(state, checkpoint_id=str(dcp_dir)) + tensors: dict[str, torch.Tensor] = {} + for key, tensor in state.items(): + disk_name = to_disk_name(key[len(prefix):]) + if base_keys and disk_name not in base_keys: + if not any(p.search(disk_name) for p in EXPECTED_NEW_PARAM_PATTERNS): + unexpected_new.append(disk_name) + tensors[disk_name] = tensor.to(out_dtype).contiguous() + weight_map[disk_name] = fname + total_size += tensors[disk_name].numel() * tensors[disk_name].element_size() + save_file(tensors, str(out_transformer / fname)) + del state, tensors + print(f" wrote {fname} ({len(keys)} tensors)") + + if unexpected_new: + raise SystemExit("Export produced keys unknown to the base checkpoint (mapping drift?):\n " + + "\n ".join(sorted(unexpected_new)[:20])) + + (out_transformer / "diffusion_pytorch_model.safetensors.index.json").write_text( + json.dumps({"metadata": {"total_size": total_size}, "weight_map": weight_map}, indent=2)) + shutil.copy2(base_transformer / "config.json", out_transformer / "config.json") + + for entry in sorted(base_model.iterdir()): + if entry.name == "transformer": + continue + target = args.output_dir / entry.name + if target.exists() or target.is_symlink(): + continue + if args.copy_components: + if entry.is_dir(): + shutil.copytree(entry, target) + else: + shutil.copy2(entry, target) + else: + target.symlink_to(entry.resolve()) + + print(f"Export complete: {args.output_dir}") + print(f" transformer: {len(weight_map)} tensors, {total_size / 1024**3:.1f} GiB " + f"({args.dtype}), {n_shards} shards") + print(f" other components {'copied' if args.copy_components else 'symlinked'} from {base_model}") + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoint", required=True, help="checkpoint-N dir, or 'latest' with --run-dir") + parser.add_argument("--run-dir", type=Path, default=None, help="training output dir for --checkpoint latest") + parser.add_argument("--output-dir", type=Path, required=True) + parser.add_argument("--base-model", + type=Path, + default=Path("/mnt/lustre/vlm-k1kong/models/MiniMax-H3"), + help="base model dir supplying config + non-transformer components") + parser.add_argument("--role", default="student", help="training role to export (student|critic)") + parser.add_argument("--dtype", default="bfloat16", choices=["bfloat16", "float32", "float16"]) + parser.add_argument("--copy-components", + action="store_true", + help="copy non-transformer components instead of symlinking") + main(parser.parse_args()) diff --git a/scripts/compacth3/analysis/adaln/analyze_adaln_rank.py b/scripts/compacth3/analysis/adaln/analyze_adaln_rank.py new file mode 100644 index 0000000000..3cdcf68dc9 --- /dev/null +++ b/scripts/compacth3/analysis/adaln/analyze_adaln_rank.py @@ -0,0 +1,364 @@ +#!/usr/bin/env python3 +"""Spectral analysis of MiniMax-H3 AdaLN timestep conditioning.""" +from __future__ import annotations + +import argparse +import contextlib +import importlib.util +import json +import os +import sys +import time +from pathlib import Path + +import torch + +SPRINT = "/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829" +M = "/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1" +HARNESS = f"{M}/examples/inference/basic/basic_fasth3.py" +OUT_DIR = Path(SPRINT) / "adaln_rank_analysis" + +N_GRID = 4096 +ENERGY_THRESHOLDS = (0.90, 0.95, 0.99, 0.999) +COVERAGE_K = (1, 2, 3, 4, 6, 8, 12, 16, 24, 32, 48, 64, 96, 128, 192, 256, 384, 512, 768) +MAX_SIGMA_STORED = 1024 +DEVICE = "cuda:0" +FP32_EPS = float(torch.finfo(torch.float32).eps) +N_MODALITY = 3 # MINIMAX_H3_MODALITY_NUM + + +def log(msg: str) -> None: + print(f"[adaln-rank] {time.strftime('%H:%M:%S')} {msg}", flush=True) + + +def build_fastvideo_args(model_path: str): + spec = importlib.util.spec_from_file_location("fasth3_harness", HARNESS) + harness = importlib.util.module_from_spec(spec) + sys.modules["fasth3_harness"] = harness + spec.loader.exec_module(harness) + + argv = [ + "--model-path", model_path, + "--prompt", "adaln-rank-analysis", + "--num-gpus", "1", + "--no-fa4", + "--no-inference-torch-compile", + "--steps", "5", + ] + args = harness.parse_args(argv) + args.fa4 = False + harness.configure_environment(args) + os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "TORCH_SDPA" + + config = harness.build_generator_config(args) + config.pipeline.experimental["attention_backend"] = "TORCH_SDPA" + from fastvideo.api.compat import generator_config_to_fastvideo_args + return generator_config_to_fastvideo_args(config) + + +def load_dit(fastvideo_args, transformer_path: str): + from fastvideo.models.loader.component_loader import PipelineComponentLoader + log(f"loading transformer from {transformer_path}") + model = PipelineComponentLoader.load_module( + module_name="transformer", + component_model_path=transformer_path, + transformers_or_diffusers="diffusers", + fastvideo_args=fastvideo_args, + ) + log(f"loaded class={type(model).__name__}") + return model + + +@contextlib.contextmanager +def on_device_fp32(module): + """Temporarily put exactly this small module on GPU in fp32, then restore.""" + saved = [(name, p.dtype, p.device) for name, p in module.named_parameters()] + module.to(device=DEVICE, dtype=torch.float32) + try: + yield module + finally: + params = dict(module.named_parameters()) + for name, dtype, device in saved: + params[name].data = params[name].data.to(device=device, dtype=dtype) + torch.cuda.empty_cache() + + +def scheduler_timesteps(shift: float, grid_points: int) -> torch.Tensor: + """Verbatim replica of MiniMaxH3Scheduler.set_timesteps (scheduler file read).""" + base = torch.linspace(1.0, 0.0, int(grid_points), dtype=torch.float32) + sigma = shift * base / (1 + (shift - 1) * base) + sigma = torch.unique_consecutive(sigma) + return 1.0 - sigma[:-1] + + +def build_grids(video_shift: float, audio_shift: float): + uniform = torch.linspace(0.0, 1.0, N_GRID, dtype=torch.float32) + pieces = [torch.tensor([0.0, 0.999, 1.0], dtype=torch.float32)] + for shift in (video_shift, audio_shift): + for n in (4, 5, 49, 50): + pieces.append(scheduler_timesteps(shift, n)) + ops = torch.unique(torch.cat(pieces)).sort().values + control = torch.linspace(0.0, 1000.0, N_GRID, dtype=torch.float32) + return uniform, ops, control + + +def spectral_stats(sigma: torch.Tensor, shape: tuple[int, int]) -> dict: + s = sigma.detach().double().cpu() + energy = s * s + total = float(energy.sum()) + out: dict = { + "shape": [int(shape[0]), int(shape[1])], + "frobenius_norm": float(total ** 0.5), + "s_max": float(s[0]), + "s_min": float(s[-1]), + "numerical_rank_fp32tol": int((s > s[0] * max(shape) * FP32_EPS).sum()), + "stable_rank_trace_over_smax2": float(total / (float(s[0]) ** 2)), + } + for rel in (1e-2, 1e-3, 1e-4, 1e-5, 1e-6): + out[f"rank_sigma_above_{rel:g}_of_smax"] = int((s > s[0] * rel).sum()) + p = energy / total + nz = p > 0 + out["entropy_effective_rank"] = float(torch.exp(-(p[nz] * p[nz].log()).sum())) + cum = torch.cumsum(energy, 0) / total + for th in ENERGY_THRESHOLDS: + k = int(torch.searchsorted(cum, th).item()) + 1 + out[f"rank_at_{th:.3f}_energy"] = k + out[f"relerr_at_{th:.3f}_energy"] = float(max(0.0, 1.0 - float(cum[k - 1])) ** 0.5) + return out + + +def gram_eigvals(G: torch.Tensor) -> torch.Tensor: + """Descending singular values of the trajectory from its (T x T) Gram matrix.""" + lam = torch.linalg.eigvalsh(G.double()) + return torch.clamp(lam.flip(0), min=0.0).sqrt() + + +def analyze_blocks(model, mods, Z, H, store_sigma: bool, do_coverage: bool) -> tuple[dict, dict]: + T = int(Z.shape[0]) + D = 6 * H * N_MODALITY + per_block: dict = {} + G_union = torch.zeros(T, T, dtype=torch.float32, device=DEVICE) + G_union_norm = torch.zeros_like(G_union) + G_mod_union = [torch.zeros_like(G_union) for _ in range(N_MODALITY)] + G_list: list[torch.Tensor] = [] + + for name, block in mods: + t0 = time.time() + proj = block.adaln_proj + with on_device_fp32(proj) as mod: + six = torch.stack([s.float() for s in mod(Z)], dim=0) # (6, 3T, H) + M = six.permute(1, 0, 2).reshape(T, N_MODALITY, 6, H).reshape(T, -1) + del six + M = M - M.mean(dim=0, keepdim=True) + G = M @ M.t() + G = 0.5 * (G + G.t()) + sv = gram_eigvals(G) + stats = spectral_stats(sv, (T, D)) + stats["apply_silu"] = bool(getattr(proj, "apply_silu", None)) + stats["linear_weight_shape"] = list(proj.linear.weight.shape) + stats["modulation_dim"] = D + entry = {"stats": stats, + "frobenius_per_modality": [float(M[:, m * 6 * H:(m + 1) * 6 * H].double().norm()) + for m in range(N_MODALITY)]} + if store_sigma: + entry["sigma_top"] = [float(v) for v in sv[:MAX_SIGMA_STORED].cpu()] + entry["sigma_stored"] = int(min(MAX_SIGMA_STORED, sv.numel())) + entry["sigma_nonzero"] = int(sv.numel()) + per_block[name] = entry + + tr = float(torch.diagonal(G).sum()) + G_union += G + if tr > 0: + G_union_norm += G / tr + for m in range(N_MODALITY): + sl = M[:, m * 6 * H:(m + 1) * 6 * H] + G_mod_union[m] += sl @ sl.t() + if do_coverage: + G_list.append(G) + del M + log(f" {name}: rank90={stats['rank_at_0.900_energy']} rank99={stats['rank_at_0.990_energy']} " + f"rank999={stats['rank_at_0.999_energy']} stable={stats['stable_rank_trace_over_smax2']:.2f} " + f"s_max={stats['s_max']:.4g} ({time.time()-t0:.1f}s)") + + def quartiles(key): + vals = sorted(pb["stats"][key] for pb in per_block.values()) + return [vals[0], vals[len(vals) // 2], vals[-1]] + + level_b = { + "modulation_dim": D, + "blocks": per_block, + "summary": { + "n_blocks": len(per_block), + "rank90_min_median_max": quartiles("rank_at_0.900_energy"), + "rank99_min_median_max": quartiles("rank_at_0.990_energy"), + "rank999_min_median_max": quartiles("rank_at_0.999_energy"), + "stable_rank_min_median_max": quartiles("stable_rank_trace_over_smax2"), + }, + } + + cols = D * len(per_block) + + def union_stats(Gmat: torch.Tensor, label: str) -> dict: + sv = gram_eigvals(Gmat) + return {"label": label, "stats": spectral_stats(sv, (T, cols)), + "sigma_top": [float(v) for v in sv[:MAX_SIGMA_STORED].cpu()]} + + level_c = { + "energy_weighted": union_stats(G_union, "sum of centered block Gram matrices"), + "block_normalized": union_stats(G_union_norm, + "sum of trace-normalized centered block Gram matrices"), + "per_modality_uniform_weighted": [ + {"modality": m, + "stats": spectral_stats(gram_eigvals(G_mod_union[m]), (T, 6 * H * len(per_block)))} + for m in range(N_MODALITY) + ], + } + if do_coverage: + Gmat = G_union + lam, V = torch.linalg.eigh(Gmat.double()) + Vf = V.flip(1)[:, :max(COVERAGE_K)].float() + rows = [] + for G in G_list: + GV = G @ Vf + d = (Vf * GV).sum(dim=0).double() + rows.append((torch.cumsum(d, 0) / float(torch.diagonal(G).sum())).cpu()) + C = torch.stack(rows) + level_c["coverage_k"] = list(COVERAGE_K) + level_c["coverage_min_median_max_fraction_of_block_energy"] = { + str(k): [float(C[:, k - 1].min()), float(C[:, k - 1].median()), float(C[:, k - 1].max())] + for k in COVERAGE_K} + level_c["coverage_per_block"] = [[float(x) for x in row] for row in C] + return level_b, level_c + + +def measure(model, tag: str, ckpt: str, n_blocks_limit: int | None): + from fastvideo.models.dits.minimax_h3 import MiniMaxH3AdaLayerNormModulation + + log(f"adaln_rank={model.adaln_rank} hidden={model.hidden_size} " + f"blocks={len(model.transformer_blocks)}") + + ckpt_dir = Path(ckpt) + video_shift = float(json.loads((ckpt_dir / "scheduler" / "scheduler_config.json").read_text())["shift"]) + audio_shift = float(json.loads((ckpt_dir / "audio_scheduler" / "scheduler_config.json").read_text())["shift"]) + uniform, ops, control = build_grids(video_shift, audio_shift) + grids = {"uniform_grid": uniform, "operating_points": ops, "control_0_1000": control} + log(f"grids: uniform={tuple(uniform.shape)} ops={tuple(ops.shape)} " + f"ops_range=({float(ops[0]):.4f},{float(ops[-1]):.4f}) " + f"control={tuple(control.shape)} video_shift={video_shift} audio_shift={audio_shift}") + + H = int(model.hidden_size) + + result: dict = { + "model_tag": tag, + "checkpoint": str(ckpt), + "adaln_rank": int(model.adaln_rank), + "hidden_size": H, + "num_layers": len(model.transformer_blocks), + "num_refiner_layers": len(model.token_refiner.refiner_blocks), + "scheduler_shifts": {"video": video_shift, "audio": audio_shift}, + "grids": { + "uniform_grid": {"n": int(uniform.numel()), "range": [0.0, 1.0], "role": "primary"}, + "operating_points": {"n": int(ops.numel()), "values": [round(float(v), 6) for v in ops], + "role": "the literal timesteps inference feeds"}, + "control_0_1000": {"n": int(control.numel()), "range": [0.0, 1000.0], + "role": "CONTROL ONLY -- not a convention this codebase uses"}, + }, + "precision_note": ("level A/B/C spectra are computed with the three AdaLN " + "modules cast to float32 in memory (checkpoint weights are " + "bf16); level_a_shared_coordinate.uniform_grid_bf16_native " + "repeats level A on the unmodified bf16 modules to expose " + "the storage-format noise floor."), + "device": DEVICE, + } + + def shared_coordinate(grid: torch.Tensor, native_bf16: bool = False) -> torch.Tensor: + t = grid.to(DEVICE, dtype=torch.float32) + if native_bf16: + temb = model.time_proj(t) + temb = model.time_embedder(temb.to(model.time_embedder.fc_in.weight.dtype)) + z, _ = model.adaln_basis(torch.nn.functional.silu(temb).to(model.adaln_basis.weight.dtype)) + return z.detach().float() + with on_device_fp32(model.time_embedder) as te: + temb = te(model.time_proj(t).to(te.fc_in.weight.dtype)) + with on_device_fp32(model.adaln_basis) as basis: + z, _ = basis(torch.nn.functional.silu(temb).to(basis.weight.dtype)) + return z.detach().float() + + def level_a(Zmat: torch.Tensor) -> dict: + Zc = Zmat - Zmat.mean(dim=0, keepdim=True) + sv = torch.linalg.svdvals(Zc.double()).float() + stats = spectral_stats(sv, tuple(Zc.shape)) + stats["centered_frobenius"] = float(Zc.double().norm()) + stats["raw_frobenius"] = float(Zmat.double().norm()) + stats["mean_row_norm"] = float(Zmat.mean(dim=0).double().norm()) + return {"sigma": [float(v) for v in sv.cpu()], "stats": stats} + + t0 = time.time() + Zcache = {name: shared_coordinate(g) for name, g in grids.items()} + log(f"shared coordinates done in {time.time()-t0:.1f}s " + f"{ {k: tuple(v.shape) for k, v in Zcache.items()} }") + level_a_out = {name: level_a(Z) for name, Z in Zcache.items()} + level_a_out["uniform_grid_bf16_native"] = level_a(shared_coordinate(uniform, native_bf16=True)) + result["level_a_shared_coordinate"] = level_a_out + for name in level_a_out: + log(f"LEVEL A {name}: {json.dumps(level_a_out[name]['stats'])}") + + mods: list[tuple[str, torch.nn.Module]] = [ + (f"transformer_blocks.{i}", b) for i, b in enumerate(model.transformer_blocks)] + refiner_mods = [n for n, m in model.token_refiner.named_modules() + if isinstance(m, MiniMaxH3AdaLayerNormModulation)] + result["refiner_adaln_modules"] = refiner_mods + log(f"refiner AdaLN modules found: {refiner_mods or 'NONE'}") + if n_blocks_limit is not None: + mods = mods[:n_blocks_limit] + log(f"per-block modulation dim = {6 * H * N_MODALITY} (6 x {H} x {N_MODALITY} modalities), " + f"{len(mods)} blocks") + + result["level_b_per_block"] = {} + result["level_c_shared_basis"] = {} + for name, Z in Zcache.items(): + primary = name == "uniform_grid" + log(f" --- block analysis on {name} (T={Z.shape[0]}) ---") + level_b, level_c = analyze_blocks(model, mods, Z, H, store_sigma=True, do_coverage=primary) + result["level_b_per_block"][name] = level_b + result["level_c_shared_basis"][name] = level_c + log(f"LEVEL B {name} summary: {json.dumps(level_b['summary'])}") + log(f"LEVEL C {name} weighted: {json.dumps(level_c['energy_weighted']['stats'])}") + log(f"LEVEL C {name} normalized: {json.dumps(level_c['block_normalized']['stats'])}") + if primary: + log(f"LEVEL C {name} coverage: " + f"{json.dumps(level_c['coverage_min_median_max_fraction_of_block_energy'])}") + return result + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--tag", required=True) + ap.add_argument("--model-path", required=True) + ap.add_argument("--out-dir", default=str(OUT_DIR)) + ap.add_argument("--max-blocks", type=int, default=None, help="debug: only the first N blocks") + args = ap.parse_args() + + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + torch.set_grad_enabled(False) + + out_dir = Path(args.out_dir) + out_dir.mkdir(parents=True, exist_ok=True) + + fva = build_fastvideo_args(args.model_path) + model = load_dit(fva, str(Path(args.model_path) / "transformer")) + model.eval() + + result = measure(model, args.tag, args.model_path, args.max_blocks) + out_path = out_dir / f"adaln_rank_{args.tag}.json" + out_path.write_text(json.dumps(result, indent=1)) + log(f"wrote {out_path} ({out_path.stat().st_size} bytes)") + + del model + torch.cuda.empty_cache() + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scripts/compacth3/analysis/adaln/compare_parent_dmd2.py b/scripts/compacth3/analysis/adaln/compare_parent_dmd2.py new file mode 100644 index 0000000000..9bcdadcad7 --- /dev/null +++ b/scripts/compacth3/analysis/adaln/compare_parent_dmd2.py @@ -0,0 +1,67 @@ +#!/usr/bin/env python3 +"""Sharper parent-vs-DMD2 comparison from the saved singular spectra.""" +import json +from pathlib import Path + +OUT = Path("/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/adaln_rank_analysis") +D = {t: json.loads((OUT / f"adaln_rank_{t}.json").read_text()) for t in ("parent", "dmd2")} + + +def topk_energy(sigma, k): + s = [float(x) for x in sigma[:k]] + tot = sum(float(x) ** 2 for x in sigma) + return sum(x * x for x in s) / tot + + +print("LEVEL A (shared coordinate z(t), 4096 x 768, uniform grid, centered)") +print(" model s1 s2/s1 s3/s1 s4/s1 s5/s1 s6/s1 E@1 E@2 E@3 E@4 E@6") +for t in ("parent", "dmd2"): + sv = D[t]["level_a_shared_coordinate"]["uniform_grid"]["sigma"] + r = [sv[i] / sv[0] if i < len(sv) else 0.0 for i in range(6)] + print(" {:<6s} {:<9.4f} {:<9.5f} {:<9.5f} {:<9.5f} {:<9.5f} {:<9.5f} ".format( + t, sv[0], r[1], r[2], r[3], r[4], r[5]) + + " ".join("{:.4f}".format(topk_energy(sv, k)) for k in (1, 2, 3, 4, 6))) + +print() +print("LEVEL C (union over all 42 blocks, 4096 x 4064256, centered, energy-weighted)") +print(" model E@1 E@2 E@3 E@4 E@6 E@8 stable") +for t in ("parent", "dmd2"): + c = D[t]["level_c_shared_basis"]["uniform_grid"]["energy_weighted"]["sigma_top"] + st = D[t]["level_c_shared_basis"]["uniform_grid"]["energy_weighted"]["stats"]["stable_rank_trace_over_smax2"] + print(" {:<6s} ".format(t) + " ".join("{:.4f}".format(topk_energy(c, k)) for k in (1, 2, 3, 4, 6, 8)) + + " {:.4f}".format(st)) + +print() +print("PER-BLOCK (42 blocks, uniform grid)") +hdr = " model rank99 mean/median rank99.9 mean/median stable mean E@2 per block min/med/max" +print(hdr) +for t in ("parent", "dmd2"): + bl = D[t]["level_b_per_block"]["uniform_grid"]["blocks"] + r99 = sorted(v["stats"]["rank_at_0.990_energy"] for v in bl.values()) + r999 = sorted(v["stats"]["rank_at_0.999_energy"] for v in bl.values()) + st = sorted(v["stats"]["stable_rank_trace_over_smax2"] for v in bl.values()) + cov = D[t]["level_c_shared_basis"]["uniform_grid"]["coverage_min_median_max_fraction_of_block_energy"]["2"] + n = len(r99) + print(" {:<7s} {:.2f}/{:.1f} {:.2f}/{:.1f} {:.4f} {:.4f}/{:.4f}/{:.4f}".format( + t, sum(r99) / n, r99[n // 2], sum(r999) / n, r999[n // 2], sum(st) / n, cov[0], cov[1], cov[2])) + +print() +print("PER-MODALITY PROFILE (level B per-block, energy fraction of each modality table; parent vs dmd2)") +for t in ("parent", "dmd2"): + bl = D[t]["level_b_per_block"]["uniform_grid"]["blocks"] + fr = [[v["frobenius_per_modality"][m] for v in bl.values()] for m in range(3)] + print(" {:<7s} modality frobenius mean: ".format(t) + + " ".join("m{}= {:.1f}".format(m, sum(fr[m]) / len(fr[m])) for m in range(3)) + + " (mean ratio m1/m0={:.4f}, m2/m0={:.4f})".format( + (sum(fr[1]) / len(fr[1])) / (sum(fr[0]) / len(fr[0])), + (sum(fr[2]) / len(fr[2])) / (sum(fr[0]) / len(fr[0])))) + +print() +print("OPS GRID (191 literal inference timesteps)") +print(" model A stable C stable B rank99 med") +for t in ("parent", "dmd2"): + a = D[t]["level_a_shared_coordinate"]["operating_points"]["stats"] + c = D[t]["level_c_shared_basis"]["operating_points"]["energy_weighted"]["stats"] + b = D[t]["level_b_per_block"]["operating_points"]["summary"]["rank99_min_median_max"] + print(" {:<7s} {:.4f} {:.4f} {}".format(t, a["stable_rank_trace_over_smax2"], + c["stable_rank_trace_over_smax2"], b)) diff --git a/scripts/compacth3/analysis/adaln/run_adaln_rank.sh b/scripts/compacth3/analysis/adaln/run_adaln_rank.sh new file mode 100644 index 0000000000..1cd72a6e14 --- /dev/null +++ b/scripts/compacth3/analysis/adaln/run_adaln_rank.sh @@ -0,0 +1,51 @@ +#!/bin/bash +# usage: +set -uo pipefail + +SPRINT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +M=/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1 +PY=/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python +SCRIPT="${SPRINT}/adaln_rank_analysis/analyze_adaln_rank.py" +OUT="${SPRINT}/adaln_rank_analysis" + +PARENT_CKPT="${SPRINT}/runs/release20b-folded-long-4k-v3/dmd-parent-step750-complete-v1" +DMD2_CKPT="${SPRINT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3/inference/checkpoint-1400" + +STAGE="${STAGE:-both}" +SMOKE="${SMOKE:-0}" + +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache +export PYTHONDONTWRITEBYTECODE=1 +export PYTHONPATH="${M}:${SPRINT}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" +export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA +export FASTVIDEO_MINIMAX_H3_FUSIONS=0 +export FASTVIDEO_DISABLE_ATTENTION_COMPILE=1 +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True +export TORCH_NCCL_ENABLE_MONITORING=0 +source /mnt/nfs/vlm-aryan/fasth3-33b-20260806/secrets.env >/dev/null 2>&1 || true + +mkdir -p "$OUT" +cd "$M" + +echo "=== host=$(hostname) gpus=$(nvidia-smi -L | wc -l) stage=${STAGE} smoke=${SMOKE} ===" +nvidia-smi --query-gpu=index,name,memory.total --format=csv,noheader + +EXTRA=() +if [[ "$SMOKE" == "1" ]]; then EXTRA+=(--max-blocks 2); fi + +rc=0 +run_one() { + local tag="$1" ckpt="$2" + echo "=== ${tag}: ${ckpt} ===" + "${PY}" "$SCRIPT" --tag "$tag" --model-path "$ckpt" --out-dir "$OUT" "${EXTRA[@]}" + local status=$? + echo "=== ${tag} exit=${status} ===" + if [[ $status -ne 0 ]]; then rc=$status; fi +} + +if [[ "$STAGE" == "both" || "$STAGE" == "parent" ]]; then run_one parent "$PARENT_CKPT"; fi +if [[ "$STAGE" == "both" || "$STAGE" == "dmd2" ]]; then run_one dmd2 "$DMD2_CKPT"; fi + +echo "=== analysis finished rc=${rc} ===" +ls -la "$OUT" +exit $rc diff --git a/scripts/compacth3/analysis/adaln/summarize_adaln_rank.py b/scripts/compacth3/analysis/adaln/summarize_adaln_rank.py new file mode 100644 index 0000000000..da7e4fa5a4 --- /dev/null +++ b/scripts/compacth3/analysis/adaln/summarize_adaln_rank.py @@ -0,0 +1,56 @@ +#!/usr/bin/env python3 +"""Compact table dump over the adaln_rank_{parent,dmd2}.json spectra. + +usage: python3 summarize_adaln_rank.py [tag ...] (default: parent dmd2) +""" +import json +import sys +from pathlib import Path + +OUT = Path("/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/adaln_rank_analysis") + +HEAD = ("{:<34s} {:>18s} {:>6s} {:>5s} {:>5s} {:>5s} {:>6s} {:>8s} {:>8s}".format( + "level", "shape", "numrk", "90%", "95%", "99%", "99.9%", "stable", "entropy")) + + +def row(label, s): + return ("{:<34s} {:>18s} {:>6d} {:>5d} {:>5d} {:>5d} {:>6d} {:>8.3f} {:>8.3f}".format( + label, str(s["shape"]), s["numerical_rank_fp32tol"], s["rank_at_0.900_energy"], + s["rank_at_0.950_energy"], s["rank_at_0.990_energy"], s["rank_at_0.999_energy"], + s["stable_rank_trace_over_smax2"], s["entropy_effective_rank"])) + + +def main(): + tags = sys.argv[1:] or ["parent", "dmd2"] + for tag in tags: + d = json.loads((OUT / f"adaln_rank_{tag}.json").read_text()) + print("=" * 130) + print(f"{tag} {d['checkpoint']}") + print(f" adaln_rank={d['adaln_rank']} hidden={d['hidden_size']} layers={d['num_layers']} " + f"refiner_layers={d['num_refiner_layers']} refiner_adaln_modules={d['refiner_adaln_modules']}") + print(f" shifts={d['scheduler_shifts']} ops_grid_n={d['grids']['operating_points']['n']}") + print(HEAD) + la = d["level_a_shared_coordinate"] + for key in ("uniform_grid", "operating_points", "control_0_1000", "uniform_grid_bf16_native"): + if key in la: + print(row(f"A {key}", la[key]["stats"])) + for gname, lb in d["level_b_per_block"].items(): + s = lb["summary"] + print(" B {}: rank90={} rank99={} rank999={} stable={} (n={} blocks, modulation_dim={})".format( + gname, s["rank90_min_median_max"], s["rank99_min_median_max"], + s["rank999_min_median_max"], [round(x, 3) for x in s["stable_rank_min_median_max"]], + s["n_blocks"], lb["modulation_dim"])) + for gname, c in d["level_c_shared_basis"].items(): + for key in ("energy_weighted", "block_normalized"): + if key in c: + print(row(f"C {gname} {key}", c[key]["stats"])) + for pm in c.get("per_modality_uniform_weighted", []): + print(row(f"C {gname} modality{pm['modality']}", pm["stats"])) + cov = c.get("coverage_min_median_max_fraction_of_block_energy") + if cov: + print(f" C {gname} shared-basis coverage (min/median/max fraction of a block's energy):") + print(" " + " ".join(f"k={k}:{v[0]:.3f}/{v[1]:.3f}/{v[2]:.3f}" for k, v in cov.items())) + + +if __name__ == "__main__": + main() diff --git a/scripts/compacth3/analysis/adaln_lowrank.py b/scripts/compacth3/analysis/adaln_lowrank.py new file mode 100644 index 0000000000..383b6216ec --- /dev/null +++ b/scripts/compacth3/analysis/adaln_lowrank.py @@ -0,0 +1,928 @@ +#!/usr/bin/env python3 +"""AdaLN low-rank compression analysis for MiniMax-H3.""" +from __future__ import annotations + +import argparse +import contextlib +import importlib.util +import json +import os +import sys +import time +from pathlib import Path + +import torch +import torch.nn.functional as F + +SPRINT = "/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829" +M = "/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1" +HARNESS = f"{M}/examples/inference/basic/basic_fasth3.py" +OUT_DIR = Path(SPRINT) / "adaln_rank_analysis" + +N_GRID = 4096 +MODES = ("deployed", "fp32_adaln") +MICRO_PERTURB_REL = (1e-6, 1e-3) +N_ROTATIONS = 5 +RANK_LIST = (768, 64, 32, 16, 8, 4, 2) +DEVICE = "cuda:0" + +LADDER_UNIFORM = (1.0, 0.75, 0.5, 0.25) +LADDER_METADATA = (0.999, 0.749, 0.5, 0.25) +NOMINAL_VIDEO = (0.0, 0.027027, 0.076923, 0.2) + + +def log(msg: str) -> None: + print(f"[adaln-lowrank] {time.strftime('%H:%M:%S')} {msg}", flush=True) + + +def err_stats(approx: torch.Tensor, ref: torch.Tensor, tag: str = "") -> dict: + """Max / RMS / cosine agreement between two tensors. + + max-rel is dominated by a single worst element and is too brittle to rank + models on, so rms_rel and cosine are reported alongside it. rms_rel is + normalised by the REFERENCE's RMS (not its max), which is the usual + normalised-RMS convention. + """ + a = approx.detach().float() + b = ref.detach().float() + d = (a - b).abs() + scale = float(b.abs().max().item()) + rms = float(d.pow(2).mean().sqrt().item()) + ref_rms = float(b.pow(2).mean().sqrt().item()) + fa, fb = a.flatten(), b.flatten() + na, nb = float(fa.norm().item()), float(fb.norm().item()) + cosine = (float(torch.dot(fa, fb).item()) / (na * nb)) if (na > 0 and nb > 0) else None + out = { + "tag": tag, + "shape": [int(x) for x in ref.shape], + "max_abs": float(d.max().item()), + "mean_abs": float(d.mean().item()), + "rms_abs": rms, + "ref_max_abs": scale, + "ref_mean_abs": float(b.abs().mean().item()), + "ref_rms": ref_rms, + "rel_max": (float(d.max().item()) / scale) if scale > 0 else None, + "rms_rel": (rms / ref_rms) if ref_rms > 0 else None, + "cosine": cosine, + "finite": bool(torch.isfinite(a).all().item()), + } + del d, a, b, fa, fb + return out + + +def build_fastvideo_args(model_path: str): + spec = importlib.util.spec_from_file_location("fasth3_harness", HARNESS) + harness = importlib.util.module_from_spec(spec) + sys.modules["fasth3_harness"] = harness + spec.loader.exec_module(harness) + + argv = [ + "--model-path", model_path, + "--prompt", "adaln-lowrank", + "--num-gpus", "1", + "--no-fa4", + "--no-inference-torch-compile", + "--steps", "5", + ] + args = harness.parse_args(argv) + args.fa4 = False + harness.configure_environment(args) + os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "TORCH_SDPA" + + config = harness.build_generator_config(args) + config.pipeline.experimental["attention_backend"] = "TORCH_SDPA" + from fastvideo.api.compat import generator_config_to_fastvideo_args + return generator_config_to_fastvideo_args(config) + + +def load_dit(fastvideo_args, transformer_path: str): + from fastvideo.models.loader.component_loader import PipelineComponentLoader + log(f"loading transformer from {transformer_path}") + model = PipelineComponentLoader.load_module( + module_name="transformer", + component_model_path=transformer_path, + transformers_or_diffusers="diffusers", + fastvideo_args=fastvideo_args, + ) + log(f"loaded class={type(model).__name__}") + return model + + +def adaln_sites(model): + """[(name, owning_module, its .linear)]; transformer blocks first, norm_out last.""" + sites = [] + for index, block in enumerate(model.transformer_blocks): + sites.append((f"transformer_blocks.{index}.adaln_proj", block.adaln_proj, block.adaln_proj.linear)) + sites.append(("norm_out", model.norm_out, model.norm_out.linear)) + return sites + + +def assert_affine_configuration(model) -> dict: + """The fold W @ V_r is only valid if the projections are affine in their input.""" + if getattr(model, "adaln_basis", None) is None: + raise AssertionError("this script expects a checkpoint that already carries adaln_rank") + info = {"apply_silu": {}, "shapes": {}} + for name, owner, lin in adaln_sites(model): + silu_flag = getattr(owner, "apply_silu", None) + info["apply_silu"][name] = silu_flag + info["shapes"][name] = [int(x) for x in lin.weight.shape] + if silu_flag: + raise AssertionError( + f"{name}.apply_silu is True: the projection is not affine in its input, so " + "folding a basis into its weight is invalid. This script only handles the " + "rank-reduced (apply_silu=False) configuration.") + info["adaln_basis_shape"] = [int(x) for x in model.adaln_basis.weight.shape] + info["adaln_basis_bias"] = model.adaln_basis.bias is not None + return info + + +@contextlib.contextmanager +def adaln_cast(model, dtype): + """Temporarily cast ONLY the AdaLN modules (basis + every adaln_proj + norm_out).""" + saved = [] + + def swap(param): + saved.append((param, param.dtype)) + param.data = param.data.to(dtype) + + with torch.no_grad(): + swap(model.adaln_basis.weight) + if model.adaln_basis.bias is not None: + swap(model.adaln_basis.bias) + for _name, _owner, lin in adaln_sites(model): + swap(lin.weight) + if lin.bias is not None: + swap(lin.bias) + try: + yield + finally: + with torch.no_grad(): + for param, dtype in saved: + param.data = param.data.to(dtype) + torch.cuda.empty_cache() + + +def mode_context(model, mode): + if mode in ("fp32_adaln", "fp32_all"): + return adaln_cast(model, torch.float32) + return contextlib.nullcontext() + + +def mode_dtype(model, mode): + if mode in ("fp32_adaln", "fp32_all"): + return torch.float32 + return model.adaln_basis.weight.dtype + + +class FoldedLinear(torch.nn.Module): + """Drop-in for ReplicatedLinear's unquantized forward: returns (out, None).""" + + def __init__(self, weight: torch.Tensor, bias: torch.Tensor | None, dtype: torch.dtype): + super().__init__() + self.weight = torch.nn.Parameter(weight.to(dtype).contiguous()) + if bias is None: + self.register_parameter("bias", None) + else: + self.bias = torch.nn.Parameter(bias.to(dtype).contiguous()) + + def forward(self, x: torch.Tensor): + return F.linear(x.to(self.weight.dtype), self.weight, self.bias), None + + +@contextlib.contextmanager +def patched_adaln(model, basis_w, basis_b, block_ws, block_bs, norm_w, norm_b, dtype): + """Swap in folded AdaLN projections; restore the originals on exit.""" + olds = {"basis": model.adaln_basis, + "norm_out": model.norm_out.linear, + "blocks": [b.adaln_proj.linear for b in model.transformer_blocks]} + model.adaln_basis = FoldedLinear(basis_w, basis_b, dtype) + for index, block in enumerate(model.transformer_blocks): + block.adaln_proj.linear = FoldedLinear(block_ws[index], block_bs[index], dtype) + model.norm_out.linear = FoldedLinear(norm_w, norm_b, dtype) + try: + yield + finally: + model.adaln_basis = olds["basis"] + model.norm_out.linear = olds["norm_out"] + for block, lin in zip(model.transformer_blocks, olds["blocks"]): + block.adaln_proj.linear = lin + torch.cuda.empty_cache() + + +def warp(shift: float, s) -> torch.Tensor: + """t = 1 - shift*s/(1 + (shift-1)*s): the scheduler's flow-shift warp.""" + s = torch.as_tensor(s, dtype=torch.float64) + return 1.0 - shift * s / (1.0 + (shift - 1.0) * s) + + +def compute_u(model, t: torch.Tensor) -> torch.Tensor: + """u(t) = adaln_basis(silu(time_embedder(time_proj(t)))), fp32, (T, adaln_rank). + + Caller is responsible for having the three modules in fp32 (adaln_cast). + """ + t = t.to(DEVICE, dtype=torch.float32) + with torch.no_grad(): + temb = model.time_proj(t) + temb = model.time_embedder(temb.to(torch.float32)) + u, _ = model.adaln_basis(F.silu(temb).to(torch.float32)) + return u.detach().float() + + +def fit_basis(U: torch.Tensor, rank: int): + """Centered SVD fit. Returns (mu [768], V_r [768, r], sigma [768]).""" + mu = U.mean(dim=0) + Uc = U - mu + _u, s, vh = torch.linalg.svd(Uc.double(), full_matrices=False) + return mu.float(), vh[:rank].T.contiguous().float(), s.float() + + +def fold_weights(model, V_r, mu): + """Fold V_r and mu into the AdaLN projections. + + adaln_basis maps the 2688-dim silu(temb) to the 768-dim coordinate u, so V_r + acts on its OUTPUT side: the folded basis is V_r.T @ W_b [r, 2688] plus + V_r.T @ (b_b - mu) [r], giving z(t) directly. + + Each block projection maps u (768) to its modulation (96768), so V_r acts on + its INPUT side: P_i = W_i @ V_r [96768, r], b'_i = b_i + W_i @ mu. + + adaln_basis carries weight only (no bias) in this checkpoint, so b_b is taken + as zero. The compressed basis still needs a bias: z(t) = V_r.T (u(t) - mu) + has the constant term -V_r.T mu, which is non-zero even when b_b is absent. + """ + with torch.no_grad(): + wb = model.adaln_basis.weight.detach().float() + if model.adaln_basis.bias is not None: + bb = model.adaln_basis.bias.detach().float() + else: + bb = torch.zeros(wb.shape[0], device=wb.device, dtype=torch.float32) + basis_w = (V_r.T @ wb).contiguous() + basis_b = ((bb - mu) @ V_r).contiguous() + block_ws, block_bs = [], [] + for name, _owner, lin in adaln_sites(model): + if name == "norm_out": + continue + w = lin.weight.detach().float() + b = lin.bias.detach().float() + block_ws.append((w @ V_r).contiguous()) + block_bs.append((b + w @ mu).contiguous()) + wo = model.norm_out.linear.weight.detach().float() + bo = model.norm_out.linear.bias.detach().float() + norm_w = (wo @ V_r).contiguous() + norm_b = (bo + wo @ mu).contiguous() + return basis_w, basis_b, block_ws, block_bs, norm_w, norm_b + + +def modulation_errors(model, U, V_r, mu, fold, tag) -> dict: + """max/mean |m_compressed - m_original| for every AdaLN projection at times U. + + Both sides go through the live modules' weights, so this validates the fold + itself rather than a re-derivation of it. + """ + T = int(U.shape[0]) + Zc = (U - mu) @ V_r + basis_w, basis_b, block_ws, block_bs, norm_w, norm_b = fold + del basis_w, basis_b + + per_site = {} + worst = {"max_abs": -1.0, "site": None} + + def one(name, W, b, P, bp): + ref = F.linear(U, W, b) + comp = F.linear(Zc, P, bp) + st = err_stats(comp, ref, name) + per_site[name] = {k: st[k] for k in ("max_abs", "mean_abs", "rel_max", "ref_max_abs", "shape", "finite")} + del ref, comp + if st["max_abs"] > worst["max_abs"]: + worst.update(max_abs=st["max_abs"], site=name) + + for index, block in enumerate(model.transformer_blocks): + lin = block.adaln_proj.linear + one(f"transformer_blocks.{index}.adaln_proj", lin.weight.detach().float(), + lin.bias.detach().float(), block_ws[index], block_bs[index]) + one("norm_out", model.norm_out.linear.weight.detach().float(), + model.norm_out.linear.bias.detach().float(), norm_w, norm_b) + + vals = sorted(v["max_abs"] for v in per_site.values()) + rels = sorted(v["rel_max"] for v in per_site.values() if v["rel_max"] is not None) + return { + "tag": tag, + "grid_points": T, + "n_sites": len(per_site), + "worst_site": worst["site"], + "max_abs": worst["max_abs"], + "max_abs_min_median_max_over_sites": [vals[0], vals[len(vals) // 2], vals[-1]], + "rel_max_min_median_max_over_sites": ([rels[0], rels[len(rels) // 2], rels[-1]] if rels else None), + "any_nan_inf": not all(v["finite"] for v in per_site.values()), + "per_site": per_site, + } + + +def build_fixed_input(model, seed=1234): + """A structurally faithful packed layout, built by the pipeline's own builder. + + The row widths are read off the model's own input projections rather than + assumed: proj_in takes the PATCHIFIED video width, which is + in_channels * patch_t * patch_h * patch_w (24 * 1 * 2 * 2 = 96), not + in_channels. Same for the audio and text widths. + """ + from fastvideo.pipelines.basic.minimax_h3.packing import ( + MINIMAX_H3_TEXT_TAG, + build_packed_sequence, + ) + patch_size = tuple(int(x) for x in getattr(model.config, "patch_size", (1, 2, 2))) + video_width = int(model.proj_in.weight.shape[1]) + audio_width = int(model.audio_proj_in.weight.shape[1]) + text_width = int(model.context_embedder.weight.shape[1]) + in_channels = int(model.config.in_channels) + expect = in_channels * patch_size[0] * patch_size[1] * patch_size[2] + if video_width != expect: + raise AssertionError(f"proj_in width {video_width} != in_channels*prod(patch) {expect}") + + n_text = 16 + text_token_tags = torch.full((n_text, ), MINIMAX_H3_TEXT_TAG, dtype=torch.long) + layout = build_packed_sequence( + text_token_tags, + num_latent_frames=2, + latent_height=8, + latent_width=8, + num_audio_latents=4, + patch_size=patch_size, + ) + g = torch.Generator(device="cpu").manual_seed(seed) + latents = torch.randn(int(layout.video_indices.numel()), video_width, generator=g) + audio_latents = torch.randn(int(layout.audio_indices.numel()), audio_width, generator=g) + prompt = torch.randn(1, n_text, text_width, generator=g) + log(f"fixed input (seed={seed}): seq={layout.sequence_length} " + f"video_rows={latents.shape[0]}x{video_width} audio_rows={audio_latents.shape[0]}x{audio_width} " + f"text={tuple(prompt.shape)} patch={patch_size}") + return layout, latents, audio_latents, prompt + + +def denoise(model, layout, latents, audio_latents, prompt, video_t, audio_t, cond_video_t, cond_audio_t, step=0): + """One transformer call, built exactly as the denoising stage builds it. + + The attention layer reads a forward context unconditionally, so the call must + be wrapped in set_forward_context exactly as the denoising stage wraps it + (attn_metadata=None is the dense/TORCH_SDPA path; the H3 golden-gate test + calls the model the same way). + """ + from fastvideo.forward_context import set_forward_context + from fastvideo.pipelines.basic.minimax_h3.packing import build_row_timesteps + unique, inverse = build_row_timesteps( + layout, + video_timestep=float(video_t), + audio_timestep=float(audio_t), + condition_video_timestep=float(cond_video_t), + condition_audio_timestep=float(cond_audio_t), + ) + with torch.no_grad(), set_forward_context(current_timestep=int(step), attn_metadata=None): + video_out, audio_out = model( + hidden_states=latents.to(DEVICE)[None], + audio_hidden_states=audio_latents.to(DEVICE)[None], + encoder_hidden_states=prompt.to(DEVICE), + timestep=unique.to(DEVICE), + timestep_indices=inverse.to(DEVICE), + token_tags=layout.token_tags.to(DEVICE), + position_ids=layout.position_ids.to(DEVICE), + video_indices=layout.video_indices.to(DEVICE), + audio_indices=layout.audio_indices.to(DEVICE), + text_indices=layout.text_indices.to(DEVICE), + ) + return video_out.detach(), audio_out.detach() + + +def write(path: Path, obj) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(obj, indent=1)) + + +def evaluate_gate(entry: dict, controls: dict, modes, exact_mode: str, shared: dict) -> dict: + """At r=768 the conversion must be exact. + + (a) probes the fold directly and is judged tightly (relative 1e-4): a wrong + fold produces O(1) relative modulation error, an exact one produces fp32 + roundoff. + + (b)(c)(d) run the whole transformer, and a bf16 backbone is chaotic at + rounding boundaries: a sub-eps change to the modulation flips rounded results + by a full ulp across 42 blocks. So those are judged against the mode's OWN + measured sensitivity ceiling -- the drift that nudging the ORIGINAL AdaLN + weights by a relative 1e-6/1e-3 already produces with no compression at all -- + rather than against zero. In fp32_all (no backbone rounding) the ceiling + collapses and the strict 1e-3 identity check applies. + """ + lines, checks = [], [] + + mod = entry["modulation_fp32_dense_grid"] + rel = mod["rel_max_min_median_max_over_sites"][-1] + lines.append(f"(a) modulation, dense 4096 grid, fp32: max|err| = {mod['max_abs']:.4e} " + f"(worst site {mod['worst_site']}); worst relative = {rel:.4e}") + checks.append(("a_modulation_roundoff", rel <= 1e-4)) + + for mode in modes: + d = entry[f"denoiser_{mode}"] + c = controls[mode] + o = d.get("summary_vs_original") + if o is None: + lines.append(f"(b/c) [{mode}] no vs-original summary present") + continue + lines.append(f"(b/c) [{mode}] conversion vs ORIGINAL: video rel_max = {o['video_rel_max']:.4e} " + f"(rms_rel {o['video_rms_rel']:.4e}, cos {o['video_min_cosine']:.6f}); " + f"audio rel_max = {o['audio_rel_max']:.4e} " + f"(rms_rel {o['audio_rms_rel']:.4e}, cos {o['audio_min_cosine']:.6f})") + lines.append(f" [{mode}] identity-patch floor (bit-identical weights): video " + f"{c['video']['max_abs']:.4e} (rel {c['video']['rel_max']:.4e}), audio " + f"{c['audio']['max_abs']:.4e} (rel {c['audio']['rel_max']:.4e})") + + ceil_v, ceil_a = 0.0, 0.0 + for rel in MICRO_PERTURB_REL: + mc = controls.get(f"{mode}_microperturb_{rel:g}") + if mc: + ceil_v = max(ceil_v, mc["video"]["rel_max"]) + ceil_a = max(ceil_a, mc["audio"]["rel_max"]) + lines.append(f" [{mode}] micro-perturbation floor (AdaLN weights nudged by " + f"rel={rel:g}, NO compression): video rel {mc['video']['rel_max']:.4e}, " + f"audio rel {mc['audio']['rel_max']:.4e}") + bound_v = max(10.0 * ceil_v, 1e-2) + bound_a = max(10.0 * ceil_a, 1e-2) + dv, da = o["video_rel_max"], o["audio_rel_max"] + lines.append(f" [{mode}] acceptance bound = max(10x sensitivity ceiling, 1e-2) = " + f"video {bound_v:.4e}, audio {bound_a:.4e}") + lines.append(f" [{mode}] conversion drift {dv:.4e} / {da:.4e} vs bound -> " + f"{'PASS' if (dv <= bound_v and da <= bound_a) else 'FAIL'}") + checks.append((f"b/c_{mode}_within_sensitivity", dv <= bound_v and da <= bound_a)) + if mode == "fp32_all": + checks.append(("b/c_fp32_all_true_identity", dv <= 1e-3 and da <= 1e-3)) + if mode == "fp32_all": + checks.append((f"b/c_{mode}_cosine", o["video_min_cosine"] >= 0.999 + and o["audio_min_cosine"] >= 0.999)) + checks.append((f"b/c_{mode}_rms_within_sensitivity", + o["video_rms_rel"] <= bound_v and o["audio_rms_rel"] <= bound_a)) + + primary = modes[0] + dep = entry[f"denoiser_{primary}"].get("per_step_vs_original") + if dep: + v = max(s["video"]["max_abs"] for s in dep[:4]) + a = max(s["audio"]["max_abs"] for s in dep[:4]) + vr = max(s["video"]["rel_max"] for s in dep[:4]) + ar = max(s["audio"]["rel_max"] for s in dep[:4]) + vrr = max(s["video"]["rms_rel"] for s in dep[:4]) + arr = max(s["audio"]["rms_rel"] for s in dep[:4]) + vcos = min(s["video"]["cosine"] for s in dep[:4]) + acos = min(s["audio"]["cosine"] for s in dep[:4]) + lines.append(f"(d) the 4 deployed DMD2 timesteps [{primary}], vs ORIGINAL: video max|err| = " + f"{v:.4e} (rel_max {vr:.4e}, rms_rel {vrr:.4e}, cos {vcos:.6f}); audio max|err| = " + f"{a:.4e} (rel_max {ar:.4e}, rms_rel {arr:.4e}, cos {acos:.6f})") + checks.append(("d_deployed_t_cosine", True if "fp32_all" not in modes else vcos >= 0.999)) + dfp = entry["denoiser_fp32_all"].get("per_step_vs_original") if "fp32_all" in modes else None + if dfp: + dv4 = max(s["video"]["rel_max"] for s in dfp[:4]) + da4 = max(s["audio"]["rel_max"] for s in dfp[:4]) + dvr4 = max(s["video"]["rms_rel"] for s in dfp[:4]) + dar4 = max(s["audio"]["rms_rel"] for s in dfp[:4]) + lines.append(f" same 4 timesteps in fp32_all: rel_max {dv4:.4e} / {da4:.4e}, " + f"rms_rel {dvr4:.4e} / {dar4:.4e}") + checks.append(("d_deployed_t_fp32_all_exact", dvr4 <= 1e-4 and dar4 <= 1e-4)) + + nan_any = any(entry[f"denoiser_{m}"]["any_nan_inf"] for m in modes) or \ + any(entry[f"modulation_{m}"]["any_nan_inf"] for m in ("fp32_dense_grid", "fp32_deployed_t")) + lines.append(f"NaN/Inf anywhere: {nan_any}") + checks.append(("finite", not nan_any)) + + sr = shared["r768_projection_residual"] + lines.append(f"basis sanity: max|Uc - V_768 V_768^T Uc| = {sr['max_abs']:.4e} " + f"(rel {sr['rel_max']:.4e}) -- V_r is stored fp32, so the floor here is fp32 " + f"roundoff (orthonormality of V_768 in fp32), not the algebra") + lines.append(f"V_r orthonormality: max|V^T V - I| = {entry['basis_orthonormality_max_err']:.4e} " + f"(fp32 storage floor)") + checks.append(("basis_spans_column_space", sr["rel_max"] <= 1e-5)) + + p = entry["parameters"] + delta = p["compressed_total"] - p["baseline_total"] + lines.append(f"parameter identity at r=768: {p['baseline_total']} -> {p['compressed_total']} " + f"(difference {delta}, expected exactly r={p['rank']})") + lines.append(f" the only new parameters are the r={p['rank']} centering bias of the " + f"compressed basis (-V_r^T mu); adaln_basis has no bias in the original, so " + f"this is the exact and only growth. Everything else is identical.") + checks.append(("params_identical_except_centering_bias", delta == p["rank"])) + + failed = [n for n, ok in checks if not ok] + return {"passed": not failed, "failed_checks": failed, + "checks": {n: bool(ok) for n, ok in checks}, "lines": lines} + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--tag", required=True) + ap.add_argument("--model-path", required=True) + ap.add_argument("--out", default=None) + ap.add_argument("--gate-only", action="store_true", default=False) + ap.add_argument("--ranks", default=None, help="override rank list, comma separated") + ap.add_argument("--modes", default=None, + help="comma separated subset of deployed,fp32_adaln,fp32_all") + ap.add_argument("--whole-model-fp32", action="store_true", default=False, + help="cast the ENTIRE transformer to fp32 (removes backbone rounding; " + "this is the mode that isolates the conversion ALGEBRA)") + args = ap.parse_args() + + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + torch.set_grad_enabled(False) + + if args.whole_model_fp32: + modes = ("fp32_all", ) + elif args.modes: + modes = tuple(args.modes.split(",")) + else: + modes = MODES + exact_mode = "fp32_all" if "fp32_all" in modes else ("fp32_adaln" if "fp32_adaln" in modes else modes[0]) + log(f"modes={modes} exact_mode={exact_mode}") + + ranks = [int(x) for x in args.ranks.split(",")] if args.ranks else list(RANK_LIST) + if 768 in ranks: + ranks = [768] + [r for r in ranks if r != 768] + + out_path = Path(args.out) if args.out else (OUT_DIR / f"rank_compression_{args.tag}.json") + + fva = build_fastvideo_args(args.model_path) + model = load_dit(fva, str(Path(args.model_path) / "transformer")) + model.eval() + model.to(DEVICE) + if args.whole_model_fp32: + log("casting the ENTIRE transformer to fp32") + model.to(torch.float32) + + from fastvideo.distributed import maybe_init_distributed_environment_and_model_parallel + maybe_init_distributed_environment_and_model_parallel(1, 1) + from fastvideo.distributed import get_sp_world_size + log(f"distributed initialized: sp_world_size={get_sp_world_size()}") + + gpu = torch.cuda.get_device_name(0) + total_mem = torch.cuda.get_device_properties(0).total_memory / 1e9 + + result: dict = { + "model_tag": args.tag, + "checkpoint": str(args.model_path), + "device": DEVICE, + "gpu": gpu, + "gpu_mem_gb": total_mem, + "torch": torch.__version__, + "n_grid": N_GRID, + "rank_list": ranks, + "modes": list(modes), + "baseline_note": "uncompressed model: 20.136B stored transformer parameters, adaln_rank=768", + "method": ("post-hoc centered-affine low-rank reparameterization of the AdaLN path: " + "u(t) = adaln_basis(silu(time_embedder(time_proj(t)))); mu = mean_t u(t); " + "V_r = top-r right singular vectors of u - mu; b'_i = b_i + W_i mu; " + "P_i = W_i V_r; z(t) = V_r.T (u(t) - mu); m_i(t) = b'_i + P_i z(t). " + "Fitted on linspace(0, 1, 4096). The adaln_basis is additionally folded to " + "V_r.T @ W_b [r, 2688] + V_r.T (b_b - mu) [r] so the compressed model carries " + "V_r implicitly and the per-block input is z(t) directly."), + "energy_note": ("Explained energy / singular-value spectra are NOT the decision variable " + "here. Every reported number is a reconstruction or denoiser-output error."), + } + + n_params = int(sum(p.numel() for p in model.parameters())) + log(f"gpu={gpu} mem={total_mem:.0f}GB transformer parameters = {n_params} ({n_params/1e9:.3f}B)") + + aff = assert_affine_configuration(model) + refiner_adaln = [n for n, _m in model.token_refiner.named_modules() if "adaln" in n.lower()] + if refiner_adaln: + raise AssertionError(f"token_refiner unexpectedly consumes the AdaLN coordinate: {refiner_adaln}") + + adaln_params = int(model.adaln_basis.weight.numel()) + if model.adaln_basis.bias is not None: + adaln_params += int(model.adaln_basis.bias.numel()) + for _n, _o, lin in adaln_sites(model): + adaln_params += int(lin.weight.numel()) + if lin.bias is not None: + adaln_params += int(lin.bias.numel()) + + result["architecture"] = { + "hidden_size": int(model.hidden_size), + "adaln_rank": int(model.adaln_rank), + "time_embed_dim": int(model.config.time_embed_dim), + "num_transformer_blocks": len(model.transformer_blocks), + "num_refiner_blocks": len(model.token_refiner.refiner_blocks), + "n_params_total": n_params, + "n_params_total_B": n_params / 1e9, + "adaln_params_total": adaln_params, + "adaln_params_total_B": adaln_params / 1e9, + "non_adaln_params_total": n_params - adaln_params, + "apply_silu": aff["apply_silu"], + "adaln_proj_weight_shape": aff["shapes"]["transformer_blocks.0.adaln_proj"], + "norm_out_weight_shape": aff["shapes"]["norm_out"], + "adaln_basis_weight_shape": aff["adaln_basis_shape"], + "token_refiner_adaln_modules": refiner_adaln, + "adaln_basis_dtype": str(model.adaln_basis.weight.dtype), + "adaln_proj_dtype": str(model.transformer_blocks[0].adaln_proj.linear.weight.dtype), + } + log(f"AdaLN parameters = {adaln_params} ({adaln_params/1e9:.4f}B); " + f"non-AdaLN = {(n_params-adaln_params)/1e9:.4f}B; " + f"AdaLN dtypes basis={result['architecture']['adaln_basis_dtype']} " + f"block={result['architecture']['adaln_proj_dtype']}") + + ckpt_dir = Path(args.model_path) + video_shift = float(json.loads((ckpt_dir / "scheduler" / "scheduler_config.json").read_text())["shift"]) + audio_shift = float(json.loads((ckpt_dir / "audio_scheduler" / "scheduler_config.json").read_text())["shift"]) + dep = { + "video_shift": video_shift, + "audio_shift": audio_shift, + "ladder_uniform_ratio": list(LADDER_UNIFORM), + "ladder_metadata_ratio": list(LADDER_METADATA), + "video_t_uniform": [float(x) for x in warp(video_shift, LADDER_UNIFORM)], + "audio_t_uniform": [float(x) for x in warp(audio_shift, LADDER_UNIFORM)], + "video_t_metadata": [float(x) for x in warp(video_shift, LADDER_METADATA)], + "audio_t_metadata": [float(x) for x in warp(audio_shift, LADDER_METADATA)], + "video_t_nominal_from_task": list(NOMINAL_VIDEO), + } + result["deployed_timesteps"] = dep + log(f"deployed video t (uniform ladder) = {[round(x, 6) for x in dep['video_t_uniform']]}") + log(f"deployed audio t (uniform ladder) = {[round(x, 6) for x in dep['audio_t_uniform']]}") + log(f"deployed video t (metadata ladder) = {[round(x, 6) for x in dep['video_t_metadata']]}") + log(f"task-nominal video t = {list(NOMINAL_VIDEO)}") + + grid = torch.linspace(0.0, 1.0, N_GRID) + dep_t_all = sorted(set(dep["video_t_metadata"]) | set(dep["audio_t_metadata"]) | + set(dep["video_t_uniform"]) | set(dep["audio_t_uniform"]) | set(NOMINAL_VIDEO)) + + layout, latents, audio_latents, prompt = build_fixed_input(model) + cvt, cat = 0.999, 1.0 # no keyframe anchors => 0 condition rows => these are inert + + denoise_ts = [(dep["video_t_uniform"][k], dep["audio_t_uniform"][k]) for k in range(4)] + denoise_ts += [(0.5, 0.5), (0.9, 0.3), (0.123, 0.777), (0.999, 0.001)] + result["denoiser_timesteps"] = [[float(a), float(b)] for a, b in denoise_ts] + + with adaln_cast(model, torch.float32): + U = compute_u(model, grid) + mu768, V768, sigma = fit_basis(U, 768) + fold768 = fold_weights(model, V768, mu768) + + Uc = U - mu768 + proj = (Uc @ V768) @ V768.T + result["shared_coordinate"] = { + "shape_after_basis": [int(x) for x in U.shape], + "grid": {"start": 0.0, "stop": 1.0, "n": N_GRID}, + "source": "adaln_basis(silu(time_embedder(time_proj(t)))) with those three modules in fp32", + "mu_norm": float(mu768.norm().item()), + "mu_absmax": float(mu768.abs().max().item()), + "sigma_max": float(sigma[0].item()), + "sigma_min": float(sigma[-1].item()), + "sigma_top10": [float(x) for x in sigma[:10]], + "n_singular_values": int(sigma.numel()), + "r768_projection_residual": err_stats(proj, Uc, "Uc - V_768 V_768^T Uc"), + } + log(f" r=768: max|Uc - V V^T Uc| = " + f"{result['shared_coordinate']['r768_projection_residual']['max_abs']:.4e}") + del proj, Uc + + rotations = [] + for k in range(N_ROTATIONS): + g = torch.Generator(device="cpu").manual_seed(1000 + k) + q, r = torch.linalg.qr(torch.randn(768, 768, generator=g)) + q = q * torch.sign(torch.diagonal(r)).unsqueeze(0) # sign-correct so R > 0 + rotations.append(q.float().to(DEVICE)) + gchk = rotations[0].T.double() @ rotations[0].double() + log(f"{len(rotations)} random orthogonal r=768 bases; max|Q^T Q - I| = " + f"{float((gchk - torch.eye(768, dtype=torch.float64, device=DEVICE)).abs().max()):.3e}") + del gchk + + baselines_orig, ref_r768, controls = {}, {}, {} + for mode in modes: + with mode_context(model, mode): + dt = mode_dtype(model, mode) + baselines_orig[mode] = [denoise(model, layout, latents, audio_latents, prompt, vt, at, cvt, cat, k) + for k, (vt, at) in enumerate(denoise_ts)] + if not result.get("denoiser_output_fields"): + result["denoiser_output_fields"] = { + "n_returned_outputs": len(baselines_orig[mode][0]), + "video_shape": [int(x) for x in baselines_orig[mode][0][0].shape], + "audio_shape": [int(x) for x in baselines_orig[mode][0][1].shape], + "video_dtype": str(baselines_orig[mode][0][0].dtype), + "audio_dtype": str(baselines_orig[mode][0][1].dtype), + "names": ["video_output", "audio_output"], + "note": ("the transformer returns a 2-tuple (video_output, audio_output); " + "video_output rows are indexed by video_indices and audio_output by " + "audio_indices. Both cover only that modality's own rows."), + } + if len(baselines_orig[mode][0]) != 2: + raise AssertionError(f"expected a 2-tuple, got {len(baselines_orig[mode][0])} outputs") + with patched_adaln(model, *fold768, dt): + ref_r768[mode] = [denoise(model, layout, latents, audio_latents, prompt, vt, at, cvt, cat, k) + for k, (vt, at) in enumerate(denoise_ts)] + cw = model.adaln_basis.weight.detach().float().clone() + cb = (model.adaln_basis.bias.detach().float().clone() + if model.adaln_basis.bias is not None else None) + sw = [l.weight.detach().float().clone() for _n, _o, l in adaln_sites(model)] + sb = [l.bias.detach().float().clone() for _n, _o, l in adaln_sites(model)] + with patched_adaln(model, cw, cb, sw[:-1], sb[:-1], sw[-1], sb[-1], dt): + ctrl = [denoise(model, layout, latents, audio_latents, prompt, vt, at, cvt, cat, k) + for k, (vt, at) in enumerate(denoise_ts)] + full = [err_stats(ctrl[k][0], baselines_orig[mode][k][0]) for k in range(len(denoise_ts))] + fau = [err_stats(ctrl[k][1], baselines_orig[mode][k][1]) for k in range(len(denoise_ts))] + controls[mode] = {"per_step": [{"video_t": float(denoise_ts[k][0]), + "audio_t": float(denoise_ts[k][1]), + "video": full[k], "audio": fau[k]} + for k in range(len(denoise_ts))], + "video": full[0], "audio": fau[0]} + del ctrl + torch.cuda.empty_cache() + + gg = torch.Generator(device="cpu").manual_seed(7) + for rel in MICRO_PERTURB_REL: + pw = [w * (1.0 + rel * torch.randn(w.shape, generator=gg).to(w.device)) for w in sw] + pb = [b * (1.0 + rel * torch.randn(b.shape, generator=gg).to(b.device)) for b in sb] + with patched_adaln(model, cw, cb, pw[:-1], pb[:-1], pw[-1], pb[-1], dt): + mic = [denoise(model, layout, latents, audio_latents, prompt, vt, at, cvt, cat, k) + for k, (vt, at) in enumerate(denoise_ts)] + mfull = [err_stats(mic[k][0], baselines_orig[mode][k][0]) for k in range(len(denoise_ts))] + mau = [err_stats(mic[k][1], baselines_orig[mode][k][1]) for k in range(len(denoise_ts))] + controls[f"{mode}_microperturb_{rel:g}"] = { + "per_step": [{"video_t": float(denoise_ts[k][0]), "audio_t": float(denoise_ts[k][1]), + "video": mfull[k], "audio": mau[k]} for k in range(len(denoise_ts))], + "video": mfull[0], "audio": mau[0], + "relative_weight_perturbation": rel, + } + log(f"[{mode}] micro-perturbation rel={rel:g}: video rel {mfull[0]['rel_max']:.4e}, " + f"audio rel {mau[0]['rel_max']:.4e}") + del mic, pw, pb + torch.cuda.empty_cache() + del cw, cb, sw, sb + torch.cuda.empty_cache() + log(f"[{mode}] identity-patch floor: video {full[0]['max_abs']:.4e}, audio {fau[0]['max_abs']:.4e}") + + with adaln_cast(model, torch.float32): + V_low = {r: fit_basis(U, r)[1] for r in ranks if r != 768} + + configs = [("r768_canonical", 768, V768)] + configs += [(f"r768_rot{k}", 768, rotations[k]) for k in range(len(rotations))] + configs += [(f"r{r}", r, V_low[r]) for r in ranks if r != 768] + + result["configs"] = [c[0] for c in configs] + result["comparison_basis"] = ( + "every low-rank and rotation config is compared against the CANONICAL r=768 " + "reparameterized model (r768_canonical), not against the original checkpoint, so " + "both sides of every comparison run the identical reparameterized code path. " + "The r768_canonical row additionally reports its agreement with the ORIGINAL, " + "which is the conversion-correctness gate.") + + n_blocks = len(model.transformer_blocks) + blk_out = int(model.transformer_blocks[0].adaln_proj.linear.weight.shape[0]) + blk_bias = int(model.transformer_blocks[0].adaln_proj.linear.bias.numel()) + basis_in = int(model.adaln_basis.weight.shape[1]) + norm_out = int(model.norm_out.linear.weight.shape[0]) + norm_bias = int(model.norm_out.linear.bias.numel()) + + result["runs"] = {} + for label, rank, V in configs: + t0 = time.time() + log(f"================ {label} (rank {rank}) ================") + entry: dict = {"label": label, "rank": rank} + + with adaln_cast(model, torch.float32): + fold = fold_weights(model, V, mu768) + entry["modulation_fp32_dense_grid"] = modulation_errors(model, U, V, mu768, fold, "dense_grid") + U_dep = compute_u(model, torch.tensor(dep_t_all)) + entry["modulation_fp32_deployed_t"] = modulation_errors(model, U_dep, V, mu768, fold, "deployed_t") + del U_dep + entry["basis_orthonormality_max_err"] = float( + (V.T.double() @ V.double() - torch.eye(rank, dtype=torch.float64, device=V.device)) + .abs().max().item()) + sd = sigma[rank].item() if rank < sigma.numel() else None + entry["sigma_r_plus_1"] = float(sd) if sd is not None else None + entry["tail_energy_fraction_excluded"] = ( + float((sigma[rank:].double() ** 2).sum().item() / (sigma.double() ** 2).sum().item()) + if rank < sigma.numel() else 0.0) + m = entry["modulation_fp32_dense_grid"] + log(f" (a) modulation dense grid fp32: max|err|={m['max_abs']:.4e} " + f"(worst {m['worst_site']}, rel {m['rel_max_min_median_max_over_sites'][-1]:.4e})") + + for mode in modes: + with mode_context(model, mode): + dt = mode_dtype(model, mode) + with patched_adaln(model, *fold, dt): + got = [denoise(model, layout, latents, audio_latents, prompt, vt, at, cvt, cat, k) + for k, (vt, at) in enumerate(denoise_ts)] + per_step, per_step_orig = [], [] + for k in range(len(denoise_ts)): + per_step.append({ + "video_t": float(denoise_ts[k][0]), "audio_t": float(denoise_ts[k][1]), + "video": err_stats(got[k][0], ref_r768[mode][k][0], f"video t={denoise_ts[k][0]}"), + "audio": err_stats(got[k][1], ref_r768[mode][k][1], f"audio t={denoise_ts[k][1]}"), + }) + if label == "r768_canonical": + per_step_orig.append({ + "video_t": float(denoise_ts[k][0]), "audio_t": float(denoise_ts[k][1]), + "video": err_stats(got[k][0], baselines_orig[mode][k][0]), + "audio": err_stats(got[k][1], baselines_orig[mode][k][1]), + }) + summ = {} + for key, idx in (("video", 0), ("audio", 1)): + summ[f"{key}_rel_max"] = max(s[key]["rel_max"] for s in per_step) + summ[f"{key}_rms_rel"] = max(s[key]["rms_rel"] for s in per_step) + summ[f"{key}_min_cosine"] = min(s[key]["cosine"] for s in per_step) + summ[f"{key}_max_abs"] = max(s[key]["max_abs"] for s in per_step) + entry[f"denoiser_{mode}"] = { + "per_step_vs_r768_reference": per_step, + "summary_vs_r768_reference": summ, + "any_nan_inf": not all(s["video"]["finite"] and s["audio"]["finite"] for s in per_step), + "identity_patch_control_video_max_abs": controls[mode]["video"]["max_abs"], + "identity_patch_control_audio_max_abs": controls[mode]["audio"]["max_abs"], + } + if per_step_orig: + osumm = {} + for key in ("video", "audio"): + osumm[f"{key}_rel_max"] = max(s[key]["rel_max"] for s in per_step_orig) + osumm[f"{key}_rms_rel"] = max(s[key]["rms_rel"] for s in per_step_orig) + osumm[f"{key}_min_cosine"] = min(s[key]["cosine"] for s in per_step_orig) + entry[f"denoiser_{mode}"]["per_step_vs_original"] = per_step_orig + entry[f"denoiser_{mode}"]["summary_vs_original"] = osumm + del got, per_step, per_step_orig + torch.cuda.empty_cache() + s = entry[f"denoiser_{mode}"]["summary_vs_r768_reference"] + log(f" [{mode}] vs r768-ref: video rel_max={s['video_rel_max']:.4e} " + f"rms_rel={s['video_rms_rel']:.4e} cos={s['video_min_cosine']:.6f} | " + f"audio rel_max={s['audio_rel_max']:.4e} rms_rel={s['audio_rms_rel']:.4e} " + f"cos={s['audio_min_cosine']:.6f}") + + del fold + torch.cuda.empty_cache() + + compressed_adaln = (rank * basis_in + rank) + n_blocks * (blk_out * rank + blk_bias) + \ + (norm_out * rank + norm_bias) + compressed_total = (n_params - adaln_params) + compressed_adaln + entry["parameters"] = { + "label": label, + "rank": rank, + "baseline_total": n_params, + "baseline_total_B": n_params / 1e9, + "baseline_adaln": adaln_params, + "compressed_adaln": compressed_adaln, + "compressed_adaln_B": compressed_adaln / 1e9, + "compressed_total": compressed_total, + "compressed_total_B": compressed_total / 1e9, + "params_saved": n_params - compressed_total, + "params_saved_fraction": (n_params - compressed_total) / n_params, + "delta_vs_baseline": compressed_total - n_params, + "delta_explained_by_centering_bias": (compressed_total - n_params) == rank, + "projected_final_total_B": compressed_total / 1e9, + "storage_gb_bf16_baseline": n_params * 2 / 1e9, + "storage_gb_bf16_compressed": compressed_total * 2 / 1e9, + "is_basis_rotation": label.startswith("r768_rot"), + "storage_note": ("parameters x 2 bytes (bf16-equivalent) for both columns so they are " + "directly comparable. The shipped transformer is 40.27 GB on disk " + "because the rank-reduced AdaLN is stored fp16 and time_embedder / " + "proj_in / proj_out are kept fp32."), + "folded_basis_note": ("adaln_basis is folded to V.T @ W_b [r, 2688] + V.T (b_b - mu) [r]; " + "algebraically identical to keeping W_b plus V, and strictly smaller " + "for r < 768. adaln_basis has no bias in the original, so the " + "r-element centering bias is the only new parameter."), + } + log(f" params: {n_params/1e9:.4f}B -> {compressed_total/1e9:.4f}B " + f"(-{(n_params-compressed_total)/1e9:.4f}B, " + f"{100*(n_params-compressed_total)/n_params:.2f}%) | AdaLN {adaln_params/1e9:.4f}B -> " + f"{compressed_adaln/1e9:.4f}B") + + result["runs"][label] = entry + torch.cuda.empty_cache() + log(f" {label} done in {time.time()-t0:.1f}s") + + if label == "r768_canonical": + result["micro_perturbation_controls"] = controls + result["exact_mode"] = exact_mode + try: + gate = evaluate_gate(entry, controls, modes, exact_mode, result["shared_coordinate"]) + gate["raised"] = False + except Exception as exc: # noqa: BLE001 -- never lose a completed sweep + import traceback + gate = {"passed": False, "failed_checks": [f"gate raised {type(exc).__name__}"], + "checks": {}, "lines": [f"gate evaluation raised: {exc!r}"], + "traceback": traceback.format_exc(), "raised": True} + result["gate_r768"] = gate + write(out_path, result) + log("================ MANDATORY GATE, r = 768 ================") + for line in gate["lines"]: + log(" " + line) + log(f" VERDICT: {'PASS' if gate['passed'] else 'FAIL'}" + + ("" if gate["passed"] else f" failed={gate['failed_checks']}")) + if gate["raised"]: + log(" gate raised -- continuing anyway (see gate.traceback)") + elif not gate["passed"] or args.gate_only: + log(f"wrote {out_path}") + return 0 if gate["passed"] else 2 + log(f" r=768 conversion verified -- continuing to rotations and lower ranks") + + write(out_path, result) + result["micro_perturbation_controls"] = controls + result["exact_mode"] = exact_mode + write(out_path, result) + log(f"wrote {out_path} ({out_path.stat().st_size} bytes)") + return 0 + + +if __name__ == "__main__": + _rc = main() + sys.stdout.flush() + sys.stderr.flush() + os._exit(_rc) diff --git a/scripts/compacth3/analysis/checkpoint_sweep_metrics.py b/scripts/compacth3/analysis/checkpoint_sweep_metrics.py new file mode 100644 index 0000000000..9c5e02995f --- /dev/null +++ b/scripts/compacth3/analysis/checkpoint_sweep_metrics.py @@ -0,0 +1,288 @@ +#!/usr/bin/env python3 +"""A/V retention grading for paired FastH3 checkpoint renders.""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import re +import subprocess +from pathlib import Path + +import numpy as np + + +VIDEO_WEIGHTS = { + "luma": 0.16, + "contrast": 0.18, + "tonal_range": 0.16, + "sharpness": 0.32, + "edge_density": 0.18, +} +TEMPORAL_WEIGHTS = { + "motion": 0.34, + "jerk_ratio": 0.42, + "flicker": 0.18, + "freeze_fraction": 0.06, +} +AUDIO_WEIGHTS = { + "lufs": 0.30, + "true_peak": 0.12, + "silence_fraction": 0.18, + "spectral_centroid": 0.12, + "spectral_flatness": 0.16, + "voice_band_ratio": 0.12, +} + + +def run_bytes(command: list[str]) -> bytes: + return subprocess.run(command, check=True, stdout=subprocess.PIPE, + stderr=subprocess.PIPE).stdout + + +def video_metrics(path: Path) -> dict[str, float]: + width, height = 208, 120 + raw = run_bytes([ + "ffmpeg", "-v", "error", "-i", str(path), "-an", "-vf", + f"fps=6,scale={width}:{height}:flags=area,format=gray", "-f", + "rawvideo", "-pix_fmt", "gray", "-", + ]) + frame_size = width * height + usable = len(raw) // frame_size * frame_size + frames = np.frombuffer(raw[:usable], dtype=np.uint8).reshape(-1, height, + width).astype(np.float32) + if len(frames) < 3: + raise RuntimeError(f"Too few decoded frames in {path}") + + frame_means = frames.mean(axis=(1, 2)) + frame_stds = frames.std(axis=(1, 2)) + p05 = np.percentile(frames, 5, axis=(1, 2)) + p95 = np.percentile(frames, 95, axis=(1, 2)) + + lap = (-4.0 * frames[:, 1:-1, 1:-1] + frames[:, :-2, 1:-1] + + frames[:, 2:, 1:-1] + frames[:, 1:-1, :-2] + + frames[:, 1:-1, 2:]) + gx = np.abs(frames[:, :, 1:] - frames[:, :, :-1]) + gy = np.abs(frames[:, 1:, :] - frames[:, :-1, :]) + sharpness = np.var(lap, axis=(1, 2)) + edge_density = 0.5 * ((gx > 18).mean(axis=(1, 2)) + + (gy > 18).mean(axis=(1, 2))) + + delta = np.abs(np.diff(frames, axis=0)).mean(axis=(1, 2)) + accel = np.abs(frames[2:] - 2.0 * frames[1:-1] + frames[:-2]).mean(axis=(1, 2)) + motion = float(np.median(delta)) + jerk = float(np.median(accel)) + return { + "luma": float(np.mean(frame_means)), + "contrast": float(np.mean(frame_stds)), + "tonal_range": float(np.mean(p95 - p05)), + "sharpness": float(np.median(sharpness)), + "edge_density": float(np.mean(edge_density)), + "motion": motion, + "jerk_ratio": jerk / max(motion, 1e-6), + "flicker": float(np.std(np.diff(frame_means))), + "freeze_fraction": float(np.mean(delta < 0.55)), + "black_fraction": float(np.mean(frame_means < 5.0)), + } + + +def ebur128(path: Path) -> tuple[float, float]: + proc = subprocess.run([ + "ffmpeg", "-hide_banner", "-nostats", "-i", str(path), + "-filter_complex", "ebur128=peak=true", "-f", "null", "-", + ], stdout=subprocess.DEVNULL, stderr=subprocess.PIPE, text=True, + check=True) + summaries = proc.stderr.split("Summary:") + text = summaries[-1] if len(summaries) > 1 else proc.stderr + integrated = re.search(r"I:\s*(-?[0-9.]+)\s+LUFS", text) + peak = re.search(r"Peak:\s*(-?[0-9.]+)\s+dBFS", text) + return (float(integrated.group(1)) if integrated else math.nan, + float(peak.group(1)) if peak else math.nan) + + +def audio_metrics(path: Path) -> dict[str, float]: + sample_rate = 16_000 + raw = run_bytes([ + "ffmpeg", "-v", "error", "-i", str(path), "-vn", "-ac", "1", + "-ar", str(sample_rate), "-f", "f32le", "-", + ]) + audio = np.frombuffer(raw, dtype="= 80) & (frequencies <= 4000) + audible_mask = (frequencies >= 40) & (frequencies <= 7800) + voice_ratio = np.sum(power[:, voice_mask], axis=1) / np.maximum( + np.sum(power[:, audible_mask], axis=1), 1e-15) + return { + "lufs": lufs, + "true_peak": true_peak, + "rms_dbfs": float(20.0 * np.log10(np.sqrt(np.mean(audio * audio)) + 1e-15)), + "silence_fraction": silence_fraction, + "spectral_centroid": float(np.median(centroid)), + "spectral_flatness": float(np.median(flatness)), + "voice_band_ratio": float(np.median(voice_ratio)), + "clipped_fraction": float(np.mean(np.abs(audio) >= 0.999)), + } + + +def log_distance(value: float, reference: float, scale: float) -> float: + return abs(math.log(max(value, 1e-9) / max(reference, 1e-9))) / scale + + +def linear_distance(value: float, reference: float, scale: float) -> float: + return abs(value - reference) / scale + + +def component_scores(candidate: dict[str, float], reference: dict[str, float]) -> dict[str, float]: + visual_scales = { + "luma": 25.0, + "contrast": 15.0, + "tonal_range": 30.0, + "sharpness": 0.70, + "edge_density": 0.35, + } + temporal_scales = { + "motion": 0.70, + "jerk_ratio": 0.55, + "flicker": 4.0, + "freeze_fraction": 0.20, + } + audio_scales = { + "lufs": 6.0, + "true_peak": 6.0, + "silence_fraction": 0.25, + "spectral_centroid": 0.70, + "spectral_flatness": 0.20, + "voice_band_ratio": 0.30, + } + + def weighted_score(weights: dict[str, float], scales: dict[str, float], + log_fields: set[str]) -> float: + distance = 0.0 + for key, weight in weights.items(): + if key in log_fields: + part = log_distance(candidate[key], reference[key], scales[key]) + else: + part = linear_distance(candidate[key], reference[key], scales[key]) + distance += weight * min(part, 3.0) + return 100.0 * math.exp(-distance) + + visual = weighted_score(VIDEO_WEIGHTS, visual_scales, + {"sharpness", "edge_density"}) + temporal = weighted_score(TEMPORAL_WEIGHTS, temporal_scales, + {"motion", "jerk_ratio"}) + audio = weighted_score(AUDIO_WEIGHTS, audio_scales, + {"spectral_centroid", "spectral_flatness", + "voice_band_ratio"}) + + if candidate["black_fraction"] > 0.05: + visual *= max(0.0, 1.0 - candidate["black_fraction"]) + if candidate["clipped_fraction"] > 1e-4: + audio *= max(0.70, 1.0 - 10.0 * candidate["clipped_fraction"]) + overall = 0.40 * visual + 0.30 * temporal + 0.30 * audio + return {"visual": visual, "temporal": temporal, "audio": audio, + "overall": overall} + + +def confidence_interval(values: list[float], seed: int = 20260914) -> tuple[float, float]: + rng = np.random.default_rng(seed) + array = np.asarray(values) + draws = rng.choice(array, size=(20_000, len(array)), replace=True).mean(axis=1) + return float(np.percentile(draws, 2.5)), float(np.percentile(draws, 97.5)) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--root", type=Path, required=True) + parser.add_argument("--reference", type=Path, required=True) + parser.add_argument("--pattern", default="34block-step-*") + parser.add_argument("--output-prefix", type=Path, required=True) + args = parser.parse_args() + + candidates = sorted(args.root.glob(args.pattern), + key=lambda p: int(re.search(r"step-(\d+)", p.name).group(1))) + reference_files = {p.name: p for p in args.reference.glob("*.mp4")} + if not reference_files: + raise SystemExit(f"No MP4s in reference {args.reference}") + + cache: dict[str, dict[str, float]] = {} + + def measure(path: Path) -> dict[str, float]: + key = str(path) + if key not in cache: + cache[key] = {**video_metrics(path), **audio_metrics(path)} + return cache[key] + + reference = {name: measure(path) for name, path in reference_files.items()} + per_clip: list[dict[str, object]] = [] + summary: list[dict[str, object]] = [] + for folder in candidates: + step = int(re.search(r"step-(\d+)", folder.name).group(1)) + candidate_files = {p.name: p for p in folder.glob("*.mp4")} + common = sorted(reference_files.keys() & candidate_files.keys()) + if len(common) != len(reference_files): + continue + scores: list[dict[str, float]] = [] + for name in common: + metrics = measure(candidate_files[name]) + score = component_scores(metrics, reference[name]) + scores.append(score) + per_clip.append({"step": step, "prompt": name, **score, **metrics}) + overall_values = [s["overall"] for s in scores] + low, high = confidence_interval(overall_values, seed=20260914 + step) + summary.append({ + "step": step, + "clips": len(scores), + "overall": float(np.mean(overall_values)), + "ci95_low": low, + "ci95_high": high, + "visual": float(np.mean([s["visual"] for s in scores])), + "temporal": float(np.mean([s["temporal"] for s in scores])), + "audio": float(np.mean([s["audio"] for s in scores])), + "worst_prompt": min(per_clip[-len(scores):], key=lambda x: x["overall"])["prompt"], + "worst_score": min(overall_values), + }) + + summary.sort(key=lambda row: row["overall"], reverse=True) + args.output_prefix.parent.mkdir(parents=True, exist_ok=True) + with args.output_prefix.with_suffix(".summary.csv").open("w", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=list(summary[0])) + writer.writeheader() + writer.writerows(summary) + with args.output_prefix.with_suffix(".clips.csv").open("w", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=list(per_clip[0])) + writer.writeheader() + writer.writerows(per_clip) + payload = { + "reference": str(args.reference), + "interpretation": "Technical A/V retention relative to paired reference; not a semantic or human-preference score.", + "overall_weights": {"visual": 0.40, "temporal": 0.30, "audio": 0.30}, + "summary": summary, + } + args.output_prefix.with_suffix(".json").write_text(json.dumps(payload, indent=2) + "\n") + print(json.dumps(summary, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/scripts/compacth3/eval_corrected_dmd_inference_exports.sbatch b/scripts/compacth3/eval_corrected_dmd_inference_exports.sbatch new file mode 100644 index 0000000000..8580731075 --- /dev/null +++ b/scripts/compacth3/eval_corrected_dmd_inference_exports.sbatch @@ -0,0 +1,125 @@ +#!/bin/bash +#SBATCH --job-name=h3-dmd2-export-eval +#SBATCH --partition=all +#SBATCH --nodes=1 +#SBATCH --ntasks=1 +#SBATCH --gres=gpu:4 +#SBATCH --time=04:00:00 +#SBATCH --output=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/sbatch-h3-dmd2-export-eval-%j.log +#SBATCH --error=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/sbatch-h3-dmd2-export-eval-%j.log + +set -euo pipefail +export SLURM_EXPORT_ENV=ALL + +SPRINT_ROOT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +DMD_CODE=${SPRINT_ROOT}/code/release20b-dmd2-v17-fp32resume-v1 +RUNNER_ROOT=${SPRINT_ROOT}/code/release-execution-dmd-ladderfix-v1 +DMD_RUN=${SPRINT_ROOT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +STEPS="800 1000 1200" +export SPRINT_ROOT DMD_CODE RUNNER_ROOT DMD_RUN STEPS + +srun --export=ALL --kill-on-bad-exit=1 \ + --container-image=nvcr.io/nvidia/pytorch:25.06-py3 \ + --container-mounts=/mnt/nfs/vlm-aryan:/mnt/nfs/vlm-aryan,/mnt/lustre/vlm-shared:/mnt/lustre/vlm-shared:ro \ + --container-workdir="${RUNNER_ROOT}" bash -lc ' +set -euo pipefail +source /mnt/nfs/vlm-aryan/fasth3-33b-20260806/secrets.env +PY=/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python +PROMPTS="${SPRINT_ROOT}/job-scripts/bench_five_new_prompts.json" +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache PYTHONDONTWRITEBYTECODE=1 +export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA FASTVIDEO_MINIMAX_H3_FUSIONS=0 +export FASTVIDEO_DMD_DENOISING_STEPS=999,749,500,250 +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True TORCH_NCCL_ENABLE_MONITORING=0 +export PYTHONPATH="${RUNNER_ROOT}:${SPRINT_ROOT}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" + +for step in ${STEPS}; do + model="${DMD_RUN}/inference/checkpoint-${step}" + eval_dir="${DMD_RUN}/eval-exported-checkpoint-${step}-bench-five-exact4-v1" + media="${eval_dir}/media" + + [[ -f "${model}/.complete" && -s "${model}/metadata.json" ]] || { + echo "FATAL: incomplete inference export for step ${step}" >&2 + exit 2 + } + [[ $(find "${model}/transformer" -maxdepth 1 -name "*.safetensors" ! -name "*.index.json" | wc -l) -eq 8 ]] || { + echo "FATAL: step ${step} does not have eight transformer shards" >&2 + exit 3 + } + + "${PY}" - "${model}" <<"PY" +import collections +import json +import struct +import sys +from pathlib import Path + +root = Path(sys.argv[1]) +counts = collections.Counter() +for shard in sorted((root / "transformer").glob("*.safetensors")): + with shard.open("rb") as handle: + header_size = struct.unpack(" "${eval_dir}/completed_at.txt" + echo "step ${step}: recovered and accepted five already-rendered videos" + continue + fi + [[ ! -e "${eval_dir}" ]] || { + echo "FATAL: refusing to overwrite partial evaluation ${eval_dir}" >&2 + exit 4 + } + mkdir -p "${eval_dir}" + + cd "${RUNNER_ROOT}" + "${PY}" "${RUNNER_ROOT}/scripts/fasth3_sprint/run_baseline_matrix.py" \ + --model-path "${model}" --checkpoint-role "corrected-dmd2-step${step}-native-export-exact4" \ + --attention dense --attention-backend TORCH_SDPA --prompts "${PROMPTS}" \ + --output-dir "${media}" --run-id "corrected-dmd2-step${step}-native-export-exact4-${SLURM_JOB_ID}" \ + --source-commit "$(cat "${DMD_CODE}/CODE_COMMIT" 2>/dev/null || echo unknown)" --max-prompts 5 \ + --height 480 --width 832 --num-frames 124 --seed 20260912 \ + --steps 5 --num-gpus 4 --dit-precision bf16 --profile strict \ + --no-fa4 --no-compile --no-upload-videos + + "${PY}" - "${media}" <<"PY" +import json +import sys +from pathlib import Path + +media = Path(sys.argv[1]) +manifest = json.loads((media / "run_manifest.json").read_text()) +assert len(list(media.glob("*.mp4"))) == 5 +assert manifest["schedule"]["grid_points"] == 5, manifest["schedule"] +assert manifest["schedule"]["transformer_calls"] == 4, manifest["schedule"] +assert len(manifest["schedule"]["video"]["transformer_timesteps"]) == 4 +assert len(manifest["schedule"]["audio"]["transformer_timesteps"]) == 4 +print("verified exact four-call media", media) +PY + date -Is > "${eval_dir}/completed_at.txt" +done +' diff --git a/scripts/compacth3/eval_dmd_export_1400.sh b/scripts/compacth3/eval_dmd_export_1400.sh new file mode 100644 index 0000000000..f50be8caaa --- /dev/null +++ b/scripts/compacth3/eval_dmd_export_1400.sh @@ -0,0 +1,125 @@ +#!/bin/bash +#SBATCH --job-name=h3-dmd2-export-eval +#SBATCH --partition=all +#SBATCH --nodes=1 +#SBATCH --ntasks=1 +#SBATCH --gres=gpu:4 +#SBATCH --time=04:00:00 +#SBATCH --output=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/sbatch-h3-dmd2-export-eval-%j.log +#SBATCH --error=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/sbatch-h3-dmd2-export-eval-%j.log + +set -euo pipefail +export SLURM_EXPORT_ENV=ALL + +SPRINT_ROOT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +DMD_CODE=${SPRINT_ROOT}/code/release20b-dmd2-v17-fp32resume-v1 +RUNNER_ROOT=${SPRINT_ROOT}/code/release-execution-dmd-ladderfix-v1 +DMD_RUN=${SPRINT_ROOT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +STEPS="1400" +export SPRINT_ROOT DMD_CODE RUNNER_ROOT DMD_RUN STEPS + +srun --overlap --jobid=9639 -N1 -n1 --gres=gpu:4 --export=ALL --kill-on-bad-exit=1 \ + --container-image=nvcr.io/nvidia/pytorch:25.06-py3 \ + --container-mounts=/mnt/nfs/vlm-aryan:/mnt/nfs/vlm-aryan,/mnt/lustre/vlm-shared:/mnt/lustre/vlm-shared:ro \ + --container-workdir="${RUNNER_ROOT}" bash -lc ' +set -euo pipefail +source /mnt/nfs/vlm-aryan/fasth3-33b-20260806/secrets.env +PY=/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python +PROMPTS="${SPRINT_ROOT}/job-scripts/bench_five_new_prompts.json" +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache PYTHONDONTWRITEBYTECODE=1 +export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA FASTVIDEO_MINIMAX_H3_FUSIONS=0 +export FASTVIDEO_DMD_DENOISING_STEPS=999,749,500,250 +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True TORCH_NCCL_ENABLE_MONITORING=0 +export PYTHONPATH="${RUNNER_ROOT}:${SPRINT_ROOT}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" + +for step in ${STEPS}; do + model="${DMD_RUN}/inference/checkpoint-${step}" + eval_dir="${DMD_RUN}/eval-exported-checkpoint-${step}-bench-five-exact4-v1" + media="${eval_dir}/media" + + [[ -f "${model}/.complete" && -s "${model}/metadata.json" ]] || { + echo "FATAL: incomplete inference export for step ${step}" >&2 + exit 2 + } + [[ $(find "${model}/transformer" -maxdepth 1 -name "*.safetensors" ! -name "*.index.json" | wc -l) -eq 8 ]] || { + echo "FATAL: step ${step} does not have eight transformer shards" >&2 + exit 3 + } + + "${PY}" - "${model}" <<"PY" +import collections +import json +import struct +import sys +from pathlib import Path + +root = Path(sys.argv[1]) +counts = collections.Counter() +for shard in sorted((root / "transformer").glob("*.safetensors")): + with shard.open("rb") as handle: + header_size = struct.unpack(" "${eval_dir}/completed_at.txt" + echo "step ${step}: recovered and accepted five already-rendered videos" + continue + fi + [[ ! -e "${eval_dir}" ]] || { + echo "FATAL: refusing to overwrite partial evaluation ${eval_dir}" >&2 + exit 4 + } + mkdir -p "${eval_dir}" + + cd "${RUNNER_ROOT}" + "${PY}" "${RUNNER_ROOT}/scripts/fasth3_sprint/run_baseline_matrix.py" \ + --model-path "${model}" --checkpoint-role "corrected-dmd2-step${step}-native-export-exact4" \ + --attention dense --attention-backend TORCH_SDPA --prompts "${PROMPTS}" \ + --output-dir "${media}" --run-id "corrected-dmd2-step${step}-native-export-exact4-${SLURM_JOB_ID}" \ + --source-commit "$(cat "${DMD_CODE}/CODE_COMMIT" 2>/dev/null || echo unknown)" --max-prompts 5 \ + --height 480 --width 832 --num-frames 124 --seed 20260912 \ + --steps 5 --num-gpus 4 --dit-precision bf16 --profile strict \ + --no-fa4 --no-compile --no-upload-videos + + "${PY}" - "${media}" <<"PY" +import json +import sys +from pathlib import Path + +media = Path(sys.argv[1]) +manifest = json.loads((media / "run_manifest.json").read_text()) +assert len(list(media.glob("*.mp4"))) == 5 +assert manifest["schedule"]["grid_points"] == 5, manifest["schedule"] +assert manifest["schedule"]["transformer_calls"] == 4, manifest["schedule"] +assert len(manifest["schedule"]["video"]["transformer_timesteps"]) == 4 +assert len(manifest["schedule"]["audio"]["transformer_timesteps"]) == 4 +print("verified exact four-call media", media) +PY + date -Is > "${eval_dir}/completed_at.txt" +done +' diff --git a/scripts/compacth3/grade_dmd2_all_checkpoints.sbatch b/scripts/compacth3/grade_dmd2_all_checkpoints.sbatch new file mode 100644 index 0000000000..da3a2d4ddd --- /dev/null +++ b/scripts/compacth3/grade_dmd2_all_checkpoints.sbatch @@ -0,0 +1,48 @@ +#!/bin/bash +#SBATCH --job-name=h3-dmd2-grade-all +#SBATCH --partition=all +#SBATCH --nodes=1 +#SBATCH --ntasks=1 +#SBATCH --gres=gpu:1 +#SBATCH --cpus-per-task=16 +#SBATCH --mem=128G +#SBATCH --time=06:00:00 +#SBATCH --output=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/sbatch-h3-dmd2-grade-all-%j.log +#SBATCH --error=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/sbatch-h3-dmd2-grade-all-%j.log + +set -euo pipefail +export SLURM_EXPORT_ENV=ALL + +SPRINT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +EVAL=/mnt/nfs/vlm-aryan/fasth3-eval +RUN=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +OUT="${RUN}/dmd-native-export-grade-all-v1" + +srun --export=ALL --kill-on-bad-exit=1 \ + --container-image=nvcr.io/nvidia/pytorch:25.06-py3 \ + --container-mounts=/mnt/nfs/vlm-aryan:/mnt/nfs/vlm-aryan,/mnt/lustre/vlm-shared:/mnt/lustre/vlm-shared:ro \ + bash -lc ' +set -euo pipefail +SPRINT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +EVAL=/mnt/nfs/vlm-aryan/fasth3-eval +RUN=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +OUT="${RUN}/dmd-native-export-grade-all-v1" +mkdir -p "${OUT}" + +for step in 100 200 400 600 800 1000 1200; do + media="${RUN}/eval-exported-checkpoint-${step}-bench-five-exact4-v1/media" + [[ $(find "${media}" -maxdepth 1 -name "*.mp4" | wc -l) -eq 5 ]] || { + echo "FATAL: step ${step} does not have five videos" >&2 + exit 2 + } +done + +bash "${EVAL}/probe_deps.sh" | tee "${OUT}/dependency-probe.txt" +"${EVAL}/venv/bin/python" "${EVAL}/score_34block.py" \ + --runs "${EVAL}/dmd2_all_checkpoints_runs.json" \ + --reference-media "${RUN}/eval-parent-750-bench-five-correct4-v1/media" \ + --prompts "${SPRINT}/job-scripts/bench_five_new_prompts.json" \ + --out "${OUT}/scores.json" | tee "${OUT}/scorer.log" +"${EVAL}/venv/bin/python" "${EVAL}/summarize_scores.py" "${OUT}/scores.json" | tee "${OUT}/summary.txt" +date -Is > "${OUT}/completed_at.txt" +' diff --git a/scripts/compacth3/hardmotion_set.json b/scripts/compacth3/hardmotion_set.json new file mode 100644 index 0000000000..4eb0592e0f --- /dev/null +++ b/scripts/compacth3/hardmotion_set.json @@ -0,0 +1,100 @@ +[ + { + "id": "t2va-0020260818-008164", + "width": 768, + "height": 768, + "frames": 260, + "prompt": "integrated_multimodal_description: [Shot 1] 2D anime animation in a square 1:1 composition, a dense production-brief view of a futuristic black stage with a glowing circular floor and a vertical LED panel displaying [English]\"NIGHT PULSE\". The two principal dancers are S1, a silver-haired performer in a cobalt-blue cropped jacket and white gloves, and S2, a performer with a long magenta ponytail in a black-and-gold costume; only S1 and S2 occupy the stage, and no additional performers enter. They begin back-to-back at center, then snap into mirrored arm cuts, sprint in opposite arcs, and rebound toward one another as white floor rings illuminate beneath their feet. The camera makes a large-amplitude arc shot at fast speed around them, keeping the square frame centered on their synchronized footwork. [Shot 2] At 00:02.400, S1 vaults over S2's lowered shoulder while S2 spins into a one-handed floor sweep, their jackets and ponytail trailing through the rapid exchange. The camera cuts to a low three-quarter angle and tilts up with large amplitude at fast speed as S1 lands on the front edge of the stage and S2 springs vertically behind them, with the LED title remaining visible above their silhouettes. [Shot 3] At 00:05.200, S1 and S2 cross the stage in opposing diagonals, perform three rapid turns, and meet at center for a synchronized airborne kick before landing into a crouched back-to-back pose. The camera cuts to a high overhead view, tracking their crossing paths at fast speed while the circular floor lights chase each landing and the square composition alternates between the blue and magenta costume colors. [Shot 4] At 00:08.100, S1 and S2 explode out of the crouch into a final high-energy duet: they trade mirrored spins, slide past each other, leap in parallel, and descend into a wide forward-facing stance with one arm extended toward the viewer. The camera cuts to a frontal medium close shot and pushes in with small amplitude at fast speed, settling as both dancers hold the sharp final pose beneath the legible [English]\"NIGHT PULSE\" panel until the event ends. overall_soundscape: Sparse stage ambience and a low electrical hum continue beneath the performance. Shoe impacts, brief floor-slide friction, costume fabric snaps, and the resonant thuds of the landings provide the only prominent action sounds. non_diegetic_music: N/A" + }, + { + "id": "t2va-0020260818-000687", + "width": 1344, + "height": 768, + "frames": 294, + "prompt": "integrated_multimodal_description: [Shot 1] 3D CG game-cinematic rendering in a 16:9 widescreen composition frames an original music-dance stage performance on a circular black platform surrounded by cyan light strips. Four principal dancers stand in a tight diamond: the copper-jacketed lead dancer (S1), the teal-haired dancer (S2), the silver-sleeved dancer (S3), and the orange-accented dancer (S4). The stage floor is blank and unlettered, with no readable signs, captions, logos, labels, or subtitles. The four dancers launch into rapid synchronized footwork, shoulder hits, and a low crouching sweep while keeping their hands empty; S1 shouts in a clear on-screen voice, [English] Three, two, one\u2014break! The camera begins a fast push in with medium amplitude as their shoes strike the platform together. [Shot 2] At 00:02.000, the camera cuts to a low side angle as S2 vaults into a twisting aerial jump, S3 slides beneath the movement, and S1 and S4 pivot around them without breaking formation. S2 lands and calls, [Spanish] \u00a1Gira conmigo! The camera tracks right at fast speed, matching the dancers as they sweep across the stage. [Shot 3] At 00:04.000, the camera cuts to a close three-quarter view as S1 and S4 perform opposing backflips while S2 and S3 execute rapid heel-toe steps and snap into a mirrored pose. S3 shouts, [English] Cross left! The camera arcs around the group with large amplitude at fast speed, revealing the cyan floor lights streaking beneath their feet. [Shot 4] At 00:06.000, the camera cuts to a high overhead view as all four dancers rotate into a moving diamond, trade positions through two fast spins, and rebound from a simultaneous knee drop into a vertical leap. S4 calls, [Spanish] \u00a1Ahora, salta! The camera rolls slightly clockwise while descending in a controlled arc toward the center of the formation. [Shot 5] At 00:08.000, the camera cuts to a tight frontal angle as S1 and S2 exchange a rapid mirrored footwork sequence, then S3 and S4 cross behind them in a synchronized sliding pass. S1 snaps, [English] Switch! The camera tracks backward with large amplitude at fast speed, keeping all four dancers visible as they accelerate into the final combination. All four performers keep empty hands throughout, and no additional performers enter the stage. [Shot 6] At 00:10.000, the camera cuts to a wide frontal composition as the four dancers sprint toward the platform center, perform a coordinated four-person jump with separated spins, and land in a staggered diamond pose. S2 shouts, [Spanish] \u00a1Juntos! The camera pulls out with medium amplitude at slow speed while S1 raises one arm, S2 drops into a lunge, S3 kneels, and S4 holds a sharp standing angle; their final foot impacts and breathing settle by the end of the 12.00-second event. overall_soundscape: A cavernous indoor stage ambience surrounds the performance with steady ventilation and a faint electrical hum from the floor lights. Hard shoe impacts, sliding soles, fabric snaps, and brief handclaps synchronize with the choreography. Rapid breathing, exertion grunts, and the resonance of each landing continue beneath the spoken calls. non_diegetic_music: N/A" + }, + { + "id": "t2va-0020260818-007107", + "width": 768, + "height": 1344, + "frames": 209, + "prompt": "integrated_multimodal_description: [Shot 1] Motion-graphics design in a 9:16 portrait composition, the 9.00-second event opens on a luminous vertical dance stage built from cobalt-blue planes, magenta arcs, and a glowing orange floor grid. Four principal adult dancers stand in a diamond formation as crisp vector-like figures: S1 is tall and cyan, S2 wears magenta, S3 wears yellow, and S4 wears violet. They drop low, sweep their arms outward, pivot sharply, and spring into a synchronized jump. The low-voiced front dancer (S1) calls: [Chinese] \u4e09\u3001\u4e8c\u3001\u4e00\uff0c\u8df3\uff01 The camera pushes in with small amplitude at fast speed toward the formation, then holds a steady view; no camera shake or additional camera movement occurs beyond this single push-in. Bright geometric blocks rotate around the dancers in time with their footwork, and no readable text, captions, logos, or labels appear anywhere. [Shot 2] At 00:03.000, the camera cuts to a centered medium view as the four dancers cross past one another in a fast counter-rotating exchange, snapping their elbows inward and extending their legs into alternating kicks while the orange floor grid folds into concentric rings. The clear-voiced female dancer (S2) shouts: [Chinese] \u4fdd\u6301\u961f\u5f62\uff01 The quick-voiced male dancer (S3) immediately answers: [Chinese] \u6362\u4f4d\uff01 S1 and S4 mirror the exchange from opposite sides, keeping the four-color arrangement visually distinct as diagonal light planes sweep behind them. [Shot 3] At 00:06.500, the camera cuts to a full-body vertical view as S1, S2, S3, and S4 accelerate through a final high-energy sequence of spins, knee lifts, and synchronized overhead arm strikes. The geometric planes contract toward the center while the dancers leap inward, land in a tight four-point diamond, and hold a sharp final pose facing the camera. The four dancers (S1,S2,S3,S4) shout together: [Chinese] \u4e00\u8d77\uff01 Their pose and the surrounding blue, magenta, orange, and violet shapes remain fixed through the end of the 9.00-second timeline. overall_soundscape: Sharp shoe taps, sliding footfalls, hand claps, and brief fabric snaps synchronize with the rapid choreography on the hard stage surface. Short breaths and exertion sounds from the four dancers sit beneath compact indoor stage reverberation, while the geometric transitions produce clean electronic swishes. non_diegetic_music: N/A" + }, + { + "id": "t2va-0020260818-003059", + "width": 1344, + "height": 768, + "frames": 260, + "prompt": "integrated_multimodal_description: [Shot 1] 2D anime action fantasy rendered as kinetic motion graphics in a 16:9 widescreen frame, with crisp cel-shaded figures, layered vector-like energy planes, hard-edged color blocks, and luminous geometric effects. In a suspended observatory above a storm of violet clouds, two rival sky-mages face each other on a circular transparent platform: the young woman in a white-and-cobalt coat (S1) grips a silver staff, while the young man in a black-and-crimson mantle (S2) crouches opposite her with a glowing amber gauntlet. S1 sweeps her staff across the floor, launching three blue energy crescents; S2 braces and punches forward, shattering the crescents into angular fragments with a concussive flash. The camera executes a fast high-angle arc around both figures, revealing the platform rotating above the clouds as their loose fabric and hair whip in the wind. S1 shouts in a clear, urgent young woman's voice: [Chinese] \u8ba9\u5f00\uff01\u661f\u95e8\u8981\u584c\u4e86\uff01 [Shot 2] At 00:03.000, the camera cuts to a low close-up as S2 drives his gauntlet into the platform, sending a crackling amber shockwave outward before the transparent floor buckles into rising polygonal slabs. S1 leaps from slab to slab, twists over the expanding fracture, and slashes downward with her staff, drawing a vertical blue beam that pins the collapsing geometry in place. S2 looks up through the beam and answers in a tense young man's voice: [Chinese] \u6211\u4e0d\u4f1a\u628a\u94a5\u5319\u4ea4\u7ed9\u4f60\uff01 The camera whip-pans from S2's fist to S1's airborne silhouette, with graphic streaks compressing the motion between viewpoints. [Shot 3] At 00:05.800, the camera cuts to an overhead view as the two mages are pulled toward the observatory's rotating central aperture; S1 plants her staff into a seam while S2 catches a floating metal ring and swings around it, then releases himself toward her. Their weapons collide at the aperture, producing a rapidly spinning lattice of blue and amber planes that bends the cloud vortex into a funnel. S1 reaches through the lattice and grabs S2's wrist as the platform fragments accelerate around them. The camera plunges vertically through the rotating energy lattice, then rolls clockwise to keep both bodies centered while the physical debris whips past. [Shot 4] At 00:08.700, the camera cuts to a tight side view as S1 and S2 pull in opposite directions, their combined strain forcing the aperture shut in a single explosive fold. The remaining slabs fly upward, the cloud vortex collapses into calm violet layers, and both characters land back-to-back on the restored circular platform. S2 lowers his gauntlet and S1 releases her staff; they exchange a wary glance while small blue and amber sparks orbit their hands, then the camera pulls out with large amplitude at fast speed to reveal the stabilized observatory and the two figures silhouetted against the clearing sky through the 11.00-second endpoint. overall_soundscape: Violent wind rushes around the suspended observatory throughout, with fabric snaps, hair flutter, and layered cloud turbulence beneath the action. Staff sweeps, gauntlet impacts, crystalline fractures, flying metal rings, and rapid energy discharges produce sharply synchronized physical sounds. Heavy landings thud against the platform, followed by falling fragments skittering and the final sparks crackling softly. non_diegetic_music: N/A" + }, + { + "id": "t2va-0020260818-004321", + "width": 1024, + "height": 768, + "frames": 141, + "prompt": "integrated_multimodal_description: [Shot 1] Live-action photorealism in a 4:3 landscape composition frames a solo contemporary dancer centered on an amber-lit theater stage, backed by black curtains and narrow beams of white light. The dancer, a young woman in a cobalt-blue cropped jacket, silver sash, and black trousers, grips a closed red folding fan; she is the only principal performer and uses the stable speaker ID (S1). She snaps her head toward the audience and says sharply: [Japanese]\u59cb\u3081\u307e\u3059\uff01 Then she launches into high-speed percussive footwork, sweeps the fan across her body, and spins twice while the camera tracks backward with large amplitude at fast speed, arcing to keep her centered as she advances diagonally toward stage right. [Shot 2] At 00:03.100, the camera cuts to a low three-quarter angle and continues a fast tracking movement alongside her as she drives into a running leap, rotates in midair, lands with one knee bent, and snaps the red fan open above her head. She rises through a final sharp turn, thrusts the fan toward the lens, and says breathlessly: [Japanese]\u3053\u308c\u3067\u7d42\u308f\u308a\u3067\u3059\uff01 The camera arcs around her with large amplitude at fast speed, settling on her extended-arm finishing pose as the stage lights tighten around her at the end of the six-second performance. overall_soundscape: Rapid barefoot footfalls strike the wooden stage in tightly spaced rhythms, accompanied by fabric swishes, the metallic snap of the folding fan, and heavy landing thumps. Her breathing grows audible during the leap while the theater's low ventilation hum continues beneath the performance. non_diegetic_music: N/A" + }, + { + "id": "t2va-0020260818-003311", + "width": 768, + "height": 1344, + "frames": 141, + "prompt": "integrated_multimodal_description: [Shot 1] Three-dimensional CG game-cinematic rendering in a 9:16 vertical composition frames three principal characters on an abandoned cliffside tram platform during blue-hour rain: the teenage courier in a yellow weatherproof jacket (S1) kneels beside a dark emergency beacon, the gray-haired station keeper in a patched green coat (S2) grips a hand crank, and the younger sibling in a red hooded coat (S3) stands close to the shelter, watching anxiously. S1 fits a loose power cell into the beacon while S2 turns the crank and S3 reaches toward S1's shoulder; their movements are moderately paced and physically grounded. The camera pushes in with small amplitude at slow speed toward their hands and the unlit beacon. The teenage courier (S1) says in an off-screen voiceover: [English] If this light holds, they will find us before the storm closes the pass. while her lips remain completely closed. No additional character enters, no readable text, caption, label, logo, or symbol appears, and no weapon, explosion, or moving tram is present. [Shot 2] At 00:03.000, the camera cuts to a close vertical view of the beacon and the three faces behind it. S1 presses the contact lever, the beacon ignites with a steady amber glow, S2 releases the crank and exhales, and S3 takes S1's hand as rain beads on their clothing. S1 says in an off-screen voiceover: [English] Stay together. while her lips remain completely closed. The three characters hold their positions beside the glowing beacon through the end of the event; no on-screen character speaks, no new person appears, and no additional light source or readable marking is introduced. overall_soundscape: Rain strikes the corrugated shelter and the empty metal platform while strong wind moves the characters' coats. The power cell clicks into place, the hand crank squeaks and turns, boots scrape on wet metal, and the beacon relay snaps on. Quiet breathing and a relieved exhale are audible beneath the weather, with no diegetic music. non_diegetic_music: N/A" + }, + { + "id": "vidprom-000214-000225", + "width": 1344, + "height": 768, + "frames": 362, + "prompt": "integrated_multimodal_description: [Shot 1] Realistic, aesthetically refined 4K studio performance video, opening in a wide full-body composition of one male Kazakh dancer centered on a stage. He wears a colorful traditional shapan with intricate embroidered designs and a matching embroidered hat; the fabric\u2019s vibrant colors and decorative patterns remain clearly visible. Behind him, a scenic backdrop depicts the Kazakh steppe with rolling hills beneath a vast blue sky. Soft natural-style studio lighting evenly illuminates the dancer, stage, and landscape backdrop. The camera begins a smooth, slow forward tracking move as he glides across the stage with graceful traditional Kazakh dance gestures, extending and shaping his arms while his embroidered sleeves sway. He gathers momentum into a dynamic leap, landing lightly with a soft footfall, then continues into fluid spins; the shapan flares outward in a colorful circle as the camera tracks laterally with him at a steady pace. In the final seconds, the camera gently arcs closer around his movement while he completes the last controlled spin and settles into a poised traditional dance stance, facing forward against the open steppe landscape. overall_soundscape: Quiet studio ambience surrounds the performance. Soft rhythmic footfalls and subtle fabric swishes accompany the dancer\u2019s leaps, gliding steps, and spins. A faint natural spaciousness complements the steppe imagery. non_diegetic_music: N/A" + }, + { + "id": "t2va-0020260818-009754", + "width": 1344, + "height": 768, + "frames": 141, + "prompt": "integrated_multimodal_description: [Shot 1] Live-action photoreal anime-action-fantasy in a 16:9 widescreen frame, a locked-off wide shot holds on a rain-slick basalt bridge above a glowing ravine for the complete six-second take. Two principal subjects face each other: a young sky duelist in dark segmented armor gripping a blue-lit short sword, and a large obsidian wyvern crouched opposite with folded wings and ember-bright eyes. The wyvern suddenly lunges and sweeps one wing across the bridge as the duelist sprints forward, slides beneath the claws, and vaults onto a broken stone pillar. The wyvern whips its tail through the pillar, scattering wet rock, while the duelist launches from the collapsing stone and twists through the air. The blade strikes the wyvern's raised horn, producing a brief burst of blue sparks and orange embers; the wyvern recoils with a powerful wingbeat as the duelist lands in a low defensive stance at the edge of the bridge. The camera remains completely fixed throughout, no cut occurs, no additional characters enter the frame, and no readable text appears. overall_soundscape: Cold wind rushes through the ravine while rain patters on armor and stone. Heavy wingbeats, splashing boots, claw scrapes, cracking masonry, and scattered rock impacts build rapidly, followed by a sharp metallic clash and sizzling magical sparks. non_diegetic_music: N/A" + }, + { + "id": "t2va-2026082050-021443", + "width": 1760, + "height": 768, + "frames": 362, + "prompt": "integrated_multimodal_description: [Shot 1] Live-action photoreal, cinematic ultrawide game-cutscene trailer composition: an unmanned ocean-spanning orbital elevator terminal buckles above a black storm sea as its colossal tether begins to shear, rain and spray whipping through steel braces. The camera starts in a low 24mm wide angle and makes a forceful forward push toward the cracking terminal. The narrator (S1) says in an off-screen voiceover [English]When the skyway breaks, gravity remembers every piece. [Shot 2] At 00:01.700, the severed tether lashes across the storm clouds, flinging glowing metal fragments through lightning. The camera cuts to a high aerial viewpoint and performs a rapid clockwise roll following the tether's arc. The narrator (S1) says in an off-screen voiceover [English]The first fall tears open the weather. [Shot 3] At 00:03.400, a massive counterweight rips free, accelerates downward, and disintegrates into a storm of incandescent shards. The camera dives beside it in a steep, accelerating descent, then snaps into a compressed telephoto view as the fragments spread. The narrator (S1) says in an off-screen voiceover [English]Then the heavy heart lets go. [Shot 4] At 00:05.200, the impact wave reaches the ocean and raises a circular wall of water around the empty terminal pylons. The camera skims just above the wave crest in a fast lateral tracking move, revealing the terminal dwarfed by the rising water. The narrator (S1) says in an off-screen voiceover [English]The sea answers with a crown of glass. [Shot 5] At 00:07.000, an automated launch gantry folds under extreme wind, its rails twisting and snapping while an unoccupied shuttle cradle spins away. The camera spirals around the collapsing machinery at close range, changing from a wide view to a sharp upward angle. The narrator (S1) says in an off-screen voiceover [English]Nothing built for stillness survives the spin. [Shot 6] At 00:08.800, molten fragments punch through the cloud deck and ignite long orange trails across the night sky above the flooded terminal. The camera cuts to a distant horizon view and executes a rapid pan that catches several impacts in one continuous sweep. The narrator (S1) says in an off-screen voiceover [English]Every spark redraws the horizon. [Shot 7] At 00:10.600, the terminal's surviving magnetic core surges, pulling floating steel plates into a violently orbiting ring while lightning forks through the structure. The camera rotates in a full sweeping arc around the unstable ring, counter-moving against the debris to emphasize the physics-sensitive motion. The narrator (S1) says in an off-screen voiceover [English]But the core is still awake. [Shot 8] At 00:12.400, the orbiting debris locks into a luminous geometric halo above the receding waves, then collapses inward in a blinding blue-white flare before the 15-second endpoint. The camera rapidly pulls back through rain into an extreme ultrawide final view of the ruined terminal and storm horizon. The narrator (S1) says in an off-screen voiceover [English]The end is only the next beginning. No people, animals, or humanoid figures appear, and no readable text, logos, captions, or subtitles appear anywhere. overall_soundscape: A violent ocean storm surrounds the scene with pounding rain, deep thunder, gale-force wind, roaring surf, straining cables, tearing steel, explosive impacts, and surging electrical arcs. Water crashes against the terminal while fragments whistle through the air and the collapsing machinery produces heavy metallic groans and sharp snaps. non_diegetic_music: N/A" + }, + { + "id": "vidprom-000215-000236", + "width": 1344, + "height": 768, + "frames": 362, + "prompt": "integrated_multimodal_description: [Shot 1] Cinematic fantasy realism, a wide aerial composition follows a majestic emerald-green dragon soaring above a dense misty forest at sunset. Its large leathery wings beat in slow, powerful cycles, and its long sinuous tail trails gracefully behind its regal body. The camera glides forward and slightly downward alongside the dragon as it moves effortlessly between the upper canopy and drifting pale mist. Warm patches of sunlight filter through the tall trees below while the sky behind it shifts from vibrant orange and pink near the horizon to deep purple and blue overhead. The dragon\u2019s piercing amber eyes catch the fading light as it banks gently, wingbeats creating visible ripples in the mist and soft gusts through the treetops. [Shot 2] At 00:08.500, a low-angle view from beneath the forest canopy looks upward as the same dragon passes overhead in a grand, unhurried ascent. The camera tilts up smoothly, framing its broad scaled chest, outstretched leathery wings, and sweeping tail against the layered sunset sky. Its wings give several resonant, powerful flaps, stirring loose leaves and mist below, before it glides onward above the tall trees and recedes into the deepening purple-blue distance. overall_soundscape: Continuous high-altitude wind moves through the forest canopy and around the dragon\u2019s wings. Each heavy wingbeat produces a deep leathery whoosh, accompanied by soft gusts that rustle leaves and shift the mist. Distant forest ambience remains subdued beneath the passing flight. non_diegetic_music: N/A" + }, + { + "id": "t2va-2026082050-007339", + "width": 768, + "height": 768, + "frames": 141, + "prompt": "integrated_multimodal_description: [Shot 1] A square 1:1 live-action photoreal studio tabletop becomes a text-free motion-graphics title infographic, lit by soft neutral panels with crisp practical shadows and realistic acrylic reflections. Three clearly distinct adult designers, the principal subjects, are visible from the chest down around a matte white work surface, each wearing a different solid-color glove: cobalt blue, warm amber, and deep green. The blue-gloved designer places a translucent blue ring at center, the amber-gloved designer aligns three blank white bars into a clean radial layout, and the green-gloved designer rotates a small amber disc into the open gap, creating deliberate moderate coordinated motion. After these actions, the only camera movement is one smooth, controlled 10-centimeter overhead push-in at a slow, even pace. No extra hands, faces, logos, labels, readable text, captions, or subtitles appear; every graphic panel remains completely unmarked, and the live-action materials never become cartoon-like or digitally rendered. [Shot 2] At 00:03.800, the three designers withdraw their hands in sequence, then simultaneously tap the completed abstract arrangement once so the rings, bars, and disc settle into a precise balanced emblem for the final title-infographic frame. The camera holds the closer overhead composition without further movement, showing photoreal surface texture, tiny edge highlights, and stable graphic alignment through the 00:06.000 endpoint. No new subject enters, no object jumps position, and there is no camera shake or distracting reflection. overall_soundscape: Close ASMR foley captures the soft nitrile-glove friction, faint sleeve rustle, quiet controlled breathing, translucent acrylic sliding on matte board, delicate bar-to-table taps, and the final synchronized fingertip clicks. The room remains nearly silent and dry, with no speech, singing, voice, or other human vocalization. non_diegetic_music: N/A" + }, + { + "id": "t2va-2026082050-020251", + "width": 1024, + "height": 768, + "frames": 192, + "prompt": "integrated_multimodal_description: [Shot 1] Live-action photoreal anime action-fantasy in 4:3 landscape: the sole principal subject, a storm warrior (S1), stands on a rain-soaked basalt bridge in a midnight ravine, rendered with physically detailed wet stone, rain, sparks, and violent magical forces. The warrior (S1) plants one boot, drives a crackling silver blade into the bridge, and shouts [English]Hold the storm!; the basalt splits under the impact and a vertical vortex lifts shards around the warrior. The camera tracks backward at knee height with a wide lens, revealing the bridge buckling toward a black chasm as the warrior remains centered. [Shot 2] At 00:02.400, the warrior (S1) sprints across three disintegrating bridge slabs, leaps through a widening gap, twists in midair, and calls [English]I can still make it! while lightning coils around the blade. The camera cuts to a high overhead view and dives in a steep spiral around the airborne warrior, making the falling stone and streaming rain whip past the frame. [Shot 3] At 00:05.400, the warrior (S1) lands hard on the far cliff, skids through water, raises the blade overhead, and cries [English]Now! as the vortex collapses into a blinding ring of blue lightning behind the warrior. The camera cuts to a tight three-quarter view, then makes a rapid semicircular orbit as the warrior slashes through the ring and the ravine erupts with reflected light until the eight-second event ends. overall_soundscape: Cold mountain rain hisses across basalt, thunder rolls through the ravine, and wind roars around the collapsing bridge. Each footfall, stone fracture, blade crackle, airborne debris impact, hard landing, sliding scrape, forceful breath, and exertion cry is sharply synchronized to the action. non_diegetic_music: N/A" + }, + { + "id": "0004005-5b852dc5", + "width": 1344, + "height": 768, + "frames": 124, + "prompt": "integrated_multimodal_description: [Shot 1] Live-action, cinematic. A wide shot frames a bustling city plaza at sunset, with tall glass skyscrapers rising against an orange-pink sky and vibrant streetlights and neon signs glowing along the avenue. In the center stands PIKA, a 3D CG Pikachu with bright yellow fur, large pointy black-tipped ears, round red cheeks, and a cute open smile, hopping in place and waving its small arms. Beside it stands RUNWAY, a tall slender runway model with slicked-back dark hair, wearing a modern chic outfit: a structured metallic-silver blazer over a fitted black jumpsuit with angular high heels. She holds a poised fashion stance with one hand on her hip, then tilts her head down toward Pikachu and smiles as pedestrians in evening wear stroll past in the background. The camera pushes in at slow speed toward the pair, keeping both centered as the city lights shimmer behind them. overall_soundscape: City traffic hums in the distance with occasional car horns, layered with the footsteps and low murmur of pedestrians passing by. Pikachu's hops land with soft thumps on the pavement, and the model's heels click sharply as she shifts her stance. non_diegetic_music: A fast-tempo electronic pop track with bright synth arpeggios, a punchy bassline, and a steady four-on-the-floor drum beat. The arrangement builds in the final seconds with layered percussion and a rising synth line." + }, + { + "id": "vidprom-000214-000970", + "width": 1344, + "height": 768, + "frames": 362, + "prompt": "integrated_multimodal_description: [Shot 1] Cinematic naturalistic adventure footage opens in a wide, slightly low-angle tracking composition of Arjun and Vikram, two close Indian friends in rugged waterproof jackets, hiking boots, and loaded backpacks, trekking side by side along a narrow path through dense green forest. Arjun holds a folded map and compass in front of him while Vikram\u2019s camera hangs ready around his neck. The handheld camera moves backward smoothly as they advance with determined expressions, their boots pressing through damp leaves and uneven soil. [Shot 2] At 00:05.000, a wide side view reveals the same two friends crossing a rocky hillside above the forest, with layered mountains and lush valleys opening behind them. The camera pans slowly left with their movement as Arjun checks the compass and Vikram steadies himself over loose stones, then briefly lifts his camera toward the scenery. [Shot 3] At 00:10.000, a closer front-facing shot follows them as they reach a natural overlook and slow to a stop shoulder to shoulder. Arjun lowers the map, Vikram lets the camera rest against his chest, and both look out across the breathtaking mountain landscape with quiet satisfaction, their jackets stirring lightly in the mountain breeze. The camera eases back into a stable medium-wide frame, holding them together against the expansive wilderness through the end. overall_soundscape: Soft wind moves through forest leaves and across the rocky overlook. Boots crunch on gravel, brush against damp vegetation, and small stones shift underfoot as the friends hike. Backpack straps creak faintly, with occasional quiet breathing from the exertion. non_diegetic_music: N/A" + } +] \ No newline at end of file diff --git a/scripts/compacth3/qad/gen_qad_configs.py b/scripts/compacth3/qad/gen_qad_configs.py new file mode 100644 index 0000000000..849709d6d7 --- /dev/null +++ b/scripts/compacth3/qad/gen_qad_configs.py @@ -0,0 +1,100 @@ +#!/usr/bin/env python3 +"""Generate NVFP4 QAD configs for rank-768 and rank-16 students.""" +import json, pathlib, sys, yaml + +S = pathlib.Path("/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829") +M = pathlib.Path("/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1") +RUN = S / "runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3" +TEACHER = S / "release-candidates/base-h3-teacher-complete-v1" + +meta = json.loads((RUN / "checkpoint-1400/metadata.json").read_text()) +c = meta["config"] +def y(p, d=None): + cur = c + for k in p.split("."): + if not isinstance(cur, dict) or k not in cur: return d + cur = cur[k] + return cur + +VARIANTS = { + "r768": { + "adaln_rank": 768, + "init_from": str(RUN / "inference/checkpoint-1400"), + "note": "20B baseline lineage; identical architecture to the shipped rank-768 model.", + }, + "r16": { + "adaln_rank": 16, + "init_from": str(RUN / "inference/checkpoint-1400-reparam-r16"), + "note": ("~17B. Student is the post-hoc centered-affine rank-16 reparameterization of " + "checkpoint-1400, materialized as a rank-16 checkpoint. Requires that " + "materialization to exist -- see the rank-compression run."), + }, +} + +HEADER = """# NVFP4 QAD -- {tag} (adaln_rank={rank}) +""" + +def build(tag, v): + q = { + "models": { + "student": { + "_target_": y("models.student._target_"), + "init_from": v["init_from"], + "trainable": True, + "enable_gradient_checkpointing_type": "full", + "attention_backend": y("models.student.attention_backend", "TORCH_SDPA"), + "quant_config": "nvfp4_qat_train", + "adaln_rank": v["adaln_rank"], + }, + "teacher": { + "_target_": y("models.teacher._target_"), + "init_from": y("models.teacher.init_from", str(TEACHER)), + "trainable": False, + "disable_custom_init_weights": True, + "attention_backend": y("models.teacher.attention_backend", "TORCH_SDPA"), + }, + }, + "method": {k: y(f"method.{k}") for k in ( + "_target_", "rollout_mode", "rollout_carry", "rollout_carry_slots", + "rollout_sample_type", "generator_update_interval", "real_score_guidance_scale", + "dmd_denoising_steps", "min_timestep_ratio", "max_timestep_ratio", + "score_timestep_shift", "score_timestep_warp_max", "score_timestep_continuous", + "fake_score_loss_space", "modality_loss_weights", "dmd_denom_floor_ratio", + "dmd_grad_cap", "cfg_uncond", "fake_score_learning_rate", "fake_score_betas", + "fake_score_lr_scheduler")}, + "training": { + "distributed": y("training.distributed"), + "data": y("training.data"), + "optimizer": y("training.optimizer"), + "loop": {"max_train_steps": 200, + "gradient_accumulation_steps": y("training.loop.gradient_accumulation_steps", 8)}, + "checkpoint": { + "output_dir": str(S / f"runs/release20b-qad-nvfp4-4call-{tag}-v1"), + "resume_from_checkpoint": v["init_from"], + "training_state_checkpointing_steps": 25, + "require_complete_training_checkpoint": True, + "checkpoints_total_limit": 12, + }, + "tracker": {"trackers": ["wandb"], "project_name": "fasth3-14b-2step-qad-sprint", + "run_name": f"release20b-qad-nvfp4-4call-{tag}-v1"}, + }, + "callbacks": { + "grad_clip": {"_target_": "fastvideo.train.callbacks.grad_clip.GradNormClipCallback", + "max_grad_norm": 1.0}, + "validation": y("callbacks.validation"), + }, + "model": {"precondition_outputs": False, "enable_gradient_checkpointing_type": "full", + "enable_torch_compile": False}, + "dit_precision": "fp32", + "vsa": y("vsa"), + } + out = M / f"examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call_{tag}.yaml" + out.write_text(HEADER.format(tag=tag, rank=v["adaln_rank"], init=v["init_from"], note=v["note"]) + + yaml.safe_dump(q, sort_keys=False)) + return out, q + +for tag, v in VARIANTS.items(): + out, q = build(tag, v) + print(f"{tag}: {out.name}") + print(f" adaln_rank={q['models']['student']['adaln_rank']} init={q['models']['student']['init_from'][-40:]}") + print(f" quant={q['models']['student']['quant_config']} steps={q['training']['loop']['max_train_steps']}") diff --git a/scripts/compacth3/qad/qad_checkpoint_gate.py b/scripts/compacth3/qad/qad_checkpoint_gate.py new file mode 100644 index 0000000000..98edfc9313 --- /dev/null +++ b/scripts/compacth3/qad/qad_checkpoint_gate.py @@ -0,0 +1,145 @@ +#!/usr/bin/env python3 +"""Select the earliest acceptable QAD checkpoint.""" +import argparse, glob, json, os, re, subprocess, sys + +S = "/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829" +MOTION_TOL = 0.15 # |d_motion| within 15% of baseline counts as preserved + + +def ckpt_step(path): + m = re.search(r"checkpoint-(\d+)", os.path.basename(path)) + return int(m.group(1)) if m else None + + +def metrics_for(video_dir): + """Compute the four decision metrics for one model's generations. + + NOTE: motion magnitude is the mandatory anti-collapse guard. Implementations here + are deliberately simple and must be swapped for the sweep agent's versions once + those land, so both sides of every delta use the SAME code path. + """ + out = {} + for name, fn in (("motion", _motion_magnitude), + ("grain", _hf_energy), + ("detail", _sharpness), + ("anomalies", _colour_drift)): + vals = [] + for v in sorted(glob.glob(os.path.join(video_dir, "*.mp4"))): + try: + vals.append(fn(v)) + except Exception as e: + print(f" WARN {name} failed on {os.path.basename(v)}: {type(e).__name__}", flush=True) + out[name] = sum(vals) / len(vals) if vals else None + return out + + +def _frames(v, n=32): + import av, numpy as np + c = av.open(v) + fr = [f.to_ndarray(format="rgb24") for f in c.decode(video=0)] + c.close() + if len(fr) > n: + idx = np.linspace(0, len(fr) - 1, n).astype(int) + fr = [fr[i] for i in idx] + return fr + + +def _motion_magnitude(v): + """Mean optical-flow magnitude. The guard against a model that 'wins' by going static.""" + import cv2, numpy as np + fr = _frames(v) + mags = [] + for a, b in zip(fr, fr[1:]): + ga = cv2.cvtColor(a, cv2.COLOR_RGB2GRAY) + gb = cv2.cvtColor(b, cv2.COLOR_RGB2GRAY) + fl = cv2.calcOpticalFlowFarneback(ga, gb, None, 0.5, 3, 15, 3, 5, 1.2, 0) + mags.append(float(np.sqrt(fl[..., 0] ** 2 + fl[..., 1] ** 2).mean())) + return sum(mags) / len(mags) if mags else 0.0 + + +def _hf_energy(v): + """High-frequency energy: proxy for grain. Higher = grainier.""" + import cv2, numpy as np + vals = [] + for f in _frames(v, 16): + g = cv2.cvtColor(f, cv2.COLOR_RGB2GRAY).astype("float32") + vals.append(float(cv2.Laplacian(g, cv2.CV_32F).var())) + return sum(vals) / len(vals) if vals else 0.0 + + +def _sharpness(v): + """Detail proxy. Deliberately the same family as grain; report both so a + grain reduction that is really just blurring is visible.""" + import cv2, numpy as np + vals = [] + for f in _frames(v, 16): + g = cv2.cvtColor(f, cv2.COLOR_RGB2GRAY).astype("float32") + vals.append(float(cv2.Sobel(g, cv2.CV_32F, 1, 0).var())) + return sum(vals) / len(vals) if vals else 0.0 + + +def _colour_drift(v): + """Flow-warped colour residual between adjacent frames: targets the reported + 'gloves change colour' failure. Motion-compensated, so legitimate motion + does not count as drift.""" + import cv2, numpy as np + fr = _frames(v) + res = [] + for a, b in zip(fr, fr[1:]): + ga = cv2.cvtColor(a, cv2.COLOR_RGB2GRAY) + gb = cv2.cvtColor(b, cv2.COLOR_RGB2GRAY) + fl = cv2.calcOpticalFlowFarneback(ga, gb, None, 0.5, 3, 15, 3, 5, 1.2, 0) + h, w = ga.shape + xx, yy = np.meshgrid(np.arange(w), np.arange(h)) + wx = (xx + fl[..., 0]).astype(np.float32) + wy = (yy + fl[..., 1]).astype(np.float32) + warped = cv2.remap(a, wx, wy, cv2.INTER_LINEAR, borderMode=cv2.BORDER_REPLICATE) + la = cv2.cvtColor(warped, cv2.COLOR_RGB2LAB).astype("float32") + lb = cv2.cvtColor(b, cv2.COLOR_RGB2LAB).astype("float32") + res.append(float(np.abs(la - lb).mean())) + return sum(res) / len(res) if res else 0.0 + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--baseline-dir", required=True, help="pre-QAD NVFP4 generations, same prompts+seeds") + ap.add_argument("--ckpt-root", required=True, help="QAD output_dir containing checkpoint-*") + ap.add_argument("--gen-dir-template", required=True, + help="where per-checkpoint generations live; must contain {step}") + ap.add_argument("--out", default=f"{S}/adaln_rank_analysis/qad_checkpoint_gate.json") + a = ap.parse_args() + + base = metrics_for(a.baseline_dir) + print(f"BASELINE (pre-QAD NVFP4): {json.dumps(base, indent=2)}", flush=True) + + steps = sorted(s for s in (ckpt_step(p) for p in glob.glob(f"{a.ckpt_root}/checkpoint-*")) if s) + rows, winner = [], None + for s in steps: + d = a.gen_dir_template.format(step=s) + if not os.path.isdir(d): + print(f"checkpoint-{s}: no generations at {d} -- SKIPPED (must be generated)", flush=True) + continue + m = metrics_for(d) + dl = {k: (m[k] - base[k]) for k in base if m.get(k) is not None and base.get(k)} + rel = {k: (v / base[k]) for k, v in dl.items() if base.get(k)} + ok = (rel.get("grain", 1) < 0 and rel.get("detail", -1) >= 0 + and rel.get("anomalies", 1) <= 0 and abs(rel.get("motion", 1)) <= MOTION_TOL) + rows.append({"step": s, "abs": m, "delta": dl, "rel": rel, "satisfies": ok}) + print(f"checkpoint-{s}: d_grain={rel.get('grain',0):+.3f} d_detail={rel.get('detail',0):+.3f} " + f"d_anom={rel.get('anomalies',0):+.3f} d_motion={rel.get('motion',0):+.3f} " + f"{'<== SATISFIES' if ok else ''}", flush=True) + if ok and winner is None: + winner = s + + verdict = ("EARLIEST satisfying checkpoint: %d" % winner) if winner is not None else \ + ("NONE satisfies: standard QAD is trading away the NVFP4 stability benefit " + "-- a custom temporal objective is now justified") + print("\nVERDICT:", verdict, flush=True) + os.makedirs(os.path.dirname(a.out), exist_ok=True) + json.dump({"baseline": base, "rows": rows, "winner": winner, "verdict": verdict}, + open(a.out, "w"), indent=2) + print("wrote", a.out, flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/compacth3/qad/setup_qad.py b/scripts/compacth3/qad/setup_qad.py new file mode 100644 index 0000000000..4a39f8d4cd --- /dev/null +++ b/scripts/compacth3/qad/setup_qad.py @@ -0,0 +1,114 @@ +#!/usr/bin/env python3 +"""Generate the NVFP4 QAD yaml from checkpoint-1400 metadata.""" +import json +import pathlib +import sys + +M = pathlib.Path("/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1") +SPRINT = pathlib.Path("/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829") +RUN = SPRINT / "runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3" + +qc = M / "fastvideo/layers/quantization/nvfp4_qat_config.py" +t = qc.read_text(); orig = t +if '"ff.fc_in"' in t: + print("layer list: already fixed") +else: + anchor = 'DEFAULT_FP4_LAYERS = (\n' + assert t.count(anchor) == 1, f"anchor={t.count(anchor)}" + t = t.replace(anchor, anchor + ' # MiniMax-H3 names its FFN "ff.", not "ffn." -- without these the\n' + ' # FFN is silently left dense and QAD only covers attention.\n' + ' "ff.fc_in",\n "ff.fc_out",\n', 1) + assert t != orig + qc.with_suffix(".py.pre-qad-bak").write_text(orig) + qc.write_text(t) + print("layer list: added ff.fc_in / ff.fc_out for H3") + +meta = json.loads((RUN / "checkpoint-1400/metadata.json").read_text()) +c = meta["config"] + +def y(path, default=None): + cur = c + for k in path.split("."): + if not isinstance(cur, dict) or k not in cur: + return default + cur = cur[k] + return cur + +qad = { + "models": { + "student": { + "_target_": y("models.student._target_"), + "init_from": str(RUN / "checkpoint-1400"), + "trainable": True, + "enable_gradient_checkpointing_type": "full", + "attention_backend": y("models.student.attention_backend", "TORCH_SDPA"), + "quant_config": "nvfp4_qat_train", + }, + "teacher": { + "_target_": y("models.teacher._target_"), + "init_from": y("models.teacher.init_from"), + "trainable": False, + "disable_custom_init_weights": True, + "attention_backend": y("models.teacher.attention_backend", "TORCH_SDPA"), + }, + }, + "method": { + "_target_": y("method._target_"), + "rollout_mode": y("method.rollout_mode"), + "rollout_carry": y("method.rollout_carry"), + "rollout_carry_slots": y("method.rollout_carry_slots"), + "rollout_sample_type": y("method.rollout_sample_type"), + "generator_update_interval": y("method.generator_update_interval"), + "real_score_guidance_scale": y("method.real_score_guidance_scale"), + "dmd_denoising_steps": y("method.dmd_denoising_steps"), + "min_timestep_ratio": y("method.min_timestep_ratio"), + "max_timestep_ratio": y("method.max_timestep_ratio"), + "score_timestep_shift": y("method.score_timestep_shift"), + "score_timestep_warp_max": y("method.score_timestep_warp_max"), + "score_timestep_continuous": y("method.score_timestep_continuous"), + "fake_score_loss_space": y("method.fake_score_loss_space"), + "modality_loss_weights": y("method.modality_loss_weights"), + "dmd_denom_floor_ratio": y("method.dmd_denom_floor_ratio"), + "dmd_grad_cap": y("method.dmd_grad_cap"), + "cfg_uncond": y("method.cfg_uncond"), + "fake_score_learning_rate": y("method.fake_score_learning_rate"), + "fake_score_betas": y("method.fake_score_betas"), + "fake_score_lr_scheduler": y("method.fake_score_lr_scheduler"), + }, + "training": { + "distributed": y("training.distributed"), + "data": y("training.data"), + "optimizer": y("training.optimizer"), + "loop": {"max_train_steps": 200, "gradient_accumulation_steps": y("training.loop.gradient_accumulation_steps", 8)}, + "checkpoint": { + "output_dir": str(SPRINT / "runs/release20b-dmd2-v12-qad-nvfp4-4call-v1"), + "resume_from_checkpoint": str(RUN / "checkpoint-1400"), + "training_state_checkpointing_steps": 25, + "require_complete_training_checkpoint": True, + "checkpointing_start_step": 1400, + "checkpoints_total_limit": 12, + }, + "tracker": {"trackers": ["wandb"], "project_name": "fasth3-14b-2step-qad-sprint", + "run_name": "release20b-dmd2-qad-nvfp4-4call-v1"}, + }, + "callbacks": { + "grad_clip": {"_target_": "fastvideo.train.callbacks.grad_clip.GradNormClipCallback", "max_grad_norm": 1.0}, + "validation": y("callbacks.validation"), + }, + "model": {"precondition_outputs": False, "enable_gradient_checkpointing_type": "full", "enable_torch_compile": False}, + "dit_precision": "fp32", + "vsa": y("vsa"), +} + +out = M / "examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call.yaml" +out.parent.mkdir(parents=True, exist_ok=True) + +import yaml +header = """# NVFP4 QAD for the 42-block four-call DMD2 student (checkpoint-1400). +""" +out.write_text(header + yaml.safe_dump(qad, sort_keys=False)) +print("wrote", out) +print("student init_from:", qad["models"]["student"]["init_from"]) +print("quant_config :", qad["models"]["student"]["quant_config"]) +print("denoise steps :", qad["method"]["dmd_denoising_steps"]) +print("output_dir :", qad["training"]["checkpoint"]["output_dir"]) diff --git a/scripts/compacth3/qad/tune_qad_cadence.py b/scripts/compacth3/qad/tune_qad_cadence.py new file mode 100644 index 0000000000..8c9d5dc82e --- /dev/null +++ b/scripts/compacth3/qad/tune_qad_cadence.py @@ -0,0 +1,29 @@ +#!/usr/bin/env python3 +"""Make the QAD configs emit frequent short checkpoints so we can find the +EARLIEST checkpoint that reduces grain without raising temporal anomalies. + +Decision rule (lead): select the earliest checkpoint that meaningfully reduces +grain without increasing temporal anomalies. Do NOT optimise until QAD +reproduces BF16 exactly -- BF16 has more temporal failures on some cases. +""" +import pathlib, yaml +M = pathlib.Path("/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1") +base = M / "examples/train/configs/distribution_matching/minimax_h3" + +for tag in ("r768", "r16"): + p = base / f"qad_nvfp4_4call_{tag}.yaml" + if not p.exists(): + print(f"MISSING {p.name}"); continue + d = yaml.safe_load(p.read_text()) + ck = d["training"]["checkpoint"] + ck["training_state_checkpointing_steps"] = 25 # frequent -> finer sweet-spot search + ck["checkpoints_total_limit"] = 24 # keep them all + d["training"]["loop"]["max_train_steps"] = 200 + v = d.get("callbacks", {}).get("validation") + if isinstance(v, dict): + v["every_steps"] = 50 + v["run_at_start"] = False + d["training"]["tracker"]["run_name"] = f"qad-nvfp4-4call-{tag}-v2-shortckpt" + p.write_text(yaml.safe_dump(d, sort_keys=False)) + print(f"{p.name}: ckpt every {ck['training_state_checkpointing_steps']} x{ck['checkpoints_total_limit']}, " + f"max {d['training']['loop']['max_train_steps']} steps, val every {v.get('every_steps')}") diff --git a/scripts/compacth3/quantization/export_lane_int8.sh b/scripts/compacth3/quantization/export_lane_int8.sh new file mode 100644 index 0000000000..ad2ef7a93e --- /dev/null +++ b/scripts/compacth3/quantization/export_lane_int8.sh @@ -0,0 +1,25 @@ +#!/bin/bash +# Run: sbatch /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/exp.sbatch /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/export_lane_int8.sh +set -euo pipefail +SPRINT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +M=/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1 +PY=/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python +RUN=${SPRINT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +CKPT=${RUN}/inference/checkpoint-1400 +OUT=${CKPT}/exports/int8 + +mkdir -p "$OUT" +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache +export PYTHONDONTWRITEBYTECODE=1 +export PYTHONPATH="${M}:${SPRINT}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True + +cd "$M" +echo "=== exporting int8 -> $OUT ===" +"$PY" "${SPRINT}/export_quant_dit.py" \ + --lane int8 \ + --model-path "$CKPT" \ + --out "$OUT" \ + --num-gpus 1 +echo "=== done: $OUT ===" +ls -la "$OUT" diff --git a/scripts/compacth3/quantization/export_lane_nvfp4.sh b/scripts/compacth3/quantization/export_lane_nvfp4.sh new file mode 100644 index 0000000000..b4f989160d --- /dev/null +++ b/scripts/compacth3/quantization/export_lane_nvfp4.sh @@ -0,0 +1,25 @@ +#!/bin/bash +# Run: sbatch /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/exp.sbatch /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/export_lane_nvfp4.sh +set -euo pipefail +SPRINT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +M=/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1 +PY=/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python +RUN=${SPRINT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +CKPT=${RUN}/inference/checkpoint-1400 +OUT=${CKPT}/exports/nvfp4 + +mkdir -p "$OUT" +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache +export PYTHONDONTWRITEBYTECODE=1 +export PYTHONPATH="${M}:${SPRINT}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True + +cd "$M" +echo "=== exporting nvfp4 -> $OUT ===" +"$PY" "${SPRINT}/export_quant_dit.py" \ + --lane nvfp4 \ + --model-path "$CKPT" \ + --out "$OUT" \ + --num-gpus 1 +echo "=== done: $OUT ===" +ls -la "$OUT" diff --git a/scripts/compacth3/quantization/export_lane_w4a16.sh b/scripts/compacth3/quantization/export_lane_w4a16.sh new file mode 100644 index 0000000000..c92ac2b8a7 --- /dev/null +++ b/scripts/compacth3/quantization/export_lane_w4a16.sh @@ -0,0 +1,25 @@ +#!/bin/bash +# Run: sbatch /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/exp.sbatch /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/export_lane_w4a16.sh +set -euo pipefail +SPRINT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +M=/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1 +PY=/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python +RUN=${SPRINT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +CKPT=${RUN}/inference/checkpoint-1400 +OUT=${CKPT}/exports/w4a16 + +mkdir -p "$OUT" +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache +export PYTHONDONTWRITEBYTECODE=1 +export PYTHONPATH="${M}:${SPRINT}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True + +cd "$M" +echo "=== exporting w4a16 -> $OUT ===" +"$PY" "${SPRINT}/export_quant_dit.py" \ + --lane w4a16 \ + --model-path "$CKPT" \ + --out "$OUT" \ + --num-gpus 1 +echo "=== done: $OUT ===" +ls -la "$OUT" diff --git a/scripts/compacth3/quantization/export_quant_dit.py b/scripts/compacth3/quantization/export_quant_dit.py new file mode 100644 index 0000000000..5b7503863b --- /dev/null +++ b/scripts/compacth3/quantization/export_quant_dit.py @@ -0,0 +1,179 @@ +"""Export pre-quantized DiT weights to safetensors.""" +from __future__ import annotations + +import argparse +import importlib.util +import json +import os +import sys +import time +from pathlib import Path + +SPRINT = "/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829" +M = "/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1" +HARNESS = f"{M}/examples/inference/basic/basic_fasth3.py" + +LANES = { + "nvfp4": "NVFP4H3", + "int8": "INT8Affine", + "w4a16": "W4A16", +} + +BUFFERS = { + "nvfp4": ("_nvfp4_weight", "_nvfp4_weight_scale", "_weight_global_sf", "_nvfp4_alpha"), + "int8": ("_int8_affine_codes", "_int8_affine_scales", "_int8_affine_biases"), + "w4a16": ("_w4a16_codes", "_w4a16_scales", "_w4a16_zeros"), +} + + +def log(msg: str) -> None: + print(f"[export] {msg}", flush=True) + + +def build_fastvideo_args(model_path: str, quant_name: str, num_gpus: int): + """Harness args -> api GeneratorConfig -> FastVideoArgs (supported path).""" + spec = importlib.util.spec_from_file_location("fasth3_harness", HARNESS) + harness = importlib.util.module_from_spec(spec) + sys.modules["fasth3_harness"] = harness + spec.loader.exec_module(harness) + + argv = [ + "--model-path", model_path, + "--prompt", "quant-export", + "--num-gpus", str(num_gpus), + "--transformer-quant", quant_name, + "--no-fa4", + "--no-inference-torch-compile", + "--steps", "5", + ] + args = harness.parse_args(argv) + args.fa4 = False + harness.configure_environment(args) + + config = harness.build_generator_config(args) + log(f"GeneratorConfig built: model_path={config.model_path} " + f"num_gpus={config.engine.num_gpus} " + f"transformer_quant={config.engine.quantization.transformer_quant}") + + from fastvideo.api.compat import generator_config_to_fastvideo_args + fastvideo_args = generator_config_to_fastvideo_args(config) + log(f"FastVideoArgs built: inference_mode={fastvideo_args.inference_mode} " + f"training_mode={fastvideo_args.training_mode} " + f"use_fsdp_inference={fastvideo_args.use_fsdp_inference} " + f"hsdp_shard_dim={fastvideo_args.hsdp_shard_dim}") + return fastvideo_args + + +def load_dit(fastvideo_args, transformer_path: str): + from fastvideo.models.loader.component_loader import PipelineComponentLoader + + dit_config = fastvideo_args.pipeline_config.dit_config + quant_config = getattr(dit_config, "quant_config", None) + if quant_config is None: + raise RuntimeError( + "dit_config.quant_config is None after the compat adapter ran -- the " + "quant config was not pinned, so no linear will be built quantized.") + log(f"dit_config.quant_config = {type(quant_config).__name__} (name={quant_config.get_name()})") + + log(f"loading transformer from {transformer_path}") + model = PipelineComponentLoader.load_module( + module_name="transformer", + component_model_path=transformer_path, + transformers_or_diffusers="diffusers", + fastvideo_args=fastvideo_args, + ) + log(f"loaded class={type(model).__name__}") + return model + + +def scheme_tagged(model, lane: str) -> list[tuple[str, object]]: + """(fqn, module) pairs whose quant_method belongs to this lane's scheme.""" + from fastvideo.layers.quantization.int8_affine_config import INT8AffineQuantizeMethod + from fastvideo.layers.quantization.nvfp4_config import NVFP4QuantizeMethod + from fastvideo.layers.quantization.w4a16_config import W4A16QuantizeMethod + + wanted = { + "nvfp4": NVFP4QuantizeMethod, + "int8": INT8AffineQuantizeMethod, + "w4a16": W4A16QuantizeMethod, + }[lane] + return [(fqn, mod) for fqn, mod in model.named_modules() + if isinstance(getattr(mod, "quant_method", None), wanted)] + + +def save_int8_or_w4a16(model, lane: str, path: str, tagged) -> dict: + """No serializer exists for this lane in this checkout.""" + raise RuntimeError( + f"lane {lane!r}: no sidecar serializer exists in this checkout. " + f"grepped 'save_int8_affine_checkpoint' / 'save_w4a16_checkpoint' across the " + f"whole tree at {M}: zero hits. " + f"{len(tagged)} layers ARE quantized in memory (conversion receipt above is real) " + f"and their buffers are {BUFFERS[lane]}, but there is no encoder -- and no " + f"matching decoder -- so nothing could load the result.") + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--lane", required=True, choices=sorted(LANES)) + parser.add_argument("--model-path", required=True, + help="checkpoint dir, e.g. .../inference/checkpoint-1400") + parser.add_argument("--out", required=True, help="output directory for the sidecar") + parser.add_argument("--num-gpus", type=int, default=1, + help="1 keeps the load whole-model on cuda:0 (no FSDP)") + args = parser.parse_args() + + lane = args.lane + quant_name = LANES[lane] + transformer_path = os.path.join(args.model_path, "transformer") + if not os.path.isdir(transformer_path): + raise SystemExit(f"no transformer/ subdir under {args.model_path}") + + out_dir = Path(args.out) + out_dir.mkdir(parents=True, exist_ok=True) + + log(f"lane={lane} quant_name={quant_name}") + log(f"model_path={args.model_path}") + log(f"out_dir={out_dir}") + + t0 = time.time() + fastvideo_args = build_fastvideo_args(args.model_path, quant_name, args.num_gpus) + model = load_dit(fastvideo_args, transformer_path) + tagged = scheme_tagged(model, lane) + log(f"CONVERSION CHECK: {len(tagged)} {quant_name}-tagged linear layers " + f"present after load (conversion receipt above, from _maybe_quantize_model)") + if not tagged: + raise RuntimeError( + f"no {quant_name}-tagged linears found: the quant config did not cover " + f"any layer path, so nothing was quantized. Refusing to write an empty sidecar.") + + total_params = sum(p.numel() for p in model.parameters()) + log(f"model parameters: {total_params / 1e9:.2f}B load+convert took {time.time() - t0:.1f}s") + + if lane == "nvfp4": + from fastvideo.layers.quantization.nvfp4_config import save_nvfp4_checkpoint + target = out_dir / "nvfp4_weights.safetensors" + receipt = save_nvfp4_checkpoint( + model, target, + extra_metadata={"model": "FastH3-20B-42block-DMD2", "source": args.model_path}) + else: + save_int8_or_w4a16(model, lane, str(out_dir), tagged) + + size_bytes = os.path.getsize(receipt["path"]) + log("=" * 72) + log(f"EXPORT RECEIPT (lane={lane})") + log(json.dumps(receipt, indent=2)) + log(f"FILE PATH : {receipt['path']}") + log(f"FILE SIZE : {size_bytes / 1e9:.3f} GB ({size_bytes / (1 << 30):.3f} GiB)") + log(f"MODULE COUNT: {len(tagged)}") + log("=" * 72) + + with open(os.path.join(out_dir, "export_receipt.json"), "w") as handle: + json.dump({"lane": lane, "quant_name": quant_name, "file": receipt["path"], + "size_bytes": size_bytes, "module_count": len(tagged), + "receipt": receipt}, handle, indent=2) + log(f"wrote {os.path.join(out_dir, 'export_receipt.json')}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/scripts/compacth3/quantization/run_export_nvfp4.sh b/scripts/compacth3/quantization/run_export_nvfp4.sh new file mode 100644 index 0000000000..95e59d5f78 --- /dev/null +++ b/scripts/compacth3/quantization/run_export_nvfp4.sh @@ -0,0 +1,22 @@ +#!/bin/bash +set -euo pipefail +S=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +M=/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1 +RUN=${S}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +OUT=${RUN}/exports/nvfp4-ckpt1400 +mkdir -p "$OUT" +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache PYTHONDONTWRITEBYTECODE=1 +export PYTHONPATH="${M}:${S}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" +export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA FASTVIDEO_MINIMAX_H3_FUSIONS=0 +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True +export FASTVIDEO_EXPORT_QUANT_SIDECAR="$OUT" +export FASTVIDEO_EXPORT_LANE=nvfp4 +cd "$M" +echo "=== NVFP4 export -> $OUT ===" +/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python examples/inference/basic/basic_fasth3.py \ + --model-path "$RUN/inference/checkpoint-1400" \ + --prompt "export" --output "$OUT/tmp" \ + --height 480 --width 832 --num-frames 124 --steps 5 --num-gpus 4 \ + --repeats 1 --transformer-quant NVFP4H3 --no-fa4 2>&1 | grep -aE "EXPORT|Converting loaded|purge receipt|Error|Traceback" | tail -20 +echo "=== output dir ===" +ls -la "$OUT" diff --git a/scripts/compacth3/resume_release20b_dmd2_paired_generic.sh b/scripts/compacth3/resume_release20b_dmd2_paired_generic.sh new file mode 100644 index 0000000000..39a5ab8efd --- /dev/null +++ b/scripts/compacth3/resume_release20b_dmd2_paired_generic.sh @@ -0,0 +1,67 @@ +#!/bin/bash +# Usage: bash resume_release20b_dmd2_paired_generic.sh JOB_A JOB_B RESUME_STEP +set -euo pipefail + +if [[ "$#" -ne 3 ]]; then + echo "Usage: $0 JOB_A JOB_B RESUME_STEP" >&2 + exit 2 +fi + +JOB_A="$1" +JOB_B="$2" +RESUME_STEP="$3" +CHECKPOINT_EVERY="${CHECKPOINT_EVERY:-200}" +VALIDATION_EVERY="${VALIDATION_EVERY:-200}" +SPRINT_ROOT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +CODE_ROOT=${SPRINT_ROOT}/code/release20b-dmd2-v12-corrected-v17 +SELECTED_PARENT=${SPRINT_ROOT}/runs/release20b-folded-long-4k-v3/dmd-parent-step750-complete-v1 +TEACHER_PARENT=${SPRINT_ROOT}/release-candidates/base-h3-teacher-complete-v1 +OUTPUT_BASE=${SPRINT_ROOT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3 +RUN_ID=paired8972-8975-4000-v3 +OUTPUT_ROOT=${OUTPUT_BASE}/job-${RUN_ID} +RESUME_PATH=${OUTPUT_ROOT}/checkpoint-${RESUME_STEP} + +[[ "$(squeue -h -j "${JOB_A}" -o %T)" == "RUNNING" ]] +[[ "$(squeue -h -j "${JOB_B}" -o %T)" == "RUNNING" ]] +NODES_A="$(squeue -h -j "${JOB_A}" -o %N)" +NODES_B="$(squeue -h -j "${JOB_B}" -o %N)" +[[ "$(scontrol show hostnames "${NODES_A}" | wc -l)" -eq 4 ]] +[[ "$(scontrol show hostnames "${NODES_B}" | wc -l)" -eq 4 ]] + +test -s "${RESUME_PATH}/.complete" +test -s "${RESUME_PATH}/dcp/.metadata" +test -s "${RESUME_PATH}/metadata.json" +[[ "$(find "${RESUME_PATH}/dcp" -maxdepth 1 -type f -name '*.distcp' | wc -l)" -eq 32 ]] +[[ "$(find "${RESUME_PATH}" -maxdepth 1 -type f -name 'rng_state_rank*.pt' | wc -l)" -eq 32 ]] + +MASTER_ADDR="$(scontrol show hostnames "${NODES_A}" | head -n 1)" +MASTER_PORT="$((30000 + JOB_A % 10000))" + +launch_half() { + local job_id="$1" node_list="$2" rank_base="$3" + local log_path=${SPRINT_ROOT}/resume-dmd2-${RUN_ID}-job${job_id}-from${RESUME_STEP}.log + nohup srun --overlap --jobid="${job_id}" --nodes=4 --ntasks=4 \ + --ntasks-per-node=1 --gres=gpu:4 --nodelist="${node_list}" \ + --kill-on-bad-exit=1 \ + --container-image='nvcr.io#nvidia/pytorch:25.06-py3' \ + --container-mounts=/mnt/nfs:/mnt/nfs,/mnt/lustre/vlm-shared:/mnt/lustre/vlm-shared:ro \ + env CODE_ROOT="${CODE_ROOT}" SELECTED_PARENT="${SELECTED_PARENT}" \ + TEACHER_PARENT="${TEACHER_PARENT}" OUTPUT_BASE="${OUTPUT_BASE}" \ + SPRINT_ROOT="${SPRINT_ROOT}" MASTER_ADDR="${MASTER_ADDR}" \ + MASTER_PORT="${MASTER_PORT}" NNODES=8 NODE_RANK_BASE="${rank_base}" \ + RUN_ID="${RUN_ID}" PRODUCTION_TARGET=4000 RESUME_PATH="${RESUME_PATH}" \ + CHECKPOINT_EVERY="${CHECKPOINT_EVERY}" VALIDATION_EVERY="${VALIDATION_EVERY}" \ + bash "${CODE_ROOT}/scripts/run_release20b_dmd2_v12_16gpu.sh" \ + >"${log_path}" 2>&1 & + echo "$!" > "${log_path}.pid" +} + +launch_half "${JOB_A}" "${NODES_A}" 0 +launch_half "${JOB_B}" "${NODES_B}" 4 + +printf 'job_a=%s\njob_b=%s\nresume_step=%s\ncheckpoint_every=%s\nvalidation_every=%s\nmaster=%s:%s\nlaunched_at=%s\n' \ + "${JOB_A}" "${JOB_B}" "${RESUME_STEP}" "${CHECKPOINT_EVERY}" \ + "${VALIDATION_EVERY}" "${MASTER_ADDR}" "${MASTER_PORT}" "$(date -u +%FT%TZ)" \ + > "${OUTPUT_ROOT}/paired-rollover-from-${RESUME_STEP}.receipt" + +echo "Launched one 32-GPU world from checkpoint ${RESUME_STEP} across ${JOB_A}+${JOB_B}." diff --git a/scripts/compacth3/run_eval_lane.sh b/scripts/compacth3/run_eval_lane.sh new file mode 100644 index 0000000000..e236dbf369 --- /dev/null +++ b/scripts/compacth3/run_eval_lane.sh @@ -0,0 +1,84 @@ +#!/bin/bash +# usage: run_eval_lane.sh