gpt-o3 / tritonfaf01e
gpt-o3_triton_faf01e · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 212 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-faf01e?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:622709b8a0f1dc1f37c14a612059950bb7a54de2893ab30481dfc2a244db262d
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20
Kernel source
main.py212 lines
import math
import torch
import triton
import triton.language as tl
###############################################################################
# Constants
###############################################################################
VOCAB_SIZE: int = 129_280 # fixed DeepSeek-V3 vocabulary size
BLOCK_SIZE: int = 1024 # tokens processed per kernel iteration
###############################################################################
# Triton kernel
###############################################################################
@triton.jit
def _sample_kernel(
probs_ptr, # float32 [batch, vocab]
rand_ptr, # float32 [batch]
sample_ptr, # int64 [batch]
stride_row, # stride between consecutive rows (vocab_size)
vocab_size: tl.constexpr, # 129_280 (compile–time constant)
BLOCK: tl.constexpr # 1 024 (compile–time constant)
):
"""
Each Triton program samples ONE sequence (one distribution / row).
The probabilities in `probs_ptr` MUST already be:
• filtered (top-k / top-p) and
• re-normalised so that they sum to 1.
"""
pid = tl.program_id(axis=0) # sequence id
row_offset = pid * stride_row # start of this row
row_ptr = probs_ptr + row_offset # pointer to first prob
rng = tl.load(rand_ptr + pid) # U(0,1) for this row
# Running state ----------------------------------------------------------
cumsum_before = tl.zeros((), dtype=tl.float32) # cumulative mass processed
found = tl.zeros((), dtype=tl.int32) # 0 -> still searching
chosen_idx = tl.full((), -1, dtype=tl.int32) # result placeholder
# Utility: thread-local contiguous indices 0 … BLOCK-1
idx_in_block = tl.arange(0, BLOCK)
# Iterate over the vocabulary ------------------------------------------------
for offs in range(0, vocab_size, BLOCK):
global_idx = offs + idx_in_block
block_mask = global_idx < vocab_size
# load current chunk of probabilities
probs = tl.load(row_ptr + global_idx, mask=block_mask, other=0.0)
# sum of this BLOCK across all threads
block_sum = tl.sum(probs, axis=0)
# If we haven’t found the token yet and the running cumulative mass
# crosses our random number *inside* this block, we must identify it.
search_block = (found == 0) & (cumsum_before + block_sum > rng)
# Prefix sums of probs within the block (only matters when searching)
prefix = tl.cumsum(probs, axis=0)
# Candidate positions: where prefix exceeds the residual mass
residual = rng - cumsum_before
in_prefix = prefix > residual
candidate = tl.where(search_block & in_prefix, idx_in_block,
BLOCK) # sentinel
# First index in this block that satisfies the predicate
first_in_blk = tl.min(candidate, axis=0)
# If a valid index was found, record the global position
is_valid = first_in_blk < BLOCK
chosen_idx = tl.where(is_valid & (found == 0),
offs + first_in_blk, chosen_idx)
found = tl.where(is_valid, 1, found)
# advance cumulative mass (only while still searching)
cumsum_before += tl.where(found == 0, block_sum,
tl.zeros((), dtype=tl.float32))
# Fallback (numerical safety) – never happens in theory
chosen_idx = tl.where(found == 0, vocab_size - 1, chosen_idx)
# Write result as int64
tl.store(sample_ptr + pid, chosen_idx.to(tl.int64))
###############################################################################
# Python wrapper
###############################################################################
def _ensure_cuda(t: torch.Tensor) -> torch.Tensor:
"""Move tensor to CUDA if it is on CPU. Raises if CUDA is unavailable."""
if t.device.type == "cuda":
return t
if not torch.cuda.is_available():
raise RuntimeError("CUDA device is required for this kernel.")
return t.cuda()
@torch.no_grad()
def run(probs: torch.Tensor,
top_k: torch.Tensor,
top_p: torch.Tensor) -> torch.Tensor:
"""
Top-k / Top-p sampling implemented with a mix of high-level PyTorch
primitives (for filtering) and a custom Triton kernel (for the final
draw). The output exactly matches the reference implementation.
"""
# --------------------------------------------------------------------- #
# 1. Device management & dtype normalisation
# --------------------------------------------------------------------- #
orig_device = probs.device
probs = probs.to(torch.float32)
top_k = top_k.to(torch.int32)
top_p = top_p.to(torch.float32)
probs_gpu = _ensure_cuda(probs)
k_gpu = _ensure_cuda(top_k)
p_gpu = _ensure_cuda(top_p)
batch, vocab = probs_gpu.shape
if vocab != VOCAB_SIZE:
raise ValueError(f"Expected vocab_size={VOCAB_SIZE}, got {vocab}")
samples = torch.empty(batch, dtype=torch.int64, device=probs_gpu.device)
# --------------------------------------------------------------------- #
# 2. Per-row filtering (top-k, top-p) — executed in PyTorch
# --------------------------------------------------------------------- #
rows_for_kernel = []
for i in range(batch):
row = probs_gpu[i]
k = int(k_gpu[i].item())
p = float(p_gpu[i].item())
# -------- top-k --------------------------------------------------
if 0 < k < VOCAB_SIZE:
vals, idx = torch.topk(row, k, largest=True, sorted=False)
mask = torch.zeros_like(row, dtype=torch.bool)
mask[idx] = True
row = row * mask.float()
row /= row.sum()
# deterministic maximum if nucleus threshold <= 0
if p <= 0.0:
samples[i] = torch.argmax(row).to(torch.int64)
probs_gpu[i] = row # store (normalised) for completeness
continue
# -------- top-p --------------------------------------------------
if p < 1.0:
vals, sidx = torch.sort(row, descending=True)
cdf = torch.cumsum(vals, 0)
remove = cdf > p
if VOCAB_SIZE > 1:
remove[1:] = remove[:-1].clone()
remove[0] = False
keep = sidx[~remove]
mask = torch.zeros_like(row, dtype=torch.bool)
mask[keep] = True
row = row * mask.float()
row /= row.sum()
# row now sums to 1 → store back
probs_gpu[i] = row
rows_for_kernel.append(i)
# --------------------------------------------------------------------- #
# 3. Sampling rows with stochastic nucleus — Triton kernel
# --------------------------------------------------------------------- #
if rows_for_kernel:
idx_tensor = torch.tensor(rows_for_kernel,
dtype=torch.int64,
device=probs_gpu.device)
sub_probs = probs_gpu.index_select(0, idx_tensor).contiguous()
rand_vec = torch.rand(len(rows_for_kernel),
dtype=torch.float32,
device=probs_gpu.device)
out_buf = torch.empty(len(rows_for_kernel),
dtype=torch.int64,
device=probs_gpu.device)
grid = (sub_probs.shape[0],)
_sample_kernel[grid](
sub_probs, rand_vec, out_buf,
sub_probs.stride(0), # stride between rows
vocab_size=VOCAB_SIZE,
BLOCK=BLOCK_SIZE,
)
samples.index_copy_(0, idx_tensor, out_buf)
# --------------------------------------------------------------------- #
# 4. Move result back to original device (if needed)
# --------------------------------------------------------------------- #
if orig_device.type == "cuda":
return samples.to(orig_device)
return samples.cpu()
# When this file is executed directly (not imported) -------------------------
if __name__ == "__main__":
# Quick sanity check
torch.manual_seed(0)
bs = 4
logits = torch.randn(bs, VOCAB_SIZE, dtype=torch.float32)
probs = torch.softmax(logits, dim=-1)
top_k = torch.tensor([50, 0, 10, VOCAB_SIZE], dtype=torch.int32)
top_p = torch.tensor([0.95, -1.0, 0.9, 0.0], dtype=torch.float32)
samples_out = run(probs, top_k, top_p)
print("Sampled indices:", samples_out)scrolls · 212 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON