gpt-5-2025-08-07_triton_e65787
gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 277 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-e65787?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:01b9edb886e80161050e8ab89f93ebbf954341eee5da6c910adee45652f54f91
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.
Kernel source
main.py277 lines
import math
import torch
import triton
import triton.language as tl
V_CONST = 129280 # DeepSeek V3 vocabulary size (constant)
@triton.jit
def sample_full_kernel(
probs_ptr, # *f32 [B, V_CONST]
rand_ptr, # *f32 [n_rows]
rows_idx_ptr, # *i32 [n_rows] mapping from local row-id -> global row-id
out_ptr, # *i64 [B]
stride_probs, # i64 stride between rows in probs (in elements)
n_rows, # i32 number of rows to process in this launch
V: tl.constexpr, # vocab size (constexpr)
BLOCK: tl.constexpr # tile width along vocab dimension
):
pid = tl.program_id(axis=0)
if pid >= n_rows:
return
# Load the mapped global row index and compute base pointer for that row
row_global_i32 = tl.load(rows_idx_ptr + pid)
row_global = row_global_i32.to(tl.int64)
row_ptr = probs_ptr + row_global * stride_probs
# Load the uniform random number in [0, 1)
u = tl.load(rand_ptr + pid, eviction_policy='evict_last')
# Pass 1: compute total mass (sum of probabilities/weights)
total = tl.zeros((), dtype=tl.float32)
for off in range(0, V, BLOCK):
offs = off + tl.arange(0, BLOCK)
mask = offs < V
p = tl.load(row_ptr + offs, mask=mask, other=0.0)
total += tl.sum(p, axis=0)
# Threshold in [0, total]
threshold = u * total
# Pass 2: scan CDF and find first index where CDF >= threshold
cdf = tl.zeros((), dtype=tl.float32)
chosen = tl.full((), -1, dtype=tl.int64)
large = tl.full((), V + 1, dtype=tl.int64)
for off in range(0, V, BLOCK):
offs = off + tl.arange(0, BLOCK)
mask = offs < V
p = tl.load(row_ptr + offs, mask=mask, other=0.0)
pref = tl.cumsum(p, axis=0) + cdf
hit = pref >= threshold
idxs = (offs).to(tl.int64)
hit_idxs = tl.where(hit & mask, idxs, large)
first = tl.min(hit_idxs, axis=0)
found = first < large
chosen = tl.where((chosen < 0) & found, first, chosen)
cdf += tl.sum(p, axis=0)
# Fallback: if no element found due to numerical issues, select the last index
chosen = tl.where(chosen < 0, tl.full((), V - 1, dtype=tl.int64), chosen)
# Write result to the correct global row position
tl.store(out_ptr + row_global, chosen)
@triton.jit
def sample_topk_kernel(
vals_ptr, # *f32 [G, K]
inds_ptr, # *i64 [G, K]
rand_ptr, # *f32 [G]
rows_idx_ptr, # *i32 [G] mapping to global rows
out_ptr, # *i64 [B]
stride_vals, # i64 stride between rows in vals (in elements)
stride_inds, # i64 stride between rows in inds (in elements)
n_rows, # i32 number of rows in this group
K: tl.constexpr, # number of columns (top-k) for this group (constexpr)
BLOCK: tl.constexpr # tile width along K
):
pid = tl.program_id(axis=0)
if pid >= n_rows:
return
# Pointers to this local row
row_vals_ptr = vals_ptr + pid * stride_vals
row_inds_ptr = inds_ptr + pid * stride_inds
# Mapped global row id for storing final answer
row_global_i32 = tl.load(rows_idx_ptr + pid)
row_global = row_global_i32.to(tl.int64)
# Random uniform in [0, 1)
u = tl.load(rand_ptr + pid, eviction_policy='evict_last')
# Pass 1: total mass
total = tl.zeros((), dtype=tl.float32)
for off in range(0, K, BLOCK):
offs = off + tl.arange(0, BLOCK)
mask = offs < K
v = tl.load(row_vals_ptr + offs, mask=mask, other=0.0)
total += tl.sum(v, axis=0)
# Threshold in [0, total]
threshold = u * total
# Pass 2: scan CDF across K and select first where CDF >= threshold
cdf = tl.zeros((), dtype=tl.float32)
chosen_local = tl.full((), -1, dtype=tl.int64)
large = tl.full((), K + 1, dtype=tl.int64)
for off in range(0, K, BLOCK):
offs = off + tl.arange(0, BLOCK)
mask = offs < K
v = tl.load(row_vals_ptr + offs, mask=mask, other=0.0)
pref = tl.cumsum(v, axis=0) + cdf
hit = pref >= threshold
idxs = offs.to(tl.int64)
hit_idxs = tl.where(hit & mask, idxs, large)
first = tl.min(hit_idxs, axis=0)
found = first < large
chosen_local = tl.where((chosen_local < 0) & found, first, chosen_local)
cdf += tl.sum(v, axis=0)
# If not found (extreme numerical edge), choose last position
chosen_local = tl.where(chosen_local < 0, tl.full((), K - 1, dtype=tl.int64), chosen_local)
# Map to original vocab index using inds_ptr
orig_idx = tl.load(row_inds_ptr + chosen_local)
tl.store(out_ptr + row_global, orig_idx)
def _ensure_cuda_tensor(t: torch.Tensor, like: torch.device) -> torch.Tensor:
if t.is_cuda:
if t.device != like:
return t.to(like)
return t
else:
return t.to(like)
def run(probs, top_k):
"""
Triton-accelerated top-k sampling from probability rows.
Inputs:
probs: [batch_size, 129280] float32 (probabilities after softmax)
top_k: [batch_size] int32, per-row top-k to consider. If k <= 0 or k >= 129280, no filtering.
Output:
samples: [batch_size] int64 sampled indices per row
"""
# Basic validation
if not isinstance(probs, torch.Tensor) or not isinstance(top_k, torch.Tensor):
raise TypeError("probs and top_k must be torch.Tensor")
if probs.dim() != 2:
raise ValueError(f"probs must be 2D [B, V], got shape {tuple(probs.shape)}")
B, V = probs.shape
if V != V_CONST:
raise AssertionError(f"Expected vocab_size == {V_CONST}, got {V}")
# DType checks/conversions
if probs.dtype != torch.float32:
probs = probs.to(torch.float32)
if top_k.dtype != torch.int32:
top_k = top_k.to(torch.int32)
# Device management
want_cuda = True # We must run Triton; ensure we are on CUDA
if want_cuda and not torch.cuda.is_available():
raise RuntimeError("CUDA is required to run Triton kernels, but torch.cuda.is_not_available().")
# Track original device for returning output
orig_device = probs.device
# Move inputs to CUDA if needed
device = torch.device("cuda") if not probs.is_cuda else probs.device
probs = _ensure_cuda_tensor(probs, device)
top_k = _ensure_cuda_tensor(top_k, device)
# Prepare output on device
samples = torch.empty(B, dtype=torch.int64, device=device)
# Common strides
stride_probs = probs.stride(0)
# Determine which rows use filtering
valid_mask = (top_k > 0) & (top_k < V_CONST)
invalid_mask = ~valid_mask
# 1) Invalid k: sample directly from full distribution
if invalid_mask.any():
rows_invalid = torch.nonzero(invalid_mask, as_tuple=False).squeeze(1).to(torch.int32)
n_invalid = rows_invalid.numel()
if n_invalid > 0:
rand = torch.rand(n_invalid, device=device, dtype=torch.float32)
grid = (triton.cdiv(n_invalid, 1),)
sample_full_kernel[grid](
probs,
rand,
rows_invalid,
samples,
stride_probs,
n_invalid,
V=V_CONST,
BLOCK=2048,
num_warps=8,
num_stages=4,
)
# 2) Valid k: group by unique k and process each group
if valid_mask.any():
unique_k = torch.unique(top_k[valid_mask], sorted=False)
# Ensure unique_k on device
unique_k = unique_k.to(device)
for k_val in unique_k.tolist():
k_int = int(k_val)
group_mask = valid_mask & (top_k == k_int)
rows_group = torch.nonzero(group_mask, as_tuple=False).squeeze(1)
if rows_group.numel() == 0:
continue
# Gather rows and compute top-k per row using PyTorch (highly-optimized)
sub_probs = probs.index_select(0, rows_group)
# topk returns values and indices along dim=1; order within top-k doesn't affect sampling correctness
vals, inds = torch.topk(sub_probs, k=k_int, dim=1, largest=True, sorted=False)
# Normalize to probabilities (avoid division-by-zero by adding tiny eps)
sums = vals.sum(dim=1, keepdim=True)
# In case of extreme edge (row all zeros) - keep numeric safety
eps = 0.0
vals = vals / (sums + eps)
G = rows_group.numel()
rows_group_i32 = rows_group.to(torch.int32)
rand = torch.rand(G, device=device, dtype=torch.float32)
grid = (triton.cdiv(G, 1),)
# Choose a practical block for K scanning; process in tiles if needed
BLOCK_K = 256
sample_topk_kernel[grid](
vals,
inds,
rand,
rows_group_i32,
samples,
vals.stride(0),
inds.stride(0),
G,
K=k_int,
BLOCK=BLOCK_K,
num_warps=4,
num_stages=3,
)
# Return samples on original device
if samples.device != orig_device:
return samples.to(orig_device)
return samples
if __name__ == "__main__":
# Minimal sanity check (not exhaustive)
B = 4
V = V_CONST
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if device.type != "cuda":
raise RuntimeError("This script requires CUDA to run.")
torch.manual_seed(0)
probs = torch.randn(B, V, device=device, dtype=torch.float32)
probs = torch.softmax(probs, dim=1)
top_k = torch.tensor([0, 1, 32, V_CONST], device=device, dtype=torch.int32)
out = run(probs, top_k)
print("Samples:", out)scrolls · 277 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON