"""Validate CuTe DSL E2M1/E4M3 conversion boundaries used by the streamed producer.""" from __future__ import annotations import json from pathlib import Path import cutlass import cutlass.cute as cute import cutlass.torch as cutlass_torch import torch from cutlass.cute.runtime import from_dlpack VALUES = ( -6.1, -6.0, -5.0, -4.999, -3.5, -3.499, -2.5, -2.499, -1.75, -1.749, -1.25, -1.249, -0.75, -0.749, -0.25, -0.249, -0.0, 0.0, 0.249, 0.25, 0.251, 0.749, 0.75, 0.751, 1.249, 1.25, 1.251, 1.749, 1.75, 1.751, 2.499, 2.5, 2.501, 3.499, 3.5, 3.501, 4.999, 5.0, 5.001, 6.0, 6.1, ) PADDED_COUNT = ((len(VALUES) + 15) // 16) * 16 @cute.kernel def conversion_kernel(source: cute.Tensor, fp4: cute.Tensor, fp8: cute.Tensor): if cute.arch.thread_idx()[0] == 0: fp4_tiles = cute.zipped_divide(fp4, (8,)) fp8_tiles = cute.zipped_divide(fp8, (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((8,), cutlass.Float32) fp4_fragment = cute.make_rmem_tensor((8,), cutlass.Float4E2M1FN) fp8_fragment = cute.make_rmem_tensor((8,), cutlass.Float8E4M3FN) for fragment_index in cutlass.range_constexpr(PADDED_COUNT // 8): offset = fragment_index * 8 for element in cutlass.range_constexpr(8): source_fragment[element] = source[offset + element] values = source_fragment.load() fp4_fragment.store(values.to(cutlass.Float4E2M1FN)) fp8_fragment.store(values.to(cutlass.Float8E4M3FN)) cute.copy(fp4_store, fp4_fragment, fp4_tiles[(None, fragment_index)]) cute.copy(fp8_store, fp8_fragment, fp8_tiles[(None, fragment_index)]) @cute.jit def convert(source: cute.Tensor, fp4: cute.Tensor, fp8: cute.Tensor): conversion_kernel(source, fp4, fp8).launch(grid=(1, 1, 1), block=(1, 1, 1)) def fp4_code(value: float) -> int: negative = torch.signbit(torch.tensor(value)).item() magnitude = abs(value) if magnitude > 5.0: code = 7 elif magnitude >= 3.5: code = 6 elif magnitude > 2.5: code = 5 elif magnitude >= 1.75: code = 4 elif magnitude > 1.25: code = 3 elif magnitude >= 0.75: code = 2 elif magnitude > 0.25: code = 1 else: code = 0 return code | (8 if negative else 0) def main() -> None: padded_values = VALUES + (0.0,) * (PADDED_COUNT - len(VALUES)) source_torch = torch.tensor(padded_values, device="cuda", dtype=torch.float32) source = from_dlpack(source_torch, assumed_align=4).mark_layout_dynamic() fp4, fp4_torch = cutlass_torch.cute_tensor_like( torch.zeros_like(source_torch), cutlass.Float4E2M1FN, is_dynamic_layout=True, assumed_align=16, ) fp8, fp8_torch = cutlass_torch.cute_tensor_like( torch.zeros_like(source_torch), cutlass.Float8E4M3FN, is_dynamic_layout=True, assumed_align=16, ) compiled = cute.compile(convert, source, fp4, fp8) compiled(source, fp4, fp8) torch.cuda.synchronize() fp4_bytes = fp4_torch.view(torch.uint8).flatten().cpu().tolist() expected_codes = [fp4_code(value) for value in padded_values] expected_low_first = [ expected_codes[index] | ((expected_codes[index + 1] if index + 1 < len(expected_codes) else 0) << 4) for index in range(0, len(expected_codes), 2) ] fp8_bytes = fp8_torch.view(torch.uint8).flatten().cpu() expected_fp8 = source_torch.to(torch.float8_e4m3fn).view(torch.uint8).cpu() report = { "device": torch.cuda.get_device_name(), "cutlass_dsl": "4.6.2", "value_count": len(VALUES), "padded_value_count": len(padded_values), "fp4_torch_shape": list(fp4_torch.shape), "fp4_bytes": fp4_bytes, "expected_low_first_bytes": expected_low_first, "fp4_prefix_equal": fp4_bytes[:len(expected_low_first)] == expected_low_first, "fp8_equal": torch.equal(fp8_bytes[:len(padded_values)], expected_fp8), "fp8_difference_count": int((fp8_bytes[:len(padded_values)] != expected_fp8).sum().item()), } output = Path("/output/h3-blackwell-runtime/benchmarks/gb10-cute-nvfp4-conversion-contract.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()