diff --git a/tools/validate_cute_nvfp4_ring_full.py b/tools/validate_cute_nvfp4_ring_full.py new file mode 100644 index 0000000..f8774fc --- /dev/null +++ b/tools/validate_cute_nvfp4_ring_full.py @@ -0,0 +1,278 @@ +"""Validate complete H3 projections through a reusable NVFP4 row ring.""" + +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( + "--linears", + nargs="+", + choices=("attn_qkv_proj", "attn_out_proj", "mlp_fc1"), + default=("attn_qkv_proj", "attn_out_proj", "mlp_fc1"), + ) + parser.add_argument("--capacity", type=int, default=2048) + parser.add_argument("--warmup", type=int, default=1) + parser.add_argument("--iterations", type=int, default=3) + 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 validate_linear(args, example, block, inputs, name: str) -> dict: + from comfy_kitchen.tensor import TensorCoreNVFP4Layout + + linear = module_for_name(block, name) + activation = inputs[name].reshape(-1, linear.in_features).contiguous() + rows, features = activation.shape + global_scale = nvfp4_activation_scale(activation).float() + packed_weight = linear._packed_weight() + b_qdata, tensor_scale_b, b_block_scales = ( + TensorCoreNVFP4Layout.get_plain_tensors(packed_weight) + ) + + with torch.inference_mode(): + reference_packed = vortex_quantize_nvfp4( + activation, scale=global_scale, + ) + reference = functional.linear(reference_packed, packed_weight, None) + ring_seed = vortex_native_quantize_nvfp4( + activation[: args.capacity], scale=global_scale, + ) + _, ring_tensor_scale, ring_block_scales = ( + TensorCoreNVFP4Layout.get_plain_tensors(ring_seed) + ) + + a, a_backing = cutlass_torch.cute_tensor_like( + torch.zeros( + args.capacity, + features, + 1, + device="cuda", + dtype=torch.float32, + ), + cutlass.Float4E2M1FN, + is_dynamic_layout=True, + assumed_align=16, + ) + ring_qdata = a_backing.view(torch.uint8).flatten()[ + : args.capacity * features // 2 + ].reshape(args.capacity, features // 2) + b, _ = fp4_tensor(b_qdata, swap_nibbles=False, reencode=True) + sfa = scale_tensor(ring_block_scales) + sfb = scale_tensor(b_block_scales) + chunks = [ + (start, min(start + args.capacity, rows)) + for start in range(0, rows, args.capacity) + ] + padded_rows = len(chunks) * args.capacity + candidate_padded = torch.zeros( + padded_rows, b_qdata.shape[0], device="cuda", dtype=torch.bfloat16, + ) + c_chunks = [ + output_tensor( + candidate_padded[ + index * args.capacity : (index + 1) * args.capacity + ] + ) + for index in range(len(chunks)) + ] + 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_chunks[0], alpha_argument, + max_active_clusters, stream, + ) + + def run_chunks(): + for index, (start, end) in enumerate(chunks): + vortex_native_quantize_nvfp4_into( + activation[start:end], + global_scale, + ring_qdata, + ring_block_scales, + hi_first=False, + ) + compiled_gemm( + a, b, sfa, sfb, c_chunks[index], alpha_argument, stream, + ) + return candidate_padded + + def run_complete_ring(): + nvfp4_activation_scale(activation) + return run_chunks() + + chunk_reports = [] + for index, (start, end) in enumerate(chunks): + vortex_native_quantize_nvfp4_into( + activation[start:end], + global_scale, + ring_qdata, + ring_block_scales, + hi_first=False, + ) + compiled_gemm( + a, b, sfa, sfb, c_chunks[index], alpha_argument, stream, + ) + torch.cuda.synchronize() + candidate = candidate_padded[start:end, : linear.out_features] + expected = reference[start:end] + delta = candidate.float() - expected.float() + chunk_reports.append({ + "start": start, + "rows": end - start, + "equal": torch.equal(candidate, expected), + "max_abs": delta.abs().max().item(), + "mean_abs": delta.abs().mean().item(), + }) + + complete_candidate = candidate_padded[:rows, : linear.out_features] + reference_checksum = reference.float().sum().item() + candidate_checksum = complete_candidate.float().sum().item() + output_bytes = candidate_padded.numel() * candidate_padded.element_size() + del reference + del reference_packed + torch.cuda.empty_cache() + + _, ring_ms = measure_cuda( + run_complete_ring, warmup=args.warmup, iterations=args.iterations, + ) + del candidate + del complete_candidate + del c_chunks + del candidate_padded + torch.cuda.empty_cache() + _, reference_ms = measure_cuda( + lambda: functional.linear( + vortex_quantize_nvfp4(activation), 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() + return { + "name": name, + "mnk": [rows, linear.out_features, features], + "capacity": args.capacity, + "chunk_count": len(chunks), + "final_chunk_rows": chunks[-1][1] - chunks[-1][0], + "ring_bytes": qdata_bytes + sfa_bytes, + "parity": { + "all_chunks_equal": all(chunk["equal"] for chunk in chunk_reports), + "max_abs": max(chunk["max_abs"] for chunk in chunk_reports), + "mean_abs_max": max(chunk["mean_abs"] for chunk in chunk_reports), + "reference_checksum": reference_checksum, + "candidate_checksum": candidate_checksum, + }, + "timing": { + "warmup": args.warmup, + "iterations": args.iterations, + "ring_complete_ms": ring_ms, + "reference_complete_ms": reference_ms, + "ring_vs_reference": ring_ms / reference_ms, + "improvement_percent": (1.0 - ring_ms / reference_ms) * 100.0, + }, + "chunks": chunk_reports, + "output_bytes": output_bytes, + } + + +def main() -> None: + args = parse_args() + if args.capacity <= 0 or args.capacity % 128: + raise ValueError("--capacity must be a positive multiple of 128") + if args.warmup < 0 or args.iterations <= 0: + raise ValueError("--warmup must be non-negative and --iterations positive") + + example = load_cutlass_example(args.cutlass_example, fuse_alpha=True) + block, inputs, metadata = representative_inputs(args) + results = [ + validate_linear(args, example, block, inputs, name) + for name in args.linears + ] + report = { + "device": torch.cuda.get_device_name(), + "cutlass_dsl": "4.6.2", + "metadata": metadata, + "block_index": args.block_index, + "capacity": args.capacity, + "all_equal": all(result["parity"]["all_chunks_equal"] for result in results), + "projection_reference_total_ms": sum( + result["timing"]["reference_complete_ms"] for result in results + ), + "projection_ring_total_ms": sum( + result["timing"]["ring_complete_ms"] for result in results + ), + "results": results, + } + report["projection_total_improvement_percent"] = ( + 1.0 + - report["projection_ring_total_ms"] + / report["projection_reference_total_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()