Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton2b4be8

gpt-o3_triton_2b4be8 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

2 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
4.71ms
#7 of 7
2025-10-21
NVIDIA B200
4.71ms
#7 of 7
2025-10-21

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:705a8da98a6637b5c21abc5fe05c27c52c17e49f4805e8c7d80cabd35c62144f
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_i, m_curr)
tile-k = 64BLOCK_K = 64

Kernel source

main.py229 lines
import math
from typing import Optional

import torch
import triton
import triton.language as tl


# --------------------------- Triton Kernel --------------------------- #
@triton.jit
def _gqa_paged_prefill_kernel(
    q_ptr, k_ptr, v_ptr,                       # bf16
    out_ptr, lse_ptr,                          # bf16 / fp32
    sm_scale,                                  # fp32 scalar
    L_q, L_k, delta,                           # int32
    q_st0, q_st1, q_st2,                       # int32
    k_st0, k_st1, k_st2,                       # int32
    v_st0, v_st1, v_st2,                       # int32
    o_st0, o_st1, o_st2,                       # int32
    lse_st0, lse_st1,                          # int32
    BLOCK_K: tl.constexpr,                    # 64
    HEAD_DIM: tl.constexpr,                   # 128
    GQA_RATIO: tl.constexpr,                  # 4
):
    # --------------------- Program IDs ---------------------- #
    pid_q = tl.program_id(0)   # query token  (0 .. L_q-1)
    pid_h = tl.program_id(1)   # qo head      (0 .. 31)

    if pid_q >= L_q:
        return

    # ---------------- Constant Offsets ---------------------- #
    offs_d = tl.arange(0, HEAD_DIM)                # [128]
    offs_d_brd = offs_d[None, :]                   # [1,128]

    # -------------------- Load Q ---------------------------- #
    q_ptrs = q_ptr + pid_q * q_st0 + pid_h * q_st1 + offs_d
    q_vec = tl.load(q_ptrs).to(tl.float32)         # [128]

    # --------------- Map to KV Head (GQA) ------------------- #
    kv_head = pid_h // GQA_RATIO                   # int32

    # -------------- Causal visible keys --------------------- #
    kv_max = pid_q + 1 + delta
    kv_max = tl.minimum(kv_max, L_k)

    if kv_max <= 0:
        # no visible keys -> output zeros, lse -inf
        out_ptrs = out_ptr + pid_q * o_st0 + pid_h * o_st1 + offs_d
        tl.store(out_ptrs, tl.zeros((HEAD_DIM,), dtype=tl.bfloat16))
        lse_ptrs = lse_ptr + pid_q * lse_st0 + pid_h * lse_st1
        tl.store(lse_ptrs, tl.full((), float("-inf"), dtype=tl.float32))
        return

    NEG_INF = -1.0e30

    # --------------- Accumulators --------------------------- #
    m_i = tl.full((), NEG_INF, dtype=tl.float32)
    l_i = tl.zeros((), dtype=tl.float32)
    out_acc = tl.zeros((HEAD_DIM,), dtype=tl.float32)

    # ------------------- Main Loop -------------------------- #
    start = tl.zeros((), dtype=tl.int32)
    while start < kv_max:
        kv_idx = start + tl.arange(0, BLOCK_K)          # [B]
        mask_k = kv_idx < kv_max                        # [B]

        # ------------------- Load K ------------------------- #
        k_ptrs = (
            k_ptr
            + kv_idx[:, None] * k_st0
            + kv_head * k_st1
            + offs_d_brd
        )
        k_chunk = tl.load(k_ptrs, mask=mask_k[:, None], other=0).to(tl.float32)  # [B,128]

        # ------------------ Q.K^T --------------------------- #
        dots = tl.sum(k_chunk * q_vec[None, :], axis=1) * sm_scale  # [B]
        dots = tl.where(mask_k, dots, NEG_INF)

        # ----------------- Softmax -------------------------- #
        m_curr = tl.max(dots, axis=0)
        exp_curr = tl.exp(dots - m_curr)
        l_curr = tl.sum(exp_curr, axis=0)

        # ------------------- Load V ------------------------- #
        v_ptrs = (
            v_ptr
            + kv_idx[:, None] * v_st0
            + kv_head * v_st1
            + offs_d_brd
        )
        v_chunk = tl.load(v_ptrs, mask=mask_k[:, None], other=0).to(tl.float32)  # [B,128]
        pv = tl.sum(exp_curr[:, None] * v_chunk, axis=0)                         # [128]

        # ------------- Update running stats ---------------- #
        m_new = tl.maximum(m_i, m_curr)
        out_acc = out_acc * tl.exp(m_i - m_new) + pv * tl.exp(m_curr - m_new)
        l_i = l_i * tl.exp(m_i - m_new) + l_curr * tl.exp(m_curr - m_new)
        m_i = m_new

        start += BLOCK_K

    # ------------------- Write Back ------------------------- #
    out_vec = out_acc / l_i
    out_ptrs = out_ptr + pid_q * o_st0 + pid_h * o_st1 + offs_d
    tl.store(out_ptrs, out_vec.to(tl.bfloat16))

    inv_ln2 = 1.4426950408889634   # 1 / ln(2)
    lse_val = (m_i + tl.log(l_i)) * inv_ln2
    lse_ptrs = lse_ptr + pid_q * lse_st0 + pid_h * lse_st1
    tl.store(lse_ptrs, lse_val)


# --------------------------- Python Wrapper --------------------------- #
def run(
    q: torch.Tensor,
    k_cache: torch.Tensor,
    v_cache: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    kv_indices: torch.Tensor,
    sm_scale: Optional[float] = None,
):
    """
    Optimised GQA paged-prefill causal attention kernel.
    """

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

    orig_device = q.device
    device = torch.device("cuda")

    def _to_cuda(t: torch.Tensor):
        return t.to(device) if t.device != device else t

    # Move tensors to GPU
    q = _to_cuda(q)
    k_cache = _to_cuda(k_cache)
    v_cache = _to_cuda(v_cache)
    qo_indptr = _to_cuda(qo_indptr)
    kv_indptr = _to_cuda(kv_indptr)
    kv_indices = _to_cuda(kv_indices)

    total_q, num_qo_heads, head_dim = q.shape
    num_pages, page_size, num_kv_heads, _ = k_cache.shape

    # ------------------- Sanity Checks --------------------- #
    assert num_qo_heads == 32, "num_qo_heads must be 32"
    assert num_kv_heads == 8, "num_kv_heads must be 8"
    assert head_dim == 128, "head_dim must be 128"
    assert page_size == 1, "page_size must be 1"
    assert total_q == qo_indptr[-1].item(), "total_q mismatch"
    assert kv_indices.shape[0] == kv_indptr[-1].item(), "kv_indices mismatch"

    gqa_ratio = num_qo_heads // num_kv_heads  # 4

    if sm_scale is None:
        sm_scale = 1.0 / math.sqrt(head_dim)
    if isinstance(sm_scale, torch.Tensor):
        sm_scale = float(sm_scale.item())

    # Flatten page dimension (page_size = 1)
    k_cache_flat = k_cache.squeeze(1)  # [num_pages, 8, 128]
    v_cache_flat = v_cache.squeeze(1)

    # Outputs with correct initialization
    output = torch.zeros_like(q)
    lse = torch.full((total_q, num_qo_heads), float("-inf"), dtype=torch.float32, device=device)

    BLOCK_K = 64
    HEAD_DIM = 128

    def _strides(t: torch.Tensor):
        return tuple(int(s) for s in t.stride())

    len_indptr = qo_indptr.numel()

    for b in range(len_indptr - 1):
        q_start = int(qo_indptr[b].item())
        q_end = int(qo_indptr[b + 1].item())
        kv_start = int(kv_indptr[b].item())
        kv_end = int(kv_indptr[b + 1].item())

        if (q_end - q_start) == 0 or (kv_end - kv_start) == 0:
            continue

        # Gather pages for this sequence
        page_ids = kv_indices[kv_start:kv_end].long()
        k_seq = k_cache_flat.index_select(0, page_ids).contiguous()  # [L_k, 8, 128]
        v_seq = v_cache_flat.index_select(0, page_ids).contiguous()
        q_seq = q[q_start:q_end].contiguous()                        # [L_q, 32, 128]

        L_q = q_seq.shape[0]
        L_k = k_seq.shape[0]
        delta = L_k - L_q

        # Strides
        q_st0, q_st1, q_st2 = _strides(q_seq)
        k_st0, k_st1, k_st2 = _strides(k_seq)
        v_st0, v_st1, v_st2 = _strides(v_seq)
        o_st0, o_st1, o_st2 = _strides(output[q_start:q_end])
        lse_st0, lse_st1 = _strides(lse[q_start:q_end])

        grid = (L_q, num_qo_heads)

        _gqa_paged_prefill_kernel[grid](
            q_seq, k_seq, v_seq,
            output[q_start:q_end], lse[q_start:q_end],
            sm_scale,
            L_q, L_k, delta,
            q_st0, q_st1, q_st2,
            k_st0, k_st1, k_st2,
            v_st0, v_st1, v_st2,
            o_st0, o_st1, o_st2,
            lse_st0, lse_st1,
            BLOCK_K=BLOCK_K,
            HEAD_DIM=HEAD_DIM,
            GQA_RATIO=gqa_ratio,
            num_warps=4,
        )

    # Move outputs back to original device if required
    if orig_device.type != "cuda":
        output = output.to(orig_device)
        lse = lse.to(orig_device)

    return output, lse
scrolls · 229 lines total

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

Best evidence level for this revision: reported

JSON