Skip to content
KernelIndex
Search⌘K

submission 755113

Nut · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1e28e9949f7a88a01e50cce3191e727ebb0eccd4bb0307f109de1417ef8e5959
license declaredunknown
license concludedunknown
authorsNut
imported2026-08-26

Techniques

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

fp4QKV_DTYPE = "mxfp4"
mmascore1 = tl.dot(q1_even, tl.trans(low1_f16)) + tl.dot(q1_odd, tl.trans(high1_f16))
num-warps = 8num_warps=8,
online-softmaxm_i_new = tl.maximum(m_i, tl.max(score, axis=1))
persistent-kernelnum_head_groups = tl.num_programs(1)
split-kdef mla_decode_mxfp4_kernel_splitk(
stages = 3num_stages=3
tile-m = 4BLOCK_M = 4

Kernel source

submission.py446 lines
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes

# QKV dtype for custom_kernel dispatch: "bf16", "fp8", or "mxfp4"
QKV_DTYPE = "mxfp4"

@triton.jit
def fast_decompress(nibble):
    x = (nibble & 7).to(tl.float32)
    val = tl.where(x < 4, x * 0.5, tl.where(x == 7, 6.0, x - 2.0))
    return tl.where((nibble & 8) != 0, -val, val)

@triton.jit
def mla_decode_mxfp4_kernel_splitk(
    q_ptr, kv_ptr, kv_scale_ptr,
    acc_even_ptr, acc_odd_ptr, m_ptr, l_ptr,
    qo_indptr, kv_indptr,
    stride_q_t, stride_q_h, stride_q_d,
    stride_kv_t, stride_kv_h, stride_kv_d,
    stride_kvs_t, stride_kvs_d,
    stride_ws_b, stride_ws_hg, stride_ws_split, stride_ws_mh,
    sm_scale, num_heads,
    q_seq_len: tl.constexpr, 
    BLOCK_KV: tl.constexpr, 
    BLOCK_M: tl.constexpr,
    BLOCK_H: tl.constexpr,
    NUM_KV_SPLITS: tl.constexpr
):
    batch_idx = tl.program_id(0)
    head_group_idx = tl.program_id(1)
    split_idx = tl.program_id(2)

    q_start = tl.load(qo_indptr + batch_idx)
    q_end = tl.load(qo_indptr + batch_idx + 1)
    q_len = q_end - q_start

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

    idx_mh = tl.arange(0, BLOCK_M * BLOCK_H)
    idx_m = idx_mh // BLOCK_H
    idx_h = head_group_idx * BLOCK_H + (idx_mh % BLOCK_H)
    mask_mh = (idx_m < q_len) & (idx_h < num_heads)

    offs_even_512 = tl.arange(0, 256) * 2
    offs_odd_512  = tl.arange(0, 256) * 2 + 1
    offs_even_64 = tl.arange(0, 32) * 2 + 512
    offs_odd_64  = tl.arange(0, 32) * 2 + 513

    q1_even = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_even_512[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
    q1_odd  = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_odd_512[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
    q2_even = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_even_64[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
    q2_odd  = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_odd_64[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
    
    q1_even = (q1_even * sm_scale).to(tl.float16)
    q1_odd  = (q1_odd * sm_scale).to(tl.float16)
    q2_even = (q2_even * sm_scale).to(tl.float16)
    q2_odd  = (q2_odd * sm_scale).to(tl.float16)

    m_i = tl.zeros([BLOCK_M * BLOCK_H], dtype=tl.float32) - float('inf')
    l_i = tl.zeros([BLOCK_M * BLOCK_H], dtype=tl.float32)
    acc_even = tl.zeros([BLOCK_M * BLOCK_H, 256], dtype=tl.float32)
    acc_odd  = tl.zeros([BLOCK_M * BLOCK_H, 256], dtype=tl.float32)

    offs_kv = tl.arange(0, BLOCK_KV)
    offs_v_uint8 = tl.arange(0, 256)
    offs_r_uint8 = tl.arange(0, 32)
    offs_scale1 = tl.arange(0, 16)
    offs_scale2 = tl.arange(16, 18)
    
    kv_buf_base = kv_ptr + offs_v_uint8[None, :] * stride_kv_d
    kv_rot_base = kv_ptr + 256 * stride_kv_d + offs_r_uint8[None, :] * stride_kv_d
    kvs_buf_base = kv_scale_ptr + offs_scale1[None, :] * stride_kvs_d
    kvs_rot_base = kv_scale_ptr + offs_scale2[None, :] * stride_kvs_d

    # SPLIT-K bounds
    chunk_size = (kv_len + NUM_KV_SPLITS - 1) // NUM_KV_SPLITS
    chunk_size = (chunk_size + BLOCK_KV - 1) // BLOCK_KV * BLOCK_KV
    start_n = split_idx * chunk_size
    end_n = start_n + chunk_size
    if end_n > kv_len:
        end_n = kv_len
    if start_n > kv_len:
        start_n = kv_len

    # Pre-compute cursor pointers outside loop!
    k1_ptrs = kv_buf_base + (kv_start + start_n + offs_kv)[:, None] * stride_kv_t
    k2_ptrs = kv_rot_base + (kv_start + start_n + offs_kv)[:, None] * stride_kv_t
    scale1_ptrs = kvs_buf_base + (kv_start + start_n + offs_kv)[:, None] * stride_kvs_t
    scale2_ptrs = kvs_rot_base + (kv_start + start_n + offs_kv)[:, None] * stride_kvs_t

    for k_step in range(start_n, end_n, BLOCK_KV):
        mask_kv = (k_step + offs_kv) < kv_len
        
        k1_uint8 = tl.load(k1_ptrs, mask=mask_kv[:, None], other=0)
        k2_uint8 = tl.load(k2_ptrs, mask=mask_kv[:, None], other=0)
        scale1_e8m0 = tl.load(scale1_ptrs, mask=mask_kv[:, None], other=127)
        scale2_e8m0 = tl.load(scale2_ptrs, mask=mask_kv[:, None], other=127)
        
        scale1_f32 = tl.exp2(scale1_e8m0.to(tl.float32) - 127.0)
        scale2_f32 = tl.exp2(scale2_e8m0.to(tl.float32) - 127.0)
        
        # Unpack via mathematical multi-polynomial to avoid array lookup
        k1_low_f32 = fast_decompress(k1_uint8 & 0x0F)
        k1_high_f32 = fast_decompress((k1_uint8 >> 4) & 0x0F)
        k2_low_f32 = fast_decompress(k2_uint8 & 0x0F)
        k2_high_f32 = fast_decompress((k2_uint8 >> 4) & 0x0F)
        
        # Scaling using view broadcasts to avoid massive registry bloat
        low1_scaled = tl.reshape(k1_low_f32, (BLOCK_KV, 16, 16)) * scale1_f32[:, :, None]
        high1_scaled = tl.reshape(k1_high_f32, (BLOCK_KV, 16, 16)) * scale1_f32[:, :, None]
        low1_f16 = tl.reshape(low1_scaled, (BLOCK_KV, 256)).to(tl.float16)
        high1_f16 = tl.reshape(high1_scaled, (BLOCK_KV, 256)).to(tl.float16)
        
        low2_scaled = tl.reshape(k2_low_f32, (BLOCK_KV, 2, 16)) * scale2_f32[:, :, None]
        high2_scaled = tl.reshape(k2_high_f32, (BLOCK_KV, 2, 16)) * scale2_f32[:, :, None]
        low2_f16 = tl.reshape(low2_scaled, (BLOCK_KV, 32)).to(tl.float16)
        high2_f16 = tl.reshape(high2_scaled, (BLOCK_KV, 32)).to(tl.float16)
        
        score1 = tl.dot(q1_even, tl.trans(low1_f16)) + tl.dot(q1_odd, tl.trans(high1_f16))
        score2 = tl.dot(q2_even, tl.trans(low2_f16)) + tl.dot(q2_odd, tl.trans(high2_f16))
        score = score1 + score2
        
        score = tl.where(mask_kv[None, :], score, float('-inf'))
        
        m_i_new = tl.maximum(m_i, tl.max(score, axis=1))
        alpha = tl.exp(m_i - m_i_new)
        p = tl.exp(score - m_i_new[:, None])
        l_i_new = alpha * l_i + tl.sum(p, axis=1)
        
        p_f16 = p.to(tl.float16)
        acc_even = acc_even * alpha[:, None] + tl.dot(p_f16, low1_f16)
        acc_odd  = acc_odd  * alpha[:, None] + tl.dot(p_f16, high1_f16)
        
        m_i = m_i_new
        l_i = l_i_new

        # Advance pointer cursors simply saving calculation overhead
        k1_ptrs += BLOCK_KV * stride_kv_t
        k2_ptrs += BLOCK_KV * stride_kv_t
        scale1_ptrs += BLOCK_KV * stride_kvs_t
        scale2_ptrs += BLOCK_KV * stride_kvs_t

    num_head_groups = tl.num_programs(1)
    offs_d = tl.arange(0, 256)
    ws_base_acc = batch_idx * stride_ws_b + head_group_idx * stride_ws_hg + split_idx * stride_ws_split + idx_mh[:, None] * stride_ws_mh + offs_d[None, :]
    ws_base_ml = batch_idx * (num_head_groups * NUM_KV_SPLITS * BLOCK_M * BLOCK_H) + head_group_idx * (NUM_KV_SPLITS * BLOCK_M * BLOCK_H) + split_idx * (BLOCK_M * BLOCK_H) + idx_mh
    
    tl.store(acc_even_ptr + ws_base_acc, acc_even, mask=mask_mh[:, None])
    tl.store(acc_odd_ptr  + ws_base_acc, acc_odd, mask=mask_mh[:, None])
    tl.store(m_ptr + ws_base_ml, m_i, mask=mask_mh)
    tl.store(l_ptr + ws_base_ml, l_i, mask=mask_mh)

@triton.jit
def mla_decode_mxfp4_reduce(
    acc_even_ptr, acc_odd_ptr, m_ptr, l_ptr, out_ptr,
    qo_indptr,
    stride_out_t, stride_out_h, stride_out_d,
    stride_ws_b, stride_ws_hg, stride_ws_split, stride_ws_mh,
    sm_scale, num_heads,
    BLOCK_M: tl.constexpr, 
    BLOCK_H: tl.constexpr,
    NUM_KV_SPLITS: tl.constexpr
):
    batch_idx = tl.program_id(0)
    head_group_idx = tl.program_id(1)
    
    q_start = tl.load(qo_indptr + batch_idx)
    q_len = tl.load(qo_indptr + batch_idx + 1) - q_start

    idx_mh = tl.arange(0, BLOCK_M * BLOCK_H)
    idx_m = idx_mh // BLOCK_H
    idx_h_local = idx_mh % BLOCK_H
    idx_h = head_group_idx * BLOCK_H + idx_h_local
    mask_mh = (idx_m < q_len) & (idx_h < num_heads)

    m_new = tl.zeros([BLOCK_M * BLOCK_H], dtype=tl.float32) - float('inf')
    l_new = tl.zeros([BLOCK_M * BLOCK_H], dtype=tl.float32)
    acc_even = tl.zeros([BLOCK_M * BLOCK_H, 256], dtype=tl.float32)
    acc_odd  = tl.zeros([BLOCK_M * BLOCK_H, 256], dtype=tl.float32)
    
    num_head_groups = tl.num_programs(1)
    offs_d = tl.arange(0, 256)
    
    for split_idx in range(NUM_KV_SPLITS):
        ws_base_acc = batch_idx * stride_ws_b + head_group_idx * stride_ws_hg + split_idx * stride_ws_split + idx_mh[:, None] * stride_ws_mh + offs_d[None, :]
        ws_base_ml = batch_idx * (num_head_groups * NUM_KV_SPLITS * BLOCK_M * BLOCK_H) + head_group_idx * (NUM_KV_SPLITS * BLOCK_M * BLOCK_H) + split_idx * (BLOCK_M * BLOCK_H) + idx_mh
        
        acc_even_k = tl.load(acc_even_ptr + ws_base_acc, mask=mask_mh[:, None], other=0.0)
        acc_odd_k  = tl.load(acc_odd_ptr  + ws_base_acc, mask=mask_mh[:, None], other=0.0)
        m_k = tl.load(m_ptr + ws_base_ml, mask=mask_mh, other=-float('inf'))
        l_k = tl.load(l_ptr + ws_base_ml, mask=mask_mh, other=0.0)
        
        m_next = tl.maximum(m_new, m_k)
        alpha_new = tl.exp(m_new - m_next)
        alpha_k = tl.exp(m_k - m_next)
        
        alpha_new = tl.where(m_new == float('-inf'), 0.0, alpha_new)
        alpha_k = tl.where(m_k == float('-inf'), 0.0, alpha_k)
        
        l_new = l_new * alpha_new + l_k * alpha_k
        
        acc_even = acc_even * alpha_new[:, None] + acc_even_k * alpha_k[:, None]
        acc_odd  = acc_odd  * alpha_new[:, None] + acc_odd_k  * alpha_k[:, None]
        
        m_new = m_next

    acc_even = acc_even / l_new[:, None]
    acc_odd  = acc_odd  / l_new[:, None]
    
    offs_even_512 = tl.arange(0, 256) * 2
    offs_odd_512  = tl.arange(0, 256) * 2 + 1
    out_ptrs_even = out_ptr + (q_start + idx_m)[:, None] * stride_out_t + idx_h[:, None] * stride_out_h + offs_even_512[None, :] * stride_out_d
    out_ptrs_odd  = out_ptr + (q_start + idx_m)[:, None] * stride_out_t + idx_h[:, None] * stride_out_h + offs_odd_512[None, :] * stride_out_d
    
    tl.store(out_ptrs_even, acc_even.to(tl.bfloat16), mask=mask_mh[:, None])
    tl.store(out_ptrs_odd,  acc_odd.to(tl.bfloat16),  mask=mask_mh[:, None])

@triton.jit
def mla_decode_mxfp4_kernel(
    q_ptr, kv_ptr, kv_scale_ptr, out_ptr,
    qo_indptr, kv_indptr,
    stride_q_t, stride_q_h, stride_q_d,
    stride_kv_t, stride_kv_h, stride_kv_d,
    stride_kvs_t, stride_kvs_d,
    stride_out_t, stride_out_h, stride_out_d,
    sm_scale, num_heads,
    q_seq_len: tl.constexpr, 
    BLOCK_KV: tl.constexpr, 
    BLOCK_M: tl.constexpr,
    BLOCK_H: tl.constexpr
):
    batch_idx = tl.program_id(0)
    head_group_idx = tl.program_id(1)

    q_start = tl.load(qo_indptr + batch_idx)
    q_len = tl.load(qo_indptr + batch_idx + 1) - q_start

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

    idx_mh = tl.arange(0, BLOCK_M * BLOCK_H)
    idx_m = idx_mh // BLOCK_H
    idx_h = head_group_idx * BLOCK_H + (idx_mh % BLOCK_H)
    mask_mh = (idx_m < q_len) & (idx_h < num_heads)

    offs_even_512 = tl.arange(0, 256) * 2
    offs_odd_512  = tl.arange(0, 256) * 2 + 1
    offs_even_64 = tl.arange(0, 32) * 2 + 512
    offs_odd_64  = tl.arange(0, 32) * 2 + 513

    q1_even = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_even_512[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
    q1_odd  = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_odd_512[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
    q2_even = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_even_64[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
    q2_odd  = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_odd_64[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
    
    q1_even = (q1_even * sm_scale).to(tl.float16)
    q1_odd  = (q1_odd * sm_scale).to(tl.float16)
    q2_even = (q2_even * sm_scale).to(tl.float16)
    q2_odd  = (q2_odd * sm_scale).to(tl.float16)

    m_i = tl.zeros([BLOCK_M * BLOCK_H], dtype=tl.float32) - float('inf')
    l_i = tl.zeros([BLOCK_M * BLOCK_H], dtype=tl.float32)
    acc_even = tl.zeros([BLOCK_M * BLOCK_H, 256], dtype=tl.float32)
    acc_odd  = tl.zeros([BLOCK_M * BLOCK_H, 256], dtype=tl.float32)

    offs_kv = tl.arange(0, BLOCK_KV)
    offs_v_uint8 = tl.arange(0, 256)
    offs_r_uint8 = tl.arange(0, 32)
    offs_scale1 = tl.arange(0, 16)
    offs_scale2 = tl.arange(16, 18)
    
    kv_buf_base = kv_ptr + offs_v_uint8[None, :] * stride_kv_d
    kv_rot_base = kv_ptr + 256 * stride_kv_d + offs_r_uint8[None, :] * stride_kv_d
    kvs_buf_base = kv_scale_ptr + offs_scale1[None, :] * stride_kvs_d
    kvs_rot_base = kv_scale_ptr + offs_scale2[None, :] * stride_kvs_d

    # Pre-compute cursor pointers
    k1_ptrs = kv_buf_base + (kv_start + offs_kv)[:, None] * stride_kv_t
    k2_ptrs = kv_rot_base + (kv_start + offs_kv)[:, None] * stride_kv_t
    scale1_ptrs = kvs_buf_base + (kv_start + offs_kv)[:, None] * stride_kvs_t
    scale2_ptrs = kvs_rot_base + (kv_start + offs_kv)[:, None] * stride_kvs_t

    for k_step in range(0, kv_len, BLOCK_KV):
        mask_kv = (k_step + offs_kv) < kv_len
        
        k1_uint8 = tl.load(k1_ptrs, mask=mask_kv[:, None], other=0)
        k2_uint8 = tl.load(k2_ptrs, mask=mask_kv[:, None], other=0)
        scale1_e8m0 = tl.load(scale1_ptrs, mask=mask_kv[:, None], other=127)
        scale2_e8m0 = tl.load(scale2_ptrs, mask=mask_kv[:, None], other=127)
        
        scale1_f32 = tl.exp2(scale1_e8m0.to(tl.float32) - 127.0)
        scale2_f32 = tl.exp2(scale2_e8m0.to(tl.float32) - 127.0)
        
        k1_low_f32 = fast_decompress(k1_uint8 & 0x0F)
        k1_high_f32 = fast_decompress((k1_uint8 >> 4) & 0x0F)
        k2_low_f32 = fast_decompress(k2_uint8 & 0x0F)
        k2_high_f32 = fast_decompress((k2_uint8 >> 4) & 0x0F)
        
        low1_scaled = tl.reshape(k1_low_f32, (BLOCK_KV, 16, 16)) * scale1_f32[:, :, None]
        high1_scaled = tl.reshape(k1_high_f32, (BLOCK_KV, 16, 16)) * scale1_f32[:, :, None]
        low1_f16 = tl.reshape(low1_scaled, (BLOCK_KV, 256)).to(tl.float16)
        high1_f16 = tl.reshape(high1_scaled, (BLOCK_KV, 256)).to(tl.float16)
        
        low2_scaled = tl.reshape(k2_low_f32, (BLOCK_KV, 2, 16)) * scale2_f32[:, :, None]
        high2_scaled = tl.reshape(k2_high_f32, (BLOCK_KV, 2, 16)) * scale2_f32[:, :, None]
        low2_f16 = tl.reshape(low2_scaled, (BLOCK_KV, 32)).to(tl.float16)
        high2_f16 = tl.reshape(high2_scaled, (BLOCK_KV, 32)).to(tl.float16)
        
        score1 = tl.dot(q1_even, tl.trans(low1_f16)) + tl.dot(q1_odd, tl.trans(high1_f16))
        score2 = tl.dot(q2_even, tl.trans(low2_f16)) + tl.dot(q2_odd, tl.trans(high2_f16))
        score = score1 + score2
        
        score = tl.where(mask_kv[None, :], score, float('-inf'))
        
        m_i_new = tl.maximum(m_i, tl.max(score, axis=1))
        alpha = tl.exp(m_i - m_i_new)
        p = tl.exp(score - m_i_new[:, None])
        l_i_new = alpha * l_i + tl.sum(p, axis=1)
        
        p_f16 = p.to(tl.float16)
        acc_even = acc_even * alpha[:, None] + tl.dot(p_f16, low1_f16)
        acc_odd  = acc_odd  * alpha[:, None] + tl.dot(p_f16, high1_f16)
        
        m_i = m_i_new
        l_i = l_i_new

        k1_ptrs += BLOCK_KV * stride_kv_t
        k2_ptrs += BLOCK_KV * stride_kv_t
        scale1_ptrs += BLOCK_KV * stride_kvs_t
        scale2_ptrs += BLOCK_KV * stride_kvs_t

    acc_even = acc_even / l_i[:, None]
    acc_odd  = acc_odd / l_i[:, None]
    
    out_ptrs_even = out_ptr + (q_start + idx_m)[:, None] * stride_out_t + idx_h[:, None] * stride_out_h + offs_even_512[None, :] * stride_out_d
    out_ptrs_odd  = out_ptr + (q_start + idx_m)[:, None] * stride_out_t + idx_h[:, None] * stride_out_h + offs_odd_512[None, :] * stride_out_d
    
    tl.store(out_ptrs_even, acc_even.to(tl.bfloat16), mask=mask_mh[:, None])
    tl.store(out_ptrs_odd,  acc_odd.to(tl.bfloat16),  mask=mask_mh[:, None])


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

    num_heads = config["num_heads"]
    sm_scale = config["sm_scale"]

    kv_buffer_mxfp4, kv_scale_mxfp4 = kv_data["mxfp4"]
    kv_buf = kv_buffer_mxfp4.view(torch.uint8)
    kv_scale = kv_scale_mxfp4.view(torch.uint8)

    batch_size = qo_indptr.shape[0] - 1
    total_q = q.shape[0]
    out = torch.empty((total_q, num_heads, 512), dtype=torch.bfloat16, device="cuda")
    
    stride_q_t, stride_q_h, stride_q_d = q.stride()
    stride_kv_t, stride_kv_h, stride_kv_d = kv_buf.stride()
    stride_kvs_t, stride_kvs_d = kv_scale.stride()
    stride_out_t, stride_out_h, stride_out_d = out.stride()
    
    BLOCK_M = 4
    BLOCK_H = 16
    num_head_groups = (num_heads + BLOCK_H - 1) // BLOCK_H
    
    num_blocks_base = batch_size * num_head_groups
    NUM_KV_SPLITS = 1
    if num_blocks_base < 64:
        NUM_KV_SPLITS = 32
    elif num_blocks_base < 128:
        NUM_KV_SPLITS = 16
    elif num_blocks_base < 256:
        NUM_KV_SPLITS = 8
        
    kv_seq_len = config["kv_seq_len"]
    max_splits_by_len = max(1, kv_seq_len // (64 * 2)) 
    NUM_KV_SPLITS = min(NUM_KV_SPLITS, max_splits_by_len)
        
    if NUM_KV_SPLITS > 1:
        # Split-K mapping
        ws_shape_acc = (batch_size, num_head_groups, NUM_KV_SPLITS, BLOCK_M * BLOCK_H, 256)
        ws_shape_ml = (batch_size, num_head_groups, NUM_KV_SPLITS, BLOCK_M * BLOCK_H)
        
        acc_even_ws = torch.empty(ws_shape_acc, dtype=torch.float32, device="cuda")
        acc_odd_ws = torch.empty(ws_shape_acc, dtype=torch.float32, device="cuda")
        m_ws = torch.empty(ws_shape_ml, dtype=torch.float32, device="cuda")
        l_ws = torch.empty(ws_shape_ml, dtype=torch.float32, device="cuda")
        
        stride_ws_b, stride_ws_hg, stride_ws_split, stride_ws_mh, _ = acc_even_ws.stride()
        
        grid_splitk = (batch_size, num_head_groups, NUM_KV_SPLITS)
        mla_decode_mxfp4_kernel_splitk[grid_splitk](
            q, kv_buf, kv_scale, 
            acc_even_ws, acc_odd_ws, m_ws, l_ws,
            qo_indptr, kv_indptr,
            stride_q_t, stride_q_h, stride_q_d,
            stride_kv_t, stride_kv_h, stride_kv_d,
            stride_kvs_t, stride_kvs_d,
            stride_ws_b, stride_ws_hg, stride_ws_split, stride_ws_mh,
            sm_scale, num_heads,
            q_seq_len=BLOCK_M,
            BLOCK_KV=64, 
            BLOCK_M=BLOCK_M,
            BLOCK_H=BLOCK_H,
            NUM_KV_SPLITS=NUM_KV_SPLITS,
            num_warps=8,
            num_stages=3
        )
        
        grid_reduce = (batch_size, num_head_groups)
        mla_decode_mxfp4_reduce[grid_reduce](
            acc_even_ws, acc_odd_ws, m_ws, l_ws, out,
            qo_indptr,
            stride_out_t, stride_out_h, stride_out_d,
            stride_ws_b, stride_ws_hg, stride_ws_split, stride_ws_mh,
            sm_scale, num_heads,
            BLOCK_M=BLOCK_M,
            BLOCK_H=BLOCK_H,
            NUM_KV_SPLITS=NUM_KV_SPLITS,
            num_warps=8
        )
    else:
        grid = (batch_size, num_head_groups)
        mla_decode_mxfp4_kernel[grid](
            q, kv_buf, kv_scale, out,
            qo_indptr, kv_indptr,
            stride_q_t, stride_q_h, stride_q_d,
            stride_kv_t, stride_kv_h, stride_kv_d,
            stride_kvs_t, stride_kvs_d,
            stride_out_t, stride_out_h, stride_out_d,
            sm_scale, num_heads,
            q_seq_len=BLOCK_M,
            BLOCK_KV=128, 
            BLOCK_M=BLOCK_M,
            BLOCK_H=BLOCK_H,
            num_warps=8,
            num_stages=3
        )

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