gpt-5 / triton41ae45
gpt-5_triton_41ae45 · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 264 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-41ae45?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:dceb94f9dc4230e0173844b6e4def21143d27131ec2fb6499eb03ee2e21e1200
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
qk = tl.dot(q_tile, tl.trans(k_tile))num-warps = 8
num_warps = 8stages = 2
num_stages = 2tile-m = 32
BLOCK_M = 32tile-n = 128
BLOCK_N = 128Kernel source
main.py264 lines
import math
import torch
import triton
import triton.language as tl
@triton.jit
def gqa_ragged_prefill_causal_h32_kv4_d128_kernel(
Q_ptr, K_ptr, V_ptr, O_ptr, LSE_ptr,
qo_indptr_ptr, kv_indptr_ptr,
total_q, total_kv,
sm_scale,
stride_q0, stride_q1, stride_q2,
stride_k0, stride_k1, stride_k2,
stride_v0, stride_v1, stride_v2,
stride_o0, stride_o1, stride_o2,
stride_lse0, stride_lse1,
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
):
pid_seq = tl.program_id(0) # sequence id
pid_kvh = tl.program_id(1) # kv head id
pid_mblk = tl.program_id(2) # query block id within sequence
# Load sequence boundaries
q_start = tl.load(qo_indptr_ptr + pid_seq)
q_end = tl.load(qo_indptr_ptr + pid_seq + 1)
kv_start = tl.load(kv_indptr_ptr + pid_seq)
kv_end = tl.load(kv_indptr_ptr + pid_seq + 1)
q_len = q_end - q_start
kv_len = kv_end - kv_start
# Offsets within sequence for queries
m_offsets = pid_mblk * BLOCK_M + tl.arange(0, BLOCK_M)
m_mask = m_offsets < q_len
q_abs = q_start + m_offsets
# Causal delta: kv_len - q_len
delta = kv_len - q_len
# Per-row kv cap: n_cap = m + 1 + delta
n_cap = m_offsets + (1 + delta)
has_attn_row = n_cap > 0
# Dimension offsets
d_offsets = tl.arange(0, HEAD_DIM)
kv_h = pid_kvh
qo_h_base = pid_kvh * GQA_RATIO
NEG_INF = float("-inf")
INV_LN2 = 1.4426950408889634 # 1 / ln(2)
sm_scale_f32 = tl.full([1], sm_scale, tl.float32)[0]
# Iterate over the 8 Qo-heads mapped to this kv head
for h in tl.static_range(GQA_RATIO):
qo_h = qo_h_base + h
# Load Q tile [M, D] in f32
q_ptrs = Q_ptr + q_abs[:, None] * stride_q0 + qo_h * stride_q1 + d_offsets[None, :] * stride_q2
q_tile = tl.load(q_ptrs, mask=m_mask[:, None], other=0.0).to(tl.float32)
# Pass 1: compute per-row max (m_i) over all K tiles with causal mask
m_i = tl.full([BLOCK_M], NEG_INF, tl.float32)
n_start = 0
while n_start < kv_len:
n_offsets = n_start + tl.arange(0, BLOCK_N)
n_inbounds = n_offsets < kv_len
# Load K tile [N, D] for this kv head
k_ptrs = K_ptr + (kv_start + n_offsets)[:, None] * stride_k0 + kv_h * stride_k1 + d_offsets[None, :] * stride_k2
k_tile = tl.load(k_ptrs, mask=n_inbounds[:, None], other=0.0).to(tl.float32)
# QK^T
qk = tl.dot(q_tile, tl.trans(k_tile))
qk_scaled = qk * sm_scale_f32
# Causal mask for this tile
n_base = n_offsets[None, :] # [1, N]
n_cap_broadcast = n_cap[:, None] # [M, 1]
causal_mask = (n_base < n_cap_broadcast) & n_inbounds[None, :] & m_mask[:, None]
# Compute tile max with mask
qk_masked = tl.where(causal_mask, qk_scaled, NEG_INF)
tile_max = tl.max(qk_masked, axis=1)
m_i = tl.maximum(m_i, tile_max)
n_start += BLOCK_N
# Pass 2: compute sum of exp and weighted value accumulation
l_i = tl.zeros([BLOCK_M], tl.float32)
acc = tl.zeros([BLOCK_M, HEAD_DIM], tl.float32)
n_start = 0
while n_start < kv_len:
n_offsets = n_start + tl.arange(0, BLOCK_N)
n_inbounds = n_offsets < kv_len
# Load K and V tiles for this kv head [N, D]
k_ptrs = K_ptr + (kv_start + n_offsets)[:, None] * stride_k0 + kv_h * stride_k1 + d_offsets[None, :] * stride_k2
v_ptrs = V_ptr + (kv_start + n_offsets)[:, None] * stride_v0 + kv_h * stride_v1 + d_offsets[None, :] * stride_v2
k_tile = tl.load(k_ptrs, mask=n_inbounds[:, None], other=0.0).to(tl.float32)
v_tile = tl.load(v_ptrs, mask=n_inbounds[:, None], other=0.0).to(tl.float32)
# QK^T
qk = tl.dot(q_tile, tl.trans(k_tile))
qk_scaled = qk * sm_scale_f32
# Causal mask
n_base = n_offsets[None, :] # [1, N]
n_cap_broadcast = n_cap[:, None] # [M, 1]
causal_mask = (n_base < n_cap_broadcast) & n_inbounds[None, :] & m_mask[:, None]
# Stable logits with global row max m_i
stable_logits = qk_scaled - m_i[:, None]
stable_logits = tl.where(causal_mask, stable_logits, NEG_INF)
# Probabilities and accumulation
p = tl.exp(stable_logits)
l_i += tl.sum(p, axis=1)
acc += tl.dot(p, v_tile)
n_start += BLOCK_N
# Build store mask: only rows with queries and at least one valid key
m_store_mask = m_mask & has_attn_row
# Normalize output
l_i_safe = tl.where(m_store_mask, l_i, 1.0)
out_tile = acc / l_i_safe[:, None]
# Store output
o_ptrs = O_ptr + q_abs[:, None] * stride_o0 + qo_h * stride_o1 + d_offsets[None, :] * stride_o2
tl.store(o_ptrs, out_tile.to(tl.bfloat16), mask=m_store_mask[:, None])
# LSE base-2: (log(sum(exp)) + m_i) / ln(2)
lse_vals = (tl.log(l_i) + m_i) * INV_LN2
lse_ptrs = LSE_ptr + q_abs * stride_lse0 + qo_h * stride_lse1
tl.store(lse_ptrs, lse_vals, mask=m_store_mask)
def _ceil_div_int(a: int, b: int) -> int:
return (a + b - 1) // b
def run(q, k, v, qo_indptr, kv_indptr, sm_scale=None):
# Validate CUDA availability
cuda_available = torch.cuda.is_available()
devices = {
"q": q.device,
"k": k.device,
"v": v.device,
"qo_indptr": qo_indptr.device,
"kv_indptr": kv_indptr.device,
}
target_device = devices["q"]
if not cuda_available:
if any(t.is_cuda for t in [q, k, v, qo_indptr, kv_indptr]):
raise RuntimeError("CUDA is not available but GPU tensors were provided.")
raise RuntimeError("CUDA is required to run Triton kernels.")
# Shapes and checks
total_q, num_qo_heads, head_dim = q.shape
total_kv, num_kv_heads, _ = k.shape
len_indptr = qo_indptr.shape[0]
assert num_qo_heads == 32, "num_qo_heads must be 32"
assert num_kv_heads == 4, "num_kv_heads must be 4"
assert head_dim == 128, "head_dim must be 128"
assert total_q == int(qo_indptr[-1].item()), "total_q must equal qo_indptr[-1]"
assert total_kv == int(kv_indptr[-1].item()), "total_kv must equal kv_indptr[-1]"
assert k.shape == v.shape, "k and v must have same shape"
assert qo_indptr.shape[0] == kv_indptr.shape[0], "qo_indptr and kv_indptr must have same length"
# Default sm_scale
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(head_dim)
# Cast to float32 to avoid 64-bit scalar promotion differences
sm_scale = float(torch.tensor(sm_scale, dtype=torch.float32).item())
# Dtype checks
if q.dtype != torch.bfloat16 or k.dtype != torch.bfloat16 or v.dtype != torch.bfloat16:
raise TypeError("q, k, v must be torch.bfloat16")
if qo_indptr.dtype != torch.int32 or kv_indptr.dtype != torch.int32:
raise TypeError("qo_indptr and kv_indptr must be torch.int32")
compute_device = torch.device("cuda")
# Move to CUDA
q_dev = q if q.device.type == "cuda" else q.to(compute_device, non_blocking=True)
k_dev = k if k.device.type == "cuda" else k.to(compute_device, non_blocking=True)
v_dev = v if v.device.type == "cuda" else v.to(compute_device, non_blocking=True)
qo_indptr_dev = qo_indptr if qo_indptr.device.type == "cuda" else qo_indptr.to(compute_device, non_blocking=True)
kv_indptr_dev = kv_indptr if kv_indptr.device.type == "cuda" else kv_indptr.to(compute_device, non_blocking=True)
# Prepare outputs on device
out_dev = torch.zeros((total_q, num_qo_heads, head_dim), dtype=torch.bfloat16, device=compute_device)
lse_dev = torch.full((total_q, num_qo_heads), float("-inf"), dtype=torch.float32, device=compute_device)
# Early exit if no sequences
num_seqs = len_indptr - 1
if num_seqs <= 0 or total_q == 0 or total_kv == 0:
target_out = out_dev if target_device.type == "cuda" else out_dev.to(target_device, non_blocking=True)
target_lse = lse_dev if target_device.type == "cuda" else lse_dev.to(target_device, non_blocking=True)
return target_out, target_lse
# Constants
GQA_RATIO = 8
NUM_QO_HEADS = 32
NUM_KV_HEADS = 4
HEAD_DIM = 128
# Block sizes tuned conservatively for B200
BLOCK_M = 32
BLOCK_N = 128
# Number of M blocks per sequence, use max across sequences for grid; masking handles others
qo_indptr_cpu = qo_indptr_dev.detach().cpu()
q_lengths = (qo_indptr_cpu[1:] - qo_indptr_cpu[:-1]).to(torch.int64)
if q_lengths.numel() > 0:
max_q_blocks = int(((q_lengths + (BLOCK_M - 1)) // BLOCK_M).max().item())
if max_q_blocks <= 0:
max_q_blocks = 1
else:
max_q_blocks = 1
# Strides
stride_q0, stride_q1, stride_q2 = q_dev.stride()
stride_k0, stride_k1, stride_k2 = k_dev.stride()
stride_v0, stride_v1, stride_v2 = v_dev.stride()
stride_o0, stride_o1, stride_o2 = out_dev.stride()
stride_lse0, stride_lse1 = lse_dev.stride()
grid = (num_seqs, NUM_KV_HEADS, max_q_blocks)
num_warps = 8
num_stages = 2
gqa_ragged_prefill_causal_h32_kv4_d128_kernel[grid](
q_dev, k_dev, v_dev, out_dev, lse_dev,
qo_indptr_dev, kv_indptr_dev,
total_q, total_kv,
sm_scale,
stride_q0, stride_q1, stride_q2,
stride_k0, stride_k1, stride_k2,
stride_v0, stride_v1, stride_v2,
stride_o0, stride_o1, stride_o2,
stride_lse0, stride_lse1,
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=num_stages
)
# Move outputs back to original device of q
if target_device.type != "cuda":
out_host = out_dev.to(target_device, non_blocking=True)
lse_host = lse_dev.to(target_device, non_blocking=True)
else:
out_host = out_dev
lse_host = lse_dev
return out_host, lse_hostscrolls · 264 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON