Skip to content
KernelIndex
Search⌘K

gpt-5 / tritona41cd4

gpt-5_triton_a41cd4 · gpt-5-2025-08-07 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

47 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=8
NVIDIA B200
101.2µs
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=108
NVIDIA B200
358.4µs
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=457
NVIDIA B200
374.7µs
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=208
NVIDIA B200
617.3µs
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=308
NVIDIA B200
870.8µs
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=408
NVIDIA B200
1.12ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=1857
NVIDIA B200
1.25ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=508
NVIDIA B200
1.38ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=7257
NVIDIA B200
1.48ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=608
NVIDIA B200
1.64ms
#4 of 4
2025-10-16
Show all 47 measurements ›
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=8057
NVIDIA B200
1.76ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=5057
NVIDIA B200
1.85ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=708
NVIDIA B200
1.92ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=5857
NVIDIA B200
1.97ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=6657
NVIDIA B200
2.13ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=808
NVIDIA B200
2.18ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=1008
NVIDIA B200
2.70ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=1108
NVIDIA B200
2.95ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=1208
NVIDIA B200
3.21ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=4357
NVIDIA B200
3.55ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=2757
NVIDIA B200
3.72ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=3557
NVIDIA B200
3.86ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=8857
NVIDIA B200
4.24ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=9945
NVIDIA B200
4.57ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=16345
NVIDIA B200
4.83ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=1908
NVIDIA B200
5.04ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=22745
NVIDIA B200
5.04ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=27545
NVIDIA B200
5.27ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=9657
NVIDIA B200
5.41ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=30745
NVIDIA B200
5.46ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=33945
NVIDIA B200
5.52ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=10857
NVIDIA B200
5.62ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=37145
NVIDIA B200
5.72ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=40345
NVIDIA B200
5.79ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=12857
NVIDIA B200
6.02ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=2408
NVIDIA B200
6.34ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=14857
NVIDIA B200
6.39ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=17257
NVIDIA B200
6.95ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=2708
NVIDIA B200
7.11ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=44845
NVIDIA B200
8.98ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=48045
NVIDIA B200
9.15ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=51245
NVIDIA B200
9.30ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=54445
NVIDIA B200
9.44ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=57645
NVIDIA B200
9.61ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=62345
NVIDIA B200
23.5ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=75145
NVIDIA B200
24.0ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=68745
NVIDIA B200
27.3ms
#4 of 4
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:c4d597804add5bfbd3c24de72e5a7c66c9f333e568fe8cdfcc9b7346fcc4cbd0
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20

Techniques

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

num-warps = 4num_warps=4,
online-softmaxm_new = tl.maximum(m, l)
stages = 2num_stages=2,

Kernel source

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


@triton.jit
def mla_paged_decode_h16_ckv512_kpe64_ps1_kernel(
    q_nope_ptr,
    q_pe_ptr,
    ckv_ptr,
    kpe_ptr,
    kv_indptr_ptr,
    kv_indices_ptr,
    output_ptr,
    lse_ptr,
    B,
    H: tl.constexpr,
    DCKV: tl.constexpr,
    DKPE: tl.constexpr,
    sm_scale,
    stride_qn_b,
    stride_qn_h,
    stride_qn_d,
    stride_qp_b,
    stride_qp_h,
    stride_qp_d,
    stride_ckv_p,
    stride_ckv_d,
    stride_kpe_p,
    stride_kpe_d,
    stride_out_b,
    stride_out_h,
    stride_out_d,
    stride_lse_b,
    stride_lse_h,
    BLOCK_TOK: tl.constexpr,
    BLOCK_DCKV: tl.constexpr,
    BLOCK_DKPE: tl.constexpr,
):
    pid = tl.program_id(0)
    b = pid // H
    h = pid % H
    if b >= B:
        return

    page_start = tl.load(kv_indptr_ptr + b, mask=True, other=0).to(tl.int32)
    page_end = tl.load(kv_indptr_ptr + (b + 1), mask=True, other=0).to(tl.int32)
    L = page_end - page_start

    qn_base = q_nope_ptr + b * stride_qn_b + h * stride_qn_h
    qp_base = q_pe_ptr + b * stride_qp_b + h * stride_qp_h
    out_base = output_ptr + b * stride_out_b + h * stride_out_h
    lse_off = lse_ptr + b * stride_lse_b + h * stride_lse_h

    # Early exit
    if L <= 0:
        offs_d = tl.arange(0, BLOCK_DCKV)
        zero_bf16 = tl.zeros([BLOCK_DCKV], dtype=tl.bfloat16)
        for t in range(0, DCKV, BLOCK_DCKV):
            d = t + offs_d
            mask_d = d < DCKV
            tl.store(out_base + d * stride_out_d, zero_bf16, mask=mask_d)
        tl.store(lse_off, -float("inf"))
        return

    # Preload q vectors
    offs_dckv = tl.arange(0, BLOCK_DCKV)
    qn0 = tl.load(qn_base + offs_dckv * stride_qn_d, mask=offs_dckv < DCKV, other=0.0).to(tl.float32)
    d1 = offs_dckv + BLOCK_DCKV
    qn1 = tl.load(qn_base + d1 * stride_qn_d, mask=d1 < DCKV, other=0.0).to(tl.float32)
    d2 = offs_dckv + 2 * BLOCK_DCKV
    qn2 = tl.load(qn_base + d2 * stride_qn_d, mask=d2 < DCKV, other=0.0).to(tl.float32)
    d3 = offs_dckv + 3 * BLOCK_DCKV
    qn3 = tl.load(qn_base + d3 * stride_qn_d, mask=d3 < DCKV, other=0.0).to(tl.float32)

    offs_kpe = tl.arange(0, BLOCK_DKPE)
    qp_vec = tl.load(qp_base + offs_kpe * stride_qp_d, mask=offs_kpe < DKPE, other=0.0).to(tl.float32)

    # Streaming softmax stats (natural log domain)
    m = tl.full([], -float("inf"), dtype=tl.float32)
    S = tl.full([], 0.0, dtype=tl.float32)

    # Numerator accumulators for output (float32)
    O0 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
    O1 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
    O2 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
    O3 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)

    start = tl.zeros([], dtype=tl.int32)
    while start < L:
        # Process BLOCK_TOK tokens sequentially for better numerical stability
        for i in range(BLOCK_TOK):
            t = start + i
            valid = t < L

            # Load token index
            tok = tl.load(kv_indices_ptr + page_start + t, mask=valid, other=0).to(tl.int32)

            # Compute logits for this token
            # Kc dot qn
            K0 = tl.load(
                ckv_ptr + tok * stride_ckv_p + offs_dckv * stride_ckv_d,
                mask=valid & (offs_dckv < DCKV),
                other=0.0,
            ).to(tl.float32)
            l = tl.sum(K0 * qn0, axis=0)

            K1 = tl.load(
                ckv_ptr + tok * stride_ckv_p + d1 * stride_ckv_d,
                mask=valid & (d1 < DCKV),
                other=0.0,
            ).to(tl.float32)
            l += tl.sum(K1 * qn1, axis=0)

            K2 = tl.load(
                ckv_ptr + tok * stride_ckv_p + d2 * stride_ckv_d,
                mask=valid & (d2 < DCKV),
                other=0.0,
            ).to(tl.float32)
            l += tl.sum(K2 * qn2, axis=0)

            K3 = tl.load(
                ckv_ptr + tok * stride_ckv_p + d3 * stride_ckv_d,
                mask=valid & (d3 < DCKV),
                other=0.0,
            ).to(tl.float32)
            l += tl.sum(K3 * qn3, axis=0)

            # Kp dot qp
            KP = tl.load(
                kpe_ptr + tok * stride_kpe_p + offs_kpe * stride_kpe_d,
                mask=valid & (offs_kpe < DKPE),
                other=0.0,
            ).to(tl.float32)
            l += tl.sum(KP * qp_vec, axis=0)

            # Scale logits and mask invalid
            l = l * sm_scale
            l = tl.where(valid, l, -float("inf"))

            # Streaming softmax update for a single token
            m_new = tl.maximum(m, l)
            scale_prev = tl.exp(m - m_new)
            p = tl.exp(l - m_new)

            # Update denominator
            S = S * scale_prev + p
            # Update numerators
            O0 = O0 * scale_prev + K0 * p
            O1 = O1 * scale_prev + K1 * p
            O2 = O2 * scale_prev + K2 * p
            O3 = O3 * scale_prev + K3 * p

            m = m_new

        start += BLOCK_TOK

    inv_S = 1.0 / S
    O0 = O0 * inv_S
    O1 = O1 * inv_S
    O2 = O2 * inv_S
    O3 = O3 * inv_S

    # Store output
    tl.store(out_base + offs_dckv * stride_out_d, O0.to(tl.bfloat16), mask=offs_dckv < DCKV)
    tl.store(out_base + d1 * stride_out_d, O1.to(tl.bfloat16), mask=d1 < DCKV)
    tl.store(out_base + d2 * stride_out_d, O2.to(tl.bfloat16), mask=d2 < DCKV)
    tl.store(out_base + d3 * stride_out_d, O3.to(tl.bfloat16), mask=d3 < DCKV)

    # Base-2 LSE: logsumexp(logits_scaled) / log(2)
    ln2 = tl.log(2.0)
    lse_val = (m + tl.log(S)) / ln2
    tl.store(lse_off, lse_val)


def run(q_nope, q_pe, ckv_cache, kpe_cache, kv_indptr, kv_indices, sm_scale):
    # Validate dtypes
    assert q_nope.dtype == torch.bfloat16, "q_nope must be bfloat16"
    assert q_pe.dtype == torch.bfloat16, "q_pe must be bfloat16"
    assert ckv_cache.dtype == torch.bfloat16, "ckv_cache must be bfloat16"
    assert kpe_cache.dtype == torch.bfloat16, "kpe_cache must be bfloat16"
    assert kv_indptr.dtype == torch.int32, "kv_indptr must be int32"
    assert kv_indices.dtype == torch.int32, "kv_indices must be int32"

    # Shapes and constants
    batch_size, num_qo_heads, head_dim_ckv = q_nope.shape
    head_dim_kpe = q_pe.shape[-1]
    page_size = ckv_cache.shape[1]
    len_indptr = kv_indptr.shape[0]
    num_kv_indices = kv_indices.shape[0]

    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 page_size == 1, "page_size must be 1"

    assert len_indptr == batch_size + 1, "len_indptr must equal batch_size + 1"
    assert num_kv_indices == int(kv_indptr[-1].item()), "num_kv_indices must equal kv_indptr[-1]"

    # Device handling
    orig_device = q_nope.device
    if q_nope.is_cuda:
        device = q_nope.device
    else:
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available, but Triton kernel requires a GPU.")
        device = torch.device("cuda")

    # Move tensors to device
    def to_dev(t):
        return t.to(device, non_blocking=True)

    q_nope_dev = to_dev(q_nope.contiguous())
    q_pe_dev = to_dev(q_pe.contiguous())
    # Squeeze page dimension (ps=1)
    ckv_dev = to_dev(ckv_cache.squeeze(1).contiguous())  # [num_pages, 512]
    kpe_dev = to_dev(kpe_cache.squeeze(1).contiguous())  # [num_pages, 64]
    kv_indptr_dev = to_dev(kv_indptr.contiguous())
    kv_indices_dev = to_dev(kv_indices.contiguous())

    # Outputs
    output_dev = torch.empty((batch_size, num_qo_heads, head_dim_ckv), dtype=torch.bfloat16, device=device)
    lse_dev = torch.empty((batch_size, num_qo_heads), dtype=torch.float32, device=device)

    # Strides (elements)
    stride_qn_b, stride_qn_h, stride_qn_d = q_nope_dev.stride()
    stride_qp_b, stride_qp_h, stride_qp_d = q_pe_dev.stride()
    stride_ckv_p, stride_ckv_d = ckv_dev.stride()
    stride_kpe_p, stride_kpe_d = kpe_dev.stride()
    stride_out_b, stride_out_h, stride_out_d = output_dev.stride()
    stride_lse_b, stride_lse_h = lse_dev.stride()

    # Launch configuration
    B = batch_size
    H = 16
    DCKV = 512
    DKPE = 64
    # Token block and vector block sizes
    # Smaller BLOCK_TOK for better numerical stability and register pressure
    BLOCK_TOK = 32
    BLOCK_DCKV = 128
    BLOCK_DKPE = 64

    grid = (B * H,)

    mla_paged_decode_h16_ckv512_kpe64_ps1_kernel[grid](
        q_nope_dev,
        q_pe_dev,
        ckv_dev,
        kpe_dev,
        kv_indptr_dev,
        kv_indices_dev,
        output_dev,
        lse_dev,
        B,
        H,
        DCKV,
        DKPE,
        float(sm_scale),
        stride_qn_b,
        stride_qn_h,
        stride_qn_d,
        stride_qp_b,
        stride_qp_h,
        stride_qp_d,
        stride_ckv_p,
        stride_ckv_d,
        stride_kpe_p,
        stride_kpe_d,
        stride_out_b,
        stride_out_h,
        stride_out_d,
        stride_lse_b,
        stride_lse_h,
        BLOCK_TOK,
        BLOCK_DCKV,
        BLOCK_DKPE,
        num_warps=4,
        num_stages=2,
    )

    # Move outputs back to original device
    output = output_dev.to(orig_device, non_blocking=True)
    lse = lse_dev.to(orig_device, non_blocking=True)
    return output, lse
scrolls · 286 lines total

Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0

Best evidence level for this revision: reported

JSON