h3-blackwell-runtime/tools/validate_cute_nvfp4_conversion.py

115 lines
4.5 KiB
Python
Raw Normal View History

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