gpt-o3 / tritonb8d9a9
gpt-o3_triton_b8d9a9 · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 202 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-b8d9a9?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:4e2c83a08b372ccbbe6b2e981838b2e91af813a0858c01b2173d37695f2d3f54
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, # execution parametersstages = 2
num_stages=2,Kernel source
main.py202 lines
import math
from typing import Sequence, Mapping, Any, List
import torch
import triton
import triton.language as tl
# ----------------------------------------------------------------------
# Problem-level constants (compile-time)
# ----------------------------------------------------------------------
VOCAB_SIZE: int = 128_256 # fixed for Llama-3.1
BLOCK_SIZE: int = 1_024 # number of tokens processed per loop
N_BLOCKS: int = (VOCAB_SIZE + BLOCK_SIZE - 1) // BLOCK_SIZE # = 126
# ----------------------------------------------------------------------
# Triton kernel – inverse-CDF sampling of ONE distribution row
# ----------------------------------------------------------------------
@triton.jit
def _sample_kernel(
probs_ptr, # *f32 [batch, VOCAB_SIZE]
rand_ptr, # *f32 [batch] – uniform[0,1)
out_ptr, # *i64 [batch]
stride_row, # i32 leading stride between rows
vocab_size: tl.constexpr,
BLOCK: tl.constexpr,
N_BLKS: tl.constexpr,
):
"""
One kernel instance (= program) handles ONE row of probabilities.
We scan the cumulative distribution until it crosses a random
threshold `r` and return the corresponding index.
"""
pid = tl.program_id(axis=0) # row id
row_ptr = probs_ptr + pid * stride_row # pointer to first element in row
r = tl.load(rand_ptr + pid) # threshold in (0, 1)
# running cumulative probability *before* current block
cum_sum = tl.zeros((), dtype=tl.float32)
# best index found so far (init to sentinel > vocab_size-1)
sentinel = vocab_size
best_ix = tl.full((), sentinel, dtype=tl.int32)
# ------------------------------------------------------------------
# iterate over blocks of size `BLOCK`
# ------------------------------------------------------------------
for blk in tl.static_range(N_BLKS):
start = blk * BLOCK
offs = tl.arange(0, BLOCK)
idxs = start + offs # absolute token indices
valid = idxs < vocab_size # mask for short last block
# load probabilities
probs = tl.load(row_ptr + idxs, mask=valid, other=0.0) # [BLOCK]
# inclusive prefix inside the block + previous cum_sum
cdf_blk = tl.cumsum(probs, axis=0) + cum_sum
# first positions where CDF ≥ r
crosses = (cdf_blk >= r) & valid
cand = tl.where(crosses, idxs, sentinel).to(tl.int32)
# first crossing inside the block
first_in_blk = tl.min(cand, axis=0)
# keep leftmost crossing overall
best_ix = tl.where(first_in_blk < best_ix, first_in_blk, best_ix)
# advance cumulative sum
cum_sum += tl.sum(probs, axis=0)
# safeguard – if nothing selected (due to tiny numerical error) pick last vocab
best_ix = tl.where(best_ix == sentinel, vocab_size - 1, best_ix)
# write result
tl.store(out_ptr + pid, best_ix.to(tl.int64))
# ----------------------------------------------------------------------
# Helper – build per-row nucleus (top-p) distribution 100 % on GPU
# ----------------------------------------------------------------------
def _build_nucleus_distribution(row: torch.Tensor, p_thresh: float) -> torch.Tensor:
"""
Keep the minimal prefix whose cumulative probability reaches `p_thresh`
(== nucleus / top-p). Returns a re-normalised probability vector.
All operations happen on `row.device` (GPU for performance).
"""
if p_thresh >= 1.0:
return row
# sort in descending order
vals, idx = torch.sort(row, descending=True)
cdf = torch.cumsum(vals, dim=0)
# mask: remove everything AFTER (not incl.) the first entry that makes CDF > p
to_remove = cdf > p_thresh
to_remove[1:] = to_remove[:-1].clone()
to_remove[0] = False
keep_idx = idx[~to_remove]
filtered = torch.zeros_like(row)
filtered[keep_idx] = row[keep_idx]
total = filtered.sum()
if total > 0:
filtered /= total
return filtered
# ----------------------------------------------------------------------
# Public API
# ----------------------------------------------------------------------
def run(
probs: torch.Tensor,
top_p: torch.Tensor,
*args: Sequence[Any],
**kwargs: Mapping[str, Any],
) -> torch.Tensor:
"""
Parameters
----------
probs : [batch, 128256] float32 – soft-maxed probabilities
top_p : [batch] float32 – per-row nucleus threshold
Returns
-------
samples : [batch] int64 – sampled token indices
"""
# --------------------------- validation ---------------------------
if probs.ndim != 2:
raise ValueError("`probs` must be 2-D [batch, vocab]")
batch, vocab = probs.shape
if vocab != VOCAB_SIZE:
raise ValueError(f"vocab_size must be {VOCAB_SIZE}, got {vocab}")
if top_p.shape != (batch,):
raise ValueError("`top_p` must have shape [batch]")
if not torch.cuda.is_available():
raise RuntimeError("CUDA device required but not available")
# ---------------------- device management ------------------------
orig_device = probs.device
dev = probs.device if probs.is_cuda else torch.device("cuda")
probs_gpu = probs.to(dev, dtype=torch.float32, non_blocking=True)
top_p_gpu = top_p.to(dev, dtype=torch.float32, non_blocking=True)
# ----------------------- pre-processing --------------------------
samples = torch.empty(batch, dtype=torch.int64, device=dev)
# indices of rows that NEED sampling through the kernel
rows_to_sample: List[int] = []
nucleus_rows = []
for i in range(batch):
p_thr = float(top_p_gpu[i].item())
row = probs_gpu[i]
# p ≤ 0 → greedy argmax
if p_thr <= 0.0:
samples[i] = torch.argmax(row).to(torch.int64)
continue
filt_row = _build_nucleus_distribution(row, p_thr)
# extremely rare – if nucleus empty fall back to argmax
if filt_row.sum() == 0:
samples[i] = torch.argmax(row).to(torch.int64)
continue
rows_to_sample.append(i)
nucleus_rows.append(filt_row)
# ---------------------- call Triton kernel -----------------------
if rows_to_sample:
sel_idx = torch.tensor(rows_to_sample, device=dev, dtype=torch.int64)
# stack selected rows into a single 2-D tensor for the kernel
probs_sel = torch.stack(nucleus_rows, dim=0).contiguous()
rand = torch.rand(len(rows_to_sample), device=dev, dtype=torch.float32)
out_buf = torch.empty(len(rows_to_sample), device=dev, dtype=torch.int64)
grid = (len(rows_to_sample),)
_sample_kernel[grid](
probs_sel, # *f32
rand, # *f32
out_buf, # *i64
probs_sel.stride(0), # i32 stride between rows
vocab_size=VOCAB_SIZE,
BLOCK=BLOCK_SIZE,
N_BLKS=N_BLOCKS,
num_warps=4, # execution parameters
num_stages=2,
)
samples.index_copy_(0, sel_idx, out_buf)
# ----------------------- return to origin ------------------------
return samples if probs.is_cuda else samples.to(orig_device)scrolls · 202 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON