gpt-o3_triton_861794
gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 196 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-861794?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32
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:839570e32823689b47dd611afab3a8592ac897feb1f95c95a768017b124f279e
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 = 4
num_warps=4, # 4 × 32 = 128 threads – matches BLOCKKernel source
main.py196 lines
import math
from typing import Any, Dict
import torch
import triton
import triton.language as tl
# ---------------------------------------------------------------------------- #
# Kernel #
# ---------------------------------------------------------------------------- #
@triton.jit
def _top_p_kernel(
probs_ptr, # *fp32 – [batch, vocab]
top_p_ptr, # *fp32 – [batch]
rand_ptr, # *fp32 – [batch] (one U(0,1) per row)
out_ptr, # *int64 – [batch]
vocab_size: tl.constexpr,
BLOCK: tl.constexpr, # 128
):
"""
One Triton program = exactly one sequence (= one matrix row).
Behaviour (per-row, decided from `top_p`):
• p ≤ 0 -> greedy arg-max
• 0 < p < 1 -> nucleus sampling (row was pre-filtered)
• p ≥ 1 -> vanilla multinomial sampling
"""
pid = tl.program_id(axis=0) # row id
row_ptr = probs_ptr + pid * vocab_size # start of this row
# Per-row params ----------------------------------------------------------------
p_val = tl.load(top_p_ptr + pid) # nucleus threshold
r_val = tl.load(rand_ptr + pid) # uniform random in [0,1)
greedy = p_val <= 0.0 # bool – take arg-max?
# Running state ------------------------------------------------------------------
best_val = tl.full((), -math.inf, tl.float32) # running maximum value
best_idx = tl.full((), 0, tl.int32) # arg-max index
sample_idx = tl.full((), -1, tl.int32) # -1 → not chosen yet
cum_prob = tl.zeros((), tl.float32) # running CDF (for sampling)
finished = tl.zeros((), tl.int1) # exits early when sampled
# Static column offsets (0 … BLOCK-1)
offs = tl.arange(0, BLOCK)
# Loop over the vocabulary -------------------------------------------------------
start = tl.zeros((), tl.int32)
while (start < vocab_size) & (finished == 0):
idxs = start + offs # absolute indices
mask = idxs < vocab_size # boundary mask
vals = tl.load(row_ptr + idxs,
mask=mask,
other=0.0) # [BLOCK] – fp32
# -------- greedy path: track block maximum ----------------------------------
block_max = tl.max(vals, axis=0)
same_max = vals == block_max
first_max = tl.where(same_max, offs, BLOCK)
block_arg = tl.min(first_max, axis=0) + start
is_better = block_max > best_val
best_val = tl.where(is_better, block_max, best_val)
best_idx = tl.where(is_better, block_arg, best_idx)
# -------- sampling path: inverse CDF scan -----------------------------------
prefix = tl.cumsum(vals) # inclusive scan over block
hit = (sample_idx < 0) & (cum_prob + prefix >= r_val)
hit_off = tl.where(hit, offs, BLOCK)
firstHit = tl.min(hit_off, axis=0)
got_it = firstHit < BLOCK
sample_idx = tl.where(got_it & (sample_idx < 0),
start + firstHit,
sample_idx)
finished = finished | ((~greedy) & got_it)
# advance --------------------------------------------------------------------
cum_prob += tl.sum(vals, axis=0)
start += BLOCK
# -------------------------------------------------------------------------------
chosen = tl.where(greedy, best_idx, sample_idx)
chosen = tl.where(chosen < 0, best_idx, chosen) # numerical-safety fallback
tl.store(out_ptr + pid, chosen.to(tl.int64))
# ---------------------------------------------------------------------------- #
# GPU-side nucleus (top-p) filter #
# ---------------------------------------------------------------------------- #
def _gpu_top_p_filter(probs: torch.Tensor,
top_p: torch.Tensor) -> torch.Tensor:
"""
In-place style (returns a clone) nucleus filtering on GPU.
Only rows with 0 < p < 1 are processed, others are unchanged.
Every processed row is re-normalised to sum to 1.
"""
out = probs.clone() # keeps dtype / device
mask = (top_p > 0.0) & (top_p < 1.0)
if not mask.any():
return out
rows = out[mask] # [M, V]
p_thr = top_p[mask].unsqueeze(1) # [M, 1]
# Full sort – still the simplest + fastest for very large vocab on GPU
vals, idx = torch.sort(rows, dim=-1, descending=True)
cdf = vals.cumsum(dim=-1)
remove = cdf > p_thr
shift = torch.zeros_like(remove, dtype=torch.bool)
shift[:, 1:] = remove[:, :-1] # keep first token ≥ threshold
keep = ~shift
kept_vals = torch.where(keep, vals, torch.zeros_like(vals))
filtered = torch.zeros_like(rows)
filtered.scatter_(1, idx, kept_vals)
# Re-normalise (row-wise)
row_sum = filtered.sum(dim=1, keepdim=True)
filtered /= row_sum
out[mask] = filtered
return out
# ---------------------------------------------------------------------------- #
# Public entry point #
# ---------------------------------------------------------------------------- #
def run(
probs: torch.Tensor,
top_p: torch.Tensor,
*kernel_args: Any,
**kernel_kwargs: Dict[str, Any],
) -> torch.Tensor:
"""
Fast top-p / multinomial sampler (B200-optimised).
Steps:
1. Device housekeeping.
2. Optional nucleus filtering (GPU).
3. Generate one U(0,1) number per sequence.
4. Launch Triton kernel (1 program / sequence, 128 threads, 4 warps).
5. Return samples on original device.
"""
# ---- device handling ----------------------------------------------------------
if not torch.cuda.is_available():
raise RuntimeError("CUDA device required to run the Triton kernel")
orig_device = probs.device # remember caller’s device
probs_gpu = probs.to(torch.float32)
top_p_gpu = top_p.to(torch.float32)
if not probs_gpu.is_cuda:
probs_gpu = probs_gpu.cuda()
if not top_p_gpu.is_cuda:
top_p_gpu = top_p_gpu.cuda()
batch, vocab = probs_gpu.shape
if vocab != 151_936:
raise ValueError(f"Expected vocab_size == 151 936, got {vocab}")
# ---- pre-processing: nucleus filter ------------------------------------------
probs_ready = _gpu_top_p_filter(probs_gpu, top_p_gpu)
# ---- random numbers -----------------------------------------------------------
rand_row = torch.rand(batch, dtype=torch.float32, device=probs_ready.device)
# ---- output buffer ------------------------------------------------------------
out = torch.empty(batch, dtype=torch.int64, device=probs_ready.device)
# ---- kernel launch ------------------------------------------------------------
BLOCK = 128
grid = (batch,)
_top_p_kernel[grid](
probs_ready,
top_p_gpu,
rand_row,
out,
vocab_size=vocab,
BLOCK=BLOCK,
num_warps=4, # 4 × 32 = 128 threads – matches BLOCK
*kernel_args,
**kernel_kwargs,
)
# ---- bring back to caller’s device -------------------------------------------
if not probs.is_cuda:
out = out.to(orig_device)
return outscrolls · 196 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON