gpt-5-2025-08-07 / tritonaf4b72
gpt-5-2025-08-07_triton_af4b72 · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 410 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-af4b72?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:7039f1a2e1cbda09006e854886751531ef63687d1b5dbee05c8fc1f463545b6b
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=4,Kernel source
main.py410 lines
import math
import torch
import triton
import triton.language as tl
VOCAB_SIZE = 128256
@triton.jit
def sample_from_packed_kernel(
probs_ptr, # float32 [total_kept]
idxs_ptr, # int32 [total_kept]
starts_ptr, # int32 [n_rows]
lens_ptr, # int32 [n_rows]
rand_ptr, # float32 [n_rows]
out_ptr, # int64 [n_rows]
BLOCK_SIZE: tl.constexpr,
MAX_TILES: tl.constexpr,
):
pid = tl.program_id(axis=0)
# Load start, length, and per-row random
start = tl.load(starts_ptr + pid)
length = tl.load(lens_ptr + pid)
u = tl.load(rand_ptr + pid)
# Clamp u to [0, 1 - eps) to avoid edge-case where cdf never exceeds u
eps = 1e-7
one_minus_eps = 1.0 - eps
u = tl.where(u < one_minus_eps, u, one_minus_eps)
running = tl.zeros((), dtype=tl.float32)
found = tl.full((), 0, tl.int32)
found_tile = tl.full((), 0, tl.int32)
carry = tl.zeros((), dtype=tl.float32)
ar = tl.arange(0, BLOCK_SIZE)
# Scan tiles to locate the tile containing u
for t in tl.static_range(0, MAX_TILES):
col_base = t * BLOCK_SIZE
rem = length - col_base
has = rem > 0
offs = start + col_base + ar
valid = has & (ar < rem)
vals = tl.load(probs_ptr + offs, mask=valid, other=0.0)
tile_sum = tl.sum(vals, axis=0)
not_found = found == 0
crosses = not_found & has & ((running + tile_sum) > u)
carry = tl.where(crosses, running, carry)
found_tile = tl.where(crosses, tl.full((), t, tl.int32), found_tile)
found = tl.where(crosses, 1, found)
running = running + tl.where(has, tile_sum, 0.0)
# Now search inside the found tile sequentially
col_base2 = found_tile * BLOCK_SIZE
base2 = start + col_base2
rem2 = length - col_base2
target = u - carry
acc = tl.zeros((), dtype=tl.float32)
j = tl.full((), -1, tl.int32)
for i in tl.static_range(0, BLOCK_SIZE):
valid_i = i < rem2
vi = tl.load(probs_ptr + base2 + i, mask=valid_i, other=0.0)
acc = acc + tl.where(valid_i, vi, 0.0)
take = (j < 0) & valid_i & (acc >= target)
j = tl.where(take, tl.full((), i, tl.int32), j)
last_idx = tl.where(rem2 > 0, rem2 - 1, 0)
j = tl.where(j >= 0, j, last_idx)
sel_off = base2 + j
tok = tl.load(idxs_ptr + sel_off).to(tl.int64)
tl.store(out_ptr + pid, tok)
@triton.jit
def sample_from_dense_kernel(
probs_ptr, # float32 [n_rows, VOCAB_SIZE] base pointer
stride_row, # int32 stride in elements between rows
rand_ptr, # float32 [n_rows]
out_ptr, # int64 [n_rows]
N_COLS: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
MAX_TILES: tl.constexpr,
):
pid = tl.program_id(axis=0)
# Pointer to the start of this row
row_ptr = probs_ptr + pid * stride_row
u = tl.load(rand_ptr + pid)
# Clamp u to [0, 1 - eps)
eps = 1e-7
one_minus_eps = 1.0 - eps
u = tl.where(u < one_minus_eps, u, one_minus_eps)
running = tl.zeros((), dtype=tl.float32)
found = tl.full((), 0, tl.int32)
found_tile = tl.full((), 0, tl.int32)
carry = tl.zeros((), dtype=tl.float32)
ar = tl.arange(0, BLOCK_SIZE)
# Scan tiles across the row
for t in tl.static_range(0, MAX_TILES):
col_base = t * BLOCK_SIZE
rem = N_COLS - col_base
has = rem > 0
offs = col_base + ar
valid = has & (ar < rem)
vals = tl.load(row_ptr + offs, mask=valid, other=0.0)
tile_sum = tl.sum(vals, axis=0)
not_found = found == 0
crosses = not_found & has & ((running + tile_sum) > u)
carry = tl.where(crosses, running, carry)
found_tile = tl.where(crosses, tl.full((), t, tl.int32), found_tile)
found = tl.where(crosses, 1, found)
running = running + tl.where(has, tile_sum, 0.0)
# Search within found tile sequentially
col_base2 = found_tile * BLOCK_SIZE
rem2 = N_COLS - col_base2
target = u - carry
acc = tl.zeros((), dtype=tl.float32)
j = tl.full((), -1, tl.int32)
for i in tl.static_range(0, BLOCK_SIZE):
valid_i = i < rem2
vi = tl.load(row_ptr + col_base2 + i, mask=valid_i, other=0.0)
acc = acc + tl.where(valid_i, vi, 0.0)
take = (j < 0) & valid_i & (acc >= target)
j = tl.where(take, tl.full((), i, tl.int32), j)
last_idx = tl.where(rem2 > 0, rem2 - 1, 0)
j = tl.where(j >= 0, j, last_idx)
tok = (col_base2 + j).to(tl.int64)
tl.store(out_ptr + pid, tok)
def run(*args, **kwargs):
"""
Entry point: top_k_top_p_sampling_from_probs_v128256
Inputs:
probs: [batch, 128256] float32
top_k: [batch] int32
top_p: [batch] float32
Output:
samples: [batch] int64
"""
# Accept args or kwargs
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 types
if probs.ndim != 2:
raise ValueError("probs must be 2D [batch, vocab_size]")
if probs.shape[1] != VOCAB_SIZE:
raise AssertionError(f"vocab_size must be {VOCAB_SIZE}, got {probs.shape[1]}")
if top_k.ndim != 1 or top_k.shape[0] != probs.shape[0]:
raise ValueError("top_k must be 1D with length equal to batch size")
if top_p.ndim != 1 or top_p.shape[0] != probs.shape[0]:
raise ValueError("top_p must be 1D with length equal to batch size")
# Device management: ensure CUDA
original_device = probs.device
if original_device.type == "cuda":
device = original_device
else:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required for Triton kernels, but not available.")
device = torch.device("cuda")
# Move inputs to device
probs_dev = probs.to(device=device, dtype=torch.float32, copy=False)
top_k_dev = top_k.to(device=device, dtype=torch.int32, copy=False)
top_p_dev = top_p.to(device=device, dtype=torch.float32, copy=False)
batch = probs_dev.shape[0]
samples_dev = torch.empty(batch, dtype=torch.int64, device=device)
# Rows where p <= 0: select argmax (top-k filtering doesn't change argmax)
mask_p_le_zero = (top_p_dev <= 0.0)
if mask_p_le_zero.any():
idx_rows = mask_p_le_zero.nonzero(as_tuple=False).squeeze(1)
if idx_rows.numel() > 0:
argmax_idx = torch.argmax(probs_dev.index_select(0, idx_rows), dim=1)
samples_dev.index_copy_(0, idx_rows, argmax_idx.to(torch.int64))
# Remaining rows (p > 0) need sampling
mask_remaining = ~mask_p_le_zero
if mask_remaining.any():
rows_remaining = mask_remaining.nonzero(as_tuple=False).squeeze(1)
k_all = top_k_dev.index_select(0, rows_remaining)
p_all = top_p_dev.index_select(0, rows_remaining)
# Case A: No top-k (k <= 0 or k >= vocab) and p >= 1 -> dense sampling from full distribution
mask_no_topk = (k_all <= 0) | (k_all >= VOCAB_SIZE)
mask_p_ge_one = p_all >= 1.0
mask_dense = mask_no_topk & mask_p_ge_one
if mask_dense.any():
dense_rows_local = rows_remaining.index_select(0, mask_dense.nonzero(as_tuple=False).squeeze(1))
if dense_rows_local.numel() > 0:
dense_probs = probs_dev.index_select(0, dense_rows_local).contiguous()
# Normalize to guard against small numerical drift
row_sums = dense_probs.sum(dim=1, keepdim=True)
dense_probs = dense_probs / torch.clamp(row_sums, min=1e-12)
n_dense = dense_probs.shape[0]
rand_u = torch.rand(n_dense, dtype=torch.float32, device=device)
# Configure kernel
BLOCK = 512
MAX_TILES = max(1, triton.cdiv(VOCAB_SIZE, BLOCK))
grid = (n_dense,)
sample_out = torch.empty(n_dense, dtype=torch.int64, device=device)
stride_row = dense_probs.stride(0)
sample_from_dense_kernel[grid](
dense_probs, # probs_ptr
stride_row, # stride_row
rand_u, # rand_ptr
sample_out, # out_ptr
N_COLS=VOCAB_SIZE,
BLOCK_SIZE=BLOCK,
MAX_TILES=MAX_TILES,
num_warps=8,
num_stages=4,
)
samples_dev.index_copy_(0, dense_rows_local, sample_out)
# Case B: Other rows -> build packed candidates via top-k and/or nucleus (top-p), then sample via packed kernel
packed_probs_list = []
packed_idxs_list = []
packed_lens = []
packed_row_ids = []
# Helper for packing row-wise tensors with a boolean keep mask; expects sorted in descending probability
def _pack_kept(vals_sorted: torch.Tensor, idx_sorted: torch.Tensor, keep_mask: torch.Tensor, row_ids: torch.Tensor):
# vals_sorted, idx_sorted, keep_mask: [n, L]
n = vals_sorted.shape[0]
if n == 0:
return
# Per-row lengths
lens = keep_mask.sum(dim=1) # int64
# Masked flatten
flat_vals = vals_sorted[keep_mask]
flat_idx = idx_sorted[keep_mask]
# Renormalize per-row
row_ids_expand = torch.repeat_interleave(torch.arange(n, device=device, dtype=torch.long), lens)
sums = torch.zeros(n, dtype=torch.float32, device=device)
sums.index_add_(0, row_ids_expand, flat_vals)
scales = 1.0 / torch.clamp(sums, min=1e-12)
flat_vals = flat_vals * scales.index_select(0, row_ids_expand)
# Record
packed_lens.extend(lens.to(torch.int32).tolist())
packed_row_ids.extend(row_ids.to(torch.int32).tolist())
packed_probs_list.append(flat_vals)
packed_idxs_list.append(flat_idx.to(torch.int32))
# Precompute masks within 'rows_remaining'
mask_topk = (k_all > 0) & (k_all < VOCAB_SIZE)
mask_no_topk_pcut = mask_no_topk & (p_all < 1.0)
# Process top-k rows grouped by unique k
if mask_topk.any():
rows_topk_local = rows_remaining.index_select(0, mask_topk.nonzero(as_tuple=False).squeeze(1))
k_topk_local = top_k_dev.index_select(0, rows_topk_local)
p_topk_local = top_p_dev.index_select(0, rows_topk_local)
# Split by p >= 1 and p < 1
mask_topk_p_ge_one = (p_topk_local >= 1.0)
mask_topk_p_lt_one = ~mask_topk_p_ge_one
# Group by unique k for p >= 1.0
if mask_topk_p_ge_one.any():
rows_ge1 = rows_topk_local.index_select(0, mask_topk_p_ge_one.nonzero(as_tuple=False).squeeze(1))
k_ge1 = k_topk_local.index_select(0, mask_topk_p_ge_one.nonzero(as_tuple=False).squeeze(1))
uniq_k = torch.unique(k_ge1, sorted=True)
for kv in uniq_k.tolist():
sel = (k_ge1 == kv)
if not sel.any():
continue
grp_rows = rows_ge1.index_select(0, sel.nonzero(as_tuple=False).squeeze(1))
grp_probs = probs_dev.index_select(0, grp_rows)
# topk sorted descending
vals, idx = torch.topk(grp_probs, kv, dim=1, largest=True, sorted=True)
# Normalize top-k distribution
sums = torch.clamp(vals.sum(dim=1, keepdim=True), min=1e-12)
vals = vals / sums
keep_mask = torch.ones_like(vals, dtype=torch.bool)
_pack_kept(vals, idx, keep_mask, grp_rows)
# Group by unique k for p < 1.0
if mask_topk_p_lt_one.any():
rows_lt1 = rows_topk_local.index_select(0, mask_topk_p_lt_one.nonzero(as_tuple=False).squeeze(1))
k_lt1 = k_topk_local.index_select(0, mask_topk_p_lt_one.nonzero(as_tuple=False).squeeze(1))
p_lt1 = p_topk_local.index_select(0, mask_topk_p_lt_one.nonzero(as_tuple=False).squeeze(1))
uniq_k = torch.unique(k_lt1, sorted=True)
for kv in uniq_k.tolist():
sel = (k_lt1 == kv)
if not sel.any():
continue
grp_rows = rows_lt1.index_select(0, sel.nonzero(as_tuple=False).squeeze(1))
grp_probs = probs_dev.index_select(0, grp_rows)
grp_p = p_lt1.index_select(0, sel.nonzero(as_tuple=False).squeeze(1))
# topk sorted descending
vals, idx = torch.topk(grp_probs, kv, dim=1, largest=True, sorted=True)
# Normalize
sums = torch.clamp(vals.sum(dim=1, keepdim=True), min=1e-12)
vals = vals / sums
# Nucleus (top-p) selection
cdf = torch.cumsum(vals, dim=1)
pcol = grp_p.unsqueeze(1)
to_remove = cdf > pcol
if kv > 1:
to_remove[:, 1:] = to_remove[:, :-1].clone()
to_remove[:, 0] = False
elif kv == 1:
to_remove[:, 0] = False
keep_mask = ~to_remove
_pack_kept(vals, idx, keep_mask, grp_rows)
# Process no-topk rows with p in (0,1) -> full sort + nucleus
if mask_no_topk_pcut.any():
rows_ntk = rows_remaining.index_select(0, mask_no_topk_pcut.nonzero(as_tuple=False).squeeze(1))
if rows_ntk.numel() > 0:
p_ntk = top_p_dev.index_select(0, rows_ntk)
probs_ntk = probs_dev.index_select(0, rows_ntk)
# Sort full vocab descending
vals, idx = torch.sort(probs_ntk, dim=1, descending=True)
# Nucleus selection per row
cdf = torch.cumsum(vals, dim=1)
pcol = p_ntk.unsqueeze(1)
to_remove = cdf > pcol
if VOCAB_SIZE > 1:
to_remove[:, 1:] = to_remove[:, :-1].clone()
to_remove[:, 0] = False
else:
to_remove[:, 0] = False
keep_mask = ~to_remove
_pack_kept(vals, idx, keep_mask, rows_ntk)
# Launch packed sampling kernel if we have any rows to sample
if len(packed_lens) > 0:
if len(packed_probs_list) == 1:
packed_probs = packed_probs_list[0]
packed_idxs = packed_idxs_list[0]
else:
packed_probs = torch.cat(packed_probs_list, dim=0)
packed_idxs = torch.cat(packed_idxs_list, dim=0)
lens_tensor = torch.tensor(packed_lens, dtype=torch.int32, device=device)
row_ids_tensor = torch.tensor(packed_row_ids, dtype=torch.int32, device=device)
# Build starts array
starts_tensor = torch.zeros_like(lens_tensor)
if lens_tensor.numel() > 1:
starts_tensor[1:] = torch.cumsum(lens_tensor[:-1], dim=0)
# Generate random uniforms per row
rand_u = torch.rand(lens_tensor.shape[0], dtype=torch.float32, device=device)
# Configure kernel
BLOCK = 512
max_len = int(lens_tensor.max().item())
MAX_TILES = max(1, triton.cdiv(max_len, BLOCK))
grid = (lens_tensor.shape[0],)
sample_out = torch.empty(lens_tensor.shape[0], dtype=torch.int64, device=device)
sample_from_packed_kernel[grid](
packed_probs, # probs_ptr
packed_idxs, # idxs_ptr
starts_tensor, # starts_ptr
lens_tensor, # lens_ptr
rand_u, # rand_ptr
sample_out, # out_ptr
BLOCK_SIZE=BLOCK,
MAX_TILES=MAX_TILES,
num_warps=8,
num_stages=4,
)
# Scatter back to global samples
samples_dev.index_copy_(0, row_ids_tensor.to(torch.long), sample_out)
# Move results to original device if needed
if original_device.type != "cuda":
samples = samples_dev.to(device=original_device)
else:
samples = samples_dev
return samplesscrolls · 410 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON