Skip to content
KernelIndex
Search⌘K

submission 688800

SSS · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:07e93b5c535a52737f4863379c1142cec61fbb010ceafa965e88e36605c9c411
license declaredunknown
license concludedunknown
authorsSSS
imported2026-08-15

Techniques

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

autotune@triton.autotune(
mmaqk = tl.dot(q_lora_fp8, tl.trans(kv_lora_fp8))
num-warps = 4triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64}, num_stages=2, num_warps=4),
online-softmaxm_i_new = tl.maximum(m_i, m_ij)
split-kuse_split_kv = base_blocks < target_blocks
stages = 2triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64}, num_stages=2, num_warps=4),

Kernel source

vsota.py939 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
# vsota
"""
```
seed: 4217; qseqlen: 1; kvseqlen: 1024; batchsize: 4
 ⏱ 20.7 ± 0.02 µs
 ⚡ 19.9 µs 🐌 25.3 µs

seed: 4220; qseqlen: 1; kvseqlen: 8192; batchsize: 4
 ⏱ 35.1 ± 0.04 µs
 ⚡ 34.3 µs 🐌 39.2 µs

seed: 5412; qseqlen: 1; kvseqlen: 1024; batchsize: 32
 ⏱ 22.6 ± 0.02 µs
 ⚡ 21.8 µs 🐌 28.0 µs

seed: 5415; qseqlen: 1; kvseqlen: 8192; batchsize: 32
 ⏱ 65.3 ± 0.07 µs
 ⚡ 64.2 µs 🐌 68.9 µs

seed: 1357; qseqlen: 1; kvseqlen: 1024; batchsize: 64
 ⏱ 27.7 ± 0.03 µs
 ⚡ 26.8 µs 🐌 32.4 µs

seed: 1360; qseqlen: 1; kvseqlen: 8192; batchsize: 64
 ⏱ 113 ± 0.1 µs
 ⚡ 111 µs 🐌 116 µs

seed: 9823; qseqlen: 1; kvseqlen: 1024; batchsize: 256
 ⏱ 60.2 ± 0.06 µs
 ⚡ 59.0 µs 🐌 63.3 µs

seed: 9826; qseqlen: 1; kvseqlen: 8192; batchsize: 256
 ⏱ 303 ± 0.3 µs
 ⚡ 294 µs 🐌 316 µs
``` 
```
"""
"""
## Benchmarks:
  # To be filled by benchmark run
"""

import torch
import triton
import triton.language as tl
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes

FP8_DTYPE = aiter_dtypes.fp8
SHORT_KV_MAX = 1024
SHORT_KV_TARGET_BLOCKS_SMALL = 128
SHORT_KV_TARGET_BLOCKS = 256
SHORT_KV_MAX_SPLITS = 32
SHORT_KV_SMALL_REDUCE_MAX_TOTAL_TH = 64
SHORT_KV_SMALL_REDUCE_BLOCK_D = 32

def custom_kernel(data: input_t) -> output_t:
    return custom_kernel_v41(data)

# ==============================================================================
# No-Split Kernel 
# ==============================================================================
@triton.autotune(
    configs=[
        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64},  num_stages=2, num_warps=4),
        triton.Config({'BLOCK_M': 32, 'BLOCK_N': 64},  num_stages=2, num_warps=4),
        triton.Config({'BLOCK_M': 64, 'BLOCK_N': 64},  num_stages=2, num_warps=8),
        triton.Config({'BLOCK_M': 32, 'BLOCK_N': 128}, num_stages=2, num_warps=8),
        triton.Config({'BLOCK_M': 64, 'BLOCK_N': 128}, num_stages=2, num_warps=8),
    ],
    key=['KV_LORA_RANK', 'QK_ROPE_DIM'],
)
@triton.jit
def _mla_v11_no_split(
    Q, KV, Out, qo_indptr, kv_indptr, kv_scale_ptr, sm_scale,
    stride_q_t, stride_q_h, stride_q_d,
    stride_kv_t, stride_kv_d,
    stride_o_t, stride_o_h, stride_o_d,
    num_heads: tl.constexpr, KV_LORA_RANK: tl.constexpr, QK_ROPE_DIM: tl.constexpr, v_head_dim: tl.constexpr,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
    batch_idx   = tl.program_id(0)
    m_block_idx = tl.program_id(1)
    
    q_start  = tl.load(qo_indptr + batch_idx)
    q_end    = tl.load(qo_indptr + batch_idx + 1)
    kv_start = tl.load(kv_indptr + batch_idx)
    kv_end   = tl.load(kv_indptr + batch_idx + 1)
    
    seq_kv   = kv_end - kv_start
    actual_m = (q_end - q_start) * num_heads
    if actual_m <= 0 or seq_kv <= 0: return
    if m_block_idx * BLOCK_M >= actual_m: return

    offs_m = m_block_idx * BLOCK_M + tl.arange(0, BLOCK_M)
    mask_m = offs_m < actual_m
    sq_idx = q_start + (offs_m // num_heads)
    h_idx  = offs_m % num_heads
    q_ptrs_base = Q + sq_idx * stride_q_t + h_idx * stride_q_h

    kv_scale_val   = tl.load(kv_scale_ptr)
    combined_scale = kv_scale_val * sm_scale

    offs_d_lora = tl.arange(0, KV_LORA_RANK)
    q_lora = tl.load(q_ptrs_base[:, None] + offs_d_lora[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)
    offs_d_rope = KV_LORA_RANK + tl.arange(0, QK_ROPE_DIM)
    q_rope = tl.load(q_ptrs_base[:, None] + offs_d_rope[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)

    # -------------------------------------------------------------
    # Quantize Q to FP8 on-the-fly to unlock MI355X FP8 Tensor Cores
    # -------------------------------------------------------------
    fp8_dtype = KV.dtype.element_ty
    
    q_lora_amax = tl.maximum(tl.max(tl.abs(q_lora), axis=1), 1e-12)
    q_lora_scale = q_lora_amax / 240.0
    q_lora_fp8 = (q_lora / q_lora_scale[:, None]).to(fp8_dtype)

    q_rope_amax = tl.maximum(tl.max(tl.abs(q_rope), axis=1), 1e-12)
    q_rope_scale = q_rope_amax / 240.0
    q_rope_fp8 = (q_rope / q_rope_scale[:, None]).to(fp8_dtype)

    m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float('inf')
    l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
    acc = tl.zeros([BLOCK_M, KV_LORA_RANK], dtype=tl.float32)

    kv_base = KV + kv_start * stride_kv_t
    offs_n = tl.arange(0, BLOCK_N)
    kv_lora_ptrs = kv_base + offs_n[:, None] * stride_kv_t + offs_d_lora[None, :] * stride_kv_d
    kv_rope_ptrs = kv_base + offs_n[:, None] * stride_kv_t + offs_d_rope[None, :] * stride_kv_d

    for start_n in range(0, seq_kv, BLOCK_N):
        curr_n = start_n + offs_n
        mask_n = curr_n < seq_kv
        
        kv_lora_fp8 = tl.load(kv_lora_ptrs, mask=mask_n[:, None], other=0.0)
        kv_rope_fp8 = tl.load(kv_rope_ptrs, mask=mask_n[:, None], other=0.0)

        # QK Matrix Multiply via FP8 Tensor Cores (4x faster than BF16)
        qk = tl.dot(q_lora_fp8, tl.trans(kv_lora_fp8))
        qk = qk.to(tl.float32) * q_lora_scale[:, None]
        
        qk_rope = tl.dot(q_rope_fp8, tl.trans(kv_rope_fp8))
        qk_rope = qk_rope.to(tl.float32) * q_rope_scale[:, None]

        qk += qk_rope
        qk  = qk * combined_scale

        qk = tl.where(mask_m[:, None] & mask_n[None, :], qk, float('-inf'))
        m_ij    = tl.max(qk, 1)
        m_i_new = tl.maximum(m_i, m_ij)
        alpha   = tl.exp(m_i - m_i_new)
        p       = tl.exp(qk - m_i_new[:, None])
        l_i_new = alpha * l_i + tl.sum(p, 1)

        acc = acc * alpha[:, None]
        kv_lora_bf16 = kv_lora_fp8.to(tl.bfloat16)
        acc += tl.dot(p.to(tl.bfloat16), kv_lora_bf16)
        m_i = m_i_new
        l_i = l_i_new

        kv_lora_ptrs += BLOCK_N * stride_kv_t
        kv_rope_ptrs += BLOCK_N * stride_kv_t

    acc = acc / l_i[:, None]
    acc = acc * kv_scale_val

    offs_d_v = tl.arange(0, KV_LORA_RANK)
    mask_v   = offs_d_v < v_head_dim
    out_ptrs = Out + sq_idx[:, None]*stride_o_t + h_idx[:, None]*stride_o_h + offs_d_v[None, :]*stride_o_d
    tl.store(out_ptrs, acc.to(Out.dtype.element_ty), mask=mask_m[:, None] & mask_v[None, :])


# ==============================================================================
# Split-KV Kernel
# ==============================================================================
@triton.autotune(
    configs=[
        # BLOCK_N=32: kv_lora tile (32,512) = 16+32=48 VGPRs vs (64,512) = 32+64=96 VGPRs
        # Target: total VGPR ~120 → Occupancy=2 (2 blocks/CU)
        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32},  num_stages=2, num_warps=4),
        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32},  num_stages=1, num_warps=4),
        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64},  num_stages=2, num_warps=4),
    ],
    key=['KV_LORA_RANK', 'QK_ROPE_DIM', 'route_bucket'],
)
@triton.jit
def _mla_v11_split(
    Q, KV, Partial_O, Partial_LSE,
    qo_indptr, kv_indptr, kv_scale_ptr, sm_scale, route_bucket,
    stride_q_t, stride_q_h, stride_q_d,
    stride_kv_t, stride_kv_d,
    stride_po_s, stride_po_t, stride_po_h, stride_po_d,
    stride_plse_s, stride_plse_t, stride_plse_h,
    num_heads: tl.constexpr, KV_LORA_RANK: tl.constexpr, QK_ROPE_DIM: tl.constexpr,
    SPLIT_SIZE: tl.constexpr, NUM_KV_SPLITS: tl.constexpr,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
    batch_idx   = tl.program_id(0)
    m_block_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)
    kv_start = tl.load(kv_indptr + batch_idx)
    kv_end   = tl.load(kv_indptr + batch_idx + 1)
    seq_kv   = kv_end - kv_start
    actual_m = (q_end - q_start) * num_heads

    if actual_m <= 0 or seq_kv <= 0: return
    if m_block_idx * BLOCK_M >= actual_m: return

    offs_m = m_block_idx * BLOCK_M + tl.arange(0, BLOCK_M)
    mask_m = offs_m < actual_m
    sq_idx = q_start + (offs_m // num_heads)
    h_idx  = offs_m % num_heads

    split_start = split_idx * SPLIT_SIZE
    split_end   = tl.minimum(split_start + SPLIT_SIZE, seq_kv)
    split_len   = split_end - split_start

    if split_start >= seq_kv:
        lse_ptrs = Partial_LSE + split_idx*stride_plse_s + sq_idx*stride_plse_t + h_idx*stride_plse_h
        tl.store(lse_ptrs, float('-inf'), mask=mask_m)
        return

    q_ptrs_base = Q + sq_idx * stride_q_t + h_idx * stride_q_h
    kv_scale_val   = tl.load(kv_scale_ptr)
    combined_scale = kv_scale_val * sm_scale

    offs_d_lora = tl.arange(0, KV_LORA_RANK)
    q_lora = tl.load(q_ptrs_base[:, None] + offs_d_lora[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)
    offs_d_rope = KV_LORA_RANK + tl.arange(0, QK_ROPE_DIM)
    q_rope = tl.load(q_ptrs_base[:, None] + offs_d_rope[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)

    # -------------------------------------------------------------
    # Quantize Q to FP8 on-the-fly to unlock MI355X FP8 Tensor Cores
    # -------------------------------------------------------------
    fp8_dtype = KV.dtype.element_ty
    
    q_lora_amax = tl.maximum(tl.max(tl.abs(q_lora), axis=1), 1e-12)
    q_lora_scale = q_lora_amax / 240.0
    q_lora_fp8 = (q_lora / q_lora_scale[:, None]).to(fp8_dtype)

    q_rope_amax = tl.maximum(tl.max(tl.abs(q_rope), axis=1), 1e-12)
    q_rope_scale = q_rope_amax / 240.0
    q_rope_fp8 = (q_rope / q_rope_scale[:, None]).to(fp8_dtype)

    m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float('inf')
    l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
    acc = tl.zeros([BLOCK_M, KV_LORA_RANK], dtype=tl.float32)

    kv_base_ptr = KV + (kv_start + split_start) * stride_kv_t
    offs_n = tl.arange(0, BLOCK_N)
    kv_lora_ptrs = kv_base_ptr + offs_n[:, None] * stride_kv_t + offs_d_lora[None, :] * stride_kv_d
    kv_rope_ptrs = kv_base_ptr + offs_n[:, None] * stride_kv_t + offs_d_rope[None, :] * stride_kv_d

    for start_n in range(0, SPLIT_SIZE, BLOCK_N):
        curr_n = start_n + offs_n
        mask_n = curr_n < split_len
        
        kv_lora_fp8 = tl.load(kv_lora_ptrs, mask=mask_n[:, None], other=0.0)
        kv_rope_fp8 = tl.load(kv_rope_ptrs, mask=mask_n[:, None], other=0.0)

        # QK Matrix Multiply via FP8 Tensor Cores
        qk = tl.dot(q_lora_fp8, tl.trans(kv_lora_fp8))
        qk = qk.to(tl.float32) * q_lora_scale[:, None]
        
        qk_rope = tl.dot(q_rope_fp8, tl.trans(kv_rope_fp8))
        qk_rope = qk_rope.to(tl.float32) * q_rope_scale[:, None]

        qk += qk_rope
        qk  = qk * combined_scale

        qk = tl.where(mask_m[:, None] & mask_n[None, :], qk, float('-inf'))
        m_ij    = tl.max(qk, 1)
        m_i_new = tl.maximum(m_i, m_ij)
        alpha   = tl.exp(m_i - m_i_new)
        p       = tl.exp(qk - m_i_new[:, None])
        l_i_new = alpha * l_i + tl.sum(p, 1)

        acc = acc * alpha[:, None]
        # Inline conversion: avoid named temporary to hint compiler for shorter liveness
        acc += tl.dot(p.to(tl.bfloat16), kv_lora_fp8.to(tl.bfloat16))

        m_i = m_i_new
        l_i = l_i_new

        kv_lora_ptrs += BLOCK_N * stride_kv_t
        kv_rope_ptrs += BLOCK_N * stride_kv_t

    acc = acc / l_i[:, None]
    po_ptrs = Partial_O + split_idx*stride_po_s + sq_idx[:, None]*stride_po_t + h_idx[:, None]*stride_po_h + offs_d_lora[None, :]*stride_po_d
    tl.store(po_ptrs, acc, mask=mask_m[:, None])

    lse = m_i + tl.log(l_i)
    lse_ptrs = Partial_LSE + split_idx*stride_plse_s + sq_idx*stride_plse_t + h_idx*stride_plse_h
    tl.store(lse_ptrs, lse, mask=mask_m)


@triton.autotune(
    configs=[
        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64},  num_stages=1, num_warps=4),
        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64},  num_stages=2, num_warps=4),
        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 128}, num_stages=2, num_warps=4),
    ],
    key=['KV_LORA_RANK', 'QK_ROPE_DIM', 'route_bucket'],
)
@triton.jit
def _mla_v28_short_no_split(
    Q, KV, Out, qo_indptr, kv_indptr, kv_scale_ptr, sm_scale, route_bucket,
    stride_q_t, stride_q_h, stride_q_d,
    stride_kv_t, stride_kv_d,
    stride_o_t, stride_o_h, stride_o_d,
    num_heads: tl.constexpr, KV_LORA_RANK: tl.constexpr, QK_ROPE_DIM: tl.constexpr, v_head_dim: tl.constexpr,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
    batch_idx   = tl.program_id(0)
    m_block_idx = tl.program_id(1)

    q_start  = tl.load(qo_indptr + batch_idx)
    q_end    = tl.load(qo_indptr + batch_idx + 1)
    kv_start = tl.load(kv_indptr + batch_idx)
    kv_end   = tl.load(kv_indptr + batch_idx + 1)

    seq_kv   = kv_end - kv_start
    actual_m = (q_end - q_start) * num_heads
    if actual_m <= 0 or seq_kv <= 0: return
    if m_block_idx * BLOCK_M >= actual_m: return

    offs_m = m_block_idx * BLOCK_M + tl.arange(0, BLOCK_M)
    mask_m = offs_m < actual_m
    sq_idx = q_start + (offs_m // num_heads)
    h_idx  = offs_m % num_heads
    q_ptrs_base = Q + sq_idx * stride_q_t + h_idx * stride_q_h

    kv_scale_val = tl.load(kv_scale_ptr)
    combined_scale = kv_scale_val * sm_scale

    offs_d_lora = tl.arange(0, KV_LORA_RANK)
    q_lora = tl.load(q_ptrs_base[:, None] + offs_d_lora[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)
    offs_d_rope = KV_LORA_RANK + tl.arange(0, QK_ROPE_DIM)
    q_rope = tl.load(q_ptrs_base[:, None] + offs_d_rope[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)

    fp8_dtype = KV.dtype.element_ty

    q_lora_amax = tl.maximum(tl.max(tl.abs(q_lora), axis=1), 1e-12)
    q_lora_scale = q_lora_amax / 240.0
    q_lora_fp8 = (q_lora / q_lora_scale[:, None]).to(fp8_dtype)

    q_rope_amax = tl.maximum(tl.max(tl.abs(q_rope), axis=1), 1e-12)
    q_rope_scale = q_rope_amax / 240.0
    q_rope_fp8 = (q_rope / q_rope_scale[:, None]).to(fp8_dtype)

    m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float('inf')
    l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
    acc = tl.zeros([BLOCK_M, KV_LORA_RANK], dtype=tl.float32)

    kv_base = KV + kv_start * stride_kv_t
    offs_n = tl.arange(0, BLOCK_N)
    kv_lora_ptrs = kv_base + offs_n[:, None] * stride_kv_t + offs_d_lora[None, :] * stride_kv_d
    kv_rope_ptrs = kv_base + offs_n[:, None] * stride_kv_t + offs_d_rope[None, :] * stride_kv_d

    for start_n in range(0, seq_kv, BLOCK_N):
        curr_n = start_n + offs_n
        mask_n = curr_n < seq_kv

        kv_lora_fp8 = tl.load(kv_lora_ptrs, mask=mask_n[:, None], other=0.0)
        kv_rope_fp8 = tl.load(kv_rope_ptrs, mask=mask_n[:, None], other=0.0)

        qk = tl.dot(q_lora_fp8, tl.trans(kv_lora_fp8))
        qk = qk.to(tl.float32) * q_lora_scale[:, None]

        qk_rope = tl.dot(q_rope_fp8, tl.trans(kv_rope_fp8))
        qk_rope = qk_rope.to(tl.float32) * q_rope_scale[:, None]

        qk += qk_rope
        qk = qk * combined_scale

        qk = tl.where(mask_m[:, None] & mask_n[None, :], qk, float('-inf'))
        m_ij = tl.max(qk, 1)
        m_i_new = tl.maximum(m_i, m_ij)
        alpha = tl.exp(m_i - m_i_new)
        p = tl.exp(qk - m_i_new[:, None])
        l_i_new = alpha * l_i + tl.sum(p, 1)

        acc = acc * alpha[:, None]
        acc += tl.dot(p.to(tl.bfloat16), kv_lora_fp8.to(tl.bfloat16))
        m_i = m_i_new
        l_i = l_i_new

        kv_lora_ptrs += BLOCK_N * stride_kv_t
        kv_rope_ptrs += BLOCK_N * stride_kv_t

    acc = acc / l_i[:, None]
    acc = acc * kv_scale_val

    offs_d_v = tl.arange(0, KV_LORA_RANK)
    mask_v = offs_d_v < v_head_dim
    out_ptrs = Out + sq_idx[:, None]*stride_o_t + h_idx[:, None]*stride_o_h + offs_d_v[None, :]*stride_o_d
    tl.store(out_ptrs, acc.to(Out.dtype.element_ty), mask=mask_m[:, None] & mask_v[None, :])


@triton.autotune(
    configs=[
        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32},  num_stages=1, num_warps=4),
        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32},  num_stages=2, num_warps=4),
        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64},  num_stages=1, num_warps=4),
        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64},  num_stages=2, num_warps=4),
        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 128}, num_stages=2, num_warps=4),
    ],
    key=['KV_LORA_RANK', 'QK_ROPE_DIM', 'route_bucket'],
)
@triton.jit
def _mla_v28_short_split(
    Q, KV, Partial_O, Partial_LSE,
    qo_indptr, kv_indptr, kv_scale_ptr, sm_scale, route_bucket,
    stride_q_t, stride_q_h, stride_q_d,
    stride_kv_t, stride_kv_d,
    stride_po_s, stride_po_t, stride_po_h, stride_po_d,
    stride_plse_s, stride_plse_t, stride_plse_h,
    num_heads: tl.constexpr, KV_LORA_RANK: tl.constexpr, QK_ROPE_DIM: tl.constexpr,
    SPLIT_SIZE: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
    batch_idx   = tl.program_id(0)
    m_block_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)
    kv_start = tl.load(kv_indptr + batch_idx)
    kv_end   = tl.load(kv_indptr + batch_idx + 1)
    seq_kv   = kv_end - kv_start
    actual_m = (q_end - q_start) * num_heads

    if actual_m <= 0 or seq_kv <= 0: return
    if m_block_idx * BLOCK_M >= actual_m: return

    offs_m = m_block_idx * BLOCK_M + tl.arange(0, BLOCK_M)
    mask_m = offs_m < actual_m
    sq_idx = q_start + (offs_m // num_heads)
    h_idx  = offs_m % num_heads

    split_start = split_idx * SPLIT_SIZE
    split_end   = tl.minimum(split_start + SPLIT_SIZE, seq_kv)
    split_len   = split_end - split_start

    if split_start >= seq_kv:
        lse_ptrs = Partial_LSE + split_idx*stride_plse_s + sq_idx*stride_plse_t + h_idx*stride_plse_h
        tl.store(lse_ptrs, float('-inf'), mask=mask_m)
        return

    q_ptrs_base = Q + sq_idx * stride_q_t + h_idx * stride_q_h
    kv_scale_val = tl.load(kv_scale_ptr)
    combined_scale = kv_scale_val * sm_scale

    offs_d_lora = tl.arange(0, KV_LORA_RANK)
    q_lora = tl.load(q_ptrs_base[:, None] + offs_d_lora[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)
    offs_d_rope = KV_LORA_RANK + tl.arange(0, QK_ROPE_DIM)
    q_rope = tl.load(q_ptrs_base[:, None] + offs_d_rope[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)

    fp8_dtype = KV.dtype.element_ty

    q_lora_amax = tl.maximum(tl.max(tl.abs(q_lora), axis=1), 1e-12)
    q_lora_scale = q_lora_amax / 240.0
    q_lora_fp8 = (q_lora / q_lora_scale[:, None]).to(fp8_dtype)

    q_rope_amax = tl.maximum(tl.max(tl.abs(q_rope), axis=1), 1e-12)
    q_rope_scale = q_rope_amax / 240.0
    q_rope_fp8 = (q_rope / q_rope_scale[:, None]).to(fp8_dtype)

    m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float('inf')
    l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
    acc = tl.zeros([BLOCK_M, KV_LORA_RANK], dtype=tl.float32)

    kv_base_ptr = KV + (kv_start + split_start) * stride_kv_t
    offs_n = tl.arange(0, BLOCK_N)
    kv_lora_ptrs = kv_base_ptr + offs_n[:, None] * stride_kv_t + offs_d_lora[None, :] * stride_kv_d
    kv_rope_ptrs = kv_base_ptr + offs_n[:, None] * stride_kv_t + offs_d_rope[None, :] * stride_kv_d

    for start_n in range(0, SPLIT_SIZE, BLOCK_N):
        curr_n = start_n + offs_n
        mask_n = curr_n < split_len

        kv_lora_fp8 = tl.load(kv_lora_ptrs, mask=mask_n[:, None], other=0.0)
        kv_rope_fp8 = tl.load(kv_rope_ptrs, mask=mask_n[:, None], other=0.0)

        qk = tl.dot(q_lora_fp8, tl.trans(kv_lora_fp8))
        qk = qk.to(tl.float32) * q_lora_scale[:, None]

        qk_rope = tl.dot(q_rope_fp8, tl.trans(kv_rope_fp8))
        qk_rope = qk_rope.to(tl.float32) * q_rope_scale[:, None]

        qk += qk_rope
        qk = qk * combined_scale

        qk = tl.where(mask_m[:, None] & mask_n[None, :], qk, float('-inf'))
        m_ij = tl.max(qk, 1)
        m_i_new = tl.maximum(m_i, m_ij)
        alpha = tl.exp(m_i - m_i_new)
        p = tl.exp(qk - m_i_new[:, None])
        l_i_new = alpha * l_i + tl.sum(p, 1)

        acc = acc * alpha[:, None]
        acc += tl.dot(p.to(tl.bfloat16), kv_lora_fp8.to(tl.bfloat16))

        m_i = m_i_new
        l_i = l_i_new
        kv_lora_ptrs += BLOCK_N * stride_kv_t
        kv_rope_ptrs += BLOCK_N * stride_kv_t

    acc = acc / l_i[:, None]
    po_ptrs = Partial_O + split_idx*stride_po_s + sq_idx[:, None]*stride_po_t + h_idx[:, None]*stride_po_h + offs_d_lora[None, :]*stride_po_d
    tl.store(po_ptrs, acc, mask=mask_m[:, None])

    lse = m_i + tl.log(l_i)
    lse_ptrs = Partial_LSE + split_idx*stride_plse_s + sq_idx*stride_plse_t + h_idx*stride_plse_h
    tl.store(lse_ptrs, lse, mask=mask_m)

@triton.jit
def _reduce_kernel(
    Partial_O, Partial_LSE, Out, kv_scale_ptr, total_th,
    stride_po_s, stride_po_t, stride_po_h, stride_po_d,
    stride_plse_s, stride_plse_t, stride_plse_h,
    stride_o_t, stride_o_h, stride_o_d,
    NUM_HEADS: tl.constexpr, NUM_KV_SPLITS: tl.constexpr, V_HEAD_DIM: tl.constexpr, BLOCK_D: tl.constexpr,
):
    pid = tl.program_id(0)
    if pid >= total_th: return
    t_idx = pid // NUM_HEADS
    h_idx = pid % NUM_HEADS

    kv_scale_val = tl.load(kv_scale_ptr)

    offs_s = tl.arange(0, NUM_KV_SPLITS)
    lse_ptrs = Partial_LSE + offs_s*stride_plse_s + t_idx*stride_plse_t + h_idx*stride_plse_h
    lse_vals = tl.load(lse_ptrs)
    
    lse_max  = tl.max(lse_vals)
    lse_max  = tl.where(lse_max == float('-inf'), 0.0, lse_max)
    weights = tl.exp(lse_vals - lse_max)
    w_sum   = tl.sum(weights)

    for d_start in range(0, V_HEAD_DIM, BLOCK_D):
        offs_d = d_start + tl.arange(0, BLOCK_D)
        mask_d = offs_d < V_HEAD_DIM
        po_ptrs = Partial_O + offs_s[:, None]*stride_po_s + t_idx*stride_po_t + h_idx*stride_po_h + offs_d[None, :]*stride_po_d
        po_vals = tl.load(po_ptrs, mask=mask_d[None, :], other=0.0)
        
        weighted = weights[:, None] * po_vals
        acc = tl.sum(weighted, axis=0)
        acc = acc / tl.maximum(w_sum, 1e-12) * kv_scale_val
        
        out_ptrs = Out + t_idx*stride_o_t + h_idx*stride_o_h + offs_d*stride_o_d
        tl.store(out_ptrs, acc.to(tl.bfloat16), mask=mask_d)


@triton.jit
def _reduce_kernel_dsplit(
    Partial_O, Partial_LSE, Out, kv_scale_ptr, total_th,
    stride_po_s, stride_po_t, stride_po_h, stride_po_d,
    stride_plse_s, stride_plse_t, stride_plse_h,
    stride_o_t, stride_o_h, stride_o_d,
    NUM_HEADS: tl.constexpr, NUM_KV_SPLITS: tl.constexpr, V_HEAD_DIM: tl.constexpr, BLOCK_D: tl.constexpr,
):
    pid_th = tl.program_id(0)
    pid_d = tl.program_id(1)
    if pid_th >= total_th:
        return

    t_idx = pid_th // NUM_HEADS
    h_idx = pid_th % NUM_HEADS

    kv_scale_val = tl.load(kv_scale_ptr)

    offs_s = tl.arange(0, NUM_KV_SPLITS)
    lse_ptrs = Partial_LSE + offs_s*stride_plse_s + t_idx*stride_plse_t + h_idx*stride_plse_h
    lse_vals = tl.load(lse_ptrs)

    lse_max = tl.max(lse_vals)
    lse_max = tl.where(lse_max == float('-inf'), 0.0, lse_max)
    weights = tl.exp(lse_vals - lse_max)
    scale = kv_scale_val / tl.maximum(tl.sum(weights), 1e-12)

    d_start = pid_d * BLOCK_D
    offs_d = d_start + tl.arange(0, BLOCK_D)
    mask_d = offs_d < V_HEAD_DIM

    po_ptrs = Partial_O + offs_s[:, None]*stride_po_s + t_idx*stride_po_t + h_idx*stride_po_h + offs_d[None, :]*stride_po_d
    po_vals = tl.load(po_ptrs, mask=mask_d[None, :], other=0.0)

    acc = tl.sum(weights[:, None] * po_vals, axis=0)
    acc = acc * scale

    out_ptrs = Out + t_idx*stride_o_t + h_idx*stride_o_h + offs_d*stride_o_d
    tl.store(out_ptrs, acc.to(tl.bfloat16), mask=mask_d)

def _short_kv_target_blocks(batch_size: int) -> int:
    if batch_size <= 4:
        return SHORT_KV_TARGET_BLOCKS_SMALL
    return SHORT_KV_TARGET_BLOCKS


def _short_split_route_bucket(split_size: int) -> int:
    if split_size <= 64:
        return 0
    if split_size <= 128:
        return 1
    return 2


def _long_split_route_bucket(num_kv_splits: int) -> int:
    if num_kv_splits >= 32:
        return 0
    if num_kv_splits >= 8:
        return 1
    return 2


# ==============================================================================
# v29: 短 KV 独立 autotune cache + 长 KV split 分桶
# ==============================================================================
def custom_kernel_v29(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    num_heads        = config["num_heads"]
    qk_head_dim      = config["qk_head_dim"]
    v_head_dim       = config["v_head_dim"]
    sm_scale         = config["sm_scale"]
    kv_lora_rank     = config.get("kv_lora_rank", 512)
    qk_rope_head_dim = config.get("qk_rope_head_dim", 64)

    kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
    kv_fp8_2d = kv_buffer_fp8.view(-1, qk_head_dim)

    batch_size = qo_indptr.shape[0] - 1
    total_q    = q.shape[0]
    total_kv   = kv_buffer_fp8.shape[0]

    q_seq_len = int(config.get("q_seq_len", max(1, (total_q + batch_size - 1) // batch_size)))
    kv_seq_len = int(config.get("kv_seq_len", max(1, (total_kv + batch_size - 1) // batch_size)))

    short_kv_special = (
        q_seq_len == 1
        and kv_seq_len <= SHORT_KV_MAX
    )

    if not short_kv_special:
        return custom_kernel_v14(data)

    max_m = q_seq_len * num_heads
    block_m_est = 16
    m_blocks = (max_m + block_m_est - 1) // block_m_est
    base_blocks = batch_size * m_blocks

    target_blocks = _short_kv_target_blocks(batch_size)
    use_split_kv = base_blocks < target_blocks

    if not use_split_kv:
        out = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=q.device)
        def grid_fn(META): return (batch_size, triton.cdiv(max_m, META['BLOCK_M']))
        _mla_v28_short_no_split[grid_fn](
            q, kv_fp8_2d, out, qo_indptr, kv_indptr, kv_scale_fp8, sm_scale, 0,
            q.stride(0), q.stride(1), q.stride(2),
            kv_fp8_2d.stride(0), kv_fp8_2d.stride(1),
            out.stride(0), out.stride(1), out.stride(2),
            num_heads=num_heads, KV_LORA_RANK=kv_lora_rank,
            QK_ROPE_DIM=qk_rope_head_dim, v_head_dim=v_head_dim,
        )
        return out

    needed_splits = max(1, (target_blocks + base_blocks - 1) // base_blocks)
    num_kv_splits = 1
    while num_kv_splits < needed_splits:
        num_kv_splits *= 2

    num_kv_splits = min(num_kv_splits, SHORT_KV_MAX_SPLITS)
    raw_split = (kv_seq_len + num_kv_splits - 1) // num_kv_splits

    BLOCK_N_max = 32
    split_size = max(BLOCK_N_max, ((raw_split + BLOCK_N_max - 1) // BLOCK_N_max) * BLOCK_N_max)
    route_bucket = _short_split_route_bucket(split_size)

    partial_o = torch.empty(
        (num_kv_splits, total_q, num_heads, kv_lora_rank),
        dtype=torch.float32, device=q.device,
    )
    partial_lse = torch.empty(
        (num_kv_splits, total_q, num_heads),
        dtype=torch.float32, device=q.device,
    )

    def grid_fn_split(META):
        return (batch_size, triton.cdiv(max_m, META['BLOCK_M']), num_kv_splits)
    
    _mla_v28_short_split[grid_fn_split](
        q, kv_fp8_2d, partial_o, partial_lse,
        qo_indptr, kv_indptr, kv_scale_fp8, sm_scale, route_bucket,
        q.stride(0), q.stride(1), q.stride(2),
        kv_fp8_2d.stride(0), kv_fp8_2d.stride(1),
        partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
        partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
        num_heads=num_heads, KV_LORA_RANK=kv_lora_rank, QK_ROPE_DIM=qk_rope_head_dim,
        SPLIT_SIZE=split_size,
    )

    out = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=q.device)
    total_th = total_q * num_heads
    _reduce_kernel[(total_th,)](
        partial_o, partial_lse, out, kv_scale_fp8, total_th,
        partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
        partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
        out.stride(0), out.stride(1), out.stride(2),
        NUM_HEADS=num_heads, NUM_KV_SPLITS=num_kv_splits, V_HEAD_DIM=v_head_dim,
        BLOCK_D=128, num_warps=1,
    )
    return out


# ==============================================================================
# v41: 在 v37 基础上只测试 small case 的 partial_o 带宽压缩
# ==============================================================================
def custom_kernel_v41(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    num_heads        = config["num_heads"]
    qk_head_dim      = config["qk_head_dim"]
    v_head_dim       = config["v_head_dim"]
    sm_scale         = config["sm_scale"]
    kv_lora_rank     = config.get("kv_lora_rank", 512)
    qk_rope_head_dim = config.get("qk_rope_head_dim", 64)

    kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
    kv_fp8_2d = kv_buffer_fp8.view(-1, qk_head_dim)

    batch_size = qo_indptr.shape[0] - 1
    total_q    = q.shape[0]
    total_kv   = kv_buffer_fp8.shape[0]

    q_seq_len = int(config.get("q_seq_len", max(1, (total_q + batch_size - 1) // batch_size)))
    kv_seq_len = int(config.get("kv_seq_len", max(1, (total_kv + batch_size - 1) // batch_size)))

    short_kv_special = (
        q_seq_len == 1
        and kv_seq_len <= SHORT_KV_MAX
    )

    if not short_kv_special:
        return custom_kernel_v14(data)

    max_m = q_seq_len * num_heads
    block_m_est = 16
    m_blocks = (max_m + block_m_est - 1) // block_m_est
    base_blocks = batch_size * m_blocks

    target_blocks = _short_kv_target_blocks(batch_size)
    use_split_kv = base_blocks < target_blocks

    if not use_split_kv:
        out = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=q.device)
        def grid_fn(META): return (batch_size, triton.cdiv(max_m, META['BLOCK_M']))
        _mla_v28_short_no_split[grid_fn](
            q, kv_fp8_2d, out, qo_indptr, kv_indptr, kv_scale_fp8, sm_scale, 0,
            q.stride(0), q.stride(1), q.stride(2),
            kv_fp8_2d.stride(0), kv_fp8_2d.stride(1),
            out.stride(0), out.stride(1), out.stride(2),
            num_heads=num_heads, KV_LORA_RANK=kv_lora_rank,
            QK_ROPE_DIM=qk_rope_head_dim, v_head_dim=v_head_dim,
        )
        return out

    needed_splits = max(1, (target_blocks + base_blocks - 1) // base_blocks)
    num_kv_splits = 1
    while num_kv_splits < needed_splits:
        num_kv_splits *= 2

    num_kv_splits = min(num_kv_splits, SHORT_KV_MAX_SPLITS)
    raw_split = (kv_seq_len + num_kv_splits - 1) // num_kv_splits

    BLOCK_N_max = 32
    split_size = max(BLOCK_N_max, ((raw_split + BLOCK_N_max - 1) // BLOCK_N_max) * BLOCK_N_max)
    route_bucket = _short_split_route_bucket(split_size)

    total_th = total_q * num_heads
    use_small_bf16_partial_o = total_th <= SHORT_KV_SMALL_REDUCE_MAX_TOTAL_TH
    partial_o_dtype = torch.bfloat16 if use_small_bf16_partial_o else torch.float32

    partial_o = torch.empty(
        (num_kv_splits, total_q, num_heads, kv_lora_rank),
        dtype=partial_o_dtype, device=q.device,
    )
    partial_lse = torch.empty(
        (num_kv_splits, total_q, num_heads),
        dtype=torch.float32, device=q.device,
    )

    def grid_fn_split(META):
        return (batch_size, triton.cdiv(max_m, META['BLOCK_M']), num_kv_splits)

    _mla_v28_short_split[grid_fn_split](
        q, kv_fp8_2d, partial_o, partial_lse,
        qo_indptr, kv_indptr, kv_scale_fp8, sm_scale, route_bucket,
        q.stride(0), q.stride(1), q.stride(2),
        kv_fp8_2d.stride(0), kv_fp8_2d.stride(1),
        partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
        partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
        num_heads=num_heads, KV_LORA_RANK=kv_lora_rank, QK_ROPE_DIM=qk_rope_head_dim,
        SPLIT_SIZE=split_size,
    )

    out = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=q.device)
    use_small_dsplit_reduce = (
        total_th <= SHORT_KV_SMALL_REDUCE_MAX_TOTAL_TH
        and v_head_dim >= SHORT_KV_SMALL_REDUCE_BLOCK_D
    )

    if use_small_dsplit_reduce:
        def grid_fn_reduce(META):
            return (total_th, triton.cdiv(v_head_dim, META['BLOCK_D']))
        _reduce_kernel_dsplit[grid_fn_reduce](
            partial_o, partial_lse, out, kv_scale_fp8, total_th,
            partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
            partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
            out.stride(0), out.stride(1), out.stride(2),
            NUM_HEADS=num_heads, NUM_KV_SPLITS=num_kv_splits, V_HEAD_DIM=v_head_dim,
            BLOCK_D=SHORT_KV_SMALL_REDUCE_BLOCK_D, num_warps=1,
        )
        return out

    _reduce_kernel[(total_th,)](
        partial_o, partial_lse, out, kv_scale_fp8, total_th,
        partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
        partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
        out.stride(0), out.stride(1), out.stride(2),
        NUM_HEADS=num_heads, NUM_KV_SPLITS=num_kv_splits, V_HEAD_DIM=v_head_dim,
        BLOCK_D=128, num_warps=1,
    )
    return out


# ==============================================================================
# v14: BLOCK_N=32 压缩 VGPR → Occupancy=2
# ==============================================================================
def custom_kernel_v14(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    num_heads        = config["num_heads"]
    qk_head_dim      = config["qk_head_dim"]      # 576
    v_head_dim       = config["v_head_dim"]        # 512
    sm_scale         = config["sm_scale"]
    kv_lora_rank     = config.get("kv_lora_rank", 512)
    qk_rope_head_dim = config.get("qk_rope_head_dim", 64)

    kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
    kv_fp8_2d = kv_buffer_fp8.view(-1, qk_head_dim)

    batch_size = qo_indptr.shape[0] - 1
    total_q    = q.shape[0]
    total_kv   = kv_buffer_fp8.shape[0]

    q_seq_avg = (total_q + batch_size - 1) // batch_size
    q_seq_len_est = 4 if q_seq_avg <= 4 else q_seq_avg + 4 # Padding bound
    max_m = q_seq_len_est * num_heads
    
    kv_seq_avg = (total_kv + batch_size - 1) // batch_size
    kv_seq_max = kv_seq_avg + 384 

    # ===================================
    # V11 激进优化:取消短序列 No-Split 惩罚,允许极端切片
    # ===================================
    BLOCK_M_est = 16
    m_blocks = (max_m + BLOCK_M_est - 1) // BLOCK_M_est
    base_blocks = batch_size * m_blocks
    
    # 只要 base_blocks < 1024 (为了撑满极高 VGPR 所需的 CU Occupancy),就必须切分!
    use_split_kv = base_blocks < 1024

    if not use_split_kv:
        out = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=q.device)
        def grid_fn(META): return (batch_size, triton.cdiv(max_m, META['BLOCK_M']))
        _mla_v11_no_split[grid_fn](
            q, kv_fp8_2d, out, qo_indptr, kv_indptr, kv_scale_fp8, sm_scale,
            q.stride(0), q.stride(1), q.stride(2),
            kv_fp8_2d.stride(0), kv_fp8_2d.stride(1),
            out.stride(0), out.stride(1), out.stride(2),
            num_heads=num_heads, KV_LORA_RANK=kv_lora_rank,
            QK_ROPE_DIM=qk_rope_head_dim, v_head_dim=v_head_dim,
        )
        return out
        
    # ====== Split-KV =======
    # 由于我们需要保障每一个 CU 尽可能运转,目标总块数达到 1024
    needed_splits = max(1, (1024 + base_blocks - 1) // base_blocks)
    num_kv_splits = 1
    while num_kv_splits < needed_splits:
        num_kv_splits *= 2
    
    # 彻底解除切分粒度封印,但最高至 32 以防 Reduce 发生降速反噬
    num_kv_splits = min(num_kv_splits, 32)
    raw_split = (kv_seq_max + num_kv_splits - 1) // num_kv_splits

    BLOCK_N_max = 32
    split_size = max(BLOCK_N_max, ((raw_split + BLOCK_N_max - 1) // BLOCK_N_max) * BLOCK_N_max)

    partial_o = torch.empty(
        (num_kv_splits, total_q, num_heads, kv_lora_rank),
        dtype=torch.float32, device=q.device,
    )
    partial_lse = torch.empty(
        (num_kv_splits, total_q, num_heads),
        dtype=torch.float32, device=q.device,
    )

    def grid_fn_split(META):
        return (batch_size, triton.cdiv(max_m, META['BLOCK_M']), num_kv_splits)
    
    long_route_bucket = _long_split_route_bucket(num_kv_splits)

    _mla_v11_split[grid_fn_split](
        q, kv_fp8_2d, partial_o, partial_lse,
        qo_indptr, kv_indptr, kv_scale_fp8, sm_scale, long_route_bucket,
        q.stride(0), q.stride(1), q.stride(2),
        kv_fp8_2d.stride(0), kv_fp8_2d.stride(1),
        partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
        partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
        num_heads=num_heads, KV_LORA_RANK=kv_lora_rank, QK_ROPE_DIM=qk_rope_head_dim,
        SPLIT_SIZE=split_size, NUM_KV_SPLITS=num_kv_splits,
    )

    out = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=q.device)
    total_th = total_q * num_heads
    _reduce_kernel[(total_th,)](
        partial_o, partial_lse, out, kv_scale_fp8, total_th,
        partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
        partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
        out.stride(0), out.stride(1), out.stride(2),
        NUM_HEADS=num_heads, NUM_KV_SPLITS=num_kv_splits, V_HEAD_DIM=v_head_dim,
        BLOCK_D=128, num_warps=1,
    )
    return out
scrolls · 939 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