Skip to content
KernelIndex
Search⌘K

gpt-5 / triton13eb4b

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-13eb4b?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
209.0µs
#6 of 6
2025-10-21
NVIDIA B200
214.4µs
#4 of 4
2025-10-21

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:ba93b915c131adb50e3c0eb79fbaed587264ecd5cdae12d02954c5b71f85d551
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.

mmaS = tl.dot(q, tl.trans(K)) # [BLOCK_Q, BLOCK_K]
num-warps = 4num_warps = 4
online-softmaxm_new = tl.maximum(m_i, s_max)
stages = 2num_stages = 2
tile-k = 32BLOCK_K = 32

Kernel source

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


# Kernel: Paged prefill attention, gqa 32->4, head_dim=128, page_size=1
@triton.jit
def gqa_paged_prefill_causal_h32_kv4_d128_ps1_kernel(
    q_ptr,                      # *bf16 [total_q, 32, 128]
    k_ptr,                      # *bf16 [num_pages, 1, 4, 128]
    v_ptr,                      # *bf16 [num_pages, 1, 4, 128]
    kv_indices_ptr,             # *int32 [num_kv_indices]
    tiles_q_global_start_ptr,   # *int32 [num_tiles]
    tiles_q_pos_start_ptr,      # *int32 [num_tiles]
    tiles_q_len_ptr,            # *int32 [num_tiles]
    tiles_kv_start_ptr,         # *int32 [num_tiles]
    tiles_kv_len_ptr,           # *int32 [num_tiles]
    tiles_q_seq_len_ptr,        # *int32 [num_tiles]
    out_ptr,                    # *bf16 [total_q, 32, 128]
    lse_ptr,                    # *fp32 [total_q, 32]
    sm_scale,                   # fp32 scalar
    total_q,                    # int32
    q_stride_q, q_stride_h, q_stride_d,  # int64 strides for q
    k_stride_0, k_stride_1, k_stride_2, k_stride_3,  # int64 strides for k_cache
    v_stride_0, v_stride_1, v_stride_2, v_stride_3,  # int64 strides for v_cache
    out_stride_q, out_stride_h, out_stride_d,        # int64 strides for out
    lse_stride_q, lse_stride_h,                      # int64 strides for lse
    MAX_K_STEPS: tl.constexpr,   # maximum number of K-chunk steps across tiles
    BLOCK_Q: tl.constexpr,       # queries per program tile
    BLOCK_K: tl.constexpr,       # keys per chunk
    HEAD_DIM: tl.constexpr       # 128
):
    pid_tile = tl.program_id(0)
    pid_head = tl.program_id(1)  # 0..31

    # Load tile metadata
    q_gstart = tl.load(tiles_q_global_start_ptr + pid_tile, mask=True, other=0).to(tl.int32)
    q_pos_start = tl.load(tiles_q_pos_start_ptr + pid_tile, mask=True, other=0).to(tl.int32)
    tile_q_len = tl.load(tiles_q_len_ptr + pid_tile, mask=True, other=0).to(tl.int32)
    kv_start = tl.load(tiles_kv_start_ptr + pid_tile, mask=True, other=0).to(tl.int32)
    kv_len = tl.load(tiles_kv_len_ptr + pid_tile, mask=True, other=0).to(tl.int32)
    q_seq_len = tl.load(tiles_q_seq_len_ptr + pid_tile, mask=True, other=0).to(tl.int32)

    # Offsets
    q_offsets = tl.arange(0, BLOCK_Q)
    d_offsets = tl.arange(0, HEAD_DIM)
    k_offsets = tl.arange(0, BLOCK_K)

    # Masks
    q_mask = q_offsets < tile_q_len

    # Global q indices and positions within sequence
    gq_idx = (q_gstart + q_offsets).to(tl.int32)
    q_pos = (q_pos_start + q_offsets).to(tl.int32)

    # Compute allowed KV length per query due to causal masking: min(kv_len, q_pos + (kv_len - q_seq_len) + 1)
    delta = kv_len - q_seq_len
    allowed = q_pos + delta + 1
    zero = tl.zeros([BLOCK_Q], dtype=tl.int32)
    allowed = tl.maximum(allowed, zero)
    allowed = tl.minimum(allowed, kv_len)

    # Head mapping
    head_idx = pid_head  # 0..31
    kv_head = head_idx // 8  # 0..3

    # Load Q for this head
    # Pointer arithmetic in elements
    q_ptrs = (
        q_ptr
        + gq_idx[:, None].to(tl.int64) * q_stride_q
        + (head_idx.to(tl.int64)) * q_stride_h
        + d_offsets[None, :].to(tl.int64) * q_stride_d
    )
    q = tl.load(q_ptrs, mask=q_mask[:, None], other=0).to(tl.float32)

    # Initialize streaming softmax state
    neg_inf = tl.full([BLOCK_Q], -float("inf"), dtype=tl.float32)
    m_i = neg_inf
    l_i = tl.zeros([BLOCK_Q], dtype=tl.float32)
    acc = tl.zeros([BLOCK_Q, HEAD_DIM], dtype=tl.float32)

    # Iterate over K/V in chunks. MAX_K_STEPS is a compile-time constant; we mask steps beyond kv_len
    for step in range(MAX_K_STEPS):
        k0 = step * BLOCK_K
        # key index within sequence
        k_idx = k0 + k_offsets  # [BLOCK_K]
        key_valid_vec = k_idx < kv_len
        # Load page IDs for this chunk
        kv_ptrs = kv_indices_ptr + (kv_start + k_idx)
        page_ids = tl.load(kv_ptrs, mask=key_valid_vec, other=0).to(tl.int32)

        # Compute K/V pointers for each page id, for this kv_head
        # K shape per row: [HEAD_DIM]
        base_k = (
            page_ids[:, None].to(tl.int64) * k_stride_0
            + kv_head.to(tl.int64) * k_stride_2
            + d_offsets[None, :].to(tl.int64) * k_stride_3
        )
        base_v = (
            page_ids[:, None].to(tl.int64) * v_stride_0
            + kv_head.to(tl.int64) * v_stride_2
            + d_offsets[None, :].to(tl.int64) * v_stride_3
        )

        # Load K and V
        k_mask_2d = key_valid_vec[:, None]
        K = tl.load(k_ptr + base_k, mask=k_mask_2d, other=0).to(tl.float32)
        V = tl.load(v_ptr + base_v, mask=k_mask_2d, other=0).to(tl.float32)

        # Compute logits S = Q * K^T
        S = tl.dot(q, tl.trans(K))  # [BLOCK_Q, BLOCK_K]
        S = S * sm_scale

        # Apply causal + bounds mask: key position within this block is k_idx; mask if k_idx >= allowed[q]
        allowed_broadcast = allowed[:, None]  # [BLOCK_Q, 1]
        keys_broadcast = k_idx[None, :]       # [1, BLOCK_K]
        mask_ca = keys_broadcast < allowed_broadcast  # [BLOCK_Q, BLOCK_K]
        mask_keys = key_valid_vec[None, :]            # [1, BLOCK_K]
        full_mask = (mask_ca & mask_keys) & q_mask[:, None]

        S = tl.where(full_mask, S, -float("inf"))

        # Update streaming softmax statistics
        s_max = tl.max(S, axis=1)  # [BLOCK_Q]
        m_new = tl.maximum(m_i, s_max)
        p = tl.exp(S - m_new[:, None])  # masked positions become exp(-inf)=0
        alpha = tl.exp(m_i - m_new)
        l_i = l_i * alpha + tl.sum(p, axis=1)
        # Update accumulator for output numerator
        PV = tl.dot(p, V)  # [BLOCK_Q, HEAD_DIM]
        acc = acc * alpha[:, None] + PV
        m_i = m_new

    # Finalize output: out = acc / l_i; lse = (log(l_i) + m_i) / log(2)
    inv_l = tl.where(l_i > 0, 1.0 / l_i, 0.0)
    out = acc * inv_l[:, None]

    ln2 = 0.6931471805599453
    lse_nat = tl.where(l_i > 0, tl.log(l_i) + m_i, -float("inf"))
    lse_base2 = lse_nat / ln2

    # Store output
    out_ptrs = (
        out_ptr
        + gq_idx[:, None].to(tl.int64) * out_stride_q
        + head_idx.to(tl.int64) * out_stride_h
        + d_offsets[None, :].to(tl.int64) * out_stride_d
    )
    tl.store(out_ptrs, out.to(tl.bfloat16), mask=q_mask[:, None])

    # Store lse
    lse_ptrs = lse_ptr + gq_idx.to(tl.int64) * lse_stride_q + head_idx.to(tl.int64) * lse_stride_h
    tl.store(lse_ptrs, lse_base2, mask=q_mask)


