Skip to content

vllm.model_executor.kernels.linear.mxfp8.rocm_block32_gemm

MXFP8 GEMM on 32x32 block-scaled weights for gfx950 (tl.dot_scaled).

y = x @ w.T with an MXFP8 activation (e4m3 values, one E8M0 scale per row and 32 K, [M, K / 32]) and an e4m3 weight whose E8M0 scales stay in the checkpoint's 32x32 blocks, [N / 32, K / 32], instead of being expanded to every row. Each weight scale byte is then read once per 32 output rows, and small-M shapes can use the packed kernel below.

Two kernels, picked per shape from a table tuned on MI355X:

  • a tiled kernel for larger M;
  • a packed kernel for small M, which fills the MFMA's rows with K panels instead of tokens and keeps only the matching-panel (block-diagonal) products, so a handful of tokens still streams the weight at full width.

Either can split K. The partials are normally summed in the same launch by the last program of each output tile to finish, in split order, so the result does not depend on scheduling. Otherwise a second launch reduces them.

The packed kernel and the in-launch split-K reduction on one XCD (_split_tile, _sum_splits, _split_counters, _k_partition) are adapted from the group32 GEMM in ROCm/aiter#5750.

Functions:

_default_config(M, N, K)

Untuned shapes: the tiers most tuned shapes settle on.

Source code in vllm/model_executor/kernels/linear/mxfp8/rocm_block32_gemm.py
def _default_config(M: int, N: int, K: int) -> _Config:
    """Untuned shapes: the tiers most tuned shapes settle on."""
    if M <= 8:
        return _packed(8, 16, 256, 2, 2, 3, 16, 0)
    if M <= 32:
        return _packed(32, 32, 256, 1, 2, 3, 16, 0)
    if M <= 256:
        return _tiled(64, 64, 256, 1, True, 4, 2, 16, 0)
    return _tiled(128, 128, 128, 1, True, 4, 2, 16, 0)

_k_partition(K, splits, step)

K per split, rounded up to whole steps, and the number of non-empty splits.

Source code in vllm/model_executor/kernels/linear/mxfp8/rocm_block32_gemm.py
def _k_partition(K: int, splits: int, step: int) -> tuple[int, int]:
    """K per split, rounded up to whole steps, and the number of non-empty
    splits."""
    size = triton.cdiv(triton.cdiv(K, splits), step) * step
    return size, triton.cdiv(K, size)

_split_counters(device)

Zeroed arrival counters, one set per stream so concurrent launches never share a tile's counter. Each launch leaves the counters it used at zero.

None when a stream first needs them inside a CUDA graph capture: that allocation would come from the graph's pool, where it can take the address of an intermediate that every replay overwrites. vLLM warms up on its capture streams first, so this only affects captures without a warmup.

Source code in vllm/model_executor/kernels/linear/mxfp8/rocm_block32_gemm.py
def _split_counters(device: torch.device) -> torch.Tensor | None:
    """Zeroed arrival counters, one set per stream so concurrent launches never
    share a tile's counter. Each launch leaves the counters it used at zero.

    None when a stream first needs them inside a CUDA graph capture: that
    allocation would come from the graph's pool, where it can take the address
    of an intermediate that every replay overwrites. vLLM warms up on its
    capture streams first, so this only affects captures without a warmup.
    """
    key = (device.index, torch.accelerator.current_stream(device).stream_id)
    counters = _split_counters_by_stream.get(key)
    if counters is None:
        if torch.cuda.is_current_stream_capturing():
            return None
        counters = torch.zeros(_SPLIT_COUNTERS, dtype=torch.int32, device=device)
        _split_counters_by_stream[key] = counters
    return counters

_split_tile(N, BLOCK_N, SPLITS)

(pid_m, pid_n, tile, split) of a 1D grid that keeps a tile's splits on one XCD, so they meet in that XCD's L2.

