h3-blackwell-runtime/tools/profile_nvfp4_pack.py
2026-08-15 02:56:19 +07:00

157 lines
6.4 KiB
Python

"""Profile NVFP4 activation pack kernels without GEMM timing noise."""
from __future__ import annotations
import argparse
import json
import os
import time
import warnings
from pathlib import Path
warnings.filterwarnings("ignore", message="Found GPU0 NVIDIA GB10 which is of cuda capability 12.1.*", category=UserWarning)
import torch
import comfy_kitchen as ck
from comfy_kitchen.tensor import TensorCoreNVFP4Layout
from h3_blackwell_runtime.nvfp4_quant import _vortex_scale_extension, nvfp4_activation_scale
from tools.profile_nvfp4_linear import module_for_name, representative_inputs
def sync() -> None:
if torch.cuda.is_available():
torch.cuda.synchronize()
def summarize(values: list[float]) -> dict[str, float]:
ordered = sorted(values)
def percentile(percent: float) -> float:
if len(ordered) == 1:
return ordered[0]
rank = (len(ordered) - 1) * percent
low = int(rank)
high = min(low + 1, len(ordered) - 1)
weight = rank - low
return ordered[low] * (1.0 - weight) + ordered[high] * weight
return {
"count": len(values),
"mean_s": sum(values) / len(values),
"p50_s": percentile(0.50),
"p90_s": percentile(0.90),
"p95_s": percentile(0.95),
"p99_s": percentile(0.99),
"min_s": ordered[0],
"max_s": ordered[-1],
}
def timed(fn, iterations: int) -> dict[str, float]:
values = []
for _ in range(iterations):
sync()
started = time.perf_counter()
fn()
sync()
values.append(time.perf_counter() - started)
return summarize(values)
def bytes_touched(tensor: torch.Tensor, qdata: torch.Tensor, block_scale: torch.Tensor) -> int:
return tensor.numel() * tensor.element_size() + qdata.numel() * qdata.element_size() + block_scale.numel()
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--model-path", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
parser.add_argument("--output", type=Path, default=Path("/output/h3-blackwell-runtime/benchmarks/nvfp4-pack-profile.json"))
parser.add_argument("--width", type=int, default=960)
parser.add_argument("--height", type=int, default=544)
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=440407)
parser.add_argument("--text-tokens", type=int, default=93)
parser.add_argument("--block-index", type=int, default=24)
parser.add_argument("--attention", default="sage2")
parser.add_argument("--linears", nargs="+", choices=("mlp_fc1", "mlp_fc2", "attn_qkv_proj", "attn_out_proj"), default=("mlp_fc1", "mlp_fc2"))
parser.add_argument("--threads", nargs="+", type=int, default=(64, 128, 256, 512))
parser.add_argument("--warmup", type=int, default=20)
parser.add_argument("--iterations", type=int, default=80)
parser.add_argument("--device", default="cuda")
return parser.parse_args()
def main() -> None:
args = parse_args()
os.environ.setdefault("H3_NVFP4_SCALE_BACKEND", "vortex")
os.environ.setdefault("H3_NVFP4_SCALE_VERSION", "1")
block, inputs, metadata = representative_inputs(args)
extension = _vortex_scale_extension()
results = []
with torch.inference_mode():
for name in args.linears:
module = module_for_name(block, name)
tensor = inputs[name].reshape(-1, module.in_features).contiguous()
scale = nvfp4_activation_scale(tensor).float()
pad = TensorCoreNVFP4Layout.get_padded_shape(tuple(tensor.shape)) != tuple(tensor.shape)
q_ref, b_ref = ck.quantize_nvfp4(tensor, scale, pad_16x=pad)
touched = bytes_touched(tensor, q_ref, b_ref)
for _ in range(args.warmup):
ck.quantize_nvfp4(tensor, scale, pad_16x=pad)
ck_timing = timed(lambda: ck.quantize_nvfp4(tensor, scale, pad_16x=pad), args.iterations)
results.append(
{
"name": name,
"packer": "ck",
"threads": None,
"input_shape": list(tensor.shape),
"qdata_shape": list(q_ref.shape),
"block_scale_shape": list(b_ref.shape),
"bytes_touched": touched,
"effective_gbps_p50": touched / ck_timing["p50_s"] / 1e9,
"effective_gbps_p95": touched / ck_timing["p95_s"] / 1e9,
"timing": ck_timing,
"qdata_diff": 0,
"block_scale_diff": 0,
}
)
for threads in args.threads:
for _ in range(args.warmup):
extension.quantize_nvfp4_bf16(tensor, scale, pad, threads)
q_native, b_native = extension.quantize_nvfp4_bf16(tensor, scale, pad, threads)
qdiff = int((q_ref != q_native).sum().item())
bdiff = int((b_ref.view(torch.uint8) != b_native.view(torch.uint8)).sum().item())
timing = timed(lambda: extension.quantize_nvfp4_bf16(tensor, scale, pad, threads), args.iterations)
results.append(
{
"name": name,
"packer": "vortex_native",
"threads": threads,
"input_shape": list(tensor.shape),
"qdata_shape": list(q_native.shape),
"block_scale_shape": list(b_native.shape),
"bytes_touched": touched,
"effective_gbps_p50": touched / timing["p50_s"] / 1e9,
"effective_gbps_p95": touched / timing["p95_s"] / 1e9,
"timing": timing,
"qdata_diff": qdiff,
"block_scale_diff": bdiff,
}
)
print(name, "threads", threads, "p50_ms", round(timing["p50_s"] * 1000, 3), "p95_ms", round(timing["p95_s"] * 1000, 3), "qdiff", qdiff, "bdiff", bdiff, flush=True)
output = {"metadata": metadata, "results": results}
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(output, indent=2), encoding="utf-8")
print(json.dumps(output, indent=2), flush=True)
if __name__ == "__main__":
main()