From 80c76eeca013ccc96a5b280c4ff0a9dbff0f1d3e Mon Sep 17 00:00:00 2001 From: Hwanseo Choi Date: Fri, 28 Aug 2026 15:56:31 -0700 Subject: [PATCH] grouped gemm: accept canonical (sum_m,k)/(l,n,k) layouts and flat SF buffers The contiguous grouped GEMM SwiGLU/dSwiGLU wrappers and API classes now additionally accept natural row-major inputs -- A (sum_m, k), B (l, n, k) C-contiguous, dense C-contiguous SFA/SFB buffers (flat or physical atom shape), prob (sum_m,) fp32/bf16 -- and normalize them internally. Canonical SF buffers compile as flat 1-D pointers: the kernels already rebuild the MMA-tiled SF layouts from the GEMM shapes and read only the base pointer. Canonical calls return natural-shaped outputs; alpha defaults to cached ones. Pre-permuted kernel-facing inputs keep working unchanged. --- ...nch_grouped_gemm_canonical_host_latency.py | 110 +++++++ .../gemm_fusions/grouped_gemm_dswiglu.md | 19 ++ .../gemm_fusions/grouped_gemm_swiglu.md | 18 ++ .../cudnn/gemm/cutedsl/grouped/canonical.py | 103 +++++++ .../cudnn/gemm/cutedsl/grouped/dswiglu/api.py | 273 +++++++++++------ .../dswiglu/grouped_gemm_dswiglu_quant.py | 2 +- .../cudnn/gemm/cutedsl/grouped/swiglu/api.py | 283 ++++++++++++------ .../swiglu/grouped_gemm_swiglu_quant.py | 2 +- .../test_grouped_gemm_canonical_layouts.py | 282 +++++++++++++++++ 9 files changed, 916 insertions(+), 176 deletions(-) create mode 100644 benchmark/gemm/bench_grouped_gemm_canonical_host_latency.py create mode 100644 python/cudnn/gemm/cutedsl/grouped/canonical.py create mode 100644 test/python/fe_api/grouped_gemm/test_grouped_gemm_canonical_layouts.py diff --git a/benchmark/gemm/bench_grouped_gemm_canonical_host_latency.py b/benchmark/gemm/bench_grouped_gemm_canonical_host_latency.py new file mode 100644 index 000000000..be0368880 --- /dev/null +++ b/benchmark/gemm/bench_grouped_gemm_canonical_host_latency.py @@ -0,0 +1,110 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Host-latency comparison for the grouped GEMM SwiGLU wrapper: legacy pre-permuted +call (including the TransformerEngine-style per-call view/permute gymnastics the +legacy contract forces on the caller) vs the canonical natural-layout call. + +DSv3 fc1 shape from MLPerf MoE: (sum_m, n, k) = (24576, 7168, 2048), MXFP8 inputs, +first dim overallocated 1.5x-4x for EP routing slack. + +Usage: python bench_grouped_gemm_canonical_host_latency.py +""" + +import time + +import torch + +from cudnn import grouped_gemm_swiglu_wrapper_sm100 +from cudnn.api_base import ceil_div + +VALID_M, N, K = 24576, 7168, 2048 +EXPERTS = 8 +SF_VEC = 32 +WARMUP, ITERS = 20, 200 + + +def make_buffers(tensor_m): + dev = "cuda" + rest_k = ceil_div(ceil_div(K, SF_VEC), 4) + # Natural (canonical) buffers, as a framework owns them. + a = torch.randint(0, 200, (tensor_m, K), dtype=torch.uint8, device=dev).view(torch.float8_e4m3fn) + b = torch.randint(0, 200, (EXPERTS, N, K), dtype=torch.uint8, device=dev).view(torch.float8_e4m3fn) + sfa = torch.randint(118, 132, (1, ceil_div(tensor_m, 128), rest_k, 32, 4, 4), dtype=torch.uint8, device=dev).view(torch.float8_e8m0fnu) + sfb = torch.randint(118, 132, (EXPERTS, ceil_div(N, 128), rest_k, 32, 4, 4), dtype=torch.uint8, device=dev).view(torch.float8_e8m0fnu) + group = VALID_M // EXPERTS + offsets = torch.arange(group, VALID_M + 1, group, dtype=torch.int32, device=dev) + alpha = torch.ones(EXPERTS, dtype=torch.float32, device=dev) + prob = torch.rand(tensor_m, dtype=torch.float32, device=dev) + norm_const = torch.tensor([0.01], dtype=torch.float32, device=dev) + return dict(a=a, b=b, sfa=sfa, sfb=sfb, offsets=offsets, alpha=alpha, prob=prob, norm_const=norm_const) + + +def call_legacy(buf): + # TE-style per-call layout gymnastics required by the legacy contract + # (see transformer_engine grouped_mlp.py: 6-D SF view+permute, B permute). + m = buf["a"].shape[0] + a3d = buf["a"].view(m, K, 1) + b_nkl = buf["b"].permute(1, 2, 0) + sfa6d = buf["sfa"].view(torch.float8_e8m0fnu).view(1, ceil_div(m, 128), ceil_div(ceil_div(K, SF_VEC), 4), 32, 4, 4).permute(3, 4, 1, 5, 2, 0) + sfb6d = buf["sfb"].view(torch.float8_e8m0fnu).view(EXPERTS, ceil_div(N, 128), ceil_div(ceil_div(K, SF_VEC), 4), 32, 4, 4).permute(3, 4, 1, 5, 2, 0) + prob3d = buf["prob"].view(m, 1, 1) + return grouped_gemm_swiglu_wrapper_sm100( + a_tensor=a3d, + b_tensor=b_nkl, + sfa_tensor=sfa6d, + sfb_tensor=sfb6d, + padded_offsets=buf["offsets"], + alpha_tensor=buf["alpha"], + norm_const_tensor=buf["norm_const"], + prob_tensor=prob3d, + d_dtype=torch.float8_e4m3fn, + sf_vec_size=SF_VEC, + ) + + +def call_canonical(buf): + return grouped_gemm_swiglu_wrapper_sm100( + a_tensor=buf["a"], + b_tensor=buf["b"], + sfa_tensor=buf["sfa"], + sfb_tensor=buf["sfb"], + padded_offsets=buf["offsets"], + alpha_tensor=None, + norm_const_tensor=buf["norm_const"], + prob_tensor=buf["prob"], + d_dtype=torch.float8_e4m3fn, + sf_vec_size=SF_VEC, + ) + + +def bench(fn, buf): + for _ in range(WARMUP): + fn(buf) + torch.cuda.synchronize() + times = [] + for _ in range(ITERS): + torch.cuda.synchronize() + t0 = time.perf_counter() + fn(buf) + times.append(time.perf_counter() - t0) + torch.cuda.synchronize() + times.sort() + n = len(times) + return times[n // 2] * 1e6, times[int(n * 0.9)] * 1e6 + + +def main(): + torch.manual_seed(0) + print(f"grouped_gemm_swiglu host latency, (sum_m, n, k)=({VALID_M}, {N}, {K}), {EXPERTS} experts, MXFP8") + print(f"{'overalloc':>9} | {'tensor_m':>8} | {'legacy p50/p90 (us)':>22} | {'canonical p50/p90 (us)':>22}") + for factor in (1.5, 2.0, 4.0): + tensor_m = ceil_div(int(VALID_M * factor), 256) * 256 + buf = make_buffers(tensor_m) + legacy_p50, legacy_p90 = bench(call_legacy, buf) + canon_p50, canon_p90 = bench(call_canonical, buf) + print(f"{factor:>8}x | {tensor_m:>8} | {legacy_p50:>10.1f} / {legacy_p90:>7.1f} | {canon_p50:>11.1f} / {canon_p90:>7.1f}") + + +if __name__ == "__main__": + main() diff --git a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dswiglu.md b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dswiglu.md index 11dc7cbf2..1d0033fff 100644 --- a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dswiglu.md +++ b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dswiglu.md @@ -363,6 +363,25 @@ Returns a `TupleDict` - a dictionary-like object that also supports tuple unpack - `C`, `D_row`, and `D_col` must be **N-major** (contiguous along N dimension) - All tensors must be **16-byte aligned** along the contiguous dimension +### Canonical layouts (additive) + +Each input is also accepted in its natural row-major form, normalized internally; +the pre-permuted kernel-facing forms above keep working unchanged: + +- `A`: `(valid_m, K)` row-major +- `B`: `(L, N, K)` C-contiguous +- `C`: `(valid_m, 2N)` row-major +- `SFA`/`SFB`: any dense C-contiguous buffer with the MMA-tiled element count, + e.g. flat 1-D or the physical `(L, ceil(mn/128), ceil(ceil(K/sf_vec_size)/4), 32, 4, 4)` + allocation. The kernel rebuilds the MMA-tiled SF layouts from the GEMM shapes and + reads only the base pointer. +- `prob`: `(valid_m,)`, `float32` or `bfloat16` +- `alpha_tensor` may be omitted (defaults to cached ones) + +When `A` is canonical (2-D), the wrapper returns natural-shaped outputs: +`d_row`/`d_col (valid_m, 2N)` row-major, `dprob (valid_m,)`, and +`sfd_row`/`sfd_col` as C-contiguous physical `(1, ceil(mn/128), rest, 32, 4, 4)` buffers. + ### Data Types #### Input/Weight Types (ab_dtype) diff --git a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md index 493244140..dfe252b30 100644 --- a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md +++ b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md @@ -340,6 +340,24 @@ Returns a `TupleDict` - a dictionary-like object that also supports tuple unpack - `C`, `D`, and `D_col` must be **N-major** (contiguous along N dimension) - All tensors must be **16-byte aligned** along the contiguous dimension +### Canonical layouts (additive) + +Each input is also accepted in its natural row-major form, normalized internally; +the pre-permuted kernel-facing forms above keep working unchanged: + +- `A`: `(valid_m, K)` row-major +- `B`: `(L, N, K)` C-contiguous +- `SFA`/`SFB`: any dense C-contiguous buffer with the MMA-tiled element count, + e.g. flat 1-D or the physical `(L, ceil(mn/128), ceil(ceil(K/sf_vec_size)/4), 32, 4, 4)` + allocation — no `.view().permute()` gymnastics required. The kernel rebuilds the + MMA-tiled SF layouts from the GEMM shapes and reads only the base pointer. +- `prob`: `(valid_m,)`, `float32` or `bfloat16` +- `alpha_tensor` may be omitted (defaults to cached ones) + +When `A` is canonical (2-D), the wrapper returns natural-shaped outputs: +`c (valid_m, N)`, `d`/`d_col (valid_m, N/2)` row-major, and `sfd_row`/`sfd_col` as +C-contiguous physical `(1, ceil(mn/128), rest, 32, 4, 4)` buffers. + ### Data Types #### Input/Weight Types (ab_dtype) diff --git a/python/cudnn/gemm/cutedsl/grouped/canonical.py b/python/cudnn/gemm/cutedsl/grouped/canonical.py new file mode 100644 index 000000000..96f6c7074 --- /dev/null +++ b/python/cudnn/gemm/cutedsl/grouped/canonical.py @@ -0,0 +1,103 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Host-side helpers for canonical (natural row-major) grouped GEMM tensor layouts. + +The contiguous grouped GEMM kernels historically required callers to pre-permute +every operand into the kernel-facing form: A/C/D as (m, x, 1) with a trailing unit +L mode, B as (n, k, l) k-major strided views, prob as (m, 1, 1), and the block +scale factors as 6-D MMA-tiled (32, 4, mn//128, 4, rest_k, l) strided views of a +dense buffer. These helpers additively accept the natural buffers instead -- +A (sum_m, k) row-major, B (l, n, k) row-major, prob (sum_m,), and flat/dense +C-contiguous scale-factor buffers -- and normalize them to the kernel-facing form +with zero-copy views. Kernel-facing inputs pass through unchanged, so existing +callers are unaffected. + +The scale-factor kernels rebuild the MMA-tiled SF layouts from the A/B/D shapes on +device and consume only the SF base pointers, so a canonical (C-contiguous) SF +buffer is compiled as a flat 1-D tensor: no MMA-permuted view is ever materialized. +""" + +from __future__ import annotations + +import cutlass.cute as cute + +_cache_of_alpha_ones = {} + + +def default_alpha_ones(l: int, device): + """Cached all-ones per-group scale for callers that don't scale per group.""" + import torch + + key = (l, str(device)) + alpha = _cache_of_alpha_ones.get(key) + if alpha is None: + alpha = torch.ones(l, dtype=torch.float32, device=device) + _cache_of_alpha_ones[key] = alpha + return alpha + + +def unsqueeze_l_dim(tensor): + """Canonical (m, x) row-major -> kernel-facing (m, x, 1); 3-D passes through.""" + if tensor is not None and tensor.ndim == 2: + return tensor.unsqueeze(-1) + return tensor + + +def is_canonical_b(tensor) -> bool: + """True for a canonical (l, n, k) row-major weight tensor. + + The kernel-facing forms keep a stride-1 k mode at dim 1 (k-major (n, k, l)) or a + stride-1 n mode at dim 0 (n-major), so a stride-1 innermost dim 2 identifies the + canonical form. + """ + if tensor is None or tensor.ndim != 3: + return False + stride = tensor.stride() + return stride[2] == 1 and stride[1] != 1 and stride[0] != 1 + + +def to_kernel_b(tensor): + """Canonical (l, n, k) row-major -> kernel-facing (n, k, l); other forms pass through.""" + if is_canonical_b(tensor): + return tensor.permute(1, 2, 0) + return tensor + + +def to_kernel_prob(tensor): + """Canonical (m,) -> kernel-facing (m, 1, 1); other ranks pass through.""" + if tensor is not None and tensor.ndim == 1: + return tensor.view(-1, 1, 1) + return tensor + + +def is_flat_sf(tensor) -> bool: + """True when a scale-factor tensor is a dense C-contiguous buffer (canonical form). + + The legacy MMA-tiled 6-D views are non-contiguous except for degenerate unit-dim + cases where both interpretations address identical memory. + """ + return tensor is not None and tensor.is_contiguous() + + +def to_kernel_sf(tensor, flat: bool): + """Flatten a canonical scale-factor buffer to 1-D; legacy MMA views pass through.""" + if tensor is None or not flat: + return tensor + return tensor if tensor.ndim == 1 else tensor.view(-1) + + +def make_flat_sf_fake(api, desc): + """Fake cute tensor for a flat (1-D, dynamic-length) scale-factor buffer. + + The kernels rebuild the SF layout from the GEMM operand shapes and read only the + base pointer, so the compiled signature needs nothing beyond dtype and a dynamic + length (always a multiple of the 32x4x4 = 512-element SF atom). + """ + if desc is None: + return None + return api._make_fake_cute_tensor( + dtype=desc.dtype, + shape=(cute.sym_int(divisibility=512),), + stride=(1,), + ) diff --git a/python/cudnn/gemm/cutedsl/grouped/dswiglu/api.py b/python/cudnn/gemm/cutedsl/grouped/dswiglu/api.py index 752f717c9..e6e413276 100644 --- a/python/cudnn/gemm/cutedsl/grouped/dswiglu/api.py +++ b/python/cudnn/gemm/cutedsl/grouped/dswiglu/api.py @@ -28,6 +28,15 @@ framework_dtype, get_compute_capability, ) +from ..canonical import ( + is_canonical_b, + is_flat_sf, + make_flat_sf_fake, + to_kernel_b, + to_kernel_prob, + to_kernel_sf, + unsqueeze_l_dim, +) class GroupedGemmDswigluSm100(APIBase): @@ -125,6 +134,26 @@ def __init__( self._warn_experimental_api() self._logger.debug("Entering __init__") + # Canonical (natural row-major) inputs normalize to the kernel-facing views; + # pre-permuted kernel-facing inputs pass through unchanged. + sample_a = unsqueeze_l_dim(sample_a) + sample_b = to_kernel_b(sample_b) + sample_c = unsqueeze_l_dim(sample_c) + sample_d_row = unsqueeze_l_dim(sample_d_row) + sample_d_col = unsqueeze_l_dim(sample_d_col) + sample_prob = to_kernel_prob(sample_prob) + sample_dprob = to_kernel_prob(sample_dprob) + # Canonical (dense C-contiguous) SF buffers compile as flat 1-D pointers; the + # kernel rebuilds the MMA-tiled SF layouts from the A/B/D shapes. + self.sfa_is_flat = is_flat_sf(sample_sfa) + self.sfb_is_flat = is_flat_sf(sample_sfb) + self.sfd_row_is_flat = is_flat_sf(sample_sfd_row) + self.sfd_col_is_flat = is_flat_sf(sample_sfd_col) + sample_sfa = to_kernel_sf(sample_sfa, self.sfa_is_flat) + sample_sfb = to_kernel_sf(sample_sfb, self.sfb_is_flat) + sample_sfd_row = to_kernel_sf(sample_sfd_row, self.sfd_row_is_flat) + sample_sfd_col = to_kernel_sf(sample_sfd_col, self.sfd_col_is_flat) + # Store sample tensor descriptors self.a_desc = self._make_tensor_desc(sample_a, name="sample_a", canonical=True) self.b_desc = self._make_tensor_desc(sample_b, name="sample_b", canonical=True) @@ -203,14 +232,23 @@ def check_support(self) -> bool: self._check_tensor_shape(self.d_row_desc, (tensor_m, n * 2, 1), "D_row") self._check_tensor_shape(self.d_col_desc, (tensor_m, n * 2, 1), "D_col") + def check_sf_shape(desc, is_flat, mma_shape, name): + if is_flat: + numel = 1 + for dim in mma_shape: + numel *= dim + self._check_tensor_shape(desc, (numel,), name) + else: + self._check_tensor_shape(desc, mma_shape, name) + rest_k = ceil_div(ceil_div(k, self.sf_vec_size), 4) - self._check_tensor_shape(self.sfa_desc, (32, 4, ceil_div(tensor_m, 128), 4, rest_k, 1), "SFA") - self._check_tensor_shape(self.sfb_desc, (32, 4, ceil_div(n, 128), 4, rest_k, l), "SFB") + check_sf_shape(self.sfa_desc, self.sfa_is_flat, (32, 4, ceil_div(tensor_m, 128), 4, rest_k, 1), "SFA") + check_sf_shape(self.sfb_desc, self.sfb_is_flat, (32, 4, ceil_div(n, 128), 4, rest_k, l), "SFB") # SFD uses full n dimension since D has n columns (interleaved ab and dswiglu) rest_n2_full = ceil_div(ceil_div(n * 2, self.sf_vec_size), 4) - self._check_tensor_shape(self.sfd_row_desc, (32, 4, ceil_div(tensor_m, 128), 4, rest_n2_full, 1), "SFD_row") + check_sf_shape(self.sfd_row_desc, self.sfd_row_is_flat, (32, 4, ceil_div(tensor_m, 128), 4, rest_n2_full, 1), "SFD_row") rest_m = ceil_div(ceil_div(tensor_m, self.sf_vec_size), 4) - self._check_tensor_shape(self.sfd_col_desc, (32, 4, ceil_div(n * 2, 128), 4, rest_m, 1), "SFD_col") + check_sf_shape(self.sfd_col_desc, self.sfd_col_is_flat, (32, 4, ceil_div(n * 2, 128), 4, rest_m, 1), "SFD_col") self._check_tensor_shape(self.padded_offsets_desc, (l,), "padded_offsets") self._check_tensor_shape(self.alpha_desc, (l,), "alpha") @@ -262,7 +300,7 @@ def check_support(self) -> bool: raise ValueError(f"ab_dtype {self.ab_dtype} and sf_vec_size {self.sf_vec_size} combination is not supported") self._check_dtype(self.acc_dtype, dtype=cutlass.Float32, name="Accumulator", extra_error_msg="Accumulator must be float32") - self._check_dtype(self.prob_desc, dtype=cutlass.Float32, name="Prob", extra_error_msg="Prob must be float32") + self._check_dtype(self.prob_desc, dtype=[cutlass.Float32, cutlass.BFloat16], name="Prob", extra_error_msg="Prob must be float32 or bfloat16") self._check_dtype(self.dprob_desc, dtype=cutlass.Float32, name="Dprob", extra_error_msg="Dprob must be float32") self.c_dtype = self._check_dtype( self.c_desc, @@ -426,31 +464,43 @@ def compile(self) -> None: tensor_m_128 = cute.sym_int() stride_tensor_m_128 = cute.sym_int(divisibility=32 * 4 * 4) - sfa_cute_fake = self._make_fake_cute_tensor( - dtype=self.sfa_desc.dtype, - shape=(32, 4, tensor_m_128, 4, self.sfa_desc.shape[4], 1), - stride=(16, 4, self.sfa_desc.stride[2], 1, 512, stride_tensor_m_128), - ) - sfb_cute_fake = self._make_fake_cute_tensor_from_desc(self.sfb_desc, assumed_align=16) + if self.sfa_is_flat: + sfa_cute_fake = make_flat_sf_fake(self, self.sfa_desc) + else: + sfa_cute_fake = self._make_fake_cute_tensor( + dtype=self.sfa_desc.dtype, + shape=(32, 4, tensor_m_128, 4, self.sfa_desc.shape[4], 1), + stride=(16, 4, self.sfa_desc.stride[2], 1, 512, stride_tensor_m_128), + ) + if self.sfb_is_flat: + sfb_cute_fake = make_flat_sf_fake(self, self.sfb_desc) + else: + sfb_cute_fake = self._make_fake_cute_tensor_from_desc(self.sfb_desc, assumed_align=16) sfd_row_fake = None sfd_col_fake = None if self.sfd_row_desc is not None: - stride_sfd_m = cute.sym_int(divisibility=32 * 4 * 4) - sfd_row_fake = self._make_fake_cute_tensor( - dtype=self.sfd_row_desc.dtype, - shape=(32, 4, tensor_m_128, 4, self.sfd_row_desc.shape[4], 1), - stride=(16, 4, self.sfd_row_desc.stride[2], 1, 512, stride_sfd_m), - ) + if self.sfd_row_is_flat: + sfd_row_fake = make_flat_sf_fake(self, self.sfd_row_desc) + else: + stride_sfd_m = cute.sym_int(divisibility=32 * 4 * 4) + sfd_row_fake = self._make_fake_cute_tensor( + dtype=self.sfd_row_desc.dtype, + shape=(32, 4, tensor_m_128, 4, self.sfd_row_desc.shape[4], 1), + stride=(16, 4, self.sfd_row_desc.stride[2], 1, 512, stride_sfd_m), + ) if self.sfd_col_desc is not None: - rest_m = cute.sym_int(divisibility=1) - stride_sfd_n = cute.sym_int(divisibility=32 * 4 * 4) - stride_rest_m = cute.sym_int(divisibility=32 * 4 * 4) - sfd_col_fake = self._make_fake_cute_tensor( - dtype=self.sfd_col_desc.dtype, - shape=(32, 4, self.sfd_col_desc.shape[2], 4, rest_m, 1), - stride=(16, 4, stride_rest_m, 1, 512, stride_sfd_n), - ) + if self.sfd_col_is_flat: + sfd_col_fake = make_flat_sf_fake(self, self.sfd_col_desc) + else: + rest_m = cute.sym_int(divisibility=1) + stride_sfd_n = cute.sym_int(divisibility=32 * 4 * 4) + stride_rest_m = cute.sym_int(divisibility=32 * 4 * 4) + sfd_col_fake = self._make_fake_cute_tensor( + dtype=self.sfd_col_desc.dtype, + shape=(32, 4, self.sfd_col_desc.shape[2], 4, rest_m, 1), + stride=(16, 4, stride_rest_m, 1, 512, stride_sfd_n), + ) else: valid_m = cute.sym_int(divisibility=256) n_2 = cute.sym_int() @@ -503,45 +553,59 @@ def compile(self) -> None: stride_order=self.dprob_desc.stride_order, ) - tensor_m_128 = cute.sym_int() - rest_k = cute.sym_int() - stride_rest_k = cute.sym_int(divisibility=32 * 4 * 4) - stride_tensor_m_128 = cute.sym_int(divisibility=32 * 4 * 4) - sfa_cute_fake = self._make_fake_cute_tensor( - dtype=self.sfa_desc.dtype, - shape=(32, 4, tensor_m_128, 4, rest_k, 1), - stride=(16, 4, stride_rest_k, 1, 512, stride_tensor_m_128), - ) - tensor_n_128 = cute.sym_int() - stride_sfb_rest_k = cute.sym_int(divisibility=32 * 4 * 4) - stride_sfb_tensor_n_128 = cute.sym_int(divisibility=32 * 4 * 4) - sfb_cute_fake = self._make_fake_cute_tensor( - dtype=self.sfb_desc.dtype, - shape=(32, 4, tensor_n_128, 4, rest_k, l), - stride=(16, 4, stride_sfb_tensor_n_128, 1, 512, stride_sfb_rest_k), - ) + if self.sfa_is_flat: + sfa_cute_fake = make_flat_sf_fake(self, self.sfa_desc) + else: + tensor_m_128 = cute.sym_int() + rest_k = cute.sym_int() + stride_rest_k = cute.sym_int(divisibility=32 * 4 * 4) + stride_tensor_m_128 = cute.sym_int(divisibility=32 * 4 * 4) + sfa_cute_fake = self._make_fake_cute_tensor( + dtype=self.sfa_desc.dtype, + shape=(32, 4, tensor_m_128, 4, rest_k, 1), + stride=(16, 4, stride_rest_k, 1, 512, stride_tensor_m_128), + ) + if self.sfb_is_flat: + sfb_cute_fake = make_flat_sf_fake(self, self.sfb_desc) + else: + tensor_n_128 = cute.sym_int() + sfb_rest_k = cute.sym_int() + stride_sfb_rest_k = cute.sym_int(divisibility=32 * 4 * 4) + stride_sfb_tensor_n_128 = cute.sym_int(divisibility=32 * 4 * 4) + sfb_cute_fake = self._make_fake_cute_tensor( + dtype=self.sfb_desc.dtype, + shape=(32, 4, tensor_n_128, 4, sfb_rest_k, l), + stride=(16, 4, stride_sfb_tensor_n_128, 1, 512, stride_sfb_rest_k), + ) sfd_row_fake = None sfd_col_fake = None if self.sfd_row_desc is not None: - rest_n2 = cute.sym_int() - stride_sfd_rest_n2 = cute.sym_int(divisibility=32 * 4 * 4) - stride_sfd_rest_tensor_m_128 = cute.sym_int(divisibility=32 * 4 * 4) - sfd_row_fake = self._make_fake_cute_tensor( - dtype=self.sfd_row_desc.dtype, - shape=(32, 4, tensor_m_128, 4, rest_n2, 1), - stride=(16, 4, stride_sfd_rest_n2, 1, 512, stride_sfd_rest_tensor_m_128), - ) + if self.sfd_row_is_flat: + sfd_row_fake = make_flat_sf_fake(self, self.sfd_row_desc) + else: + sfd_tensor_m_128 = cute.sym_int() + rest_n2 = cute.sym_int() + stride_sfd_rest_n2 = cute.sym_int(divisibility=32 * 4 * 4) + stride_sfd_rest_tensor_m_128 = cute.sym_int(divisibility=32 * 4 * 4) + sfd_row_fake = self._make_fake_cute_tensor( + dtype=self.sfd_row_desc.dtype, + shape=(32, 4, sfd_tensor_m_128, 4, rest_n2, 1), + stride=(16, 4, stride_sfd_rest_n2, 1, 512, stride_sfd_rest_tensor_m_128), + ) if self.sfd_col_desc is not None: - tensor_n2_128 = cute.sym_int() - rest_m = cute.sym_int() - stride_sfd_rest_m = cute.sym_int(divisibility=32 * 4 * 4) - stride_sfd_n2 = cute.sym_int(divisibility=32 * 4 * 4) - sfd_col_fake = self._make_fake_cute_tensor( - dtype=self.sfd_col_desc.dtype, - shape=(32, 4, tensor_n2_128, 4, rest_m, 1), - stride=(16, 4, stride_sfd_rest_m, 1, 512, stride_sfd_n2), - ) + if self.sfd_col_is_flat: + sfd_col_fake = make_flat_sf_fake(self, self.sfd_col_desc) + else: + tensor_n2_128 = cute.sym_int() + rest_m = cute.sym_int() + stride_sfd_rest_m = cute.sym_int(divisibility=32 * 4 * 4) + stride_sfd_n2 = cute.sym_int(divisibility=32 * 4 * 4) + sfd_col_fake = self._make_fake_cute_tensor( + dtype=self.sfd_col_desc.dtype, + shape=(32, 4, tensor_n2_128, 4, rest_m, 1), + stride=(16, 4, stride_sfd_rest_m, 1, 512, stride_sfd_n2), + ) _compiled_kernel = cute.compile( gemm_dswiglu, @@ -661,22 +725,22 @@ def execute( raise RuntimeError("Kernel not compiled; call compile() first") self._logger.debug("Executing grouped_gemm_dswiglu kernel") self._compiled_kernel( - a_tensor=a_tensor, - b_tensor=b_tensor, - c_tensor=c_tensor, - d_row_tensor=d_row_tensor, - d_col_tensor=d_col_tensor, - sfa_tensor=sfa_tensor, - sfb_tensor=sfb_tensor, - sfd_row_tensor=sfd_row_tensor, - sfd_col_tensor=sfd_col_tensor, + a_tensor=unsqueeze_l_dim(a_tensor), + b_tensor=to_kernel_b(b_tensor), + c_tensor=unsqueeze_l_dim(c_tensor), + d_row_tensor=unsqueeze_l_dim(d_row_tensor), + d_col_tensor=unsqueeze_l_dim(d_col_tensor), + sfa_tensor=to_kernel_sf(sfa_tensor, self.sfa_is_flat), + sfb_tensor=to_kernel_sf(sfb_tensor, self.sfb_is_flat), + sfd_row_tensor=to_kernel_sf(sfd_row_tensor, self.sfd_row_is_flat), + sfd_col_tensor=to_kernel_sf(sfd_col_tensor, self.sfd_col_is_flat), amax_tensor=amax_tensor, norm_const_tensor=norm_const_tensor, padded_offsets=padded_offsets, alpha_tensor=alpha_tensor, beta_tensor=beta_tensor, - prob_tensor=prob_tensor, - dprob_tensor=dprob_tensor, + prob_tensor=to_kernel_prob(prob_tensor), + dprob_tensor=to_kernel_prob(dprob_tensor), stream=current_stream, ) @@ -697,9 +761,9 @@ def grouped_gemm_dswiglu_wrapper_sm100( sfa_tensor: torch.Tensor, sfb_tensor: torch.Tensor, padded_offsets: torch.Tensor, - alpha_tensor: torch.Tensor, - beta_tensor: Optional[torch.Tensor], - prob_tensor: torch.Tensor, + alpha_tensor: Optional[torch.Tensor] = None, + beta_tensor: Optional[torch.Tensor] = None, + prob_tensor: Optional[torch.Tensor] = None, norm_const_tensor: Optional[torch.Tensor] = None, acc_dtype: Optional[torch.dtype] = None, d_dtype: Optional[torch.dtype] = None, @@ -720,16 +784,26 @@ def grouped_gemm_dswiglu_wrapper_sm100( This function creates the API, compiles, and executes in one call. Compiled kernels are cached for reuse when called with the same configuration. + Canonical layouts (additive): each input is also accepted in its natural + row-major form and normalized internally -- A as (valid_m, k), B as (l, n, k) + C-contiguous, C as (valid_m, 2n), SFA/SFB as dense C-contiguous buffers of any + shape with the MMA-tiled element count (e.g. flat 1-D, or physical + (l, mn//128, ceil(ceil(k/sf_vec_size)/4), 32, 4, 4)), and prob as (valid_m,) + float32 or bfloat16. When A is canonical (2-D), outputs come back natural-shaped: + d_row/d_col (valid_m, 2n) row-major, dprob (valid_m,), and sfd_row/sfd_col as + C-contiguous physical (1, mn//128, rest, 32, 4, 4) buffers. The pre-permuted + kernel-facing forms below keep working unchanged. + Args: - a_tensor: Input A tensor (valid_m, k, 1) - b_tensor: Weight B tensor (n, k, l) - c_tensor: Intermediate C tensor from forward pass (valid_m, 2n, 1) - sfa_tensor: Scale factor A - sfb_tensor: Scale factor B + a_tensor: Input A tensor (valid_m, k, 1), or canonical (valid_m, k) row-major + b_tensor: Weight B tensor (n, k, l) k-major, or canonical (l, n, k) row-major + c_tensor: Intermediate C tensor from forward pass (valid_m, 2n, 1) or (valid_m, 2n) + sfa_tensor: Scale factor A (MMA-tiled view, or canonical dense buffer) + sfb_tensor: Scale factor B (MMA-tiled view, or canonical dense buffer) padded_offsets: End offset per expert after padding (l,) - alpha_tensor: Per-group alpha scaling + alpha_tensor: Per-group alpha scaling; None defaults to ones (cached) beta_tensor: Per-group beta scaling - prob_tensor: Per-row probability tensor + prob_tensor: Per-row probability tensor (required) norm_const_tensor: Optional normalization constant acc_dtype: Accumulator data type d_dtype: Output D tensor data type @@ -767,7 +841,19 @@ def grouped_gemm_dswiglu_wrapper_sm100( acc_dtype = _convert_to_cutlass_data_type(acc_dtype) if acc_dtype is not None else cutlass.Float32 d_dtype = _convert_to_cutlass_data_type(d_dtype) if d_dtype is not None else cutlass.BFloat16 valid_m = a_tensor.shape[0] - n, _, l = b_tensor.shape + if is_canonical_b(b_tensor): + l, n, _ = b_tensor.shape + else: + n, _, l = b_tensor.shape + + # Canonical (sum_m, k) A selects natural-shaped outputs: (m, x) row-major D, + # (m,) dprob, and dense C-contiguous SFD buffers. + canonical_outputs = a_tensor.ndim == 2 + + if prob_tensor is None: + raise ValueError("prob_tensor is required for grouped_gemm_dswiglu_wrapper_sm100") + if alpha_tensor is None: + alpha_tensor = default_alpha_ones(l, a_tensor.device) if cd_major != "n": raise ValueError(f"cd_major must be 'n', got {cd_major}") @@ -806,15 +892,27 @@ def stride_order(tensor: torch.Tensor) -> Tuple[int, ...]: m_aligned, discrete_col_sfd, epilogue_op, + # Canonical-vs-kernel-facing input forms compile different signatures. + prob_tensor.dtype, + prob_tensor.ndim, + is_flat_sf(sfa_tensor), + is_flat_sf(sfb_tensor), + canonical_outputs, ) # Allocate M-dependent output tensors fresh every call (M varies across MoE steps). # Only M-independent tensors (amax, beta) are cached to avoid repeated allocation. _logger.debug("grouped_gemm_dswiglu_wrapper_sm100: Allocating M-dependent output tensors") d_torch_dtype = framework_dtype(d_dtype, "torch") - d_row_tensor = torch.empty_strided((valid_m, n * 2, 1), (n * 2, 1, valid_m * n * 2), dtype=d_torch_dtype, device=a_tensor.device) - d_col_tensor = torch.empty_strided((valid_m, n * 2, 1), (n * 2, 1, valid_m * n * 2), dtype=d_torch_dtype, device=a_tensor.device) - dprob_tensor = dprob_tensor_buf.zero_() if dprob_tensor_buf is not None else torch.zeros((valid_m, 1, 1), dtype=torch.float32, device=a_tensor.device) + if canonical_outputs: + d_row_tensor = torch.empty((valid_m, n * 2), dtype=d_torch_dtype, device=a_tensor.device) + d_col_tensor = torch.empty((valid_m, n * 2), dtype=d_torch_dtype, device=a_tensor.device) + dprob_default_shape = (valid_m,) + else: + d_row_tensor = torch.empty_strided((valid_m, n * 2, 1), (n * 2, 1, valid_m * n * 2), dtype=d_torch_dtype, device=a_tensor.device) + d_col_tensor = torch.empty_strided((valid_m, n * 2, 1), (n * 2, 1, valid_m * n * 2), dtype=d_torch_dtype, device=a_tensor.device) + dprob_default_shape = (valid_m, 1, 1) + dprob_tensor = dprob_tensor_buf.zero_() if dprob_tensor_buf is not None else torch.zeros(dprob_default_shape, dtype=torch.float32, device=a_tensor.device) if valid_m == 0: amax_tensor = None @@ -841,10 +939,13 @@ def stride_order(tensor: torch.Tensor) -> Tuple[int, ...]: mma_permute_order = (3, 4, 1, 5, 2, 0) sf_k_row = ceil_div(n * 2, sf_vec_size) mma_shape_row = (1, ceil_div(valid_m, 128), ceil_div(sf_k_row, 4), 32, 4, 4) - sfd_row_tensor = torch.empty(mma_shape_row, dtype=sf_dtype, device=a_tensor.device).permute(mma_permute_order) + sfd_row_tensor = torch.empty(mma_shape_row, dtype=sf_dtype, device=a_tensor.device) sf_k_col = ceil_div(valid_m, sf_vec_size) mma_shape_col = (1, ceil_div(n * 2, 128), ceil_div(sf_k_col, 4), 32, 4, 4) - sfd_col_tensor = torch.empty(mma_shape_col, dtype=sf_dtype, device=a_tensor.device).permute(mma_permute_order) + sfd_col_tensor = torch.empty(mma_shape_col, dtype=sf_dtype, device=a_tensor.device) + if not canonical_outputs: + sfd_row_tensor = sfd_row_tensor.permute(mma_permute_order) + sfd_col_tensor = sfd_col_tensor.permute(mma_permute_order) if cache_key in _cache_of_GroupedGemmDswigluSm100Objects: _logger.debug("group_gemm_dswiglu_wrapper_sm100: Using previously cached GroupedGemmDswigluSm100 object") diff --git a/python/cudnn/gemm/cutedsl/grouped/dswiglu/grouped_gemm_dswiglu_quant.py b/python/cudnn/gemm/cutedsl/grouped/dswiglu/grouped_gemm_dswiglu_quant.py index 7190f3496..134b404a8 100644 --- a/python/cudnn/gemm/cutedsl/grouped/dswiglu/grouped_gemm_dswiglu_quant.py +++ b/python/cudnn/gemm/cutedsl/grouped/dswiglu/grouped_gemm_dswiglu_quant.py @@ -2580,7 +2580,7 @@ def kernel( # Note, it always assumes T2R_M/EPI_M is 1, otherwise it will break the result. # mPosition = tile_info[0] * self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape) + tidx - mProb = prob[mPosition, 0, 0] + mProb = prob[mPosition, 0, 0].to(cutlass.Float32) if cutlass.const_expr(self.generate_dprob): dProbVal = cutlass.Float32(0.0) diff --git a/python/cudnn/gemm/cutedsl/grouped/swiglu/api.py b/python/cudnn/gemm/cutedsl/grouped/swiglu/api.py index ccb11c4ea..3f010d003 100644 --- a/python/cudnn/gemm/cutedsl/grouped/swiglu/api.py +++ b/python/cudnn/gemm/cutedsl/grouped/swiglu/api.py @@ -29,6 +29,16 @@ framework_dtype, get_compute_capability, ) +from ..canonical import ( + default_alpha_ones, + is_canonical_b, + is_flat_sf, + make_flat_sf_fake, + to_kernel_b, + to_kernel_prob, + to_kernel_sf, + unsqueeze_l_dim, +) _JAX_SF_LAYOUT_ERROR = ( "the block scale-factor tensors (sfa/sfb and the sfd outputs) are MMA-tiled " @@ -122,6 +132,25 @@ def __init__( self._warn_experimental_api() self._logger.debug("Entering __init__") + # Canonical (natural row-major) inputs normalize to the kernel-facing views; + # pre-permuted kernel-facing inputs pass through unchanged. + sample_a = unsqueeze_l_dim(sample_a) + sample_b = to_kernel_b(sample_b) + sample_c = unsqueeze_l_dim(sample_c) + sample_d = unsqueeze_l_dim(sample_d) + sample_d_col = unsqueeze_l_dim(sample_d_col) + sample_prob = to_kernel_prob(sample_prob) + # Canonical (dense C-contiguous) SF buffers compile as flat 1-D pointers; the + # kernel rebuilds the MMA-tiled SF layouts from the A/B/D shapes. + self.sfa_is_flat = is_flat_sf(sample_sfa) + self.sfb_is_flat = is_flat_sf(sample_sfb) + self.sfd_row_is_flat = is_flat_sf(sample_sfd_row) + self.sfd_col_is_flat = is_flat_sf(sample_sfd_col) + sample_sfa = to_kernel_sf(sample_sfa, self.sfa_is_flat) + sample_sfb = to_kernel_sf(sample_sfb, self.sfb_is_flat) + sample_sfd_row = to_kernel_sf(sample_sfd_row, self.sfd_row_is_flat) + sample_sfd_col = to_kernel_sf(sample_sfd_col, self.sfd_col_is_flat) + # Store sample tensor descriptors self.a_desc = self._make_tensor_desc(sample_a, name="sample_a", canonical=True) self.b_desc = self._make_tensor_desc(sample_b, name="sample_b", canonical=True) @@ -198,17 +227,27 @@ def check_support(self) -> bool: self._check_tensor_shape(self.d_col_desc, (tensor_m, n // 2, 1), "D_col") + def check_sf_shape(desc, is_flat, mma_shape, name): + if is_flat: + numel = 1 + for dim in mma_shape: + numel *= dim + self._check_tensor_shape(desc, [(numel,)], name) + else: + self._check_tensor_shape(desc, mma_shape, name) + rest_k = ceil_div(ceil_div(k, self.sf_vec_size), 4) - self._check_tensor_shape(self.sfa_desc, (32, 4, ceil_div(tensor_m, 128), 4, rest_k, 1), "SFA") - self._check_tensor_shape(self.sfb_desc, (32, 4, ceil_div(n, 128), 4, rest_k, l), "SFB") + check_sf_shape(self.sfa_desc, self.sfa_is_flat, (32, 4, ceil_div(tensor_m, 128), 4, rest_k, 1), "SFA") + check_sf_shape(self.sfb_desc, self.sfb_is_flat, (32, 4, ceil_div(n, 128), 4, rest_k, l), "SFB") rest_n2 = ceil_div(ceil_div(n // 2, self.sf_vec_size), 4) - self._check_tensor_shape( + check_sf_shape( self.sfd_row_desc, + self.sfd_row_is_flat, (32, 4, ceil_div(tensor_m, 128), 4, rest_n2, 1), "SFD_row", ) rest_m = ceil_div(ceil_div(tensor_m, self.sf_vec_size), 4) - self._check_tensor_shape(self.sfd_col_desc, (32, 4, ceil_div(n // 2, 128), 4, rest_m, 1), "SFD_col") + check_sf_shape(self.sfd_col_desc, self.sfd_col_is_flat, (32, 4, ceil_div(n // 2, 128), 4, rest_m, 1), "SFD_col") self._check_tensor_shape(self.alpha_desc, (l,), "alpha") self._check_tensor_shape(self.prob_desc, (tensor_m, 1, 1), "prob") @@ -303,6 +342,12 @@ def check_support(self) -> bool: name="Accumulator", extra_error_msg="Accumulator must be float32", ) + self._check_dtype( + self.prob_desc, + dtype=[cutlass.Float32, cutlass.BFloat16], + name="Prob", + extra_error_msg="Prob must be float32 or bfloat16", + ) self.c_dtype = self._check_dtype( self.c_desc, dtype=[ @@ -505,13 +550,19 @@ def compile(self) -> None: tensor_m_128 = cute.sym_int() stride_tensor_m_128 = cute.sym_int(divisibility=32 * 4 * 4) - sfa_cute_fake = self._make_fake_cute_tensor( - dtype=self.sfa_desc.dtype, - shape=(32, 4, tensor_m_128, 4, self.sfa_desc.shape[4], 1), - stride=(16, 4, self.sfa_desc.stride[2], 1, 512, stride_tensor_m_128), - ) + if self.sfa_is_flat: + sfa_cute_fake = make_flat_sf_fake(self, self.sfa_desc) + else: + sfa_cute_fake = self._make_fake_cute_tensor( + dtype=self.sfa_desc.dtype, + shape=(32, 4, tensor_m_128, 4, self.sfa_desc.shape[4], 1), + stride=(16, 4, self.sfa_desc.stride[2], 1, 512, stride_tensor_m_128), + ) - sfb_cute_fake = self._make_fake_cute_tensor_from_desc(self.sfb_desc, assumed_align=16) + if self.sfb_is_flat: + sfb_cute_fake = make_flat_sf_fake(self, self.sfb_desc) + else: + sfb_cute_fake = self._make_fake_cute_tensor_from_desc(self.sfb_desc, assumed_align=16) prob_cute_fake = None if self.prob_desc is not None: @@ -524,21 +575,27 @@ def compile(self) -> None: sfd_row_fake = None sfd_col_fake = None if self.sfd_row_desc is not None: - stride_sfd_m = cute.sym_int(divisibility=32 * 4 * 4) - sfd_row_fake = self._make_fake_cute_tensor( - dtype=self.sfd_row_desc.dtype, - shape=(32, 4, tensor_m_128, 4, self.sfd_row_desc.shape[4], 1), - stride=(16, 4, self.sfd_row_desc.stride[2], 1, 512, stride_sfd_m), - ) + if self.sfd_row_is_flat: + sfd_row_fake = make_flat_sf_fake(self, self.sfd_row_desc) + else: + stride_sfd_m = cute.sym_int(divisibility=32 * 4 * 4) + sfd_row_fake = self._make_fake_cute_tensor( + dtype=self.sfd_row_desc.dtype, + shape=(32, 4, tensor_m_128, 4, self.sfd_row_desc.shape[4], 1), + stride=(16, 4, self.sfd_row_desc.stride[2], 1, 512, stride_sfd_m), + ) if self.sfd_col_desc is not None: - rest_m = cute.sym_int(divisibility=1) - stride_sfd_n = cute.sym_int(divisibility=32 * 4 * 4) - stride_rest_m = cute.sym_int(divisibility=32 * 4 * 4) - sfd_col_fake = self._make_fake_cute_tensor( - dtype=self.sfd_col_desc.dtype, - shape=(32, 4, self.sfd_col_desc.shape[2], 4, rest_m, 1), - stride=(16, 4, stride_rest_m, 1, 512, stride_sfd_n), - ) + if self.sfd_col_is_flat: + sfd_col_fake = make_flat_sf_fake(self, self.sfd_col_desc) + else: + rest_m = cute.sym_int(divisibility=1) + stride_sfd_n = cute.sym_int(divisibility=32 * 4 * 4) + stride_rest_m = cute.sym_int(divisibility=32 * 4 * 4) + sfd_col_fake = self._make_fake_cute_tensor( + dtype=self.sfd_col_desc.dtype, + shape=(32, 4, self.sfd_col_desc.shape[2], 4, rest_m, 1), + stride=(16, 4, stride_rest_m, 1, 512, stride_sfd_n), + ) else: valid_m = cute.sym_int(divisibility=256) n = cute.sym_int() @@ -582,23 +639,30 @@ def compile(self) -> None: divisibility=8 if self._is_f16(self.d_col_desc.dtype) else 16, ) - tensor_m_128 = cute.sym_int() - rest_k = cute.sym_int() - stride_rest_k = cute.sym_int(divisibility=32 * 4 * 4) - stride_tensor_m_128 = cute.sym_int(divisibility=32 * 4 * 4) - sfa_cute_fake = self._make_fake_cute_tensor( - dtype=self.sfa_desc.dtype, - shape=(32, 4, tensor_m_128, 4, rest_k, 1), - stride=(16, 4, stride_rest_k, 1, 512, stride_tensor_m_128), - ) - tensor_n_128 = cute.sym_int() - stride_sfb_rest_k = cute.sym_int(divisibility=32 * 4 * 4) - stride_sfb_tensor_n_128 = cute.sym_int(divisibility=32 * 4 * 4) - sfb_cute_fake = self._make_fake_cute_tensor( - dtype=self.sfb_desc.dtype, - shape=(32, 4, tensor_n_128, 4, rest_k, l), - stride=(16, 4, stride_sfb_tensor_n_128, 1, 512, stride_sfb_rest_k), - ) + if self.sfa_is_flat: + sfa_cute_fake = make_flat_sf_fake(self, self.sfa_desc) + else: + tensor_m_128 = cute.sym_int() + rest_k = cute.sym_int() + stride_rest_k = cute.sym_int(divisibility=32 * 4 * 4) + stride_tensor_m_128 = cute.sym_int(divisibility=32 * 4 * 4) + sfa_cute_fake = self._make_fake_cute_tensor( + dtype=self.sfa_desc.dtype, + shape=(32, 4, tensor_m_128, 4, rest_k, 1), + stride=(16, 4, stride_rest_k, 1, 512, stride_tensor_m_128), + ) + if self.sfb_is_flat: + sfb_cute_fake = make_flat_sf_fake(self, self.sfb_desc) + else: + tensor_n_128 = cute.sym_int() + sfb_rest_k = cute.sym_int() + stride_sfb_rest_k = cute.sym_int(divisibility=32 * 4 * 4) + stride_sfb_tensor_n_128 = cute.sym_int(divisibility=32 * 4 * 4) + sfb_cute_fake = self._make_fake_cute_tensor( + dtype=self.sfb_desc.dtype, + shape=(32, 4, tensor_n_128, 4, sfb_rest_k, l), + stride=(16, 4, stride_sfb_tensor_n_128, 1, 512, stride_sfb_rest_k), + ) prob_cute_fake = None if self.prob_desc is not None: @@ -611,31 +675,38 @@ def compile(self) -> None: sfd_row_fake = None sfd_col_fake = None if self.sfd_row_desc is not None: - rest_n2 = cute.sym_int() - stride_sfd_rest_n2 = cute.sym_int(divisibility=32 * 4 * 4) - stride_sfd_rest_tensor_m_128 = cute.sym_int(divisibility=32 * 4 * 4) - sfd_row_fake = self._make_fake_cute_tensor( - dtype=self.sfd_row_desc.dtype, - shape=(32, 4, tensor_m_128, 4, rest_n2, 1), - stride=( - 16, - 4, - stride_sfd_rest_n2, - 1, - 512, - stride_sfd_rest_tensor_m_128, - ), - ) + if self.sfd_row_is_flat: + sfd_row_fake = make_flat_sf_fake(self, self.sfd_row_desc) + else: + sfd_tensor_m_128 = cute.sym_int() + rest_n2 = cute.sym_int() + stride_sfd_rest_n2 = cute.sym_int(divisibility=32 * 4 * 4) + stride_sfd_rest_tensor_m_128 = cute.sym_int(divisibility=32 * 4 * 4) + sfd_row_fake = self._make_fake_cute_tensor( + dtype=self.sfd_row_desc.dtype, + shape=(32, 4, sfd_tensor_m_128, 4, rest_n2, 1), + stride=( + 16, + 4, + stride_sfd_rest_n2, + 1, + 512, + stride_sfd_rest_tensor_m_128, + ), + ) if self.sfd_col_desc is not None: - tensor_n2_128 = cute.sym_int() - rest_m = cute.sym_int() - stride_sfd_rest_m = cute.sym_int(divisibility=32 * 4 * 4) - stride_sfd_n2 = cute.sym_int(divisibility=32 * 4 * 4) - sfd_col_fake = self._make_fake_cute_tensor( - dtype=self.sfd_col_desc.dtype, - shape=(32, 4, tensor_n2_128, 4, rest_m, 1), - stride=(16, 4, stride_sfd_rest_m, 1, 512, stride_sfd_n2), - ) + if self.sfd_col_is_flat: + sfd_col_fake = make_flat_sf_fake(self, self.sfd_col_desc) + else: + tensor_n2_128 = cute.sym_int() + rest_m = cute.sym_int() + stride_sfd_rest_m = cute.sym_int(divisibility=32 * 4 * 4) + stride_sfd_n2 = cute.sym_int(divisibility=32 * 4 * 4) + sfd_col_fake = self._make_fake_cute_tensor( + dtype=self.sfd_col_desc.dtype, + shape=(32, 4, tensor_n2_128, 4, rest_m, 1), + stride=(16, 4, stride_sfd_rest_m, 1, 512, stride_sfd_n2), + ) _compiled_kernel = cute.compile( gemm_swiglu, @@ -749,20 +820,20 @@ def execute( ) self._logger.debug("Executing grouped_gemm_swiglu kernel") self._compiled_kernel( - a_tensor=a_tensor, - b_tensor=b_tensor, - c_tensor=c_tensor, - d_tensor=d_tensor, - d_col_tensor=d_col_tensor, - sfa_tensor=sfa_tensor, - sfb_tensor=sfb_tensor, - sfd_row_tensor=sfd_row_tensor, - sfd_col_tensor=sfd_col_tensor, + a_tensor=unsqueeze_l_dim(a_tensor), + b_tensor=to_kernel_b(b_tensor), + c_tensor=unsqueeze_l_dim(c_tensor), + d_tensor=unsqueeze_l_dim(d_tensor), + d_col_tensor=unsqueeze_l_dim(d_col_tensor), + sfa_tensor=to_kernel_sf(sfa_tensor, self.sfa_is_flat), + sfb_tensor=to_kernel_sf(sfb_tensor, self.sfb_is_flat), + sfd_row_tensor=to_kernel_sf(sfd_row_tensor, self.sfd_row_is_flat), + sfd_col_tensor=to_kernel_sf(sfd_col_tensor, self.sfd_col_is_flat), amax_tensor=amax_tensor, norm_const_tensor=norm_const_tensor, padded_offsets=padded_offsets, alpha_tensor=alpha_tensor, - prob_tensor=prob_tensor, + prob_tensor=to_kernel_prob(prob_tensor), stream=current_stream, ) @@ -781,7 +852,7 @@ def grouped_gemm_swiglu_wrapper_sm100( sfa_tensor: torch.Tensor, sfb_tensor: torch.Tensor, padded_offsets: torch.Tensor, - alpha_tensor: torch.Tensor, + alpha_tensor: Optional[torch.Tensor] = None, norm_const_tensor: Optional[torch.Tensor] = None, prob_tensor: Optional[torch.Tensor] = None, acc_dtype: Optional[torch.dtype] = None, @@ -801,13 +872,23 @@ def grouped_gemm_swiglu_wrapper_sm100( This function creates the API, compiles, and executes in one call. Compiled kernels are cached for reuse when called with the same configuration. + Canonical layouts (additive): each input is also accepted in its natural + row-major form and normalized internally -- A as (valid_m, k), B as (l, n, k) + C-contiguous, SFA/SFB as dense C-contiguous buffers of any shape with the + MMA-tiled element count (e.g. flat 1-D, or physical + (l, mn//128, ceil(ceil(k/sf_vec_size)/4), 32, 4, 4)), and prob as (valid_m,) + float32 or bfloat16. When A is canonical (2-D), outputs come back natural-shaped: + c (valid_m, n), d/d_col (valid_m, n//2) row-major, and sfd_row/sfd_col as + C-contiguous physical (1, mn//128, rest, 32, 4, 4) buffers. The pre-permuted + kernel-facing forms below keep working unchanged. + Args: - a_tensor: Input A tensor (valid_m, k, 1) - b_tensor: Weight B tensor (n, k, l) - sfa_tensor: Scale factor A - sfb_tensor: Scale factor B + a_tensor: Input A tensor (valid_m, k, 1), or canonical (valid_m, k) row-major + b_tensor: Weight B tensor (n, k, l) k-major, or canonical (l, n, k) row-major + sfa_tensor: Scale factor A (MMA-tiled view, or canonical dense buffer) + sfb_tensor: Scale factor B (MMA-tiled view, or canonical dense buffer) padded_offsets: End offset per expert after padding (l,) - alpha_tensor: Per-group scaling + alpha_tensor: Per-group scaling; None defaults to ones (cached) norm_const_tensor: Optional normalization constant. Required when using FP8 input configurations (i.e., when a_tensor.dtype is FP8 and sfa_tensor.dtype is FP8). Should be None for FP4/BF16 input configurations. @@ -857,13 +938,29 @@ def grouped_gemm_swiglu_wrapper_sm100( acc_dtype = _convert_to_cutlass_data_type(acc_dtype) if acc_dtype is not None else cutlass.Float32 c_dtype = _convert_to_cutlass_data_type(c_dtype) if c_dtype is not None else cutlass.BFloat16 d_dtype = _convert_to_cutlass_data_type(d_dtype) if d_dtype is not None else cutlass.BFloat16 - valid_m, k, _ = a_tensor.shape - n, _, l = b_tensor.shape + valid_m = a_tensor.shape[0] + if is_canonical_b(b_tensor): + l, n, _ = b_tensor.shape + else: + n, _, l = b_tensor.shape n_out = n // 2 # After SwiGLU + # Canonical (sum_m, k) A selects natural-shaped outputs: (m, x) row-major C/D and + # dense C-contiguous SFD buffers instead of the pre-permuted kernel-facing views. + canonical_outputs = a_tensor.ndim == 2 + + if alpha_tensor is None: + alpha_tensor = default_alpha_ones(l, a_tensor.device) + _logger.debug("grouped_gemm_swiglu_wrapper_sm100: Creating output tensors c_tensor, d_tensor, d_col_tensor") - if cd_major == "n": + if cd_major != "n": + raise ValueError(f"cd_major must be 'n', got {cd_major}") + if canonical_outputs: + c_tensor = torch.empty((valid_m, n), dtype=framework_dtype(c_dtype, "torch"), device=a_tensor.device) + d_tensor = torch.empty((valid_m, n_out), dtype=framework_dtype(d_dtype, "torch"), device=a_tensor.device) + d_col_tensor = torch.empty((valid_m, n_out), dtype=framework_dtype(d_dtype, "torch"), device=a_tensor.device) + else: # 1, m, n, permute (1, 2, 0) -> (m, n, 1) c_tensor = torch.empty_strided((valid_m, n, 1), (n, 1, valid_m * n), dtype=framework_dtype(c_dtype, "torch"), device=a_tensor.device) d_tensor = torch.empty_strided( @@ -878,8 +975,6 @@ def grouped_gemm_swiglu_wrapper_sm100( dtype=framework_dtype(d_dtype, "torch"), device=a_tensor.device, ) - else: - raise ValueError(f"cd_major must be 'n', got {cd_major}") sfd_row_tensor = None sfd_col_tensor = None @@ -906,7 +1001,7 @@ def grouped_gemm_swiglu_wrapper_sm100( 4, 4, ) - sfd_row_tensor = torch.empty(mma_shape_row, dtype=sf_dtype, device=a_tensor.device).permute(mma_permute_order) + sfd_row_tensor = torch.empty(mma_shape_row, dtype=sf_dtype, device=a_tensor.device) # sfd_col: l=1, mn=n_out, k=valid_m sf_k_col = ceil_div(valid_m, sf_vec_size) @@ -918,7 +1013,10 @@ def grouped_gemm_swiglu_wrapper_sm100( 4, 4, ) - sfd_col_tensor = torch.empty(mma_shape_col, dtype=sf_dtype, device=a_tensor.device).permute(mma_permute_order) + sfd_col_tensor = torch.empty(mma_shape_col, dtype=sf_dtype, device=a_tensor.device) + if not canonical_outputs: + sfd_row_tensor = sfd_row_tensor.permute(mma_permute_order) + sfd_col_tensor = sfd_col_tensor.permute(mma_permute_order) if valid_m == 0: if d_dtype in (cutlass.BFloat16, cutlass.Float16): @@ -967,6 +1065,15 @@ def stride_order(tensor: torch.Tensor) -> Tuple[int, ...]: m_aligned, discrete_col_sfd, prob_tensor is not None, + # Canonical-vs-kernel-facing input forms compile different signatures. + prob_tensor.dtype if prob_tensor is not None else None, + prob_tensor.ndim if prob_tensor is not None else None, + is_flat_sf(sfa_tensor), + is_flat_sf(sfb_tensor), + canonical_outputs, + # The compiled signature binds the SF dtype (e8m0 vs e4m3). + sfa_tensor.dtype, + sfb_tensor.dtype, ) if cache_key in _cache_of_GroupedGemmSwigluSm100Objects: diff --git a/python/cudnn/gemm/cutedsl/grouped/swiglu/grouped_gemm_swiglu_quant.py b/python/cudnn/gemm/cutedsl/grouped/swiglu/grouped_gemm_swiglu_quant.py index 3637c72d2..9986b9a2f 100644 --- a/python/cudnn/gemm/cutedsl/grouped/swiglu/grouped_gemm_swiglu_quant.py +++ b/python/cudnn/gemm/cutedsl/grouped/swiglu/grouped_gemm_swiglu_quant.py @@ -2311,7 +2311,7 @@ def kernel( # if cutlass.const_expr(prob is not None): mPosition = tile_info[0] * self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape) + tidx - mProb = prob[mPosition, 0, 0] + mProb = prob[mPosition, 0, 0].to(cutlass.Float32) else: mProb = cutlass.Float32(1.0) diff --git a/test/python/fe_api/grouped_gemm/test_grouped_gemm_canonical_layouts.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_canonical_layouts.py new file mode 100644 index 000000000..985e9fedc --- /dev/null +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_canonical_layouts.py @@ -0,0 +1,282 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +""" +Canonical (natural row-major) layout acceptance for the contiguous grouped GEMM +SwiGLU forward and dSwiGLU backward wrappers. + +Each test runs the same problem twice on identical device buffers -- once through +the pre-permuted kernel-facing views (legacy) and once through the canonical +forms: A (sum_m, k) row-major, B (l, n, k) C-contiguous, SFA/SFB as dense +C-contiguous buffers (physical atom shape or flat 1-D), prob (sum_m,) -- and +requires identical results. Canonical calls return natural-shaped outputs. +""" + +import pytest +import torch +from test_utils import torch_fork_set_rng + +from fe_api.grouped_gemm.test_grouped_gemm_swiglu_utils import ( + grouped_gemm_swiglu_init, + allocate_grouped_gemm_input_tensors, +) +from fe_api.grouped_gemm.test_grouped_gemm_dswiglu_utils import ( + allocate_grouped_gemm_dswiglu_tensors, +) + +# Inverse of the (3, 4, 1, 5, 2, 0) MMA permute: recovers the physical C-contiguous +# (l, mn//128, rest_k, 32, 4, 4) allocation underneath a kernel-facing SF view. +SF_PHYSICAL_PERMUTE = (5, 2, 4, 0, 1, 3) + + +def swiglu_case(request, ab_dtype, d_dtype, sf_vec_size, sf_dtype): + cfg = grouped_gemm_swiglu_init( + request=request, + ab_dtype=ab_dtype, + c_dtype=torch.bfloat16, + d_dtype=d_dtype, + cd_major="n", + acc_dtype=torch.float32, + mma_tiler_mn=(256, 256), + cluster_shape_mn=(2, 1), + sf_vec_size=sf_vec_size, + sf_dtype=sf_dtype, + ) + inputs = allocate_grouped_gemm_input_tensors( + n=cfg["n"], + k=cfg["k"], + l=cfg["l"], + group_m_list=cfg["group_m_list"], + ab_dtype=cfg["ab_dtype"], + sf_dtype=cfg["sf_dtype"], + sf_vec_size=cfg["sf_vec_size"], + m_aligned=cfg["m_aligned"], + ) + return cfg, inputs + + +def run_swiglu(cfg, inputs, canonical, flat_sf=False, prob_dtype=None, alpha=...): + from cudnn import grouped_gemm_swiglu_wrapper_sm100 + + prob = inputs["prob_tensor"] + if alpha is ...: + alpha = inputs["alpha_tensor"] + if canonical: + a = inputs["a_tensor"].squeeze(-1) + b = inputs["b_tensor"].permute(2, 0, 1) + assert b.is_contiguous() + sfa = inputs["sfa_tensor"].permute(*SF_PHYSICAL_PERMUTE) + sfb = inputs["sfb_tensor"].permute(*SF_PHYSICAL_PERMUTE) + assert sfa.is_contiguous() and sfb.is_contiguous() + if flat_sf: + sfa = sfa.reshape(-1) + sfb = sfb.reshape(-1) + prob = prob.view(-1) + else: + a, b = inputs["a_tensor"], inputs["b_tensor"] + sfa, sfb = inputs["sfa_tensor"], inputs["sfb_tensor"] + if prob_dtype is not None: + prob = prob.to(prob_dtype) + return grouped_gemm_swiglu_wrapper_sm100( + a_tensor=a, + b_tensor=b, + sfa_tensor=sfa, + sfb_tensor=sfb, + padded_offsets=inputs["padded_offsets_tensor"], + alpha_tensor=alpha, + norm_const_tensor=inputs.get("norm_const_tensor"), + prob_tensor=prob, + d_dtype=cfg["d_dtype"], + sf_vec_size=cfg["sf_vec_size"], + ) + + +def assert_same(name, canonical_t, legacy_t): + if legacy_t is None and canonical_t is None: + return + legacy_flat = legacy_t.reshape(-1) if legacy_t.is_contiguous() else legacy_t.contiguous().reshape(-1) + canonical_flat = canonical_t.reshape(-1) if canonical_t.is_contiguous() else canonical_t.contiguous().reshape(-1) + if canonical_flat.dtype in (torch.float8_e8m0fnu, torch.float8_e4m3fn, torch.float8_e5m2, torch.float4_e2m1fn_x2): + legacy_flat = legacy_flat.view(torch.uint8) + canonical_flat = canonical_flat.view(torch.uint8) + torch.testing.assert_close(canonical_flat, legacy_flat, rtol=0, atol=0, msg=lambda m: f"{name}: {m}") + + +SWIGLU_CASES = [ + pytest.param(torch.float8_e4m3fn, torch.float8_e4m3fn, 32, torch.float8_e8m0fnu, id="fp8"), + pytest.param(torch.float4_e2m1fn_x2, torch.bfloat16, 16, torch.float8_e8m0fnu, id="fp4"), +] + + +@pytest.mark.L0 +@torch_fork_set_rng(seed=0) +@pytest.mark.parametrize("flat_sf", [False, True], ids=["physical_sf", "flat_sf"]) +@pytest.mark.parametrize("ab_dtype,d_dtype,sf_vec_size,sf_dtype", SWIGLU_CASES) +def test_grouped_gemm_swiglu_canonical_matches_legacy(request, ab_dtype, d_dtype, sf_vec_size, sf_dtype, flat_sf): + try: + cfg, inputs = swiglu_case(request, ab_dtype, d_dtype, sf_vec_size, sf_dtype) + legacy = run_swiglu(cfg, inputs, canonical=False) + canonical = run_swiglu(cfg, inputs, canonical=True, flat_sf=flat_sf) + except ImportError: + pytest.skip("Environment not supported: cudnn optional dependencies not installed") + + m = inputs["tensor_m"] + n, n_out = cfg["n"], cfg["n"] // 2 + assert canonical["c_tensor"].shape == (m, n) + assert canonical["d_tensor"].shape == (m, n_out) + assert canonical["d_col_tensor"].shape == (m, n_out) + assert_same("c", canonical["c_tensor"], legacy["c_tensor"]) + assert_same("d", canonical["d_tensor"], legacy["d_tensor"]) + if legacy["sfd_row_tensor"] is not None: + # d_col is only written on the quantized (generate_sfd) path + assert_same("d_col", canonical["d_col_tensor"], legacy["d_col_tensor"]) + assert canonical["sfd_row_tensor"].is_contiguous() + assert canonical["sfd_col_tensor"].is_contiguous() + assert_same("sfd_row", canonical["sfd_row_tensor"], legacy["sfd_row_tensor"].permute(*SF_PHYSICAL_PERMUTE)) + assert_same("sfd_col", canonical["sfd_col_tensor"], legacy["sfd_col_tensor"].permute(*SF_PHYSICAL_PERMUTE)) + if legacy["amax_tensor"] is not None: + assert_same("amax", canonical["amax_tensor"], legacy["amax_tensor"]) + + +@pytest.mark.L0 +@torch_fork_set_rng(seed=0) +def test_grouped_gemm_swiglu_canonical_bf16_prob(request): + try: + cfg, inputs = swiglu_case(request, torch.float8_e4m3fn, torch.float8_e4m3fn, 32, torch.float8_e8m0fnu) + # prob values are small integers, exactly representable in bf16, so the + # bf16-prob run must match the fp32-prob run bitwise. + legacy = run_swiglu(cfg, inputs, canonical=False) + canonical = run_swiglu(cfg, inputs, canonical=True, prob_dtype=torch.bfloat16) + except ImportError: + pytest.skip("Environment not supported: cudnn optional dependencies not installed") + assert_same("d", canonical["d_tensor"], legacy["d_tensor"]) + assert_same("c", canonical["c_tensor"], legacy["c_tensor"]) + + +@pytest.mark.L0 +@torch_fork_set_rng(seed=0) +def test_grouped_gemm_swiglu_alpha_defaults_to_ones(request): + try: + cfg, inputs = swiglu_case(request, torch.float8_e4m3fn, torch.float8_e4m3fn, 32, torch.float8_e8m0fnu) + ones = torch.ones_like(inputs["alpha_tensor"]) + explicit = run_swiglu(cfg, inputs, canonical=True, alpha=ones) + defaulted = run_swiglu(cfg, inputs, canonical=True, alpha=None) + except ImportError: + pytest.skip("Environment not supported: cudnn optional dependencies not installed") + assert_same("d", defaulted["d_tensor"], explicit["d_tensor"]) + assert_same("c", defaulted["c_tensor"], explicit["c_tensor"]) + + +def dswiglu_case(request): + cfg = grouped_gemm_swiglu_init( + request=request, + ab_dtype=torch.float8_e4m3fn, + c_dtype=torch.bfloat16, + d_dtype=torch.float8_e4m3fn, + cd_major="n", + acc_dtype=torch.float32, + mma_tiler_mn=(256, 256), + cluster_shape_mn=(2, 1), + sf_vec_size=32, + sf_dtype=torch.float8_e8m0fnu, + ) + inputs = allocate_grouped_gemm_input_tensors( + n=cfg["n"], + k=cfg["k"], + l=cfg["l"], + group_m_list=cfg["group_m_list"], + ab_dtype=cfg["ab_dtype"], + sf_dtype=cfg["sf_dtype"], + sf_vec_size=cfg["sf_vec_size"], + m_aligned=cfg["m_aligned"], + ) + inputs, _ = allocate_grouped_gemm_dswiglu_tensors( + tensor_m=inputs["tensor_m"], + n=cfg["n"], + l=cfg["l"], + ab_dtype=cfg["ab_dtype"], + c_dtype=cfg["c_dtype"], + d_dtype=cfg["d_dtype"], + cd_major=cfg["cd_major"], + sf_dtype=cfg["sf_dtype"], + sf_vec_size=cfg["sf_vec_size"], + input_tensors=inputs, + ) + return cfg, inputs + + +def run_dswiglu(cfg, inputs, canonical, flat_sf=False, prob_dtype=None): + from cudnn import grouped_gemm_dswiglu_wrapper_sm100 + + prob = inputs["prob_tensor"] + if canonical: + a = inputs["a_tensor"].squeeze(-1) + b = inputs["b_tensor"].permute(2, 0, 1) + c = inputs["c_tensor"].squeeze(-1) + sfa = inputs["sfa_tensor"].permute(*SF_PHYSICAL_PERMUTE) + sfb = inputs["sfb_tensor"].permute(*SF_PHYSICAL_PERMUTE) + assert b.is_contiguous() and sfa.is_contiguous() and sfb.is_contiguous() + if flat_sf: + sfa = sfa.reshape(-1) + sfb = sfb.reshape(-1) + prob = prob.view(-1) + else: + a, b, c = inputs["a_tensor"], inputs["b_tensor"], inputs["c_tensor"] + sfa, sfb = inputs["sfa_tensor"], inputs["sfb_tensor"] + if prob_dtype is not None: + prob = prob.to(prob_dtype) + return grouped_gemm_dswiglu_wrapper_sm100( + a_tensor=a, + b_tensor=b, + c_tensor=c, + sfa_tensor=sfa, + sfb_tensor=sfb, + padded_offsets=inputs["padded_offsets_tensor"], + alpha_tensor=inputs["alpha_tensor"], + beta_tensor=inputs["beta_tensor"], + prob_tensor=prob, + norm_const_tensor=inputs.get("norm_const_tensor"), + d_dtype=cfg["d_dtype"], + sf_vec_size=cfg["sf_vec_size"], + ) + + +@pytest.mark.L0 +@torch_fork_set_rng(seed=0) +@pytest.mark.parametrize("flat_sf", [False, True], ids=["physical_sf", "flat_sf"]) +def test_grouped_gemm_dswiglu_canonical_matches_legacy(request, flat_sf): + try: + cfg, inputs = dswiglu_case(request) + legacy = run_dswiglu(cfg, inputs, canonical=False) + canonical = run_dswiglu(cfg, inputs, canonical=True, flat_sf=flat_sf) + except ImportError: + pytest.skip("Environment not supported: cudnn optional dependencies not installed") + + m = inputs["tensor_m"] + n2 = cfg["n"] * 2 + assert canonical["d_row_tensor"].shape == (m, n2) + assert canonical["d_col_tensor"].shape == (m, n2) + assert canonical["dprob_tensor"].shape == (m,) + assert_same("d_row", canonical["d_row_tensor"], legacy["d_row_tensor"]) + assert_same("d_col", canonical["d_col_tensor"], legacy["d_col_tensor"]) + # dprob accumulates with atomic float adds; ordering differs between launches. + torch.testing.assert_close(canonical["dprob_tensor"].reshape(-1), legacy["dprob_tensor"].reshape(-1), rtol=1e-4, atol=1e-4) + if legacy["sfd_row_tensor"] is not None: + assert canonical["sfd_row_tensor"].is_contiguous() + assert_same("sfd_row", canonical["sfd_row_tensor"], legacy["sfd_row_tensor"].permute(*SF_PHYSICAL_PERMUTE)) + assert_same("sfd_col", canonical["sfd_col_tensor"], legacy["sfd_col_tensor"].permute(*SF_PHYSICAL_PERMUTE)) + if legacy["amax_tensor"] is not None: + assert_same("amax", canonical["amax_tensor"], legacy["amax_tensor"]) + + +@pytest.mark.L0 +@torch_fork_set_rng(seed=0) +def test_grouped_gemm_dswiglu_canonical_bf16_prob(request): + try: + cfg, inputs = dswiglu_case(request) + legacy = run_dswiglu(cfg, inputs, canonical=False) + canonical = run_dswiglu(cfg, inputs, canonical=True, prob_dtype=torch.bfloat16) + except ImportError: + pytest.skip("Environment not supported: cudnn optional dependencies not installed") + assert_same("d_row", canonical["d_row_tensor"], legacy["d_row_tensor"]) + torch.testing.assert_close(canonical["dprob_tensor"].reshape(-1), legacy["dprob_tensor"].reshape(-1), rtol=1e-4, atol=1e-4)