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
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 = 4
num_warps=4,online-softmax
m_i_new = tl.maximum(m_i, tl.max(x, axis=0))stages = 2
num_stages=2,tile-k = 64
BK = 64Kernel 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, lsescrolls · 249 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON