gpt-o3 / triton579f5d
gpt-o3_triton_579f5d · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 185 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-579f5d?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:9f352056176db8bb1f80b6a5181347f1182f1f03bed077d27f031c4192720937
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 = 1
num_warps=1, # 1 warp is enough for the serial scanKernel source
main.py185 lines
import math
from typing import Any
import torch
import triton
import triton.language as tl
# ------------------------------------------------------------------------------
# Constants
# ------------------------------------------------------------------------------
VOCAB_SIZE: int = 129_280 # DeepSeek-V3 vocabulary size
# ------------------------------------------------------------------------------
# Triton kernel
# ------------------------------------------------------------------------------
@triton.jit
def _top_p_sampling_kernel(
probs_ptr, # *f32 – [batch, vocab]
top_p_ptr, # *f32 – [batch]
rand_ptr, # *f32 – [batch] (uniform in [0, 1])
out_ptr, # *i64 – [batch]
stride_probs, # int – leading dim of probs
stride_top_p, # int
stride_rand, # int
stride_out, # int
vocab_size: tl.constexpr, # compile-time constant (=129 280)
):
"""
One Triton program (thread-block) handles exactly one sequence.
The vocabulary is scanned linearly; despite being simple, this is
already much faster than launching a separate kernel per token
thanks to Triton’s fused control-flow.
"""
pid = tl.program_id(0) # sequence id
# ------------------------------------------------------------------
# Load per-row scalars
# ------------------------------------------------------------------
row_ptr = probs_ptr + pid * stride_probs # *f32 to row[0]
p_thresh = tl.load(top_p_ptr + pid * stride_top_p) # float32
rand_val = tl.load(rand_ptr + pid * stride_rand) # float32
is_greedy = p_thresh <= 0.0 # bool tensor
# ------------------------------------------------------------------
# Running state initialisation (Triton scalars)
# ------------------------------------------------------------------
best_val = tl.full((), -1.0, tl.float32) # best prob for greedy path
best_idx = tl.full((), 0, tl.int32)
running = tl.zeros((), dtype=tl.float32) # running CDF for sampling
chosen_idx = tl.zeros((), dtype=tl.int32) # sampled index
found_flag = tl.zeros((), dtype=tl.int32) # 0 → not yet, 1 → found
idx = tl.zeros((), dtype=tl.int32) # vocabulary pointer
# ------------------------------------------------------------------
# Linear scan over the vocabulary
# ------------------------------------------------------------------
while idx < vocab_size:
prob = tl.load(row_ptr + idx)
# ---- greedy argmax -----------------------------------------------------
is_better = prob > best_val
best_val = tl.where(is_better, prob, best_val)
best_idx = tl.where(is_better, idx, best_idx)
# ---- multinomial prefix-sum sampling -----------------------------------
next_running = running + prob
hit = (found_flag == 0) & (next_running >= rand_val)
chosen_idx = tl.where(hit, idx, chosen_idx)
found_flag = tl.where(hit, 1, found_flag)
running = next_running
idx += 1
# ------------------------------------------------------------------
# Write result
# ------------------------------------------------------------------
final_idx = tl.where(is_greedy, best_idx, chosen_idx)
tl.store(out_ptr + pid * stride_out, final_idx.to(tl.int64))
# ------------------------------------------------------------------------------
# Helper: vectorised top-p filtering (GPU, PyTorch)
# ------------------------------------------------------------------------------
def _filter_probs_top_p(probs: torch.Tensor, top_p: torch.Tensor) -> torch.Tensor:
"""
Applies nucleus (top-p) filtering row-wise.
The logic exactly matches the reference implementation.
"""
vals_sorted, idx_sorted = torch.sort(probs, dim=1, descending=True)
cdf = vals_sorted.cumsum(dim=1)
# mask out everything AFTER the first value that makes CDF > p
to_remove = cdf > top_p.unsqueeze(1)
shifted = torch.zeros_like(to_remove)
shifted[:, 1:] = to_remove[:, :-1]
to_remove = shifted
to_remove[:, 0] = False
keep = ~to_remove
filtered = torch.zeros_like(probs)
filtered.scatter_(1, idx_sorted, vals_sorted * keep.float())
row_sums = filtered.sum(dim=1, keepdim=True)
row_sums = torch.where(row_sums == 0.0, torch.ones_like(row_sums), row_sums)
return filtered / row_sums
# ------------------------------------------------------------------------------
# Utility
# ------------------------------------------------------------------------------
def _to_gpu(t: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
"""
Ensure `t` resides on a CUDA device and has the requested dtype.
"""
if t.device.type == "cpu":
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required but not available.")
return t.cuda().to(dtype=dtype, copy=False)
return t.to(dtype=dtype, copy=False)
# ------------------------------------------------------------------------------
# Public entry point
# ------------------------------------------------------------------------------
def run(probs: torch.Tensor, top_p: torch.Tensor, *args: Any, **kwargs: Any) -> torch.Tensor:
"""
Parameters
----------
probs : (batch, VOCAB_SIZE) float32
Probability distributions (already softmax-normalised).
top_p : (batch,) float32
Cumulative probability threshold per sequence.
Returns
-------
(batch,) int64 tensor – sampled token indices (on same device as `probs`)
"""
# ------------- sanity checks ------------------------------------------------
if probs.ndim != 2:
raise ValueError("`probs` must be 2-D [batch, vocab]")
if top_p.ndim != 1:
raise ValueError("`top_p` must be 1-D [batch]")
batch_size, vocab_size = probs.shape
if vocab_size != VOCAB_SIZE:
raise ValueError(f"vocab_size must be {VOCAB_SIZE}, got {vocab_size}")
if batch_size != top_p.shape[0]:
raise ValueError("Batch size mismatch between `probs` and `top_p`")
original_device = probs.device
# ------------- move tensors to GPU -----------------------------------------
probs_gpu = _to_gpu(probs, torch.float32)
top_p_gpu = _to_gpu(top_p, torch.float32)
# ------------- apply top-p filtering ---------------------------------------
filtered_probs = probs_gpu.clone()
mid_mask = (top_p_gpu > 0.0) & (top_p_gpu < 1.0)
if mid_mask.any():
filtered_probs_mid = _filter_probs_top_p(filtered_probs[mid_mask],
top_p_gpu[mid_mask])
filtered_probs[mid_mask] = filtered_probs_mid
# ------------- prepare RNG + output ----------------------------------------
rand_vec = torch.rand(batch_size, dtype=torch.float32, device=filtered_probs.device)
out_gpu = torch.empty(batch_size, dtype=torch.int64, device=filtered_probs.device)
# ------------- launch Triton kernel ----------------------------------------
grid = (batch_size,)
_top_p_sampling_kernel[grid](
filtered_probs,
top_p_gpu,
rand_vec,
out_gpu,
filtered_probs.stride(0),
top_p_gpu.stride(0),
rand_vec.stride(0),
out_gpu.stride(0),
vocab_size=vocab_size,
num_warps=1, # 1 warp is enough for the serial scan
)
# ------------- return result on original device ----------------------------
return out_gpu.to(original_device, non_blocking=True)scrolls · 185 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON