Skip to content
KernelIndex
Search⌘K

submission 607604

Haolin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

test_0321_1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-607604?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
640.3µs
#719 of 766
2026-03-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b6ec742b5cb0a9cdb7675d784cf61b8ab0f5638da9e7ee383ffed36ea3ba09f1
license declaredunknown
license concludedunknown
authorsHaolin
imported2026-08-26

Techniques

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

fp4kv_buffer, kv_scale = kv_data["mxfp4"]
mmaS += tl.dot(q_even_256, tl.trans(K_even_256))
online-softmaxm_new = tl.maximum(m_global, m_i)
split-kMAX_SPLIT_K,

Kernel source

test_0321_1.py265 lines
import torch
import triton
import triton.language as tl

# ---------------------------------------------------------------------------
# MXFP4 Dequantization Kernel Utils (E2M1)
# ---------------------------------------------------------------------------
@triton.jit
def dequant_nibble(n):
    sign = (n >> 3) & 1
    abs_val = n & 0x7
    
    # FIX: Remove dtype= kwarg and use .to() explicit casting
    f = tl.zeros_like(abs_val).to(tl.float32)
    
    f = tl.where(abs_val == 1, 0.5, f)
    f = tl.where(abs_val == 2, 1.0, f)
    f = tl.where(abs_val == 3, 1.5, f)
    f = tl.where(abs_val == 4, 2.0, f)
    f = tl.where(abs_val == 5, 3.0, f)
    f = tl.where(abs_val == 6, 4.0, f)
    f = tl.where(abs_val == 7, 6.0, f)
    
    return tl.where(sign, -f, f)

# ---------------------------------------------------------------------------
# Phase 1: Chunked Attention computation (Split-K)
# ---------------------------------------------------------------------------
@triton.jit
def mla_decode_mxfp4_kernel(
    q_even_ptr, q_odd_ptr,
    kv_buffer_ptr, kv_scale_ptr,
    kv_indptr,
    workspace_O_even, workspace_O_odd,
    workspace_m, workspace_l,
    stride_q_even_b, stride_q_even_h, stride_q_even_d,
    stride_q_odd_b, stride_q_odd_h, stride_q_odd_d,
    stride_kv_row, stride_scale_row,
    sm_scale,
    MAX_SPLIT_K,
    BLOCK_KV: tl.constexpr
):
    seq_idx = tl.program_id(0)
    split_idx = tl.program_id(1)

    kv_start = tl.load(kv_indptr + seq_idx)
    kv_end = tl.load(kv_indptr + seq_idx + 1)
    kv_len = kv_end - kv_start

    chunk_start = split_idx * BLOCK_KV
    if chunk_start >= kv_len:
        return

    # MQA sharing - load 16 heads. Split 288-dim loads into 256 + 32 to satisfy Triton power-of-2 rules
    heads = tl.arange(0, 16)
    cols_256 = tl.arange(0, 256)
    cols_32 = tl.arange(0, 32)

    # --- Load Q (256 chunk) ---
    q_even_ptrs_256 = q_even_ptr + seq_idx * stride_q_even_b + heads[:, None] * stride_q_even_h + cols_256[None, :] * stride_q_even_d
    q_odd_ptrs_256 = q_odd_ptr + seq_idx * stride_q_odd_b + heads[:, None] * stride_q_odd_h + cols_256[None, :] * stride_q_odd_d
    q_even_256 = tl.load(q_even_ptrs_256)
    q_odd_256 = tl.load(q_odd_ptrs_256)

    # --- Load Q (32 chunk) ---
    q_even_ptrs_32 = q_even_ptr + seq_idx * stride_q_even_b + heads[:, None] * stride_q_even_h + (256 + cols_32)[None, :] * stride_q_even_d
    q_odd_ptrs_32 = q_odd_ptr + seq_idx * stride_q_odd_b + heads[:, None] * stride_q_odd_h + (256 + cols_32)[None, :] * stride_q_odd_d
    q_even_32 = tl.load(q_even_ptrs_32)
    q_odd_32 = tl.load(q_odd_ptrs_32)

    # --- Setup KV Offsets ---
    offs_kv = tl.arange(0, BLOCK_KV)
    mask_kv = (chunk_start + offs_kv) < kv_len
    row_offsets = kv_start + chunk_start + offs_kv

    # --- Load KV and Scales (256 chunk) ---
    kv_ptrs_256 = kv_buffer_ptr + row_offsets[:, None] * stride_kv_row + cols_256[None, :]
    K_byte_256 = tl.load(kv_ptrs_256, mask=mask_kv[:, None])

    scale_cols_256 = cols_256 // 16
    scale_ptrs_256 = kv_scale_ptr + row_offsets[:, None] * stride_scale_row + scale_cols_256[None, :]
    scale_uint8_256 = tl.load(scale_ptrs_256, mask=mask_kv[:, None])
    scale_f32_256 = (scale_uint8_256.to(tl.uint32) << 23).to(tl.float32, bitcast=True)

    low_256 = K_byte_256 & 0x0F
    high_256 = (K_byte_256 >> 4) & 0x0F
    K_even_256 = (dequant_nibble(low_256) * scale_f32_256).to(tl.bfloat16)
    K_odd_256 = (dequant_nibble(high_256) * scale_f32_256).to(tl.bfloat16)

    # --- Load KV and Scales (32 chunk) ---
    kv_ptrs_32 = kv_buffer_ptr + row_offsets[:, None] * stride_kv_row + (256 + cols_32)[None, :]
    K_byte_32 = tl.load(kv_ptrs_32, mask=mask_kv[:, None])

    scale_cols_32 = (256 + cols_32) // 16
    scale_ptrs_32 = kv_scale_ptr + row_offsets[:, None] * stride_scale_row + scale_cols_32[None, :]
    scale_uint8_32 = tl.load(scale_ptrs_32, mask=mask_kv[:, None])
    scale_f32_32 = (scale_uint8_32.to(tl.uint32) << 23).to(tl.float32, bitcast=True)

    low_32 = K_byte_32 & 0x0F
    high_32 = (K_byte_32 >> 4) & 0x0F
    K_even_32 = (dequant_nibble(low_32) * scale_f32_32).to(tl.bfloat16)
    K_odd_32 = (dequant_nibble(high_32) * scale_f32_32).to(tl.bfloat16)

    # --- Compute attention scores (GEMM) ---
    S = tl.zeros((16, BLOCK_KV), dtype=tl.float32)
    S += tl.dot(q_even_256, tl.trans(K_even_256))
    S += tl.dot(q_odd_256, tl.trans(K_odd_256))
    S += tl.dot(q_even_32, tl.trans(K_even_32))
    S += tl.dot(q_odd_32, tl.trans(K_odd_32))
    S = S * sm_scale

    # --- Masking and Softmax ---
    S = tl.where(mask_kv[None, :], S, float('-inf'))
    m_i = tl.max(S, axis=1)
    p = tl.exp(S - m_i[:, None])
    l_i = tl.sum(p, axis=1)
    p_bf16 = p.to(tl.bfloat16)

    # --- Compute Output ---
    O_even = tl.dot(p_bf16, K_even_256)
    O_odd = tl.dot(p_bf16, K_odd_256)

    # Store intermediate results
    cols_out = tl.arange(0, 256)
    out_offs = seq_idx * (MAX_SPLIT_K * 16 * 256) + split_idx * (16 * 256) + heads[:, None] * 256 + cols_out[None, :]

    tl.store(workspace_O_even + out_offs, O_even)
    tl.store(workspace_O_odd + out_offs, O_odd)

    ml_offs = seq_idx * (MAX_SPLIT_K * 16) + split_idx * 16 + heads
    tl.store(workspace_m + ml_offs, m_i)
    tl.store(workspace_l + ml_offs, l_i)

# ---------------------------------------------------------------------------
# Phase 2: Reduction Kernel
# ---------------------------------------------------------------------------
@triton.jit
def mla_reduce_kernel(
    workspace_O_even, workspace_O_odd, workspace_m, workspace_l,
    out_even, out_odd,
    kv_indptr,
    MAX_SPLIT_K, BLOCK_KV: tl.constexpr
):
    seq_idx = tl.program_id(0)
    head_idx = tl.program_id(1)

    kv_start = tl.load(kv_indptr + seq_idx)
    kv_end = tl.load(kv_indptr + seq_idx + 1)
    kv_len = kv_end - kv_start

    if kv_len == 0:
        return

    num_splits = (kv_len + BLOCK_KV - 1) // BLOCK_KV
    m_global = float('-inf')
    l_global = 0.0

    ml_base = workspace_m + seq_idx * (MAX_SPLIT_K * 16) + head_idx

    # Pass 1: Max m and sum l over all sequence splits
    for i in range(num_splits):
        m_i = tl.load(ml_base + i * 16)
        l_i = tl.load(workspace_l + seq_idx * (MAX_SPLIT_K * 16) + i * 16 + head_idx)
        m_new = tl.maximum(m_global, m_i)
        l_new = l_global * tl.exp(m_global - m_new) + l_i * tl.exp(m_i - m_new)
        m_global = m_new
        l_global = l_new

    # Pass 2: Combine partial Attention Outputs
    acc_even = tl.zeros((256,), dtype=tl.float32)
    acc_odd = tl.zeros((256,), dtype=tl.float32)
    cols = tl.arange(0, 256)
    
    O_base_even = workspace_O_even + seq_idx * (MAX_SPLIT_K * 16 * 256) + head_idx * 256 + cols
    O_base_odd = workspace_O_odd + seq_idx * (MAX_SPLIT_K * 16 * 256) + head_idx * 256 + cols

    for i in range(num_splits):
        m_i = tl.load(ml_base + i * 16)
        weight = tl.exp(m_i - m_global)

        O_even_i = tl.load(O_base_even + i * (16 * 256))
        O_odd_i = tl.load(O_base_odd + i * (16 * 256))

        acc_even += O_even_i * weight
        acc_odd += O_odd_i * weight

    acc_even = acc_even / l_global
    acc_odd = acc_odd / l_global

    out_even_ptr = out_even + seq_idx * (16 * 256) + head_idx * 256 + cols
    out_odd_ptr = out_odd + seq_idx * (16 * 256) + head_idx * 256 + cols

    tl.store(out_even_ptr, acc_even.to(tl.bfloat16))
    tl.store(out_odd_ptr, acc_odd.to(tl.bfloat16))

# ---------------------------------------------------------------------------
# Host Function Wrapper
# ---------------------------------------------------------------------------
def custom_kernel(data) -> torch.Tensor:
    q, kv_data, qo_indptr, kv_indptr, config = data
    batch_size = config["batch_size"]
    
    # 1. Extract MXFP4 cache 
    kv_buffer, kv_scale = kv_data["mxfp4"]
    
    # View uint8 cache flat
    kv_buffer = kv_buffer.view(torch.uint8).view(-1, 288)
    
    # E8M0 Scale normalization
    kv_scale = kv_scale.view(torch.uint8)
    if kv_scale.dim() == 3:
        kv_scale = kv_scale.view(kv_scale.shape[0], kv_scale.shape[-1])

    # 2. Re-stride Query
    q_even = q[:, :, 0::2].contiguous()
    q_odd = q[:, :, 1::2].contiguous()

    # 3. Dynamic Flash-Decoding Sizing
    max_kv_len = 0
    if batch_size > 0:
        max_kv_len = int((kv_indptr[1:] - kv_indptr[:-1]).max().item())

    BLOCK_KV = 64
    MAX_SPLIT_K = max(1, (max_kv_len + BLOCK_KV - 1) // BLOCK_KV)

    # 4. Global Workspace Allocation
    workspace_O_even = torch.empty((batch_size, MAX_SPLIT_K, 16, 256), dtype=torch.float32, device='cuda')
    workspace_O_odd = torch.empty((batch_size, MAX_SPLIT_K, 16, 256), dtype=torch.float32, device='cuda')
    workspace_m = torch.empty((batch_size, MAX_SPLIT_K, 16), dtype=torch.float32, device='cuda')
    workspace_l = torch.empty((batch_size, MAX_SPLIT_K, 16), dtype=torch.float32, device='cuda')

    out_even = torch.empty((batch_size, 16, 256), dtype=torch.bfloat16, device='cuda')
    out_odd = torch.empty((batch_size, 16, 256), dtype=torch.bfloat16, device='cuda')

    # 5. Dispatch
    if batch_size > 0 and max_kv_len > 0:
        grid_1 = (batch_size, MAX_SPLIT_K)
        mla_decode_mxfp4_kernel[grid_1](
            q_even, q_odd,
            kv_buffer, kv_scale,
            kv_indptr,
            workspace_O_even, workspace_O_odd,
            workspace_m, workspace_l,
            q_even.stride(0), q_even.stride(1), q_even.stride(2),
            q_odd.stride(0), q_odd.stride(1), q_odd.stride(2),
            kv_buffer.stride(0), kv_scale.stride(0),
            config["sm_scale"],
            MAX_SPLIT_K,
            BLOCK_KV=BLOCK_KV
        )

        grid_2 = (batch_size, 16)
        mla_reduce_kernel[grid_2](
            workspace_O_even, workspace_O_odd, workspace_m, workspace_l,
            out_even, out_odd,
            kv_indptr,
            MAX_SPLIT_K, BLOCK_KV
        )

    # 6. Reconstruct sequence dimensions seamlessly
    out = torch.zeros((batch_size, 16, 512), dtype=torch.bfloat16, device='cuda')
    out[:, :, 0::2] = out_even
    out[:, :, 1::2] = out_odd

    return out
scrolls · 265 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