gpt-o3 / tritonad56c1
gpt-o3_triton_ad56c1 · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 206 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-ad56c1?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
38 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Show all 38 measurements ›Showing all 38 measurements ⌄
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [123, 16, 64]
NVIDIA B200
455.9µs
#4 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [138, 16, 64]
NVIDIA B200
646.8µs
#4 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [199, 16, 64]
NVIDIA B200
707.2µs
#4 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [1954, 16, 64]
NVIDIA B200
8.83ms
#4 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [1028, 16, 64]
NVIDIA B200
13.3ms
#4 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [1187, 16, 64]
NVIDIA B200
14.6ms
#4 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [3842, 16, 64]
NVIDIA B200
34.7ms
#4 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [6053, 16, 64]
NVIDIA B200
85.3ms
#4 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [3024, 16, 64]
NVIDIA B200
89.9ms
#4 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [8987, 16, 64]
NVIDIA B200
148.6ms
#4 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [15092, 16, 64]
NVIDIA B200
329.5ms
#4 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [15883, 16, 64]
NVIDIA B200
462.1ms
#4 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [10870, 16, 64]
NVIDIA B200
713.5ms
#4 of 4
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [16384, 16, 64]
NVIDIA B200
3.23s
#4 of 4
2025-10-16
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:ad057faa231033181e08f35d96fbd208a61042e2dd0ef32835e77da5486c5ba8
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 = 4
num_warps=4, num_stages=2,online-softmax
m_new = tl.maximum(m_prev, m_blk)stages = 2
num_warps=4, num_stages=2,tile-k = 32
BLOCK_K: tl.constexpr = 32,Kernel source
main.py206 lines
import math
import torch
import triton
import triton.language as tl
# ============================================================================ #
# TRITON KERNEL #
# ============================================================================ #
@triton.jit
def _mla_paged_prefill_kernel(
q_nope_ptr, # *bf16 [TOTAL_Q, 16, 512]
q_pe_ptr, # *bf16 [TOTAL_Q, 16, 64]
kc_ptr, # *bf16 [KV_LEN, 512] (sequence-contiguous)
kp_ptr, # *bf16 [KV_LEN, 64]
out_ptr, # *bf16 [TOTAL_Q, 16, 512]
lse_ptr, # *fp32 [TOTAL_Q, 16]
kv_len, # i32 – #tokens in this sequence’s KV buffer
prefix_len, # i32 – kv_len - q_len
sm_scale, # fp32 – soft-max scale
q_global_offset, # i32 – start row of this sequence in Q tensors
BLOCK_K: tl.constexpr = 32,
HEAD_C: tl.constexpr = 512,
HEAD_P: tl.constexpr = 64,
):
"""
One program instance = (one query token, one head).
Grid = (q_len, 16)
Implements streaming softmax with causal masking and fused output.
"""
# --------------------------------------------------------------------- #
# INDICES #
# --------------------------------------------------------------------- #
pid_q = tl.program_id(0) # query index inside the sequence
pid_h = tl.program_id(1) # head (0‥15)
q_row = q_global_offset + pid_q # absolute query row inside Q tensors
# --------------------------------------------------------------------- #
# LOAD QUERY VECTORS #
# --------------------------------------------------------------------- #
qn_off = (q_row * 16 + pid_h) * HEAD_C + tl.arange(0, HEAD_C)
qp_off = (q_row * 16 + pid_h) * HEAD_P + tl.arange(0, HEAD_P)
qn = tl.load(q_nope_ptr + qn_off).to(tl.float32) # [512]
qp = tl.load(q_pe_ptr + qp_off).to(tl.float32) # [ 64]
# --------------------------------------------------------------------- #
# STREAMING SOFTMAX ACCUMULATORS #
# --------------------------------------------------------------------- #
m_prev = tl.full((), -float("inf"), tl.float32) # running max
l_prev = tl.zeros((), tl.float32) # running sum(exp)
acc_out = tl.zeros((HEAD_C,), tl.float32) # running numerator
query_abs_pos = prefix_len + pid_q # absolute position
num_iters = (kv_len + BLOCK_K - 1) // BLOCK_K
INV_LN2 = 1.4426950408889634 # 1 / ln(2)
iter_idx = 0
while iter_idx < num_iters:
k_start = iter_idx * BLOCK_K
tok_offs = k_start + tl.arange(0, BLOCK_K) # [B]
valid_m = tok_offs < kv_len # [B]
causal_m = tok_offs > query_abs_pos # [B]
keep_m = valid_m & ~causal_m # [B]
# ------------------------ LOAD KC / KP BLOCK --------------------- #
kc_ptrs = kc_ptr + tok_offs[:, None] * HEAD_C + tl.arange(0, HEAD_C)[None, :]
kp_ptrs = kp_ptr + tok_offs[:, None] * HEAD_P + tl.arange(0, HEAD_P)[None, :]
kc_blk = tl.load(kc_ptrs, mask=valid_m[:, None], other=0).to(tl.float32) # [B,512]
kp_blk = tl.load(kp_ptrs, mask=valid_m[:, None], other=0).to(tl.float32) # [B, 64]
# ---------------------------- DOTS ------------------------------- #
dotkc = tl.sum(kc_blk * qn[None, :], 1) # [B]
dotkp = tl.sum(kp_blk * qp[None, :], 1) # [B]
logits = (dotkc + dotkp) * sm_scale # [B]
neg_inf = -float("inf")
logits = tl.where(keep_m, logits, neg_inf)
# ----------------------- STABLE SOFTMAX -------------------------- #
m_blk = tl.max(logits, 0)
m_new = tl.maximum(m_prev, m_blk)
exp_logits = tl.exp(logits - m_new)
exp_logits = tl.where(keep_m, exp_logits, 0.0)
alpha_prev = tl.exp(m_prev - m_new)
l_prev = l_prev * alpha_prev + tl.sum(exp_logits, 0)
acc_out = acc_out * alpha_prev + tl.sum(exp_logits[:, None] * kc_blk, 0)
m_prev = m_new
iter_idx += 1
# --------------------------- WRITE-BACK ------------------------------ #
out_vec = acc_out / l_prev
out_ptrs = out_ptr + (q_row * 16 + pid_h) * HEAD_C + tl.arange(0, HEAD_C)
tl.store(out_ptrs, out_vec.to(tl.bfloat16))
lse_val = (tl.log(l_prev) + m_prev) * INV_LN2 # base-2 log-sum-exp
tl.store(lse_ptr + q_row * 16 + pid_h, lse_val)
# ============================================================================ #
# PYTHON ENTRY #
# ============================================================================ #
def run(
q_nope, # bf16 [total_q , 16, 512]
q_pe, # bf16 [total_q , 16, 64]
ckv_cache, # bf16 [num_pages, 1, 512]
kpe_cache, # bf16 [num_pages, 1, 64]
qo_indptr, # int32 [len_indptr]
kv_indptr, # int32 [len_indptr]
kv_indices, # int32 [num_kv_indices]
sm_scale=None, # optional float
):
"""
Optimized paged-KV prefill for B200 GPUs.
Handles device transfers transparently.
"""
if not torch.cuda.is_available():
raise RuntimeError("CUDA device is required but not available.")
NUM_HEADS = 16
HEAD_C = 512
HEAD_P = 64
PAGE_SIZE = 1
BLOCK_K = 32 # must stay in sync with kernel default
# --------------------------- SHAPE CHECKS --------------------------- #
assert q_nope.shape[1:] == (NUM_HEADS, HEAD_C)
assert q_pe.shape[1:] == (NUM_HEADS, HEAD_P)
assert ckv_cache.shape[1] == PAGE_SIZE
assert kpe_cache.shape[1] == PAGE_SIZE
assert q_nope.shape[0] == qo_indptr[-1].item()
assert kv_indices.shape[0] == kv_indptr[-1].item()
# ---------------------------- DEVICE I/O --------------------------- #
def _to_cuda(t: torch.Tensor):
return t.cuda(non_blocking=True) if t.device.type != "cuda" else t
def _back(t: torch.Tensor, ref: torch.Tensor):
return t.cpu() if ref.device.type != "cuda" else t
q_nope_c = _to_cuda(q_nope)
q_pe_c = _to_cuda(q_pe)
kc_all = _to_cuda(ckv_cache).squeeze(1).contiguous() # [pages,512]
kp_all = _to_cuda(kpe_cache).squeeze(1).contiguous() # [pages, 64]
qo_ind_c = _to_cuda(qo_indptr)
kv_ind_c = _to_cuda(kv_indptr)
kv_idx_c = _to_cuda(kv_indices)
total_q = q_nope_c.shape[0]
batch = qo_ind_c.shape[0] - 1
# ------------------------- OUTPUT BUFFERS --------------------------- #
out_c = torch.empty_like(q_nope_c)
lse_c = torch.full(
(total_q, NUM_HEADS), -float("inf"), dtype=torch.float32, device=q_nope_c.device
)
# --------------------------- SM SCALE ------------------------------ #
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(HEAD_C)
sm_scale = float(sm_scale)
# ------------------------- SEQUENCE LOOP --------------------------- #
for b in range(batch):
q_beg, q_end = int(qo_ind_c[b].item()), int(qo_ind_c[b + 1].item())
if q_beg >= q_end:
continue
p_beg, p_end = int(kv_ind_c[b].item()), int(kv_ind_c[b + 1].item())
if p_beg >= p_end:
continue
kv_pages = kv_idx_c[p_beg:p_end].long()
kv_len = kv_pages.numel()
q_len = q_end - q_beg
prefix = kv_len - q_len
if prefix < 0:
raise RuntimeError("KV length must be ≥ query length (causal)")
# Gather contiguous KC / KP for this sequence
kc_seq = kc_all.index_select(0, kv_pages).contiguous()
kp_seq = kp_all.index_select(0, kv_pages).contiguous()
grid = (q_len, NUM_HEADS) # (pid_q, pid_h)
_mla_paged_prefill_kernel[grid](
q_nope_c, q_pe_c,
kc_seq, kp_seq,
out_c, lse_c,
kv_len, prefix,
sm_scale, q_beg,
num_warps=4, num_stages=2,
BLOCK_K=BLOCK_K,
HEAD_C=HEAD_C,
HEAD_P=HEAD_P,
)
return _back(out_c, q_nope), _back(lse_c, q_nope)scrolls · 206 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON