Skip to content
KernelIndex
Search⌘K

gpt-o3_triton_861794

gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 196 lines, Apache-2.0, pinned at da91508.

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-861794?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:839570e32823689b47dd611afab3a8592ac897feb1f95c95a768017b124f279e
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 = 4num_warps=4, # 4 × 32 = 128 threads – matches BLOCK

Kernel source

main.py196 lines
import math
from typing import Any, Dict

import torch
import triton
import triton.language as tl


# ---------------------------------------------------------------------------- #
#  Kernel                                                                      #
# ---------------------------------------------------------------------------- #
@triton.jit
def _top_p_kernel(
    probs_ptr,        # *fp32  – [batch, vocab]
    top_p_ptr,        # *fp32  – [batch]
    rand_ptr,         # *fp32  – [batch]  (one U(0,1) per row)
    out_ptr,          # *int64 – [batch]
    vocab_size: tl.constexpr,
    BLOCK:     tl.constexpr,        # 128
):
    """
    One Triton program = exactly one sequence (= one matrix row).

    Behaviour (per-row, decided from `top_p`):
      •  p ≤ 0           -> greedy arg-max
      •  0 < p < 1       -> nucleus sampling (row was pre-filtered)
      •  p ≥ 1           -> vanilla multinomial sampling
    """
    pid       = tl.program_id(axis=0)                 # row id
    row_ptr   = probs_ptr + pid * vocab_size          # start of this row

    # Per-row params ----------------------------------------------------------------
    p_val     = tl.load(top_p_ptr + pid)              # nucleus threshold
    r_val     = tl.load(rand_ptr  + pid)              # uniform random in [0,1)
    greedy    = p_val <= 0.0                          # bool – take arg-max?

    # Running state ------------------------------------------------------------------
    best_val  = tl.full((), -math.inf, tl.float32)    # running maximum value
    best_idx  = tl.full((),  0,        tl.int32)      # arg-max index

    sample_idx = tl.full((), -1,       tl.int32)      # -1 → not chosen yet
    cum_prob   = tl.zeros((), tl.float32)             # running CDF (for sampling)
    finished   = tl.zeros((), tl.int1)                # exits early when sampled

    # Static column offsets (0 … BLOCK-1)
    offs = tl.arange(0, BLOCK)

    # Loop over the vocabulary -------------------------------------------------------
    start = tl.zeros((), tl.int32)
    while (start < vocab_size) & (finished == 0):
        idxs = start + offs                           # absolute indices
        mask = idxs < vocab_size                      # boundary mask
        vals = tl.load(row_ptr + idxs,
                       mask=mask,
                       other=0.0)                     # [BLOCK] – fp32

        # -------- greedy path: track block maximum ----------------------------------
        block_max  = tl.max(vals, axis=0)
        same_max   = vals == block_max
        first_max  = tl.where(same_max, offs, BLOCK)
        block_arg  = tl.min(first_max, axis=0) + start

        is_better  = block_max > best_val
        best_val   = tl.where(is_better, block_max, best_val)
        best_idx   = tl.where(is_better, block_arg, best_idx)

        # -------- sampling path: inverse CDF scan -----------------------------------
        prefix   = tl.cumsum(vals)                    # inclusive scan over block
        hit      = (sample_idx < 0) & (cum_prob + prefix >= r_val)
        hit_off  = tl.where(hit, offs, BLOCK)
        firstHit = tl.min(hit_off, axis=0)
        got_it   = firstHit < BLOCK

        sample_idx = tl.where(got_it & (sample_idx < 0),
                              start + firstHit,
                              sample_idx)

        finished   = finished | ((~greedy) & got_it)

        # advance --------------------------------------------------------------------
        cum_prob += tl.sum(vals, axis=0)
        start    += BLOCK

    # -------------------------------------------------------------------------------
    chosen = tl.where(greedy, best_idx, sample_idx)
    chosen = tl.where(chosen < 0,  best_idx, chosen)   # numerical-safety fallback

    tl.store(out_ptr + pid, chosen.to(tl.int64))


