diff --git a/tools/validate_cute_qkv_block.py b/tools/validate_cute_qkv_block.py new file mode 100644 index 0000000..84811c1 --- /dev/null +++ b/tools/validate_cute_qkv_block.py @@ -0,0 +1,157 @@ +"""Alternate baseline and CuTe-QKV execution inside one loaded H3 block.""" + +from __future__ import annotations + +import argparse +import json +import os +import time +from pathlib import Path + +import torch + +from h3_blackwell_runtime.adaln import H3CurveAdaLN +from h3_blackwell_runtime.block import H3DiTBlock +from h3_blackwell_runtime.checkpoint import H3Checkpoint +from h3_blackwell_runtime.packing import H3PromptPacker +from h3_blackwell_runtime.rope import h3_rope_rotation +from h3_blackwell_runtime.sampler import _audio_sigma, _model_sigma, beta_sigmas +from h3_blackwell_runtime.t2v import random_av_latents + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--model-path", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors") + parser.add_argument("--block-index", type=int, required=True) + parser.add_argument("--width", type=int, default=1344) + parser.add_argument("--height", type=int, default=768) + parser.add_argument("--frames", type=int, default=124) + parser.add_argument("--steps", type=int, default=12) + parser.add_argument("--sampler-step", type=int, default=1) + parser.add_argument("--seed", type=int, default=440420) + parser.add_argument("--text-tokens", type=int, default=100) + parser.add_argument("--attention", default="sage2") + parser.add_argument( + "--feature", choices=("qkv_ring", "modulate_fusion", "swiglu_fusion"), default="qkv_ring", + ) + parser.add_argument("--warmup", type=int, default=2) + parser.add_argument("--iterations", type=int, default=10) + parser.add_argument("--device", default="cuda") + return parser.parse_args() + + +def sync() -> None: + torch.cuda.synchronize() + + +def summarize(values: list[float]) -> dict[str, float]: + ordered = sorted(values) + middle = len(ordered) // 2 + median = ( + ordered[middle] + if len(ordered) % 2 + else (ordered[middle - 1] + ordered[middle]) / 2 + ) + return { + "mean_s": sum(values) / len(values), + "p50_s": median, + "min_s": ordered[0], + "max_s": ordered[-1], + } + + +def main() -> None: + args = parse_args() + torch.manual_seed(args.seed) + checkpoint = H3Checkpoint(args.model_path, device=args.device) + block = H3DiTBlock.from_checkpoint( + checkpoint, args.block_index, attention_backend=args.attention, + ).eval() + adaln = H3CurveAdaLN.from_checkpoint( + checkpoint, f"blocks.{args.block_index}.adaln_proj", + ).eval() + packer = H3PromptPacker(checkpoint) + video, audio, _ = random_av_latents( + args.width, args.height, args.frames, args.seed, device=args.device, + ) + sigmas = beta_sigmas(args.steps, device=args.device) + sigma = sigmas[args.sampler_step - 1] + native_audio = audio.to(torch.bfloat16) * (_audio_sigma(sigma) / sigma) + text = torch.randn( + 1, args.text_tokens, 5376, device=args.device, dtype=torch.bfloat16, + ) + hidden, timesteps, segments, positions, _, _ = packer( + text, video, native_audio, _model_sigma(sigma), + ) + rotation = h3_rope_rotation( + positions.to(args.device), + checkpoint.tensor("rope.inv_freq", dtype=torch.float32), + hidden.dtype, + ) + adaln_values = tuple(value.detach() for value in adaln(timesteps)) + + def run(enabled: bool): + if args.feature == "modulate_fusion": + block.fused_nvfp4_modulation = enabled + elif args.feature == "swiglu_fusion": + block.mlp.fused_nvfp4_swiglu = enabled + else: + if enabled: + os.environ["H3_CUTE_QKV_RING"] = "1" + else: + os.environ.pop("H3_CUTE_QKV_RING", None) + return block(hidden, rotation, *adaln_values, segments) + + with torch.inference_mode(): + reference = run(False) + candidate = run(True) + sync() + delta = candidate.float() - reference.float() + for _ in range(args.warmup): + run(False) + run(True) + sync() + baseline_times = [] + candidate_times = [] + last_reference = reference + last_candidate = candidate + for _ in range(args.iterations): + sync() + started = time.perf_counter() + last_reference = run(False) + sync() + baseline_times.append(time.perf_counter() - started) + + sync() + started = time.perf_counter() + last_candidate = run(True) + sync() + candidate_times.append(time.perf_counter() - started) + + baseline = summarize(baseline_times) + candidate_timing = summarize(candidate_times) + report = { + "device": torch.cuda.get_device_name(), + "block_index": args.block_index, + "feature": args.feature, + "hidden_shape": list(hidden.shape), + "iterations": args.iterations, + "equal": torch.equal(reference, candidate), + "max_abs": delta.abs().max().item(), + "mean_abs": delta.abs().mean().item(), + "reference_checksum": last_reference.float().sum().item(), + "candidate_checksum": last_candidate.float().sum().item(), + "baseline": baseline, + "candidate": candidate_timing, + "p50_improvement_percent": ( + 1.0 - candidate_timing["p50_s"] / baseline["p50_s"] + ) * 100.0, + } + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8") + print(json.dumps(report, indent=2), flush=True) + + +if __name__ == "__main__": + main()