Skip to content
KernelIndex
Search⌘K

gpt-5 / tritone289b9

gpt-5_triton_e289b9 · gpt-5-2025-08-07 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-e289b9?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
259.7µs
#5 of 7
2025-10-21
NVIDIA B200
262.3µs
#5 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:18f9297b90a0c41a25d01256381f966827bcdf10556ffb291720c33ff46d239f
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
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_i_new = tl.maximum(m_i, tl.max(x, axis=0))
stages = 2num_stages=2,

Kernel source

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


@triton.jit
def gqa_paged_prefill_causal_h32_kv8_d128_ps1_kernel(
    q_ptr,                          # bfloat16* [total_q, 32, 128]
    k_ptr,                          # bfloat16* [num_pages, 1, 8, 128]
    v_ptr,                          # bfloat16* [num_pages, 1, 8, 128]
    qo_indptr_ptr,                  # int32* [len_indptr]
    kv_indptr_ptr,                  # int32* [len_indptr]
    kv_indices_ptr,                 # int32* [num_kv_indices]
    q_seq_ids_ptr,                  # int32* [total_q]
    q_pos_ptr,                      # int32* [total_q]
    out_ptr,                        # bfloat16* [total_q, 32, 128]
    lse_ptr,                        # float32* [total_q, 32]
    sm_scale,                       # float32 scalar
    total_q: tl.constexpr,          # int
    head_dim: tl.constexpr,         # int (128)
    BK: tl.constexpr,               # block size along K
    NUM_K_BLOCKS: tl.constexpr,     # global upper bound for K tiles
    # strides (in elements)
    q_stride_0, q_stride_1, q_stride_2,
    k_stride_0, k_stride_1, k_stride_2, k_stride_3,
    v_stride_0, v_stride_1, v_stride_2, v_stride_3,
    out_stride_0, out_stride_1, out_stride_2,
    lse_stride_0, lse_stride_1,
):
    q_idx = tl.program_id(0)  # 0..total_q-1
    h_idx = tl.program_id(1)  # 0..31

    d = tl.arange(0, head_dim)
    offs_k = tl.arange(0, BK)

    # Load sequence id and position for this query
    seq_id = tl.load(q_seq_ids_ptr + q_idx).to(tl.int32)
    q_pos = tl.load(q_pos_ptr + q_idx).to(tl.int32)

    # Load q_len and kv_len using indptr
    q_start = tl.load(qo_indptr_ptr + seq_id).to(tl.int32)
    q_end = tl.load(qo_indptr_ptr + seq_id + 1).to(tl.int32)
    kv_start = tl.load(kv_indptr_ptr + seq_id).to(tl.int32)
    kv_end = tl.load(kv_indptr_ptr + seq_id + 1).to(tl.int32)

    q_len = q_end - q_start
    kv_len = kv_end - kv_start
    delta_len = kv_len - q_len
    max_k = q_pos + 1 + delta_len
    max_k = tl.where(max_k < 0, 0, max_k)
    max_k = tl.where(max_k > kv_len, kv_len, max_k)
    has_any = max_k > 0

    # Compute kv head index for GQA (32 / 8 = 4)
    kvh = (h_idx // 4).to(tl.int32)

    # Load Q vector
    q_ptrs = q_ptr + q_idx * q_stride_0 + h_idx * q_stride_1 + d * q_stride_2
    q_vec_bf16 = tl.load(q_ptrs, mask=d < head_dim, other=0)
    q_vec = q_vec_bf16.to(tl.float32)

    # Streaming softmax variables (scalars)
    m_i = -float("inf")                       # running max
    l_i = 0.0                                 # running sum of exp
    acc = tl.zeros([head_dim], dtype=tl.float32)  # accumulated output

    # Iterate over K/V in tiles
    for blk in range(NUM_K_BLOCKS):
        start = blk * BK
        kv_pos = start + offs_k  # [BK]
        tile_mask = kv_pos < max_k  # [BK]

        # Load page_ids for this tile
        page_ids = tl.load(kv_indices_ptr + kv_start + kv_pos, mask=tile_mask, other=0).to(tl.int32)

        # Prepare pointer matrices for K and V loads
        # Shape after broadcasting: [BK, head_dim]
        k_ptrs = (
            k_ptr
            + page_ids[:, None] * k_stride_0
            + kvh * k_stride_2
            + d[None, :] * k_stride_3
        )
        v_ptrs = (
            v_ptr
            + page_ids[:, None] * v_stride_0
            + kvh * v_stride_2
            + d[None, :] * v_stride_3
        )

        # Load K and V tiles
        k_tile_bf16 = tl.load(k_ptrs, mask=tile_mask[:, None], other=0)
        v_tile_bf16 = tl.load(v_ptrs, mask=tile_mask[:, None], other=0)
        k_tile = k_tile_bf16.to(tl.float32)
        v_tile = v_tile_bf16.to(tl.float32)

        # Compute logits for this tile: [BK]
        logits = tl.sum(k_tile * q_vec[None, :], axis=1) * sm_scale

        # Mask invalid positions with -inf for max update
        x = tl.where(tile_mask, logits, -float("inf"))
        m_i_new = tl.maximum(m_i, tl.max(x, axis=0))

        # Compute exp only for valid lanes; invalid lanes are -inf -> exp=0
        logits_shift = tl.where(tile_mask, logits - m_i_new, -float("inf"))
        p = tl.exp(logits_shift)

        # alpha factor for running sum/max
        alpha = tl.exp(m_i - m_i_new)

        l_i = l_i * alpha + tl.sum(p, axis=0)
        acc = acc * alpha + tl.sum(v_tile * p[:, None], axis=0)

        m_i = m_i_new

    # Finalize output and LSE
    l_i_safe = tl.where(l_i > 0.0, l_i, 1.0)
    out_vec = acc / l_i_safe
    # Store output
    out_ptrs = out_ptr + q_idx * out_stride_0 + h_idx * out_stride_1 + d * out_stride_2
    tl.store(out_ptrs, out_vec.to(tl.bfloat16), mask=d < head_dim)

    # LSE base-2: (log(l_i) + m_i) / ln(2) if has_any else -inf
    ln2 = 0.6931471805599453
    lse_valid = (tl.log(l_i) + m_i) / ln2
    lse_val = tl.where(has_any, lse_valid, -float("inf"))
    lse_ptrs = lse_ptr + q_idx * lse_stride_0 + h_idx * lse_stride_1
    tl.store(lse_ptrs, lse_val)


def run(q, k_cache, v_cache, qo_indptr, kv_indptr, kv_indices, sm_scale=None):
    # Validate constants and dtypes
    if q.dtype != torch.bfloat16:
        raise TypeError("q must be bfloat16")
    if not (k_cache.dtype == torch.bfloat16 and v_cache.dtype == torch.bfloat16):
        raise TypeError("k_cache and v_cache must be bfloat16")
    if not (qo_indptr.dtype == torch.int32 and kv_indptr.dtype == torch.int32 and kv_indices.dtype == torch.int32):
        raise TypeError("qo_indptr, kv_indptr, kv_indices must be int32")

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

    if head_dim != 128 or hd2 != 128:
        raise ValueError("head_dim must be 128")
    if num_qo_heads != 32:
        raise ValueError("num_qo_heads must be 32")
    if num_kv_heads != 8:
        raise ValueError("num_kv_heads must be 8")
    if page_size != 1:
        raise ValueError("page_size must be 1")

    len_indptr = qo_indptr.shape[0]
    if total_q != int(qo_indptr[-1].item()):
        raise ValueError("total_q must equal qo_indptr[-1]")
    if int(kv_indptr.shape[0]) != len_indptr:
        raise ValueError("qo_indptr and kv_indptr must have the same length")

    num_kv_indices = kv_indices.shape[0]
    if num_kv_indices != int(kv_indptr[-1].item()):
        raise ValueError("num_kv_indices must equal kv_indptr[-1]")

    # Device management
    orig_device = q.device
    if q.is_cuda:
        device = q.device
    else:
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is required to run Triton kernel but is not available.")
        device = torch.device("cuda")

    def to_dev(x):
        return x.to(device, non_blocking=True) if x.device != device else x

    q_dev = to_dev(q)
    k_cache_dev = to_dev(k_cache)
    v_cache_dev = to_dev(v_cache)
    qo_indptr_dev = to_dev(qo_indptr)
    kv_indptr_dev = to_dev(kv_indptr)
    kv_indices_dev = to_dev(kv_indices)

    # Prepare helper arrays: q_seq_ids and q_pos_in_seq
    B = len_indptr - 1
    if B > 0 and total_q > 0:
        q_lens = (qo_indptr_dev[1:] - qo_indptr_dev[:-1]).to(torch.int32)
        seq_ids = torch.arange(B, device=device, dtype=torch.int32)
        q_seq_ids = torch.repeat_interleave(seq_ids, q_lens)

        q_seq_starts = torch.repeat_interleave(qo_indptr_dev[:-1].to(torch.int32), q_lens)
        q_positions = torch.arange(total_q, device=device, dtype=torch.int32) - q_seq_starts

        kv_lens = (kv_indptr_dev[1:] - kv_indptr_dev[:-1]).to(torch.int32)
        max_kv_len = int(kv_lens.max().item()) if kv_lens.numel() > 0 else 0
    else:
        q_seq_ids = torch.empty((0,), device=device, dtype=torch.int32)
        q_positions = torch.empty((0,), device=device, dtype=torch.int32)
        max_kv_len = 0

    # Allocate outputs on device
    out_dev = torch.zeros((total_q, num_qo_heads, head_dim), dtype=torch.bfloat16, device=device)
    lse_dev = torch.full((total_q, num_qo_heads), -float("inf"), dtype=torch.float32, device=device)

    # Softmax scale
    sm_scale_val = float(1.0 / math.sqrt(head_dim)) if sm_scale is None else float(sm_scale)

    # Strides (in elements)
    q_s0, q_s1, q_s2 = q_dev.stride()
    k_s0, k_s1, k_s2, k_s3 = k_cache_dev.stride()
    v_s0, v_s1, v_s2, v_s3 = v_cache_dev.stride()
    out_s0, out_s1, out_s2 = out_dev.stride()
    lse_s0, lse_s1 = lse_dev.stride()

    # Kernel launch configuration
    BLOCK_D = 128  # head_dim
    BK = 64
    num_k_blocks = (max_kv_len + BK - 1) // BK if max_kv_len > 0 else 1

    grid = (total_q, num_qo_heads)

    if total_q > 0:
        gqa_paged_prefill_causal_h32_kv8_d128_ps1_kernel[grid](
            q_dev,
            k_cache_dev,
            v_cache_dev,
            qo_indptr_dev,
            kv_indptr_dev,
            kv_indices_dev,
            q_seq_ids,
            q_positions,
            out_dev,
            lse_dev,
            sm_scale_val,
            total_q,
            BLOCK_D,
            BK,
            num_k_blocks,
            q_s0, q_s1, q_s2,
            k_s0, k_s1, k_s2, k_s3,
            v_s0, v_s1, v_s2, v_s3,
            out_s0, out_s1, out_s2,
            lse_s0, lse_s1,
            num_warps=4,
            num_stages=2,
        )

    out = out_dev.to(orig_device, non_blocking=True) if orig_device != device else out_dev
    lse = lse_dev.to(orig_device, non_blocking=True) if orig_device != device else lse_dev

    return out, lse
scrolls · 249 lines total

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

Best evidence level for this revision: reported

JSON