Skip to content
KernelIndex
Search⌘K

submission 745087

HorizonLiang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 1304 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-745087?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, int32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD Instinct MI355X
38.3µs
#105 of 766
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3bcfcb2de30cc10ca2b9b49ef7cd55971cee7336194599ad396882c60285b419
license declaredunknown
license concludedunknown
authorsHorizonLiang
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fused-epilogueos.environ.setdefault("OPTIMIZE_EPILOGUE", "1")
num-warps = 4TRITON_REDUCE_NUM_WARPS = 4
stages = 1TRITON_REDUCE_NUM_STAGES = 1

Kernel source

submission.py1304 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""Submission version 033."""

from __future__ import annotations

import os
import sys

os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("OPTIMIZE_EPILOGUE", "1")

import torch
import triton
import triton.language as tl
from task import input_t, output_t

from aiter import (
    dtypes as aiter_dtypes,
    dynamic_per_tensor_quant,
    get_mla_metadata_info_v1,
    get_mla_metadata_v1,
    mla_decode_stage1_asm_fwd,
)
from aiter.mla import mla_decode_fwd
from aiter.ops.triton.utils._triton.pid_preprocessing import remap_xcd


PAGE_SIZE = 1
SMALLKV_PAGE_SIZE = 2
LONGKV_PAGE_SIZE = 8
DEFAULT_KV_GRANULARITY = 16
DEFAULT_SPLIT_PER_BATCH = 32
FAST_MODE = True
INTRA_BATCH_MODE = False
FP8_DTYPE = aiter_dtypes.fp8
ZERO_REDUCE_PARTIAL_SLOTS = 1

NUM_HEADS = 16
V_HEAD_DIM = 512
TRITON_REDUCE_MAX_SPLITS = 6
TRITON_REDUCE_BLOCK_H = 8
TRITON_REDUCE_BLOCK_D = 256
TRITON_REDUCE_NUM_WARPS = 4
TRITON_REDUCE_NUM_STAGES = 1
TRITON_NUM_XCDS = 8

_META_CACHE: dict[tuple, dict[str, torch.Tensor]] = {}
_PAGED_META_CACHE: dict[tuple, dict[str, torch.Tensor]] = {}
_KV_INDEX_CACHE: dict[tuple, torch.Tensor] = {}
_KV_MISC_CACHE: dict[tuple, torch.Tensor] = {}
_OUTPUT_CACHE: dict[tuple, torch.Tensor] = {}
_PARTIAL_CACHE: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
_Q_QUANT_CACHE: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
_SEQ_INDEX_CACHE: dict[tuple, torch.Tensor] = {}
_DENSE_FASTPATH_PROBE_SEEN: set[tuple[int, int, str]] = set()


@triton.jit
def _mla_reduce_q1_longkv_kernel(
    partial_output_ptr,
    partial_lse_ptr,
    reduce_indptr_ptr,
    reduce_q_indices_ptr,
    reduce_partial_map_ptr,
    output_ptr,
    stride_po0,
    stride_po1,
    stride_po2,
    stride_pl0,
    stride_pl1,
    stride_o0,
    stride_o1,
    stride_o2,
    grid_tiles,
    tiles_per_group,
    BLOCK_H: tl.constexpr,
    BLOCK_D: tl.constexpr,
    MAX_SPLITS: tl.constexpr,
    NUM_XCDS: tl.constexpr,
):
    pid = tl.program_id(0)
    pid = remap_xcd(pid, grid_tiles, NUM_XCDS=NUM_XCDS)

    pid_group = pid // tiles_per_group
    pid_tile = pid % tiles_per_group
    num_pid_d = tl.cdiv(512, BLOCK_D)
    pid_h = pid_tile // num_pid_d
    pid_d = pid_tile % num_pid_d

    offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
    offs_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D)
    mask_h = offs_h < 16
    mask_d = offs_d < 512

    reduce_start = tl.load(reduce_indptr_ptr + pid_group)
    reduce_end = tl.load(reduce_indptr_ptr + pid_group + 1)
    split_count = reduce_end - reduce_start
    q_row = tl.load(reduce_q_indices_ptr + pid_group)

    max_lse = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32)
    sum_e_lse = tl.zeros((BLOCK_H,), dtype=tl.float32)
    acc = tl.zeros((BLOCK_H, BLOCK_D), dtype=tl.float32)

    for split_idx in tl.static_range(MAX_SPLITS):
        active = split_idx < split_count
        partial_idx = tl.load(
            reduce_partial_map_ptr + reduce_start + split_idx,
            mask=active,
            other=0,
        )
        lse = tl.load(
            partial_lse_ptr + partial_idx * stride_pl0 + offs_h * stride_pl1,
            mask=active & mask_h,
            other=-float("inf"),
        )
        new_max_lse = tl.maximum(max_lse, lse)
        old_scale = tl.exp(max_lse - new_max_lse)
        new_scale = tl.exp(lse - new_max_lse)

        partial = tl.load(
            partial_output_ptr
            + partial_idx * stride_po0
            + offs_h[:, None] * stride_po1
            + offs_d[None, :] * stride_po2,
            mask=(active & mask_h)[:, None] & mask_d[None, :],
            other=0.0,
        )
        acc = acc * old_scale[:, None] + partial * new_scale[:, None]
        sum_e_lse = sum_e_lse * old_scale + new_scale
        max_lse = new_max_lse

    denom = tl.where(sum_e_lse > 0.0, sum_e_lse, 1.0)
    out = acc / denom[:, None]
    tl.store(
        output_ptr
        + q_row * stride_o0
        + offs_h[:, None] * stride_o1
        + offs_d[None, :] * stride_o2,
        out.to(output_ptr.type.element_ty),
        mask=mask_h[:, None] & mask_d[None, :],
    )


@triton.jit
def _mla_reduce_q1_dense_exact_kernel(
    partial_output_ptr,
    partial_lse_ptr,
    output_ptr,
    stride_po0,
    stride_po1,
    stride_po2,
    stride_pl0,
    stride_pl1,
    stride_o0,
    stride_o1,
    stride_o2,
    grid_tiles,
    tiles_per_group,
    BLOCK_H: tl.constexpr,
    BLOCK_D: tl.constexpr,
    EXACT_SPLITS: tl.constexpr,
    NUM_XCDS: tl.constexpr,
):
    pid = tl.program_id(0)
    pid = remap_xcd(pid, grid_tiles, NUM_XCDS=NUM_XCDS)

    pid_group = pid // tiles_per_group
    pid_tile = pid % tiles_per_group
    num_pid_d = tl.cdiv(512, BLOCK_D)
    pid_h = pid_tile // num_pid_d
    pid_d = pid_tile % num_pid_d

    offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
    offs_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D)
    mask_h = offs_h < 16
    mask_d = offs_d < 512

    max_lse = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32)
    sum_e_lse = tl.zeros((BLOCK_H,), dtype=tl.float32)
    acc = tl.zeros((BLOCK_H, BLOCK_D), dtype=tl.float32)

    base_partial_idx = pid_group * EXACT_SPLITS
    q_row = pid_group

    for split_idx in tl.static_range(EXACT_SPLITS):
        partial_idx = base_partial_idx + split_idx
        lse = tl.load(
            partial_lse_ptr + partial_idx * stride_pl0 + offs_h * stride_pl1,
            mask=mask_h,
            other=-float("inf"),
        )
        new_max_lse = tl.maximum(max_lse, lse)
        old_scale = tl.exp(max_lse - new_max_lse)
        new_scale = tl.exp(lse - new_max_lse)

        partial = tl.load(
            partial_output_ptr
            + partial_idx * stride_po0
            + offs_h[:, None] * stride_po1
            + offs_d[None, :] * stride_po2,
            mask=mask_h[:, None] & mask_d[None, :],
            other=0.0,
        )
        acc = acc * old_scale[:, None] + partial * new_scale[:, None]
        sum_e_lse = sum_e_lse * old_scale + new_scale
        max_lse = new_max_lse

    denom = tl.where(sum_e_lse > 0.0, sum_e_lse, 1.0)
    out = acc / denom[:, None]
    tl.store(
        output_ptr
        + q_row * stride_o0
        + offs_h[:, None] * stride_o1
        + offs_d[None, :] * stride_o2,
        out.to(output_ptr.type.element_ty),
        mask=mask_h[:, None] & mask_d[None, :],
    )


@triton.jit
def _mla_reduce_q1_dense_upto_kernel(
    partial_output_ptr,
    partial_lse_ptr,
    reduce_indptr_ptr,
    output_ptr,
    stride_po0,
    stride_po1,
    stride_po2,
    stride_pl0,
    stride_pl1,
    stride_o0,
    stride_o1,
    stride_o2,
    grid_tiles,
    tiles_per_group,
    BLOCK_H: tl.constexpr,
    BLOCK_D: tl.constexpr,
    MAX_SPLITS: tl.constexpr,
    NUM_XCDS: tl.constexpr,
):
    pid = tl.program_id(0)
    pid = remap_xcd(pid, grid_tiles, NUM_XCDS=NUM_XCDS)

    pid_group = pid // tiles_per_group
    pid_tile = pid % tiles_per_group
    num_pid_d = tl.cdiv(512, BLOCK_D)
    pid_h = pid_tile // num_pid_d
    pid_d = pid_tile % num_pid_d

    offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
    offs_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D)
    mask_h = offs_h < 16
    mask_d = offs_d < 512

    reduce_start = tl.load(reduce_indptr_ptr + pid_group)
    reduce_end = tl.load(reduce_indptr_ptr + pid_group + 1)
    split_count = reduce_end - reduce_start
    q_row = pid_group

    max_lse = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32)
    sum_e_lse = tl.zeros((BLOCK_H,), dtype=tl.float32)
    acc = tl.zeros((BLOCK_H, BLOCK_D), dtype=tl.float32)

    for split_idx in tl.static_range(MAX_SPLITS):
        active = split_idx < split_count
        partial_idx = reduce_start + split_idx
        lse = tl.load(
            partial_lse_ptr + partial_idx * stride_pl0 + offs_h * stride_pl1,
            mask=active & mask_h,
            other=-float("inf"),
        )
        new_max_lse = tl.maximum(max_lse, lse)
        old_scale = tl.exp(max_lse - new_max_lse)
        new_scale = tl.exp(lse - new_max_lse)

        partial = tl.load(
            partial_output_ptr
            + partial_idx * stride_po0
            + offs_h[:, None] * stride_po1
            + offs_d[None, :] * stride_po2,
            mask=(active & mask_h)[:, None] & mask_d[None, :],
            other=0.0,
        )
        acc = acc * old_scale[:, None] + partial * new_scale[:, None]
        sum_e_lse = sum_e_lse * old_scale + new_scale
        max_lse = new_max_lse

    denom = tl.where(sum_e_lse > 0.0, sum_e_lse, 1.0)
    out = acc / denom[:, None]
    tl.store(
        output_ptr
        + q_row * stride_o0
        + offs_h[:, None] * stride_o1
        + offs_d[None, :] * stride_o2,
        out.to(output_ptr.type.element_ty),
        mask=mask_h[:, None] & mask_d[None, :],
    )


def _device_key(device: torch.device) -> tuple[str, int | None]:
    return device.type, device.index


def _get_cached_seq_indices(length: int, device: torch.device) -> torch.Tensor:
    key = (_device_key(device), length)
    cached = _SEQ_INDEX_CACHE.get(key)
    if cached is None:
        cached = torch.arange(length, dtype=torch.int32, device=device)
        _SEQ_INDEX_CACHE[key] = cached
    return cached


def _probe_dense_fastpath_once(batch_size: int, kv_seq_len: int, message: str) -> None:
    key = (batch_size, kv_seq_len, message)
    if key in _DENSE_FASTPATH_PROBE_SEEN:
        return
    print(f"[probe] dense_tail_reduce batch={batch_size} kv={kv_seq_len} {message}", file=sys.stderr, flush=True)
    _DENSE_FASTPATH_PROBE_SEEN.add(key)


def _use_bf16_q(kv_seq_len: int) -> bool:
    return kv_seq_len == 1024


def _choose_num_kv_splits(batch_size: int, kv_seq_len: int) -> int:
    split_table = {
        (4, 8192): 24,
        (32, 8192): 6,
        (64, 1024): 16,
        (64, 8192): 6,
        (256, 1024): 1,
        (256, 8192): 6,
    }
    return split_table.get((batch_size, kv_seq_len), DEFAULT_SPLIT_PER_BATCH)


def _choose_kv_granularity(batch_size: int, kv_seq_len: int) -> int:
    if batch_size == 4 and kv_seq_len == 8192:
        return 32
    if batch_size == 64 and kv_seq_len == 1024:
        return 8
    if batch_size == 256 and kv_seq_len == 8192:
        return 8
    if kv_seq_len == 8192 and batch_size >= 32:
        return 16
    return DEFAULT_KV_GRANULARITY


def _use_paged_fast_mode(batch_size: int, kv_seq_len: int, page_size: int) -> bool:
    return FAST_MODE


def _get_alt_page_size(batch_size: int, kv_seq_len: int) -> int | None:
    if kv_seq_len == 1024 and batch_size in (4, 32, 64, 256):
        return SMALLKV_PAGE_SIZE
    if kv_seq_len == 8192 and batch_size >= 4:
        return LONGKV_PAGE_SIZE
    return None


def _get_smallkv_dense_exact_splits(
    batch_size: int,
    kv_seq_len: int,
    page_size: int | None,
    num_kv_splits: int,
) -> int | None:
    if kv_seq_len != 1024 or page_size != SMALLKV_PAGE_SIZE:
        return None
    exact_split_table = {
        (4, 32): 16,
        (32, 32): 7,
        (64, 16): 4,
    }
    return exact_split_table.get((batch_size, num_kv_splits))


def _use_stage1_direct(batch_size: int, kv_seq_len: int, num_kv_splits: int) -> bool:
    return batch_size == 256 and kv_seq_len == 1024 and num_kv_splits == 1


def _use_paged_stage1_only(
    batch_size: int,
    kv_seq_len: int,
    page_size: int | None,
    num_kv_splits: int,
) -> bool:
    if batch_size != 256:
        return False
    if kv_seq_len == 1024:
        return page_size == SMALLKV_PAGE_SIZE and num_kv_splits == 1
    return kv_seq_len == 8192 and page_size == LONGKV_PAGE_SIZE


def _use_triton_longkv_reduce(
    batch_size: int,
    kv_seq_len: int,
    page_size: int | None,
    num_kv_splits: int,
) -> bool:
    return (
        batch_size in (32, 64)
        and kv_seq_len == 8192
        and page_size == LONGKV_PAGE_SIZE
        and num_kv_splits == TRITON_REDUCE_MAX_SPLITS
    )


def _choose_stage1_partial_slots(
    batch_size: int,
    kv_seq_len: int,
    page_size: int,
    partial_tiles: int,
) -> int:
    if batch_size == 256 and kv_seq_len == 8192 and page_size == LONGKV_PAGE_SIZE:
        return ZERO_REDUCE_PARTIAL_SLOTS
    return max(partial_tiles, 1)


def _get_used_group_count(indptr: torch.Tensor) -> int:
    if indptr.numel() <= 1:
        return 0
    diffs = indptr[1:] - indptr[:-1]
    nonzero = torch.nonzero(diffs, as_tuple=False)
    if nonzero.numel() == 0:
        return 0
    return int(nonzero[-1, 0].item()) + 1


def _get_reduce_compact_metadata(metadata: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
    cached = metadata.get("_reduce_compact_metadata")
    if cached is not None:
        return cached

    reduce_indptr = metadata["reduce_indptr"]
    used_reduce_groups = _get_used_group_count(reduce_indptr)
    compact_reduce_indptr = reduce_indptr[: used_reduce_groups + 1]
    num_partial = (
        int(compact_reduce_indptr[-1].item())
        if compact_reduce_indptr.numel() != 0
        else 0
    )
    reduce_final_map = metadata["reduce_final_map"][:used_reduce_groups]
    reduce_q_indices = reduce_final_map[:, 0].contiguous()

    compact = dict(metadata)
    compact.update(
        {
            "reduce_indptr": compact_reduce_indptr,
            "reduce_final_map": reduce_final_map,
            "reduce_q_indices": reduce_q_indices,
            "reduce_partial_map": metadata["reduce_partial_map"][: max(num_partial, 1)],
        }
    )
    metadata["_reduce_compact_metadata"] = compact
    return compact


def _get_dense_reduce_layout(metadata: dict[str, torch.Tensor]) -> dict[str, int | bool]:
    cached = metadata.get("_dense_reduce_layout")
    if cached is not None:
        return cached

    reduce_indptr = metadata["reduce_indptr"]
    reduce_q_indices = metadata["reduce_q_indices"]
    reduce_partial_map = metadata["reduce_partial_map"]
    num_groups = int(reduce_q_indices.numel())
    num_partial = int(reduce_partial_map.numel())

    q_identity = torch.equal(
        reduce_q_indices, _get_cached_seq_indices(num_groups, reduce_q_indices.device)
    )
    partial_identity = torch.equal(
        reduce_partial_map,
        _get_cached_seq_indices(num_partial, reduce_partial_map.device),
    )

    if num_groups > 0:
        split_counts = reduce_indptr[1:] - reduce_indptr[:-1]
        split_min = int(split_counts.min().item())
        split_max = int(split_counts.max().item())
    else:
        split_min = 0
        split_max = 0

    cached = {
        "num_groups": num_groups,
        "num_partial": num_partial,
        "q_identity": q_identity,
        "partial_identity": partial_identity,
        "split_min": split_min,
        "split_max": split_max,
    }
    metadata["_dense_reduce_layout"] = cached
    return cached


def _run_triton_longkv_reduce_q1(
    partial_output: torch.Tensor,
    partial_lse: torch.Tensor,
    reduce_indptr: torch.Tensor,
    reduce_q_indices: torch.Tensor,
    reduce_partial_map: torch.Tensor,
    output: torch.Tensor,
) -> None:
    num_groups = int(reduce_q_indices.numel())
    if num_groups == 0:
        return

    partial_output_view = partial_output[:, 0]
    partial_lse_view = partial_lse[:, 0, :, 0]
    tiles_per_group = triton.cdiv(NUM_HEADS, TRITON_REDUCE_BLOCK_H) * triton.cdiv(
        V_HEAD_DIM, TRITON_REDUCE_BLOCK_D
    )
    grid_tiles = num_groups * tiles_per_group
    _mla_reduce_q1_longkv_kernel[(grid_tiles,)](
        partial_output_view,
        partial_lse_view,
        reduce_indptr,
        reduce_q_indices,
        reduce_partial_map,
        output,
        partial_output_view.stride(0),
        partial_output_view.stride(1),
        partial_output_view.stride(2),
        partial_lse_view.stride(0),
        partial_lse_view.stride(1),
        output.stride(0),
        output.stride(1),
        output.stride(2),
        grid_tiles,
        tiles_per_group,
        BLOCK_H=TRITON_REDUCE_BLOCK_H,
        BLOCK_D=TRITON_REDUCE_BLOCK_D,
        MAX_SPLITS=TRITON_REDUCE_MAX_SPLITS,
        NUM_XCDS=TRITON_NUM_XCDS,
        num_warps=TRITON_REDUCE_NUM_WARPS,
        num_stages=TRITON_REDUCE_NUM_STAGES,
    )


def _run_dense_longkv_reduce_q1_exact(
    partial_output: torch.Tensor,
    partial_lse: torch.Tensor,
    output: torch.Tensor,
    exact_splits: int,
) -> None:
    num_groups = output.shape[0]
    if num_groups == 0:
        return

    partial_output_view = partial_output[:, 0]
    partial_lse_view = partial_lse[:, 0, :, 0]
    tiles_per_group = triton.cdiv(NUM_HEADS, TRITON_REDUCE_BLOCK_H) * triton.cdiv(
        V_HEAD_DIM, TRITON_REDUCE_BLOCK_D
    )
    grid_tiles = num_groups * tiles_per_group
    _mla_reduce_q1_dense_exact_kernel[(grid_tiles,)](
        partial_output_view,
        partial_lse_view,
        output,
        partial_output_view.stride(0),
        partial_output_view.stride(1),
        partial_output_view.stride(2),
        partial_lse_view.stride(0),
        partial_lse_view.stride(1),
        output.stride(0),
        output.stride(1),
        output.stride(2),
        grid_tiles,
        tiles_per_group,
        BLOCK_H=TRITON_REDUCE_BLOCK_H,
        BLOCK_D=TRITON_REDUCE_BLOCK_D,
        EXACT_SPLITS=exact_splits,
        NUM_XCDS=TRITON_NUM_XCDS,
        num_warps=TRITON_REDUCE_NUM_WARPS,
        num_stages=TRITON_REDUCE_NUM_STAGES,
    )


def _run_dense_longkv_reduce_q1_upto(
    partial_output: torch.Tensor,
    partial_lse: torch.Tensor,
    reduce_indptr: torch.Tensor,
    output: torch.Tensor,
    max_splits: int,
) -> None:
    num_groups = output.shape[0]
    if num_groups == 0:
        return

    partial_output_view = partial_output[:, 0]
    partial_lse_view = partial_lse[:, 0, :, 0]
    tiles_per_group = triton.cdiv(NUM_HEADS, TRITON_REDUCE_BLOCK_H) * triton.cdiv(
        V_HEAD_DIM, TRITON_REDUCE_BLOCK_D
    )
    grid_tiles = num_groups * tiles_per_group
    _mla_reduce_q1_dense_upto_kernel[(grid_tiles,)](
        partial_output_view,
        partial_lse_view,
        reduce_indptr,
        output,
        partial_output_view.stride(0),
        partial_output_view.stride(1),
        partial_output_view.stride(2),
        partial_lse_view.stride(0),
        partial_lse_view.stride(1),
        output.stride(0),
        output.stride(1),
        output.stride(2),
        grid_tiles,
        tiles_per_group,
        BLOCK_H=TRITON_REDUCE_BLOCK_H,
        BLOCK_D=TRITON_REDUCE_BLOCK_D,
        MAX_SPLITS=max_splits,
        NUM_XCDS=TRITON_NUM_XCDS,
        num_warps=TRITON_REDUCE_NUM_WARPS,
        num_stages=TRITON_REDUCE_NUM_STAGES,
    )


def _get_cached_kv_indices(total_kv: int, device: torch.device) -> torch.Tensor:
    key = (_device_key(device), total_kv)
    cached = _KV_INDEX_CACHE.get(key)
    if cached is None:
        cached = torch.arange(total_kv, dtype=torch.int32, device=device)
        _KV_INDEX_CACHE[key] = cached
    return cached


def _get_kv_last_page_len(
    batch_size: int, kv_seq_len: int, device: torch.device
) -> torch.Tensor:
    key = (_device_key(device), batch_size, kv_seq_len)
    cached = _KV_MISC_CACHE.get(key)
    if cached is None:
        cached = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
        _KV_MISC_CACHE[key] = cached
    return cached


def _get_cached_output(
    total_q: int,
    num_heads: int,
    v_head_dim: int,
    device: torch.device,
    batch_size: int,
    kv_seq_len: int,
) -> torch.Tensor:
    key = (_device_key(device), total_q, num_heads, v_head_dim, batch_size, kv_seq_len)
    cached = _OUTPUT_CACHE.get(key)
    if cached is None:
        cached = torch.empty(
            (total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=device
        )
        _OUTPUT_CACHE[key] = cached
    return cached


def _get_cached_partials(
    partial_tiles: int,
    q_seq_len: int,
    num_heads: int,
    v_head_dim: int,
    device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
    slots = max(partial_tiles, 1)
    key = (_device_key(device), slots, q_seq_len, num_heads, v_head_dim)
    cached = _PARTIAL_CACHE.get(key)
    if cached is None:
        cached = (
            torch.empty(
                (slots * q_seq_len, 1, num_heads, v_head_dim),
                dtype=torch.float32,
                device=device,
            ),
            torch.empty(
                (slots * q_seq_len, 1, num_heads, 1),
                dtype=torch.float32,
                device=device,
            ),
        )
        _PARTIAL_CACHE[key] = cached
    return cached


def _get_cached_q_quant(
    total_q: int, num_heads: int, qk_head_dim: int, device: torch.device
) -> tuple[torch.Tensor, torch.Tensor]:
    key = (_device_key(device), total_q, num_heads, qk_head_dim)
    cached = _Q_QUANT_CACHE.get(key)
    if cached is None:
        cached = (
            torch.empty(
                (total_q, num_heads, qk_head_dim), dtype=FP8_DTYPE, device=device
            ),
            torch.empty((1,), dtype=torch.float32, device=device),
        )
        _Q_QUANT_CACHE[key] = cached
    return cached


def _quantize_q_fp8(
    q: torch.Tensor, num_heads: int, qk_head_dim: int
) -> tuple[torch.Tensor, torch.Tensor]:
    q_fp8, q_scale = _get_cached_q_quant(q.shape[0], num_heads, qk_head_dim, q.device)
    dynamic_per_tensor_quant(
        q_fp8.view(-1, qk_head_dim),
        q.view(-1, qk_head_dim),
        q_scale,
    )
    return q_fp8, q_scale


def _get_cached_metadata(
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    batch_size: int,
    q_seq_len: int,
    kv_seq_len: int,
    num_heads: int,
    num_kv_heads: int,
    num_kv_splits: int,
    kv_granularity: int,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
) -> dict[str, torch.Tensor]:
    device = qo_indptr.device
    key = (
        _device_key(device),
        batch_size,
        q_seq_len,
        kv_seq_len,
        num_heads,
        num_kv_heads,
        num_kv_splits,
        kv_granularity,
        str(q_dtype),
        str(kv_dtype),
    )
    cached = _META_CACHE.get(key)
    if cached is not None:
        return cached

    kv_last_page_len = _get_kv_last_page_len(batch_size, kv_seq_len, device)
    info = get_mla_metadata_info_v1(
        batch_size,
        q_seq_len,
        num_heads,
        q_dtype,
        kv_dtype,
        is_sparse=False,
        fast_mode=FAST_MODE,
        num_kv_splits=num_kv_splits,
        intra_batch_mode=INTRA_BATCH_MODE,
    )
    work = [torch.empty(shape, dtype=dtype, device=device) for shape, dtype in info]
    (
        work_metadata,
        work_indptr,
        work_info_set,
        reduce_indptr,
        reduce_final_map,
        reduce_partial_map,
    ) = work

    get_mla_metadata_v1(
        qo_indptr,
        kv_indptr,
        kv_last_page_len,
        num_heads // num_kv_heads,
        num_kv_heads,
        True,
        work_metadata,
        work_info_set,
        work_indptr,
        reduce_indptr,
        reduce_final_map,
        reduce_partial_map,
        page_size=PAGE_SIZE,
        kv_granularity=kv_granularity,
        max_seqlen_qo=q_seq_len,
        uni_seqlen_qo=q_seq_len,
        fast_mode=FAST_MODE,
        max_split_per_batch=num_kv_splits,
        intra_batch_mode=INTRA_BATCH_MODE,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )

    cached = {
        "kv_last_page_len": kv_last_page_len,
        "work_meta_data": work_metadata,
        "work_indptr": work_indptr,
        "work_info_set": work_info_set,
        "reduce_indptr": reduce_indptr,
        "reduce_final_map": reduce_final_map,
        "reduce_partial_map": reduce_partial_map,
    }
    _META_CACHE[key] = cached
    return cached


def _get_cached_paged_kv_indices(
    batch_size: int, kv_seq_len: int, page_size: int, device: torch.device
) -> torch.Tensor:
    pages_per_batch = kv_seq_len // page_size
    total_pages = batch_size * pages_per_batch
    return _get_cached_kv_indices(total_pages, device)


def _get_cached_paged_kv_indptr(
    batch_size: int, kv_seq_len: int, page_size: int, device: torch.device
) -> torch.Tensor:
    pages_per_batch = kv_seq_len // page_size
    key = (_device_key(device), "paged_kv_indptr", batch_size, kv_seq_len, page_size)
    cached = _KV_MISC_CACHE.get(key)
    if cached is None:
        cached = torch.arange(
            0,
            (batch_size + 1) * pages_per_batch,
            pages_per_batch,
            dtype=torch.int32,
            device=device,
        )
        _KV_MISC_CACHE[key] = cached
    return cached


def _get_cached_paged_last_page_len(
    batch_size: int, page_size: int, device: torch.device
) -> torch.Tensor:
    key = (_device_key(device), "paged_last_page_len", batch_size, page_size)
    cached = _KV_MISC_CACHE.get(key)
    if cached is None:
        cached = torch.full((batch_size,), page_size, dtype=torch.int32, device=device)
        _KV_MISC_CACHE[key] = cached
    return cached


def _get_cached_paged_metadata(
    qo_indptr: torch.Tensor,
    batch_size: int,
    q_seq_len: int,
    kv_seq_len: int,
    page_size: int,
    num_heads: int,
    num_kv_heads: int,
    num_kv_splits: int,
    kv_granularity: int,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
) -> dict[str, torch.Tensor]:
    device = qo_indptr.device
    paged_fast_mode = _use_paged_fast_mode(batch_size, kv_seq_len, page_size)
    key = (
        _device_key(device),
        "paged",
        batch_size,
        q_seq_len,
        kv_seq_len,
        page_size,
        num_heads,
        num_kv_heads,
        num_kv_splits,
        kv_granularity,
        paged_fast_mode,
        str(q_dtype),
        str(kv_dtype),
    )
    cached = _PAGED_META_CACHE.get(key)
    if cached is not None:
        return cached

    paged_kv_indptr = _get_cached_paged_kv_indptr(
        batch_size, kv_seq_len, page_size, device
    )
    paged_last_page_len = _get_cached_paged_last_page_len(batch_size, page_size, device)
    info = get_mla_metadata_info_v1(
        batch_size,
        q_seq_len,
        num_heads,
        q_dtype,
        kv_dtype,
        is_sparse=False,
        fast_mode=paged_fast_mode,
        num_kv_splits=num_kv_splits,
        intra_batch_mode=INTRA_BATCH_MODE,
    )
    work = [torch.empty(shape, dtype=dtype, device=device) for shape, dtype in info]
    (
        work_metadata,
        work_indptr,
        work_info_set,
        reduce_indptr,
        reduce_final_map,
        reduce_partial_map,
    ) = work

    get_mla_metadata_v1(
        qo_indptr,
        paged_kv_indptr,
        paged_last_page_len,
        num_heads // num_kv_heads,
        num_kv_heads,
        True,
        work_metadata,
        work_info_set,
        work_indptr,
        reduce_indptr,
        reduce_final_map,
        reduce_partial_map,
        page_size=page_size,
        kv_granularity=kv_granularity,
        max_seqlen_qo=q_seq_len,
        uni_seqlen_qo=q_seq_len,
        fast_mode=paged_fast_mode,
        max_split_per_batch=num_kv_splits,
        intra_batch_mode=INTRA_BATCH_MODE,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )

    cached = {
        "kv_indptr": paged_kv_indptr,
        "kv_last_page_len": paged_last_page_len,
        "work_meta_data": work_metadata,
        "work_indptr": work_indptr,
        "work_info_set": work_info_set,
        "reduce_indptr": reduce_indptr,
        "reduce_final_map": reduce_final_map,
        "reduce_partial_map": reduce_partial_map,
    }
    _PAGED_META_CACHE[key] = cached
    return cached


def _view_kv_paged(
    kv_buffer_fp8: torch.Tensor,
    batch_size: int,
    kv_seq_len: int,
    num_kv_heads: int,
    page_size: int,
) -> torch.Tensor:
    return kv_buffer_fp8.view(
        batch_size,
        kv_seq_len // page_size,
        page_size,
        num_kv_heads,
        kv_buffer_fp8.shape[-1],
    ).reshape(-1, page_size, num_kv_heads, kv_buffer_fp8.shape[-1])


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data

    batch_size = config["batch_size"]
    num_heads = config["num_heads"]
    num_kv_heads = config["num_kv_heads"]
    qk_head_dim = config["qk_head_dim"]
    v_head_dim = config["v_head_dim"]
    sm_scale = config["sm_scale"]
    q_seq_len = config["q_seq_len"]
    kv_seq_len = config["kv_seq_len"]

    num_kv_splits = _choose_num_kv_splits(batch_size, kv_seq_len)
    kv_granularity = _choose_kv_granularity(batch_size, kv_seq_len)
    kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]

    if _use_bf16_q(kv_seq_len):
        q_input = q
        q_scale = None
    else:
        q_input, q_scale = _quantize_q_fp8(q, num_heads, qk_head_dim)

    output = _get_cached_output(
        q.shape[0],
        num_heads,
        v_head_dim,
        q.device,
        batch_size,
        kv_seq_len,
    )
    alt_page_size = _get_alt_page_size(batch_size, kv_seq_len)

    if alt_page_size is not None:
        paged_kv = _view_kv_paged(
            kv_buffer_fp8,
            batch_size,
            kv_seq_len,
            num_kv_heads,
            alt_page_size,
        )
        paged_metadata = _get_cached_paged_metadata(
            qo_indptr,
            batch_size,
            q_seq_len,
            kv_seq_len,
            alt_page_size,
            num_heads,
            num_kv_heads,
            num_kv_splits,
            kv_granularity,
            q_input.dtype,
            kv_buffer_fp8.dtype,
        )
        paged_kv_indices = _get_cached_paged_kv_indices(
            batch_size, kv_seq_len, alt_page_size, q.device
        )

        if _use_paged_stage1_only(batch_size, kv_seq_len, alt_page_size, num_kv_splits):
            partial_output, partial_lse = _get_cached_partials(
                _choose_stage1_partial_slots(
                    batch_size,
                    kv_seq_len,
                    alt_page_size,
                    int(paged_metadata["reduce_partial_map"].numel()),
                ),
                q_seq_len,
                num_heads,
                v_head_dim,
                q.device,
            )
            mla_decode_stage1_asm_fwd(
                q_input.view(-1, num_heads, qk_head_dim),
                paged_kv,
                qo_indptr,
                paged_metadata["kv_indptr"],
                paged_kv_indices,
                paged_metadata["kv_last_page_len"],
                None,
                paged_metadata["work_meta_data"],
                paged_metadata["work_indptr"],
                paged_metadata["work_info_set"],
                q_seq_len,
                alt_page_size,
                num_kv_heads,
                sm_scale,
                partial_output,
                partial_lse,
                output,
                q_scale,
                kv_scale_fp8,
            )
            return output

        if _use_triton_longkv_reduce(batch_size, kv_seq_len, alt_page_size, num_kv_splits):
            compact_paged_metadata = _get_reduce_compact_metadata(paged_metadata)
            partial_output, partial_lse = _get_cached_partials(
                int(compact_paged_metadata["reduce_partial_map"].numel()),
                q_seq_len,
                num_heads,
                v_head_dim,
                q.device,
            )
            mla_decode_stage1_asm_fwd(
                q_input.view(-1, num_heads, qk_head_dim),
                paged_kv,
                qo_indptr,
                paged_metadata["kv_indptr"],
                paged_kv_indices,
                paged_metadata["kv_last_page_len"],
                None,
                paged_metadata["work_meta_data"],
                paged_metadata["work_indptr"],
                paged_metadata["work_info_set"],
                q_seq_len,
                alt_page_size,
                num_kv_heads,
                sm_scale,
                partial_output,
                partial_lse,
                output,
                q_scale,
                kv_scale_fp8,
            )
            dense_layout = _get_dense_reduce_layout(compact_paged_metadata)
            use_dense_exact6 = (
                batch_size == 32
                and dense_layout["q_identity"]
                and dense_layout["partial_identity"]
                and dense_layout["num_groups"] == batch_size
                and dense_layout["num_partial"] == batch_size * 6
                and dense_layout["split_min"] == 6
                and dense_layout["split_max"] == 6
            )
            use_dense_upto5 = (
                batch_size == 64
                and dense_layout["q_identity"]
                and dense_layout["partial_identity"]
                and dense_layout["num_groups"] == batch_size
                and dense_layout["num_partial"] == 304
                and dense_layout["split_min"] == 4
                and dense_layout["split_max"] == 5
            )

            if use_dense_exact6:
                _probe_dense_fastpath_once(
                    batch_size,
                    kv_seq_len,
                    "dense_exact6 enabled",
                )
                _run_dense_longkv_reduce_q1_exact(
                    partial_output,
                    partial_lse,
                    output,
                    exact_splits=6,
                )
            elif use_dense_upto5:
                _probe_dense_fastpath_once(
                    batch_size,
                    kv_seq_len,
                    "dense_upto5 enabled",
                )
                _run_dense_longkv_reduce_q1_upto(
                    partial_output,
                    partial_lse,
                    compact_paged_metadata["reduce_indptr"],
                    output,
                    max_splits=5,
                )
            else:
                _probe_dense_fastpath_once(
                    batch_size,
                    kv_seq_len,
                    "fallback generic",
                )
                _run_triton_longkv_reduce_q1(
                    partial_output,
                    partial_lse,
                    compact_paged_metadata["reduce_indptr"],
                    compact_paged_metadata["reduce_q_indices"],
                    compact_paged_metadata["reduce_partial_map"],
                    output,
                )
            return output

        smallkv_dense_exact_splits = _get_smallkv_dense_exact_splits(
            batch_size,
            kv_seq_len,
            alt_page_size,
            num_kv_splits,
        )
        if smallkv_dense_exact_splits is not None:
            compact_paged_metadata = _get_reduce_compact_metadata(paged_metadata)
            dense_layout = _get_dense_reduce_layout(compact_paged_metadata)
            use_dense_exact = (
                dense_layout["q_identity"]
                and dense_layout["partial_identity"]
                and dense_layout["num_groups"] == batch_size
                and dense_layout["num_partial"] == batch_size * smallkv_dense_exact_splits
                and dense_layout["split_min"] == smallkv_dense_exact_splits
                and dense_layout["split_max"] == smallkv_dense_exact_splits
            )
            if use_dense_exact:
                partial_output, partial_lse = _get_cached_partials(
                    int(compact_paged_metadata["reduce_partial_map"].numel()),
                    q_seq_len,
                    num_heads,
                    v_head_dim,
                    q.device,
                )
                mla_decode_stage1_asm_fwd(
                    q_input.view(-1, num_heads, qk_head_dim),
                    paged_kv,
                    qo_indptr,
                    paged_metadata["kv_indptr"],
                    paged_kv_indices,
                    paged_metadata["kv_last_page_len"],
                    None,
                    paged_metadata["work_meta_data"],
                    paged_metadata["work_indptr"],
                    paged_metadata["work_info_set"],
                    q_seq_len,
                    alt_page_size,
                    num_kv_heads,
                    sm_scale,
                    partial_output,
                    partial_lse,
                    output,
                    q_scale,
                    kv_scale_fp8,
                )
                _probe_dense_fastpath_once(
                    batch_size,
                    kv_seq_len,
                    f"dense_exact{smallkv_dense_exact_splits} enabled",
                )
                _run_dense_longkv_reduce_q1_exact(
                    partial_output,
                    partial_lse,
                    output,
                    exact_splits=smallkv_dense_exact_splits,
                )
                return output

        mla_decode_fwd(
            q_input.view(-1, num_heads, qk_head_dim),
            paged_kv,
            output,
            qo_indptr,
            paged_metadata["kv_indptr"],
            paged_kv_indices,
            paged_metadata["kv_last_page_len"],
            q_seq_len,
            page_size=alt_page_size,
            nhead_kv=num_kv_heads,
            sm_scale=sm_scale,
            logit_cap=0.0,
            num_kv_splits=num_kv_splits,
            q_scale=q_scale,
            kv_scale=kv_scale_fp8,
            intra_batch_mode=INTRA_BATCH_MODE,
            work_meta_data=paged_metadata["work_meta_data"],
            work_indptr=paged_metadata["work_indptr"],
            work_info_set=paged_metadata["work_info_set"],
            reduce_indptr=paged_metadata["reduce_indptr"],
            reduce_final_map=paged_metadata["reduce_final_map"],
            reduce_partial_map=paged_metadata["reduce_partial_map"],
        )
        return output

    metadata = _get_cached_metadata(
        qo_indptr,
        kv_indptr,
        batch_size,
        q_seq_len,
        kv_seq_len,
        num_heads,
        num_kv_heads,
        num_kv_splits,
        kv_granularity,
        q_input.dtype,
        kv_buffer_fp8.dtype,
    )
    kv_indices = _get_cached_kv_indices(kv_buffer_fp8.shape[0], kv_buffer_fp8.device)

    if _use_stage1_direct(batch_size, kv_seq_len, num_kv_splits):
        partial_output, partial_lse = _get_cached_partials(
            int(metadata["reduce_partial_map"].numel()),
            q_seq_len,
            num_heads,
            v_head_dim,
            q.device,
        )
        mla_decode_stage1_asm_fwd(
            q_input.view(-1, num_heads, qk_head_dim),
            kv_buffer_fp8.view(
                kv_buffer_fp8.shape[0],
                PAGE_SIZE,
                num_kv_heads,
                kv_buffer_fp8.shape[-1],
            ),
            qo_indptr,
            kv_indptr,
            kv_indices,
            metadata["kv_last_page_len"],
            None,
            metadata["work_meta_data"],
            metadata["work_indptr"],
            metadata["work_info_set"],
            q_seq_len,
            PAGE_SIZE,
            num_kv_heads,
            sm_scale,
            partial_output,
            partial_lse,
            output,
            q_scale,
            kv_scale_fp8,
        )
        return output

    mla_decode_fwd(
        q_input.view(-1, num_heads, qk_head_dim),
        kv_buffer_fp8.view(
            kv_buffer_fp8.shape[0],
            PAGE_SIZE,
            num_kv_heads,
            kv_buffer_fp8.shape[-1],
        ),
        output,
        qo_indptr,
        kv_indptr,
        kv_indices,
        metadata["kv_last_page_len"],
        q_seq_len,
        page_size=PAGE_SIZE,
        nhead_kv=num_kv_heads,
        sm_scale=sm_scale,
        logit_cap=0.0,
        num_kv_splits=num_kv_splits,
        q_scale=q_scale,
        kv_scale=kv_scale_fp8,
        intra_batch_mode=INTRA_BATCH_MODE,
        work_meta_data=metadata["work_meta_data"],
        work_indptr=metadata["work_indptr"],
        work_info_set=metadata["work_info_set"],
        reduce_indptr=metadata["reduce_indptr"],
        reduce_final_map=metadata["reduce_final_map"],
        reduce_partial_map=metadata["reduce_partial_map"],
    )
    return output
scrolls · 1304 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON