From df5c281cb6aa320534f66d692443580955c70c82 Mon Sep 17 00:00:00 2001 From: Kaining Zhong Date: Tue, 29 Sep 2026 20:53:01 +0000 Subject: [PATCH 01/14] tmp Signed-off-by: Kaining Zhong --- .../cast/mxfp8/group_quantize_mxfp8.py | 1136 +++++++++++++++++ .../common/cast/dispatch/quantize.cuh | 15 +- .../mxfp8/group_quantize_mxfp8_cutedsl.cuh | 312 +++++ 3 files changed, 1460 insertions(+), 3 deletions(-) create mode 100644 transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py create mode 100644 transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh 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..11c288a7d36 --- /dev/null +++ b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py @@ -0,0 +1,1136 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Grouped MXFP8 quantization kernel implemented in CuTeDSL. + +Strategy-aligned port of group_quantize_mxfp8.cuh. The scheduling, descriptor +management and per-tensor scale addressing mirror the CUDA kernel one-for-one: + + * `is_single_tensor` reps (SAME_BOTH_DIMS, VARYING_FIRST_DIM) launch ONE CTA per + 128x128 chunk and address the group through ONE static TMA descriptor with + global block offsets -- the CUDA `tensor_map_*_static` "direct mapper" path. + For SAME_BOTH_DIMS the CUDA grid is linearized per tensor (X, Y-in-tensor, + tensor) while this kernel linearizes it flat over the stacked rows; both + require every member's row count to be a multiple of CHUNK_DIM_Y, and under + that precondition the two decode to the identical (block_offset_Y, block_id_X) + for every block index. + * the other reps launch grid=(workers_per_tensor, num_tensors) and bind + tensor_id to blockIdx.y, so a CTA grid-strides only within its own tensor and + never re-resolves which tensor a chunk belongs to. They get per-tensor + descriptors written by a prologue kernel (the CuTeDSL analog of + update_tma_descriptors filling g_tensor_maps) and acquired with a tensormap + proxy fence. + * per-tensor scale bases/strides follow the CUDA formulas: + scales_* += is_single_tensor ? 0 : tensor_base / 32 + stride_rowwise = roundup(cols/32, 4) stride_colwise = roundup(cols, 128) + +Both kernels dropped the older flat persistent grid that strided across tensor +boundaries (CUDA's decode_job / advance_to_next_job, removed in #3483); the +per-tensor grid above is what replaced it on both sides. + +Mechanics that provably yield the same bytes may differ: the mbarrier pipeline is +expressed with PipelineTmaAsync instead of hand-rolled mbarriers, and +out-of-bounds scale padding is skipped rather than explicitly zeroed (the CUDA +kernel writes 0 there; downstream only consumes the meaningful region). + +Scope: cast-only (no dbias / activation / dact / amax), compact (non-swizzled) +scales, rowwise and/or colwise. VARYING_BOTH_DIMS is not handled here (its +logical shape [1, total] is not tileable); the C++ bridge falls back to CUDA. + +Known gap vs CUDA: the CUDA kernel validates on device that every member's first +dim is a multiple of 128 (get_tensor_rows_num -> NVTE_DEVICE_ERROR). This kernel +shares the precondition but cannot check it -- CuTeDSL has no device-side assert, +and for VARYING_FIRST_DIM the per-tensor extents live in device memory, so the +C++ bridge cannot check them either without a sync. A violating group mis-tiles +silently here where CUDA raises. +""" + +""" +Measured, deliberately NOT changed: + - sO_row and sO_col are both allocated unconditionally, where CUDA sizes only the + direction in use. Sizing them conditionally does work -- ncu confirms the shared-memory + occupancy limit goes 6 -> 9 CTAs/SM for a single-direction bf16 config -- but it is a + small LOSS on GB200, not a win: rep_med_sbd bf16 colwise 54.9 -> 56.1 us, rowwise + 56.5 -> 57.0 us (fp32 rowwise gains ~1%). The kernel is DRAM-bandwidth-bound at + ~6.3 TB/s, so extra resident CTAs only add contention. Verified by a control that kept + the conditional code but padded SMEM back to the old size: timings returned exactly to + the unconditional numbers, so the effect is the occupancy, not codegen. Revisit if a + future variant (dbias / activation) makes this kernel latency- rather than + bandwidth-bound. +""" + +# pylint: disable=missing-class-docstring + +import logging +import os +from typing import Type + +import cutlass +from cutlass import cute +from cutlass import pipeline +from cutlass import Boolean, Int32, Int64, Float8E8M0FNU +from cutlass.cute.nvgpu import cpasync +from cutlass.tensor_utils import TensorMapManager, TensorMapUpdateMode +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, + FP8E4M3_MAX_NORM_RCP, + FP8E5M2_MAX_NORM_RCP, + 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. +NUM_TENSORMAPS = 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" +SUPPORTED_SHAPE_REPS = (SAME_BOTH_DIMS, VARYING_FIRST_DIM, VARYING_LAST_DIM) + + +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): + if dtype not in ("fp32", "fp16", "bf16"): + raise ValueError(f"unknown input dtype {dtype!r}; expected fp32|fp16|bf16") + self.DTYPE = str_to_cutlass_dtype(dtype) + self.DTYPE_STR = dtype + if fp8_dtype not in ("fp8_e4m3fn", "fp8_e5m2"): + raise ValueError(f"unknown FP8 dtype {fp8_dtype!r}; expected fp8_e4m3fn|fp8_e5m2") + 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 + # Mirrors `is_single_tensor` in group_quantize_mxfp8.cuh. + self.IS_SINGLE_TENSOR = shape_rep in (SAME_BOTH_DIMS, VARYING_FIRST_DIM) + self.MAX_NORM_RCP = ( + FP8E4M3_MAX_NORM_RCP if fp8_dtype == "fp8_e4m3fn" else FP8E5M2_MAX_NORM_RCP + ) + + 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})" + ) + + __repr__ = __str__ + + +class MXFP8GroupQuantizeKernel: + """Grouped MXFP8 quantize mirroring group_quantize_mxfp8_kernel's strategy.""" + + # TunableConfig / derived constants from group_quantize_mxfp8.cuh. + CHUNK_DIM_Y = 128 + CHUNK_DIM_X = 128 + THREADS_PER_CHUNK = 128 + STATIC_PERSISTENT_BLOCKS_PER_SM = 24 + ELTS_PER_CHUNK = CHUNK_DIM_Y * CHUNK_DIM_X + THREADS_X = CHUNK_DIM_X // MXFP8_BLOCK_SCALING_SIZE # 4 + THREADS_Y = THREADS_PER_CHUNK // THREADS_X # 32 + BUFF_DIM_Y = THREADS_Y # 32 + BUFF_DIM_X = CHUNK_DIM_X # 128 + # Each block of (CHUNK_DIM_Y, CHUNK_DIM_X) consists of STAGES tiles of (BUFF_DIM_X, BUFF_DIM_Y) stacked vertically + STAGES = CHUNK_DIM_Y // BUFF_DIM_Y # 4 + PIPELINE_DEPTH = 2 # PREFETCH_STAGES(1) + 1 + NUM_WARPS = THREADS_PER_CHUNK // THREADS_PER_WARP # 4 + + # Rowwise vectorization constants (mirror MXFP8QuantizeKernel / CUDA PACK_SIZE). + PACK_SIZE = 4 + WAVES = MXFP8_BLOCK_SCALING_SIZE // PACK_SIZE # 8 + THREADS_PER_BANK = (32 * 4) // MXFP8_BLOCK_SCALING_SIZE # 4 + + def __init__(self, cfg: MXFP8GroupQuantizeConfig, SM_COUNT: int): + self.cfg = cfg + self.SM_COUNT = SM_COUNT + + # ---------------------------------------------------------------- helpers + @cute.jit + def _tensor_rows_cols( + self, tensor_id, mFirstDims, mLastDims, first_logical_dim, last_logical_dim + ): + """Get the shape (rows, cols) of the tensor by tensor_id.""" + cfg = self.cfg + if cutlass.const_expr(cfg.SHAPE_REP == VARYING_FIRST_DIM): + rows = Int32(mFirstDims[tensor_id]) + else: + rows = Int32(first_logical_dim) + if cutlass.const_expr(cfg.SHAPE_REP == VARYING_LAST_DIM): + cols = Int32(mLastDims[tensor_id]) + else: + cols = Int32(last_logical_dim) + return rows, cols + + # ------------------------------------------------------------ entry point + @cute.jit + def __call__( + self, + mX: cute.Tensor, + mO_row: cute.Tensor, + mO_col: cute.Tensor, + mS_row: cute.Tensor, + mS_col: cute.Tensor, + mOffsets: cute.Tensor, # int64[num_tensors + 1], CSR element offsets + mFirstDims: cute.Tensor, # int64[num_tensors] (VARYING_FIRST_DIM) + mLastDims: cute.Tensor, # int64[num_tensors] (VARYING_LAST_DIM) + mTensormaps: cute.Tensor, # int64[num_tensors, NUM_TENSORMAPS, 16] + stream: CUstream, + ): + 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] + # Number of group members. Do NOT derive this from mOffsets: a caller is free to + # pass a length-num_tensors stub for SAME_BOTH_DIMS (where the offsets array is + # unused), which would make `mOffsets.shape[0] - 1` read one too few and divide by + # zero at num_tensors == 1. The per-tensor descriptor workspace is num_tensors long + # by construction -- one slot set per member -- so it is the reliable source, and it + # is only consulted on the multi-tensor path that actually needs the workspace. + num_tensors = mTensormaps.shape[0] + + smem_tile_layout = cute.make_ordered_layout( + (self.BUFF_DIM_Y, self.BUFF_DIM_X), order=(1, 0) + ) + cta_tiler = (self.BUFF_DIM_Y, self.BUFF_DIM_X) + print(f"mx={mX}, smem_tile_layout={smem_tile_layout}, cta_tiler={cta_tiler}\n") + + 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 + ) + print(f"tma_atom_x={tma_atom_x}\n") + print(f"tma_src={tma_src}\n") + op_store = cpasync.CopyBulkTensorTileS2GOp() + 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, 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 blocks does the grouped tensor have in both directions + work_blocks_X = cute.ceil_div(Int32(last_logical_dim), self.CHUNK_DIM_X) + work_blocks_Y = cute.ceil_div(Int32(first_logical_dim), self.CHUNK_DIM_Y) + # Each CTA handles a block from one individual tensor + grid = [work_blocks_X * work_blocks_Y, 1, 1] + else: + # The work-block grid is per-tensor here, not global: each CTA derives its own + # block range from its tensor's extents (see the kernel's `else` branch), so + # work_blocks_X/Y are dead on this path -- the kernel reads them only under + # IS_SINGLE_TENSOR. Pass dummies rather than computing first*last, which is an + # ELEMENT count and would wrap Int32 past 2^31 elements. + work_blocks_Y = Int32(1) + work_blocks_X = Int32(1) + # Persistent worker count, mirroring get_launch_config() in + # group_quantize_mxfp8.cuh. There are SM_COUNT * STATIC_PERSISTENT_BLOCKS_PER_SM + # workers in total, split evenly across tensors -- but clamped to the average + # number of chunks a tensor actually holds, so a group of small tensors does not + # launch CTAs that only reach the `first_block_id >= blocks_in_tensor` early-out. + n_tensors = cutlass.max(Int32(num_tensors), Int32(1)) # never divide by zero + # CUDA's DIVUP(elts_total, CHUNK_DIM_Y * TILE_DIM_X) / STAGES_X, where TILE_DIM_X + # is BUFF_DIM_X here and STAGES_X is 1 for every rep this kernel supports. The + # element product would wrap Int32 past 2^31 elements, so divide the first extent + # by CHUNK_DIM_Y up front -- it is exact, the kernel is compiled with + # sym_int32(divisibility=CHUNK_DIM_Y) on that extent -- and keep the whole + # estimate in Int32. Feeding an Int64 into the grid poisons the tile arithmetic + # downstream ('cute.make_tile' expects width=32). + estimated_work_blocks = cute.ceil_div( + (Int32(first_logical_dim) // self.CHUNK_DIM_Y) * Int32(last_logical_dim), + self.ELTS_PER_CHUNK // self.CHUNK_DIM_Y, + ) + requested_workers_per_tensor = cutlass.max( + Int32(1), + Int32(self.SM_COUNT * self.STATIC_PERSISTENT_BLOCKS_PER_SM) // n_tensors, + ) + average_work_blocks_per_tensor = cutlass.max( + Int32(1), cute.ceil_div(estimated_work_blocks, n_tensors) + ) + workers_per_tensor = cutlass.min( + requested_workers_per_tensor, average_work_blocks_per_tensor + ) + # Each group of workers_per_tensor CTAs serves one tensor, and we launch + # num_tensors such groups to cover the whole group. + grid = [workers_per_tensor, Int32(num_tensors), 1] + + # Only manually create descriptors for the non-single-tensor case because we will need to manually + # overwrite the descriptors as we visit different groups + if cutlass.const_expr(not cfg.IS_SINGLE_TENSOR): + self.update_descriptors_kernel( + mX, + mO_row, + mO_col, + mOffsets, + mFirstDims, + mLastDims, + mTensormaps, + first_logical_dim, + last_logical_dim, + mX.element_type, + tma_atom_x, + tma_atom_out_row, + tma_atom_out_col, + ).launch(grid=[num_tensors, 1, 1], block=[THREADS_PER_WARP, 1, 1], stream=stream) + + self.kernel( + mS_row, + mS_col, + mOffsets, + mFirstDims, + mLastDims, + mTensormaps, + first_logical_dim, + last_logical_dim, + num_tensors, + work_blocks_X, + work_blocks_Y, + mX.element_type, + tma_atom_x, + tma_src, + tma_atom_out_row, + tma_dst_out_row, + tma_atom_out_col, + tma_dst_out_col, + ).launch( + grid=grid, + block=[self.THREADS_PER_CHUNK, 1, 1], + stream=stream, + ) + + # ------------------------------------------------- descriptor prologue + @cute.kernel + def update_descriptors_kernel( + self, + mX, + mO_row, + mO_col, + mOffsets, + mFirstDims, + mLastDims, + mTensormaps, + first_logical_dim, + last_logical_dim, + dtype: cutlass.Constexpr[Type[cutlass.Numeric]], + tma_atom_x, + tma_atom_orow, + tma_atom_ocol, + ): + """One CTA per tensor: point that tensor's TMA descriptors at its own block. + + CuTeDSL analog of common::update_tma_descriptors writing g_tensor_maps[]. + """ + cfg = self.cfg + tensor_id, _, _ = cute.arch.block_idx() + rows, cols = self._tensor_rows_cols( + tensor_id, mFirstDims, mLastDims, first_logical_dim, last_logical_dim + ) + base_elts = Int64(mOffsets[tensor_id]) + + # Publish this tensor's geometry for the main kernel (written even when empty). + meta = mTensormaps[(tensor_id, META_SLOT, None)] + meta[0] = Int64(rows) + meta[1] = Int64(cols) + meta[2] = base_elts + + tmap = TensorMapManager(TensorMapUpdateMode.GMEM, BYTES_PER_TENSORMAP) + 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) + + # Zero-sized groups: creating a descriptor with a zero extent is invalid, + # so skip (the main kernel skips these jobs via job_has_work). + if rows > 0 and cols > 0: + gX = cute.make_tensor( + cute.make_ptr( + dtype, + mX.iterator.toint() + base_elts * (dtype.width // 8), + cute.AddressSpace.gmem, + assumed_align=16, + ), + cute.make_layout((rows, cols), stride=(cols, 1)), + ) + gO_row = cute.make_tensor( + cute.make_ptr( + cfg.FP8_DTYPE, + mO_row.iterator.toint() + base_elts, + cute.AddressSpace.gmem, + assumed_align=16, + ), + cute.make_layout((rows, cols), stride=(cols, 1)), + ) + gO_col = cute.make_tensor( + cute.make_ptr( + cfg.FP8_DTYPE, + mO_col.iterator.toint() + base_elts, + cute.AddressSpace.gmem, + assumed_align=16, + ), + cute.make_layout((rows, cols), stride=(cols, 1)), + ) + + tmap.init_tensormap_from_atom(tma_atom_x, desc_x, 0) + if cutlass.const_expr(cfg.ROWWISE): + tmap.init_tensormap_from_atom(tma_atom_orow, desc_orow, 0) + if cutlass.const_expr(cfg.COLWISE): + tmap.init_tensormap_from_atom(tma_atom_ocol, desc_ocol, 0) + tmap.fence_tensormap_initialization() + + if cutlass.const_expr(cfg.ROWWISE and cfg.COLWISE): + tmap.update_tensormap( + (gX, gO_row, gO_col), + (tma_atom_x, tma_atom_orow, tma_atom_ocol), + (desc_x, desc_orow, desc_ocol), + 0, + (), # smem staging is unused in GMEM update mode + ) + elif cutlass.const_expr(cfg.ROWWISE): + tmap.update_tensormap( + (gX, gO_row), + (tma_atom_x, tma_atom_orow), + (desc_x, desc_orow), + 0, + (), # smem staging is unused in GMEM update mode + ) + else: + tmap.update_tensormap( + (gX, gO_col), + (tma_atom_x, tma_atom_ocol), + (desc_x, desc_ocol), + 0, + (), # smem staging is unused in GMEM update mode + ) + + # ------------------------------------------------------------ main kernel + @cute.kernel + def kernel( + self, + mS_row, + mS_col, + mOffsets, + mFirstDims, + mLastDims, + mTensormaps, + first_logical_dim, + last_logical_dim, + num_tensors, + work_blocks_X, + work_blocks_Y, + dtype: cutlass.Constexpr[Type[cutlass.Numeric]], + tma_atom_x, + tma_src, + tma_atom_out_row, + tma_dst_out_row, + tma_atom_out_col, + tma_dst_out_col, + ): + 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()) + + # --- shared memory (allocated once, reused across jobs) --- + @cute.struct + class SharedStorage: + mbar: cute.struct.MemRange[cute.Int64, 2 * self.PIPELINE_DEPTH] + sX: cute.struct.Align[ + cute.struct.MemRange[ + dtype, self.BUFF_DIM_Y * self.BUFF_DIM_X * self.PIPELINE_DEPTH + ], + 128, + ] + sO_row: cute.struct.Align[ + cute.struct.MemRange[ + FP8_DTYPE, self.BUFF_DIM_Y * self.BUFF_DIM_X * self.PIPELINE_DEPTH + ], + 128, + ] + sO_col: cute.struct.Align[ + cute.struct.MemRange[ + FP8_DTYPE, self.BUFF_DIM_Y * self.BUFF_DIM_X * self.PIPELINE_DEPTH + ], + 128, + ] + + storage = cutlass.utils.SmemAllocator().allocate(SharedStorage) + tile_layout = cute.make_layout( + ((self.BUFF_DIM_Y, self.BUFF_DIM_X), self.PIPELINE_DEPTH), + stride=((self.BUFF_DIM_X, 1), self.BUFF_DIM_Y * self.BUFF_DIM_X), + ) + sX = storage.sX.get_tensor(tile_layout) + print(f"sX={sX}\n") + sO_row = storage.sO_row.get_tensor(tile_layout) + sO_col = storage.sO_col.get_tensor(tile_layout) + + mainloop_pipeline = pipeline.PipelineTmaAsync.create( + barrier_storage=storage.mbar.data_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=self.BUFF_DIM_Y * self.BUFF_DIM_X * dtype.width // 8, + 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 + ) + + # TMA partitions built from the representative views. For multi-tensor reps + # the descriptor is swapped per tensor and the tile coords are tensor-local. + gX_tiled = cute.zipped_divide(tma_src, (self.BUFF_DIM_Y, self.BUFF_DIM_X)) + tXsX, tXgX = cpasync.tma_partition(tma_atom_x, 0, cute.make_layout(1), sX, gX_tiled) + gO_row_tiled = cute.zipped_divide(tma_dst_out_row, (self.BUFF_DIM_Y, self.BUFF_DIM_X)) + tXsO_row, tXgO_row = cpasync.tma_partition( + tma_atom_out_row, 0, cute.make_layout(1), sO_row, gO_row_tiled + ) + gO_col_tiled = cute.zipped_divide(tma_dst_out_col, (self.BUFF_DIM_Y, self.BUFF_DIM_X)) + tXsO_col, tXgO_col = cpasync.tma_partition( + tma_atom_out_col, 0, cute.make_layout(1), sO_col, gO_col_tiled + ) + print(f"tma_atom_x={tma_atom_x}\n") + print(f"tma_src={tma_src}\n") + print(f"gX_tiled={gX_tiled}\n") + print(f"tXsX={tXsX}\n") + print(f"tXgX={tXgX}\n") + + 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 block + 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. + # 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) + # Block's offset and id in this individual tensor / global single tensor + block_offset_Y = Int32(0) + block_id_X = Int32(0) + + first_block_id = Int32(0) + blocks_in_tensor = Int32(1) + block_stride = Int32(1) + block_columns_in_tensor = Int32(1) + + if cutlass.const_expr(cfg.IS_SINGLE_TENSOR): + # grid = [work_blocks_X * work_blocks_Y, 1, 1] + block_id_Y = Int32(bidx) // work_blocks_X + block_id_X = Int32(bidx) % work_blocks_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 block start from + block_offset_Y = block_id_Y * self.CHUNK_DIM_Y + if cutlass.const_expr(cfg.SHAPE_REP == VARYING_FIRST_DIM): + # Check if this block contains any valid tokens since the last offset <= logical_first_dim * logical_last_dim + # Last CSR offset == the group's total element count. Indexed off the + # array's own length so it does not depend on num_tensors. + total_elts = Int64(mOffsets[mOffsets.shape[0] - 1]) + if Int64(block_offset_Y) * Int64(last_logical_dim) >= total_elts: + has_work = Boolean(False) + 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 blocks does this tensor have in both directions + block_columns_in_tensor = cute.ceil_div(tensor_cols, self.CHUNK_DIM_X) + block_rows_in_tensor = cute.ceil_div(tensor_rows, self.CHUNK_DIM_Y) + # How many blocks does this tensor have + blocks_in_tensor = block_columns_in_tensor * block_rows_in_tensor + # Which block (1D index) does this CTA start from + first_block_id = Int32(bidx) + # gdx is workers_per_tensor (how many CTAs are assigned to this tensor) + block_stride = Int32(gdx) + # If my first block_id is already beyond the tensor's last block, I have no work to do + if first_block_id >= blocks_in_tensor: + has_work = Boolean(False) + else: + # This tensor is empty, so this CTA has no work to do + has_work = Boolean(False) + + if has_work: + if cutlass.const_expr(cfg.IS_SINGLE_TENSOR): + # For single tensor case we don't use tensor descriptors + desc_x = desc_out_row = desc_out_col = None + + cute.arch.sync_threads() + + self._process_block( + block_offset_Y, + block_id_X, + tensor_rows, + tensor_cols, + tensor_base, + desc_x, + desc_out_row, + desc_out_col, + mS_row, + mS_col, + first_logical_dim, + tmap, + warp_idx, + tidx, + sX, + sO_row, + sO_col, + tXsX, + tXgX, + tXsO_row, + tXgO_row, + tXsO_col, + tXgO_col, + tma_atom_x, + tma_atom_out_row, + tma_atom_out_col, + mainloop_pipeline, + prod_state, + cons_state, + ) + else: + # For multi-tensor case, retrieve tensor descriptors we processed early in the prologue kernel + 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) + # Acquire the descriptors on ONE thread, as the CUDA kernel does + # (`leading_thread` in group_quantize_mxfp8.cuh); the sync_threads below + # publishes it CTA-wide. Running the tensormap acquire fence on all 128 + # threads is correct but very expensive -- it more than doubles the + # kernel time on the multi-tensor path (4096x14336 bidirectional: + # 133 us -> 58 us), since the cost scales with threads x descriptors. + if tidx == 0: + tmap.fence_tensormap_update(desc_x) + 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) + + cute.arch.sync_threads() + + # Grid-stride over this tensor's own chunks; the descriptors never change. + block_id = first_block_id + job_finished = Boolean(False) + while not job_finished: + block_id_Y_in_tensor = block_id // block_columns_in_tensor + block_id_X_in_tensor = block_id % block_columns_in_tensor + self._process_block( + block_id_Y_in_tensor * self.CHUNK_DIM_Y, + block_id_X_in_tensor, + tensor_rows, + tensor_cols, + tensor_base, + desc_x, + desc_out_row, + desc_out_col, + mS_row, + mS_col, + first_logical_dim, + tmap, + warp_idx, + tidx, + sX, + sO_row, + sO_col, + tXsX, + tXgX, + tXsO_row, + tXgO_row, + tXsO_col, + tXgO_col, + tma_atom_x, + tma_atom_out_row, + tma_atom_out_col, + mainloop_pipeline, + prod_state, + cons_state, + ) + # Find the next block to process + block_id = block_id + block_stride + if block_id >= blocks_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, + prod_state, + tile_y, + tile_x, + tma_atom_x, + tXgX, + tXsX, + tmap, + desc_x, + ): + """Emit one 32x128 TMA load 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. + """ + # Wait for the consumer to finish using this SMEM buffer + pipeline_obj.producer_acquire(prod_state) + if cutlass.const_expr(self.cfg.IS_SINGLE_TENSOR): + cute.copy( + tma_atom_x, + tXgX[(None, (tile_y, tile_x))], + tXsX[(None, prod_state.index)], + tma_bar_ptr=pipeline_obj.producer_get_barrier(prod_state), + ) + else: + # Every member shares tXgX's tile-coordinate arithmetic (the coefficients are + # just the tile size); tma_desc_ptr supplies this member's geometry. + cute.copy( + tma_atom_x, + tXgX[(None, (tile_y, tile_x))], + tXsX[(None, prod_state.index)], + tma_bar_ptr=pipeline_obj.producer_get_barrier(prod_state), + tma_desc_ptr=tmap.get_tensormap_ptr(desc_x, cute.AddressSpace.generic), + ) + # Notify the consumer that this SMEM buffer is ready for consumption + pipeline_obj.producer_commit(prod_state) + + @cute.jit + def _process_block( + self, + block_offset_Y_in_tensor, # Row offset of this chunk (global for single-tensor, else tensor-local) + block_id_X_in_tensor, # Column-chunk index within the tensor + rows, # Number of rows in this tensor (logical shape) + cols, # Number of columns in this tensor (logical shape) + tensor_base, # Int64 element offset of this tensor in the group (0 for single-tensor) + desc_x, # Per-tensor descriptors, already acquired by the caller (None if single-tensor) + desc_out_row, + desc_out_col, + mS_row, # Grouped rowwise scales + mS_col, # Grouped colwise scales + first_logical_dim, # First dimension of the grouped tensor + tmap, # TensorMapManager for managing TMA descriptors + warp_idx, + tidx, + sX, # SMEM input for this block + sO_row, # SMEM rowwise output for this block + sO_col, # SMEM colwise output for this block + tXsX, # Tiled sX for TMA + tXgX, # Tiled gX for TMA + tXsO_row, # Tiled sO_row for TMA + tXgO_row, # Tiled gO_row for TMA + tXsO_col, # Tiled sO_col for TMA + tXgO_col, # Tiled gO_col for TMA + tma_atom_x, # TMA atom for input + tma_atom_out_row, # TMA atom for rowwise output + tma_atom_out_col, # TMA atom for colwise output + mainloop_pipeline: cutlass.pipeline.PipelineTmaAsync, + prod_state, + cons_state, + ): + """Quantize one 128x128 chunk in STAGES slices of BUFF_DIM_Y rows.""" + cfg = self.cfg + block_offset_X = block_id_X_in_tensor * self.CHUNK_DIM_X + + # This tensor's rowwise scales + scale_rows = Int32(first_logical_dim) if cutlass.const_expr(cfg.IS_SINGLE_TENSOR) else rows + scale_base = ( + Int64(0) + if cutlass.const_expr(cfg.IS_SINGLE_TENSOR) + else tensor_base // Int64(MXFP8_BLOCK_SCALING_SIZE) + ) + + if cutlass.const_expr(cfg.ROWWISE): + # Rowwise scale's divisibility guarantee: (128, 4) + scale_row_stride = cute.round_up(cute.ceil_div(cols, MXFP8_BLOCK_SCALING_SIZE), 4) + # Advance to this tensor's rowwise scales + mS_row_t = cute.make_tensor( + cute.make_ptr( + Float8E8M0FNU, + mS_row.iterator.toint() + scale_base, + cute.AddressSpace.gmem, + assumed_align=4, + ), + cute.make_layout((scale_rows, scale_row_stride), stride=(scale_row_stride, 1)), + ) + mS_row_tiled = cute.zipped_divide( + mS_row_t, (self.BUFF_DIM_Y, self.CHUNK_DIM_X // MXFP8_BLOCK_SCALING_SIZE) + ) + + if cutlass.const_expr(cfg.COLWISE): + # Colwise scale's divisibility guarantee: (4, 128) + scale_col_stride = cute.round_up(cols, 128) + # Advance to this tensor's colwise scales + mS_col_t = cute.make_tensor( + cute.make_ptr( + Float8E8M0FNU, + mS_col.iterator.toint() + scale_base, + cute.AddressSpace.gmem, + assumed_align=4, + ), + cute.make_layout( + (scale_rows // MXFP8_BLOCK_SCALING_SIZE, scale_col_stride), + stride=(scale_col_stride, 1), + ), + ) + mS_col_tiled = cute.zipped_divide( + mS_col_t, (self.BUFF_DIM_Y // MXFP8_BLOCK_SCALING_SIZE, self.CHUNK_DIM_X) + ) + + cute.arch.sync_threads() + + # This chunk's coordinates in the tile grid (32x128 TMA boxes, not elements). + tile_id_Y = block_offset_Y_in_tensor // self.BUFF_DIM_Y + tile_id_X = block_id_X_in_tensor + + # 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, + tile_id_Y + prologue_stage, + tile_id_X, + tma_atom_x, + tXgX, + tXsX, + tmap, + desc_x, + ) + prod_state.advance() + + for stage in cutlass.range_constexpr(self.STAGES): + # 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)] + row_tile = tile_id_Y + stage + tile_row_start = block_offset_Y_in_tensor + stage * self.BUFF_DIM_Y + + if cutlass.const_expr(cfg.COLWISE): + quantize_colwise_mxfp8( + sX_tile, + None, + sO_col[(None, cons_state.index)], + cute.flatten(mS_col_tiled[(None, (row_tile, tile_id_X))]), + cfg.MAX_NORM_RCP, + tile_row_start, + block_offset_X, + scale_rows, + cols, + ACTIVATION=None, + DTYPE=cfg.DTYPE, + FP8_DTYPE=cfg.FP8_DTYPE, + SWIZZLE=False, + TILE_X=self.BUFF_DIM_X, + TILE_Y=self.BUFF_DIM_Y, + SKIP_MASKING=False, + ) + if cutlass.const_expr(cfg.ROWWISE): + quantize_rowwise_mxfp8( + sX_tile, + None, + sO_row[(None, cons_state.index)], + cute.flatten(mS_row_tiled[(None, (row_tile, tile_id_X))]), + cfg.MAX_NORM_RCP, + tile_row_start, + block_offset_X, + scale_rows, + cols, + ACTIVATION=None, + DTYPE=cfg.DTYPE, + FP8_DTYPE=cfg.FP8_DTYPE, + TILE_X=self.BUFF_DIM_X, + TILE_Y=self.BUFF_DIM_Y, + WAVES=self.WAVES, + THREADS_PER_BANK=self.THREADS_PER_BANK, + PACK_SIZE=self.PACK_SIZE, + SKIP_MASKING=False, + ) + + # 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 cutlass.const_expr(stage + self.PIPELINE_DEPTH < self.STAGES): + if warp_idx == 0: + self._issue_load( + mainloop_pipeline, + prod_state, + tile_id_Y + stage + self.PIPELINE_DEPTH, + tile_id_X, + tma_atom_x, + tXgX, + tXsX, + tmap, + desc_x, + ) + prod_state.advance() + + # Write result to GMEM via TMA + if warp_idx == 0: + if cutlass.const_expr(cfg.ROWWISE): + if cutlass.const_expr(cfg.IS_SINGLE_TENSOR): + cute.copy( + tma_atom_out_row, + tXsO_row[(None, cons_state.index)], + tXgO_row[(None, (row_tile, tile_id_X))], + ) + else: + cute.copy( + tma_atom_out_row, + tXsO_row[(None, cons_state.index)], + tXgO_row[(None, (row_tile, tile_id_X))], + tma_desc_ptr=tmap.get_tensormap_ptr( + desc_out_row, cute.AddressSpace.generic + ), + ) + if cutlass.const_expr(cfg.COLWISE): + if cutlass.const_expr(cfg.IS_SINGLE_TENSOR): + cute.copy( + tma_atom_out_col, + tXsO_col[(None, cons_state.index)], + tXgO_col[(None, (row_tile, tile_id_X))], + ) + else: + cute.copy( + tma_atom_out_col, + tXsO_col[(None, cons_state.index)], + tXgO_col[(None, (row_tile, tile_id_X))], + tma_desc_ptr=tmap.get_tensormap_ptr( + desc_out_col, cute.AddressSpace.generic + ), + ) + # Commit all TMA operations of this iteration + cute.arch.cp_async_bulk_commit_group() + + cons_state.advance() + + +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/cols likewise); MXFP8 needs the last dim divisible by 32. + sym_M = cute.sym_int32(divisibility=128) + sym_N = cute.sym_int32(divisibility=MXFP8_BLOCK_SCALING_SIZE) + logical_shape = (sym_M, sym_N) + + out_dtype = cfg.FP8_DTYPE + scale_dtype = cutlass.Float8E8M0FNU + + def g2d(dtype, align=16): + return cute.runtime.make_fake_compact_tensor( + dtype, + logical_shape, + stride_order=(1, 0), + memspace=cute.AddressSpace.gmem, + assumed_align=align, + ) + + def g1d(dtype, align=4): + return cute.runtime.make_fake_compact_tensor( + dtype, + (cute.sym_int32(),), + stride_order=(0,), + memspace=cute.AddressSpace.gmem, + assumed_align=align, + ) + + # 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. + in_fake = g2d(cfg.DTYPE) + out_row_fake = g2d(out_dtype) + out_col_fake = g2d(out_dtype) + scale_row_fake = g1d(scale_dtype) + scale_col_fake = g1d(scale_dtype) + offsets_fake = g1d(cutlass.Int64, align=8) + first_dims_fake = g1d(cutlass.Int64, align=8) + last_dims_fake = g1d(cutlass.Int64, align=8) + 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, + ) + + from cutlass.utils import HardwareInfo # pylint: disable=import-outside-toplevel + + sm_count = HardwareInfo().get_device_multiprocessor_count() + kernel_obj = MXFP8GroupQuantizeKernel(cfg, sm_count) + return cute.compile( + kernel_obj, + in_fake, + out_row_fake, + out_col_fake, + scale_row_fake, + scale_col_fake, + offsets_fake, + first_dims_fake, + last_dims_fake, + tensormaps_fake, + cute.runtime.make_fake_stream(), + options="--enable-tvm-ffi", + ) + + +# TEMPORARY (demo only): same as compile_cutedsl_function_from_cfg but with every +# extent pinned to a compile-time constant, so traced layouts print concrete numbers +# (e.g. `(384,256):(1@1,1@0)`) instead of `?{div=128}`. Delete once done inspecting. +def compile_cutedsl_function_from_cfg_static( + cfg: MXFP8GroupQuantizeConfig, + M_total: int, + N: int, + num_tensors: int, + scale_row_numel: int, + scale_col_numel: int, +): + """Compile with fully static shapes. Only accepts inputs of exactly these extents.""" + logical_shape = (M_total, N) + + def g2d(dtype, align=16): + return cute.runtime.make_fake_compact_tensor( + dtype, + logical_shape, + stride_order=(1, 0), + memspace=cute.AddressSpace.gmem, + assumed_align=align, + ) + + def g1d(dtype, numel, align=4): + return cute.runtime.make_fake_compact_tensor( + dtype, + (numel,), + stride_order=(0,), + memspace=cute.AddressSpace.gmem, + assumed_align=align, + ) + + scale_dtype = cutlass.Float8E8M0FNU + tensormaps_fake = cute.runtime.make_fake_compact_tensor( + cutlass.Int64, + (num_tensors, NUM_WORKSPACE_SLOTS, BYTES_PER_TENSORMAP // 8), + stride_order=(2, 1, 0), + memspace=cute.AddressSpace.gmem, + assumed_align=128, + ) + + from cutlass.utils import HardwareInfo # pylint: disable=import-outside-toplevel + + sm_count = HardwareInfo().get_device_multiprocessor_count() + return cute.compile( + MXFP8GroupQuantizeKernel(cfg, sm_count), + g2d(cfg.DTYPE), + g2d(cfg.FP8_DTYPE), + g2d(cfg.FP8_DTYPE), + g1d(scale_dtype, scale_row_numel), + g1d(scale_dtype, scale_col_numel), + g1d(cutlass.Int64, num_tensors + 1, align=8), + g1d(cutlass.Int64, num_tensors, align=8), + g1d(cutlass.Int64, num_tensors, align=8), + tensormaps_fake, + 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, +) -> bool: + """Compile the grouped MXFP8 quantize kernel for this config and register it in the + TVM-FFI global registry under EXACTLY `fn_name`. Returns True on success (the C++ + dispatcher then fetches it with GetGlobal(fn_name)); False if unsupported, so the + caller falls back to the CUDA C++ grouped kernel. + """ + 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, + ) + 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) + try: + compiled = compile_cutedsl_function_from_cfg(cfg) + except Exception as e: # pylint: disable=broad-exception-caught + logger.error( + "CuTeDSL grouped MXFP8 kernel compilation failed, " + "falling back to the CUDA C++ kernel: %s", + e, + ) + return False + tvm_ffi.register_global_func(fn_name, compiled, override=True) + return True diff --git a/transformer_engine/common/cast/dispatch/quantize.cuh b/transformer_engine/common/cast/dispatch/quantize.cuh index d31f18f0d8b..b38b6fc12d1 100644 --- a/transformer_engine/common/cast/dispatch/quantize.cuh +++ b/transformer_engine/common/cast/dispatch/quantize.cuh @@ -524,9 +524,18 @@ 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); + #ifdef NVTE_WITH_CUTEDSL + quantized_with_cutedsl = + cutedsl_backend::mxfp8_group_quantize_cutedsl( + input_tensor, noop_tensor, output_tensor, quant_config_cpp.mxfp8_2d_quantization, + 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: { 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..ca33f0b25ef --- /dev/null +++ b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh @@ -0,0 +1,312 @@ +/************************************************************************* + * 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 "../../common.h" +#include "../../tvm_ffi_bridge.h" +#include "../../utils.cuh" // ShapeRepresentation +#include "../core/common.cuh" +#include "group_quantize_mxfp8.cuh" + +namespace transformer_engine { +namespace cutedsl_backend { + +// te_dtype_to_str, DLTensorWrapper, TVMFFICentral all live in +// transformer_engine::tvm_ffi_bridge (tvm_ffi_bridge.h). +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 + + constexpr uint32_t to_id() const { + static_assert(static_cast(DType::kNumTypes) <= 256, + "DType no longer fits in the 8 bits to_id() gives it."); + return static_cast(dtype) | (static_cast(fp8_dtype) << 8) | + (static_cast(rowwise) << 16) | (static_cast(colwise) << 17) | + (static_cast(shape_rep) << 18); + } + + 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; + key.reserve(64); // longest: cutedsl_group_mxfp8_bf16_fp8_e4m3fn_1_1_varying_first_dim + key.append("cutedsl_group_mxfp8_") + .append(te_dtype_to_str(dtype)) + .append("_") + .append(te_dtype_to_str(fp8_dtype)) + .append("_") + .append(rowwise ? "1" : "0") + .append("_") + .append(colwise ? "1" : "0") + .append("_") + .append(shape_rep_to_str(shape_rep)); + 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(te_dtype_to_str(dtype)), + tvm::ffi::String(te_dtype_to_str(fp8_dtype)), rowwise, colwise, + tvm::ffi::String(shape_rep_to_str(shape_rep))); + return result.try_cast().value_or(false); + } +}; + +// Descriptor slots per group member: input, rowwise output, colwise output, 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 = 4; +constexpr size_t kInt64PerTensorMap = 128 / sizeof(int64_t); + +struct alignas(128) GroupDescriptorWorkspace { + alignas(128) int64_t tensor_maps[dispatch::common::MAX_SUPPORTED_TENSOR_DESCRIPTORS] + [kGroupTensorMapSlots][kInt64PerTensorMap]; + // Stand-in for the offsets / first_dims / last_dims arrays a given shape representation + // does not carry: the kernel takes all three unconditionally but only dereferences the + // ones its representation uses, so the contents are never read. Sized num_tensors + 1 + // for the CSR offsets array, the longest of the three. + int64_t unused_dims[dispatch::common::MAX_SUPPORTED_TENSOR_DESCRIPTORS + 1]; +}; + +// One workspace per translation unit, mirroring `g_tensor_maps` on the CUDA path -- and +// sharing its caveat that two grouped quantize calls in flight on different streams would +// overwrite each other's descriptors. +static __device__ GroupDescriptorWorkspace g_group_descriptor_workspace; + +inline GroupDescriptorWorkspace *group_descriptor_workspace_ptr() { + static GroupDescriptorWorkspace *const ptr = [] { + void *p = nullptr; + NVTE_CHECK_CUDA(cudaGetSymbolAddress(&p, g_group_descriptor_workspace)); + return static_cast(p); + }(); + return ptr; +} + +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, output, stream) for the subset the +// CuTeDSL kernel covers. Returns false to fall back to the CUDA kernel. +inline bool mxfp8_group_quantize_cutedsl(const MXFP8GroupQuantConfig &config, + const GroupedTensor *input_tensor, + GroupedTensor *output_tensor, cudaStream_t stream) { + using namespace dispatch::mxfp8::group_quantize_kernel; + + 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 own chunk + // height and MXFP8 block size, which it tiles independently of the CUDA kernel's + // CastTraits. + constexpr size_t kChunkDimY = 128; + constexpr size_t kScaleDimX = 32; + if (first_logical_dim % kChunkDimY != 0 || last_logical_dim % kScaleDimX != 0) { + maybe_warn_cutedsl_not_chosen("the grouped logical shape is not a multiple of (", kChunkDimY, + ", ", kScaleDimX, ")."); + return false; + } + + std::optional group_quant_func_opt = config.get_kernel(); + if (!group_quant_func_opt.has_value()) { + return false; + } + + GroupDescriptorWorkspace *const workspace = group_descriptor_workspace_ptr(); + + // Both output directions are handed to the kernel unconditionally: the compiled + // signature has no optional tensors, and building a TMA descriptor needs a real + // address for each. The disabled direction is never read or written, so it points at + // the enabled one instead of at a buffer that would have to be allocated. + const SimpleTensor &data_row = + config.rowwise ? output_tensor->data : output_tensor->columnwise_data; + const SimpleTensor &data_col = + config.colwise ? output_tensor->columnwise_data : output_tensor->data; + 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}; + const NVTEBasicTensor x_bt = + make_basic_tensor(input_tensor->data.dptr, input_tensor->dtype(), logical_shape); + const NVTEBasicTensor o_row_bt = make_basic_tensor(data_row.dptr, data_row.dtype, logical_shape); + const NVTEBasicTensor o_col_bt = make_basic_tensor(data_col.dptr, data_col.dtype, logical_shape); + DLTensorWrapper mX(x_bt), mO_row(o_row_bt), mO_col(o_col_bt); + + // 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. + const NVTEBasicTensor s_row_bt = + make_basic_tensor(scale_row.dptr, scale_row.dtype, {scale_row.numel()}); + const NVTEBasicTensor s_col_bt = + make_basic_tensor(scale_col.dptr, scale_col.dtype, {scale_col.numel()}); + DLTensorWrapper mS_row(s_row_bt, /*flatten_2D=*/false), mS_col(s_col_bt, /*flatten_2D=*/false); + + const SimpleTensor &offsets = output_tensor->tensor_offsets; + const SimpleTensor &first_dims = output_tensor->first_dims; + const SimpleTensor &last_dims = output_tensor->last_dims; + const NVTEBasicTensor offsets_bt = make_basic_tensor( + offsets.has_data() ? offsets.dptr : static_cast(workspace->unused_dims), + DType::kInt64, {num_tensors + 1}); + const NVTEBasicTensor first_dims_bt = make_basic_tensor( + first_dims.has_data() ? first_dims.dptr : static_cast(workspace->unused_dims), + DType::kInt64, {num_tensors}); + const NVTEBasicTensor last_dims_bt = make_basic_tensor( + last_dims.has_data() ? last_dims.dptr : static_cast(workspace->unused_dims), + DType::kInt64, {num_tensors}); + DLTensorWrapper mOffsets(offsets_bt, /*flatten_2D=*/false), + mFirstDims(first_dims_bt, /*flatten_2D=*/false), + mLastDims(last_dims_bt, /*flatten_2D=*/false); + + // 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. + const NVTEBasicTensor tensormaps_bt = + make_basic_tensor(static_cast(workspace->tensor_maps), DType::kInt64, + {num_tensors, kGroupTensorMapSlots, kInt64PerTensorMap}); + DLTensorWrapper mTensormaps(tensormaps_bt, /*flatten_2D=*/false); + + // stream is a tvm-ffi opaque "handle"; pass it as void*. + (*group_quant_func_opt)(&mX, &mO_row, &mO_col, &mS_row, &mS_col, &mOffsets, &mFirstDims, + &mLastDims, &mTensormaps, static_cast(stream)); + return true; +} + +template +bool mxfp8_group_quantize_cutedsl(const GroupedTensor *input_tensor, const Tensor *noop_tensor, + GroupedTensor *output_tensor, const bool use_2d_quantization, + 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; + } + // The CuTeDSL grouped kernel is cast-only: no dbias, no fused (derivative) activation. + if constexpr (IS_DBIAS || IS_DACT || IS_ACT || OP != nullptr) { + maybe_warn_cutedsl_not_chosen( + "grouped quantization with dbias or a fused activation is not supported."); + return false; + } else { + // TODO(kainingz): port 2D quantization to CuTeDSL + if (use_2d_quantization) { + maybe_warn_cutedsl_not_chosen("2D quantization is not supported."); + return false; + } + // The kernel takes no noop flag, no amax accumulator, and writes compact scales only. + if (noop_tensor != nullptr && noop_tensor->data.dptr != nullptr) { + maybe_warn_cutedsl_not_chosen("the cast-noop flag is not supported."); + return false; + } + if (output_tensor->amax.dptr != nullptr) { + maybe_warn_cutedsl_not_chosen("amax computation is not supported."); + return false; + } + if (output_tensor->with_gemm_swizzled_scales) { + maybe_warn_cutedsl_not_chosen("GEMM-swizzled scales are not supported."); + return false; + } + + // 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 { + // VARYING_BOTH_DIMS: the logical shape is [1, total], which is not tileable. + maybe_warn_cutedsl_not_chosen("groups with both dimensions varying are not supported."); + return false; + } + // 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. + if (input_tensor->num_tensors > dispatch::common::MAX_SUPPORTED_TENSOR_DESCRIPTORS) { + maybe_warn_cutedsl_not_chosen("the group has more than ", + dispatch::common::MAX_SUPPORTED_TENSOR_DESCRIPTORS, + " tensors."); + 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; + } + + checkCuDriverContext(stream); + // Sanity checks, mirroring mxfp8::group_quantize + 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."); + + const MXFP8GroupQuantConfig config{/*dtype=*/input_tensor->dtype(), + /*fp8_dtype=*/output_tensor->dtype(), + /*rowwise=*/rowwise, + /*colwise=*/colwise, + /*shape_rep=*/shape_rep}; + return mxfp8_group_quantize_cutedsl(config, input_tensor, output_tensor, stream); + } +} + +} // namespace cutedsl_backend +} // namespace transformer_engine + +#endif // TRANSFORMER_ENGINE_COMMON_CAST_MXFP8_GROUP_QUANTIZE_MXFP8_CUTEDSL_CUH_ From 0185ada30b5b5679197b4435bb846f12e6907808 Mon Sep 17 00:00:00 2001 From: Kaining Zhong Date: Tue, 29 Sep 2026 21:15:15 +0000 Subject: [PATCH 02/14] fix Signed-off-by: Kaining Zhong --- transformer_engine/common/CuTeDSL/__init__.py | 8 + .../common/CuTeDSL/cast/mxfp8/__init__.py | 1 + .../cast/mxfp8/group_quantize_mxfp8.py | 155 ++++++----------- .../common/cast/dispatch/quantize.cuh | 4 +- .../mxfp8/group_quantize_mxfp8_cutedsl.cuh | 158 ++++++++++-------- 5 files changed, 155 insertions(+), 171 deletions(-) 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 index 11c288a7d36..116b07200f6 100644 --- a/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py +++ b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py @@ -110,12 +110,14 @@ 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): - if dtype not in ("fp32", "fp16", "bf16"): - raise ValueError(f"unknown input dtype {dtype!r}; expected fp32|fp16|bf16") + 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 ("fp8_e4m3fn", "fp8_e5m2"): - raise ValueError(f"unknown FP8 dtype {fp8_dtype!r}; expected fp8_e4m3fn|fp8_e5m2") + 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): @@ -131,7 +133,7 @@ def __init__(self, dtype: str, fp8_dtype: str, rowwise: bool, colwise: bool, sha # Mirrors `is_single_tensor` in group_quantize_mxfp8.cuh. self.IS_SINGLE_TENSOR = shape_rep in (SAME_BOTH_DIMS, VARYING_FIRST_DIM) self.MAX_NORM_RCP = ( - FP8E4M3_MAX_NORM_RCP if fp8_dtype == "fp8_e4m3fn" else FP8E5M2_MAX_NORM_RCP + FP8E4M3_MAX_NORM_RCP if fp8_dtype == "Float8E4M3" else FP8E5M2_MAX_NORM_RCP ) def __str__(self): @@ -1020,66 +1022,6 @@ def g1d(dtype, align=4): ) -# TEMPORARY (demo only): same as compile_cutedsl_function_from_cfg but with every -# extent pinned to a compile-time constant, so traced layouts print concrete numbers -# (e.g. `(384,256):(1@1,1@0)`) instead of `?{div=128}`. Delete once done inspecting. -def compile_cutedsl_function_from_cfg_static( - cfg: MXFP8GroupQuantizeConfig, - M_total: int, - N: int, - num_tensors: int, - scale_row_numel: int, - scale_col_numel: int, -): - """Compile with fully static shapes. Only accepts inputs of exactly these extents.""" - logical_shape = (M_total, N) - - def g2d(dtype, align=16): - return cute.runtime.make_fake_compact_tensor( - dtype, - logical_shape, - stride_order=(1, 0), - memspace=cute.AddressSpace.gmem, - assumed_align=align, - ) - - def g1d(dtype, numel, align=4): - return cute.runtime.make_fake_compact_tensor( - dtype, - (numel,), - stride_order=(0,), - memspace=cute.AddressSpace.gmem, - assumed_align=align, - ) - - scale_dtype = cutlass.Float8E8M0FNU - tensormaps_fake = cute.runtime.make_fake_compact_tensor( - cutlass.Int64, - (num_tensors, NUM_WORKSPACE_SLOTS, BYTES_PER_TENSORMAP // 8), - stride_order=(2, 1, 0), - memspace=cute.AddressSpace.gmem, - assumed_align=128, - ) - - from cutlass.utils import HardwareInfo # pylint: disable=import-outside-toplevel - - sm_count = HardwareInfo().get_device_multiprocessor_count() - return cute.compile( - MXFP8GroupQuantizeKernel(cfg, sm_count), - g2d(cfg.DTYPE), - g2d(cfg.FP8_DTYPE), - g2d(cfg.FP8_DTYPE), - g1d(scale_dtype, scale_row_numel), - g1d(scale_dtype, scale_col_numel), - g1d(cutlass.Int64, num_tensors + 1, align=8), - g1d(cutlass.Int64, num_tensors, align=8), - g1d(cutlass.Int64, num_tensors, align=8), - tensormaps_fake, - cute.runtime.make_fake_stream(), - options="--enable-tvm-ffi", - ) - - def get_mxfp8_group_quantization_function( fn_name: str, dtype: str, @@ -1088,49 +1030,58 @@ def get_mxfp8_group_quantization_function( colwise: bool, shape_rep: str, ) -> bool: - """Compile the grouped MXFP8 quantize kernel for this config and register it in the - TVM-FFI global registry under EXACTLY `fn_name`. Returns True on success (the C++ - dispatcher then fetches it with GetGlobal(fn_name)); False if unsupported, so the - caller falls back to the CUDA C++ grouped kernel. + """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. """ - 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, - ) - 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 + # 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, + ) + 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) - try: + 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 failed, " - "falling back to the CUDA C++ kernel: %s", + "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 - tvm_ffi.register_global_func(fn_name, compiled, override=True) - return True diff --git a/transformer_engine/common/cast/dispatch/quantize.cuh b/transformer_engine/common/cast/dispatch/quantize.cuh index b38b6fc12d1..7443d174bfc 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,7 +525,8 @@ void group_quantize_fwd_helper(const NVTEGroupedTensor input, NVTEGroupedTensor break; } case NVTE_MXFP8_1D_SCALING: { - #ifdef NVTE_WITH_CUTEDSL + bool quantized_with_cutedsl = false; +#ifdef NVTE_WITH_CUTEDSL quantized_with_cutedsl = cutedsl_backend::mxfp8_group_quantize_cutedsl( diff --git a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh index ca33f0b25ef..b5b72ad20fc 100644 --- a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh +++ b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh @@ -12,21 +12,21 @@ #include #include +#include #include #include #include #include "../../common.h" #include "../../tvm_ffi_bridge.h" -#include "../../utils.cuh" // ShapeRepresentation -#include "../core/common.cuh" -#include "group_quantize_mxfp8.cuh" +#include "../../util/cuda_runtime.h" +#include "../../utils.cuh" // ShapeRepresentation +#include "../core/grouped_tma.cuh" // dispatch::common::MAX_SUPPORTED_TENSOR_DESCRIPTORS namespace transformer_engine { namespace cutedsl_backend { -// te_dtype_to_str, DLTensorWrapper, TVMFFICentral all live in -// transformer_engine::tvm_ffi_bridge (tvm_ffi_bridge.h). +// 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) { @@ -50,16 +50,23 @@ struct MXFP8GroupQuantConfig { bool rowwise; // If quantize rowwisely bool colwise; // If quantize columnwisely ShapeRepresentation shape_rep; // How the member shapes vary across the group - - constexpr uint32_t to_id() const { - static_assert(static_cast(DType::kNumTypes) <= 256, - "DType no longer fits in the 8 bits to_id() gives it."); - return static_cast(dtype) | (static_cast(fp8_dtype) << 8) | - (static_cast(rowwise) << 16) | (static_cast(colwise) << 17) | - (static_cast(shape_rep) << 18); + 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 [9:8] (2 used), shape_rep [11:10] (2 used), and SM architecture [20:12] + // (9 used/reserved). Bits [31:21] are 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."); + 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(shape_rep) << 10) | (sm_arch << 12); } - std::optional get_kernel() const { + std::optional get_kernel() const { static TVMFFIConfigCache &cache = TVMFFIConfigCache::create(); return cache.get_or_load(*this); } @@ -68,11 +75,14 @@ struct MXFP8GroupQuantConfig { // compiled and registered on a cache miss. std::string to_key() const { std::string key; - key.reserve(64); // longest: cutedsl_group_mxfp8_bf16_fp8_e4m3fn_1_1_varying_first_dim - key.append("cutedsl_group_mxfp8_") - .append(te_dtype_to_str(dtype)) + key.reserve( + 80); // longest: cutedsl_group_mxfp8_smXXX_BFloat16_Float8E4M3_1_1_varying_first_dim + key.append("cutedsl_group_mxfp8_sm") + .append(std::to_string(sm_arch)) + .append("_") + .append(to_string(dtype)) .append("_") - .append(te_dtype_to_str(fp8_dtype)) + .append(to_string(fp8_dtype)) .append("_") .append(rowwise ? "1" : "0") .append("_") @@ -88,8 +98,8 @@ struct MXFP8GroupQuantConfig { return false; } tvm::ffi::Any result = - (*entrypoint)(tvm::ffi::String(fn_name), tvm::ffi::String(te_dtype_to_str(dtype)), - tvm::ffi::String(te_dtype_to_str(fp8_dtype)), rowwise, colwise, + (*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))); return result.try_cast().value_or(false); } @@ -100,29 +110,39 @@ struct MXFP8GroupQuantConfig { // CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py. constexpr size_t kGroupTensorMapSlots = 4; 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[dispatch::common::MAX_SUPPORTED_TENSOR_DESCRIPTORS] - [kGroupTensorMapSlots][kInt64PerTensorMap]; + alignas(128) int64_t tensor_maps[kMaxGroupTensors][kGroupTensorMapSlots][kInt64PerTensorMap]; // Stand-in for the offsets / first_dims / last_dims arrays a given shape representation // does not carry: the kernel takes all three unconditionally but only dereferences the // ones its representation uses, so the contents are never read. Sized num_tensors + 1 // for the CSR offsets array, the longest of the three. - int64_t unused_dims[dispatch::common::MAX_SUPPORTED_TENSOR_DESCRIPTORS + 1]; + int64_t unused_dims[kMaxGroupTensors + 1]; }; -// One workspace per translation unit, mirroring `g_tensor_maps` on the CUDA path -- and -// sharing its caveat that two grouped quantize calls in flight on different streams would -// overwrite each other's descriptors. +// 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; -inline GroupDescriptorWorkspace *group_descriptor_workspace_ptr() { - static GroupDescriptorWorkspace *const ptr = [] { +// 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)); - return static_cast(p); - }(); - return ptr; + cache[device_id] = static_cast(p); + }); + return cache[device_id]; } inline NVTEBasicTensor make_basic_tensor(void *dptr, DType dtype, @@ -136,8 +156,6 @@ inline NVTEBasicTensor make_basic_tensor(void *dptr, DType dtype, inline bool mxfp8_group_quantize_cutedsl(const MXFP8GroupQuantConfig &config, const GroupedTensor *input_tensor, GroupedTensor *output_tensor, cudaStream_t stream) { - using namespace dispatch::mxfp8::group_quantize_kernel; - 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]; @@ -154,12 +172,19 @@ inline bool mxfp8_group_quantize_cutedsl(const MXFP8GroupQuantConfig &config, ", ", kScaleDimX, ")."); 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; + } - std::optional group_quant_func_opt = config.get_kernel(); + 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(); // Both output directions are handed to the kernel unconditionally: the compiled @@ -177,42 +202,36 @@ inline bool mxfp8_group_quantize_cutedsl(const MXFP8GroupQuantConfig &config, // 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}; - const NVTEBasicTensor x_bt = - make_basic_tensor(input_tensor->data.dptr, input_tensor->dtype(), logical_shape); - const NVTEBasicTensor o_row_bt = make_basic_tensor(data_row.dptr, data_row.dtype, logical_shape); - const NVTEBasicTensor o_col_bt = make_basic_tensor(data_col.dptr, data_col.dtype, logical_shape); - DLTensorWrapper mX(x_bt), mO_row(o_row_bt), mO_col(o_col_bt); + DLTensorWrapper mX( + make_basic_tensor(input_tensor->data.dptr, input_tensor->dtype(), logical_shape), true, + device_index); + DLTensorWrapper mO_row(make_basic_tensor(data_row.dptr, data_row.dtype, logical_shape), true, + device_index); + DLTensorWrapper mO_col(make_basic_tensor(data_col.dptr, data_col.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. - const NVTEBasicTensor s_row_bt = - make_basic_tensor(scale_row.dptr, scale_row.dtype, {scale_row.numel()}); - const NVTEBasicTensor s_col_bt = - make_basic_tensor(scale_col.dptr, scale_col.dtype, {scale_col.numel()}); - DLTensorWrapper mS_row(s_row_bt, /*flatten_2D=*/false), mS_col(s_col_bt, /*flatten_2D=*/false); - - const SimpleTensor &offsets = output_tensor->tensor_offsets; - const SimpleTensor &first_dims = output_tensor->first_dims; - const SimpleTensor &last_dims = output_tensor->last_dims; - const NVTEBasicTensor offsets_bt = make_basic_tensor( - offsets.has_data() ? offsets.dptr : static_cast(workspace->unused_dims), - DType::kInt64, {num_tensors + 1}); - const NVTEBasicTensor first_dims_bt = make_basic_tensor( - first_dims.has_data() ? first_dims.dptr : static_cast(workspace->unused_dims), - DType::kInt64, {num_tensors}); - const NVTEBasicTensor last_dims_bt = make_basic_tensor( - last_dims.has_data() ? last_dims.dptr : static_cast(workspace->unused_dims), - DType::kInt64, {num_tensors}); - DLTensorWrapper mOffsets(offsets_bt, /*flatten_2D=*/false), - mFirstDims(first_dims_bt, /*flatten_2D=*/false), - mLastDims(last_dims_bt, /*flatten_2D=*/false); + 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 dims_or_unused = [&](const SimpleTensor &t, size_t numel) { + void *dptr = t.has_data() ? t.dptr : static_cast(workspace->unused_dims); + return DLTensorWrapper(make_basic_tensor(dptr, DType::kInt64, {numel}), false, device_index); + }; + DLTensorWrapper mOffsets = dims_or_unused(output_tensor->tensor_offsets, num_tensors + 1); + DLTensorWrapper mFirstDims = dims_or_unused(output_tensor->first_dims, num_tensors); + DLTensorWrapper mLastDims = dims_or_unused(output_tensor->last_dims, num_tensors); // 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. - const NVTEBasicTensor tensormaps_bt = + DLTensorWrapper mTensormaps( make_basic_tensor(static_cast(workspace->tensor_maps), DType::kInt64, - {num_tensors, kGroupTensorMapSlots, kInt64PerTensorMap}); - DLTensorWrapper mTensormaps(tensormaps_bt, /*flatten_2D=*/false); + {num_tensors, kGroupTensorMapSlots, kInt64PerTensorMap}), + false, device_index); // stream is a tvm-ffi opaque "handle"; pass it as void*. (*group_quant_func_opt)(&mX, &mO_row, &mO_col, &mS_row, &mS_col, &mOffsets, &mFirstDims, @@ -262,17 +281,20 @@ bool mxfp8_group_quantize_cutedsl(const GroupedTensor *input_tensor, const Tenso shape_rep = ShapeRepresentation::VARYING_LAST_DIM; } else if (output_tensor->all_same_last_dim()) { shape_rep = ShapeRepresentation::VARYING_FIRST_DIM; - } else { - // VARYING_BOTH_DIMS: the logical shape is [1, total], which is not tileable. + } else if (output_tensor->varying_both_dims()) { + shape_rep = ShapeRepresentation::VARYING_BOTH_DIMS; + } + if (shape_rep == ShapeRepresentation::VARYING_BOTH_DIMS) { + // The logical shape is [1, total], which is not tileable. maybe_warn_cutedsl_not_chosen("groups with both dimensions varying are not supported."); return false; } + // 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. - if (input_tensor->num_tensors > dispatch::common::MAX_SUPPORTED_TENSOR_DESCRIPTORS) { - maybe_warn_cutedsl_not_chosen("the group has more than ", - dispatch::common::MAX_SUPPORTED_TENSOR_DESCRIPTORS, - " tensors."); + if (input_tensor->num_tensors == 0 || input_tensor->num_tensors > kMaxGroupTensors) { + maybe_warn_cutedsl_not_chosen("the group size ", input_tensor->num_tensors, + " is not between 1 and ", kMaxGroupTensors, "."); return false; } From 221143e67a376fcaa5b17cbcbd53107dc86322aa Mon Sep 17 00:00:00 2001 From: Kaining Zhong Date: Tue, 29 Sep 2026 21:41:49 +0000 Subject: [PATCH 03/14] fix Signed-off-by: Kaining Zhong --- .../cast/mxfp8/group_quantize_mxfp8.py | 24 ++++--------------- .../CuTeDSL/cast/mxfp8/quantize_mxfp8.py | 16 +++++++++++++ .../mxfp8/group_quantize_mxfp8_cutedsl.cuh | 21 +++++++++++++--- 3 files changed, 39 insertions(+), 22 deletions(-) diff --git a/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py index 116b07200f6..c0d03d37bc7 100644 --- a/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py +++ b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py @@ -30,9 +30,8 @@ per-tensor grid above is what replaced it on both sides. Mechanics that provably yield the same bytes may differ: the mbarrier pipeline is -expressed with PipelineTmaAsync instead of hand-rolled mbarriers, and -out-of-bounds scale padding is skipped rather than explicitly zeroed (the CUDA -kernel writes 0 there; downstream only consumes the meaningful region). +expressed with PipelineTmaAsync instead of hand-rolled mbarriers. As in CUDA, the +scales of out-of-bounds columns in a chunk (the scale-row padding) are written as 0. Scope: cast-only (no dbias / activation / dact / amax), compact (non-swizzled) scales, rowwise and/or colwise. VARYING_BOTH_DIMS is not handled here (its @@ -44,9 +43,7 @@ and for VARYING_FIRST_DIM the per-tensor extents live in device memory, so the C++ bridge cannot check them either without a sync. A violating group mis-tiles silently here where CUDA raises. -""" -""" Measured, deliberately NOT changed: - sO_row and sO_col are both allocated unconditionally, where CUDA sizes only the direction in use. Sizing them conditionally does work -- ncu confirms the shared-memory @@ -71,7 +68,7 @@ from cutlass import pipeline from cutlass import Boolean, Int32, Int64, Float8E8M0FNU from cutlass.cute.nvgpu import cpasync -from cutlass.tensor_utils import TensorMapManager, TensorMapUpdateMode +from cutlass.utils import TensorMapManager, TensorMapUpdateMode from cuda.bindings.driver import CUstream # pylint: disable=no-name-in-module import tvm_ffi @@ -222,14 +219,11 @@ def __call__( (self.BUFF_DIM_Y, self.BUFF_DIM_X), order=(1, 0) ) cta_tiler = (self.BUFF_DIM_Y, self.BUFF_DIM_X) - print(f"mx={mX}, smem_tile_layout={smem_tile_layout}, cta_tiler={cta_tiler}\n") 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 ) - print(f"tma_atom_x={tma_atom_x}\n") - print(f"tma_src={tma_src}\n") op_store = cpasync.CopyBulkTensorTileS2GOp() 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 @@ -489,7 +483,6 @@ class SharedStorage: stride=((self.BUFF_DIM_X, 1), self.BUFF_DIM_Y * self.BUFF_DIM_X), ) sX = storage.sX.get_tensor(tile_layout) - print(f"sX={sX}\n") sO_row = storage.sO_row.get_tensor(tile_layout) sO_col = storage.sO_col.get_tensor(tile_layout) @@ -520,11 +513,6 @@ class SharedStorage: tXsO_col, tXgO_col = cpasync.tma_partition( tma_atom_out_col, 0, cute.make_layout(1), sO_col, gO_col_tiled ) - print(f"tma_atom_x={tma_atom_x}\n") - print(f"tma_src={tma_src}\n") - print(f"gX_tiled={gX_tiled}\n") - print(f"tXsX={tXsX}\n") - print(f"tXgX={tXgX}\n") tmap = TensorMapManager(TensorMapUpdateMode.GMEM, BYTES_PER_TENSORMAP) cute.arch.sync_threads() @@ -536,8 +524,6 @@ class SharedStorage: 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. - # 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) # Block's offset and id in this individual tensor / global single tensor block_offset_Y = Int32(0) @@ -867,7 +853,7 @@ def _process_block( SWIZZLE=False, TILE_X=self.BUFF_DIM_X, TILE_Y=self.BUFF_DIM_Y, - SKIP_MASKING=False, + ZERO_OOB_SCALES=True, ) if cutlass.const_expr(cfg.ROWWISE): quantize_rowwise_mxfp8( @@ -888,7 +874,7 @@ def _process_block( WAVES=self.WAVES, THREADS_PER_BANK=self.THREADS_PER_BANK, PACK_SIZE=self.PACK_SIZE, - SKIP_MASKING=False, + ZERO_OOB_SCALES=True, ) # Force consumer's write to SMEM to be visible to TMA stores later diff --git a/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py b/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py index 824ed18aa59..a59b9745a24 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,9 @@ 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, ): """Quantize one SMEM tile colwise to MXFP8 (per-column 32-elt block scales); returns (amax, dbias_partial).""" tidx, _, _ = cute.arch.thread_idx() @@ -541,6 +552,11 @@ 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): + # Swizzled layouts pad through derive_swizzled_scale_layout instead. + assert not SWIZZLE + if tile_row_start < M and scale_col >= N: + 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 diff --git a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh index b5b72ad20fc..41d4b23cd17 100644 --- a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh +++ b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh @@ -292,9 +292,24 @@ bool mxfp8_group_quantize_cutedsl(const GroupedTensor *input_tensor, const Tenso // 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. - if (input_tensor->num_tensors == 0 || input_tensor->num_tensors > kMaxGroupTensors) { - maybe_warn_cutedsl_not_chosen("the group size ", input_tensor->num_tensors, - " is not between 1 and ", kMaxGroupTensors, "."); + 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; } From aa29f69349d9d22ff9b53fee9c58c8ed77a51c99 Mon Sep 17 00:00:00 2001 From: Kaining Zhong Date: Wed, 30 Sep 2026 18:47:15 +0000 Subject: [PATCH 04/14] fix Signed-off-by: Kaining Zhong --- .../mxfp8/test_mxfp8_cutedsl_backend.py | 219 ++++ .../cast/mxfp8/group_quantize_mxfp8.py | 1015 +++++++++++------ .../CuTeDSL/cast/mxfp8/quantize_mxfp8.py | 12 +- .../common/cast/dispatch/quantize.cuh | 20 +- .../mxfp8/group_quantize_mxfp8_cutedsl.cuh | 199 +++- 5 files changed, 1075 insertions(+), 390 deletions(-) diff --git a/tests/pytorch/mxfp8/test_mxfp8_cutedsl_backend.py b/tests/pytorch/mxfp8/test_mxfp8_cutedsl_backend.py index f409f76993d..c3ab0bfa8c2 100644 --- a/tests/pytorch/mxfp8/test_mxfp8_cutedsl_backend.py +++ b/tests/pytorch/mxfp8/test_mxfp8_cutedsl_backend.py @@ -370,3 +370,222 @@ 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), + # 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)]), + # N ends in a partial 32-element scale block. + ("varying_first_n144", "varying_first_dim", [(128, 144), (384, 144)]), + ("varying_last", "varying_last_dim", [(256, 128), (256, 384), (256, 256)]), + ("varying_both", "varying_both_dims", [(128, 256), (256, 128), (384, 512)]), +] +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, dbias_cuda), 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 + if block_size[0] != 1 and shape_rep in ("varying_last_dim", "varying_both_dims"): + pytest.skip( + "The CUDA kernel's GEMM-swizzled colwise scale index double-counts the tensor base" + " for varying last dims; the CuTeDSL backend leaves these configs to it." + ) + 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", [torch.bfloat16, torch.float32], 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/cast/mxfp8/group_quantize_mxfp8.py b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py index c0d03d37bc7..26e36db8100 100644 --- a/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py +++ b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py @@ -33,16 +33,23 @@ expressed with PipelineTmaAsync instead of hand-rolled mbarriers. As in CUDA, the scales of out-of-bounds columns in a chunk (the scale-row padding) are written as 0. -Scope: cast-only (no dbias / activation / dact / amax), compact (non-swizzled) -scales, rowwise and/or colwise. VARYING_BOTH_DIMS is not handled here (its -logical shape [1, total] is not tileable); the C++ bridge falls back to CUDA. - -Known gap vs CUDA: the CUDA kernel validates on device that every member's first -dim is a multiple of 128 (get_tensor_rows_num -> NVTE_DEVICE_ERROR). This kernel -shares the precondition but cannot check it -- CuTeDSL has no device-side assert, -and for VARYING_FIRST_DIM the per-tensor extents live in device memory, so the -C++ bridge cannot check them either without a sync. A violating group mis-tiles -silently here where CUDA raises. +Scope: everything group_quantize_mxfp8.cuh covers except 2D block scaling -- the +cast-noop flag, fused activation (IS_ACT) and activation derivative (IS_DACT), dbias, +compact and GEMM-swizzled scales, rowwise and/or colwise, and all four shape +representations. Differences from CUDA: + * the grouped amax pointer is accepted and left untouched, as the CUDA kernel does; + * GEMM-swizzled colwise scales are only produced for the single-tensor reps. For + VARYING_LAST_DIM / VARYING_BOTH_DIMS the CUDA kernel adds the tensor base to the + colwise swizzled index twice, which is only in bounds for the first member, so the + C++ bridge falls back to CUDA instead of reproducing it; + * dbias partial sums are accumulated in the CUDA kernel's order (a running column sum + with colwise output, otherwise per-thread sums reduced across the CTA), and the C++ + bridge reduces the workspace with the same grouped_reduce_dbias. + +Like the CUDA kernel, every member's first dim must be a multiple of 128 (and, for the +varying-last reps, its last dim too). The kernel prints the same diagnostics as +get_tensor_rows_num / get_tensor_cols_num when a group violates this and, like +NVTE_DEVICE_ERROR in a release build, carries on. Measured, deliberately NOT changed: - sO_row and sO_col are both allocated unconditionally, where CUDA sizes only the @@ -61,12 +68,12 @@ import logging import os -from typing import Type +from typing import Optional, Type import cutlass from cutlass import cute from cutlass import pipeline -from cutlass import Boolean, Int32, Int64, Float8E8M0FNU +from cutlass import Boolean, Float32, Int32, Int64, Float8E8M0FNU from cutlass.cute.nvgpu import cpasync from cutlass.utils import TensorMapManager, TensorMapUpdateMode from cuda.bindings.driver import CUstream # pylint: disable=no-name-in-module @@ -78,8 +85,13 @@ ) 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, quantize_rowwise_mxfp8, quantize_colwise_mxfp8, ) @@ -89,8 +101,9 @@ THREADS_PER_WARP = 32 BYTES_PER_TENSORMAP = 128 -# Descriptor slots per tensor: input, rowwise output, colwise output. -NUM_TENSORMAPS = 3 +# 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 @@ -100,13 +113,30 @@ SAME_BOTH_DIMS = "same_both_dims" VARYING_FIRST_DIM = "varying_first_dim" VARYING_LAST_DIM = "varying_last_dim" -SUPPORTED_SHAPE_REPS = (SAME_BOTH_DIMS, VARYING_FIRST_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): + 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) @@ -133,10 +163,48 @@ def __init__(self, dtype: str, fp8_dtype: str, rowwise: bool, colwise: bool, sha FP8E4M3_MAX_NORM_RCP if fp8_dtype == "Float8E4M3" else FP8E5M2_MAX_NORM_RCP ) + self.WITH_GEMM_SWIZZLED_SCALES = with_gemm_swizzled_scales + if with_gemm_swizzled_scales and colwise and not self.IS_SINGLE_TENSOR: + # The CUDA kernel offsets the colwise swizzled scales of these representations by + # the tensor base twice, which only lands inside the buffer for the first member. + raise ValueError( + "GEMM-swizzled colwise scales are only supported for single-tensor representations" + ) + 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"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__ @@ -168,6 +236,24 @@ class MXFP8GroupQuantizeKernel: def __init__(self, cfg: MXFP8GroupQuantizeConfig, SM_COUNT: int): self.cfg = cfg self.SM_COUNT = SM_COUNT + # CastConfig widens the chunk to 128x256, i.e. STAGES_X = 2 tiles of + # BUFF_DIM_X columns, each traversed in STAGES row stages before moving right. + self.STAGES_X = 2 if cfg.SHAPE_REP == VARYING_BOTH_DIMS else 1 + self.CHUNK_WIDTH = self.CHUNK_DIM_X * self.STAGES_X + # 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) + # Like IS_CACHED_ACT_OP in CUDA: with both directions, the colwise pass caches the + # activation in the input tile for the rowwise pass. 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" + ) + # CUDA reduces dbias in the colwise pass when there is one, 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 # ---------------------------------------------------------------- helpers @cute.jit @@ -176,16 +262,80 @@ def _tensor_rows_cols( ): """Get the shape (rows, cols) of the tensor by tensor_id.""" cfg = self.cfg - if cutlass.const_expr(cfg.SHAPE_REP == VARYING_FIRST_DIM): + 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 == VARYING_LAST_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) return rows, cols + @cute.jit + def _find_tensor_from_offsets(self, mOffsets, num_tensors, offset: Int64): + """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, base: Int64, layout): + """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, base: Int64, rows, cols): + """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.BUFF_DIM_Y, self.BUFF_DIM_X // MXFP8_BLOCK_SCALING_SIZE) + ) + + @cute.jit + def _colwise_scales(self, mS_col, base: Int64, rows, cols): + """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.BUFF_DIM_Y // MXFP8_BLOCK_SCALING_SIZE, self.BUFF_DIM_X) + ) + # ------------------------------------------------------------ entry point @cute.jit def __call__( @@ -196,9 +346,12 @@ def __call__( mS_row: cute.Tensor, mS_col: cute.Tensor, mOffsets: cute.Tensor, # int64[num_tensors + 1], CSR element offsets - mFirstDims: cute.Tensor, # int64[num_tensors] (VARYING_FIRST_DIM) - mLastDims: cute.Tensor, # int64[num_tensors] (VARYING_LAST_DIM) - mTensormaps: cute.Tensor, # int64[num_tensors, NUM_TENSORMAPS, 16] + mFirstDims: cute.Tensor, # int64[num_tensors] (VARYING_FIRST_DIM / VARYING_BOTH_DIMS) + mLastDims: 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, ): if cutlass.const_expr(CUTEDSL_DEBUG_LOGGING): @@ -211,8 +364,7 @@ def __call__( # pass a length-num_tensors stub for SAME_BOTH_DIMS (where the offsets array is # unused), which would make `mOffsets.shape[0] - 1` read one too few and divide by # zero at num_tensors == 1. The per-tensor descriptor workspace is num_tensors long - # by construction -- one slot set per member -- so it is the reliable source, and it - # is only consulted on the multi-tensor path that actually needs the workspace. + # by construction -- one slot set per member -- so it is the reliable source. num_tensors = mTensormaps.shape[0] smem_tile_layout = cute.make_ordered_layout( @@ -224,6 +376,12 @@ def __call__( 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 + ) op_store = cpasync.CopyBulkTensorTileS2GOp() 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 @@ -236,33 +394,30 @@ def __call__( # How many blocks does the grouped tensor have in both directions work_blocks_X = cute.ceil_div(Int32(last_logical_dim), self.CHUNK_DIM_X) work_blocks_Y = cute.ceil_div(Int32(first_logical_dim), self.CHUNK_DIM_Y) - # Each CTA handles a block from one individual tensor + # Each CTA handles one chunk. With every member's rows a multiple of CHUNK_DIM_Y, + # this linear order is the CUDA kernel's (X, Y-in-tensor, tensor) order. grid = [work_blocks_X * work_blocks_Y, 1, 1] else: - # The work-block grid is per-tensor here, not global: each CTA derives its own - # block range from its tensor's extents (see the kernel's `else` branch), so - # work_blocks_X/Y are dead on this path -- the kernel reads them only under - # IS_SINGLE_TENSOR. Pass dummies rather than computing first*last, which is an - # ELEMENT count and would wrap Int32 past 2^31 elements. + # The work-block grid is per-tensor here: each CTA derives its own block range + # from its tensor's extents, so work_blocks_X/Y are unused on this path. work_blocks_Y = Int32(1) work_blocks_X = Int32(1) # Persistent worker count, mirroring get_launch_config() in - # group_quantize_mxfp8.cuh. There are SM_COUNT * STATIC_PERSISTENT_BLOCKS_PER_SM - # workers in total, split evenly across tensors -- but clamped to the average - # number of chunks a tensor actually holds, so a group of small tensors does not - # launch CTAs that only reach the `first_block_id >= blocks_in_tensor` early-out. + # group_quantize_mxfp8.cuh: SM_COUNT * STATIC_PERSISTENT_BLOCKS_PER_SM workers + # split evenly across tensors, clamped to the average number of chunks a tensor + # holds. The element count would wrap Int32, so the estimate + # DIVUP(elts_total, CHUNK_DIM_Y * TILE_DIM_X) is formed without it. n_tensors = cutlass.max(Int32(num_tensors), Int32(1)) # never divide by zero - # CUDA's DIVUP(elts_total, CHUNK_DIM_Y * TILE_DIM_X) / STAGES_X, where TILE_DIM_X - # is BUFF_DIM_X here and STAGES_X is 1 for every rep this kernel supports. The - # element product would wrap Int32 past 2^31 elements, so divide the first extent - # by CHUNK_DIM_Y up front -- it is exact, the kernel is compiled with - # sym_int32(divisibility=CHUNK_DIM_Y) on that extent -- and keep the whole - # estimate in Int32. Feeding an Int64 into the grid poisons the tile arithmetic - # downstream ('cute.make_tile' expects width=32). - estimated_work_blocks = cute.ceil_div( - (Int32(first_logical_dim) // self.CHUNK_DIM_Y) * Int32(last_logical_dim), - self.ELTS_PER_CHUNK // self.CHUNK_DIM_Y, - ) + if cutlass.const_expr(cfg.SHAPE_REP == VARYING_BOTH_DIMS): + # The logical shape is [1, total]. + estimated_work_blocks = cute.ceil_div(Int32(last_logical_dim), self.ELTS_PER_CHUNK) + else: + # Exact: the first extent is compiled with divisibility CHUNK_DIM_Y. + estimated_work_blocks = cute.ceil_div( + (Int32(first_logical_dim) // self.CHUNK_DIM_Y) * Int32(last_logical_dim), + self.ELTS_PER_CHUNK // self.CHUNK_DIM_Y, + ) + estimated_work_blocks = cute.ceil_div(estimated_work_blocks, self.STAGES_X) requested_workers_per_tensor = cutlass.max( Int32(1), Int32(self.SM_COUNT * self.STATIC_PERSISTENT_BLOCKS_PER_SM) // n_tensors, @@ -273,17 +428,15 @@ def __call__( workers_per_tensor = cutlass.min( requested_workers_per_tensor, average_work_blocks_per_tensor ) - # Each group of workers_per_tensor CTAs serves one tensor, and we launch - # num_tensors such groups to cover the whole group. grid = [workers_per_tensor, Int32(num_tensors), 1] - # Only manually create descriptors for the non-single-tensor case because we will need to manually - # overwrite the descriptors as we visit different groups + # 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, @@ -294,6 +447,7 @@ def __call__( 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( @@ -301,16 +455,18 @@ def __call__( mS_col, mOffsets, mFirstDims, - mLastDims, mTensormaps, + mNoop, + mWorkspace, first_logical_dim, last_logical_dim, num_tensors, work_blocks_X, - work_blocks_Y, 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, @@ -328,6 +484,7 @@ def update_descriptors_kernel( mX, mO_row, mO_col, + mActInput, mOffsets, mFirstDims, mLastDims, @@ -338,6 +495,7 @@ def update_descriptors_kernel( tma_atom_x, tma_atom_orow, tma_atom_ocol, + tma_atom_act, ): """One CTA per tensor: point that tensor's TMA descriptors at its own block. @@ -345,11 +503,28 @@ def update_descriptors_kernel( """ cfg = self.cfg tensor_id, _, _ = cute.arch.block_idx() + tidx, _, _ = cute.arch.thread_idx() rows, cols = self._tensor_rows_cols( tensor_id, mFirstDims, mLastDims, first_logical_dim, last_logical_dim ) base_elts = 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. + 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, + ) + # Publish this tensor's geometry for the main kernel (written even when empty). meta = mTensormaps[(tensor_id, META_SLOT, None)] meta[0] = Int64(rows) @@ -360,69 +535,51 @@ def update_descriptors_kernel( 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 jobs via job_has_work). + # so skip (the main kernel skips these tensors as well). if rows > 0 and cols > 0: - gX = cute.make_tensor( - cute.make_ptr( - dtype, - mX.iterator.toint() + base_elts * (dtype.width // 8), - cute.AddressSpace.gmem, - assumed_align=16, - ), - cute.make_layout((rows, cols), stride=(cols, 1)), - ) - gO_row = cute.make_tensor( - cute.make_ptr( - cfg.FP8_DTYPE, - mO_row.iterator.toint() + base_elts, - cute.AddressSpace.gmem, - assumed_align=16, - ), - cute.make_layout((rows, cols), stride=(cols, 1)), - ) - gO_col = cute.make_tensor( - cute.make_ptr( - cfg.FP8_DTYPE, - mO_col.iterator.toint() + base_elts, - cute.AddressSpace.gmem, - assumed_align=16, - ), - cute.make_layout((rows, cols), stride=(cols, 1)), - ) + member_layout = cute.make_layout((rows, cols), stride=(cols, 1)) + + def member_view(tensor, elt_dtype): + return cute.make_tensor( + cute.make_ptr( + elt_dtype, + tensor.iterator.toint() + base_elts * (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) + descs.append(desc_orow) tmap.init_tensormap_from_atom(tma_atom_orow, desc_orow, 0) if cutlass.const_expr(cfg.COLWISE): + views.append(member_view(mO_col, cfg.FP8_DTYPE)) + atoms.append(tma_atom_ocol) + descs.append(desc_ocol) tmap.init_tensormap_from_atom(tma_atom_ocol, desc_ocol, 0) + if cutlass.const_expr(cfg.WITH_DACT): + views.append(member_view(mActInput, dtype)) + atoms.append(tma_atom_act) + descs.append(desc_act) + tmap.init_tensormap_from_atom(tma_atom_act, desc_act, 0) tmap.fence_tensormap_initialization() - - if cutlass.const_expr(cfg.ROWWISE and cfg.COLWISE): - tmap.update_tensormap( - (gX, gO_row, gO_col), - (tma_atom_x, tma_atom_orow, tma_atom_ocol), - (desc_x, desc_orow, desc_ocol), - 0, - (), # smem staging is unused in GMEM update mode - ) - elif cutlass.const_expr(cfg.ROWWISE): - tmap.update_tensormap( - (gX, gO_row), - (tma_atom_x, tma_atom_orow), - (desc_x, desc_orow), - 0, - (), # smem staging is unused in GMEM update mode - ) - else: - tmap.update_tensormap( - (gX, gO_col), - (tma_atom_x, tma_atom_ocol), - (desc_x, desc_ocol), - 0, - (), # smem staging is unused in GMEM update mode - ) + tmap.update_tensormap( + tuple(views), + tuple(atoms), + tuple(descs), + 0, + (), # smem staging is unused in GMEM update mode + ) # ------------------------------------------------------------ main kernel @cute.kernel @@ -432,16 +589,68 @@ def kernel( mS_col, mOffsets, mFirstDims, - mLastDims, mTensormaps, + mNoop, + mWorkspace, + first_logical_dim, + last_logical_dim, + num_tensors, + work_blocks_X, + dtype: cutlass.Constexpr[Type[cutlass.Numeric]], + 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, + ): + """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, + work_blocks_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, + mS_col, + mOffsets, + mFirstDims, + mTensormaps, + mWorkspace, first_logical_dim, last_logical_dim, num_tensors, work_blocks_X, - work_blocks_Y, dtype: cutlass.Constexpr[Type[cutlass.Numeric]], tma_atom_x, tma_src, + tma_atom_act, + tma_src_act, tma_atom_out_row, tma_dst_out_row, tma_atom_out_col, @@ -454,6 +663,18 @@ def kernel( 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, + ) + # --- shared memory (allocated once, reused across jobs) --- @cute.struct class SharedStorage: @@ -477,7 +698,8 @@ class SharedStorage: 128, ] - storage = cutlass.utils.SmemAllocator().allocate(SharedStorage) + smem = cutlass.utils.SmemAllocator() + storage = smem.allocate(SharedStorage) tile_layout = cute.make_layout( ((self.BUFF_DIM_Y, self.BUFF_DIM_X), self.PIPELINE_DEPTH), stride=((self.BUFF_DIM_X, 1), self.BUFF_DIM_Y * self.BUFF_DIM_X), @@ -486,12 +708,43 @@ class SharedStorage: sO_row = storage.sO_row.get_tensor(tile_layout) sO_col = storage.sO_col.get_tensor(tile_layout) + sActInput = None + if cutlass.const_expr(cfg.WITH_DACT): + + @cute.struct + class DactStorage: + sActInput: cute.struct.Align[ + cute.struct.MemRange[ + dtype, self.BUFF_DIM_Y * self.BUFF_DIM_X * 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.BUFF_DIM_X), stride=(DBIAS_BUFF_WIDTH, 1)) + ) + + # Grad and activation input share each stage's barrier. + tx_count = self.BUFF_DIM_Y * self.BUFF_DIM_X * dtype.width // 8 + if cutlass.const_expr(cfg.WITH_DACT): + tx_count *= 2 mainloop_pipeline = pipeline.PipelineTmaAsync.create( barrier_storage=storage.mbar.data_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=self.BUFF_DIM_Y * self.BUFF_DIM_X * dtype.width // 8, + tx_count=tx_count, cta_layout_vmnk=None, ) prod_state = pipeline.make_pipeline_state( @@ -505,6 +758,13 @@ class SharedStorage: # the descriptor is swapped per tensor and the tile coords are tensor-local. gX_tiled = cute.zipped_divide(tma_src, (self.BUFF_DIM_Y, self.BUFF_DIM_X)) 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.BUFF_DIM_Y, self.BUFF_DIM_X)) + tXsA, tXgA = cpasync.tma_partition( + tma_atom_act, 0, cute.make_layout(1), sActInput, gA_tiled + ) gO_row_tiled = cute.zipped_divide(tma_dst_out_row, (self.BUFF_DIM_Y, self.BUFF_DIM_X)) tXsO_row, tXgO_row = cpasync.tma_partition( tma_atom_out_row, 0, cute.make_layout(1), sO_row, gO_row_tiled @@ -544,9 +804,8 @@ class SharedStorage: # Which row does this block start from block_offset_Y = block_id_Y * self.CHUNK_DIM_Y if cutlass.const_expr(cfg.SHAPE_REP == VARYING_FIRST_DIM): - # Check if this block contains any valid tokens since the last offset <= logical_first_dim * logical_last_dim - # Last CSR offset == the group's total element count. Indexed off the - # array's own length so it does not depend on num_tensors. + # logical_shape may describe graph-safe capacity beyond the active tensors, + # whose total element count is the last CSR offset. total_elts = Int64(mOffsets[mOffsets.shape[0] - 1]) if Int64(block_offset_Y) * Int64(last_logical_dim) >= total_elts: has_work = Boolean(False) @@ -560,7 +819,7 @@ class SharedStorage: tensor_base = Int64(meta[2]) if tensor_rows > 0 and tensor_cols > 0: # How many blocks does this tensor have in both directions - block_columns_in_tensor = cute.ceil_div(tensor_cols, self.CHUNK_DIM_X) + block_columns_in_tensor = cute.ceil_div(tensor_cols, self.CHUNK_WIDTH) block_rows_in_tensor = cute.ceil_div(tensor_rows, self.CHUNK_DIM_Y) # How many blocks does this tensor have blocks_in_tensor = block_columns_in_tensor * block_rows_in_tensor @@ -575,10 +834,51 @@ class SharedStorage: # 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: if cutlass.const_expr(cfg.IS_SINGLE_TENSOR): # For single tensor case we don't use tensor descriptors - desc_x = desc_out_row = desc_out_col = None + 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) + col_scales = None + col_scale_row0 = block_offset_Y + 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 chunk. + member_rows = Int32(0) + member_row0 = Int32(0) + if cutlass.const_expr(cfg.SHAPE_REP == SAME_BOTH_DIMS): + member_rows = tensor_rows // Int32(num_tensors) + member_row0 = block_offset_Y // member_rows * member_rows + else: + member_id = self._find_tensor_from_offsets( + mOffsets, + num_tensors, + Int64(block_offset_Y) * Int64(tensor_cols), + ) + member_rows = Int32(mFirstDims[member_id]) + member_row0 = Int32(Int64(mOffsets[member_id]) // Int64(tensor_cols)) + col_scale_base = ( + Int64(member_row0) + * Int64(cute.round_up(tensor_cols, 128)) + // MXFP8_BLOCK_SCALING_SIZE + ) + col_scale_row0 = block_offset_Y - member_row0 + col_scale_rows = member_rows + col_scales = self._colwise_scales( + mS_col, col_scale_base, col_scale_rows, tensor_cols + ) cute.arch.sync_threads() @@ -587,28 +887,23 @@ class SharedStorage: block_id_X, tensor_rows, tensor_cols, - tensor_base, - desc_x, - desc_out_row, - desc_out_col, - mS_row, - mS_col, - first_logical_dim, + row_scales, + col_scales, + col_scale_row0, + col_scale_rows, + block_offset_Y // self.CHUNK_DIM_Y, + mWorkspace, + sDbias, + descs, tmap, warp_idx, tidx, sX, + sActInput, sO_row, sO_col, - tXsX, - tXgX, - tXsO_row, - tXgO_row, - tXsO_col, - tXgO_col, - tma_atom_x, - tma_atom_out_row, - tma_atom_out_col, + partitions, + atoms, mainloop_pipeline, prod_state, cons_state, @@ -618,6 +913,9 @@ class SharedStorage: 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 + ) # Acquire the descriptors on ONE thread, as the CUDA kernel does # (`leading_thread` in group_quantize_mxfp8.cuh); the sync_threads below # publishes it CTA-wide. Running the tensormap acquire fence on all 128 @@ -626,10 +924,22 @@ class SharedStorage: # 133 us -> 58 us), since the cost scales with threads x descriptors. 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) cute.arch.sync_threads() @@ -639,37 +949,69 @@ class SharedStorage: while not job_finished: block_id_Y_in_tensor = block_id // block_columns_in_tensor block_id_X_in_tensor = block_id % block_columns_in_tensor - self._process_block( - block_id_Y_in_tensor * self.CHUNK_DIM_Y, - block_id_X_in_tensor, - tensor_rows, - tensor_cols, - tensor_base, - desc_x, - desc_out_row, - desc_out_col, - mS_row, - mS_col, - first_logical_dim, - tmap, - warp_idx, - tidx, - sX, - sO_row, - sO_col, - tXsX, - tXgX, - tXsO_row, - tXgO_row, - tXsO_col, - tXgO_col, - tma_atom_x, - tma_atom_out_row, - tma_atom_out_col, - mainloop_pipeline, - prod_state, - cons_state, - ) + block_offset_Y_in_tensor = block_id_Y_in_tensor * self.CHUNK_DIM_Y + if cutlass.const_expr(self.STAGES_X == 1): + self._process_block( + block_offset_Y_in_tensor, + block_id_X_in_tensor, + tensor_rows, + tensor_cols, + row_scales, + col_scales, + block_offset_Y_in_tensor, + tensor_rows, + Int32(0), # dbias is only supported for single-tensor reps + mWorkspace, + sDbias, + descs, + tmap, + warp_idx, + tidx, + sX, + sActInput, + sO_row, + sO_col, + partitions, + atoms, + mainloop_pipeline, + prod_state, + cons_state, + ) + else: + # The chunk's column tiles in order, stopping at the tensor's last + # column like the CUDA kernel's stages_X = DIVUP(chunk_cols, TILE_DIM_X). + chunk_col0 = block_id_X_in_tensor * self.CHUNK_WIDTH + tiles_X = cutlass.min( + Int32(self.STAGES_X), + cute.ceil_div(tensor_cols - chunk_col0, self.BUFF_DIM_X), + ) + for stage_X in cutlass.range(tiles_X, unroll=1): + self._process_block( + block_offset_Y_in_tensor, + block_id_X_in_tensor * self.STAGES_X + stage_X, + tensor_rows, + tensor_cols, + row_scales, + col_scales, + block_offset_Y_in_tensor, + tensor_rows, + Int32(0), # dbias is only supported for single-tensor reps + mWorkspace, + sDbias, + descs, + tmap, + warp_idx, + tidx, + sX, + sActInput, + sO_row, + sO_col, + partitions, + atoms, + mainloop_pipeline, + prod_state, + cons_state, + ) # Find the next block to process block_id = block_id + block_stride if block_id >= blocks_in_tensor: @@ -686,126 +1028,98 @@ def _issue_load( prod_state, tile_y, tile_x, - tma_atom_x, - tXgX, - tXsX, + atoms, + partitions, tmap, - desc_x, + descs, ): - """Emit one 32x128 TMA load into the current pipeline buffer. + """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) - if cutlass.const_expr(self.cfg.IS_SINGLE_TENSOR): - cute.copy( - tma_atom_x, - tXgX[(None, (tile_y, tile_x))], - tXsX[(None, prod_state.index)], - tma_bar_ptr=pipeline_obj.producer_get_barrier(prod_state), - ) - else: - # Every member shares tXgX's tile-coordinate arithmetic (the coefficients are - # just the tile size); tma_desc_ptr supplies this member's geometry. - cute.copy( - tma_atom_x, - tXgX[(None, (tile_y, tile_x))], - tXsX[(None, prod_state.index)], - tma_bar_ptr=pipeline_obj.producer_get_barrier(prod_state), - tma_desc_ptr=tmap.get_tensormap_ptr(desc_x, cute.AddressSpace.generic), - ) + 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_block( self, - block_offset_Y_in_tensor, # Row offset of this chunk (global for single-tensor, else tensor-local) - block_id_X_in_tensor, # Column-chunk index within the tensor - rows, # Number of rows in this tensor (logical shape) - cols, # Number of columns in this tensor (logical shape) - tensor_base, # Int64 element offset of this tensor in the group (0 for single-tensor) - desc_x, # Per-tensor descriptors, already acquired by the caller (None if single-tensor) - desc_out_row, - desc_out_col, - mS_row, # Grouped rowwise scales - mS_col, # Grouped colwise scales - first_logical_dim, # First dimension of the grouped tensor + block_offset_Y, # Row offset of this chunk (global for single-tensor, else tensor-local) + block_id_X, # Column-chunk index within the tensor + rows, # Rows of the rowwise-scale view (the group for single-tensor, else the tensor) + cols, # Number of columns in this tensor + row_scales, # Rowwise scales tiled per stage, rows counted like block_offset_Y + col_scales, # Colwise scales tiled per stage + col_scale_row0, # Row of this chunk in the colwise-scale view + col_scale_rows, # Rows of the colwise-scale view + dbias_row, # Row of the dbias workspace this chunk reduces into + mWorkspace, # f32 partial dbias workspace (WITH_DBIAS) + sDbias, # SMEM buffer for the rowwise dbias reduction (rowwise-only dbias) + descs, # Per-tensor descriptors (x, act, out_row, out_col), None if single-tensor tmap, # TensorMapManager for managing TMA descriptors warp_idx, tidx, - sX, # SMEM input for this block - sO_row, # SMEM rowwise output for this block - sO_col, # SMEM colwise output for this block - tXsX, # Tiled sX for TMA - tXgX, # Tiled gX for TMA - tXsO_row, # Tiled sO_row for TMA - tXgO_row, # Tiled gO_row for TMA - tXsO_col, # Tiled sO_col for TMA - tXgO_col, # Tiled gO_col for TMA - tma_atom_x, # TMA atom for input - tma_atom_out_row, # TMA atom for rowwise output - tma_atom_out_col, # TMA atom for colwise output + sX, # SMEM input ring + sActInput, # SMEM activation input ring (WITH_DACT) + sO_row, # SMEM rowwise output ring + sO_col, # SMEM colwise output ring + partitions, # TMA partitions (x, act, out_row, out_col) + atoms, # TMA atoms (x, act, out_row, out_col) mainloop_pipeline: cutlass.pipeline.PipelineTmaAsync, prod_state, cons_state, ): - """Quantize one 128x128 chunk in STAGES slices of BUFF_DIM_Y rows.""" + """Quantize one 128x128 tile of a chunk in STAGES slices of BUFF_DIM_Y rows.""" cfg = self.cfg - block_offset_X = block_id_X_in_tensor * self.CHUNK_DIM_X - - # This tensor's rowwise scales - scale_rows = Int32(first_logical_dim) if cutlass.const_expr(cfg.IS_SINGLE_TENSOR) else rows - scale_base = ( - Int64(0) - if cutlass.const_expr(cfg.IS_SINGLE_TENSOR) - else tensor_base // Int64(MXFP8_BLOCK_SCALING_SIZE) - ) - - if cutlass.const_expr(cfg.ROWWISE): - # Rowwise scale's divisibility guarantee: (128, 4) - scale_row_stride = cute.round_up(cute.ceil_div(cols, MXFP8_BLOCK_SCALING_SIZE), 4) - # Advance to this tensor's rowwise scales - mS_row_t = cute.make_tensor( - cute.make_ptr( - Float8E8M0FNU, - mS_row.iterator.toint() + scale_base, - cute.AddressSpace.gmem, - assumed_align=4, - ), - cute.make_layout((scale_rows, scale_row_stride), stride=(scale_row_stride, 1)), - ) - mS_row_tiled = cute.zipped_divide( - mS_row_t, (self.BUFF_DIM_Y, self.CHUNK_DIM_X // MXFP8_BLOCK_SCALING_SIZE) - ) - - if cutlass.const_expr(cfg.COLWISE): - # Colwise scale's divisibility guarantee: (4, 128) - scale_col_stride = cute.round_up(cols, 128) - # Advance to this tensor's colwise scales - mS_col_t = cute.make_tensor( - cute.make_ptr( - Float8E8M0FNU, - mS_col.iterator.toint() + scale_base, - cute.AddressSpace.gmem, - assumed_align=4, - ), - cute.make_layout( - (scale_rows // MXFP8_BLOCK_SCALING_SIZE, scale_col_stride), - stride=(scale_col_stride, 1), - ), - ) - mS_col_tiled = cute.zipped_divide( - mS_col_t, (self.BUFF_DIM_Y // MXFP8_BLOCK_SCALING_SIZE, self.CHUNK_DIM_X) - ) - - cute.arch.sync_threads() + _, _, 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 + block_offset_X = block_id_X * self.CHUNK_DIM_X # This chunk's coordinates in the tile grid (32x128 TMA boxes, not elements). - tile_id_Y = block_offset_Y_in_tensor // self.BUFF_DIM_Y - tile_id_X = block_id_X_in_tensor + tile_id_Y = block_offset_Y // self.BUFF_DIM_Y + tile_id_X = block_id_X + col_scale_tile_Y = col_scale_row0 // self.BUFF_DIM_Y + + # Per-chunk dbias accumulators, in the CUDA kernel's summation order: a running + # column sum over the chunk's rows (colwise), or per-thread partial sums over its + # stages that the whole CTA reduces afterwards (rowwise-only). + dbias_col = 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) # 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): @@ -815,11 +1129,10 @@ def _process_block( prod_state, tile_id_Y + prologue_stage, tile_id_X, - tma_atom_x, - tXgX, - tXsX, + atoms, + partitions, tmap, - desc_x, + descs, ) prod_state.advance() @@ -833,40 +1146,50 @@ def _process_block( 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 = tile_id_Y + stage - tile_row_start = block_offset_Y_in_tensor + stage * self.BUFF_DIM_Y if cutlass.const_expr(cfg.COLWISE): - quantize_colwise_mxfp8( + _, dbias_col = quantize_colwise_mxfp8( sX_tile, - None, + sAct_tile, sO_col[(None, cons_state.index)], - cute.flatten(mS_col_tiled[(None, (row_tile, tile_id_X))]), + cute.flatten(col_scales[(None, (col_scale_tile_Y + stage, tile_id_X))]), cfg.MAX_NORM_RCP, - tile_row_start, + (col_scale_tile_Y + stage) * self.BUFF_DIM_Y, block_offset_X, - scale_rows, + col_scale_rows, cols, - ACTIVATION=None, + ACTIVATION=cfg.ACTIVATION, DTYPE=cfg.DTYPE, FP8_DTYPE=cfg.FP8_DTYPE, - SWIZZLE=False, + SWIZZLE=cfg.WITH_GEMM_SWIZZLED_SCALES, TILE_X=self.BUFF_DIM_X, TILE_Y=self.BUFF_DIM_Y, + 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=dbias_col, ) + 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, + None if self.CACHE_ACTIVATION else sAct_tile, sO_row[(None, cons_state.index)], - cute.flatten(mS_row_tiled[(None, (row_tile, tile_id_X))]), + cute.flatten(row_scales[(None, (row_tile, tile_id_X))]), cfg.MAX_NORM_RCP, - tile_row_start, + row_tile * self.BUFF_DIM_Y, block_offset_X, - scale_rows, + rows, cols, - ACTIVATION=None, + ACTIVATION=None if self.CACHE_ACTIVATION else cfg.ACTIVATION, DTYPE=cfg.DTYPE, FP8_DTYPE=cfg.FP8_DTYPE, TILE_X=self.BUFF_DIM_X, @@ -874,6 +1197,10 @@ def _process_block( 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, ) @@ -893,69 +1220,97 @@ def _process_block( prod_state, tile_id_Y + stage + self.PIPELINE_DEPTH, tile_id_X, - tma_atom_x, - tXgX, - tXsX, + atoms, + partitions, tmap, - desc_x, + descs, ) prod_state.advance() # Write result to GMEM via TMA if warp_idx == 0: + stores = [] if cutlass.const_expr(cfg.ROWWISE): - if cutlass.const_expr(cfg.IS_SINGLE_TENSOR): - cute.copy( - tma_atom_out_row, - tXsO_row[(None, cons_state.index)], - tXgO_row[(None, (row_tile, tile_id_X))], - ) - else: - cute.copy( - tma_atom_out_row, - tXsO_row[(None, cons_state.index)], - tXgO_row[(None, (row_tile, tile_id_X))], - tma_desc_ptr=tmap.get_tensormap_ptr( - desc_out_row, cute.AddressSpace.generic - ), - ) + 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( - tma_atom_out_col, - tXsO_col[(None, cons_state.index)], - tXgO_col[(None, (row_tile, tile_id_X))], + atom, + tXs[(None, cons_state.index)], + tXg[(None, (row_tile, tile_id_X))], ) else: cute.copy( - tma_atom_out_col, - tXsO_col[(None, cons_state.index)], - tXgO_col[(None, (row_tile, tile_id_X))], - tma_desc_ptr=tmap.get_tensormap_ptr( - desc_out_col, cute.AddressSpace.generic - ), + 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): + dbias_col = self._reduce_rowwise_dbias(sDbias, tidx, dbias_row_acc) + # One partial-dbias row per chunk, as in the CUDA kernel's dbias_workspace. + dbias_x = block_offset_X + tidx + if dbias_x < cols: + mWorkspace[(dbias_row, dbias_x)] = dbias_col + + @cute.jit + def _reduce_rowwise_dbias(self, sDbias, tidx, dbias_row_acc): + """Reduce the per-thread rowwise partial sums to one sum per column, in the order of the + CUDA kernel's partial_dbias_rowwise reduction.""" + _, tv_write = cute.make_layout_tv( + thr_layout=cute.make_layout( + (self.THREADS_Y, self.THREADS_X), stride=(self.THREADS_X, 1) + ), + val_layout=cute.make_layout( + (1, MXFP8_BLOCK_SCALING_SIZE), stride=(MXFP8_BLOCK_SCALING_SIZE, 1) + ), + ) + sDbias_write = cute.composition(sDbias, tv_write) + bank_group = (tidx % THREADS_PER_WARP) // self.THREADS_PER_BANK + offset = bank_group * self.PACK_SIZE + for w in cutlass.range_constexpr(self.WAVES): + # Undo the bank-conflict rotation quantize_rowwise_mxfp8 accumulated in. + start = (w * self.PACK_SIZE + offset) % MXFP8_BLOCK_SCALING_SIZE + for i in cutlass.range_constexpr(self.PACK_SIZE): + sDbias_write[(tidx, start + i)] = dbias_row_acc[w * self.PACK_SIZE + i] + cute.arch.sync_threads() + # Thread tidx sums column tidx over the THREADS_Y partial rows. + dbias = Float32(0.0) + for i in cutlass.range_constexpr(self.THREADS_Y): + dbias += sDbias[(i, tidx)] + # The buffer is rewritten by the next chunk. + cute.arch.sync_threads() + return 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/cols likewise); MXFP8 needs the last dim divisible by 32. - sym_M = cute.sym_int32(divisibility=128) - sym_N = cute.sym_int32(divisibility=MXFP8_BLOCK_SCALING_SIZE) + # 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 - def g2d(dtype, align=16): + def g2d(dtype, shape=logical_shape, align=16): return cute.runtime.make_fake_compact_tensor( dtype, - logical_shape, + shape, stride_order=(1, 0), memspace=cute.AddressSpace.gmem, assumed_align=align, @@ -972,14 +1327,6 @@ def g1d(dtype, align=4): # 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. - in_fake = g2d(cfg.DTYPE) - out_row_fake = g2d(out_dtype) - out_col_fake = g2d(out_dtype) - scale_row_fake = g1d(scale_dtype) - scale_col_fake = g1d(scale_dtype) - offsets_fake = g1d(cutlass.Int64, align=8) - first_dims_fake = g1d(cutlass.Int64, align=8) - last_dims_fake = g1d(cutlass.Int64, align=8) tensormaps_fake = cute.runtime.make_fake_compact_tensor( cutlass.Int64, (cute.sym_int32(), NUM_WORKSPACE_SLOTS, BYTES_PER_TENSORMAP // 8), @@ -987,6 +1334,15 @@ def g1d(dtype, align=4): 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 = g2d(cfg.DTYPE) if cfg.WITH_DACT else None + workspace_fake = ( + g2d(Float32, shape=(cute.sym_int32(), cute.sym_int32()), align=4) + if cfg.WITH_DBIAS + else None + ) from cutlass.utils import HardwareInfo # pylint: disable=import-outside-toplevel @@ -994,15 +1350,18 @@ def g1d(dtype, align=4): kernel_obj = MXFP8GroupQuantizeKernel(cfg, sm_count) return cute.compile( kernel_obj, - in_fake, - out_row_fake, - out_col_fake, - scale_row_fake, - scale_col_fake, - offsets_fake, - first_dims_fake, - last_dims_fake, - tensormaps_fake, + g2d(cfg.DTYPE), # mX + g2d(out_dtype), # mO_row + g2d(out_dtype), # mO_col + g1d(scale_dtype), # mS_row + g1d(scale_dtype), # mS_col + g1d(cutlass.Int64, align=8), # mOffsets + g1d(cutlass.Int64, align=8), # mFirstDims + g1d(cutlass.Int64, align=8), # mLastDims + tensormaps_fake, # mTensormaps + noop_fake, # mNoop + act_input_fake, # mActInput + workspace_fake, # mWorkspace cute.runtime.make_fake_stream(), options="--enable-tvm-ffi", ) @@ -1015,6 +1374,11 @@ def get_mxfp8_group_quantization_function( 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 @@ -1044,6 +1408,11 @@ def get_mxfp8_group_quantization_function( 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( diff --git a/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py b/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py index a59b9745a24..9880276d396 100644 --- a/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py +++ b/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py @@ -432,6 +432,9 @@ def quantize_colwise_mxfp8( # 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() @@ -452,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 @@ -553,10 +556,11 @@ def quantize_colwise_mxfp8( else: mS_col_stage[(0, tidx)] = biased_exp_c if cutlass.const_expr(ZERO_OOB_SCALES): - # Swizzled layouts pad through derive_swizzled_scale_layout instead. - assert not SWIZZLE if tile_row_start < M and scale_col >= N: - mS_col_stage[(0, tidx)] = Uint8(0).bitcast(Float8E8M0FNU) + 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 diff --git a/transformer_engine/common/cast/dispatch/quantize.cuh b/transformer_engine/common/cast/dispatch/quantize.cuh index 7443d174bfc..4a6a4d4387c 100644 --- a/transformer_engine/common/cast/dispatch/quantize.cuh +++ b/transformer_engine/common/cast/dispatch/quantize.cuh @@ -530,8 +530,8 @@ void group_quantize_fwd_helper(const NVTEGroupedTensor input, NVTEGroupedTensor quantized_with_cutedsl = cutedsl_backend::mxfp8_group_quantize_cutedsl( - input_tensor, noop_tensor, output_tensor, quant_config_cpp.mxfp8_2d_quantization, - stream); + input_tensor, activations_tensor, noop_tensor, output_tensor, dbias_tensor, + workspace_tensor, quant_config_cpp.mxfp8_2d_quantization, stream); #endif if (!quantized_with_cutedsl) { mxfp8::group_quantize( @@ -630,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.mxfp8_2d_quantization, 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 index 41d4b23cd17..c1b1aae0edd 100644 --- a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh +++ b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh @@ -20,8 +20,9 @@ #include "../../common.h" #include "../../tvm_ffi_bridge.h" #include "../../util/cuda_runtime.h" -#include "../../utils.cuh" // ShapeRepresentation -#include "../core/grouped_tma.cuh" // dispatch::common::MAX_SUPPORTED_TENSOR_DESCRIPTORS +#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 { @@ -50,20 +51,30 @@ struct MXFP8GroupQuantConfig { 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 [9:8] (2 used), shape_rep [11:10] (2 used), and SM architecture [20:12] - // (9 used/reserved). Bits [31:21] are unused. + // 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(shape_rep) << 10) | (sm_arch << 12); + (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 { @@ -75,8 +86,8 @@ struct MXFP8GroupQuantConfig { // compiled and registered on a cache miss. std::string to_key() const { std::string key; - key.reserve( - 80); // longest: cutedsl_group_mxfp8_smXXX_BFloat16_Float8E4M3_1_1_varying_first_dim + // 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("_") @@ -88,7 +99,17 @@ struct MXFP8GroupQuantConfig { .append("_") .append(colwise ? "1" : "0") .append("_") - .append(shape_rep_to_str(shape_rep)); + .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; } @@ -100,15 +121,16 @@ struct MXFP8GroupQuantConfig { 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))); + 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, 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 = 4; +// 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); @@ -151,25 +173,29 @@ inline NVTEBasicTensor make_basic_tensor(void *dptr, DType dtype, nvte_make_shape(shape.data(), shape.size())}; } -// Signature mirrors mxfp8::group_quantize (input, output, stream) for the subset the -// CuTeDSL kernel covers. Returns false to fall back to the CUDA kernel. +// 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, - GroupedTensor *output_tensor, cudaStream_t stream) { + 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 own chunk - // height and MXFP8 block size, which it tiles independently of the CUDA kernel's - // CastTraits. + // 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 kScaleDimX = 32; - if (first_logical_dim % kChunkDimY != 0 || last_logical_dim % kScaleDimX != 0) { + 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, - ", ", kScaleDimX, ")."); + ", ", kLastDimAlignment, ")."); return false; } // The same extents are sym_int32 in the compiled kernel. @@ -179,6 +205,15 @@ inline bool mxfp8_group_quantize_cutedsl(const MXFP8GroupQuantConfig &config, 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; @@ -188,7 +223,7 @@ inline bool mxfp8_group_quantize_cutedsl(const MXFP8GroupQuantConfig &config, GroupDescriptorWorkspace *const workspace = group_descriptor_workspace_ptr(); // Both output directions are handed to the kernel unconditionally: the compiled - // signature has no optional tensors, and building a TMA descriptor needs a real + // signature has no optional outputs, and building a TMA descriptor needs a real // address for each. The disabled direction is never read or written, so it points at // the enabled one instead of at a buffer that would have to be allocated. const SimpleTensor &data_row = @@ -233,46 +268,62 @@ inline bool mxfp8_group_quantize_cutedsl(const MXFP8GroupQuantConfig &config, {num_tensors, kGroupTensorMapSlots, kInt64PerTensorMap}), false, device_index); - // stream is a tvm-ffi opaque "handle"; pass it as void*. + // 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, static_cast(stream)); + &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 Tensor *noop_tensor, - GroupedTensor *output_tensor, const bool use_2d_quantization, +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 bool use_2d_quantization, 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; } - // The CuTeDSL grouped kernel is cast-only: no dbias, no fused (derivative) activation. - if constexpr (IS_DBIAS || IS_DACT || IS_ACT || OP != nullptr) { - maybe_warn_cutedsl_not_chosen( - "grouped quantization with dbias or a fused activation is not supported."); + // TODO(kainingz): port 2D quantization to CuTeDSL + if (use_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 { - // TODO(kainingz): port 2D quantization to CuTeDSL - if (use_2d_quantization) { - maybe_warn_cutedsl_not_chosen("2D quantization is not supported."); - return false; - } - // The kernel takes no noop flag, no amax accumulator, and writes compact scales only. - if (noop_tensor != nullptr && noop_tensor->data.dptr != nullptr) { - maybe_warn_cutedsl_not_chosen("the cast-noop flag is not supported."); - return false; - } - if (output_tensor->amax.dptr != nullptr) { - maybe_warn_cutedsl_not_chosen("amax computation is not supported."); - return false; - } - if (output_tensor->with_gemm_swizzled_scales) { - maybe_warn_cutedsl_not_chosen("GEMM-swizzled scales are not supported."); - return false; - } - // Mirrors the shape-representation selection in mxfp8::group_quantize. ShapeRepresentation shape_rep = ShapeRepresentation::SAME_BOTH_DIMS; if (output_tensor->all_same_shape()) { @@ -284,11 +335,9 @@ bool mxfp8_group_quantize_cutedsl(const GroupedTensor *input_tensor, const Tenso } else if (output_tensor->varying_both_dims()) { shape_rep = ShapeRepresentation::VARYING_BOTH_DIMS; } - if (shape_rep == ShapeRepresentation::VARYING_BOTH_DIMS) { - // The logical shape is [1, total], which is not tileable. - maybe_warn_cutedsl_not_chosen("groups with both dimensions varying are not supported."); - return false; - } + 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. @@ -312,6 +361,11 @@ bool mxfp8_group_quantize_cutedsl(const GroupedTensor *input_tensor, const Tenso 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(); @@ -319,9 +373,18 @@ bool mxfp8_group_quantize_cutedsl(const GroupedTensor *input_tensor, const Tenso // mxfp8::group_quantize raises a proper error for this. return false; } + const bool swizzled = output_tensor->with_gemm_swizzled_scales; + if (swizzled && colwise && !is_single_tensor) { + // For these representations the CUDA kernel adds the tensor base to the colwise + // swizzled scale index twice, so leave them to it rather than reproduce that. + maybe_warn_cutedsl_not_chosen( + "GEMM-swizzled colwise scales are only supported for a common last dimension."); + return false; + } - checkCuDriverContext(stream); // 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"); } @@ -333,13 +396,33 @@ bool mxfp8_group_quantize_cutedsl(const GroupedTensor *input_tensor, const Tenso "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}; - return mxfp8_group_quantize_cutedsl(config, input_tensor, output_tensor, stream); + /*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); } } From 04c12ca04dc939cfd8b433544fdde2c34ba75fdc Mon Sep 17 00:00:00 2001 From: Kaining Zhong Date: Fri, 2 Oct 2026 03:16:22 +0000 Subject: [PATCH 05/14] tmp Signed-off-by: Kaining Zhong --- .../cast/mxfp8/group_quantize_mxfp8.py | 713 +++++++++++------- 1 file changed, 430 insertions(+), 283 deletions(-) diff --git a/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py index 26e36db8100..b86476c5824 100644 --- a/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py +++ b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py @@ -12,8 +12,8 @@ global block offsets -- the CUDA `tensor_map_*_static` "direct mapper" path. For SAME_BOTH_DIMS the CUDA grid is linearized per tensor (X, Y-in-tensor, tensor) while this kernel linearizes it flat over the stacked rows; both - require every member's row count to be a multiple of CHUNK_DIM_Y, and under - that precondition the two decode to the identical (block_offset_Y, block_id_X) + require every member's row count to be a multiple of 128, and under + that precondition the two decode to the identical (block_start_row, block_id_X) for every block index. * the other reps launch grid=(workers_per_tensor, num_tensors) and bind tensor_id to blockIdx.y, so a CTA grid-strides only within its own tensor and @@ -38,10 +38,6 @@ compact and GEMM-swizzled scales, rowwise and/or colwise, and all four shape representations. Differences from CUDA: * the grouped amax pointer is accepted and left untouched, as the CUDA kernel does; - * GEMM-swizzled colwise scales are only produced for the single-tensor reps. For - VARYING_LAST_DIM / VARYING_BOTH_DIMS the CUDA kernel adds the tensor base to the - colwise swizzled index twice, which is only in bounds for the first member, so the - C++ bridge falls back to CUDA instead of reproducing it; * dbias partial sums are accumulated in the CUDA kernel's order (a running column sum with colwise output, otherwise per-thread sums reduced across the CTA), and the C++ bridge reduces the workspace with the same grouped_reduce_dbias. @@ -75,7 +71,9 @@ 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 TensorMapManager, TensorMapUpdateMode +from cutlass.utils import HardwareInfo from cuda.bindings.driver import CUstream # pylint: disable=no-name-in-module import tvm_ffi @@ -157,19 +155,13 @@ def __init__( f" {SUPPORTED_SHAPE_REPS}" ) self.SHAPE_REP = shape_rep - # Mirrors `is_single_tensor` in group_quantize_mxfp8.cuh. + # 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_gemm_swizzled_scales and colwise and not self.IS_SINGLE_TENSOR: - # The CUDA kernel offsets the colwise swizzled scales of these representations by - # the tensor base twice, which only lands inside the buffer for the first member. - raise ValueError( - "GEMM-swizzled colwise scales are only supported for single-tensor representations" - ) 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") @@ -213,65 +205,48 @@ def __str__(self): class MXFP8GroupQuantizeKernel: """Grouped MXFP8 quantize mirroring group_quantize_mxfp8_kernel's strategy.""" - # TunableConfig / derived constants from group_quantize_mxfp8.cuh. - CHUNK_DIM_Y = 128 - CHUNK_DIM_X = 128 - THREADS_PER_CHUNK = 128 - STATIC_PERSISTENT_BLOCKS_PER_SM = 24 - ELTS_PER_CHUNK = CHUNK_DIM_Y * CHUNK_DIM_X - THREADS_X = CHUNK_DIM_X // MXFP8_BLOCK_SCALING_SIZE # 4 - THREADS_Y = THREADS_PER_CHUNK // THREADS_X # 32 - BUFF_DIM_Y = THREADS_Y # 32 - BUFF_DIM_X = CHUNK_DIM_X # 128 - # Each block of (CHUNK_DIM_Y, CHUNK_DIM_X) consists of STAGES tiles of (BUFF_DIM_X, BUFF_DIM_Y) stacked vertically - STAGES = CHUNK_DIM_Y // BUFF_DIM_Y # 4 - PIPELINE_DEPTH = 2 # PREFETCH_STAGES(1) + 1 - NUM_WARPS = THREADS_PER_CHUNK // THREADS_PER_WARP # 4 - - # Rowwise vectorization constants (mirror MXFP8QuantizeKernel / CUDA PACK_SIZE). + # Target persistent CTA count per SM for sizing the grid + STATIC_PERSISTENT_WORK_UNITS_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): self.cfg = cfg self.SM_COUNT = SM_COUNT - # CastConfig widens the chunk to 128x256, i.e. STAGES_X = 2 tiles of - # BUFF_DIM_X columns, each traversed in STAGES row stages before moving right. - self.STAGES_X = 2 if cfg.SHAPE_REP == VARYING_BOTH_DIMS else 1 - self.CHUNK_WIDTH = self.CHUNK_DIM_X * self.STAGES_X + # 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) - # Like IS_CACHED_ACT_OP in CUDA: with both directions, the colwise pass caches the - # activation in the input tile for the rowwise pass. ReLU is fused into the conversion - # instead, which yields the same bytes. + # 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" ) - # CUDA reduces dbias in the colwise pass when there is one, else in the rowwise pass. + # 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 - # ---------------------------------------------------------------- helpers - @cute.jit - def _tensor_rows_cols( - self, tensor_id, mFirstDims, mLastDims, first_logical_dim, last_logical_dim - ): - """Get the shape (rows, cols) of the tensor by tensor_id.""" - cfg = self.cfg - 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) - return rows, cols - @cute.jit def _find_tensor_from_offsets(self, mOffsets, num_tensors, offset: Int64): """Index of the tensor whose element range holds `offset` (find_tensor_from_offsets).""" @@ -314,7 +289,7 @@ def _rowwise_scales(self, mS_row, base: Int64, rows, cols): mS_row, base, cute.make_layout((rows, stride), stride=(stride, 1)) ) return cute.zipped_divide( - mS_t, (self.BUFF_DIM_Y, self.BUFF_DIM_X // MXFP8_BLOCK_SCALING_SIZE) + mS_t, (self.TILE_ROWS, self.TILE_COLS // MXFP8_BLOCK_SCALING_SIZE) ) @cute.jit @@ -333,21 +308,24 @@ def _colwise_scales(self, mS_col, base: Int64, rows, cols): cute.make_layout((rows // MXFP8_BLOCK_SCALING_SIZE, stride), stride=(stride, 1)), ) return cute.zipped_divide( - mS_t, (self.BUFF_DIM_Y // MXFP8_BLOCK_SCALING_SIZE, self.BUFF_DIM_X) + mS_t, (self.TILE_ROWS // MXFP8_BLOCK_SCALING_SIZE, self.TILE_COLS) ) - # ------------------------------------------------------------ entry point @cute.jit def __call__( self, mX: cute.Tensor, - mO_row: cute.Tensor, - mO_col: 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: cute.Tensor, # int64[num_tensors] (VARYING_FIRST_DIM / VARYING_BOTH_DIMS) - mLastDims: cute.Tensor, # int64[num_tensors] (VARYING_LAST_DIM / VARYING_BOTH_DIMS) + 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 @@ -360,18 +338,19 @@ def __call__( cfg = self.cfg first_logical_dim = mX.shape[0] last_logical_dim = mX.shape[1] - # Number of group members. Do NOT derive this from mOffsets: a caller is free to - # pass a length-num_tensors stub for SAME_BOTH_DIMS (where the offsets array is - # unused), which would make `mOffsets.shape[0] - 1` read one too few and divide by - # zero at num_tensors == 1. The per-tensor descriptor workspace is num_tensors long - # by construction -- one slot set per member -- so it is the reliable source. + 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") - smem_tile_layout = cute.make_ordered_layout( - (self.BUFF_DIM_Y, self.BUFF_DIM_X), order=(1, 0) - ) - cta_tiler = (self.BUFF_DIM_Y, self.BUFF_DIM_X) + # 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 @@ -382,57 +361,72 @@ def __call__( 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, 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, tma_dst_out_col = cpasync.make_tiled_tma_atom( - op_store, mO_col, smem_tile_layout, cta_tiler, num_multicast=1 - ) + 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 blocks does the grouped tensor have in both directions - work_blocks_X = cute.ceil_div(Int32(last_logical_dim), self.CHUNK_DIM_X) - work_blocks_Y = cute.ceil_div(Int32(first_logical_dim), self.CHUNK_DIM_Y) - # Each CTA handles one chunk. With every member's rows a multiple of CHUNK_DIM_Y, - # this linear order is the CUDA kernel's (X, Y-in-tensor, tensor) order. + # How many CTAs does the grouped tensor have in both directions + work_blocks_Y = cute.ceil_div( + Int32(first_logical_dim), (self.TILE_ROWS * self.NUM_TILES_Y) + ) + work_blocks_X = cute.ceil_div( + Int32(last_logical_dim), self.TILE_COLS * self.NUM_TILES_X + ) + # Flatten it to an 1D grid grid = [work_blocks_X * work_blocks_Y, 1, 1] else: - # The work-block grid is per-tensor here: each CTA derives its own block range - # from its tensor's extents, so work_blocks_X/Y are unused on this path. - work_blocks_Y = Int32(1) - work_blocks_X = Int32(1) - # Persistent worker count, mirroring get_launch_config() in - # group_quantize_mxfp8.cuh: SM_COUNT * STATIC_PERSISTENT_BLOCKS_PER_SM workers - # split evenly across tensors, clamped to the average number of chunks a tensor - # holds. The element count would wrap Int32, so the estimate - # DIVUP(elts_total, CHUNK_DIM_Y * TILE_DIM_X) is formed without it. - n_tensors = cutlass.max(Int32(num_tensors), Int32(1)) # never divide by zero + # A placeholder for the kernel signature only; we won't use it in non-single tensor cases + work_blocks_X = None + # Estimate how many CTA work unit we need to do, where one work unit is NUM_TILES_X * NUM_TILES_Y tiles if cutlass.const_expr(cfg.SHAPE_REP == VARYING_BOTH_DIMS): - # The logical shape is [1, total]. - estimated_work_blocks = cute.ceil_div(Int32(last_logical_dim), self.ELTS_PER_CHUNK) - else: - # Exact: the first extent is compiled with divisibility CHUNK_DIM_Y. - estimated_work_blocks = cute.ceil_div( - (Int32(first_logical_dim) // self.CHUNK_DIM_Y) * Int32(last_logical_dim), - self.ELTS_PER_CHUNK // self.CHUNK_DIM_Y, + # Note: when VARYING_BOTH_DIMS, the first_logical_dim must be 1 + estimated_work_units = cute.ceil_div( + Int32(first_logical_dim) * Int32(last_logical_dim), self.ELTS_PER_CTA ) - estimated_work_blocks = cute.ceil_div(estimated_work_blocks, self.STAGES_X) - requested_workers_per_tensor = cutlass.max( + 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_work_units = 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}") + + # All SMs can do SM_COUNT * STATIC_PERSISTENT_WORK_UNITS_PER_SM units of work, which are distributed evenly + # to all tensors in the group. + # If there are more tensors than work units, we will launch one CTA per tensor since we have enough parallelism here + requested_CTAs_per_tensor = cutlass.max( Int32(1), - Int32(self.SM_COUNT * self.STATIC_PERSISTENT_BLOCKS_PER_SM) // n_tensors, - ) - average_work_blocks_per_tensor = cutlass.max( - Int32(1), cute.ceil_div(estimated_work_blocks, n_tensors) + Int32(self.SM_COUNT * self.STATIC_PERSISTENT_WORK_UNITS_PER_SM) + // Int32(num_tensors), ) - workers_per_tensor = cutlass.min( - requested_workers_per_tensor, average_work_blocks_per_tensor + # In average how many work units per tensor (only an average, the actual work units per tensor may vary) + average_work_units_per_tensor = cutlass.max( + Int32(1), cute.ceil_div(estimated_work_units, Int32(num_tensors)) ) - grid = [workers_per_tensor, Int32(num_tensors), 1] + # Don't launch more CTAs than the average work units per tensor in case + # STATIC_PERSISTENT_WORK_UNITS_PER_SM causes redundancy + CTAs_per_tensor = cutlass.min(requested_CTAs_per_tensor, average_work_units_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( + self._update_descriptors_kernel( mX, mO_row, mO_col, @@ -473,13 +467,12 @@ def __call__( tma_dst_out_col, ).launch( grid=grid, - block=[self.THREADS_PER_CHUNK, 1, 1], + block=[self.THREADS_PER_CTA, 1, 1], stream=stream, ) - # ------------------------------------------------- descriptor prologue @cute.kernel - def update_descriptors_kernel( + def _update_descriptors_kernel( self, mX, mO_row, @@ -497,20 +490,33 @@ def update_descriptors_kernel( tma_atom_ocol, tma_atom_act, ): - """One CTA per tensor: point that tensor's TMA descriptors at its own block. - - CuTeDSL analog of common::update_tma_descriptors writing g_tensor_maps[]. + """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() - tidx, _, _ = cute.arch.thread_idx() - rows, cols = self._tensor_rows_cols( - tensor_id, mFirstDims, mLastDims, first_logical_dim, last_logical_dim - ) - base_elts = Int64(mOffsets[tensor_id]) + # 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( @@ -525,13 +531,13 @@ def update_descriptors_kernel( tensor_id, ) - # Publish this tensor's geometry for the main kernel (written even when empty). meta = mTensormaps[(tensor_id, META_SLOT, None)] meta[0] = Int64(rows) meta[1] = Int64(cols) - meta[2] = base_elts + 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) @@ -546,7 +552,7 @@ def member_view(tensor, elt_dtype): return cute.make_tensor( cute.make_ptr( elt_dtype, - tensor.iterator.toint() + base_elts * (elt_dtype.width // 8), + tensor.iterator.toint() + base_offset * (elt_dtype.width // 8), cute.AddressSpace.gmem, assumed_align=16, ), @@ -560,19 +566,20 @@ def member_view(tensor, elt_dtype): if cutlass.const_expr(cfg.ROWWISE): views.append(member_view(mO_row, cfg.FP8_DTYPE)) atoms.append(tma_atom_orow) - descs.append(desc_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) - descs.append(desc_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) - descs.append(desc_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), @@ -581,7 +588,125 @@ def member_view(tensor, elt_dtype): (), # smem staging is unused in GMEM update mode ) - # ------------------------------------------------------------ main kernel + @cute.jit + def _make_shared_storage( + self, + smem: cutlass.Constexpr, + dtype: cutlass.Constexpr[Type[cutlass.Numeric]], + ): + """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, @@ -645,7 +770,7 @@ def _kernel_main( first_logical_dim, last_logical_dim, num_tensors, - work_blocks_X, + work_units_X, dtype: cutlass.Constexpr[Type[cutlass.Numeric]], tma_atom_x, tma_src, @@ -675,72 +800,18 @@ def _kernel_main( tidx, ) - # --- shared memory (allocated once, reused across jobs) --- - @cute.struct - class SharedStorage: - mbar: cute.struct.MemRange[cute.Int64, 2 * self.PIPELINE_DEPTH] - sX: cute.struct.Align[ - cute.struct.MemRange[ - dtype, self.BUFF_DIM_Y * self.BUFF_DIM_X * self.PIPELINE_DEPTH - ], - 128, - ] - sO_row: cute.struct.Align[ - cute.struct.MemRange[ - FP8_DTYPE, self.BUFF_DIM_Y * self.BUFF_DIM_X * self.PIPELINE_DEPTH - ], - 128, - ] - sO_col: cute.struct.Align[ - cute.struct.MemRange[ - FP8_DTYPE, self.BUFF_DIM_Y * self.BUFF_DIM_X * self.PIPELINE_DEPTH - ], - 128, - ] - smem = cutlass.utils.SmemAllocator() - storage = smem.allocate(SharedStorage) - tile_layout = cute.make_layout( - ((self.BUFF_DIM_Y, self.BUFF_DIM_X), self.PIPELINE_DEPTH), - stride=((self.BUFF_DIM_X, 1), self.BUFF_DIM_Y * self.BUFF_DIM_X), - ) - 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) - - sActInput = None - if cutlass.const_expr(cfg.WITH_DACT): - - @cute.struct - class DactStorage: - sActInput: cute.struct.Align[ - cute.struct.MemRange[ - dtype, self.BUFF_DIM_Y * self.BUFF_DIM_X * 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.BUFF_DIM_X), stride=(DBIAS_BUFF_WIDTH, 1)) - ) + # 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) - # Grad and activation input share each stage's barrier. - tx_count = self.BUFF_DIM_Y * self.BUFF_DIM_X * dtype.width // 8 + 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=storage.mbar.data_ptr(), + 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), @@ -754,25 +825,32 @@ class DbiasStorage: pipeline.PipelineUserType.Consumer, self.PIPELINE_DEPTH ) - # TMA partitions built from the representative views. For multi-tensor reps - # the descriptor is swapped per tensor and the tile coords are tensor-local. - gX_tiled = cute.zipped_divide(tma_src, (self.BUFF_DIM_Y, self.BUFF_DIM_X)) + 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.BUFF_DIM_Y, self.BUFF_DIM_X)) + 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 ) - gO_row_tiled = cute.zipped_divide(tma_dst_out_row, (self.BUFF_DIM_Y, self.BUFF_DIM_X)) - tXsO_row, tXgO_row = cpasync.tma_partition( - tma_atom_out_row, 0, cute.make_layout(1), sO_row, gO_row_tiled - ) - gO_col_tiled = cute.zipped_divide(tma_dst_out_col, (self.BUFF_DIM_Y, self.BUFF_DIM_X)) - tXsO_col, tXgO_col = cpasync.tma_partition( - tma_atom_out_col, 0, cute.make_layout(1), sO_col, gO_col_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() @@ -786,29 +864,30 @@ class DbiasStorage: # group can exceed 2^31 elements even when every individual extent is small. tensor_base = Int64(0) # Block's offset and id in this individual tensor / global single tensor - block_offset_Y = Int32(0) + block_start_row = Int32(0) block_id_X = Int32(0) first_block_id = Int32(0) - blocks_in_tensor = Int32(1) + units_in_tensor = Int32(1) block_stride = Int32(1) - block_columns_in_tensor = Int32(1) + units_X_in_tensor = Int32(1) if cutlass.const_expr(cfg.IS_SINGLE_TENSOR): - # grid = [work_blocks_X * work_blocks_Y, 1, 1] - block_id_Y = Int32(bidx) // work_blocks_X - block_id_X = Int32(bidx) % work_blocks_X + # grid = [work_units_X * work_units_Y, 1, 1] + block_id_Y = Int32(bidx) // work_units_X + block_id_X = Int32(bidx) % work_units_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 block start from - block_offset_Y = block_id_Y * self.CHUNK_DIM_Y + block_start_row = block_id_Y * (self.TILE_ROWS * self.NUM_TILES_Y) if cutlass.const_expr(cfg.SHAPE_REP == VARYING_FIRST_DIM): - # logical_shape may describe graph-safe capacity beyond the active tensors, - # whose total element count is the last CSR offset. total_elts = Int64(mOffsets[mOffsets.shape[0] - 1]) - if Int64(block_offset_Y) * Int64(last_logical_dim) >= total_elts: + # If the starting element is already beyond the last element of the group, this CTA can stop + if Int64(block_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 block's starting row being beyond the last row of the tensor. else: # grid = [workers_per_tensor, Int32(num_tensors), 1] tensor_id = Int32(bidy) @@ -818,17 +897,17 @@ class DbiasStorage: tensor_cols = Int32(meta[1]) tensor_base = Int64(meta[2]) if tensor_rows > 0 and tensor_cols > 0: - # How many blocks does this tensor have in both directions - block_columns_in_tensor = cute.ceil_div(tensor_cols, self.CHUNK_WIDTH) - block_rows_in_tensor = cute.ceil_div(tensor_rows, self.CHUNK_DIM_Y) - # How many blocks does this tensor have - blocks_in_tensor = block_columns_in_tensor * block_rows_in_tensor - # Which block (1D index) does this CTA start from + # How many work units does this tensor have in both directions + units_X_in_tensor = cute.ceil_div(tensor_cols, (self.TILE_COLS * self.NUM_TILES_X)) + units_Y_in_tensor = cute.ceil_div(tensor_rows, (self.TILE_ROWS * self.NUM_TILES_Y)) + # How many work units does this tensor have + units_in_tensor = units_X_in_tensor * units_Y_in_tensor + # Which work unit (1D index) does this CTA start from, which is also their worker ID first_block_id = Int32(bidx) - # gdx is workers_per_tensor (how many CTAs are assigned to this tensor) + # gdx is how many CTAs are assigned to this tensor block_stride = Int32(gdx) - # If my first block_id is already beyond the tensor's last block, I have no work to do - if first_block_id >= blocks_in_tensor: + # If my first work unit is already beyond the tensor's last work unit, I have no work to do + if first_block_id >= units_in_tensor: has_work = Boolean(False) else: # This tensor is empty, so this CTA has no work to do @@ -838,6 +917,7 @@ class DbiasStorage: 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 work unit if cutlass.const_expr(cfg.IS_SINGLE_TENSOR): # For single tensor case we don't use tensor descriptors descs = (None, None, None, None) @@ -847,8 +927,9 @@ class DbiasStorage: row_scales = None if cutlass.const_expr(cfg.ROWWISE): row_scales = self._rowwise_scales(mS_row, Int64(0), tensor_rows, tensor_cols) + col_scales = None - col_scale_row0 = block_offset_Y + col_scale_row_start = block_start_row col_scale_rows = tensor_rows if cutlass.const_expr(cfg.COLWISE): col_scale_base = Int64(0) @@ -860,12 +941,12 @@ class DbiasStorage: member_row0 = Int32(0) if cutlass.const_expr(cfg.SHAPE_REP == SAME_BOTH_DIMS): member_rows = tensor_rows // Int32(num_tensors) - member_row0 = block_offset_Y // member_rows * member_rows + member_row0 = block_start_row // member_rows * member_rows else: member_id = self._find_tensor_from_offsets( mOffsets, num_tensors, - Int64(block_offset_Y) * Int64(tensor_cols), + Int64(block_start_row) * Int64(tensor_cols), ) member_rows = Int32(mFirstDims[member_id]) member_row0 = Int32(Int64(mOffsets[member_id]) // Int64(tensor_cols)) @@ -874,7 +955,7 @@ class DbiasStorage: * Int64(cute.round_up(tensor_cols, 128)) // MXFP8_BLOCK_SCALING_SIZE ) - col_scale_row0 = block_offset_Y - member_row0 + col_scale_row_start = block_start_row - member_row0 col_scale_rows = member_rows col_scales = self._colwise_scales( mS_col, col_scale_base, col_scale_rows, tensor_cols @@ -883,15 +964,15 @@ class DbiasStorage: cute.arch.sync_threads() self._process_block( - block_offset_Y, + block_start_row, block_id_X, tensor_rows, tensor_cols, row_scales, col_scales, - col_scale_row0, + col_scale_row_start, col_scale_rows, - block_offset_Y // self.CHUNK_DIM_Y, + block_start_row // (self.TILE_ROWS * self.NUM_TILES_Y), mWorkspace, sDbias, descs, @@ -909,19 +990,14 @@ class DbiasStorage: cons_state, ) else: - # For multi-tensor case, retrieve tensor descriptors we processed early in the prologue kernel + # For non-single tensor case, we use persistent kernel so each CTA keeps processing work units 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 ) - # Acquire the descriptors on ONE thread, as the CUDA kernel does - # (`leading_thread` in group_quantize_mxfp8.cuh); the sync_threads below - # publishes it CTA-wide. Running the tensormap acquire fence on all 128 - # threads is correct but very expensive -- it more than doubles the - # kernel time on the multi-tensor path (4096x14336 bidirectional: - # 133 us -> 58 us), since the cost scales with threads x descriptors. + if tidx == 0: tmap.fence_tensormap_update(desc_x) if cutlass.const_expr(cfg.WITH_DACT): @@ -941,24 +1017,27 @@ class DbiasStorage: 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 blocks. cute.arch.sync_threads() # Grid-stride over this tensor's own chunks; the descriptors never change. block_id = first_block_id job_finished = Boolean(False) while not job_finished: - block_id_Y_in_tensor = block_id // block_columns_in_tensor - block_id_X_in_tensor = block_id % block_columns_in_tensor - block_offset_Y_in_tensor = block_id_Y_in_tensor * self.CHUNK_DIM_Y - if cutlass.const_expr(self.STAGES_X == 1): + block_id_Y_in_tensor = block_id // units_X_in_tensor + block_id_X_in_tensor = block_id % units_X_in_tensor + block_start_row_in_tensor = block_id_Y_in_tensor * ( + self.TILE_ROWS * self.NUM_TILES_Y + ) + if cutlass.const_expr(self.NUM_TILES_X == 1): self._process_block( - block_offset_Y_in_tensor, + block_start_row_in_tensor, block_id_X_in_tensor, tensor_rows, tensor_cols, row_scales, col_scales, - block_offset_Y_in_tensor, + block_start_row_in_tensor, tensor_rows, Int32(0), # dbias is only supported for single-tensor reps mWorkspace, @@ -980,20 +1059,20 @@ class DbiasStorage: else: # The chunk's column tiles in order, stopping at the tensor's last # column like the CUDA kernel's stages_X = DIVUP(chunk_cols, TILE_DIM_X). - chunk_col0 = block_id_X_in_tensor * self.CHUNK_WIDTH + chunk_col0 = block_id_X_in_tensor * (self.TILE_COLS * self.NUM_TILES_X) tiles_X = cutlass.min( - Int32(self.STAGES_X), - cute.ceil_div(tensor_cols - chunk_col0, self.BUFF_DIM_X), + Int32(self.NUM_TILES_X), + cute.ceil_div(tensor_cols - chunk_col0, self.TILE_COLS), ) for stage_X in cutlass.range(tiles_X, unroll=1): self._process_block( - block_offset_Y_in_tensor, - block_id_X_in_tensor * self.STAGES_X + stage_X, + block_start_row_in_tensor, + block_id_X_in_tensor * self.NUM_TILES_X + stage_X, tensor_rows, tensor_cols, row_scales, col_scales, - block_offset_Y_in_tensor, + block_start_row_in_tensor, tensor_rows, Int32(0), # dbias is only supported for single-tensor reps mWorkspace, @@ -1014,7 +1093,7 @@ class DbiasStorage: ) # Find the next block to process block_id = block_id + block_stride - if block_id >= blocks_in_tensor: + if block_id >= units_in_tensor: job_finished = Boolean(True) # Drain every TMA store before the CTA releases its shared-memory source buffers. @@ -1071,13 +1150,13 @@ def _issue_load( @cute.jit def _process_block( self, - block_offset_Y, # Row offset of this chunk (global for single-tensor, else tensor-local) + block_start_row, # Row offset of this chunk (global for single-tensor, else tensor-local) block_id_X, # Column-chunk index within the tensor rows, # Rows of the rowwise-scale view (the group for single-tensor, else the tensor) cols, # Number of columns in this tensor - row_scales, # Rowwise scales tiled per stage, rows counted like block_offset_Y + row_scales, # Rowwise scales tiled per stage, rows counted like block_start_row col_scales, # Colwise scales tiled per stage - col_scale_row0, # Row of this chunk in the colwise-scale view + col_scale_row_start, # Row of this chunk in the colwise-scale view col_scale_rows, # Rows of the colwise-scale view dbias_row, # Row of the dbias workspace this chunk reduces into mWorkspace, # f32 partial dbias workspace (WITH_DBIAS) @@ -1096,17 +1175,17 @@ def _process_block( prod_state, cons_state, ): - """Quantize one 128x128 tile of a chunk in STAGES slices of BUFF_DIM_Y rows.""" + """Quantize NUM_TILES_Y vertically stacked tiles in one column strip of a work item.""" 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 - block_offset_X = block_id_X * self.CHUNK_DIM_X + block_offset_X = block_id_X * self.TILE_COLS # This chunk's coordinates in the tile grid (32x128 TMA boxes, not elements). - tile_id_Y = block_offset_Y // self.BUFF_DIM_Y + tile_id_Y = block_start_row // self.TILE_ROWS tile_id_X = block_id_X - col_scale_tile_Y = col_scale_row0 // self.BUFF_DIM_Y + col_scale_tile_Y = col_scale_row_start // self.TILE_ROWS # Per-chunk dbias accumulators, in the CUDA kernel's summation order: a running # column sum over the chunk's rows (colwise), or per-thread partial sums over its @@ -1136,7 +1215,7 @@ def _process_block( ) prod_state.advance() - for stage in cutlass.range_constexpr(self.STAGES): + for stage in cutlass.range_constexpr(self.NUM_TILES_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) @@ -1158,7 +1237,7 @@ def _process_block( sO_col[(None, cons_state.index)], cute.flatten(col_scales[(None, (col_scale_tile_Y + stage, tile_id_X))]), cfg.MAX_NORM_RCP, - (col_scale_tile_Y + stage) * self.BUFF_DIM_Y, + (col_scale_tile_Y + stage) * self.TILE_ROWS, block_offset_X, col_scale_rows, cols, @@ -1166,8 +1245,8 @@ def _process_block( DTYPE=cfg.DTYPE, FP8_DTYPE=cfg.FP8_DTYPE, SWIZZLE=cfg.WITH_GEMM_SWIZZLED_SCALES, - TILE_X=self.BUFF_DIM_X, - TILE_Y=self.BUFF_DIM_Y, + 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, @@ -1185,15 +1264,15 @@ def _process_block( sO_row[(None, cons_state.index)], cute.flatten(row_scales[(None, (row_tile, tile_id_X))]), cfg.MAX_NORM_RCP, - row_tile * self.BUFF_DIM_Y, + row_tile * self.TILE_ROWS, block_offset_X, rows, cols, ACTIVATION=None if self.CACHE_ACTIVATION else cfg.ACTIVATION, DTYPE=cfg.DTYPE, FP8_DTYPE=cfg.FP8_DTYPE, - TILE_X=self.BUFF_DIM_X, - TILE_Y=self.BUFF_DIM_Y, + 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, @@ -1213,7 +1292,7 @@ def _process_block( # 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 cutlass.const_expr(stage + self.PIPELINE_DEPTH < self.STAGES): + if cutlass.const_expr(stage + self.PIPELINE_DEPTH < self.NUM_TILES_Y): if warp_idx == 0: self._issue_load( mainloop_pipeline, @@ -1307,23 +1386,29 @@ def compile_cutedsl_function_from_cfg(cfg: MXFP8GroupQuantizeConfig): out_dtype = cfg.FP8_DTYPE scale_dtype = cutlass.Float8E8M0FNU - def g2d(dtype, shape=logical_shape, align=16): - return cute.runtime.make_fake_compact_tensor( - dtype, - shape, + out_col_fake = ( + cute.runtime.make_fake_compact_tensor( + out_dtype, + logical_shape, stride_order=(1, 0), memspace=cute.AddressSpace.gmem, - assumed_align=align, + assumed_align=16, ) + if cfg.COLWISE + else None + ) - def g1d(dtype, align=4): - return cute.runtime.make_fake_compact_tensor( - dtype, - (cute.sym_int32(),), - stride_order=(0,), + out_row_fake = ( + cute.runtime.make_fake_compact_tensor( + out_dtype, + logical_shape, + stride_order=(1, 0), memspace=cute.AddressSpace.gmem, - assumed_align=align, + 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. @@ -1337,27 +1422,89 @@ def g1d(dtype, align=4): # 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 = g2d(cfg.DTYPE) if cfg.WITH_DACT else None + 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 = ( - g2d(Float32, shape=(cute.sym_int32(), cute.sym_int32()), align=4) + 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 ) - from cutlass.utils import HardwareInfo # pylint: disable=import-outside-toplevel + 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, - g2d(cfg.DTYPE), # mX - g2d(out_dtype), # mO_row - g2d(out_dtype), # mO_col - g1d(scale_dtype), # mS_row - g1d(scale_dtype), # mS_col - g1d(cutlass.Int64, align=8), # mOffsets - g1d(cutlass.Int64, align=8), # mFirstDims - g1d(cutlass.Int64, align=8), # mLastDims + 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 From 674c41a32c10abcaaaa6e71c34fc1e433989ce42 Mon Sep 17 00:00:00 2001 From: Kaining Zhong Date: Fri, 2 Oct 2026 18:10:49 +0000 Subject: [PATCH 06/14] add tests Signed-off-by: Kaining Zhong --- tests/cpp/CMakeLists.txt | 1 + tests/cpp/operator/test_cast_mxfp8_grouped.cu | 59 ++++++++++++++++ tests/jax/test_mxfp8_cutedsl_backend.py | 67 +++++++++++++++++++ .../mxfp8/test_mxfp8_cutedsl_backend.py | 20 +++--- 4 files changed, 138 insertions(+), 9 deletions(-) diff --git a/tests/cpp/CMakeLists.txt b/tests/cpp/CMakeLists.txt index af074be1c78..9a1e0bae18f 100644 --- a/tests/cpp/CMakeLists.txt +++ b/tests/cpp/CMakeLists.txt @@ -57,6 +57,7 @@ if(NVTE_WITH_CUTEDSL) # Find the interpreter too so older CMake versions use its version to locate libpython. find_package(Python COMPONENTS Interpreter Development.Embed REQUIRED) foreach(test_target test_operator test_util) + target_compile_definitions(${test_target} PRIVATE NVTE_WITH_CUTEDSL) target_link_libraries(${test_target} PRIVATE "-Wl,--no-as-needed" Python::Python diff --git a/tests/cpp/operator/test_cast_mxfp8_grouped.cu b/tests/cpp/operator/test_cast_mxfp8_grouped.cu index c2309f7aa4b..13c74491a11 100644 --- a/tests/cpp/operator/test_cast_mxfp8_grouped.cu +++ b/tests/cpp/operator/test_cast_mxfp8_grouped.cu @@ -9,6 +9,11 @@ #include #include +#ifdef NVTE_WITH_CUTEDSL +#include +#include +#endif + #include #include #include "../test_common.h" @@ -1030,3 +1035,57 @@ INSTANTIATE_TEST_SUITE_P( ::testing::Values(DType::kBFloat16), ::testing::Values(DType::kFloat8E4M3)), MakeGroupedFusedCastMXFP8TestName); + +// Exercise the grouped C API with NVTE_ENABLE_CUTEDSL_BACKEND=1 as well as CUDA. +// These cases cover both column strips, a final half-width chunk, and every fused +// activation on all four shape representations against the independent CPU reference. +INSTANTIATE_TEST_SUITE_P( + OperatorTest_GroupedFusedCastMXFP8_MultiChunkActivations, + GroupedFusedCastMXFP8TestSuite, + ::testing::Combine( + ::testing::Values(ProcessingMethod::CAST_ACT, ProcessingMethod::CAST_DACT, + ProcessingMethod::CAST_DBIAS_DACT), + ::testing::Values(ActivationKind::GeLU, ActivationKind::SiLU, ActivationKind::ReLU, + ActivationKind::QGeLU, ActivationKind::SReLU), + ::testing::ValuesIn(scaling_directions), + ::testing::ValuesIn(input_config_multichunk), + ::testing::Values(DType::kBFloat16), + ::testing::Values(DType::kFloat8E4M3)), + MakeGroupedFusedCastMXFP8TestName); + +#ifdef NVTE_WITH_CUTEDSL +TEST(OperatorTest_GroupedFusedCastMXFP8, TestCuTeDSLRegistration) { + const char* enabled = std::getenv("NVTE_ENABLE_CUTEDSL_BACKEND"); + if (enabled == nullptr || enabled[0] == '0' || + getDeviceComputeCapability() < blackwellComputeCapability) { + GTEST_SKIP() << "Requires Blackwell and NVTE_ENABLE_CUTEDSL_BACKEND=1"; + } + + // Cover both column strips and a final half-width chunk. Verify registration + // after the C API call so a silent CUDA fallback cannot satisfy this test. + const std::vector first_dims = {128, 256}; + const std::vector last_dims = {128, 384}; + const std::vector offsets = {0, 128 * 128, 128 * 128 + 256 * 384}; + performTest(CAST_ONLY, &identity, VARYING_BOTH_DIMS, 2, + {1, offsets.back()}, first_dims, last_dims, offsets, + /*rowwise=*/true, /*colwise=*/true); + + ASSERT_TRUE(Py_IsInitialized()) << "CuTeDSL did not initialize embedded Python"; + const std::string key = + "cutedsl_group_mxfp8_sm" + std::to_string(getDeviceComputeCapability()) + + "_BFloat16_Float8E4M3_1_1_varying_both_dims_0_0_0_0_none"; + const PyGILState_STATE gil = PyGILState_Ensure(); + PyObject* module = PyImport_ImportModule("tvm_ffi"); + PyObject* kernel = module == nullptr ? nullptr : + PyObject_CallMethod(module, "get_global_func", "s", key.c_str()); + const bool registered = kernel != nullptr && kernel != Py_None; + if (PyErr_Occurred() != nullptr) { + PyErr_Print(); + } + Py_XDECREF(kernel); + Py_XDECREF(module); + PyGILState_Release(gil); + EXPECT_TRUE(registered) << "CuTeDSL kernel not registered for " << key + << "; the grouped C API fell back to CUDA"; +} +#endif 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 c3ab0bfa8c2..0847421383e 100644 --- a/tests/pytorch/mxfp8/test_mxfp8_cutedsl_backend.py +++ b/tests/pytorch/mxfp8/test_mxfp8_cutedsl_backend.py @@ -381,14 +381,19 @@ def test_dtypes(swizzled, method, act, fp8_dtype, in_dtype): # (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)]), - # N ends in a partial 32-element scale block. - ("varying_first_n144", "varying_first_dim", [(128, 144), (384, 144)]), + # 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") @@ -494,7 +499,9 @@ def run_group_test_case( cutedsl_output[name], cuda_bytes ), f"{tag}: {name} differ between backends" if dbias: - assert torch.equal(dbias_cutedsl, dbias_cuda), f"{tag}: dbias differs between backends" + 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) @@ -510,11 +517,6 @@ def test_group_cast_only(fp8_dtype, in_dtype, block_size, case): @pytest.mark.parametrize("block_size", BLOCK_SIZES, ids=get_block_id) def test_group_swizzled(block_size, case): _, shape_rep, members = case - if block_size[0] != 1 and shape_rep in ("varying_last_dim", "varying_both_dims"): - pytest.skip( - "The CUDA kernel's GEMM-swizzled colwise scale index double-counts the tensor base" - " for varying last dims; the CuTeDSL backend leaves these configs to it." - ) run_group_test_case( members, shape_rep, block_size, torch.bfloat16, tex.DType.kFloat8E4M3, swizzled=True ) @@ -522,7 +524,7 @@ def test_group_swizzled(block_size, case): @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", [torch.bfloat16, torch.float32], ids=get_dtype_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 From 6f8d5c513a1873bfaa72e255cad67ac33ee0e4af Mon Sep 17 00:00:00 2001 From: Kaining Zhong Date: Fri, 2 Oct 2026 18:14:29 +0000 Subject: [PATCH 07/14] fix Signed-off-by: Kaining Zhong --- .../cast/mxfp8/group_quantize_mxfp8.py | 216 +++++++++--------- .../cast/mxfp8/group_quantize_mxfp8.cuh | 3 +- .../mxfp8/group_quantize_mxfp8_cutedsl.cuh | 52 ++--- 3 files changed, 134 insertions(+), 137 deletions(-) diff --git a/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py index b86476c5824..d4c61f5730e 100644 --- a/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py +++ b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py @@ -8,16 +8,16 @@ management and per-tensor scale addressing mirror the CUDA kernel one-for-one: * `is_single_tensor` reps (SAME_BOTH_DIMS, VARYING_FIRST_DIM) launch ONE CTA per - 128x128 chunk and address the group through ONE static TMA descriptor with - global block offsets -- the CUDA `tensor_map_*_static` "direct mapper" path. + 128x128 job and address the group through ONE static TMA descriptor with + global job offsets -- the CUDA `tensor_map_*_static` "direct mapper" path. For SAME_BOTH_DIMS the CUDA grid is linearized per tensor (X, Y-in-tensor, tensor) while this kernel linearizes it flat over the stacked rows; both require every member's row count to be a multiple of 128, and under - that precondition the two decode to the identical (block_start_row, block_id_X) - for every block index. + that precondition the two decode to the identical (job_start_row, job_id_X) + for every job index. * the other reps launch grid=(workers_per_tensor, num_tensors) and bind tensor_id to blockIdx.y, so a CTA grid-strides only within its own tensor and - never re-resolves which tensor a chunk belongs to. They get per-tensor + never re-resolves which tensor a job belongs to. They get per-tensor descriptors written by a prologue kernel (the CuTeDSL analog of update_tma_descriptors filling g_tensor_maps) and acquired with a tensormap proxy fence. @@ -31,7 +31,7 @@ Mechanics that provably yield the same bytes may differ: the mbarrier pipeline is expressed with PipelineTmaAsync instead of hand-rolled mbarriers. As in CUDA, the -scales of out-of-bounds columns in a chunk (the scale-row padding) are written as 0. +scales of out-of-bounds columns in a job (the scale-row padding) are written as 0. Scope: everything group_quantize_mxfp8.cuh covers except 2D block scaling -- the cast-noop flag, fused activation (IS_ACT) and activation derivative (IS_DACT), dbias, @@ -206,7 +206,7 @@ class MXFP8GroupQuantizeKernel: """Grouped MXFP8 quantize mirroring group_quantize_mxfp8_kernel's strategy.""" # Target persistent CTA count per SM for sizing the grid - STATIC_PERSISTENT_WORK_UNITS_PER_SM = 24 + STATIC_PERSISTENT_WORKERS_PER_SM = 24 # The shape of one pipeline stage processed by a CTA TILE_ROWS = 32 TILE_COLS = 128 @@ -379,49 +379,43 @@ def __call__( if cutlass.const_expr(cfg.IS_SINGLE_TENSOR): # How many CTAs does the grouped tensor have in both directions - work_blocks_Y = cute.ceil_div( - Int32(first_logical_dim), (self.TILE_ROWS * self.NUM_TILES_Y) - ) - work_blocks_X = cute.ceil_div( - Int32(last_logical_dim), self.TILE_COLS * self.NUM_TILES_X - ) + 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 = [work_blocks_X * work_blocks_Y, 1, 1] + 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 - work_blocks_X = None - # Estimate how many CTA work unit we need to do, where one work unit is NUM_TILES_X * NUM_TILES_Y tiles + 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_work_units = cute.ceil_div( + 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_work_units = cute.ceil_div( + 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}") - # All SMs can do SM_COUNT * STATIC_PERSISTENT_WORK_UNITS_PER_SM units of work, which are distributed evenly - # to all tensors in the group. - # If there are more tensors than work units, we will launch one CTA per tensor since we have enough parallelism here + # 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_WORK_UNITS_PER_SM) - // Int32(num_tensors), + Int32(self.SM_COUNT * self.STATIC_PERSISTENT_WORKERS_PER_SM) // Int32(num_tensors), ) - # In average how many work units per tensor (only an average, the actual work units per tensor may vary) - average_work_units_per_tensor = cutlass.max( - Int32(1), cute.ceil_div(estimated_work_units, 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 work units per tensor in case - # STATIC_PERSISTENT_WORK_UNITS_PER_SM causes redundancy - CTAs_per_tensor = cutlass.min(requested_CTAs_per_tensor, average_work_units_per_tensor) + # 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. @@ -455,7 +449,7 @@ def __call__( first_logical_dim, last_logical_dim, num_tensors, - work_blocks_X, + jobs_X, mX.element_type, tma_atom_x, tma_src, @@ -720,7 +714,7 @@ def kernel( first_logical_dim, last_logical_dim, num_tensors, - work_blocks_X, + jobs_X, dtype: cutlass.Constexpr[Type[cutlass.Numeric]], tma_atom_x, tma_src, @@ -746,7 +740,7 @@ def kernel( first_logical_dim, last_logical_dim, num_tensors, - work_blocks_X, + jobs_X, dtype, tma_atom_x, tma_src, @@ -770,7 +764,7 @@ def _kernel_main( first_logical_dim, last_logical_dim, num_tensors, - work_units_X, + jobs_X, dtype: cutlass.Constexpr[Type[cutlass.Numeric]], tma_atom_x, tma_src, @@ -857,37 +851,37 @@ def _kernel_main( # If the CTA has work to do has_work = Boolean(True) - # Metadata of the tensor that owns this block + # 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) - # Block's offset and id in this individual tensor / global single tensor - block_start_row = Int32(0) - block_id_X = Int32(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_block_id = Int32(0) - units_in_tensor = Int32(1) - block_stride = Int32(1) - units_X_in_tensor = Int32(1) + 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 = [work_units_X * work_units_Y, 1, 1] - block_id_Y = Int32(bidx) // work_units_X - block_id_X = Int32(bidx) % work_units_X + # 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 block start from - block_start_row = block_id_Y * (self.TILE_ROWS * self.NUM_TILES_Y) + # 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(block_start_row) * Int64(last_logical_dim) >= total_elts: + 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 block's starting row being beyond the last row of the tensor. + # 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) @@ -897,17 +891,17 @@ def _kernel_main( tensor_cols = Int32(meta[1]) tensor_base = Int64(meta[2]) if tensor_rows > 0 and tensor_cols > 0: - # How many work units does this tensor have in both directions - units_X_in_tensor = cute.ceil_div(tensor_cols, (self.TILE_COLS * self.NUM_TILES_X)) - units_Y_in_tensor = cute.ceil_div(tensor_rows, (self.TILE_ROWS * self.NUM_TILES_Y)) - # How many work units does this tensor have - units_in_tensor = units_X_in_tensor * units_Y_in_tensor - # Which work unit (1D index) does this CTA start from, which is also their worker ID - first_block_id = Int32(bidx) + # 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 - block_stride = Int32(gdx) - # If my first work unit is already beyond the tensor's last work unit, I have no work to do - if first_block_id >= units_in_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 @@ -917,7 +911,7 @@ def _kernel_main( 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 work unit + # 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) @@ -929,33 +923,35 @@ def _kernel_main( row_scales = self._rowwise_scales(mS_row, Int64(0), tensor_rows, tensor_cols) col_scales = None - col_scale_row_start = block_start_row + 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 chunk. + # owns this job. member_rows = Int32(0) - member_row0 = 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_row0 = block_start_row // member_rows * member_rows + member_row_start = job_start_row // member_rows * member_rows else: member_id = self._find_tensor_from_offsets( mOffsets, num_tensors, - Int64(block_start_row) * Int64(tensor_cols), + Int64(job_start_row) * Int64(tensor_cols), ) member_rows = Int32(mFirstDims[member_id]) - member_row0 = Int32(Int64(mOffsets[member_id]) // Int64(tensor_cols)) + member_row_start = Int32( + Int64(mOffsets[member_id]) // Int64(tensor_cols) + ) col_scale_base = ( - Int64(member_row0) + Int64(member_row_start) * Int64(cute.round_up(tensor_cols, 128)) // MXFP8_BLOCK_SCALING_SIZE ) - col_scale_row_start = block_start_row - member_row0 + 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 @@ -963,16 +959,16 @@ def _kernel_main( cute.arch.sync_threads() - self._process_block( - block_start_row, - block_id_X, + self._process_job_strip( + job_start_row, + job_id_X, tensor_rows, tensor_cols, row_scales, col_scales, col_scale_row_start, col_scale_rows, - block_start_row // (self.TILE_ROWS * self.NUM_TILES_Y), + job_start_row // (self.TILE_ROWS * self.NUM_TILES_Y), mWorkspace, sDbias, descs, @@ -990,7 +986,7 @@ def _kernel_main( cons_state, ) else: - # For non-single tensor case, we use persistent kernel so each CTA keeps processing work units until none is left + # 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) @@ -1017,27 +1013,27 @@ def _kernel_main( 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 blocks. + # 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 chunks; the descriptors never change. - block_id = first_block_id + # 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: - block_id_Y_in_tensor = block_id // units_X_in_tensor - block_id_X_in_tensor = block_id % units_X_in_tensor - block_start_row_in_tensor = block_id_Y_in_tensor * ( + 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 ) if cutlass.const_expr(self.NUM_TILES_X == 1): - self._process_block( - block_start_row_in_tensor, - block_id_X_in_tensor, + self._process_job_strip( + job_start_row_in_tensor, + job_id_X_in_tensor, tensor_rows, tensor_cols, row_scales, col_scales, - block_start_row_in_tensor, + job_start_row_in_tensor, tensor_rows, Int32(0), # dbias is only supported for single-tensor reps mWorkspace, @@ -1057,22 +1053,22 @@ def _kernel_main( cons_state, ) else: - # The chunk's column tiles in order, stopping at the tensor's last + # The job's column tiles in order, stopping at the tensor's last # column like the CUDA kernel's stages_X = DIVUP(chunk_cols, TILE_DIM_X). - chunk_col0 = block_id_X_in_tensor * (self.TILE_COLS * self.NUM_TILES_X) + job_start_col = job_id_X_in_tensor * (self.TILE_COLS * self.NUM_TILES_X) tiles_X = cutlass.min( Int32(self.NUM_TILES_X), - cute.ceil_div(tensor_cols - chunk_col0, self.TILE_COLS), + cute.ceil_div(tensor_cols - job_start_col, self.TILE_COLS), ) for stage_X in cutlass.range(tiles_X, unroll=1): - self._process_block( - block_start_row_in_tensor, - block_id_X_in_tensor * self.NUM_TILES_X + stage_X, + self._process_job_strip( + job_start_row_in_tensor, + job_id_X_in_tensor * self.NUM_TILES_X + stage_X, tensor_rows, tensor_cols, row_scales, col_scales, - block_start_row_in_tensor, + job_start_row_in_tensor, tensor_rows, Int32(0), # dbias is only supported for single-tensor reps mWorkspace, @@ -1091,9 +1087,9 @@ def _kernel_main( prod_state, cons_state, ) - # Find the next block to process - block_id = block_id + block_stride - if block_id >= units_in_tensor: + # 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. @@ -1148,17 +1144,17 @@ def _issue_load( pipeline_obj.producer_commit(prod_state) @cute.jit - def _process_block( + def _process_job_strip( self, - block_start_row, # Row offset of this chunk (global for single-tensor, else tensor-local) - block_id_X, # Column-chunk index within the tensor + job_start_row, # Row offset of this job (global for single-tensor, else tensor-local) + column_tile_id, # Column-tile index within the tensor rows, # Rows of the rowwise-scale view (the group for single-tensor, else the tensor) cols, # Number of columns in this tensor - row_scales, # Rowwise scales tiled per stage, rows counted like block_start_row + row_scales, # Rowwise scales tiled per stage, rows counted like job_start_row col_scales, # Colwise scales tiled per stage - col_scale_row_start, # Row of this chunk in the colwise-scale view + col_scale_row_start, # Row of this job in the colwise-scale view col_scale_rows, # Rows of the colwise-scale view - dbias_row, # Row of the dbias workspace this chunk reduces into + dbias_row, # Row of the dbias workspace this job reduces into mWorkspace, # f32 partial dbias workspace (WITH_DBIAS) sDbias, # SMEM buffer for the rowwise dbias reduction (rowwise-only dbias) descs, # Per-tensor descriptors (x, act, out_row, out_col), None if single-tensor @@ -1175,20 +1171,20 @@ def _process_block( prod_state, cons_state, ): - """Quantize NUM_TILES_Y vertically stacked tiles in one column strip of a work item.""" + """Quantize NUM_TILES_Y vertically stacked tiles in one column strip of a job.""" 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 - block_offset_X = block_id_X * self.TILE_COLS + job_start_col = column_tile_id * self.TILE_COLS - # This chunk's coordinates in the tile grid (32x128 TMA boxes, not elements). - tile_id_Y = block_start_row // self.TILE_ROWS - tile_id_X = block_id_X + # This job's coordinates in the tile grid (32x128 TMA boxes, not elements). + tile_id_Y = job_start_row // self.TILE_ROWS + tile_id_X = column_tile_id col_scale_tile_Y = col_scale_row_start // self.TILE_ROWS - # Per-chunk dbias accumulators, in the CUDA kernel's summation order: a running - # column sum over the chunk's rows (colwise), or per-thread partial sums over its + # Per-job dbias accumulators, in the CUDA kernel's summation order: a running + # column sum over the job's rows (colwise), or per-thread partial sums over its # stages that the whole CTA reduces afterwards (rowwise-only). dbias_col = Float32(0.0) dbias_row_acc = None @@ -1238,7 +1234,7 @@ def _process_block( cute.flatten(col_scales[(None, (col_scale_tile_Y + stage, tile_id_X))]), cfg.MAX_NORM_RCP, (col_scale_tile_Y + stage) * self.TILE_ROWS, - block_offset_X, + job_start_col, col_scale_rows, cols, ACTIVATION=cfg.ACTIVATION, @@ -1265,7 +1261,7 @@ def _process_block( cute.flatten(row_scales[(None, (row_tile, tile_id_X))]), cfg.MAX_NORM_RCP, row_tile * self.TILE_ROWS, - block_offset_X, + job_start_col, rows, cols, ACTIVATION=None if self.CACHE_ACTIVATION else cfg.ACTIVATION, @@ -1335,8 +1331,8 @@ def _process_block( if cutlass.const_expr(cfg.WITH_DBIAS): if cutlass.const_expr(self.DBIAS_IN_ROWWISE): dbias_col = self._reduce_rowwise_dbias(sDbias, tidx, dbias_row_acc) - # One partial-dbias row per chunk, as in the CUDA kernel's dbias_workspace. - dbias_x = block_offset_X + tidx + # One partial-dbias row per job, as in the CUDA kernel's dbias_workspace. + dbias_x = job_start_col + tidx if dbias_x < cols: mWorkspace[(dbias_row, dbias_x)] = dbias_col @@ -1365,7 +1361,7 @@ def _reduce_rowwise_dbias(self, sDbias, tidx, dbias_row_acc): dbias = Float32(0.0) for i in cutlass.range_constexpr(self.THREADS_Y): dbias += sDbias[(i, tidx)] - # The buffer is rewritten by the next chunk. + # The buffer is rewritten by the next job. cute.arch.sync_threads() return dbias diff --git a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8.cuh b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8.cuh index bd8c3052f60..c39624889ca 100644 --- a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8.cuh +++ b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8.cuh @@ -882,7 +882,8 @@ __global__ void __launch_bounds__(CastTraits::THREADS_PER_CHUNK) group_quantize_ DIVUP_TO_MULTIPLE(DIVUP(cols, static_cast(SCALE_DIM_X)), scale_alignment_X_rowwise); const size_t scale_stride_colwise = DIVUP_TO_MULTIPLE(cols, scale_alignment_X_colwise); - const size_t tensor_base_for_scales = is_single_tensor ? tensor_start_offset : tensor_base; + // Non-single-tensor scale pointers already include the member offset. + const size_t tensor_base_for_scales = is_single_tensor ? tensor_start_offset : 0; e8m0_t *const scales_rowwise = scales_rowwise_ptr + (is_single_tensor ? 0 : tensor_base / SCALE_DIM_X); e8m0_t *const scales_colwise = diff --git a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh index c1b1aae0edd..71f660f3748 100644 --- a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh +++ b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh @@ -137,10 +137,8 @@ constexpr size_t kMaxGroupTensors = struct alignas(128) GroupDescriptorWorkspace { alignas(128) int64_t tensor_maps[kMaxGroupTensors][kGroupTensorMapSlots][kInt64PerTensorMap]; - // Stand-in for the offsets / first_dims / last_dims arrays a given shape representation - // does not carry: the kernel takes all three unconditionally but only dereferences the - // ones its representation uses, so the contents are never read. Sized num_tensors + 1 - // for the CSR offsets array, the longest of the three. + // 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_dims[kMaxGroupTensors + 1]; }; @@ -222,14 +220,6 @@ inline bool mxfp8_group_quantize_cutedsl(const MXFP8GroupQuantConfig &config, const int32_t device_index = transformer_engine::cuda::current_device(); GroupDescriptorWorkspace *const workspace = group_descriptor_workspace_ptr(); - // Both output directions are handed to the kernel unconditionally: the compiled - // signature has no optional outputs, and building a TMA descriptor needs a real - // address for each. The disabled direction is never read or written, so it points at - // the enabled one instead of at a buffer that would have to be allocated. - const SimpleTensor &data_row = - config.rowwise ? output_tensor->data : output_tensor->columnwise_data; - const SimpleTensor &data_col = - config.colwise ? output_tensor->columnwise_data : output_tensor->data; const SimpleTensor &scale_row = config.rowwise ? output_tensor->scale_inv : output_tensor->columnwise_scale_inv; const SimpleTensor &scale_col = @@ -240,10 +230,17 @@ inline bool mxfp8_group_quantize_cutedsl(const MXFP8GroupQuantConfig &config, DLTensorWrapper mX( make_basic_tensor(input_tensor->data.dptr, input_tensor->dtype(), logical_shape), true, device_index); - DLTensorWrapper mO_row(make_basic_tensor(data_row.dptr, data_row.dtype, logical_shape), true, - device_index); - DLTensorWrapper mO_col(make_basic_tensor(data_col.dptr, data_col.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. @@ -258,8 +255,19 @@ inline bool mxfp8_group_quantize_cutedsl(const MXFP8GroupQuantConfig &config, return DLTensorWrapper(make_basic_tensor(dptr, DType::kInt64, {numel}), false, device_index); }; DLTensorWrapper mOffsets = dims_or_unused(output_tensor->tensor_offsets, num_tensors + 1); - DLTensorWrapper mFirstDims = dims_or_unused(output_tensor->first_dims, num_tensors); - DLTensorWrapper mLastDims = dims_or_unused(output_tensor->last_dims, num_tensors); + DLTensorWrapper mFirstDims, mLastDims; + if (config.shape_rep == ShapeRepresentation::VARYING_FIRST_DIM || + config.shape_rep == ShapeRepresentation::VARYING_BOTH_DIMS) { + 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) { + 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. @@ -374,14 +382,6 @@ bool mxfp8_group_quantize_cutedsl(const GroupedTensor *input_tensor, return false; } const bool swizzled = output_tensor->with_gemm_swizzled_scales; - if (swizzled && colwise && !is_single_tensor) { - // For these representations the CUDA kernel adds the tensor base to the colwise - // swizzled scale index twice, so leave them to it rather than reproduce that. - maybe_warn_cutedsl_not_chosen( - "GEMM-swizzled colwise scales are only supported for a common last dimension."); - return false; - } - // Sanity checks, mirroring mxfp8::group_quantize checkCuDriverContext(stream); CheckNoopTensor(*noop_tensor, "cast_noop"); From c23e068b0ae9c57610b092f1d1c0e755d0f71cb5 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 18:16:10 +0000 Subject: [PATCH 08/14] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .../common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh index 71f660f3748..11d6500ff71 100644 --- a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh +++ b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh @@ -259,14 +259,14 @@ inline bool mxfp8_group_quantize_cutedsl(const MXFP8GroupQuantConfig &config, if (config.shape_rep == ShapeRepresentation::VARYING_FIRST_DIM || config.shape_rep == ShapeRepresentation::VARYING_BOTH_DIMS) { mFirstDims = DLTensorWrapper( - make_basic_tensor(output_tensor->first_dims.dptr, DType::kInt64, {num_tensors}), - false, device_index); + 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) { mLastDims = DLTensorWrapper( - make_basic_tensor(output_tensor->last_dims.dptr, DType::kInt64, {num_tensors}), - false, device_index); + 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 From c498ebbea736b10e8a2793e2bd47053dcc20354d Mon Sep 17 00:00:00 2001 From: Kaining Zhong Date: Fri, 2 Oct 2026 18:24:51 +0000 Subject: [PATCH 09/14] nit Signed-off-by: Kaining Zhong --- .../cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh index 11d6500ff71..fee0dc8a7d9 100644 --- a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh +++ b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh @@ -139,7 +139,7 @@ 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_dims[kMaxGroupTensors + 1]; + int64_t unused_offsets[kMaxGroupTensors + 1]; }; // Like `g_tensor_maps` on the CUDA path, this has internal linkage, so every translation @@ -250,20 +250,24 @@ inline bool mxfp8_group_quantize_cutedsl(const MXFP8GroupQuantConfig &config, false, device_index); // Offsets and member dims are read from the output, as in mxfp8::group_quantize. - auto dims_or_unused = [&](const SimpleTensor &t, size_t numel) { - void *dptr = t.has_data() ? t.dptr : static_cast(workspace->unused_dims); + 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 = dims_or_unused(output_tensor->tensor_offsets, num_tensors + 1); + 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); From 11375169be86181647b65f011d01b89c2fe612dd Mon Sep 17 00:00:00 2001 From: Kaining Zhong Date: Fri, 2 Oct 2026 22:11:12 +0000 Subject: [PATCH 10/14] nit Signed-off-by: Kaining Zhong --- .../cast/mxfp8/group_quantize_mxfp8.py | 419 ++++++++---------- .../CuTeDSL/cast/mxfp8/quantize_mxfp8.py | 86 ++-- 2 files changed, 238 insertions(+), 267 deletions(-) diff --git a/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py index d4c61f5730e..fc65638faf2 100644 --- a/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py +++ b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py @@ -90,6 +90,7 @@ SUPPORTED_DACTIVATIONS, derive_swizzled_scale_layout, noop_flag_is_set, + reduce_rowwise_dbias, quantize_rowwise_mxfp8, quantize_colwise_mxfp8, ) @@ -922,6 +923,7 @@ def _kernel_main( 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 @@ -959,16 +961,15 @@ def _kernel_main( cute.arch.sync_threads() - self._process_job_strip( + self._process_job( job_start_row, - job_id_X, + job_id_X * self.TILE_COLS, tensor_rows, tensor_cols, row_scales, col_scales, col_scale_row_start, col_scale_rows, - job_start_row // (self.TILE_ROWS * self.NUM_TILES_Y), mWorkspace, sDbias, descs, @@ -1025,68 +1026,32 @@ def _kernel_main( job_start_row_in_tensor = job_id_Y_in_tensor * ( self.TILE_ROWS * self.NUM_TILES_Y ) - if cutlass.const_expr(self.NUM_TILES_X == 1): - self._process_job_strip( - job_start_row_in_tensor, - job_id_X_in_tensor, - tensor_rows, - tensor_cols, - row_scales, - col_scales, - job_start_row_in_tensor, - tensor_rows, - Int32(0), # dbias is only supported for single-tensor reps - mWorkspace, - sDbias, - descs, - tmap, - warp_idx, - tidx, - sX, - sActInput, - sO_row, - sO_col, - partitions, - atoms, - mainloop_pipeline, - prod_state, - cons_state, - ) - else: - # The job's column tiles in order, stopping at the tensor's last - # column like the CUDA kernel's stages_X = DIVUP(chunk_cols, TILE_DIM_X). - job_start_col = job_id_X_in_tensor * (self.TILE_COLS * self.NUM_TILES_X) - tiles_X = cutlass.min( - Int32(self.NUM_TILES_X), - cute.ceil_div(tensor_cols - job_start_col, self.TILE_COLS), - ) - for stage_X in cutlass.range(tiles_X, unroll=1): - self._process_job_strip( - job_start_row_in_tensor, - job_id_X_in_tensor * self.NUM_TILES_X + stage_X, - tensor_rows, - tensor_cols, - row_scales, - col_scales, - job_start_row_in_tensor, - tensor_rows, - Int32(0), # dbias is only supported for single-tensor reps - mWorkspace, - sDbias, - descs, - tmap, - warp_idx, - tidx, - sX, - sActInput, - sO_row, - sO_col, - partitions, - atoms, - mainloop_pipeline, - prod_state, - cons_state, - ) + 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: @@ -1144,17 +1109,16 @@ def _issue_load( pipeline_obj.producer_commit(prod_state) @cute.jit - def _process_job_strip( + def _process_job( self, job_start_row, # Row offset of this job (global for single-tensor, else tensor-local) - column_tile_id, # Column-tile index within the tensor + job_start_col, # Column offset of this job within the tensor rows, # Rows of the rowwise-scale view (the group for single-tensor, else the tensor) cols, # Number of columns in this tensor row_scales, # Rowwise scales tiled per stage, rows counted like job_start_row col_scales, # Colwise scales tiled per stage col_scale_row_start, # Row of this job in the colwise-scale view col_scale_rows, # Rows of the colwise-scale view - dbias_row, # Row of the dbias workspace this job reduces into mWorkspace, # f32 partial dbias workspace (WITH_DBIAS) sDbias, # SMEM buffer for the rowwise dbias reduction (rowwise-only dbias) descs, # Per-tensor descriptors (x, act, out_row, out_col), None if single-tensor @@ -1171,30 +1135,19 @@ def _process_job_strip( prod_state, cons_state, ): - """Quantize NUM_TILES_Y vertically stacked tiles in one column strip of a job.""" + """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_start_col = column_tile_id * self.TILE_COLS - # This job's coordinates in the tile grid (32x128 TMA boxes, not elements). - tile_id_Y = job_start_row // self.TILE_ROWS - tile_id_X = column_tile_id + 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 - - # Per-job dbias accumulators, in the CUDA kernel's summation order: a running - # column sum over the job's rows (colwise), or per-thread partial sums over its - # stages that the whole CTA reduces afterwards (rowwise-only). - dbias_col = 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) + 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): @@ -1202,8 +1155,8 @@ def _process_job_strip( self._issue_load( mainloop_pipeline, prod_state, - tile_id_Y + prologue_stage, - tile_id_X, + job_tile_Y + prologue_stage % self.NUM_TILES_Y, + job_tile_X + prologue_stage // self.NUM_TILES_Y, atoms, partitions, tmap, @@ -1211,159 +1164,159 @@ def _process_job_strip( ) prod_state.advance() - for stage in cutlass.range_constexpr(self.NUM_TILES_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 = tile_id_Y + stage - - if cutlass.const_expr(cfg.COLWISE): - _, dbias_col = quantize_colwise_mxfp8( - sX_tile, - sAct_tile, - sO_col[(None, cons_state.index)], - cute.flatten(col_scales[(None, (col_scale_tile_Y + stage, tile_id_X))]), - cfg.MAX_NORM_RCP, - (col_scale_tile_Y + stage) * self.TILE_ROWS, - job_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=dbias_col, + 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, ) - if cutlass.const_expr(self.CACHE_ACTIVATION): - # The rowwise pass reads the activation the colwise pass cached in sX. + 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() - 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, - job_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, - ) + 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 - # 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 cutlass.const_expr(stage + self.PIPELINE_DEPTH < self.NUM_TILES_Y): - if warp_idx == 0: - self._issue_load( - mainloop_pipeline, - prod_state, - tile_id_Y + stage + self.PIPELINE_DEPTH, - tile_id_X, - atoms, - partitions, - tmap, - descs, + 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, ) - prod_state.advance() + if cutlass.const_expr(self.CACHE_ACTIVATION): + # The rowwise pass reads the activation the colwise pass cached in sX. + cute.arch.sync_threads() - # 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() + 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, + ) - if cutlass.const_expr(cfg.WITH_DBIAS): - if cutlass.const_expr(self.DBIAS_IN_ROWWISE): - dbias_col = self._reduce_rowwise_dbias(sDbias, tidx, dbias_row_acc) - # One partial-dbias row per job, as in the CUDA kernel's dbias_workspace. - dbias_x = job_start_col + tidx - if dbias_x < cols: - mWorkspace[(dbias_row, dbias_x)] = dbias_col + # 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() - @cute.jit - def _reduce_rowwise_dbias(self, sDbias, tidx, dbias_row_acc): - """Reduce the per-thread rowwise partial sums to one sum per column, in the order of the - CUDA kernel's partial_dbias_rowwise reduction.""" - _, tv_write = cute.make_layout_tv( - thr_layout=cute.make_layout( - (self.THREADS_Y, self.THREADS_X), stride=(self.THREADS_X, 1) - ), - val_layout=cute.make_layout( - (1, MXFP8_BLOCK_SCALING_SIZE), stride=(MXFP8_BLOCK_SCALING_SIZE, 1) - ), - ) - sDbias_write = cute.composition(sDbias, tv_write) - bank_group = (tidx % THREADS_PER_WARP) // self.THREADS_PER_BANK - offset = bank_group * self.PACK_SIZE - for w in cutlass.range_constexpr(self.WAVES): - # Undo the bank-conflict rotation quantize_rowwise_mxfp8 accumulated in. - start = (w * self.PACK_SIZE + offset) % MXFP8_BLOCK_SCALING_SIZE - for i in cutlass.range_constexpr(self.PACK_SIZE): - sDbias_write[(tidx, start + i)] = dbias_row_acc[w * self.PACK_SIZE + i] - cute.arch.sync_threads() - # Thread tidx sums column tidx over the THREADS_Y partial rows. - dbias = Float32(0.0) - for i in cutlass.range_constexpr(self.THREADS_Y): - dbias += sDbias[(i, tidx)] - # The buffer is rewritten by the next job. - cute.arch.sync_threads() - return dbias + # 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): diff --git a/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py b/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py index 9880276d396..c267638bbd7 100644 --- a/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py +++ b/transformer_engine/common/CuTeDSL/cast/mxfp8/quantize_mxfp8.py @@ -815,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. @@ -1645,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( From 76631afb7c03f3e373ee46deff6ea0d0d503c9d6 Mon Sep 17 00:00:00 2001 From: Kaining Zhong Date: Fri, 2 Oct 2026 23:17:46 +0000 Subject: [PATCH 11/14] type annotation Signed-off-by: Kaining Zhong --- .../cast/mxfp8/group_quantize_mxfp8.py | 271 ++++++++---------- 1 file changed, 117 insertions(+), 154 deletions(-) diff --git a/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py index fc65638faf2..7d90d56da93 100644 --- a/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py +++ b/transformer_engine/common/CuTeDSL/cast/mxfp8/group_quantize_mxfp8.py @@ -2,63 +2,7 @@ # # See LICENSE for license information. -"""Grouped MXFP8 quantization kernel implemented in CuTeDSL. - -Strategy-aligned port of group_quantize_mxfp8.cuh. The scheduling, descriptor -management and per-tensor scale addressing mirror the CUDA kernel one-for-one: - - * `is_single_tensor` reps (SAME_BOTH_DIMS, VARYING_FIRST_DIM) launch ONE CTA per - 128x128 job and address the group through ONE static TMA descriptor with - global job offsets -- the CUDA `tensor_map_*_static` "direct mapper" path. - For SAME_BOTH_DIMS the CUDA grid is linearized per tensor (X, Y-in-tensor, - tensor) while this kernel linearizes it flat over the stacked rows; both - require every member's row count to be a multiple of 128, and under - that precondition the two decode to the identical (job_start_row, job_id_X) - for every job index. - * the other reps launch grid=(workers_per_tensor, num_tensors) and bind - tensor_id to blockIdx.y, so a CTA grid-strides only within its own tensor and - never re-resolves which tensor a job belongs to. They get per-tensor - descriptors written by a prologue kernel (the CuTeDSL analog of - update_tma_descriptors filling g_tensor_maps) and acquired with a tensormap - proxy fence. - * per-tensor scale bases/strides follow the CUDA formulas: - scales_* += is_single_tensor ? 0 : tensor_base / 32 - stride_rowwise = roundup(cols/32, 4) stride_colwise = roundup(cols, 128) - -Both kernels dropped the older flat persistent grid that strided across tensor -boundaries (CUDA's decode_job / advance_to_next_job, removed in #3483); the -per-tensor grid above is what replaced it on both sides. - -Mechanics that provably yield the same bytes may differ: the mbarrier pipeline is -expressed with PipelineTmaAsync instead of hand-rolled mbarriers. As in CUDA, the -scales of out-of-bounds columns in a job (the scale-row padding) are written as 0. - -Scope: everything group_quantize_mxfp8.cuh covers except 2D block scaling -- the -cast-noop flag, fused activation (IS_ACT) and activation derivative (IS_DACT), dbias, -compact and GEMM-swizzled scales, rowwise and/or colwise, and all four shape -representations. Differences from CUDA: - * the grouped amax pointer is accepted and left untouched, as the CUDA kernel does; - * dbias partial sums are accumulated in the CUDA kernel's order (a running column sum - with colwise output, otherwise per-thread sums reduced across the CTA), and the C++ - bridge reduces the workspace with the same grouped_reduce_dbias. - -Like the CUDA kernel, every member's first dim must be a multiple of 128 (and, for the -varying-last reps, its last dim too). The kernel prints the same diagnostics as -get_tensor_rows_num / get_tensor_cols_num when a group violates this and, like -NVTE_DEVICE_ERROR in a release build, carries on. - -Measured, deliberately NOT changed: - - sO_row and sO_col are both allocated unconditionally, where CUDA sizes only the - direction in use. Sizing them conditionally does work -- ncu confirms the shared-memory - occupancy limit goes 6 -> 9 CTAs/SM for a single-direction bf16 config -- but it is a - small LOSS on GB200, not a win: rep_med_sbd bf16 colwise 54.9 -> 56.1 us, rowwise - 56.5 -> 57.0 us (fp32 rowwise gains ~1%). The kernel is DRAM-bandwidth-bound at - ~6.3 TB/s, so extra resident CTAs only add contention. Verified by a control that kept - the conditional code but padded SMEM back to the old size: timings returned exactly to - the unconditional numbers, so the effect is the occupancy, not codegen. Revisit if a - future variant (dbias / activation) makes this kernel latency- rather than - bandwidth-bound. -""" +"""Grouped MXFP8 quantization kernel implemented in CuTeDSL.""" # pylint: disable=missing-class-docstring @@ -72,8 +16,8 @@ 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 TensorMapManager, TensorMapUpdateMode -from cutlass.utils import HardwareInfo +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 @@ -226,7 +170,7 @@ class MXFP8GroupQuantizeKernel: # 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): + 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 @@ -249,7 +193,9 @@ def __init__(self, cfg: MXFP8GroupQuantizeConfig, SM_COUNT: int): self.DBIAS_IN_ROWWISE = cfg.WITH_DBIAS and not cfg.COLWISE @cute.jit - def _find_tensor_from_offsets(self, mOffsets, num_tensors, offset: Int64): + 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) @@ -264,7 +210,7 @@ def _find_tensor_from_offsets(self, mOffsets, num_tensors, offset: Int64): return low - 1 @cute.jit - def _scale_tensor(self, mS, base: Int64, layout): + 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( @@ -277,7 +223,9 @@ def _scale_tensor(self, mS, base: Int64, layout): ) @cute.jit - def _rowwise_scales(self, mS_row, base: Int64, rows, cols): + 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( @@ -294,7 +242,9 @@ def _rowwise_scales(self, mS_row, base: Int64, rows, cols): ) @cute.jit - def _colwise_scales(self, mS_col, base: Int64, rows, cols): + 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( @@ -332,7 +282,7 @@ def __call__( 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") @@ -469,22 +419,22 @@ def __call__( @cute.kernel def _update_descriptors_kernel( self, - mX, - mO_row, - mO_col, - mActInput, - mOffsets, - mFirstDims, - mLastDims, - mTensormaps, - first_logical_dim, - last_logical_dim, + 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, - tma_atom_orow, - tma_atom_ocol, - tma_atom_act, - ): + 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: @@ -543,7 +493,7 @@ def _update_descriptors_kernel( if rows > 0 and cols > 0: member_layout = cute.make_layout((rows, cols), stride=(cols, 1)) - def member_view(tensor, elt_dtype): + def member_view(tensor: cute.Tensor, elt_dtype: Type[cutlass.Numeric]) -> cute.Tensor: return cute.make_tensor( cute.make_ptr( elt_dtype, @@ -586,9 +536,16 @@ def member_view(tensor, elt_dtype): @cute.jit def _make_shared_storage( self, - smem: cutlass.Constexpr, + 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( @@ -705,27 +662,27 @@ class DbiasStorage: @cute.kernel def kernel( self, - mS_row, - mS_col, - mOffsets, - mFirstDims, - mTensormaps, - mNoop, - mWorkspace, - first_logical_dim, - last_logical_dim, - num_tensors, - jobs_X, + 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, - tma_src, - tma_atom_act, - tma_src_act, - tma_atom_out_row, - tma_dst_out_row, - tma_atom_out_col, - tma_dst_out_col, - ): + 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): @@ -756,26 +713,26 @@ def kernel( @cute.jit def _kernel_main( self, - mS_row, - mS_col, - mOffsets, - mFirstDims, - mTensormaps, - mWorkspace, - first_logical_dim, - last_logical_dim, - num_tensors, - jobs_X, + 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, - tma_src, - tma_atom_act, - tma_src_act, - tma_atom_out_row, - tma_dst_out_row, - tma_atom_out_col, - tma_dst_out_col, - ): + 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() @@ -1064,15 +1021,15 @@ def _kernel_main( def _issue_load( self, - pipeline_obj, - prod_state, - tile_y, - tile_x, - atoms, - partitions, - tmap, - descs, - ): + 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 @@ -1111,30 +1068,36 @@ def _issue_load( @cute.jit def _process_job( self, - job_start_row, # Row offset of this job (global for single-tensor, else tensor-local) - job_start_col, # Column offset of this job within the tensor - rows, # Rows of the rowwise-scale view (the group for single-tensor, else the tensor) - cols, # Number of columns in this tensor - row_scales, # Rowwise scales tiled per stage, rows counted like job_start_row - col_scales, # Colwise scales tiled per stage - col_scale_row_start, # Row of this job in the colwise-scale view - col_scale_rows, # Rows of the colwise-scale view - mWorkspace, # f32 partial dbias workspace (WITH_DBIAS) - sDbias, # SMEM buffer for the rowwise dbias reduction (rowwise-only dbias) - descs, # Per-tensor descriptors (x, act, out_row, out_col), None if single-tensor - tmap, # TensorMapManager for managing TMA descriptors - warp_idx, - tidx, - sX, # SMEM input ring - sActInput, # SMEM activation input ring (WITH_DACT) - sO_row, # SMEM rowwise output ring - sO_col, # SMEM colwise output ring - partitions, # TMA partitions (x, act, out_row, out_col) - atoms, # TMA atoms (x, act, out_row, out_col) + 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, - cons_state, - ): + 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 From 25ba81f4140d9187f799d64489fcbb5a22fa9022 Mon Sep 17 00:00:00 2001 From: Kaining Zhong Date: Fri, 2 Oct 2026 23:21:33 +0000 Subject: [PATCH 12/14] Remove unrelated C++ test additions from CuTeDSL changes Signed-off-by: Kaining Zhong --- tests/cpp/CMakeLists.txt | 1 - tests/cpp/operator/test_cast_mxfp8_grouped.cu | 59 ------------------- 2 files changed, 60 deletions(-) diff --git a/tests/cpp/CMakeLists.txt b/tests/cpp/CMakeLists.txt index 9a1e0bae18f..af074be1c78 100644 --- a/tests/cpp/CMakeLists.txt +++ b/tests/cpp/CMakeLists.txt @@ -57,7 +57,6 @@ if(NVTE_WITH_CUTEDSL) # Find the interpreter too so older CMake versions use its version to locate libpython. find_package(Python COMPONENTS Interpreter Development.Embed REQUIRED) foreach(test_target test_operator test_util) - target_compile_definitions(${test_target} PRIVATE NVTE_WITH_CUTEDSL) target_link_libraries(${test_target} PRIVATE "-Wl,--no-as-needed" Python::Python diff --git a/tests/cpp/operator/test_cast_mxfp8_grouped.cu b/tests/cpp/operator/test_cast_mxfp8_grouped.cu index 13c74491a11..c2309f7aa4b 100644 --- a/tests/cpp/operator/test_cast_mxfp8_grouped.cu +++ b/tests/cpp/operator/test_cast_mxfp8_grouped.cu @@ -9,11 +9,6 @@ #include #include -#ifdef NVTE_WITH_CUTEDSL -#include -#include -#endif - #include #include #include "../test_common.h" @@ -1035,57 +1030,3 @@ INSTANTIATE_TEST_SUITE_P( ::testing::Values(DType::kBFloat16), ::testing::Values(DType::kFloat8E4M3)), MakeGroupedFusedCastMXFP8TestName); - -// Exercise the grouped C API with NVTE_ENABLE_CUTEDSL_BACKEND=1 as well as CUDA. -// These cases cover both column strips, a final half-width chunk, and every fused -// activation on all four shape representations against the independent CPU reference. -INSTANTIATE_TEST_SUITE_P( - OperatorTest_GroupedFusedCastMXFP8_MultiChunkActivations, - GroupedFusedCastMXFP8TestSuite, - ::testing::Combine( - ::testing::Values(ProcessingMethod::CAST_ACT, ProcessingMethod::CAST_DACT, - ProcessingMethod::CAST_DBIAS_DACT), - ::testing::Values(ActivationKind::GeLU, ActivationKind::SiLU, ActivationKind::ReLU, - ActivationKind::QGeLU, ActivationKind::SReLU), - ::testing::ValuesIn(scaling_directions), - ::testing::ValuesIn(input_config_multichunk), - ::testing::Values(DType::kBFloat16), - ::testing::Values(DType::kFloat8E4M3)), - MakeGroupedFusedCastMXFP8TestName); - -#ifdef NVTE_WITH_CUTEDSL -TEST(OperatorTest_GroupedFusedCastMXFP8, TestCuTeDSLRegistration) { - const char* enabled = std::getenv("NVTE_ENABLE_CUTEDSL_BACKEND"); - if (enabled == nullptr || enabled[0] == '0' || - getDeviceComputeCapability() < blackwellComputeCapability) { - GTEST_SKIP() << "Requires Blackwell and NVTE_ENABLE_CUTEDSL_BACKEND=1"; - } - - // Cover both column strips and a final half-width chunk. Verify registration - // after the C API call so a silent CUDA fallback cannot satisfy this test. - const std::vector first_dims = {128, 256}; - const std::vector last_dims = {128, 384}; - const std::vector offsets = {0, 128 * 128, 128 * 128 + 256 * 384}; - performTest(CAST_ONLY, &identity, VARYING_BOTH_DIMS, 2, - {1, offsets.back()}, first_dims, last_dims, offsets, - /*rowwise=*/true, /*colwise=*/true); - - ASSERT_TRUE(Py_IsInitialized()) << "CuTeDSL did not initialize embedded Python"; - const std::string key = - "cutedsl_group_mxfp8_sm" + std::to_string(getDeviceComputeCapability()) + - "_BFloat16_Float8E4M3_1_1_varying_both_dims_0_0_0_0_none"; - const PyGILState_STATE gil = PyGILState_Ensure(); - PyObject* module = PyImport_ImportModule("tvm_ffi"); - PyObject* kernel = module == nullptr ? nullptr : - PyObject_CallMethod(module, "get_global_func", "s", key.c_str()); - const bool registered = kernel != nullptr && kernel != Py_None; - if (PyErr_Occurred() != nullptr) { - PyErr_Print(); - } - Py_XDECREF(kernel); - Py_XDECREF(module); - PyGILState_Release(gil); - EXPECT_TRUE(registered) << "CuTeDSL kernel not registered for " << key - << "; the grouped C API fell back to CUDA"; -} -#endif From 05dfba817dfe1cb5343082ba96ce574eb9cd0b60 Mon Sep 17 00:00:00 2001 From: Kaining Zhong Date: Fri, 2 Oct 2026 23:57:58 +0000 Subject: [PATCH 13/14] nit Signed-off-by: Kaining Zhong --- transformer_engine/common/cast/dispatch/quantize.cuh | 4 ++-- .../common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/transformer_engine/common/cast/dispatch/quantize.cuh b/transformer_engine/common/cast/dispatch/quantize.cuh index 4a6a4d4387c..b4394975dc3 100644 --- a/transformer_engine/common/cast/dispatch/quantize.cuh +++ b/transformer_engine/common/cast/dispatch/quantize.cuh @@ -531,7 +531,7 @@ void group_quantize_fwd_helper(const NVTEGroupedTensor input, NVTEGroupedTensor cutedsl_backend::mxfp8_group_quantize_cutedsl( input_tensor, activations_tensor, noop_tensor, output_tensor, dbias_tensor, - workspace_tensor, quant_config_cpp.mxfp8_2d_quantization, stream); + workspace_tensor, &quant_config_cpp, stream); #endif if (!quantized_with_cutedsl) { mxfp8::group_quantize( @@ -636,7 +636,7 @@ void group_quantize_bwd_helper(const NVTEGroupedTensor grad, const NVTEGroupedTe cutedsl_backend::mxfp8_group_quantize_cutedsl( grad_tensor, input_tensor, noop_tensor, output_tensor, dbias_tensor, workspace_tensor, - quant_config_cpp.mxfp8_2d_quantization, stream); + &quant_config_cpp, stream); #endif if (!quantized_with_cutedsl) { mxfp8::group_quantize( diff --git a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh index fee0dc8a7d9..edb674a09c0 100644 --- a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh +++ b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8_cutedsl.cuh @@ -320,14 +320,14 @@ template mxfp8_2d_quantization) { maybe_warn_cutedsl_not_chosen("2D quantization is not supported."); return false; } From 9b6bb1d366845ad80342bf7d4d4d515ec25bf7bf Mon Sep 17 00:00:00 2001 From: Kaining Zhong Date: Fri, 2 Oct 2026 23:59:49 +0000 Subject: [PATCH 14/14] nit Signed-off-by: Kaining Zhong --- transformer_engine/common/cast/mxfp8/group_quantize_mxfp8.cuh | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8.cuh b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8.cuh index c39624889ca..bd8c3052f60 100644 --- a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8.cuh +++ b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8.cuh @@ -882,8 +882,7 @@ __global__ void __launch_bounds__(CastTraits::THREADS_PER_CHUNK) group_quantize_ DIVUP_TO_MULTIPLE(DIVUP(cols, static_cast(SCALE_DIM_X)), scale_alignment_X_rowwise); const size_t scale_stride_colwise = DIVUP_TO_MULTIPLE(cols, scale_alignment_X_colwise); - // Non-single-tensor scale pointers already include the member offset. - const size_t tensor_base_for_scales = is_single_tensor ? tensor_start_offset : 0; + const size_t tensor_base_for_scales = is_single_tensor ? tensor_start_offset : tensor_base; e8m0_t *const scales_rowwise = scales_rowwise_ptr + (is_single_tensor ? 0 : tensor_base / SCALE_DIM_X); e8m0_t *const scales_colwise =