"""Validate the opt-in Nvfp4Linear CuTe QKV runtime dispatch.""" from __future__ import annotations import argparse import json import os from pathlib import Path import torch from profile_nvfp4_linear import representative_inputs def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--output", type=Path, required=True) parser.add_argument("--rows", type=int, default=2048) parser.add_argument("--warmup", type=int, default=3) parser.add_argument("--iterations", type=int, default=10) parser.add_argument("--model-path", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors") parser.add_argument("--block-index", type=int, default=24) 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("--device", default="cuda") return parser.parse_args() def measure(fn, warmup: int, iterations: int) -> float: for _ in range(warmup): fn() torch.cuda.synchronize() start = torch.cuda.Event(enable_timing=True) end = torch.cuda.Event(enable_timing=True) start.record() for _ in range(iterations): fn() end.record() end.synchronize() return start.elapsed_time(end) / iterations def main() -> None: args = parse_args() block, inputs, metadata = representative_inputs(args) linear = block.attention.qkv_proj x = inputs["attn_qkv_proj"][: args.rows].contiguous() if linear.role != "h3_attn_qkv": raise RuntimeError(f"Expected h3_attn_qkv role, got {linear.role!r}") os.environ.pop("H3_CUTE_QKV_RING", None) with torch.inference_mode(): reference = linear(x) os.environ["H3_CUTE_QKV_RING"] = "1" with torch.inference_mode(): candidate = linear(x) torch.cuda.synchronize() delta = candidate.float() - reference.float() with torch.inference_mode(): ring_ms = measure(lambda: linear(x), args.warmup, args.iterations) os.environ.pop("H3_CUTE_QKV_RING", None) reference_ms = measure(lambda: linear(x), args.warmup, args.iterations) report = { "device": torch.cuda.get_device_name(), "metadata": metadata, "block_index": args.block_index, "rows": args.rows, "role": linear.role, "equal": torch.equal(candidate, reference), "max_abs": delta.abs().max().item(), "mean_abs": delta.abs().mean().item(), "ring_ms": ring_ms, "reference_ms": reference_ms, "improvement_percent": (1.0 - ring_ms / reference_ms) * 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()