Source code in vllm/model_executor/kernels/linear/mxfp8/rocm_block32_gemm.py
@triton.jit
def _split_tile(N: tl.constexpr, BLOCK_N: tl.constexpr, SPLITS: tl.constexpr):
    """(pid_m, pid_n, tile, split) of a 1D grid that keeps a tile's splits on
    one XCD, so they meet in that XCD's L2."""
    pid = tl.program_id(0)
    tile = pid // (8 * SPLITS) * 8 + pid % 8
    grid_n: tl.constexpr = (N + BLOCK_N - 1) // BLOCK_N
    return tile // grid_n, tile % grid_n, tile, pid // 8 % SPLITS

_sum_splits(acc, out_ptrs, out_mask, slot_ptr, count_ptr, split, SPLITS)

Publish this split's partial; the tile's last arrival sums all of them in split order and re-arms the counter.

Source code in vllm/model_executor/kernels/linear/mxfp8/rocm_block32_gemm.py
@triton.jit
def _sum_splits(
    acc, out_ptrs, out_mask, slot_ptr, count_ptr, split, SPLITS: tl.constexpr
):
    """Publish this split's partial; the tile's last arrival sums all of them
    in split order and re-arms the counter."""
    BM: tl.constexpr = acc.shape[0]
    BN: tl.constexpr = acc.shape[1]
    local = tl.arange(0, BM)[:, None] * BN + tl.arange(0, BN)[None, :]
    tl.store(slot_ptr + split * BM * BN + local, acc)
    # The partial must be in L2 before the arrival is counted. The splits
    # share an XCD, so no device-scope release (an L2 writeback) is needed.
    tl.inline_asm_elementwise(
        "s_waitcnt vmcnt(0)", "=v,v", [split], dtype=tl.int32, is_pure=False, pack=1
    )
    tl.debug_barrier()
    if tl.atomic_add(count_ptr, 1, sem="acq_rel", scope="cta") == SPLITS - 1:
        total = tl.zeros((BM, BN), dtype=tl.float32)
        for s in tl.static_range(SPLITS):
            total += tl.load(slot_ptr + s * BM * BN + local, cache_modifier=".cv")
        tl.store(out_ptrs, total.to(out_ptrs.dtype.element_ty), mask=out_mask)
        tl.store(count_ptr, 0)

rocm_mxfp8_block32_gemm(x, x_scale, weight, weight_scale, out_dtype)

x @ weight.T for MXFP8 x and a 32x32 block-scaled MXFP8 weight.

Parameters:

  • x

    (Tensor) –

    [M, K] e4m3 activation, contiguous.

  • x_scale

    (Tensor) –

    [M, K / 32] E8M0 (uint8) activation scales.

  • weight

    (Tensor) –

    [N, K] e4m3 weight, contiguous.

  • weight_scale

    (Tensor) –

    [ceil(N / 32), K / 32] E8M0 (uint8) weight block scales.

  • out_dtype

    (dtype) –

    Output dtype.

Returns:

  • Tensor –

    The [M, N] product.

Source code in vllm/model_executor/kernels/linear/mxfp8/rocm_block32_gemm.py
def rocm_mxfp8_block32_gemm(
    x: torch.Tensor,
    x_scale: torch.Tensor,
    weight: torch.Tensor,
    weight_scale: torch.Tensor,
    out_dtype: torch.dtype,
) -> torch.Tensor:
    """``x @ weight.T`` for MXFP8 ``x`` and a 32x32 block-scaled MXFP8 weight.

    Args:
        x: [M, K] e4m3 activation, contiguous.
        x_scale: [M, K / 32] E8M0 (uint8) activation scales.
        weight: [N, K] e4m3 weight, contiguous.
        weight_scale: [ceil(N / 32), K / 32] E8M0 (uint8) weight block scales.
        out_dtype: Output dtype.

    Returns:
        The [M, N] product.

    """
    return torch.ops.vllm.rocm_mxfp8_block32_gemm(
        x, x_scale, weight, weight_scale, out_dtype
    )