Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton6fd1ef

gpt-o3_triton_6fd1ef · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-6fd1ef?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
143.7µs
#2 of 4
2025-10-21
NVIDIA B200
155.7µs
#3 of 6
2025-10-21

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:4a7bcf97e2019110a403ee7cdecb75144f576f4b0e60324a07796189aab45fe7
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,
tile-k = 64BLOCK_K = 64 # empirically good for H100/B200

Kernel source

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


################################################################################
#                                 Triton kernel                                #
################################################################################
@triton.jit
def _gqa_prefill_kernel(
    Q, K, V,                                   # bf16
    O, LSE,                                    # O: bf16, LSE: fp32
    stride_q_tok, stride_q_hd,                 # int32
    stride_k_tok, stride_k_hd,                 # int32
    stride_v_tok, stride_v_hd,                 # int32
    q_len: tl.constexpr,                       # int32
    kv_len: tl.constexpr,                      # int32
    delta: tl.constexpr,                       # int32 (kv_len - q_len)
    sm_scale: tl.constexpr,                    # fp32
    BLOCK_K: tl.constexpr,
    HEAD_DIM: tl.constexpr,
):
    """
    One program computes attention for a single (query_token, qo_head) pair.

    grid = (q_len, 32)
      pid0 -> query token index in the sequence          [0 .. q_len)
      pid1 -> query/output head index (32 heads total)   [0 .. 31]
    """
    q_idx = tl.program_id(0)       # query token
    h_idx = tl.program_id(1)       # qo head

    # Only launch the work-items we actually need
    if (q_idx >= q_len) | (h_idx >= 32):
        return

    # GQA: map 32 qo-heads → 4 kv-heads
    kv_head = h_idx // 8  # 32 / 4 = 8 qo-heads per kv-head

    # -------------------------------------------------------------------------
    # Load query vector [HEAD_DIM] (bf16 → fp32)
    # -------------------------------------------------------------------------
    offs_d = tl.arange(0, HEAD_DIM)
    q_ptrs = Q + q_idx * stride_q_tok + h_idx * stride_q_hd + offs_d
    q = tl.load(q_ptrs).to(tl.float32)

    # -------------------------------------------------------------------------
    # Streaming soft-max initialisation
    # -------------------------------------------------------------------------
    acc = tl.zeros([HEAD_DIM], dtype=tl.float32)    # output accumulator
    m_prev = tl.full((), -float("inf"), dtype=tl.float32)
    l_prev = tl.zeros((), dtype=tl.float32)

    # Number of KV tokens visible to this query (causal mask)
    kv_allowed = tl.minimum(kv_len, q_idx + 1 + delta)

    # -------------------------------------------------------------------------
    # Iterate over KV tokens in blocks of BLOCK_K
    # -------------------------------------------------------------------------
    offs_k = tl.arange(0, BLOCK_K)

    for kv_start in range(0, kv_len, BLOCK_K):
        curr_k_ids = kv_start + offs_k                     # [BLOCK_K]
        mask_tok = curr_k_ids < kv_allowed                 # causal / length mask

        # ---------------------------------------------------------------------
        # Load K / V blocks (bf16 → fp32)
        # ---------------------------------------------------------------------
        k_ptrs = (
            K + curr_k_ids[:, None] * stride_k_tok
              + kv_head * stride_k_hd
              + offs_d[None, :]
        )
        v_ptrs = (
            V + curr_k_ids[:, None] * stride_v_tok
              + kv_head * stride_v_hd
              + offs_d[None, :]
        )
        k_block = tl.load(k_ptrs, mask=mask_tok[:, None]).to(tl.float32)  # [B, D]
        v_block = tl.load(v_ptrs, mask=mask_tok[:, None]).to(tl.float32)  # [B, D]

        # ---------------------------------------------------------------------
        # Dot-product q · k and scale
        # ---------------------------------------------------------------------
        logits = tl.sum(k_block * q[None, :], axis=1) * sm_scale          # [B]
        logits = tl.where(mask_tok, logits, -float("inf"))

        # ---------------------------------------------------------------------
        # Numerically-stable online soft-max
        # ---------------------------------------------------------------------
        m_curr = tl.maximum(m_prev, tl.max(logits, axis=0))
        exp_logits = tl.exp(logits - m_curr)
        l_curr = tl.exp(m_prev - m_curr) * l_prev + tl.sum(exp_logits, axis=0)

        p = exp_logits / l_curr                                           # [B]

        factor = tl.exp(m_prev - m_curr) * l_prev / l_curr
        acc = acc * factor + tl.sum(p[:, None] * v_block, axis=0)         # [D]

        m_prev = m_curr
        l_prev = l_curr

    # -------------------------------------------------------------------------
    # Write output
    # -------------------------------------------------------------------------
    o_ptrs = O + q_idx * stride_q_tok + h_idx * stride_q_hd + offs_d
    tl.store(o_ptrs, tl.cast(acc, tl.bfloat16))

    log2e = 1.4426950408889634  # 1 / ln(2)
    lse_val = (m_prev + tl.log(l_prev)) * log2e
    lse_ptr = LSE + q_idx * 32 + h_idx
    tl.store(lse_ptr, lse_val)


################################################################################
#                              Python entry point                              #
################################################################################
def run(
    q, k_cache, v_cache,
    qo_indptr, kv_indptr, kv_indices,
    sm_scale=None,
):
    """
    Optimised paged-KV GQA pre-fill kernel
    (page_size = 1, 32 qo-heads / 4 kv-heads, head_dim = 128).
    """
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device is required to run Triton kernels")

    # -------------------------------------------------------------------------
    # Constants
    # -------------------------------------------------------------------------
    NUM_QO_HEADS = 32
    NUM_KV_HEADS = 4
    HEAD_DIM = 128
    PAGE_SIZE = 1
    BLOCK_K = 64  # empirically good for H100/B200

    # -------------------------------------------------------------------------
    # Soft-max scale
    # -------------------------------------------------------------------------
    if sm_scale is None:
        sm_scale = 1.0 / math.sqrt(float(HEAD_DIM))
    if torch.is_tensor(sm_scale):
        sm_scale = float(sm_scale.item())

    # -------------------------------------------------------------------------
    # Device management helpers
    # -------------------------------------------------------------------------
    orig_device = q.device
    to_cuda = lambda x: x.cuda() if not x.is_cuda else x

    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)

    # -------------------------------------------------------------------------
    # Validations
    # -------------------------------------------------------------------------
    assert q.shape[1:] == (NUM_QO_HEADS, HEAD_DIM)
    assert k_cache.shape[1:] == (PAGE_SIZE, NUM_KV_HEADS, HEAD_DIM)
    assert v_cache.shape == k_cache.shape
    assert PAGE_SIZE == 1
    total_q = q.shape[0]
    assert total_q == qo_indptr[-1].item()
    assert kv_indices.shape[0] == kv_indptr[-1].item()

    # -------------------------------------------------------------------------
    # Flatten page dimension (since page_size == 1)
    # -------------------------------------------------------------------------
    k_flat = k_cache.squeeze(1).contiguous()  # [num_pages, 4, 128]
    v_flat = v_cache.squeeze(1).contiguous()

    # -------------------------------------------------------------------------
    # Allocate outputs
    # -------------------------------------------------------------------------
    output = torch.empty_like(q)
    lse = torch.empty((total_q, NUM_QO_HEADS), dtype=torch.float32, device=q.device)

    # Strides (in elements, not bytes)
    stride_q_tok = NUM_QO_HEADS * HEAD_DIM
    stride_q_hd = HEAD_DIM
    stride_k_tok = NUM_KV_HEADS * HEAD_DIM
    stride_k_hd = HEAD_DIM
    stride_v_tok = stride_k_tok
    stride_v_hd = HEAD_DIM

    # -------------------------------------------------------------------------
    # Launch kernel sequence-by-sequence
    # -------------------------------------------------------------------------
    batch_size = qo_indptr.numel() - 1
    for b in range(batch_size):
        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())

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

        delta = kv_len - q_len

        # Gather the relevant KV pages for this sequence
        page_ids = kv_indices[kv_start:kv_end].long()
        k_seq = k_flat.index_select(0, page_ids).contiguous()
        v_seq = v_flat.index_select(0, page_ids).contiguous()

        q_seq = q[q_start:q_end].contiguous()
        o_seq = output[q_start:q_end]
        lse_seq = lse[q_start:q_end]

        grid = (q_len, NUM_QO_HEADS)

        _gqa_prefill_kernel[grid](
            q_seq, k_seq, v_seq,
            o_seq, lse_seq,
            stride_q_tok, stride_q_hd,
            stride_k_tok, stride_k_hd,
            stride_v_tok, stride_v_hd,
            q_len, kv_len, delta,
            sm_scale,
            BLOCK_K=BLOCK_K,
            HEAD_DIM=HEAD_DIM,
            num_warps=8,
            num_stages=4,
        )

    # -------------------------------------------------------------------------
    # Return results on original device
    # -------------------------------------------------------------------------
    if orig_device.type != "cuda":
        output = output.to(orig_device)
        lse = lse.to(orig_device)

    return output, lse
scrolls · 241 lines total

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

Best evidence level for this revision: reported

JSON