gpt-o3 / triton1d8355
gpt-o3_triton_1d8355 · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 188 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-1d8355?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:70cee49e1836c5334fe6ead0e3ca1bb91f9d8de2a576027ea94d86160f04ce07
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, # 8×32 = 256 threads per CTAstages = 2
num_stages=2,Kernel source
main.py188 lines
import math
import torch
import triton
import triton.language as tl
###############################################################################
# Constants – tuned for good compile-time and runtime on Hopper / B200
###############################################################################
VOCAB_SIZE: int = 128_256
# Making the tile wide (2 048) means we need to iterate only 63 times over a
# 128 256-token row – this keeps kernel size and compile time small, but still
# uses a very modest amount of registers / shared memory per block.
BLOCK_SIZE: int = 2_048
N_BLOCKS: int = (VOCAB_SIZE + BLOCK_SIZE - 1) // BLOCK_SIZE # 63
###############################################################################
# Triton kernel – one CTA (“program”) samples one full distribution row
###############################################################################
@triton.jit
def _sample_kernel(
probs_ptr, # *f32 [batch, VOCAB]
rand_ptr, # *f32 [batch]
out_ptr, # *i64 [batch]
stride_row, # ld stride ( = VOCAB_SIZE )
n_rows, # batch size
BLOCK_SIZE: tl.constexpr,
N_BLOCKS: tl.constexpr,
VOCAB_SIZE: tl.constexpr,
):
"""
Parameters
----------
probs_ptr : pointer to row-major tensor [batch, vocab] (float32)
rand_ptr : uniform random numbers in [0,1) (float32)
out_ptr : output indices (int64)
The kernel performs a streaming prefix-sum (CDF) over the probability
vector and returns the first index whose prefix exceeds the random number.
"""
pid = tl.program_id(axis=0)
if pid >= n_rows:
return
# ---------------------------------------------------------------------
# Per-row state
# ---------------------------------------------------------------------
row_ptr = probs_ptr + pid * stride_row
u = tl.load(rand_ptr + pid) # threshold in [0,1)
running = tl.zeros((), dtype=tl.float32) # prefix before current tile
found = tl.zeros((), dtype=tl.int1) # whether we already found
chosen = tl.zeros((), dtype=tl.int32) # resulting token id
# ---------------------------------------------------------------------
# Tile-wise scan over the 128 256-token row
# ---------------------------------------------------------------------
for b in tl.static_range(N_BLOCKS):
offset = b * BLOCK_SIZE
idx_vec = tl.arange(0, BLOCK_SIZE) + offset # [B]
lane_ok = idx_vec < VOCAB_SIZE # guard tail
# If we have not found the token yet, read this tile – otherwise skip
p = tl.load(row_ptr + idx_vec,
mask = lane_ok & (found == 0),
other = 0.0)
# Inclusive prefix sum inside the tile (only meaningful if !found)
cdf_local = running + tl.cumsum(p, axis=0)
# Lanes whose CDF crosses threshold
hit_mask = (u <= cdf_local) & lane_ok & (found == 0)
# Convert to candidate index, use big sentinel for “no hit”
big_val = tl.full([BLOCK_SIZE], VOCAB_SIZE, dtype=tl.int32)
cand_idx = tl.where(hit_mask, idx_vec.to(tl.int32), big_val)
# Reduction to obtain the left-most hit in this tile
cand_min = tl.min(cand_idx.to(tl.float32), axis=0).to(tl.int32)
# Update state
is_hit = cand_min < VOCAB_SIZE
chosen = tl.where(is_hit, cand_min, chosen)
found = tl.where(is_hit, 1, found)
running += tl.sum(p, axis=0) # advance
# Numerical safety – fall back to last token if nothing matched
chosen = tl.where(found == 0, VOCAB_SIZE - 1, chosen)
tl.store(out_ptr + pid, chosen.to(tl.int64))
###############################################################################
# Fast top-k filtering (in-place, GPU only)
###############################################################################
@torch.no_grad()
def _topk_filter_inplace(probs: torch.Tensor, top_k: torch.Tensor) -> None:
"""
In-place retains only the k largest entries of each row and re-normalises.
Rows with k ≤0 or k ≥ vocab_size are left unchanged.
"""
vocab = probs.size(1)
valid = (top_k > 0) & (top_k < vocab)
if not torch.any(valid):
return
rows = torch.nonzero(valid, as_tuple=False).squeeze(1)
sub_probs = probs[rows] # view into `probs`
sub_k = top_k[rows]
k_max = int(sub_k.max().item()) # <= vocab
vals, idx = torch.topk(sub_probs, k_max,
dim=1, largest=True, sorted=False)
keep_mask = torch.arange(k_max, device=probs.device).unsqueeze(0) \
< sub_k.unsqueeze(1)
vals = vals * keep_mask
sub_probs.zero_()
sub_probs.scatter_(1, idx, vals)
sub_probs.div_(sub_probs.sum(dim=1, keepdim=True).clamp_min(1e-20))
###############################################################################
# Public entry point
###############################################################################
@torch.no_grad()
def run(probs: torch.Tensor, top_k: torch.Tensor):
"""
Parameters
----------
probs : [batch, 128 256] float32 – probability distributions (softmaxed)
top_k : [batch] int32 – per-row k
Returns
-------
samples : [batch] int64 – sampled token indices
"""
# --------------- Basic sanity checks ----------------------------------
if probs.ndim != 2:
raise ValueError("`probs` has to 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_k.numel() != batch:
raise ValueError("len(top_k) must equal batch size")
if not torch.cuda.is_available():
raise RuntimeError("CUDA device required but not available")
# --------------- Move tensors to GPU (no-copy when already on GPU) ----
orig_device = probs.device
probs_gpu = probs.to('cuda', dtype=torch.float32, copy=False)
topk_gpu = top_k.to('cuda', dtype=torch.int32, copy=False)
# --------------- Optional top-k filtering -----------------------------
_topk_filter_inplace(probs_gpu, topk_gpu)
# --------------- Prepare RNG & output ---------------------------------
rand = torch.rand(batch, device='cuda', dtype=torch.float32)
out = torch.empty(batch, device='cuda', dtype=torch.int64)
# --------------- Launch Triton kernel ---------------------------------
_sample_kernel[(batch,)](
probs_gpu, rand, out,
probs_gpu.stride(0), batch,
BLOCK_SIZE=BLOCK_SIZE,
N_BLOCKS=N_BLOCKS,
VOCAB_SIZE=VOCAB_SIZE,
num_warps=8, # 8×32 = 256 threads per CTA
num_stages=2,
)
# --------------- Return on original device ----------------------------
return out.to(orig_device)
###############################################################################
# Lightweight smoke-test
###############################################################################
if __name__ == "__main__":
torch.manual_seed(0)
bs = 8
p = torch.randn(bs, VOCAB_SIZE, dtype=torch.float32)
p = torch.softmax(p, dim=-1)
k = torch.tensor([40, 0, VOCAB_SIZE, 10, 7, 50, 0, VOCAB_SIZE],
dtype=torch.int32)
print("Samples:", run(p, k))scrolls · 188 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON