gpt-o3 / triton2a200e
gpt-o3_triton_2a200e · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 203 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-2a200e?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:e90a96015efd0c0b754e66ea6f326414d7820289e6fc4d849001f1da2781c9d5
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20
Kernel source
main.py203 lines
import math
from typing import Any, Dict
import torch
import triton
import triton.language as tl
# ─────────────────────────────────────────────────────────────────────────────
# TRITON KERNEL
# ─────────────────────────────────────────────────────────────────────────────
@triton.jit
def _inverse_cdf_sample_kernel(
probs_ptr, # *f32 – flattened [rows, vocab]
rand_ptr, # *f32 – one random number per row
out_ptr, # *i64 – output indices
stride_row: tl.constexpr,
vocab_size: tl.constexpr,
BLOCK_SIZE: tl.constexpr = 2048, # 128 256 / 2 048 = 63 blocks
):
"""
One Triton program handles one probability row (already top-k/top-p filtered
and re-normalised). It draws a single sample from that categorical
distribution using an inverse-CDF search that is
• vectorised inside each block (cumsum),
• block-wise across the vocabulary (≤ 63 iterations).
"""
pid = tl.program_id(axis=0) # row id
row_ptr = probs_ptr + pid * stride_row
rand_val = tl.load(rand_ptr + pid) # uniform in (0, 1]
running_cdf = tl.full((), 0.0, dtype=tl.float32)
found_idx = tl.full((), -1, dtype=tl.int32) # “not found” sentinel
offs = tl.arange(0, BLOCK_SIZE) # 0 … BLOCK_SIZE-1
# Search block-by-block (compile-time unrolled – only 63 steps)
for block_start in tl.static_range(0, vocab_size, BLOCK_SIZE):
g_idx = block_start + offs
in_vocab = g_idx < vocab_size
# Load a vector of probabilities
vals = tl.load(row_ptr + g_idx, mask=in_vocab, other=0.0)
# Inclusive scan within the vector
cdf_block = tl.cumsum(vals, axis=0)
# Does the sample fall into this block?
hit_vec = (found_idx < 0) & in_vocab & (rand_val < running_cdf + cdf_block)
hit_int = hit_vec.to(tl.int32)
hit_any = tl.sum(hit_int, axis=0) # scalar ∈ {0, …}
# Earliest index inside the block where the CDF exceeds rand_val
hit_pos = tl.argmax(hit_int, axis=0) # 0 … BLOCK_SIZE-1
found_idx = tl.where(
(found_idx < 0) & (hit_any > 0),
tl.full((), block_start, dtype=tl.int32) + hit_pos,
found_idx,
)
running_cdf += tl.sum(vals, axis=0)
# Numerical corner case (due to fp rounding): still not found → last token
found_idx = tl.where(
found_idx < 0,
tl.full((), vocab_size - 1, dtype=tl.int32),
found_idx,
)
tl.store(out_ptr + pid, found_idx.to(tl.int64))
# ─────────────────────────────────────────────────────────────────────────────
# HOST-SIDE HELPER FUNCTIONS
# ─────────────────────────────────────────────────────────────────────────────
def _ensure_cuda(t: torch.Tensor, name: str) -> torch.Tensor:
"""
Move a tensor to GPU if required; raise a clear error when CUDA is absent.
"""
if t.is_cuda:
return t
if not torch.cuda.is_available():
raise RuntimeError(
f"CUDA is required for kernel execution, but tensor '{name}' is on CPU "
"and no GPU is available."
)
return t.cuda(non_blocking=True)
@torch.no_grad()
def _top_k_top_p_filter(
probs: torch.Tensor,
top_k: torch.Tensor,
top_p: torch.Tensor,
) -> torch.Tensor:
"""
Row-wise top-k / nucleus (top-p) filtering.
Implemented with plain Torch ops; runs on GPU when inputs are CUDA tensors.
"""
B, V = probs.shape
out = torch.zeros_like(probs)
for r in range(B):
row = probs[r]
# --------------------------- top-k ---------------------------
k = int(top_k[r].item())
if 0 < k < V:
vals, idx = torch.topk(row, k, largest=True, sorted=False)
masked = torch.zeros_like(row)
masked.scatter_(0, idx, vals)
row = masked / masked.sum()
# --------------------------- top-p ---------------------------
p = float(top_p[r].item())
if 0.0 < p < 1.0:
vals, idx = torch.sort(row, descending=True)
cdf = torch.cumsum(vals, 0)
remove = cdf > p
if V > 1:
remove[1:] = remove[:-1].clone()
remove[0] = False
keep_idx = idx[~remove]
masked = torch.zeros_like(row)
masked[keep_idx] = row[keep_idx]
row = masked / masked.sum()
out[r] = row
return out
# ─────────────────────────────────────────────────────────────────────────────
# PUBLIC API
# ─────────────────────────────────────────────────────────────────────────────
@torch.no_grad()
def run(
probs: torch.Tensor,
top_k: torch.Tensor,
top_p: torch.Tensor,
**kwargs: Dict[str, Any],
) -> torch.Tensor:
"""
Optimised implementation of `top_k_top_p_sampling_from_probs_v128256`.
Preserves reference behaviour while off-loading the expensive sampling
step to a Triton kernel geared towards B200 GPUs.
"""
# --------------------------- argument checks ----------------------------
if probs.ndim != 2:
raise ValueError("`probs` must be 2-D with shape [batch_size, vocab_size].")
batch, vocab = probs.shape
if vocab != 128_256:
raise ValueError(f"vocab_size must be 128 256, got {vocab}.")
# --------------------------- device handling ----------------------------
orig_device = probs.device
probs = _ensure_cuda(probs.to(torch.float32), "probs")
top_k = _ensure_cuda(top_k.to(torch.int32), "top_k")
top_p = _ensure_cuda(top_p.to(torch.float32), "top_p")
# --------------------------- filtering ----------------------------------
filtered = _top_k_top_p_filter(probs, top_k, top_p)
# ---------------------- greedy vs stochastic rows -----------------------
greedy_mask = top_p <= 0.0
samples = torch.empty(batch, dtype=torch.int64, device=probs.device)
# Greedy rows (argmax)
if greedy_mask.any():
samples[greedy_mask] = torch.argmax(filtered[greedy_mask], dim=1)
# Stochastic rows (inverse-CDF sample via Triton)
stoch_mask = ~greedy_mask
if stoch_mask.any():
probs_stoch = filtered[stoch_mask].contiguous()
n_rows = probs_stoch.shape[0]
# Uniform random numbers in (0, 1]; clamp to avoid exact 0
rnd = torch.rand(
n_rows,
dtype=torch.float32,
device=probs.device,
).clamp_min_(1e-30)
out_buf = torch.empty(n_rows, dtype=torch.int64, device=probs.device)
_inverse_cdf_sample_kernel[(n_rows,)](
probs_stoch,
rnd,
out_buf,
stride_row=vocab,
vocab_size=vocab,
)
samples[stoch_mask] = out_buf
# ------------------------------ done ------------------------------------
return samples if orig_device.type == "cuda" else samples.cpu()scrolls · 203 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON