Skip to content
KernelIndex
Search⌘K

submission 660382

Maxwell Cipher · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

mla_v41.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-660382?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
73.5µs
#352 of 766
2026-03-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8f8071424466b6b9a927d8d92f793f8a8b7f93cdf2aa1766cf93ceeca569bab1
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15

Techniques

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

persistent-kernel"""v41: Pure bf16 non-persistent — the dgavriloff insight.

Kernel source

mla_v41.py52 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""v41: Pure bf16 non-persistent — the dgavriloff insight.
Eliminates ALL overhead: no Q quantization, no metadata, no reduce.
Only 1 kernel launch per call. Trades 2x bandwidth for zero overhead."""

import torch
from task import input_t, output_t
from aiter.mla import mla_decode_fwd

NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)

_cache = {}


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

    batch_size = int(config["batch_size"])
    kv_seq_len = int(config["kv_seq_len"])
    q_total = q.shape[0]

    kv_bf16 = kv_data["bf16"]
    q_bf16 = q.view(-1, NUM_HEADS, QK_HEAD_DIM)
    kv_4d = kv_bf16.view(-1, 1, NUM_KV_HEADS, kv_bf16.shape[-1])

    key = (batch_size, kv_seq_len)
    if key not in _cache:
        total_kv = batch_size * kv_seq_len
        _cache[key] = (
            torch.arange(total_kv, dtype=torch.int32, device="cuda"),
            torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda"),
        )

    kv_indices, kv_last_page_len = _cache[key]
    output = torch.empty((q_total, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")

    mla_decode_fwd(
        q_bf16, kv_4d, output,
        qo_indptr, kv_indptr,
        kv_indices, kv_last_page_len,
        1,
        page_size=1, nhead_kv=NUM_KV_HEADS, sm_scale=SM_SCALE,
        intra_batch_mode=False,
    )
    return output
scrolls · 52 lines total

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

Changes from previous submission

Against this author's previous submission submission 646958.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- import torch
+ """v41: Pure bf16 non-persistent — the dgavriloff insight.
+ Eliminates ALL overhead: no Q quantization, no metadata, no reduce.
+ Only 1 kernel launch per call. Trades 2x bandwidth for zero overhead."""
+ import torch
from task import input_t, output_t
-
from aiter.mla import mla_decode_fwd
- from aiter import dtypes as aiter_dtypes
- from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
- try:
- from aiter import per_tensor_quant_hip
- HAS_QUANT_HIP = True
- except ImportError:
- HAS_QUANT_HIP = False
-
-
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
- PAGE_SIZE = 1
- DEFAULT_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
- FP8_DTYPE = aiter_dtypes.fp8
- LOW_PAYLOAD_TOKEN_LIMIT = 65536
- PERSISTENT_TOKEN_LIMIT = 65536
- PERSISTENT_SPLITS = 32
+ SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
+ _cache = {}
- def _quantize_query(q_view: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
- if HAS_QUANT_HIP:
- q_fp8, q_scale = per_tensor_quant_hip(q_view, quant_dtype=FP8_DTYPE)
- return q_fp8, q_scale.reshape(1)
- finfo = torch.finfo(FP8_DTYPE)
- amax = q_view.abs().amax().clamp(min=1e-12)
- q_scale = (amax / finfo.max).to(torch.float32).reshape(1)
- q_fp8 = (q_view / q_scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
- return q_fp8, q_scale
+ def custom_kernel(data: input_t) -> output_t:
+ q, kv_data, qo_indptr, kv_indptr, config = data
+ batch_size = int(config["batch_size"])
+ kv_seq_len = int(config["kv_seq_len"])
+ q_total = q.shape[0]
- def _prefer_low_overhead_path(total_kv: int) -> bool:
- return total_kv <= LOW_PAYLOAD_TOKEN_LIMIT
+ kv_bf16 = kv_data["bf16"]
+ q_bf16 = q.view(-1, NUM_HEADS, QK_HEAD_DIM)
+ kv_4d = kv_bf16.view(-1, 1, NUM_KV_HEADS, kv_bf16.shape[-1])
+ key = (batch_size, kv_seq_len)
+ if key not in _cache:
+ total_kv = batch_size * kv_seq_len
+ _cache[key] = (
+ torch.arange(total_kv, dtype=torch.int32, device="cuda"),
+ torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda"),
+ )
- def _prefer_persistent_fp8_path(total_kv: int) -> bool:
- return total_kv > PERSISTENT_TOKEN_LIMIT
+ kv_indices, kv_last_page_len = _cache[key]
+ output = torch.empty((q_total, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
-
- def _run_nonpersistent_path(
- q_view: torch.Tensor,
- kv_data: dict,
- qo_indptr: torch.Tensor,
- kv_indptr: torch.Tensor,
- batch_size: int,
- kv_seq_len: int,
- sm_scale: float,
- ) -> torch.Tensor:
- total_q = q_view.shape[0]
- total_kv = batch_size * kv_seq_len
- device = q_view.device
-
- page_ids = torch.arange(total_kv, dtype=torch.int32, device=device)
- last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
- output = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)
-
- if _prefer_low_overhead_path(total_kv):
- kv_tensor = kv_data["bf16"]
- q_input = q_view
- q_scale = None
- kv_scale = None
- else:
- kv_tensor, kv_scale = kv_data["fp8"]
- q_input, q_scale = _quantize_query(q_view)
-
- kv_pages = kv_tensor.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_tensor.shape[-1])
-
mla_decode_fwd(
- q_input,
- kv_pages,
- output,
- qo_indptr,
- kv_indptr,
- page_ids,
- last_page_len,
+ q_bf16, kv_4d, output,
+ qo_indptr, kv_indptr,
+ kv_indices, kv_last_page_len,
1,
- page_size=PAGE_SIZE,
- nhead_kv=NUM_KV_HEADS,
- sm_scale=sm_scale,
- q_scale=q_scale,
- kv_scale=kv_scale,
+ page_size=1, nhead_kv=NUM_KV_HEADS, sm_scale=SM_SCALE,
intra_batch_mode=False,
)
-
return output
-
-
- def _build_persistent_metadata(
- batch_size: int,
- qo_indptr: torch.Tensor,
- kv_indptr: torch.Tensor,
- q_dtype: torch.dtype,
- kv_dtype: torch.dtype,
- kv_last_page_len: torch.Tensor,
- ):
- info = get_mla_metadata_info_v1(
- batch_size,
- 1,
- NUM_HEADS,
- q_dtype,
- kv_dtype,
- is_sparse=False,
- fast_mode=False,
- num_kv_splits=PERSISTENT_SPLITS,
- intra_batch_mode=True,
- )
-
- work_meta_data, work_indptr, work_info_set, reduce_indptr, reduce_final_map, reduce_partial_map = [
- torch.empty(shape, dtype=dtype, device=qo_indptr.device) for shape, dtype in info
- ]
-
- 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=PERSISTENT_SPLITS,
- intra_batch_mode=True,
- dtype_q=q_dtype,
- dtype_kv=kv_dtype,
- )
-
- return {
- "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,
- }
-
-
- def _run_persistent_fp8_path(
- q_view: torch.Tensor,
- kv_data: dict,
- qo_indptr: torch.Tensor,
- kv_indptr: torch.Tensor,
- batch_size: int,
- kv_seq_len: int,
- sm_scale: float,
- ) -> torch.Tensor:
- total_q = q_view.shape[0]
- total_kv = batch_size * kv_seq_len
- device = q_view.device
-
- q_fp8, q_scale = _quantize_query(q_view)
- kv_fp8, kv_scale = kv_data["fp8"]
- kv_pages = kv_fp8.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_fp8.shape[-1])
-
- page_ids = torch.arange(total_kv, dtype=torch.int32, device=device)
- last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
- metadata = _build_persistent_metadata(
- batch_size,
- qo_indptr,
- kv_indptr,
- q_fp8.dtype,
- kv_fp8.dtype,
- last_page_len,
- )
-
- output = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)
-
- mla_decode_fwd(
- q_fp8,
- kv_pages,
- output,
- qo_indptr,
- kv_indptr,
- page_ids,
- last_page_len,
- 1,
- page_size=PAGE_SIZE,
- nhead_kv=NUM_KV_HEADS,
- sm_scale=sm_scale,
- logit_cap=0.0,
- num_kv_splits=PERSISTENT_SPLITS,
- q_scale=q_scale,
- kv_scale=kv_scale,
- intra_batch_mode=True,
- **metadata,
- )
-
- return output
-
-
- def custom_kernel(data: input_t) -> output_t:
- q, kv_data, qo_indptr, kv_indptr, config = data
-
- batch_size = int(config["batch_size"])
- kv_seq_len = int(config["kv_seq_len"])
- total_kv = batch_size * kv_seq_len
- sm_scale = float(config.get("sm_scale", DEFAULT_SCALE))
- q_view = q.view(-1, NUM_HEADS, QK_HEAD_DIM)
-
- if _prefer_persistent_fp8_path(total_kv):
- return _run_persistent_fp8_path(
- q_view,
- kv_data,
- qo_indptr,
- kv_indptr,
- batch_size,
- kv_seq_len,
- sm_scale,
- )
-
- return _run_nonpersistent_path(
- q_view,
- kv_data,
- qo_indptr,
- kv_indptr,
- batch_size,
- kv_seq_len,
- sm_scale,
- )
scrolls · 265 diff lines total

Best evidence level for this revision: reported

JSON