def run(q, k_cache, v_cache, qo_indptr, kv_indptr, kv_indices, sm_scale):
    # Validate CUDA availability and move tensors to GPU if needed
    if not torch.cuda.is_available():
        # Ensure all inputs are on CPU or raise
        devices = {t.device.type for t in [q, k_cache, v_cache, qo_indptr, kv_indptr, kv_indices]}
        if "cuda" in devices:
            raise RuntimeError("CUDA is not available but some inputs are on GPU.")
        device = torch.device("cpu")
        raise RuntimeError("CUDA device is required to run the Triton kernel.")
    else:
        device = torch.device("cuda")

    # Constants per spec
    NUM_QO_HEADS = 32
    NUM_KV_HEADS = 4
    HEAD_DIM = 128
    PAGE_SIZE = 1

    # Checks
    assert q.dtype == torch.bfloat16
    assert k_cache.dtype == torch.bfloat16
    assert v_cache.dtype == torch.bfloat16
    assert qo_indptr.dtype in (torch.int32, torch.int64)
    assert kv_indptr.dtype in (torch.int32, torch.int64)
    assert kv_indices.dtype in (torch.int32, torch.int64)

    total_q, num_qo_heads, head_dim = q.shape
    num_pages, page_size, num_kv_heads, head_dim_k = k_cache.shape
    assert num_qo_heads == NUM_QO_HEADS, "num_qo_heads must be 32"
    assert num_kv_heads == NUM_KV_HEADS, "num_kv_heads must be 4"
    assert head_dim == HEAD_DIM and head_dim_k == HEAD_DIM, "head_dim must be 128"
    assert page_size == PAGE_SIZE, "page_size must be 1"
    assert qo_indptr[-1].item() == total_q, "Constraint violated: total_q == qo_indptr[-1]"
    assert kv_indptr[-1].item() == kv_indices.shape[0], "Constraint violated: num_kv_indices == kv_indptr[-1]"

    # Remember original devices to restore outputs
    orig_device = q.device

    # Move to GPU if necessary
    def to_cuda(t):
        return t if t.is_cuda else t.cuda(device=device, non_blocking=True)

    q = to_cuda(q)
    k_cache = to_cuda(k_cache)
    v_cache = to_cuda(v_cache)
    qo_indptr = to_cuda(qo_indptr.to(torch.int32))
    kv_indptr = to_cuda(kv_indptr.to(torch.int32))
    kv_indices = to_cuda(kv_indices.to(torch.int32))

    # Prepare outputs
    output = torch.zeros((total_q, NUM_QO_HEADS, HEAD_DIM), dtype=torch.bfloat16, device=q.device)
    lse = torch.full((total_q, NUM_QO_HEADS), -float("inf"), dtype=torch.float32, device=q.device)

    # Create tile metadata
    BLOCK_Q = 64
    BLOCK_K = 32

    # Build tiles: one tile is up to BLOCK_Q queries within a sequence
    len_indptr = qo_indptr.shape[0]
    num_seqs = len_indptr - 1
    # Guard no sequences
    if num_seqs <= 0:
        if orig_device.type != "cuda":
            return output.to(orig_device), lse.to(orig_device)
        return output, lse

    tiles_q_global_start = []
    tiles_q_pos_start = []
    tiles_q_len = []
    tiles_kv_start = []
    tiles_kv_len = []
    tiles_q_seq_len = []

    max_kv_len = 0

    # Build tiles on CPU for ease, then move to GPU
    qo_indptr_cpu = qo_indptr.cpu()
    kv_indptr_cpu = kv_indptr.cpu()

    for b in range(num_seqs):
        q_start = int(qo_indptr_cpu[b].item())
        q_end = int(qo_indptr_cpu[b + 1].item())
        kv_start = int(kv_indptr_cpu[b].item())
        kv_end = int(kv_indptr_cpu[b + 1].item())
        q_len = q_end - q_start
        kv_len = kv_end - kv_start
        if q_len <= 0 or kv_len <= 0:
            continue
        max_kv_len = max(max_kv_len, kv_len)
        t = 0
        while t < q_len:
            t_len = min(BLOCK_Q, q_len - t)
            tiles_q_global_start.append(q_start + t)
            tiles_q_pos_start.append(t)
            tiles_q_len.append(t_len)
            tiles_kv_start.append(kv_start)
            tiles_kv_len.append(kv_len)
            tiles_q_seq_len.append(q_len)
            t += t_len

    num_tiles = len(tiles_q_global_start)
    if num_tiles == 0:
        # No work to do
        if orig_device.type != "cuda":
            return output.to(orig_device), lse.to(orig_device)
        return output, lse

    # Compute max steps
    max_k_steps = (max_kv_len + BLOCK_K - 1) // BLOCK_K
    if max_k_steps <= 0:
        if orig_device.type != "cuda":
            return output.to(orig_device), lse.to(orig_device)
        return output, lse

    # Move tile metadata to GPU
    tiles_q_global_start = torch.tensor(tiles_q_global_start, dtype=torch.int32, device=q.device)
    tiles_q_pos_start = torch.tensor(tiles_q_pos_start, dtype=torch.int32, device=q.device)
    tiles_q_len = torch.tensor(tiles_q_len, dtype=torch.int32, device=q.device)
    tiles_kv_start = torch.tensor(tiles_kv_start, dtype=torch.int32, device=q.device)
    tiles_kv_len = torch.tensor(tiles_kv_len, dtype=torch.int32, device=q.device)
    tiles_q_seq_len = torch.tensor(tiles_q_seq_len, dtype=torch.int32, device=q.device)

    # Prepare stride information (in elements)
    q_s0, q_s1, q_s2 = q.stride()
    k_s0, k_s1, k_s2, k_s3 = k_cache.stride()
    v_s0, v_s1, v_s2, v_s3 = v_cache.stride()
    out_s0, out_s1, out_s2 = output.stride()
    lse_s0, lse_s1 = lse.stride()

    # Convert sm_scale
    if isinstance(sm_scale, (float, int)):
        sm_scale_val = float(sm_scale)
    elif torch.is_tensor(sm_scale):
        sm_scale_val = float(sm_scale.item())
    else:
        sm_scale_val = float(sm_scale)

    # Launch kernel
    grid = (num_tiles, NUM_QO_HEADS)
    num_warps = 4
    num_stages = 2

    gqa_paged_prefill_causal_h32_kv4_d128_ps1_kernel[grid](
        q,
        k_cache,
        v_cache,
        kv_indices,
        tiles_q_global_start,
        tiles_q_pos_start,
        tiles_q_len,
        tiles_kv_start,
        tiles_kv_len,
        tiles_q_seq_len,
        output,
        lse,
        sm_scale_val,
        total_q,
        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,
        MAX_K_STEPS=max_k_steps,
        BLOCK_Q=BLOCK_Q,
        BLOCK_K=BLOCK_K,
        HEAD_DIM=HEAD_DIM,
        num_warps=num_warps,
        num_stages=num_stages,
    )

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

    return output, lse
scrolls · 333 lines total

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

Best evidence level for this revision: reported

JSON