h3-blackwell-runtime/research/cute_nvfp4_ring/patches/0009-qkv-runtime-gate.patch
2026-08-25 20:30:22 +07:00

96 lines
3.4 KiB
Diff

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()