h3-blackwell-runtime/research/cute_nvfp4_ring/patches/0004-ring-validator.patch
2026-08-25 20:30:22 +07:00

277 lines
11 KiB
Diff

diff --git a/tools/validate_cute_nvfp4_ring.py b/tools/validate_cute_nvfp4_ring.py
new file mode 100644
index 0000000..51217ec
--- /dev/null
+++ b/tools/validate_cute_nvfp4_ring.py
@@ -0,0 +1,271 @@
+"""Validate and time a bounded 128-row NVFP4 packed-tile ring prototype."""
+
+from __future__ import annotations
+
+import argparse
+import json
+import math
+from pathlib import Path
+
+import cutlass
+import cutlass.cute as cute
+import cutlass.torch as cutlass_torch
+import torch
+import torch.nn.functional as functional
+from cutlass.cute.runtime import from_dlpack
+
+from h3_blackwell_runtime.nvfp4_quant import (
+ nvfp4_activation_scale,
+ vortex_native_quantize_nvfp4,
+ vortex_native_quantize_nvfp4_into,
+ vortex_quantize_nvfp4,
+)
+from profile_nvfp4_linear import module_for_name, representative_inputs
+from validate_cute_nvfp4_h3 import (
+ fp4_tensor,
+ load_cutlass_example,
+ output_tensor,
+ scale_tensor,
+)
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser(description=__doc__)
+ parser.add_argument("--cutlass-example", type=Path, required=True)
+ parser.add_argument("--output", type=Path, required=True)
+ parser.add_argument(
+ "--linear",
+ choices=("attn_qkv_proj", "attn_out_proj", "mlp_fc1"),
+ required=True,
+ )
+ parser.add_argument("--warmup", type=int, default=5)
+ parser.add_argument("--iterations", type=int, default=20)
+ parser.add_argument("--rows", type=int, default=128)
+ 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_cuda(fn, *, warmup: int, iterations: int) -> tuple[object, float]:
+ result = None
+ for _ in range(warmup):
+ result = fn()
+ torch.cuda.synchronize()
+ started = torch.cuda.Event(enable_timing=True)
+ finished = torch.cuda.Event(enable_timing=True)
+ started.record()
+ for _ in range(iterations):
+ result = fn()
+ finished.record()
+ finished.synchronize()
+ return result, started.elapsed_time(finished) / iterations
+
+
+def main() -> None:
+ from comfy_kitchen.tensor import TensorCoreNVFP4Layout
+ import comfy_kitchen as ck
+
+ args = parse_args()
+ if args.warmup < 0 or args.iterations <= 0:
+ raise ValueError("--warmup must be non-negative and --iterations must be positive")
+ if args.rows <= 0 or args.rows % 128:
+ raise ValueError("--rows must be a positive multiple of 128")
+
+ example = load_cutlass_example(args.cutlass_example, fuse_alpha=True)
+ block, inputs, metadata = representative_inputs(args)
+ linear = module_for_name(block, args.linear)
+ full_activation = inputs[args.linear].reshape(-1, linear.in_features).contiguous()
+ if args.rows > full_activation.shape[0]:
+ raise ValueError(f"--rows exceeds the available {full_activation.shape[0]} rows")
+ activation = full_activation[:args.rows].contiguous()
+ if activation.shape[1] % 128:
+ raise ValueError(f"Ring prototype requires K divisible by 128, got {activation.shape}")
+ global_scale = nvfp4_activation_scale(full_activation).float()
+
+ with torch.inference_mode():
+ expected_packed = vortex_quantize_nvfp4(activation, scale=global_scale)
+ ring_packed = vortex_native_quantize_nvfp4(activation, scale=global_scale)
+ packed_weight = linear._packed_weight()
+ expected_qdata, expected_tensor_scale, expected_block_scales = (
+ TensorCoreNVFP4Layout.get_plain_tensors(expected_packed)
+ )
+ ring_qdata, ring_tensor_scale, ring_block_scales = (
+ TensorCoreNVFP4Layout.get_plain_tensors(ring_packed)
+ )
+ b_qdata, tensor_scale_b, b_block_scales = (
+ TensorCoreNVFP4Layout.get_plain_tensors(packed_weight)
+ )
+ reference = functional.linear(expected_packed, packed_weight, None)
+
+ a, _ = fp4_tensor(ring_qdata, swap_nibbles=False, reencode=True)
+ b, _ = fp4_tensor(b_qdata, swap_nibbles=False, reencode=True)
+ sfa = scale_tensor(ring_block_scales)
+ sfb = scale_tensor(b_block_scales)
+ output_bf16 = torch.zeros(
+ args.rows, b_qdata.shape[0], device="cuda", dtype=torch.bfloat16,
+ )
+ c = output_tensor(output_bf16)
+ alpha = ring_tensor_scale.float() * tensor_scale_b.float()
+ alpha_argument = from_dlpack(alpha.reshape(1).contiguous(), assumed_align=4)
+ gemm = example.Sm120BlockScaledGemmKernel(
+ cutlass.Float32, 16, (128, 128, 128), (128, 128),
+ )
+ max_active_clusters = cutlass.utils.HardwareInfo().get_max_active_clusters(1)
+ stream = cutlass_torch.default_stream()
+ compiled_gemm = cute.compile(
+ gemm, a, b, sfa, sfb, c, alpha_argument, max_active_clusters, stream,
+ )
+
+ def run_consumer():
+ return compiled_gemm(a, b, sfa, sfb, c, alpha_argument, stream)
+
+ def run_pack_into():
+ return vortex_native_quantize_nvfp4_into(
+ activation, global_scale, ring_qdata, ring_block_scales,
+ )
+
+ def run_cute_ring():
+ run_pack_into()
+ return run_consumer()
+
+ def run_comfy_consumer():
+ return ck.scaled_mm_nvfp4(
+ ring_qdata,
+ b_qdata,
+ tensor_scale_a=ring_tensor_scale,
+ tensor_scale_b=tensor_scale_b,
+ block_scale_a=ring_block_scales,
+ block_scale_b=b_block_scales,
+ out_dtype=torch.bfloat16,
+ alpha=ring_tensor_scale.float() * tensor_scale_b.float(),
+ )
+
+ def run_comfy_ring():
+ run_pack_into()
+ return run_comfy_consumer()
+
+ run_consumer()
+ torch.cuda.synchronize()
+ candidate = output_bf16[:, : linear.out_features]
+ delta = candidate.float() - reference.float()
+
+ _, scale_ms = measure_cuda(
+ lambda: nvfp4_activation_scale(full_activation),
+ warmup=args.warmup,
+ iterations=args.iterations,
+ )
+ _, producer_ms = measure_cuda(
+ run_pack_into,
+ warmup=args.warmup,
+ iterations=args.iterations,
+ )
+ _, consumer_ms = measure_cuda(
+ run_consumer,
+ warmup=args.warmup,
+ iterations=args.iterations,
+ )
+ _, actual_ring_ms = measure_cuda(
+ run_cute_ring,
+ warmup=args.warmup,
+ iterations=args.iterations,
+ )
+ comfy_candidate, comfy_consumer_ms = measure_cuda(
+ run_comfy_consumer,
+ warmup=args.warmup,
+ iterations=args.iterations,
+ )
+ _, comfy_ring_ms = measure_cuda(
+ run_comfy_ring,
+ warmup=args.warmup,
+ iterations=args.iterations,
+ )
+ _, reference_ms = measure_cuda(
+ lambda: functional.linear(
+ vortex_quantize_nvfp4(activation, scale=global_scale),
+ packed_weight,
+ None,
+ ),
+ warmup=args.warmup,
+ iterations=args.iterations,
+ )
+
+ qdata_bytes = ring_qdata.numel() * ring_qdata.element_size()
+ sfa_bytes = ring_block_scales.numel() * ring_block_scales.element_size()
+ chunk_count = math.ceil(full_activation.shape[0] / args.rows)
+ modeled_cute_chunk_ms = producer_ms + consumer_ms
+ modeled_comfy_chunk_ms = producer_ms + comfy_consumer_ms
+ modeled_reference_canonical_ms = scale_ms + chunk_count * reference_ms
+ report = {
+ "device": torch.cuda.get_device_name(),
+ "cutlass_dsl": "4.6.2",
+ "metadata": metadata,
+ "linear": args.linear,
+ "mnk": [args.rows, linear.out_features, activation.shape[1]],
+ "full_activation_rows": full_activation.shape[0],
+ "modeled_chunk_count": chunk_count,
+ "ring": {
+ "row_capacity": args.rows,
+ "producer": "vortex_native_quantize_nvfp4",
+ "qdata_bytes": qdata_bytes,
+ "sfa_bytes": sfa_bytes,
+ "logical_bytes": qdata_bytes + sfa_bytes,
+ },
+ "parity": {
+ "tensor_scale_equal": torch.equal(
+ ring_tensor_scale, expected_tensor_scale,
+ ),
+ "fp4_difference_count": int(
+ (ring_qdata != expected_qdata).sum().item()
+ ),
+ "block_scale_difference_count": int(
+ (ring_block_scales.view(torch.uint8)
+ != expected_block_scales.view(torch.uint8)).sum().item()
+ ),
+ "output_equal": torch.equal(candidate, reference),
+ "max_abs": delta.abs().max().item(),
+ "mean_abs": delta.abs().mean().item(),
+ "comfy_output_equal": torch.equal(
+ comfy_candidate[: args.rows, : linear.out_features], reference,
+ ),
+ },
+ "timing": {
+ "warmup": args.warmup,
+ "iterations": args.iterations,
+ "producer_ms": producer_ms,
+ "global_scale_ms": scale_ms,
+ "cute_consumer_ms": consumer_ms,
+ "modeled_cute_chunk_ms": modeled_cute_chunk_ms,
+ "modeled_comfy_chunk_ms": modeled_comfy_chunk_ms,
+ "actual_into_ring_cute_gemm_ms": actual_ring_ms,
+ "comfy_consumer_ms": comfy_consumer_ms,
+ "actual_into_ring_comfy_gemm_ms": comfy_ring_ms,
+ "reference_vortex_scale_comfy_pack_gemm_ms": reference_ms,
+ "modeled_canonical_reference_ms": modeled_reference_canonical_ms,
+ "modeled_canonical_cute_ring_ms": scale_ms + chunk_count * actual_ring_ms,
+ "modeled_canonical_comfy_ring_ms": scale_ms + chunk_count * comfy_ring_ms,
+ "actual_ring_vs_reference": actual_ring_ms / reference_ms,
+ "comfy_ring_vs_reference": comfy_ring_ms / reference_ms,
+ "modeled_canonical_cute_ring_vs_reference": (
+ scale_ms + chunk_count * actual_ring_ms
+ ) / modeled_reference_canonical_ms,
+ "modeled_canonical_comfy_ring_vs_reference": (
+ scale_ms + chunk_count * comfy_ring_ms
+ ) / modeled_reference_canonical_ms,
+ },
+ }
+ 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()