Skip to content
KernelIndex
Search⌘K

submission 589468

_radna · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.after.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-589468?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
68.3µs
#304 of 766
2026-03-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:19c194c332648ec6efc9eb8adb4a9d2d69a3851cb22e7d1cd92136c630b6d1af
license declaredunknown
license concludedunknown
authors_radna
imported2026-08-26

Techniques

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

persistent-kernelThe primary fast path uses aiter's persistent a8w8 decode kernel with cached
tile-n = 128FP8_Q1_MIN_BLOCK_N = 128

Kernel source

submission.after.py718 lines
"""
MLA decode submission that prioritizes the real aiter asm q=1 path.

The primary fast path uses aiter's persistent a8w8 decode kernel with cached
metadata, quantization, and scratch buffers. The accepted grouped Triton bf16
path stays as the first fallback, followed by the safe fp8 template path.
"""

import aiter
import torch
import torch.nn.functional as F
from task import input_t, output_t

from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.ops.triton.attention.mla_decode_rope import decode_attention_fwd_grouped_rope

try:
    from aiter.jit.utils.chip_info import get_cu_num as _aiter_get_cu_num
except Exception:
    _aiter_get_cu_num = None


FP8_DTYPE = aiter_dtypes.fp8
DEFAULT_NUM_KV_SPLITS = 32
MAX_AITER_NUM_KV_SPLITS = 16
PAGE_SIZE = 1
TRITON_Q1_WORK_THRESHOLD = 32768
TRITON_TARGET_STAGE1_PROGRAMS = 512
TRITON_SHORT_KV_LEN = 1024
TRITON_SHORT_KV_BATCH_FLOOR = 32
TRITON_SHORT_KV_TARGET_STAGE1_PROGRAMS = 384
AITER_TARGET_STAGE1_PROGRAMS = 256
FP8_Q1_MIN_BLOCK_N = 128
AITER_LONG_KV_SPLIT_OVERHEAD = 84.1
_AITER_Q1_STATE_CACHE: dict[str, object] = {}
_TRITON_Q1_BUFFER_CACHE: dict[str, object] = {}
_TRITON_KV_INDEX_CACHE: dict[str, object] = {}
_FP8_Q_BUFFER_CACHE: dict[str, object] = {}
_FP8_DYNAMIC_QUANT_STATE: dict[str, object] = {"enabled": True}
_AITER_BF16_FP8_SUPPORT: dict[str, object] = {}
_AITER_CU_NUM_STATE: dict[str, object] = {"resolved": False, "value": None}


def _quantize_fp8_torch(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Reference dynamic per-tensor FP8 quantization."""
    finfo = torch.finfo(FP8_DTYPE)
    amax = tensor.abs().amax().clamp(min=1e-12)
    scale = amax / finfo.max
    fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
    return fp8_tensor, scale.to(torch.float32).reshape(1)


def _get_fp8_q_buffers(
    shape: tuple[int, ...],
    input_dtype: torch.dtype,
    device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
    key = (
        shape,
        str(input_dtype),
        -1 if device.index is None else int(device.index),
    )
    cached = _cache_lookup(_FP8_Q_BUFFER_CACHE, key)
    if cached is not None:
        return cached

    q_fp8 = torch.empty(shape, dtype=FP8_DTYPE, device=device)
    q_scale = torch.empty((1,), dtype=torch.float32, device=device)
    return _cache_store(_FP8_Q_BUFFER_CACHE, key, (q_fp8, q_scale))


def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Prefer compiled aiter FP8 quantization, then fall back to the reference path."""
    if bool(_FP8_DYNAMIC_QUANT_STATE.get("enabled", True)):
        q_fp8, q_scale = _get_fp8_q_buffers(tensor.shape, tensor.dtype, tensor.device)
        try:
            aiter.dynamic_per_tensor_quant(q_fp8, tensor, q_scale)
            return q_fp8, q_scale
        except Exception:
            _FP8_DYNAMIC_QUANT_STATE["enabled"] = False
    return _quantize_fp8_torch(tensor)


def _scaled_mm_qk(
    q_i_fp8: torch.Tensor,
    k_i_fp8: torch.Tensor,
    q_scale: torch.Tensor,
    kv_scale_fp8: torch.Tensor,
) -> torch.Tensor:
    try:
        return torch._scaled_mm(
            q_i_fp8,
            k_i_fp8.t(),
            scale_a=q_scale,
            scale_b=kv_scale_fp8,
            out_dtype=torch.bfloat16,
        )
    except RuntimeError:
        return torch._scaled_mm(
            q_i_fp8,
            k_i_fp8.t(),
            scale_a=q_scale,
            scale_b=kv_scale_fp8,
            out_dtype=torch.float32,
        )


def _cache_lookup(cache: dict[str, object], key: tuple) -> object | None:
    entry = cache.get("entry")
    if entry is None or entry["key"] != key:
        return None
    return entry["value"]


def _cache_store(cache: dict[str, object], key: tuple, value: object) -> object:
    cache.clear()
    cache["entry"] = {"key": key, "value": value}
    return value


def _uniform_segment_length(indptr: torch.Tensor) -> int | None:
    lengths = indptr[1:] - indptr[:-1]
    if lengths.numel() == 0:
        return 0
    first = int(lengths[0].item())
    if not bool(torch.all(lengths == first).item()):
        return None
    return first


def _supports_q1(qo_indptr: torch.Tensor) -> bool:
    return _uniform_segment_length(qo_indptr) == 1


def _get_q1_router_shape(config: dict) -> tuple[int, int]:
    batch_size = int(config.get("batch_size", 0))
    q_seq_len = int(config.get("q_seq_len", -1))
    kv_seq_len = int(config.get("kv_seq_len", 0))
    if batch_size <= 0 or q_seq_len != 1 or kv_seq_len <= 0:
        raise RuntimeError("Config router only handles q_len == 1")
    return batch_size, kv_seq_len


def _prefer_triton_q1(batch_size: int, kv_len: int) -> bool:
    if batch_size <= 4 and kv_len > TRITON_SHORT_KV_LEN:
        return False
    return batch_size * kv_len <= TRITON_Q1_WORK_THRESHOLD


def _ceil_div(num: int, den: int) -> int:
    return (num + den - 1) // den


def _get_aiter_cu_num_cached() -> int | None:
    if not bool(_AITER_CU_NUM_STATE["resolved"]):
        value = None
        if _aiter_get_cu_num is not None:
            try:
                value = int(_aiter_get_cu_num())
            except Exception:
                value = None
        _AITER_CU_NUM_STATE["value"] = value
        _AITER_CU_NUM_STATE["resolved"] = True
    return _AITER_CU_NUM_STATE["value"]


def _select_triton_num_kv_splits(batch_size: int, num_heads: int, kv_len: int) -> int:
    head_blocks = max(1, _ceil_div(num_heads, 16))
    target_stage1_programs = TRITON_TARGET_STAGE1_PROGRAMS
    if kv_len <= TRITON_SHORT_KV_LEN and batch_size >= TRITON_SHORT_KV_BATCH_FLOOR:
        target_stage1_programs = TRITON_SHORT_KV_TARGET_STAGE1_PROGRAMS
    desired = _ceil_div(
        target_stage1_programs,
        max(1, batch_size * head_blocks),
    )
    desired = max(4, min(DEFAULT_NUM_KV_SPLITS, desired))
    max_valid_splits = max(1, _ceil_div(kv_len, 32))
    return min(desired, max_valid_splits)


def _select_aiter_num_kv_splits(batch_size: int, num_heads: int, kv_len: int) -> int:
    head_groups = max(1, num_heads // 16)
    fallback = _ceil_div(
        AITER_TARGET_STAGE1_PROGRAMS,
        max(1, batch_size * head_groups),
    )
    fallback = max(1, min(MAX_AITER_NUM_KV_SPLITS, fallback))
    max_valid_splits = max(1, _ceil_div(kv_len, FP8_Q1_MIN_BLOCK_N))
    if kv_len <= TRITON_SHORT_KV_LEN:
        return min(fallback, max_valid_splits)

    cu_num = _get_aiter_cu_num_cached()
    if cu_num is None or cu_num <= 0:
        return min(fallback, max_valid_splits)

    max_candidate = min(MAX_AITER_NUM_KV_SPLITS, max_valid_splits)
    avg_kv = float(kv_len)

    def _split_score(split_count: int) -> float:
        occupied_cus = _ceil_div(batch_size * split_count, cu_num) * cu_num
        return (
            (batch_size * split_count / max(1, occupied_cus))
            * avg_kv
            / (avg_kv + AITER_LONG_KV_SPLIT_OVERHEAD * split_count)
        )

    return max(range(1, max_candidate + 1), key=_split_score)


def _prefer_aiter_bf16_fp8_shortkv(batch_size: int, kv_len: int) -> bool:
    return kv_len == TRITON_SHORT_KV_LEN and not _prefer_triton_q1(batch_size, kv_len)


def _get_aiter_q1_state(
    batch_size: int,
    num_heads: int,
    num_kv_heads: int,
    qk_head_dim: int,
    v_head_dim: int,
    kv_len: int,
    num_kv_splits: int,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
    device: torch.device,
) -> dict[str, torch.Tensor]:
    key = (
        batch_size,
        num_heads,
        num_kv_heads,
        qk_head_dim,
        v_head_dim,
        kv_len,
        num_kv_splits,
        str(q_dtype),
        str(kv_dtype),
        -1 if device.index is None else int(device.index),
    )
    cached = _cache_lookup(_AITER_Q1_STATE_CACHE, key)
    if cached is not None:
        return cached

    if num_heads not in (16, 32) or num_heads % 16 != 0:
        raise RuntimeError("Unsupported aiter q=1 head shape")

    qo_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device=device)
    kv_indptr = qo_indptr * kv_len
    kv_last_page_len = torch.full(
        (batch_size,),
        kv_len,
        dtype=torch.int32,
        device=device,
    )
    kv_indices = torch.arange(batch_size * kv_len, dtype=torch.int32, device=device)

    info = get_mla_metadata_info_v1(
        batch_size,
        1,
        num_heads,
        q_dtype,
        kv_dtype,
        is_sparse=False,
        fast_mode=False,
        num_kv_splits=num_kv_splits,
        intra_batch_mode=True,
    )
    work = [torch.empty(shape, dtype=dtype, device=device) for shape, dtype in info]
    (
        work_meta_data,
        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_meta_data,
        work_info_set,
        work_indptr,
        reduce_indptr,
        reduce_final_map,
        reduce_partial_map,
        page_size=PAGE_SIZE,
        kv_granularity=max(PAGE_SIZE, 16),
        max_seqlen_qo=1,
        uni_seqlen_qo=1,
        fast_mode=False,
        max_split_per_batch=num_kv_splits,
        intra_batch_mode=True,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )

    effective_heads = 16 if num_heads > 16 else num_heads
    partial_tiles = int(reduce_partial_map.size(0))
    output = torch.empty(
        (batch_size, num_heads, v_head_dim),
        dtype=torch.bfloat16,
        device=device,
    )
    logits = torch.empty(
        (partial_tiles, 1, effective_heads, v_head_dim),
        dtype=torch.float32,
        device=device,
    )
    attn_lse = torch.empty(
        (partial_tiles, 1, effective_heads, 1),
        dtype=torch.float32,
        device=device,
    )

    return _cache_store(
        _AITER_Q1_STATE_CACHE,
        key,
        {
            "qo_indptr": qo_indptr,
            "kv_indptr": kv_indptr,
            "kv_last_page_len": kv_last_page_len,
            "kv_indices": kv_indices,
            "work_meta_data": work_meta_data,
            "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,
            "output": output,
            "logits": logits,
            "attn_lse": attn_lse,
        },
    )


def _get_triton_q1_buffers(
    batch_size: int,
    num_heads: int,
    kv_lora_rank: int,
    num_kv_splits: int,
    q_dtype: torch.dtype,
    device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
    key = (
        batch_size,
        num_heads,
        kv_lora_rank,
        num_kv_splits,
        str(q_dtype),
        -1 if device.index is None else int(device.index),
    )
    cached = _cache_lookup(_TRITON_Q1_BUFFER_CACHE, key)
    if cached is not None:
        return cached

    output = torch.empty(
        (batch_size, num_heads, kv_lora_rank),
        dtype=torch.bfloat16,
        device=device,
    )
    attn_buffer = torch.empty(
        (batch_size, num_heads, num_kv_splits, kv_lora_rank + 1),
        dtype=torch.float32,
        device=device,
    )
    dummy_k_pe = torch.empty((1,), dtype=q_dtype, device=device)
    dummy_cos_sin = torch.empty((1, 1), dtype=q_dtype, device=device)
    dummy_positions = torch.zeros((batch_size,), dtype=torch.int32, device=device)
    return _cache_store(
        _TRITON_Q1_BUFFER_CACHE,
        key,
        (output, attn_buffer, dummy_k_pe, dummy_cos_sin, dummy_positions),
    )


def _get_kv_indices(total_kv: int, device: torch.device) -> torch.Tensor:
    key = (total_kv, -1 if device.index is None else int(device.index))
    cached = _cache_lookup(_TRITON_KV_INDEX_CACHE, key)
    if cached is not None:
        return cached
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
    return _cache_store(_TRITON_KV_INDEX_CACHE, key, kv_indices)


def _custom_kernel_aiter_asm_q1_impl(
    data: input_t,
    batch_size: int,
    kv_len: int,
) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data

    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"]
    num_kv_splits = _select_aiter_num_kv_splits(batch_size, num_heads, kv_len)
    kv_buffer_fp8, kv_scale = kv_data["fp8"]

    kv_buffer_4d = kv_buffer_fp8.view(
        kv_buffer_fp8.shape[0],
        PAGE_SIZE,
        num_kv_heads,
        qk_head_dim,
    )

    def _run_stage1(
        q_input: torch.Tensor,
        q_scale: torch.Tensor | None,
    ) -> output_t:
        state = _get_aiter_q1_state(
            batch_size,
            num_heads,
            num_kv_heads,
            qk_head_dim,
            v_head_dim,
            kv_len,
            num_kv_splits,
            q_input.dtype,
            kv_buffer_fp8.dtype,
            q.device,
        )

        q_view = q_input.view(batch_size, num_heads, qk_head_dim)
        output = state["output"]
        stage_q = q_view
        stage_output = output
        if num_heads != 16:
            stage_q = q_view.view(batch_size * (num_heads // 16), 16, qk_head_dim)
            stage_output = output.view(batch_size * (num_heads // 16), 16, v_head_dim)

        aiter.mla_decode_stage1_asm_fwd(
            stage_q,
            kv_buffer_4d,
            state["qo_indptr"],
            state["kv_indptr"],
            state["kv_indices"],
            state["kv_last_page_len"],
            None,
            state["work_meta_data"],
            state["work_indptr"],
            state["work_info_set"],
            1,
            PAGE_SIZE,
            num_kv_heads,
            config["sm_scale"],
            state["logits"],
            state["attn_lse"],
            stage_output,
            q_scale,
            kv_scale,
        )
        aiter.mla_reduce_v1(
            state["logits"],
            state["attn_lse"],
            state["reduce_indptr"],
            state["reduce_final_map"],
            state["reduce_partial_map"],
            1,
            stage_output,
            None,
        )
        return output

    q_fp8, q_scale = quantize_fp8(q)
    return _run_stage1(q_fp8, q_scale)


def _custom_kernel_aiter_asm_q1_bf16_fp8_impl(
    data: input_t,
    batch_size: int,
    kv_len: int,
) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data

    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"]
    num_kv_splits = _select_aiter_num_kv_splits(batch_size, num_heads, kv_len)
    kv_buffer_fp8, kv_scale = kv_data["fp8"]
    kv_buffer_4d = kv_buffer_fp8.view(
        kv_buffer_fp8.shape[0],
        PAGE_SIZE,
        num_kv_heads,
        qk_head_dim,
    )

    support_key = (
        batch_size,
        kv_len,
        num_heads,
        num_kv_heads,
        qk_head_dim,
        v_head_dim,
        num_kv_splits,
        -1 if q.device.index is None else int(q.device.index),
    )
    if _AITER_BF16_FP8_SUPPORT.get(support_key) is False:
        raise RuntimeError("AITER bf16/fp8 q=1 asm path disabled for this shape")

    q_bf16 = q if q.is_contiguous() else q.contiguous()
    try:
        state = _get_aiter_q1_state(
            batch_size,
            num_heads,
            num_kv_heads,
            qk_head_dim,
            v_head_dim,
            kv_len,
            num_kv_splits,
            q_bf16.dtype,
            kv_buffer_fp8.dtype,
            q.device,
        )

        q_view = q_bf16.view(batch_size, num_heads, qk_head_dim)
        output = state["output"]
        stage_q = q_view
        stage_output = output
        if num_heads != 16:
            stage_q = q_view.view(batch_size * (num_heads // 16), 16, qk_head_dim)
            stage_output = output.view(batch_size * (num_heads // 16), 16, v_head_dim)

        aiter.mla_decode_stage1_asm_fwd(
            stage_q,
            kv_buffer_4d,
            state["qo_indptr"],
            state["kv_indptr"],
            state["kv_indices"],
            state["kv_last_page_len"],
            None,
            state["work_meta_data"],
            state["work_indptr"],
            state["work_info_set"],
            1,
            PAGE_SIZE,
            num_kv_heads,
            config["sm_scale"],
            state["logits"],
            state["attn_lse"],
            stage_output,
            None,
            kv_scale,
        )
        aiter.mla_reduce_v1(
            state["logits"],
            state["attn_lse"],
            state["reduce_indptr"],
            state["reduce_final_map"],
            state["reduce_partial_map"],
            1,
            stage_output,
            None,
        )
        _AITER_BF16_FP8_SUPPORT[support_key] = True
        return output
    except Exception:
        _AITER_BF16_FP8_SUPPORT[support_key] = False
        raise


def _custom_kernel_aiter_asm_q1(data: input_t) -> output_t:
    _, _, qo_indptr, kv_indptr, _ = data

    batch_size = qo_indptr.shape[0] - 1
    q_len = _uniform_segment_length(qo_indptr)
    kv_len = _uniform_segment_length(kv_indptr)
    if batch_size == 0 or q_len != 1 or kv_len is None:
        raise RuntimeError("Aiter asm path only handles uniform q_len == 1")
    return _custom_kernel_aiter_asm_q1_impl(data, batch_size, kv_len)


def _custom_kernel_triton_q1_impl(
    data: input_t,
    batch_size: int,
    kv_len: int,
) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data

    num_heads = config["num_heads"]
    kv_lora_rank = config["kv_lora_rank"]
    qk_rope_head_dim = config["qk_rope_head_dim"]
    sm_scale = config["sm_scale"]
    num_kv_splits = _select_triton_num_kv_splits(batch_size, num_heads, kv_len)
    total_kv = batch_size * kv_len

    q_batch = q.reshape(batch_size, num_heads, config["qk_head_dim"])
    kv_buffer = kv_data["bf16"]
    v_buffer = kv_buffer[:, :, :kv_lora_rank]
    output, attn_buffer, dummy_k_pe, dummy_cos_sin, dummy_positions = (
        _get_triton_q1_buffers(
            batch_size,
            num_heads,
            kv_lora_rank,
            num_kv_splits,
            q.dtype,
            q.device,
        )
    )
    kv_indices = _get_kv_indices(total_kv, q.device)

    decode_attention_fwd_grouped_rope(
        q=q_batch,
        k_buffer=kv_buffer,
        v_buffer=v_buffer,
        o=output,
        kv_indptr=kv_indptr,
        kv_indices=kv_indices,
        k_pe_tokens=dummy_k_pe,
        kv_lora_rank=kv_lora_rank,
        rotary_dim=qk_rope_head_dim,
        cos_sin_cache=dummy_cos_sin,
        positions=dummy_positions,
        attn_logits=attn_buffer,
        num_kv_splits=num_kv_splits,
        sm_scale=sm_scale,
        logit_cap=0.0,
        use_rope=False,
        is_neox_style=False,
    )
    return output


def _custom_kernel_triton_q1(data: input_t) -> output_t:
    _, _, qo_indptr, kv_indptr, _ = data

    batch_size = qo_indptr.shape[0] - 1
    kv_len = _uniform_segment_length(kv_indptr)
    if batch_size == 0 or not _supports_q1(qo_indptr) or kv_len is None:
        raise RuntimeError("Triton path only handles q_len == 1")
    return _custom_kernel_triton_q1_impl(data, batch_size, kv_len)


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

    num_heads = config["num_heads"]
    kv_lora_rank = config["kv_lora_rank"]
    qk_head_dim = config["qk_head_dim"]
    sm_scale = config["sm_scale"]

    kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
    kv_fp8_2d = kv_buffer_fp8.view(-1, qk_head_dim)
    kv_values_bf16 = kv_data["bf16"][:, 0, :kv_lora_rank]

    q_fp8, q_scale = quantize_fp8(q)

    batch_size = qo_indptr.shape[0] - 1
    out_list = []

    for i in range(batch_size):
        q_s = int(qo_indptr[i].item())
        q_e = int(qo_indptr[i + 1].item())
        kv_s = int(kv_indptr[i].item())
        kv_e = int(kv_indptr[i + 1].item())
        seq_q = q_e - q_s

        q_i_fp8 = q_fp8[q_s:q_e].reshape(seq_q * num_heads, qk_head_dim)
        k_i_fp8 = kv_fp8_2d[kv_s:kv_e]
        v_i = kv_values_bf16[kv_s:kv_e]

        raw_scores = _scaled_mm_qk(q_i_fp8, k_i_fp8, q_scale, kv_scale_fp8)

        if seq_q == 1:
            scores = F.softmax(raw_scores.float() * sm_scale, dim=-1)
            output_i = torch.matmul(scores.to(torch.bfloat16), v_i)
            out_list.append(output_i.unsqueeze(0))
            continue

        scores = raw_scores.view(seq_q, num_heads, -1).permute(1, 0, 2)
        scores = F.softmax(scores.float() * sm_scale, dim=-1)
        output_i = torch.matmul(scores.to(torch.bfloat16), v_i).permute(1, 0, 2)
        out_list.append(output_i)

    return torch.cat(out_list, dim=0)


def _custom_kernel_hybrid_q1(data: input_t) -> output_t:
    _, _, _, _, config = data

    batch_size, kv_len = _get_q1_router_shape(config)
    if _prefer_aiter_bf16_fp8_shortkv(batch_size, kv_len):
        try:
            return _custom_kernel_aiter_asm_q1_bf16_fp8_impl(data, batch_size, kv_len)
        except Exception:
            return _custom_kernel_aiter_asm_q1_impl(data, batch_size, kv_len)
    if _prefer_triton_q1(batch_size, kv_len):
        try:
            return _custom_kernel_triton_q1_impl(data, batch_size, kv_len)
        except Exception:
            return _custom_kernel_aiter_asm_q1_impl(data, batch_size, kv_len)

    try:
        return _custom_kernel_aiter_asm_q1_impl(data, batch_size, kv_len)
    except Exception:
        return _custom_kernel_triton_q1_impl(data, batch_size, kv_len)


def custom_kernel(data: input_t) -> output_t:
    try:
        return _custom_kernel_hybrid_q1(data)
    except Exception:
        try:
            return _custom_kernel_triton_q1(data)
        except Exception:
            try:
                return _custom_kernel_aiter_asm_q1(data)
            except Exception:
                return _custom_kernel_safe_fp8(data)


#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
scrolls · 718 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