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
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 = 8
num_warps=8,stages = 4
num_stages=4,tile-k = 64
BLOCK_K = 64 # empirically good for H100/B200Kernel 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, lsescrolls · 241 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON