Skip to content
KernelIndex
Search⌘K

submission 629089

michael ma · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission4.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-629089?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
188.6µs
#593 of 766
2026-03-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:921eef02b58e28f8b1747f4f1e952911560f094178c0a36671d03b7713796187
license declaredunknown
license concludedunknown
authorsmichael ma
imported2026-08-26

Techniques

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

split-k基于理论模型选择最优 split-K 数量:

Kernel source

submission4.py91 lines
import torch
from aiter.mla import mla_decode_fwd
from reference import (
    NUM_HEADS, NUM_KV_HEADS, QK_HEAD_DIM, V_HEAD_DIM,
    SM_SCALE, PAGE_SIZE, _make_mla_decode_metadata, quantize_fp8
)

def select_num_kv_splits(kv_seq_len: int, batch_size: int) -> int:
    """
    基于理论模型选择最优 split-K 数量:
        S_opt ≈ sqrt(L * B_kv / overhead)
    其中 B_kv = 576 (FP8 每 token 字节数),overhead ≈ 2KB(合并开销)
    结果再根据 batch 规模调整,并取 2 的幂以便 kernel 调度。
    """
    L = kv_seq_len
    B_kv = 576
    overhead = 2048   # 2KB
    s_theory = int((L * B_kv / overhead) ** 0.5)
    # 限制范围
    s = max(1, min(128, s_theory))
    # 针对大 batch 长序列可适当增加
    if batch_size >= 64 and L >= 8192:
        s = max(s, 64)
    # 量化到 2 的幂
    if s < 16:
        return 16
    elif s < 32:
        return 32
    elif s < 64:
        return 64
    else:
        return 128

def custom_kernel(data):
    q, kv_data, qo_indptr, kv_indptr, config = data
    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]

    # 自适应 split-K
    num_kv_splits = select_num_kv_splits(kv_seq_len, batch_size)

    # 使用 FP8 加速(aiter 最优路径)
    kv_buffer, kv_scale = kv_data["fp8"]
    q_fp8, q_scale = quantize_fp8(q)

    # 准备 4D KV buffer
    kv_buffer_4d = kv_buffer.view(kv_buffer.shape[0], PAGE_SIZE, NUM_KV_HEADS, kv_buffer.shape[-1])

    total_kv_len = int(kv_indptr[-1].item())
    kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
    kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)

    # 构建 metadata(不使用 fast_mode,保持与参考一致)
    meta = _make_mla_decode_metadata(
        batch_size,
        config["q_seq_len"],
        NUM_HEADS,
        NUM_KV_HEADS,
        q_fp8.dtype,
        kv_buffer.dtype,
        qo_indptr,
        kv_indptr,
        kv_last_page_len,
        num_kv_splits=num_kv_splits,
    )

    o = torch.empty(
        (q.shape[0], NUM_HEADS, V_HEAD_DIM),
        dtype=torch.bfloat16, device="cuda"
    )

    mla_decode_fwd(
        q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM),
        kv_buffer_4d,
        o,
        qo_indptr,
        kv_indptr,
        kv_indices,
        kv_last_page_len,
        config["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,
        intra_batch_mode=True,
        **meta,
    )
    return o
scrolls · 91 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