gpt-5-2025-08-07_triton_657308
gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 440 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-657308?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32, 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:c03e9d523983ade21c790201f164f3a02af2324a9046fbe226112145c89a703f
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.
num-warps = 8
num_warps=8,stages = 4
num_stages=4Kernel source
main.py440 lines
import math
from typing import Any, Dict, Tuple
import torch
import triton
import triton.language as tl
VOCAB_SIZE = 129280 # constant per spec
@triton.jit
def _argmax_kernel(
probs_ptr, # *const float32
row_ids_ptr, # *const int32 (indices into batch)
out_ptr, # *mut int64 (write results at absolute row index)
V: tl.constexpr, # vocab size
BLOCK_SIZE: tl.constexpr # tile size along vocab
):
pid = tl.program_id(0)
rid = tl.load(row_ids_ptr + pid, mask=True, other=0).to(tl.int32)
base = rid * V
best_val = tl.float32(-1.0e30)
best_idx = tl.int32(0)
for off in range(0, V, BLOCK_SIZE):
idxs = off + tl.arange(0, BLOCK_SIZE)
mask = idxs < V
vals = tl.load(probs_ptr + base + idxs, mask=mask, other=tl.float32(-1.0e30))
local_max = tl.max(vals, axis=0)
local_arg = tl.argmax(vals, axis=0) # index within the tile
g_idx = off + local_arg
better = local_max > best_val
best_val = tl.where(better, local_max, best_val)
best_idx = tl.where(better, g_idx, best_idx)
tl.store(out_ptr + rid, best_idx.to(tl.int64))
@triton.jit
def _sample_full_kernel(
probs_ptr, # *const float32
row_ids_ptr, # *const int32
rand_ptr, # *const float32 (one uniform [0,1) per row)
out_ptr, # *mut int64
V: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
SEG: tl.constexpr
):
pid = tl.program_id(0)
rid = tl.load(row_ids_ptr + pid, mask=True, other=0).to(tl.int32)
base = rid * V
# Pass 1: total sum
total = tl.float32(0.0)
for off in range(0, V, BLOCK_SIZE):
idxs = off + tl.arange(0, BLOCK_SIZE)
mask = idxs < V
vals = tl.load(probs_ptr + base + idxs, mask=mask, other=tl.float32(0.0))
total += tl.sum(vals, axis=0)
u = tl.load(rand_ptr + rid)
# Clamp to strictly less than total to avoid boundary issues
eps = tl.float32(1e-7)
target = u * total
target = tl.where(target >= total, total - eps * total, target)
prefix = tl.float32(0.0)
res_idx = tl.int32(-1)
for off in range(0, V, BLOCK_SIZE):
idxs_block = off + tl.arange(0, BLOCK_SIZE)
mask_block = idxs_block < V
vals_block = tl.load(probs_ptr + base + idxs_block, mask=mask_block, other=tl.float32(0.0))
block_sum = tl.sum(vals_block, axis=0)
# If threshold is not in this block, skip
in_block = (res_idx < 0) & (prefix + block_sum > target)
# If in this block, find the exact index
if in_block:
# segmented search to limit inner unroll
for so in range(0, BLOCK_SIZE, SEG):
j = so + tl.arange(0, SEG)
mask_seg = mask_block & (j < BLOCK_SIZE)
vals_seg = tl.load(probs_ptr + base + off + j, mask=mask_seg, other=tl.float32(0.0))
seg_sum = tl.sum(vals_seg, axis=0)
in_seg = (res_idx < 0) & (prefix + seg_sum > target)
# If the target is in this segment, do a linear search in the small segment
if in_seg:
# Linear search within the small segment (SEG is small, e.g., 128)
for t in range(SEG):
v = vals_seg[t]
prefix = prefix + v
found_now = (res_idx < 0) & (prefix > target)
idx_found = off + so + t
res_idx = tl.where(found_now, idx_found.to(tl.int32), res_idx)
else:
# target not in this segment
prefix = prefix + seg_sum
else:
prefix = prefix + block_sum
# Safety: if due to numeric precision res_idx is still -1, choose last valid index
res_idx = tl.where(res_idx < 0, (V - 1).to(tl.int32), res_idx)
tl.store(out_ptr + rid, res_idx.to(tl.int64))
@triton.jit
def _topk_sample_kernel(
probs_ptr, # *const float32
topk_ptr, # *const int32
topp_ptr, # *const float32
rand_ptr, # *const float32
row_ids_ptr, # *const int32
out_ptr, # *mut int64
V: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
K_MAX: tl.constexpr
):
pid = tl.program_id(0)
rid = tl.load(row_ids_ptr + pid, mask=True, other=0).to(tl.int32)
base = rid * V
k_val = tl.load(topk_ptr + rid).to(tl.int32)
p_val = tl.load(topp_ptr + rid)
u_val = tl.load(rand_ptr + rid)
# Buffers for top-k
sel_idx = tl.full([K_MAX], -1, dtype=tl.int32)
sel_val = tl.zeros([K_MAX], dtype=tl.float32)
# Iteratively select global top-k (k <= K_MAX) using masked argmax
for t in range(K_MAX):
active = t < k_val
# initialize best for this iteration
best_val = tl.float32(-1.0e30)
best_gidx = tl.int32(0)
for off in range(0, V, BLOCK_SIZE):
idxs = off + tl.arange(0, BLOCK_SIZE)
mask = idxs < V
vals = tl.load(probs_ptr + base + idxs, mask=mask, other=tl.float32(-1.0e30))
# Mask out previously selected indices
# Build a blocked mask: True if idx equals any sel_idx[j] for j < t
blocked = tl.zeros([BLOCK_SIZE], dtype=tl.int1)
for j in range(K_MAX):
if j < t:
sidx = sel_idx[j]
blocked = blocked | (idxs == sidx)
masked_vals = tl.where(blocked, tl.float32(-1.0e30), vals)
local_max = tl.max(masked_vals, axis=0)
local_arg = tl.argmax(masked_vals, axis=0)
g_idx = off + local_arg
better = (local_max > best_val) & active
best_val = tl.where(better, local_max, best_val)
best_gidx = tl.where(better, g_idx, best_gidx)
# Write selected
if active:
sel_idx[t] = best_gidx
sel_val[t] = best_val
# Sum of top-k values
sum_k = tl.float32(0.0)
for t in range(K_MAX):
if t < k_val:
sum_k += sel_val[t]
# Determine allowed prefix count under top-p
use_topp = p_val < 1.0
threshold = tl.where(use_topp, p_val * sum_k, sum_k)
# Compute minimal m such that prefix >= threshold; guarantee at least one
m_count = tl.int32(0)
cum = tl.float32(0.0)
for t in range(K_MAX):
if t < k_val:
cum = cum + sel_val[t]
set_now = use_topp & (m_count == 0) & (cum >= threshold)
m_count = tl.where(set_now, (t + 1).to(tl.int32), m_count)
allowed_count = tl.where(use_topp, m_count, k_val)
# sum over allowed prefix for sampling
sum_allowed = tl.float32(0.0)
for t in range(K_MAX):
if t < allowed_count:
sum_allowed += sel_val[t]
# Guard against degenerate sums
eps = tl.float32(1e-7)
sum_allowed = tl.where(sum_allowed <= tl.float32(0.0), eps, sum_allowed)
r = u_val * sum_allowed
# Clamp r to strictly less than sum_allowed
r = tl.where(r >= sum_allowed, sum_allowed - eps * sum_allowed, r)
# Sample within the allowed prefix
pref = tl.float32(0.0)
drawn_pos = tl.int32(0)
taken = False
for t in range(K_MAX):
if t < allowed_count:
pref = pref + sel_val[t]
take_now = (not taken) & (pref > r)
drawn_pos = tl.where(take_now, t.to(tl.int32), drawn_pos)
# "taken" cannot be changed directly as Python bool; emulate
taken = taken | (pref > r)
out_idx = sel_idx[drawn_pos]
tl.store(out_ptr + rid, out_idx.to(tl.int64))
def _to_cuda(t: torch.Tensor) -> torch.Tensor:
if t.is_cuda:
return t
if torch.cuda.is_available():
return t.cuda()
raise RuntimeError("CUDA is required but not available; received a CPU tensor.")
def _ensure_dtype(t: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
if t.dtype != dtype:
return t.to(dtype)
return t
def _call_argmax_kernel(probs: torch.Tensor, row_ids: torch.Tensor, out: torch.Tensor):
assert probs.is_cuda and row_ids.is_cuda and out.is_cuda
BLOCK_SIZE = 4096
grid = (row_ids.numel(),)
_argmax_kernel[grid](
probs, row_ids, out,
V=VOCAB_SIZE,
BLOCK_SIZE=BLOCK_SIZE,
num_warps=8,
num_stages=4
)
def _call_sample_full_kernel(probs: torch.Tensor, row_ids: torch.Tensor, rand: torch.Tensor, out: torch.Tensor):
assert probs.is_cuda and row_ids.is_cuda and rand.is_cuda and out.is_cuda
BLOCK_SIZE = 4096
SEG = 128
grid = (row_ids.numel(),)
_sample_full_kernel[grid](
probs, row_ids, rand, out,
V=VOCAB_SIZE,
BLOCK_SIZE=BLOCK_SIZE,
SEG=SEG,
num_warps=8,
num_stages=4
)
def _call_topk_sample_kernel(
probs: torch.Tensor,
top_k: torch.Tensor,
top_p: torch.Tensor,
row_ids: torch.Tensor,
rand: torch.Tensor,
out: torch.Tensor,
k_max: int
):
assert probs.is_cuda and top_k.is_cuda and top_p.is_cuda and row_ids.is_cuda and rand.is_cuda and out.is_cuda
BLOCK_SIZE = 4096
grid = (row_ids.numel(),)
_topk_sample_kernel[grid](
probs, top_k, top_p, rand, row_ids, out,
V=VOCAB_SIZE,
BLOCK_SIZE=BLOCK_SIZE,
K_MAX=k_max,
num_warps=8,
num_stages=4
)
@torch.no_grad()
def run(*args, **kwargs):
"""
Entry point:
run(probs, top_k, top_p) -> samples
Implements top-k then top-p sampling as specified, optimized with Triton kernels on B200.
"""
# Handle both positional and keyword forms
if len(args) == 3 and not kwargs:
probs, top_k, top_p = args
else:
probs = kwargs.get("probs", args[0] if len(args) > 0 else None)
top_k = kwargs.get("top_k", args[1] if len(args) > 1 else None)
top_p = kwargs.get("top_p", args[2] if len(args) > 2 else None)
if probs is None or top_k is None or top_p is None:
raise ValueError("Missing required arguments: probs, top_k, top_p")
# Validate shapes and dtypes
if probs.dim() != 2:
raise ValueError(f"probs must be 2D [batch_size, vocab_size], got {tuple(probs.shape)}")
batch_size, vocab_size = probs.shape
if vocab_size != VOCAB_SIZE:
raise AssertionError(f"vocab_size must be {VOCAB_SIZE}, got {vocab_size}")
if top_k.shape != (batch_size,):
raise ValueError(f"top_k must be shape [{batch_size}], got {tuple(top_k.shape)}")
if top_p.shape != (batch_size,):
raise ValueError(f"top_p must be shape [{batch_size}], got {tuple(top_p.shape)}")
# Keep originals to restore device
out_device = probs.device
# Ensure on CUDA
probs_dev = _to_cuda(_ensure_dtype(probs, torch.float32))
top_k_dev = _to_cuda(_ensure_dtype(top_k, torch.int32))
top_p_dev = _to_cuda(_ensure_dtype(top_p, torch.float32))
# Allocate output on device
samples_dev = torch.empty(batch_size, dtype=torch.int64, device=probs_dev.device)
# Random uniforms per row for sampling kernels
rand = torch.rand(batch_size, dtype=torch.float32, device=probs_dev.device)
# Build row masks
with torch.no_grad():
k = top_k_dev
p = top_p_dev
# Masks
mask_p_le_zero = p <= 0.0
mask_use_topk_kernel = (p > 0.0) & (k > 0) & (k < VOCAB_SIZE) # further restricted by K_MAX later
mask_full_sampling = (p >= 1.0) & ((k <= 0) | (k >= VOCAB_SIZE))
# The remainder will use a GPU PyTorch fallback for exactness
# We will use a reasonable K_MAX for the Triton top-k kernel
K_MAX = 128
# Split the top-k mask based on K_MAX
mask_topk_small = mask_use_topk_kernel & (k <= K_MAX)
mask_topk_large = mask_use_topk_kernel & (k > K_MAX)
# Category 1: p <= 0 -> argmax (no need to apply top-k since argmax is invariant)
idxs = torch.nonzero(mask_p_le_zero, as_tuple=False).flatten()
if idxs.numel() > 0:
row_ids = idxs.to(torch.int32).contiguous()
_call_argmax_kernel(probs_dev, row_ids, samples_dev)
# Category 2: 0 < p, 0 < k < V, k <= K_MAX -> Triton top-k + top-p selection + sampling
idxs = torch.nonzero(mask_topk_small, as_tuple=False).flatten()
if idxs.numel() > 0:
row_ids = idxs.to(torch.int32).contiguous()
_call_topk_sample_kernel(
probs_dev,
top_k_dev,
top_p_dev,
row_ids=row_ids,
rand=rand,
out=samples_dev,
k_max=K_MAX
)
# Category 3: p >= 1 and (k <= 0 or k >= V) -> sample from full distribution
idxs = torch.nonzero(mask_full_sampling, as_tuple=False).flatten()
if idxs.numel() > 0:
row_ids = idxs.to(torch.int32).contiguous()
_call_sample_full_kernel(probs_dev, row_ids, rand, samples_dev)
# Category 4: Fallback exact GPU path using PyTorch ops for all remaining rows
remaining_mask = ~(mask_p_le_zero | mask_topk_small | mask_full_sampling)
idxs = torch.nonzero(remaining_mask, as_tuple=False).flatten()
if idxs.numel() > 0:
# Process each row independently for exactness, on GPU
for rid in idxs.tolist():
row = probs_dev[rid]
ki = int(k[rid].item())
pi = float(p[rid].item())
# Apply top-k filtering if needed
if 0 < ki < VOCAB_SIZE:
vals, idx_sorted = torch.sort(row, descending=True)
keep_idx_k = idx_sorted[:ki]
filtered_k = torch.zeros_like(row)
filtered_k[keep_idx_k] = row[keep_idx_k]
row_work = filtered_k
else:
row_work = row
# Apply top-p if needed
if pi <= 0.0:
# This shouldn't happen due to mask, but keep for safety
samples_dev[rid] = torch.argmax(row_work).to(torch.int64)
continue
if pi < 1.0:
vals, idx_sorted = torch.sort(row_work, descending=True)
cdf = torch.cumsum(vals, dim=0)
to_remove = cdf > pi
if VOCAB_SIZE > 1:
to_remove[1:] = to_remove[:-1].clone()
to_remove[0] = False
keep_idx_p = idx_sorted[~to_remove]
filtered_p = torch.zeros_like(row_work)
filtered_p[keep_idx_p] = row_work[keep_idx_p]
row_final = filtered_p
else:
row_final = row_work
# Renormalize and sample
s = row_final.sum()
if s.item() <= 0.0:
# Degenerate: pick argmax
samples_dev[rid] = torch.argmax(row).to(torch.int64)
else:
probs_vec = row_final / s
draw = torch.multinomial(probs_vec, 1, replacement=True).squeeze(0)
samples_dev[rid] = draw.to(torch.int64)
# Move to original device if needed
if samples_dev.device != out_device:
samples = samples_dev.to(out_device)
else:
samples = samples_dev
return samples
if __name__ == "__main__":
# Simple sanity check (will run on CUDA if available)
bs = 4
V = VOCAB_SIZE
device = "cuda" if torch.cuda.is_available() else "cpu"
torch.manual_seed(0)
probs = torch.rand(bs, V, device=device, dtype=torch.float32)
probs = probs / probs.sum(dim=1, keepdim=True)
top_k = torch.tensor([50, 0, 100, 10], device=device, dtype=torch.int32)
top_p = torch.tensor([0.9, 1.0, 0.95, 0.0], device=device, dtype=torch.float32)
out = run(probs, top_k, top_p)
print("Samples:", out)scrolls · 440 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON