Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
67 changes: 67 additions & 0 deletions tests/jax/test_mxfp8_cutedsl_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -283,3 +283,70 @@ def test_dtypes(method, act_type, act_desc, fp8_dtype, in_dtype):
in_dtype,
fp8_dtype,
)


# JAX's V2 grouped quantize emits GEMM-swizzled scales and supports common last
# dimensions, represented by device row counts even for equal-size groups.
# Varying last dimensions and grouped fused activations are covered
# through the C API / PyTorch, since JAX has no wrappers for those operations.
# (name, shape representation, input shape, optional device-resident row counts)
GROUP_CASES = [
("single_member", "varying_first_dim", (1, 128, 128), None),
("same_both_multichunk", "varying_first_dim", (3, 384, 384), None),
("varying_first", "varying_first_dim", (768, 256), (128, 384, 256)),
("varying_first_multichunk", "varying_first_dim", (1024, 384), (128, 256, 384, 256)),
("varying_first_empty", "varying_first_dim", (512, 256), (128, 0, 384)),
]
GROUP_Q_LAYOUTS = [QuantizeLayout.ROWWISE, QuantizeLayout.COLWISE, QuantizeLayout.ROWWISE_COLWISE]


def get_group_cfg_key(shape_rep, in_dtype, fp8_dtype, q_layout):
"""Registry key for the grouped V2 kernel (swizzled, without dbias or activation)."""
major, minor = device_compute_capability()
return (
f"cutedsl_group_mxfp8_sm{major * 10 + minor}_{DTYPE_TO_STR[in_dtype]}_"
f"{FP8_TO_KEY[fp8_dtype]}_{int(q_layout.has_rowwise)}_{int(q_layout.has_colwise)}_"
f"{shape_rep}_1_0_0_0_none"
)


@pytest.mark.parametrize("case", GROUP_CASES, ids=lambda c: c[0])
@pytest.mark.parametrize("q_layout", GROUP_Q_LAYOUTS, ids=lambda l: l.name.lower())
@pytest.mark.parametrize("in_dtype", IN_DTYPES, ids=get_dtype_id)
@pytest.mark.parametrize("fp8_dtype", FP8_DTYPES, ids=get_fp8_id)
def test_group_cast_only(case, q_layout, in_dtype, fp8_dtype):
"""Compare every grouped data/scale byte, including swizzled scales, across backends."""
name, shape_rep, shape, row_counts = case
x, _ = generate_inputs(int(np.prod(shape[:-1])), shape[-1], in_dtype)
x = x.reshape(shape)
group_sizes = None if row_counts is None else jnp.asarray(row_counts, dtype=jnp.int32)
n_groups = shape[0] if row_counts is None else len(row_counts)
quantizer = QuantizerFactory.create(
scaling_mode=ScalingMode.MXFP8_1D_SCALING,
q_dtype=fp8_dtype,
q_layout=q_layout,
n_groups=n_groups,
)

def run():
out = tex.grouped_quantize(x, quantizer=quantizer, group_sizes=group_sizes)
# Host materialization completes the FFI call before changing the backend flag.
return extract_quantized_output(out, None)[0]

set_cutedsl_backend(False)
cuda_output = run()
set_cutedsl_backend(True)
try:
cutedsl_output = run()
finally:
set_cutedsl_backend(False)

key = get_group_cfg_key(shape_rep, in_dtype, fp8_dtype, q_layout)
assert (
tvm_ffi.get_global_func(key, allow_missing=True) is not None
), f"CuTeDSL kernel not registered for {key}; grouped quantization fell back to CUDA"
for part, cuda_bytes in cuda_output.items():
assert np.array_equal(cutedsl_output[part], cuda_bytes), (
f"group/{name}/{get_dtype_id(in_dtype)}/{get_fp8_id(fp8_dtype)}/{q_layout.name}: "
f"{part} differ between backends"
)
221 changes: 221 additions & 0 deletions tests/pytorch/mxfp8/test_mxfp8_cutedsl_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -370,3 +370,224 @@ def test_sizes(swizzled, method, act, block_size, shape):
@pytest.mark.parametrize("swizzled", SWIZZLE_MODES, ids=get_swizzle_id)
def test_dtypes(swizzled, method, act, fp8_dtype, in_dtype):
run_test_case(method, act, (256, 384), (32, 32), in_dtype, fp8_dtype, swizzled)


# Grouped quantization (nvte_group_quantize / nvte_group_quantize_dbias). Every member's first
# dim is a multiple of 128, which both grouped kernels require. VARYING_LAST_DIM and
# VARYING_BOTH_DIMS members also keep their last dims 128-aligned, as both kernels' per-member
# scale layout assumes. The fused activations have no PyTorch binding; the CuTeDSL path for
# them is covered by tests/cpp/operator/test_cast_mxfp8_grouped.cu run with
# NVTE_ENABLE_CUTEDSL_BACKEND=1.
# (name, shape representation in the config key, per-member shapes)
GROUP_CASES = [
("same_both", "same_both_dims", [(256, 512)] * 3),
("single_member", "same_both_dims", [(384, 384)]),
# N is 32- but not 128-divisible, so the rowwise and colwise scales carry zeroed padding.
("same_both_n96", "same_both_dims", [(128, 96)] * 2),
("varying_first", "varying_first_dim", [(128, 256), (384, 256), (256, 256)]),
("varying_first_n160", "varying_first_dim", [(128, 160), (256, 160)]),
# Partial 32-element blocks (e.g. N=144) are covered by the C++ grouped tests;
# the PyTorch MXFP8 quantizer requires dimensions divisible by 32.
("varying_last", "varying_last_dim", [(256, 128), (256, 384), (256, 256)]),
("varying_both", "varying_both_dims", [(128, 256), (256, 128), (384, 512)]),
# Multiple chunks in both dimensions, including a final half-width 128-column strip.
("varying_both_multichunk", "varying_both_dims", [(128, 128), (256, 384), (384, 640)]),
("varying_first_empty", "varying_first_dim", [(128, 256), (0, 256), (384, 256)]),
("varying_both_empty", "varying_both_dims", [(128, 128), (0, 256), (256, 384)]),
]
SINGLE_TENSOR_GROUP_CASES = [
c for c in GROUP_CASES if c[1] in ("same_both_dims", "varying_first_dim")
]
get_group_case_id = lambda c: c[0]


def group_dims(members):
"""(logical shape, first_dims, last_dims) of the grouped tensor made of `members`."""
first_dims = [m for m, _ in members]
last_dims = [n for _, n in members]
same_first = len(set(first_dims)) == 1
same_last = len(set(last_dims)) == 1
if same_last:
logical_shape = (sum(first_dims), last_dims[0])
elif same_first:
logical_shape = (first_dims[0], sum(last_dims))
else:
logical_shape = (1, sum(m * n for m, n in members))
to_tensor = lambda dims: torch.tensor(dims, dtype=torch.int64, device="cuda")
return (
logical_shape,
None if same_first else to_tensor(first_dims),
None if same_last else to_tensor(last_dims),
)


def run_group_quantize(members, in_dtype, fp8_dtype, rowwise, columnwise, swizzled, dbias):
"""Quantize the concatenated members; returns (grouped output, dbias or None)."""
logical_shape, first_dims, last_dims = group_dims(members)
x, _ = generate_inputs(*logical_shape, in_dtype)
q = MXFP8Quantizer(fp8_dtype=fp8_dtype, rowwise=rowwise, columnwise=columnwise)
q.optimize_for_gemm = swizzled
if dbias:
return tex.bgrad_group_quantize(x, q, len(members), first_dims, last_dims)
return tex.group_quantize(x, q, len(members), first_dims, last_dims), None


def extract_group_quantized_output(out, rowwise, columnwise):
"""Extract the bytes to compare between backends.

Unlike the single-tensor kernels, both grouped backends zero the scale padding, so the
whole data and scale buffers are compared.
"""
parts = {}
if rowwise:
parts["rowwise data"] = out.rowwise_data.view(torch.uint8).clone()
parts["rowwise scales"] = out.scale_inv.view(torch.uint8).clone()
if columnwise:
parts["colwise data"] = out.columnwise_data.view(torch.uint8).clone()
parts["colwise scales"] = out.columnwise_scale_inv.view(torch.uint8).clone()
return parts


def get_group_cfg_key(shape_rep, in_dtype, fp8_dtype, rowwise, colwise, swizzled, dbias):
"""Mirror of MXFP8GroupQuantConfig::to_key (group_quantize_mxfp8_cutedsl.cuh) for the
configs reachable from PyTorch (no fused activation)."""
major, minor = device_compute_capability()
flags = (swizzled, dbias, False, False) # swizzled, with_dbias, with_dact, with_act
return (
f"cutedsl_group_mxfp8_sm{major * 10 + minor}_{DTYPE_TO_STR[in_dtype]}_"
f"{FP8_TO_KEY[fp8_dtype]}_{int(rowwise)}_{int(colwise)}_{shape_rep}_"
+ "_".join("1" if f else "0" for f in flags)
+ "_none"
)


def assert_group_cutedsl_registered(*key_args):
"""Guard against a silent CUDA fallback; see run_test_case."""
key = get_group_cfg_key(*key_args)
assert tvm_ffi.get_global_func(key, allow_missing=True) is not None, (
f"CuTeDSL kernel not registered for {key}; the CuTeDSL backend fell back "
"to CUDA and this case compared CUDA against itself"
)


def run_group_test_case(
members, shape_rep, block_size, in_dtype, fp8_dtype, swizzled=False, dbias=False
):
"""Assert the CuTeDSL and CUDA grouped backends produce bit-identical outputs, including
dbias, which both accumulate in the same order."""
rowwise = block_size[1] != 1
columnwise = block_size[0] != 1
args = (members, in_dtype, fp8_dtype, rowwise, columnwise, swizzled, dbias)

set_cutedsl_backend(False)
out_cuda, dbias_cuda = run_group_quantize(*args)
cuda_output = extract_group_quantized_output(out_cuda, rowwise, columnwise)

set_cutedsl_backend(True)
try:
out_cutedsl, dbias_cutedsl = run_group_quantize(*args)
cutedsl_output = extract_group_quantized_output(out_cutedsl, rowwise, columnwise)
finally:
set_cutedsl_backend(False)

assert_group_cutedsl_registered(
shape_rep, in_dtype, fp8_dtype, rowwise, columnwise, swizzled, dbias
)
tag = f"group/{members}/{DTYPE_TO_STR[in_dtype]}/{FP8_TO_STR[fp8_dtype]}"
for name, cuda_bytes in cuda_output.items():
assert torch.equal(
cutedsl_output[name], cuda_bytes
), f"{tag}: {name} differ between backends"
if dbias:
assert torch.equal(
dbias_cutedsl.view(torch.uint8), dbias_cuda.view(torch.uint8)
), f"{tag}: dbias differs between backends"


@pytest.mark.parametrize("case", GROUP_CASES, ids=get_group_case_id)
@pytest.mark.parametrize("block_size", BLOCK_SIZES, ids=get_block_id)
@pytest.mark.parametrize("in_dtype", IN_DTYPES, ids=get_dtype_id)
@pytest.mark.parametrize("fp8_dtype", FP8_DTYPES, ids=get_fp8_id)
def test_group_cast_only(fp8_dtype, in_dtype, block_size, case):
_, shape_rep, members = case
run_group_test_case(members, shape_rep, block_size, in_dtype, fp8_dtype)


@pytest.mark.parametrize("case", GROUP_CASES, ids=get_group_case_id)
@pytest.mark.parametrize("block_size", BLOCK_SIZES, ids=get_block_id)
def test_group_swizzled(block_size, case):
_, shape_rep, members = case
run_group_test_case(
members, shape_rep, block_size, torch.bfloat16, tex.DType.kFloat8E4M3, swizzled=True
)


