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
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.
mma
S = tl.dot(q, tl.trans(K)) # [BLOCK_Q, BLOCK_K]num-warps = 4
num_warps = 4online-softmax
m_new = tl.maximum(m_i, s_max)stages = 2
num_stages = 2tile-k = 32
BLOCK_K = 32Kernel 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, lsescrolls · 333 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON