gpt-5-2025-08-07 / tritonda906d
gpt-5-2025-08-07_triton_da906d · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 194 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-da906d?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:638e4458bb0f554ac090b6793ad99dc5013e5c0e419e14c4c40ee7035e1a65e3
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.py194 lines
import math
import torch
import triton
import triton.language as tl
# Kernel: per-row sampling
# - If k <= 0 or k >= vocab_size: sample from full distribution (unnormalized probs)
# - Else: write sentinel -1 to signal host-side top-k sampling fallback
@triton.jit
def _sample_full_or_signal_topk(
probs_ptr, # float32* [B, V]
topk_ptr, # int32* [B]
rand_ptr, # float32* [B], uniform in [0, 1)
samples_ptr, # int64* [B]
stride_row, # int32 row stride in elements
vocab_size: tl.constexpr, # constexpr = 151936
BLOCK_N: tl.constexpr, # tile size along vocab dimension
):
pid = tl.program_id(axis=0)
row_ptr = probs_ptr + pid * stride_row
# Load k and random u for this row
k = tl.load(topk_ptr + pid)
u = tl.load(rand_ptr + pid)
# Decide path
k_no_filter = (k <= 0) | (k >= vocab_size)
# Default: signal fallback (sentinel -1)
selected_idx = tl.full((), -1, dtype=tl.int32)
if k_no_filter:
# First pass: total sum across the row
total_sum = tl.zeros((), dtype=tl.float32)
for start in range(0, vocab_size, BLOCK_N):
offs = start + tl.arange(0, BLOCK_N)
mask = offs < vocab_size
vals = tl.load(row_ptr + offs, mask=mask, other=0.0)
total_sum += tl.sum(vals, axis=0)
# Handle degenerate case: if total_sum <= 0, return argmax index (tie -> smallest idx)
if total_sum <= 0:
neg_inf = tl.full((), -float("inf"), dtype=tl.float32)
best_val = neg_inf
best_idx = tl.full((), 0, dtype=tl.int32)
big_int = tl.full((), 2147483647, dtype=tl.int32)
for start in range(0, vocab_size, BLOCK_N):
offs = start + tl.arange(0, BLOCK_N)
mask = offs < vocab_size
vals = tl.load(row_ptr + offs, mask=mask, other=0.0)
# invalid lanes -> -inf
vals = tl.where(mask, vals, neg_inf)
tile_max = tl.max(vals, axis=0)
is_eq = vals == tile_max
tie_idx = tl.min(tl.where(is_eq, offs, big_int), axis=0)
take = tile_max > best_val
best_val = tl.where(take, tile_max, best_val)
best_idx = tl.where(take, tie_idx, best_idx)
selected_idx = best_idx
else:
# Second pass: sample categorical by threshold t = u * total_sum
t = u * total_sum
prefix = tl.zeros((), dtype=tl.float32)
found = tl.full((), 0, dtype=tl.int32)
big_int = tl.full((), 2147483647, dtype=tl.int32)
for start in range(0, vocab_size, BLOCK_N):
offs = start + tl.arange(0, BLOCK_N)
mask = offs < vocab_size
vals = tl.load(row_ptr + offs, mask=mask, other=0.0)
block_sum = tl.sum(vals, axis=0)
need = (found == 0) & (prefix + block_sum >= t)
if need:
# Find the first index within this block where cumsum crosses t
v = tl.where(mask, vals, 0.0)
csum = tl.cumsum(v, axis=0)
thr = t - prefix
cross = csum >= thr
cand = tl.where(cross, offs, big_int)
pick = tl.min(cand, axis=0)
selected_idx = pick
found = tl.full((), 1, dtype=tl.int32)
else:
# update prefix only if not yet found
prefix = tl.where(found == 0, prefix + block_sum, prefix)
# Fallback in case of numerical issues: last index
selected_idx = tl.where(found == 1, selected_idx, vocab_size - 1)
# Store result as int64
tl.store(samples_ptr + pid, tl.cast(selected_idx, tl.int64))
def _ensure_cuda_device(t: torch.Tensor, name: str) -> torch.device:
if t.is_cuda:
return t.device
if torch.cuda.is_available():
return torch.device("cuda")
raise RuntimeError(f"CUDA is required to run this kernel, but {name} is on CPU and no CUDA device is available.")
def _move_to_device(t: torch.Tensor, device: torch.device) -> torch.Tensor:
if t.device != device:
return t.to(device, non_blocking=True)
return t
@torch.no_grad()
def run(*args, **kwargs):
# Parse inputs
if len(args) >= 2:
probs, top_k = args[0], args[1]
else:
probs = kwargs.get("probs", None)
top_k = kwargs.get("top_k", None)
if probs is None or top_k is None:
raise ValueError("Missing required arguments: probs and top_k")
if probs.dim() != 2:
raise ValueError("probs must be a 2D tensor of shape [batch_size, vocab_size]")
batch_size, vocab_size = probs.shape
if vocab_size != 151936:
raise AssertionError(f"Expected vocab_size=151936, got {vocab_size}")
# Ensure dtype
probs = probs.to(dtype=torch.float32)
# Device management
device = _ensure_cuda_device(probs, "probs")
_ = _ensure_cuda_device(top_k, "top_k") # just to validate availability
probs_gpu = _move_to_device(probs, device)
top_k_gpu = _move_to_device(top_k, device).to(dtype=torch.int32)
# Output and RNG
samples_gpu = torch.empty((batch_size,), dtype=torch.int64, device=device)
rand_gpu = torch.rand((batch_size,), dtype=torch.float32, device=device)
# Kernel launch params
stride_row = probs_gpu.stride(0)
grid = (batch_size,)
# BLOCK_N tuned for large V on B200; adjust if needed
BLOCK_N = 2048
_sample_full_or_signal_topk[grid](
probs_gpu,
top_k_gpu,
rand_gpu,
samples_gpu,
stride_row,
vocab_size=vocab_size,
BLOCK_N=BLOCK_N,
num_warps=8,
num_stages=3,
)
# Host-side fallback for rows requiring top-k filtering (0 < k < vocab_size)
# We detect those rows by sentinel -1 written by the kernel.
with torch.no_grad():
sentinel = (samples_gpu == -1)
if sentinel.any():
idx_rows = torch.nonzero(sentinel, as_tuple=False).squeeze(1)
for r in idx_rows.tolist():
row = probs_gpu[r]
k = int(top_k_gpu[r].item())
# Only do top-k filtering if 0 < k < vocab_size; otherwise, keep kernel result
if 0 < k < vocab_size:
# Select top-k tokens and sample among them
top_vals, top_idx = torch.topk(row, k, largest=True, sorted=False)
denom = top_vals.sum()
if float(denom.item()) > 0.0:
weights = top_vals / denom
sel_local = torch.multinomial(weights, 1, replacement=True).item()
token = int(top_idx[sel_local].item())
else:
token = int(torch.argmax(row).item())
samples_gpu[r] = token
else:
# Shouldn't happen because kernel handled this case, but in case it does:
# sample from full distribution via torch.multinomial
denom = row.sum()
if float(denom.item()) > 0.0:
weights = row / denom
samples_gpu[r] = int(torch.multinomial(weights, 1, replacement=True).item())
else:
samples_gpu[r] = int(torch.argmax(row).item())
# Move to original device if needed
if probs.device.type != "cuda":
samples_out = samples_gpu.cpu()
else:
samples_out = samples_gpu
return samples_outscrolls · 194 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON