Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton4c17a1

gpt-o3_triton_4c17a1 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-4c17a1?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
59.1µs
#3 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=108
NVIDIA B200
68.5µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=208
NVIDIA B200
75.5µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=457
NVIDIA B200
78.8µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=308
NVIDIA B200
86.7µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=408
NVIDIA B200
99.0µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=508
NVIDIA B200
104.2µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=1857
NVIDIA B200
117.5µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=608
NVIDIA B200
119.2µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=708
NVIDIA B200
129.3µs
#2 of 4
2025-10-16
Show all 47 measurements ›
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=808
NVIDIA B200
147.8µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=1008
NVIDIA B200
162.4µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=5057
NVIDIA B200
166.4µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=5857
NVIDIA B200
167.1µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=1108
NVIDIA B200
173.1µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=7257
NVIDIA B200
177.3µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=8057
NVIDIA B200
183.6µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=1208
NVIDIA B200
186.5µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=6657
NVIDIA B200
190.2µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=2757
NVIDIA B200
204.3µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=3557
NVIDIA B200
208.6µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=4357
NVIDIA B200
231.9µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=8857
NVIDIA B200
236.5µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=9945
NVIDIA B200
251.2µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=1908
NVIDIA B200
257.9µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=9657
NVIDIA B200
270.8µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=10857
NVIDIA B200
284.0µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=12857
NVIDIA B200
309.9µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=16345
NVIDIA B200
311.0µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=2408
NVIDIA B200
315.5µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=14857
NVIDIA B200
337.4µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=2708
NVIDIA B200
356.6µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=22745
NVIDIA B200
365.3µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=17257
NVIDIA B200
376.7µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=27545
NVIDIA B200
437.0µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=30745
NVIDIA B200
483.8µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=33945
NVIDIA B200
532.5µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=37145
NVIDIA B200
574.9µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=40345
NVIDIA B200
598.5µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=44845
NVIDIA B200
716.3µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=48045
NVIDIA B200
739.0µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=51245
NVIDIA B200
804.5µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=54445
NVIDIA B200
822.0µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=57645
NVIDIA B200
831.9µs
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=62345
NVIDIA B200
1.13ms
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=68745
NVIDIA B200
1.16ms
#2 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=75145
NVIDIA B200
1.20ms
#2 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:3da822530f4925b23d07b281c971e733bdc1067be986982bc2e1bdbdbd7a38f0
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20

Techniques

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

num-warps = 8num_warps=8,
stages = 4num_stages=4,

Kernel source

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


@triton.jit
def _paged_decode_kernel(
    QN,                     # (B, H, 512)       bf16
    QP,                     # (B, H, 64)        bf16
    KC,                     # (P, 512)          bf16
    KP,                     # (P, 64)           bf16
    KV_INDICES,             # (N)               int32
    KV_INDPTR,              # (B + 1)           int32
    SM_SCALE,               # scalar            fp32
    OUT,                    # (B, H, 512)       bf16
    LSE,                    # (B, H)            fp32
    B: tl.constexpr,
    H: tl.constexpr,
    D_CKV: tl.constexpr,
    D_KPE: tl.constexpr,
    BLOCK_TOK: tl.constexpr,
):
    """
    One Triton program computes a single (batch, head) pair.
    page_size == 1, num_qo_heads == 16, D_CKV == 512, D_KPE == 64
    """

    pid = tl.program_id(axis=0)
    b = pid // H                      # batch index
    h = pid % H                       # head  index

    # ------------------- offsets ------------------- #
    offs_ckv = tl.arange(0, D_CKV)            # (512,)
    offs_kpe = tl.arange(0, D_KPE)            # (64,)
    offs_t   = tl.arange(0, BLOCK_TOK)        # (T,)

    # ------------------- KV range ------------------ #
    kv_beg = tl.load(KV_INDPTR + b)
    kv_end = tl.load(KV_INDPTR + b + 1)
    kv_len = kv_end - kv_beg                  # scalar int32

    # Pointers to output locations
    ptr_out = OUT + (b * H + h) * D_CKV + offs_ckv
    ptr_lse = LSE + b * H + h

    # If there is no KV data, write zeros / -inf and exit.
    if kv_len <= 0:
        tl.store(ptr_out, tl.zeros([D_CKV], dtype=tl.bfloat16))
        tl.store(ptr_lse, -float("inf"))
        return

    # ------------------- load queries -------------- #
    qn_ptr = QN + (b * H + h) * D_CKV + offs_ckv
    qp_ptr = QP + (b * H + h) * D_KPE + offs_kpe
    qn = tl.load(qn_ptr).to(tl.float32)        # (512,)
    qp = tl.load(qp_ptr).to(tl.float32)        # (64,)

    # ------------------- accumulators -------------- #
    s_sum = tl.zeros([], dtype=tl.float32)     # scalar
    w_sum = tl.zeros([D_CKV], dtype=tl.float32)

    tok_start = tl.zeros([], dtype=tl.int32)   # current token pointer

    while tok_start < kv_len:
        remaining = kv_len - tok_start
        block_n = tl.where(remaining < BLOCK_TOK, remaining, BLOCK_TOK)  # scalar int32
        mask_t = offs_t < block_n                                        # (T,)

        # --- gather token indices -------------------------------------- #
        idx_ptr  = KV_INDICES + kv_beg + tok_start + offs_t
        tok_idx  = tl.load(idx_ptr, mask=mask_t, other=0)                # (T,)

        # --- gather KC, KP --------------------------------------------- #
        kc_ptr = KC + tok_idx[:, None] * D_CKV + offs_ckv[None, :]
        kp_ptr = KP + tok_idx[:, None] * D_KPE + offs_kpe[None, :]

        kc_blk = tl.load(kc_ptr, mask=mask_t[:, None], other=0).to(tl.float32)  # (T,512)
        kp_blk = tl.load(kp_ptr, mask=mask_t[:, None], other=0).to(tl.float32)  # (T,64)

        # --- compute logits -------------------------------------------- #
        l_ckv  = tl.sum(kc_blk * qn[None, :], axis=1)          # (T,)
        l_kpe  = tl.sum(kp_blk * qp[None, :], axis=1)          # (T,)
        logits = (l_ckv + l_kpe) * SM_SCALE                   # (T,)

        exp_logits = tl.exp(logits)
        exp_logits = tl.where(mask_t, exp_logits, 0.0)

        # --- accumulate ------------------------------------------------- #
        s_sum += tl.sum(exp_logits, axis=0)                               # scalar
        w_sum += tl.sum(exp_logits[:, None] * kc_blk, axis=0)             # (512,)

        tok_start += BLOCK_TOK

    # ------------------- write back ------------------------------------ #
    inv_ln2 = 1.4426950408889634  # 1 / ln(2)
    out_vec = w_sum / s_sum
    log_s   = tl.log(s_sum) * inv_ln2

    tl.store(ptr_out, out_vec.to(tl.bfloat16))
    tl.store(ptr_lse, log_s)


def run(
    q_nope: torch.Tensor,
    q_pe: torch.Tensor,
    ckv_cache: torch.Tensor,
    kpe_cache: torch.Tensor,
    kv_indptr: torch.Tensor,
    kv_indices: torch.Tensor,
    sm_scale: float,
):
    """
    Optimized paged-decode kernel for (H=16, D_CKV=512, D_KPE=64, page_size=1).

    Inputs:
        q_nope     : (B, 16, 512)  bfloat16
        q_pe       : (B, 16, 64)   bfloat16
        ckv_cache  : (P, 1, 512)   bfloat16
        kpe_cache  : (P, 1, 64)    bfloat16
        kv_indptr  : (B + 1)       int32
        kv_indices : (N)           int32
        sm_scale   : float (fp32)

    Returns:
        dict(output=(B,16,512) bf16, lse=(B,16) fp32)
    """
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device is required to run Triton kernels.")

    # ------------------- validation ------------------- #
    assert q_nope.dtype == torch.bfloat16 and q_pe.dtype == torch.bfloat16
    B, H, D_CKV = q_nope.shape
    assert H == 16 and D_CKV == 512
    assert q_pe.shape == (B, 16, 64)
    assert ckv_cache.shape[1] == 1 and kpe_cache.shape[1] == 1            # page_size = 1
    assert kv_indptr.shape[0] == B + 1
    assert kv_indices.shape[0] == kv_indptr[-1].item()

    # ---------------- device handling ----------------- #
    orig_device = q_nope.device
    cuda_dev = torch.cuda.current_device()

    def _to_cuda(t: torch.Tensor):
        return t.to(device=cuda_dev, non_blocking=True) if not t.is_cuda else t

    q_nope_d = _to_cuda(q_nope)
    q_pe_d   = _to_cuda(q_pe)
    kc_d     = _to_cuda(ckv_cache.squeeze(1))
    kp_d     = _to_cuda(kpe_cache.squeeze(1))
    indptr_d = _to_cuda(kv_indptr)
    indices_d= _to_cuda(kv_indices)

    # ---------------- output buffers ------------------ #
    out_d = torch.empty((B, H, 512), dtype=torch.bfloat16, device=cuda_dev)
    lse_d = torch.empty((B, H), dtype=torch.float32, device=cuda_dev)

    # ---------------- kernel launch ------------------- #
    BLOCK_TOK = 128
    grid = (B * H,)

    _paged_decode_kernel[grid](
        q_nope_d,
        q_pe_d,
        kc_d,
        kp_d,
        indices_d,
        indptr_d,
        float(sm_scale),
        out_d,
        lse_d,
        B=B,
        H=H,
        D_CKV=512,
        D_KPE=64,
        BLOCK_TOK=BLOCK_TOK,
        num_warps=8,
        num_stages=4,
    )

    # --------------- move outputs back --------------- #
    if orig_device.type == "cpu":
        out_d = out_d.cpu()
        lse_d = lse_d.cpu()
    elif orig_device != out_d.device:
        out_d = out_d.to(orig_device)
        lse_d = lse_d.to(orig_device)

    return out_d, lse_d
scrolls · 189 lines total

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

Best evidence level for this revision: reported

JSON