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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
Extension: moe_sched_extension.py (WgradDense / WgradDiscrete)
"""

from importlib.metadata import PackageNotFoundError, version
from typing import Type, Tuple, Optional

import cuda.bindings.driver as cuda
Expand Down Expand Up @@ -57,6 +58,17 @@
)


def _using_internal_cutlass_dsl() -> bool:
try:
version("nvidia-cutlass-dsl-internal")
except PackageNotFoundError:
return False
return True


_USING_INTERNAL_CUTLASS_DSL = _using_internal_cutlass_dsl()


class BlockScaledMoEGroupedGemmWgradKernel:
"""Block-scaled grouped GEMM kernel for MoE weight gradient (2Dx2D).

Expand Down Expand Up @@ -315,14 +327,22 @@ def __call__(
out_single_expert: Optional[cute.Tensor] = None,
) -> None:

# SM100 still needs the packed-FP4 from_dlpack layout workaround.
# Rubin consumes the native 4-bit layout directly.
if cutlass.const_expr(self.architecture != "sm_107" and mat_a.iterator.dtype.width < 8):
# Public CUTLASS DSL 4.5 needs the packed-FP4 from_dlpack layout
# workaround. Rubin and the internal DSL wheel consume the native
# 4-bit layout directly.
needs_fp4_layout_workaround = (
self.architecture != "sm_107" and not _USING_INTERNAL_CUTLASS_DSL
)
if cutlass.const_expr(
needs_fp4_layout_workaround and mat_a.iterator.dtype.width < 8
):
mat_a = cute.make_tensor(
mat_a.iterator,
cute.recast_layout(mat_a.iterator.dtype.width, 8, mat_a.layout),
)
if cutlass.const_expr(self.architecture != "sm_107" and mat_b.iterator.dtype.width < 8):
if cutlass.const_expr(
needs_fp4_layout_workaround and mat_b.iterator.dtype.width < 8
):
mat_b = cute.make_tensor(
mat_b.iterator,
cute.recast_layout(mat_b.iterator.dtype.width, 8, mat_b.layout),
Expand Down Expand Up @@ -469,6 +489,23 @@ def __call__(
a_smem_layout_helper = None
b_smem_layout_helper = None

# Blackwell's epilogue tile is an MLIR-backed layout and must be
# created outside the isolated helper-kernel region. Rubin uses a
# static tuple and rebuilds the C TMA metadata inside the helper.
if cutlass.const_expr(
self.weight_mode == MoEWeightMode.DISCRETE
and self.architecture != "sm_107"
):
c_tma_op_helper = c_tma_op
epi_smem_layout_helper = cute.select(
self.c_smem_layout_staged, mode=[0, 1]
)
epi_tile_helper = self.epi_tile
else:
c_tma_op_helper = None
epi_smem_layout_helper = None
epi_tile_helper = None

self.helper_kernel(
sfa_gemm,
sfb_gemm,
Expand All @@ -484,6 +521,9 @@ def __call__(
self.cluster_layout_sfb_vmnk.shape,
out if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else None,
c_gemm if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE) else None,
c_tma_op_helper,
epi_smem_layout_helper,
epi_tile_helper,
a_gemm_helper,
b_gemm_helper,
a_op_helper,
Expand Down Expand Up @@ -612,6 +652,9 @@ def helper_kernel(
cluster_layout_sfb_vmnk_shape: cutlass.Constexpr,
c_ptrs=None,
c_single_expert=None,
c_tma_op: cutlass.Constexpr = None,
epi_smem_layout=None,
epi_tile=None,
a_tensor=None,
b_tensor=None,
a_tma_op: cutlass.Constexpr = None,
Expand All @@ -626,11 +669,13 @@ def helper_kernel(
"""
from ..moe_utils import WgradSfTensormapConstructor

# Build C's operation in the helper-kernel IR context. In particular,
# TMA reduce requires its SMEM layout and CTA V-map to be static in
# that context; reusing the host-created operation leaves a dynamic
# CTA V-map for discrete accumulated wgrad.
if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE):
# Rubin requires C's TMA operation and static layout to be built in
# the helper-kernel IR context. Blackwell receives host-built values
# because its MLIR-backed epilogue tile cannot cross region isolation.
if cutlass.const_expr(
self.weight_mode == MoEWeightMode.DISCRETE
and self.architecture == "sm_107"
):
if cutlass.const_expr(self.accumulate_on_output):
c_tma_op = cpasync.CopyReduceBulkTensorTileS2GOp()
else:
Expand All @@ -643,10 +688,6 @@ def helper_kernel(
)
epi_smem_layout = cute.select(c_smem_layout_staged, mode=[0, 1])
epi_tile = self.epi_tile
else:
c_tma_op = None
epi_smem_layout = None
epi_tile = None

ctor = WgradSfTensormapConstructor(
sf_vec_size=self.sf_vec_size,
Expand Down
6 changes: 5 additions & 1 deletion test/python/fe_api/test_rubin_kernel_dispatch.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,7 +51,7 @@
]

_REPO_ROOT = Path(__file__).resolve().parents[3]
_GROUPED_GEMM_ROOT = _REPO_ROOT / "python" / "cudnn" / "grouped_gemm"
_GROUPED_GEMM_ROOT = _REPO_ROOT / "python" / "cudnn" / "gemm" / "cutedsl" / "grouped"


def _import_api_module(module_path: str):
Expand Down Expand Up @@ -214,6 +214,10 @@ def test_grouped_gemm_wgrad_rubin_quantization_validation(

@pytest.mark.L0
def test_grouped_gemm_wgrad_rubin_tmem_plan_rejects_invalid_sf_vector():
pytest.importorskip(
"cutlass.utils.rubin_helpers",
reason="Rubin helpers are unavailable in this CUTLASS DSL wheel",
)
rubin_mod = importlib.import_module(
"cudnn.gemm.cutedsl.grouped.wgrad.moe_blockscaled_grouped_gemm_wgrad_rubin"
)
Expand Down