91 lines
3.2 KiB
Python
91 lines
3.2 KiB
Python
|
|
"""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()
|