Skip to content
KernelIndex
Search⌘K

gpt-o3 / tritonb8d9a9

gpt-o3_triton_b8d9a9 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-b8d9a9?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:4e2c83a08b372ccbbe6b2e981838b2e91af813a0858c01b2173d37695f2d3f54
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, # execution parameters
stages = 2num_stages=2,

Kernel source

main.py202 lines
import math
from typing import Sequence, Mapping, Any, List

import torch
import triton
import triton.language as tl

# ----------------------------------------------------------------------
# Problem-level constants (compile-time)
# ----------------------------------------------------------------------
VOCAB_SIZE: int = 128_256            # fixed for Llama-3.1
BLOCK_SIZE: int = 1_024              # number of tokens processed per loop
N_BLOCKS:   int = (VOCAB_SIZE + BLOCK_SIZE - 1) // BLOCK_SIZE   # = 126


# ----------------------------------------------------------------------
# Triton kernel – inverse-CDF sampling of ONE distribution row
# ----------------------------------------------------------------------
@triton.jit
def _sample_kernel(
    probs_ptr,             # *f32  [batch, VOCAB_SIZE]
    rand_ptr,              # *f32  [batch] – uniform[0,1)
    out_ptr,               # *i64  [batch]
    stride_row,            # i32   leading stride between rows
    vocab_size: tl.constexpr,
    BLOCK:      tl.constexpr,
    N_BLKS:     tl.constexpr,
):
    """
    One kernel instance (= program) handles ONE row of probabilities.
    We scan the cumulative distribution until it crosses a random
    threshold `r` and return the corresponding index.
    """
    pid = tl.program_id(axis=0)                      # row id
    row_ptr = probs_ptr + pid * stride_row           # pointer to first element in row
    r = tl.load(rand_ptr + pid)                      # threshold in (0, 1)

    # running cumulative probability *before* current block
    cum_sum = tl.zeros((), dtype=tl.float32)

    # best index found so far (init to sentinel > vocab_size-1)
    sentinel = vocab_size
    best_ix = tl.full((), sentinel, dtype=tl.int32)

    # ------------------------------------------------------------------
    # iterate over blocks of size `BLOCK`
    # ------------------------------------------------------------------
    for blk in tl.static_range(N_BLKS):
        start = blk * BLOCK
        offs = tl.arange(0, BLOCK)
        idxs = start + offs                          # absolute token indices
        valid = idxs < vocab_size                    # mask for short last block

        # load probabilities
        probs = tl.load(row_ptr + idxs, mask=valid, other=0.0)   # [BLOCK]

        # inclusive prefix inside the block + previous cum_sum
        cdf_blk = tl.cumsum(probs, axis=0) + cum_sum

        # first positions where CDF ≥ r
        crosses = (cdf_blk >= r) & valid
        cand = tl.where(crosses, idxs, sentinel).to(tl.int32)

        # first crossing inside the block
        first_in_blk = tl.min(cand, axis=0)

        # keep leftmost crossing overall
        best_ix = tl.where(first_in_blk < best_ix, first_in_blk, best_ix)

        # advance cumulative sum
        cum_sum += tl.sum(probs, axis=0)

    # safeguard – if nothing selected (due to tiny numerical error) pick last vocab
    best_ix = tl.where(best_ix == sentinel, vocab_size - 1, best_ix)

    # write result
    tl.store(out_ptr + pid, best_ix.to(tl.int64))


# ----------------------------------------------------------------------
# Helper – build per-row nucleus (top-p) distribution 100 % on GPU
# ----------------------------------------------------------------------
def _build_nucleus_distribution(row: torch.Tensor, p_thresh: float) -> torch.Tensor:
    """
    Keep the minimal prefix whose cumulative probability reaches `p_thresh`
    (== nucleus / top-p).  Returns a re-normalised probability vector.
    All operations happen on `row.device` (GPU for performance).
    """
    if p_thresh >= 1.0:
        return row

    # sort in descending order
    vals, idx = torch.sort(row, descending=True)
    cdf = torch.cumsum(vals, dim=0)

    # mask: remove everything AFTER (not incl.) the first entry that makes CDF > p
    to_remove = cdf > p_thresh
    to_remove[1:] = to_remove[:-1].clone()
    to_remove[0] = False

    keep_idx = idx[~to_remove]

    filtered = torch.zeros_like(row)
    filtered[keep_idx] = row[keep_idx]

    total = filtered.sum()
    if total > 0:
        filtered /= total
    return filtered


# ----------------------------------------------------------------------
# Public API
# ----------------------------------------------------------------------
def run(
    probs: torch.Tensor,
    top_p: torch.Tensor,
    *args: Sequence[Any],
    **kwargs: Mapping[str, Any],
) -> torch.Tensor:
    """
    Parameters
    ----------
    probs : [batch, 128256] float32 – soft-maxed probabilities
    top_p : [batch]         float32 – per-row nucleus threshold

    Returns
    -------
    samples : [batch] int64 – sampled token indices
    """
    # --------------------------- validation ---------------------------
    if probs.ndim != 2:
        raise ValueError("`probs` must 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_p.shape != (batch,):
        raise ValueError("`top_p` must have shape [batch]")

    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device required but not available")

    # ---------------------- device management ------------------------
    orig_device = probs.device
    dev = probs.device if probs.is_cuda else torch.device("cuda")

    probs_gpu = probs.to(dev, dtype=torch.float32, non_blocking=True)
    top_p_gpu = top_p.to(dev,  dtype=torch.float32, non_blocking=True)

    # ----------------------- pre-processing --------------------------
    samples = torch.empty(batch, dtype=torch.int64, device=dev)

    # indices of rows that NEED sampling through the kernel
    rows_to_sample: List[int] = []
    nucleus_rows = []

    for i in range(batch):
        p_thr = float(top_p_gpu[i].item())
        row = probs_gpu[i]

        # p ≤ 0   → greedy argmax
        if p_thr <= 0.0:
            samples[i] = torch.argmax(row).to(torch.int64)
            continue

        filt_row = _build_nucleus_distribution(row, p_thr)

        # extremely rare – if nucleus empty fall back to argmax
        if filt_row.sum() == 0:
            samples[i] = torch.argmax(row).to(torch.int64)
            continue

        rows_to_sample.append(i)
        nucleus_rows.append(filt_row)

    # ---------------------- call Triton kernel -----------------------
    if rows_to_sample:
        sel_idx = torch.tensor(rows_to_sample, device=dev, dtype=torch.int64)

        # stack selected rows into a single 2-D tensor for the kernel
        probs_sel = torch.stack(nucleus_rows, dim=0).contiguous()
        rand = torch.rand(len(rows_to_sample), device=dev, dtype=torch.float32)
        out_buf = torch.empty(len(rows_to_sample), device=dev, dtype=torch.int64)

        grid = (len(rows_to_sample),)

        _sample_kernel[grid](
            probs_sel,                       # *f32
            rand,                            # *f32
            out_buf,                         # *i64
            probs_sel.stride(0),             # i32 stride between rows
            vocab_size=VOCAB_SIZE,
            BLOCK=BLOCK_SIZE,
            N_BLKS=N_BLOCKS,
            num_warps=4,                     # execution parameters
            num_stages=2,
        )

        samples.index_copy_(0, sel_idx, out_buf)

    # ----------------------- return to origin ------------------------
    return samples if probs.is_cuda else samples.to(orig_device)
scrolls · 202 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON