230 lines
9.1 KiB
Diff
230 lines
9.1 KiB
Diff
diff --git a/tools/validate_cute_nvfp4_tile_producer.py b/tools/validate_cute_nvfp4_tile_producer.py
|
|
new file mode 100644
|
|
index 0000000..d03c79e
|
|
--- /dev/null
|
|
+++ b/tools/validate_cute_nvfp4_tile_producer.py
|
|
@@ -0,0 +1,224 @@
|
|
+"""Validate a CuTe BF16-to-NVFP4 producer on one 128x128 activation tile."""
|
|
+
|
|
+from __future__ import annotations
|
|
+
|
|
+import json
|
|
+from pathlib import Path
|
|
+from typing import Optional
|
|
+
|
|
+import cutlass
|
|
+import cutlass.cute as cute
|
|
+import cutlass.torch as cutlass_torch
|
|
+import torch
|
|
+from cutlass import Float32
|
|
+from cutlass.cutlass_dsl import dsl_user_op
|
|
+from cutlass._mlir import ir
|
|
+from cutlass._mlir.dialects import llvm
|
|
+from cutlass.cute.runtime import from_dlpack
|
|
+
|
|
+from h3_blackwell_runtime.nvfp4_quant import nvfp4_activation_scale
|
|
+
|
|
+
|
|
+TILE = 128
|
|
+BLOCK = 16
|
|
+BLOCKS_PER_ROW = TILE // BLOCK
|
|
+JOBS = TILE * BLOCKS_PER_ROW
|
|
+
|
|
+
|
|
+@dsl_user_op
|
|
+def rcp_approx_ftz_f32(
|
|
+ x: Float32,
|
|
+ *,
|
|
+ loc: Optional[ir.Location] = None,
|
|
+ ip: Optional[ir.InsertionPoint] = None,
|
|
+) -> Float32:
|
|
+ result = llvm.inline_asm(
|
|
+ Float32.mlir_type,
|
|
+ [x.ir_value(loc=loc, ip=ip)],
|
|
+ "rcp.approx.ftz.f32 $0, $1;",
|
|
+ "=f,f",
|
|
+ has_side_effects=False,
|
|
+ asm_dialect=0,
|
|
+ loc=loc,
|
|
+ ip=ip,
|
|
+ )
|
|
+ return Float32(result)
|
|
+
|
|
+
|
|
+@cute.kernel
|
|
+def tile_producer_kernel(
|
|
+ source: cute.Tensor,
|
|
+ tensor_scale: cute.Tensor,
|
|
+ fp4: cute.Tensor,
|
|
+ block_scales: cute.Tensor,
|
|
+ scalar_scales: cute.Tensor,
|
|
+):
|
|
+ row = cute.arch.thread_idx()[0]
|
|
+ fp4_linear = cute.make_tensor(
|
|
+ fp4.iterator,
|
|
+ cute.make_layout((TILE * TILE,), stride=(1,)),
|
|
+ )
|
|
+ fp4_tiles = cute.zipped_divide(fp4_linear, (8,))
|
|
+ scale_tiles = cute.zipped_divide(block_scales, (8,))
|
|
+ fp4_store = cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), cutlass.Float4E2M1FN)
|
|
+ fp8_store = cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), cutlass.Float8E4M3FN)
|
|
+ source_fragment = cute.make_rmem_tensor((16,), cutlass.Float32)
|
|
+ normalized_fragment = cute.make_rmem_tensor((8,), cutlass.Float32)
|
|
+ fp4_fragment = cute.make_rmem_tensor((8,), cutlass.Float4E2M1FN)
|
|
+ scale_source = cute.make_rmem_tensor((8,), cutlass.Float32)
|
|
+ scale_fragment = cute.make_rmem_tensor((8,), cutlass.Float8E4M3FN)
|
|
+ decoded_scale_fragment = cute.make_rmem_tensor((8,), cutlass.Float32)
|
|
+ scale = tensor_scale[0]
|
|
+
|
|
+ for block_column in cutlass.range_constexpr(BLOCKS_PER_ROW):
|
|
+ job = row * BLOCKS_PER_ROW + block_column
|
|
+ column = block_column * BLOCK
|
|
+ maximum = cutlass.Float32(0.0)
|
|
+ for element in cutlass.range_constexpr(BLOCK):
|
|
+ value = source[row, column + element]
|
|
+ source_fragment[element] = value
|
|
+ maximum = cutlass.max(cutlass.max(value, -value), maximum)
|
|
+
|
|
+ raw_block_scale = (maximum / cutlass.Float32(6.0)) / scale
|
|
+ for element in cutlass.range_constexpr(8):
|
|
+ scale_source[element] = raw_block_scale
|
|
+ scale_values = scale_source.load()
|
|
+ scale_values = cute.where(
|
|
+ scale_values <= cutlass.Float32(448.0),
|
|
+ scale_values,
|
|
+ cutlass.Float32(448.0),
|
|
+ )
|
|
+ scale_fragment.store(scale_values.to(cutlass.Float8E4M3FN))
|
|
+ cute.copy(fp8_store, scale_fragment, scale_tiles[(None, job)])
|
|
+ scalar_scales[job] = scale_fragment[0]
|
|
+ decoded_scale_fragment.store(scale_fragment.load().to(cutlass.Float32))
|
|
+ decoded_scale = decoded_scale_fragment[0]
|
|
+ raw_encode_scale = rcp_approx_ftz_f32(decoded_scale * scale)
|
|
+ for element in cutlass.range_constexpr(8):
|
|
+ scale_source[element] = raw_encode_scale
|
|
+ encode_scale_values = scale_source.load()
|
|
+ encode_scale_values = cute.where(
|
|
+ encode_scale_values <= cutlass.Float32(3.402823466e38),
|
|
+ encode_scale_values,
|
|
+ cutlass.Float32(3.402823466e38),
|
|
+ )
|
|
+ scale_source.store(encode_scale_values)
|
|
+ encode_scale = scale_source[0]
|
|
+
|
|
+ for half in cutlass.range_constexpr(2):
|
|
+ for element in cutlass.range_constexpr(8):
|
|
+ normalized = source_fragment[half * 8 + element] * encode_scale
|
|
+ normalized_fragment[element] = normalized
|
|
+ fp4_fragment.store(normalized_fragment.load().to(cutlass.Float4E2M1FN))
|
|
+ output_tile = job * 2 + half
|
|
+ cute.copy(fp4_store, fp4_fragment, fp4_tiles[(None, output_tile)])
|
|
+
|
|
+
|
|
+@cute.jit
|
|
+def produce_tile(
|
|
+ source: cute.Tensor,
|
|
+ tensor_scale: cute.Tensor,
|
|
+ fp4: cute.Tensor,
|
|
+ block_scales: cute.Tensor,
|
|
+ scalar_scales: cute.Tensor,
|
|
+):
|
|
+ tile_producer_kernel(source, tensor_scale, fp4, block_scales, scalar_scales).launch(
|
|
+ grid=(1, 1, 1), block=(TILE, 1, 1),
|
|
+ )
|
|
+
|
|
+
|
|
+def unswizzle_scales(physical: torch.Tensor) -> torch.Tensor:
|
|
+ rows, scale_columns = physical.shape
|
|
+ row = torch.arange(rows, device=physical.device).view(-1, 1)
|
|
+ block_column = torch.arange(scale_columns, device=physical.device).view(1, -1)
|
|
+ row_in_tile = row % 128
|
|
+ tile = (row // 128) * (scale_columns // 4) + block_column // 4
|
|
+ within = (
|
|
+ ((row_in_tile % 32) // 2) * 32
|
|
+ + block_column % 4
|
|
+ + (row_in_tile // 32) * 4
|
|
+ + (row_in_tile % 2) * 16
|
|
+ )
|
|
+ return physical.flatten()[(tile * 512 + within).long()]
|
|
+
|
|
+
|
|
+def main() -> None:
|
|
+ import comfy_kitchen as ck
|
|
+
|
|
+ torch.manual_seed(440420)
|
|
+ random_values = torch.randn(TILE, TILE, device="cuda", dtype=torch.bfloat16)
|
|
+ zero_values = torch.zeros_like(random_values)
|
|
+ sparse_values = torch.zeros_like(random_values)
|
|
+ sparse_values.flatten()[:16] = torch.tensor(
|
|
+ [-100, -6, -4, -3, -2, -1.5, -1, -0.5, 0, 0.5, 1, 1.5, 2, 3, 6, 100],
|
|
+ device="cuda",
|
|
+ dtype=torch.bfloat16,
|
|
+ )
|
|
+ source_torch = torch.empty_like(random_values)
|
|
+ scale_torch = torch.empty(1, device="cuda", dtype=torch.float32)
|
|
+
|
|
+ source = from_dlpack(source_torch, assumed_align=16).mark_layout_dynamic(leading_dim=1)
|
|
+ scale = from_dlpack(scale_torch, assumed_align=4).mark_layout_dynamic()
|
|
+ fp4, fp4_torch = cutlass_torch.cute_tensor_like(
|
|
+ torch.zeros_like(source_torch, dtype=torch.float32),
|
|
+ cutlass.Float4E2M1FN,
|
|
+ is_dynamic_layout=True,
|
|
+ assumed_align=16,
|
|
+ )
|
|
+ block_scales_torch = torch.zeros(JOBS, 8, device="cuda", dtype=torch.uint8)
|
|
+ block_scales = from_dlpack(block_scales_torch.flatten(), assumed_align=16)
|
|
+ block_scales.element_type = cutlass.Float8E4M3FN
|
|
+ block_scales = block_scales.mark_layout_dynamic()
|
|
+ scalar_scales_torch = torch.zeros(JOBS, device="cuda", dtype=torch.uint8)
|
|
+ scalar_scales = from_dlpack(scalar_scales_torch, assumed_align=16)
|
|
+ scalar_scales.element_type = cutlass.Float8E4M3FN
|
|
+ scalar_scales = scalar_scales.mark_layout_dynamic()
|
|
+
|
|
+ compiled = cute.compile(produce_tile, source, scale, fp4, block_scales, scalar_scales)
|
|
+ cases = []
|
|
+ for name, values in (("random", random_values), ("zeros", zero_values), ("sparse_extremes", sparse_values)):
|
|
+ source_torch.copy_(values)
|
|
+ scale_torch.copy_(nvfp4_activation_scale(source_torch).float().reshape(1))
|
|
+ expected_qdata, expected_physical_scales = ck.quantize_nvfp4(
|
|
+ source_torch, scale_torch, pad_16x=False,
|
|
+ )
|
|
+ expected_low_first = ((expected_qdata & 0x0F) << 4) | ((expected_qdata & 0xF0) >> 4)
|
|
+ expected_scales = unswizzle_scales(expected_physical_scales.view(torch.uint8))
|
|
+ compiled(source, scale, fp4, block_scales, scalar_scales)
|
|
+ torch.cuda.synchronize()
|
|
+ actual_fp4 = fp4_torch.view(torch.uint8).flatten()[:TILE * TILE // 2].reshape(TILE, TILE // 2)
|
|
+ actual_scales = block_scales_torch[:, 0].reshape(TILE, BLOCKS_PER_ROW)
|
|
+ cases.append({
|
|
+ "name": name,
|
|
+ "tensor_scale": scale_torch.item(),
|
|
+ "fp4_difference_count": int((actual_fp4 != expected_low_first).sum().item()),
|
|
+ "fp4_equal": torch.equal(actual_fp4, expected_low_first),
|
|
+ "block_scale_difference_count": int((actual_scales != expected_scales).sum().item()),
|
|
+ "block_scales_equal": torch.equal(actual_scales, expected_scales),
|
|
+ "scalar_block_scales_equal": torch.equal(
|
|
+ scalar_scales_torch.reshape(TILE, BLOCKS_PER_ROW), expected_scales,
|
|
+ ),
|
|
+ "actual_block_scale_bytes": actual_scales.unique().tolist(),
|
|
+ "expected_block_scale_bytes": expected_scales.unique().tolist(),
|
|
+ })
|
|
+
|
|
+ report = {
|
|
+ "device": torch.cuda.get_device_name(),
|
|
+ "cutlass_dsl": "4.6.2",
|
|
+ "tile": [TILE, TILE],
|
|
+ "all_equal": all(
|
|
+ case["fp4_equal"]
|
|
+ and case["block_scales_equal"]
|
|
+ and case["scalar_block_scales_equal"]
|
|
+ for case in cases
|
|
+ ),
|
|
+ "cases": cases,
|
|
+ }
|
|
+ output = Path("/output/h3-blackwell-runtime/benchmarks/gb10-cute-nvfp4-tile-producer.json")
|
|
+ output.parent.mkdir(parents=True, exist_ok=True)
|
|
+ output.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8")
|
|
+ print(json.dumps(report, indent=2), flush=True)
|
|
+
|
|
+
|
|
+if __name__ == "__main__":
|
|
+ main()
|