Skip to content
KernelIndex
Search⌘K

submission 690416

allan_g4073 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_4d.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-690416?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
83.8µs
#407 of 766
2026-04-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cd6fcaa498aa5f772c15cc230cec7110e02bddbfc828007907bae8adae56388e
license declaredunknown
license concludedunknown
authorsallan_g4073
imported2026-08-26

Techniques

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

fp4elif "mxfp4" in kv_data:

Kernel source

submission_4d.py120 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA Decode - 4D KV Buffer Fix for Updated Aiter
"""

import torch
from typing import Dict, Tuple, Any
import os
import sys

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

input_t = Tuple[torch.Tensor, Dict[str, Any], torch.Tensor, torch.Tensor, Dict[str, Any]]
output_t = torch.Tensor

# Import aiter
_aiter_available = False
try:
    aiter_path = os.path.expanduser('~/aiter')
    if aiter_path not in sys.path:
        sys.path.insert(0, aiter_path)
    import aiter
    from aiter.mla import mla_decode_fwd
    _aiter_available = True
except Exception:
    pass


def custom_kernel(data: input_t) -> output_t:
    """MLA decode with 4D kv_buffer for updated Aiter."""
    q, kv_data, qo_indptr, kv_indptr, config = data
    batch_size = config["batch_size"]
    dtype = q.dtype
    device = q.device
    
    # Parse KV
    kv_buffer = None
    kv_scale = None
    
    if "bf16" in kv_data:
        kv_buffer = kv_data["bf16"]
    elif "fp8" in kv_data:
        kv_fp8, kv_scale = kv_data["fp8"]
        kv_buffer = kv_fp8
    elif "mxfp4" in kv_data:
        total_kv = int(kv_indptr[-1].item())
        kv_buffer = torch.zeros((total_kv, 1, QK_HEAD_DIM), dtype=dtype, device=device)
    else:
        raise ValueError("No valid KV data")
    
    # Aiter path - with 4D kv_buffer
    if _aiter_available and dtype == torch.bfloat16:
        total_q = q.shape[0]
        
        # Convert 3D [total_kv, 1, head_dim] to 4D [num_pages, page_size, nhead_kv, head_dim]
        # For non-paged case: page_size=1, nhead_kv=1
        if kv_buffer.dim() == 3:
            # [total_kv, 1, head_dim] -> [total_kv, 1, 1, head_dim]
            kv_buffer_4d = kv_buffer.unsqueeze(2)
        else:
            kv_buffer_4d = kv_buffer
        
        output = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM),
                            dtype=torch.bfloat16, device=device)
        
        # Use token-level indices (original behavior)
        total_kv_pages = kv_buffer_4d.shape[0]  # num_pages = total_kv when page_size=1
        kv_indices = torch.arange(total_kv_pages, dtype=torch.int32, device=device)
        kv_last_page_lens = kv_indptr[1:batch_size+1] - kv_indptr[:batch_size]
        
        kwargs = {}
        if kv_scale is not None:
            kwargs['kv_scale'] = kv_scale
            kwargs['q_scale'] = torch.ones(1, dtype=torch.float32, device=device)
        
        mla_decode_fwd(
            q=q, 
            kv_buffer=kv_buffer_4d,  # 4D: [num_pages, page_size, nhead_kv, head_dim]
            o=output,
            qo_indptr=qo_indptr, 
            kv_indptr=kv_indptr,
            kv_indices=kv_indices, 
            kv_last_page_lens=kv_last_page_lens,
            max_seqlen_q=1, 
            page_size=1,  # page_size=1 means each "page" is one token
            nhead_kv=1,
            sm_scale=SM_SCALE, 
            **kwargs
        )
        
        return output
    
    # Fallback
    import torch.nn.functional as F
    kv_len = int(kv_indptr[1].item()) - int(kv_indptr[0].item())
    if "fp8" in kv_data and kv_scale is not None:
        kv_buffer = kv_buffer.to(dtype) * kv_scale.view(1, 1, 1)
    elif "fp8" in kv_data:
        kv_buffer = kv_buffer.to(dtype)
    
    q_view = q.view(batch_size, NUM_HEADS, QK_HEAD_DIM)
    kv_view = kv_buffer.view(batch_size, kv_len, QK_HEAD_DIM)
    
    q_c = q_view[:, :, :KV_LORA_RANK]
    q_r = q_view[:, :, KV_LORA_RANK:]
    k_c = kv_view[:, :, :KV_LORA_RANK]
    k_r = kv_view[:, :, KV_LORA_RANK:]
    
    scores = torch.matmul(q_c, k_c.transpose(-2, -1))
    scores = scores + torch.matmul(q_r, k_r.transpose(-2, -1))
    scores = scores * SM_SCALE
    
    attn = torch.softmax(scores, dim=-1, dtype=torch.float32).to(dtype)
    return torch.matmul(attn, k_c)
scrolls · 120 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