gpt-o3 / triton7d588b
gpt-o3_triton_7d588b · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 256 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-7d588b?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:02c4d08c502ec1dc81e3850702ab6a67f61294e495bab707d8b186c8cc69add4
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, # Tuned empirically – good starting point for B200Kernel source
main.py256 lines
import math
from typing import Any, Tuple, Union
import torch
import triton
import triton.language as tl
################################################################################
# TRITON KERNELS #
################################################################################
@triton.jit
def _sample_topk_kernel( # noqa: N802
probs_ptr, # float32 [num_rows, K_MAX]
idx_ptr, # int32 [num_rows, K_MAX]
rand_ptr, # float32 [num_rows]
out_ptr, # int64 [num_rows]
stride_probs, # int32 leading dimension for probs
stride_idx, # int32 leading dimension for indices
K: tl.constexpr, # compile–time: maximum K across this launch
):
"""
One program = one row (sequence).
Each program draws exactly one sample from the provided probability rows.
The probability rows are assumed to
1. already be limited to top-k tokens (zeros elsewhere)
2. already be normalised (sum == 1)
Kernel launches with grid = (num_rows,)
"""
row_id = tl.program_id(axis=0)
# Base pointers for this row
probs_row_ptr = probs_ptr + row_id * stride_probs
idx_row_ptr = idx_ptr + row_id * stride_idx
# Remaining probability mass before we “hit” the sample
remaining = tl.load(rand_ptr + row_id) # uniform in [0, 1)
sample_idx_in_row = tl.full((), -1, tl.int32) # sentinel (-1 -> not chosen yet)
# Sequential (compile-time) scan over the *fixed* number of columns K
for j in tl.static_range(K):
p_val = tl.load(probs_row_ptr + j, eviction_policy='evict_last')
# If we haven’t picked a token yet, check whether the current position
# crosses the remaining cumulative mass.
not_found = sample_idx_in_row < 0
take_token = not_found & (remaining <= p_val)
sample_idx_in_row = tl.where(
take_token,
tl.full((), j, tl.int32),
sample_idx_in_row,
)
# If not picked yet, subtract this probability and keep scanning
remaining = tl.where(
not_found,
remaining - p_val,
remaining,
)
# ------------------------------------------------------------------
# Resolve the *vocabulary* index that corresponds to sample_idx_in_row
# Fallback to the *last* candidate if, due to numerical issues,
# no token was selected (should be extremely rare).
# ------------------------------------------------------------------
safe_idx = tl.where(sample_idx_in_row >= 0,
sample_idx_in_row,
tl.full((), K - 1, tl.int32))
token_id = tl.load(idx_row_ptr + safe_idx).to(tl.int64)
tl.store(out_ptr + row_id, token_id)
################################################################################
# PYTHON / HOST SIDE #
################################################################################
def _prepare_topk_tensors(
probs: torch.Tensor,
top_k: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
"""
Utility that extracts + normalises the per-row top-k probability mass.
Returns
-------
(rows_kept,
topk_probs, # float32 [R, K_MAX] (row-major, contiguous)
topk_indices, # int32 [R, K_MAX] (row-major, contiguous)
K_MAX) # python int
"""
vocab_size = probs.shape[1]
mask_valid = (top_k > 0) & (top_k < vocab_size)
if not torch.any(mask_valid):
# No row needs top-k filtering; callers can skip Triton completely.
return (
torch.empty(0, dtype=torch.int64, device=probs.device),
torch.empty(0, dtype=torch.float32, device=probs.device),
torch.empty(0, dtype=torch.int32, device=probs.device),
0,
)
rows_kept: torch.Tensor = torch.nonzero(mask_valid, as_tuple=False).squeeze(1)
# Maximum k across *selected* rows
K_MAX: int = int(top_k[rows_kept].max().item())
# Use torch.topk (efficient & GPU-accelerated) to fetch the candidates
topk_vals, topk_indices = torch.topk(
probs[rows_kept], K_MAX, dim=1, largest=True, sorted=True,
)
# Normalise probabilities *inside* each row up to its own k_i
topk_probs = torch.zeros_like(topk_vals)
for row_local, row_global in enumerate(rows_kept):
k_i = int(top_k[row_global].item())
if k_i == 0: # should not happen (mask_valid) but stay safe
continue
vals_slice = topk_vals[row_local, :k_i]
row_sum = vals_slice.sum()
topk_probs[row_local, :k_i] = vals_slice / row_sum
# Remaining positions stay zero, which is what the kernel expects.
return (
rows_kept,
topk_probs.contiguous(),
topk_indices.to(torch.int32).contiguous(),
K_MAX,
)
def _device_guard(t: torch.Tensor) -> torch.device:
"""Utility – also ensures CUDA availability."""
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required but `torch.cuda.is_available()` is False.")
return t.device if t.is_cuda else torch.device("cuda")
################################################################################
# ENTRY POINT #
################################################################################
def run(
*args: Any,
**kwargs: Any,
) -> torch.Tensor:
"""
Public entry point that mirrors the reference API:
samples = run(probs, top_k)
The function takes care of
• transferring to GPU (if needed),
• launching the Triton kernel,
• handling edge-cases (`k <= 0` or `k >= vocab_size`),
• moving results back to the original device.
"""
# ------------------------------------------------------------------
# Parse arguments
# ------------------------------------------------------------------
if len(args) >= 1:
probs = args[0]
top_k = args[1] if len(args) >= 2 else kwargs.get("top_k", None)
else:
probs = kwargs.get("probs", None)
top_k = kwargs.get("top_k", None)
if probs is None or top_k is None:
raise ValueError("`run` expects two tensors: `probs` and `top_k`.")
# Make sure dtypes / shapes are as expected
probs = probs.to(torch.float32)
top_k = top_k.to(torch.int32)
batch_size, vocab_size = probs.shape
if vocab_size != 151_936:
raise ValueError(
f"Expected vocab_size = 151,936 but got {vocab_size}"
)
# ------------------------------------------------------------------
# Device management – transfer to GPU if needed
# ------------------------------------------------------------------
orig_device = probs.device
cuda_device = _device_guard(probs)
if not probs.is_cuda:
probs_cuda = probs.to(cuda_device, non_blocking=True)
else:
probs_cuda = probs
if not top_k.is_cuda:
top_k_cuda = top_k.to(cuda_device, non_blocking=True)
else:
top_k_cuda = top_k
# Output tensor (on CUDA for now, moved back later if needed)
samples_cuda = torch.empty(batch_size, dtype=torch.int64, device=cuda_device)
# ------------------------------------------------------------------
# Rows that do *not* need top-k filtering
# (k <= 0 OR k >= vocab_size) -> vanilla multinomial
# ------------------------------------------------------------------
full_rows_mask = (top_k_cuda <= 0) | (top_k_cuda >= vocab_size)
if torch.any(full_rows_mask):
rows_full = torch.nonzero(full_rows_mask, as_tuple=False).squeeze(1)
sub = torch.multinomial(
probs_cuda[rows_full], 1, replacement=True,
).squeeze(1)
samples_cuda[rows_full] = sub.to(torch.int64)
# ------------------------------------------------------------------
# Rows that *do* need top-k filtering → Triton
# ------------------------------------------------------------------
(
rows_kept,
topk_probs,
topk_indices,
K_MAX,
) = _prepare_topk_tensors(probs_cuda, top_k_cuda)
if rows_kept.numel() > 0:
# Random numbers for each selected row
rand_uniform = torch.rand(
rows_kept.shape[0], dtype=torch.float32, device=cuda_device
)
# Output buffer for the Triton kernel
out_subset = torch.empty(
rows_kept.shape[0], dtype=torch.int64, device=cuda_device
)
grid = (rows_kept.shape[0],)
stride_probs = topk_probs.stride(0)
stride_idx = topk_indices.stride(0)
_sample_topk_kernel[grid](
topk_probs,
topk_indices,
rand_uniform,
out_subset,
stride_probs,
stride_idx,
K=K_MAX,
num_warps=4, # Tuned empirically – good starting point for B200
)
# Scatter back to their original rows
samples_cuda[rows_kept] = out_subset
# ------------------------------------------------------------------
# Move back to the original device (if needed) & return
# ------------------------------------------------------------------
if not probs.is_cuda:
return samples_cuda.to(orig_device, non_blocking=True)
return samples_cudascrolls · 256 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON