Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton25db20

gpt-o3_triton_25db20 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

21 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #cb4b00
NVIDIA B200
93.6µs
#4 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #5fa6fa
NVIDIA B200
94.8µs
#4 of 5
2025-10-19
NVIDIA B200
94.9µs
#3 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
94.9µs
#7 of 10
2025-10-19
NVIDIA B200
94.9µs
#4 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
94.9µs
#5 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #f3d59b
NVIDIA B200
95.0µs
#2 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
95.2µs
#8 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
95.3µs
#6 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
95.4µs
#6 of 10
2025-10-19
Show all 21 measurements ›
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
95.7µs
#7 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
96.8µs
#8 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
98.8µs
#7 of 10
2025-10-19
NVIDIA B200
109.6µs
#3 of 10
2025-10-19
NVIDIA B200
109.9µs
#4 of 10
2025-10-19
NVIDIA B200
180.7µs
#2 of 5
2025-10-19
NVIDIA B200
425.4µs
#5 of 5
2025-10-19
NVIDIA B200
1.91ms
#3 of 5
2025-10-19
NVIDIA B200
90.5ms
#3 of 5
2025-10-19
NVIDIA B200
92.4ms
#3 of 5
2025-10-19
NVIDIA B200
93.0ms
#3 of 5
2025-10-19

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:ff9215c61d71f0e0d14eff780af2c632149307e0fdfd3d2668689d759e121cff
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
online-softmaxm_new = tl.maximum(m_prev, block_max)
stages = 2NUM_STAGES = 2
tile-k = 64BLOCK_K = 64

Kernel source

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


# -----------------------------------------------------------------------------#
#                               Triton  Kernel                                  #
# -----------------------------------------------------------------------------#
@triton.jit
def gqa_ragged_prefill_causal_kernel(
    Q_ptr,                    # *bf16  [num_q, 32, 128]
    K_ptr,                    # *bf16  [num_kv, 8, 128]
    V_ptr,                    # *bf16  [num_kv, 8, 128]
    OUT_ptr,                  # *bf16  [num_q, 32, 128]
    LSE_ptr,                  # *fp32  [num_q, 32]
    NUM_Q: tl.constexpr,      # number of query tokens in this sequence
    NUM_KV: tl.constexpr,     # number of kv tokens in this sequence
    DELTA: tl.constexpr,      # NUM_KV - NUM_Q
    SM_SCALE: tl.constexpr,   # softmax scale (float32)
    gqa_ratio: tl.constexpr,  # 4             (32 / 8)
    BLOCK_D: tl.constexpr,    # 128
    BLOCK_K: tl.constexpr,    # 64  / 128
    N_BLOCKS_K: tl.constexpr, # ceil_div(NUM_KV, BLOCK_K)
):
    """
    One program = one (query_token, qo_head) pair.
    program_id(0) := query token   in [0, NUM_Q)
    program_id(1) := qo head index in [0, 32)
    """

    pid_q = tl.program_id(0)
    pid_h = tl.program_id(1)

    # Out-of-bounds queries are ignored (host pads the launch grid if needed).
    if pid_q >= NUM_Q:
        return

    d = tl.arange(0, BLOCK_D)  # [0 .. 127]

    # -------------------- Load Q ------------------------------------------------
    q_ptrs = Q_ptr + (pid_q * 32 + pid_h) * BLOCK_D + d            # [D] strides
    q_vec = tl.load(q_ptrs).to(tl.float32)                         # [D] (fp32)

    kv_head = pid_h // gqa_ratio                                   # 0 .. 7

    # -------------------- Running accumulators -------------------------------
    m_prev = tl.full((), -float("inf"), tl.float32)  # running max
    s_prev = tl.zeros((), tl.float32)                # running sum(exp)
    acc_prev = tl.zeros((BLOCK_D,), tl.float32)      # running weighted value sum

    allowed_k = tl.minimum(pid_q + 1 + DELTA, NUM_KV)  # causal upper-bound
    ln2_const = 0.6931471805599453                      # ln(2)

    for block_idx in tl.static_range(N_BLOCKS_K):
        kv_idx_base = block_idx * BLOCK_K
        kv_offsets = kv_idx_base + tl.arange(0, BLOCK_K)           # [BLOCK_K]
        mask_k = kv_offsets < allowed_k                            # bool

        # ------------- Load  K ------------------------------------------------
        k_ptrs = (
            K_ptr
            + ((kv_offsets[:, None] * 8 + kv_head) * BLOCK_D)
            + d[None, :]
        )
        k_tile = tl.load(k_ptrs, mask=mask_k[:, None], other=0.0).to(tl.float32)
        # k_tile: [BLOCK_K, D]

        # ------------- Compute Scores ----------------------------------------
        scores = tl.sum(k_tile * q_vec[None, :], axis=1)           # [BLOCK_K]
        scores = scores * SM_SCALE
        scores = tl.where(mask_k, scores, -float("inf"))

        block_max = tl.max(scores, axis=0)
        m_new = tl.maximum(m_prev, block_max)

        exp_scores = tl.exp(scores - m_new)
        exp_scores = tl.where(mask_k, exp_scores, 0.0)

        alpha = tl.exp(m_prev - m_new)
        s_new = s_prev * alpha + tl.sum(exp_scores, axis=0)

        # ------------- Load  V -----------------------------------------------
        v_ptrs = (
            V_ptr
            + ((kv_offsets[:, None] * 8 + kv_head) * BLOCK_D)
            + d[None, :]
        )
        v_tile = tl.load(v_ptrs, mask=mask_k[:, None], other=0.0).to(tl.float32)
        # v_tile: [BLOCK_K, D]

        attn_v = tl.sum(v_tile * exp_scores[:, None], axis=0)      # [D]
        acc_new = acc_prev * alpha + attn_v

        # update running state
        m_prev = m_new
        s_prev = s_new
        acc_prev = acc_new

    # -------------------- Finalize & Store ------------------------------------
    zero_mask = s_prev == 0.0
    out_vec = tl.where(zero_mask, tl.zeros_like(acc_prev), acc_prev / s_prev)

    lse_val = tl.where(
        zero_mask,
        -float("inf"),
        (tl.log(s_prev) + m_prev) / ln2_const,
    )

    # store output
    out_ptrs = OUT_ptr + (pid_q * 32 + pid_h) * BLOCK_D + d
    tl.store(out_ptrs, out_vec.to(tl.bfloat16))

    lse_ptr = LSE_ptr + pid_q * 32 + pid_h
    tl.store(lse_ptr, lse_val)


# -----------------------------------------------------------------------------#
#                                Python Wrapper                                #
# -----------------------------------------------------------------------------#
@torch.no_grad()
def run(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    sm_scale: float | None = None,
):
    """
    Entry point that mimics the reference interface.
    Handles device placement, per-sequence kernel launches, and result gathering.
    """
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device is required to run the Triton kernel.")

    # ---------------------------- constants -----------------------------------
    if sm_scale is None:
        sm_scale = 1.0 / math.sqrt(128.0)
    sm_scale = float(sm_scale)

    BLOCK_D = 128
    BLOCK_K = 64
    GQA_RATIO = 4
    NUM_WARPS = 4
    NUM_STAGES = 2

    # ------------------------ helpers -----------------------------------------
    def _to_cuda(t: torch.Tensor):
        return t.cuda() if not t.is_cuda else t

    def _maybe_cpu(t: torch.Tensor, ref: torch.Tensor):
        return t.cpu() if not ref.is_cuda else t

    # -------------------- move inputs to GPU ----------------------------------
    q_d = _to_cuda(q)
    k_d = _to_cuda(k)
    v_d = _to_cuda(v)
    qo_indptr_d = _to_cuda(qo_indptr)
    kv_indptr_d = _to_cuda(kv_indptr)

    total_q = int(q_d.shape[0])
    total_kv = int(k_d.shape[0])

    # -------------------- allocate outputs ------------------------------------
    output_d = torch.empty(
        (total_q, 32, 128), dtype=torch.bfloat16, device=q_d.device
    )
    lse_d = torch.empty((total_q, 32), dtype=torch.float32, device=q_d.device)

    # -------------------- per-sequence launch ---------------------------------
    len_indptr = int(qo_indptr_d.shape[0])
    for b in range(len_indptr - 1):
        q_start = int(qo_indptr_d[b].item())
        q_end   = int(qo_indptr_d[b + 1].item())
        kv_start = int(kv_indptr_d[b].item())
        kv_end   = int(kv_indptr_d[b + 1].item())

        num_q  = q_end - q_start
        num_kv = kv_end - kv_start
        if num_q <= 0 or num_kv <= 0:
            continue

        delta = num_kv - num_q
        n_blocks_k = (num_kv + BLOCK_K - 1) // BLOCK_K

        q_seq  = q_d[q_start:q_end].contiguous()
        k_seq  = k_d[kv_start:kv_end].contiguous()
        v_seq  = v_d[kv_start:kv_end].contiguous()
        out_seq = output_d[q_start:q_end]
        lse_seq = lse_d[q_start:q_end]

        grid = (triton.cdiv(num_q, 1), 32)

        gqa_ragged_prefill_causal_kernel[grid](
            q_seq,
            k_seq,
            v_seq,
            out_seq,
            lse_seq,
            num_q,
            num_kv,
            delta,
            sm_scale,
            gqa_ratio=GQA_RATIO,
            BLOCK_D=BLOCK_D,
            BLOCK_K=BLOCK_K,
            N_BLOCKS_K=n_blocks_k,
            num_warps=NUM_WARPS,
            num_stages=NUM_STAGES,
        )

    # -------------------- restore to CPU if needed ----------------------------
    output = _maybe_cpu(output_d, q)
    lse    = _maybe_cpu(lse_d, q)

    return output, lse
scrolls · 217 lines total

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

Best evidence level for this revision: reported

JSON