Skip to content
KernelIndex
Search⌘K

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 = 8num_warps=8, # 8×32 = 256 threads per CTA
stages = 2num_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