Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton75f9e9

gpt-o3_triton_75f9e9 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-75f9e9?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:07455d1f83ccc9b3cbc19f5e96f4e61eb152bb27d3b7c03ad6f2b4bea486f793
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, # 256 threads / 8 warps

Kernel source

main.py179 lines
import math
from typing import Any

import torch
import triton
import triton.language as tl


###############################################################################
# Kernel : draw ONE sample from ONE categorical distribution                  #
###############################################################################
@triton.jit
def _cdf_sample_kernel(
    probs_ptr,                # *f32  – [batch , vocab]
    rand_ptr,                 # *f32  – [batch]  (0 ≤ r < 1)
    out_ptr,                  # *i64  – [batch]
    VOCAB_SIZE: tl.constexpr, # 129 280
    BLOCK_SIZE: tl.constexpr  # 128 / 256 / …
):
    """
    Each Triton program processes exactly ONE row.

    We iterate over the row BLOCK_SIZE tokens at a time while
    maintaining a running prefix-sum. The first entry whose
    cumulative probability strictly exceeds the random threshold
    is selected.
    """
    pid       = tl.program_id(axis=0)           # row index
    row_start = probs_ptr + pid * VOCAB_SIZE
    thresh    = tl.load(rand_ptr + pid)         # 0 ≤ thresh < 1

    lane_off  = tl.arange(0, BLOCK_SIZE)        # 0 … BLOCK_SIZE-1
    running   = tl.zeros((), tl.float32)        # prefix sum of previous blocks
    chosen    = tl.full((), -1, tl.int32)       # –1  ⇒ not found yet
    base_idx  = tl.zeros((), tl.int32)          # first token handled by block

    # ---------------------------------------------------------------- main scan
    while (base_idx < VOCAB_SIZE) & (chosen < 0):
        idx  = base_idx + lane_off
        mask = idx < VOCAB_SIZE                 # guard against OOB accesses

        # 1. load probabilities of the current chunk
        p    = tl.load(row_start + idx, mask=mask, other=0.0)

        # 2. cumulative sum *inside* this block + running prefix
        local_cdf = tl.cumsum(p, axis=0) + running

        # NOTE: we need a STRICT comparison here.  If `thresh` is 0
        #       we must pick the first *positive* probability entry,
        #       not a zero-probability token that happens to precede it.
        crossed   = mask & (local_cdf > thresh)

        # 3. first index in this block that crosses the threshold
        INF       = tl.full((BLOCK_SIZE,), BLOCK_SIZE, idx.dtype)
        cand_off  = tl.where(crossed, lane_off, INF)
        min_off   = tl.min(cand_off, axis=0)

        found     = min_off < BLOCK_SIZE
        first_idx = base_idx + min_off
        chosen    = tl.where(found & (chosen < 0), first_idx, chosen)

        # 4. advance to next block
        running += tl.sum(p, axis=0)
        base_idx += BLOCK_SIZE

    # Numerical fallback – should never trigger
    chosen = tl.where(chosen < 0, VOCAB_SIZE - 1, chosen)

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


###############################################################################
# Fast batched top-k filtering (host side, PyTorch)                           #
###############################################################################
def _vectorised_topk_filter(
    probs: torch.Tensor,
    top_k: torch.Tensor,
    vocab_size: int,
) -> torch.Tensor:
    """
    For every row i with 0 < k_i < vocab_size:
        • keep exactly the k_i largest probabilities
        • set all remaining entries to 0

    The rows are NOT renormalised here – the caller does that afterwards.
    """
    need = (top_k > 0) & (top_k < vocab_size)
    if not torch.any(need):
        return probs

    filtered = probs.clone()
    rows     = torch.nonzero(need, as_tuple=False).squeeze(1)
    ks       = top_k[rows]

    k_max = int(ks.max().item())
    # sorted=True guarantees that the first k_i entries
    # correspond to the k_i largest tokens of each row
    vals, idxs = torch.topk(
        filtered[rows],
        k_max,
        dim=1,
        largest=True,
        sorted=True,
    )

    col_idx   = torch.arange(k_max, device=probs.device)
    keep_mask = col_idx.unsqueeze(0) < ks.unsqueeze(1)
    vals      = vals * keep_mask

    filtered[rows].zero_()
    filtered[rows].scatter_(1, idxs, vals)

    return filtered


###############################################################################
# Public API                                                                  #
###############################################################################
def run(
    probs: torch.Tensor,
    top_k:  torch.Tensor,
    *args: Any,
    **kwargs: Any,
) -> torch.Tensor:
    """
    Parameters
    ----------
    probs : [batch , 129280] – soft-maxed probabilities (float16/bfloat16/float32)
    top_k : [batch] int32    – per-row top-k
                                 (0 or ≥ vocab_size  ⇒ keep row unchanged)

    Returns
    -------
    samples : [batch] int64 – one sampled token id per input row
    """
    if not torch.cuda.is_available():
        raise RuntimeError("A CUDA-capable device is required to run this kernel.")

    # ---------------------------------------------------------------- device juggling
    src_device = probs.device
    cuda_dev   = torch.device("cuda")

    probs_fp32 = probs.to(device=cuda_dev, dtype=torch.float32, copy=False)
    topk_i32   = top_k.to(device=cuda_dev, dtype=torch.int32,   copy=False)

    batch, vocab = probs_fp32.shape
    if vocab != 129_280:
        raise ValueError(f"vocab_size must be exactly 129 280, got {vocab}")

    # ---------------------------------------------------------------- top-k filter
    probs_filt = _vectorised_topk_filter(probs_fp32, topk_i32, vocab)

    # final normalisation (guards against FP drift)
    row_sums   = probs_filt.sum(dim=1, keepdim=True)
    # If a row became all-zero (should not happen), fall back to the original row
    probs_norm = torch.where(
        row_sums > 0,
        probs_filt / row_sums.clamp(min=1e-7),
        probs_fp32,
    )

    # ---------------------------------------------------------------- RNG   (uniform in [0, 1))
    rnd = torch.rand(batch, device=cuda_dev, dtype=torch.float32)

    # ---------------------------------------------------------------- launch kernel
    out = torch.empty(batch, device=cuda_dev, dtype=torch.int64)

    BLOCK = 256  # empirically a good fit for B200
    _cdf_sample_kernel[(batch,)](
        probs_norm,
        rnd,
        out,
        VOCAB_SIZE=vocab,
        BLOCK_SIZE=BLOCK,
        num_warps=8,  # 256 threads / 8 warps
    )

    # ---------------------------------------------------------------- restore device
    return out.to(src_device, non_blocking=True)
scrolls · 179 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON