diff --git a/tools/validate_cute_qkv_runtime.py b/tools/validate_cute_qkv_runtime.py new file mode 100644 index 0000000..b6b2e74 --- /dev/null +++ b/tools/validate_cute_qkv_runtime.py @@ -0,0 +1,90 @@ +"""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()