Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro_triton_dorbxs

gemini-2.5-pro · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 249 lines, Apache-2.0, pinned at da91508.

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-dorbxs?include=source"
interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32

Benchmark evidence

No published measurement for this revision.

No evidence · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:79e60611256991bae2f3af286e994b0729ac438273ff289ee29aaf1e51df6cba
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20

Techniques

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

autotune@triton.autotune(
mmas_update = tl.dot(k_ckv_tile, q_nope_tile_2d)
num-warps = 4triton.Config({'BLOCK_L': 32, 'BLOCK_DCKV': 64, 'BLOCK_DKPE': 64}, num_warps=4, num_stages=3),
stages = 3triton.Config({'BLOCK_L': 32, 'BLOCK_DCKV': 64, 'BLOCK_DKPE': 64}, num_warps=4, num_stages=3),

Kernel source

main.py249 lines
import math
import torch
import triton
import triton.language as tl


@triton.autotune(
    configs=[
        triton.Config({'BLOCK_L': 32, 'BLOCK_DCKV': 64, 'BLOCK_DKPE': 64}, num_warps=4, num_stages=3),
        triton.Config({'BLOCK_L': 64, 'BLOCK_DCKV': 64, 'BLOCK_DKPE': 64}, num_warps=4, num_stages=3),
        triton.Config({'BLOCK_L': 32, 'BLOCK_DCKV': 128, 'BLOCK_DKPE': 64}, num_warps=4, num_stages=2),
        triton.Config({'BLOCK_L': 64, 'BLOCK_DCKV': 128, 'BLOCK_DKPE': 64}, num_warps=8, num_stages=2),
        triton.Config({'BLOCK_L': 128, 'BLOCK_DCKV': 128, 'BLOCK_DKPE': 64}, num_warps=8, num_stages=2),
        triton.Config({'BLOCK_L': 32, 'BLOCK_DCKV': 256, 'BLOCK_DKPE': 64}, num_warps=8, num_stages=2),
    ],
    key=['HEAD_DIM_CKV', 'HEAD_DIM_KPE'],
)
@triton.jit
def mla_paged_decode_h16_ckv512_kpe64_ps1_kernel(
    # Pointers to tensors
    q_nope_ptr, q_pe_ptr, ckv_cache_ptr, kpe_cache_ptr,
    kv_indptr_ptr, kv_indices_ptr,
    output_ptr, lse_ptr,
    # Scalar inputs
    sm_scale,
    # Strides
    q_nope_stride_bs, q_nope_stride_h,
    q_pe_stride_bs, q_pe_stride_h,
    ckv_cache_stride_n,
    kpe_cache_stride_n,
    output_stride_bs, output_stride_h,
    lse_stride_bs, lse_stride_h,
    # Compile-time constants
    HEAD_DIM_CKV: tl.constexpr,
    HEAD_DIM_KPE: tl.constexpr,
    # Tuning parameters
    BLOCK_L: tl.constexpr,
    BLOCK_DCKV: tl.constexpr,
    BLOCK_DKPE: tl.constexpr,
):
    """
    Triton kernel for paged multi-level attention decode.
    Each program instance computes one head for one batch element.
    """
    # Grid computes (batch_size, num_qo_heads)
    b_idx = tl.program_id(0)
    h_idx = tl.program_id(1)
    log2 = 1.4426950408889634  # 1.0 / math.log(2.0)

    # 1. --- Get sequence length for this batch element ---
    page_beg = tl.load(kv_indptr_ptr + b_idx)
    page_end = tl.load(kv_indptr_ptr + b_idx + 1)
    L_tokens = page_end - page_beg

    # 2. --- Initialize pointers and accumulators ---
    q_nope_ptr += b_idx * q_nope_stride_bs + h_idx * q_nope_stride_h
    q_pe_ptr += b_idx * q_pe_stride_bs + h_idx * q_pe_stride_h
    output_ptr += b_idx * output_stride_bs + h_idx * output_stride_h
    lse_ptr += b_idx * lse_stride_bs + h_idx * lse_stride_h

    m_i = -float('inf')
    l_i = 0.0

    NUM_CHUNKS = HEAD_DIM_CKV // BLOCK_DCKV
    acc0 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
    acc1 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
    acc2 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
    acc3 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
    acc4 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
    acc5 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
    acc6 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
    acc7 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)

    # 3. --- Handle empty sequences ---
    if L_tokens <= 0:
        out_dtype = output_ptr.dtype.element_ty
        zero_chunk = tl.zeros([BLOCK_DCKV], dtype=tl.float32).to(out_dtype)
        for i in range(NUM_CHUNKS):
            tl.store(output_ptr + i * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), zero_chunk)
        tl.store(lse_ptr, m_i)
        return

    # 4. --- Main loop over the KV sequence in blocks ---
    for l_start in range(0, L_tokens, BLOCK_L):
        l_offs = l_start + tl.arange(0, BLOCK_L)
        l_mask = l_offs < L_tokens
        indices = tl.load(kv_indices_ptr + page_beg + l_offs, mask=l_mask, other=0)

        # --- Compute logits for the block ---
        s_block = tl.zeros([BLOCK_L], dtype=tl.float32)
        # Contribution from ckv
        for d_start in range(0, HEAD_DIM_CKV, BLOCK_DCKV):
            d_offs = d_start + tl.arange(0, BLOCK_DCKV)
            q_nope_tile = tl.load(q_nope_ptr + d_offs)
            k_ckv_ptrs = ckv_cache_ptr + indices[:, None] * ckv_cache_stride_n + d_offs[None, :]
            k_ckv_tile = tl.load(k_ckv_ptrs, mask=l_mask[:, None], other=0.0)
            # CORRECTNESS FIX: Reshape 1D q_nope_tile to 2D for tl.dot, then squeeze result
            q_nope_tile_2d = tl.reshape(q_nope_tile, (BLOCK_DCKV, 1))
            s_update = tl.dot(k_ckv_tile, q_nope_tile_2d)
            s_block += tl.squeeze(s_update, axis=1)

        # Contribution from kpe
        for d_start in range(0, HEAD_DIM_KPE, BLOCK_DKPE):
            d_offs = d_start + tl.arange(0, BLOCK_DKPE)
            q_pe_tile = tl.load(q_pe_ptr + d_offs)
            k_kpe_ptrs = kpe_cache_ptr + indices[:, None] * kpe_cache_stride_n + d_offs[None, :]
            k_kpe_tile = tl.load(k_kpe_ptrs, mask=l_mask[:, None], other=0.0)
            # CORRECTNESS FIX: Reshape 1D q_pe_tile to 2D for tl.dot, then squeeze result
            q_pe_tile_2d = tl.reshape(q_pe_tile, (BLOCK_DKPE, 1))
            s_update = tl.dot(k_kpe_tile, q_pe_tile_2d)
            s_block += tl.squeeze(s_update, axis=1)

        # --- Online softmax update ---
        s_block = tl.where(l_mask, s_block * sm_scale, -float('inf'))
        m_i_old = m_i
        m_i = tl.maximum(m_i, tl.max(s_block, axis=0))

        # NUMERICAL STABILITY: guard against nan from exp(-inf - (-inf))
        s_block_shifted = s_block - m_i
        s_block_shifted = tl.where(m_i == -float('inf'), -float('inf'), s_block_shifted)
        p_block = tl.exp(s_block_shifted)

        l_i_new = tl.sum(p_block, axis=0)
        
        alpha = tl.exp(m_i_old - m_i)
        # NUMERICAL STABILITY: if m_i_old == m_i, alpha should be 1.0. Handles -inf case.
        alpha = tl.where(m_i_old == m_i, 1.0, alpha)
        
        l_i = alpha * l_i + l_i_new
        p_block = p_block.to(ckv_cache_ptr.dtype.element_ty)

        # --- Update output accumulator ---
        if NUM_CHUNKS == 8:
            acc0 *= alpha; acc1 *= alpha; acc2 *= alpha; acc3 *= alpha
            acc4 *= alpha; acc5 *= alpha; acc6 *= alpha; acc7 *= alpha
        elif NUM_CHUNKS == 4:
            acc0 *= alpha; acc1 *= alpha; acc2 *= alpha; acc3 *= alpha
        elif NUM_CHUNKS == 2:
            acc0 *= alpha; acc1 *= alpha

        # Add contribution from the current block (p_block @ v_block)
        for i in range(NUM_CHUNKS):
            d_offs = i * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV)
            v_ckv_ptrs = ckv_cache_ptr + indices[:, None] * ckv_cache_stride_n + d_offs[None, :]
            v_ckv_tile = tl.load(v_ckv_ptrs, mask=l_mask[:, None], other=0.0)
            # CORRECTNESS FIX: Reshape 1D p_block to 2D for tl.dot, then squeeze result
            p_block_2d = tl.reshape(p_block, (1, BLOCK_L))
            update_2d = tl.dot(p_block_2d, v_ckv_tile)
            update = tl.squeeze(update_2d, axis=0)
            if i == 0: acc0 += update
            elif i == 1: acc1 += update
            elif i == 2: acc2 += update
            elif i == 3: acc3 += update
            elif i == 4: acc4 += update
            elif i == 5: acc5 += update
            elif i == 6: acc6 += update
            elif i == 7: acc7 += update

    # 5. --- Finalize and store results ---
    l_i_reciprocal = tl.where(l_i > 0.0, 1.0 / l_i, 0.0)
    out_dtype = output_ptr.dtype.element_ty
    if NUM_CHUNKS == 8:
        tl.store(output_ptr + 0 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc0 * l_i_reciprocal).to(out_dtype))
        tl.store(output_ptr + 1 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc1 * l_i_reciprocal).to(out_dtype))
        tl.store(output_ptr + 2 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc2 * l_i_reciprocal).to(out_dtype))
        tl.store(output_ptr + 3 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc3 * l_i_reciprocal).to(out_dtype))
        tl.store(output_ptr + 4 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc4 * l_i_reciprocal).to(out_dtype))
        tl.store(output_ptr + 5 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc5 * l_i_reciprocal).to(out_dtype))
        tl.store(output_ptr + 6 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc6 * l_i_reciprocal).to(out_dtype))
        tl.store(output_ptr + 7 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc7 * l_i_reciprocal).to(out_dtype))
    elif NUM_CHUNKS == 4:
        tl.store(output_ptr + 0 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc0 * l_i_reciprocal).to(out_dtype))
        tl.store(output_ptr + 1 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc1 * l_i_reciprocal).to(out_dtype))
        tl.store(output_ptr + 2 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc2 * l_i_reciprocal).to(out_dtype))
        tl.store(output_ptr + 3 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc3 * l_i_reciprocal).to(out_dtype))
    elif NUM_CHUNKS == 2:
        tl.store(output_ptr + 0 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc0 * l_i_reciprocal).to(out_dtype))
        tl.store(output_ptr + 1 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc1 * l_i_reciprocal).to(out_dtype))

    final_lse = (m_i + tl.log(l_i)) * log2
    # handle case where l_i is 0, which makes tl.log(l_i) -> -inf
    final_lse = tl.where(l_i > 0.0, final_lse, -float('inf'))
    tl.store(lse_ptr, final_lse)


def mla_paged_decode_h16_ckv512_kpe64_ps1(q_nope, q_pe, ckv_cache, kpe_cache, kv_indptr, kv_indices, sm_scale):
    """
    Wrapper function for the Triton kernel.
    Handles device management, grid computation, and kernel launch.
    """
    # 1. --- Check inputs and constants ---
    batch_size, num_qo_heads, head_dim_ckv = q_nope.shape
    head_dim_kpe = q_pe.shape[-1]

    assert num_qo_heads == 16, "num_qo_heads must be 16"
    assert head_dim_ckv == 512, "head_dim_ckv must be 512"
    assert head_dim_kpe == 64, "head_dim_kpe must be 64"
    assert ckv_cache.shape[1] == 1, "page_size must be 1"

    # 2. --- Device Management ---
    input_device = q_nope.device
    is_cpu = input_device.type == 'cpu'
    if is_cpu:
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available, but input tensors are on CPU.")
        q_nope = q_nope.cuda()
        q_pe = q_pe.cuda()
        ckv_cache = ckv_cache.cuda()
        kpe_cache = kpe_cache.cuda()
        kv_indptr = kv_indptr.cuda()
        kv_indices = kv_indices.cuda()

    # 3. --- Prepare outputs and grid ---
    output = torch.empty_like(q_nope)
    lse = torch.empty((batch_size, num_qo_heads), dtype=torch.float32, device=q_nope.device)

    ckv_cache_squeezed = ckv_cache.squeeze(1)
    kpe_cache_squeezed = kpe_cache.squeeze(1)

    grid = (batch_size, num_qo_heads)

    # 4. --- Launch kernel ---
    mla_paged_decode_h16_ckv512_kpe64_ps1_kernel[grid](
        q_nope, q_pe, ckv_cache_squeezed, kpe_cache_squeezed,
        kv_indptr, kv_indices,
        output, lse,
        sm_scale,
        q_nope.stride(0), q_nope.stride(1),
        q_pe.stride(0), q_pe.stride(1),
        ckv_cache_squeezed.stride(0),
        kpe_cache_squeezed.stride(0),
        output.stride(0), output.stride(1),
        lse.stride(0), lse.stride(1),
        HEAD_DIM_CKV=head_dim_ckv,
        HEAD_DIM_KPE=head_dim_kpe,
    )

    # 5. --- Restore device and return ---
    if is_cpu:
        output = output.to(input_device)
        lse = lse.to(input_device)

    return {"output": output, "lse": lse}

def run(*args, **kwargs):
    """
    Public entry point. Handles both args and kwargs for flexibility.
    """
    return mla_paged_decode_h16_ckv512_kpe64_ps1(*args, **kwargs)
scrolls · 249 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON