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