gpt-o3 / triton25db20
gpt-o3_triton_25db20 · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 217 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-25db20?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
21 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #cb4b00
NVIDIA B200
93.6µs
#4 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #5fa6fa
NVIDIA B200
94.8µs
#4 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
94.9µs
#7 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
94.9µs
#5 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #f3d59b
NVIDIA B200
95.0µs
#2 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
95.2µs
#8 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
95.3µs
#6 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
95.4µs
#6 of 10
2025-10-19
Show all 21 measurements ›Showing all 21 measurements ⌄
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
95.7µs
#7 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
96.8µs
#8 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
98.8µs
#7 of 10
2025-10-19
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:ff9215c61d71f0e0d14eff780af2c632149307e0fdfd3d2668689d759e121cff
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 = 4online-softmax
m_new = tl.maximum(m_prev, block_max)stages = 2
NUM_STAGES = 2tile-k = 64
BLOCK_K = 64Kernel source
main.py217 lines
import math
import torch
import triton
import triton.language as tl
# -----------------------------------------------------------------------------#
# Triton Kernel #
# -----------------------------------------------------------------------------#
@triton.jit
def gqa_ragged_prefill_causal_kernel(
Q_ptr, # *bf16 [num_q, 32, 128]
K_ptr, # *bf16 [num_kv, 8, 128]
V_ptr, # *bf16 [num_kv, 8, 128]
OUT_ptr, # *bf16 [num_q, 32, 128]
LSE_ptr, # *fp32 [num_q, 32]
NUM_Q: tl.constexpr, # number of query tokens in this sequence
NUM_KV: tl.constexpr, # number of kv tokens in this sequence
DELTA: tl.constexpr, # NUM_KV - NUM_Q
SM_SCALE: tl.constexpr, # softmax scale (float32)
gqa_ratio: tl.constexpr, # 4 (32 / 8)
BLOCK_D: tl.constexpr, # 128
BLOCK_K: tl.constexpr, # 64 / 128
N_BLOCKS_K: tl.constexpr, # ceil_div(NUM_KV, BLOCK_K)
):
"""
One program = one (query_token, qo_head) pair.
program_id(0) := query token in [0, NUM_Q)
program_id(1) := qo head index in [0, 32)
"""
pid_q = tl.program_id(0)
pid_h = tl.program_id(1)
# Out-of-bounds queries are ignored (host pads the launch grid if needed).
if pid_q >= NUM_Q:
return
d = tl.arange(0, BLOCK_D) # [0 .. 127]
# -------------------- Load Q ------------------------------------------------
q_ptrs = Q_ptr + (pid_q * 32 + pid_h) * BLOCK_D + d # [D] strides
q_vec = tl.load(q_ptrs).to(tl.float32) # [D] (fp32)
kv_head = pid_h // gqa_ratio # 0 .. 7
# -------------------- Running accumulators -------------------------------
m_prev = tl.full((), -float("inf"), tl.float32) # running max
s_prev = tl.zeros((), tl.float32) # running sum(exp)
acc_prev = tl.zeros((BLOCK_D,), tl.float32) # running weighted value sum
allowed_k = tl.minimum(pid_q + 1 + DELTA, NUM_KV) # causal upper-bound
ln2_const = 0.6931471805599453 # ln(2)
for block_idx in tl.static_range(N_BLOCKS_K):
kv_idx_base = block_idx * BLOCK_K
kv_offsets = kv_idx_base + tl.arange(0, BLOCK_K) # [BLOCK_K]
mask_k = kv_offsets < allowed_k # bool
# ------------- Load K ------------------------------------------------
k_ptrs = (
K_ptr
+ ((kv_offsets[:, None] * 8 + kv_head) * BLOCK_D)
+ d[None, :]
)
k_tile = tl.load(k_ptrs, mask=mask_k[:, None], other=0.0).to(tl.float32)
# k_tile: [BLOCK_K, D]
# ------------- Compute Scores ----------------------------------------
scores = tl.sum(k_tile * q_vec[None, :], axis=1) # [BLOCK_K]
scores = scores * SM_SCALE
scores = tl.where(mask_k, scores, -float("inf"))
block_max = tl.max(scores, axis=0)
m_new = tl.maximum(m_prev, block_max)
exp_scores = tl.exp(scores - m_new)
exp_scores = tl.where(mask_k, exp_scores, 0.0)
alpha = tl.exp(m_prev - m_new)
s_new = s_prev * alpha + tl.sum(exp_scores, axis=0)
# ------------- Load V -----------------------------------------------
v_ptrs = (
V_ptr
+ ((kv_offsets[:, None] * 8 + kv_head) * BLOCK_D)
+ d[None, :]
)
v_tile = tl.load(v_ptrs, mask=mask_k[:, None], other=0.0).to(tl.float32)
# v_tile: [BLOCK_K, D]
attn_v = tl.sum(v_tile * exp_scores[:, None], axis=0) # [D]
acc_new = acc_prev * alpha + attn_v
# update running state
m_prev = m_new
s_prev = s_new
acc_prev = acc_new
# -------------------- Finalize & Store ------------------------------------
zero_mask = s_prev == 0.0
out_vec = tl.where(zero_mask, tl.zeros_like(acc_prev), acc_prev / s_prev)
lse_val = tl.where(
zero_mask,
-float("inf"),
(tl.log(s_prev) + m_prev) / ln2_const,
)
# store output
out_ptrs = OUT_ptr + (pid_q * 32 + pid_h) * BLOCK_D + d
tl.store(out_ptrs, out_vec.to(tl.bfloat16))
lse_ptr = LSE_ptr + pid_q * 32 + pid_h
tl.store(lse_ptr, lse_val)
# -----------------------------------------------------------------------------#
# Python Wrapper #
# -----------------------------------------------------------------------------#
@torch.no_grad()
def run(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
sm_scale: float | None = None,
):
"""
Entry point that mimics the reference interface.
Handles device placement, per-sequence kernel launches, and result gathering.
"""
if not torch.cuda.is_available():
raise RuntimeError("CUDA device is required to run the Triton kernel.")
# ---------------------------- constants -----------------------------------
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(128.0)
sm_scale = float(sm_scale)
BLOCK_D = 128
BLOCK_K = 64
GQA_RATIO = 4
NUM_WARPS = 4
NUM_STAGES = 2
# ------------------------ helpers -----------------------------------------
def _to_cuda(t: torch.Tensor):
return t.cuda() if not t.is_cuda else t
def _maybe_cpu(t: torch.Tensor, ref: torch.Tensor):
return t.cpu() if not ref.is_cuda else t
# -------------------- move inputs to GPU ----------------------------------
q_d = _to_cuda(q)
k_d = _to_cuda(k)
v_d = _to_cuda(v)
qo_indptr_d = _to_cuda(qo_indptr)
kv_indptr_d = _to_cuda(kv_indptr)
total_q = int(q_d.shape[0])
total_kv = int(k_d.shape[0])
# -------------------- allocate outputs ------------------------------------
output_d = torch.empty(
(total_q, 32, 128), dtype=torch.bfloat16, device=q_d.device
)
lse_d = torch.empty((total_q, 32), dtype=torch.float32, device=q_d.device)
# -------------------- per-sequence launch ---------------------------------
len_indptr = int(qo_indptr_d.shape[0])
for b in range(len_indptr - 1):
q_start = int(qo_indptr_d[b].item())
q_end = int(qo_indptr_d[b + 1].item())
kv_start = int(kv_indptr_d[b].item())
kv_end = int(kv_indptr_d[b + 1].item())
num_q = q_end - q_start
num_kv = kv_end - kv_start
if num_q <= 0 or num_kv <= 0:
continue
delta = num_kv - num_q
n_blocks_k = (num_kv + BLOCK_K - 1) // BLOCK_K
q_seq = q_d[q_start:q_end].contiguous()
k_seq = k_d[kv_start:kv_end].contiguous()
v_seq = v_d[kv_start:kv_end].contiguous()
out_seq = output_d[q_start:q_end]
lse_seq = lse_d[q_start:q_end]
grid = (triton.cdiv(num_q, 1), 32)
gqa_ragged_prefill_causal_kernel[grid](
q_seq,
k_seq,
v_seq,
out_seq,
lse_seq,
num_q,
num_kv,
delta,
sm_scale,
gqa_ratio=GQA_RATIO,
BLOCK_D=BLOCK_D,
BLOCK_K=BLOCK_K,
N_BLOCKS_K=n_blocks_k,
num_warps=NUM_WARPS,
num_stages=NUM_STAGES,
)
# -------------------- restore to CPU if needed ----------------------------
output = _maybe_cpu(output_d, q)
lse = _maybe_cpu(lse_d, q)
return output, lsescrolls · 217 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON