> For clean Markdown of any page, append .md to the page URL.
> For a complete documentation index, see https://docs.nvidia.com/cudnn/llms.txt.
> For AI client integration (Claude Code, Cursor, etc.), connect to the MCP server at https://docs.nvidia.com/cudnn/_mcp/server.

# Grouped GEMM + GLU (SM100)

**This is an experimental API and subject to change.**

## JAX support

Supports **JAX arrays** on the BF16 backend in discrete weight mode (swiglu and geglu): `b_ptrs` as a packed little-endian uint8 pointer array (8 bytes per pointer; int64 accepted with jax x64 mode), outputs allocated as n-major C-contiguous `jnp` arrays. Dense `b_tensor` (expert-outermost strides), column-major `bias_tensor`, and the block-scaled backend (MMA-interleaved scale-factor layouts) are not expressible as JAX arrays and raise clear errors. The wrapper is eager, on the CUDA legacy default stream: `block_until_ready` inputs, synchronize before reading outputs; keep weight arrays alive until the kernel completes.

For jitted JAX programs use the `jax.jit`-compatible XLA custom-call entry point `grouped_gemm_glu_jax_sm100` (built on `cudnn.jax.call`; discrete mode, no bias, `b_major="k"`): outputs are fresh XLA-managed arrays with rows at/past `padded_offsets[-1]` zero-filled, no manual synchronization needed. `linear_offset` is a compile-time constant (each distinct value compiles a new specialization). Under tracing the `padded_offsets` *values* cannot be host-validated, and the per-expert weight buffers behind `b_ptrs` must stay alive and unmoved across every execution of the traced computation.

## Overview

**Unified Grouped GEMM + GLU fusion**: one public class and wrapper select a
plain BF16 or legacy block-scaled grouped GEMM fused with a GLU epilogue
(SwiGLU, GeGLU, or block-scaled SiTU-GLU) on NVIDIA Blackwell GPUs (SM100/SM103),
with a block-scaled forward path on Rubin (SM107). The operation is
implemented with CUTLASS/CuTe DSL.

This is a **unified API** that supports both weight layout modes:
- **Dense mode**: All expert weights packed into a single contiguous `(N, K, L)` tensor
- **Discrete mode**: Per-expert weight pointers (no weight stacking required)

Supported activation functions:
- **SwiGLU**: `act_func="swiglu"` (default)
- **GeGLU**: `act_func="geglu"`
- **SiTU-GLU**: `act_func="situglu"` (block-scaled SM100/SM103 only)

Groups are contiguous in the M dimension and described by `padded_offsets` (cumulative aligned end offsets).

### Backend dispatch

| Operand contract | Selected backend |
| --- | --- |
| `A` and `B` are BF16 | BF16 |
| matching supported FP4/FP8 `A` and `B` plus scale descriptors | block-scaled |

Mixed families and unsupported pairs are rejected before allocation or
compilation. Each backend's argument contract is described below.
SiTU-GLU is not available on the BF16 or Rubin backends.
The block-scaled GeGLU forward path supports runtime activation alpha, clamp
limits, and linear offset on both Blackwell and Rubin. The corresponding
[dGeGLU backward path](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm_dglu) supports the same activation
configuration, with the parameter values included in the backward compiled-kernel
cache key.

## BF16 contract

Pass `sfa_tensor=None`, `sfb_tensor=None` (or `sfb_ptrs=None`), and
`norm_const_tensor=None`; keep `sf_vec_size=16` and `discrete_col_sfd=False`.
Non-`None` scale controls are an error.

### Tensors, layouts, and equation

For padded rows `M`, reduction dimension `K`, pre-GLU width `N`, and `L`
experts, BF16 uses:

- `A`: `(M, K, 1)`, stride `(K, 1, M*K)`, BF16;
- dense `B`: `(N, K, L)`, K-major stride `(K, 1, N*K)`, BF16;
- discrete `b_ptrs`: contiguous CUDA int64 pointers to expert `(N, K)` BF16
  matrices, with `n=N`, `b_dtype=torch.bfloat16`, and `b_major="k"` or `"n"`;
- `padded_offsets`: `(L,)`, stride `(1,)`, int32 cumulative 256-aligned ends;
- `alpha`: `(L,)`, FP32; `prob`: `(M, 1, 1)`, stride `(1, 1, 1)`, FP32;
- optional bias: `(N, L)`, stride `(1, N)`, BF16/FP16/FP32;
- `C`: `(M, N, 1)`, stride `(N, 1, M*N)`;
- `D`: `(M, N/2, 1)`, stride `(N/2, 1, M*N/2)`.

For expert `g`, first compute

$$
C_g = \alpha_g A_g B_g^T + \mathrm{bias}_g.
$$

Columns are paired as alternating 32-wide gate/up blocks. For SwiGLU,

$$
D_g = \mathrm{prob}_g \cdot \mathrm{up}(C_g) \cdot
      \mathrm{silu}(\mathrm{gate}(C_g)).
$$

For GeGLU, let `gate = min(gate(C), 7)`,
`up = clamp(up(C), -7, 7)`, `geglu_alpha=1.702`, and the default
`linear_offset=1`:

$$
D_g = \mathrm{prob}_g \cdot (\mathrm{up}+\mathrm{linear\_offset})
      \cdot \mathrm{gate} \cdot \sigma(1.702\,\mathrm{gate}).
$$

`C`/`D` may be BF16, FP16, or FP32. `N` is divisible by 64. The pointer-array
tensor is stream-recorded; every pointed allocation must remain alive and
unchanged until the launch stream completes.

The wrapper return order is exactly `c_tensor`, `d_tensor`, `d_col_tensor`,
`amax_tensor`, `sfd_row_tensor`, `sfd_col_tensor`. On BF16,
`d_col_tensor`, `amax_tensor`, `sfd_row_tensor`, and `sfd_col_tensor` are
always `None`; `c_tensor` is `None` unless `generate_c=True`.

## Block-scaled contract

The block-scaled backend performs:
1. **Block-scaled grouped GEMM**: Low-precision GEMM (FP4, FP8) with per-block scale factors across multiple expert groups
2. **GLU activation**: Fused SwiGLU, GeGLU, or SiTU-GLU activation applied to the GEMM output
3. **Optional quantized output**: Produces row and column scale factors for downstream quantization

### Shapes

### Equations

For SiTU-GLU, with gate branch `G` and up branch `U`, the fused epilogue computes

$$
D = \mathrm{prob}\,
    \left[\beta_1\tanh(G/\beta_1)\sigma(G)\right]
    \left[\beta_2\tanh(U/\beta_2)\right].
$$

where `beta_1 = situ_beta1` and `beta_2 = situ_beta2`, with defaults
`beta_1 = 4.0` and `beta_2 = 25.0`. `situ_beta1` specializes the compiled kernel
and is part of its cache key; `situ_beta2` is a runtime FP32 scalar and does not
create a new compiled-kernel cache entry.

- **Inputs**
  - `A`: contiguous activation tensor across all groups, shape `(valid_m, K, 1)`
  - `B` (dense): weight tensor across all groups, shape `(N, K, L)`
  - `B` (discrete): per-expert weight pointers, `b_ptrs` shape `(num_experts,)` of int64
  - `SFA`: scale factor tensor for A, shape `(32, 4, ceil(valid_m/128), 4, ceil(ceil(K/sf_vec_size)/4), 1)`
  - `SFB` (dense): scale factor tensor for B, shape `(32, 4, ceil(N/128), 4, ceil(ceil(K/sf_vec_size)/4), L)`
  - `SFB` (discrete): per-expert SFB pointers, `sfb_ptrs` shape `(num_experts,)` of int64
  - `padded_offsets`: cumulative sum of aligned group M sizes, shape `(L,)`. `valid_m = padded_offsets[-1]`
  - `alpha`: per-group scaling factors, shape `(L,)`
  - `bias` (optional): per-expert bias tensor, shape `(N, L)` with stride `(1, N)`
  - `prob`: per-row gating probabilities, shape `(valid_m, 1, 1)`
  - `norm_const`: normalization constant for FP8 quantization, shape `(1,)`
- **Outputs**
  - `C`: intermediate GEMM result, shape `(valid_m, N, 1)`
  - `D`: row-quantized GLU output, shape `(valid_m, N/2, 1)`
  - `D_col`: column-quantized GLU output, shape `(valid_m, N/2, 1)`
  - `SFD_row`: row scale factors (when `d_dtype` is FP8), shape `(32, 4, ceil(valid_m/128), 4, ceil(ceil((N/2)/sf_vec_size)/4), 1)`
  - `SFD_col`: column scale factors (when `d_dtype` is FP8), shape `(32, 4, ceil((N/2)/128), 4, ceil(ceil(valid_m/sf_vec_size)/4), 1)`
  - `amax`: per-group amax (when `d_dtype` is bf16/fp16), shape `(L, 1)`

**Step 1: Block-scaled grouped GEMM** (per group `g` with rows `m` in `[padded_offsets[g-1], padded_offsets[g])`):

$$
C[m, n] = \alpha_g \sum_{k} \text{dequantize}(A[m, k], \text{SFA}) \cdot \text{dequantize}(B[n, k, g], \text{SFB})
$$

**Step 2: GLU epilogue** (performed by pairing 32-column blocks along `N`):

Let block size `G = 32`. For each pair of consecutive 32-wide column blocks:
- Gate block: `G_b = C[:, 2·b·G : 2·b·G + G]`
- Up block: `U_b = C[:, 2·b·G + G : 2·b·G + 2·G]`

For **SwiGLU** (`act_func="swiglu"`):

$$
D[:, bG:(b+1)G] = \text{prob} \cdot U_b \cdot \text{swish}(G_b), \quad \text{swish}(x) = x \cdot \sigma(x)
$$

For **GeGLU** (`act_func="geglu"`):

$$
\widehat{G}_b = \min(G_b, \text{glu\_clamp\_max}), \qquad
\widehat{U}_b = \operatorname{clamp}(U_b, \text{glu\_clamp\_min}, \text{glu\_clamp\_max})
$$

$$
D[:, bG:(b+1)G] = \text{prob} \cdot
    (\widehat{U}_b + \text{linear\_offset}) \cdot
    \widehat{G}_b \cdot \sigma(\text{geglu\_alpha} \cdot \widehat{G}_b)
$$

Only the gate's upper bound is clamped. The offset is added after clamping
the up branch. The gate nonlinearity is `g * sigmoid(geglu_alpha * g)`;
`silu(geglu_alpha * g)` would introduce an extra factor of `geglu_alpha`.
The optional `C` output stores the GEMM result before clamping. The epilogue
uses the FP32 accumulator, without rounding it to the `C` output dtype first.

The defaults are `geglu_alpha=1.702`, `glu_clamp_max=7.0`,
`glu_clamp_min=-7.0`, and `linear_offset=1.0`. For DeepSeek V4 clamped SwiGLU,
select `act_func="geglu"` with `geglu_alpha=1.0`, `linear_offset=0.0`,
`glu_clamp_max=L`, and `glu_clamp_min=-L`, where `L` is the model's clamp
limit. `act_func="swiglu"` does not apply these clamp parameters.

For block-scaled dense and discrete calls, these four activation parameters
are runtime FP32 scalars: changing their values reuses the same compiled
kernel. `geglu_alpha` scales the sigmoid input and is independent of the
per-expert GEMM scaling tensor `alpha_tensor`.

**Step 3: Optional output quantization** (when SFD outputs are generated):

$$
\text{SFD\_row}[m, n] = \text{norm\_const} \cdot \max_{k \in \text{block}} |D[m, k]| \cdot \text{rcp\_max}
$$

$$
D_{\text{quantized}}[m, n] = D[m, n] \cdot \frac{\text{norm\_const}}{\text{SFD\_row}[m, n]}
$$

### Diagram

```text
 A (valid_m×K×1)    B (N×K×L) or b_ptrs    padded_offsets
 SFA                SFB or sfb_ptrs              |
   |                 |                           |
   |    +------------+                           |
   |    |                                        |
   v    v                                        v
  Dequantize → Grouped GEMM (per group ranges) → Select B[:,:,group_idx]
                    |
                    | × alpha[group_idx]
                    v
               C (valid_m×N×1)
                    |
                    | Pair 32-col blocks: [G0|U0|G1|U1|...]
                    |     U_b × swish(G_b)  [SwiGLU]
                    |     clamp gate/up, then
                    |     (U_b+offset) × G_b·σ(alpha·G_b)  [GeGLU]
                    v
                    | × prob
                    v
               D (valid_m×N/2×1)
                    |
         +----------+-----------+
         |                      |
         v                      v
    Row Quantize           Col Quantize
         |                      |
         v                      v
    D_row, SFD_row        D_col, SFD_col
```

---

## API usage

### BF16

#### High-level wrapper

```python
import cudnn
import torch

# Dense wrapper. The required scale positions are explicitly None for BF16.
out = cudnn.grouped_gemm_glu_wrapper_sm100(
    a_tensor=a,
    sfa_tensor=None,
    padded_offsets=padded_offsets,
    alpha_tensor=alpha,
    b_tensor=b,
    sfb_tensor=None,
    bias_tensor=bias,
    prob_tensor=prob,
    act_func="swiglu",
    generate_c=True,
    use_dynamic_sched=True,
)
c, d, d_col, amax, sfd_row, sfd_col = out

# Discrete wrapper.
out = cudnn.grouped_gemm_glu_wrapper_sm100(
    a_tensor=a,
    sfa_tensor=None,
    padded_offsets=padded_offsets,
    alpha_tensor=alpha,
    b_ptrs=b_ptrs,
    sfb_ptrs=None,
    n=N,
    b_dtype=torch.bfloat16,
    prob_tensor=prob,
    act_func="geglu",
)
```

#### Class API

```python
op = cudnn.GroupedGemmGluSm100(
    sample_a=a,
    sample_c=c,
    sample_d=d,
    sample_d_col=None,
    sample_sfa=None,
    sample_padded_offsets=padded_offsets,
    sample_alpha=alpha,
    sample_b=b,
    sample_sfb=None,
    sample_prob=prob,
    act_func="swiglu",
    generate_c=True,
)
assert op.check_support()
op.compile()
op.execute(
    a_tensor=a, c_tensor=c, d_tensor=d, sfa_tensor=None,
    padded_offsets=padded_offsets, alpha_tensor=alpha,
    b_tensor=b, sfb_tensor=None, prob_tensor=prob,
)
```

`use_dynamic_sched=False` uses static scheduling; `True` caches a dynamic-M
callable for compatible shapes. Cache keys include compile-sensitive layouts,
dtypes, features, activation, scheduler, tile/cluster, output policy, and
overlap margin, but not the runtime GeGLU `linear_offset`.

### Block-scaled

#### High-level wrapper

**Dense mode:**

```python
from cudnn import grouped_gemm_glu_wrapper_sm100
from cuda.bindings import driver as cuda

stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream)

outputs = grouped_gemm_glu_wrapper_sm100(
    a_tensor=a,
    sfa_tensor=sfa,
    padded_offsets=padded_offsets,
    alpha_tensor=alpha,
    bias_tensor=bias,
    # Dense mode weights:
    b_tensor=b,
    sfb_tensor=sfb,
    # Common:
    norm_const_tensor=norm_const,
    prob_tensor=prob,
    acc_dtype=torch.float32,
    c_dtype=torch.bfloat16,
    d_dtype=torch.bfloat16,
    cd_major="n",
    mma_tiler_mn=(256, 256),
    cluster_shape_mn=(2, 1),
    sf_vec_size=32,
    vector_f32=False,
    m_aligned=256,
    discrete_col_sfd=False,
    act_func="swiglu",
    current_stream=stream,
)

# dictionary access:
c = outputs["c_tensor"]
d = outputs["d_tensor"]
d_col = outputs["d_col_tensor"]
amax = outputs["amax_tensor"]
sfd_row = outputs["sfd_row_tensor"]
sfd_col = outputs["sfd_col_tensor"]

# or tuple unpacking:
c, d, d_col, amax, sfd_row, sfd_col = outputs
```

**Discrete mode:**

```python
outputs = grouped_gemm_glu_wrapper_sm100(
    a_tensor=a,
    sfa_tensor=sfa,
    padded_offsets=padded_offsets,
    alpha_tensor=alpha,
    # Discrete mode weights:
    b_ptrs=b_ptrs,       # int64 tensor of per-expert B data pointers
    sfb_ptrs=sfb_ptrs,   # int64 tensor of per-expert SFB data pointers
    n=n_dim,             # B weight N dimension
    b_dtype=torch.uint8, # B weight data type
    b_major="k",         # B tensor major dimension
    # Common:
    norm_const_tensor=norm_const,
    prob_tensor=prob,
    act_func="geglu",    # GeGLU activation
    current_stream=stream,
)
```

`bias_tensor` must use the kernel layout expected by the fused bias path: shape `(N, L)` and stride `(1, N)`.

#### Class API

**Dense mode:**

```python
from cudnn import GroupedGemmGluSm100
from cuda.bindings import driver as cuda

stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream)

api = GroupedGemmGluSm100(
    sample_a=a,
    sample_c=c,
    sample_d=d,
    sample_sfa=sfa,
    sample_padded_offsets=padded_offsets,
    sample_alpha=alpha,
    sample_d_col=d_col,
    sample_bias=bias,
    # Dense mode:
    sample_b=b,
    sample_sfb=sfb,
    # Optional quantization outputs
    sample_sfd_row=sfd_row,
    sample_sfd_col=sfd_col,
    sample_amax=amax,
    sample_norm_const=norm_const,
    sample_prob=prob,
    # Configuration
    acc_dtype=torch.float32,
    mma_tiler_mn=(256, 256),
    cluster_shape_mn=(2, 1),
    sf_vec_size=32,
    act_func="swiglu",
)
assert api.check_support()
api.compile()
api.execute(
    a_tensor=a, c_tensor=c, d_tensor=d,
    sfa_tensor=sfa, padded_offsets=padded_offsets, alpha_tensor=alpha,
    b_tensor=b, sfb_tensor=sfb, bias_tensor=bias,
    d_col_tensor=d_col, sfd_row_tensor=sfd_row, sfd_col_tensor=sfd_col,
    amax_tensor=amax, norm_const_tensor=norm_const, prob_tensor=prob,
    current_stream=stream,
)
```

`sample_bias` and runtime `bias_tensor` must both use shape `(N, L)` and stride `(1, N)`.

**Discrete mode:**

```python
api = GroupedGemmGluSm100(
    sample_a=a,
    sample_c=c,
    sample_d=d,
    sample_sfa=sfa,
    sample_padded_offsets=padded_offsets,
    sample_alpha=alpha,
    sample_d_col=d_col,
    # Discrete mode:
    num_experts=num_experts,
    b_shape=(n, k),
    b_dtype=torch.uint8,
    # Configuration
    act_func="geglu",
    b_major="k",
)
assert api.check_support()
api.compile()
api.execute(
    a_tensor=a, c_tensor=c, d_tensor=d,
    sfa_tensor=sfa, padded_offsets=padded_offsets, alpha_tensor=alpha,
    b_ptrs=b_ptrs, sfb_ptrs=sfb_ptrs,
    d_col_tensor=d_col, prob_tensor=prob,
    current_stream=stream,
)
```

---

## Parameters

### Weight Mode

The weight mode is auto-detected from constructor arguments:
- **Dense**: Provide `sample_b` and `sample_sfb` (contiguous weight tensors)
- **Discrete**: Provide `num_experts`, `b_shape`, and `b_dtype` (per-expert pointer mode)

Providing both or neither raises `ValueError`.

### Input/Output Tensors

- **Input tensor A**: `a_tensor` / `sample_a`
  - Shape: `(valid_m, K, 1)`
  - Stride: `(K, 1, valid_m·K)` -- must be K-major
  - Dtype (`ab_dtype`): `{float4_e2m1fn_x2, uint8, float8_e4m3fn, float8_e5m2}`

- **Input tensor B** (dense mode): `b_tensor` / `sample_b`
  - Shape: `(N, K, L)` where `L = num_groups`
  - Stride: `(K, 1, N·K)` -- must be K-major
  - Dtype: Must match A

- **Input B pointers** (discrete mode): `b_ptrs`
  - Shape: `(num_experts,)` -- 1-D int64 device tensor of per-expert B data pointers
  - Build via: `torch.tensor([b.data_ptr() for b in experts], dtype=torch.int64, device="cuda")`

- **Output tensor C**: returned in wrapper dict or `c_tensor` in class
  - Shape: `(valid_m, N, 1)`
  - Stride: `(N, 1, valid_m·N)` -- **must be N-major**
  - Dtype (`c_dtype`): `{float16, bfloat16}` for FP4 inputs; `{float32, float16, bfloat16, float8_e4m3fn, float8_e5m2, float4_e2m1fn_x2}` otherwise

- **Output tensor D**: `d_tensor` / `sample_d`
  - Shape: `(valid_m, N/2, 1)`
  - Stride: `(N/2, 1, valid_m·N/2)` -- **must be N-major**
  - Dtype (`d_dtype`): `{bfloat16, float32}` for FP4 inputs; `{float16, bfloat16, float8_e4m3fn, float8_e5m2, float4_e2m1fn_x2}` otherwise

- **Output tensor D_col**: `d_col_tensor` / `sample_d_col`
  - Shape: `(valid_m, N/2, 1)` -- must match D dtype and stride

- **Scale factor tensors**: Same as contiguous swiglu (SFA, SFB, SFD_row, SFD_col)
  - SFB (discrete mode): use `sfb_ptrs` (1-D int64 device tensor of per-expert SFB pointers)

- **Group offsets**: `padded_offsets` -- shape `(L,)`, dtype `int32`

- **Scaling tensors**: `alpha` shape `(L,)`, `prob` shape `(valid_m, 1, 1)`, `amax` shape `(L, 1)`, `norm_const` shape `(1,)`

### Common Parameters

- `acc_dtype`: Must be `torch.float32`
- `mma_tiler_mn`: Kernel tile size `(TILE_M, TILE_N)`. Default: `(256, 256)`
  - `TILE_M ∈ {128, 256}`
  - `TILE_N = 256`
- `cluster_shape_mn`: Thread Block cluster shape. Default: `(2, 1)` when `TILE_M=256`, `(1, 1)` otherwise
- `sf_vec_size`: Scale factor vector size. `{16, 32}`. Default: `16`
- `vector_f32`: Enable packed f32 operations. Default: `False`
- `m_aligned`: Must be `256` (FIX_PAD_SIZE). Default: `256`
- `discrete_col_sfd`: Generate discrete col-major scale factors. Default: `False`
- `act_func`: Activation function. `"swiglu"` (default), `"geglu"`, or block-scaled `"situglu"`
- `linear_offset`: Offset added to the clamped GeGLU up branch. Default: `1.0` for GeGLU
- `geglu_alpha`: GeGLU sigmoid input scale. Default: `1.702`
- `glu_clamp_max`: GeGLU upper bound for gate and up. Default: `7.0`
- `glu_clamp_min`: GeGLU lower bound for up only. Default: `-7.0`
- The activation alpha and clamp controls above are supported by the block-scaled forward backend on Blackwell and Rubin; the BF16 contract retains its fixed values
- `situ_beta1`: Positive finite gate tanh scale for SiTU-GLU. Default: `4.0`
- `situ_beta2`: Positive finite up-branch tanh scale for SiTU-GLU. Default: `25.0`
- `b_major` (discrete only): B tensor major dimension. `"k"` (default) or `"n"`. Must be `"k"` for FP4.

### Wrapper-specific Parameters

- `c_dtype`: Intermediate C tensor data type. Default: `torch.bfloat16`
- `d_dtype`: Output D tensor data type. Default: `torch.bfloat16`
- `cd_major`: Must be `"n"`. Default: `"n"`
- `n` (discrete only): B weight N dimension (full N before GLU split)
- `b_dtype` (discrete only): B weight data type

### Wrapper Return Values

Returns a `TupleDict` (dictionary + tuple unpacking):
- `c_tensor`: Intermediate GEMM result
- `d_tensor`: Row-quantized GLU output
- `d_col_tensor`: Column-quantized GLU output
- `amax_tensor`: Per-group amax (when `d_dtype` is bf16/fp16)
- `sfd_row_tensor`: Row scale factors (when SFD enabled)
- `sfd_col_tensor`: Column scale factors (when SFD enabled)

---

## Support Surface and Constraints

### Layouts and Strides

- `A` must be **K-major**
- `B` must be **K-major** (dense mode). For discrete mode: K-major or N-major (K-major required for FP4)
- `C`, `D`, `D_col` must be **N-major**
- All tensors must be **16-byte aligned** along the contiguous dimension

### Data Types

| Format | ab_dtype | sf_dtype | sf_vec_size | d_dtype |
|--------|----------|----------|-------------|---------|
| **MXFP8** | `float8_e4m3fn` or `float8_e5m2` | `float8_e8m0fnu` | 32 | `{float16, bfloat16, float8_e4m3fn, float8_e5m2}` |
| **NVF4** | `float4_e2m1fn_x2` or `uint8` | \{`float8_e4m3fn`, `float8_e8m0fnu`\} | \{16, 32\} | `{bfloat16, float32}` |

### Additional Type Constraints

- `A` and `B` must have the same dtype
- Scale factor tensors (SFA, SFB, SFD_row, SFD_col) must have the same dtype
- `D` and `D_col` must have the same dtype
- `bias` must be one of `{float16, bfloat16, float32}`
- `bias` must have shape `(N, L)` and stride `(1, N)`
- For non-bias paths, FP4 `ab_dtype` with `sf_vec_size=16` and `d_dtype=float32` is not supported
- FP4 `ab_dtype` requires `c_dtype` in `{float16, bfloat16}`

### Shapes and Divisibility

- `N` must be divisible by 64 (two consecutive 32-column blocks for GLU pairing)
- Expert count must be `<= 1024`
- Each group's M dimension is aligned to `m_aligned` (256)
- All supported kernel configurations require `mma_tiler_mn[1] == 256`
- `use_single_group_runtime_offsets=True` is supported only by the block-scaled
  kernel with exactly one expert. In this mode the kernel derives
  `padded_offsets[0]` from runtime `A.shape[0]` and does not load its value from
  device memory; the argument must still be an int32 tensor with shape `(1,)`.

### Environment

- Requires CUDA with **SM100/SM103** (Blackwell), or **SM107** (Rubin) for the block-scaled forward backend
- Rubin MXFP8 uses matching FP8 A/B operands, E8M0 scale factors, and `sf_vec_size=32`

---

## Usage Examples

For usage examples, see test cases in `test/python/fe_api/grouped_gemm/test_grouped_gemm_glu.py` (dense mode, unified API) and `test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_swiglu.py` (discrete mode).