# ---------------------------------------------------------------------------- #
#  GPU-side nucleus (top-p) filter                                             #
# ---------------------------------------------------------------------------- #
def _gpu_top_p_filter(probs: torch.Tensor,
                      top_p: torch.Tensor) -> torch.Tensor:
    """
    In-place style (returns a clone) nucleus filtering on GPU.

    Only rows with 0 < p < 1 are processed, others are unchanged.
    Every processed row is re-normalised to sum to 1.
    """
    out  = probs.clone()                              # keeps dtype / device
    mask = (top_p > 0.0) & (top_p < 1.0)
    if not mask.any():
        return out

    rows        = out[mask]                           # [M, V]
    p_thr       = top_p[mask].unsqueeze(1)            # [M, 1]

    # Full sort – still the simplest + fastest for very large vocab on GPU
    vals, idx   = torch.sort(rows, dim=-1, descending=True)
    cdf         = vals.cumsum(dim=-1)

    remove      = cdf > p_thr
    shift       = torch.zeros_like(remove, dtype=torch.bool)
    shift[:, 1:] = remove[:, :-1]                     # keep first token ≥ threshold
    keep        = ~shift

    kept_vals   = torch.where(keep, vals, torch.zeros_like(vals))
    filtered    = torch.zeros_like(rows)
    filtered.scatter_(1, idx, kept_vals)

    # Re-normalise (row-wise)
    row_sum     = filtered.sum(dim=1, keepdim=True)
    filtered   /= row_sum

    out[mask]   = filtered
    return out


# ---------------------------------------------------------------------------- #
#  Public entry point                                                           #
# ---------------------------------------------------------------------------- #
def run(
    probs: torch.Tensor,
    top_p: torch.Tensor,
    *kernel_args: Any,
    **kernel_kwargs: Dict[str, Any],
) -> torch.Tensor:
    """
    Fast top-p / multinomial sampler (B200-optimised).

    Steps:
      1. Device housekeeping.
      2. Optional nucleus filtering (GPU).
      3. Generate one U(0,1) number per sequence.
      4. Launch Triton kernel (1 program / sequence, 128 threads, 4 warps).
      5. Return samples on original device.
    """
    # ---- device handling ----------------------------------------------------------
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device required to run the Triton kernel")

    orig_device = probs.device                       # remember caller’s device

    probs_gpu = probs.to(torch.float32)
    top_p_gpu = top_p.to(torch.float32)

    if not probs_gpu.is_cuda:
        probs_gpu = probs_gpu.cuda()
    if not top_p_gpu.is_cuda:
        top_p_gpu = top_p_gpu.cuda()

    batch, vocab = probs_gpu.shape
    if vocab != 151_936:
        raise ValueError(f"Expected vocab_size == 151 936, got {vocab}")

    # ---- pre-processing: nucleus filter ------------------------------------------
    probs_ready = _gpu_top_p_filter(probs_gpu, top_p_gpu)

    # ---- random numbers -----------------------------------------------------------
    rand_row = torch.rand(batch, dtype=torch.float32, device=probs_ready.device)

    # ---- output buffer ------------------------------------------------------------
    out = torch.empty(batch, dtype=torch.int64, device=probs_ready.device)

    # ---- kernel launch ------------------------------------------------------------
    BLOCK = 128
    grid  = (batch,)

    _top_p_kernel[grid](
        probs_ready,
        top_p_gpu,
        rand_row,
        out,
        vocab_size=vocab,
        BLOCK=BLOCK,
        num_warps=4,          # 4 × 32 = 128 threads – matches BLOCK
        *kernel_args,
        **kernel_kwargs,
    )

    # ---- bring back to caller’s device -------------------------------------------
    if not probs.is_cuda:
        out = out.to(orig_device)
    return out
scrolls · 196 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON