diff --git a/python/cudnn/api_base.py b/python/cudnn/api_base.py index 2a2a1741a..79c283593 100644 --- a/python/cudnn/api_base.py +++ b/python/cudnn/api_base.py @@ -33,6 +33,16 @@ def is_power_of_2(n: int) -> bool: return n > 0 and (n & (n - 1)) == 0 +def is_sm107_device() -> bool: + """Return True when the current CUDA device is Rubin (SM107).""" + return torch.cuda.is_available() and torch.cuda.get_device_capability(torch.cuda.current_device()) == (10, 7) + + +def get_device_type() -> str: + """Return the architecture family used by SM100 grouped GEMM wrappers.""" + return "rubin" if is_sm107_device() else "blackwell" + + _experimental_api_warnings_emitted = set() _experimental_api_warnings_lock = threading.Lock() @@ -379,12 +389,16 @@ def __init__(self): - self._is_supported: Flag indicating if configuration is validated - self._kernel: Kernel instance - self._compiled_kernel: Cache for compiled kernel + - self._device_type: Architecture family used for dispatch/cache keys + - self._is_rubin_kernel: True when running on Rubin (SM107) - self._logger: Logger instance for this class """ self._is_supported = False self._kernel = None self._compiled_kernel = None self._interpret_uint8_as_fp4x2 = False + self._device_type = get_device_type() + self._is_rubin_kernel = self._device_type == "rubin" self._logger = logging.getLogger(self.__class__.__name__) def _warn_experimental_api(self) -> None: diff --git a/python/cudnn/grouped_gemm/grouped_gemm_dglu/_blockscaled_api.py b/python/cudnn/grouped_gemm/grouped_gemm_dglu/_blockscaled_api.py index f412d4140..ce5c4ac1b 100644 --- a/python/cudnn/grouped_gemm/grouped_gemm_dglu/_blockscaled_api.py +++ b/python/cudnn/grouped_gemm/grouped_gemm_dglu/_blockscaled_api.py @@ -45,6 +45,7 @@ from .moe_blockscaled_grouped_gemm_dglu_dbias import BlockScaledMoEGroupedGemmDgluDbiasKernel from ..moe_utils import MoEWeightMode +from ..grouped_gemm_utils import rubin_single_group_offsets_kwarg from cuda.bindings import driver as cuda import os import torch @@ -59,6 +60,29 @@ from cudnn.api_base import APIBase, ceil_div, is_power_of_2 +def _get_rubin_kernel(): + from .moe_blockscaled_grouped_gemm_dglu_rubin import ( + BlockScaledMoEGroupedGemmDgluKernel as RubinBlockScaledMoEGroupedGemmDgluKernel, + ) + + return RubinBlockScaledMoEGroupedGemmDgluKernel + + +_GEGGLU_ALPHA_DEFAULT = 1.702 +_GLU_CLAMP_MAX_DEFAULT = 7.0 +_GLU_CLAMP_MIN_DEFAULT = -7.0 + + +def _reject_unsupported_rubin_glu_tune_params( + is_rubin_kernel: bool, + geglu_alpha: float, + glu_clamp_max: float, + glu_clamp_min: float, +) -> None: + if is_rubin_kernel and (geglu_alpha != _GEGGLU_ALPHA_DEFAULT or glu_clamp_max != _GLU_CLAMP_MAX_DEFAULT or glu_clamp_min != _GLU_CLAMP_MIN_DEFAULT): + raise NotImplementedError("Rubin grouped GEMM dGLU does not support geglu_alpha, glu_clamp_max, or glu_clamp_min tuning") + + class GroupedGemmDgluBlockScaledAPI(APIBase): """Unified API for grouped GEMM dGLU backward operation on SM100+ GPUs. @@ -279,10 +303,16 @@ def __init__( self.geglu_alpha = geglu_alpha self.glu_clamp_max = glu_clamp_max self.glu_clamp_min = glu_clamp_min + _reject_unsupported_rubin_glu_tune_params( + self._is_rubin_kernel, + self.geglu_alpha, + self.glu_clamp_max, + self.glu_clamp_min, + ) self._interpret_uint8_as_fp4x2 = True self._has_dbias = self.dbias_desc is not None - self._kernel = BlockScaledMoEGroupedGemmDgluDbiasKernel + self._kernel = _get_rubin_kernel() if self._is_rubin_kernel else BlockScaledMoEGroupedGemmDgluDbiasKernel self.num_cluster_overlap_margin = int(os.getenv("CUDNNFE_CLUSTER_OVERLAP_MARGIN", "0")) self._logger.debug(f"setting num_cluster_overlap_margin: {self.num_cluster_overlap_margin}") @@ -605,8 +635,8 @@ def check_support(self) -> bool: f"m_aligned must be divisible by mma_tiler_mn[0], got {self.m_aligned} % {self.mma_tiler_mn[0]} != 0", ) self._value_error_if( - self.m_aligned != BlockScaledMoEGroupedGemmDgluDbiasKernel.FIX_PAD_SIZE, - f"m_aligned must be {BlockScaledMoEGroupedGemmDgluDbiasKernel.FIX_PAD_SIZE} (FIX_PAD_SIZE), got {self.m_aligned}", + self.m_aligned != self._kernel.FIX_PAD_SIZE, + f"m_aligned must be {self._kernel.FIX_PAD_SIZE} (FIX_PAD_SIZE), got {self.m_aligned}", ) # ---- Tensor alignment ---- @@ -688,7 +718,7 @@ def compile(self) -> None: weight_mode=self.weight_mode, act_func=self.act_func, use_dynamic_sched=self.use_dynamic_sched, - use_single_group_runtime_offsets=self.use_single_group_runtime_offsets, + **rubin_single_group_offsets_kwarg(self._is_rubin_kernel, self.use_single_group_runtime_offsets), ) hardware_info = cutlass.utils.HardwareInfo() @@ -903,8 +933,7 @@ def _compile_dense(self, gemm_dglu, max_active_clusters, fake_stream) -> None: # Compile with keyword args (dense mode uses the unified __call__ positional order). dbias_fake = self._make_fake_cute_tensor_from_desc(self.dbias_desc, assumed_align=16) - _compiled_kernel = cute.compile( - gemm_dglu, + compile_kwargs = dict( a=a_cute_fake, b=b_cute_fake, sfb=sfb_cute_fake, @@ -930,12 +959,18 @@ def _compile_dense(self, gemm_dglu, max_active_clusters, fake_stream) -> None: max_active_clusters=max_active_clusters, stream=fake_stream, epilogue_op=self.epilogue_op, - linear_offset=self.linear_offset, - geglu_alpha=self.geglu_alpha, - glu_clamp_max=self.glu_clamp_max, - glu_clamp_min=self.glu_clamp_min, + linear_offset=cutlass.Float32(self.linear_offset) if self._is_rubin_kernel else self.linear_offset, options="--enable-tvm-ffi", ) + if not self._is_rubin_kernel: + compile_kwargs.update( + { + "geglu_alpha": self.geglu_alpha, + "glu_clamp_max": self.glu_clamp_max, + "glu_clamp_min": self.glu_clamp_min, + } + ) + _compiled_kernel = cute.compile(gemm_dglu, **compile_kwargs) # Cache workspace pointer for the tensor_api closure cached_workspace_ptr = from_dlpack(self._workspace, assumed_align=128).iterator @@ -961,7 +996,7 @@ def tensor_api( stream: cuda.CUstream, ) -> None: norm_const_tensor = self._unpad_tensor_to_ndim(norm_const_tensor, 1, "norm_const") - _compiled_kernel( + kernel_args = ( a_tensor, b_tensor, sfb_tensor, @@ -982,9 +1017,11 @@ def tensor_api( beta_tensor, prob_tensor, dprob_tensor, - dbias_tensor, - stream, ) + if self._is_rubin_kernel: + _compiled_kernel(*kernel_args, cutlass.Float32(self.linear_offset), dbias_tensor, stream) + else: + _compiled_kernel(*kernel_args, dbias_tensor, stream) self._compiled_kernel = tensor_api @@ -1093,8 +1130,7 @@ def _compile_discrete(self, gemm_dglu, max_active_clusters, fake_stream) -> None workspace_ptr_cute = from_dlpack(self._workspace, assumed_align=128).iterator self._logger.debug("Compiling discrete grouped GEMM dGLU kernel") - _compiled_kernel = cute.compile( - gemm_dglu, + compile_kwargs = dict( a=a_tensor, b=b_ptrs_cute, sfb=sfb_ptrs_cute, @@ -1120,12 +1156,18 @@ def _compile_discrete(self, gemm_dglu, max_active_clusters, fake_stream) -> None max_active_clusters=max_active_clusters, stream=fake_stream, epilogue_op=self.epilogue_op, - linear_offset=self.linear_offset, - geglu_alpha=self.geglu_alpha, - glu_clamp_max=self.glu_clamp_max, - glu_clamp_min=self.glu_clamp_min, + linear_offset=cutlass.Float32(self.linear_offset) if self._is_rubin_kernel else self.linear_offset, options="--enable-tvm-ffi", ) + if not self._is_rubin_kernel: + compile_kwargs.update( + { + "geglu_alpha": self.geglu_alpha, + "glu_clamp_max": self.glu_clamp_max, + "glu_clamp_min": self.glu_clamp_min, + } + ) + _compiled_kernel = cute.compile(gemm_dglu, **compile_kwargs) self._n = n self._k = k @@ -1161,7 +1203,7 @@ def tensor_api( b_ptrs_addr = int(b_ptrs_device.data_ptr()) sfb_ptrs_addr = int(sfb_ptrs_device.data_ptr()) - _compiled_kernel( + kernel_args = ( a_tensor, b_ptrs_addr, sfb_ptrs_addr, @@ -1182,9 +1224,11 @@ def tensor_api( beta_tensor, prob_tensor, dprob_tensor, - dbias_tensor, - stream, ) + if self._is_rubin_kernel: + _compiled_kernel(*kernel_args, cutlass.Float32(self.linear_offset), dbias_tensor, stream) + else: + _compiled_kernel(*kernel_args, dbias_tensor, stream) self._compiled_kernel = tensor_api diff --git a/python/cudnn/grouped_gemm/grouped_gemm_dglu/api.py b/python/cudnn/grouped_gemm/grouped_gemm_dglu/api.py index 1ddaea944..55fe69d56 100644 --- a/python/cudnn/grouped_gemm/grouped_gemm_dglu/api.py +++ b/python/cudnn/grouped_gemm/grouped_gemm_dglu/api.py @@ -58,7 +58,7 @@ import torch from typing import Any, Tuple, Optional, overload -from cudnn.api_base import APIBase, TupleDict, ceil_div +from cudnn.api_base import APIBase, TupleDict, ceil_div, get_device_type _BLOCK_SCALED_DTYPE_PAIRS = { (dtype, dtype) @@ -72,7 +72,10 @@ from ._bf16_api import GroupedGemmDgluBf16API -from ._blockscaled_api import GroupedGemmDgluBlockScaledAPI +from ._blockscaled_api import ( + GroupedGemmDgluBlockScaledAPI, + _reject_unsupported_rubin_glu_tune_params, +) @dataclass(frozen=True) @@ -521,6 +524,13 @@ def _grouped_gemm_dglu_block_scaled_call(call: DgluCall) -> TupleDict: # default" (1.0 for dgeglu, 0.0 for dswiglu). if linear_offset is None: linear_offset = 1.0 if act_func == "dgeglu" else 0.0 + device_type = get_device_type() + _reject_unsupported_rubin_glu_tune_params( + device_type == "rubin", + geglu_alpha, + glu_clamp_max, + glu_clamp_min, + ) dgeglu_cache_signature = None if act_func == "dgeglu": dgeglu_cache_signature = ( @@ -632,6 +642,7 @@ def dynamic_m_tensor_signature( if is_dense: cache_key = ( + device_type, weight_mode, act_func, dgeglu_cache_signature, @@ -677,6 +688,7 @@ def dynamic_m_tensor_signature( ) else: cache_key = ( + device_type, weight_mode, act_func, dgeglu_cache_signature, diff --git a/python/cudnn/grouped_gemm/grouped_gemm_dglu/moe_blockscaled_grouped_gemm_dglu_dbias.py b/python/cudnn/grouped_gemm/grouped_gemm_dglu/moe_blockscaled_grouped_gemm_dglu_dbias.py index 21e405097..77ac4c8d9 100644 --- a/python/cudnn/grouped_gemm/grouped_gemm_dglu/moe_blockscaled_grouped_gemm_dglu_dbias.py +++ b/python/cudnn/grouped_gemm/grouped_gemm_dglu/moe_blockscaled_grouped_gemm_dglu_dbias.py @@ -588,8 +588,10 @@ def _setup_attributes(self): # Overlap and double buffer accumulator when num_acc_stage == 1 for cta_tile_n = 256 case self.overlapping_accum = self.num_acc_stage == 1 and self.mma_tiler[1] == 256 - # To prefetch more accumulator when overlapping_accum is enabled in epilogue - self.epilogue_prefetch_more = self.d_dtype.width == 8 and self.a_dtype.width == 8 + # The ping-pong prefetch path selects between two rmem tensor objects in + # the epilogue loop. Recent CuTe DSL lowers that to an arith.select over + # memrefs and leaves an unrealized_conversion_cast during LLVM lowering. + self.epilogue_prefetch_more = False # Use ptx fp8 fp32 convert self.use_fp8_ptx_cvt = True diff --git a/python/cudnn/grouped_gemm/grouped_gemm_dglu/moe_blockscaled_grouped_gemm_dglu_rubin.py b/python/cudnn/grouped_gemm/grouped_gemm_dglu/moe_blockscaled_grouped_gemm_dglu_rubin.py new file mode 100644 index 000000000..d92a33b96 --- /dev/null +++ b/python/cudnn/grouped_gemm/grouped_gemm_dglu/moe_blockscaled_grouped_gemm_dglu_rubin.py @@ -0,0 +1,4286 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: BSD-3-Clause + +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: + +# 1. Redistributions of source code must retain the above copyright notice, this +# list of conditions and the following disclaimer. + +# 2. Redistributions in binary form must reproduce the above copyright notice, +# this list of conditions and the following disclaimer in the documentation +# and/or other materials provided with the distribution. + +# 3. Neither the name of the copyright holder nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. + +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +""" +MoE Block-Scaled Grouped GEMM Kernel with dGLU (dSwiGLU/dGeGLU) Backward Fusion. + +Supports: + - Static / Dynamic persistent tile scheduling (MoEPersistentTileScheduler) + - Dense (contiguous 3-D B) / Discrete (per-expert pointer array B) weight layout + - FP8/FP4 output quantization with row/column scale factors (SFD) + - AMAX reduction for FP8 calibration + - dGLU backward activation fusion (dSwiGLU / dGeGLU) + +This module contains only the kernel class. +MoE scheduler components live in moe_persistent_scheduler.py / moe_sched_extension.py / moe_utils.py. +""" + +from typing import Type, Tuple, Union, Optional +from functools import partial + +import cuda.bindings.driver as cuda + +import cutlass +import cutlass.cute as cute +from cutlass.cute.nvgpu import cpasync, tcgen05 +from cutlass.cute.nvgpu.tcgen05 import OperandMajorMode, CollectorOp +from cutlass.cutlass_dsl import T +import cutlass.utils as utils +import cutlass.pipeline as pipeline +import cutlass.utils.blackwell_helpers as sm100_utils +import cutlass.utils.rubin_helpers as sm107_utils +import cutlass.utils.blockscaled_layout as blockscaled_utils +from cutlass.utils.gemm.sm100 import transform_partitioned_tensor_layout + +from cutlass.cute.typing import Float32, Int32, AddressSpace +from cutlass._mlir import ir +from cutlass._mlir.dialects import math, llvm +from cutlass._mlir.dialects import vector, arith + +from ..moe_persistent_scheduler import ( + MoEPersistentTileScheduler, + MoESchedulerParams, + MoEWorkTileInfo, +) +from ..moe_utils import ( + MoEWeightMode, + TensormapWorkspace, + store_tma_desc, +) +from ..moe_sched_extension import ( + DiscreteWeightScaledGemmSchedExtension, + ContiguousAndConsistentGroupedGemmSchedExtension, +) +from ..moe_kernel_helpers import ( + fmin, + fmax, + fmin_bf16x2, + fmax_bf16x2, + atomic_max_float32, + atomic_add_float32, + sigmoid_f32, + get_dtype_rcp_limits, + get_amax_smem_size, + can_implement, + is_valid_dtypes_and_scale_factor_vec_size, + is_valid_layouts, + is_valid_mma_tiler_and_cluster_shape, + is_valid_tensor_alignment, +) + + +def atomic_add_bf16x2(ptr, val_fp32_lo, val_fp32_hi, *, loc=None, ip=None): + """Packed BF16x2 atomic reduction to global memory.""" + lo_ir = val_fp32_lo.ir_value(loc=loc, ip=ip) + hi_ir = val_fp32_hi.ir_value(loc=loc, ip=ip) + llvm.inline_asm( + None, + [ptr, hi_ir, lo_ir], + "{ .reg .b32 packed; cvt.rn.bf16x2.f32 packed, $1, $2; red.global.add.noftz.bf16x2 [$0], packed; }", + "l,f,f", + has_side_effects=True, + is_align_stack=False, + asm_dialect=llvm.AsmDialect.AD_ATT, + ) + + +class BlockScaledMoEGroupedGemmDgluKernel: + """Block-scaled grouped GEMM kernel with MoE tile scheduling and dGLU backward fusion. + + Supports both dense and discrete weight layouts, static and dynamic + scheduling, and quantized output with row/column scale factors. + + This version uses a fixed padding size (FIX_PAD_SIZE=256) that is decoupled from the kernel's tile size, + allowing users to pad their tensors without knowing the specific kernel implementation details. + + :param sf_vec_size: Scalefactor vector size. + :type sf_vec_size: int + :param mma_tiler_mn: Shape of the Matrix Multiply-Accumulate (MMA) tile (M,N) + :type mma_tiler_mn: Tuple[int, int] + :param cluster_shape_mn: Cluster dimensions (M,N) for parallel processing + :type cluster_shape_mn: Tuple[int, int] + + :note: In current version, A and B tensor must have the same data type + - i.e., Float8E4M3FN for A and Float8E5M2 for B is not supported + + :note: Supported combinations of A/B data types, SF data typs and SF vector size: + - MXF8: A/B: Float8E5M2/Float8E4M3FN + SF: Float8E8M0FNU + sf_vec_size: 32 + - MXF4: A/B: Float4E2M1FN + SF: Float8E8M0FNU + sf_vec_size: 32 + - NVF4: A/B: Float4E2M1FN + SF: Float8E8M0FNU/Float8E4M3FN + sf_vec_size: 16 + + :note: Supported accumulator data types: + - Float32 + + :note: Supported D data types: + - BFloat16 + - Float8E4M3FN/Float8E5M2 + + :note: Constraints: + - MMA tiler M must be 128 or 256 (use_2cta_instrs) + - MMA tiler N must be 64/128/192/256 + - Cluster shape M must be multiple of 2 if Mma tiler M is 256 + - Cluster shape M/N must be positive and power of 2, total cluster size <= 16 + - Also, Cluster shape M/N must be <= 4 for scale factor multicasts due to limited size of scale factors + - FIX_PAD_SIZE (256) must be divisible by mma_tiler_mn[0] + - m_aligned parameter in create_mask() MUST equal FIX_PAD_SIZE (256) + - Each padded_offsets[i] will be a multiple of FIX_PAD_SIZE (guaranteed by m_aligned == FIX_PAD_SIZE) + + :note: New Interface (padded_offsets): + Instead of tile_idx_to_expert_idx, num_non_exiting_tiles, and m_split_cumsum, users now provide: + - padded_offsets: shape (expert_cnt,), where padded_offsets[i] is the end position + of expert[i] in the padded A tensor. + - Expert i processes A[padded_offsets[i-1]:padded_offsets[i], :] (with padded_offsets[-1]=0) + + """ + + # Fixed pad size for user-side padding (decoupled from kernel tile size) + FIX_PAD_SIZE = 256 + + @staticmethod + def can_implement( + ab_dtype: Type[cutlass.Numeric], + sf_dtype: Type[cutlass.Numeric], + sf_vec_size: int, + acc_dtype: Type[cutlass.Numeric], + d_dtype: Type[cutlass.Numeric], + use_2cta_instrs: bool, + mma_tiler_mn: Tuple[int, int], + cluster_shape_mn: Tuple[int, int], + m: int, + n: int, + k: int, + l: int, + a_major: str, + b_major: str, + cd_major: str, + m_aligned: int, + act_func: str, + ) -> bool: + FPS = BlockScaledMoEGroupedGemmDgluKernel.FIX_PAD_SIZE + # B-reuse case: 2CTA + mma_tiler_mn[0] = 512 (two 256-M instructions per tile) + if use_2cta_instrs and mma_tiler_mn[0] == 512: + if not is_valid_dtypes_and_scale_factor_vec_size(ab_dtype, sf_dtype, sf_vec_size, acc_dtype, d_dtype): + return False + if not is_valid_layouts(ab_dtype, d_dtype, a_major, b_major, cd_major): + return False + if not is_valid_tensor_alignment(m, n, k, l, ab_dtype, d_dtype, a_major, b_major, cd_major): + return False + # Tile N must be 192 or 256 + if mma_tiler_mn[1] not in {192, 256}: + return False + # Culster M must be a multiple of 2. For cluster M=1, we don't have test it yet. + if cluster_shape_mn[0] % 2 != 0: + return False + if m_aligned % mma_tiler_mn[0] != 0: + return False + if m % mma_tiler_mn[0] != 0: + return False + if act_func not in ["dswiglu", "dgeglu"]: + return False + return True + # Allow N=192 in addition to the shared helper's N=256-only constraint. + if mma_tiler_mn[1] == 192: + if m_aligned != FPS: + return False + if ab_dtype.width == 8: + return False + if not is_valid_dtypes_and_scale_factor_vec_size(ab_dtype, sf_dtype, sf_vec_size, acc_dtype, d_dtype): + return False + if not is_valid_layouts(ab_dtype, d_dtype, a_major, b_major, cd_major): + return False + if not is_valid_tensor_alignment(m, n, k, l, ab_dtype, d_dtype, a_major, b_major, cd_major): + return False + if not (a_major == "k" and b_major == "k"): + return False + if n % 64 != 0 or m % 256 != 0: return False + if not (use_2cta_instrs and mma_tiler_mn[0] == 256): + return False + if cluster_shape_mn[0] % 2 != 0: + return False + return True + result = True + if m_aligned != FPS: + result = False + if not is_valid_dtypes_and_scale_factor_vec_size( + ab_dtype, sf_dtype, sf_vec_size, acc_dtype, d_dtype + ): + result = False + if not is_valid_layouts( + ab_dtype, d_dtype, a_major, b_major, cd_major + ): + result = False + if not is_valid_mma_tiler_and_cluster_shape( + use_2cta_instrs, mma_tiler_mn, cluster_shape_mn, m_aligned, FPS + ): + result = False + if not is_valid_tensor_alignment( + m, n, k, l, ab_dtype, d_dtype, a_major, b_major, cd_major + ): + result = False + if not (a_major == "k"): + result = False + if act_func not in ["dswiglu", "dgeglu"]: + result = False + return result + + def __init__( + self, + sf_vec_size: int, + acc_dtype: Type[cutlass.Numeric], + use_2cta_instrs: bool, + mma_tiler_mn: Tuple[int, int], + cluster_shape_mn: Tuple[int, int], + vectorized_f32: bool, + discrete_col_sfd: bool, + expert_cnt: int, + weight_mode: MoEWeightMode = MoEWeightMode.DISCRETE, + use_dynamic_sched: bool = False, + act_func: str = "dswiglu", + ): + """Initializes the configuration for a Blackwell blockscaled grouped GEMM dGLU kernel. + + This configuration includes several key aspects: + + 1. MMA Instruction Settings (tcgen05): + - acc_dtype: Data types for MMA accumulator. + - mma_tiler_mn: The (M, N) shape of the MMA instruction tiler. + - use_2cta_instrs: Boolean indicating if the tcgen05 MMA variant + with cta_group=2 should be used. + + 2. Cluster Shape: + - cluster_shape_mn: The (ClusterM, ClusterN) shape of the CTA cluster. + + 3. Expert Count: + - expert_cnt: Number of experts for MoE grouped GEMM. + + 4. MoE Tile Scheduling: + - Uses MoEPersistentTileScheduler for tile iteration across experts + - Expert lookup is handled by the scheduler (cached O(1) fast path) + + :param acc_dtype: Data type of the accumulator. + :type acc_dtype: type[cutlass.Numeric] + :param mma_tiler_mn: Tuple (M, N) shape of the MMA instruction. + :type mma_tiler_mn: Tuple[int, int] + :param use_2cta_instrs: Boolean, True to use cta_group=2 MMA variant. + :type use_2cta_instrs: bool + :param cluster_shape_mn: Tuple (ClusterM, ClusterN) shape of the cluster. + :type cluster_shape_mn: Tuple[int, int] + :param expert_cnt: Number of experts (compile-time constant). + :type expert_cnt: int + + :raises ValueError: If FIX_PAD_SIZE is not divisible by mma_tiler_mn[0]. + """ + # Hardware MMA instruction M: 2CTA → 256, 1CTA → 128 + mma_inst_m = 256 if use_2cta_instrs else 128 + enable_breuse = (mma_tiler_mn[0] // mma_inst_m == 2) + if enable_breuse: + if mma_tiler_mn[0] % self.FIX_PAD_SIZE != 0: + raise ValueError( + f"mma_tiler_mn[0] ({mma_tiler_mn[0]}) must be a multiple of " + f"FIX_PAD_SIZE ({self.FIX_PAD_SIZE}) for breuse. " + f"Also ensure callers use m_aligned=mma_tiler_mn[0] (={mma_tiler_mn[0]})." + ) + else: + cta_tile_m = mma_tiler_mn[0] // (2 if use_2cta_instrs else 1) + if self.FIX_PAD_SIZE % cta_tile_m != 0: + raise ValueError( + f"FIX_PAD_SIZE ({self.FIX_PAD_SIZE}) must be divisible by " + f"cta_tile_m ({cta_tile_m}). " + f"Supported mma_tiler_mn[0] values: 128, 256, 512." + ) + if expert_cnt > 1024: + raise ValueError("Expert count > 1024 is not supported.") + if not isinstance(weight_mode, MoEWeightMode): + raise TypeError(f"weight_mode must be a MoEWeightMode, got {type(weight_mode)}") + + self.sf_vec_size = sf_vec_size + self.expert_cnt = expert_cnt + self.acc_dtype: Type[cutlass.Numeric] = acc_dtype + self.use_2cta_instrs = use_2cta_instrs + self.cluster_shape_mn = cluster_shape_mn + # K dimension is deferred in _setup_attributes + self.mma_tiler = (*mma_tiler_mn, 1) + # B-reuse: enabled when the mma_tiler M is 2× the hardware instruction M + self.enable_breuse = enable_breuse + + self.cta_group = ( + tcgen05.CtaGroup.TWO if use_2cta_instrs else tcgen05.CtaGroup.ONE + ) + + self.occupancy = 1 + self.epilog_warp_id = (0, 1, 2, 3) + self.mma_warp_id = 4 + self.tma_warp_id = 5 + self.epilog_load_tma_id = 6 + self.sched_warp_id = 7 + self.threads_per_warp = 32 + self.threads_per_cta = self.threads_per_warp * len( + ( + *self.epilog_warp_id, + self.mma_warp_id, + self.tma_warp_id, + self.epilog_load_tma_id, + self.sched_warp_id, + ) + ) + self.threads_wo_sched = self.threads_per_warp * len( + ( + *self.epilog_warp_id, + self.mma_warp_id, + self.tma_warp_id, + self.epilog_load_tma_id, + ) + ) + + # Set barrier for cta sync, epilogue sync and tmem ptr sync + self.cta_sync_barrier = pipeline.NamedBarrier( + barrier_id=1, + num_threads=self.threads_per_cta, + ) + self.epilog_sync_barrier = pipeline.NamedBarrier( + barrier_id=2, + num_threads=32 * len(self.epilog_warp_id), + ) + self.tmem_alloc_barrier = pipeline.NamedBarrier( + barrier_id=3, + num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id)), + ) + self.sched_sync_barrier = pipeline.NamedBarrier( + barrier_id=4, + num_threads=self.threads_per_warp, + ) + self.num_smem_capacity = utils.get_smem_capacity_in_bytes("sm_107") + self.num_tmem_alloc_cols = cute.arch.get_max_tmem_alloc_cols("sm_107") + + self.vectorized_f32 = vectorized_f32 + self.use_dynamic_sched = use_dynamic_sched + + # Amax reduction configuration + self.num_epilog_warps = len(self.epilog_warp_id) + + self.discrete_col_sfd = discrete_col_sfd + self.weight_mode = weight_mode + + self.act_func = act_func + + def _get_mma_permutation_mnk(self, mma_inst_shape_mnk): + """Return MMA permutation for the Bkeep-Breuse pattern (2CTA only). + Only active when enable_breuse=True to avoid breaking the TMA atom + setup for the non-breuse 2CTA case. + """ + if cutlass.const_expr(self.use_2cta_instrs and self.enable_breuse): + m_layout = cute.make_layout( + shape=(mma_inst_shape_mnk[0] // 2, 2, 2), + stride=(1, mma_inst_shape_mnk[0], mma_inst_shape_mnk[0] // 2), + ) + return (m_layout, mma_inst_shape_mnk[1], mma_inst_shape_mnk[2]) + return (1, 1, 1) + + def _setup_attributes(self): + """Set up configurations that are dependent on GEMM inputs + + This method configures various attributes based on the input tensor properties + (data types, leading dimensions) and kernel settings: + - Configuring tiled MMA + - Computing MMA/cluster/tile shapes + - Computing cluster layout + - Computing multicast CTAs for A/B + - Computing epilogue subtile + - Setting up A/B/D stage counts in shared memory + - Computing A/B/D shared memory layout + - Computing tensor memory allocation columns + """ + + # Hardware MMA instruction M: 2CTA → 256, 1CTA → 128 + mma_inst_m = 256 if self.use_2cta_instrs else 128 + self.mma_inst_shape_mn = (mma_inst_m, self.mma_tiler[1]) + # (CTA_Tile_Shape_M, Round_Up(MMA_Tile_Shape_N, 128), MMA_Inst_Shape_K) + self.mma_inst_shape_mn_sfb = ( + self.mma_inst_shape_mn[0] // (2 if self.use_2cta_instrs else 1), + cute.round_up(self.mma_inst_shape_mn[1], 128), + ) + + # K dim of the MMA instruction shape on Rubin sm107: + # - SM107MmaMXF4NVF4Op (FP4 x FP4) requires K=128 + # - SM107BlockScaledMmaMXF8F6F4Op (FP8 or mixed) requires K=64 + mma_inst_k = ( + 128 + if (self.a_dtype.width == 4 and self.b_dtype.width == 4) + else 64 + ) + mma_inst_shape_mnk = (*self.mma_inst_shape_mn, mma_inst_k) + mma_inst_shape_mnk_sfb = (*self.mma_inst_shape_mn_sfb, mma_inst_k) + + atom_layout_mnk = (1, 1, 1) + permutation_mnk = self._get_mma_permutation_mnk(mma_inst_shape_mnk) + + # Configure tiled mma (Rubin sm107) + tiled_mma = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + self.cta_group, + mma_inst_shape_mnk, + a_collector_op=CollectorOp.DISCARD, + b_collector_op=CollectorOp.DISCARD, + atom_layout_mnk=atom_layout_mnk, + permutation_mnk=permutation_mnk, + ) + + tiled_mma_sfb = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + cute.nvgpu.tcgen05.CtaGroup.ONE, + mma_inst_shape_mnk_sfb, + ) + + # Compute mma/cluster/tile shapes + mma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2]) + mma_inst_tile_k = 2 if (self.a_dtype.width == 4 and self.sf_vec_size == 16) else 4 + self.mma_tiler = ( + self.mma_tiler[0], + self.mma_tiler[1], + mma_inst_shape_k * mma_inst_tile_k, + ) + + self.mma_tiler_sfb = ( + self.mma_inst_shape_mn_sfb[0], + self.mma_inst_shape_mn_sfb[1], + mma_inst_shape_k * mma_inst_tile_k, + ) + + self.cta_tile_shape_mnk = ( + self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape), + self.mma_tiler[1], + self.mma_tiler[2], + ) + self.cta_tile_shape_mnk_sfb = ( + self.mma_tiler_sfb[0] // cute.size(tiled_mma.thr_id.shape), + self.mma_tiler_sfb[1], + self.mma_tiler_sfb[2], + ) + + self.mma_tiler_d = ( + self.mma_tiler[0], + self.mma_inst_shape_mn[1], + mma_inst_shape_k * mma_inst_tile_k, + ) + self.cta_tile_shape_mnk_d = ( + self.mma_tiler_d[0] // cute.size(tiled_mma.thr_id.shape), + self.mma_tiler_d[1], + self.mma_tiler_d[2], + ) + # Compute cluster layout + self.cluster_layout_vmnk = cute.tiled_divide( + cute.make_layout((*self.cluster_shape_mn, 1)), + (tiled_mma.thr_id.shape,), + ) + + self.cluster_layout_sfb_vmnk = cute.tiled_divide( + cute.make_layout((*self.cluster_shape_mn, 1)), + (tiled_mma_sfb.thr_id.shape,), + ) + + # Compute number of multicast CTAs for A/B + self.num_mcast_ctas_a = cute.size(self.cluster_layout_vmnk.shape[2]) + self.num_mcast_ctas_b = cute.size(self.cluster_layout_vmnk.shape[1]) + self.is_a_mcast = self.num_mcast_ctas_a > 1 + self.is_b_mcast = self.num_mcast_ctas_b > 1 + + # Set epilogue subtile + self.epi_tile = (128, 32) + self.epi_tile_cnt = ( + self.cta_tile_shape_mnk_d[0] // self.epi_tile[0], + self.cta_tile_shape_mnk_d[1] // self.epi_tile[1], + ) + + # enable direct store D. Currently, it is disabled, could be enabled in the future. + self.store_d_directly = ( + False + ) + + # Setup A/B/D/Scale stage count in shared memory and ACC stage count in tensor memory + ( + self.num_acc_stage, + self.num_ab_stage, + self.num_c_stage, + self.num_d_stage, + self.num_tile_stage, + ) = self._compute_stages( + tiled_mma, + self.mma_tiler, + self.a_dtype, + self.b_dtype, + self.epi_tile, + self.c_dtype, + self.c_layout, + self.d_dtype, + self.d_layout, + self.sf_dtype, + self.sf_vec_size, + self.num_smem_capacity, + self.occupancy, + self.store_d_directly, + self.generate_dbias, + self.enable_breuse, + ) + + # Compute A/B/D/Scale shared memory layout + self.a_smem_layout_staged = sm100_utils.make_smem_layout_a( + tiled_mma, + self.mma_tiler, + self.a_dtype, + self.num_ab_stage, + ) + self.b_smem_layout_staged = sm100_utils.make_smem_layout_b( + tiled_mma, + self.mma_tiler, + self.b_dtype, + self.num_ab_stage, + ) + self.sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa( + tiled_mma, + self.mma_tiler, + self.sf_vec_size, + self.num_ab_stage, + ) + self.sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb( + tiled_mma, + self.mma_tiler, + self.sf_vec_size, + self.num_ab_stage, + ) + + self.c_smem_layout_staged = sm100_utils.make_smem_layout_epi( + self.c_dtype, + self.c_layout, + self.epi_tile, + self.num_c_stage, + ) + + if cutlass.const_expr(not self.store_d_directly): + self.d_smem_layout_staged = sm100_utils.make_smem_layout_epi( + self.d_dtype, + self.d_layout, + self.epi_tile, + self.num_d_stage, + ) + else: + self.d_smem_layout_staged = sm100_utils.make_smem_layout_epi( + self.d_dtype, + self.d_layout, + self.epi_tile, + 1, + ) + + # For breuse+N=192 or N=256, override num_acc_stage=1 + if self.enable_breuse and self.mma_tiler[1] in {192, 256}: + self.num_acc_stage = 1 + + # Overlap and double buffer accumulator when num_acc_stage == 1. (Reserved for N=192 / NVFP4) + self.overlapping_accum = ( + self.num_acc_stage == 1 and self.mma_tiler[1] == 256 and not self.enable_breuse + ) + + # Use ptx fp8 fp32 convert + self.use_fp8_ptx_cvt = True + + # Compute number of TMEM columns for SFA/SFB/Accumulator. + sf_atom_mn = 32 + sf_pack_factor = 32 // self.sf_vec_size + self.num_sfa_tmem_cols = ( + self.cta_tile_shape_mnk[0] // sf_atom_mn + ) * mma_inst_tile_k * sf_pack_factor + self.num_sfb_tmem_cols = ( + self.cta_tile_shape_mnk_sfb[1] // sf_atom_mn + ) * mma_inst_tile_k * sf_pack_factor + self.num_sf_tmem_cols = self.num_sfa_tmem_cols + self.num_sfb_tmem_cols + if self.enable_breuse: + self.num_accumulator_tmem_cols = ( + self.cta_tile_shape_mnk[1] * self.num_acc_stage * 2 + ) + elif self.overlapping_accum: + self.num_accumulator_tmem_cols = ( + self.cta_tile_shape_mnk[1] * 2 - self.num_sf_tmem_cols + ) + else: + self.num_accumulator_tmem_cols = ( + self.cta_tile_shape_mnk[1] * self.num_acc_stage + ) + if self.cta_tile_shape_mnk[1] == 192 and not self.enable_breuse: + self.num_accumulator_tmem_cols = self.num_tmem_alloc_cols - self.num_sf_tmem_cols + self.num_accumulator_tmem_stride = self.num_accumulator_tmem_cols - 192 + else: + self.num_accumulator_tmem_stride = self.num_accumulator_tmem_cols + + self.epi_tile_n_required = cute.size(self.epi_tile[1]) + # Only when overlapping_accum is enabled, we need to release accumulator buffer early in epilogue + self.iter_acc_early_release_in_epilogue = ( + self.num_sf_tmem_cols + self.epi_tile_n_required - 1 + ) // self.epi_tile_n_required - 1 + + def get_desc_workspace_bytes(self) -> int: + """Return descriptor workspace size in bytes.""" + if self.weight_mode == MoEWeightMode.DISCRETE: + from ..moe_utils import DiscreteWeightTensormapConstructor + return DiscreteWeightTensormapConstructor.get_workspace_size(self.expert_cnt) + return 0 + + def get_workspace_bytes(self) -> int: + """Return descriptor workspace plus optional dynamic scheduler state.""" + desc_workspace_bytes = self.get_desc_workspace_bytes() + dynamic_sched_bytes = 4 if self.use_dynamic_sched else 0 + return desc_workspace_bytes + dynamic_sched_bytes + + @cute.jit + def _get_sched_counter_ptr(self, workspace_ptr): + counter_addr = workspace_ptr.toint() + self.get_desc_workspace_bytes() + return cute.make_ptr( + cutlass.Int32, + counter_addr, + AddressSpace.gmem, + assumed_align=4, + ) + + @cute.kernel + def helper_kernel( + self, + ptrs_b: cute.Pointer, # Device pointer to int64[expert_cnt] array of B addresses + ptrs_sfb: cute.Pointer, # Device pointer to int64[expert_cnt] array of SFB addresses + n: Int32, + k: Int32, + b_stride_size: cutlass.Int64, + b_major_mode: cutlass.Constexpr, + workspace_ptr, + tiled_mma_arg: cute.TiledMma, + tiled_mma_sfb_arg: cute.TiledMma, + b_smem_layout_arg, + sfb_smem_layout_arg, + cluster_layout_vmnk_shape_arg: cutlass.Constexpr, + cluster_layout_sfb_vmnk_shape_arg: cutlass.Constexpr, + ): + """Pre-main-kernel initialization. + + Launched with grid=(expert_cnt, 1, 1) for discrete mode, or + grid=(1, 1, 1) for dense+dynamic mode. + + Discrete weight: each block builds B/SFB TMA descriptors for one expert. + Dynamic sched: block 0 resets the atomic tile counter to 0. + """ + expert_idx = cute.arch.block_idx()[0] + + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE): + b_tma_op_arg = sm100_utils.cluster_shape_to_tma_atom_B( + self.cluster_shape_mn, tiled_mma_arg.thr_id + ) + sfb_tma_op_arg = sm100_utils.cluster_shape_to_tma_atom_SFB( + self.cluster_shape_mn, tiled_mma_arg.thr_id + ) + + b_ptr_tensor = cute.make_tensor( + cute.make_ptr(cutlass.Int64, ptrs_b.toint(), AddressSpace.gmem, assumed_align=8), + cute.make_layout((self.expert_cnt,)) + ) + sfb_ptr_tensor = cute.make_tensor( + cute.make_ptr(cutlass.Int64, ptrs_sfb.toint(), AddressSpace.gmem, assumed_align=8), + cute.make_layout((self.expert_cnt,)) + ) + + c0 = cutlass.Int64(0) + c1_64 = 1 + if cutlass.const_expr(b_major_mode == OperandMajorMode.K): + stride_n = b_stride_size + stride_k = c1_64 + else: + stride_n = c1_64 + stride_k = b_stride_size + + b_ptr_val = b_ptr_tensor[expert_idx] + b_ptr = cute.make_ptr(self.b_dtype, b_ptr_val, AddressSpace.gmem) + b_tensor_i = cute.make_tensor( + b_ptr, + cute.make_layout((n, k, cutlass.Int32(1)), stride=(stride_n, stride_k, c0)), + ) + tma_atom_b, _ = cute.nvgpu.make_tiled_tma_atom_B( + b_tma_op_arg, b_tensor_i, b_smem_layout_arg, + self.mma_tiler, tiled_mma_arg, + cluster_layout_vmnk_shape_arg, + ) + + workspace = TensormapWorkspace(workspace_ptr, ["b", "sfb"]) + store_tma_desc(tma_atom_b, workspace.get_ptr("b", expert_idx)) + + sfb_ptr_val = sfb_ptr_tensor[expert_idx] + sfb_ptr = cute.make_ptr(self.sf_dtype, sfb_ptr_val, AddressSpace.gmem) + sfb_layout = blockscaled_utils.tile_atom_to_shape_SF( + (n, k, cutlass.Int32(1)), self.sf_vec_size + ) + sfb_tensor_i = cute.make_tensor(sfb_ptr, sfb_layout) + tma_atom_sfb, _ = cute.nvgpu.make_tiled_tma_atom_B( + sfb_tma_op_arg, sfb_tensor_i, sfb_smem_layout_arg, + self.mma_tiler_sfb, tiled_mma_sfb_arg, + cluster_layout_sfb_vmnk_shape_arg, + internal_type=cutlass.Uint64, + ) + store_tma_desc(tma_atom_sfb, workspace.get_ptr("sfb", expert_idx)) + + if cutlass.const_expr(self.use_dynamic_sched): + if expert_idx == cutlass.Int32(0): + sched_counter = cute.make_tensor( + self._get_sched_counter_ptr(workspace_ptr), + cute.make_layout(1), + ) + sched_counter[0] = cutlass.Int32(0) + + @cute.jit + def __call__( + self, + a: cute.Tensor, + b, # Dense: cute.Tensor (N,K,L) | Discrete: cute.Pointer to int64[] + sfb, # Dense: cute.Tensor | Discrete: cute.Pointer to int64[] + n: Int32, # Ignored for dense mode + k: Int32, # Ignored for dense mode + b_stride_size: cutlass.Int64, # Ignored for dense mode + b_major_mode: cutlass.Constexpr, # Ignored for dense mode + workspace_ptr, # Descriptor workspace, plus dynamic scheduler counter when enabled + c: cute.Tensor, + d: cute.Tensor, + d_col: Optional[cute.Tensor], + sfa: cute.Tensor, + sfd_row_tensor: Optional[cute.Tensor], + sfd_col_tensor: Optional[cute.Tensor], + amax_tensor: Optional[cute.Tensor], + norm_const_tensor: Optional[cute.Tensor], + padded_offsets: cute.Tensor, + alpha: cute.Tensor, + beta: cute.Tensor, + prob: cute.Tensor, + dprob: cute.Tensor, + linear_offset: Float32, + dbias_tensor: Optional[cute.Tensor], + max_active_clusters: cutlass.Constexpr, + stream: cuda.CUstream, + epilogue_op: cutlass.Constexpr = lambda x: x, + ): + """Execute the GEMM. + + Dense mode: ``b`` and ``sfb`` are 3-D cute.Tensor (N, K, L). + Discrete mode: ``b`` and ``sfb`` are cute.Pointer to device int64[] + arrays of per-expert base addresses; ``n``, ``k``, ``b_stride_size``, + ``b_major_mode`` describe the uniform per-expert layout. + """ + # Setup static attributes before smem/grid/tma computation + self.a_dtype: Type[cutlass.Numeric] = a.element_type + self.b_dtype: Type[cutlass.Numeric] = a.element_type # B must match A dtype + self.c_dtype: Type[cutlass.Numeric] = c.element_type + self.d_dtype: Type[cutlass.Numeric] = d.element_type + self.sf_dtype: Type[cutlass.Numeric] = sfa.element_type + self.a_major_mode = utils.LayoutEnum.from_tensor(a).mma_major_mode() + + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DENSE): + self.b_major_mode = utils.LayoutEnum.from_tensor(b).mma_major_mode() + else: + self.b_major_mode = b_major_mode + self.c_layout = utils.LayoutEnum.from_tensor(c) + self.d_layout = utils.LayoutEnum.from_tensor(d) + + # dBias configuration + self.generate_dbias = dbias_tensor is not None + self.dbias_cross_warp_reduce = self.generate_dbias # always cross-warp reduce + self.has_prob = prob is not None + self.generate_dprob = dprob is not None + + # Check if input data types are compatible with MMA instruction + if cutlass.const_expr(self.a_dtype != self.b_dtype): + raise TypeError(f"Type must match: {self.a_dtype} != {self.b_dtype}") + + # Setup attributes that dependent on gemm inputs + self._setup_attributes() + + # ---- B / SFB setup (mode-dependent) ---- + b_from_call_arg = b + sfb_from_call_arg = sfb + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DENSE): + sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b.shape, self.sf_vec_size) + sfb = cute.make_tensor(sfb.iterator, sfb_layout) + else: + c0 = cutlass.Int64(0) + c1_64 = 1 + if cutlass.const_expr(b_major_mode == OperandMajorMode.K): + b_template_stride = (b_stride_size, c1_64, c0) + else: + b_template_stride = (c1_64, b_stride_size, c0) + b_template_layout = cute.make_layout((n, k, cutlass.Int32(1)), stride=b_template_stride) + b_ptr_typed = cute.make_ptr( + self.b_dtype, b.toint(), AddressSpace.gmem, assumed_align=16 + ) + b = cute.make_tensor(b_ptr_typed, b_template_layout) + + sfb_ptr_typed = cute.make_ptr( + self.sf_dtype, sfb.toint(), AddressSpace.gmem, assumed_align=16 + ) + sfb_layout = blockscaled_utils.tile_atom_to_shape_SF( + (n, k, cutlass.Int32(1)), self.sf_vec_size + ) + sfb = cute.make_tensor(sfb_ptr_typed, sfb_layout) + + # Setup sfa tensor by filling A tensor to scale factor atom layout + # ((Atom_M, Rest_M),(Atom_K, Rest_K),RestL) + sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a.shape, self.sf_vec_size) + sfa = cute.make_tensor(sfa.iterator, sfa_layout) + + # Compute grid size + m, n_d, l = cute.shape(d) + + # Setup sfd tensor by filling D tensor to scale factor atom layout + self.generate_sfd = ( + self.a_dtype in (cutlass.Float8E5M2, cutlass.Float8E4M3FN) + and self.sf_dtype == cutlass.Float8E8M0FNU + and self.d_dtype in (cutlass.Float8E5M2, cutlass.Float8E4M3FN) + ) + if cutlass.const_expr(self.generate_sfd == False): + self.discrete_col_sfd = False + if cutlass.const_expr(self.generate_sfd): + output_sfd_shape = (m, n_d, l) + sfd_layout = blockscaled_utils.tile_atom_to_shape_SF( + output_sfd_shape, self.sf_vec_size + ) + sfd_row_tensor = cute.make_tensor(sfd_row_tensor.iterator, sfd_layout) + sfd_col_quant_layout = cute.tile_to_shape( + blockscaled_utils.BlockScaledBasicChunk( + self.sf_vec_size, OperandMajorMode.MN + ).layout, + output_sfd_shape, + (1, 2, 3), + ) + if cutlass.const_expr(self.discrete_col_sfd): + sfd_col_quant_layout = sfd_layout + sfd_col_tensor = cute.make_tensor( + sfd_col_tensor.iterator, sfd_col_quant_layout + ) + + self.generate_amax = amax_tensor is not None + + # K dim of the MMA instruction shape on Rubin sm107: + # - SM107MmaMXF4NVF4Op (FP4 x FP4) requires K=128 + # - SM107BlockScaledMmaMXF8F6F4Op (FP8 or mixed) requires K=64 + mma_inst_k = ( + 128 + if (self.a_dtype.width == 4 and self.b_dtype.width == 4) + else 64 + ) + mma_inst_shape_mnk = (*self.mma_inst_shape_mn, mma_inst_k) + mma_inst_shape_mnk_sfb = (*self.mma_inst_shape_mn_sfb, mma_inst_k) + + atom_layout_mnk = (1, 1, 1) + permutation_mnk = self._get_mma_permutation_mnk(mma_inst_shape_mnk) + + tiled_mma = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + self.cta_group, + mma_inst_shape_mnk, + a_collector_op=CollectorOp.DISCARD, + b_collector_op=CollectorOp.DISCARD, + atom_layout_mnk=atom_layout_mnk, + permutation_mnk=permutation_mnk, + ) + + tiled_mma_sfb = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + cute.nvgpu.tcgen05.CtaGroup.ONE, + mma_inst_shape_mnk_sfb, + ) + + tiled_mma_bkeep = None + tiled_mma_breuse = None + if cutlass.const_expr(self.enable_breuse): + tiled_mma_bkeep = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, self.b_dtype, + self.a_major_mode, self.b_major_mode, + self.sf_dtype, self.sf_vec_size, + self.cta_group, + mma_inst_shape_mnk, + a_collector_op=CollectorOp.DISCARD, + b_collector_op=CollectorOp.FILL, + atom_layout_mnk=atom_layout_mnk, + permutation_mnk=permutation_mnk, + ) + tiled_mma_bkeep.set(tcgen05.Field.NEGATE_A, False) + tiled_mma_bkeep.set(tcgen05.Field.NEGATE_B, False) + tiled_mma_breuse = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, self.b_dtype, + self.a_major_mode, self.b_major_mode, + self.sf_dtype, self.sf_vec_size, + self.cta_group, + mma_inst_shape_mnk, + a_collector_op=CollectorOp.DISCARD, + b_collector_op=CollectorOp.LASTUSE, + atom_layout_mnk=atom_layout_mnk, + permutation_mnk=permutation_mnk, + ) + tiled_mma_breuse.set(tcgen05.Field.NEGATE_A, False) + tiled_mma_breuse.set(tcgen05.Field.NEGATE_B, False) + + atom_thr_size = cute.size(tiled_mma.thr_id.shape) + + # Setup TMA load for A + a_op = sm100_utils.cluster_shape_to_tma_atom_A( + self.cluster_shape_mn, tiled_mma.thr_id + ) + a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0)) + tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A( + a_op, + a, + a_smem_layout, + self.mma_tiler, + tiled_mma, + self.cluster_layout_vmnk.shape, + ) + + # Setup TMA load for B + b_op = sm100_utils.cluster_shape_to_tma_atom_B( + self.cluster_shape_mn, tiled_mma.thr_id + ) + b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0)) + tma_atom_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B( + b_op, + b, + b_smem_layout, + self.mma_tiler, + tiled_mma, + self.cluster_layout_vmnk.shape, + ) + + # Setup TMA load for SFA + sfa_op = sm100_utils.cluster_shape_to_tma_atom_A( + self.cluster_shape_mn, tiled_mma.thr_id + ) + sfa_smem_layout = cute.slice_( + self.sfa_smem_layout_staged, (None, None, None, 0) + ) + tma_atom_sfa, tma_tensor_sfa = cute.nvgpu.make_tiled_tma_atom_A( + sfa_op, + sfa, + sfa_smem_layout, + self.mma_tiler, + tiled_mma, + self.cluster_layout_vmnk.shape, + internal_type=cutlass.Uint16, + ) + + # Setup TMA load for SFB + sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB( + self.cluster_shape_mn, tiled_mma.thr_id + ) + sfb_smem_layout = cute.slice_( + self.sfb_smem_layout_staged, (None, None, None, 0) + ) + tma_atom_sfb, tma_tensor_sfb = cute.nvgpu.make_tiled_tma_atom_B( + sfb_op, + sfb, + sfb_smem_layout, + self.mma_tiler_sfb, + tiled_mma_sfb, + self.cluster_layout_sfb_vmnk.shape, + internal_type=cutlass.Uint64, + ) + + # N=192: do NOT reshape tma_tensor_sfb here; reshape real_sfb in mainloop. + + a_copy_size = cute.size_in_bytes(self.a_dtype, a_smem_layout) + b_copy_size = cute.size_in_bytes(self.b_dtype, b_smem_layout) + sfa_copy_size = cute.size_in_bytes(self.sf_dtype, sfa_smem_layout) + sfb_copy_size = cute.size_in_bytes(self.sf_dtype, sfb_smem_layout) + self.num_tma_load_bytes = ( + a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size + ) * atom_thr_size + + # Setup TMA store for C + c_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0)) + self.tma_c_load_bytes = cute.size_in_bytes(self.c_dtype, c_smem_layout) + tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom( + cpasync.CopyBulkTensorTileG2SOp(), + c, + c_smem_layout, + self.epi_tile, + ) + + # Setup TMA store for D + tma_atom_d = None + tma_tensor_d = None + if cutlass.const_expr(not self.store_d_directly): + d_smem_layout = cute.slice_(self.d_smem_layout_staged, (None, None, 0)) + tma_atom_d, tma_tensor_d = cpasync.make_tiled_tma_atom( + cpasync.CopyBulkTensorTileS2GOp(), + d, + d_smem_layout, + self.epi_tile, + ) + else: + tma_tensor_d = d + + tma_atom_d_col = None + tma_tensor_d_col = None + if cutlass.const_expr(self.generate_sfd): + tma_atom_d_col, tma_tensor_d_col = cpasync.make_tiled_tma_atom( + cpasync.CopyBulkTensorTileS2GOp(), + d_col, + d_smem_layout, + self.epi_tile, + ) + + n_half = n_d // 2 + sched_params = MoESchedulerParams( + scenario="2Dx3D", + expert_shape=(self.expert_cnt, n_half, cute.size(a.shape, mode=[1])), + cta_tile_shape_mnk=self.cta_tile_shape_mnk_d, + cluster_shape_mn=self.cluster_shape_mn, + use_dynamic_sched=self.use_dynamic_sched, + ) + grid = MoESchedulerParams.get_grid_shape( + sched_params, max_active_clusters + ) + + self.buffer_align_bytes = 1024 + + # Define shared storage for kernel + # sD_col is only needed when generating SFD; use size 0 to avoid wasting smem + sD_col_size = 0 if not self.generate_sfd else cute.cosize(self.d_smem_layout_staged.outer) + # sD is not needed when storing D directly (and not generating SFD) + sD_size = 0 if (not self.generate_sfd and self.store_d_directly) else cute.cosize(self.d_smem_layout_staged.outer) + SchedulerStorage = MoEPersistentTileScheduler.make_storage_struct( + self.num_tile_stage, self.use_dynamic_sched + ) + + @cute.struct + class SharedStorage: + ab_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage * 2] + acc_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage * 2] + scheduler: SchedulerStorage + c_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_c_stage] + c_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_c_stage] + tmem_dealloc_mbar_ptr: cutlass.Int64 + tmem_holding_buf: cutlass.Int32 + # (EPI_TILE_M, EPI_TILE_N, STAGE) + sC: cute.struct.Align[ + cute.struct.MemRange[ + self.c_dtype, + cute.cosize(self.c_smem_layout_staged.outer), + ], + self.buffer_align_bytes, + ] + sD: cute.struct.Align[ + cute.struct.MemRange[self.d_dtype, sD_size], + self.buffer_align_bytes, + ] + sD_col: cute.struct.Align[ + cute.struct.MemRange[self.d_dtype, sD_col_size], + self.buffer_align_bytes, + ] + # (MMA, MMA_M, MMA_K, STAGE) + sA: cute.struct.Align[ + cute.struct.MemRange[ + self.a_dtype, cute.cosize(self.a_smem_layout_staged.outer) + ], + self.buffer_align_bytes, + ] + # (MMA, MMA_N, MMA_K, STAGE) + sB: cute.struct.Align[ + cute.struct.MemRange[ + self.b_dtype, cute.cosize(self.b_smem_layout_staged.outer) + ], + self.buffer_align_bytes, + ] + # (granularity_m, repeat_m), (granularity_k, repeat_k), num_scale_stage) + sSFA: cute.struct.Align[ + cute.struct.MemRange[ + self.sf_dtype, cute.cosize(self.sfa_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + # (granularity_n, repeat_n), (granularity_k, repeat_k), num_scale_stage) + sSFB: cute.struct.Align[ + cute.struct.MemRange[ + self.sf_dtype, cute.cosize(self.sfb_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + # Amax reduction shared memory (one FP32 per epilogue warp) + sAmax: cute.struct.Align[ + cute.struct.MemRange[cutlass.Float32, self.num_epilog_warps], + 4, + ] + if cutlass.const_expr(self.generate_dbias): + # dBias SMEM transpose buffer: (128, epi_tile_n*2) col-major FP32 + sDbias: cute.struct.Align[ + cute.struct.MemRange[ + cutlass.Float32, + 128 * self.epi_tile[1] * 2 if self.generate_dbias else 1, + ], + 128 if self.generate_dbias else 4, + ] + + self.shared_storage = SharedStorage + + # Initialize per-expert B/SFB TMA descriptors in workspace + b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0)) + sfb_smem_layout = cute.slice_(self.sfb_smem_layout_staged, (None, None, None, 0)) + _need_helper = cutlass.const_expr( + self.weight_mode == MoEWeightMode.DISCRETE or self.use_dynamic_sched + ) + if cutlass.const_expr(_need_helper): + _helper_grid_x = self.expert_cnt if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else 1 + _helper_args = ( + b_from_call_arg if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cute.make_ptr(cutlass.Int64, 0, AddressSpace.gmem), + sfb_from_call_arg if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cute.make_ptr(cutlass.Int64, 0, AddressSpace.gmem), + n if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cutlass.Int32(0), + k if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cutlass.Int32(0), + b_stride_size if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cutlass.Int64(0), + b_major_mode if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else self.b_major_mode, + workspace_ptr, + tiled_mma, tiled_mma_sfb, + b_smem_layout, sfb_smem_layout, + self.cluster_layout_vmnk.shape, + self.cluster_layout_sfb_vmnk.shape, + ) + self.helper_kernel(*_helper_args).launch( + grid=(_helper_grid_x, 1, 1), block=(1, 1, 1), + stream=stream, min_blocks_per_mp=1, + ) + + # Launch the main kernel + self.kernel( + tiled_mma, + tiled_mma_bkeep, + tiled_mma_breuse, + tiled_mma_sfb, + tma_atom_a, + tma_tensor_a, + tma_atom_b, + tma_tensor_b, + tma_atom_sfa, + tma_tensor_sfa, + tma_atom_sfb, + tma_tensor_sfb, + tma_atom_c, + tma_tensor_c, + tma_atom_d, + tma_tensor_d, + tma_atom_d_col, + tma_tensor_d_col, + sfd_row_tensor, + sfd_col_tensor, + norm_const_tensor, + amax_tensor, + padded_offsets, + alpha, + beta, + prob, + dprob, + linear_offset, + dbias_tensor, + workspace_ptr, + self.cluster_layout_vmnk, + self.cluster_layout_sfb_vmnk, + self.a_smem_layout_staged, + self.b_smem_layout_staged, + self.sfa_smem_layout_staged, + self.sfb_smem_layout_staged, + self.c_smem_layout_staged, + self.d_smem_layout_staged, + self.epi_tile, + sched_params, + epilogue_op, + ).launch( + grid=grid, + block=[self.threads_per_cta, 1, 1], + cluster=(*self.cluster_shape_mn, 1), + max_number_threads=[self.threads_per_cta, 1, 1], + smem=self.shared_storage.size_in_bytes(), + stream=stream, + min_blocks_per_mp=1, + ) + return + + @cute.jit + def _make_extension(self, workspace_ptr): + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE): + desc_workspace = TensormapWorkspace(workspace_ptr, ["b", "sfb"]) + return DiscreteWeightScaledGemmSchedExtension( + tensormap_ctor=desc_workspace, sf_vec_size=self.sf_vec_size, + ) + else: + return ContiguousAndConsistentGroupedGemmSchedExtension( + sf_vec_size=self.sf_vec_size, + ) + + def mainloop_s2t_copy_and_partition( + self, + sSF: cute.Tensor, + tSF: cute.Tensor, + ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]: + """ + Make tiledCopy for smem to tmem load for scale factor tensor, then use it to partition smem memory (source) and tensor memory (destination). + + :param sSF: The scale factor tensor in smem + :type sSF: cute.Tensor + :param tSF: The scale factor tensor in tmem + :type tSF: cute.Tensor + + :return: A tuple containing (tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t) where: + - tiled_copy_s2t: The tiled copy operation for smem to tmem load for scale factor tensor(s2t) + - tCsSF_compact_s2t: The partitioned scale factor tensor in smem + - tSF_compact_s2t: The partitioned scale factor tensor in tmem + :rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor] + """ + # (MMA, MMA_MN, MMA_K, STAGE) + tCsSF_compact = cute.filter_zeros(sSF) + # (MMA, MMA_MN, MMA_K) + tCtSF_compact = cute.filter_zeros(tSF) + + # Make S2T CopyAtom and tiledCopy + copy_atom_s2t = cute.make_copy_atom( + tcgen05.Cp4x32x128bOp(self.cta_group), + self.sf_dtype, + ) + tiled_copy_s2t = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSF_compact) + thr_copy_s2t = tiled_copy_s2t.get_slice(0) + + def _append_mn_broadcast_mode(smem_layout: cute.Layout): + mn_dim = cute.get(smem_layout, mode=[0, 0]) + mn_dim = cute.append(mn_dim, cute.make_layout((4), stride=(0))) + layout = cute.append( + cute.group_modes(mn_dim, 0), cute.get(smem_layout, mode=[0, 1]) + ) + layout = cute.append( + cute.group_modes(layout, 0), cute.get(smem_layout, mode=[1]) + ) + layout = cute.append(layout, cute.get(smem_layout, mode=[2])) + layout = cute.append(layout, cute.get(smem_layout, mode=[3])) + return layout + + tCsSF_compact_bcast = cute.make_tensor( + tCsSF_compact.iterator, _append_mn_broadcast_mode(tCsSF_compact.layout) + ) + + # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE) + tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact_bcast) + # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE) + tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor( + tiled_copy_s2t, tCsSF_compact_s2t_ + ) + # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K) + tCtSF_compact_s2t = thr_copy_s2t.partition_D(tCtSF_compact) + + return tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t + + @cute.jit + def amax_reduction_per_thread(self, vec_fp32, amax_fp32) -> None: + vec_fp32_ssa = vec_fp32 + abs_acc_values_ir = cutlass._mlir.dialects.math.absf(vec_fp32_ssa.ir_value()) + abs_acc_values = type(vec_fp32_ssa)( + abs_acc_values_ir, vec_fp32_ssa.shape, vec_fp32_ssa.dtype + ) + subtile_amax = abs_acc_values.reduce( + cute.ReductionOp.MAX, cutlass.Float32(0.0), 0 + ) + return cute.arch.fmax(amax_fp32, subtile_amax) + + @cute.jit + def amax_reduction_per_warp_and_cta( + self, amax_fp32, warp_idx, amax_smem, amax_gmem + ) -> None: + # Warp-level reduction using wrapper function + warp_amax = cute.arch.warp_redux_sync( + value=amax_fp32, + kind="fmax", + mask_and_clamp=0xFFFFFFFF, + nan=True, + ) + # Each epilogue warp's lane 0 writes warp amax to shared memory + lane_idx = cute.arch.thread_idx()[0] % 32 + if lane_idx == 0: + amax_smem[warp_idx] = cutlass.Float32(warp_amax) + + # Ensure all epilogue warps complete their writes before block reduction + self.epilog_sync_barrier.arrive_and_wait() + + # Block-level reduction: only first epilogue warp's lane 0 handles this + if warp_idx == self.epilog_warp_id[0] and lane_idx == 0: + block_amax = cutlass.Float32(0.0) + for i in cutlass.range(self.num_epilog_warps): + warp_amax_val = amax_smem[i] + block_amax = cute.arch.fmax(block_amax, warp_amax_val) + + # Global atomic max (accumulates across all tiles for final tensor amax) + _ = atomic_max_float32(ptr=amax_gmem, value=block_amax) + + # Ensure all epilogue warps complete their writes before global reduction + self.epilog_sync_barrier.arrive_and_wait() + + @cute.jit + def cvt_f32x4_to_f8x4_pack_i32(self, fp32x4, fp8_type, loc=None, ip=None): + fp32x4 = fp32x4.load() + src_vec4 = ( + fp32x4.ir_value(loc=loc, ip=ip) if hasattr(fp32x4, "ir_value") else fp32x4 + ) + + src0 = Float32(vector.extract(src_vec4, [], [0])).ir_value(loc=loc, ip=ip) + src1 = Float32(vector.extract(src_vec4, [], [1])).ir_value(loc=loc, ip=ip) + src2 = Float32(vector.extract(src_vec4, [], [2])).ir_value(loc=loc, ip=ip) + src3 = Float32(vector.extract(src_vec4, [], [3])).ir_value(loc=loc, ip=ip) + + cvt_instruction = "" + if cutlass.const_expr(fp8_type == cutlass.Float8E8M0FNU): + cvt_instruction = "cvt.rp.satfinite.ue8m0x2.f32" + elif cutlass.const_expr(fp8_type == cutlass.Float8E4M3FN): + cvt_instruction = "cvt.rn.satfinite.e4m3x2.f32" + else: + with cute.arch.elect_one(): + cute.printf("error: unsupported fp8 element type") + return + + asm_tmpl = ( + "{\n" + " .reg .b16 lo;\n" + " .reg .b16 hi;\n" + f" {cvt_instruction} lo, $2, $1;\n" + f" {cvt_instruction} hi, $4, $3;\n" + " mov.b32 $0, {lo, hi};\n" + "}" + ) + packed_i32 = llvm.inline_asm( + T.i32(), + [src0, src1, src2, src3], + asm_tmpl, + "=r,f,f,f,f", + has_side_effects=True, + is_align_stack=False, + asm_dialect=llvm.AsmDialect.AD_ATT, + ) + + return packed_i32 + + @cute.jit + def cvt_f32x4_to_f8x4(self, fp32x4, fp8x4, loc=None, ip=None): + packed_i32 = self.cvt_f32x4_to_f8x4_pack_i32(fp32x4, fp8x4.element_type) + fp8x4_i32 = cute.recast_tensor(fp8x4, cutlass.Int32) + fp8x4_i32[0] = cutlass.Int32(packed_i32) + return + + @cute.jit + def cvt_f32_to_f8_to_f32(self, fp32x1, fp8_type, loc=None, ip=None): + src_fp32 = Float32(fp32x1).ir_value(loc=loc, ip=ip) + + cvt_instruction_downcast = "" + cvt_instruction_upcast = "" + if cutlass.const_expr(fp8_type == cutlass.Float8E8M0FNU): + cvt_instruction_downcast = "cvt.rp.satfinite.ue8m0x2.f32" + cvt_instruction_upcast = "cvt.rn.bf16x2.ue8m0x2" + elif cutlass.const_expr(fp8_type == cutlass.Float8E4M3FN): + cvt_instruction_downcast = "cvt.rn.satfinite.e4m3x2.f32" + cvt_instruction_upcast = "cvt.rn.bf16x2.e4m3x2" + else: + with cute.arch.elect_one(): + cute.printf("error: unsupported fp8 element type") + return + + asm_tmpl = ( + "{\n" + " .reg .b16 bf_lo;\n" + f" {cvt_instruction_downcast} bf_lo, 0f00000000, $1;\n" + f" {cvt_instruction_upcast} $0, bf_lo;\n" + "}" + ) + packed_i32 = llvm.inline_asm( + T.i32(), + [src_fp32], + asm_tmpl, + "=r,f", + has_side_effects=True, + is_align_stack=False, + asm_dialect=llvm.AsmDialect.AD_ATT, + ) + + vec_bf16_ty = ir.Type.parse("vector<2xbf16>") + bf2_lo = llvm.bitcast(vec_bf16_ty, packed_i32, loc=loc, ip=ip) + h0 = vector.extract(bf2_lo, [], [0], loc=loc, ip=ip) + dst_f32 = arith.extf(Float32.mlir_type, h0, loc=loc, ip=ip) + + return dst_f32 + + @cute.jit + def quant_sfd_row( + self, + tile_idx, + tiled_copy_r2s, + src, + pvscale, + norm_const, + rcp_limit, + tRSrD, + ) -> None: + # Get absolute max across a vector and Compute SFD + tCompute = cute.make_rmem_tensor(src.shape, self.acc_dtype) + tCompute.store(src) + tTR_rAcc_frg = cute.logical_divide(tCompute, cute.make_layout(self.sf_vec_size)) + acc_frg = tTR_rAcc_frg.load() + abs_acc_frg_ir = cutlass._mlir.dialects.math.absf(acc_frg.ir_value()) + abs_acc_frg = type(acc_frg)(abs_acc_frg_ir, acc_frg.shape, acc_frg.dtype) + + avg_fp32 = ( + abs_acc_frg[None, 0].reduce( + cute.ReductionOp.MAX, + cutlass.Float32(0.0), + 0, # Use 0.0 as init for abs values + ) + * rcp_limit + * norm_const + ) + # + # Manually store pvscale to avoid spilling + # + if tile_idx == 0: + pvscale[None, None, 0][0] = avg_fp32 + elif tile_idx == 1: + pvscale[None, None, 1][0] = avg_fp32 + elif tile_idx == 2: + pvscale[None, None, 2][0] = avg_fp32 + elif tile_idx == 3: + pvscale[None, None, 3][0] = avg_fp32 + + # + # Compute quantized output values and convert to D type + # + qpvscale_up = self.cvt_f32_to_f8_to_f32(avg_fp32, self.sf_dtype) + + fp32_max = cutlass.Float32(3.40282346638528859812e38) + acc_scale = norm_const * cute.arch.rcp_approx(qpvscale_up) + acc_scale = fmin(acc_scale, fp32_max, nan=True) + vec = tTR_rAcc_frg[None, 0] + if cutlass.const_expr(self.vectorized_f32): + for ei in cutlass.range_constexpr(0, self.sf_vec_size, 2): + ( + vec[ei], + vec[ei + 1], + ) = cute.arch.mul_packed_f32x2( + (vec[ei], vec[ei + 1]), + (acc_scale, acc_scale), + rnd='rn', + ftz=False, + ) + else: + for ei in cutlass.range_constexpr(self.sf_vec_size): + vec[ei] = vec[ei] * acc_scale + + acc_vec = tiled_copy_r2s.retile(tCompute).load() + if cutlass.const_expr(not self.use_fp8_ptx_cvt): + tRSrD.store(acc_vec.to(self.d_dtype)) + else: + tRSrD_i32 = cute.recast_tensor(tRSrD, cutlass.Int32) + for ei in cutlass.range_constexpr(0, self.sf_vec_size, 4): + fp32x4 = cute.make_rmem_tensor(4, cutlass.Float32) + fp32x4[0] = acc_vec[ei + 0] + fp32x4[1] = acc_vec[ei + 1] + fp32x4[2] = acc_vec[ei + 2] + fp32x4[3] = acc_vec[ei + 3] + fp8x4_i32 = self.cvt_f32x4_to_f8x4_pack_i32(fp32x4, self.d_dtype) + tRSrD_i32[ei // 4] = cutlass.Int32(fp8x4_i32) + + @cute.jit + def quant_sfd_col( + self, + tile_idx, + tiled_copy_r2s, + src, + pvscale, + norm_const, + rcp_limit, + tRSrD, + ): + # Get absolute max across a vector and Compute SFD + tCompute = cute.make_rmem_tensor(src.shape, self.acc_dtype) + tCompute.store(src) + tTR_rAcc_frg = cute.logical_divide(tCompute, cute.make_layout(self.sf_vec_size)) + acc_frg = tTR_rAcc_frg.load() + abs_acc_frg_ir = cutlass._mlir.dialects.math.absf(acc_frg.ir_value()) + acc_frg = type(acc_frg)(abs_acc_frg_ir, acc_frg.shape, acc_frg.dtype) + + tmp_fp32 = cutlass.Float32(0.0) + fp32_max = cutlass.Float32(3.40282346638528859812e38) + tidx, _, _ = cute.arch.thread_idx() + + for vi in cutlass.range_constexpr(0, acc_frg.shape[0], 4): + max_value0 = cutlass.Float32( + cute.arch.warp_redux_sync( + value=acc_frg[vi, 0], + kind="fmax", + mask_and_clamp=0xFFFFFFFF, + nan=True, + ) + ) + max_value1 = cutlass.Float32( + cute.arch.warp_redux_sync( + value=acc_frg[vi + 1, 0], + kind="fmax", + mask_and_clamp=0xFFFFFFFF, + nan=True, + ) + ) + max_value2 = cutlass.Float32( + cute.arch.warp_redux_sync( + value=acc_frg[vi + 2, 0], + kind="fmax", + mask_and_clamp=0xFFFFFFFF, + nan=True, + ) + ) + max_value3 = cutlass.Float32( + cute.arch.warp_redux_sync( + value=acc_frg[vi + 3, 0], + kind="fmax", + mask_and_clamp=0xFFFFFFFF, + nan=True, + ) + ) + + scale = rcp_limit * norm_const + (max_value0, max_value1) = cute.arch.mul_packed_f32x2( + (max_value0, max_value1), + (scale, scale), + rnd='rn', + ftz=False, + ) + (max_value2, max_value3) = cute.arch.mul_packed_f32x2( + (max_value2, max_value3), + (scale, scale), + rnd='rn', + ftz=False, + ) + + if tidx % 32 == vi: + tmp_fp32 = max_value0 + if tidx % 32 == vi + 1: + tmp_fp32 = max_value1 + if tidx % 32 == vi + 2: + tmp_fp32 = max_value2 + if tidx % 32 == vi + 3: + tmp_fp32 = max_value3 + + max_value_tensor = cute.make_rmem_tensor(4, cutlass.Float32) + max_value_tensor[0] = max_value0 + max_value_tensor[1] = max_value1 + max_value_tensor[2] = max_value2 + max_value_tensor[3] = max_value3 + + if cutlass.const_expr(not self.use_fp8_ptx_cvt): + max_value_vec_f8 = max_value_tensor.load().to(self.sf_dtype) + else: + max_value_vec_f8 = cute.make_rmem_tensor(4, self.sf_dtype) + self.cvt_f32x4_to_f8x4(max_value_tensor, max_value_vec_f8) + max_value_vec_f8 = max_value_vec_f8.load() + max_value_vec_f32_chunked = max_value_vec_f8.to(cutlass.Float32) + max_value0 = max_value_vec_f32_chunked[0] + max_value1 = max_value_vec_f32_chunked[1] + max_value2 = max_value_vec_f32_chunked[2] + max_value3 = max_value_vec_f32_chunked[3] + + max_value_rcp0 = cute.arch.rcp_approx(max_value0) + max_value_rcp1 = cute.arch.rcp_approx(max_value1) + max_value_rcp2 = cute.arch.rcp_approx(max_value2) + max_value_rcp3 = cute.arch.rcp_approx(max_value3) + + max_value_rcp0 = fmin(max_value_rcp0, fp32_max, nan=True) + max_value_rcp1 = fmin(max_value_rcp1, fp32_max, nan=True) + max_value_rcp2 = fmin(max_value_rcp2, fp32_max, nan=True) + max_value_rcp3 = fmin(max_value_rcp3, fp32_max, nan=True) + + (acc_scale_col0, acc_scale_col1) = cute.arch.mul_packed_f32x2( + (norm_const, norm_const), + (max_value_rcp0, max_value_rcp1), + rnd='rn', + ftz=False, + ) + (acc_scale_col2, acc_scale_col3) = cute.arch.mul_packed_f32x2( + (norm_const, norm_const), + (max_value_rcp2, max_value_rcp3), + rnd='rn', + ftz=False, + ) + + (tTR_rAcc_frg[vi], tTR_rAcc_frg[vi + 1]) = cute.arch.mul_packed_f32x2( + (tTR_rAcc_frg[vi], tTR_rAcc_frg[vi + 1]), + (acc_scale_col0, acc_scale_col1), + rnd='rn', + ftz=False, + ) + (tTR_rAcc_frg[vi + 2], tTR_rAcc_frg[vi + 3]) = cute.arch.mul_packed_f32x2( + (tTR_rAcc_frg[vi + 2], tTR_rAcc_frg[vi + 3]), + (acc_scale_col2, acc_scale_col3), + rnd='rn', + ftz=False, + ) + + pvscale[None, None, tile_idx][0] = tmp_fp32 + + acc_vec = tiled_copy_r2s.retile(tCompute).load() + if cutlass.const_expr(not self.use_fp8_ptx_cvt): + tRSrD.store(acc_vec.to(self.d_dtype)) + else: + tRSrD_i32 = cute.recast_tensor(tRSrD, cutlass.Int32) + for ei in cutlass.range_constexpr(0, self.sf_vec_size, 4): + fp32x4 = cute.make_rmem_tensor(4, cutlass.Float32) + fp32x4[0] = acc_vec[ei + 0] + fp32x4[1] = acc_vec[ei + 1] + fp32x4[2] = acc_vec[ei + 2] + fp32x4[3] = acc_vec[ei + 3] + fp8x4_i32 = self.cvt_f32x4_to_f8x4_pack_i32(fp32x4, self.d_dtype) + tRSrD_i32[ei // 4] = cutlass.Int32(fp8x4_i32) + + @cute.jit + def stg_256(self, ptr, vec8_f32, *, loc=None, ip=None): + """ + Store 8xf32 (256b) to global memory with L1::no_allocate. + ptr: pointer (byte addressable) + vec8_f32: vector<8xf32> to store + """ + dst = ptr.ir_value(loc=loc, ip=ip) if hasattr(ptr, "ir_value") else ptr + src = ( + vec8_f32.ir_value(loc=loc, ip=ip) + if hasattr(vec8_f32, "ir_value") + else vec8_f32 + ) + dummy = llvm.inline_asm( + T.i32(), + [ + dst, + vector.extract(src, [], [0], loc=loc, ip=ip), + vector.extract(src, [], [1], loc=loc, ip=ip), + vector.extract(src, [], [2], loc=loc, ip=ip), + vector.extract(src, [], [3], loc=loc, ip=ip), + vector.extract(src, [], [4], loc=loc, ip=ip), + vector.extract(src, [], [5], loc=loc, ip=ip), + vector.extract(src, [], [6], loc=loc, ip=ip), + vector.extract(src, [], [7], loc=loc, ip=ip), + ], + "st.global.L1::no_allocate.v8.f32 [$1], {$2, $3, $4, $5, $6, $7, $8, $9}; mov.u32 $0, 0;", + "=r,l,f,f,f,f,f,f,f,f", + has_side_effects=True, + is_align_stack=False, + asm_dialect=llvm.AsmDialect.AD_ATT, + loc=loc, + ip=ip, + ) + + @cute.jit + def store_global_memory_256b(self, dst: cute.Tensor, src: cute.Tensor): + vec_shape = cute.make_layout(8) + dst_f32 = cute.flatten(cute.recast_tensor(dst, cutlass.Float32)) + src_f32 = cute.flatten(cute.recast_tensor(src, cutlass.Float32)) + dst_vf32x8 = cute.logical_divide(dst_f32, vec_shape) + src_vf32x8 = cute.logical_divide(src_f32, vec_shape) + for ei in cutlass.range_constexpr(dst_vf32x8.shape[1]): + self.stg_256( + dst_vf32x8[None, ei].iterator.llvm_ptr, src_vf32x8[None, ei].load() + ) + + @cute.jit + def create_and_partition_new_SFDCol( + self, + mSFDCol_gemm_domain: cute.Tensor, + ): + """Partition SFD Col tensor that is already in per-expert GEMM domain. + + Reinterprets the tensor with BlockScaledBasicChunk layout, then + does local_tile + partition for epilogue store. + """ + current_m = cute.size(mSFDCol_gemm_domain, mode=[0]) + current_n = cute.size(mSFDCol_gemm_domain, mode=[1]) + + sfd_col_quant_layout = cute.tile_to_shape( + blockscaled_utils.BlockScaledBasicChunk( + self.sf_vec_size, OperandMajorMode.MN + ).layout, + (current_m, current_n, mSFDCol_gemm_domain.shape[2]), + (1, 2, 3), + ) + regPerSubtile = 4 + sfd_tile = ( + cute.make_layout(128), + cute.make_layout(32 * regPerSubtile), + ) + mSFDCol_reinterp = cute.make_tensor(mSFDCol_gemm_domain.iterator, sfd_col_quant_layout) + gSFDCol_new = cute.local_tile(mSFDCol_reinterp, sfd_tile, (None, None, None)) + + thr_layout = cute.make_ordered_layout((4, 32), order=(1, 0)) + val_layout = cute.make_ordered_layout((1,), order=(0,)) + copy_atom_sfd_col_quant = cute.make_copy_atom( + cute.nvgpu.CopyUniversalOp(), + gSFDCol_new.element_type, + num_bits_per_copy=8, + ) + tiled_copy_sfd_col_quant = cute.make_tiled_copy_tv( + copy_atom_sfd_col_quant, thr_layout, val_layout + ) + tidx = cute.arch.thread_idx()[0] + thr_copy_sfd_col_quant = tiled_copy_sfd_col_quant.get_slice(tidx) + tCgSFDCol_mnl = thr_copy_sfd_col_quant.partition_D( + cute.filter_zeros(gSFDCol_new) + ) + tCgSFDCol_mnl = cute.filter_zeros(tCgSFDCol_mnl) + return tCgSFDCol_mnl + + @cute.jit + def dbias_reduction( + self, d1_vec, d2_vec, warp_idx, sDbias, + dbias_gmem_2d, expert_idx, n_base_d1, n_base_d2, + dbias_n_total, + ) -> None: + """Merged dy1+dy2 dbias reduction via SMEM transpose.""" + epi_n = self.epi_tile[1] + lane_idx = cute.arch.lane_idx() + warp_local = warp_idx - self.epilog_warp_id[0] + + for n in cutlass.range(epi_n, unroll_full=True): + sDbias[(n, lane_idx, warp_local)] = d1_vec[n] + sDbias[(epi_n + n, lane_idx, warp_local)] = d2_vec[n] + + self.epilog_sync_barrier.arrive_and_wait() + + col_a = 2 * lane_idx if lane_idx < 16 else epi_n + 2 * (lane_idx - 16) + col_b = col_a + 1 + + copy_128bit_atom = cute.make_copy_atom( + cute.nvgpu.CopyUniversalOp(), cutlass.Float32, num_bits_per_copy=128 + ) + warp_base_ptr = sDbias.iterator + warp_local * epi_n * 2 * 32 + swizzle_a = ((col_a >> 1) & 0x7) << 2 + swizzle_b = ((col_b >> 1) & 0x7) << 2 + + sum_a = cutlass.Float32(0.0) + sum_b = cutlass.Float32(0.0) + rDst_a = cute.make_rmem_tensor(cute.make_layout((4,)), cutlass.Float32) + rDst_b = cute.make_rmem_tensor(cute.make_layout((4,)), cutlass.Float32) + for g in cutlass.range(8, unroll_full=True): + m_base = g * 4 + sw_offset_a = col_a * 32 + (m_base ^ swizzle_a) + sSrc_a = cute.make_tensor(warp_base_ptr + sw_offset_a, cute.make_layout((4,))) + cute.copy_atom_call(copy_128bit_atom, sSrc_a, rDst_a) + + sw_offset_b = col_b * 32 + (m_base ^ swizzle_b) + sSrc_b = cute.make_tensor(warp_base_ptr + sw_offset_b, cute.make_layout((4,))) + cute.copy_atom_call(copy_128bit_atom, sSrc_b, rDst_b) + + for i in cutlass.range(4, unroll_full=True): + sum_a = sum_a + rDst_a[i] + sum_b = sum_b + rDst_b[i] + + n_offset = (n_base_d1 + 2 * lane_idx) if lane_idx < 16 else (n_base_d2 + 2 * (lane_idx - 16)) + + if cutlass.const_expr(self.dbias_cross_warp_reduce): + reduce_base = sDbias.iterator + copy_64bit_atom = cute.make_copy_atom( + cute.nvgpu.CopyUniversalOp(), cutlass.Float32, num_bits_per_copy=64 + ) + + self.epilog_sync_barrier.arrive_and_wait() + rSrc_partial = cute.make_rmem_tensor(cute.make_layout((2,)), cutlass.Float32) + rSrc_partial[0] = sum_a + rSrc_partial[1] = sum_b + sDst_partial = cute.make_tensor( + reduce_base + warp_local * 64 + lane_idx * 2, + cute.make_layout((2,)) + ) + cute.copy_atom_call(copy_64bit_atom, rSrc_partial, sDst_partial) + self.epilog_sync_barrier.arrive_and_wait() + + if warp_idx == self.epilog_warp_id[0]: + cta_sum_a = cutlass.Float32(0.0) + cta_sum_b = cutlass.Float32(0.0) + rDst_w = cute.make_rmem_tensor(cute.make_layout((2,)), cutlass.Float32) + for w in cutlass.range(self.num_epilog_warps): + sSrc_w = cute.make_tensor( + reduce_base + w * 64 + lane_idx * 2, + cute.make_layout((2,)) + ) + cute.copy_atom_call(copy_64bit_atom, sSrc_w, rDst_w) + cta_sum_a = cta_sum_a + rDst_w[0] + cta_sum_b = cta_sum_b + rDst_w[1] + if n_offset < dbias_n_total: + gmem_ptr = dbias_gmem_2d[(expert_idx, n_offset, None)].iterator.llvm_ptr + atomic_add_bf16x2(gmem_ptr, cta_sum_a, cta_sum_b) + else: + if n_offset < dbias_n_total: + gmem_ptr = dbias_gmem_2d[(expert_idx, n_offset, None)].iterator.llvm_ptr + atomic_add_bf16x2(gmem_ptr, sum_a, sum_b) + + @cute.jit + def dswiglu(self, + acc_vec: cute.Tensor, + ab1_vec_load: cute.Tensor, + ab2_vec_load: cute.Tensor, + mProb: cute.Tensor, + beta_val: Float32, + square_alpha: Float32, + dprob_swiglu: Optional[cute.Tensor] = None + ): + LOG2_E = cutlass.Float32(1.4426950408889634) + if cutlass.const_expr(self.vectorized_f32): + d1_vec = cute.make_rmem_tensor(acc_vec.shape, cutlass.Float32) + d2_vec = cute.make_rmem_tensor(acc_vec.shape, cutlass.Float32) + for i in cutlass.range( + 0, cute.size(acc_vec), 2, unroll_full=True + ): + # Apply scaling factors for FP8 + ( + acc_vec[i + 0], + acc_vec[i + 1], + ) = cute.arch.mul_packed_f32x2( + (acc_vec[i + 0], acc_vec[i + 1]), + (square_alpha, square_alpha), + rnd='rn', + ftz=False, + ) + ab1_vec_acc_type = cute.arch.mul_packed_f32x2( + ( + ab1_vec_load[i + 0].to(self.acc_dtype), + ab1_vec_load[i + 1].to(self.acc_dtype), + ), + (beta_val, beta_val), + rnd='rn', + ftz=False, + ) + ab2_vec_acc_type = cute.arch.mul_packed_f32x2( + ( + ab2_vec_load[i + 0].to(self.acc_dtype), + ab2_vec_load[i + 1].to(self.acc_dtype), + ), + (beta_val, beta_val), + rnd='rn', + ftz=False, + ) + (sig_rcp_0, sig_rcp_1) = cute.arch.mul_packed_f32x2( + (ab1_vec_acc_type), + (-LOG2_E, -LOG2_E), + rnd='rn', + ftz=False, + ) + (sig_rcp_0, sig_rcp_1) = cute.arch.add_packed_f32x2( + ( + cute.math.exp2(sig_rcp_0, fastmath=True), + cute.math.exp2(sig_rcp_1, fastmath=True), + ), + (1.0, 1.0), + rnd='rn', + ftz=False, + ) + sig = ( + cute.arch.rcp_approx(sig_rcp_0), + cute.arch.rcp_approx(sig_rcp_1), + ) + swish = cute.arch.mul_packed_f32x2( + ab1_vec_acc_type, + sig, + rnd='rn', + ftz=False, + ) + # calculate dprob + if cutlass.const_expr(self.generate_dprob): + ( + dprob_swiglu[i + 0], + dprob_swiglu[i + 1], + ) = cute.arch.mul_packed_f32x2( + (ab2_vec_acc_type[0], ab2_vec_acc_type[1]), + swish, + ) + ( + dprob_swiglu[i + 0], + dprob_swiglu[i + 1], + ) = cute.arch.mul_packed_f32x2( + (dprob_swiglu[i + 0], dprob_swiglu[i + 1]), + (acc_vec[i + 0], acc_vec[i + 1]), + ) + # calculate dswiglu + acc_vec_prob = (acc_vec[i + 0], acc_vec[i + 1]) + if cutlass.const_expr(self.has_prob): + acc_vec_prob = cute.arch.mul_packed_f32x2( + acc_vec_prob, + (mProb, mProb), + ) + # calculate d2_vec + ( + d2_vec[i + 0], + d2_vec[i + 1], + ) = cute.arch.mul_packed_f32x2( + (acc_vec_prob[0], acc_vec_prob[1]), + swish, + rnd='rn', + ftz=False, + ) + # calculate d1_vec + ( + d1_vec[i + 0], + d1_vec[i + 1], + ) = cute.arch.mul_packed_f32x2( + (acc_vec_prob[0], acc_vec_prob[1]), + (ab2_vec_acc_type[0], ab2_vec_acc_type[1]), + rnd='rn', + ftz=False, + ) + ( + d1_vec[i + 0], + d1_vec[i + 1], + ) = cute.arch.mul_packed_f32x2( + (d1_vec[i + 0], d1_vec[i + 1]), + sig, + rnd='rn', + ftz=False, + ) + one_minus_sig = cute.arch.add_packed_f32x2( + (1.0, 1.0), + (-sig[0], -sig[1]), + rnd='rn', + ftz=False, + ) + dsig = cute.arch.mul_packed_f32x2( + ab1_vec_acc_type, + one_minus_sig, + rnd='rn', + ftz=False, + ) + dsig_add_1 = cute.arch.add_packed_f32x2( + (dsig[0], dsig[1]), + (1.0, 1.0), + rnd='rn', + ftz=False, + ) + ( + d1_vec[i + 0], + d1_vec[i + 1], + ) = cute.arch.mul_packed_f32x2( + (d1_vec[i + 0], d1_vec[i + 1]), + dsig_add_1, + rnd='rn', + ftz=False, + ) + d1_vec = d1_vec.load() + d2_vec = d2_vec.load() + if cutlass.const_expr(self.generate_dprob): + dprob_swiglu = dprob_swiglu.load() + return d1_vec, d2_vec, dprob_swiglu + else: + acc_vec = acc_vec.load() + ab1_vec_load = ab1_vec_load.load() + ab2_vec_load = ab2_vec_load.load() + + acc_vec = acc_vec * square_alpha # apply scale for A*B + ab1_vec_load = ab1_vec_load * beta_val # apply scale for C + ab2_vec_load = ab2_vec_load * beta_val # apply scale for C + + sig_rcp = (1 + cute.math.exp(-1 * ab1_vec_load, True)).to( + self.acc_dtype + ) + res = cute.make_rmem_tensor(sig_rcp.shape, cutlass.Float32) + res.store(sig_rcp) + # let every res[?] be cute.arch.rcp_approx(res[?]) + [ + res.__setitem__(i, cute.arch.rcp_approx(res[i])) + for i in range(cute.size(res.shape)) + ] + sig = res.load() + swish = ab1_vec_load * sig + + # calculate dprob + if cutlass.const_expr(self.generate_dprob): + dprob_swiglu = ab2_vec_load * swish + dprob_swiglu = acc_vec * dprob_swiglu + + # calculate dswiglu + acc_vec_prob = acc_vec + if cutlass.const_expr(self.has_prob): + acc_vec_prob = acc_vec * mProb + d1_vec = ( + acc_vec_prob + * ab2_vec_load + * sig + * (1 + ab1_vec_load * (1 - sig)) + ) + d2_vec = acc_vec_prob * swish + return d1_vec, d2_vec, dprob_swiglu + + @cute.jit + def dgeglu(self, + acc_vec: cute.Tensor, + x1_vec_load: cute.Tensor, + x2_vec_load: cute.Tensor, + mProb: cute.Tensor, + linear_offset: Float32, + dprob_swiglu: Optional[cute.Tensor] = None + ): + LOG2_E = cutlass.Float32(1.4426950408889634) + x_dtype = x1_vec_load.element_type + geglu_max_value = x_dtype(7.0) + geglu_min_value = x_dtype(-7.0) + zero_x_dtype = x_dtype(0.0) + fmul2 = partial(cute.arch.mul_packed_f32x2, rnd='rn', ftz=False) + fadd2 = partial(cute.arch.add_packed_f32x2, rnd='rn', ftz=False) + scale_1702 = (1.702, 1.702) + ones2 = (1.0, 1.0) + mprob2 = (mProb, mProb) + linear_offset2 = (linear_offset, linear_offset) + + if cutlass.const_expr(self.vectorized_f32): + dx1_vec = cute.make_rmem_tensor(acc_vec.shape, cutlass.Float32) + dx2_vec = cute.make_rmem_tensor(acc_vec.shape, cutlass.Float32) + for i in cutlass.range( + 0, cute.size(acc_vec), 2, unroll_full=True + ): + acc = (acc_vec[i], acc_vec[i + 1]) + x1_0 = x1_vec_load[i] + x1_1 = x1_vec_load[i + 1] + x2_0 = x2_vec_load[i] + x2_1 = x2_vec_load[i + 1] + + y1_0 = 0.0; y1_1 = 0.0; y2_0 = 0.0; y2_1 = 0.0 + if cutlass.const_expr(x_dtype == cutlass.BFloat16): + y1_0, y1_1 = fmin_bf16x2(x1_0, x1_1, geglu_max_value, geglu_max_value) + y2_0, y2_1 = fmax_bf16x2(x2_0, x2_1, geglu_min_value, geglu_min_value) + y2_0, y2_1 = fmin_bf16x2(y2_0, y2_1, geglu_max_value, geglu_max_value) + y1_0 = Float32(y1_0) + y1_1 = Float32(y1_1) + y2_0 = Float32(y2_0) + y2_1 = Float32(y2_1) + else: + y1_0 = fmin(x1_0, geglu_max_value) + y1_1 = fmin(x1_1, geglu_max_value) + y2_0 = fmin(x2_0, geglu_max_value) + y2_1 = fmin(x2_1, geglu_max_value) + y2_0 = fmax(y2_0, geglu_min_value) + y2_1 = fmax(y2_1, geglu_min_value) + y1_0 = Float32(y1_0) + y1_1 = Float32(y1_1) + y2_0 = Float32(y2_0) + y2_1 = Float32(y2_1) + + y1 = (y1_0, y1_1) + y2 = (y2_0, y2_1) + + # y1 = 1.702 * x1 + y1_scaled = fmul2(y1, scale_1702) + + sigmoid_out_0 = sigmoid_f32(y1_scaled[0], fastmath=True) + sigmoid_out_1 = sigmoid_f32(y1_scaled[1], fastmath=True) + + # g * sigmoid_out + acc_mul_sigmoid_out = fmul2( + acc, (sigmoid_out_0, sigmoid_out_1) + ) + acc_mul_sigmoid_prob = acc_mul_sigmoid_out + if cutlass.const_expr(self.has_prob): + acc_mul_sigmoid_prob = fmul2( + acc_mul_sigmoid_out, mprob2 + ) + + # y1 = 1 + 1.702 * y1 * (1 - sigmoid_out) + one_minus_sigmoid_0, one_minus_sigmoid_1 = fadd2( + ones2, (-sigmoid_out_0, -sigmoid_out_1) + ) + y1_scaled = fadd2( + fmul2(y1_scaled, (one_minus_sigmoid_0, one_minus_sigmoid_1)), + ones2, + ) + + # y2 + linear_offset + y2_with_linear_offset_0, y2_with_linear_offset_1 = fadd2( + y2, linear_offset2 + ) + + # dy1 = g * sigmoid_out * (y2 + linear_offset) + dy1_pre_0, dy1_pre_1 = fmul2( + (y2_with_linear_offset_0, y2_with_linear_offset_1), + acc_mul_sigmoid_out, + ) + # dy1 = g * sigmoid_out * (y2 + linear_offset) * (1 + 1.702 * y1 * (1 - sigmoid_out)) * mProb + dy1_0, dy1_1 = fmul2((dy1_pre_0, dy1_pre_1), y1_scaled) + if cutlass.const_expr(self.has_prob): + dy1_0, dy1_1 = fmul2((dy1_0, dy1_1), mprob2) + + x1_filter_0 = y1_0 if x1_0 <= geglu_max_value else cutlass.Float32(0.0) + x1_filter_1 = y1_1 if x1_1 <= geglu_max_value else cutlass.Float32(0.0) + + (dx1_vec[i], dx1_vec[i+1]) = fmul2((dy1_0, dy1_1), (cutlass.Float32(x1_filter_0), cutlass.Float32(x1_filter_1))) + + # dy2 = g * y1 * sigmoid_out * mProb + dy2_0, dy2_1 = fmul2( + y1, acc_mul_sigmoid_prob + ) + x2_filter_0 = x2_0 if x2_0 <= geglu_max_value else x_dtype(0.0) + x2_filter_1 = x2_1 if x2_1 <= geglu_max_value else x_dtype(0.0) + x2_filter_0 = y2_0 if x2_filter_0 >= geglu_min_value else cutlass.Float32(0.0) + x2_filter_1 = y2_1 if x2_filter_1 >= geglu_min_value else cutlass.Float32(0.0) + (dx2_vec[i], dx2_vec[i+1]) = fmul2((dy2_0, dy2_1), (cutlass.Float32(x2_filter_0), cutlass.Float32(x2_filter_1))) + + if cutlass.const_expr(self.generate_dprob): + (prob_grad, prob_grad_1) = fmul2( + (dy1_pre_0, dy1_pre_1), + y1, + ) + dprob_swiglu[i] = prob_grad + dprob_swiglu[i+1] = prob_grad_1 + dx1_vec = dx1_vec.load() + dx2_vec = dx2_vec.load() + if cutlass.const_expr(self.generate_dprob): + dprob_swiglu = dprob_swiglu.load() + return dx1_vec, dx2_vec, dprob_swiglu + else: + element_count = cute.size(x1_vec_load) + acc_vec = acc_vec.load() + x1_vec_load = x1_vec_load.load().to(cutlass.Float32) + x2_vec_load = x2_vec_load.load().to(cutlass.Float32) + dx1_vec = cute.make_rmem_tensor(acc_vec.shape, cutlass.Float32) + dx2_vec = cute.make_rmem_tensor(acc_vec.shape, cutlass.Float32) + + # y1 = clamp(x1, max=7.0); y2 = clamp(x2, min=-7.0, max=7.0) + for i in cutlass.range_constexpr(element_count): + fc2_dgrad = acc_vec[i] + g = fc2_dgrad + if cutlass.const_expr(self.has_prob): + g = fc2_dgrad * mProb + y1 = min(x1_vec_load[i], 7.0) + y2 = min(x2_vec_load[i], 7.0) + y2 = max(y2, -7.0) + + sigmoid_out = sigmoid_f32(y1 * 1.702, fastmath=True) + + dy1 = g * sigmoid_out * (1 + 1.702 * y1 * (1 - sigmoid_out)) * (y2 + linear_offset) + dy2 = g * y1 * sigmoid_out + + x1_filter = x1_vec_load[i] if x1_vec_load[i] <= 7.0 else 0.0 + x2_filter = x2_vec_load[i] if x2_vec_load[i] <= 7.0 else 0.0 + x2_filter = x2_filter if x2_filter >= -7.0 else 0.0 + + dx1_vec[i] = x1_filter * dy1 + dx2_vec[i] = x2_filter * dy2 + + if cutlass.const_expr(self.generate_dprob): + prob_grad = y1 * sigmoid_out * (y2 + linear_offset) * fc2_dgrad + dprob_swiglu[i] = prob_grad + + if cutlass.const_expr(self.generate_dprob): + dprob_swiglu = dprob_swiglu.load() + return dx1_vec.load(), dx2_vec.load(), dprob_swiglu + + + # GPU device kernel + @cute.kernel + def kernel( + self, + tiled_mma: cute.TiledMma, + tiled_mma_bkeep: Optional[cute.TiledMma], + tiled_mma_breuse: Optional[cute.TiledMma], + tiled_mma_sfb: cute.TiledMma, + tma_atom_a: cute.CopyAtom, + mA_mkl: cute.Tensor, + tma_atom_b: cute.CopyAtom, + mB_nkl: cute.Tensor, + tma_atom_sfa: cute.CopyAtom, + mSFA_mkl: cute.Tensor, + tma_atom_sfb: cute.CopyAtom, + mSFB_nkl: cute.Tensor, + tma_atom_c: cute.CopyAtom, + mC_mnl: cute.Tensor, + tma_atom_d: cute.CopyAtom, + mD_mnl: cute.Tensor, + tma_atom_d_col: Optional[cute.CopyAtom], + mD_col_mnl: Optional[cute.Tensor], + mSFDRow_mnl: Optional[cute.Tensor], + mSFDCol_mnl: Optional[cute.Tensor], + norm_const_tensor: Optional[cute.Tensor], + mAmax_tensor: Optional[cute.Tensor], + padded_offsets: cute.Tensor, + alpha: cute.Tensor, + beta: cute.Tensor, + prob: cute.Tensor, + dprob: cute.Tensor, + linear_offset: Float32, + mDbias_tensor: Optional[cute.Tensor], + workspace_ptr, + cluster_layout_vmnk: cute.Layout, + cluster_layout_sfb_vmnk: cute.Layout, + a_smem_layout_staged: cute.ComposedLayout, + b_smem_layout_staged: cute.ComposedLayout, + sfa_smem_layout_staged: cute.Layout, + sfb_smem_layout_staged: cute.Layout, + c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout, None], + d_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout, None], + epi_tile: cute.Tile, + sched_params: MoESchedulerParams, + epilogue_op: cutlass.Constexpr, + ): + """ + GPU device kernel performing the Persistent batched GEMM computation. + """ + tidx, _, _ = cute.arch.thread_idx() + warp_idx = tidx // 32 + warp_idx = cute.arch.make_warp_uniform(warp_idx) + + total_tokens = padded_offsets[self.expert_cnt - 1] + + # + # Prefetch tma desc + # + if warp_idx == self.tma_warp_id: + cpasync.prefetch_descriptor(tma_atom_a) + cpasync.prefetch_descriptor(tma_atom_sfa) + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DENSE): + cpasync.prefetch_descriptor(tma_atom_b) + cpasync.prefetch_descriptor(tma_atom_sfb) + cpasync.prefetch_descriptor(tma_atom_c) + if cutlass.const_expr(not self.store_d_directly): + cpasync.prefetch_descriptor(tma_atom_d) + if cutlass.const_expr(self.generate_sfd): + cpasync.prefetch_descriptor(tma_atom_d_col) + + use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2 + + # + # Setup cta/thread coordinates + # + bidx, bidy, bidz = cute.arch.block_idx() + mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape) + is_leader_cta = mma_tile_coord_v == 0 + cta_rank_in_cluster = cute.arch.make_warp_uniform( + cute.arch.block_idx_in_cluster() + ) + block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord( + cta_rank_in_cluster + ) + + block_in_cluster_coord_sfb_vmnk = cluster_layout_sfb_vmnk.get_flat_coord( + cta_rank_in_cluster + ) + + smem = utils.SmemAllocator() + storage = smem.allocate(self.shared_storage) + sched_storage = storage.scheduler + + # Initialize mainloop ab_pipeline (barrier) and states + ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1 + ab_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, num_tma_producer + ) + ab_pipeline = pipeline.PipelineTmaUmma.create( + barrier_storage=storage.ab_mbar_ptr.data_ptr(), + num_stages=self.num_ab_stage, + producer_group=ab_pipeline_producer_group, + consumer_group=ab_pipeline_consumer_group, + tx_count=self.num_tma_load_bytes, + cta_layout_vmnk=cluster_layout_vmnk, + defer_sync=True + ) + + # Initialize acc_pipeline (barrier) and states + acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + num_acc_consumer_threads = len(self.epilog_warp_id) * ( + 2 if use_2cta_instrs else 1 + ) + acc_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, num_acc_consumer_threads + ) + acc_pipeline = pipeline.PipelineUmmaAsync.create( + barrier_storage=storage.acc_mbar_ptr.data_ptr(), + num_stages=self.num_acc_stage, + producer_group=acc_pipeline_producer_group, + consumer_group=acc_pipeline_consumer_group, + cta_layout_vmnk=cluster_layout_vmnk, + defer_sync=True + ) + + # Load C pipeline + # Threads/warps participating in tma store pipeline + c_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + c_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + len(self.epilog_warp_id), + ) + c_pipeline = pipeline.PipelineTmaAsync.create( + barrier_storage=storage.c_full_mbar_ptr.data_ptr(), + num_stages=self.num_c_stage, + producer_group=c_producer_group, + consumer_group=c_consumer_group, + tx_count=self.tma_c_load_bytes, + defer_sync=True + ) + + # Initialize tile info pipeline (barrier) and states + tile_info_pipeline_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + self.threads_per_warp * 1, + ) + tile_info_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + self.threads_wo_sched, + ) + tile_info_pipeline = pipeline.PipelineAsync.create( + barrier_storage=sched_storage.tile_info_mbar.data_ptr(), + num_stages=self.num_tile_stage, + producer_group=tile_info_pipeline_producer_group, + consumer_group=tile_info_pipeline_consumer_group, + ) + + scheduler = MoEPersistentTileScheduler.create( + sched_params, + padded_offsets, + cute.arch.block_idx(), + cute.arch.grid_dim(), + counter_ptr=self._get_sched_counter_ptr(workspace_ptr), + sched_storage=sched_storage, + ) + scheduler.internal_init() + + # dBias SMEM setup + if cutlass.const_expr(self.generate_dbias): + sDbias = storage.sDbias.get_tensor(cute.make_layout( + (self.epi_tile[1] * 2, 32, len(self.epilog_warp_id)), + stride=(32, 1, self.epi_tile[1] * 2 * 32), + )) + + # Tensor memory dealloc barrier init + tmem = utils.TmemAllocator( + storage.tmem_holding_buf.ptr, + barrier_for_retrieve=self.tmem_alloc_barrier, + allocator_warp_id=self.epilog_warp_id[0], + is_two_cta=use_2cta_instrs, + two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr.ptr, + arch="sm_107", + ) + + # Cluster arrive after barrier init + if cute.size(self.cluster_shape_mn) > 1: + cute.arch.cluster_arrive_relaxed() + + # + # Setup smem tensor A/B/D/Scale + # + # (EPI_TILE_M, EPI_TILE_N, STAGE) + sC = storage.sC.get_tensor( + c_smem_layout_staged.outer, swizzle=c_smem_layout_staged.inner + ) + sD = None + if cutlass.const_expr(not self.store_d_directly): + sD = storage.sD.get_tensor( + d_smem_layout_staged.outer, swizzle=d_smem_layout_staged.inner + ) + sD_col = None + if cutlass.const_expr(self.generate_sfd): + sD_col = storage.sD_col.get_tensor( + d_smem_layout_staged.outer, swizzle=d_smem_layout_staged.inner + ) + # (MMA, MMA_M, MMA_K, STAGE) + sA = storage.sA.get_tensor( + a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner + ) + # (MMA, MMA_N, MMA_K, STAGE) + sB = storage.sB.get_tensor( + b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner + ) + # (granularity_m, repeat_m), (granularity_k, repeat_k), num_scale_stage) + sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged) + # (granularity_n, repeat_n), (granularity_k, repeat_k), num_scale_stage) + sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged) + # Shared memory for amax reduction (one FP32 per epilogue warp) + amax_layout = cute.make_layout((self.num_epilog_warps,)) + sAmax = storage.sAmax.get_tensor(amax_layout) + # (expert_idx, tile_m_idx, tile_n_idx, k_tile_cnt) + info_layout = cute.make_layout((4, self.num_tile_stage), stride=(1, 4)) + sInfo = sched_storage.sInfo.get_tensor(info_layout) + + # + # Compute multicast mask for A/B buffer full + # + a_full_mcast_mask = None + b_full_mcast_mask = None + sfa_full_mcast_mask = None + sfb_full_mcast_mask = None + if cutlass.const_expr(self.is_a_mcast or self.is_b_mcast or use_2cta_instrs): + a_full_mcast_mask = cpasync.create_tma_multicast_mask( + cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2 + ) + b_full_mcast_mask = cpasync.create_tma_multicast_mask( + cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1 + ) + sfa_full_mcast_mask = cpasync.create_tma_multicast_mask( + cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2 + ) + sfb_full_mcast_mask = cpasync.create_tma_multicast_mask( + cluster_layout_sfb_vmnk, block_in_cluster_coord_sfb_vmnk, mcast_mode=1 + ) + + # + # Partition shared/tensor memory tensor for TiledMMA_A/B/D + # + # (MMA, MMA_M, MMA_K, STAGE) + tCrA = tiled_mma.make_fragment_A(sA) + # (MMA, MMA_N, MMA_K, STAGE) + tCrB = tiled_mma.make_fragment_B(sB) + # (MMA, MMA_M, MMA_N) + acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2]) + # (MMA, MMA_M, MMA_N, STAGE) + if cutlass.const_expr(self.overlapping_accum): + num_acc_stage_overlapped = 2 + tCtAcc_fake = tiled_mma.make_fragment_C( + cute.append(acc_shape, num_acc_stage_overlapped) + ) + tCtAcc_fake = cute.make_tensor( + tCtAcc_fake.iterator, + cute.make_layout( + tCtAcc_fake.shape, + stride=( + tCtAcc_fake.stride[0], + tCtAcc_fake.stride[1], + tCtAcc_fake.stride[2], + (256 - self.num_sf_tmem_cols) * tCtAcc_fake.stride[0][1], + ), + ), + ) + elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192): + tCtAcc_fake = tiled_mma.make_fragment_C( + cute.append(acc_shape, self.num_acc_stage) + ) + tCtAcc_fake = cute.make_tensor( + tCtAcc_fake.iterator, + cute.make_layout( + tCtAcc_fake.shape, + stride=( + tCtAcc_fake.stride[0], + tCtAcc_fake.stride[1], + tCtAcc_fake.stride[2], + self.num_accumulator_tmem_stride, + ), + ), + ) + else: + tCtAcc_fake = tiled_mma.make_fragment_C( + cute.append(acc_shape, self.num_acc_stage) + ) + + # + # Cluster wait before tensor memory alloc + # + if cute.size(self.cluster_shape_mn) > 1: + cute.arch.cluster_wait() + else: + self.cta_sync_barrier.arrive_and_wait() + + if total_tokens <= 0: + cute.arch.nvvm.exit() + k_tile_cnt = cute.ceil_div(cute.size(mB_nkl, mode=[1]), self.mma_tiler[2]) + + # + # Specialized Schedule warp (MoE Persistent Tile Scheduler) + # + if warp_idx == self.sched_warp_id: + work_tile_info = scheduler.initial_work_tile_info() + + tile_info_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_tile_stage + ) + + while work_tile_info.is_valid_tile: + # sInfo format: (expert_idx, tile_m_idx, tile_n_idx, k_tile_cnt) + tile_info_pipeline.producer_acquire(tile_info_producer_state) + with cute.arch.elect_one(): + sInfo[(0, tile_info_producer_state.index)] = work_tile_info.expert_idx + sInfo[(1, tile_info_producer_state.index)] = work_tile_info.tile_m_idx + sInfo[(2, tile_info_producer_state.index)] = work_tile_info.tile_n_idx + sInfo[(3, tile_info_producer_state.index)] = work_tile_info.k_tile_cnt + cute.arch.fence_proxy("async.shared", space="cta") + + self.sched_sync_barrier.arrive_and_wait() + tile_info_pipeline.producer_commit(tile_info_producer_state) + tile_info_producer_state.advance() + + work_tile_info = scheduler.advance_to_next_work() + + # Send invalid tile signal: expert_idx = -1 + tile_info_pipeline.producer_acquire(tile_info_producer_state) + with cute.arch.elect_one(): + sInfo[(0, tile_info_producer_state.index)] = cutlass.Int32(-1) + sInfo[(1, tile_info_producer_state.index)] = cutlass.Int32(0) + sInfo[(2, tile_info_producer_state.index)] = cutlass.Int32(0) + sInfo[(3, tile_info_producer_state.index)] = cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + self.sched_sync_barrier.arrive_and_wait() + tile_info_pipeline.producer_commit(tile_info_producer_state) + tile_info_producer_state.advance() + tile_info_pipeline.producer_tail(tile_info_producer_state) + + # + # Specialized TMA load warp + # + if warp_idx == self.tma_warp_id: + ext = self._make_extension(workspace_ptr) + + ab_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_ab_stage + ) + + tile_info_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_tile_stage + ) + + # Get the first tile info + tile_info = cute.make_rmem_tensor((4,), cutlass.Int32) + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + while is_valid_tile: + work_tile_info = MoEWorkTileInfo( + expert_idx=tile_info[0], + tile_m_idx=tile_info[1], + tile_n_idx=tile_info[2], + k_tile_cnt=tile_info[3], + ) + # assert(k_tile_cnt == work_tile_info.k_tile_cnt) + ext.update_expert_info(padded_offsets, work_tile_info.expert_idx) + + # Get per-expert real tensors + TMA desc ptrs via extension + real_a, _ = ext.get_gmem_tensor( + "a", mA_mkl, padded_offsets, work_tile_info + ) + real_b, desc_ptr_b = ext.get_gmem_tensor( + "b", mB_nkl, padded_offsets, work_tile_info + ) + real_sfa, _ = ext.get_gmem_tensor( + "sfa", mSFA_mkl, padded_offsets, work_tile_info + ) + real_sfb, desc_ptr_sfb = ext.get_gmem_tensor( + "sfb", mSFB_nkl, padded_offsets, work_tile_info + ) + + if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192): + x = real_sfb.stride[0][1] + y = cute.ceil_div(real_sfb.shape[0][1], 4) + new_shape = ( + (real_sfb.shape[0][0], ((2, 2), y)), + real_sfb.shape[1], real_sfb.shape[2], + ) + x_times_3 = 3 * x + new_stride = ( + (real_sfb.stride[0][0], ((x, x), x_times_3)), + real_sfb.stride[1], real_sfb.stride[2], + ) + real_sfb = cute.make_tensor( + real_sfb.iterator, + cute.make_layout(new_shape, stride=new_stride), + ) + + # local_tile on per-expert tensors + gA_mkl = cute.local_tile( + real_a, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None) + ) + gB_nkl = cute.local_tile( + real_b, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None) + ) + gSFA_mkl = cute.local_tile( + real_sfa, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None) + ) + gSFB_nkl = cute.local_tile( + real_sfb, cute.slice_(self.mma_tiler_sfb, (0, None, None)), (None, None, None) + ) + + # MMA partition + thr_mma = tiled_mma.get_slice(mma_tile_coord_v) + thr_mma_sfb = tiled_mma_sfb.get_slice(mma_tile_coord_v) + tCgA = thr_mma.partition_A(gA_mkl) + tCgB = thr_mma.partition_B(gB_nkl) + tCgSFA = thr_mma.partition_A(gSFA_mkl) + tCgSFB = thr_mma_sfb.partition_B(gSFB_nkl) + + # TMA partition A + a_cta_layout = cute.make_layout( + cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape + ) + tAsA, tAgA = cpasync.tma_partition( + tma_atom_a, + block_in_cluster_coord_vmnk[2], + a_cta_layout, + cute.group_modes(sA, 0, 3), + cute.group_modes(tCgA, 0, 3), + ) + # TMA partition B + b_cta_layout = cute.make_layout( + cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape + ) + tBsB, tBgB = cpasync.tma_partition( + tma_atom_b, + block_in_cluster_coord_vmnk[1], + b_cta_layout, + cute.group_modes(sB, 0, 3), + cute.group_modes(tCgB, 0, 3), + ) + # TMA partition SFA + sfa_cta_layout = a_cta_layout + tAsSFA, tAgSFA = cpasync.tma_partition( + tma_atom_sfa, + block_in_cluster_coord_vmnk[2], + sfa_cta_layout, + cute.group_modes(sSFA, 0, 3), + cute.group_modes(tCgSFA, 0, 3), + ) + tAsSFA = cute.filter_zeros(tAsSFA) + tAgSFA = cute.filter_zeros(tAgSFA) + # TMA partition SFB + sfb_cta_layout = cute.make_layout( + cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape + ) + tBsSFB, tBgSFB = cpasync.tma_partition( + tma_atom_sfb, + block_in_cluster_coord_sfb_vmnk[1], + sfb_cta_layout, + cute.group_modes(sSFB, 0, 3), + cute.group_modes(tCgSFB, 0, 3), + ) + tBsSFB = cute.filter_zeros(tBsSFB) + tBgSFB = cute.filter_zeros(tBgSFB) + + # Slice to per mma tile index (L=0 since domain already offset'd) + mma_tile_coord_m = work_tile_info.tile_m_idx // cute.size(tiled_mma.thr_id.shape) + mma_tile_coord_n = work_tile_info.tile_n_idx + tAgA_slice = tAgA[(None, mma_tile_coord_m, None, 0)] + tBgB_slice = tBgB[(None, mma_tile_coord_n, None, 0)] + tAgSFA_slice = tAgSFA[(None, mma_tile_coord_m, None, 0)] + slice_n = mma_tile_coord_n + if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64): + slice_n = mma_tile_coord_n // 2 + tBgSFB_slice = tBgSFB[(None, slice_n, None, 0)] + + # Peek (try_wait) AB buffer empty + peek_ab_empty_status = cutlass.Boolean(1) + if k_tile_cnt > 0: + peek_ab_empty_status = ab_pipeline.producer_try_acquire( + ab_producer_state + ) + + # + # Tma load loop + # + for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1): + tAgA_k = tAgA_slice[(None, k_tile)] + tBgB_k = tBgB_slice[(None, k_tile)] + tAgSFA_k = tAgSFA_slice[(None, k_tile)] + tBgSFB_k = tBgSFB_slice[(None, k_tile)] + tAsA_pipe = tAsA[(None, ab_producer_state.index)] + tBsB_pipe = tBsB[(None, ab_producer_state.index)] + tAsSFA_pipe = tAsSFA[(None, ab_producer_state.index)] + tBsSFB_pipe = tBsSFB[(None, ab_producer_state.index)] + + tma_bar = ab_pipeline.producer_get_barrier(ab_producer_state) + + # Conditionally wait for AB buffer empty + ab_pipeline.producer_acquire( + ab_producer_state, peek_ab_empty_status + ) + ab_producer_state_next = ab_producer_state.clone() + ab_producer_state_next.advance() + if k_tile < k_tile_cnt - 1: + peek_ab_empty_status = ab_pipeline.producer_try_acquire( + ab_producer_state_next + ) + + # TMA load A (contiguous, global desc) + cute.copy( + tma_atom_a, + tAgA_k, + tAsA_pipe, + tma_bar_ptr=tma_bar, + mcast_mask=a_full_mcast_mask, + ) + # TMA load B (discrete, per-expert desc from workspace) + cute.copy( + tma_atom_b, + tBgB_k, + tBsB_pipe, + tma_bar_ptr=tma_bar, + mcast_mask=b_full_mcast_mask, + tma_desc_ptr=desc_ptr_b, + ) + # TMA load SFA (contiguous, global desc) + cute.copy( + tma_atom_sfa, + tAgSFA_k, + tAsSFA_pipe, + tma_bar_ptr=tma_bar, + mcast_mask=sfa_full_mcast_mask, + ) + # TMA load SFB (discrete, per-expert desc from workspace) + cute.copy( + tma_atom_sfb, + tBgSFB_k, + tBsSFB_pipe, + tma_bar_ptr=tma_bar, + mcast_mask=sfb_full_mcast_mask, + tma_desc_ptr=desc_ptr_sfb, + ) + + # Peek (try_wait) AB buffer empty for next k_tile + ab_producer_state = ab_producer_state_next + + # + # Advance to next tile + # + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + # + # Wait A/B buffer empty + # + ab_pipeline.producer_tail(ab_producer_state) + + # + # Specialized MMA warp + # + if warp_idx == self.mma_warp_id: + # + # Bar sync for retrieve tensor memory ptr from shared mem + # + tmem.wait_for_alloc() + + # + # Retrieving tensor memory ptr and make accumulator tensor + # + acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype) + # (MMA, MMA_M, MMA_N, STAGE) + tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout) + + # Make SFA tmem tensor + sfa_tmem_ptr = cute.recast_ptr( + acc_tmem_ptr + self.num_accumulator_tmem_cols, + dtype=self.sf_dtype, + ) + # (MMA, MMA_M, MMA_K) + tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa( + tiled_mma, + self.mma_tiler, + self.sf_vec_size, + cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)), + ) + tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout) + + # Make SFB tmem tensor + sfb_tmem_ptr = cute.recast_ptr( + acc_tmem_ptr + self.num_accumulator_tmem_cols + self.num_sfa_tmem_cols, + dtype=self.sf_dtype, + ) + # (MMA, MMA_N, MMA_K) + tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb( + tiled_mma, + self.mma_tiler, + self.sf_vec_size, + cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)), + ) + tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout) + + # Partition for S2T copy of SFA/SFB + # + ( + tiled_copy_s2t_sfa, + tCsSFA_compact_s2t, + tCtSFA_compact_s2t, + ) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA) + ( + tiled_copy_s2t_sfb, + tCsSFB_compact_s2t, + tCtSFB_compact_s2t, + ) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB) + + ab_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_ab_stage + ) + acc_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_acc_stage + ) + + tile_info_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_tile_stage + ) + + # Get the first tile info (sInfo format: expert_idx, tile_m_idx, tile_n_idx, k_tile_cnt) + tile_info = cute.make_rmem_tensor((4,), cutlass.Int32) + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + while is_valid_tile: + # Peek (try_wait) AB buffer full for k_tile = 0 + peek_ab_full_status = cutlass.Boolean(1) + if k_tile_cnt > 0 and is_leader_cta: + peek_ab_full_status = ab_pipeline.consumer_try_wait( + ab_consumer_state + ) + + # sInfo: (expert_idx, tile_m_idx, tile_n_idx, k_tile_cnt) + mma_tile_coord_mnl = ( + tile_info[1] // cute.size(tiled_mma.thr_id.shape), + tile_info[2], + cutlass.Int32(0), + ) + + # Get accumulator stage index + if cutlass.const_expr(self.overlapping_accum): + acc_stage_index = acc_producer_state.phase ^ 1 + else: + acc_stage_index = acc_producer_state.index + + tCtAcc = tCtAcc_base[(None, None, None, acc_stage_index)] + tCtSFB_mma = tCtSFB + if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192): + # If an ODD tile, shift the TMEM start address for cta_tile_shape_n=192 case by two words + offset = ( + cutlass.Int32(2) + if mma_tile_coord_mnl[1] % 2 == 1 + else cutlass.Int32(0) + ) + shifted_ptr = cute.recast_ptr( + acc_tmem_ptr + + self.num_accumulator_tmem_cols + + self.num_sfa_tmem_cols + + offset, + dtype=self.sf_dtype, + ) + tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout) + elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64): + # Move in increments of 64 columns of SFB + offset = cutlass.Int32((mma_tile_coord_mnl[1] % 2) * 2) + shifted_ptr = cute.recast_ptr( + acc_tmem_ptr + + self.num_accumulator_tmem_cols + + self.num_sfa_tmem_cols + + offset, + dtype=self.sf_dtype, + ) + tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout) + # + # Wait for accumulator buffer empty + # + if is_leader_cta: + acc_pipeline.producer_acquire(acc_producer_state) + if cutlass.const_expr( + self.enable_breuse + and cute.size(tCtAcc.layout, mode=[1]) == 2 + and cute.size(tCtAcc.layout, mode=[2]) == 1 + ): + tCtAcc_bkeep = tCtAcc[(None, 0, 0)] + tCtAcc_breuse = tCtAcc[(None, 1, 0)] + tiled_mma_bkeep.set(tcgen05.Field.ACCUMULATE, False) + tiled_mma_breuse.set(tcgen05.Field.ACCUMULATE, False) + else: + tiled_mma.set(tcgen05.Field.ACCUMULATE, False) + for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1): + # Conditionally wait for AB buffer full + ab_pipeline.consumer_wait( + ab_consumer_state, peek_ab_full_status + ) + ab_consumer_state_next = ab_consumer_state.clone() + ab_consumer_state_next.advance() + if k_tile < k_tile_cnt - 1: + peek_ab_full_status = ab_pipeline.consumer_try_wait( + ab_consumer_state_next + ) + + # Copy SFA/SFB from smem to tmem + s2t_stage_coord = ( + None, + None, + None, + None, + ab_consumer_state.index, + ) + tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[s2t_stage_coord] + tCsSFB_compact_s2t_staged = tCsSFB_compact_s2t[s2t_stage_coord] + cute.copy( + tiled_copy_s2t_sfa, + tCsSFA_compact_s2t_staged, + tCtSFA_compact_s2t, + ) + cute.copy( + tiled_copy_s2t_sfb, + tCsSFB_compact_s2t_staged, + tCtSFB_compact_s2t, + ) + + # tCtAcc += tCrA * tCrSFA * tCrB * tCrSFB + num_kblocks = cute.size(tCrA, mode=[2]) + + for kblock_idx in cutlass.range(num_kblocks, unroll_full=True): + # Set SFA/SFB tensor to tiled_mma + sf_kblock_coord = (None, None, kblock_idx) + + if cutlass.const_expr( + self.enable_breuse + and cute.size(tCtAcc.layout, mode=[1]) == 2 + and cute.size(tCtAcc.layout, mode=[2]) == 1 + ): + a_kblk_crd_keep = (None, 0, kblock_idx, ab_consumer_state.index) + a_kblk_crd_reuse = (None, 1, kblock_idx, ab_consumer_state.index) + b_kblk_crd = (None, 0, kblock_idx, ab_consumer_state.index) + sfa_kblk_crd_keep = (None, 0, kblock_idx) + sfa_kblk_crd_reuse = (None, 1, kblock_idx) + sfb_kblk_crd = (None, 0, kblock_idx) + # Bkeep: accumulate A_upper × B + tiled_mma_bkeep.set( + tcgen05.Field.ACCUMULATE, k_tile != 0 or kblock_idx != 0 + ) + cute.gemm( + tiled_mma_bkeep, + tCtAcc_bkeep, + [tCrA[a_kblk_crd_keep], tCtSFA[sfa_kblk_crd_keep]], + [tCrB[b_kblk_crd], tCtSFB_mma[sfb_kblk_crd]], + tCtAcc_bkeep, + ) + # Breuse: accumulate A_lower × B (reuses B) + tiled_mma_breuse.set( + tcgen05.Field.ACCUMULATE, k_tile != 0 or kblock_idx != 0 + ) + cute.gemm( + tiled_mma_breuse, + tCtAcc_breuse, + [tCrA[a_kblk_crd_reuse], tCtSFA[sfa_kblk_crd_reuse]], + [tCrB[b_kblk_crd], tCtSFB_mma[sfb_kblk_crd]], + tCtAcc_breuse, + ) + else: + kblock_coord = ( + None, + None, + kblock_idx, + ab_consumer_state.index, + ) + tiled_mma.set( + tcgen05.Field.SFA, + tCtSFA[sf_kblock_coord].iterator, + ) + tiled_mma.set( + tcgen05.Field.SFB, + tCtSFB_mma[sf_kblock_coord].iterator, + ) + + cute.gemm( + tiled_mma, + tCtAcc, + tCrA[kblock_coord], + tCrB[kblock_coord], + tCtAcc, + ) + # Enable accumulate on tCtAcc after first kblock + tiled_mma.set(tcgen05.Field.ACCUMULATE, True) + + # Async arrive AB buffer empty + ab_pipeline.consumer_release(ab_consumer_state) + ab_consumer_state = ab_consumer_state_next + acc_pipeline.producer_commit(acc_producer_state) + + # Peek (try_wait) Acc buffer empty for k_tile = k_tile + 1 + acc_producer_state.advance() + + # + # Advance to next tile + # + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + # + # Wait for accumulator buffer empty + # + acc_pipeline.producer_tail(acc_producer_state) + + # + # Specialized epilogue warps + # + if warp_idx < self.mma_warp_id: + # + # Alloc tensor memory buffer + # + tmem.allocate(self.num_tmem_alloc_cols) + + # + # Bar sync for retrieve tensor memory ptr from shared memory + # + tmem.wait_for_alloc() + + # + # Retrieving tensor memory ptr and make accumulator tensor + # + tmem_ptr = tmem.retrieve_ptr(self.acc_dtype) + # (MMA, MMA_M, MMA_N, STAGE) + tCtAcc_base = cute.make_tensor(tmem_ptr, tCtAcc_fake.layout) + + # + # Partition for epilogue (SMEM/TMEM/register - invariant across experts) + # + epi_tidx = tidx + + # Shape-only partition on global tensor (invariant setup for t2r copy atom) + thr_mma_epi = tiled_mma.get_slice(mma_tile_coord_v) + gD_mnl_shape = cute.local_tile( + mD_mnl, cute.slice_(self.mma_tiler_d, (None, None, 0)), (None, None, None) + ) + tCgD_shape = thr_mma_epi.partition_C(gD_mnl_shape) + + if cutlass.const_expr(self.enable_breuse): + # Merge bkeep/breuse M-split (mode1=2) into the M dimension for epilogue + tCtAcc_epi_input = transform_partitioned_tensor_layout(tCtAcc_base) + tCgD_epi_input = transform_partitioned_tensor_layout(tCgD_shape) + else: + tCtAcc_epi_input = tCtAcc_base + tCgD_epi_input = tCgD_shape + + ( + tiled_copy_t2r, + tTR_tAcc_base, + tTR_rAcc, + ) = self.epilog_tmem_copy_and_partition( + epi_tidx, tCtAcc_epi_input, tCgD_epi_input, epi_tile, use_2cta_instrs + ) + + tTR_rC1 = cute.make_rmem_tensor(tTR_rAcc.shape, self.c_dtype) + tTR_rC2 = cute.make_rmem_tensor(tTR_rAcc.shape, self.c_dtype) + tiled_copy_s2r, tRS_rC1, tRS_rC2, tRS_sC = ( + self.epilog_smem_copy_and_partition_load( + tiled_copy_t2r, tTR_rC1, tTR_rC2, epi_tidx, sC + ) + ) + + tTR_rD1 = cute.make_rmem_tensor(tTR_rAcc.shape, self.d_dtype) + tTR_rD2 = cute.make_rmem_tensor(tTR_rAcc.shape, self.d_dtype) + tiled_copy_r2s, tRS_rD1, tRS_rD2, tRS_sD = ( + self.epilog_smem_copy_and_partition_store( + tiled_copy_t2r, tTR_rD1, tTR_rD2, epi_tidx, sD + ) + ) + if cutlass.const_expr(self.generate_sfd): + tTR_rD1_col = cute.make_rmem_tensor(tTR_rAcc.shape, self.d_dtype) + tTR_rD2_col = cute.make_rmem_tensor(tTR_rAcc.shape, self.d_dtype) + ( + tiled_copy_r2s, + tRS_rD1_col, + tRS_rD2_col, + tRS_sD_col, + ) = self.epilog_smem_copy_and_partition_store( + tiled_copy_t2r, tTR_rD1_col, tTR_rD2_col, epi_tidx, sD_col + ) + + norm_const = cutlass.Float32(norm_const_tensor[0]) + d_rcp_limits = get_dtype_rcp_limits(self.d_dtype) + + # Extension for per-expert domain conversion in epilogue + epi_ext = self._make_extension(workspace_ptr) + + acc_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_acc_stage + ) + + # Load C pipeline + c_pipeline_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_c_stage + ) + + # Threads/warps participating in tma store pipeline + d_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + 32 * len(self.epilog_warp_id), + ) + d_pipeline = None + if cutlass.const_expr(not self.store_d_directly): + num_d_stages = self.num_d_stage // 2 + d_pipeline = pipeline.PipelineTmaStore.create( + num_stages=num_d_stages, + producer_group=d_producer_group, + ) + + tile_info_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_tile_stage + ) + + # Get the first tile info + tile_info = cute.make_rmem_tensor((4,), cutlass.Int32) + + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + num_prev_subtiles = cutlass.Int32(0) + while is_valid_tile: + # sInfo: (expert_idx, tile_m_idx, tile_n_idx, k_tile_cnt) + epi_work_tile_info = MoEWorkTileInfo( + expert_idx=tile_info[0], + tile_m_idx=tile_info[1], + tile_n_idx=tile_info[2], + k_tile_cnt=tile_info[3], + ) + expert_idx = epi_work_tile_info.expert_idx + # N is doubled for dGLU dual output + mma_tile_coord_mnl = ( + epi_work_tile_info.tile_m_idx // cute.size(tiled_mma.thr_id.shape), + epi_work_tile_info.tile_n_idx * 2, + cutlass.Int32(0), + ) + + # + # Get alpha/beta for current expert + # + alpha_val = alpha[expert_idx] + beta_val = beta[expert_idx] + epi_ext.update_expert_info(padded_offsets, expert_idx) + + # + # Per-expert gmem tensor setup via extension + # + real_d, _ = epi_ext.get_gmem_tensor( + "d", mD_mnl, padded_offsets, epi_work_tile_info + ) + gD_mnl_loop = cute.local_tile( + real_d, cute.slice_(self.mma_tiler_d, (None, None, 0)), (None, None, None) + ) + thr_mma_epi = tiled_mma.get_slice(mma_tile_coord_v) + tCgD_loop = thr_mma_epi.partition_C(gD_mnl_loop) + + if cutlass.const_expr(not self.store_d_directly): + if cutlass.const_expr(self.enable_breuse): + tCgD_loop_epi = transform_partitioned_tensor_layout(tCgD_loop) + gD_epi_tr = cute.flat_divide(tCgD_loop_epi, epi_tile) + sD_for_tma = cute.group_modes(sD, 0, 2) + gD_for_tma_tr = cute.group_modes(gD_epi_tr, 0, 2) + bSG_sD, bSG_gD_partitioned = cpasync.tma_partition( + tma_atom_d, 0, cute.make_layout(1), sD_for_tma, gD_for_tma_tr, + ) + else: + bSG_sD, bSG_gD_partitioned = self.epilog_gmem_copy_and_partition( + epi_tidx, tma_atom_d, tCgD_loop, epi_tile, sD + ) + bSG_gD = bSG_gD_partitioned[ + (None, None, None, mma_tile_coord_mnl[0], mma_tile_coord_mnl[1], 0) + ] + bSG_gD = cute.group_modes(bSG_gD, 1, cute.rank(bSG_gD)) + + if cutlass.const_expr(self.generate_sfd): + real_d_col, _ = epi_ext.get_gmem_tensor( + "d_col", mD_col_mnl, padded_offsets, epi_work_tile_info + ) + gD_col_mnl_loop = cute.local_tile( + real_d_col, cute.slice_(self.mma_tiler_d, (None, None, 0)), (None, None, None) + ) + tCgD_col_loop = thr_mma_epi.partition_C(gD_col_mnl_loop) + if cutlass.const_expr(self.enable_breuse): + tCgD_col_loop_epi = transform_partitioned_tensor_layout(tCgD_col_loop) + gD_col_epi_tr = cute.flat_divide(tCgD_col_loop_epi, epi_tile) + sD_col_for_tma = cute.group_modes(sD_col, 0, 2) + gD_col_for_tma_tr = cute.group_modes(gD_col_epi_tr, 0, 2) + bSG_sD_col, bSG_gD_col_partitioned = cpasync.tma_partition( + tma_atom_d_col, 0, cute.make_layout(1), sD_col_for_tma, gD_col_for_tma_tr, + ) + else: + bSG_sD_col, bSG_gD_col_partitioned = self.epilog_gmem_copy_and_partition( + epi_tidx, tma_atom_d_col, tCgD_col_loop, epi_tile, sD_col + ) + bSG_gD_col = bSG_gD_col_partitioned[ + (None, None, None, mma_tile_coord_mnl[0], mma_tile_coord_mnl[1], 0) + ] + bSG_gD_col = cute.group_modes(bSG_gD_col, 1, cute.rank(bSG_gD_col)) + + # Get accumulator stage index + if cutlass.const_expr(self.overlapping_accum): + acc_stage_index = acc_consumer_state.phase + reverse_subtile = ( + cutlass.Boolean(True) + if acc_stage_index == 0 + else cutlass.Boolean(False) + ) + else: + acc_stage_index = acc_consumer_state.index + + # Set tensor memory buffer for current tile + # (T2R, T2R_M, T2R_N, EPI_M, EPI_M) + tTR_tAcc = tTR_tAcc_base[ + (None, None, None, None, None, acc_stage_index) + ] + + if cutlass.const_expr(self.generate_sfd): + regPerSubtile = 4 + sfd_row_tile = ( + cute.make_layout(128), + cute.make_layout(32 * regPerSubtile), + ) + # SFD Row: tile_atom_to_shape_SF layout, same path as SFA + real_sfd_row, _ = epi_ext.get_gmem_tensor( + "sfd", mSFDRow_mnl, padded_offsets, epi_work_tile_info + ) + gSFDRow_mnl = cute.local_tile( + real_sfd_row, sfd_row_tile, (None, None, None) + ) + # Don't ask why, AST is shit tracking the constexpr values to loop args. + tiled_copy_t2r_local, _, _ = self.epilog_tmem_copy_and_partition(epi_tidx, tCtAcc_epi_input, tCgD_epi_input, epi_tile, use_2cta_instrs) + thr_copy_t2r_local = tiled_copy_t2r_local.get_slice(tidx) + tCgSFDRow_mnl = thr_copy_t2r_local.partition_D(gSFDRow_mnl) + tCgSFDRow_mnl = cute.filter_zeros(tCgSFDRow_mnl) + tCrSFDRow = cute.make_rmem_tensor( + tCgSFDRow_mnl[(None, None, None, 0, 0, 0)].layout, self.sf_dtype + ) + tCrSFDRow_pvscale = cute.make_rmem_tensor_like( + tCrSFDRow, cutlass.Float32 + ) + tCgSFDRow_mn = tCgSFDRow_mnl[ + (None, None, None, None, None, 0) + ] + + # SFD Col: layout depends on discrete_col_sfd + if cutlass.const_expr(self.discrete_col_sfd): + # discrete_col_sfd uses tile_atom_to_shape_SF layout + real_sfd_col, _ = epi_ext.get_gmem_tensor( + "sfd", mSFDCol_mnl, padded_offsets, epi_work_tile_info + ) + else: + # non-discrete uses BlockScaledBasicChunk layout + real_sfd_col, _ = epi_ext.get_gmem_tensor( + "sfd_col", mSFDCol_mnl, padded_offsets, epi_work_tile_info + ) + gSFDCol_mnl = cute.local_tile( + real_sfd_col, sfd_row_tile, (None, None, None) + ) + thr_layout = cute.make_ordered_layout((4, 32), order=(1, 0)) + val_layout = cute.make_ordered_layout((1,), order=(0,)) + copy_atom_sfd_col_quant = cute.make_copy_atom( + cute.nvgpu.CopyUniversalOp(), + gSFDCol_mnl.element_type, + num_bits_per_copy=8, + ) + tiled_copy_sfd_col_quant = cute.make_tiled_copy_tv( + copy_atom_sfd_col_quant, thr_layout, val_layout + ) + thr_copy_sfd_col_quant = tiled_copy_sfd_col_quant.get_slice(tidx) + tCgSFDCol_mnl = thr_copy_sfd_col_quant.partition_D( + cute.filter_zeros(gSFDCol_mnl) + ) + tCgSFDCol_mnl = cute.filter_zeros(tCgSFDCol_mnl) + if cutlass.const_expr(self.discrete_col_sfd): + tCgSFDCol_mnl = self.create_and_partition_new_SFDCol( + real_sfd_col + ) + tCrSFDCol = cute.make_rmem_tensor( + tCrSFDRow.layout, tCrSFDRow.element_type + ) + tCrSFDCol_pvscale = cute.make_rmem_tensor_like( + tCrSFDRow_pvscale.layout, cutlass.Float32 + ) + tCgSFDCol_mn = tCgSFDCol_mnl[ + (None, None, None, None, None, 0) + ] + + if cutlass.const_expr(self.generate_amax): + thread_tile_amax = cutlass.Float32(0.0) + + # + # Get PROB (per-expert local M position) + # + mPosition = ( + (epi_work_tile_info.tile_m_idx // cute.size(tiled_mma.thr_id.shape)) + * self.mma_tiler[0] + + mma_tile_coord_v * (self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape)) + + tidx + ) + mProb = cutlass.Float32(1.0) + mProb_bk = cutlass.Float32(1.0) + mProb_br = cutlass.Float32(1.0) + if cutlass.const_expr(self.has_prob): + real_prob, _ = epi_ext.get_gmem_tensor( + "prob", prob, padded_offsets, epi_work_tile_info + ) + mProb = real_prob[mPosition, 0, 0] + if cutlass.const_expr(self.enable_breuse): + # Two M halves: bkeep (rows 0..cta_m/2-1) and breuse (rows cta_m/2..cta_m-1) + mPosition_bk = mPosition + mPosition_br = mPosition + (self.cta_tile_shape_mnk[0] // 2) + if cutlass.const_expr(self.has_prob): + mProb_bk = real_prob[mPosition_bk, 0, 0] + mProb_br = real_prob[mPosition_br, 0, 0] + mProb = mProb_bk + if cutlass.const_expr(self.generate_dprob): + dProbVal = cutlass.Float32(0.0) + if cutlass.const_expr(self.enable_breuse): + dProbVal_bk = cutlass.Float32(0.0) + dProbVal_br = cutlass.Float32(0.0) + + # + # Wait for accumulator buffer full + # + acc_pipeline.consumer_wait(acc_consumer_state) + tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc)) + + # Initialize thread-local amax accumulator for this tile + if cutlass.const_expr(self.generate_amax): + thread_tile_amax_1 = cutlass.Float32(0.0) + thread_tile_amax_2 = cutlass.Float32(0.0) + + # + # Store accumulator to global memory in subtiles + # + copy_atom_t2r = sm100_utils.get_tmem_load_op( + self.cta_tile_shape_mnk, + self.d_layout, + self.d_dtype, + self.acc_dtype, + epi_tile, + use_2cta_instrs, + ) + subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3]) + for subtile_idx in cutlass.range(0, subtile_cnt, 1, unroll=1): + real_subtile_idx = subtile_idx + if cutlass.const_expr(self.overlapping_accum): + if reverse_subtile: + real_subtile_idx = ( + self.cta_tile_shape_mnk[1] // self.epi_tile_n_required + - 1 + - subtile_idx + ) + + tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)] + cute.copy(copy_atom_t2r, tTR_tAcc_mn, tTR_rAcc) + + # Update mProb based on which M half this subtile belongs to. + if cutlass.const_expr(self.enable_breuse): + if real_subtile_idx % 2 == 0: + mProb = mProb_bk + else: + mProb = mProb_br + + # + # Async arrive accumulator buffer empty ealier when overlapping_accum is enabled + # + if cutlass.const_expr(self.overlapping_accum): + bReleaseAcc = subtile_idx == ( + self.iter_acc_early_release_in_epilogue + ) + if bReleaseAcc: + # Fence for TMEM load + cute.arch.fence_view_async_tmem_load() + with cute.arch.elect_one(): + acc_pipeline.consumer_release(acc_consumer_state) + acc_consumer_state.advance() + + # Wait for C1/C2 load to complete + c_pipeline.consumer_wait(c_pipeline_consumer_state) + cute.copy( + tiled_copy_s2r, + tRS_sC[(None, None, None, c_pipeline_consumer_state.index)], + tRS_rC1, + ) + cute.arch.fence_proxy("async.shared", space="cta") + c_pipeline.consumer_release(c_pipeline_consumer_state) + c_pipeline_consumer_state.advance() + c_pipeline.consumer_wait(c_pipeline_consumer_state) + cute.copy( + tiled_copy_s2r, + tRS_sC[(None, None, None, c_pipeline_consumer_state.index)], + tRS_rC2, + ) + cute.arch.fence_proxy("async.shared", space="cta") + c_pipeline.consumer_release(c_pipeline_consumer_state) + c_pipeline_consumer_state.advance() + + acc_vec = tiled_copy_r2s.retile(tTR_rAcc) + ab1_vec_load = tiled_copy_r2s.retile(tRS_rC1) + ab2_vec_load = tiled_copy_r2s.retile(tRS_rC2) + if cutlass.const_expr(self.generate_dprob): + dprob_swiglu = cute.make_rmem_tensor( + acc_vec.shape, cutlass.Float32 + ) + else: + dprob_swiglu = None + + # + # Apply alpha, act, and prob + # + square_alpha = alpha_val * alpha_val + if cutlass.const_expr(self.act_func == "dswiglu"): + d1_vec, d2_vec, dprob_swiglu = self.dswiglu(acc_vec, ab1_vec_load, ab2_vec_load, mProb, beta_val, square_alpha, dprob_swiglu) + elif cutlass.const_expr(self.act_func == "dgeglu"): + d1_vec, d2_vec, dprob_swiglu = self.dgeglu(acc_vec, ab1_vec_load, ab2_vec_load, mProb, linear_offset, dprob_swiglu) + + if cutlass.const_expr(self.generate_dprob): + # dprob sum reduction + if cutlass.const_expr(self.vectorized_f32): + dprob_pair_0 = cutlass.Float32(0.0) + dprob_pair_1 = cutlass.Float32(0.0) + for j in cutlass.range( + 0, cute.size(dprob_swiglu.shape), 2, unroll_full=True + ): + ( + dprob_pair_0, + dprob_pair_1, + ) = cute.arch.add_packed_f32x2( + (dprob_pair_0, dprob_pair_1), + (dprob_swiglu[j], dprob_swiglu[j + 1]), + rnd='rn', + ftz=False, + ) + subtile_dprob = dprob_pair_0 + dprob_pair_1 + else: + subtile_dprob = dprob_swiglu.reduce( + cute.ReductionOp.ADD, + cutlass.Float32(0.0), + 0, + ) + if cutlass.const_expr(self.enable_breuse): + # Accumulate bk/br separately based on even/odd subtile + if real_subtile_idx % 2 == 0: + dProbVal_bk += subtile_dprob + else: + dProbVal_br += subtile_dprob + else: + dProbVal += subtile_dprob + + # + # Generate dBias + # + if cutlass.const_expr(self.generate_dbias): + n_base_d1 = epi_work_tile_info.tile_n_idx * (self.mma_tiler[1] * 2) + (2 * real_subtile_idx + 0) * self.epi_tile[1] + n_base_d2 = epi_work_tile_info.tile_n_idx * (self.mma_tiler[1] * 2) + (2 * real_subtile_idx + 1) * self.epi_tile[1] + dbias_n_total = cute.size(mDbias_tensor, mode=[1]) + self.dbias_reduction( + d1_vec, d2_vec, warp_idx, sDbias, + mDbias_tensor, expert_idx, n_base_d1, n_base_d2, + dbias_n_total, + ) + + # + # Generate amax + # + if cutlass.const_expr(self.generate_amax): + thread_tile_amax_1 = self.amax_reduction_per_thread( + d1_vec, thread_tile_amax_1 + ) + thread_tile_amax_2 = self.amax_reduction_per_thread( + d2_vec, thread_tile_amax_2 + ) + + # + # Generate SFD + # + if cutlass.const_expr(self.generate_sfd): + # + # Generate row major SFD + # + self.quant_sfd_row( + (real_subtile_idx * 2 + 0) % 4, + tiled_copy_r2s, + d1_vec, + tCrSFDRow_pvscale, + norm_const, + d_rcp_limits, + tRS_rD1, + ) + self.quant_sfd_col( + (real_subtile_idx * 2 + 0) % 4, + tiled_copy_r2s, + d1_vec, + tCrSFDCol_pvscale, + norm_const, + d_rcp_limits, + tRS_rD1_col, + ) + self.quant_sfd_row( + (real_subtile_idx * 2 + 1) % 4, + tiled_copy_r2s, + d2_vec, + tCrSFDRow_pvscale, + norm_const, + d_rcp_limits, + tRS_rD2, + ) + self.quant_sfd_col( + (real_subtile_idx * 2 + 1) % 4, + tiled_copy_r2s, + d2_vec, + tCrSFDCol_pvscale, + norm_const, + d_rcp_limits, + tRS_rD2_col, + ) + + if subtile_idx % 2 == 1: + local_m_tile = epi_work_tile_info.tile_m_idx + local_n_tile = epi_work_tile_info.tile_n_idx + sfd_row_idx_mn = ( + local_m_tile * self.epi_tile_cnt[0] + 0, + local_n_tile * self.epi_tile_cnt[1] // 2 + + (real_subtile_idx // 2), + ) + sfd_col_idx_mn = sfd_row_idx_mn + if cutlass.const_expr(self.discrete_col_sfd): + sfd_col_idx_mn = ( + local_m_tile * self.epi_tile_cnt[0] + 0, + local_n_tile * self.epi_tile_cnt[1] // 2 + + (real_subtile_idx // 2), + ) + + tCgSFDRow = tCgSFDRow_mn[ + ( + None, + None, + None, + *sfd_row_idx_mn, + ) + ] + tCgSFDCol = tCgSFDCol_mn[ + ( + None, + None, + None, + *sfd_col_idx_mn, + ) + ] + if cutlass.const_expr(not self.use_fp8_ptx_cvt): + tCrSFDRow.store( + tCrSFDRow_pvscale.load().to(self.sf_dtype) + ) + tCrSFDCol.store( + tCrSFDCol_pvscale.load().to(self.sf_dtype) + ) + else: + self.cvt_f32x4_to_f8x4(tCrSFDRow_pvscale, tCrSFDRow) + self.cvt_f32x4_to_f8x4(tCrSFDCol_pvscale, tCrSFDCol) + if sfd_row_idx_mn[1] * 32 * regPerSubtile < cute.size(cute.shape(mSFDRow_mnl.layout, mode=[1])): + cute.autovec_copy(tCrSFDRow, tCgSFDRow) + if sfd_col_idx_mn[1] * 32 * regPerSubtile < cute.size(cute.shape(mSFDCol_mnl.layout, mode=[1])): + cute.autovec_copy(tCrSFDCol, tCgSFDCol) + else: + # + # Convert to D type + # + tRS_rD1.store(d1_vec.to(self.d_dtype)) + tRS_rD2.store(d2_vec.to(self.d_dtype)) + + # + # Store D + # + if cutlass.const_expr(self.store_d_directly): + self.epilog_sync_barrier.arrive_and_wait() + d_idx_mn = (mma_tile_coord_mnl[0], mma_tile_coord_mnl[1]) + d_epilogue_subtile = ( + cute.make_layout(128), + cute.make_layout(self.mma_tiler[1] * 2), + ) + gD_sub_loop = cute.local_tile( + real_d, d_epilogue_subtile, (None, None, None) + ) + tCgD_mnl_loop = thr_copy_t2r.partition_D(gD_sub_loop) + tCgD_mnl_loop = cute.filter_zeros(tCgD_mnl_loop) + tCgD1 = tCgD_mnl_loop[ + ( + None, + 0, # T2R_M + 2 * real_subtile_idx + 0, # T2R_N + *d_idx_mn, # RestM/N + 0, # RestL + ) + ] + tCgD2 = tCgD_mnl_loop[ + ( + None, + 0, # T2R_M + 2 * real_subtile_idx + 1, # T2R_N + *d_idx_mn, # RestM/N + 0, # RestL + ) + ] + self.store_global_memory_256b(tCgD1, tRS_rD1) + self.store_global_memory_256b(tCgD2, tRS_rD2) + else: + if warp_idx == self.epilog_warp_id[0]: + d_pipeline.producer_acquire() + self.epilog_sync_barrier.arrive_and_wait() + d1_buffer = num_prev_subtiles % self.num_d_stage + num_prev_subtiles = num_prev_subtiles + 1 + cute.copy( + tiled_copy_r2s, + tRS_rD1, + tRS_sD[(None, None, None, d1_buffer)], + ) + if cutlass.const_expr(self.generate_sfd): + cute.copy( + tiled_copy_r2s, + tRS_rD1_col, + tRS_sD_col[(None, None, None, d1_buffer)], + ) + d2_buffer = num_prev_subtiles % self.num_d_stage + num_prev_subtiles = num_prev_subtiles + 1 + cute.copy( + tiled_copy_r2s, + tRS_rD2, + tRS_sD[(None, None, None, d2_buffer)], + ) + if cutlass.const_expr(self.generate_sfd): + cute.copy( + tiled_copy_r2s, + tRS_rD2_col, + tRS_sD_col[(None, None, None, d2_buffer)], + ) + # Fence and barrier to make sure shared memory store is visible to TMA store + cute.arch.fence_proxy("async.shared", space="cta") + self.epilog_sync_barrier.arrive_and_wait() + # + # TMA store D to global memory + # + if warp_idx == self.epilog_warp_id[0]: + cute.copy( + tma_atom_d, + bSG_sD[(None, d1_buffer)], + bSG_gD[(None, 2 * real_subtile_idx + 0)], + ) + cute.copy( + tma_atom_d, + bSG_sD[(None, d2_buffer)], + bSG_gD[(None, 2 * real_subtile_idx + 1)], + ) + if cutlass.const_expr(self.generate_sfd): + cute.copy( + tma_atom_d_col, + bSG_sD_col[(None, d1_buffer)], + bSG_gD_col[(None, 2 * real_subtile_idx + 0)], + ) + cute.copy( + tma_atom_d_col, + bSG_sD_col[(None, d2_buffer)], + bSG_gD_col[(None, 2 * real_subtile_idx + 1)], + ) + # Fence and barrier to make sure shared memory store is visible to TMA store + d_pipeline.producer_commit() + self.epilog_sync_barrier.arrive_and_wait() + + # + # Async arrive accumulator buffer empty + # + if cutlass.const_expr(not self.overlapping_accum): + with cute.arch.elect_one(): + acc_pipeline.consumer_release(acc_consumer_state) + acc_consumer_state.advance() + + # + # Advance to next tile + # + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + # Perform amax reduction after all subtiles are processed + if cutlass.const_expr(self.generate_amax): + gAmax1 = mAmax_tensor[ + (expert_idx, 0, None) + ].iterator.llvm_ptr + gAmax2 = mAmax_tensor[ + (expert_idx, 1, None) + ].iterator.llvm_ptr + self.amax_reduction_per_warp_and_cta( + thread_tile_amax_1, + warp_idx, + sAmax, + gAmax1, + ) + self.amax_reduction_per_warp_and_cta( + thread_tile_amax_2, + warp_idx, + sAmax, + gAmax2, + ) + + if cutlass.const_expr(self.generate_dprob): + real_dprob, _ = epi_ext.get_gmem_tensor( + "dprob", dprob, padded_offsets, epi_work_tile_info + ) + if cutlass.const_expr(self.enable_breuse): + _ = atomic_add_float32( + ptr=real_dprob[(mPosition_bk, None, None)].iterator.llvm_ptr, + value=dProbVal_bk, + ) + _ = atomic_add_float32( + ptr=real_dprob[(mPosition_br, None, None)].iterator.llvm_ptr, + value=dProbVal_br, + ) + else: + _ = atomic_add_float32( + ptr=real_dprob[(mPosition, None, None)].iterator.llvm_ptr, + value=dProbVal, + ) + + # + # Dealloc the tensor memory buffer + # + tmem.relinquish_alloc_permit() + self.epilog_sync_barrier.arrive_and_wait() + tmem.free(tmem_ptr) + # + # Wait for D store complete + # + if cutlass.const_expr(not self.store_d_directly): + d_pipeline.producer_tail() + # + # Specialized epilog load warp (loads C from GMEM to SMEM via TMA) + # + if warp_idx == self.epilog_load_tma_id: + c_load_ext = self._make_extension(workspace_ptr) + + tile_info_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_tile_stage + ) + tile_info = cute.make_rmem_tensor((4,), cutlass.Int32) + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + c_pipeline_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_c_stage + ) + is_reverse = True + while is_valid_tile: + if cutlass.const_expr(self.overlapping_accum): + reverse_subtile = is_reverse + is_reverse = not is_reverse + + c_work_tile_info = MoEWorkTileInfo( + expert_idx=tile_info[0], + tile_m_idx=tile_info[1], + tile_n_idx=tile_info[2], + k_tile_cnt=tile_info[3], + ) + mma_tile_coord_mnl = ( + c_work_tile_info.tile_m_idx // cute.size(tiled_mma.thr_id.shape), + c_work_tile_info.tile_n_idx * 2, + cutlass.Int32(0), + ) + + # Per-expert C tensor via extension + real_c, _ = c_load_ext.get_gmem_tensor( + "c", mC_mnl, padded_offsets, c_work_tile_info + ) + gC_mnl_loop = cute.local_tile( + real_c, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None) + ) + thr_mma_c_load = tiled_mma.get_slice(mma_tile_coord_v) + tCgC_loop = thr_mma_c_load.partition_C(gC_mnl_loop) + + # Build bGS_gC for bkeep half (mode1=0, same as non-breuse) + bGS_sC, bGS_gC_partitioned_bk = self.epilog_gmem_copy_and_partition( + tidx, tma_atom_c, tCgC_loop, epi_tile, sC + ) + bGS_gC_bk = bGS_gC_partitioned_bk[ + (None, None, None, mma_tile_coord_mnl[0], mma_tile_coord_mnl[1], 0) + ] + bGS_gC_bk = cute.group_modes(bGS_gC_bk, 1, cute.rank(bGS_gC_bk)) + bGS_gC_br = bGS_gC_bk # default; overridden for breuse + if cutlass.const_expr(self.enable_breuse): + # Breuse half: select mode1=1 from tCgC_loop, build separate partition + gC_br_half = tCgC_loop[((None, None), 1, 0, None, None, None)] + gC_br_epi = cute.flat_divide(gC_br_half, epi_tile) + sC_for_tma = cute.group_modes(sC, 0, 2) + bGS_sC_br, bGS_gC_partitioned_br = cpasync.tma_partition( + tma_atom_c, 0, cute.make_layout(1), sC_for_tma, + cute.group_modes(gC_br_epi, 0, 2), + ) + bGS_gC_br = bGS_gC_partitioned_br[ + (None, None, None, mma_tile_coord_mnl[0], mma_tile_coord_mnl[1], 0) + ] + bGS_gC_br = cute.group_modes(bGS_gC_br, 1, cute.rank(bGS_gC_br)) + # Default for JIT scoping + bGS_gC = bGS_gC_bk + else: + bGS_gC = bGS_gC_bk + subtile_cnt = cute.size(bGS_gC_bk.shape, mode=[1]) + # For breuse: produce C pipeline entries in INTERLEAVED order matching + # the transform-based main epilogue (M0,N0),(M1,N0),(M0,N1),(M1,N1),... + # Non-breuse: standard loop. + for subtile_idx in cutlass.range(subtile_cnt, unroll=1): + real_subtile_idx = subtile_idx + if cutlass.const_expr(self.overlapping_accum): + if reverse_subtile: + real_subtile_idx = subtile_cnt - 1 - subtile_idx + + # For each N-subtile, load C1+C2 for bkeep M group (same as non-breuse) + c_pipeline.producer_acquire(c_pipeline_producer_state) + cute.copy( + tma_atom_c, + bGS_gC_bk[(None, 2 * real_subtile_idx + 0)], + bGS_sC[(None, c_pipeline_producer_state.index)], + tma_bar_ptr=c_pipeline.producer_get_barrier( + c_pipeline_producer_state + ), + ) + c_pipeline_producer_state.advance() + c_pipeline.producer_acquire(c_pipeline_producer_state) + cute.copy( + tma_atom_c, + bGS_gC_bk[(None, 2 * real_subtile_idx + 1)], + bGS_sC[(None, c_pipeline_producer_state.index)], + tma_bar_ptr=c_pipeline.producer_get_barrier( + c_pipeline_producer_state + ), + ) + c_pipeline_producer_state.advance() + + # For breuse: immediately follow with C1+C2 for breuse M group + # (interleaved: bk,N0 → br,N0 → bk,N1 → br,N1 → ...) + if cutlass.const_expr(self.enable_breuse): + c_pipeline.producer_acquire(c_pipeline_producer_state) + cute.copy( + tma_atom_c, + bGS_gC_br[(None, 2 * real_subtile_idx + 0)], + bGS_sC[(None, c_pipeline_producer_state.index)], + tma_bar_ptr=c_pipeline.producer_get_barrier( + c_pipeline_producer_state + ), + ) + c_pipeline_producer_state.advance() + c_pipeline.producer_acquire(c_pipeline_producer_state) + cute.copy( + tma_atom_c, + bGS_gC_br[(None, 2 * real_subtile_idx + 1)], + bGS_sC[(None, c_pipeline_producer_state.index)], + tma_bar_ptr=c_pipeline.producer_get_barrier( + c_pipeline_producer_state + ), + ) + c_pipeline_producer_state.advance() + + # + # Advance to next tile + # + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + # + # Wait C buffer tail complete + # + c_pipeline.producer_tail(c_pipeline_producer_state) + + def epilog_tmem_copy_and_partition( + self, + tidx: cutlass.Int32, + tAcc: cute.Tensor, + gD_mnl: cute.Tensor, + epi_tile: cute.Tile, + use_2cta_instrs: Union[cutlass.Boolean, bool], + ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]: + """ + Make tiledCopy for tensor memory load, then use it to partition tensor memory (source) + and derive register array shape from the TMEM partition (no gmem dependency). + + For breuse: tAcc and gD_mnl have been through transform_partitioned_tensor_layout + so the M-split (mode1=2) is already merged into the M dimension. + For non-breuse: tAcc has shape (MMA, 1, 1, STAGE); strip with [0,0] internally. + + :param tidx: The thread index in epilogue warp groups + :type tidx: cutlass.Int32 + :param tAcc: The accumulator tensor to be copied and partitioned + :type tAcc: cute.Tensor + :param gD_mnl: The global D tensor (shape-only, for copy atom setup) + :type gD_mnl: cute.Tensor + :param epi_tile: The epilogue tiler + :type epi_tile: cute.Tile + :param use_2cta_instrs: Whether use_2cta_instrs is enabled + :type use_2cta_instrs: bool + + :return: A tuple containing (tiled_copy_t2r, tTR_tAcc, tTR_rAcc) where: + - tiled_copy_t2r: The tiled copy operation for tmem to register copy(t2r) + - tTR_tAcc: The partitioned accumulator tensor in TMEM + - tTR_rAcc: The register tensor for accumulator (shape derived from TMEM partition) + :rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor] + """ + copy_atom_t2r = sm100_utils.get_tmem_load_op( + self.cta_tile_shape_mnk, + self.d_layout, + self.d_dtype, + self.acc_dtype, + epi_tile, + use_2cta_instrs, + ) + + if cutlass.const_expr(self.enable_breuse): + # After transform_partitioned_tensor_layout, tAcc has merged M-split; + # gD_mnl is also already transformed — divide directly. + tAcc_epi = cute.flat_divide(tAcc, epi_tile) + gD_mnl_epi = cute.flat_divide(gD_mnl, epi_tile) + else: + # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, STAGE) + tAcc_epi = cute.flat_divide( + tAcc[((None, None), 0, 0, None)], + epi_tile, + ) + gD_mnl_epi = cute.flat_divide( + gD_mnl[((None, None), 0, 0, None, None, None)], + epi_tile, + ) + # (EPI_TILE_M, EPI_TILE_N) + tiled_copy_t2r = tcgen05.make_tmem_copy( + copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)] + ) + + thr_copy_t2r = tiled_copy_t2r.get_slice(tidx) + # (T2R, T2R_M, T2R_N, EPI_M, EPI_N, STAGE) + tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi) + tTR_gD = thr_copy_t2r.partition_D(gD_mnl_epi) + + # Derive register shape from gmem D partition + tTR_rAcc = cute.make_rmem_tensor(tTR_gD[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype) + return tiled_copy_t2r, tTR_tAcc, tTR_rAcc + + def epilog_smem_copy_and_partition_load( + self, + tiled_copy_t2r: cute.TiledCopy, + tTR_rC: cute.Tensor, + tTR_rC1: cute.Tensor, + tidx: cutlass.Int32, + sC: cute.Tensor, + ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]: + """ + Make tiledCopy for shared memory load, then use it to partition register array (destination) and shared memory (source). + + :param tiled_copy_t2r: The tiled copy operation for tmem to register copy(t2r) + :type tiled_copy_t2r: cute.TiledCopy + :param tTR_rC: The partitioned accumulator tensor + :type tTR_rC: cute.Tensor + :param tidx: The thread index in epilogue warp groups + :type tidx: cutlass.Int32 + :param sC: The shared memory tensor to be copied and partitioned + :type sC: cute.Tensor + + :return: A tuple containing (tiled_copy_s2r, tSR_rC, tSR_sC) where: + - tiled_copy_s2r: The tiled copy operation for smem to register copy(s2r) + - tSR_rC: The partitioned tensor C (register destination) + - tSR_sC: The partitioned tensor C (smem source) + :rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor] + """ + copy_atom_s2r = cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), self.c_dtype) + tiled_copy_s2r = cute.make_tiled_copy_D(copy_atom_s2r, tiled_copy_t2r) + # (S2R, S2R_M, S2R_N, PIPE_C) + thr_copy_s2r = tiled_copy_s2r.get_slice(tidx) + tSR_sC = thr_copy_s2r.partition_D(sC) + # (S2R, S2R_M, S2R_N) + tSR_rC = tiled_copy_s2r.retile(tTR_rC) + tSR_rC1 = tiled_copy_s2r.retile(tTR_rC1) + return tiled_copy_s2r, tSR_rC, tSR_rC1, tSR_sC + + def epilog_smem_copy_and_partition_store( + self, + tiled_copy_t2r: cute.TiledCopy, + tTR_rD1: cute.Tensor, + tTR_rD2: cute.Tensor, + tidx: cutlass.Int32, + sD: cute.Tensor, + ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]: + """ + Make tiledCopy for shared memory store, then use it to partition register array (source) and shared memory (destination). + + :param tiled_copy_t2r: The tiled copy operation for tmem to register copy(t2r) + :type tiled_copy_t2r: cute.TiledCopy + :param tTR_rD1: The partitioned accumulator tensor + :type tTR_rD1: cute.Tensor + :param tTR_rD2: The partitioned accumulator tensor + :type tTR_rD2: cute.Tensor + :param tidx: The thread index in epilogue warp groups + :type tidx: cutlass.Int32 + :param sD: The shared memory tensor to be copied and partitioned + :type sD: cute.Tensor + + :return: A tuple containing (tiled_copy_r2s, tRS_rD, tRS_sD) where: + - tiled_copy_r2s: The tiled copy operation for register to smem copy(r2s) + - tRS_rD: The partitioned tensor D (register source) + - tRS_sD: The partitioned tensor D (smem destination) + :rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor] + """ + copy_atom_r2s = sm100_utils.get_smem_store_op( + self.d_layout, self.d_dtype, self.acc_dtype, tiled_copy_t2r + ) + tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r) + # (R2S, R2S_M, R2S_N, PIPE_D) + thr_copy_r2s = tiled_copy_r2s.get_slice(tidx) + tRS_sD = None + if cutlass.const_expr(sD is not None): + tRS_sD = thr_copy_r2s.partition_D(sD) + # (R2S, R2S_M, R2S_N) + tRS_rD1 = tiled_copy_r2s.retile(tTR_rD1) + tRS_rD2 = tiled_copy_r2s.retile(tTR_rD2) + return tiled_copy_r2s, tRS_rD1, tRS_rD2, tRS_sD + + def epilog_gmem_copy_and_partition( + self, + tidx: cutlass.Int32, + atom: Union[cute.CopyAtom, cute.TiledCopy], + gD_mnl: cute.Tensor, + epi_tile: cute.Tile, + sD: cute.Tensor, + ) -> Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor]: + """Make tiledCopy for global memory store, then use it to: + - partition register array (source) and global memory (destination) for none TMA store version; + - partition shared memory (source) and global memory (destination) for TMA store version. + + :param tidx: The thread index in epilogue warp groups + :type tidx: cutlass.Int32 + :param atom: The copy_atom_c to be used for TMA store version, or tiled_copy_t2r for none TMA store version + :type atom: cute.CopyAtom or cute.TiledCopy + :param gD_mnl: The global tensor D + :type gD_mnl: cute.Tensor + :param epi_tile: The epilogue tiler + :type epi_tile: cute.Tile + :param sD: The shared memory tensor to be copied and partitioned + :type sD: cute.Tensor + + :return: A tuple containing : + - For TMA store: (tma_atom_d, bSG_sD, bSG_gD) where: + - tma_atom_d: The TMA copy atom + - bSG_sD: The partitioned shared memory tensor D + - bSG_gD: The partitioned global tensor D + :rtype: Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor] + """ + # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, loopM, loopN, loopL) + gD_epi = cute.flat_divide( + gD_mnl[((None, None), 0, 0, None, None, None)], epi_tile + ) + tma_atom_d = atom + sD_for_tma_partition = cute.group_modes(sD, 0, 2) + gD_for_tma_partition = cute.group_modes(gD_epi, 0, 2) + # ((ATOM_V, REST_V), EPI_M, EPI_N) + # ((ATOM_V, REST_V), EPI_M, EPI_N, loopM, loopN, loopL) + bSG_sD, bSG_gD = cpasync.tma_partition( + tma_atom_d, + 0, + cute.make_layout(1), + sD_for_tma_partition, + gD_for_tma_partition, + ) + return bSG_sD, bSG_gD + + @staticmethod + def _compute_stages( + tiled_mma: cute.TiledMma, + mma_tiler_mnk: Tuple[int, int, int], + a_dtype: Type[cutlass.Numeric], + b_dtype: Type[cutlass.Numeric], + epi_tile: cute.Tile, + c_dtype: Type[cutlass.Numeric], + c_layout: utils.LayoutEnum, + d_dtype: Type[cutlass.Numeric], + d_layout: utils.LayoutEnum, + sf_dtype: Type[cutlass.Numeric], + sf_vec_size: int, + num_smem_capacity: int, + occupancy: int, + store_d_directly: bool, + generate_dbias: bool = False, + enable_breuse: bool = False, + ) -> Tuple[int, int, int]: + """Computes the number of stages for A/B/D operands based on heuristics. + + :param tiled_mma: The tiled MMA object defining the core computation. + :type tiled_mma: cute.TiledMma + :param mma_tiler_mnk: The shape (M, N, K) of the MMA tiler. + :type mma_tiler_mnk: tuple[int, int, int] + :param a_dtype: Data type of operand A. + :type a_dtype: type[cutlass.Numeric] + :param b_dtype: Data type of operand B. + :type b_dtype: type[cutlass.Numeric] + :param epi_tile: The epilogue tile shape. + :type epi_tile: cute.Tile + :param c_dtype: Data type of operand C (output). + :type c_dtype: type[cutlass.Numeric] + :param d_layout: Layout of operand D. + :type d_layout: utils.LayoutEnum + :param sf_dtype: Data type of scale factor. + :type sf_dtype: type[cutlass.Numeric] + :param sf_vec_size: Vector size of scale factor. + :type sf_vec_size: int + :param num_smem_capacity: Total available shared memory capacity in bytes. + :type num_smem_capacity: int + :param occupancy: Target number of CTAs per SM (occupancy). + :type occupancy: int + + :return: A tuple containing the computed number of stages for: + (ACC stages, A/B operand stages, D stages) + :rtype: tuple[int, int, int] + """ + # Default ACC stages + num_acc_stage = 1 if (mma_tiler_mnk[1] == 256 or (enable_breuse and mma_tiler_mnk[1] == 192)) else 2 + + # Default C/D stages + num_c_stage = 4 if a_dtype.width == 8 else (4 if store_d_directly else 2) + num_d_stage = 2 if a_dtype.width == 8 else (0 if store_d_directly else 2) + + # Default Tile info stages + num_tile_stage = 2 + + # Calculate smem layout and size for one stage of A, B, and D + a_smem_layout_stage_one = sm100_utils.make_smem_layout_a( + tiled_mma, + mma_tiler_mnk, + a_dtype, + 1, # a tmp 1 stage is provided + ) + b_smem_layout_staged_one = sm100_utils.make_smem_layout_b( + tiled_mma, + mma_tiler_mnk, + b_dtype, + 1, # a tmp 1 stage is provided + ) + + sfa_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfa( + tiled_mma, + mma_tiler_mnk, + sf_vec_size, + 1, # a tmp 1 stage is provided + ) + + sfb_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfb( + tiled_mma, + mma_tiler_mnk, + sf_vec_size, + 1, # a tmp 1 stage is provided + ) + + c_smem_layout_staged_one = sm100_utils.make_smem_layout_epi( + c_dtype, + c_layout, + epi_tile, + 1, + ) + + d_smem_layout_staged_one = sm100_utils.make_smem_layout_epi( + d_dtype, + d_layout, + epi_tile, + 1, + ) + + ab_bytes_per_stage = ( + cute.size_in_bytes(a_dtype, a_smem_layout_stage_one) + + cute.size_in_bytes(b_dtype, b_smem_layout_staged_one) + + cute.size_in_bytes(sf_dtype, sfa_smem_layout_staged_one) + + cute.size_in_bytes(sf_dtype, sfb_smem_layout_staged_one) + ) + # Mbar bytes + mbar_helpers_bytes = 1024 + # Sinfo bytes + sinfo_bytes = 4 * 4 * num_tile_stage + # C/D bytes + c_bytes_per_stage = cute.size_in_bytes(c_dtype, c_smem_layout_staged_one) + c_bytes = c_bytes_per_stage * num_c_stage + d_bytes_per_stage = cute.size_in_bytes(d_dtype, d_smem_layout_staged_one) + d_bytes = d_bytes_per_stage * num_d_stage + if d_dtype == cutlass.Float8E5M2 or d_dtype == cutlass.Float8E4M3FN: + d_bytes = d_bytes * 2 + # AMAX bytes + amax_bytes = ( + get_amax_smem_size() + if d_dtype == cutlass.BFloat16 + else 0 + ) + # dBias transpose buffer: (128, 64) column-major FP32 = 32 KB + dbias_bytes = ( + 128 * 64 * cute.size_in_bytes(cutlass.Float32, cute.make_layout((1,))) + if generate_dbias + else 0 + ) + # Epilogue bytes + epi_bytes = c_bytes + d_bytes + amax_bytes + dbias_bytes + + # Calculate A/B stages: + # Start with total smem per CTA (capacity / occupancy) + # Subtract reserved bytes and initial D stages bytes + # Divide remaining by bytes needed per A/B stage + num_ab_stage = ( + num_smem_capacity // occupancy + - (mbar_helpers_bytes + epi_bytes + sinfo_bytes) + ) // ab_bytes_per_stage + + total_bytes = occupancy * ( + ab_bytes_per_stage * num_ab_stage + + epi_bytes + + sinfo_bytes + + mbar_helpers_bytes + ) + + return num_acc_stage, num_ab_stage, num_c_stage, num_d_stage, num_tile_stage diff --git a/python/cudnn/grouped_gemm/grouped_gemm_glu/_blockscaled_api.py b/python/cudnn/grouped_gemm/grouped_gemm_glu/_blockscaled_api.py index c739e485e..c3f6e0e28 100644 --- a/python/cudnn/grouped_gemm/grouped_gemm_glu/_blockscaled_api.py +++ b/python/cudnn/grouped_gemm/grouped_gemm_glu/_blockscaled_api.py @@ -44,7 +44,7 @@ """ from .moe_blockscaled_grouped_gemm_glu_bias import BlockScaledMoEGroupedGemmGluBiasKernel -from ..grouped_gemm_utils import _torch_stream_context +from ..grouped_gemm_utils import _torch_stream_context, rubin_single_group_offsets_kwarg from ..moe_utils import MoEWeightMode from cuda.bindings import driver as cuda import os @@ -60,6 +60,29 @@ from cudnn.api_base import APIBase, ceil_div, is_power_of_2 +def _get_rubin_kernel(): + from .moe_blockscaled_grouped_gemm_glu_rubin import ( + BlockScaledMoEGroupedGemmGluKernel as RubinBlockScaledMoEGroupedGemmGluKernel, + ) + + return RubinBlockScaledMoEGroupedGemmGluKernel + + +_GEGGLU_ALPHA_DEFAULT = 1.702 +_GLU_CLAMP_MAX_DEFAULT = 7.0 +_GLU_CLAMP_MIN_DEFAULT = -7.0 + + +def _reject_unsupported_rubin_glu_tune_params( + is_rubin_kernel: bool, + geglu_alpha: float, + glu_clamp_max: float, + glu_clamp_min: float, +) -> None: + if is_rubin_kernel and (geglu_alpha != _GEGGLU_ALPHA_DEFAULT or glu_clamp_max != _GLU_CLAMP_MAX_DEFAULT or glu_clamp_min != _GLU_CLAMP_MIN_DEFAULT): + raise NotImplementedError("Rubin grouped GEMM GLU does not support geglu_alpha, glu_clamp_max, or glu_clamp_min tuning") + + class GroupedGemmGluBlockScaledAPI(APIBase): """Unified API for grouped GEMM GLU forward operation on SM100+ GPUs. @@ -242,7 +265,7 @@ def __init__( self._interpret_uint8_as_fp4x2 = True self._has_bias = self.bias_desc is not None - self._kernel = BlockScaledMoEGroupedGemmGluBiasKernel + self._kernel = _get_rubin_kernel() if self._is_rubin_kernel else BlockScaledMoEGroupedGemmGluBiasKernel self.num_cluster_overlap_margin = int(os.getenv("CUDNNFE_CLUSTER_OVERLAP_MARGIN", "0")) self._logger.debug(f"setting num_cluster_overlap_margin: {self.num_cluster_overlap_margin}") @@ -538,8 +561,8 @@ def check_support(self) -> bool: f"m_aligned must be divisible by mma_tiler_mn[0], got {self.m_aligned} % {self.mma_tiler_mn[0]} != 0", ) self._value_error_if( - self.m_aligned != BlockScaledMoEGroupedGemmGluBiasKernel.FIX_PAD_SIZE, - f"m_aligned must be {BlockScaledMoEGroupedGemmGluBiasKernel.FIX_PAD_SIZE} (FIX_PAD_SIZE), got {self.m_aligned}", + self.m_aligned != self._kernel.FIX_PAD_SIZE, + f"m_aligned must be {self._kernel.FIX_PAD_SIZE} (FIX_PAD_SIZE), got {self.m_aligned}", ) # ---- Tensor alignment ---- @@ -628,7 +651,7 @@ def compile(self) -> None: act_func=self.act_func, enable_bias=self._has_bias, use_dynamic_sched=self.use_dynamic_sched, - use_single_group_runtime_offsets=self.use_single_group_runtime_offsets, + **rubin_single_group_offsets_kwarg(self._is_rubin_kernel, self.use_single_group_runtime_offsets), ) hardware_info = cutlass.utils.HardwareInfo() @@ -830,12 +853,9 @@ def _compile_dense(self, gemm_glu, max_active_clusters, fake_stream) -> None: ) # Compile with keyword args (dense mode uses the unified __call__ positional order). - # linear_offset, geglu_alpha, glu_clamp_max, and glu_clamp_min are runtime - # cutlass.Float32 (not Constexpr), so the compile-time placeholders below are - # irrelevant -- the values passed through tensor_api() at execute() time are - # what the kernel actually uses. - _compiled_kernel = cute.compile( - gemm_glu, + # linear_offset is a runtime cutlass.Float32 on both paths. geglu_alpha and the + # clamp limits are only supported on the SM100 kernel. + compile_kwargs = dict( a=a_cute_fake, b=b_cute_fake, sfb=sfb_cute_fake, @@ -860,11 +880,17 @@ def _compile_dense(self, gemm_glu, max_active_clusters, fake_stream) -> None: stream=fake_stream, epilogue_op=lambda x: x, linear_offset=cutlass.Float32(0.0), - geglu_alpha=cutlass.Float32(1.702), - glu_clamp_max=cutlass.Float32(7.0), - glu_clamp_min=cutlass.Float32(-7.0), options="--enable-tvm-ffi", ) + if not self._is_rubin_kernel: + compile_kwargs.update( + { + "geglu_alpha": cutlass.Float32(_GEGGLU_ALPHA_DEFAULT), + "glu_clamp_max": cutlass.Float32(_GLU_CLAMP_MAX_DEFAULT), + "glu_clamp_min": cutlass.Float32(_GLU_CLAMP_MIN_DEFAULT), + } + ) + _compiled_kernel = cute.compile(gemm_glu, **compile_kwargs) # Cache workspace pointer for the tensor_api closure cached_workspace_ptr = from_dlpack(self._workspace, assumed_align=128).iterator @@ -892,7 +918,7 @@ def tensor_api( glu_clamp_min: float = -7.0, ) -> None: norm_const_tensor = self._unpad_tensor_to_ndim(norm_const_tensor, 1, "norm_const") - _compiled_kernel( + kernel_args = ( a_tensor, b_tensor, sfb_tensor, @@ -914,10 +940,16 @@ def tensor_api( bias_tensor, stream, cutlass.Float32(linear_offset), - cutlass.Float32(geglu_alpha), - cutlass.Float32(glu_clamp_max), - cutlass.Float32(glu_clamp_min), ) + if self._is_rubin_kernel: + _compiled_kernel(*kernel_args) + else: + _compiled_kernel( + *kernel_args, + cutlass.Float32(geglu_alpha), + cutlass.Float32(glu_clamp_max), + cutlass.Float32(glu_clamp_min), + ) self._compiled_kernel = tensor_api @@ -1016,12 +1048,9 @@ def _compile_discrete(self, gemm_glu, max_active_clusters, fake_stream) -> None: workspace_ptr_cute = from_dlpack(self._workspace, assumed_align=128).iterator - # linear_offset, geglu_alpha, glu_clamp_max, and glu_clamp_min are runtime - # cutlass.Float32 (not Constexpr), so the compile-time placeholders below are - # irrelevant -- the values passed through tensor_api() at execute() time are - # what the kernel actually uses. + # linear_offset is runtime on both paths; clamp/alpha tuning kwargs are SM100-only. self._logger.debug("Compiling discrete grouped GEMM GLU kernel") - _compiled_kernel = cute.compile( + discrete_compile_args = ( gemm_glu, a_tensor, b_ptrs_cute, @@ -1047,11 +1076,17 @@ def _compile_discrete(self, gemm_glu, max_active_clusters, fake_stream) -> None: fake_stream, lambda x: x, # epilogue_op (Constexpr, baked in) cutlass.Float32(0.0), - cutlass.Float32(1.702), - cutlass.Float32(7.0), - cutlass.Float32(-7.0), - options="--enable-tvm-ffi", ) + if self._is_rubin_kernel: + _compiled_kernel = cute.compile(*discrete_compile_args, options="--enable-tvm-ffi") + else: + _compiled_kernel = cute.compile( + *discrete_compile_args, + cutlass.Float32(_GEGGLU_ALPHA_DEFAULT), + cutlass.Float32(_GLU_CLAMP_MAX_DEFAULT), + cutlass.Float32(_GLU_CLAMP_MIN_DEFAULT), + options="--enable-tvm-ffi", + ) self._n = n self._k = k @@ -1089,7 +1124,7 @@ def tensor_api( b_ptrs_addr = int(b_ptrs_device.data_ptr()) sfb_ptrs_addr = int(sfb_ptrs_device.data_ptr()) - _compiled_kernel( + kernel_args = ( a_tensor, b_ptrs_addr, sfb_ptrs_addr, @@ -1111,10 +1146,16 @@ def tensor_api( bias_tensor, stream, cutlass.Float32(linear_offset), - cutlass.Float32(geglu_alpha), - cutlass.Float32(glu_clamp_max), - cutlass.Float32(glu_clamp_min), ) + if self._is_rubin_kernel: + _compiled_kernel(*kernel_args) + else: + _compiled_kernel( + *kernel_args, + cutlass.Float32(geglu_alpha), + cutlass.Float32(glu_clamp_max), + cutlass.Float32(glu_clamp_min), + ) self._compiled_kernel = tensor_api @@ -1209,6 +1250,12 @@ def execute( # that pre-date the explicit linear_offset kwarg. if linear_offset is None: linear_offset = 1.0 if self.act_func == "geglu" else 0.0 + _reject_unsupported_rubin_glu_tune_params( + self._is_rubin_kernel, + geglu_alpha, + glu_clamp_max, + glu_clamp_min, + ) self._logger.debug("Executing grouped GEMM GLU kernel") if self._has_bias: diff --git a/python/cudnn/grouped_gemm/grouped_gemm_glu/api.py b/python/cudnn/grouped_gemm/grouped_gemm_glu/api.py index 12cd2e00c..6d7bfd578 100644 --- a/python/cudnn/grouped_gemm/grouped_gemm_glu/api.py +++ b/python/cudnn/grouped_gemm/grouped_gemm_glu/api.py @@ -57,7 +57,7 @@ import torch from typing import Any, Tuple, Optional, overload -from cudnn.api_base import APIBase, TupleDict, ceil_div +from cudnn.api_base import APIBase, TupleDict, ceil_div, get_device_type _BLOCK_SCALED_DTYPE_PAIRS = { (dtype, dtype) @@ -71,7 +71,10 @@ from ._bf16_api import GroupedGemmGluBf16API -from ._blockscaled_api import GroupedGemmGluBlockScaledAPI +from ._blockscaled_api import ( + GroupedGemmGluBlockScaledAPI, + _reject_unsupported_rubin_glu_tune_params, +) @dataclass(frozen=True) @@ -611,8 +614,11 @@ def dynamic_m_tensor_signature( use_full_dynamic = is_dense and os.environ.get("CUDNN_FE_GROUPED_GEMM_DYNAMIC_MNKL", "1") != "0" + device_type = get_device_type() + if is_dense: cache_key = ( + device_type, weight_mode, act_func, use_full_dynamic, @@ -655,6 +661,7 @@ def dynamic_m_tensor_signature( ) else: cache_key = ( + device_type, weight_mode, act_func, a_tensor.shape[1:], @@ -1135,6 +1142,12 @@ def grouped_gemm_glu_wrapper_sm100( generate_c: bool = False, ) -> TupleDict: """Dispatch grouped GEMM GLU once from an immutable normalized call.""" + _reject_unsupported_rubin_glu_tune_params( + get_device_type() == "rubin", + geglu_alpha, + glu_clamp_max, + glu_clamp_min, + ) call = GluCall( a_tensor=a_tensor, sfa_tensor=sfa_tensor, diff --git a/python/cudnn/grouped_gemm/grouped_gemm_glu/moe_blockscaled_grouped_gemm_glu_bias.py b/python/cudnn/grouped_gemm/grouped_gemm_glu/moe_blockscaled_grouped_gemm_glu_bias.py index 67a3a90ef..f031ded2f 100644 --- a/python/cudnn/grouped_gemm/grouped_gemm_glu/moe_blockscaled_grouped_gemm_glu_bias.py +++ b/python/cudnn/grouped_gemm/grouped_gemm_glu/moe_blockscaled_grouped_gemm_glu_bias.py @@ -1228,13 +1228,13 @@ def store_c( real_subtile_idx, ) -> None: c_buffer = prev_subtile_idx % self.num_c_stage - tRS_rC.store(tTR_rAcc.load().to(self.c_dtype)) + tRS_rC.store(tiled_copy_r2s.retile(tTR_rAcc).load().to(self.c_dtype)) cute.copy( tiled_copy_r2s, tRS_rC[(None, None, 0)], tRS_sC[(None, None, 0, c_buffer)], ) - tRS_rC.store(tTR_rAcc_up.load().to(self.c_dtype)) + tRS_rC.store(tiled_copy_r2s.retile(tTR_rAcc_up).load().to(self.c_dtype)) cute.copy( tiled_copy_r2s, tRS_rC[(None, None, 0)], diff --git a/python/cudnn/grouped_gemm/grouped_gemm_glu/moe_blockscaled_grouped_gemm_glu_rubin.py b/python/cudnn/grouped_gemm/grouped_gemm_glu/moe_blockscaled_grouped_gemm_glu_rubin.py new file mode 100644 index 000000000..46cde9615 --- /dev/null +++ b/python/cudnn/grouped_gemm/grouped_gemm_glu/moe_blockscaled_grouped_gemm_glu_rubin.py @@ -0,0 +1,3364 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: BSD-3-Clause + +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: + +# 1. Redistributions of source code must retain the above copyright notice, this +# list of conditions and the following disclaimer. + +# 2. Redistributions in binary form must reproduce the above copyright notice, +# this list of conditions and the following disclaimer in the documentation +# and/or other materials provided with the distribution. + +# 3. Neither the name of the copyright holder nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. + +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +""" +MoE Block-Scaled Grouped GEMM Kernel with GLU (SwiGLU/GeGLU) Fusion. + +Supports: + - Static / Dynamic persistent tile scheduling (MoEPersistentTileScheduler) + - Dense (contiguous 3-D B) / Discrete (per-expert pointer array B) weight layout + - FP8/FP4 output quantization with row/column scale factors (SFD) + - Optional C output (generate_c) + - AMAX reduction for FP8 calibration + - GLU activation fusion (SwiGLU / GeGLU) + +This module contains only the kernel class. +MoE scheduler components live in moe_persistent_scheduler.py / moe_sched_extension.py / moe_utils.py. +""" + +from typing import Type, Tuple, Union, Optional + +import cuda.bindings.driver as cuda + +import cutlass +import cutlass.cute as cute +from cutlass.cute.nvgpu import cpasync, tcgen05 +from cutlass.cute.nvgpu.tcgen05 import OperandMajorMode, CollectorOp +from cutlass.utils.gemm.sm100 import transform_partitioned_tensor_layout +import cutlass.utils as utils +import cutlass.pipeline as pipeline +import cutlass.utils.blackwell_helpers as sm100_utils +import cutlass.utils.rubin_helpers as sm107_utils +import cutlass.utils.blockscaled_layout as blockscaled_utils +from cutlass.cute.typing import Float32, Int32, AddressSpace +from ..moe_persistent_scheduler import ( + MoEPersistentTileScheduler, + MoESchedulerParams, + MoEWorkTileInfo, +) +from ..moe_utils import ( + compute_expert_token_range, + MoEWeightMode, + TensormapWorkspace, + store_tma_desc, +) +from ..moe_sched_extension import ( + DiscreteWeightScaledGemmSchedExtension, + ContiguousAndConsistentGroupedGemmSchedExtension, +) +from ..moe_kernel_helpers import ( + fmin, + fmax, + atomic_max_float32, + silu_f32, + silu_f32_geglu_scaled, + compute_stages, + compute_grid, + get_dtype_rcp_limits, + can_implement, + amax_reduction_per_thread, + epilog_gmem_copy_and_partition, +) + + +class BlockScaledMoEGroupedGemmGluKernel: + """Block-scaled grouped GEMM kernel with MoE tile scheduling and GLU fusion. + + Supports both dense and discrete weight layouts, static and dynamic + scheduling, and quantized output with row/column scale factors. + + This version uses a fixed padding size (FIX_PAD_SIZE=256) that is decoupled from the kernel's tile size, + allowing users to pad their tensors without knowing the specific kernel implementation details. + + :param sf_vec_size: Scalefactor vector size. + :type sf_vec_size: int + :param mma_tiler_mn: Shape of the Matrix Multiply-Accumulate (MMA) tile (M,N) + :type mma_tiler_mn: Tuple[int, int] + :param cluster_shape_mn: Cluster dimensions (M,N) for parallel processing + :type cluster_shape_mn: Tuple[int, int] + :param expert_cnt: Number of experts (compile-time constant) + :type expert_cnt: int + + :note: In current version, A and B tensor must have the same data type + - i.e., Float8E4M3FN for A and Float8E5M2 for B is not supported + + :note: Supported combinations of A/B data types, SF data typs and SF vector size: + - MXF8: A/B: Float8E5M2/Float8E4M3FN + SF: Float8E8M0FNU + sf_vec_size: 32 + - MXF4: A/B: Float4E2M1FN + SF: Float8E8M0FNU + sf_vec_size: 32 + - NVF4: A/B: Float4E2M1FN + SF: Float8E8M0FNU/Float8E4M3FN + sf_vec_size: 16 + + :note: Supported accumulator data types: + - Float32 + + :note: Supported D data types: + - BFloat16 + - Float8E4M3FN/Float8E5M2 + + :note: Constraints: + - MMA tiler M must be 128 or 256 (use_2cta_instrs) + - MMA tiler N must be 64/128/192/256 + - Cluster shape M must be multiple of 2 if Mma tiler M is 256 + - Cluster shape M/N must be positive and power of 2, total cluster size <= 16 + - Also, Cluster shape M/N must be <= 4 for scale factor multicasts due to limited size of scale factors + - FIX_PAD_SIZE (256) must be divisible by mma_tiler_mn[0] + - m_aligned parameter in create_mask() MUST equal FIX_PAD_SIZE (256) + - Each padded_offsets[i] will be a multiple of FIX_PAD_SIZE (guaranteed by m_aligned == FIX_PAD_SIZE) + + :note: New Interface (padded_offsets): + Instead of tile_idx_to_expert_idx, num_non_exiting_tiles, and m_split_cumsum, users now provide: + - padded_offsets: shape (expert_cnt,), where padded_offsets[i] is the end position + of expert[i] in the padded A tensor. + - Expert i processes A[padded_offsets[i-1]:padded_offsets[i], :] (with padded_offsets[-1]=0) + + """ + + # Fixed pad size for user-side padding (decoupled from kernel tile size) + FIX_PAD_SIZE = 256 + + @staticmethod + def can_implement( + ab_dtype: Type[cutlass.Numeric], + sf_dtype: Type[cutlass.Numeric], + sf_vec_size: int, + acc_dtype: Type[cutlass.Numeric], + d_dtype: Type[cutlass.Numeric], + use_2cta_instrs: bool, + mma_tiler_mn: Tuple[int, int], + cluster_shape_mn: Tuple[int, int], + m: int, + n: int, + k: int, + l: int, + a_major: str, + b_major: str, + cd_major: str, + m_aligned: int, + ) -> bool: + from ..moe_kernel_helpers import ( + is_valid_dtypes_and_scale_factor_vec_size, + is_valid_layouts, is_valid_tensor_alignment, FIX_PAD_SIZE, + ) + # B-reuse: 2CTA + mma_tiler_mn[0]=512, N in {192, 256} + if use_2cta_instrs and mma_tiler_mn[0] == 512: + if not is_valid_dtypes_and_scale_factor_vec_size(ab_dtype, sf_dtype, sf_vec_size, acc_dtype, d_dtype): + return False + if not is_valid_layouts(ab_dtype, d_dtype, a_major, b_major, cd_major): + return False + if not is_valid_tensor_alignment(m, n, k, l, ab_dtype, d_dtype, a_major, b_major, cd_major): + return False + # Tile N must be 192 or 256 + if mma_tiler_mn[1] not in {192, 256}: + return False + # Culster M must be a multiple of 2. For cluster M=1, we don't have test it yet. + if cluster_shape_mn[0] % 2 != 0: + return False + return True + # Allow N=192 in addition to the shared helper's N=256-only constraint. + if mma_tiler_mn[1] == 192: + if m_aligned != FIX_PAD_SIZE: + return False + if ab_dtype.width == 8: + return False + if not is_valid_dtypes_and_scale_factor_vec_size(ab_dtype, sf_dtype, sf_vec_size, acc_dtype, d_dtype): + return False + if not is_valid_layouts(ab_dtype, d_dtype, a_major, b_major, cd_major): + return False + if not is_valid_tensor_alignment(m, n, k, l, ab_dtype, d_dtype, a_major, b_major, cd_major): + return False + if not (a_major == "k" and b_major == "k"): + return False + if n % 64 != 0 or m % 256 != 0: + return False + if not (use_2cta_instrs and mma_tiler_mn[0] == 256): + return False + if cluster_shape_mn[0] % 2 != 0: + return False + return True + return can_implement( + ab_dtype, sf_dtype, sf_vec_size, acc_dtype, d_dtype, + use_2cta_instrs, mma_tiler_mn, cluster_shape_mn, + m, n, k, l, a_major, b_major, cd_major, m_aligned, + fix_pad_size=BlockScaledMoEGroupedGemmGluKernel.FIX_PAD_SIZE, + ) + + def __init__( + self, + sf_vec_size: int, + acc_dtype: Type[cutlass.Numeric], + use_2cta_instrs: bool, + mma_tiler_mn: Tuple[int, int], + cluster_shape_mn: Tuple[int, int], + vectorized_f32: bool, + generate_sfd: bool, + discrete_col_sfd: bool, + expert_cnt: int, + weight_mode: MoEWeightMode = MoEWeightMode.DISCRETE, + use_dynamic_sched: bool = False, + act_func: str = "swiglu", + enable_bias: bool = False, + generate_c: bool = True, + ): + """Initializes the configuration for a Blackwell blockscaled grouped GEMM GLU kernel. + + This configuration includes several key aspects: + + 1. MMA Instruction Settings (tcgen05): + - acc_dtype: Data types for MMA accumulator. + - mma_tiler_mn: The (M, N) shape of the MMA instruction tiler. + - use_2cta_instrs: Boolean indicating if the tcgen05 MMA variant + with cta_group=2 should be used. + + 2. Cluster Shape: + - cluster_shape_mn: The (ClusterM, ClusterN) shape of the CTA cluster. + + 3. Expert Count: + - expert_cnt: Number of experts for MoE grouped GEMM. + + 4. MoE Tile Scheduling: + - Uses MoEPersistentTileScheduler for tile iteration across experts + - Expert lookup is handled by the scheduler (cached O(1) fast path) + + :param acc_dtype: Data type of the accumulator. + :type acc_dtype: type[cutlass.Numeric] + :param mma_tiler_mn: Tuple (M, N) shape of the MMA instruction. + :type mma_tiler_mn: Tuple[int, int] + :param use_2cta_instrs: Boolean, True to use cta_group=2 MMA variant. + :type use_2cta_instrs: bool + :param cluster_shape_mn: Tuple (ClusterM, ClusterN) shape of the cluster. + :type cluster_shape_mn: Tuple[int, int] + :param expert_cnt: Number of experts (compile-time constant). + :type expert_cnt: int + + :raises ValueError: If FIX_PAD_SIZE is not divisible by mma_tiler_mn[0]. + """ + mma_inst_m = 256 if use_2cta_instrs else 128 + enable_breuse = (mma_tiler_mn[0] // mma_inst_m == 2) + if enable_breuse: + if mma_tiler_mn[0] % self.FIX_PAD_SIZE != 0: + raise ValueError( + f"mma_tiler_mn[0] ({mma_tiler_mn[0]}) must be a multiple of " + f"FIX_PAD_SIZE ({self.FIX_PAD_SIZE}) for breuse. " + f"Also use m_aligned=mma_tiler_mn[0] (={mma_tiler_mn[0]})." + ) + else: + cta_tile_m = mma_tiler_mn[0] // (2 if use_2cta_instrs else 1) + if self.FIX_PAD_SIZE % cta_tile_m != 0: + raise ValueError( + f"FIX_PAD_SIZE ({self.FIX_PAD_SIZE}) must be divisible by " + f"cta_tile_m ({cta_tile_m}). " + f"Supported mma_tiler_mn[0] values: 128, 256, 512." + ) + if expert_cnt > 1024: + raise ValueError("Expert count > 1024 is not supported.") + if not isinstance(weight_mode, MoEWeightMode): + raise TypeError(f"weight_mode must be a MoEWeightMode, got {type(weight_mode)}") + + self.sf_vec_size = sf_vec_size + self.expert_cnt = expert_cnt + self.acc_dtype: Type[cutlass.Numeric] = acc_dtype + self.use_2cta_instrs = use_2cta_instrs + self.cluster_shape_mn = cluster_shape_mn + self.mma_tiler = (*mma_tiler_mn, 1) + self.enable_breuse = enable_breuse + + self.cta_group = ( + tcgen05.CtaGroup.TWO if use_2cta_instrs else tcgen05.CtaGroup.ONE + ) + + self.enable_bias = enable_bias + self.generate_c = generate_c + self.occupancy = 1 + self.epilog_warp_id = (0, 1, 2, 3) + self.mma_warp_id = 4 + self.tma_warp_id = 5 + self.sched_warp_id = 6 + self.bias_load_warp_id = 7 if enable_bias else None + self.threads_per_warp = 32 + + all_warps = [ + *self.epilog_warp_id, + self.mma_warp_id, + self.tma_warp_id, + self.sched_warp_id, + ] + warps_wo_sched = [ + *self.epilog_warp_id, + self.mma_warp_id, + self.tma_warp_id, + ] + if enable_bias: + all_warps.append(self.bias_load_warp_id) + warps_wo_sched.append(self.bias_load_warp_id) + self.threads_per_cta = self.threads_per_warp * len(all_warps) + self.threads_wo_sched = self.threads_per_warp * len(warps_wo_sched) + + # Set barrier for cta sync, epilogue sync and tmem ptr sync + self.cta_sync_barrier = pipeline.NamedBarrier( + barrier_id=1, + num_threads=self.threads_per_cta, + ) + self.epilog_sync_barrier = pipeline.NamedBarrier( + barrier_id=2, + num_threads=32 * len(self.epilog_warp_id), + ) + self.tmem_alloc_barrier = pipeline.NamedBarrier( + barrier_id=3, + num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id)), + ) + self.sched_sync_barrier = pipeline.NamedBarrier( + barrier_id=4, + num_threads=self.threads_per_warp, + ) + self.num_smem_capacity = utils.get_smem_capacity_in_bytes("sm_107") + self.num_tmem_alloc_cols = cute.arch.get_max_tmem_alloc_cols("sm_107") + + self.vectorized_f32 = vectorized_f32 + + self.generate_sfd = generate_sfd + self.discrete_col_sfd = discrete_col_sfd + self.weight_mode = weight_mode + self.use_dynamic_sched = use_dynamic_sched + + # Amax reduction configuration + self.num_epilog_warps = len(self.epilog_warp_id) + + self.act_func = act_func + if act_func not in ["swiglu", "geglu"]: + raise ValueError(f"Invalid activation function: {act_func}") + + + def _get_mma_permutation_mnk(self, mma_inst_shape_mnk): + """Return MMA permutation for the Bkeep-Breuse pattern (2CTA only). + Only active when enable_breuse=True to avoid breaking the TMA atom + setup for the non-breuse 2CTA case. + """ + if cutlass.const_expr(self.use_2cta_instrs and self.enable_breuse): + m_layout = cute.make_layout( + shape=(mma_inst_shape_mnk[0] // 2, 2, 2), + stride=(1, mma_inst_shape_mnk[0], mma_inst_shape_mnk[0] // 2), + ) + return (m_layout, mma_inst_shape_mnk[1], mma_inst_shape_mnk[2]) + return (1, 1, 1) + + def _setup_attributes(self): + """Set up configurations that are dependent on GEMM inputs + + This method configures various attributes based on the input tensor properties + (data types, leading dimensions) and kernel settings: + - Configuring tiled MMA + - Computing MMA/cluster/tile shapes + - Computing cluster layout + - Computing multicast CTAs for A/B + - Computing epilogue subtile + - Setting up A/B/D stage counts in shared memory + - Computing A/B/D shared memory layout + - Computing tensor memory allocation columns + """ + + # Hardware MMA instruction M: 2CTA → 256, 1CTA → 128 + mma_inst_m = 256 if self.use_2cta_instrs else 128 + self.mma_inst_shape_mn = (mma_inst_m, self.mma_tiler[1]) + # (CTA_Tile_Shape_M, Round_Up(MMA_Tile_Shape_N, 128), MMA_Inst_Shape_K) + self.mma_inst_shape_mn_sfb = ( + self.mma_inst_shape_mn[0] // (2 if self.use_2cta_instrs else 1), + cute.round_up(self.mma_inst_shape_mn[1], 128), + ) + + # K dim of the MMA instruction shape on Rubin sm107: + # - SM107MmaMXF4NVF4Op (FP4 x FP4) requires K=128 + # - SM107BlockScaledMmaMXF8F6F4Op (FP8 or mixed) requires K=64 + mma_inst_k = ( + 128 + if (self.a_dtype.width == 4 and self.b_dtype.width == 4) + else 64 + ) + mma_inst_shape_mnk = (*self.mma_inst_shape_mn, mma_inst_k) + mma_inst_shape_mnk_sfb = (*self.mma_inst_shape_mn_sfb, mma_inst_k) + + # Configure tiled mma (Rubin sm107) — use permutation always for 2CTA + # (creates bkeep/breuse M-split in the accumulator when enable_breuse=True) + atom_layout_mnk = (1, 1, 1) + permutation_mnk = self._get_mma_permutation_mnk(mma_inst_shape_mnk) + tiled_mma = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + self.cta_group, + mma_inst_shape_mnk, + a_collector_op=CollectorOp.DISCARD, + b_collector_op=CollectorOp.DISCARD, + atom_layout_mnk=atom_layout_mnk, + permutation_mnk=permutation_mnk, + ) + + tiled_mma_sfb = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + cute.nvgpu.tcgen05.CtaGroup.ONE, + mma_inst_shape_mnk_sfb, + ) + + # Compute mma/cluster/tile shapes + mma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2]) + mma_inst_tile_k = 2 if (self.a_dtype.width == 4 and self.sf_vec_size == 16) else 4 + self.mma_tiler = ( + self.mma_tiler[0], + self.mma_tiler[1], + mma_inst_shape_k * mma_inst_tile_k, + ) + + self.mma_tiler_sfb = ( + self.mma_inst_shape_mn_sfb[0], + self.mma_inst_shape_mn_sfb[1], + mma_inst_shape_k * mma_inst_tile_k, + ) + + self.cta_tile_shape_mnk = ( + self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape), + self.mma_tiler[1], + self.mma_tiler[2], + ) + self.cta_tile_shape_mnk_sfb = ( + self.mma_tiler_sfb[0] // cute.size(tiled_mma.thr_id.shape), + self.mma_tiler_sfb[1], + self.mma_tiler_sfb[2], + ) + + # For breuse, D tiler M = mma_tiler[0] (512) so each CTA writes all its rows. + self.mma_tiler_d = ( + self.mma_tiler[0], + self.mma_inst_shape_mn[1] // 2, + mma_inst_shape_k * mma_inst_tile_k, + ) + self.cta_tile_shape_mnk_d = ( + self.mma_tiler_d[0] // cute.size(tiled_mma.thr_id.shape), + self.mma_tiler_d[1], + self.mma_tiler_d[2], + ) + # Compute cluster layout + self.cluster_layout_vmnk = cute.tiled_divide( + cute.make_layout((*self.cluster_shape_mn, 1)), + (tiled_mma.thr_id.shape,), + ) + + self.cluster_layout_sfb_vmnk = cute.tiled_divide( + cute.make_layout((*self.cluster_shape_mn, 1)), + (tiled_mma_sfb.thr_id.shape,), + ) + + # Compute number of multicast CTAs for A/B + self.num_mcast_ctas_a = cute.size(self.cluster_layout_vmnk.shape[2]) + self.num_mcast_ctas_b = cute.size(self.cluster_layout_vmnk.shape[1]) + self.is_a_mcast = self.num_mcast_ctas_a > 1 + self.is_b_mcast = self.num_mcast_ctas_b > 1 + + # Set epilogue subtile + self.epi_tile = (128, 32) + self.epi_tile_cnt = ( + self.cta_tile_shape_mnk_d[0] // self.epi_tile[0], + self.cta_tile_shape_mnk_d[1] // self.epi_tile[1], + ) + self.epi_tile_c = (128, 64) + + # Setup A/B/D/Scale stage count in shared memory and ACC stage count in tensor memory + ( + self.num_acc_stage, + self.num_ab_stage, + self.num_c_stage, + self.num_d_stage, + self.num_tile_stage, + self.num_bias_stage, + ) = compute_stages( + tiled_mma, + self.mma_tiler, + self.a_dtype, + self.b_dtype, + self.epi_tile, + self.epi_tile_c, + self.c_dtype, + self.c_layout, + self.d_dtype, + self.d_layout, + self.sf_dtype, + self.sf_vec_size, + self.num_smem_capacity, + self.occupancy, + self.generate_sfd, + bias_dtype=self.bias_dtype if self.enable_bias else None, + ) + + # Compute A/B/D/Scale shared memory layout + self.a_smem_layout_staged = sm100_utils.make_smem_layout_a( + tiled_mma, + self.mma_tiler, + self.a_dtype, + self.num_ab_stage, + ) + self.b_smem_layout_staged = sm100_utils.make_smem_layout_b( + tiled_mma, + self.mma_tiler, + self.b_dtype, + self.num_ab_stage, + ) + self.sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa( + tiled_mma, + self.mma_tiler, + self.sf_vec_size, + self.num_ab_stage, + ) + self.sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb( + tiled_mma, + self.mma_tiler, + self.sf_vec_size, + self.num_ab_stage, + ) + + self.c_smem_layout_staged = sm100_utils.make_smem_layout_epi( + self.c_dtype, + self.c_layout, + self.epi_tile_c, + self.num_c_stage, + ) + + self.d_smem_layout_staged = sm100_utils.make_smem_layout_epi( + self.d_dtype, + self.d_layout, + self.epi_tile, + self.num_d_stage, + ) + + # For breuse+N=192, override num_acc_stage=1 (compute_stages uses N=256 heuristic only) + if self.enable_breuse and self.mma_tiler[1] in {192, 256}: + self.num_acc_stage = 1 + + # overlapping_accum: share SF TMEM with second acc stage; incompatible with breuse + self.overlapping_accum = ( + self.num_acc_stage == 1 and self.mma_tiler[1] == 256 and not self.enable_breuse + ) + + # Compute number of TMEM columns for SFA/SFB/Accumulator. + sf_atom_mn = 32 + sf_pack_factor = 32 // self.sf_vec_size + self.num_sfa_tmem_cols = ( + self.cta_tile_shape_mnk[0] // sf_atom_mn + ) * mma_inst_tile_k * sf_pack_factor + self.num_sfb_tmem_cols = ( + self.cta_tile_shape_mnk_sfb[1] // sf_atom_mn + ) * mma_inst_tile_k * sf_pack_factor + self.num_sf_tmem_cols = self.num_sfa_tmem_cols + self.num_sfb_tmem_cols + if self.enable_breuse: + self.num_accumulator_tmem_cols = self.cta_tile_shape_mnk[1] * self.num_acc_stage * 2 + elif self.overlapping_accum: + self.num_accumulator_tmem_cols = self.cta_tile_shape_mnk[1] * 2 - self.num_sf_tmem_cols + else: + self.num_accumulator_tmem_cols = self.cta_tile_shape_mnk[1] * self.num_acc_stage + if self.cta_tile_shape_mnk[1] == 192 and not self.enable_breuse: + self.num_accumulator_tmem_cols = self.num_tmem_alloc_cols - self.num_sf_tmem_cols + self.num_accumulator_tmem_stride = self.num_accumulator_tmem_cols - 192 + else: + self.num_accumulator_tmem_stride = self.num_accumulator_tmem_cols + + self.epi_tile_n_required = 2 * cute.size(self.epi_tile[1]) + # Only when overlapping_accum is enabled, we need to release accumulator buffer early in epilogue + self.iter_acc_early_release_in_epilogue = ( + (self.num_sf_tmem_cols + self.epi_tile_n_required - 1) + // self.epi_tile_n_required + - 1 + ) * 2 + + # Bias SMEM layout: (tile_N, num_stages) double-buffered + if self.enable_bias: + self.bias_smem_layout_staged = cute.make_layout( + (self.mma_tiler[1], self.num_bias_stage), + stride=(1, self.mma_tiler[1]), + ) + else: + self.bias_smem_layout_staged = cute.make_layout((1, 1)) + + def get_desc_workspace_bytes(self) -> int: + """Return descriptor workspace size in bytes.""" + if self.weight_mode == MoEWeightMode.DISCRETE: + from ..moe_utils import DiscreteWeightTensormapConstructor + return DiscreteWeightTensormapConstructor.get_workspace_size(self.expert_cnt) + return 0 + + def get_workspace_bytes(self) -> int: + """Return descriptor workspace plus optional dynamic scheduler state.""" + desc_workspace_bytes = self.get_desc_workspace_bytes() + dynamic_sched_bytes = 4 if self.use_dynamic_sched else 0 + return desc_workspace_bytes + dynamic_sched_bytes + + @cute.jit + def _get_sched_counter_ptr(self, workspace_ptr): + counter_addr = workspace_ptr.toint() + self.get_desc_workspace_bytes() + return cute.make_ptr( + cutlass.Int32, + counter_addr, + AddressSpace.gmem, + assumed_align=4, + ) + + @cute.kernel + def helper_kernel( + self, + ptrs_b: cute.Pointer, + ptrs_sfb: cute.Pointer, + n: Int32, + k: Int32, + b_stride_size: cutlass.Int64, + b_major_mode: cutlass.Constexpr, + workspace_ptr, + tiled_mma_arg: cute.TiledMma, + tiled_mma_sfb_arg: cute.TiledMma, + b_smem_layout_arg, + sfb_smem_layout_arg, + cluster_layout_vmnk_shape_arg: cutlass.Constexpr, + cluster_layout_sfb_vmnk_shape_arg: cutlass.Constexpr, + ): + """Pre-main-kernel initialization. + + Launched with grid=(expert_cnt, 1, 1) for discrete mode, or + grid=(1, 1, 1) for dense+dynamic mode. + + Discrete weight: each block builds B/SFB TMA descriptors for one expert. + Dynamic sched: block 0 resets the atomic tile counter to 0. + """ + expert_idx = cute.arch.block_idx()[0] + + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE): + b_tma_op_arg = sm100_utils.cluster_shape_to_tma_atom_B( + self.cluster_shape_mn, tiled_mma_arg.thr_id + ) + sfb_tma_op_arg = sm100_utils.cluster_shape_to_tma_atom_SFB( + self.cluster_shape_mn, tiled_mma_arg.thr_id + ) + + b_ptr_tensor = cute.make_tensor( + cute.make_ptr(cutlass.Int64, ptrs_b.toint(), AddressSpace.gmem, assumed_align=8), + cute.make_layout((self.expert_cnt,)) + ) + sfb_ptr_tensor = cute.make_tensor( + cute.make_ptr(cutlass.Int64, ptrs_sfb.toint(), AddressSpace.gmem, assumed_align=8), + cute.make_layout((self.expert_cnt,)) + ) + + c0 = cutlass.Int64(0) + c1_64 = 1 + if cutlass.const_expr(b_major_mode == OperandMajorMode.K): + stride_n = b_stride_size + stride_k = c1_64 + else: + stride_n = c1_64 + stride_k = b_stride_size + + b_ptr_val = b_ptr_tensor[expert_idx] + b_ptr = cute.make_ptr(self.b_dtype, b_ptr_val, AddressSpace.gmem) + b_tensor_i = cute.make_tensor( + b_ptr, + cute.make_layout((n, k, cutlass.Int32(1)), stride=(stride_n, stride_k, c0)), + ) + tma_atom_b, _ = cute.nvgpu.make_tiled_tma_atom_B( + b_tma_op_arg, b_tensor_i, b_smem_layout_arg, + self.mma_tiler, tiled_mma_arg, + cluster_layout_vmnk_shape_arg, + ) + + workspace = TensormapWorkspace(workspace_ptr, ["b", "sfb"]) + store_tma_desc(tma_atom_b, workspace.get_ptr("b", expert_idx)) + + sfb_ptr_val = sfb_ptr_tensor[expert_idx] + sfb_ptr = cute.make_ptr(self.sf_dtype, sfb_ptr_val, AddressSpace.gmem) + sfb_layout = blockscaled_utils.tile_atom_to_shape_SF( + (n, k, cutlass.Int32(1)), self.sf_vec_size + ) + sfb_tensor_i = cute.make_tensor(sfb_ptr, sfb_layout) + tma_atom_sfb, _ = cute.nvgpu.make_tiled_tma_atom_B( + sfb_tma_op_arg, sfb_tensor_i, sfb_smem_layout_arg, + self.mma_tiler_sfb, tiled_mma_sfb_arg, + cluster_layout_sfb_vmnk_shape_arg, + internal_type=cutlass.Uint64, + ) + store_tma_desc(tma_atom_sfb, workspace.get_ptr("sfb", expert_idx)) + + if cutlass.const_expr(self.use_dynamic_sched): + if expert_idx == cutlass.Int32(0): + sched_counter = cute.make_tensor( + self._get_sched_counter_ptr(workspace_ptr), + cute.make_layout(1), + ) + sched_counter[0] = cutlass.Int32(0) + + @cute.jit + def __call__( + self, + a: cute.Tensor, + b, # Dense: cute.Tensor (N,K,L) | Discrete: cute.Pointer to int64[] + sfb, # Dense: cute.Tensor | Discrete: cute.Pointer to int64[] + n: Int32, # Ignored for dense mode + k: Int32, # Ignored for dense mode + b_stride_size: cutlass.Int64, # Ignored for dense mode + b_major_mode: cutlass.Constexpr, # Ignored for dense mode + workspace_ptr, + c: cute.Tensor, + d: cute.Tensor, + d_col: cute.Tensor, + sfa: cute.Tensor, + sfd_row_tensor: Optional[cute.Tensor], + sfd_col_tensor: Optional[cute.Tensor], + amax_tensor: Optional[cute.Tensor], + norm_const_tensor: Optional[cute.Tensor], + padded_offsets: cute.Tensor, + alpha: cute.Tensor, + prob: cute.Tensor, + bias: Optional[cute.Tensor], + max_active_clusters: cutlass.Constexpr, + stream: cuda.CUstream, + epilogue_op: cutlass.Constexpr = lambda x: x, + linear_offset: cutlass.Float32 = 0.0, + ): + """Execute the GEMM. + + Dense mode: ``b`` and ``sfb`` are 3-D cute.Tensor (N, K, L). + Discrete mode: ``b`` and ``sfb`` are cute.Pointer to device int64[] + arrays of per-expert base addresses; ``n``, ``k``, ``b_stride_size``, + ``b_major_mode`` describe the uniform per-expert layout. + """ + self.a_dtype: Type[cutlass.Numeric] = a.element_type + self.b_dtype: Type[cutlass.Numeric] = a.element_type + self.c_dtype: Type[cutlass.Numeric] = c.element_type + self.d_dtype: Type[cutlass.Numeric] = d.element_type + self.sf_dtype: Type[cutlass.Numeric] = sfa.element_type + self.bias_dtype = bias.element_type if cutlass.const_expr(self.enable_bias) else cutlass.BFloat16 + self.a_major_mode = utils.LayoutEnum.from_tensor(a).mma_major_mode() + self.c_layout = utils.LayoutEnum.from_tensor(c) + self.d_layout = utils.LayoutEnum.from_tensor(d) + + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DENSE): + self.b_major_mode = utils.LayoutEnum.from_tensor(b).mma_major_mode() + else: + self.b_major_mode = b_major_mode + + if cutlass.const_expr(self.a_dtype != self.b_dtype): + raise TypeError(f"Type must match: {self.a_dtype} != {self.b_dtype}") + + self._setup_attributes() + + # ---- B / SFB setup (mode-dependent) ---- + b_from_call_arg = b + sfb_from_call_arg = sfb + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DENSE): + sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b.shape, self.sf_vec_size) + sfb = cute.make_tensor(sfb.iterator, sfb_layout) + else: + c0 = cutlass.Int64(0) + c1_64 = 1 + if cutlass.const_expr(b_major_mode == OperandMajorMode.K): + b_template_stride = (b_stride_size, c1_64, c0) + else: + b_template_stride = (c1_64, b_stride_size, c0) + b_template_layout = cute.make_layout((n, k, cutlass.Int32(1)), stride=b_template_stride) + b_ptr_typed = cute.make_ptr( + self.b_dtype, b.toint(), AddressSpace.gmem, assumed_align=16 + ) + b = cute.make_tensor(b_ptr_typed, b_template_layout) + + sfb_ptr_typed = cute.make_ptr( + self.sf_dtype, sfb.toint(), AddressSpace.gmem, assumed_align=16 + ) + sfb_layout = blockscaled_utils.tile_atom_to_shape_SF( + (n, k, cutlass.Int32(1)), self.sf_vec_size + ) + sfb = cute.make_tensor(sfb_ptr_typed, sfb_layout) + + # Setup sfa tensor by filling A tensor to scale factor atom layout + # ((Atom_M, Rest_M),(Atom_K, Rest_K),RestL) + sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a.shape, self.sf_vec_size) + sfa = cute.make_tensor(sfa.iterator, sfa_layout) + + # Setup sfd tensor by filling D tensor to scale factor atom layout + self.generate_sfd = sfd_row_tensor is not None and norm_const_tensor is not None + if cutlass.const_expr(self.generate_sfd == False): + self.discrete_col_sfd = False + if cutlass.const_expr(self.generate_sfd): + sfd_row_layout = blockscaled_utils.tile_atom_to_shape_SF( + d.shape, self.sf_vec_size + ) + sfd_row_tensor = cute.make_tensor(sfd_row_tensor.iterator, sfd_row_layout) + sfd_col_layout = cute.tile_to_shape( + blockscaled_utils.BlockScaledBasicChunk( + self.sf_vec_size, OperandMajorMode.MN + ).layout, + d.shape, + (1, 2, 3), + ) + if cutlass.const_expr(self.discrete_col_sfd): + sfd_col_layout = sfd_row_layout + sfd_col_tensor = cute.make_tensor(sfd_col_tensor.iterator, sfd_col_layout) + + self.generate_amax = amax_tensor is not None + self.has_prob = prob is not None + + # K dim of the MMA instruction shape on Rubin sm107: + # - SM107MmaMXF4NVF4Op (FP4 x FP4) requires K=128 + # - SM107BlockScaledMmaMXF8F6F4Op (FP8 or mixed) requires K=64 + mma_inst_k = ( + 128 if (self.a_dtype.width == 4 and self.b_dtype.width == 4) else 64 + ) + mma_inst_shape_mnk = (*self.mma_inst_shape_mn, mma_inst_k) + mma_inst_shape_mnk_sfb = (*self.mma_inst_shape_mn_sfb, mma_inst_k) + + atom_layout_mnk = (1, 1, 1) + permutation_mnk = self._get_mma_permutation_mnk(mma_inst_shape_mnk) + tiled_mma = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + self.cta_group, + mma_inst_shape_mnk, + a_collector_op=CollectorOp.DISCARD, + b_collector_op=CollectorOp.DISCARD, + atom_layout_mnk=atom_layout_mnk, + permutation_mnk=permutation_mnk, + ) + + tiled_mma_sfb = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + cute.nvgpu.tcgen05.CtaGroup.ONE, + mma_inst_shape_mnk_sfb, + ) + + tiled_mma_bkeep = None + tiled_mma_breuse = None + if cutlass.const_expr(self.enable_breuse): + tiled_mma_bkeep = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, self.b_dtype, + self.a_major_mode, self.b_major_mode, + self.sf_dtype, self.sf_vec_size, + self.cta_group, mma_inst_shape_mnk, + a_collector_op=CollectorOp.DISCARD, + b_collector_op=CollectorOp.FILL, + atom_layout_mnk=atom_layout_mnk, + permutation_mnk=permutation_mnk, + ) + tiled_mma_bkeep.set(tcgen05.Field.NEGATE_A, False) + tiled_mma_bkeep.set(tcgen05.Field.NEGATE_B, False) + tiled_mma_breuse = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, self.b_dtype, + self.a_major_mode, self.b_major_mode, + self.sf_dtype, self.sf_vec_size, + self.cta_group, mma_inst_shape_mnk, + a_collector_op=CollectorOp.DISCARD, + b_collector_op=CollectorOp.LASTUSE, + atom_layout_mnk=atom_layout_mnk, + permutation_mnk=permutation_mnk, + ) + tiled_mma_breuse.set(tcgen05.Field.NEGATE_A, False) + tiled_mma_breuse.set(tcgen05.Field.NEGATE_B, False) + + atom_thr_size = cute.size(tiled_mma.thr_id.shape) + + # Setup TMA load for A + a_op = sm100_utils.cluster_shape_to_tma_atom_A( + self.cluster_shape_mn, tiled_mma.thr_id + ) + a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0)) + tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A( + a_op, + a, + a_smem_layout, + self.mma_tiler, + tiled_mma, + self.cluster_layout_vmnk.shape, + ) + + # Setup TMA load for B + b_op = sm100_utils.cluster_shape_to_tma_atom_B( + self.cluster_shape_mn, tiled_mma.thr_id + ) + b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0)) + tma_atom_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B( + b_op, + b, + b_smem_layout, + self.mma_tiler, + tiled_mma, + self.cluster_layout_vmnk.shape, + ) + + # Setup TMA load for SFA + sfa_op = sm100_utils.cluster_shape_to_tma_atom_A( + self.cluster_shape_mn, tiled_mma.thr_id + ) + sfa_smem_layout = cute.slice_( + self.sfa_smem_layout_staged, (None, None, None, 0) + ) + tma_atom_sfa, tma_tensor_sfa = cute.nvgpu.make_tiled_tma_atom_A( + sfa_op, + sfa, + sfa_smem_layout, + self.mma_tiler, + tiled_mma, + self.cluster_layout_vmnk.shape, + internal_type=cutlass.Int16, + ) + + # Setup TMA load for SFB + sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB( + self.cluster_shape_mn, tiled_mma.thr_id + ) + sfb_smem_layout = cute.slice_( + self.sfb_smem_layout_staged, (None, None, None, 0) + ) + tma_atom_sfb, tma_tensor_sfb = cute.nvgpu.make_tiled_tma_atom_B( + sfb_op, + sfb, + sfb_smem_layout, + self.mma_tiler_sfb, + tiled_mma_sfb, + self.cluster_layout_sfb_vmnk.shape, + internal_type=cutlass.Uint64, + ) + + a_copy_size = cute.size_in_bytes(self.a_dtype, a_smem_layout) + b_copy_size = cute.size_in_bytes(self.b_dtype, b_smem_layout) + sfa_copy_size = cute.size_in_bytes(self.sf_dtype, sfa_smem_layout) + sfb_copy_size = cute.size_in_bytes(self.sf_dtype, sfb_smem_layout) + self.num_tma_load_bytes = ( + a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size + ) * atom_thr_size + + # Setup TMA store for C + c_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0)) + tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom( + cpasync.CopyBulkTensorTileS2GOp(), + c, + c_smem_layout, + self.epi_tile_c, + ) + + # Setup TMA store for D + d_smem_layout = cute.slice_(self.d_smem_layout_staged, (None, None, 0)) + tma_atom_d, tma_tensor_d = cpasync.make_tiled_tma_atom( + cpasync.CopyBulkTensorTileS2GOp(), + d, + d_smem_layout, + self.epi_tile, + ) + tma_atom_d_col, tma_tensor_d_col = cpasync.make_tiled_tma_atom( + cpasync.CopyBulkTensorTileS2GOp(), + d_col, + d_smem_layout, + self.epi_tile, + ) + + # ---- Helper kernel: TMA desc init (discrete) + sched counter reset (dynamic) ---- + _need_helper = cutlass.const_expr( + self.weight_mode == MoEWeightMode.DISCRETE or self.use_dynamic_sched + ) + if cutlass.const_expr(_need_helper): + _helper_grid_x = self.expert_cnt if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else 1 + _helper_args = ( + b_from_call_arg if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cute.make_ptr(cutlass.Int64, 0, AddressSpace.gmem), + sfb_from_call_arg if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cute.make_ptr(cutlass.Int64, 0, AddressSpace.gmem), + n if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cutlass.Int32(0), + k if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cutlass.Int32(0), + b_stride_size if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cutlass.Int64(0), + b_major_mode if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else self.b_major_mode, + workspace_ptr, + tiled_mma, tiled_mma_sfb, + b_smem_layout, sfb_smem_layout, + self.cluster_layout_vmnk.shape, + self.cluster_layout_sfb_vmnk.shape, + ) + self.helper_kernel(*_helper_args).launch( + grid=(_helper_grid_x, 1, 1), block=(1, 1, 1), + stream=stream, min_blocks_per_mp=1, + ) + + # ---- Grid computation via MoE scheduler ---- + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DENSE): + b_n, b_k, b_l = cute.shape(b) + sched_expert_shape = (self.expert_cnt, b_n, b_k) + else: + sched_expert_shape = (self.expert_cnt, n, k) + + sched_params = MoESchedulerParams( + scenario="2Dx3D", + expert_shape=sched_expert_shape, + cta_tile_shape_mnk=self.cta_tile_shape_mnk, + cluster_shape_mn=self.cluster_shape_mn, + use_dynamic_sched=self.use_dynamic_sched, + ) + self.sched_params, grid = compute_grid( + sched_params, max_active_clusters, self.use_2cta_instrs, + ) + + self.buffer_align_bytes = 1024 + + # Define shared storage for kernel + # sD_col is only needed when generating SFD; use size 0 to avoid wasting smem + sD_col_size = cute.cosize(self.d_smem_layout_staged.outer) if self.generate_sfd else 0 + SchedulerStorage = MoEPersistentTileScheduler.make_storage_struct( + self.num_tile_stage, self.use_dynamic_sched + ) + + @cute.struct + class SharedStorage: + ab_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage * 2] + acc_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage * 2] + bias_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_bias_stage * 2 if self.enable_bias else 1] + scheduler: SchedulerStorage + tmem_dealloc_mbar_ptr: cutlass.Int64 + tmem_holding_buf: cutlass.Int32 + # (EPI_TILE_M, EPI_TILE_N, STAGE) + sC: cute.struct.Align[ + cute.struct.MemRange[ + self.c_dtype, + cute.cosize(self.c_smem_layout_staged.outer), + ], + self.buffer_align_bytes, + ] + sD: cute.struct.Align[ + cute.struct.MemRange[ + self.d_dtype, + cute.cosize(self.d_smem_layout_staged.outer), + ], + self.buffer_align_bytes, + ] + sD_col: cute.struct.Align[ + cute.struct.MemRange[self.d_dtype, sD_col_size], + self.buffer_align_bytes, + ] + # (MMA, MMA_M, MMA_K, STAGE) + sA: cute.struct.Align[ + cute.struct.MemRange[ + self.a_dtype, cute.cosize(self.a_smem_layout_staged.outer) + ], + self.buffer_align_bytes, + ] + # (MMA, MMA_N, MMA_K, STAGE) + sB: cute.struct.Align[ + cute.struct.MemRange[ + self.b_dtype, cute.cosize(self.b_smem_layout_staged.outer) + ], + self.buffer_align_bytes, + ] + # (granularity_m, repeat_m), (granularity_k, repeat_k), num_scale_stage) + sSFA: cute.struct.Align[ + cute.struct.MemRange[ + self.sf_dtype, cute.cosize(self.sfa_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + # (granularity_n, repeat_n), (granularity_k, repeat_k), num_scale_stage) + sSFB: cute.struct.Align[ + cute.struct.MemRange[ + self.sf_dtype, cute.cosize(self.sfb_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + # Amax reduction shared memory (one FP32 per epilogue warp) + sAmax: cute.struct.Align[ + cute.struct.MemRange[cutlass.Float32, self.num_epilog_warps], + 4, + ] + if cutlass.const_expr(self.enable_bias): + # Bias SMEM: (tile_N, num_bias_stage) BF16 double-buffered + sBias: cute.struct.Align[ + cute.struct.MemRange[self.bias_dtype, cute.cosize(self.bias_smem_layout_staged)], + 16, + ] + + self.shared_storage = SharedStorage + + # Launch the kernel synchronously + self.kernel( + tiled_mma, + tiled_mma_bkeep, + tiled_mma_breuse, + tiled_mma_sfb, + tma_atom_a, + tma_tensor_a, + tma_atom_b, + tma_tensor_b, + tma_atom_sfa, + tma_tensor_sfa, + tma_atom_sfb, + tma_tensor_sfb, + tma_atom_c, + tma_tensor_c, + tma_atom_d, + tma_tensor_d, + tma_atom_d_col, + tma_tensor_d_col, + sfd_row_tensor, + sfd_col_tensor, + norm_const_tensor, + amax_tensor, + padded_offsets, + alpha, + bias, + prob, + workspace_ptr, # Contains per-expert B/SFB TMA descriptors + self.cluster_layout_vmnk, + self.cluster_layout_sfb_vmnk, + self.a_smem_layout_staged, + self.b_smem_layout_staged, + self.sfa_smem_layout_staged, + self.sfb_smem_layout_staged, + self.c_smem_layout_staged, + self.d_smem_layout_staged, + self.bias_smem_layout_staged, + self.epi_tile, + self.sched_params, + epilogue_op, + linear_offset, + ).launch( + grid=grid, + block=[self.threads_per_cta, 1, 1], + cluster=(*self.cluster_shape_mn, 1), + max_number_threads=[self.threads_per_cta, 1, 1], + smem=self.shared_storage.size_in_bytes(), + stream=stream, + min_blocks_per_mp=1, + ) + return + + # ------------------------------------------------------------------ + # Internal: create extension based on weight_mode + # ------------------------------------------------------------------ + + @cute.jit + def _make_extension(self, workspace_ptr): + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE): + desc_workspace = TensormapWorkspace(workspace_ptr, ["b", "sfb"]) + return DiscreteWeightScaledGemmSchedExtension( + tensormap_ctor=desc_workspace, sf_vec_size=self.sf_vec_size, + ) + else: + return ContiguousAndConsistentGroupedGemmSchedExtension( + sf_vec_size=self.sf_vec_size, + ) + + def mainloop_s2t_copy_and_partition( + self, + sSF: cute.Tensor, + tSF: cute.Tensor, + ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]: + """ + Make tiledCopy for smem to tmem load for scale factor tensor, then use it to partition smem memory (source) and tensor memory (destination). + + :param sSF: The scale factor tensor in smem + :type sSF: cute.Tensor + :param tSF: The scale factor tensor in tmem + :type tSF: cute.Tensor + + :return: A tuple containing (tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t) where: + - tiled_copy_s2t: The tiled copy operation for smem to tmem load for scale factor tensor(s2t) + - tCsSF_compact_s2t: The partitioned scale factor tensor in smem + - tSF_compact_s2t: The partitioned scale factor tensor in tmem + :rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor] + """ + # (MMA, MMA_MN, MMA_K, STAGE) + tCsSF_compact = cute.filter_zeros(sSF) + # (MMA, MMA_MN, MMA_K) + tCtSF_compact = cute.filter_zeros(tSF) + + # Make S2T CopyAtom and tiledCopy + copy_atom_s2t = cute.make_copy_atom( + tcgen05.Cp4x32x128bOp(self.cta_group), + self.sf_dtype, + ) + tiled_copy_s2t = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSF_compact) + thr_copy_s2t = tiled_copy_s2t.get_slice(0) + + def _append_mn_broadcast_mode(smem_layout: cute.Layout): + mn_dim = cute.get(smem_layout, mode=[0, 0]) + mn_dim = cute.append(mn_dim, cute.make_layout((4), stride=(0))) + layout = cute.append( + cute.group_modes(mn_dim, 0), cute.get(smem_layout, mode=[0, 1]) + ) + layout = cute.append( + cute.group_modes(layout, 0), cute.get(smem_layout, mode=[1]) + ) + layout = cute.append(layout, cute.get(smem_layout, mode=[2])) + layout = cute.append(layout, cute.get(smem_layout, mode=[3])) + return layout + + tCsSF_compact_bcast = cute.make_tensor( + tCsSF_compact.iterator, _append_mn_broadcast_mode(tCsSF_compact.layout) + ) + + # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE) + tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact_bcast) + # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE) + tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor( + tiled_copy_s2t, tCsSF_compact_s2t_ + ) + # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K) + tCtSF_compact_s2t = thr_copy_s2t.partition_D(tCtSF_compact) + + return tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t + + @cute.jit + def amax_reduction_per_warp_and_cta( + self, amax_fp32, warp_idx, amax_smem, amax_gmem + ) -> None: + # Warp-level reduction using wrapper function + warp_amax = cute.arch.warp_redux_sync( + value=amax_fp32, + kind="fmax", + mask_and_clamp=0xFFFFFFFF, + nan=True, + ) + # Each epilogue warp's lane 0 writes warp amax to shared memory + if cute.arch.lane_idx() == 0: + amax_smem[warp_idx] = cutlass.Float32(warp_amax) + + # Ensure all epilogue warps complete their writes before block reduction + self.epilog_sync_barrier.arrive_and_wait() + + # Block-level reduction: only first epilogue warp's lane 0 handles this + if warp_idx == self.epilog_warp_id[0] and cute.arch.lane_idx() == 0: + block_amax = cutlass.Float32(0.0) + for i in cutlass.range(self.num_epilog_warps): + warp_amax_val = amax_smem[i] + block_amax = cute.arch.fmax(block_amax, warp_amax_val) + + # Global atomic max (accumulates across all tiles for final tensor amax) + _ = atomic_max_float32(ptr=amax_gmem, value=block_amax) + + @cute.jit + def store_c( + self, + tiled_copy_r2s, + tma_atom_c, + warp_idx, + tTR_rAcc, + tTR_rAcc_up, + tTR_rC, + tRS_rC, + tRS_sC, + bSG_gC, + bSG_sC, + c_pipeline, + prev_subtile_idx, + real_subtile_idx, + ) -> None: + c_buffer = prev_subtile_idx % self.num_c_stage + tTR_rC.store(tTR_rAcc.load().to(self.c_dtype)) + cute.copy( + tiled_copy_r2s, + tRS_rC[(None, None, 0)], + tRS_sC[(None, None, 0, c_buffer)], + ) + tTR_rC.store(tTR_rAcc_up.load().to(self.c_dtype)) + cute.copy( + tiled_copy_r2s, + tRS_rC[(None, None, 0)], + tRS_sC[(None, None, 1, c_buffer)], + ) + # Fence and barrier to make sure shared memory store is visible to TMA store + cute.arch.fence_proxy("async.shared", space="cta") + self.epilog_sync_barrier.arrive_and_wait() + # TMA store smem to global memory + if warp_idx == self.epilog_warp_id[0]: + cute.copy( + tma_atom_c, + bSG_sC[(None, c_buffer)], + bSG_gC[(None, real_subtile_idx)], + ) + # Fence and barrier to make sure shared memory store is visible to TMA store + c_pipeline.producer_commit() + c_pipeline.producer_acquire() + self.epilog_sync_barrier.arrive_and_wait() + + @cute.jit + def quant_sfd_row( + self, + tile_idx, + tiled_copy_r2s, + src, + pvscale, + norm_const, + rcp_limit, + tRSrD, + tile_info, + ) -> None: + # Get absolute max across a vector and Compute SFD + tTR_rAcc_frg = cute.logical_divide(src, cute.make_layout(self.sf_vec_size)) + acc_frg = tTR_rAcc_frg.load() + abs_acc_frg_ir = cutlass._mlir.dialects.math.absf(acc_frg.ir_value()) + abs_acc_frg = type(acc_frg)(abs_acc_frg_ir, acc_frg.shape, acc_frg.dtype) + + pvscale_f32x4 = cute.make_rmem_tensor(4, cutlass.Float32) + sfd_f8x4 = cute.make_rmem_tensor(4, self.sf_dtype) + tmp_f32 = ( + abs_acc_frg[None, 0].reduce( + cute.ReductionOp.MAX, + cutlass.Float32(0.0), + 0, # Use 0.0 as init for abs values + ) + * rcp_limit + * norm_const + ) + # + # Manually store pvscale to avoid spilling + # + if tile_idx == 0: + pvscale[0] = tmp_f32 + elif tile_idx == 1: + pvscale[1] = tmp_f32 + elif tile_idx == 2: + pvscale[2] = tmp_f32 + elif tile_idx == 3: + pvscale[3] = tmp_f32 + + # + # Compute quantized output values and convert to D type + # + pvscale_f32x4[0] = tmp_f32 + sfd_f8x4.store(pvscale_f32x4.load().to(self.sf_dtype)) + pvscale_f32x4.store(sfd_f8x4.load().to(cutlass.Float32)) + qpvscale_up = pvscale_f32x4[0] + + fp32_max = cutlass.Float32(3.40282346638528859812e38) + acc_scale = norm_const * cute.arch.rcp_approx(qpvscale_up) + acc_scale = fmin(acc_scale, fp32_max, nan=True) + if cutlass.const_expr(self.vectorized_f32): + vec = tTR_rAcc_frg[None, 0] + for ei in cutlass.range_constexpr(0, self.sf_vec_size, 2): + ( + vec[ei], + vec[ei + 1], + ) = cute.arch.mul_packed_f32x2( + (vec[ei], vec[ei + 1]), + (acc_scale, acc_scale), + rnd='rn', + ftz=False, + ) + else: + vec = tTR_rAcc_frg[None, 0] + for ei in cutlass.range_constexpr(self.sf_vec_size): + vec[ei] = vec[ei] * acc_scale + + acc_vec = tiled_copy_r2s.retile(src).load() + tRSrD.store(acc_vec.to(self.d_dtype)) + + @cute.jit + def quant_sfd_col( + self, + tile_idx, + tiled_copy_r2s, + src, + pvscale, + norm_const, + rcp_limit, + tRSrD, + tile_info, + ) -> None: + # Get absolute max across a vector and Compute SFD + tTR_rAcc_frg = cute.logical_divide(src, cute.make_layout(self.sf_vec_size)) + acc_frg = tTR_rAcc_frg.load() + abs_acc_frg_ir = cutlass._mlir.dialects.math.absf(acc_frg.ir_value()) + acc_frg = type(acc_frg)(abs_acc_frg_ir, acc_frg.shape, acc_frg.dtype) + + tmp_f32 = cutlass.Float32(0.0) + for vi in cutlass.range_constexpr(acc_frg.shape[0]): + max_value_original = ( + cutlass.Float32( + cute.arch.warp_redux_sync( + value=acc_frg[vi, 0], + kind="fmax", + mask_and_clamp=0xFFFFFFFF, + nan=True, + ) + ) + * rcp_limit + * norm_const + ) + max_value_vec = cute.full(4, max_value_original, dtype=cutlass.Float32) + max_value_vec_f8 = max_value_vec.to(cutlass.Float8E8M0FNU) + max_value_vec_f32_chunked = max_value_vec_f8.to(cutlass.Float32) + max_value = max_value_vec_f32_chunked[0] + tidx = cute.arch.thread_idx()[0] + if tidx % 32 == vi: + tmp_f32 = max_value + + acc_scale_col = cutlass.Float32(0.0) + if max_value_vec_f32_chunked[0] == 0.000000: + acc_scale_col = cutlass.Float32(0.0) + else: + acc_scale_col = norm_const * cute.arch.rcp_approx( + max_value_vec_f32_chunked[0] + ) + fp32_max = cutlass.Float32(3.40282346638528859812e38) + acc_scale_col = fmin(acc_scale_col, fp32_max) + tTR_rAcc_frg[vi] = tTR_rAcc_frg[vi] * acc_scale_col + pvscale[None, None, tile_idx][0] = tmp_f32 + + acc_vec = tiled_copy_r2s.retile(src).load() + tRSrD.store(acc_vec.to(self.d_dtype)) + + + @cute.jit + def tile_info_to_mn_idx( + self, + tile_info: cute.Tensor, + ): + m_idx = tile_info[1] * cute.size(self.cta_tile_shape_mnk[0]) + n_idx = tile_info[2] * cute.size(self.cta_tile_shape_mnk[1]) + return m_idx, n_idx + + @cute.jit + def create_and_partition_new_SFDCol( + self, + tile_info: cute.Tensor, + mSFDCol_mnl: cute.Tensor, + padded_offsets: cute.Tensor, + ): + m_idx, n_idx = self.tile_info_to_mn_idx(tile_info) + expert_idx = tile_info[0] + cumsum_tokens, tokens_this_group = compute_expert_token_range( + padded_offsets, expert_idx + ) + n_total = cute.size(mSFDCol_mnl.shape[1]) + + sf_tile_idx_begin = cumsum_tokens // cute.size(mSFDCol_mnl.shape[0][0]) + mSFDCol_mnl_new_ptr = mSFDCol_mnl[(None, sf_tile_idx_begin), None, 0].iterator + + sfd_col_quant_layout = cute.tile_to_shape( + blockscaled_utils.BlockScaledBasicChunk( + self.sf_vec_size, OperandMajorMode.MN + ).layout, + (tokens_this_group, n_total, mSFDCol_mnl.shape[2]), + (1, 2, 3), + ) + regPerSubtile = 4 + sfd_tile = ( + cute.make_layout(128), + cute.make_layout(32 * regPerSubtile), + ) + mSFDCol_mnl_new = cute.make_tensor(mSFDCol_mnl_new_ptr, sfd_col_quant_layout) + gSFDCol_mnl_new = cute.local_tile(mSFDCol_mnl_new, sfd_tile, (None, None, None)) + + thr_layout = cute.make_ordered_layout((4, 32), order=(1, 0)) + val_layout = cute.make_ordered_layout((1,), order=(0,)) + copy_atom_sfd_col_quant = cute.make_copy_atom( + cute.nvgpu.CopyUniversalOp(), + gSFDCol_mnl_new.element_type, + num_bits_per_copy=8, + ) + tiled_copy_sfd_col_quant = cute.make_tiled_copy_tv( + copy_atom_sfd_col_quant, thr_layout, val_layout + ) + tidx = cute.arch.thread_idx()[0] + thr_copy_sfd_col_quant = tiled_copy_sfd_col_quant.get_slice(tidx) + tCgSFDCol_mnl = thr_copy_sfd_col_quant.partition_D( + cute.filter_zeros(gSFDCol_mnl_new) + ) + tCgSFDCol_mnl = cute.filter_zeros(tCgSFDCol_mnl) + return tCgSFDCol_mnl + + + @cute.jit + def geglu_act(self, tCompute: cute.Tensor, acc_vec_up: cute.Tensor, acc_vec_gate: cute.Tensor, mProb: cute.Tensor, linear_offset: cutlass.Float32 = 1.0): + if cutlass.const_expr(self.vectorized_f32): + # GeGlu Packed Version + LOG2_E = cutlass.Float32(1.4426950408889634) + for i in cutlass.range_constexpr(0, cute.size(tCompute), 2): + + (scaled_gate_0, scaled_gate_1) = cute.arch.mul_packed_f32x2( + (acc_vec_gate[i], acc_vec_gate[i + 1]), + (1.702, 1.702), + rnd='rn', + ftz=False, + ) + + tCompute_log2e = cute.arch.mul_packed_f32x2( + (scaled_gate_0, scaled_gate_1), + (-LOG2_E, -LOG2_E), + rnd='rn', + ftz=False, + ) + + ( + tCompute[i], + tCompute[i + 1], + ) = cute.arch.add_packed_f32x2( + ( + cute.math.exp2(tCompute_log2e[0], fastmath=True), + cute.math.exp2(tCompute_log2e[1], fastmath=True), + ), + (1.0, 1.0), + ) + + tCompute[i] = cute.arch.rcp_approx(tCompute[i]) + tCompute[i + 1] = cute.arch.rcp_approx(tCompute[i + 1]) + ( + tCompute[i], + tCompute[i + 1], + ) = cute.arch.mul_packed_f32x2( + (tCompute[i], tCompute[i + 1]), + (acc_vec_gate[i + 0], acc_vec_gate[i + 1]), + rnd='rn', + ftz=False, + ) + ( + up_with_offset0, + up_with_offset1, + ) = cute.arch.add_packed_f32x2( + (linear_offset, linear_offset), + (acc_vec_up[i + 0], acc_vec_up[i + 1]), + rnd='rn', + ftz=False, + ) + ( + tCompute[i], + tCompute[i + 1], + ) = cute.arch.mul_packed_f32x2( + (tCompute[i], tCompute[i + 1]), + (up_with_offset0, up_with_offset1), + rnd='rn', + ftz=False, + ) + if cutlass.const_expr(self.has_prob): + ( + tCompute[i], + tCompute[i + 1], + ) = cute.arch.mul_packed_f32x2( + (tCompute[i], tCompute[i + 1]), + (mProb, mProb), + rnd='rn', + ftz=False, + ) + else: + # GeGlu Unpacked Version + for i in cutlass.range_constexpr(cute.size(tCompute)): + tCompute[i] = (acc_vec_up[i] + linear_offset) * silu_f32_geglu_scaled( + acc_vec_gate[i], fastmath=True + ) + if cutlass.const_expr(self.has_prob): + tCompute[i] = tCompute[i] * mProb + + @cute.jit + def swiglu_act(self, tCompute: cute.Tensor, acc_vec_up: cute.Tensor, acc_vec_gate: cute.Tensor, mProb: cute.Tensor): + if cutlass.const_expr(self.vectorized_f32): + # SwiGlu Packed Version + LOG2_E = cutlass.Float32(1.4426950408889634) + for i in cutlass.range_constexpr(0, cute.size(tCompute), 2): + tCompute_log2e = cute.arch.mul_packed_f32x2( + (acc_vec_gate[i], acc_vec_gate[i + 1]), + (-LOG2_E, -LOG2_E), + rnd='rn', + ftz=False, + ) + ( + tCompute[i], + tCompute[i + 1], + ) = cute.arch.add_packed_f32x2( + ( + cute.math.exp2(tCompute_log2e[0], fastmath=True), + cute.math.exp2(tCompute_log2e[1], fastmath=True), + ), + (1.0, 1.0), + ) + tCompute[i] = cute.arch.rcp_approx(tCompute[i]) + tCompute[i + 1] = cute.arch.rcp_approx(tCompute[i + 1]) + ( + tCompute[i], + tCompute[i + 1], + ) = cute.arch.mul_packed_f32x2( + (tCompute[i], tCompute[i + 1]), + (acc_vec_gate[i + 0], acc_vec_gate[i + 1]), + rnd='rn', + ftz=False, + ) + ( + tCompute[i], + tCompute[i + 1], + ) = cute.arch.mul_packed_f32x2( + (tCompute[i], tCompute[i + 1]), + (acc_vec_up[i], acc_vec_up[i + 1]), + rnd='rn', + ftz=False, + ) + if cutlass.const_expr(self.has_prob): + ( + tCompute[i], + tCompute[i + 1], + ) = cute.arch.mul_packed_f32x2( + (tCompute[i], tCompute[i + 1]), + (mProb, mProb), + rnd='rn', + ftz=False, + ) + else: + # SwiGlu Unpacked Version + for i in cutlass.range_constexpr(cute.size(tCompute)): + tCompute[i] = acc_vec_up[i] * silu_f32( + acc_vec_gate[i], fastmath=True + ) + if cutlass.const_expr(self.has_prob): + tCompute[i] = tCompute[i] * mProb + + + # GPU device kernel + @cute.kernel + def kernel( + self, + tiled_mma: cute.TiledMma, + tiled_mma_bkeep: Optional[cute.TiledMma], + tiled_mma_breuse: Optional[cute.TiledMma], + tiled_mma_sfb: cute.TiledMma, + tma_atom_a: cute.CopyAtom, + mA_mkl: cute.Tensor, + tma_atom_b: cute.CopyAtom, + mB_nkl: cute.Tensor, + tma_atom_sfa: cute.CopyAtom, + mSFA_mkl: cute.Tensor, + tma_atom_sfb: cute.CopyAtom, + mSFB_nkl: cute.Tensor, + tma_atom_c: cute.CopyAtom, + mC_mnl: cute.Tensor, + tma_atom_d: cute.CopyAtom, + mD_mnl: cute.Tensor, + tma_atom_d_col: cute.CopyAtom, + mD_col_mnl: cute.Tensor, + mSFDRow_mnl: Optional[cute.Tensor], + mSFDCol_mnl: Optional[cute.Tensor], + norm_const_tensor: Optional[cute.Tensor], + mAmax_tensor: Optional[cute.Tensor], + padded_offsets: cute.Tensor, + alpha: cute.Tensor, + mBias_nl: Optional[cute.Tensor], + prob: cute.Tensor, + workspace_ptr, # Pointer to TMA descriptor workspace (from desc_init_kernel) + cluster_layout_vmnk: cute.Layout, + cluster_layout_sfb_vmnk: cute.Layout, + a_smem_layout_staged: cute.ComposedLayout, + b_smem_layout_staged: cute.ComposedLayout, + sfa_smem_layout_staged: cute.Layout, + sfb_smem_layout_staged: cute.Layout, + c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout, None], + d_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout, None], + bias_smem_layout_staged: cute.Layout, + epi_tile: cute.Tile, + sched_params: MoESchedulerParams, + epilogue_op: cutlass.Constexpr, + linear_offset: cutlass.Float32 = 0.0, + ): + """ + GPU device kernel performing the Persistent batched GEMM computation. + """ + warp_idx = cute.arch.warp_idx() + warp_idx = cute.arch.make_warp_uniform(warp_idx) + lane_idx = cute.arch.lane_idx() + + # + # Prefetch tma desc + # + if warp_idx == self.tma_warp_id: + cpasync.prefetch_descriptor(tma_atom_a) + cpasync.prefetch_descriptor(tma_atom_sfa) + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DENSE): + cpasync.prefetch_descriptor(tma_atom_b) + cpasync.prefetch_descriptor(tma_atom_sfb) + cpasync.prefetch_descriptor(tma_atom_c) + cpasync.prefetch_descriptor(tma_atom_d) + if cutlass.const_expr(self.generate_sfd): + cpasync.prefetch_descriptor(tma_atom_d_col) + + use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2 + total_token = padded_offsets[self.expert_cnt - 1] + + # + # Setup cta/thread coordinates + # + bidx, bidy, bidz = cute.arch.block_idx() + mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape) + is_leader_cta = mma_tile_coord_v == 0 + cta_rank_in_cluster = cute.arch.make_warp_uniform( + cute.arch.block_idx_in_cluster() + ) + block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord( + cta_rank_in_cluster + ) + + block_in_cluster_coord_sfb_vmnk = cluster_layout_sfb_vmnk.get_flat_coord( + cta_rank_in_cluster + ) + + # Coord inside cta + tidx, _, _ = cute.arch.thread_idx() + + # + # Alloc and init: a+b full/empty, accumulator full/empty, tensor memory dealloc barrier + # + smem = utils.SmemAllocator() + storage = smem.allocate(self.shared_storage) + sched_storage = storage.scheduler + + # Initialize mainloop ab_pipeline (barrier) and states + ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1 + ab_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, num_tma_producer + ) + ab_pipeline = pipeline.PipelineTmaUmma.create( + barrier_storage=storage.ab_mbar_ptr.data_ptr(), + num_stages=self.num_ab_stage, + producer_group=ab_pipeline_producer_group, + consumer_group=ab_pipeline_consumer_group, + tx_count=self.num_tma_load_bytes, + cta_layout_vmnk=cluster_layout_vmnk, + ) + + # Initialize acc_pipeline (barrier) and states + acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + num_acc_consumer_threads = len(self.epilog_warp_id) * ( + 2 if use_2cta_instrs else 1 + ) + acc_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, num_acc_consumer_threads + ) + acc_pipeline = pipeline.PipelineUmmaAsync.create( + barrier_storage=storage.acc_mbar_ptr.data_ptr(), + num_stages=self.num_acc_stage, + producer_group=acc_pipeline_producer_group, + consumer_group=acc_pipeline_consumer_group, + cta_layout_vmnk=cluster_layout_vmnk, + ) + + # Initialize tile info pipeline (barrier) and states + tile_info_pipeline_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + self.threads_per_warp * 1, + ) + tile_info_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + self.threads_wo_sched, + ) + tile_info_pipeline = pipeline.PipelineAsync.create( + barrier_storage=sched_storage.tile_info_mbar.data_ptr(), + num_stages=self.num_tile_stage, + producer_group=tile_info_pipeline_producer_group, + consumer_group=tile_info_pipeline_consumer_group, + ) + + scheduler = MoEPersistentTileScheduler.create( + sched_params, + padded_offsets, + cute.arch.block_idx(), + cute.arch.grid_dim(), + counter_ptr=self._get_sched_counter_ptr(workspace_ptr), + sched_storage=sched_storage, + ) + scheduler.internal_init() + + # Bias pipeline + SMEM + if cutlass.const_expr(self.enable_bias): + bias_pipeline_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, self.threads_per_warp, + ) + bias_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + self.threads_per_warp * len(self.epilog_warp_id), + ) + bias_pipeline = pipeline.PipelineCpAsync.create( + barrier_storage=storage.bias_mbar_ptr.data_ptr(), + num_stages=self.num_bias_stage, + producer_group=bias_pipeline_producer_group, + consumer_group=bias_pipeline_consumer_group, + ) + sBias = storage.sBias.get_tensor(bias_smem_layout_staged) + # (MMA_N, loopN, loopL) + gBias_nl = cute.local_tile( + mBias_nl, cute.slice_(self.mma_tiler[:2], (0, None)), (None, None) + ) + + # Tensor memory dealloc barrier init + tmem = utils.TmemAllocator( + storage.tmem_holding_buf.ptr, + barrier_for_retrieve=self.tmem_alloc_barrier, + allocator_warp_id=self.epilog_warp_id[0], + is_two_cta=use_2cta_instrs, + two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr.ptr, + arch="sm_107", + ) + + # Cluster arrive after barrier init + if cute.size(self.cluster_shape_mn) > 1: + cute.arch.cluster_arrive_relaxed() + + # + # Setup smem tensor A/B/D/Scale + # + # (EPI_TILE_M, EPI_TILE_N, STAGE) + sC = storage.sC.get_tensor( + c_smem_layout_staged.outer, swizzle=c_smem_layout_staged.inner + ) + sD = storage.sD.get_tensor( + d_smem_layout_staged.outer, swizzle=d_smem_layout_staged.inner + ) + # (EPI_TILE_M, EPI_TILE_N, STAGE) + sD_col = sD + if cutlass.const_expr(self.generate_sfd): + sD_col = storage.sD_col.get_tensor( + d_smem_layout_staged.outer, swizzle=d_smem_layout_staged.inner + ) + # (MMA, MMA_M, MMA_K, STAGE) + sA = storage.sA.get_tensor( + a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner + ) + # (MMA, MMA_N, MMA_K, STAGE) + sB = storage.sB.get_tensor( + b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner + ) + # (granularity_m, repeat_m), (granularity_k, repeat_k), num_scale_stage) + sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged) + # (granularity_n, repeat_n), (granularity_k, repeat_k), num_scale_stage) + sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged) + # Shared memory for amax reduction (one FP32 per epilogue warp) + amax_layout = cute.make_layout((self.num_epilog_warps,)) + sAmax = storage.sAmax.get_tensor(amax_layout) + # (expert_idx, tile_m_idx, tile_n_idx, k_tile_cnt) + info_layout = cute.make_layout((4, self.num_tile_stage), stride=(1, 4)) + sInfo = sched_storage.sInfo.get_tensor(info_layout) + + # + # Compute multicast mask for A/B buffer full + # + a_full_mcast_mask = None + b_full_mcast_mask = None + sfa_full_mcast_mask = None + sfb_full_mcast_mask = None + if cutlass.const_expr(self.is_a_mcast or self.is_b_mcast or use_2cta_instrs): + a_full_mcast_mask = cpasync.create_tma_multicast_mask( + cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2 + ) + b_full_mcast_mask = cpasync.create_tma_multicast_mask( + cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1 + ) + sfa_full_mcast_mask = cpasync.create_tma_multicast_mask( + cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2 + ) + sfb_full_mcast_mask = cpasync.create_tma_multicast_mask( + cluster_layout_sfb_vmnk, block_in_cluster_coord_sfb_vmnk, mcast_mode=1 + ) + + # + # Partition shared/tensor memory tensor for TiledMMA_A/B/D + # + # (MMA, MMA_M, MMA_K, STAGE) + tCrA = tiled_mma.make_fragment_A(sA) + # (MMA, MMA_N, MMA_K, STAGE) + tCrB = tiled_mma.make_fragment_B(sB) + # (MMA, MMA_M, MMA_N) + acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2]) + # (MMA, MMA_M, MMA_N, STAGE) + if cutlass.const_expr(self.overlapping_accum): + num_acc_stage_overlapped = 2 + tCtAcc_fake = tiled_mma.make_fragment_C( + cute.append(acc_shape, num_acc_stage_overlapped) + ) + tCtAcc_fake = cute.make_tensor( + tCtAcc_fake.iterator, + cute.make_layout( + tCtAcc_fake.shape, + stride=( + tCtAcc_fake.stride[0], + tCtAcc_fake.stride[1], + tCtAcc_fake.stride[2], + (256 - self.num_sf_tmem_cols) * tCtAcc_fake.stride[0][1], + ), + ), + ) + elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192): + tCtAcc_fake = tiled_mma.make_fragment_C( + cute.append(acc_shape, self.num_acc_stage) + ) + tCtAcc_fake = cute.make_tensor( + tCtAcc_fake.iterator, + cute.make_layout( + tCtAcc_fake.shape, + stride=( + tCtAcc_fake.stride[0], + tCtAcc_fake.stride[1], + tCtAcc_fake.stride[2], + self.num_accumulator_tmem_stride, + ), + ), + ) + else: + tCtAcc_fake = tiled_mma.make_fragment_C( + cute.append(acc_shape, self.num_acc_stage) + ) + + # + # Cluster wait before tensor memory alloc + # + if cute.size(self.cluster_shape_mn) > 1: + cute.arch.cluster_wait() + else: + self.cta_sync_barrier.arrive_and_wait() + + if total_token <= 0: + cute.arch.nvvm.exit() + + # + # Specialized Schedule warp (MoE Persistent Tile Scheduler) + # + if warp_idx == self.sched_warp_id: + work_tile_info = scheduler.initial_work_tile_info() + + tile_info_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_tile_stage + ) + + while work_tile_info.is_valid_tile: + tile_info_pipeline.producer_acquire(tile_info_producer_state) + with cute.arch.elect_one(): + sInfo[(0, tile_info_producer_state.index)] = work_tile_info.expert_idx + sInfo[(1, tile_info_producer_state.index)] = work_tile_info.tile_m_idx + sInfo[(2, tile_info_producer_state.index)] = work_tile_info.tile_n_idx + sInfo[(3, tile_info_producer_state.index)] = work_tile_info.k_tile_cnt + cute.arch.fence_proxy("async.shared", space="cta") + + self.sched_sync_barrier.arrive_and_wait() + tile_info_pipeline.producer_commit(tile_info_producer_state) + tile_info_producer_state.advance() + + work_tile_info = scheduler.advance_to_next_work() + + # Send invalid tile signal: expert_idx = -1 + tile_info_pipeline.producer_acquire(tile_info_producer_state) + with cute.arch.elect_one(): + sInfo[(0, tile_info_producer_state.index)] = cutlass.Int32(-1) + sInfo[(1, tile_info_producer_state.index)] = cutlass.Int32(0) + sInfo[(2, tile_info_producer_state.index)] = cutlass.Int32(0) + sInfo[(3, tile_info_producer_state.index)] = cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + self.sched_sync_barrier.arrive_and_wait() + tile_info_pipeline.producer_commit(tile_info_producer_state) + tile_info_producer_state.advance() + tile_info_pipeline.producer_tail(tile_info_producer_state) + + # + # Specialized TMA load warp + # + if warp_idx == self.tma_warp_id: + ext = self._make_extension(workspace_ptr) + + ab_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_ab_stage + ) + + tile_info_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_tile_stage + ) + + # Get the first tile info + tile_info = cute.make_rmem_tensor((4,), cutlass.Int32) + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + while is_valid_tile: + work_tile_info = MoEWorkTileInfo( + expert_idx=tile_info[0], + tile_m_idx=tile_info[1], + tile_n_idx=tile_info[2], + k_tile_cnt=tile_info[3], + ) + k_tile_cnt = work_tile_info.k_tile_cnt + ext.update_expert_info(padded_offsets, work_tile_info.expert_idx) + + # Get per-expert real tensors + TMA desc ptrs via extension + real_a, _ = ext.get_gmem_tensor( + "a", mA_mkl, padded_offsets, work_tile_info + ) + real_b, desc_ptr_b = ext.get_gmem_tensor( + "b", mB_nkl, padded_offsets, work_tile_info + ) + real_sfa, _ = ext.get_gmem_tensor( + "sfa", mSFA_mkl, padded_offsets, work_tile_info + ) + real_sfb, desc_ptr_sfb = ext.get_gmem_tensor( + "sfb", mSFB_nkl, padded_offsets, work_tile_info + ) + + # N=192: reshape SFB gmem layout to nested ((2,2),y) form for + # correct TMEM mapping (matches the SFB TMEM pointer offset at + # cta_tile_shape_mnk[1]==192 below). + if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192): + x = real_sfb.stride[0][1] + y = cute.ceil_div(real_sfb.shape[0][1], 4) + new_shape = ( + (real_sfb.shape[0][0], ((2, 2), y)), + real_sfb.shape[1], real_sfb.shape[2], + ) + x_times_3 = 3 * x + new_stride = ( + (real_sfb.stride[0][0], ((x, x), x_times_3)), + real_sfb.stride[1], real_sfb.stride[2], + ) + real_sfb = cute.make_tensor( + real_sfb.iterator, + cute.make_layout(new_shape, stride=new_stride), + ) + + # local_tile on per-expert tensors + gA_mkl = cute.local_tile( + real_a, + cute.slice_(self.mma_tiler, (None, 0, None)), + (None, None, None), + ) + gB_nkl = cute.local_tile( + real_b, + cute.slice_(self.mma_tiler, (0, None, None)), + (None, None, None), + ) + gSFA_mkl = cute.local_tile( + real_sfa, + cute.slice_(self.mma_tiler, (None, 0, None)), + (None, None, None), + ) + gSFB_nkl = cute.local_tile( + real_sfb, + cute.slice_(self.mma_tiler_sfb, (0, None, None)), + (None, None, None), + ) + + # MMA partition + thr_mma = tiled_mma.get_slice(mma_tile_coord_v) + thr_mma_sfb = tiled_mma_sfb.get_slice(mma_tile_coord_v) + tCgA = thr_mma.partition_A(gA_mkl) + tCgB = thr_mma.partition_B(gB_nkl) + tCgSFA = thr_mma.partition_A(gSFA_mkl) + tCgSFB = thr_mma_sfb.partition_B(gSFB_nkl) + + # TMA partition A + a_cta_layout = cute.make_layout( + cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape + ) + tAsA, tAgA = cpasync.tma_partition( + tma_atom_a, + block_in_cluster_coord_vmnk[2], + a_cta_layout, + cute.group_modes(sA, 0, 3), + cute.group_modes(tCgA, 0, 3), + ) + # TMA partition B + b_cta_layout = cute.make_layout( + cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape + ) + tBsB, tBgB = cpasync.tma_partition( + tma_atom_b, + block_in_cluster_coord_vmnk[1], + b_cta_layout, + cute.group_modes(sB, 0, 3), + cute.group_modes(tCgB, 0, 3), + ) + # TMA partition SFA + sfa_cta_layout = a_cta_layout + tAsSFA, tAgSFA = cpasync.tma_partition( + tma_atom_sfa, + block_in_cluster_coord_vmnk[2], + sfa_cta_layout, + cute.group_modes(sSFA, 0, 3), + cute.group_modes(tCgSFA, 0, 3), + ) + tAsSFA = cute.filter_zeros(tAsSFA) + tAgSFA = cute.filter_zeros(tAgSFA) + # TMA partition SFB + sfb_cta_layout = cute.make_layout( + cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape + ) + tBsSFB, tBgSFB = cpasync.tma_partition( + tma_atom_sfb, + block_in_cluster_coord_sfb_vmnk[1], + sfb_cta_layout, + cute.group_modes(sSFB, 0, 3), + cute.group_modes(tCgSFB, 0, 3), + ) + tBsSFB = cute.filter_zeros(tBsSFB) + tBgSFB = cute.filter_zeros(tBgSFB) + + # Convert CTA tile index to MMA tile index + mma_tile_coord_m = work_tile_info.tile_m_idx // cute.size(tiled_mma.thr_id.shape) + mma_tile_coord_n = work_tile_info.tile_n_idx + tAgA_slice = tAgA[(None, mma_tile_coord_m, None, 0)] + tBgB_slice = tBgB[(None, mma_tile_coord_n, None, 0)] + tAgSFA_slice = tAgSFA[(None, mma_tile_coord_m, None, 0)] + slice_n = mma_tile_coord_n + if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64): + slice_n = mma_tile_coord_n // 2 + tBgSFB_slice = tBgSFB[(None, slice_n, None, 0)] + + # Peek (try_wait) AB buffer empty + ab_producer_state.reset_count() + peek_ab_empty_status = cutlass.Boolean(1) + if ab_producer_state.count < k_tile_cnt: + peek_ab_empty_status = ab_pipeline.producer_try_acquire( + ab_producer_state + ) + + # + # Tma load loop + # + for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1): + tAgA_k = tAgA_slice[(None, ab_producer_state.count)] + tBgB_k = tBgB_slice[(None, ab_producer_state.count)] + tAgSFA_k = tAgSFA_slice[(None, ab_producer_state.count)] + tBgSFB_k = tBgSFB_slice[(None, ab_producer_state.count)] + tAsA_pipe = tAsA[(None, ab_producer_state.index)] + tBsB_pipe = tBsB[(None, ab_producer_state.index)] + tAsSFA_pipe = tAsSFA[(None, ab_producer_state.index)] + tBsSFB_pipe = tBsSFB[(None, ab_producer_state.index)] + + tma_bar = ab_pipeline.producer_get_barrier(ab_producer_state) + + # Conditionally wait for AB buffer empty + ab_pipeline.producer_acquire( + ab_producer_state, peek_ab_empty_status + ) + ab_producer_state_next = ab_producer_state.clone() + ab_producer_state_next.advance() + if ab_producer_state_next.count < k_tile_cnt: + peek_ab_empty_status = ab_pipeline.producer_try_acquire( + ab_producer_state_next + ) + else: + peek_ab_empty_status = cutlass.Boolean(1) + + # TMA load A (contiguous, global desc via domain_offset) + cute.copy( + tma_atom_a, + tAgA_k, + tAsA_pipe, + tma_bar_ptr=tma_bar, + mcast_mask=a_full_mcast_mask, + ) + # TMA load B (discrete, per-expert desc from workspace) + cute.copy( + tma_atom_b, + tBgB_k, + tBsB_pipe, + tma_bar_ptr=tma_bar, + mcast_mask=b_full_mcast_mask, + tma_desc_ptr=desc_ptr_b, + ) + # TMA load SFA (contiguous, global desc via domain_offset) + cute.copy( + tma_atom_sfa, + tAgSFA_k, + tAsSFA_pipe, + tma_bar_ptr=tma_bar, + mcast_mask=sfa_full_mcast_mask, + ) + # TMA load SFB (discrete, per-expert desc from workspace) + cute.copy( + tma_atom_sfb, + tBgSFB_k, + tBsSFB_pipe, + tma_bar_ptr=tma_bar, + mcast_mask=sfb_full_mcast_mask, + tma_desc_ptr=desc_ptr_sfb, + ) + # Peek (try_wait) AB buffer empty for next k_tile + ab_producer_state.advance() + + # + # Advance to next tile + # + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + # + # Wait A/B buffer empty + # + ab_pipeline.producer_tail(ab_producer_state) + + # + # Specialized MMA warp + # + if warp_idx == self.mma_warp_id: + # + # Bar sync for retrieve tensor memory ptr from shared mem + # + tmem.wait_for_alloc() + # + # Retrieving tensor memory ptr and make accumulator tensor + # + acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype) + # (MMA, MMA_M, MMA_N, STAGE) + tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout) + + # Make SFA tmem tensor + sfa_tmem_ptr = cute.recast_ptr( + acc_tmem_ptr + self.num_accumulator_tmem_cols, + dtype=self.sf_dtype, + ) + # (MMA, MMA_M, MMA_K) + tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa( + tiled_mma, + self.mma_tiler, + self.sf_vec_size, + cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)), + ) + tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout) + + # Make SFB tmem tensor + sfb_tmem_ptr = cute.recast_ptr( + acc_tmem_ptr + self.num_accumulator_tmem_cols + self.num_sfa_tmem_cols, + dtype=self.sf_dtype, + ) + # (MMA, MMA_N, MMA_K) + tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb( + tiled_mma, + self.mma_tiler, + self.sf_vec_size, + cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)), + ) + tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout) + + # Partition for S2T copy of SFA/SFB + # + ( + tiled_copy_s2t_sfa, + tCsSFA_compact_s2t, + tCtSFA_compact_s2t, + ) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA) + ( + tiled_copy_s2t_sfb, + tCsSFB_compact_s2t, + tCtSFB_compact_s2t, + ) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB) + + ab_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_ab_stage + ) + acc_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_acc_stage + ) + + tile_info_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_tile_stage + ) + + # Get the first tile info from pipeline (scheduler has filtered out tiles >= num_non_exiting_tiles) + tile_info = cute.make_rmem_tensor((4,), cutlass.Int32) + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + while is_valid_tile: + k_tile_cnt = tile_info[3] + + # Peek (try_wait) AB buffer full for k_tile = 0 + ab_consumer_state.reset_count() + peek_ab_full_status = cutlass.Boolean(1) + if ab_consumer_state.count < k_tile_cnt and is_leader_cta: + peek_ab_full_status = ab_pipeline.consumer_try_wait( + ab_consumer_state + ) + + # Peek (try_wait) Acc buffer empty for k_tile = 0 + acc_producer_state.reset_count() + peek_acc_empty_status = cutlass.Boolean(1) + if ab_consumer_state.count < k_tile_cnt and is_leader_cta: + peek_acc_empty_status = acc_pipeline.producer_try_acquire( + acc_producer_state + ) + + # Convert CTA tile index to MMA tile index + mma_tile_coord_mnl = ( + tile_info[1] // cute.size(tiled_mma.thr_id.shape), + tile_info[2], # tile_n_idx + tile_info[0], # expert_idx + ) + + # Get accumulator stage index + if cutlass.const_expr(self.overlapping_accum): + acc_stage_index = acc_producer_state.phase ^ 1 + else: + acc_stage_index = acc_producer_state.index + + tCtAcc = tCtAcc_base[(None, None, None, acc_stage_index)] + tCtSFB_mma = tCtSFB + if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192): + # If an ODD tile, shift the TMEM start address for cta_tile_shape_n=192 case by two words + offset = ( + cutlass.Int32(2) + if mma_tile_coord_mnl[1] % 2 == 1 + else cutlass.Int32(0) + ) + shifted_ptr = cute.recast_ptr( + acc_tmem_ptr + + self.num_accumulator_tmem_cols + + self.num_sfa_tmem_cols + + offset, + dtype=self.sf_dtype, + ) + tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout) + elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64): + # Move in increments of 64 columns of SFB + offset = cutlass.Int32((mma_tile_coord_mnl[1] % 2) * 2) + shifted_ptr = cute.recast_ptr( + acc_tmem_ptr + + self.num_accumulator_tmem_cols + + self.num_sfa_tmem_cols + + offset, + dtype=self.sf_dtype, + ) + tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout) + # + # Wait for accumulator buffer empty + # + if is_leader_cta: + acc_pipeline.producer_acquire( + acc_producer_state, peek_acc_empty_status + ) + # + # Mma mainloop + # + + # + # Reset the ACCUMULATE field for each tile + # + tiled_mma.set(tcgen05.Field.ACCUMULATE, False) + + for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1): + # Set tensor memory buffer for current tile + # (MMA, MMA_M, MMA_N) + + if is_leader_cta: + # Conditionally wait for AB buffer full + ab_pipeline.consumer_wait( + ab_consumer_state, peek_ab_full_status + ) + + # Copy SFA/SFB from smem to tmem + s2t_stage_coord = ( + None, + None, + None, + None, + ab_consumer_state.index, + ) + tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[s2t_stage_coord] + tCsSFB_compact_s2t_staged = tCsSFB_compact_s2t[s2t_stage_coord] + cute.copy( + tiled_copy_s2t_sfa, + tCsSFA_compact_s2t_staged, + tCtSFA_compact_s2t, + ) + cute.copy( + tiled_copy_s2t_sfb, + tCsSFB_compact_s2t_staged, + tCtSFB_compact_s2t, + ) + + # tCtAcc += tCrA * tCrSFA * tCrB * tCrSFB + num_kblocks = cute.size(tCrA, mode=[2]) + ab_consumer_state_next = ab_consumer_state.clone() + ab_consumer_state_next.advance() + if ab_consumer_state_next.count < k_tile_cnt: + peek_ab_full_status = ab_pipeline.consumer_try_wait( + ab_consumer_state_next + ) + + + for kblock_idx in cutlass.range(num_kblocks, unroll_full=True): + if cutlass.const_expr( + self.enable_breuse + and cute.size(tCtAcc.layout, mode=[1]) == 2 + and cute.size(tCtAcc.layout, mode=[2]) == 1 + ): + tCtAcc_bkeep = tCtAcc[(None, 0, 0)] + tCtAcc_breuse = tCtAcc[(None, 1, 0)] + a_kblk_crd_keep = (None, 0, kblock_idx, ab_consumer_state.index) + a_kblk_crd_reuse = (None, 1, kblock_idx, ab_consumer_state.index) + b_kblk_crd = (None, 0, kblock_idx, ab_consumer_state.index) + sfa_kblk_crd_keep = (None, 0, kblock_idx) + sfa_kblk_crd_reuse = (None, 1, kblock_idx) + sfb_kblk_crd = (None, 0, kblock_idx) + # Bkeep: accumulate A_upper × B (keeps B in reuse buffer) + tiled_mma_bkeep.set(tcgen05.Field.ACCUMULATE, k_tile != 0 or kblock_idx != 0) + cute.gemm(tiled_mma_bkeep, tCtAcc_bkeep, + [tCrA[a_kblk_crd_keep], tCtSFA[sfa_kblk_crd_keep]], + [tCrB[b_kblk_crd], tCtSFB_mma[sfb_kblk_crd]], tCtAcc_bkeep) + # Breuse: accumulate A_lower × B (reuses B) + tiled_mma_breuse.set(tcgen05.Field.ACCUMULATE, k_tile != 0 or kblock_idx != 0) + cute.gemm(tiled_mma_breuse, tCtAcc_breuse, + [tCrA[a_kblk_crd_reuse], tCtSFA[sfa_kblk_crd_reuse]], + [tCrB[b_kblk_crd], tCtSFB_mma[sfb_kblk_crd]], tCtAcc_breuse) + else: + kblock_coord = (None, None, kblock_idx, ab_consumer_state.index) + sf_kblock_coord = (None, None, kblock_idx) + tiled_mma.set(tcgen05.Field.SFA, tCtSFA[sf_kblock_coord].iterator) + tiled_mma.set(tcgen05.Field.SFB, tCtSFB_mma[sf_kblock_coord].iterator) + cute.gemm(tiled_mma, tCtAcc, tCrA[kblock_coord], tCrB[kblock_coord], tCtAcc) + tiled_mma.set(tcgen05.Field.ACCUMULATE, True) + + # Async arrive AB buffer empty + ab_pipeline.consumer_release(ab_consumer_state) + ab_consumer_state = ab_consumer_state_next + + # + # Async arrive accumulator buffer full(each kblock) + # + if is_leader_cta: + acc_pipeline.producer_commit(acc_producer_state) + + # Peek (try_wait) Acc buffer empty for k_tile = k_tile + 1 + acc_producer_state.advance() + if acc_producer_state.count < k_tile_cnt: + if is_leader_cta: + peek_acc_empty_status = acc_pipeline.producer_try_acquire( + acc_producer_state + ) + + # + # Advance to next tile + # + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + # + # Wait for accumulator buffer empty + # + acc_pipeline.producer_tail(acc_producer_state) + + # + # Specialized bias load warp — cp.async 32-bit GMEM→SMEM + # + if cutlass.const_expr(self.enable_bias): + if warp_idx == self.bias_load_warp_id and total_token > 0: + bias_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_bias_stage + ) + tile_info_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_tile_stage + ) + + # 128-bit cp.async: 32 threads × (128/dtype_bits) elements = tile_N per warp + bias_elems_per_thread = 128 // self.bias_dtype.width + bias_g2s_atom = cute.make_copy_atom( + cute.nvgpu.cpasync.CopyG2SOp(), self.bias_dtype, num_bits_per_copy=128, + ) + bias_g2s_tiled = cute.make_tiled_copy_tv( + bias_g2s_atom, cute.make_layout((self.threads_per_warp,)), cute.make_layout((bias_elems_per_thread,)), + ) + thr_bias_g2s = bias_g2s_tiled.get_slice(cute.arch.lane_idx()) + tBs_sBias = thr_bias_g2s.partition_D(sBias) + + # Predicate tensor for bias cp.async + bias_n_total = mBias_nl.shape[0] + tBpBias = cute.make_rmem_tensor(cute.make_layout((1,)), cutlass.Boolean) + + # Get first tile info from pipeline + tile_info = cute.make_rmem_tensor((4,), cutlass.Int32) + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + while is_valid_tile: + bias_producer_state.reset_count() + + # sInfo format: (expert_idx, tile_m_idx, tile_n_idx, k_tile_cnt) + mma_n_coord = tile_info[2] + expert_idx = tile_info[0] + + gBias_tile = gBias_nl[(None, mma_n_coord, expert_idx)] + tBs_gBias = thr_bias_g2s.partition_S(gBias_tile) + + # Predicate: check if this thread's chunk is within N + tBpBias[0] = mma_n_coord * self.mma_tiler[1] + cute.arch.lane_idx() * bias_elems_per_thread < bias_n_total + + bias_pipeline.producer_acquire(bias_producer_state) + cute.copy(bias_g2s_tiled, tBs_gBias[(None, 0)], tBs_sBias[(None, 0, bias_producer_state.index)], pred=tBpBias) + bias_pipeline.producer_commit(bias_producer_state) + bias_producer_state.advance() + + # Get next tile info + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + bias_pipeline.producer_tail(bias_producer_state) + + # + # Specialized epilogue warps + # + if warp_idx < self.mma_warp_id: + # + # Alloc tensor memory buffer + # + tmem.allocate(self.num_tmem_alloc_cols) + + # + # Bar sync for retrieve tensor memory ptr from shared memory + # + tmem.wait_for_alloc() + + # + # Retrieving tensor memory ptr and make accumulator tensor + # + tmem_ptr = tmem.retrieve_ptr(self.acc_dtype) + # (MMA, MMA_M, MMA_N, STAGE) + tCtAcc_base = cute.make_tensor(tmem_ptr, tCtAcc_fake.layout) + + # + # Partition for epilogue (shape-only: use global tensor for invariant setup) + # + epi_tidx = tidx + thr_mma_epi = tiled_mma.get_slice(mma_tile_coord_v) + gD_mnl_shape = cute.local_tile( + mD_mnl, cute.slice_(self.mma_tiler_d, (None, None, 0)), (None, None, None) + ) + tCgD_shape = thr_mma_epi.partition_C(gD_mnl_shape) + + # For breuse: set up bkeep (m_half=0) and breuse (m_half=1) partitions. + # Each half is processed as a normal GLU subtile loop. + ( + tiled_copy_t2r, + tTR_tAcc_base, + tTR_rAcc_gate, + tTR_rAcc_up, + ) = self.epilog_tmem_copy_and_partition( + epi_tidx, tCtAcc_base, tCgD_shape, epi_tile, use_2cta_instrs, m_half=0 + ) + # Default to bk for JIT scoping; overridden when enable_breuse=True + tTR_tAcc_base_br = tTR_tAcc_base + if cutlass.const_expr(self.enable_breuse): + ( + _, + tTR_tAcc_base_br, + _tTR_rAcc_gate_br, + _tTR_rAcc_up_br, + ) = self.epilog_tmem_copy_and_partition( + epi_tidx, tCtAcc_base, tCgD_shape, epi_tile, use_2cta_instrs, m_half=1 + ) + + tTR_rC = cute.make_rmem_tensor(tTR_rAcc_gate.shape, self.c_dtype) + tiled_copy_r2s, tRS_rC, tRS_sC = self.epilog_smem_copy_and_partition( + tiled_copy_t2r, tTR_rC, epi_tidx, sC + ) + + tTR_rD = cute.make_rmem_tensor(tTR_rAcc_gate.shape, self.d_dtype) + tiled_copy_r2s, tRS_rD, tRS_sD = self.epilog_smem_copy_and_partition( + tiled_copy_t2r, tTR_rD, epi_tidx, sD + ) + + tTR_rD_col = cute.make_rmem_tensor(tTR_rAcc_gate.shape, self.d_dtype) + tiled_copy_r2s, tRS_rD_col, tRS_sD_col = ( + self.epilog_smem_copy_and_partition( + tiled_copy_t2r, tTR_rD_col, epi_tidx, sD_col + ) + ) + + epi_ext = self._make_extension(workspace_ptr) + + if cutlass.const_expr(self.generate_sfd): + norm_const = cutlass.Float32(norm_const_tensor[0]) + regPerSubtile = 4 + sfd_row_tile = ( + cute.make_layout(128), + cute.make_layout(32 * regPerSubtile), + ) + # (EPI_TILE_M, EPI_TILE_N, RestM, RestN, RestL) + gSFDRow_mnl = cute.local_tile( + mSFDRow_mnl, sfd_row_tile, (None, None, None) + ) + thr_copy_t2r = tiled_copy_t2r.get_slice(tidx) + # (T2R, T2R_M, T2R_N, RestM, RestN, RestL) + tCgSFDRow_mnl = thr_copy_t2r.partition_D(gSFDRow_mnl) + tCgSFDRow_mnl = cute.filter_zeros(tCgSFDRow_mnl) + # (T2R, T2R_M, T2R_N) + tCrSFDRow = cute.make_rmem_tensor( + tCgSFDRow_mnl[(None, None, None, 0, 0, 0)].layout, self.sf_dtype + ) + tCrSFDRow_pvscale = cute.make_rmem_tensor_like( + tCrSFDRow, cutlass.Float32 + ) + d_rcp_limits = get_dtype_rcp_limits(self.d_dtype) + + # both SFDs are stored in row major mode. + sfd_col_tile = sfd_row_tile + gSFDCol_mnl = cute.local_tile( + mSFDCol_mnl, sfd_col_tile, (None, None, None) + ) + thr_layout = cute.make_ordered_layout((4, 32), order=(1, 0)) + val_layout = cute.make_ordered_layout((1,), order=(0,)) + copy_atom_sfd_col = cute.make_copy_atom( + cute.nvgpu.CopyUniversalOp(), + gSFDCol_mnl.element_type, + num_bits_per_copy=8, + ) + tiled_copy_sfd_col = cute.make_tiled_copy_tv( + copy_atom_sfd_col, thr_layout, val_layout + ) + thr_copy_sfd_col = tiled_copy_sfd_col.get_slice(tidx) + tCgSFDCol_mnl = thr_copy_sfd_col.partition_D( + cute.filter_zeros(gSFDCol_mnl) + ) + tCgSFDCol_mnl = cute.filter_zeros(tCgSFDCol_mnl) + tCrSFDCol = cute.make_rmem_tensor( + tCgSFDRow_mnl[(None, None, None, 0, 0, 0)].shape, self.sf_dtype + ) + tCrSFDCol_pvscale = cute.make_rmem_tensor_like( + tCrSFDRow, cutlass.Float32 + ) + tCrSFDCol_qpvscale_up_fp32 = cute.make_rmem_tensor_like( + tCrSFDRow, cutlass.Float32 + ) + + acc_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_acc_stage + ) + + c_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + 32 * len(self.epilog_warp_id), + ) + c_pipeline = pipeline.PipelineTmaStore.create( + num_stages=self.num_c_stage, + producer_group=c_producer_group, + ) + + d_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + 32 * len(self.epilog_warp_id), + ) + d_pipeline = pipeline.PipelineTmaStore.create( + num_stages=self.num_d_stage, + producer_group=d_producer_group, + ) + d_col_pipeline = pipeline.PipelineTmaStore.create( + num_stages=self.num_d_stage, + producer_group=d_producer_group, + ) + + tile_info_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_tile_stage + ) + + # Get the first tile info + tile_info = cute.make_rmem_tensor((4,), cutlass.Int32) + + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + if cutlass.const_expr(self.enable_bias): + bias_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_bias_stage + ) + bias_s2r_atom = cute.make_copy_atom( + cute.nvgpu.CopyUniversalOp(), self.bias_dtype, num_bits_per_copy=128 + ) + tTR_rBias_gate = cute.make_rmem_tensor(cute.make_layout(self.epi_tile[1]), self.bias_dtype) + tTR_rBias_up = cute.make_rmem_tensor(cute.make_layout(self.epi_tile[1]), self.bias_dtype) + + num_prev_subtiles = cutlass.Int32(0) + while is_valid_tile: + # sInfo format: (expert_idx, tile_m_idx, tile_n_idx, k_tile_cnt) + epi_work_tile_info = MoEWorkTileInfo( + expert_idx=tile_info[0], + tile_m_idx=tile_info[1], + tile_n_idx=tile_info[2], + k_tile_cnt=tile_info[3], + ) + mma_tile_coord_mnl = ( + epi_work_tile_info.tile_m_idx // cute.size(tiled_mma.thr_id.shape), + epi_work_tile_info.tile_n_idx, + cutlass.Int32(0), + ) + + expert_idx = epi_work_tile_info.expert_idx + alpha_val = alpha[expert_idx] + epi_ext.update_expert_info(padded_offsets, epi_work_tile_info.expert_idx) + + if cutlass.const_expr(self.enable_bias): + bias_consumer_state.reset_count() + bias_pipeline.consumer_wait(bias_consumer_state) + sBias_stage = sBias[(None, bias_consumer_state.index)] + sBias_subtiles = cute.flat_divide(sBias_stage, cute.make_layout(2 * self.epi_tile[1])) + + # Get per-expert C/D/D_col tensors via extension + real_c, _ = epi_ext.get_gmem_tensor( + "c", mC_mnl, padded_offsets, epi_work_tile_info + ) + real_d, _ = epi_ext.get_gmem_tensor( + "d", mD_mnl, padded_offsets, epi_work_tile_info + ) + real_d_col = real_d + if cutlass.const_expr(self.generate_sfd): + real_d_col, _ = epi_ext.get_gmem_tensor( + "d_col", mD_col_mnl, padded_offsets, epi_work_tile_info + ) + + # local_tile + partition on per-expert tensors + thr_mma_epi_loop = tiled_mma.get_slice(mma_tile_coord_v) + gC_mnl = cute.local_tile( + real_c, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None) + ) + tCgC = thr_mma_epi_loop.partition_C(gC_mnl) + if cutlass.const_expr(self.enable_breuse): + # Select bkeep/breuse halves by indexing mode1 with 0 or 1 + gC_epi_bk = cute.flat_divide(tCgC[((None, None), 0, 0, None, None, None)], self.epi_tile_c) + gC_epi_br = cute.flat_divide(tCgC[((None, None), 1, 0, None, None, None)], self.epi_tile_c) + sC_for_tma = cute.group_modes(sC, 0, 2) + bSG_sC, bSG_gC_partitioned_bk = cpasync.tma_partition( + tma_atom_c, 0, cute.make_layout(1), sC_for_tma, cute.group_modes(gC_epi_bk, 0, 2)) + _, bSG_gC_partitioned_br = cpasync.tma_partition( + tma_atom_c, 0, cute.make_layout(1), sC_for_tma, cute.group_modes(gC_epi_br, 0, 2)) + else: + _, bSG_sC, bSG_gC_partitioned = epilog_gmem_copy_and_partition( + epi_tidx, tma_atom_c, tCgC, self.epi_tile_c, sC + ) + + gD_mnl_loop = cute.local_tile( + real_d, cute.slice_(self.mma_tiler_d, (None, None, 0)), (None, None, None) + ) + tCgD_loop = thr_mma_epi_loop.partition_C(gD_mnl_loop) + if cutlass.const_expr(self.enable_breuse): + gD_epi_bk = cute.flat_divide(tCgD_loop[((None, None), 0, 0, None, None, None)], epi_tile) + gD_epi_br = cute.flat_divide(tCgD_loop[((None, None), 1, 0, None, None, None)], epi_tile) + sD_for_tma = cute.group_modes(sD, 0, 2) + bSG_sD, bSG_gD_partitioned_bk = cpasync.tma_partition( + tma_atom_d, 0, cute.make_layout(1), sD_for_tma, cute.group_modes(gD_epi_bk, 0, 2)) + _, bSG_gD_partitioned_br = cpasync.tma_partition( + tma_atom_d, 0, cute.make_layout(1), sD_for_tma, cute.group_modes(gD_epi_br, 0, 2)) + else: + _, bSG_sD, bSG_gD_partitioned = epilog_gmem_copy_and_partition( + epi_tidx, tma_atom_d, tCgD_loop, epi_tile, sD) + + gD_col_mnl_loop = gD_mnl_loop + tCgD_col_loop = tCgD_loop + if cutlass.const_expr(self.generate_sfd): + gD_col_mnl_loop = cute.local_tile( + real_d_col, cute.slice_(self.mma_tiler_d, (None, None, 0)), (None, None, None) + ) + tCgD_col_loop = thr_mma_epi_loop.partition_C(gD_col_mnl_loop) + if cutlass.const_expr(self.enable_breuse): + gD_col_epi_bk = cute.flat_divide(tCgD_col_loop[((None, None), 0, 0, None, None, None)], epi_tile) + gD_col_epi_br = cute.flat_divide(tCgD_col_loop[((None, None), 1, 0, None, None, None)], epi_tile) + sD_col_for_tma = cute.group_modes(sD_col, 0, 2) + bSG_sD_col, bSG_gD_col_partitioned_bk = cpasync.tma_partition( + tma_atom_d_col, 0, cute.make_layout(1), sD_col_for_tma, cute.group_modes(gD_col_epi_bk, 0, 2)) + _, bSG_gD_col_partitioned_br = cpasync.tma_partition( + tma_atom_d_col, 0, cute.make_layout(1), sD_col_for_tma, cute.group_modes(gD_col_epi_br, 0, 2)) + else: + _, bSG_sD_col, bSG_gD_col_partitioned = epilog_gmem_copy_and_partition( + epi_tidx, tma_atom_d_col, tCgD_col_loop, epi_tile, sD_col) + + # Slice to per-expert tile coords (L=0, domain already offset'd) + def _grp(p): + q = p[(None, None, None, *mma_tile_coord_mnl)] + return cute.group_modes(q, 1, cute.rank(q)) + + if cutlass.const_expr(self.enable_breuse): + bSG_gC_bk = _grp(bSG_gC_partitioned_bk); bSG_gC_br = _grp(bSG_gC_partitioned_br) + bSG_gD_bk = _grp(bSG_gD_partitioned_bk); bSG_gD_br = _grp(bSG_gD_partitioned_br) + bSG_gD_col_bk = _grp(bSG_gD_col_partitioned_bk); bSG_gD_col_br = _grp(bSG_gD_col_partitioned_br) + # Placeholders so the JIT sees all names in non-breuse paths too + bSG_gC = bSG_gC_bk; bSG_gD = bSG_gD_bk; bSG_gD_col = bSG_gD_col_bk + else: + bSG_gC = _grp(bSG_gC_partitioned) + bSG_gD = _grp(bSG_gD_partitioned) + bSG_gD_col = _grp(bSG_gD_col_partitioned) + bSG_gC_bk = bSG_gC; bSG_gC_br = bSG_gC + bSG_gD_bk = bSG_gD; bSG_gD_br = bSG_gD + bSG_gD_col_bk = bSG_gD_col; bSG_gD_col_br = bSG_gD_col + + # Get accumulator stage index + if cutlass.const_expr(self.overlapping_accum): + acc_stage_index = acc_consumer_state.phase + reverse_subtile = ( + cutlass.Boolean(True) + if acc_stage_index == 0 + else cutlass.Boolean(False) + ) + else: + acc_stage_index = acc_consumer_state.index + + # Set tensor memory buffer for current tile + # (T2R, T2R_M, T2R_N, EPI_M, EPI_M) + tTR_tAcc = tTR_tAcc_base[ + (None, None, None, None, None, acc_stage_index) + ] + + if cutlass.const_expr(self.generate_sfd): + # (T2R, T2R_M, T2R_N, RestM, RestN) + tCgSFDRow_mn = tCgSFDRow_mnl[ + ( + None, + None, + None, + None, + None, + 0, + ) + ] + tCgSFDCol_mnl_new = tCgSFDCol_mnl + if cutlass.const_expr(self.discrete_col_sfd): + tCgSFDCol_mnl_new = self.create_and_partition_new_SFDCol( + tile_info, mSFDCol_mnl, padded_offsets + ) + tCgSFDCol_mn = tCgSFDCol_mnl_new[ + ( + None, + None, + None, + None, + None, + 0, + ) + ] + + if cutlass.const_expr(self.generate_amax): + thread_tile_amax = cutlass.Float32(0.0) + + mPosition_base = ( + (epi_work_tile_info.tile_m_idx // cute.size(tiled_mma.thr_id.shape)) + * self.mma_tiler[0] + + mma_tile_coord_v * (self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape)) + + tidx + ) + mProb = cutlass.Float32(1.0) + mProb_bk = cutlass.Float32(1.0) + mProb_br = cutlass.Float32(1.0) + if cutlass.const_expr(self.has_prob): + real_prob, _ = epi_ext.get_gmem_tensor( + "prob", prob, padded_offsets, epi_work_tile_info + ) + mProb = real_prob[mPosition_base, 0, 0] + mProb_bk = mProb + mProb_br = mProb + if cutlass.const_expr(self.enable_breuse): + mProb_br = real_prob[ + mPosition_base + (self.cta_tile_shape_mnk[0] // 2), 0, 0 + ] + + # + # Wait for accumulator buffer full + # + acc_pipeline.consumer_wait(acc_consumer_state) + + # For breuse: process bkeep (m_half=0) then breuse (m_half=1) M-halves. + # For non-breuse: one pass (m_half=0). + _breuse_halves = [0, 1] if self.enable_breuse else [0] + + _tTR_tAcc_base_h = tTR_tAcc_base + _mProb_h = mProb + _bSG_gC_h = bSG_gC + _bSG_gD_h = bSG_gD + _bSG_gD_col_h = bSG_gD_col + + for _m_half in _breuse_halves: + if self.enable_breuse: + _tTR_tAcc_base_h = tTR_tAcc_base if _m_half == 0 else tTR_tAcc_base_br + _mProb_h = mProb_bk if _m_half == 0 else mProb_br + _bSG_gC_h = bSG_gC_bk if _m_half == 0 else bSG_gC_br + _bSG_gD_h = bSG_gD_bk if _m_half == 0 else bSG_gD_br + _bSG_gD_col_h = bSG_gD_col_bk if _m_half == 0 else bSG_gD_col_br + else: + _tTR_tAcc_base_h = tTR_tAcc_base + _mProb_h = mProb + _bSG_gC_h = bSG_gC + _bSG_gD_h = bSG_gD + _bSG_gD_col_h = bSG_gD_col + + tTR_tAcc_h = _tTR_tAcc_base_h[(None, None, None, None, None, acc_stage_index)] + tTR_tAcc_h = cute.group_modes(tTR_tAcc_h, 3, cute.rank(tTR_tAcc_h)) + + # + # Store accumulator to global memory in subtiles + # + subtile_cnt = cute.size(tTR_tAcc_h.shape, mode=[3]) + for subtile_idx in cutlass.range(0, subtile_cnt, 2, unroll=1): + real_subtile_idx = subtile_idx // 2 + if cutlass.const_expr(self.overlapping_accum): + if reverse_subtile: + real_subtile_idx = ( + self.cta_tile_shape_mnk[1] // self.epi_tile_n_required + - 1 + - subtile_idx // 2 + ) + + # + # Load accumulator from tensor memory buffer to register + # + tTR_tAcc_mn_gate = tTR_tAcc_h[ + (None, None, None, real_subtile_idx * 2) + ] + tTR_tAcc_mn_up = tTR_tAcc_h[ + (None, None, None, real_subtile_idx * 2 + 1) + ] + + cute.copy(tiled_copy_t2r, tTR_tAcc_mn_gate, tTR_rAcc_gate) + cute.copy(tiled_copy_t2r, tTR_tAcc_mn_up, tTR_rAcc_up) + + # + # Async arrive accumulator buffer empty ealier when overlapping_accum is enabled + # + if cutlass.const_expr(self.overlapping_accum): + if subtile_idx == self.iter_acc_early_release_in_epilogue: + # Fence for TMEM load + cute.arch.fence_view_async_tmem_load() + with cute.arch.elect_one(): + acc_pipeline.consumer_release(acc_consumer_state) + acc_consumer_state.advance() + + # + # Apply alpha (+ bias if enabled) + # + if cutlass.const_expr(self.enable_bias): + sBias_sub = sBias_subtiles[(None, real_subtile_idx)] + for i in cutlass.range_constexpr(self.epi_tile[1]): + tTR_rBias_gate[i] = sBias_sub[i] + tTR_rBias_up[i] = sBias_sub[self.epi_tile[1] + i] + bias_vec_gate = tTR_rBias_gate.load() + bias_vec_up = tTR_rBias_up.load() + + if cutlass.const_expr(self.vectorized_f32): + for i in cutlass.range_constexpr(0, cute.size(tTR_rAcc_gate), 2): + bias_gate_f32_0 = bias_vec_gate[i].to(cutlass.Float32) + bias_gate_f32_1 = bias_vec_gate[i + 1].to(cutlass.Float32) + bias_up_f32_0 = bias_vec_up[i].to(cutlass.Float32) + bias_up_f32_1 = bias_vec_up[i + 1].to(cutlass.Float32) + (tTR_rAcc_gate[i], tTR_rAcc_gate[i + 1]) = ( + cute.arch.fma_packed_f32x2( + (tTR_rAcc_gate[i], tTR_rAcc_gate[i + 1]), + (cutlass.Float32(alpha_val), cutlass.Float32(alpha_val)), + (bias_gate_f32_0, bias_gate_f32_1), + rnd='rn', ftz=False, + ) + ) + (tTR_rAcc_up[i], tTR_rAcc_up[i + 1]) = ( + cute.arch.fma_packed_f32x2( + (tTR_rAcc_up[i], tTR_rAcc_up[i + 1]), + (cutlass.Float32(alpha_val), cutlass.Float32(alpha_val)), + (bias_up_f32_0, bias_up_f32_1), + rnd='rn', ftz=False, + ) + ) + else: + for i in cutlass.range_constexpr(cute.size(tTR_rAcc_gate)): + tTR_rAcc_gate[i] = ( + tTR_rAcc_gate[i] * cutlass.Float32(alpha_val) + + bias_vec_gate[i].to(cutlass.Float32) + ) + tTR_rAcc_up[i] = ( + tTR_rAcc_up[i] * cutlass.Float32(alpha_val) + + bias_vec_up[i].to(cutlass.Float32) + ) + + if subtile_idx == subtile_cnt - 2: + bias_pipeline.consumer_release(bias_consumer_state) + bias_consumer_state.advance() + else: + if cutlass.const_expr(self.vectorized_f32): + for i in cutlass.range_constexpr(0, cute.size(tTR_rAcc_gate), 2): + (tTR_rAcc_gate[i], tTR_rAcc_gate[i + 1]) = ( + cute.arch.mul_packed_f32x2( + (tTR_rAcc_gate[i], tTR_rAcc_gate[i + 1]), + (cutlass.Float32(alpha_val), cutlass.Float32(alpha_val)), + rnd='rn', ftz=False, + ) + ) + (tTR_rAcc_up[i], tTR_rAcc_up[i + 1]) = ( + cute.arch.mul_packed_f32x2( + (tTR_rAcc_up[i], tTR_rAcc_up[i + 1]), + (cutlass.Float32(alpha_val), cutlass.Float32(alpha_val)), + rnd='rn', ftz=False, + ) + ) + else: + for i in cutlass.range_constexpr(cute.size(tTR_rAcc_gate)): + tTR_rAcc_gate[i] = tTR_rAcc_gate[i] * cutlass.Float32(alpha_val) + tTR_rAcc_up[i] = tTR_rAcc_up[i] * cutlass.Float32( + alpha_val + ) + + # + # Store to C tensor (optional, only when generate_c=True) + # + if cutlass.const_expr(self.generate_c): + self.store_c( + tiled_copy_r2s, + tma_atom_c, + warp_idx, + tTR_rAcc_gate, + tTR_rAcc_up, + tTR_rC, + tRS_rC, + tRS_sC, + _bSG_gC_h, + bSG_sC, + c_pipeline, + num_prev_subtiles, + real_subtile_idx, + ) + + if cutlass.const_expr(self.act_func == "geglu"): + geglu_max_val = cutlass.Float32(7.0) + geglu_min_val = cutlass.Float32(-7.0) + for i in cutlass.range_constexpr(cute.size(tTR_rAcc_up)): + tTR_rAcc_gate[i] = fmin(tTR_rAcc_gate[i], geglu_max_val) + tTR_rAcc_up[i] = fmin(tTR_rAcc_up[i], geglu_max_val) + tTR_rAcc_up[i] = fmax(tTR_rAcc_up[i], geglu_min_val) + + acc_vec_gate = tTR_rAcc_gate.load() + acc_vec_up = tTR_rAcc_up.load() + + # SwiGlu or GeGLU + tCompute = cute.make_rmem_tensor(acc_vec_gate.shape, self.acc_dtype) + if cutlass.const_expr(self.act_func == "geglu"): + self.geglu_act(tCompute, acc_vec_up, acc_vec_gate, _mProb_h, linear_offset) + elif cutlass.const_expr(self.act_func == "swiglu"): + self.swiglu_act(tCompute, acc_vec_up, acc_vec_gate, _mProb_h) + + # + # Generate amax + # + if cutlass.const_expr(self.generate_amax): + thread_tile_amax = amax_reduction_per_thread( + tCompute, thread_tile_amax + ) + + if cutlass.const_expr(self.generate_sfd): + tCompute_col = cute.make_rmem_tensor( + tCompute.layout, tCompute.element_type + ) + tCompute_col.store(tCompute.load()) + # + # Generate row major SFD + # + self.quant_sfd_row( + real_subtile_idx, + tiled_copy_r2s, + tCompute, + tCrSFDRow_pvscale, + norm_const, + d_rcp_limits, + tRS_rD, + tile_info, + ) + # + # Generate col major SFD + # + self.quant_sfd_col( + real_subtile_idx, + tiled_copy_r2s, + tCompute_col, + tCrSFDCol_pvscale, + norm_const, + d_rcp_limits, + tRS_rD_col, + tile_info, + ) + + # tile_m_idx is CTA-level (like bidx), use directly as raw_tile_m + cta_group_size = cute.size(tiled_mma.thr_id.shape) + raw_tile_m = epi_work_tile_info.tile_m_idx + token_offset_sfd, _ = compute_expert_token_range( + padded_offsets, expert_idx + ) + global_tile_m_offset = token_offset_sfd * cta_group_size // self.mma_tiler[0] + sfd_row_idx_mn = ( + raw_tile_m + global_tile_m_offset, + epi_work_tile_info.tile_n_idx, + ) + sfd_col_idx_mn = sfd_row_idx_mn + if cutlass.const_expr(self.discrete_col_sfd): + sfd_col_idx_mn = ( + raw_tile_m, + epi_work_tile_info.tile_n_idx, + ) + tCgSFDRow = tCgSFDRow_mn[ + ( + None, + None, + None, + *sfd_row_idx_mn, + ) + ] + tCgSFDCol = tCgSFDCol_mn[ + ( + None, + None, + None, + *sfd_col_idx_mn, + ) + ] + + if subtile_idx == 6: + if sfd_row_idx_mn[1] * 32 * regPerSubtile < cute.size(cute.shape(mSFDRow_mnl.layout, mode=[1])): + tCrSFDRow.store(tCrSFDRow_pvscale.load().to(self.sf_dtype)) + cute.autovec_copy(tCrSFDRow, tCgSFDRow) + if sfd_col_idx_mn[1] * 32 * regPerSubtile < cute.size(cute.shape(mSFDCol_mnl.layout, mode=[1])): + tCrSFDCol.store(tCrSFDCol_pvscale.load().to(self.sf_dtype)) + cute.autovec_copy(tCrSFDCol, tCgSFDCol) + else: + # + # Convert to D type + # + acc_vec = tiled_copy_r2s.retile(tCompute).load() + tRS_rD.store(acc_vec.to(self.d_dtype)) + + # + # Store D to shared memory + # + d_buffer = num_prev_subtiles % self.num_d_stage + num_prev_subtiles = num_prev_subtiles + 1 + cute.copy( + tiled_copy_r2s, + tRS_rD, + tRS_sD[(None, None, None, d_buffer)], + ) + if cutlass.const_expr(self.generate_sfd): + cute.copy( + tiled_copy_r2s, + tRS_rD_col, + tRS_sD_col[(None, None, None, d_buffer)], + ) + # Fence and barrier to make sure shared memory store is visible to TMA store + cute.arch.fence_proxy("async.shared", space="cta") + self.epilog_sync_barrier.arrive_and_wait() + # + # TMA store D to global memory + # + if warp_idx == self.epilog_warp_id[0]: + cute.copy( + tma_atom_d, + bSG_sD[(None, d_buffer)], + _bSG_gD_h[(None, real_subtile_idx)], + ) + if cutlass.const_expr(self.generate_sfd): + cute.copy( + tma_atom_d_col, + bSG_sD_col[(None, d_buffer)], + _bSG_gD_col_h[(None, real_subtile_idx)], + ) + # Fence and barrier to make sure shared memory store is visible to TMA store + d_pipeline.producer_commit() + self.epilog_sync_barrier.arrive_and_wait() + + # + # Async arrive accumulator buffer empty + # + if cutlass.const_expr(not self.overlapping_accum): + with cute.arch.elect_one(): + acc_pipeline.consumer_release(acc_consumer_state) + acc_consumer_state.advance() + + # + # Advance to next tile + # + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + # Perform amax reduction after all subtiles are processed + if cutlass.const_expr(self.generate_amax): + gAmax = mAmax_tensor[ + (expert_idx, None) + ].iterator.llvm_ptr # First element + self.amax_reduction_per_warp_and_cta( + thread_tile_amax, warp_idx, sAmax, gAmax + ) + + # + # Dealloc the tensor memory buffer + # + tmem.relinquish_alloc_permit() + self.epilog_sync_barrier.arrive_and_wait() + tmem.free(tmem_ptr) + # + # Wait for C/D store complete + # + c_pipeline.producer_tail() + d_pipeline.producer_tail() + + def epilog_tmem_copy_and_partition( + self, + tidx: cutlass.Int32, + tAcc: cute.Tensor, + gD_mnl: cute.Tensor, + epi_tile: cute.Tile, + use_2cta_instrs: Union[cutlass.Boolean, bool], + m_half: int = 0, + ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor, cute.Tensor]: + """ + Make tiledCopy for tensor memory load, then use it to partition tensor memory (source) and register array (destination). + + :param tidx: The thread index in epilogue warp groups + :type tidx: cutlass.Int32 + :param tAcc: The accumulator tensor to be copied and partitioned + :type tAcc: cute.Tensor + :param gD_mnl: The global tensor D + :type gD_mnl: cute.Tensor + :param epi_tile: The epilogue tiler + :type epi_tile: cute.Tile + :param use_2cta_instrs: Whether use_2cta_instrs is enabled + :type use_2cta_instrs: bool + + :return: A tuple containing (tiled_copy_t2r, tTR_tAcc, tTR_rAcc_gate, tTR_rAcc_up) where: + - tiled_copy_t2r: The tiled copy operation for tmem to register copy(t2r) + - tTR_tAcc: The partitioned accumulator tensor + - tTR_rAcc_gate: The partitioned accumulator tensor for acc gate + - tTR_rAcc_up: The partitioned accumulator tensor for acc up + :rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor, cute.Tensor] + """ + # Make tiledCopy for tensor memory load + copy_atom_t2r = sm100_utils.get_tmem_load_op( + self.cta_tile_shape_mnk, + self.d_layout, + self.d_dtype, + self.acc_dtype, + epi_tile, + use_2cta_instrs, + ) + + # For breuse, select the requested m_half (0=bkeep, 1=breuse) from the split accumulator. + # The m_half parameter is a Python int (compile-time constant), so this resolves at trace time. + if cutlass.const_expr(self.enable_breuse): + tAcc_epi = cute.flat_divide(tAcc[((None, None), m_half, 0, None)], epi_tile) + gD_mnl_epi = cute.flat_divide(gD_mnl[((None, None), m_half, 0, None, None, None)], epi_tile) + else: + tAcc_epi = cute.flat_divide(tAcc[((None, None), 0, 0, None)], epi_tile) + gD_mnl_epi = cute.flat_divide(gD_mnl[((None, None), 0, 0, None, None, None)], epi_tile) + + # (EPI_TILE_M, EPI_TILE_N) + tiled_copy_t2r = tcgen05.make_tmem_copy( + copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)] + ) + + thr_copy_t2r = tiled_copy_t2r.get_slice(tidx) + # (T2R, T2R_M, T2R_N, EPI_M, EPI_M, STAGE) + tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi) + + # (T2R, T2R_M, T2R_N, EPI_M, EPI_N, loopM, loopN, loopL) + tTR_gC = thr_copy_t2r.partition_D(gD_mnl_epi) + + # (T2R, T2R_M, T2R_N) + tTR_rAcc_gate = cute.make_rmem_tensor( + tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype + ) + # (T2R, T2R_M, T2R_N) + tTR_rAcc_up = cute.make_rmem_tensor( + tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype + ) + return tiled_copy_t2r, tTR_tAcc, tTR_rAcc_gate, tTR_rAcc_up + + def epilog_smem_copy_and_partition( + self, + tiled_copy_t2r: cute.TiledCopy, + tTR_rC: cute.Tensor, + tidx: cutlass.Int32, + sD: cute.Tensor, + ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]: + """ + Make tiledCopy for shared memory store, then use it to partition register array (source) and shared memory (destination). + + :param tiled_copy_t2r: The tiled copy operation for tmem to register copy(t2r) + :type tiled_copy_t2r: cute.TiledCopy + :param tTR_rC: The partitioned accumulator tensor + :type tTR_rC: cute.Tensor + :param tidx: The thread index in epilogue warp groups + :type tidx: cutlass.Int32 + :param sD: The shared memory tensor to be copied and partitioned + :type sD: cute.Tensor + + :return: A tuple containing (tiled_copy_r2s, tRS_rD, tRS_sD) where: + - tiled_copy_r2s: The tiled copy operation for register to smem copy(r2s) + - tRS_rD: The partitioned tensor D (register source) + - tRS_sD: The partitioned tensor D (smem destination) + :rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor] + """ + copy_atom_r2s = sm100_utils.get_smem_store_op( + self.d_layout, self.d_dtype, self.acc_dtype, tiled_copy_t2r + ) + tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r) + # (R2S, R2S_M, R2S_N, PIPE_D) + thr_copy_r2s = tiled_copy_r2s.get_slice(tidx) + tRS_sD = thr_copy_r2s.partition_D(sD) + # (R2S, R2S_M, R2S_N) + tRS_rD = tiled_copy_r2s.retile(tTR_rC) + return tiled_copy_r2s, tRS_rD, tRS_sD diff --git a/python/cudnn/grouped_gemm/grouped_gemm_quant/api.py b/python/cudnn/grouped_gemm/grouped_gemm_quant/api.py index 5db03ac6f..01c7377a9 100644 --- a/python/cudnn/grouped_gemm/grouped_gemm_quant/api.py +++ b/python/cudnn/grouped_gemm/grouped_gemm_quant/api.py @@ -43,17 +43,26 @@ from cuda.bindings import driver as cuda from cutlass.cute.runtime import make_fake_stream -from cudnn.api_base import APIBase, TensorDesc, TupleDict, ceil_div, is_power_of_2 +from cudnn.api_base import APIBase, TensorDesc, TupleDict, ceil_div, get_device_type, is_power_of_2 from cudnn.datatypes import _convert_to_cutlass_data_type from .grouped_gemm_quant import ( BlockScaledMoEGroupedGemmQuantKernel, ) from ..moe_utils import MoEWeightMode +from ..grouped_gemm_utils import rubin_single_group_offsets_kwarg from cutlass.cute.nvgpu import OperandMajorMode from cutlass.cute.runtime import from_dlpack +def _get_rubin_kernel(): + from .moe_blockscaled_grouped_gemm_quant_rubin import ( + BlockScaledMoEGroupedGemmQuantKernel as RubinBlockScaledMoEGroupedGemmQuantKernel, + ) + + return RubinBlockScaledMoEGroupedGemmQuantKernel + + class GroupedGemmQuantSm100(APIBase): """Unified API for grouped GEMM quant operation on SM100+ GPUs. @@ -223,7 +232,7 @@ def __init__( self._interpret_uint8_as_fp4x2 = True self._has_bias = self.bias_desc is not None - self._kernel = BlockScaledMoEGroupedGemmQuantKernel + self._kernel = _get_rubin_kernel() if self._is_rubin_kernel else BlockScaledMoEGroupedGemmQuantKernel self.num_cluster_overlap_margin = int(os.getenv("CUDNNFE_CLUSTER_OVERLAP_MARGIN", "0")) self._logger.debug(f"setting num_cluster_overlap_margin: {self.num_cluster_overlap_margin}") @@ -290,7 +299,12 @@ def check_support(self) -> bool: "Pass a tensor of ones with shape (valid_m, 1, 1) if no gating is needed.", ) self._check_tensor_shape(self.prob_desc, (tensor_m, 1, 1), "prob") - self._check_tensor_shape(self.row_scale_desc, (tensor_m,), "row_scale") + self._not_implemented_error_if( + self._is_rubin_kernel and self.row_scale_desc is not None, + "Rubin grouped GEMM quant does not support row_scale fusion", + ) + if not self._is_rubin_kernel: + self._check_tensor_shape(self.row_scale_desc, (tensor_m,), "row_scale") self._check_tensor_shape(self.bias_desc, (n, l), "bias") self._check_tensor_shape(self.amax_desc, (self.expert_cnt, 1), "amax") self._check_tensor_shape(self.norm_const_desc, (1,), "norm_const") @@ -328,11 +342,12 @@ def check_support(self) -> bool: self.bias_desc, stride=[(1, n)], ) - _ = self._check_tensor_stride( - self.row_scale_desc, - stride=[(1,)], - extra_error_msg="row_scale must be a contiguous 1-D tensor", - ) + if not self._is_rubin_kernel: + _ = self._check_tensor_stride( + self.row_scale_desc, + stride=[(1,)], + extra_error_msg="row_scale must be a contiguous 1-D tensor", + ) self._logger.debug("Checking data types") self.ab_dtype = self._check_dtype( @@ -439,12 +454,13 @@ def check_support(self) -> bool: name="D_col", extra_error_msg="D_col must have the same dtype as D", ) - self._check_dtype( - self.row_scale_desc, - dtype=torch.float32, - name="row_scale", - extra_error_msg="row_scale must be float32", - ) + if not self._is_rubin_kernel: + self._check_dtype( + self.row_scale_desc, + dtype=torch.float32, + name="row_scale", + extra_error_msg="row_scale must be float32", + ) self._not_implemented_error_if( self._is_fp4x2(self.ab_dtype) and self.sf_vec_size == 16 and self.d_dtype == torch.float32, @@ -500,8 +516,8 @@ def check_support(self) -> bool: f"Invalid m_aligned: expected m_aligned to be divisible by mma_tiler_mn[0], got {self.m_aligned} % {self.mma_tiler_mn[0]} != 0", ) self._value_error_if( - self.m_aligned != BlockScaledMoEGroupedGemmQuantKernel.FIX_PAD_SIZE, - f"m_aligned must be {BlockScaledMoEGroupedGemmQuantKernel.FIX_PAD_SIZE} (FIX_PAD_SIZE), got {self.m_aligned}", + self.m_aligned != self._kernel.FIX_PAD_SIZE, + f"m_aligned must be {self._kernel.FIX_PAD_SIZE} (FIX_PAD_SIZE), got {self.m_aligned}", ) self._logger.debug("Checking tensor alignment") @@ -566,7 +582,7 @@ def compile(self) -> None: self._use_full_dynamic_mnkl = os.environ.get("CUDNN_FE_GROUPED_GEMM_DYNAMIC_MNKL", "1") != "0" - gemm_quant = self._kernel( + kernel_kwargs = dict( sf_vec_size=self.sf_vec_size, acc_dtype=_convert_to_cutlass_data_type(self.acc_dtype), use_2cta_instrs=self.use_2cta_instrs, @@ -579,8 +595,13 @@ def compile(self) -> None: expert_cnt=self.expert_cnt, weight_mode=self.weight_mode, use_dynamic_sched=self.use_dynamic_sched, - use_single_group_runtime_offsets=self.use_single_group_runtime_offsets, + **rubin_single_group_offsets_kwarg(self._is_rubin_kernel, self.use_single_group_runtime_offsets), ) + if self._is_rubin_kernel: + # The Rubin quant kernel supports optional C materialization, but + # this cuDNN FE wrapper only exposes quantized D/D_col outputs. + kernel_kwargs["generate_c"] = False + gemm_quant = self._kernel(**kernel_kwargs) hardware_info = cutlass.utils.HardwareInfo() max_active_clusters = hardware_info.get_max_active_clusters(self.cluster_shape_mn[0] * self.cluster_shape_mn[1]) @@ -785,8 +806,7 @@ def _compile_dense(self, gemm_quant, max_active_clusters, fake_stream) -> None: stride=(1, n_sym), ) - _compiled_kernel = cute.compile( - gemm_quant, + compile_kwargs = dict( a=a_cute_fake, b=b_cute_fake, sfb=sfb_cute_fake, @@ -804,13 +824,18 @@ def _compile_dense(self, gemm_quant, max_active_clusters, fake_stream) -> None: norm_const_tensor=self._make_fake_cute_tensor_from_desc(self.norm_const_desc, assumed_align=16), padded_offsets=self._make_fake_cute_tensor_from_desc(self.padded_offsets_desc, assumed_align=16), alpha=self._make_fake_cute_tensor_from_desc(self.alpha_desc, assumed_align=16), - row_scale=row_scale_cute_fake, bias=bias_cute_fake, prob=prob_cute_fake, max_active_clusters=max_active_clusters, stream=fake_stream, options="--enable-tvm-ffi", ) + if self._is_rubin_kernel: + compile_kwargs["c"] = d_cute_fake + compile_kwargs["epilogue_op"] = lambda x: x + else: + compile_kwargs["row_scale"] = row_scale_cute_fake + _compiled_kernel = cute.compile(gemm_quant, **compile_kwargs) cached_workspace_ptr = from_dlpack(self._workspace, assumed_align=128).iterator @@ -833,28 +858,52 @@ def tensor_api( stream: cuda.CUstream, ) -> None: norm_const_tensor = self._unpad_tensor_to_ndim(norm_const_tensor, 1, "norm_const") - _compiled_kernel( - a_tensor, - b_tensor, - sfb_tensor, - cutlass.Int32(0), - cutlass.Int32(0), - cutlass.Int64(0), - cached_workspace_ptr, - d_tensor, - d_col_tensor, - sfa_tensor, - sfd_row_tensor, - sfd_col_tensor, - amax_tensor, - norm_const_tensor, - padded_offsets, - alpha_tensor, - row_scale_tensor, - bias_tensor, - prob_tensor, - stream, - ) + if self._is_rubin_kernel: + _compiled_kernel( + a_tensor, + b_tensor, + sfb_tensor, + cutlass.Int32(0), + cutlass.Int32(0), + cutlass.Int64(0), + cached_workspace_ptr, + d_tensor, + d_tensor, + d_col_tensor, + sfa_tensor, + sfd_row_tensor, + sfd_col_tensor, + amax_tensor, + norm_const_tensor, + padded_offsets, + alpha_tensor, + bias_tensor, + prob_tensor, + stream, + ) + else: + _compiled_kernel( + a_tensor, + b_tensor, + sfb_tensor, + cutlass.Int32(0), + cutlass.Int32(0), + cutlass.Int64(0), + cached_workspace_ptr, + d_tensor, + d_col_tensor, + sfa_tensor, + sfd_row_tensor, + sfd_col_tensor, + amax_tensor, + norm_const_tensor, + padded_offsets, + alpha_tensor, + row_scale_tensor, + bias_tensor, + prob_tensor, + stream, + ) self._compiled_kernel = tensor_api @@ -941,8 +990,7 @@ def _compile_discrete(self, gemm_quant, max_active_clusters, fake_stream) -> Non workspace_ptr_cute = from_dlpack(self._workspace, assumed_align=128).iterator self._logger.debug("Compiling discrete grouped_gemm_quant kernel") - _compiled_kernel = cute.compile( - gemm_quant, + compile_kwargs = dict( a=a_tensor, b=b_ptrs_cute, sfb=sfb_ptrs_cute, @@ -960,7 +1008,6 @@ def _compile_discrete(self, gemm_quant, max_active_clusters, fake_stream) -> Non norm_const_tensor=norm_const_tensor_cute, padded_offsets=padded_offsets_tensor, alpha=alpha_tensor, - row_scale=row_scale_tensor, bias=bias_cute_fake, prob=prob_tensor, max_active_clusters=max_active_clusters, @@ -968,6 +1015,11 @@ def _compile_discrete(self, gemm_quant, max_active_clusters, fake_stream) -> Non epilogue_op=lambda x: x, options="--enable-tvm-ffi", ) + if self._is_rubin_kernel: + compile_kwargs["c"] = d_tensor + else: + compile_kwargs["row_scale"] = row_scale_tensor + _compiled_kernel = cute.compile(gemm_quant, **compile_kwargs) cached_workspace_ptr = from_dlpack(self._workspace, assumed_align=128).iterator cached_n = cutlass.Int32(n) @@ -995,28 +1047,52 @@ def tensor_api( norm_const_tensor = self._unpad_tensor_to_ndim(norm_const_tensor, 1, "norm_const") b_ptrs_addr = int(b_ptrs_device.data_ptr()) sfb_ptrs_addr = int(sfb_ptrs_device.data_ptr()) - _compiled_kernel( - a_tensor, - b_ptrs_addr, - sfb_ptrs_addr, - cached_n, - cached_k, - cached_b_stride, - cached_workspace_ptr, - d_tensor, - d_col_tensor, - sfa_tensor, - sfd_row_tensor, - sfd_col_tensor, - amax_tensor, - norm_const_tensor, - padded_offsets, - alpha_tensor, - row_scale_tensor, - bias_tensor, - prob_tensor, - stream, - ) + if self._is_rubin_kernel: + _compiled_kernel( + a_tensor, + b_ptrs_addr, + sfb_ptrs_addr, + cached_n, + cached_k, + cached_b_stride, + cached_workspace_ptr, + d_tensor, + d_tensor, + d_col_tensor, + sfa_tensor, + sfd_row_tensor, + sfd_col_tensor, + amax_tensor, + norm_const_tensor, + padded_offsets, + alpha_tensor, + bias_tensor, + prob_tensor, + stream, + ) + else: + _compiled_kernel( + a_tensor, + b_ptrs_addr, + sfb_ptrs_addr, + cached_n, + cached_k, + cached_b_stride, + cached_workspace_ptr, + d_tensor, + d_col_tensor, + sfa_tensor, + sfd_row_tensor, + sfd_col_tensor, + amax_tensor, + norm_const_tensor, + padded_offsets, + alpha_tensor, + row_scale_tensor, + bias_tensor, + prob_tensor, + stream, + ) self._compiled_kernel = tensor_api @@ -1101,7 +1177,12 @@ def execute( bias_tensor is not None, "bias_tensor must be omitted at execute() when the API was compiled without sample_bias", ) - if self.row_scale_desc is None: + if self._is_rubin_kernel: + self._value_error_if( + row_scale_tensor is not None, + "row_scale_tensor is not supported on Rubin (sm107)", + ) + elif self.row_scale_desc is None: self._value_error_if( row_scale_tensor is not None, "row_scale_tensor must be omitted at execute() when the API was compiled without sample_row_scale", @@ -1374,7 +1455,10 @@ def grouped_gemm_quant_wrapper_sm100( "prob_tensor is required: the kernel unconditionally multiplies output by per-row gating probability. " "Pass a tensor of ones with shape (valid_m, 1, 1) if no gating is needed." ) + device_type = get_device_type() if row_scale_tensor is not None: + if device_type == "rubin": + raise NotImplementedError("Rubin grouped GEMM quant does not support row_scale fusion") if row_scale_tensor.dtype != torch.float32: raise ValueError(f"row_scale_tensor must be float32, got {row_scale_tensor.dtype}") if tuple(row_scale_tensor.shape) != (valid_m,): @@ -1417,6 +1501,7 @@ def dynamic_m_tensor_signature( if is_dense: cache_key = ( + device_type, weight_mode, use_full_dynamic, a_tensor.shape[1:] if not use_full_dynamic else None, @@ -1455,6 +1540,7 @@ def dynamic_m_tensor_signature( ) else: cache_key = ( + device_type, weight_mode, a_tensor.shape[1:], stride_order(a_tensor), diff --git a/python/cudnn/grouped_gemm/grouped_gemm_quant/moe_blockscaled_grouped_gemm_quant_rubin.py b/python/cudnn/grouped_gemm/grouped_gemm_quant/moe_blockscaled_grouped_gemm_quant_rubin.py new file mode 100644 index 000000000..cf4e9448b --- /dev/null +++ b/python/cudnn/grouped_gemm/grouped_gemm_quant/moe_blockscaled_grouped_gemm_quant_rubin.py @@ -0,0 +1,2519 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: BSD-3-Clause + +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: + +# 1. Redistributions of source code must retain the above copyright notice, this +# list of conditions and the following disclaimer. + +# 2. Redistributions in binary form must reproduce the above copyright notice, +# this list of conditions and the following disclaimer in the documentation +# and/or other materials provided with the distribution. + +# 3. Neither the name of the copyright holder nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. + +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +""" +MoE Block-Scaled Grouped GEMM Kernel with Quantization Support. + +Supports: + - Static / Dynamic persistent tile scheduling (MoEPersistentTileScheduler) + - Dense (contiguous 3-D B) / Discrete (per-expert pointer array B) weight layout + - FP8/FP4 output quantization with row/column scale factors (SFD) + - Optional bias and routing-probability (prob) fusion + - Optional C output (generate_c) + - AMAX reduction for FP8 calibration + +This module contains only the kernel class. +MoE scheduler components live in moe_persistent_scheduler.py / moe_sched_extension.py / moe_utils.py. +""" + +from typing import Type, Tuple, Union, Optional +from enum import Enum + +import cuda.bindings.driver as cuda + +import cutlass +import cutlass.cute as cute +from cutlass.cute.nvgpu import cpasync, tcgen05 +from cutlass.cute.nvgpu.tcgen05 import OperandMajorMode, CollectorOp +from cutlass.utils.gemm.sm100 import transform_partitioned_tensor_layout +import cutlass.utils as utils +import cutlass.pipeline as pipeline +import cutlass.utils.blackwell_helpers as sm100_utils +import cutlass.utils.rubin_helpers as sm107_utils +import cutlass.utils.blockscaled_layout as blockscaled_utils +from cutlass.cute.typing import Float32, Int32, AddressSpace +from ..moe_persistent_scheduler import ( + MoEPersistentTileScheduler, + MoESchedulerParams, + MoEWorkTileInfo, +) +from ..moe_utils import ( + compute_expert_token_range, + MoEWeightMode, + TensormapWorkspace, + store_tma_desc, +) +from ..moe_sched_extension import ( + DiscreteWeightScaledGemmSchedExtension, + ContiguousAndConsistentGroupedGemmSchedExtension, +) +from ..moe_kernel_helpers import ( + fmin, + fmax, + atomic_max_float32, + compute_stages, + compute_grid, + can_implement, + amax_reduction_per_thread, + epilog_gmem_copy_and_partition, + get_dtype_rcp_limits, + is_valid_dtypes_and_scale_factor_vec_size, + is_valid_layouts, + is_valid_tensor_alignment, + FIX_PAD_SIZE, +) + + +class EpilogueType(Enum): + NONE = 0 + SRELU = 1 + + +class BlockScaledMoEGroupedGemmQuantKernel: + """Block-scaled grouped GEMM kernel with MoE tile scheduling and quantization. + + Supports both dense and discrete weight layouts, static and dynamic + scheduling, and quantized output with row/column scale factors. + + :param sf_vec_size: Scale-factor vector size (16 or 32). + :param acc_dtype: Accumulator data type (Float32). + :param use_2cta_instrs: Use 2-CTA MMA instructions. + :param mma_tiler_mn: MMA tile shape (M, N). + :param cluster_shape_mn: Cluster shape (M, N). + :param vectorized_f32: Use packed FP32 arithmetic. + :param generate_sfd: Generate output scale factors. + :param discrete_col_sfd: Use discrete column SFD layout. + :param generate_c: Generate C output tensor. + :param enable_bias: Fuse bias addition. + :param expert_cnt: Number of experts. + :param weight_mode: ``MoEWeightMode.DENSE`` or ``MoEWeightMode.DISCRETE``. + :param use_dynamic_sched: Enable dynamic tile scheduling. + """ + + FIX_PAD_SIZE = 256 + + @staticmethod + def can_implement( + ab_dtype: Type[cutlass.Numeric], + sf_dtype: Type[cutlass.Numeric], + sf_vec_size: int, + acc_dtype: Type[cutlass.Numeric], + d_dtype: Type[cutlass.Numeric], + use_2cta_instrs: bool, + mma_tiler_mn: Tuple[int, int], + cluster_shape_mn: Tuple[int, int], + m: int, + n: int, + k: int, + l: int, + a_major: str, + b_major: str, + cd_major: str, + m_aligned: int, + ) -> bool: + # B-reuse case: 2CTA + mma_tiler_mn[0] = 512 (two 256-M instructions per tile) + if use_2cta_instrs and mma_tiler_mn[0] == 512: + # Pad alignment: per CTA tile = 256 M rows + if not is_valid_dtypes_and_scale_factor_vec_size(ab_dtype, sf_dtype, sf_vec_size, acc_dtype, d_dtype): + return False + if not is_valid_layouts(ab_dtype, d_dtype, a_major, b_major, cd_major): + return False + if not is_valid_tensor_alignment(m, n, k, l, ab_dtype, d_dtype, a_major, b_major, cd_major): + return False + if mma_tiler_mn[1] not in {192, 256}: + return False + if cluster_shape_mn[0] % 2 != 0: + return False + if m_aligned % mma_tiler_mn[0] != 0: + return False + if m % mma_tiler_mn[0] != 0: + return False + return True + # Allow N=192 in addition to the shared helper's N=256 constraint. + if mma_tiler_mn[1] == 192: + if m_aligned != FIX_PAD_SIZE: + return False + if ab_dtype.width == 8: + return False + if not is_valid_dtypes_and_scale_factor_vec_size(ab_dtype, sf_dtype, sf_vec_size, acc_dtype, d_dtype): + return False + if not is_valid_layouts(ab_dtype, d_dtype, a_major, b_major, cd_major): + return False + if not is_valid_tensor_alignment(m, n, k, l, ab_dtype, d_dtype, a_major, b_major, cd_major): + return False + if not (a_major == "k" and b_major == "k"): + return False + if n % 64 != 0 or m % 256 != 0: + return False + if not (use_2cta_instrs and mma_tiler_mn[0] == 256): + return False + if cluster_shape_mn[0] % 2 != 0: + return False + return True + return can_implement( + ab_dtype, + sf_dtype, + sf_vec_size, + acc_dtype, + d_dtype, + use_2cta_instrs, + mma_tiler_mn, + cluster_shape_mn, + m, + n, + k, + l, + a_major, + b_major, + cd_major, + m_aligned, + fix_pad_size=BlockScaledMoEGroupedGemmQuantKernel.FIX_PAD_SIZE, + ) + + def __init__( + self, + sf_vec_size: int, + acc_dtype: Type[cutlass.Numeric], + use_2cta_instrs: bool, + mma_tiler_mn: Tuple[int, int], + cluster_shape_mn: Tuple[int, int], + vectorized_f32: bool, + generate_sfd: bool, + discrete_col_sfd: bool, + generate_c: bool, + enable_bias: bool, + expert_cnt: int, + weight_mode: MoEWeightMode = MoEWeightMode.DENSE, + use_dynamic_sched: bool = False, + epilogue_type: int = EpilogueType.NONE.value, + ): + # Hardware MMA instruction M: 2CTA → 256, 1CTA → 128 + mma_inst_m = 256 if use_2cta_instrs else 128 + enable_breuse = mma_tiler_mn[0] // mma_inst_m == 2 + # For non-breuse: FIX_PAD_SIZE must be divisible by the per-CTA tile M. + # For breuse: mma_tiler_mn[0] (the D tile span) must be a multiple of FIX_PAD_SIZE, + # so that expert padding to mma_tiler_mn[0] is also compatible with FIX_PAD_SIZE. + if enable_breuse: + if mma_tiler_mn[0] % self.FIX_PAD_SIZE != 0: + raise ValueError( + f"mma_tiler_mn[0] ({mma_tiler_mn[0]}) must be a multiple of " + f"FIX_PAD_SIZE ({self.FIX_PAD_SIZE}) for breuse. " + f"Also ensure callers use m_aligned=mma_tiler_mn[0] (={mma_tiler_mn[0]})." + ) + else: + cta_tile_m = mma_tiler_mn[0] // (2 if use_2cta_instrs else 1) + if self.FIX_PAD_SIZE % cta_tile_m != 0: + raise ValueError( + f"FIX_PAD_SIZE ({self.FIX_PAD_SIZE}) must be divisible by " + f"cta_tile_m ({cta_tile_m}). " + f"Supported mma_tiler_mn[0] values: 128, 256, 512." + ) + if expert_cnt > 1024: + raise ValueError("Expert count > 1024 is not supported.") + if not isinstance(weight_mode, MoEWeightMode): + raise TypeError(f"weight_mode must be a MoEWeightMode, got {type(weight_mode)}") + + self.sf_vec_size = sf_vec_size + self.expert_cnt = expert_cnt + self.acc_dtype: Type[cutlass.Numeric] = acc_dtype + self.use_2cta_instrs = use_2cta_instrs + self.cluster_shape_mn = cluster_shape_mn + self.mma_tiler = (*mma_tiler_mn, 1) + # B-reuse: enabled when the mma_tiler M is 2× the hardware instruction M + # (2CTA instruction M = 256; 1CTA instruction M = 128) + self.enable_breuse = mma_tiler_mn[0] // mma_inst_m == 2 + + self.cta_group = tcgen05.CtaGroup.TWO if use_2cta_instrs else tcgen05.CtaGroup.ONE + + self.occupancy = 1 + self.epilog_warp_id = (0, 1, 2, 3) + self.mma_warp_id = 4 + self.tma_warp_id = 5 + self.sched_warp_id = 6 + self.bias_load_warp_id = 7 if enable_bias else None + self.threads_per_warp = 32 + all_warps = [ + *self.epilog_warp_id, + self.mma_warp_id, + self.tma_warp_id, + self.sched_warp_id, + ] + warps_wo_sched = [*self.epilog_warp_id, self.mma_warp_id, self.tma_warp_id] + if enable_bias: + all_warps.append(self.bias_load_warp_id) + warps_wo_sched.append(self.bias_load_warp_id) + self.threads_per_cta = self.threads_per_warp * len(all_warps) + self.threads_wo_sched = self.threads_per_warp * len(warps_wo_sched) + + self.cta_sync_barrier = pipeline.NamedBarrier( + barrier_id=1, + num_threads=self.threads_per_cta, + ) + self.epilog_sync_barrier = pipeline.NamedBarrier( + barrier_id=2, + num_threads=32 * len(self.epilog_warp_id), + ) + self.tmem_alloc_barrier = pipeline.NamedBarrier( + barrier_id=3, + num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id)), + ) + self.sched_sync_barrier = pipeline.NamedBarrier( + barrier_id=4, + num_threads=self.threads_per_warp, + ) + self.num_smem_capacity = utils.get_smem_capacity_in_bytes("sm_107") + self.num_tmem_alloc_cols = cute.arch.get_max_tmem_alloc_cols("sm_107") + + self.vectorized_f32 = vectorized_f32 + self.generate_sfd = generate_sfd + self.discrete_col_sfd = discrete_col_sfd + self.generate_c = generate_c + self.enable_bias = enable_bias + + self.weight_mode = weight_mode + self.use_dynamic_sched = use_dynamic_sched + self.epilogue_type = epilogue_type + + self.epilogue_use_functor = False + + self.num_epilog_warps = len(self.epilog_warp_id) + + # ------------------------------------------------------------------ + # _setup_attributes + # ------------------------------------------------------------------ + + def _get_mma_permutation_mnk(self): + """Return MMA permutation for the Bkeep-Breuse pattern (2CTA only).""" + if cutlass.const_expr(self.use_2cta_instrs and self.enable_breuse): + mma_inst_k = 128 if (self.a_dtype.width == 4 and self.b_dtype.width == 4) else 64 + m_layout = cute.make_layout( + shape=(self.mma_inst_shape_mn[0] // 2, 2, 2), + stride=(1, self.mma_inst_shape_mn[0], self.mma_inst_shape_mn[0] // 2), + ) + return (m_layout, self.mma_inst_shape_mn[1], mma_inst_k) + else: + return (1, 1, 1) + + def _setup_attributes(self): + """Configure MMA / tile / stage / SMEM layouts from GEMM inputs.""" + + # Hardware MMA instruction M: always 256 for 2CTA, 128 for 1CTA + mma_inst_m = 256 if self.use_2cta_instrs else 128 + self.mma_inst_shape_mn = (mma_inst_m, self.mma_tiler[1]) + self.mma_inst_shape_mn_sfb = ( + self.mma_inst_shape_mn[0] // (2 if self.use_2cta_instrs else 1), + cute.round_up(self.mma_inst_shape_mn[1], 128), + ) + + # K dim: sm107 uses K=128 for FP4×FP4, K=64 for FP8/mixed + mma_inst_k = 128 if (self.a_dtype.width == 4 and self.b_dtype.width == 4) else 64 + mma_inst_shape_mnk = (*self.mma_inst_shape_mn, mma_inst_k) + mma_inst_shape_mnk_sfb = (*self.mma_inst_shape_mn_sfb, mma_inst_k) + + atom_layout_mnk = (1, 1, 1) + permutation_mnk = self._get_mma_permutation_mnk() + tiled_mma = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + self.cta_group, + mma_inst_shape_mnk, + a_collector_op=CollectorOp.DISCARD, + b_collector_op=CollectorOp.DISCARD, + atom_layout_mnk=atom_layout_mnk, + permutation_mnk=permutation_mnk, + ) + tiled_mma_sfb = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + cute.nvgpu.tcgen05.CtaGroup.ONE, + mma_inst_shape_mnk_sfb, + ) + + mma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2]) + mma_inst_tile_k = 2 if (self.a_dtype.width == 4 and self.sf_vec_size == 16) else 4 + self.mma_tiler = ( + self.mma_tiler[0], + self.mma_tiler[1], + mma_inst_shape_k * mma_inst_tile_k, + ) + self.mma_tiler_sfb = ( + self.mma_inst_shape_mn_sfb[0], + self.mma_inst_shape_mn_sfb[1], + mma_inst_shape_k * mma_inst_tile_k, + ) + + self.cta_tile_shape_mnk = ( + self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape), + self.mma_tiler[1], + self.mma_tiler[2], + ) + self.cta_tile_shape_mnk_sfb = ( + self.mma_tiler_sfb[0] // cute.size(tiled_mma.thr_id.shape), + self.mma_tiler_sfb[1], + self.mma_tiler_sfb[2], + ) + + # For breuse, D tiler M = mma_tiler[0] (512 for 2CTA+breuse) so each CTA writes + # all its output rows (bkeep + breuse = 256 per CTA). + # For non-breuse, mma_tiler[0] == mma_inst_shape_mn[0] so this is equivalent. + self.mma_tiler_d = ( + self.mma_tiler[0], + self.mma_inst_shape_mn[1], + mma_inst_shape_k * mma_inst_tile_k, + ) + self.cta_tile_shape_mnk_d = ( + self.mma_tiler_d[0] // cute.size(tiled_mma.thr_id.shape), + self.mma_tiler_d[1], + self.mma_tiler_d[2], + ) + + self.cluster_layout_vmnk = cute.tiled_divide( + cute.make_layout((*self.cluster_shape_mn, 1)), + (tiled_mma.thr_id.shape,), + ) + self.cluster_layout_sfb_vmnk = cute.tiled_divide( + cute.make_layout((*self.cluster_shape_mn, 1)), + (tiled_mma_sfb.thr_id.shape,), + ) + + self.num_mcast_ctas_a = cute.size(self.cluster_layout_vmnk.shape[2]) + self.num_mcast_ctas_b = cute.size(self.cluster_layout_vmnk.shape[1]) + self.is_a_mcast = self.num_mcast_ctas_a > 1 + self.is_b_mcast = self.num_mcast_ctas_b > 1 + + self.epi_tile = (128, 32) + + ( + self.num_acc_stage, + self.num_ab_stage, + self.num_c_stage, + self.num_d_stage, + self.num_tile_stage, + self.num_bias_stage, + ) = self._compute_stages( + tiled_mma, + self.mma_tiler, + self.a_dtype, + self.b_dtype, + self.epi_tile, + self.c_dtype, + self.c_layout, + self.d_dtype, + self.d_layout, + self.sf_dtype, + self.sf_vec_size, + self.num_smem_capacity, + self.occupancy, + self.generate_sfd, + self.generate_c, + self.bias_dtype if self.enable_bias else None, + self.enable_breuse, + ) + + self.a_smem_layout_staged = sm100_utils.make_smem_layout_a( + tiled_mma, + self.mma_tiler, + self.a_dtype, + self.num_ab_stage, + ) + self.b_smem_layout_staged = sm100_utils.make_smem_layout_b( + tiled_mma, + self.mma_tiler, + self.b_dtype, + self.num_ab_stage, + ) + self.sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa( + tiled_mma, + self.mma_tiler, + self.sf_vec_size, + self.num_ab_stage, + ) + self.sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb( + tiled_mma, + self.mma_tiler, + self.sf_vec_size, + self.num_ab_stage, + ) + self.c_smem_layout_staged = sm100_utils.make_smem_layout_epi( + self.c_dtype, + self.c_layout, + self.epi_tile, + self.num_c_stage, + ) + self.d_smem_layout_staged = sm100_utils.make_smem_layout_epi( + self.d_dtype, + self.d_layout, + self.epi_tile, + self.num_d_stage, + ) + + if self.enable_bias: + self.bias_smem_layout_staged = cute.make_layout( + (self.mma_tiler[1], self.num_bias_stage), + stride=(1, self.mma_tiler[1]), + ) + else: + self.bias_smem_layout_staged = cute.make_layout((1, 1)) + + # overlapping_accum shares SFA/SFB TMEM with second acc stage; not compatible with breuse + self.overlapping_accum = self.num_acc_stage == 1 and self.mma_tiler[1] == 256 and not self.enable_breuse + self.epilogue_prefetch_more = False + + sf_atom_mn = 32 + sf_pack_factor = 32 // self.sf_vec_size + self.num_sfa_tmem_cols = (self.cta_tile_shape_mnk[0] // sf_atom_mn) * mma_inst_tile_k * sf_pack_factor + self.num_sfb_tmem_cols = (self.cta_tile_shape_mnk_sfb[1] // sf_atom_mn) * mma_inst_tile_k * sf_pack_factor + self.num_sf_tmem_cols = self.num_sfa_tmem_cols + self.num_sfb_tmem_cols + if self.enable_breuse: + # Breuse: 2 accumulators (bkeep + breuse) active simultaneously + self.num_accumulator_tmem_cols = self.cta_tile_shape_mnk[1] * self.num_acc_stage * 2 + elif self.overlapping_accum: + self.num_accumulator_tmem_cols = self.cta_tile_shape_mnk[1] * 2 - self.num_sf_tmem_cols + else: + self.num_accumulator_tmem_cols = self.cta_tile_shape_mnk[1] * self.num_acc_stage + # N=192 non-breuse: 192 cols don't fill a full TMEM row, so pack two acc stages + # into the remaining space (overlapping with SF area isn't an option here). + # For breuse+N=192 this is skipped: breuse already sets acc=2*192=384 correctly. + if self.cta_tile_shape_mnk[1] == 192 and not self.enable_breuse: + self.num_accumulator_tmem_cols = self.num_tmem_alloc_cols - self.num_sf_tmem_cols + self.num_accumulator_tmem_stride = self.num_accumulator_tmem_cols - 192 + else: + self.num_accumulator_tmem_stride = self.num_accumulator_tmem_cols + + self.epi_tile_n_required = cute.size(self.epi_tile[1]) + self.iter_acc_early_release_in_epilogue = (self.num_sf_tmem_cols + self.epi_tile_n_required - 1) // self.epi_tile_n_required - 1 + + # ------------------------------------------------------------------ + # _compute_stages (with bias support) + # ------------------------------------------------------------------ + + @staticmethod + def _compute_stages( + tiled_mma, + mma_tiler_mnk, + a_dtype, + b_dtype, + epi_tile, + c_dtype, + c_layout, + d_dtype, + d_layout, + sf_dtype, + sf_vec_size, + num_smem_capacity, + occupancy, + generate_sfd, + generate_c, + bias_dtype, + enable_breuse=False, + ): + num_acc_stage = 1 if (mma_tiler_mnk[1] == 256 or (enable_breuse and mma_tiler_mnk[1] == 192)) else 2 + num_c_stage = 2 if generate_sfd else 1 + num_d_stage = 2 if generate_sfd else 1 + num_tile_stage = 2 + + a_smem_layout_stage_one = sm100_utils.make_smem_layout_a(tiled_mma, mma_tiler_mnk, a_dtype, 1) + b_smem_layout_staged_one = sm100_utils.make_smem_layout_b(tiled_mma, mma_tiler_mnk, b_dtype, 1) + sfa_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfa(tiled_mma, mma_tiler_mnk, sf_vec_size, 1) + sfb_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfb(tiled_mma, mma_tiler_mnk, sf_vec_size, 1) + c_smem_layout_staged_one = sm100_utils.make_smem_layout_epi(c_dtype, c_layout, epi_tile, 1) + d_smem_layout_staged_one = sm100_utils.make_smem_layout_epi(d_dtype, d_layout, epi_tile, 1) + + ab_bytes_per_stage = ( + cute.size_in_bytes(a_dtype, a_smem_layout_stage_one) + + cute.size_in_bytes(b_dtype, b_smem_layout_staged_one) + + cute.size_in_bytes(sf_dtype, sfa_smem_layout_staged_one) + + cute.size_in_bytes(sf_dtype, sfb_smem_layout_staged_one) + ) + mbar_helpers_bytes = 1024 + sinfo_bytes = 4 * 4 * num_tile_stage + c_bytes_per_stage = cute.size_in_bytes(c_dtype, c_smem_layout_staged_one) + c_bytes = c_bytes_per_stage * num_c_stage + d_bytes_per_stage = cute.size_in_bytes(d_dtype, d_smem_layout_staged_one) + d_bytes = d_bytes_per_stage * num_d_stage * (2 if generate_sfd else 1) + amax_bytes = 4 * cute.size_in_bytes(cutlass.Float32, cute.make_layout((1,))) if d_dtype == cutlass.BFloat16 else 0 + + if bias_dtype is not None: + num_bias_stage = 2 + bias_epi_tile_n = mma_tiler_mnk[1] + bias_bytes = bias_epi_tile_n * num_bias_stage * (bias_dtype.width // 8) + else: + num_bias_stage = 0 + bias_bytes = 0 + + epi_bytes = c_bytes + d_bytes + amax_bytes + bias_bytes + num_ab_stage = (num_smem_capacity // occupancy - (mbar_helpers_bytes + epi_bytes + sinfo_bytes)) // ab_bytes_per_stage + + return num_acc_stage, num_ab_stage, num_c_stage, num_d_stage, num_tile_stage, num_bias_stage + + # ------------------------------------------------------------------ + # Workspace helpers + # ------------------------------------------------------------------ + + def get_desc_workspace_bytes(self) -> int: + if self.weight_mode == MoEWeightMode.DISCRETE: + from ..moe_utils import DiscreteWeightTensormapConstructor + + return DiscreteWeightTensormapConstructor.get_workspace_size(self.expert_cnt) + return 0 + + def get_workspace_bytes(self) -> int: + desc_workspace_bytes = self.get_desc_workspace_bytes() + dynamic_sched_bytes = 4 if self.use_dynamic_sched else 0 + return desc_workspace_bytes + dynamic_sched_bytes + + @cute.jit + def _get_sched_counter_ptr(self, workspace_ptr): + counter_addr = workspace_ptr.toint() + self.get_desc_workspace_bytes() + return cute.make_ptr( + cutlass.Int32, + counter_addr, + AddressSpace.gmem, + assumed_align=4, + ) + + # ------------------------------------------------------------------ + # helper_kernel: pre-main-kernel initialization + # - discrete weight: build per-expert B/SFB TMA descriptors + # - dynamic sched: reset the atomic tile counter + # ------------------------------------------------------------------ + + @cute.kernel + def helper_kernel( + self, + # Discrete-only params (unused in dense mode, but must be present for signature) + ptrs_b: cute.Pointer, + ptrs_sfb: cute.Pointer, + n: Int32, + k: Int32, + b_stride_size: cutlass.Int64, + b_major_mode: cutlass.Constexpr, + workspace_ptr, + tiled_mma_arg: cute.TiledMma, + tiled_mma_sfb_arg: cute.TiledMma, + b_smem_layout_arg, + sfb_smem_layout_arg, + cluster_layout_vmnk_shape_arg: cutlass.Constexpr, + cluster_layout_sfb_vmnk_shape_arg: cutlass.Constexpr, + ): + """Pre-main-kernel initialization. + + Launched with grid=(expert_cnt, 1, 1) for discrete mode, or + grid=(1, 1, 1) for dense+dynamic mode. + + Discrete weight: each block builds B/SFB TMA descriptors for one expert. + Dynamic sched: block 0 resets the atomic tile counter to 0. + """ + expert_idx = cute.arch.block_idx()[0] + + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE): + b_tma_op_arg = sm100_utils.cluster_shape_to_tma_atom_B(self.cluster_shape_mn, tiled_mma_arg.thr_id) + sfb_tma_op_arg = sm100_utils.cluster_shape_to_tma_atom_SFB(self.cluster_shape_mn, tiled_mma_arg.thr_id) + + b_ptr_tensor = cute.make_tensor( + cute.make_ptr(cutlass.Int64, ptrs_b.toint(), AddressSpace.gmem, assumed_align=8), cute.make_layout((self.expert_cnt,)) + ) + sfb_ptr_tensor = cute.make_tensor( + cute.make_ptr(cutlass.Int64, ptrs_sfb.toint(), AddressSpace.gmem, assumed_align=8), cute.make_layout((self.expert_cnt,)) + ) + + c0 = cutlass.Int64(0) + c1_64 = 1 + if cutlass.const_expr(b_major_mode == OperandMajorMode.K): + stride_n = b_stride_size + stride_k = c1_64 + else: + stride_n = c1_64 + stride_k = b_stride_size + + b_ptr_val = b_ptr_tensor[expert_idx] + b_ptr = cute.make_ptr(self.b_dtype, b_ptr_val, AddressSpace.gmem) + b_tensor_i = cute.make_tensor( + b_ptr, + cute.make_layout((n, k, cutlass.Int32(1)), stride=(stride_n, stride_k, c0)), + ) + tma_atom_b, _ = cute.nvgpu.make_tiled_tma_atom_B( + b_tma_op_arg, + b_tensor_i, + b_smem_layout_arg, + self.mma_tiler, + tiled_mma_arg, + cluster_layout_vmnk_shape_arg, + ) + workspace = TensormapWorkspace(workspace_ptr, ["b", "sfb"]) + store_tma_desc(tma_atom_b, workspace.get_ptr("b", expert_idx)) + + sfb_ptr_val = sfb_ptr_tensor[expert_idx] + sfb_ptr = cute.make_ptr(self.sf_dtype, sfb_ptr_val, AddressSpace.gmem) + sfb_layout = blockscaled_utils.tile_atom_to_shape_SF((n, k, cutlass.Int32(1)), self.sf_vec_size) + sfb_tensor_i = cute.make_tensor(sfb_ptr, sfb_layout) + tma_atom_sfb, _ = cute.nvgpu.make_tiled_tma_atom_B( + sfb_tma_op_arg, + sfb_tensor_i, + sfb_smem_layout_arg, + self.mma_tiler_sfb, + tiled_mma_sfb_arg, + cluster_layout_sfb_vmnk_shape_arg, + internal_type=cutlass.Uint64, + ) + store_tma_desc(tma_atom_sfb, workspace.get_ptr("sfb", expert_idx)) + + if cutlass.const_expr(self.use_dynamic_sched): + if expert_idx == cutlass.Int32(0): + sched_counter = cute.make_tensor( + self._get_sched_counter_ptr(workspace_ptr), + cute.make_layout(1), + ) + sched_counter[0] = cutlass.Int32(0) + + @cute.jit + def __call__( + self, + a: cute.Tensor, + b, # Dense: cute.Tensor (N,K,L) | Discrete: cute.Pointer to int64[] + sfb, # Dense: cute.Tensor | Discrete: cute.Pointer to int64[] + n: Int32, # Ignored for dense mode + k: Int32, # Ignored for dense mode + b_stride_size: cutlass.Int64, # Ignored for dense mode + b_major_mode: cutlass.Constexpr, # Ignored for dense mode + workspace_ptr, + c: cute.Tensor, + d: cute.Tensor, + d_col: Optional[cute.Tensor], + sfa: cute.Tensor, + sfd_row_tensor: Optional[cute.Tensor], + sfd_col_tensor: Optional[cute.Tensor], + amax_tensor: Optional[cute.Tensor], + norm_const_tensor: Optional[cute.Tensor], + padded_offsets: cute.Tensor, + alpha: cute.Tensor, + bias: Optional[cute.Tensor], + prob: cute.Tensor, + max_active_clusters: cutlass.Constexpr, + stream: cuda.CUstream, + epilogue_op: cutlass.Constexpr = lambda x: x, + ): + """Execute the GEMM. + + Dense mode: ``b`` and ``sfb`` are 3-D cute.Tensor (N, K, L). + Discrete mode: ``b`` and ``sfb`` are cute.Pointer to device int64[] + arrays of per-expert base addresses; ``n``, ``k``, ``b_stride_size``, + ``b_major_mode`` describe the uniform per-expert layout. + """ + self.a_dtype: Type[cutlass.Numeric] = a.element_type + self.b_dtype: Type[cutlass.Numeric] = a.element_type + self.c_dtype: Type[cutlass.Numeric] = c.element_type + self.d_dtype: Type[cutlass.Numeric] = d.element_type + self.sf_dtype: Type[cutlass.Numeric] = sfa.element_type + self.a_major_mode = utils.LayoutEnum.from_tensor(a).mma_major_mode() + self.c_layout = utils.LayoutEnum.from_tensor(c) + self.d_layout = utils.LayoutEnum.from_tensor(d) + self.bias_dtype = bias.element_type if cutlass.const_expr(self.enable_bias) else cutlass.BFloat16 + + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DENSE): + self.b_major_mode = utils.LayoutEnum.from_tensor(b).mma_major_mode() + else: + self.b_major_mode = b_major_mode + + if cutlass.const_expr(self.a_dtype != self.b_dtype): + raise TypeError(f"A/B dtype must match: {self.a_dtype} != {self.b_dtype}") + + self._setup_attributes() + + # ---- SFA layout ---- + sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a.shape, self.sf_vec_size) + sfa = cute.make_tensor(sfa.iterator, sfa_layout) + + # ---- B / SFB setup (mode-dependent) ---- + # Save the call-arg b/sfb before the discrete branch overwrites them + # with template tensors. helper_kernel needs the original Pointers. + b_from_call_arg = b + sfb_from_call_arg = sfb + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DENSE): + sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b.shape, self.sf_vec_size) + sfb = cute.make_tensor(sfb.iterator, sfb_layout) + else: + c0 = cutlass.Int64(0) + c1_64 = 1 + if cutlass.const_expr(b_major_mode == OperandMajorMode.K): + b_template_stride = (b_stride_size, c1_64, c0) + else: + b_template_stride = (c1_64, b_stride_size, c0) + b_template_layout = cute.make_layout((n, k, cutlass.Int32(1)), stride=b_template_stride) + b_ptr_typed = cute.make_ptr(self.b_dtype, b.toint(), AddressSpace.gmem, assumed_align=16) + b = cute.make_tensor(b_ptr_typed, b_template_layout) + + sfb_ptr_typed = cute.make_ptr(self.sf_dtype, sfb.toint(), AddressSpace.gmem, assumed_align=16) + sfb_layout = blockscaled_utils.tile_atom_to_shape_SF((n, k, cutlass.Int32(1)), self.sf_vec_size) + sfb = cute.make_tensor(sfb_ptr_typed, sfb_layout) + + # ---- SFD setup ---- + self.generate_sfd = sfd_row_tensor is not None and norm_const_tensor is not None + if cutlass.const_expr(self.generate_sfd == False): + self.discrete_col_sfd = False + if cutlass.const_expr(self.generate_sfd): + sfd_row_layout = blockscaled_utils.tile_atom_to_shape_SF(d.shape, self.sf_vec_size) + sfd_row_tensor = cute.make_tensor(sfd_row_tensor.iterator, sfd_row_layout) + sfd_col_layout = cute.tile_to_shape( + blockscaled_utils.BlockScaledBasicChunk(self.sf_vec_size, OperandMajorMode.MN).layout, + d.shape, + (1, 2, 3), + ) + if cutlass.const_expr(self.discrete_col_sfd): + sfd_col_layout = sfd_row_layout + sfd_col_tensor = cute.make_tensor(sfd_col_tensor.iterator, sfd_col_layout) + + self.generate_amax = amax_tensor is not None + + # ---- TMA atoms ---- + mma_inst_k = 128 if (self.a_dtype.width == 4 and self.b_dtype.width == 4) else 64 + mma_inst_shape_mnk = (*self.mma_inst_shape_mn, mma_inst_k) + mma_inst_shape_mnk_sfb = (*self.mma_inst_shape_mn_sfb, mma_inst_k) + + atom_layout_mnk = (1, 1, 1) + permutation_mnk = self._get_mma_permutation_mnk() + tiled_mma = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + self.cta_group, + mma_inst_shape_mnk, + a_collector_op=CollectorOp.DISCARD, + b_collector_op=CollectorOp.DISCARD, + atom_layout_mnk=atom_layout_mnk, + permutation_mnk=permutation_mnk, + ) + tiled_mma_sfb = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + cute.nvgpu.tcgen05.CtaGroup.ONE, + mma_inst_shape_mnk_sfb, + ) + + tiled_mma_bkeep = None + tiled_mma_breuse = None + if cutlass.const_expr(self.enable_breuse): + tiled_mma_bkeep = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + self.cta_group, + mma_inst_shape_mnk, + a_collector_op=CollectorOp.DISCARD, + b_collector_op=CollectorOp.FILL, + atom_layout_mnk=atom_layout_mnk, + permutation_mnk=permutation_mnk, + ) + tiled_mma_bkeep.set(tcgen05.Field.NEGATE_A, False) + tiled_mma_bkeep.set(tcgen05.Field.NEGATE_B, False) + tiled_mma_breuse = sm107_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.b_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + self.cta_group, + mma_inst_shape_mnk, + a_collector_op=CollectorOp.DISCARD, + b_collector_op=CollectorOp.LASTUSE, + atom_layout_mnk=atom_layout_mnk, + permutation_mnk=permutation_mnk, + ) + tiled_mma_breuse.set(tcgen05.Field.NEGATE_A, False) + tiled_mma_breuse.set(tcgen05.Field.NEGATE_B, False) + + atom_thr_size = cute.size(tiled_mma.thr_id.shape) + + a_op = sm100_utils.cluster_shape_to_tma_atom_A(self.cluster_shape_mn, tiled_mma.thr_id) + a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0)) + tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A( + a_op, + a, + a_smem_layout, + self.mma_tiler, + tiled_mma, + self.cluster_layout_vmnk.shape, + ) + + b_op = sm100_utils.cluster_shape_to_tma_atom_B(self.cluster_shape_mn, tiled_mma.thr_id) + b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0)) + tma_atom_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B( + b_op, + b, + b_smem_layout, + self.mma_tiler, + tiled_mma, + self.cluster_layout_vmnk.shape, + ) + + sfa_op = sm100_utils.cluster_shape_to_tma_atom_A(self.cluster_shape_mn, tiled_mma.thr_id) + sfa_smem_layout = cute.slice_(self.sfa_smem_layout_staged, (None, None, None, 0)) + tma_atom_sfa, tma_tensor_sfa = cute.nvgpu.make_tiled_tma_atom_A( + sfa_op, + sfa, + sfa_smem_layout, + self.mma_tiler, + tiled_mma, + self.cluster_layout_vmnk.shape, + internal_type=cutlass.Int16, + ) + + sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB(self.cluster_shape_mn, tiled_mma.thr_id) + sfb_smem_layout = cute.slice_(self.sfb_smem_layout_staged, (None, None, None, 0)) + tma_atom_sfb, tma_tensor_sfb = cute.nvgpu.make_tiled_tma_atom_B( + sfb_op, + sfb, + sfb_smem_layout, + self.mma_tiler_sfb, + tiled_mma_sfb, + self.cluster_layout_sfb_vmnk.shape, + internal_type=cutlass.Uint64, + ) + + a_copy_size = cute.size_in_bytes(self.a_dtype, a_smem_layout) + b_copy_size = cute.size_in_bytes(self.b_dtype, b_smem_layout) + sfa_copy_size = cute.size_in_bytes(self.sf_dtype, sfa_smem_layout) + sfb_copy_size = cute.size_in_bytes(self.sf_dtype, sfb_smem_layout) + self.num_tma_load_bytes = (a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size) * atom_thr_size + + c_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0)) + tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom( + cpasync.CopyBulkTensorTileS2GOp(), + c, + c_smem_layout, + self.epi_tile, + ) + d_smem_layout = cute.slice_(self.d_smem_layout_staged, (None, None, 0)) + tma_atom_d, tma_tensor_d = cpasync.make_tiled_tma_atom( + cpasync.CopyBulkTensorTileS2GOp(), + d, + d_smem_layout, + self.epi_tile, + ) + tma_atom_d_col, tma_tensor_d_col = cpasync.make_tiled_tma_atom( + cpasync.CopyBulkTensorTileS2GOp(), + d_col, + d_smem_layout, + self.epi_tile, + ) + + # ---- Helper kernel: TMA desc init (discrete) + sched counter reset (dynamic) ---- + _need_helper = cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE or self.use_dynamic_sched) + if cutlass.const_expr(_need_helper): + _helper_grid_x = self.expert_cnt if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else 1 + _helper_args = ( + b_from_call_arg if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cute.make_ptr(cutlass.Int64, 0, AddressSpace.gmem), + sfb_from_call_arg if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cute.make_ptr(cutlass.Int64, 0, AddressSpace.gmem), + n if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cutlass.Int32(0), + k if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cutlass.Int32(0), + b_stride_size if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else cutlass.Int64(0), + b_major_mode if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else self.b_major_mode, + workspace_ptr, + tiled_mma, + tiled_mma_sfb, + b_smem_layout, + sfb_smem_layout, + self.cluster_layout_vmnk.shape, + self.cluster_layout_sfb_vmnk.shape, + ) + self.helper_kernel(*_helper_args).launch( + grid=(_helper_grid_x, 1, 1), + block=(1, 1, 1), + stream=stream, + min_blocks_per_mp=1, + ) + + # ---- Grid computation via MoE scheduler ---- + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DENSE): + b_n, b_k, b_l = cute.shape(b) # B is (N, K, L) + sched_expert_shape = (self.expert_cnt, b_n, b_k) + else: + sched_expert_shape = (self.expert_cnt, n, k) + + sched_params = MoESchedulerParams( + scenario="2Dx3D", + expert_shape=sched_expert_shape, + cta_tile_shape_mnk=self.cta_tile_shape_mnk, + cluster_shape_mn=self.cluster_shape_mn, + use_dynamic_sched=self.use_dynamic_sched, + ) + self.sched_params, grid = compute_grid(sched_params, max_active_clusters, self.use_2cta_instrs) + + self.buffer_align_bytes = 1024 + + # ---- Shared storage ---- + sD_col_size = cute.cosize(self.d_smem_layout_staged.outer) if self.generate_sfd else 0 + SchedulerStorage = MoEPersistentTileScheduler.make_storage_struct(self.num_tile_stage, self.use_dynamic_sched) + + @cute.struct + class SharedStorage: + ab_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage * 2] + acc_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage * 2] + scheduler: SchedulerStorage + if cutlass.const_expr(self.enable_bias): + bias_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_bias_stage * 2] + tmem_dealloc_mbar_ptr: cutlass.Int64 + tmem_holding_buf: cutlass.Int32 + sC: cute.struct.Align[ + cute.struct.MemRange[self.c_dtype, cute.cosize(self.c_smem_layout_staged.outer)], + self.buffer_align_bytes, + ] + sD: cute.struct.Align[ + cute.struct.MemRange[self.d_dtype, cute.cosize(self.d_smem_layout_staged.outer)], + self.buffer_align_bytes, + ] + sD_col: cute.struct.Align[ + cute.struct.MemRange[self.d_dtype, sD_col_size], + self.buffer_align_bytes, + ] + sA: cute.struct.Align[ + cute.struct.MemRange[self.a_dtype, cute.cosize(self.a_smem_layout_staged.outer)], + self.buffer_align_bytes, + ] + sB: cute.struct.Align[ + cute.struct.MemRange[self.b_dtype, cute.cosize(self.b_smem_layout_staged.outer)], + self.buffer_align_bytes, + ] + sSFA: cute.struct.Align[ + cute.struct.MemRange[self.sf_dtype, cute.cosize(self.sfa_smem_layout_staged)], + self.buffer_align_bytes, + ] + sSFB: cute.struct.Align[ + cute.struct.MemRange[self.sf_dtype, cute.cosize(self.sfb_smem_layout_staged)], + self.buffer_align_bytes, + ] + if cutlass.const_expr(self.enable_bias): + sBias: cute.struct.Align[ + cute.struct.MemRange[self.bias_dtype, cute.cosize(self.bias_smem_layout_staged)], + 16, + ] + sAmax: cute.struct.Align[ + cute.struct.MemRange[cutlass.Float32, self.num_epilog_warps], + 4, + ] + + self.shared_storage = SharedStorage + + # ---- Launch ---- + self.kernel( + tiled_mma, + tiled_mma_bkeep, + tiled_mma_breuse, + tiled_mma_sfb, + tma_atom_a, + tma_tensor_a, + tma_atom_b, + tma_tensor_b, + tma_atom_sfa, + tma_tensor_sfa, + tma_atom_sfb, + tma_tensor_sfb, + tma_atom_c, + tma_tensor_c, + tma_atom_d, + tma_tensor_d, + tma_atom_d_col, + tma_tensor_d_col, + sfd_row_tensor, + sfd_col_tensor, + norm_const_tensor, + amax_tensor, + padded_offsets, + alpha, + bias, + prob, + workspace_ptr, + self.cluster_layout_vmnk, + self.cluster_layout_sfb_vmnk, + self.a_smem_layout_staged, + self.b_smem_layout_staged, + self.sfa_smem_layout_staged, + self.sfb_smem_layout_staged, + self.c_smem_layout_staged, + self.d_smem_layout_staged, + self.bias_smem_layout_staged, + self.epi_tile, + self.sched_params, + epilogue_op, + ).launch( + grid=grid, + block=[self.threads_per_cta, 1, 1], + cluster=(*self.cluster_shape_mn, 1), + max_number_threads=[self.threads_per_cta, 1, 1], + smem=self.shared_storage.size_in_bytes(), + stream=stream, + min_blocks_per_mp=1, + ) + return + + # ------------------------------------------------------------------ + # Helper methods + # ------------------------------------------------------------------ + + def mainloop_s2t_copy_and_partition(self, sSF, tSF): + tCsSF_compact = cute.filter_zeros(sSF) + tCtSF_compact = cute.filter_zeros(tSF) + copy_atom_s2t = cute.make_copy_atom(tcgen05.Cp4x32x128bOp(self.cta_group), self.sf_dtype) + tiled_copy_s2t = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSF_compact) + thr_copy_s2t = tiled_copy_s2t.get_slice(0) + + # Rubin sm107 workaround: append stride-0 broadcast mode so partition_S + # produces the right shape for NVF4 (sf_vec_size=16); idempotent for sf_vec_size=32. + def _append_mn_broadcast_mode(smem_layout: cute.Layout): + mn_dim = cute.get(smem_layout, mode=[0, 0]) + mn_dim = cute.append(mn_dim, cute.make_layout((4), stride=(0))) + layout = cute.append(cute.group_modes(mn_dim, 0), cute.get(smem_layout, mode=[0, 1])) + layout = cute.append(cute.group_modes(layout, 0), cute.get(smem_layout, mode=[1])) + layout = cute.append(layout, cute.get(smem_layout, mode=[2])) + layout = cute.append(layout, cute.get(smem_layout, mode=[3])) + return layout + + tCsSF_compact_bcast = cute.make_tensor(tCsSF_compact.iterator, _append_mn_broadcast_mode(tCsSF_compact.layout)) + tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact_bcast) + tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(tiled_copy_s2t, tCsSF_compact_s2t_) + tCtSF_compact_s2t = thr_copy_s2t.partition_D(tCtSF_compact) + return tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t + + @cute.jit + def amax_reduction_per_warp_and_cta(self, amax_fp32, warp_idx, amax_smem, amax_gmem): + warp_amax = cute.arch.warp_redux_sync( + value=amax_fp32, + kind="fmax", + mask_and_clamp=0xFFFFFFFF, + nan=True, + ) + if cute.arch.lane_idx() == 0: + amax_smem[warp_idx] = cutlass.Float32(warp_amax) + self.epilog_sync_barrier.arrive_and_wait() + if warp_idx == self.epilog_warp_id[0] and cute.arch.lane_idx() == 0: + block_amax = cutlass.Float32(0.0) + for i in cutlass.range(self.num_epilog_warps): + warp_amax_val = amax_smem[i] + block_amax = cute.arch.fmax(block_amax, warp_amax_val) + _ = atomic_max_float32(ptr=amax_gmem, value=block_amax) + + @cute.jit + def store_c( + self, + tiled_copy_r2s, + tma_atom_c, + warp_idx, + tTR_rAcc, + tTR_rC, + tRS_rC, + tRS_sC, + bSG_gC, + bSG_sC, + c_pipeline, + prev_subtile_idx, + real_subtile_idx, + ): + c_buffer = prev_subtile_idx % self.num_c_stage + tTR_rC.store(tTR_rAcc.load().to(self.c_dtype)) + cute.copy(tiled_copy_r2s, tRS_rC[(None, None, 0)], tRS_sC[(None, None, 0, c_buffer)]) + cute.arch.fence_proxy("async.shared", space="cta") + self.epilog_sync_barrier.arrive_and_wait() + if warp_idx == self.epilog_warp_id[0]: + cute.copy(tma_atom_c, bSG_sC[(None, c_buffer)], bSG_gC[(None, real_subtile_idx)]) + c_pipeline.producer_commit() + c_pipeline.producer_acquire() + self.epilog_sync_barrier.arrive_and_wait() + + @cute.jit + def quant_sfd_row(self, tile_idx, tiled_copy_r2s, src, pvscale, norm_const, rcp_limit, tRSrD): + tTR_rAcc_frg = cute.logical_divide(src, cute.make_layout(self.sf_vec_size)) + acc_frg = tTR_rAcc_frg.load() + abs_acc_frg_ir = cutlass._mlir.dialects.math.absf(acc_frg.ir_value()) + abs_acc_frg = type(acc_frg)(abs_acc_frg_ir, acc_frg.shape, acc_frg.dtype) + pvscale_f32x4 = cute.make_rmem_tensor(4, cutlass.Float32) + sfd_f8x4 = cute.make_rmem_tensor(4, self.sf_dtype) + tmp_f32 = abs_acc_frg[None, 0].reduce(cute.ReductionOp.MAX, cutlass.Float32(0.0), 0) * rcp_limit * norm_const + if tile_idx == 0: + pvscale[0] = tmp_f32 + elif tile_idx == 1: + pvscale[1] = tmp_f32 + elif tile_idx == 2: + pvscale[2] = tmp_f32 + elif tile_idx == 3: + pvscale[3] = tmp_f32 + pvscale_f32x4[0] = tmp_f32 + sfd_f8x4.store(pvscale_f32x4.load().to(self.sf_dtype)) + pvscale_f32x4.store(sfd_f8x4.load().to(cutlass.Float32)) + qpvscale_up = pvscale_f32x4[0] + fp32_max = cutlass.Float32(3.40282346638528859812e38) + acc_scale = norm_const * cute.arch.rcp_approx(qpvscale_up) + acc_scale = fmin(acc_scale, fp32_max, nan=True) + if cutlass.const_expr(self.vectorized_f32): + vec = tTR_rAcc_frg[None, 0] + for ei in cutlass.range_constexpr(0, self.sf_vec_size, 2): + vec[ei], vec[ei + 1] = cute.arch.mul_packed_f32x2( + (vec[ei], vec[ei + 1]), + (acc_scale, acc_scale), + rnd="rn", + ftz=False, + ) + else: + vec = tTR_rAcc_frg[None, 0] + for ei in cutlass.range_constexpr(self.sf_vec_size): + vec[ei] = vec[ei] * acc_scale + acc_vec = tiled_copy_r2s.retile(src).load() + tRSrD.store(acc_vec.to(self.d_dtype)) + + @cute.jit + def quant_sfd_col(self, tile_idx, tiled_copy_r2s, src, pvscale, norm_const, rcp_limit, tRSrD): + tTR_rAcc_frg = cute.logical_divide(src, cute.make_layout(self.sf_vec_size)) + acc_frg = tTR_rAcc_frg.load() + abs_acc_frg_ir = cutlass._mlir.dialects.math.absf(acc_frg.ir_value()) + acc_frg = type(acc_frg)(abs_acc_frg_ir, acc_frg.shape, acc_frg.dtype) + tmp_f32 = cutlass.Float32(0.0) + for vi in cutlass.range_constexpr(acc_frg.shape[0]): + max_value_original = ( + cutlass.Float32( + cute.arch.warp_redux_sync( + value=acc_frg[vi, 0], + kind="fmax", + mask_and_clamp=0xFFFFFFFF, + nan=True, + ) + ) + * rcp_limit + * norm_const + ) + max_value_vec = cute.full(4, max_value_original, dtype=cutlass.Float32) + max_value_vec_f8 = max_value_vec.to(cutlass.Float8E8M0FNU) + max_value_vec_f32_chunked = max_value_vec_f8.to(cutlass.Float32) + max_value = max_value_vec_f32_chunked[0] + tidx = cute.arch.thread_idx()[0] + if tidx % 32 == vi: + tmp_f32 = max_value + acc_scale_col = cutlass.Float32(0.0) + if max_value_vec_f32_chunked[0] == 0.000000: + acc_scale_col = cutlass.Float32(0.0) + else: + acc_scale_col = norm_const * cute.arch.rcp_approx(max_value_vec_f32_chunked[0]) + fp32_max = cutlass.Float32(3.40282346638528859812e38) + acc_scale_col = fmin(acc_scale_col, fp32_max) + tTR_rAcc_frg[vi] = tTR_rAcc_frg[vi] * acc_scale_col + pvscale[None, None, tile_idx][0] = tmp_f32 + acc_vec = tiled_copy_r2s.retile(src).load() + tRSrD.store(acc_vec.to(self.d_dtype)) + + @cute.jit + def tile_info_to_mn_idx(self, tile_info: cute.Tensor): + m_idx = tile_info[1] * cute.size(self.cta_tile_shape_mnk[0]) + n_idx = tile_info[2] * cute.size(self.cta_tile_shape_mnk[1]) + return m_idx, n_idx + + @cute.jit + def create_and_partition_new_SFDCol(self, tile_info, mSFDCol_mnl, padded_offsets): + m_idx, n_idx = self.tile_info_to_mn_idx(tile_info) + expert_idx = tile_info[0] + cumsum_tokens, tokens_this_group = compute_expert_token_range(padded_offsets, expert_idx) + n_total = cute.size(mSFDCol_mnl.shape[1]) + + sf_tile_idx_begin = cumsum_tokens // cute.size(mSFDCol_mnl.shape[0][0]) + mSFDCol_mnl_new_ptr = mSFDCol_mnl[(None, sf_tile_idx_begin), None, 0].iterator + + sfd_col_quant_layout = cute.tile_to_shape( + blockscaled_utils.BlockScaledBasicChunk(self.sf_vec_size, OperandMajorMode.MN).layout, + (tokens_this_group, n_total, mSFDCol_mnl.shape[2]), + (1, 2, 3), + ) + regPerSubtile = 4 + sfd_tile = (cute.make_layout(128), cute.make_layout(32 * regPerSubtile)) + mSFDCol_mnl_new = cute.make_tensor(mSFDCol_mnl_new_ptr, sfd_col_quant_layout) + gSFDCol_mnl_new = cute.local_tile(mSFDCol_mnl_new, sfd_tile, (None, None, None)) + + thr_layout = cute.make_ordered_layout((4, 32), order=(1, 0)) + val_layout = cute.make_ordered_layout((1,), order=(0,)) + copy_atom_sfd_col_quant = cute.make_copy_atom( + cute.nvgpu.CopyUniversalOp(), + gSFDCol_mnl_new.element_type, + num_bits_per_copy=8, + ) + tiled_copy_sfd_col_quant = cute.make_tiled_copy_tv( + copy_atom_sfd_col_quant, + thr_layout, + val_layout, + ) + tidx = cute.arch.thread_idx()[0] + thr_copy_sfd_col_quant = tiled_copy_sfd_col_quant.get_slice(tidx) + tCgSFDCol_mnl = thr_copy_sfd_col_quant.partition_D(cute.filter_zeros(gSFDCol_mnl_new)) + tCgSFDCol_mnl = cute.filter_zeros(tCgSFDCol_mnl) + return tCgSFDCol_mnl + + def epilog_tmem_copy_and_partition(self, tidx, tAcc, gD_mnl, epi_tile, use_2cta_instrs): + # For breuse: tAcc and gD_mnl have been through transform_partitioned_tensor_layout + # and have merged M-split into the first mode. No [0,0] selection needed. + # For non-breuse: tAcc has shape (MMA, 1, 1, STAGE); strip with [0,0] internally. + copy_atom_t2r = sm100_utils.get_tmem_load_op( + self.cta_tile_shape_mnk, + self.d_layout, + self.d_dtype, + self.acc_dtype, + epi_tile, + use_2cta_instrs, + ) + if cutlass.const_expr(self.enable_breuse): + tAcc_epi = cute.flat_divide(tAcc, epi_tile) + gD_mnl_epi = cute.flat_divide(gD_mnl, epi_tile) + else: + tAcc_epi = cute.flat_divide(tAcc[((None, None), 0, 0, None)], epi_tile) + gD_mnl_epi = cute.flat_divide(gD_mnl[((None, None), 0, 0, None, None, None)], epi_tile) + tiled_copy_t2r = tcgen05.make_tmem_copy(copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)]) + thr_copy_t2r = tiled_copy_t2r.get_slice(tidx) + tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi) + tTR_gC = thr_copy_t2r.partition_D(gD_mnl_epi) + tTR_rAcc = cute.make_rmem_tensor(tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype) + return tiled_copy_t2r, tTR_tAcc, tTR_rAcc + + def epilog_smem_copy_and_partition(self, tiled_copy_t2r, tTR_rD, tidx, sD): + copy_atom_r2s = sm100_utils.get_smem_store_op(self.d_layout, self.d_dtype, self.acc_dtype, tiled_copy_t2r) + tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r) + thr_copy_r2s = tiled_copy_r2s.get_slice(tidx) + tRS_sD = thr_copy_r2s.partition_D(sD) + tRS_rD = tiled_copy_r2s.retile(tTR_rD) + return tiled_copy_r2s, tRS_rD, tRS_sD + + @cute.kernel + def kernel( + self, + tiled_mma: cute.TiledMma, + tiled_mma_bkeep: Optional[cute.TiledMma], + tiled_mma_breuse: Optional[cute.TiledMma], + tiled_mma_sfb: cute.TiledMma, + tma_atom_a: cute.CopyAtom, + mA_mkl: cute.Tensor, + tma_atom_b: cute.CopyAtom, + mB_nkl: cute.Tensor, + tma_atom_sfa: cute.CopyAtom, + mSFA_mkl: cute.Tensor, + tma_atom_sfb: cute.CopyAtom, + mSFB_nkl: cute.Tensor, + tma_atom_c: cute.CopyAtom, + mC_mnl: cute.Tensor, + tma_atom_d: cute.CopyAtom, + mD_mnl: cute.Tensor, + tma_atom_d_col: cute.CopyAtom, + mD_col_mnl: cute.Tensor, + mSFDRow_mnl: Optional[cute.Tensor], + mSFDCol_mnl: Optional[cute.Tensor], + norm_const_tensor: Optional[cute.Tensor], + mAmax_tensor: Optional[cute.Tensor], + padded_offsets: cute.Tensor, + alpha: cute.Tensor, + mBias_nl: Optional[cute.Tensor], + prob: cute.Tensor, + workspace_ptr, + cluster_layout_vmnk: cute.Layout, + cluster_layout_sfb_vmnk: cute.Layout, + a_smem_layout_staged: cute.ComposedLayout, + b_smem_layout_staged: cute.ComposedLayout, + sfa_smem_layout_staged: cute.Layout, + sfb_smem_layout_staged: cute.Layout, + c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout, None], + d_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout, None], + bias_smem_layout_staged: Optional[cute.Layout], + epi_tile: cute.Tile, + sched_params: MoESchedulerParams, + epilogue_op: cutlass.Constexpr, + ): + """GPU device kernel for persistent MoE grouped GEMM with quantization.""" + warp_idx = cute.arch.warp_idx() + warp_idx = cute.arch.make_warp_uniform(warp_idx) + lane_idx = cute.arch.lane_idx() + + if warp_idx == self.tma_warp_id: + cpasync.prefetch_descriptor(tma_atom_a) + cpasync.prefetch_descriptor(tma_atom_sfa) + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DENSE): + cpasync.prefetch_descriptor(tma_atom_b) + cpasync.prefetch_descriptor(tma_atom_sfb) + cpasync.prefetch_descriptor(tma_atom_d) + if cutlass.const_expr(self.generate_sfd): + cpasync.prefetch_descriptor(tma_atom_d_col) + if cutlass.const_expr(self.generate_c): + cpasync.prefetch_descriptor(tma_atom_c) + + use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2 + total_token = padded_offsets[self.expert_cnt - 1] + + bidx, bidy, bidz = cute.arch.block_idx() + mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape) + is_leader_cta = mma_tile_coord_v == 0 + cta_rank_in_cluster = cute.arch.make_warp_uniform(cute.arch.block_idx_in_cluster()) + block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord(cta_rank_in_cluster) + block_in_cluster_coord_sfb_vmnk = cluster_layout_sfb_vmnk.get_flat_coord(cta_rank_in_cluster) + tidx, _, _ = cute.arch.thread_idx() + + smem = utils.SmemAllocator() + storage = smem.allocate(self.shared_storage) + sched_storage = storage.scheduler + + ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1 + ab_pipeline_consumer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread, num_tma_producer) + ab_pipeline = pipeline.PipelineTmaUmma.create( + barrier_storage=storage.ab_mbar_ptr.data_ptr(), + num_stages=self.num_ab_stage, + producer_group=ab_pipeline_producer_group, + consumer_group=ab_pipeline_consumer_group, + tx_count=self.num_tma_load_bytes, + cta_layout_vmnk=cluster_layout_vmnk, + ) + + acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + num_acc_consumer_threads = len(self.epilog_warp_id) * (2 if use_2cta_instrs else 1) + acc_pipeline_consumer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread, num_acc_consumer_threads) + acc_pipeline = pipeline.PipelineUmmaAsync.create( + barrier_storage=storage.acc_mbar_ptr.data_ptr(), + num_stages=self.num_acc_stage, + producer_group=acc_pipeline_producer_group, + consumer_group=acc_pipeline_consumer_group, + cta_layout_vmnk=cluster_layout_vmnk, + ) + + tile_info_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread, self.threads_per_warp * 1) + tile_info_pipeline_consumer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread, self.threads_wo_sched) + tile_info_pipeline = pipeline.PipelineAsync.create( + barrier_storage=sched_storage.tile_info_mbar.data_ptr(), + num_stages=self.num_tile_stage, + producer_group=tile_info_pipeline_producer_group, + consumer_group=tile_info_pipeline_consumer_group, + ) + + if cutlass.const_expr(self.enable_bias): + bias_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread, self.threads_per_warp) + bias_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + self.threads_per_warp * len(self.epilog_warp_id), + ) + bias_pipeline = pipeline.PipelineCpAsync.create( + barrier_storage=storage.bias_mbar_ptr.data_ptr(), + num_stages=self.num_bias_stage, + producer_group=bias_pipeline_producer_group, + consumer_group=bias_pipeline_consumer_group, + ) + sBias = storage.sBias.get_tensor(bias_smem_layout_staged) + + scheduler = MoEPersistentTileScheduler.create( + sched_params, + padded_offsets, + cute.arch.block_idx(), + cute.arch.grid_dim(), + counter_ptr=self._get_sched_counter_ptr(workspace_ptr), + sched_storage=sched_storage, + ) + scheduler.internal_init() + + tmem = utils.TmemAllocator( + storage.tmem_holding_buf.ptr, + barrier_for_retrieve=self.tmem_alloc_barrier, + allocator_warp_id=self.epilog_warp_id[0], + is_two_cta=use_2cta_instrs, + two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr.ptr, + arch="sm_107", + ) + + if cute.size(self.cluster_shape_mn) > 1: + cute.arch.cluster_arrive_relaxed() + + sC = storage.sC.get_tensor(c_smem_layout_staged.outer, swizzle=c_smem_layout_staged.inner) + sD = storage.sD.get_tensor(d_smem_layout_staged.outer, swizzle=d_smem_layout_staged.inner) + sD_col = sD + if cutlass.const_expr(self.generate_sfd): + sD_col = storage.sD_col.get_tensor(d_smem_layout_staged.outer, swizzle=d_smem_layout_staged.inner) + sA = storage.sA.get_tensor(a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner) + sB = storage.sB.get_tensor(b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner) + sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged) + sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged) + amax_layout = cute.make_layout((self.num_epilog_warps,)) + sAmax = storage.sAmax.get_tensor(amax_layout) + info_layout = cute.make_layout((4, self.num_tile_stage), stride=(1, 4)) + sInfo = sched_storage.sInfo.get_tensor(info_layout) + + # Multicast masks — must create ALL when any mcast or 2CTA is active + a_full_mcast_mask = None + b_full_mcast_mask = None + sfa_full_mcast_mask = None + sfb_full_mcast_mask = None + if cutlass.const_expr(self.is_a_mcast or self.is_b_mcast or use_2cta_instrs): + a_full_mcast_mask = cpasync.create_tma_multicast_mask(cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2) + b_full_mcast_mask = cpasync.create_tma_multicast_mask(cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1) + sfa_full_mcast_mask = cpasync.create_tma_multicast_mask(cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2) + sfb_full_mcast_mask = cpasync.create_tma_multicast_mask(cluster_layout_sfb_vmnk, block_in_cluster_coord_sfb_vmnk, mcast_mode=1) + + # MMA partition (for tCtAcc_fake shape computation only) + thr_mma_common = tiled_mma.get_slice(0) + tCsA_common = thr_mma_common.partition_A(sA) + tCsB_common = thr_mma_common.partition_B(sB) + tCsA_common = cute.filter_zeros(tCsA_common) + tCsB_common = cute.filter_zeros(tCsB_common) + + # SMEM fragments for MMA (used by MMA warp) + tCrA = tiled_mma.make_fragment_A(sA) + tCrB = tiled_mma.make_fragment_B(sB) + + # TMEM accumulator shape + acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2]) + if cutlass.const_expr(self.overlapping_accum): + num_acc_stage_overlapped = 2 + tCtAcc_fake = tiled_mma.make_fragment_C(cute.append(acc_shape, num_acc_stage_overlapped)) + tCtAcc_fake = cute.make_tensor( + tCtAcc_fake.iterator, + cute.make_layout( + tCtAcc_fake.shape, + stride=( + tCtAcc_fake.stride[0], + tCtAcc_fake.stride[1], + tCtAcc_fake.stride[2], + (256 - self.num_sf_tmem_cols) * tCtAcc_fake.stride[0][1], + ), + ), + ) + elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192): + tCtAcc_fake = tiled_mma.make_fragment_C(cute.append(acc_shape, self.num_acc_stage)) + tCtAcc_fake = cute.make_tensor( + tCtAcc_fake.iterator, + cute.make_layout( + tCtAcc_fake.shape, + stride=( + tCtAcc_fake.stride[0], + tCtAcc_fake.stride[1], + tCtAcc_fake.stride[2], + self.num_accumulator_tmem_stride, + ), + ), + ) + else: + tCtAcc_fake = tiled_mma.make_fragment_C(cute.append(acc_shape, self.num_acc_stage)) + + # Cluster sync before warp specialization + if cute.size(self.cluster_shape_mn) > 1: + cute.arch.cluster_wait() + else: + self.cta_sync_barrier.arrive_and_wait() + + if total_token <= 0: + cute.arch.nvvm.exit() + + # ============================================================== + # Scheduler warp (MoE Persistent Tile Scheduler) + # ============================================================== + if warp_idx == self.sched_warp_id: + work_tile_info = scheduler.initial_work_tile_info() + tile_info_producer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Producer, self.num_tile_stage) + while work_tile_info.is_valid_tile: + tile_info_pipeline.producer_acquire(tile_info_producer_state) + with cute.arch.elect_one(): + sInfo[(0, tile_info_producer_state.index)] = work_tile_info.expert_idx + sInfo[(1, tile_info_producer_state.index)] = work_tile_info.tile_m_idx + sInfo[(2, tile_info_producer_state.index)] = work_tile_info.tile_n_idx + sInfo[(3, tile_info_producer_state.index)] = work_tile_info.k_tile_cnt + cute.arch.fence_proxy("async.shared", space="cta") + self.sched_sync_barrier.arrive_and_wait() + tile_info_pipeline.producer_commit(tile_info_producer_state) + tile_info_producer_state.advance() + work_tile_info = scheduler.advance_to_next_work() + + tile_info_pipeline.producer_acquire(tile_info_producer_state) + with cute.arch.elect_one(): + sInfo[(0, tile_info_producer_state.index)] = cutlass.Int32(-1) + sInfo[(1, tile_info_producer_state.index)] = cutlass.Int32(0) + sInfo[(2, tile_info_producer_state.index)] = cutlass.Int32(0) + sInfo[(3, tile_info_producer_state.index)] = cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + self.sched_sync_barrier.arrive_and_wait() + tile_info_pipeline.producer_commit(tile_info_producer_state) + tile_info_producer_state.advance() + tile_info_pipeline.producer_tail(tile_info_producer_state) + + # ============================================================== + # Bias load warp + # ============================================================== + if cutlass.const_expr(self.enable_bias): + if warp_idx == self.bias_load_warp_id: + bias_ext = self._make_extension(workspace_ptr) + bias_producer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Producer, self.num_bias_stage) + tile_info_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, self.num_tile_stage) + bias_g2s_atom = cute.make_copy_atom( + cute.nvgpu.cpasync.CopyG2SOp(cache_mode=cute.nvgpu.cpasync.LoadCacheMode.GLOBAL), + self.bias_dtype, + num_bits_per_copy=128, + ) + bias_g2s_tiled = cute.make_tiled_copy_tv( + bias_g2s_atom, + cute.make_layout((32,)), + cute.make_layout((8,)), + ) + thr_bias_g2s = bias_g2s_tiled.get_slice(cute.arch.lane_idx()) + tBs_sBias = thr_bias_g2s.partition_D(sBias) + + tile_info = cute.make_rmem_tensor((4,), cutlass.Int32) + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + while is_valid_tile: + bias_producer_state.reset_count() + work_tile_info = MoEWorkTileInfo( + expert_idx=tile_info[0], + tile_m_idx=tile_info[1], + tile_n_idx=tile_info[2], + k_tile_cnt=tile_info[3], + ) + bias_ext.update_expert_info(padded_offsets, work_tile_info.expert_idx) + real_bias, _ = bias_ext.get_gmem_tensor("bias", mBias_nl, padded_offsets, work_tile_info) + gBias_expert = cute.local_tile(real_bias, cute.slice_(self.mma_tiler[:2], (0, None)), (None, None)) + bias_tile = gBias_expert[(None, work_tile_info.tile_n_idx, 0)] + bias_identity_tensor = cute.make_identity_tensor(bias_tile.shape) + bias_partitioned_by_g2s = thr_bias_g2s.partition_S(bias_tile) + bias_coord_partitioned_by_g2s = thr_bias_g2s.partition_S(bias_identity_tensor) + + residue_n = sched_params.intermediate - work_tile_info.tile_n_idx * self.cta_tile_shape_mnk[1] + bias_pred_tensor = cute.make_rmem_tensor(bias_coord_partitioned_by_g2s[(None, 0)].shape, cutlass.Boolean) + for vi in cutlass.range_constexpr(cute.size(bias_pred_tensor)): + bias_pred_tensor[vi] = cute.elem_less(bias_coord_partitioned_by_g2s[(vi, 0)], (residue_n,)) + bias_pred_tensor = bias_pred_tensor[((0, None),)] + + bias_pipeline.producer_acquire(bias_producer_state) + cute.copy(bias_g2s_tiled, bias_partitioned_by_g2s[(None, 0)], tBs_sBias[(None, 0, bias_producer_state.index)], pred=bias_pred_tensor) + bias_pipeline.producer_commit(bias_producer_state) + bias_producer_state.advance() + + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + bias_pipeline.producer_tail(bias_producer_state) + + # ============================================================== + # DMA / TMA load warp + # ============================================================== + if warp_idx == self.tma_warp_id: + ext = self._make_extension(workspace_ptr) + ab_producer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Producer, self.num_ab_stage) + tile_info_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, self.num_tile_stage) + + tile_info = cute.make_rmem_tensor((4,), cutlass.Int32) + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + while is_valid_tile: + work_tile_info = MoEWorkTileInfo( + expert_idx=tile_info[0], + tile_m_idx=tile_info[1], + tile_n_idx=tile_info[2], + k_tile_cnt=tile_info[3], + ) + k_tile_cnt = work_tile_info.k_tile_cnt + ext.update_expert_info(padded_offsets, work_tile_info.expert_idx) + + real_a, _ = ext.get_gmem_tensor("a", mA_mkl, padded_offsets, work_tile_info) + real_b, desc_ptr_b = ext.get_gmem_tensor("b", mB_nkl, padded_offsets, work_tile_info) + real_sfa, _ = ext.get_gmem_tensor("sfa", mSFA_mkl, padded_offsets, work_tile_info) + real_sfb, desc_ptr_sfb = ext.get_gmem_tensor("sfb", mSFB_nkl, padded_offsets, work_tile_info) + + # N=192: the SFB TMEM layout uses a nested ((2,2), y) shape to + # map the 192-column tile onto TMEM correctly. + if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192): + x = real_sfb.stride[0][1] + y = cute.ceil_div(real_sfb.shape[0][1], 4) + new_shape = ( + (real_sfb.shape[0][0], ((2, 2), y)), + real_sfb.shape[1], + real_sfb.shape[2], + ) + new_stride = ( + (real_sfb.stride[0][0], ((x, x), 3 * x)), + real_sfb.stride[1], + real_sfb.stride[2], + ) + real_sfb = cute.make_tensor( + real_sfb.iterator, + cute.make_layout(new_shape, stride=new_stride), + ) + + gA_mkl = cute.local_tile(real_a, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)) + gB_nkl = cute.local_tile(real_b, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)) + gSFA_mkl = cute.local_tile(real_sfa, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)) + gSFB_nkl = cute.local_tile(real_sfb, cute.slice_(self.mma_tiler_sfb, (0, None, None)), (None, None, None)) + + # MMA partition on gmem tensors + thr_mma_dma = tiled_mma.get_slice(mma_tile_coord_v) + thr_mma_sfb_dma = tiled_mma_sfb.get_slice(mma_tile_coord_v) + tCgA = thr_mma_dma.partition_A(gA_mkl) + tCgB = thr_mma_dma.partition_B(gB_nkl) + tCgSFA = thr_mma_dma.partition_A(gSFA_mkl) + tCgSFB = thr_mma_sfb_dma.partition_B(gSFB_nkl) + + # TMA partition A + a_cta_layout = cute.make_layout(cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape) + tAsA, tAgA = cpasync.tma_partition( + tma_atom_a, + block_in_cluster_coord_vmnk[2], + a_cta_layout, + cute.group_modes(sA, 0, 3), + cute.group_modes(tCgA, 0, 3), + ) + # TMA partition B + b_cta_layout = cute.make_layout(cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape) + tBsB, tBgB = cpasync.tma_partition( + tma_atom_b, + block_in_cluster_coord_vmnk[1], + b_cta_layout, + cute.group_modes(sB, 0, 3), + cute.group_modes(tCgB, 0, 3), + ) + # TMA partition SFA + sfa_cta_layout = a_cta_layout + tAsSFA, tAgSFA = cpasync.tma_partition( + tma_atom_sfa, + block_in_cluster_coord_vmnk[2], + sfa_cta_layout, + cute.group_modes(sSFA, 0, 3), + cute.group_modes(tCgSFA, 0, 3), + ) + tAsSFA = cute.filter_zeros(tAsSFA) + tAgSFA = cute.filter_zeros(tAgSFA) + # TMA partition SFB + sfb_cta_layout = cute.make_layout(cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape) + tBsSFB, tBgSFB = cpasync.tma_partition( + tma_atom_sfb, + block_in_cluster_coord_sfb_vmnk[1], + sfb_cta_layout, + cute.group_modes(sSFB, 0, 3), + cute.group_modes(tCgSFB, 0, 3), + ) + tBsSFB = cute.filter_zeros(tBsSFB) + tBgSFB = cute.filter_zeros(tBgSFB) + + mma_tile_coord_m = work_tile_info.tile_m_idx // cute.size(tiled_mma.thr_id.shape) + mma_tile_coord_n = work_tile_info.tile_n_idx + tAgA_slice = tAgA[(None, mma_tile_coord_m, None, 0)] + tBgB_slice = tBgB[(None, mma_tile_coord_n, None, 0)] + tAgSFA_slice = tAgSFA[(None, mma_tile_coord_m, None, 0)] + slice_n = mma_tile_coord_n + if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64): + slice_n = mma_tile_coord_n // 2 + tBgSFB_slice = tBgSFB[(None, slice_n, None, 0)] + + ab_producer_state.reset_count() + peek_ab_empty_status = cutlass.Boolean(1) + if ab_producer_state.count < k_tile_cnt: + peek_ab_empty_status = ab_pipeline.producer_try_acquire(ab_producer_state) + + for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1): + tAgA_k = tAgA_slice[(None, ab_producer_state.count)] + tBgB_k = tBgB_slice[(None, ab_producer_state.count)] + tAgSFA_k = tAgSFA_slice[(None, ab_producer_state.count)] + tBgSFB_k = tBgSFB_slice[(None, ab_producer_state.count)] + tAsA_pipe = tAsA[(None, ab_producer_state.index)] + tBsB_pipe = tBsB[(None, ab_producer_state.index)] + tAsSFA_pipe = tAsSFA[(None, ab_producer_state.index)] + tBsSFB_pipe = tBsSFB[(None, ab_producer_state.index)] + + tma_bar = ab_pipeline.producer_get_barrier(ab_producer_state) + ab_pipeline.producer_acquire(ab_producer_state, peek_ab_empty_status) + + cute.copy(tma_atom_a, tAgA_k, tAsA_pipe, tma_bar_ptr=tma_bar, mcast_mask=a_full_mcast_mask) + cute.copy(tma_atom_b, tBgB_k, tBsB_pipe, tma_bar_ptr=tma_bar, mcast_mask=b_full_mcast_mask, tma_desc_ptr=desc_ptr_b) + cute.copy(tma_atom_sfa, tAgSFA_k, tAsSFA_pipe, tma_bar_ptr=tma_bar, mcast_mask=sfa_full_mcast_mask) + cute.copy(tma_atom_sfb, tBgSFB_k, tBsSFB_pipe, tma_bar_ptr=tma_bar, mcast_mask=sfb_full_mcast_mask, tma_desc_ptr=desc_ptr_sfb) + + ab_producer_state.advance() + peek_ab_empty_status = cutlass.Boolean(1) + if ab_producer_state.count < k_tile_cnt: + peek_ab_empty_status = ab_pipeline.producer_try_acquire(ab_producer_state) + + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + ab_pipeline.producer_tail(ab_producer_state) + + # ============================================================== + # MMA warp + # ============================================================== + if warp_idx == self.mma_warp_id: + tmem.wait_for_alloc() + acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype) + tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout) + + # SFA TMEM tensor + sfa_tmem_ptr = cute.recast_ptr( + acc_tmem_ptr + self.num_accumulator_tmem_cols, + dtype=self.sf_dtype, + ) + tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa( + tiled_mma, + self.mma_tiler, + self.sf_vec_size, + cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)), + ) + tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout) + + # SFB TMEM tensor + sfb_tmem_ptr = cute.recast_ptr( + acc_tmem_ptr + self.num_accumulator_tmem_cols + self.num_sfa_tmem_cols, + dtype=self.sf_dtype, + ) + tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb( + tiled_mma, + self.mma_tiler, + self.sf_vec_size, + cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)), + ) + tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout) + + # S2T copy partition for SFA/SFB + ( + tiled_copy_s2t_sfa, + tCsSFA_compact_s2t, + tCtSFA_compact_s2t, + ) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA) + ( + tiled_copy_s2t_sfb, + tCsSFB_compact_s2t, + tCtSFB_compact_s2t, + ) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB) + + ab_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, self.num_ab_stage) + acc_producer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Producer, self.num_acc_stage) + + tile_info_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, self.num_tile_stage) + tile_info = cute.make_rmem_tensor((4,), cutlass.Int32) + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + while is_valid_tile: + k_tile_cnt = tile_info[3] + + # Peek AB buffer full + ab_consumer_state.reset_count() + peek_ab_full_status = cutlass.Boolean(1) + if ab_consumer_state.count < k_tile_cnt and is_leader_cta: + peek_ab_full_status = ab_pipeline.consumer_try_wait(ab_consumer_state) + + # Peek Acc buffer empty + acc_producer_state.reset_count() + peek_acc_empty_status = cutlass.Boolean(1) + if ab_consumer_state.count < k_tile_cnt and is_leader_cta: + peek_acc_empty_status = acc_pipeline.producer_try_acquire(acc_producer_state) + + mma_tile_coord_mnl = ( + tile_info[1] // cute.size(tiled_mma.thr_id.shape), + tile_info[2], + tile_info[0], + ) + + if cutlass.const_expr(self.overlapping_accum): + acc_stage_index = acc_producer_state.phase ^ 1 + else: + acc_stage_index = acc_producer_state.index + + tCtAcc = tCtAcc_base[(None, None, None, acc_stage_index)] + + tCtSFB_mma = tCtSFB + if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192): + offset = cutlass.Int32(2) if mma_tile_coord_mnl[1] % 2 == 1 else cutlass.Int32(0) + shifted_ptr = cute.recast_ptr( + acc_tmem_ptr + self.num_accumulator_tmem_cols + self.num_sfa_tmem_cols + offset, + dtype=self.sf_dtype, + ) + tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout) + elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64): + offset = cutlass.Int32((mma_tile_coord_mnl[1] % 2) * 2) + shifted_ptr = cute.recast_ptr( + acc_tmem_ptr + self.num_accumulator_tmem_cols + self.num_sfa_tmem_cols + offset, + dtype=self.sf_dtype, + ) + tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout) + + if is_leader_cta: + acc_pipeline.producer_acquire(acc_producer_state, peek_acc_empty_status) + + tiled_mma.set(tcgen05.Field.ACCUMULATE, False) + + for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1): + if is_leader_cta: + ab_pipeline.consumer_wait(ab_consumer_state, peek_ab_full_status) + + s2t_stage_coord = (None, None, None, None, ab_consumer_state.index) + cute.copy(tiled_copy_s2t_sfa, tCsSFA_compact_s2t[s2t_stage_coord], tCtSFA_compact_s2t) + cute.copy(tiled_copy_s2t_sfb, tCsSFB_compact_s2t[s2t_stage_coord], tCtSFB_compact_s2t) + + num_kblocks = cute.size(tCrA, mode=[2]) + ab_consumer_state_next = ab_consumer_state.clone() + ab_consumer_state_next.advance() + if ab_consumer_state_next.count < k_tile_cnt: + peek_ab_full_status = ab_pipeline.consumer_try_wait(ab_consumer_state_next) + + for kblock_idx in cutlass.range(num_kblocks, unroll_full=True): + if cutlass.const_expr(self.enable_breuse and cute.size(tCtAcc.layout, mode=[1]) == 2 and cute.size(tCtAcc.layout, mode=[2]) == 1): + tCtAcc_bkeep = tCtAcc[(None, 0, 0)] + tCtAcc_breuse = tCtAcc[(None, 1, 0)] + + a_kblk_crd_keep = (None, 0, kblock_idx, ab_consumer_state.index) + a_kblk_crd_reuse = (None, 1, kblock_idx, ab_consumer_state.index) + b_kblk_crd = (None, 0, kblock_idx, ab_consumer_state.index) + sfa_kblk_crd_keep = (None, 0, kblock_idx) + sfa_kblk_crd_reuse = (None, 1, kblock_idx) + sfb_kblk_crd = (None, 0, kblock_idx) + + # Bkeep: accumulate A_upper × B (keeps B in reuse buffer) + tiled_mma_bkeep.set( + tcgen05.Field.ACCUMULATE, + k_tile != 0 or kblock_idx != 0, + ) + cute.gemm( + tiled_mma_bkeep, + tCtAcc_bkeep, + [tCrA[a_kblk_crd_keep], tCtSFA[sfa_kblk_crd_keep]], + [tCrB[b_kblk_crd], tCtSFB_mma[sfb_kblk_crd]], + tCtAcc_bkeep, + ) + # Breuse: accumulate A_lower × B (reuses B from buffer) + tiled_mma_breuse.set( + tcgen05.Field.ACCUMULATE, + k_tile != 0 or kblock_idx != 0, + ) + cute.gemm( + tiled_mma_breuse, + tCtAcc_breuse, + [tCrA[a_kblk_crd_reuse], tCtSFA[sfa_kblk_crd_reuse]], + [tCrB[b_kblk_crd], tCtSFB_mma[sfb_kblk_crd]], + tCtAcc_breuse, + ) + else: + kblock_coord = (None, None, kblock_idx, ab_consumer_state.index) + sf_kblock_coord = (None, None, kblock_idx) + tiled_mma.set(tcgen05.Field.SFA, tCtSFA[sf_kblock_coord].iterator) + tiled_mma.set(tcgen05.Field.SFB, tCtSFB_mma[sf_kblock_coord].iterator) + cute.gemm(tiled_mma, tCtAcc, tCrA[kblock_coord], tCrB[kblock_coord], tCtAcc) + tiled_mma.set(tcgen05.Field.ACCUMULATE, True) + + ab_pipeline.consumer_release(ab_consumer_state) + ab_consumer_state = ab_consumer_state_next + + if is_leader_cta: + acc_pipeline.producer_commit(acc_producer_state) + + acc_producer_state.advance() + if acc_producer_state.count < k_tile_cnt: + if is_leader_cta: + peek_acc_empty_status = acc_pipeline.producer_try_acquire(acc_producer_state) + + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + acc_pipeline.producer_tail(acc_producer_state) + + # ============================================================== + # Epilogue warps + # ============================================================== + if warp_idx < self.mma_warp_id and total_token > 0: + tmem.allocate(self.num_tmem_alloc_cols) + tmem.wait_for_alloc() + tmem_ptr = tmem.retrieve_ptr(self.acc_dtype) + tCtAcc_base = cute.make_tensor(tmem_ptr, tCtAcc_fake.layout) + + epi_tidx = tidx + thr_mma_epi = tiled_mma.get_slice(mma_tile_coord_v) + + # Shape-only partition on global tensor (invariant setup for t2r copy atom) + gD_mnl_shape = cute.local_tile(mD_mnl, cute.slice_(self.mma_tiler_d, (None, None, 0)), (None, None, None)) + tCgD_shape = thr_mma_epi.partition_C(gD_mnl_shape) + + if cutlass.const_expr(self.enable_breuse): + # Merge bkeep/breuse M-split (mode1=2) into the M dimension for epilogue + tCtAcc_epi_input = transform_partitioned_tensor_layout(tCtAcc_base) + tCgD_epi_input = transform_partitioned_tensor_layout(tCgD_shape) + else: + tCtAcc_epi_input = tCtAcc_base + tCgD_epi_input = tCgD_shape + + tiled_copy_t2r, tTR_tAcc_base, tTR_rAcc = self.epilog_tmem_copy_and_partition( + epi_tidx, + tCtAcc_epi_input, + tCgD_epi_input, + epi_tile, + use_2cta_instrs, + ) + + tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, self.c_dtype) + tiled_copy_r2s, tRS_rC, tRS_sC = self.epilog_smem_copy_and_partition( + tiled_copy_t2r, + tTR_rC, + epi_tidx, + sC, + ) + tTR_rD = cute.make_rmem_tensor(tTR_rAcc.shape, self.d_dtype) + tiled_copy_r2s, tRS_rD, tRS_sD = self.epilog_smem_copy_and_partition( + tiled_copy_t2r, + tTR_rD, + epi_tidx, + sD, + ) + tTR_rD_col = cute.make_rmem_tensor(tTR_rAcc.shape, self.d_dtype) + tiled_copy_r2s, tRS_rD_col, tRS_sD_col = self.epilog_smem_copy_and_partition( + tiled_copy_t2r, + tTR_rD_col, + epi_tidx, + sD_col, + ) + + if cutlass.const_expr(self.generate_sfd): + norm_const = norm_const_tensor[0] + regPerSubtile = 4 + sfd_row_tile = (cute.make_layout(128), cute.make_layout(32 * regPerSubtile)) + gSFDRow_mnl = cute.local_tile(mSFDRow_mnl, sfd_row_tile, (None, None, None)) + thr_copy_t2r_local = tiled_copy_t2r.get_slice(tidx) + tCgSFDRow_mnl = thr_copy_t2r_local.partition_D(gSFDRow_mnl) + tCgSFDRow_mnl = cute.filter_zeros(tCgSFDRow_mnl) + tCrSFDRow = cute.make_rmem_tensor(tCgSFDRow_mnl[(None, None, None, 0, 0, 0)].layout, self.sf_dtype) + tCrSFDRow_pvscale = cute.make_rmem_tensor_like(tCrSFDRow, cutlass.Float32) + # For breuse, bkeep and breuse M groups interleave in the subtile loop; + # they need separate pvscale arrays to avoid corrupting each other. + if cutlass.const_expr(self.enable_breuse): + tCrSFDRow_pvscale_br = cute.make_rmem_tensor_like(tCrSFDRow, cutlass.Float32) + d_rcp_limits = get_dtype_rcp_limits(self.d_dtype) + + sfd_col_tile = sfd_row_tile + gSFDCol_mnl = cute.local_tile(mSFDCol_mnl, sfd_col_tile, (None, None, None)) + thr_layout = cute.make_ordered_layout((4, 32), order=(1, 0)) + val_layout = cute.make_ordered_layout((1,), order=(0,)) + copy_atom_sfd_col = cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), gSFDCol_mnl.element_type, num_bits_per_copy=8) + tiled_copy_sfd_col = cute.make_tiled_copy_tv(copy_atom_sfd_col, thr_layout, val_layout) + thr_copy_sfd_col = tiled_copy_sfd_col.get_slice(tidx) + tCgSFDCol_mnl = thr_copy_sfd_col.partition_D(cute.filter_zeros(gSFDCol_mnl)) + tCgSFDCol_mnl = cute.filter_zeros(tCgSFDCol_mnl) + tCrSFDCol = cute.make_rmem_tensor(tCgSFDRow_mnl[(None, None, None, 0, 0, 0)].shape, self.sf_dtype) + tCrSFDCol_pvscale = cute.make_rmem_tensor_like(tCrSFDRow, cutlass.Float32) + if cutlass.const_expr(self.enable_breuse): + tCrSFDCol_pvscale_br = cute.make_rmem_tensor_like(tCrSFDRow, cutlass.Float32) + + epi_ext = self._make_extension(workspace_ptr) + + acc_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, self.num_acc_stage) + c_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread, 32 * len(self.epilog_warp_id)) + c_pipeline = pipeline.PipelineTmaStore.create(num_stages=self.num_c_stage, producer_group=c_producer_group) + d_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread, 32 * len(self.epilog_warp_id)) + d_pipeline = pipeline.PipelineTmaStore.create(num_stages=self.num_d_stage, producer_group=d_producer_group) + + tile_info_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, self.num_tile_stage) + tile_info = cute.make_rmem_tensor((4,), cutlass.Int32) + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + if cutlass.const_expr(self.enable_bias): + bias_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, self.num_bias_stage) + bias_s2r_tom = cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), self.bias_dtype, num_bits_per_copy=128) + tTR_rBias = cute.make_rmem_tensor(cute.make_layout(self.epi_tile[1]), self.bias_dtype) + + num_prev_subtiles = cutlass.Int32(0) + while is_valid_tile: + epi_work_tile_info = MoEWorkTileInfo( + expert_idx=tile_info[0], + tile_m_idx=tile_info[1], + tile_n_idx=tile_info[2], + k_tile_cnt=tile_info[3], + ) + expert_idx = epi_work_tile_info.expert_idx + epi_ext.update_expert_info(padded_offsets, expert_idx) + + alpha_val = alpha[expert_idx] + + if cutlass.const_expr(self.enable_bias): + bias_consumer_state.reset_count() + bias_pipeline.consumer_wait(bias_consumer_state) + sBias_stage = sBias[(None, bias_consumer_state.index)] + sBias_subtiles = cute.flat_divide(sBias_stage, cute.make_layout(self.epi_tile[1])) + + real_d, _ = epi_ext.get_gmem_tensor("d", mD_mnl, padded_offsets, epi_work_tile_info) + real_c, _ = epi_ext.get_gmem_tensor("c", mC_mnl, padded_offsets, epi_work_tile_info) + real_d_col = real_d + if cutlass.const_expr(self.generate_sfd): + real_d_col, _ = epi_ext.get_gmem_tensor("d_col", mD_col_mnl, padded_offsets, epi_work_tile_info) + + thr_mma_epi_loop = tiled_mma.get_slice(mma_tile_coord_v) + + gD_mnl_loop = cute.local_tile(real_d, cute.slice_(self.mma_tiler_d, (None, None, 0)), (None, None, None)) + tCgD_loop = thr_mma_epi_loop.partition_C(gD_mnl_loop) + if cutlass.const_expr(self.enable_breuse): + tCgD_loop_epi = transform_partitioned_tensor_layout(tCgD_loop) + gD_epi_tr = cute.flat_divide(tCgD_loop_epi, epi_tile) + sD_for_tma = cute.group_modes(sD, 0, 2) + gD_for_tma_tr = cute.group_modes(gD_epi_tr, 0, 2) + bSG_sD, bSG_gD_partitioned = cpasync.tma_partition( + tma_atom_d, + 0, + cute.make_layout(1), + sD_for_tma, + gD_for_tma_tr, + ) + else: + _, bSG_sD, bSG_gD_partitioned = epilog_gmem_copy_and_partition( + epi_tidx, + tma_atom_d, + tCgD_loop, + epi_tile, + sD, + ) + + gC_mnl_loop = cute.local_tile(real_c, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None)) + tCgC_loop = thr_mma_epi_loop.partition_C(gC_mnl_loop) + if cutlass.const_expr(self.enable_breuse): + tCgC_loop_epi = transform_partitioned_tensor_layout(tCgC_loop) + gC_epi_tr = cute.flat_divide(tCgC_loop_epi, epi_tile) + sC_for_tma = cute.group_modes(sC, 0, 2) + gC_for_tma_tr = cute.group_modes(gC_epi_tr, 0, 2) + bSG_sC, bSG_gC_partitioned = cpasync.tma_partition( + tma_atom_c, + 0, + cute.make_layout(1), + sC_for_tma, + gC_for_tma_tr, + ) + else: + _, bSG_sC, bSG_gC_partitioned = epilog_gmem_copy_and_partition( + epi_tidx, + tma_atom_c, + tCgC_loop, + epi_tile, + sC, + ) + + gD_col_mnl_loop = gD_mnl_loop + tCgD_col_loop = tCgD_loop + if cutlass.const_expr(self.generate_sfd): + gD_col_mnl_loop = cute.local_tile(real_d_col, cute.slice_(self.mma_tiler_d, (None, None, 0)), (None, None, None)) + tCgD_col_loop = thr_mma_epi_loop.partition_C(gD_col_mnl_loop) + if cutlass.const_expr(self.enable_breuse): + tCgD_col_loop_epi = transform_partitioned_tensor_layout(tCgD_col_loop) + gD_col_epi_tr = cute.flat_divide(tCgD_col_loop_epi, epi_tile) + sD_col_for_tma = cute.group_modes(sD_col, 0, 2) + gD_col_for_tma_tr = cute.group_modes(gD_col_epi_tr, 0, 2) + bSG_sD_col, bSG_gD_col_partitioned = cpasync.tma_partition( + tma_atom_d_col, + 0, + cute.make_layout(1), + sD_col_for_tma, + gD_col_for_tma_tr, + ) + else: + _, bSG_sD_col, bSG_gD_col_partitioned = epilog_gmem_copy_and_partition( + epi_tidx, + tma_atom_d_col, + tCgD_col_loop, + epi_tile, + sD_col, + ) + + epi_mma_tile_coord = ( + epi_work_tile_info.tile_m_idx // cute.size(tiled_mma.thr_id.shape), + epi_work_tile_info.tile_n_idx, + 0, + ) + bSG_gC = bSG_gC_partitioned[(None, None, None, *epi_mma_tile_coord)] + bSG_gD = bSG_gD_partitioned[(None, None, None, *epi_mma_tile_coord)] + bSG_gD_col = bSG_gD_col_partitioned[(None, None, None, *epi_mma_tile_coord)] + bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC)) + bSG_gD = cute.group_modes(bSG_gD, 1, cute.rank(bSG_gD)) + bSG_gD_col = cute.group_modes(bSG_gD_col, 1, cute.rank(bSG_gD_col)) + + if cutlass.const_expr(self.generate_sfd): + tCgSFDRow_mn = tCgSFDRow_mnl[(None, None, None, None, None, 0)] + tCgSFDCol_mnl_new = tCgSFDCol_mnl + if cutlass.const_expr(self.discrete_col_sfd): + tCgSFDCol_mnl_new = self.create_and_partition_new_SFDCol(tile_info, mSFDCol_mnl, padded_offsets) + tCgSFDCol_mn = tCgSFDCol_mnl_new[(None, None, None, None, None, 0)] + + if cutlass.const_expr(self.generate_amax): + thread_tile_amax = cutlass.Float32(0.0) + + real_prob, _ = epi_ext.get_gmem_tensor("prob", prob, padded_offsets, epi_work_tile_info) + if cutlass.const_expr(self.enable_breuse): + # Two M halves: bkeep (rows 0..cta_m/2-1) and breuse (rows cta_m/2..cta_m-1) + mPosition_bk = epi_work_tile_info.tile_m_idx * self.cta_tile_shape_mnk[0] + tidx + mPosition_br = mPosition_bk + (self.cta_tile_shape_mnk[0] // 2) + mProb_bk = real_prob[mPosition_bk, 0, 0] + mProb_br = real_prob[mPosition_br, 0, 0] + mProb = mProb_bk # default; will be overridden per subtile below + else: + mPosition = epi_work_tile_info.tile_m_idx * self.cta_tile_shape_mnk[0] + tidx + mProb = real_prob[mPosition, 0, 0] + + # C1 fix: phase-based acc stage indexing for overlapping_accum + if cutlass.const_expr(self.overlapping_accum): + acc_stage_index = acc_consumer_state.phase + reverse_subtile = cutlass.Boolean(True) if acc_stage_index == 0 else cutlass.Boolean(False) + else: + acc_stage_index = acc_consumer_state.index + + tTR_tAcc = tTR_tAcc_base[(None, None, None, None, None, acc_stage_index)] + tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc)) + + acc_pipeline.consumer_wait(acc_consumer_state) + + subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3]) + + for subtile_idx in cutlass.range(0, subtile_cnt, 1, unroll=1): + real_subtile_idx = subtile_idx + if cutlass.const_expr(self.overlapping_accum): + if reverse_subtile: + real_subtile_idx = self.cta_tile_shape_mnk[1] // self.epi_tile_n_required - 1 - subtile_idx + + # C1 fix: fence + early release for overlapping_accum + if cutlass.const_expr(self.overlapping_accum): + if subtile_idx == self.iter_acc_early_release_in_epilogue: + cute.arch.fence_view_async_tmem_load() + with cute.arch.elect_one(): + acc_pipeline.consumer_release(acc_consumer_state) + acc_consumer_state.advance() + + tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)] + cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc) + + # For breuse, update mProb based on which M half this subtile belongs to. + # With transform, subtiles interleave M groups: even = bkeep, odd = breuse. + if cutlass.const_expr(self.enable_breuse): + if real_subtile_idx % 2 == 0: + mProb = mProb_bk + else: + mProb = mProb_br + + if cutlass.const_expr(self.enable_bias): + # m7 fix: use real_subtile_idx directly (matches contiguous) + sBias_sub = sBias_subtiles[(None, real_subtile_idx)] + cute.copy(bias_s2r_tom, sBias_sub, tTR_rBias) + bias_vec = tTR_rBias.load() + if cutlass.const_expr(self.vectorized_f32): + for i in cutlass.range_constexpr(0, cute.size(tTR_rAcc), 2): + bias_f32_0 = bias_vec[i].to(cutlass.Float32) + bias_f32_1 = bias_vec[i + 1].to(cutlass.Float32) + bias_f32_0, bias_f32_1 = cute.arch.mul_packed_f32x2( + (mProb, mProb), + (bias_f32_0, bias_f32_1), + rnd="rn", + ftz=False, + ) + tTR_rAcc[i], tTR_rAcc[i + 1] = cute.arch.fma_packed_f32x2( + (tTR_rAcc[i], tTR_rAcc[i + 1]), + (cutlass.Float32(alpha_val), cutlass.Float32(alpha_val)), + (bias_f32_0, bias_f32_1), + rnd="rn", + ftz=False, + ) + else: + for i in cutlass.range_constexpr(cute.size(tTR_rAcc)): + tTR_rAcc[i] = tTR_rAcc[i] * cutlass.Float32(alpha_val) + bias_vec[i].to(cutlass.Float32) * mProb + else: + if cutlass.const_expr(self.vectorized_f32): + for i in cutlass.range_constexpr(0, cute.size(tTR_rAcc), 2): + tTR_rAcc[i], tTR_rAcc[i + 1] = cute.arch.mul_packed_f32x2( + (tTR_rAcc[i], tTR_rAcc[i + 1]), + (cutlass.Float32(alpha_val), cutlass.Float32(alpha_val)), + rnd="rn", + ftz=False, + ) + else: + for i in cutlass.range_constexpr(cute.size(tTR_rAcc)): + tTR_rAcc[i] = tTR_rAcc[i] * cutlass.Float32(alpha_val) + + if cutlass.const_expr(self.generate_c): + self.store_c( + tiled_copy_r2s, + tma_atom_c, + warp_idx, + tTR_rAcc, + tTR_rC, + tRS_rC, + tRS_sC, + bSG_gC, + bSG_sC, + c_pipeline, + num_prev_subtiles, + real_subtile_idx, + ) + + # sReLU: apply max(x,0)^2 before prob-scale and quantization + if cutlass.const_expr(self.epilogue_type == EpilogueType.SRELU.value): + acc_vec = tTR_rAcc.load() + acc_relu = cute.where(acc_vec > 0, acc_vec, cute.full_like(acc_vec, 0)) + for i in cutlass.range_constexpr(0, cute.size(tTR_rAcc), 2): + tTR_rAcc[i], tTR_rAcc[i + 1] = cute.arch.mul_packed_f32x2( + (acc_relu[i], acc_relu[i + 1]), + (acc_relu[i], acc_relu[i + 1]), + rnd="rn", + ftz=False, + ) + + acc_vec = tTR_rAcc.load() + if cutlass.const_expr(not self.enable_bias): + tCompute = cute.make_rmem_tensor(acc_vec.shape, self.acc_dtype) + if cutlass.const_expr(self.vectorized_f32): + for i in cutlass.range_constexpr(0, cute.size(tTR_rAcc), 2): + tCompute[i], tCompute[i + 1] = cute.arch.mul_packed_f32x2( + (acc_vec[i], acc_vec[i + 1]), + (mProb, mProb), + rnd="rn", + ftz=False, + ) + else: + for i in cutlass.range_constexpr(cute.size(tTR_rAcc)): + tCompute[i] = acc_vec[i] * mProb + else: + tCompute = tTR_rAcc + + if cutlass.const_expr(self.generate_amax): + thread_tile_amax = amax_reduction_per_thread(tCompute, thread_tile_amax) + + if cutlass.const_expr(self.generate_sfd): + tCompute_col = cute.make_rmem_tensor(tCompute.layout, tCompute.element_type) + tCompute_col.store(tCompute.load()) + if cutlass.const_expr(self.enable_breuse): + # Breuse: subtiles interleave M groups (even=bkeep, odd=breuse). + n_sub = real_subtile_idx // 2 + sfd_tile_idx = n_sub % 4 + is_bkeep = real_subtile_idx % 2 == 0 + if is_bkeep: + self.quant_sfd_row( + sfd_tile_idx, + tiled_copy_r2s, + tCompute, + tCrSFDRow_pvscale, + norm_const, + d_rcp_limits, + tRS_rD, + ) + self.quant_sfd_col( + sfd_tile_idx, + tiled_copy_r2s, + tCompute_col, + tCrSFDCol_pvscale, + norm_const, + d_rcp_limits, + tRS_rD_col, + ) + else: + self.quant_sfd_row( + sfd_tile_idx, + tiled_copy_r2s, + tCompute, + tCrSFDRow_pvscale_br, + norm_const, + d_rcp_limits, + tRS_rD, + ) + self.quant_sfd_col( + sfd_tile_idx, + tiled_copy_r2s, + tCompute_col, + tCrSFDCol_pvscale_br, + norm_const, + d_rcp_limits, + tRS_rD_col, + ) + # SFD M: two 128-row SFD tiles per 256-row CTA tile + global_sfd_m_base = epi_work_tile_info.tile_m_idx * 2 + epi_ext.token_offset // (self.cta_tile_shape_mnk[0] // 2) + m_half = real_subtile_idx % 2 + global_sfd_m = global_sfd_m_base + m_half + sfd_n = epi_work_tile_info.tile_n_idx * 2 + (n_sub >> 2) + sfd_write = n_sub % 4 == 3 + sfd_row_idx_mn = (global_sfd_m, sfd_n) + sfd_col_idx_mn = sfd_row_idx_mn + if cutlass.const_expr(self.discrete_col_sfd): + sfd_col_idx_mn = (epi_work_tile_info.tile_m_idx, sfd_n) + tCgSFDRow = tCgSFDRow_mn[(None, None, None, *sfd_row_idx_mn)] + tCgSFDCol = tCgSFDCol_mn[(None, None, None, *sfd_col_idx_mn)] + if sfd_write: + if is_bkeep: + if sfd_row_idx_mn[1] * 32 * regPerSubtile < cute.size(cute.shape(mSFDRow_mnl.layout, mode=[1])): + tCrSFDRow.store(tCrSFDRow_pvscale.load().to(self.sf_dtype)) + cute.autovec_copy(tCrSFDRow, tCgSFDRow) + if sfd_col_idx_mn[1] * 32 * regPerSubtile < cute.size(cute.shape(mSFDCol_mnl.layout, mode=[1])): + tCrSFDCol.store(tCrSFDCol_pvscale.load().to(self.sf_dtype)) + cute.autovec_copy(tCrSFDCol, tCgSFDCol) + else: + if sfd_row_idx_mn[1] * 32 * regPerSubtile < cute.size(cute.shape(mSFDRow_mnl.layout, mode=[1])): + tCrSFDRow.store(tCrSFDRow_pvscale_br.load().to(self.sf_dtype)) + cute.autovec_copy(tCrSFDRow, tCgSFDRow) + if sfd_col_idx_mn[1] * 32 * regPerSubtile < cute.size(cute.shape(mSFDCol_mnl.layout, mode=[1])): + tCrSFDCol.store(tCrSFDCol_pvscale_br.load().to(self.sf_dtype)) + cute.autovec_copy(tCrSFDCol, tCgSFDCol) + else: + self.quant_sfd_row( + real_subtile_idx % 4, + tiled_copy_r2s, + tCompute, + tCrSFDRow_pvscale, + norm_const, + d_rcp_limits, + tRS_rD, + ) + self.quant_sfd_col( + real_subtile_idx % 4, + tiled_copy_r2s, + tCompute_col, + tCrSFDCol_pvscale, + norm_const, + d_rcp_limits, + tRS_rD_col, + ) + # SFD M tile = cta_tile_m = 128; tile_m_idx is CTA-level per-expert + global_sfd_m = epi_work_tile_info.tile_m_idx + epi_ext.token_offset // self.cta_tile_shape_mnk[0] + if cutlass.const_expr(self.mma_tiler[1] == 256): + sfd_n = epi_work_tile_info.tile_n_idx * 2 + (real_subtile_idx >> 2) + else: + sfd_n = epi_work_tile_info.tile_n_idx + sfd_row_idx_mn = (global_sfd_m, sfd_n) + sfd_col_idx_mn = sfd_row_idx_mn + if cutlass.const_expr(self.discrete_col_sfd): + sfd_col_idx_mn = (epi_work_tile_info.tile_m_idx, sfd_n) + tCgSFDRow = tCgSFDRow_mn[(None, None, None, *sfd_row_idx_mn)] + tCgSFDCol = tCgSFDCol_mn[(None, None, None, *sfd_col_idx_mn)] + if subtile_idx == 3 or subtile_idx == 7: + if sfd_row_idx_mn[1] * 32 * regPerSubtile < cute.size(cute.shape(mSFDRow_mnl.layout, mode=[1])): + tCrSFDRow.store(tCrSFDRow_pvscale.load().to(self.sf_dtype)) + cute.autovec_copy(tCrSFDRow, tCgSFDRow) + if sfd_col_idx_mn[1] * 32 * regPerSubtile < cute.size(cute.shape(mSFDCol_mnl.layout, mode=[1])): + tCrSFDCol.store(tCrSFDCol_pvscale.load().to(self.sf_dtype)) + cute.autovec_copy(tCrSFDCol, tCgSFDCol) + else: + acc_vec = tiled_copy_r2s.retile(tCompute).load() + tRS_rD.store(acc_vec.to(self.d_dtype)) + + d_buffer = num_prev_subtiles % self.num_d_stage + num_prev_subtiles = num_prev_subtiles + 1 + cute.copy(tiled_copy_r2s, tRS_rD, tRS_sD[(None, None, None, d_buffer)]) + if cutlass.const_expr(self.generate_sfd): + cute.copy(tiled_copy_r2s, tRS_rD_col, tRS_sD_col[(None, None, None, d_buffer)]) + cute.arch.fence_proxy("async.shared", space="cta") + self.epilog_sync_barrier.arrive_and_wait() + if warp_idx == self.epilog_warp_id[0]: + cute.copy(tma_atom_d, bSG_sD[(None, d_buffer)], bSG_gD[(None, real_subtile_idx)]) + if cutlass.const_expr(self.generate_sfd): + cute.copy(tma_atom_d_col, bSG_sD_col[(None, d_buffer)], bSG_gD_col[(None, real_subtile_idx)]) + d_pipeline.producer_commit() + d_pipeline.producer_acquire() + self.epilog_sync_barrier.arrive_and_wait() + + if cutlass.const_expr(not self.overlapping_accum): + with cute.arch.elect_one(): + acc_pipeline.consumer_release(acc_consumer_state) + acc_consumer_state.advance() + + if cutlass.const_expr(self.enable_bias): + bias_pipeline.consumer_release(bias_consumer_state) + bias_consumer_state.advance() + + tile_info_pipeline.consumer_wait(tile_info_consumer_state) + for idx in cutlass.range(4, unroll_full=True): + tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)] + is_valid_tile = tile_info[0] >= cutlass.Int32(0) + cute.arch.fence_proxy("async.shared", space="cta") + tile_info_pipeline.consumer_release(tile_info_consumer_state) + tile_info_consumer_state.advance() + + if cutlass.const_expr(self.generate_amax): + gAmax = mAmax_tensor[(expert_idx, None)].iterator.llvm_ptr + self.amax_reduction_per_warp_and_cta(thread_tile_amax, warp_idx, sAmax, gAmax) + + tmem.relinquish_alloc_permit() + self.epilog_sync_barrier.arrive_and_wait() + tmem.free(tmem_ptr) + if cutlass.const_expr(self.generate_c): + c_pipeline.producer_tail() + d_pipeline.producer_tail() + + # ------------------------------------------------------------------ + # Internal: create extension based on weight_mode + # ------------------------------------------------------------------ + + @cute.jit + def _make_extension(self, workspace_ptr): + if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE): + desc_workspace = TensormapWorkspace(workspace_ptr, ["b", "sfb"]) + return DiscreteWeightScaledGemmSchedExtension( + tensormap_ctor=desc_workspace, + sf_vec_size=self.sf_vec_size, + ) + else: + return ContiguousAndConsistentGroupedGemmSchedExtension( + sf_vec_size=self.sf_vec_size, + ) diff --git a/python/cudnn/grouped_gemm/grouped_gemm_swiglu/grouped_gemm_swiglu_quant.py b/python/cudnn/grouped_gemm/grouped_gemm_swiglu/grouped_gemm_swiglu_quant.py index aadfe126b..091217362 100644 --- a/python/cudnn/grouped_gemm/grouped_gemm_swiglu/grouped_gemm_swiglu_quant.py +++ b/python/cudnn/grouped_gemm/grouped_gemm_swiglu/grouped_gemm_swiglu_quant.py @@ -1037,13 +1037,13 @@ def store_c( real_subtile_idx, ) -> None: c_buffer = prev_subtile_idx % self.num_c_stage - tRS_rC.store(tTR_rAcc.load().to(self.c_dtype)) + tRS_rC.store(tiled_copy_r2s.retile(tTR_rAcc).load().to(self.c_dtype)) cute.copy( tiled_copy_r2s, tRS_rC[(None, None, 0)], tRS_sC[(None, None, 1, c_buffer)], ) - tRS_rC.store(tTR_rAcc_gate.load().to(self.c_dtype)) + tRS_rC.store(tiled_copy_r2s.retile(tTR_rAcc_gate).load().to(self.c_dtype)) cute.copy( tiled_copy_r2s, tRS_rC[(None, None, 0)], diff --git a/python/cudnn/grouped_gemm/grouped_gemm_utils.py b/python/cudnn/grouped_gemm/grouped_gemm_utils.py index d0d1d8e9f..844f01224 100644 --- a/python/cudnn/grouped_gemm/grouped_gemm_utils.py +++ b/python/cudnn/grouped_gemm/grouped_gemm_utils.py @@ -56,3 +56,19 @@ def select_grouped_gemm_backend( def backend_cache_key(backend, *components): return (backend.value, *components) + + +def rubin_single_group_offsets_kwarg(is_rubin_kernel, use_single_group_runtime_offsets): + """Return the ``use_single_group_runtime_offsets`` kwarg for a kernel constructor. + + The Rubin (sm107) grouped GEMM kernels predate ``use_single_group_runtime_offsets`` + and do not accept it, so forwarding it unconditionally is a ``TypeError`` even when + it is ``False``. Send it only to kernels that implement it, and reject an explicit + request on Rubin rather than silently ignoring it and running a different schedule + than the caller asked for. + """ + if not is_rubin_kernel: + return {"use_single_group_runtime_offsets": use_single_group_runtime_offsets} + if use_single_group_runtime_offsets: + raise NotImplementedError("The Rubin grouped GEMM kernels do not support use_single_group_runtime_offsets") + return {} diff --git a/skills/cutedsl-kernel-integration/SKILL.md b/skills/cutedsl-kernel-integration/SKILL.md index de11e3506..dade68eca 100644 --- a/skills/cutedsl-kernel-integration/SKILL.md +++ b/skills/cutedsl-kernel-integration/SKILL.md @@ -17,6 +17,7 @@ Use this skill to add or update a CuTeDSL frontend-only API in cuDNN Frontend. T - Execution topology: single kernel, paired forward/backward APIs, multi-kernel orchestrator, helper-kernel setup, distributed/runtime-coordinated execution, or internal scheduler. - Public surface: class API, high-level wrapper, returned tensors, optional outputs, workspace ownership, and import/export namespace. - Internal support: source helper modules, schedulers, metadata utilities, and generated descriptors that must stay private to the package. + - Architecture variant: whether the public API needs transparent dispatch to an alternate CuTeDSL module for a newer GPU (for example Rubin `sm107` vs the default SM100 kernel). Keep the public class and wrapper unchanged when dispatch is internal. 5. Read `references/integration-pattern.md` for the detailed repo conventions before implementing. ## Integration Workflow @@ -29,9 +30,11 @@ Use this skill to add or update a CuTeDSL frontend-only API in cuDNN Frontend. T 6. Add FE OSS documentation and update the relevant overview or operation index links. 7. Add tests under `test/python/fe_api/`, including support validation and numerical/reference coverage when executable. 8. For grouped/discrete/MoE/SDPA kernels, preserve the source helper and scheduler topology; shared helper modules should be internal package files, not public `cudnn` exports. +9. When an existing SM100 kernel needs a Rubin (`sm107`) variant, follow the architecture-dispatch pattern in `references/integration-pattern.md` instead of exposing a new public API. Current examples: `grouped_gemm_quant`, `grouped_gemm_glu`, and `grouped_gemm_dglu`. ## Verification - Run focused formatting or tests for the files changed. - At minimum for skill-only edits, verify this `SKILL.md` has valid frontmatter and all referenced paths exist. - For kernel integrations, run the relevant `pytest test/python/fe_api/test_.py` target when the environment has the required GPU and optional dependencies; otherwise report the skipped verification explicitly. +- For architecture-dispatch work, also run `pytest test/python/fe_api/test_rubin_kernel_dispatch.py`. On Rubin hardware, the existing FE API e2e tests for the affected operation should still pass without API changes. diff --git a/skills/cutedsl-kernel-integration/references/integration-pattern.md b/skills/cutedsl-kernel-integration/references/integration-pattern.md index 608f29379..ddf54d8e6 100644 --- a/skills/cutedsl-kernel-integration/references/integration-pattern.md +++ b/skills/cutedsl-kernel-integration/references/integration-pattern.md @@ -24,6 +24,8 @@ Nested API families are also valid when matching existing structure, for example Choose the closest existing family before creating a new top-level package. +When an existing SM100 kernel needs a Rubin-specific CuTeDSL implementation, keep the public API unchanged and add an internal architecture dispatch layer. See [Architecture-Specific Kernel Variants](#architecture-specific-kernel-variants-rubin--sm107) below. + Use this routing table before choosing the package namespace: | Kernel shape | Preferred family | @@ -71,6 +73,68 @@ Follow the closest template instead of inventing a new lifecycle. Use existing helpers from `api_base.py`, `datatypes.py`, and family utility modules before adding new helpers. +## Architecture-Specific Kernel Variants (Rubin / SM107) + +Use this pattern when the public API stays the same but Rubin (`sm107`, compute capability `(10, 7)`) needs a different CuTeDSL kernel module than the default SM100 implementation. + +Current examples: + +- `python/cudnn/grouped_gemm/grouped_gemm_quant/` +- `python/cudnn/grouped_gemm/grouped_gemm_glu/` +- `python/cudnn/grouped_gemm/grouped_gemm_dglu/` + +### File layout + +Keep the default kernel module unchanged and add a Rubin sibling: + +```text +python/cudnn/grouped_gemm// +|-- api.py +|-- .py +`-- _rubin.py +``` + +Examples: + +- `moe_blockscaled_grouped_gemm_quant.py` + `moe_blockscaled_grouped_gemm_quant_rubin.py` +- `moe_blockscaled_grouped_gemm_glu_rubin.py` +- `moe_blockscaled_grouped_gemm_dglu_rubin.py` + +Rubin kernel modules are internal implementation details. Do not add new public exports or lazy imports for them. + +### Rubin kernel module conventions + +- Import `cutlass.utils.rubin_helpers as sm107_utils` for MMA/tile setup that differs from SM100. +- Keep the kernel class name aligned with the upstream CuTeDSL source when possible so `api.py` can alias it cleanly. +- Preserve MoE helper/scheduler topology from the default kernel unless the Rubin source explicitly changes it. + +### `api.py` dispatch conventions + +Device gating lives in `cudnn.api_base`: `is_sm107_device()` and `self._is_rubin_kernel` (set in `APIBase.__init__`). + +Add a lazy Rubin kernel loader near the top of `api.py`: + +```python +def _get_rubin_kernel(): + from . import as RubinKernelAlias + return RubinKernelAlias +``` + +In `__init__` (after `super().__init__()`): + +```python +self._kernel = _get_rubin_kernel() if self._is_rubin_kernel else DefaultKernelClass +``` + +Then: + +- Replace hard-coded references like `DefaultKernelClass.FIX_PAD_SIZE` with `self._kernel.FIX_PAD_SIZE`. +- Include `get_device_type()` (`"blackwell"` or `"rubin"`) in wrapper cache keys whenever the wrapper can dispatch to architecture-specific kernels. +- Lazy-import the Rubin module inside `_get_rubin_kernel()` so non-Rubin environments do not pay import cost up front. +- Branch in `compile()` / `execute()` only when the Rubin kernel signature or epilogue contract differs from the default kernel. `grouped_gemm_glu` and `grouped_gemm_dglu` only swap the kernel class; `grouped_gemm_quant` additionally adapts compile/execute kwargs for Rubin's optional `c` materialization path and omits `row_scale` on Rubin while keeping it on non-Rubin architectures. + +Do not expose `_is_rubin_kernel`, `_get_rubin_kernel()`, or Rubin module paths in public docs unless the user-visible contract changes. + ## Public Exports Add lazy top-level exports in `python/cudnn/__init__.py` for public APIs intended to be imported as `from cudnn import ...`. diff --git a/test/python/fe_api/test_rubin_kernel_dispatch.py b/test/python/fe_api/test_rubin_kernel_dispatch.py new file mode 100644 index 000000000..302de5372 --- /dev/null +++ b/test/python/fe_api/test_rubin_kernel_dispatch.py @@ -0,0 +1,167 @@ +""" +Unit tests for Rubin (sm107) architecture dispatch in grouped GEMM FE APIs. + +These tests verify device gating and kernel-module selection without requiring +Rubin hardware. Full numerical coverage remains in the existing grouped GEMM +FE API tests, which should pass unchanged on sm107 when cutedsl is available. +""" + +from __future__ import annotations + +import importlib +from pathlib import Path +from unittest import mock + +import pytest + +pytest.importorskip("cutlass") + +RUBIN_DISPATCH_CASES = [ + pytest.param( + "cudnn.grouped_gemm.grouped_gemm_quant.api", + "cudnn.grouped_gemm.grouped_gemm_quant.grouped_gemm_quant", + "BlockScaledMoEGroupedGemmQuantKernel", + "moe_blockscaled_grouped_gemm_quant_rubin.py", + id="grouped_gemm_quant", + ), + pytest.param( + "cudnn.grouped_gemm.grouped_gemm_glu.api", + "cudnn.grouped_gemm.grouped_gemm_glu.moe_blockscaled_grouped_gemm_glu_bias", + "BlockScaledMoEGroupedGemmGluBiasKernel", + "moe_blockscaled_grouped_gemm_glu_rubin.py", + id="grouped_gemm_glu", + ), + pytest.param( + "cudnn.grouped_gemm.grouped_gemm_dglu.api", + "cudnn.grouped_gemm.grouped_gemm_dglu.moe_blockscaled_grouped_gemm_dglu_dbias", + "BlockScaledMoEGroupedGemmDgluDbiasKernel", + "moe_blockscaled_grouped_gemm_dglu_rubin.py", + id="grouped_gemm_dglu", + ), +] + +_REPO_ROOT = Path(__file__).resolve().parents[3] +_GROUPED_GEMM_ROOT = _REPO_ROOT / "python" / "cudnn" / "grouped_gemm" + + +def _import_api_module(module_path: str): + return importlib.import_module(module_path) + + +@pytest.mark.parametrize( + "api_module_path,default_module_path,default_kernel_name,rubin_filename", + RUBIN_DISPATCH_CASES, +) +def test_rubin_kernel_module_is_present( + api_module_path, + default_module_path, + default_kernel_name, + rubin_filename, +): + family = api_module_path.split(".")[-2] + rubin_path = _GROUPED_GEMM_ROOT / family / rubin_filename + assert rubin_path.is_file(), f"Missing Rubin kernel module: {rubin_path}" + + +@pytest.mark.parametrize( + "api_module_path,default_module_path,default_kernel_name,rubin_filename", + RUBIN_DISPATCH_CASES, +) +def test_is_sm107_device_gating( + api_module_path, + default_module_path, + default_kernel_name, + rubin_filename, +): + import cudnn.api_base as api_base + + with mock.patch("torch.cuda.is_available", return_value=False): + assert api_base.is_sm107_device() is False + assert api_base.get_device_type() == "blackwell" + + with ( + mock.patch("torch.cuda.is_available", return_value=True), + mock.patch( + "torch.cuda.get_device_capability", + return_value=(10, 0), + ), + ): + assert api_base.is_sm107_device() is False + assert api_base.get_device_type() == "blackwell" + + with ( + mock.patch("torch.cuda.is_available", return_value=True), + mock.patch( + "torch.cuda.get_device_capability", + return_value=(10, 7), + ), + ): + assert api_base.is_sm107_device() is True + assert api_base.get_device_type() == "rubin" + + +@pytest.mark.parametrize( + "api_module_path,default_module_path,default_kernel_name,rubin_filename", + RUBIN_DISPATCH_CASES, +) +def test_get_rubin_kernel_lazy_import( + api_module_path, + default_module_path, + default_kernel_name, + rubin_filename, +): + api_mod = _import_api_module(api_module_path) + default_mod = importlib.import_module(default_module_path) + default_kernel = getattr(default_mod, default_kernel_name) + + rubin_kernel = api_mod._get_rubin_kernel() + + assert rubin_kernel is not default_kernel + assert rubin_kernel.__name__ != default_kernel.__name__ or rubin_kernel.__module__ != default_kernel.__module__ + assert rubin_kernel.FIX_PAD_SIZE == default_kernel.FIX_PAD_SIZE == 256 + + +@pytest.mark.parametrize( + "api_module_path,default_module_path,default_kernel_name,rubin_filename", + RUBIN_DISPATCH_CASES, +) +def test_kernel_selection_uses_rubin_on_sm107( + api_module_path, + default_module_path, + default_kernel_name, + rubin_filename, +): + api_mod = _import_api_module(api_module_path) + default_mod = importlib.import_module(default_module_path) + default_kernel = getattr(default_mod, default_kernel_name) + rubin_kernel = api_mod._get_rubin_kernel() + import cudnn.api_base as api_base + + with ( + mock.patch("cudnn.api_base.is_sm107_device", return_value=True), + mock.patch.object( + api_mod, + "_get_rubin_kernel", + return_value=rubin_kernel, + ), + ): + selected = api_mod._get_rubin_kernel() if api_base.is_sm107_device() else default_kernel + + assert selected is rubin_kernel + + with mock.patch("cudnn.api_base.is_sm107_device", return_value=False): + selected = api_mod._get_rubin_kernel() if api_base.is_sm107_device() else default_kernel + + assert selected is default_kernel + + +@pytest.mark.L0 +def test_grouped_gemm_quant_has_rubin_compile_branches(): + """Quant adapts compile/execute kwargs on Rubin; keep this contract covered.""" + api_mod = _import_api_module("cudnn.grouped_gemm.grouped_gemm_quant.api") + source = Path(api_mod.__file__).read_text(encoding="utf-8") + + assert "self._is_rubin_kernel" in source + assert 'kernel_kwargs["generate_c"] = False' in source + assert 'compile_kwargs["c"]' in source + assert "if self._is_rubin_kernel:" in source