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
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.
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, lsescrolls · 229 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON