h3-blackwell-runtime/research/cute_nvfp4_ring/patches/0003-tile-producer.patch
2026-08-25 20:30:22 +07:00

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