gpt-o3 / tritondeaf62
gpt-o3_triton_deaf62 · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 255 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-deaf62?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
No published measurement for this revision.
No evidence · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:798c396221a454a79362500b188e1f32e6600949d92486a9d4ae8ac096c6f16e
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.
mma
scores = tl.dot(q, tl.trans(k_tile)) * sm_scale # [BM, BN]num-warps = 8
NUM_WARPS = 8 # good default for B200 GPUsonline-softmax
m_new = tl.maximum(m_i, m_ij)stages = 1
num_stages=1,tile-m = 64
BLOCK_M = 64 # queries per blocktile-n = 64
BLOCK_N = 64 # keys per blockKernel source
main.py255 lines
import math
import torch
import triton
import triton.language as tl
# -----------------------------------------------------------------------------#
# Global compile-time constants #
# -----------------------------------------------------------------------------#
NUM_QO_HEADS = 32
NUM_KV_HEADS = 4
HEAD_DIM = 128
GQA_RATIO = NUM_QO_HEADS // NUM_KV_HEADS
# Tunable tile sizes for B200
BLOCK_M = 64 # queries per block
BLOCK_N = 64 # keys per block
NUM_WARPS = 8 # good default for B200 GPUs
# -----------------------------------------------------------------------------#
# Triton kernel #
# -----------------------------------------------------------------------------#
@triton.jit
def _gqa_ragged_prefill_kernel(
Q_ptr, K_ptr, V_ptr, # *bf16
O_ptr, LSE_ptr, # *bf16 / *fp32
q_start: tl.int32, # offset of first query token
kv_start: tl.int32, # offset of first kv token
q_len: tl.int32, # number of query tokens
kv_len: tl.int32, # number of kv tokens
delta: tl.int32, # kv_len - q_len
sm_scale: tl.float32, # softmax scale
inv_ln2: tl.float32, # 1 / ln(2)
*,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
HEAD_DIM: tl.constexpr,
NUM_QO_HEADS: tl.constexpr,
NUM_KV_HEADS: tl.constexpr,
GQA_RATIO: tl.constexpr,
):
# ------------------------------------------------------------------#
# Program IDs #
# ------------------------------------------------------------------#
pid_m = tl.program_id(0) # query-block id
pid_h = tl.program_id(1) # qo-head id (0 … 31)
# ------------------------------------------------------------------#
# Index computations #
# ------------------------------------------------------------------#
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) # [BM]
offs_n = tl.arange(0, BLOCK_N) # [BN]
offs_d = tl.arange(0, HEAD_DIM) # [HD]
row_mask = offs_m < q_len # [BM] bool
qo_head = pid_h
kv_head = qo_head // GQA_RATIO
stride_q_token = NUM_QO_HEADS * HEAD_DIM
stride_kv_token = NUM_KV_HEADS * HEAD_DIM
# ------------------------------------------------------------------#
# Load Q #
# ------------------------------------------------------------------#
q_ptrs = (
Q_ptr
+ (q_start + offs_m[:, None]) * stride_q_token
+ qo_head * HEAD_DIM
+ offs_d[None, :]
)
q = tl.load(q_ptrs, mask=row_mask[:, None], other=0).to(tl.float32) # [BM, HD]
# ------------------------------------------------------------------#
# Online softmax initialisation #
# ------------------------------------------------------------------#
NEG_INF = -1.0e30
m_i = tl.full((BLOCK_M,), NEG_INF, dtype=tl.float32)
l_i = tl.zeros((BLOCK_M,), dtype=tl.float32)
acc = tl.zeros((BLOCK_M, HEAD_DIM), dtype=tl.float32)
# ------------------------------------------------------------------#
# Iterate over KV tiles #
# ------------------------------------------------------------------#
kv_tile_start = tl.int32(0)
while kv_tile_start < kv_len:
k_ids = kv_tile_start + offs_n # [BN]
k_valid = k_ids < kv_len # [BN] bool
# ---- load K / V ---------------------------------------------
k_ptrs = (
K_ptr
+ (kv_start + k_ids[:, None]) * stride_kv_token
+ kv_head * HEAD_DIM
+ offs_d[None, :]
)
v_ptrs = (
V_ptr
+ (kv_start + k_ids[:, None]) * stride_kv_token
+ kv_head * HEAD_DIM
+ offs_d[None, :]
)
k_tile = tl.load(k_ptrs, mask=k_valid[:, None], other=0).to(tl.float32) # [BN, HD]
v_tile = tl.load(v_ptrs, mask=k_valid[:, None], other=0).to(tl.float32) # [BN, HD]
# ---- attention scores ----------------------------------------
scores = tl.dot(q, tl.trans(k_tile)) * sm_scale # [BM, BN]
# ---- causal masking ------------------------------------------
allowed_k = offs_m + delta + 1 # [BM]
causal_mask = k_ids[None, :] >= allowed_k[:, None] # [BM, BN]
valid_mask = k_valid[None, :] & (~causal_mask) & row_mask[:, None]
scores = tl.where(valid_mask, scores, NEG_INF)
# ---- online softmax ------------------------------------------
m_ij = tl.max(scores, axis=1) # [BM]
m_new = tl.maximum(m_i, m_ij)
exp_m_i = tl.exp(m_i - m_new)
exp_scores = tl.exp(scores - m_new[:, None]) * valid_mask.to(tl.float32)
l_new = l_i * exp_m_i + tl.sum(exp_scores, axis=1) # [BM]
# update accumulator
pv = tl.dot(exp_scores, v_tile) # [BM, HD]
acc = (acc * (l_i * exp_m_i)[:, None] + pv) / l_new[:, None]
m_i = m_new
l_i = l_new
kv_tile_start += BLOCK_N
# ------------------------------------------------------------------#
# Write back output & LSE #
# ------------------------------------------------------------------#
o_ptrs = (
O_ptr
+ (q_start + offs_m[:, None]) * stride_q_token
+ qo_head * HEAD_DIM
+ offs_d[None, :]
)
tl.store(o_ptrs, acc.to(tl.bfloat16), mask=row_mask[:, None])
lse_vals = (m_i + tl.log(l_i)) * inv_ln2 # [BM]
lse_ptrs = (
LSE_ptr
+ (q_start + offs_m) * NUM_QO_HEADS
+ qo_head
)
tl.store(lse_ptrs, lse_vals, mask=row_mask)
# -----------------------------------------------------------------------------#
# Python wrapper #
# -----------------------------------------------------------------------------#
@torch.no_grad()
def run(q, k, v, qo_indptr, kv_indptr, sm_scale=None):
"""
Optimised Triton implementation of
gqa_ragged_prefill_causal_h32_kv4_d128
"""
# ---------------------------------------------------------------#
# Device management #
# ---------------------------------------------------------------#
if not torch.cuda.is_available():
raise RuntimeError("CUDA device is required to run Triton kernels.")
orig_device = q.device
if orig_device.type == "cpu":
q, k, v = q.cuda(), k.cuda(), v.cuda()
qo_indptr, kv_indptr = qo_indptr.cuda(), kv_indptr.cuda()
elif orig_device.type != "cuda":
raise RuntimeError(f"Unsupported device type: {orig_device.type!r}")
# ---------------------------------------------------------------#
# Shape / constant checks #
# ---------------------------------------------------------------#
total_q, num_qo_heads, head_dim = q.shape
total_kv, num_kv_heads, _ = k.shape
assert num_qo_heads == NUM_QO_HEADS, "num_qo_heads mismatch"
assert num_kv_heads == NUM_KV_HEADS, "num_kv_heads mismatch"
assert head_dim == HEAD_DIM, "head_dim mismatch"
assert total_q == qo_indptr[-1].item(), "total_q != qo_indptr[-1]"
assert total_kv == kv_indptr[-1].item(), "total_kv != kv_indptr[-1]"
# ---------------------------------------------------------------#
# Soft-max scale #
# ---------------------------------------------------------------#
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(head_dim)
sm_scale = float(sm_scale)
inv_ln2 = 1.0 / math.log(2.0)
# ---------------------------------------------------------------#
# Allocate outputs #
# ---------------------------------------------------------------#
output = torch.empty(
(total_q, NUM_QO_HEADS, HEAD_DIM),
dtype=torch.bfloat16,
device=q.device,
)
lse = torch.empty(
(total_q, NUM_QO_HEADS),
dtype=torch.float32,
device=q.device,
)
# ---------------------------------------------------------------#
# Launch kernel for each 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())
if q_start >= q_end or kv_start >= kv_end:
continue # empty slice
q_len = q_end - q_start
kv_len = kv_end - kv_start
delta = kv_len - q_len
grid_m = triton.cdiv(q_len, BLOCK_M)
grid = (grid_m, NUM_QO_HEADS)
_gqa_ragged_prefill_kernel[grid](
q, k, v,
output, lse,
q_start, kv_start,
q_len, kv_len,
delta,
sm_scale,
inv_ln2,
BLOCK_M=BLOCK_M,
BLOCK_N=BLOCK_N,
HEAD_DIM=HEAD_DIM,
NUM_QO_HEADS=NUM_QO_HEADS,
NUM_KV_HEADS=NUM_KV_HEADS,
GQA_RATIO=GQA_RATIO,
num_warps=NUM_WARPS,
num_stages=1,
)
# ---------------------------------------------------------------#
# Move outputs back to original device #
# ---------------------------------------------------------------#
if orig_device.type == "cpu":
output = output.cpu()
lse = lse.cpu()
return output, lsescrolls · 255 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON