diff --git a/tests/jax/test_mxfp8_cutedsl_backend.py b/tests/jax/test_mxfp8_cutedsl_backend.py index 4671ea06869..a74baf93ec1 100644 --- a/tests/jax/test_mxfp8_cutedsl_backend.py +++ b/tests/jax/test_mxfp8_cutedsl_backend.py @@ -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" + ) diff --git a/tests/pytorch/mxfp8/test_mxfp8_cutedsl_backend.py b/tests/pytorch/mxfp8/test_mxfp8_cutedsl_backend.py index f409f76993d..0847421383e 100644 --- a/tests/pytorch/mxfp8/test_mxfp8_cutedsl_backend.py +++ b/tests/pytorch/mxfp8/test_mxfp8_cutedsl_backend.py @@ -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" diff --git a/transformer_engine/common/CuTeDSL/__init__.py b/transformer_engine/common/CuTeDSL/__init__.py index 50d536852c2..6698e22965b 100644 --- a/transformer_engine/common/CuTeDSL/__init__.py +++ b/transformer_engine/common/CuTeDSL/__init__.py @@ -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(): @@ -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, + ) diff --git a/transformer_engine/common/CuTeDSL/cast/mxfp8/__init__.py b/transformer_engine/common/CuTeDSL/cast/mxfp8/__init__.py index 6491573690f..97cc599ce29 100644 --- a/transformer_engine/common/CuTeDSL/cast/mxfp8/__init__.py +++ b/transformer_engine/common/CuTeDSL/cast/mxfp8/__init__.py @@ -5,3 +5,4 @@ """CuTeDSL MXFP8 quantization kernels.""" from . import quantize_mxfp8 +from . import group_quantize_mxfp8 diff --git a/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py new file mode 100644 index 00000000000..7d90d56da93 --- /dev/null +++ b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py @@ -0,0 +1,1501 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Grouped MXFP8 quantization kernel implemented in CuTeDSL.""" + +# pylint: disable=missing-class-docstring + +import logging +import os +from typing import Optional, Type + +import cutlass +from cutlass import cute +from cutlass import pipeline +from cutlass import Boolean, Float32, Int32, Int64, Float8E8M0FNU +from cutlass.cute.nvgpu import cpasync +from cutlass.cute.testing import assert_ as runtime_assert +from cutlass.utils import TensorMapUpdateMode, HardwareInfo +from cutlass.tensor_utils import TensorMapManager +from cuda.bindings.driver import CUstream # pylint: disable=no-name-in-module +import tvm_ffi + +from transformer_engine.common.CuTeDSL.utils import ( + str_to_cutlass_dtype, + device_compute_capability, +) +from transformer_engine.common.CuTeDSL.cast.mxfp8.quantize_mxfp8 import ( + MXFP8_BLOCK_SCALING_SIZE, + SYM_N_DIVISIBILITY, + FP8E4M3_MAX_NORM_RCP, + FP8E5M2_MAX_NORM_RCP, + SUPPORTED_ACTIVATIONS, + SUPPORTED_DACTIVATIONS, + derive_swizzled_scale_layout, + noop_flag_is_set, + reduce_rowwise_dbias, + quantize_rowwise_mxfp8, + quantize_colwise_mxfp8, +) + +CUTEDSL_DEBUG_LOGGING = os.environ.get("CUTEDSL_DEBUG_LOGGING", "0") == "1" +logger = logging.getLogger("transformer_engine.cutedsl.mxfp8") + +THREADS_PER_WARP = 32 +BYTES_PER_TENSORMAP = 128 +# Descriptor slots per tensor: input, rowwise output, colwise output, activation input. +NUM_TENSORMAPS = 4 +ACT_INPUT_SLOT = 3 +# One extra slot holds per-tensor (rows, cols, base_elts), so the main kernel never has to +# binary-search the offsets array. Mirrors TensorMapStorage::rows/cols/offsets upstream. +META_SLOT = NUM_TENSORMAPS +NUM_WORKSPACE_SLOTS = NUM_TENSORMAPS + 1 + +# Shape representations, mirroring ShapeRepresentation in common/utils.cuh. +SAME_BOTH_DIMS = "same_both_dims" +VARYING_FIRST_DIM = "varying_first_dim" +VARYING_LAST_DIM = "varying_last_dim" +VARYING_BOTH_DIMS = "varying_both_dims" +SUPPORTED_SHAPE_REPS = (SAME_BOTH_DIMS, VARYING_FIRST_DIM, VARYING_LAST_DIM, VARYING_BOTH_DIMS) + +# Upper bound on the group size (MAX_SUPPORTED_TENSOR_DESCRIPTORS in grouped_tma.cuh); sizes +# the fixed binary search over the offsets. +MAX_SUPPORTED_TENSORS = 64 + + +class MXFP8GroupQuantizeConfig: + """Compile-time config for the grouped MXFP8 quantize kernel.""" + + def __init__( + self, + dtype: str, + fp8_dtype: str, + rowwise: bool, + colwise: bool, + shape_rep: str, + with_gemm_swizzled_scales: bool = False, + with_dbias: bool = False, + with_dact: bool = False, + with_act: bool = False, + activation: str = "none", + ): + if dtype not in ("Float32", "Float16", "BFloat16"): + raise ValueError(f"unknown input dtype {dtype!r}; expected Float32|Float16|BFloat16") + self.DTYPE = str_to_cutlass_dtype(dtype) + self.DTYPE_STR = dtype + if fp8_dtype not in ("Float8E4M3", "Float8E5M2"): + raise ValueError( + f"unknown FP8 dtype {fp8_dtype!r}; expected 'Float8E4M3' or 'Float8E5M2'" + ) + self.FP8_DTYPE = str_to_cutlass_dtype(fp8_dtype) + self.FP8_DTYPE_STR = fp8_dtype + if not (rowwise or colwise): + raise ValueError("at least one of rowwise or colwise must be true") + self.ROWWISE = rowwise + self.COLWISE = colwise + if shape_rep not in SUPPORTED_SHAPE_REPS: + raise ValueError( + f"unsupported shape representation {shape_rep!r}; expected one of" + f" {SUPPORTED_SHAPE_REPS}" + ) + self.SHAPE_REP = shape_rep + # Only can view the grouped tensor as one tensor when the last dim is same across all tensors + self.IS_SINGLE_TENSOR = shape_rep in (SAME_BOTH_DIMS, VARYING_FIRST_DIM) + self.MAX_NORM_RCP = ( + FP8E4M3_MAX_NORM_RCP if fp8_dtype == "Float8E4M3" else FP8E5M2_MAX_NORM_RCP + ) + + self.WITH_GEMM_SWIZZLED_SCALES = with_gemm_swizzled_scales + if with_dbias and not self.IS_SINGLE_TENSOR: + # mxfp8::group_quantize raises for this. + raise ValueError("dbias is only supported for tensors with a common last dimension") + self.WITH_DBIAS = with_dbias + + if with_dact and with_act: + raise ValueError("with_dact and with_act cannot both be set") + if with_dact: + if activation not in SUPPORTED_DACTIVATIONS: + raise ValueError( + f"unknown activation {activation!r} for with_dact=True; expected one of" + f" {sorted(SUPPORTED_DACTIVATIONS)}" + ) + self.ACTIVATION = activation + elif with_act: + if activation not in SUPPORTED_ACTIVATIONS: + raise ValueError( + f"unknown activation {activation!r} for with_act=True; expected one of" + f" {sorted(SUPPORTED_ACTIVATIONS)}" + ) + self.ACTIVATION = activation + else: + if activation != "none": + raise ValueError("activation must be none when with_dact and with_act are False") + self.ACTIVATION = None + self.WITH_DACT = with_dact + self.WITH_ACT = with_act + + def __str__(self): + return ( + f"MXFP8GroupQuantizeConfig(dtype={self.DTYPE_STR}, fp8_dtype={self.FP8_DTYPE_STR}, " + f"rowwise={self.ROWWISE}, colwise={self.COLWISE}, shape_rep={self.SHAPE_REP}, " + f"swizzled={self.WITH_GEMM_SWIZZLED_SCALES}, with_dbias={self.WITH_DBIAS}, " + f"with_dact={self.WITH_DACT}, with_act={self.WITH_ACT}, " + f"activation={self.ACTIVATION})" + ) + + __repr__ = __str__ + + +class MXFP8GroupQuantizeKernel: + """Grouped MXFP8 quantize mirroring group_quantize_mxfp8_kernel's strategy.""" + + # Target persistent CTA count per SM for sizing the grid + STATIC_PERSISTENT_WORKERS_PER_SM = 24 + # The shape of one pipeline stage processed by a CTA + TILE_ROWS = 32 + TILE_COLS = 128 + PIPELINE_DEPTH = 2 + # The number of elements in a tile + ELTS_PER_TILE = TILE_ROWS * TILE_COLS + # CTA shape + THREADS_PER_CTA = 128 + NUM_WARPS = THREADS_PER_CTA // THREADS_PER_WARP # 4 + THREADS_X = TILE_COLS // MXFP8_BLOCK_SCALING_SIZE # 4 + THREADS_Y = THREADS_PER_CTA // THREADS_X # 32 + # How many elements a thread handles in a wave + PACK_SIZE = 4 + # How many waves needed to handle a MXFP8 block + WAVES = MXFP8_BLOCK_SCALING_SIZE // PACK_SIZE # 8 + # How many threads per bank -- for avoiding bank conflicts + THREADS_PER_BANK = (32 * 4) // MXFP8_BLOCK_SCALING_SIZE # 4 + + def __init__(self, cfg: MXFP8GroupQuantizeConfig, SM_COUNT: int) -> None: + self.cfg = cfg + self.SM_COUNT = SM_COUNT + # A CTA processes (NUM_TILES_Y, NUM_TILES_X) tiles, NUM_STAGES tiles in total + self.NUM_TILES_X = 2 if cfg.SHAPE_REP == VARYING_BOTH_DIMS else 1 + self.NUM_TILES_Y = 4 + self.NUM_STAGES = self.NUM_TILES_X * self.NUM_TILES_Y + self.ELTS_PER_CTA = self.ELTS_PER_TILE * self.NUM_STAGES + # The CUDA kernel honors the noop flag only without fused activations or dbias. + self.CHECK_NOOP_FLAG = not (cfg.WITH_ACT or cfg.WITH_DACT or cfg.WITH_DBIAS) + # The colwise pass caches the activation in the input tile for the rowwise pass. + # Note: ReLU is fused into the conversion instead, which yields the same bytes. + self.CACHE_ACTIVATION = ( + (cfg.WITH_ACT or cfg.WITH_DACT) + and cfg.ROWWISE + and cfg.COLWISE + and cfg.ACTIVATION != "relu" + ) + # Prefer to reduce dbias in the colwise pass if quantized in columnwise, else in the rowwise pass. + self.DBIAS_IN_COLWISE = cfg.WITH_DBIAS and cfg.COLWISE + self.DBIAS_IN_ROWWISE = cfg.WITH_DBIAS and not cfg.COLWISE + + @cute.jit + def _find_tensor_from_offsets( + self, mOffsets: cute.Tensor, num_tensors: Int32, offset: Int64 + ) -> Int32: + """Index of the tensor whose element range holds `offset` (find_tensor_from_offsets).""" + low = Int32(1) + hi = Int32(num_tensors) + # Enough bisection steps for any group of up to MAX_SUPPORTED_TENSORS members. + for _ in cutlass.range_constexpr(MAX_SUPPORTED_TENSORS.bit_length()): + if low < hi: + mid = low + (hi - low) // 2 + if Int64(mOffsets[mid]) <= offset: + low = mid + 1 + else: + hi = mid + return low - 1 + + @cute.jit + def _scale_tensor(self, mS: cute.Tensor, base: Int64, layout: cute.Layout) -> cute.Tensor: + """View the scale buffer from element `base` on with `layout`.""" + return cute.make_tensor( + cute.make_ptr( + Float8E8M0FNU, + mS.iterator.toint() + base, + cute.AddressSpace.gmem, + assumed_align=4, + ), + layout, + ) + + @cute.jit + def _rowwise_scales( + self, mS_row: cute.Tensor, base: Int64, rows: Int32, cols: Int32 + ) -> cute.Tensor: + """Rowwise scales of a (rows, cols) tensor at `base`, tiled per 32x128 stage.""" + if cutlass.const_expr(self.cfg.WITH_GEMM_SWIZZLED_SCALES): + mS_t, _ = derive_swizzled_scale_layout( + rows, cols, True, False, self._scale_tensor(mS_row, base, cute.make_layout(1)), None + ) + else: + # Rowwise scale's divisibility guarantee: (128, 4) + stride = cute.round_up(cute.ceil_div(cols, MXFP8_BLOCK_SCALING_SIZE), 4) + mS_t = self._scale_tensor( + mS_row, base, cute.make_layout((rows, stride), stride=(stride, 1)) + ) + return cute.zipped_divide( + mS_t, (self.TILE_ROWS, self.TILE_COLS // MXFP8_BLOCK_SCALING_SIZE) + ) + + @cute.jit + def _colwise_scales( + self, mS_col: cute.Tensor, base: Int64, rows: Int32, cols: Int32 + ) -> cute.Tensor: + """Colwise scales of a (rows, cols) tensor at `base`, tiled per 32x128 stage.""" + if cutlass.const_expr(self.cfg.WITH_GEMM_SWIZZLED_SCALES): + _, mS_t = derive_swizzled_scale_layout( + rows, cols, False, True, None, self._scale_tensor(mS_col, base, cute.make_layout(1)) + ) + else: + # Colwise scale's divisibility guarantee: (4, 128) + stride = cute.round_up(cols, 128) + mS_t = self._scale_tensor( + mS_col, + base, + cute.make_layout((rows // MXFP8_BLOCK_SCALING_SIZE, stride), stride=(stride, 1)), + ) + return cute.zipped_divide( + mS_t, (self.TILE_ROWS // MXFP8_BLOCK_SCALING_SIZE, self.TILE_COLS) + ) + + @cute.jit + def __call__( + self, + mX: cute.Tensor, + mO_row: Optional[cute.Tensor], + mO_col: Optional[cute.Tensor], + mS_row: cute.Tensor, + mS_col: cute.Tensor, + mOffsets: cute.Tensor, # int64[num_tensors + 1], CSR element offsets + mFirstDims: Optional[ + cute.Tensor + ], # int64[num_tensors] (VARYING_FIRST_DIM / VARYING_BOTH_DIMS) + mLastDims: Optional[ + cute.Tensor + ], # int64[num_tensors] (VARYING_LAST_DIM / VARYING_BOTH_DIMS) + mTensormaps: cute.Tensor, # int64[num_tensors, NUM_WORKSPACE_SLOTS, 16] + mNoop: cute.Pointer, # f32 cast_noop flag; may be null, checked on device + mActInput: Optional[cute.Tensor], # activation input, only with WITH_DACT + mWorkspace: Optional[cute.Tensor], # f32 partial dbias, only with WITH_DBIAS + stream: CUstream, + ) -> None: + if cutlass.const_expr(CUTEDSL_DEBUG_LOGGING): + cute.printf(f"[CuTeDSL] MXFP8GroupQuantizeKernel.__call__() cfg: {self.cfg}\n") + + cfg = self.cfg + first_logical_dim = mX.shape[0] + last_logical_dim = mX.shape[1] + if cutlass.const_expr(cfg.SHAPE_REP == VARYING_BOTH_DIMS): + runtime_assert( + first_logical_dim == 1, "VARYING_BOTH_DIMS requires logical shape [1, total]" + ) + + num_tensors = mTensormaps.shape[0] + runtime_assert(num_tensors > 0, "Grouped quantization requires at least one tensor") + + # A TMA atom copies a TILE at a time + smem_tile_layout = cute.make_ordered_layout((self.TILE_ROWS, self.TILE_COLS), order=(1, 0)) + cta_tiler = (self.TILE_ROWS, self.TILE_COLS) + + # TMA atom for loading the input tensor and the activation input (if WITH_DACT) + op_load = cpasync.CopyBulkTensorTileG2SOp() + tma_atom_x, tma_src = cpasync.make_tiled_tma_atom( + op_load, mX, smem_tile_layout, cta_tiler, num_multicast=1 + ) + tma_atom_act = None + tma_src_act = None + if cutlass.const_expr(cfg.WITH_DACT): + tma_atom_act, tma_src_act = cpasync.make_tiled_tma_atom( + op_load, mActInput, smem_tile_layout, cta_tiler, num_multicast=1 + ) + + # TMA atom for storing the rowwise and colwise outputs (if enabled) + op_store = cpasync.CopyBulkTensorTileS2GOp() + tma_atom_out_row = None + tma_dst_out_row = None + if cutlass.const_expr(cfg.ROWWISE): + tma_atom_out_row, tma_dst_out_row = cpasync.make_tiled_tma_atom( + op_store, mO_row, smem_tile_layout, cta_tiler, num_multicast=1 + ) + tma_atom_out_col = None + tma_dst_out_col = None + if cutlass.const_expr(cfg.COLWISE): + tma_atom_out_col, tma_dst_out_col = cpasync.make_tiled_tma_atom( + op_store, mO_col, smem_tile_layout, cta_tiler, num_multicast=1 + ) + + if cutlass.const_expr(cfg.IS_SINGLE_TENSOR): + # How many CTAs does the grouped tensor have in both directions + jobs_Y = cute.ceil_div(Int32(first_logical_dim), (self.TILE_ROWS * self.NUM_TILES_Y)) + jobs_X = cute.ceil_div(Int32(last_logical_dim), self.TILE_COLS * self.NUM_TILES_X) + # Flatten it to an 1D grid + grid = [jobs_X * jobs_Y, 1, 1] + else: + # A placeholder for the kernel signature only; we won't use it in non-single tensor cases + jobs_X = None + # Estimate the total jobs across the group; each job has NUM_TILES_X * NUM_TILES_Y tiles + if cutlass.const_expr(cfg.SHAPE_REP == VARYING_BOTH_DIMS): + # Note: when VARYING_BOTH_DIMS, the first_logical_dim must be 1 + estimated_jobs = cute.ceil_div( + Int32(first_logical_dim) * Int32(last_logical_dim), self.ELTS_PER_CTA + ) + elif cutlass.const_expr(cfg.SHAPE_REP == VARYING_LAST_DIM): + # Same as VARYING_BOTH_DIMS but we divide 128 before multiplying to avoid overflowing Int32 + # because the first logical dimension is always 128-aligned when not VARYING_BOTH_DIMS + # so they are equivalent + estimated_jobs = cute.ceil_div( + (Int32(first_logical_dim) // 128) * Int32(last_logical_dim), + self.ELTS_PER_CTA // 128, + ) + else: + raise ValueError(f"unexpected shape representation {cfg.SHAPE_REP!r}") + + # Divide the persistent worker budget evenly across tensors, with at least + # one worker per tensor. Each worker may process several jobs. + requested_CTAs_per_tensor = cutlass.max( + Int32(1), + Int32(self.SM_COUNT * self.STATIC_PERSISTENT_WORKERS_PER_SM) // Int32(num_tensors), + ) + # In average how many jobs per tensor (only an average, the actual jobs per tensor may vary) + average_jobs_per_tensor = cutlass.max( + Int32(1), cute.ceil_div(estimated_jobs, Int32(num_tensors)) + ) + # Don't launch more CTAs than the average jobs per tensor in case + # STATIC_PERSISTENT_WORKERS_PER_SM causes redundancy + CTAs_per_tensor = cutlass.min(requested_CTAs_per_tensor, average_jobs_per_tensor) + grid = [CTAs_per_tensor, Int32(num_tensors), 1] + + # Only the multi-tensor representations need per-tensor descriptors. + if cutlass.const_expr(not cfg.IS_SINGLE_TENSOR): + self._update_descriptors_kernel( + mX, + mO_row, + mO_col, + mActInput, + mOffsets, + mFirstDims, + mLastDims, + mTensormaps, + first_logical_dim, + last_logical_dim, + mX.element_type, + tma_atom_x, + tma_atom_out_row, + tma_atom_out_col, + tma_atom_act, + ).launch(grid=[num_tensors, 1, 1], block=[THREADS_PER_WARP, 1, 1], stream=stream) + + self.kernel( + mS_row, + mS_col, + mOffsets, + mFirstDims, + mTensormaps, + mNoop, + mWorkspace, + first_logical_dim, + last_logical_dim, + num_tensors, + jobs_X, + mX.element_type, + tma_atom_x, + tma_src, + tma_atom_act, + tma_src_act, + tma_atom_out_row, + tma_dst_out_row, + tma_atom_out_col, + tma_dst_out_col, + ).launch( + grid=grid, + block=[self.THREADS_PER_CTA, 1, 1], + stream=stream, + ) + + @cute.kernel + def _update_descriptors_kernel( + self, + mX: cute.Tensor, + mO_row: Optional[cute.Tensor], + mO_col: Optional[cute.Tensor], + mActInput: Optional[cute.Tensor], + mOffsets: cute.Tensor, + mFirstDims: Optional[cute.Tensor], + mLastDims: Optional[cute.Tensor], + mTensormaps: cute.Tensor, + first_logical_dim: Int32, + last_logical_dim: Int32, + dtype: cutlass.Constexpr[Type[cutlass.Numeric]], + tma_atom_x: cute.CopyAtom, + tma_atom_orow: Optional[cute.CopyAtom], + tma_atom_ocol: Optional[cute.CopyAtom], + tma_atom_act: Optional[cute.CopyAtom], + ) -> None: + """Update the per-tensor TMA descriptors for the group quantization kernel. + + mTensormaps: int64[num_tensors, NUM_WORKSPACE_SLOTS, 16], where the slots are: + - 0: input tensor + - 1: rowwise output tensor + - 2: colwise output tensor + - 3: activation input tensor (only with WITH_DACT) + - 4: metadata (rows, cols, base_offset) + """ + cfg = self.cfg + + # For the descriptor update kernel, grid=[num_tensors, 1, 1] + tensor_id, _, _ = cute.arch.block_idx() + # Figure out how many rows and columns this tensor has, and where its first element is in the group. + if cutlass.const_expr(cfg.SHAPE_REP in (VARYING_FIRST_DIM, VARYING_BOTH_DIMS)): + rows = Int32(mFirstDims[tensor_id]) + else: + rows = Int32(first_logical_dim) + if cutlass.const_expr(cfg.SHAPE_REP in (VARYING_LAST_DIM, VARYING_BOTH_DIMS)): + cols = Int32(mLastDims[tensor_id]) + else: + cols = Int32(last_logical_dim) + base_offset = Int64(mOffsets[tensor_id]) + + # Same diagnostics as get_tensor_rows_num / get_tensor_cols_num. Like NVTE_DEVICE_ERROR + # in a release build, they only print. + tidx, _, _ = cute.arch.thread_idx() + if tidx == 0: + if rows % 128 != 0: + cute.printf( + "tensor %d: First dimension of each tensor in a group must be divisible" + " by 128.\n", + tensor_id, + ) + if cols % 128 != 0: + cute.printf( + "tensor %d: For varying last dimensions support, the last dimension of each" + " tensor in a group must be divisible by 128.\n", + tensor_id, + ) + + meta = mTensormaps[(tensor_id, META_SLOT, None)] + meta[0] = Int64(rows) + meta[1] = Int64(cols) + meta[2] = base_offset + + tmap = TensorMapManager(TensorMapUpdateMode.GMEM, BYTES_PER_TENSORMAP) + # Obtain the pointers of these descriptors + desc_x = tmap.get_tensormap_ptr(mTensormaps[(tensor_id, 0, None)].iterator) + desc_orow = tmap.get_tensormap_ptr(mTensormaps[(tensor_id, 1, None)].iterator) + desc_ocol = tmap.get_tensormap_ptr(mTensormaps[(tensor_id, 2, None)].iterator) + desc_act = tmap.get_tensormap_ptr(mTensormaps[(tensor_id, ACT_INPUT_SLOT, None)].iterator) + + # Zero-sized groups: creating a descriptor with a zero extent is invalid, + # so skip (the main kernel skips these tensors as well). + if rows > 0 and cols > 0: + member_layout = cute.make_layout((rows, cols), stride=(cols, 1)) + + def member_view(tensor: cute.Tensor, elt_dtype: Type[cutlass.Numeric]) -> cute.Tensor: + return cute.make_tensor( + cute.make_ptr( + elt_dtype, + tensor.iterator.toint() + base_offset * (elt_dtype.width // 8), + cute.AddressSpace.gmem, + assumed_align=16, + ), + member_layout, + ) + + views = [member_view(mX, dtype)] + atoms = [tma_atom_x] + descs = [desc_x] + tmap.init_tensormap_from_atom(tma_atom_x, desc_x, 0) + if cutlass.const_expr(cfg.ROWWISE): + views.append(member_view(mO_row, cfg.FP8_DTYPE)) + atoms.append(tma_atom_orow) + tmap.init_tensormap_from_atom(tma_atom_orow, desc_orow, 0) + descs.append(desc_orow) + if cutlass.const_expr(cfg.COLWISE): + views.append(member_view(mO_col, cfg.FP8_DTYPE)) + atoms.append(tma_atom_ocol) + tmap.init_tensormap_from_atom(tma_atom_ocol, desc_ocol, 0) + descs.append(desc_ocol) + if cutlass.const_expr(cfg.WITH_DACT): + views.append(member_view(mActInput, dtype)) + atoms.append(tma_atom_act) + tmap.init_tensormap_from_atom(tma_atom_act, desc_act, 0) + descs.append(desc_act) + tmap.fence_tensormap_initialization() + # Update descriptors in global memory with views and atoms + tmap.update_tensormap( + tuple(views), + tuple(atoms), + tuple(descs), + 0, + (), # smem staging is unused in GMEM update mode + ) + + @cute.jit + def _make_shared_storage( + self, + smem: cutlass.Constexpr[cutlass.utils.SmemAllocator], + dtype: cutlass.Constexpr[Type[cutlass.Numeric]], + ) -> tuple[ + cute.Pointer, + cute.Tensor, + Optional[cute.Tensor], + Optional[cute.Tensor], + Optional[cute.Tensor], + Optional[cute.Tensor], + ]: + """Allocate pipeline buffers and optional activation input and dbias storage.""" + FP8_DTYPE = self.cfg.FP8_DTYPE + tile_layout = cute.make_layout( + ((self.TILE_ROWS, self.TILE_COLS), self.PIPELINE_DEPTH), + stride=((self.TILE_COLS, 1), self.TILE_ROWS * self.TILE_COLS), + ) + + sX = None + sO_row = None + sO_col = None + sActInput = None + + if cutlass.const_expr(self.cfg.ROWWISE and self.cfg.COLWISE): + + @cute.struct + class SharedStorage: + mbar: cute.struct.MemRange[cute.Int64, 2 * self.PIPELINE_DEPTH] + sX: cute.struct.Align[ + cute.struct.MemRange[ + dtype, self.TILE_ROWS * self.TILE_COLS * self.PIPELINE_DEPTH + ], + 128, + ] + sO_row: cute.struct.Align[ + cute.struct.MemRange[ + FP8_DTYPE, self.TILE_ROWS * self.TILE_COLS * self.PIPELINE_DEPTH + ], + 128, + ] + sO_col: cute.struct.Align[ + cute.struct.MemRange[ + FP8_DTYPE, self.TILE_ROWS * self.TILE_COLS * self.PIPELINE_DEPTH + ], + 128, + ] + + storage = smem.allocate(SharedStorage) + sX = storage.sX.get_tensor(tile_layout) + sO_row = storage.sO_row.get_tensor(tile_layout) + sO_col = storage.sO_col.get_tensor(tile_layout) + + elif cutlass.const_expr(self.cfg.ROWWISE): + + @cute.struct + class SharedStorage: + mbar: cute.struct.MemRange[cute.Int64, 2 * self.PIPELINE_DEPTH] + sX: cute.struct.Align[ + cute.struct.MemRange[ + dtype, self.TILE_ROWS * self.TILE_COLS * self.PIPELINE_DEPTH + ], + 128, + ] + sO_row: cute.struct.Align[ + cute.struct.MemRange[ + FP8_DTYPE, self.TILE_ROWS * self.TILE_COLS * self.PIPELINE_DEPTH + ], + 128, + ] + + storage = smem.allocate(SharedStorage) + sX = storage.sX.get_tensor(tile_layout) + sO_row = storage.sO_row.get_tensor(tile_layout) + + elif cutlass.const_expr(self.cfg.COLWISE): + + @cute.struct + class SharedStorage: + mbar: cute.struct.MemRange[cute.Int64, 2 * self.PIPELINE_DEPTH] + sX: cute.struct.Align[ + cute.struct.MemRange[ + dtype, self.TILE_ROWS * self.TILE_COLS * self.PIPELINE_DEPTH + ], + 128, + ] + sO_col: cute.struct.Align[ + cute.struct.MemRange[ + FP8_DTYPE, self.TILE_ROWS * self.TILE_COLS * self.PIPELINE_DEPTH + ], + 128, + ] + + storage = smem.allocate(SharedStorage) + sX = storage.sX.get_tensor(tile_layout) + sO_col = storage.sO_col.get_tensor(tile_layout) + + if cutlass.const_expr(self.cfg.WITH_DACT): + + @cute.struct + class DactStorage: + sActInput: cute.struct.Align[ + cute.struct.MemRange[ + dtype, self.TILE_ROWS * self.TILE_COLS * self.PIPELINE_DEPTH + ], + 128, + ] + + sActInput = smem.allocate(DactStorage).sActInput.get_tensor(tile_layout) + + sDbias = None + if cutlass.const_expr(self.DBIAS_IN_ROWWISE): + # Padded like the CUDA kernel's partial_dbias_rowwise to avoid bank conflicts. + DBIAS_BUFF_WIDTH = self.THREADS_X * (MXFP8_BLOCK_SCALING_SIZE + 1) + + @cute.struct + class DbiasStorage: + sDbias: cute.struct.MemRange[Float32, self.THREADS_Y * DBIAS_BUFF_WIDTH] + + sDbias = smem.allocate(DbiasStorage).sDbias.get_tensor( + cute.make_layout((self.THREADS_Y, self.TILE_COLS), stride=(DBIAS_BUFF_WIDTH, 1)) + ) + + return storage.mbar.data_ptr(), sX, sO_row, sO_col, sActInput, sDbias + + @cute.kernel + def kernel( + self, + mS_row: cute.Tensor, + mS_col: cute.Tensor, + mOffsets: cute.Tensor, + mFirstDims: Optional[cute.Tensor], + mTensormaps: cute.Tensor, + mNoop: cute.Pointer, + mWorkspace: Optional[cute.Tensor], + first_logical_dim: Int32, + last_logical_dim: Int32, + num_tensors: Int32, + jobs_X: Optional[Int32], + dtype: cutlass.Constexpr[Type[cutlass.Numeric]], + tma_atom_x: cute.CopyAtom, + tma_src: cute.Tensor, + tma_atom_act: Optional[cute.CopyAtom], + tma_src_act: Optional[cute.Tensor], + tma_atom_out_row: Optional[cute.CopyAtom], + tma_dst_out_row: Optional[cute.Tensor], + tma_atom_out_col: Optional[cute.CopyAtom], + tma_dst_out_col: Optional[cute.Tensor], + ) -> None: + """No-op the CTA when the noop flag is set, else run the quantize main loop.""" + skip_execution = Boolean(False) + if cutlass.const_expr(self.CHECK_NOOP_FLAG): + skip_execution = noop_flag_is_set(mNoop) + if not skip_execution: + self._kernel_main( + mS_row, + mS_col, + mOffsets, + mFirstDims, + mTensormaps, + mWorkspace, + first_logical_dim, + last_logical_dim, + num_tensors, + jobs_X, + dtype, + tma_atom_x, + tma_src, + tma_atom_act, + tma_src_act, + tma_atom_out_row, + tma_dst_out_row, + tma_atom_out_col, + tma_dst_out_col, + ) + + @cute.jit + def _kernel_main( + self, + mS_row: cute.Tensor, + mS_col: cute.Tensor, + mOffsets: cute.Tensor, + mFirstDims: Optional[cute.Tensor], + mTensormaps: cute.Tensor, + mWorkspace: Optional[cute.Tensor], + first_logical_dim: Int32, + last_logical_dim: Int32, + num_tensors: Int32, + jobs_X: Optional[Int32], + dtype: cutlass.Constexpr[Type[cutlass.Numeric]], + tma_atom_x: cute.CopyAtom, + tma_src: cute.Tensor, + tma_atom_act: Optional[cute.CopyAtom], + tma_src_act: Optional[cute.Tensor], + tma_atom_out_row: Optional[cute.CopyAtom], + tma_dst_out_row: Optional[cute.Tensor], + tma_atom_out_col: Optional[cute.CopyAtom], + tma_dst_out_col: Optional[cute.Tensor], + ) -> None: + cfg = self.cfg + FP8_DTYPE = cfg.FP8_DTYPE + tidx, _, _ = cute.arch.thread_idx() + bidx, bidy, _ = cute.arch.block_idx() + gdx, _, _ = cute.arch.grid_dim() + warp_idx = cute.arch.make_warp_uniform(cute.arch.warp_idx()) + + if cutlass.const_expr(cfg.SHAPE_REP == VARYING_FIRST_DIM): + # The first CTA validates every member's rows, as the CUDA kernel does. Like + # NVTE_DEVICE_ERROR in a release build, this only prints. + if bidx == 0: + if tidx < num_tensors: + if Int64(mFirstDims[tidx]) % 128 != 0: + cute.printf( + "tensor %d: First dimension of each tensor in a group must be" + " divisible by 128.\n", + tidx, + ) + + smem = cutlass.utils.SmemAllocator() + # Layouts for sX, sO_row, sO_col, sActInput: + # ((TILE_ROWS, TILE_COLS), PIPELINE_DEPTH):((TILE_COLS, 1), TILE_ROWS * TILE_COLS) + # Layout for sDbias: + # (THREADS_Y, TILE_COLS):(THREADS_X * (MXFP8_BLOCK_SCALING_SIZE + 1), 1) + mbar_ptr, sX, sO_row, sO_col, sActInput, sDbias = self._make_shared_storage(smem, dtype) + + tx_count = self.TILE_ROWS * self.TILE_COLS * dtype.width // 8 + if cutlass.const_expr(cfg.WITH_DACT): + tx_count *= 2 + mainloop_pipeline = pipeline.PipelineTmaAsync.create( + barrier_storage=mbar_ptr, + num_stages=self.PIPELINE_DEPTH, + producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, 1), + consumer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, self.NUM_WARPS), + tx_count=tx_count, + cta_layout_vmnk=None, + ) + prod_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.PIPELINE_DEPTH + ) + cons_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.PIPELINE_DEPTH + ) + + gX_tiled = cute.zipped_divide(tma_src, (self.TILE_ROWS, self.TILE_COLS)) + tXsX, tXgX = cpasync.tma_partition(tma_atom_x, 0, cute.make_layout(1), sX, gX_tiled) + + tXsA = None + tXgA = None + if cutlass.const_expr(cfg.WITH_DACT): + gA_tiled = cute.zipped_divide(tma_src_act, (self.TILE_ROWS, self.TILE_COLS)) + tXsA, tXgA = cpasync.tma_partition( + tma_atom_act, 0, cute.make_layout(1), sActInput, gA_tiled + ) + + tXsO_row = None + tXgO_row = None + if cutlass.const_expr(cfg.ROWWISE): + gO_row_tiled = cute.zipped_divide(tma_dst_out_row, (self.TILE_ROWS, self.TILE_COLS)) + tXsO_row, tXgO_row = cpasync.tma_partition( + tma_atom_out_row, 0, cute.make_layout(1), sO_row, gO_row_tiled + ) + + tXsO_col = None + tXgO_col = None + if cutlass.const_expr(cfg.COLWISE): + gO_col_tiled = cute.zipped_divide(tma_dst_out_col, (self.TILE_ROWS, self.TILE_COLS)) + tXsO_col, tXgO_col = cpasync.tma_partition( + tma_atom_out_col, 0, cute.make_layout(1), sO_col, gO_col_tiled + ) + + tmap = TensorMapManager(TensorMapUpdateMode.GMEM, BYTES_PER_TENSORMAP) + cute.arch.sync_threads() + + # If the CTA has work to do + has_work = Boolean(True) + # Metadata of the tensor that owns this job + tensor_rows = Int32(0) + tensor_cols = Int32(0) + # Element offset of this tensor within the group: Int64 (CUDA uses size_t), since a + # group can exceed 2^31 elements even when every individual extent is small. + tensor_base = Int64(0) + # Job's starting row and id in this individual tensor / global single tensor + job_start_row = Int32(0) + job_id_X = Int32(0) + + first_job_id = Int32(0) + jobs_in_tensor = Int32(1) + job_stride = Int32(1) + jobs_X_in_tensor = Int32(1) + + if cutlass.const_expr(cfg.IS_SINGLE_TENSOR): + # grid = [jobs_X * jobs_Y, 1, 1] + job_id_Y = Int32(bidx) // jobs_X + job_id_X = Int32(bidx) % jobs_X + # View the grouped tensor as a single tensor of shape (first_logical_dim, last_logical_dim) + tensor_rows = Int32(first_logical_dim) + tensor_cols = Int32(last_logical_dim) + # Which row does this job start from + job_start_row = job_id_Y * (self.TILE_ROWS * self.NUM_TILES_Y) + if cutlass.const_expr(cfg.SHAPE_REP == VARYING_FIRST_DIM): + total_elts = Int64(mOffsets[mOffsets.shape[0] - 1]) + # If the starting element is already beyond the last element of the group, this CTA can stop + if Int64(job_start_row) * Int64(last_logical_dim) >= total_elts: + has_work = Boolean(False) + # When SAME_BOTH_DIM, M is always divisible by 128, which is exactly TILE_ROWS * NUM_TILES_Y, + # so no need to check for the last job's starting row being beyond the last row of the tensor. + else: + # grid = [workers_per_tensor, Int32(num_tensors), 1] + tensor_id = Int32(bidy) + # Extract tensor's metadata + meta = mTensormaps[(tensor_id, META_SLOT, None)] + tensor_rows = Int32(meta[0]) + tensor_cols = Int32(meta[1]) + tensor_base = Int64(meta[2]) + if tensor_rows > 0 and tensor_cols > 0: + # How many jobs does this tensor have in both directions + jobs_X_in_tensor = cute.ceil_div(tensor_cols, (self.TILE_COLS * self.NUM_TILES_X)) + jobs_Y_in_tensor = cute.ceil_div(tensor_rows, (self.TILE_ROWS * self.NUM_TILES_Y)) + # How many jobs does this tensor have + jobs_in_tensor = jobs_X_in_tensor * jobs_Y_in_tensor + # Which job (1D index) does this CTA start from, which is also their worker ID + first_job_id = Int32(bidx) + # gdx is how many CTAs are assigned to this tensor + job_stride = Int32(gdx) + # If my first job is already beyond the tensor's last job, I have no work to do + if first_job_id >= jobs_in_tensor: + has_work = Boolean(False) + else: + # This tensor is empty, so this CTA has no work to do + has_work = Boolean(False) + + partitions = (tXsX, tXgX, tXsA, tXgA, tXsO_row, tXgO_row, tXsO_col, tXgO_col) + atoms = (tma_atom_x, tma_atom_act, tma_atom_out_row, tma_atom_out_col) + + if has_work: + # For single tensor case, each CTA only processes one job + if cutlass.const_expr(cfg.IS_SINGLE_TENSOR): + # For single tensor case we don't use tensor descriptors + descs = (None, None, None, None) + + # Rowwise scales span the whole group: members are stacked on 128-row + # boundaries, so each member's swizzled tiles follow the previous member's. + row_scales = None + if cutlass.const_expr(cfg.ROWWISE): + row_scales = self._rowwise_scales(mS_row, Int64(0), tensor_rows, tensor_cols) + + # Colwise scales require special handling for swizzled scales + col_scales = None + col_scale_row_start = job_start_row + col_scale_rows = tensor_rows + if cutlass.const_expr(cfg.COLWISE): + col_scale_base = Int64(0) + if cutlass.const_expr(cfg.WITH_GEMM_SWIZZLED_SCALES): + # Colwise swizzled scale indices restart at each member and depend + # on its rows (process_colwise_stage), so address the member that + # owns this job. + member_rows = Int32(0) + member_row_start = Int32(0) + if cutlass.const_expr(cfg.SHAPE_REP == SAME_BOTH_DIMS): + member_rows = tensor_rows // Int32(num_tensors) + member_row_start = job_start_row // member_rows * member_rows + else: + member_id = self._find_tensor_from_offsets( + mOffsets, + num_tensors, + Int64(job_start_row) * Int64(tensor_cols), + ) + member_rows = Int32(mFirstDims[member_id]) + member_row_start = Int32( + Int64(mOffsets[member_id]) // Int64(tensor_cols) + ) + col_scale_base = ( + Int64(member_row_start) + * Int64(cute.round_up(tensor_cols, 128)) + // MXFP8_BLOCK_SCALING_SIZE + ) + col_scale_row_start = job_start_row - member_row_start + col_scale_rows = member_rows + col_scales = self._colwise_scales( + mS_col, col_scale_base, col_scale_rows, tensor_cols + ) + + cute.arch.sync_threads() + + self._process_job( + job_start_row, + job_id_X * self.TILE_COLS, + tensor_rows, + tensor_cols, + row_scales, + col_scales, + col_scale_row_start, + col_scale_rows, + mWorkspace, + sDbias, + descs, + tmap, + warp_idx, + tidx, + sX, + sActInput, + sO_row, + sO_col, + partitions, + atoms, + mainloop_pipeline, + prod_state, + cons_state, + ) + else: + # For non-single tensor case, we use persistent kernel so each CTA keeps processing jobs until none is left + desc_x = tmap.get_tensormap_ptr(mTensormaps[(tensor_id, 0, None)].iterator) + desc_out_row = tmap.get_tensormap_ptr(mTensormaps[(tensor_id, 1, None)].iterator) + desc_out_col = tmap.get_tensormap_ptr(mTensormaps[(tensor_id, 2, None)].iterator) + desc_act = tmap.get_tensormap_ptr( + mTensormaps[(tensor_id, ACT_INPUT_SLOT, None)].iterator + ) + + if tidx == 0: + tmap.fence_tensormap_update(desc_x) + if cutlass.const_expr(cfg.WITH_DACT): + tmap.fence_tensormap_update(desc_act) + if cutlass.const_expr(cfg.ROWWISE): + tmap.fence_tensormap_update(desc_out_row) + if cutlass.const_expr(cfg.COLWISE): + tmap.fence_tensormap_update(desc_out_col) + descs = (desc_x, desc_act, desc_out_row, desc_out_col) + + # This tensor's scales start at tensor_base / 32 in both directions. + scale_base = tensor_base // Int64(MXFP8_BLOCK_SCALING_SIZE) + row_scales = None + if cutlass.const_expr(cfg.ROWWISE): + row_scales = self._rowwise_scales(mS_row, scale_base, tensor_rows, tensor_cols) + col_scales = None + if cutlass.const_expr(cfg.COLWISE): + col_scales = self._colwise_scales(mS_col, scale_base, tensor_rows, tensor_cols) + + # Make sure all threads see the updated descriptors and scales before processing any jobs. + cute.arch.sync_threads() + + # Grid-stride over this tensor's own jobs; the descriptors never change. + job_id = first_job_id + job_finished = Boolean(False) + while not job_finished: + job_id_Y_in_tensor = job_id // jobs_X_in_tensor + job_id_X_in_tensor = job_id % jobs_X_in_tensor + job_start_row_in_tensor = job_id_Y_in_tensor * ( + self.TILE_ROWS * self.NUM_TILES_Y + ) + job_start_col = job_id_X_in_tensor * (self.TILE_COLS * self.NUM_TILES_X) + self._process_job( + job_start_row_in_tensor, + job_start_col, + tensor_rows, + tensor_cols, + row_scales, + col_scales, + job_start_row_in_tensor, + tensor_rows, + mWorkspace, + sDbias, + descs, + tmap, + warp_idx, + tidx, + sX, + sActInput, + sO_row, + sO_col, + partitions, + atoms, + mainloop_pipeline, + prod_state, + cons_state, + ) + # Find the next job to process + job_id = job_id + job_stride + if job_id >= jobs_in_tensor: + job_finished = Boolean(True) + + # Drain every TMA store before the CTA releases its shared-memory source buffers. + if warp_idx == 0: + cute.arch.cp_async_bulk_wait_group(0, read=False) + cute.arch.sync_threads() + + def _issue_load( + self, + pipeline_obj: pipeline.PipelineTmaAsync, + prod_state: pipeline.PipelineState, + tile_y: Int32, + tile_x: Int32, + atoms: tuple[Optional[cute.CopyAtom], ...], + partitions: tuple[Optional[cute.Tensor], ...], + tmap: TensorMapManager, + descs: tuple[Optional[cute.Pointer], ...], + ) -> None: + """Emit the 32x128 TMA load(s) of one stage into the current pipeline buffer. + + Caller gates this on warp 0 and advances `prod_state` afterwards -- the advance + must happen outside the gate or the mutated SSA values stay trapped in the scf.if. + """ + tma_atom_x, tma_atom_act, _, _ = atoms + tXsX, tXgX, tXsA, tXgA, _, _, _, _ = partitions + desc_x, desc_act, _, _ = descs + # Wait for the consumer to finish using this SMEM buffer + pipeline_obj.producer_acquire(prod_state) + barrier = pipeline_obj.producer_get_barrier(prod_state) + loads = [(tma_atom_x, tXgX, tXsX, desc_x)] + if cutlass.const_expr(self.cfg.WITH_DACT): + loads.append((tma_atom_act, tXgA, tXsA, desc_act)) + for atom, tXg, tXs, desc in loads: + if cutlass.const_expr(self.cfg.IS_SINGLE_TENSOR): + cute.copy( + atom, + tXg[(None, (tile_y, tile_x))], + tXs[(None, prod_state.index)], + tma_bar_ptr=barrier, + ) + else: + # Every member shares tXg's tile-coordinate arithmetic (the coefficients are + # just the tile size); tma_desc_ptr supplies this member's geometry. + cute.copy( + atom, + tXg[(None, (tile_y, tile_x))], + tXs[(None, prod_state.index)], + tma_bar_ptr=barrier, + tma_desc_ptr=tmap.get_tensormap_ptr(desc, cute.AddressSpace.generic), + ) + # Notify the consumer that this SMEM buffer is ready for consumption + pipeline_obj.producer_commit(prod_state) + + @cute.jit + def _process_job( + self, + job_start_row: Int32, # Row offset of this job (global for single-tensor, else tensor-local) + job_start_col: Int32, # Column offset of this job within the tensor + rows: Int32, # Rows of the rowwise-scale view (the group for single-tensor, else the tensor) + cols: Int32, # Number of columns in this tensor + row_scales: Optional[ + cute.Tensor + ], # Rowwise scales tiled per stage, rows counted like job_start_row + col_scales: Optional[cute.Tensor], # Colwise scales tiled per stage + col_scale_row_start: Int32, # Row of this job in the colwise-scale view + col_scale_rows: Int32, # Rows of the colwise-scale view + mWorkspace: Optional[cute.Tensor], # f32 partial dbias workspace (WITH_DBIAS) + sDbias: Optional[ + cute.Tensor + ], # SMEM buffer for the rowwise dbias reduction (rowwise-only dbias) + descs: tuple[ + Optional[cute.Pointer], ... + ], # Per-tensor descriptors (x, act, out_row, out_col), None if single-tensor + tmap: TensorMapManager, # TensorMapManager for managing TMA descriptors + warp_idx: Int32, + tidx: Int32, + sX: cute.Tensor, # SMEM input ring + sActInput: Optional[cute.Tensor], # SMEM activation input ring (WITH_DACT) + sO_row: Optional[cute.Tensor], # SMEM rowwise output ring + sO_col: Optional[cute.Tensor], # SMEM colwise output ring + partitions: tuple[Optional[cute.Tensor], ...], # TMA partitions (x, act, out_row, out_col) + atoms: tuple[Optional[cute.CopyAtom], ...], # TMA atoms (x, act, out_row, out_col) + mainloop_pipeline: cutlass.pipeline.PipelineTmaAsync, + prod_state: pipeline.PipelineState, + cons_state: pipeline.PipelineState, + ) -> None: + """Quantize a job with one continuous pipeline across its column and row tiles.""" + cfg = self.cfg + _, _, tma_atom_out_row, tma_atom_out_col = atoms + _, _, _, _, tXsO_row, tXgO_row, tXsO_col, tXgO_col = partitions + _, _, desc_out_row, desc_out_col = descs + + job_tile_Y = job_start_row // self.TILE_ROWS + job_tile_X = job_start_col // self.TILE_COLS + col_scale_tile_Y = col_scale_row_start // self.TILE_ROWS + tiles_X = cutlass.min( + Int32(self.NUM_TILES_X), cute.ceil_div(cols - job_start_col, self.TILE_COLS) + ) + num_tiles = tiles_X * self.NUM_TILES_Y + + # Fill every buffer up front, then issue one more each time a stage is consumed. + for prologue_stage in cutlass.range_constexpr(self.PIPELINE_DEPTH): + if warp_idx == 0: + self._issue_load( + mainloop_pipeline, + prod_state, + job_tile_Y + prologue_stage % self.NUM_TILES_Y, + job_tile_X + prologue_stage // self.NUM_TILES_Y, + atoms, + partitions, + tmap, + descs, + ) + prod_state.advance() + + for tile_X in cutlass.range(tiles_X, unroll=1): + tile_id_X = job_tile_X + tile_X + tile_start_col = job_start_col + tile_X * self.TILE_COLS + partial_dbias = Float32(0.0) + dbias_row_acc = None + if cutlass.const_expr(self.DBIAS_IN_ROWWISE): + dbias_row_acc = cute.make_rmem_tensor( + layout_or_shape=cute.make_layout((MXFP8_BLOCK_SCALING_SIZE,), stride=(1,)), + dtype=Float32, + ) + for c in cutlass.range_constexpr(MXFP8_BLOCK_SCALING_SIZE): + dbias_row_acc[c] = Float32(0.0) + + for tile_Y in cutlass.range_constexpr(self.NUM_TILES_Y): + stage = tile_X * self.NUM_TILES_Y + tile_Y + # Wait for at most DEPTH-1 iters on the fly, which means the the last DEPTH iter has finished + # so we can reuse its SMEM output buffer + # (input buffer is managed by the producer and consumer pipeline states) + if warp_idx == 0: + cute.arch.cp_async_bulk_wait_group(self.PIPELINE_DEPTH - 1, read=True) + # Wait for this stage's input buffer to be filled by the producer + mainloop_pipeline.consumer_wait(cons_state) + cute.arch.sync_threads() + sX_tile = sX[(None, cons_state.index)] + sAct_tile = None + if cutlass.const_expr(cfg.WITH_DACT): + sAct_tile = sActInput[(None, cons_state.index)] + row_tile = job_tile_Y + tile_Y + + if cutlass.const_expr(cfg.COLWISE): + _, partial_dbias = quantize_colwise_mxfp8( + sX_tile, + sAct_tile, + sO_col[(None, cons_state.index)], + cute.flatten(col_scales[(None, (col_scale_tile_Y + tile_Y, tile_id_X))]), + cfg.MAX_NORM_RCP, + (col_scale_tile_Y + tile_Y) * self.TILE_ROWS, + tile_start_col, + col_scale_rows, + cols, + ACTIVATION=cfg.ACTIVATION, + DTYPE=cfg.DTYPE, + FP8_DTYPE=cfg.FP8_DTYPE, + SWIZZLE=cfg.WITH_GEMM_SWIZZLED_SCALES, + TILE_X=self.TILE_COLS, + TILE_Y=self.TILE_ROWS, + WITH_ACT=cfg.WITH_ACT, + WITH_DACT=cfg.WITH_DACT, + WITH_DBIAS=self.DBIAS_IN_COLWISE, + CACHE_ACTIVATION=self.CACHE_ACTIVATION, + ZERO_OOB_SCALES=True, + dbias_init=partial_dbias, + ) + if cutlass.const_expr(self.CACHE_ACTIVATION): + # The rowwise pass reads the activation the colwise pass cached in sX. + cute.arch.sync_threads() + + if cutlass.const_expr(cfg.ROWWISE): + quantize_rowwise_mxfp8( + sX_tile, + None if self.CACHE_ACTIVATION else sAct_tile, + sO_row[(None, cons_state.index)], + cute.flatten(row_scales[(None, (row_tile, tile_id_X))]), + cfg.MAX_NORM_RCP, + row_tile * self.TILE_ROWS, + tile_start_col, + rows, + cols, + ACTIVATION=None if self.CACHE_ACTIVATION else cfg.ACTIVATION, + DTYPE=cfg.DTYPE, + FP8_DTYPE=cfg.FP8_DTYPE, + TILE_X=self.TILE_COLS, + TILE_Y=self.TILE_ROWS, + WAVES=self.WAVES, + THREADS_PER_BANK=self.THREADS_PER_BANK, + PACK_SIZE=self.PACK_SIZE, + WITH_ACT=cfg.WITH_ACT and not self.CACHE_ACTIVATION, + WITH_DACT=cfg.WITH_DACT and not self.CACHE_ACTIVATION, + WITH_DBIAS=self.DBIAS_IN_ROWWISE, + dbias_acc=dbias_row_acc, + ZERO_OOB_SCALES=True, + ) + + # Force consumer's write to SMEM to be visible to TMA stores later + cute.arch.fence_proxy("async.shared", space="cta") + # Only after everyone finishes computation then this stage can be considered as "consumed" + cute.arch.sync_threads() + # I'm done with my input SMEM buffer, so the producer can write the next stage's data into it + mainloop_pipeline.consumer_release(cons_state) + + # I just freed my input SMEM buffer (stage), so the producer now can use it for writing + # (stage+DEPTH) stage's data if that stage exists + if stage + self.PIPELINE_DEPTH < num_tiles: + if warp_idx == 0: + self._issue_load( + mainloop_pipeline, + prod_state, + job_tile_Y + (stage + self.PIPELINE_DEPTH) % self.NUM_TILES_Y, + job_tile_X + (stage + self.PIPELINE_DEPTH) // self.NUM_TILES_Y, + atoms, + partitions, + tmap, + descs, + ) + prod_state.advance() + + # Write result to GMEM via TMA + if warp_idx == 0: + stores = [] + if cutlass.const_expr(cfg.ROWWISE): + stores.append((tma_atom_out_row, tXsO_row, tXgO_row, desc_out_row)) + if cutlass.const_expr(cfg.COLWISE): + stores.append((tma_atom_out_col, tXsO_col, tXgO_col, desc_out_col)) + for atom, tXs, tXg, desc in stores: + if cutlass.const_expr(cfg.IS_SINGLE_TENSOR): + cute.copy( + atom, + tXs[(None, cons_state.index)], + tXg[(None, (row_tile, tile_id_X))], + ) + else: + cute.copy( + atom, + tXs[(None, cons_state.index)], + tXg[(None, (row_tile, tile_id_X))], + tma_desc_ptr=tmap.get_tensormap_ptr( + desc, cute.AddressSpace.generic + ), + ) + # Commit all TMA operations of this iteration + cute.arch.cp_async_bulk_commit_group() + + cons_state.advance() + + if cutlass.const_expr(cfg.WITH_DBIAS): + if cutlass.const_expr(self.DBIAS_IN_ROWWISE): + partial_dbias = reduce_rowwise_dbias( + sDbias, + tidx, + dbias_row_acc, + self.TILE_ROWS, + self.TILE_COLS, + self.PACK_SIZE, + self.THREADS_PER_BANK, + ) + # All threads must finish reading before the next job rewrites the buffer. + cute.arch.sync_threads() + + # A job has TILE_ROWS * NUM_TILES_Y rows, and dbias reduces them to one row + dbias_row = job_start_row // (self.TILE_ROWS * self.NUM_TILES_Y) + dbias_col = tile_start_col + tidx + if dbias_col < cols: + mWorkspace[(dbias_row, dbias_col)] = partial_dbias + + +def compile_cutedsl_function_from_cfg(cfg: MXFP8GroupQuantizeConfig): + """Return the compiled CuTeDSL function object for the given grouped config.""" + # CUDA requires the group's first logical dim to be a multiple of 128 (and each + # tensor's rows likewise). VARYING_BOTH_DIMS is the exception: its logical shape is + # [1, total]. The last dim only needs the 16-byte TMA row alignment; a partial 32-element + # scale block at the end of a row is zero-filled by TMA, as in the CUDA kernel. + if cfg.SHAPE_REP == VARYING_BOTH_DIMS: + sym_M = cute.sym_int32() + else: + sym_M = cute.sym_int32(divisibility=128) + sym_N = cute.sym_int32(divisibility=SYM_N_DIVISIBILITY) + logical_shape = (sym_M, sym_N) + + out_dtype = cfg.FP8_DTYPE + scale_dtype = cutlass.Float8E8M0FNU + + out_col_fake = ( + cute.runtime.make_fake_compact_tensor( + out_dtype, + logical_shape, + stride_order=(1, 0), + memspace=cute.AddressSpace.gmem, + assumed_align=16, + ) + if cfg.COLWISE + else None + ) + + out_row_fake = ( + cute.runtime.make_fake_compact_tensor( + out_dtype, + logical_shape, + stride_order=(1, 0), + memspace=cute.AddressSpace.gmem, + assumed_align=16, + ) + if cfg.ROWWISE + else None + ) + + # The kernel only takes the base address of the scale buffers (per-tensor strides + # are derived from cols), so their fake shape is a flat 1D byte run. + tensormaps_fake = cute.runtime.make_fake_compact_tensor( + cutlass.Int64, + (cute.sym_int32(), NUM_WORKSPACE_SLOTS, BYTES_PER_TENSORMAP // 8), + stride_order=(2, 1, 0), + memspace=cute.AddressSpace.gmem, + assumed_align=128, + ) + # The cast-noop flag is an always-present f32 pointer instead of an optional tensor, so + # that one compiled kernel serves both an absent and a present flag (noop_flag_is_set). + noop_fake = cute.runtime.nullptr(Float32, mem_space=cute.AddressSpace.gmem, assumed_align=4) + act_input_fake = ( + cute.runtime.make_fake_compact_tensor( + cfg.DTYPE, + logical_shape, + stride_order=(1, 0), + memspace=cute.AddressSpace.gmem, + assumed_align=16, + ) + if cfg.WITH_DACT + else None + ) + workspace_fake = ( + cute.runtime.make_fake_compact_tensor( + Float32, + (cute.sym_int32(), cute.sym_int32()), + stride_order=(1, 0), + memspace=cute.AddressSpace.gmem, + assumed_align=4, + ) + if cfg.WITH_DBIAS + else None + ) + + first_dims_fake = ( + cute.runtime.make_fake_compact_tensor( + cutlass.Int64, + (cute.sym_int32(),), + stride_order=(0,), + memspace=cute.AddressSpace.gmem, + assumed_align=8, + ) + if cfg.SHAPE_REP in (VARYING_FIRST_DIM, VARYING_BOTH_DIMS) + else None + ) + + last_dims_fake = ( + cute.runtime.make_fake_compact_tensor( + cutlass.Int64, + (cute.sym_int32(),), + stride_order=(0,), + memspace=cute.AddressSpace.gmem, + assumed_align=8, + ) + if cfg.SHAPE_REP in (VARYING_LAST_DIM, VARYING_BOTH_DIMS) + else None + ) + + sm_count = HardwareInfo().get_device_multiprocessor_count() + kernel_obj = MXFP8GroupQuantizeKernel(cfg, sm_count) + return cute.compile( + kernel_obj, + cute.runtime.make_fake_compact_tensor( # mX + cfg.DTYPE, + logical_shape, + stride_order=(1, 0), + memspace=cute.AddressSpace.gmem, + assumed_align=16, + ), + out_row_fake, # mO_row + out_col_fake, # mO_col + cute.runtime.make_fake_compact_tensor( # mS_row + scale_dtype, + (cute.sym_int32(),), + stride_order=(0,), + memspace=cute.AddressSpace.gmem, + assumed_align=4, + ), + cute.runtime.make_fake_compact_tensor( # mS_col + scale_dtype, + (cute.sym_int32(),), + stride_order=(0,), + memspace=cute.AddressSpace.gmem, + assumed_align=4, + ), + cute.runtime.make_fake_compact_tensor( # mOffsets + cutlass.Int64, + (cute.sym_int32(),), + stride_order=(0,), + memspace=cute.AddressSpace.gmem, + assumed_align=8, + ), + first_dims_fake, # mFirstDims + last_dims_fake, # mLastDims + tensormaps_fake, # mTensormaps + noop_fake, # mNoop + act_input_fake, # mActInput + workspace_fake, # mWorkspace + cute.runtime.make_fake_stream(), + options="--enable-tvm-ffi", + ) + + +def get_mxfp8_group_quantization_function( + fn_name: str, + dtype: str, + fp8_dtype: str, + rowwise: bool, + colwise: bool, + shape_rep: str, + with_gemm_swizzled_scales: bool, + with_dbias: bool, + with_dact: bool, + with_act: bool, + activation: str, +) -> bool: + """Compile the grouped MXFP8 quantize kernel for this config and register it in the TVM-FFI + global registry under EXACTLY `fn_name` (the key the C++ dispatcher built; Python treats it as + an opaque name). Returns True if a kernel is successfully registered under `fn_name` (the C++ + side then fetches it with GetGlobal(fn_name)); False if the config is unsupported, so the caller + caches the negative result and falls back to the CUDA C++ grouped kernel. + """ + try: + # Already registered (e.g. by a prior call) -> supported. + if tvm_ffi.get_global_func(fn_name, allow_missing=True) is not None: + return True + + major, minor = device_compute_capability() + if major < 10: + logger.warning( + "CuTeDSL MXFP8 backend requires compute capability >= 10.0 (Blackwell), " + "but detected %d.%d; falling back to the CUDA C++ kernel.", + major, + minor, + ) + return False + + try: + cfg = MXFP8GroupQuantizeConfig( + dtype=dtype, + fp8_dtype=fp8_dtype, + rowwise=rowwise, + colwise=colwise, + shape_rep=shape_rep, + with_gemm_swizzled_scales=with_gemm_swizzled_scales, + with_dbias=with_dbias, + with_dact=with_dact, + with_act=with_act, + activation=activation, + ) + except ValueError as e: + logger.warning( + "CuTeDSL grouped MXFP8 backend does not support this config, " + "falling back to the CUDA C++ kernel: %s", + e, + ) + return False + + logger.debug("Compiling CuTeDSL grouped MXFP8 quantization kernel for %s", cfg) + compiled = compile_cutedsl_function_from_cfg(cfg) + # Register the native TVM-FFI function rather than its Python argument-parsing wrapper; + # see get_mxfp8_quantization_function for why. + native = getattr(compiled, "__tvm_ffi_object__", lambda: None)() + tvm_ffi.register_global_func( + fn_name, native if native is not None else compiled, override=True + ) + return True + except Exception as e: # pylint: disable=broad-exception-caught + logger.error( + "CuTeDSL grouped MXFP8 kernel compilation & registration failed, falling back to the" + " CUDA C++ kernel: %s", + e, + ) + # Unconditionally fallback to CUDA path because we can't tell if this exception is + # transient or permanent. + return False diff --git a/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py b/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py index 824ed18aa59..c267638bbd7 100644 --- a/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py +++ b/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py @@ -175,6 +175,9 @@ def quantize_rowwise_mxfp8( WITH_DACT: cutlass.Constexpr[bool] = False, WITH_DBIAS: cutlass.Constexpr[bool] = False, dbias_acc: Optional[cute.Tensor] = None, # only needed when WITH_DBIAS is True + # Write 0 to the scales of col-blocks past N instead of skipping them, so the padding of the + # scale row is zeroed (mirrors group_quantize_mxfp8.cuh). + ZERO_OOB_SCALES: cutlass.Constexpr[bool] = False, ): """Quantize one SMEM tile rowwise to MXFP8 (per-row 32-elt block scales); returns the tile amax.""" tidx, _, _ = cute.arch.thread_idx() @@ -375,6 +378,11 @@ def quantize_rowwise_mxfp8( scale_col_first_elt = tile_col_start + (tidx % CTA_THREADS_X) * MXFP8_BLOCK_SCALING_SIZE if scale_row < M and scale_col_first_elt < N: mS_row_stage[(tidx // CTA_THREADS_X, tidx % CTA_THREADS_X)] = biased_exp_r + if cutlass.const_expr(ZERO_OOB_SCALES): + if scale_row < M and scale_col_first_elt >= N: + mS_row_stage[(tidx // CTA_THREADS_X, tidx % CTA_THREADS_X)] = Uint8(0).bitcast( + Float8E8M0FNU + ) inv_scale_r = exp2f_rcp(biased_exp_r) # f32 reciprocal of the scale scale_2x = pack_f32x2(inv_scale_r, inv_scale_r) @@ -421,6 +429,12 @@ def quantize_colwise_mxfp8( WITH_DACT: cutlass.Constexpr[bool] = False, WITH_DBIAS: cutlass.Constexpr[bool] = False, CACHE_ACTIVATION: cutlass.Constexpr[bool] = False, # cache post-activation values to sX_tile + # Write 0 to the scales of columns past N instead of skipping them, so the padding of the + # scale row is zeroed (mirrors group_quantize_mxfp8.cuh). + ZERO_OOB_SCALES: cutlass.Constexpr[bool] = False, + # Running column sum to continue the dbias accumulation from, so that consecutive tiles add + # their elements in row order as the CUDA kernel does. Starts from 0 when None. + dbias_init: Optional[Float32] = None, ): """Quantize one SMEM tile colwise to MXFP8 (per-column 32-elt block scales); returns (amax, dbias_partial).""" tidx, _, _ = cute.arch.thread_idx() @@ -441,7 +455,7 @@ def quantize_colwise_mxfp8( FUSE_RELU = cutlass.const_expr(ACTIVATION == "relu") and not WITH_DBIAS # Keep input in half precision format if possible USE_HALF_PRECISION = is_packed16(DTYPE) and (ACTIVATION is None or FUSE_RELU) - dbias_partial = Float32(0.0) + dbias_partial = Float32(0.0) if cutlass.const_expr(dbias_init is None) else dbias_init if cutlass.const_expr(USE_HALF_PRECISION): max_scalar = max_scalar_f16 if DTYPE is cutlass.Float16 else max_scalar_bf16 @@ -541,6 +555,12 @@ def quantize_colwise_mxfp8( mS_col_stage[(0, tidx % 32, tidx // 32)] = biased_exp_c else: mS_col_stage[(0, tidx)] = biased_exp_c + if cutlass.const_expr(ZERO_OOB_SCALES): + if tile_row_start < M and scale_col >= N: + if cutlass.const_expr(SWIZZLE): + mS_col_stage[(0, tidx % 32, tidx // 32)] = Uint8(0).bitcast(Float8E8M0FNU) + else: + mS_col_stage[(0, tidx)] = Uint8(0).bitcast(Float8E8M0FNU) inv_scale_c = exp2f_rcp(biased_exp_c) # cvt.rn.satfinite can be vectorized to convert 2 f32 to 2 fp8 in one instruction @@ -795,6 +815,50 @@ def quantize_bidimensional_mxfp8_swizzled( cute.autovec_copy(rO_col, tXsO_col) +@cute.jit +def reduce_rowwise_dbias( + sDbias: cute.Tensor, + tidx: Int32, + rowwise_dbias_acc: cute.Tensor, + TILE_ROWS: cutlass.Constexpr[int], + TILE_COLS: cutlass.Constexpr[int], + PACK_SIZE: cutlass.Constexpr[int], + THREADS_PER_BANK: cutlass.Constexpr[int], +): + """Reduce per-thread rowwise partial sums to one sum per column using shared memory.""" + _, tv_layout_dbias_write = cute.make_layout_tv( + thr_layout=cute.make_layout( + (TILE_ROWS, TILE_COLS // MXFP8_BLOCK_SCALING_SIZE), + stride=(TILE_COLS // MXFP8_BLOCK_SCALING_SIZE, 1), + ), + val_layout=cute.make_layout( + (1, MXFP8_BLOCK_SCALING_SIZE), stride=(MXFP8_BLOCK_SCALING_SIZE, 1) + ), + ) + sDbias_write = cute.composition(sDbias, tv_layout_dbias_write) + # Undo the bank-conflict rotation used when accumulating the per-thread sums. + bank_group = (tidx % THREADS_PER_WARP) // THREADS_PER_BANK + offset = bank_group * PACK_SIZE + for w in cutlass.range_constexpr(MXFP8_BLOCK_SCALING_SIZE // PACK_SIZE): + start = (w * PACK_SIZE + offset) % MXFP8_BLOCK_SCALING_SIZE + for i in cutlass.range_constexpr(PACK_SIZE): + # All threads write their per-thread partial sum results to the shared buffer. + sDbias_write[(tidx, start + i)] = rowwise_dbias_acc[w * PACK_SIZE + i] + cute.arch.sync_threads() + # All threads reduce the cross-thread partial sums to the per-block partial sum. + _, tv_layout_dbias_reduce = cute.make_layout_tv( + thr_layout=cute.make_layout((1, TILE_COLS), stride=(TILE_COLS, 1)), + val_layout=cute.make_layout((TILE_ROWS, 1), stride=(1, 1)), + ) + sDbias_reduce = cute.composition(sDbias, tv_layout_dbias_reduce) + # make_layout_tv yields a (thread, value) layout: thread=tidx -> column tidx, + # value=i -> row i. So index [tidx, i] (thread first), summing the column's rows. + block_dbias = Float32(0.0) + for i in cutlass.range_constexpr(TILE_ROWS): + block_dbias += sDbias_reduce[tidx, i] + return block_dbias + + @cute.jit def noop_flag_is_set(mNoop: cute.Pointer) -> Boolean: """Whether the cast_noop flag says this quantization is a no-op and must be skipped. @@ -1625,41 +1689,15 @@ class DbiasStorage: sDbias = dbias_storage.sDbias.get_tensor( cute.make_layout((self._TILE_ROWS, self._TILE_COLS), stride=(DBIAS_BUFF_WIDTH, 1)), ) - _, tv_layout_dbias_write = cute.make_layout_tv( - thr_layout=cute.make_layout( - (self._TILE_ROWS, self._TILE_COLS // MXFP8_BLOCK_SCALING_SIZE), - stride=(self._TILE_COLS // MXFP8_BLOCK_SCALING_SIZE, 1), - ), - val_layout=cute.make_layout( - (1, MXFP8_BLOCK_SCALING_SIZE), stride=(MXFP8_BLOCK_SCALING_SIZE, 1) - ), + return reduce_rowwise_dbias( + sDbias, + tidx, + rowwise_dbias_acc, + self._TILE_ROWS, + self._TILE_COLS, + self._PACK_SIZE, + self._THREADS_PER_BANK, ) - sDbias_write = cute.composition(sDbias, tv_layout_dbias_write) - # Each thread start reading from the specfic bank based on its thread ID so they can do their best to access different banks - # to avoid bank conflict. - bank_group = (tidx % THREADS_PER_WARP) // self._THREADS_PER_BANK - # The offset this thread should start reading from based on what's its first bank to access. - offset = bank_group * self._PACK_SIZE - for w in cutlass.range_constexpr( - self._WAVES - ): # Each thread starts from this offset when writing into SMEM to avoid bank conflict - start = (w * self._PACK_SIZE + offset) % MXFP8_BLOCK_SCALING_SIZE - for i in cutlass.range_constexpr(self._PACK_SIZE): - # All threads write their per-thread partial sum results to the shared buffer. - sDbias_write[(tidx, start + i)] = rowwise_dbias_acc[w * self._PACK_SIZE + i] - cute.arch.sync_threads() - # All threads reduce the cross-thread partial sums to the per-block partial sum. - _, tv_layout_dbias_reduce = cute.make_layout_tv( - thr_layout=cute.make_layout((1, self._TILE_COLS), stride=(self._TILE_COLS, 1)), - val_layout=cute.make_layout((self._TILE_ROWS, 1), stride=(1, 1)), - ) - sDbias_reduce = cute.composition(sDbias, tv_layout_dbias_reduce) - # make_layout_tv yields a (thread, value) layout: thread=tidx -> column tidx, - # value=i -> row i. So index [tidx, i] (thread first), summing the column's rows. - block_dbias = Float32(0.0) - for i in cutlass.range_constexpr(self._TILE_ROWS): - block_dbias += sDbias_reduce[tidx, i] - return block_dbias @cute.jit def _amax_epilogue( diff --git a/transformer_engine/common/cast/dispatch/quantize.cuh b/transformer_engine/common/cast/dispatch/quantize.cuh index d31f18f0d8b..b4394975dc3 100644 --- a/transformer_engine/common/cast/dispatch/quantize.cuh +++ b/transformer_engine/common/cast/dispatch/quantize.cuh @@ -32,6 +32,7 @@ #include "../nvfp4/quantize_transpose_nvfp4.cuh" #ifdef NVTE_WITH_CUTEDSL +#include "../mxfp8/group_quantize_mxfp8_cutedsl.cuh" #include "../mxfp8/quantize_mxfp8_cutedsl.cuh" #endif @@ -524,9 +525,19 @@ void group_quantize_fwd_helper(const NVTEGroupedTensor input, NVTEGroupedTensor break; } case NVTE_MXFP8_1D_SCALING: { - mxfp8::group_quantize( - input_tensor, activations_tensor, noop_tensor, output_tensor, dbias_tensor, - workspace_tensor, &quant_config_cpp, stream); + bool quantized_with_cutedsl = false; +#ifdef NVTE_WITH_CUTEDSL + quantized_with_cutedsl = + cutedsl_backend::mxfp8_group_quantize_cutedsl( + input_tensor, activations_tensor, noop_tensor, output_tensor, dbias_tensor, + workspace_tensor, &quant_config_cpp, stream); +#endif + if (!quantized_with_cutedsl) { + mxfp8::group_quantize( + input_tensor, activations_tensor, noop_tensor, output_tensor, dbias_tensor, + workspace_tensor, &quant_config_cpp, stream); + } break; } case NVTE_BLOCK_SCALING_1D: { @@ -619,9 +630,19 @@ void group_quantize_bwd_helper(const NVTEGroupedTensor grad, const NVTEGroupedTe // Dispatch to quantization kernel depending on data format switch (scaling_mode) { case NVTE_MXFP8_1D_SCALING: { - mxfp8::group_quantize( - grad_tensor, input_tensor, noop_tensor, output_tensor, dbias_tensor, workspace_tensor, - &quant_config_cpp, stream); + bool quantized_with_cutedsl = false; +#ifdef NVTE_WITH_CUTEDSL + quantized_with_cutedsl = + cutedsl_backend::mxfp8_group_quantize_cutedsl( + grad_tensor, input_tensor, noop_tensor, output_tensor, dbias_tensor, workspace_tensor, + &quant_config_cpp, stream); +#endif + if (!quantized_with_cutedsl) { + mxfp8::group_quantize( + grad_tensor, input_tensor, noop_tensor, output_tensor, dbias_tensor, workspace_tensor, + &quant_config_cpp, stream); + } break; } case NVTE_BLOCK_SCALING_1D: diff --git a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh new file mode 100644 index 00000000000..edb674a09c0 --- /dev/null +++ b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh @@ -0,0 +1,436 @@ +/************************************************************************* + * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * + * See LICENSE for license information. + ************************************************************************/ + +#ifndef TRANSFORMER_ENGINE_COMMON_CAST_MXFP8_GROUP_QUANTIZE_MXFP8_CUTEDSL_CUH_ +#define TRANSFORMER_ENGINE_COMMON_CAST_MXFP8_GROUP_QUANTIZE_MXFP8_CUTEDSL_CUH_ + +#include +#include + +#include +#include +#include +#include +#include +#include + +#include "../../common.h" +#include "../../tvm_ffi_bridge.h" +#include "../../util/cuda_runtime.h" +#include "../../util/cutedsl_utils.h" +#include "../../utils.cuh" // ShapeRepresentation +#include "../core/common.cuh" // MAX_SUPPORTED_TENSOR_DESCRIPTORS, grouped_reduce_dbias + +namespace transformer_engine { +namespace cutedsl_backend { + +// DLTensorWrapper and TVMFFICentral live in transformer_engine::tvm_ffi_bridge. +using namespace tvm_ffi_bridge; + +inline const char *shape_rep_to_str(ShapeRepresentation shape_rep) { + switch (shape_rep) { + case ShapeRepresentation::SAME_BOTH_DIMS: + return "same_both_dims"; + case ShapeRepresentation::VARYING_FIRST_DIM: + return "varying_first_dim"; + case ShapeRepresentation::VARYING_LAST_DIM: + return "varying_last_dim"; + default: + return "varying_both_dims"; + } +} + +struct MXFP8GroupQuantConfig { + static constexpr const char *kEntrypointName = "get_mxfp8_group_quantization_function"; + + DType dtype; // The input format + DType fp8_dtype; // The fp8 output format + bool rowwise; // If quantize rowwisely + bool colwise; // If quantize columnwisely + ShapeRepresentation shape_rep; // How the member shapes vary across the group + bool swizzled; // If the scales are written in the GEMM-swizzled layout + bool with_dbias; // If the partial dbias is computed (via the workspace tensor) + bool with_dact; // If an activation derivative operation is fused + bool with_act; // If an activation operation is fused + Activation activation = Activation::kNone; + uint32_t sm_arch = static_cast(cuda::sm_arch()); + + // Bit layout: dtype [3:0] (4 used/reserved), fp8_dtype [7:4] (4 used/reserved), + // flags [13:8] (6 used), shape_rep [15:14] (2 used), activation [21:16] (6 used/reserved), + // and SM architecture [30:22] (9 used/reserved). Bit 31 is unused. + uint32_t to_id() const { + static_assert(static_cast(DType::kNumTypes) <= 16, + "DType no longer fits in the 4 bits to_id() gives it."); + static_assert(ShapeRepresentation::VARYING_BOTH_DIMS < 4, + "ShapeRepresentation no longer fits in the 2 bits to_id() gives it."); + static_assert(static_cast(Activation::kNumTypes) <= 64, + "Activation no longer fits in the 6 bits to_id() gives it."); + NVTE_CHECK(sm_arch < 512, "SM architecture no longer fits in the 9 bits to_id() gives it."); + return static_cast(dtype) | (static_cast(fp8_dtype) << 4) | + (static_cast(rowwise) << 8) | (static_cast(colwise) << 9) | + (static_cast(swizzled) << 10) | (static_cast(with_dbias) << 11) | + (static_cast(with_dact) << 12) | (static_cast(with_act) << 13) | + (static_cast(shape_rep) << 14) | (static_cast(activation) << 16) | + (sm_arch << 22); + } + + std::optional get_kernel() const { + static TVMFFIConfigCache &cache = TVMFFIConfigCache::create(); + return cache.get_or_load(*this); + } + + // Globally unique TVM-FFI registry key used when the CuTeDSL function is + // compiled and registered on a cache miss. + std::string to_key() const { + std::string key; + // longest: cutedsl_group_mxfp8_smXXX_BFloat16_Float8E4M3_1_1_varying_first_dim_1_1_1_0_dqgelu + key.reserve(96); + key.append("cutedsl_group_mxfp8_sm") + .append(std::to_string(sm_arch)) + .append("_") + .append(to_string(dtype)) + .append("_") + .append(to_string(fp8_dtype)) + .append("_") + .append(rowwise ? "1" : "0") + .append("_") + .append(colwise ? "1" : "0") + .append("_") + .append(shape_rep_to_str(shape_rep)) + .append("_") + .append(swizzled ? "1" : "0") + .append("_") + .append(with_dbias ? "1" : "0") + .append("_") + .append(with_dact ? "1" : "0") + .append("_") + .append(with_act ? "1" : "0") + .append("_") + .append(activation_to_str(activation)); + return key; + } + + bool retrieve_func_from_python(const std::string &fn_name) const { + auto entrypoint = tvm::ffi::Function::GetGlobal(kEntrypointName); + if (!entrypoint.has_value()) { + return false; + } + tvm::ffi::Any result = + (*entrypoint)(tvm::ffi::String(fn_name), tvm::ffi::String(to_string(dtype)), + tvm::ffi::String(to_string(fp8_dtype)), rowwise, colwise, + tvm::ffi::String(shape_rep_to_str(shape_rep)), swizzled, with_dbias, + with_dact, with_act, tvm::ffi::String(activation_to_str(activation))); + return result.try_cast().value_or(false); + } +}; + +// Descriptor slots per group member: input, rowwise output, colwise output, activation input, +// plus one carrying (rows, cols, base_elts). Mirrors NUM_WORKSPACE_SLOTS / BYTES_PER_TENSORMAP +// in CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py. +constexpr size_t kGroupTensorMapSlots = 5; +constexpr size_t kInt64PerTensorMap = 128 / sizeof(int64_t); +constexpr size_t kMaxGroupTensors = + static_cast(dispatch::common::MAX_SUPPORTED_TENSOR_DESCRIPTORS); + +struct alignas(128) GroupDescriptorWorkspace { + alignas(128) int64_t tensor_maps[kMaxGroupTensors][kGroupTensorMapSlots][kInt64PerTensorMap]; + // Stand-in for the unused offsets array in SAME_BOTH_DIMS. The kernel takes + // offsets unconditionally but does not read them for this representation. + int64_t unused_offsets[kMaxGroupTensors + 1]; +}; + +// Like `g_tensor_maps` on the CUDA path, this has internal linkage, so every translation +// unit including this header gets its own copy. It shares that path's caveat that two +// grouped quantize calls in flight on different streams would overwrite each other's +// descriptors. +static __device__ GroupDescriptorWorkspace g_group_descriptor_workspace; + +// Device address of this translation unit's workspace on the current device. The address +// is per device (each device context loads its own copy of the module), so it is cached +// per device. `static` rather than `inline`: it refers to the internal-linkage symbol above, +// so each translation unit needs its own definition and cache. +static GroupDescriptorWorkspace *group_descriptor_workspace_ptr() { + static std::vector cache(cuda::num_devices(), nullptr); + static std::vector flags(cuda::num_devices()); + const int device_id = cuda::current_device(); + NVTE_CHECK(0 <= device_id && device_id < cuda::num_devices(), "invalid CUDA device ID"); + std::call_once(flags[device_id], [&]() { + void *p = nullptr; + NVTE_CHECK_CUDA(cudaGetSymbolAddress(&p, g_group_descriptor_workspace)); + cache[device_id] = static_cast(p); + }); + return cache[device_id]; +} + +inline NVTEBasicTensor make_basic_tensor(void *dptr, DType dtype, + const std::vector &shape) { + return NVTEBasicTensor{dptr, static_cast(dtype), + nvte_make_shape(shape.data(), shape.size())}; +} + +// Signature mirrors mxfp8::group_quantize (input, act_input, noop, output, dbias, workspace, +// stream). Returns false to fall back to the CUDA kernel. +inline bool mxfp8_group_quantize_cutedsl(const MXFP8GroupQuantConfig &config, + const GroupedTensor *input_tensor, + const GroupedTensor *act_input_tensor, + const Tensor *noop_tensor, GroupedTensor *output_tensor, + GroupedTensor *dbias_tensor, Tensor *workspace_tensor, + cudaStream_t stream) { + const size_t num_tensors = input_tensor->num_tensors; + const size_t first_logical_dim = input_tensor->logical_shape.data[0]; + const size_t last_logical_dim = input_tensor->logical_shape.data[1]; + + // The kernel is compiled with cute.sym_int32(divisibility=...) on both logical extents, + // so a violating shape would silently mis-tile rather than fail. These mirror sym_M / + // sym_N in CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py -- the DSL kernel's chunk height and + // the 16-byte TMA row alignment. The logical shape of VARYING_BOTH_DIMS is [1, total]. + constexpr size_t kChunkDimY = 128; + constexpr size_t kLastDimAlignment = 16; + const bool first_dim_tiles = config.shape_rep == ShapeRepresentation::VARYING_BOTH_DIMS || + first_logical_dim % kChunkDimY == 0; + if (!first_dim_tiles || last_logical_dim % kLastDimAlignment != 0) { + maybe_warn_cutedsl_not_chosen("the grouped logical shape is not a multiple of (", kChunkDimY, + ", ", kLastDimAlignment, ")."); + return false; + } + // The same extents are sym_int32 in the compiled kernel. + if (first_logical_dim > static_cast(INT32_MAX) || + last_logical_dim > static_cast(INT32_MAX)) { + maybe_warn_cutedsl_not_chosen("the grouped logical shape does not fit in int32."); + return false; + } + + // dbias workspace-size query, mirroring mxfp8::group_quantize: the framework first calls + // with an unallocated workspace to learn its shape, allocates it, then calls again to run. + // The kernel writes one partial-dbias row per 128-row chunk. + if (config.with_dbias && workspace_tensor->data.dptr == nullptr) { + workspace_tensor->data.shape = {DIVUP(first_logical_dim, kChunkDimY), last_logical_dim}; + workspace_tensor->data.dtype = DType::kFloat32; + return true; + } + + std::optional group_quant_func_opt = config.get_kernel(); + if (!group_quant_func_opt.has_value()) { + return false; + } + + const int32_t device_index = transformer_engine::cuda::current_device(); + GroupDescriptorWorkspace *const workspace = group_descriptor_workspace_ptr(); + + const SimpleTensor &scale_row = + config.rowwise ? output_tensor->scale_inv : output_tensor->columnwise_scale_inv; + const SimpleTensor &scale_col = + config.colwise ? output_tensor->columnwise_scale_inv : output_tensor->scale_inv; + + // The group's payload is stored flat; the kernel wants it as the logical 2D view. + const std::vector logical_shape{first_logical_dim, last_logical_dim}; + DLTensorWrapper mX( + make_basic_tensor(input_tensor->data.dptr, input_tensor->dtype(), logical_shape), true, + device_index); + DLTensorWrapper mO_row, mO_col; + if (config.rowwise) { + mO_row = DLTensorWrapper( + make_basic_tensor(output_tensor->data.dptr, output_tensor->data.dtype, logical_shape), true, + device_index); + } + if (config.colwise) { + mO_col = DLTensorWrapper(make_basic_tensor(output_tensor->columnwise_data.dptr, + output_tensor->columnwise_data.dtype, logical_shape), + true, device_index); + } + + // The kernel only takes the base address of the scale buffers (per-tensor bases and + // strides are derived from the member shapes), so these stay 1D. + DLTensorWrapper mS_row(make_basic_tensor(scale_row.dptr, scale_row.dtype, {scale_row.numel()}), + false, device_index); + DLTensorWrapper mS_col(make_basic_tensor(scale_col.dptr, scale_col.dtype, {scale_col.numel()}), + false, device_index); + + // Offsets and member dims are read from the output, as in mxfp8::group_quantize. + auto offsets_or_unused = [&](const SimpleTensor &t, size_t numel) { + void *dptr = t.has_data() ? t.dptr : static_cast(workspace->unused_offsets); + return DLTensorWrapper(make_basic_tensor(dptr, DType::kInt64, {numel}), false, device_index); + }; + DLTensorWrapper mOffsets = offsets_or_unused(output_tensor->tensor_offsets, num_tensors + 1); + DLTensorWrapper mFirstDims, mLastDims; + if (config.shape_rep == ShapeRepresentation::VARYING_FIRST_DIM || + config.shape_rep == ShapeRepresentation::VARYING_BOTH_DIMS) { + NVTE_CHECK(output_tensor->first_dims.has_data(), "Grouped MXFP8 quantization with ", + shape_rep_to_str(config.shape_rep), " requires an allocated first_dims buffer."); + mFirstDims = DLTensorWrapper( + make_basic_tensor(output_tensor->first_dims.dptr, DType::kInt64, {num_tensors}), false, + device_index); + } + if (config.shape_rep == ShapeRepresentation::VARYING_LAST_DIM || + config.shape_rep == ShapeRepresentation::VARYING_BOTH_DIMS) { + NVTE_CHECK(output_tensor->last_dims.has_data(), "Grouped MXFP8 quantization with ", + shape_rep_to_str(config.shape_rep), " requires an allocated last_dims buffer."); + mLastDims = DLTensorWrapper( + make_basic_tensor(output_tensor->last_dims.dptr, DType::kInt64, {num_tensors}), false, + device_index); + } + + // The kernel reads num_tensors off this tensor's leading extent, so it must be exactly + // the group size even on the single-tensor path that leaves the descriptors untouched. + DLTensorWrapper mTensormaps( + make_basic_tensor(static_cast(workspace->tensor_maps), DType::kInt64, + {num_tensors, kGroupTensorMapSlots, kInt64PerTensorMap}), + false, device_index); + + // Optional inputs: a wrapper over a null buffer packs as TVM-FFI None. + DLTensorWrapper mActInput, mWorkspace; + if (config.with_dact) { + mActInput = DLTensorWrapper( + make_basic_tensor(act_input_tensor->data.dptr, act_input_tensor->dtype(), logical_shape), + true, device_index); + } + if (config.with_dbias) { + mWorkspace = DLTensorWrapper(workspace_tensor->data, true, device_index); + } + + // The cast-noop flag travels as a raw device pointer (not a tensor): it may be null, and the + // kernel null-checks it on device, so one compiled kernel serves both cases. + void *noop_ptr = (noop_tensor != nullptr) ? noop_tensor->data.dptr : nullptr; + + // noop and stream are tvm-ffi opaque "handles"; pass them as void*. + (*group_quant_func_opt)(&mX, &mO_row, &mO_col, &mS_row, &mS_col, &mOffsets, &mFirstDims, + &mLastDims, &mTensormaps, noop_ptr, &mActInput, &mWorkspace, + static_cast(stream)); + + // Reduce the per-chunk partial dbias per member with the CUDA kernel's reduction. + if (config.with_dbias) { + const float *workspace_ptr = reinterpret_cast(workspace_tensor->data.dptr); + TRANSFORMER_ENGINE_TYPE_SWITCH_NON_FP8ONLY( + input_tensor->dtype(), IType, + dispatch::common::grouped_reduce_dbias( + config.shape_rep, num_tensors, first_logical_dim, last_logical_dim, + reinterpret_cast(output_tensor->tensor_offsets.dptr), + reinterpret_cast(output_tensor->first_dims.dptr), + reinterpret_cast(output_tensor->last_dims.dptr), dbias_tensor, + workspace_ptr, kChunkDimY, stream);) // NOLINT(*) + } + return true; +} + +template +bool mxfp8_group_quantize_cutedsl(const GroupedTensor *input_tensor, + const GroupedTensor *act_input_tensor, const Tensor *noop_tensor, + GroupedTensor *output_tensor, GroupedTensor *dbias_tensor, + Tensor *workspace_tensor, const QuantizationConfig *quant_config, + cudaStream_t stream) { + if (!tvm_ffi_bridge::TVMFFICentral::getInstance().get_cutedsl_backend_enabled()) { + maybe_warn_cutedsl_not_chosen("the CuTeDSL backend is disabled."); + return false; + } + // TODO(kainingz): port 2D quantization to CuTeDSL + if (quant_config != nullptr && quant_config->mxfp8_2d_quantization) { + maybe_warn_cutedsl_not_chosen("2D quantization is not supported."); + return false; + } + constexpr Activation activation = activation_func_to_enum(); + if constexpr (activation == Activation::kUnsupported) { + maybe_warn_cutedsl_not_chosen("the fused activation is not supported."); + return false; + } else { + // Mirrors the shape-representation selection in mxfp8::group_quantize. + ShapeRepresentation shape_rep = ShapeRepresentation::SAME_BOTH_DIMS; + if (output_tensor->all_same_shape()) { + shape_rep = ShapeRepresentation::SAME_BOTH_DIMS; + } else if (output_tensor->all_same_first_dim()) { + shape_rep = ShapeRepresentation::VARYING_LAST_DIM; + } else if (output_tensor->all_same_last_dim()) { + shape_rep = ShapeRepresentation::VARYING_FIRST_DIM; + } else if (output_tensor->varying_both_dims()) { + shape_rep = ShapeRepresentation::VARYING_BOTH_DIMS; + } + const bool is_single_tensor = shape_rep == ShapeRepresentation::SAME_BOTH_DIMS || + shape_rep == ShapeRepresentation::VARYING_FIRST_DIM; + + // Leave invalid group sizes to mxfp8::group_quantize, which raises a proper error. + // Every member gets a descriptor slot in the fixed-size workspace, so the CUDA + // kernel's descriptor limit applies to the single-tensor representations here too. + const size_t num_tensors = input_tensor->num_tensors; + if (num_tensors == 0 || num_tensors > kMaxGroupTensors) { + maybe_warn_cutedsl_not_chosen("the group size ", num_tensors, " is not between 1 and ", + kMaxGroupTensors, "."); + return false; + } + if (shape_rep == ShapeRepresentation::SAME_BOTH_DIMS) { + // The kernel tiles the stacked rows without tensor boundaries, which matches the CUDA + // kernel's per-tensor tiling only when every member's rows are a multiple of its + // 128-row chunk. mxfp8::group_quantize raises for a non-integral row count. + const size_t first_logical_dim = input_tensor->logical_shape.data[0]; + if (first_logical_dim % num_tensors != 0 || (first_logical_dim / num_tensors) % 128 != 0) { + maybe_warn_cutedsl_not_chosen("the rows of each group member are not a multiple of 128."); + return false; + } + } else if (!output_tensor->tensor_offsets.has_data()) { + // The varying representations read per-member offsets, as the CUDA kernel does. + maybe_warn_cutedsl_not_chosen("the grouped tensor has no tensor offsets."); + return false; + } + if (IS_DBIAS && !is_single_tensor) { + // mxfp8::group_quantize raises a proper error for this. + maybe_warn_cutedsl_not_chosen("dbias is only supported for a common last dimension."); + return false; + } + + const bool rowwise = output_tensor->has_data(); + const bool colwise = output_tensor->has_columnwise_data(); + if (!rowwise && !colwise) { + // mxfp8::group_quantize raises a proper error for this. + return false; + } + const bool swizzled = output_tensor->with_gemm_swizzled_scales; + // Sanity checks, mirroring mxfp8::group_quantize + checkCuDriverContext(stream); + CheckNoopTensor(*noop_tensor, "cast_noop"); + if (rowwise) { + NVTE_CHECK(output_tensor->scale_inv.dptr != nullptr, "Scaling tensor must be allocated"); + } + if (colwise) { + NVTE_CHECK(output_tensor->columnwise_scale_inv.dptr != nullptr, + "Columnwise scaling tensor must be allocated"); + } + NVTE_CHECK(input_tensor->num_tensors == output_tensor->num_tensors, + "Number of input and output tensors must be same."); + NVTE_CHECK(input_tensor->has_data(), "Cannot quantize tensor without rowwise data."); + NVTE_CHECK(is_fp8_dtype(output_tensor->dtype()), "Output must have FP8 type."); + if constexpr (IS_DACT) { + NVTE_CHECK(act_input_tensor->has_data(), "Activations tensor must have data."); + NVTE_CHECK(input_tensor->num_tensors == act_input_tensor->num_tensors, + "Number of grad and activations tensors must be same."); + NVTE_CHECK(input_tensor->dtype() == act_input_tensor->dtype(), + "Grad and activations tensors must have the same type."); + } + if constexpr (IS_DBIAS) { + NVTE_CHECK(dbias_tensor->data.dtype == input_tensor->dtype(), + "DBias must have the same type as input_tensor."); + const Shape expected_shape_dbias_tensor = {num_tensors, input_tensor->logical_shape.data[1]}; + NVTE_CHECK(dbias_tensor->data.shape == expected_shape_dbias_tensor, "Wrong shape of DBias."); + NVTE_CHECK(workspace_tensor != nullptr, "Workspace must be a tensor."); + } + + const MXFP8GroupQuantConfig config{/*dtype=*/input_tensor->dtype(), + /*fp8_dtype=*/output_tensor->dtype(), + /*rowwise=*/rowwise, + /*colwise=*/colwise, + /*shape_rep=*/shape_rep, + /*swizzled=*/swizzled, + /*with_dbias=*/IS_DBIAS, + /*with_dact=*/IS_DACT, + /*with_act=*/IS_ACT, + /*activation=*/activation}; + return mxfp8_group_quantize_cutedsl(config, input_tensor, act_input_tensor, noop_tensor, + output_tensor, dbias_tensor, workspace_tensor, stream); + } +} + +} // namespace cutedsl_backend +} // namespace transformer_engine + +#endif // TRANSFORMER_ENGINE_COMMON_CAST_MXFP8_GROUP_QUANTIZE_MXFP8_CUTEDSL_CUH_