Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton2a200e

gpt-o3_triton_2a200e · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-2a200e?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:e90a96015efd0c0b754e66ea6f326414d7820289e6fc4d849001f1da2781c9d5
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20

Kernel source

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

import torch
import triton
import triton.language as tl


# ─────────────────────────────────────────────────────────────────────────────
#                              TRITON  KERNEL
# ─────────────────────────────────────────────────────────────────────────────
@triton.jit
def _inverse_cdf_sample_kernel(
    probs_ptr,          # *f32 – flattened [rows, vocab]
    rand_ptr,           # *f32 – one random number per row
    out_ptr,            # *i64 – output indices
    stride_row: tl.constexpr,
    vocab_size: tl.constexpr,
    BLOCK_SIZE: tl.constexpr = 2048,          # 128 256 / 2 048 = 63 blocks
):
    """
    One Triton program handles one probability row (already top-k/top-p filtered
    and re-normalised).  It draws a single sample from that categorical
    distribution using an inverse-CDF search that is
      • vectorised inside each block (cumsum),
      • block-wise across the vocabulary (≤ 63 iterations).
    """

    pid       = tl.program_id(axis=0)         # row id
    row_ptr   = probs_ptr + pid * stride_row
    rand_val  = tl.load(rand_ptr + pid)       # uniform in (0, 1]

    running_cdf = tl.full((), 0.0, dtype=tl.float32)
    found_idx   = tl.full((), -1,  dtype=tl.int32)   # “not found” sentinel

    offs = tl.arange(0, BLOCK_SIZE)           # 0 … BLOCK_SIZE-1

    # Search block-by-block (compile-time unrolled – only 63 steps)
    for block_start in tl.static_range(0, vocab_size, BLOCK_SIZE):
        g_idx    = block_start + offs
        in_vocab = g_idx < vocab_size

        # Load a vector of probabilities
        vals = tl.load(row_ptr + g_idx, mask=in_vocab, other=0.0)

        # Inclusive scan within the vector
        cdf_block = tl.cumsum(vals, axis=0)

        # Does the sample fall into this block?
        hit_vec = (found_idx < 0) & in_vocab & (rand_val < running_cdf + cdf_block)
        hit_int = hit_vec.to(tl.int32)
        hit_any = tl.sum(hit_int, axis=0)                    # scalar ∈ {0, …}

        # Earliest index inside the block where the CDF exceeds rand_val
        hit_pos = tl.argmax(hit_int, axis=0)                 # 0 … BLOCK_SIZE-1

        found_idx = tl.where(
            (found_idx < 0) & (hit_any > 0),
            tl.full((), block_start, dtype=tl.int32) + hit_pos,
            found_idx,
        )

        running_cdf += tl.sum(vals, axis=0)

    # Numerical corner case (due to fp rounding): still not found → last token
    found_idx = tl.where(
        found_idx < 0,
        tl.full((), vocab_size - 1, dtype=tl.int32),
        found_idx,
    )

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


# ─────────────────────────────────────────────────────────────────────────────
#                         HOST-SIDE HELPER FUNCTIONS
# ─────────────────────────────────────────────────────────────────────────────
def _ensure_cuda(t: torch.Tensor, name: str) -> torch.Tensor:
    """
    Move a tensor to GPU if required; raise a clear error when CUDA is absent.
    """
    if t.is_cuda:
        return t
    if not torch.cuda.is_available():
        raise RuntimeError(
            f"CUDA is required for kernel execution, but tensor '{name}' is on CPU "
            "and no GPU is available."
        )
    return t.cuda(non_blocking=True)


@torch.no_grad()
def _top_k_top_p_filter(
    probs: torch.Tensor,
    top_k: torch.Tensor,
    top_p: torch.Tensor,
) -> torch.Tensor:
    """
    Row-wise top-k / nucleus (top-p) filtering.
    Implemented with plain Torch ops; runs on GPU when inputs are CUDA tensors.
    """
    B, V = probs.shape
    out = torch.zeros_like(probs)

    for r in range(B):
        row = probs[r]

        # --------------------------- top-k ---------------------------
        k = int(top_k[r].item())
        if 0 < k < V:
            vals, idx = torch.topk(row, k, largest=True, sorted=False)
            masked = torch.zeros_like(row)
            masked.scatter_(0, idx, vals)
            row = masked / masked.sum()

        # --------------------------- top-p ---------------------------
        p = float(top_p[r].item())
        if 0.0 < p < 1.0:
            vals, idx = torch.sort(row, descending=True)
            cdf = torch.cumsum(vals, 0)

            remove = cdf > p
            if V > 1:
                remove[1:] = remove[:-1].clone()
                remove[0] = False

            keep_idx = idx[~remove]
            masked = torch.zeros_like(row)
            masked[keep_idx] = row[keep_idx]
            row = masked / masked.sum()

        out[r] = row

    return out


# ─────────────────────────────────────────────────────────────────────────────
#                                   PUBLIC API
# ─────────────────────────────────────────────────────────────────────────────
@torch.no_grad()
def run(
    probs: torch.Tensor,
    top_k: torch.Tensor,
    top_p: torch.Tensor,
    **kwargs: Dict[str, Any],
) -> torch.Tensor:
    """
    Optimised implementation of `top_k_top_p_sampling_from_probs_v128256`.
    Preserves reference behaviour while off-loading the expensive sampling
    step to a Triton kernel geared towards B200 GPUs.
    """

    # --------------------------- argument checks ----------------------------
    if probs.ndim != 2:
        raise ValueError("`probs` must be 2-D with shape [batch_size, vocab_size].")
    batch, vocab = probs.shape
    if vocab != 128_256:
        raise ValueError(f"vocab_size must be 128 256, got {vocab}.")

    # --------------------------- device handling ----------------------------
    orig_device = probs.device
    probs = _ensure_cuda(probs.to(torch.float32), "probs")
    top_k = _ensure_cuda(top_k.to(torch.int32),  "top_k")
    top_p = _ensure_cuda(top_p.to(torch.float32), "top_p")

    # --------------------------- filtering ----------------------------------
    filtered = _top_k_top_p_filter(probs, top_k, top_p)

    # ---------------------- greedy vs stochastic rows -----------------------
    greedy_mask = top_p <= 0.0
    samples = torch.empty(batch, dtype=torch.int64, device=probs.device)

    # Greedy rows (argmax)
    if greedy_mask.any():
        samples[greedy_mask] = torch.argmax(filtered[greedy_mask], dim=1)

    # Stochastic rows (inverse-CDF sample via Triton)
    stoch_mask = ~greedy_mask
    if stoch_mask.any():
        probs_stoch = filtered[stoch_mask].contiguous()
        n_rows = probs_stoch.shape[0]

        # Uniform random numbers in (0, 1]; clamp to avoid exact 0
        rnd = torch.rand(
            n_rows,
            dtype=torch.float32,
            device=probs.device,
        ).clamp_min_(1e-30)

        out_buf = torch.empty(n_rows, dtype=torch.int64, device=probs.device)

        _inverse_cdf_sample_kernel[(n_rows,)](
            probs_stoch,
            rnd,
            out_buf,
            stride_row=vocab,
            vocab_size=vocab,
        )

        samples[stoch_mask] = out_buf

    # ------------------------------ done ------------------------------------
    return samples if orig_device.type == "cuda" else samples.cpu()
scrolls · 203 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON