Skip to content
KernelIndex
Search⌘K

gpt-o3 / tritonc3c0cc

gpt-o3_triton_c3c0cc · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

48 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=18 · num_kv_indices=2
NVIDIA B200
10.3µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
10.4µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=11 · num_kv_indices=10
NVIDIA B200
10.4µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=55 · num_kv_indices=38
NVIDIA B200
11.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=63 · num_kv_indices=46
NVIDIA B200
12.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
12.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=59 · num_kv_indices=42
NVIDIA B200
12.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=71 · num_kv_indices=54
NVIDIA B200
12.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=74 · num_kv_indices=57
NVIDIA B200
12.4µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=78 · num_kv_indices=61
NVIDIA B200
12.5µs
#2 of 7
2025-10-16
Show all 48 measurements ›
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=82 · num_kv_indices=65
NVIDIA B200
12.5µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9316 · num_kv_indices=73
NVIDIA B200
12.7µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9341 · num_kv_indices=98
NVIDIA B200
14.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=223 · num_kv_indices=173
NVIDIA B200
19.9µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=406 · num_kv_indices=356
NVIDIA B200
28.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1069 · num_kv_indices=1034
NVIDIA B200
30.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1220 · num_kv_indices=1193
NVIDIA B200
30.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1556 · num_kv_indices=1529
NVIDIA B200
30.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1396 · num_kv_indices=1369
NVIDIA B200
31.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1732 · num_kv_indices=1705
NVIDIA B200
32.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1892 · num_kv_indices=1865
NVIDIA B200
35.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2052 · num_kv_indices=2025
NVIDIA B200
35.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2212 · num_kv_indices=2185
NVIDIA B200
36.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2372 · num_kv_indices=2345
NVIDIA B200
37.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2548 · num_kv_indices=2521
NVIDIA B200
38.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2708 · num_kv_indices=2681
NVIDIA B200
38.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2868 · num_kv_indices=2841
NVIDIA B200
39.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=3044 · num_kv_indices=3017
NVIDIA B200
39.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=597 · num_kv_indices=547
NVIDIA B200
41.1µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=4333 · num_kv_indices=4298
NVIDIA B200
46.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=81390 · num_kv_indices=12942
NVIDIA B200
169.7µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=30163 · num_kv_indices=20911
NVIDIA B200
204.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=60071 · num_kv_indices=50902
NVIDIA B200
277.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63053 · num_kv_indices=53782
NVIDIA B200
288.2µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=62605 · num_kv_indices=53334
NVIDIA B200
289.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63437 · num_kv_indices=54166
NVIDIA B200
291.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64205 · num_kv_indices=54934
NVIDIA B200
292.6µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63821 · num_kv_indices=54550
NVIDIA B200
292.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65037 · num_kv_indices=55766
NVIDIA B200
293.3µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65805 · num_kv_indices=56534
NVIDIA B200
295.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65421 · num_kv_indices=56150
NVIDIA B200
295.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64653 · num_kv_indices=55382
NVIDIA B200
295.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67021 · num_kv_indices=57750
NVIDIA B200
297.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66253 · num_kv_indices=56982
NVIDIA B200
298.8µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66637 · num_kv_indices=57366
NVIDIA B200
299.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67405 · num_kv_indices=58134
NVIDIA B200
300.1µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67853 · num_kv_indices=58582
NVIDIA B200
301.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=68237 · num_kv_indices=58966
NVIDIA B200
309.0µs
#3 of 7
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:55702029eab3dc4e109a991973a4c7a8b689d1790d073c27b8e9a8ad1c7a22d9
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 = 4num_warps=4,
stages = 4num_stages=4,

Kernel source

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


