Skip to content
KernelIndex
Search⌘K

submission 755180

gavin1104. · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-755180?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
331.8µs
#706 of 766
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7afdef429d4ff9dde318cd0bafd9460b093b0836bb4827f3469a0ca8099a9db8
license declaredunknown
license concludedunknown
authorsgavin1104.
imported2026-08-26

Techniques

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

persistent-kernelStrategy: Avoid aiter persistent-mode overhead by using direct PyTorch batched
split-ksplit-K reduction, persistent mode coordination) that dominates for small batches.

Kernel source

submission.py49 lines
"""
Optimized MLA decode kernel for MI355X.

Strategy: Avoid aiter persistent-mode overhead by using direct PyTorch batched
attention for all cases. For decode (q_seq_len=1), the computation is simple:
  scores = Q @ K^T * sm_scale   -> softmax -> @ V

The aiter persistent kernel has significant per-call overhead (metadata allocation,
split-K reduction, persistent mode coordination) that dominates for small batches.
Pure PyTorch bmm avoids all of this.

Uses bf16 KV cache directly (no quantization overhead).
"""

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

SM_SCALE = 1.0 / (576 ** 0.5)
NUM_HEADS = 16
HEAD_DIM = 576
V_DIM = 512


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    kv_len = config["kv_seq_len"]

    # Use bf16 KV directly - no quantization overhead
    kv_bf16 = kv_data["bf16"]  # (total_kv, 1, 576)

    # Reshape to batched format
    q_3d = q.view(bs, NUM_HEADS, HEAD_DIM)                  # (bs, 16, 576)
    kv_3d = kv_bf16.view(bs, kv_len, HEAD_DIM)              # (bs, kv_len, 576)

    # Compute attention scores: (bs, 16, kv_len)
    scores = torch.bmm(q_3d.float(), kv_3d.transpose(1, 2).float()) * SM_SCALE

    # Softmax over kv_len dimension
    probs = F.softmax(scores.float(), dim=-1)

    # Compute output using first 512 dims as values: (bs, 16, 512)
    values = kv_3d[:, :, :V_DIM].float()                    # (bs, kv_len, 512)
    output = torch.bmm(probs, values)

    return output.to(torch.bfloat16).reshape(-1, NUM_HEADS, V_DIM)
scrolls · 49 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