@pytest.mark.parametrize("case", SINGLE_TENSOR_GROUP_CASES, ids=get_group_case_id)
@pytest.mark.parametrize("block_size", BLOCK_SIZES, ids=get_block_id)
@pytest.mark.parametrize("in_dtype", IN_DTYPES, ids=get_dtype_id)
@pytest.mark.parametrize("swizzled", SWIZZLE_MODES, ids=get_swizzle_id)
def test_group_dbias(swizzled, in_dtype, block_size, case):
_, shape_rep, members = case
run_group_test_case(
members,
shape_rep,
block_size,
in_dtype,
tex.DType.kFloat8E4M3,
swizzled=swizzled,
dbias=True,
)


@pytest.mark.parametrize("cutedsl", [False, True], ids=["cuda", "cutedsl"])
def test_group_noop(cutedsl):
"""A set cast-noop flag leaves a reused grouped output untouched; a clear one does not."""
members = [(256, 512)] * 3
logical_shape, _, _ = group_dims(members)
x1, x2 = generate_inputs(*logical_shape, torch.bfloat16)
q = MXFP8Quantizer(fp8_dtype=tex.DType.kFloat8E4M3, rowwise=True, columnwise=True)
set_cutedsl_backend(cutedsl)
try:
out = tex.group_quantize(x1, q, len(members), None)
first = extract_group_quantized_output(out, True, True)
noop = torch.ones(1, dtype=torch.float32, device="cuda")
tex.group_quantize(x2, q, len(members), None, noop_flag=noop, output=out)
skipped = extract_group_quantized_output(out, True, True)
noop.zero_()
tex.group_quantize(x2, q, len(members), None, noop_flag=noop, output=out)
quantized = extract_group_quantized_output(out, True, True)
expected = extract_group_quantized_output(
tex.group_quantize(x2, q, len(members), None), True, True
)
finally:
set_cutedsl_backend(False)
if cutedsl:
assert_group_cutedsl_registered(
"same_both_dims", torch.bfloat16, tex.DType.kFloat8E4M3, True, True, False, False
)
for name, first_bytes in first.items():
assert torch.equal(skipped[name], first_bytes), f"{name} changed under a set noop flag"
assert torch.equal(quantized[name], expected[name]), f"{name} wrong under a clear flag"


def test_group_2d_quantization_fallback():
"""2D block scaling is left to the CUDA kernel, which must still be what runs."""
members = [(128, 256), (384, 256), (256, 256)]
logical_shape, first_dims, _ = group_dims(members)
x, _ = generate_inputs(*logical_shape, torch.bfloat16)
outputs = []
for cutedsl in (False, True):
q = MXFP8Quantizer(
fp8_dtype=tex.DType.kFloat8E4M3,
rowwise=True,
columnwise=True,
with_2d_quantization=True,
)
set_cutedsl_backend(cutedsl)
try:
out = tex.group_quantize(x, q, len(members), first_dims)
outputs.append(extract_group_quantized_output(out, True, True))
finally:
set_cutedsl_backend(False)
for name, cuda_bytes in outputs[0].items():
assert torch.equal(outputs[1][name], cuda_bytes), f"{name} differ between backends"
8 changes: 8 additions & 0 deletions transformer_engine/common/CuTeDSL/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,9 @@
from transformer_engine.common.CuTeDSL.cast.mxfp8.quantize_mxfp8 import (
get_mxfp8_quantization_function,
)
from transformer_engine.common.CuTeDSL.cast.mxfp8.group_quantize_mxfp8 import (
get_mxfp8_group_quantization_function,
)


def register_cutedsl_backends():
Expand All @@ -22,3 +25,8 @@ def register_cutedsl_backends():
tvm_ffi.register_global_func(
"get_mxfp8_quantization_function", get_mxfp8_quantization_function, override=True
)
tvm_ffi.register_global_func(
"get_mxfp8_group_quantization_function",
get_mxfp8_group_quantization_function,
override=True,
)
1 change: 1 addition & 0 deletions transformer_engine/common/CuTeDSL/cast/mxfp8/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,3 +5,4 @@
"""CuTeDSL MXFP8 quantization kernels."""

from . import quantize_mxfp8
from . import group_quantize_mxfp8
Loading
Loading