@triton.jit
def gqa_paged_decode_kernel(
    q_ptr,               # *bf16  [B, 32, 128]
    k_ptr,               # *bf16  [N_pages, 8, 128]  (page_size squeezed)
    v_ptr,               # *bf16  [N_pages, 8, 128]  (page_size squeezed)
    kv_indptr_ptr,       # *int32 [B + 1]
    kv_indices_ptr,      # *int32 [num_kv_indices]
    sm_scale,            # fp32 scalar
    out_ptr,             # *bf16  [B, 32, 128]
    lse_ptr,             # *fp32  [B, 32]
    BLOCK_T: tl.constexpr,
    HEAD_DIM: tl.constexpr,
    NUM_QO_HEADS: tl.constexpr,
    NUM_KV_HEADS: tl.constexpr,
):
    pid = tl.program_id(0)

    batch_idx = pid // NUM_QO_HEADS
    qo_head   = pid % NUM_QO_HEADS
    gqa_ratio = NUM_QO_HEADS // NUM_KV_HEADS
    kv_head   = qo_head // gqa_ratio

    # ---- strides (in elements, not bytes) ----
    stride_q_batch    = NUM_QO_HEADS * HEAD_DIM
    stride_q_head     = HEAD_DIM

    stride_k_page     = NUM_KV_HEADS * HEAD_DIM         # page_size = 1
    stride_k_kv_head  = HEAD_DIM

    stride_v_page     = stride_k_page
    stride_v_kv_head  = HEAD_DIM

    # ---- load query vector ----
    d_offs = tl.arange(0, HEAD_DIM)
    q_ptr_head = q_ptr + batch_idx * stride_q_batch + qo_head * stride_q_head + d_offs
    q_vec = tl.cast(tl.load(q_ptr_head), tl.float32)

    # ---- sequence token range ----
    start = tl.load(kv_indptr_ptr + batch_idx)
    end   = tl.load(kv_indptr_ptr + batch_idx + 1)
    num_tokens = end - start

    # ---- streaming softmax vars ----
    m_val   = tl.full([], -1e30, tl.float32)          # running max
    d_val   = tl.zeros([], tl.float32)                # running sum exp
    o_vec   = tl.zeros([HEAD_DIM], tl.float32)        # running output vector

    offset = tl.zeros([], tl.int32)

    while offset < num_tokens:
        t_offs      = tl.arange(0, BLOCK_T)
        remain      = num_tokens - offset
        tok_mask    = t_offs < remain

        # ---- load page indices ----
        pages = tl.load(kv_indices_ptr + start + offset + t_offs,
                        mask=tok_mask, other=0)

        # ---- gather K / V ----
        k_ptrs = k_ptr + pages[:, None] * stride_k_page + kv_head * stride_k_kv_head + d_offs[None, :]
        v_ptrs = v_ptr + pages[:, None] * stride_v_page + kv_head * stride_v_kv_head + d_offs[None, :]

        k_block = tl.cast(tl.load(k_ptrs, mask=tok_mask[:, None], other=0), tl.float32)
        v_block = tl.cast(tl.load(v_ptrs, mask=tok_mask[:, None], other=0), tl.float32)

        # ---- logits ----
        logits = tl.sum(k_block * q_vec[None, :], axis=1) * sm_scale
        logits = tl.where(tok_mask, logits, -1e30)

        # ---- block softmax ----
        m_block        = tl.max(logits, axis=0)
        exp_logits     = tl.exp(logits - m_block)
        sum_exp_block  = tl.sum(exp_logits, axis=0)
        weighted_v     = tl.sum(exp_logits[:, None] * v_block, axis=0)

        # ---- merge with running values ----
        new_m      = tl.maximum(m_val, m_block)
        alpha_prev = tl.exp(m_val - new_m)
        alpha_blk  = tl.exp(m_block - new_m)

        o_vec = o_vec * alpha_prev + weighted_v * alpha_blk
        d_val = d_val * alpha_prev + sum_exp_block * alpha_blk
        m_val = new_m

        offset += BLOCK_T

    inv_d   = tl.where(d_val == 0, 0.0, 1.0 / d_val)
    out_vec = o_vec * inv_d
    log2e   = 1.4426950408889634
    lse_val = tl.where(d_val == 0,
                       -1e30,
                       (tl.log(d_val) + m_val) * log2e)

    # ---- store ----
    out_ptr_head = out_ptr + batch_idx * stride_q_batch + qo_head * stride_q_head + d_offs
    tl.store(out_ptr_head, tl.cast(out_vec, tl.bfloat16))

    lse_ptr_head = lse_ptr + batch_idx * NUM_QO_HEADS + qo_head
    tl.store(lse_ptr_head, lse_val)


def run(q,
        k_cache,
        v_cache,
        kv_indptr,
        kv_indices,
        sm_scale: float | None = None):
    """
    Entry point for gqa_paged_decode_h32_kv8_d128_ps1.
    Returns (output, lse).
    """
    if sm_scale is None:
        sm_scale = 1.0 / math.sqrt(128.0)

    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device is required to run Triton kernels.")

    # move tensors to GPU if necessary
    tensors = [q, k_cache, v_cache, kv_indptr, kv_indices]
    device_tensors = [t.cuda() if not t.is_cuda else t for t in tensors]
    q_dev, k_dev, v_dev, iptr_dev, idx_dev = [t.contiguous() for t in device_tensors]

    batch_size = q_dev.shape[0]
    num_qo_heads = 32
    head_dim = 128

    # squeeze page dimension (=1)
    k_dev_flat = k_dev.squeeze(1).contiguous()
    v_dev_flat = v_dev.squeeze(1).contiguous()

    out_dev = torch.empty((batch_size, num_qo_heads, head_dim),
                          dtype=torch.bfloat16,
                          device=q_dev.device)
    lse_dev = torch.empty((batch_size, num_qo_heads),
                          dtype=torch.float32,
                          device=q_dev.device)

    # launch kernel
    BLOCK_T = 128
    grid = (batch_size * num_qo_heads,)

    gqa_paged_decode_kernel[grid](
        q_dev, k_dev_flat, v_dev_flat,
        iptr_dev, idx_dev,
        sm_scale,
        out_dev, lse_dev,
        BLOCK_T=BLOCK_T,
        HEAD_DIM=128,
        NUM_QO_HEADS=32,
        NUM_KV_HEADS=8,
        num_warps=4,
        num_stages=4,
    )

    # move back to original device if needed
    if not q.is_cuda:
        return out_dev.cpu(), lse_dev.cpu()
    return out_dev, lse_dev
scrolls · 164 lines total

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

Best evidence level for this revision: reported

JSON