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