gpt-o3 / triton75f9e9
gpt-o3_triton_75f9e9 · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 179 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-75f9e9?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:07455d1f83ccc9b3cbc19f5e96f4e61eb152bb27d3b7c03ad6f2b4bea486f793
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 8
num_warps=8, # 256 threads / 8 warpsKernel source
main.py179 lines
import math
from typing import Any
import torch
import triton
import triton.language as tl
###############################################################################
# Kernel : draw ONE sample from ONE categorical distribution #
###############################################################################
@triton.jit
def _cdf_sample_kernel(
probs_ptr, # *f32 – [batch , vocab]
rand_ptr, # *f32 – [batch] (0 ≤ r < 1)
out_ptr, # *i64 – [batch]
VOCAB_SIZE: tl.constexpr, # 129 280
BLOCK_SIZE: tl.constexpr # 128 / 256 / …
):
"""
Each Triton program processes exactly ONE row.
We iterate over the row BLOCK_SIZE tokens at a time while
maintaining a running prefix-sum. The first entry whose
cumulative probability strictly exceeds the random threshold
is selected.
"""
pid = tl.program_id(axis=0) # row index
row_start = probs_ptr + pid * VOCAB_SIZE
thresh = tl.load(rand_ptr + pid) # 0 ≤ thresh < 1
lane_off = tl.arange(0, BLOCK_SIZE) # 0 … BLOCK_SIZE-1
running = tl.zeros((), tl.float32) # prefix sum of previous blocks
chosen = tl.full((), -1, tl.int32) # –1 ⇒ not found yet
base_idx = tl.zeros((), tl.int32) # first token handled by block
# ---------------------------------------------------------------- main scan
while (base_idx < VOCAB_SIZE) & (chosen < 0):
idx = base_idx + lane_off
mask = idx < VOCAB_SIZE # guard against OOB accesses
# 1. load probabilities of the current chunk
p = tl.load(row_start + idx, mask=mask, other=0.0)
# 2. cumulative sum *inside* this block + running prefix
local_cdf = tl.cumsum(p, axis=0) + running
# NOTE: we need a STRICT comparison here. If `thresh` is 0
# we must pick the first *positive* probability entry,
# not a zero-probability token that happens to precede it.
crossed = mask & (local_cdf > thresh)
# 3. first index in this block that crosses the threshold
INF = tl.full((BLOCK_SIZE,), BLOCK_SIZE, idx.dtype)
cand_off = tl.where(crossed, lane_off, INF)
min_off = tl.min(cand_off, axis=0)
found = min_off < BLOCK_SIZE
first_idx = base_idx + min_off
chosen = tl.where(found & (chosen < 0), first_idx, chosen)
# 4. advance to next block
running += tl.sum(p, axis=0)
base_idx += BLOCK_SIZE
# Numerical fallback – should never trigger
chosen = tl.where(chosen < 0, VOCAB_SIZE - 1, chosen)
tl.store(out_ptr + pid, chosen.to(tl.int64))
###############################################################################
# Fast batched top-k filtering (host side, PyTorch) #
###############################################################################
def _vectorised_topk_filter(
probs: torch.Tensor,
top_k: torch.Tensor,
vocab_size: int,
) -> torch.Tensor:
"""
For every row i with 0 < k_i < vocab_size:
• keep exactly the k_i largest probabilities
• set all remaining entries to 0
The rows are NOT renormalised here – the caller does that afterwards.
"""
need = (top_k > 0) & (top_k < vocab_size)
if not torch.any(need):
return probs
filtered = probs.clone()
rows = torch.nonzero(need, as_tuple=False).squeeze(1)
ks = top_k[rows]
k_max = int(ks.max().item())
# sorted=True guarantees that the first k_i entries
# correspond to the k_i largest tokens of each row
vals, idxs = torch.topk(
filtered[rows],
k_max,
dim=1,
largest=True,
sorted=True,
)
col_idx = torch.arange(k_max, device=probs.device)
keep_mask = col_idx.unsqueeze(0) < ks.unsqueeze(1)
vals = vals * keep_mask
filtered[rows].zero_()
filtered[rows].scatter_(1, idxs, vals)
return filtered
###############################################################################
# Public API #
###############################################################################
def run(
probs: torch.Tensor,
top_k: torch.Tensor,
*args: Any,
**kwargs: Any,
) -> torch.Tensor:
"""
Parameters
----------
probs : [batch , 129280] – soft-maxed probabilities (float16/bfloat16/float32)
top_k : [batch] int32 – per-row top-k
(0 or ≥ vocab_size ⇒ keep row unchanged)
Returns
-------
samples : [batch] int64 – one sampled token id per input row
"""
if not torch.cuda.is_available():
raise RuntimeError("A CUDA-capable device is required to run this kernel.")
# ---------------------------------------------------------------- device juggling
src_device = probs.device
cuda_dev = torch.device("cuda")
probs_fp32 = probs.to(device=cuda_dev, dtype=torch.float32, copy=False)
topk_i32 = top_k.to(device=cuda_dev, dtype=torch.int32, copy=False)
batch, vocab = probs_fp32.shape
if vocab != 129_280:
raise ValueError(f"vocab_size must be exactly 129 280, got {vocab}")
# ---------------------------------------------------------------- top-k filter
probs_filt = _vectorised_topk_filter(probs_fp32, topk_i32, vocab)
# final normalisation (guards against FP drift)
row_sums = probs_filt.sum(dim=1, keepdim=True)
# If a row became all-zero (should not happen), fall back to the original row
probs_norm = torch.where(
row_sums > 0,
probs_filt / row_sums.clamp(min=1e-7),
probs_fp32,
)
# ---------------------------------------------------------------- RNG (uniform in [0, 1))
rnd = torch.rand(batch, device=cuda_dev, dtype=torch.float32)
# ---------------------------------------------------------------- launch kernel
out = torch.empty(batch, device=cuda_dev, dtype=torch.int64)
BLOCK = 256 # empirically a good fit for B200
_cdf_sample_kernel[(batch,)](
probs_norm,
rnd,
out,
VOCAB_SIZE=vocab,
BLOCK_SIZE=BLOCK,
num_warps=8, # 256 threads / 8 warps
)
# ---------------------------------------------------------------- restore device
return out.to(src_device, non_blocking=True)scrolls · 179 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON