Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton579f5d

gpt-o3_triton_579f5d · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-579f5d?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:9f352056176db8bb1f80b6a5181347f1182f1f03bed077d27f031c4192720937
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 = 1num_warps=1, # 1 warp is enough for the serial scan

Kernel source

main.py185 lines
import math
from typing import Any

import torch
import triton
import triton.language as tl

# ------------------------------------------------------------------------------
# Constants
# ------------------------------------------------------------------------------
VOCAB_SIZE: int = 129_280          # DeepSeek-V3 vocabulary size


# ------------------------------------------------------------------------------
# Triton kernel
# ------------------------------------------------------------------------------
@triton.jit
def _top_p_sampling_kernel(
    probs_ptr,                 # *f32  – [batch, vocab]
    top_p_ptr,                 # *f32  – [batch]
    rand_ptr,                  # *f32  – [batch]  (uniform in [0, 1])
    out_ptr,                   # *i64  – [batch]
    stride_probs,              # int   – leading dim of probs
    stride_top_p,              # int
    stride_rand,               # int
    stride_out,                # int
    vocab_size: tl.constexpr,  # compile-time constant (=129 280)
):
    """
    One Triton program (thread-block) handles exactly one sequence.
    The vocabulary is scanned linearly; despite being simple, this is
    already much faster than launching a separate kernel per token
    thanks to Triton’s fused control-flow.
    """
    pid = tl.program_id(0)                                    # sequence id

    # ------------------------------------------------------------------
    #  Load per-row scalars
    # ------------------------------------------------------------------
    row_ptr  = probs_ptr + pid * stride_probs                 # *f32 to row[0]
    p_thresh = tl.load(top_p_ptr + pid * stride_top_p)        # float32
    rand_val = tl.load(rand_ptr  + pid * stride_rand)         # float32

    is_greedy = p_thresh <= 0.0                               # bool tensor

    # ------------------------------------------------------------------
    #  Running state initialisation (Triton scalars)
    # ------------------------------------------------------------------
    best_val   = tl.full((), -1.0,  tl.float32)    # best prob for greedy path
    best_idx   = tl.full((),  0,    tl.int32)
    running    = tl.zeros((), dtype=tl.float32)    # running CDF for sampling
    chosen_idx = tl.zeros((), dtype=tl.int32)      # sampled index
    found_flag = tl.zeros((), dtype=tl.int32)      # 0 → not yet, 1 → found
    idx        = tl.zeros((), dtype=tl.int32)      # vocabulary pointer

    # ------------------------------------------------------------------
    #  Linear scan over the vocabulary
    # ------------------------------------------------------------------
    while idx < vocab_size:
        prob = tl.load(row_ptr + idx)

        # ---- greedy argmax -----------------------------------------------------
        is_better = prob > best_val
        best_val  = tl.where(is_better, prob, best_val)
        best_idx  = tl.where(is_better, idx,  best_idx)

        # ---- multinomial prefix-sum sampling -----------------------------------
        next_running = running + prob
        hit          = (found_flag == 0) & (next_running >= rand_val)
        chosen_idx   = tl.where(hit, idx, chosen_idx)
        found_flag   = tl.where(hit, 1,   found_flag)
        running      = next_running

        idx += 1

    # ------------------------------------------------------------------
    #  Write result
    # ------------------------------------------------------------------
    final_idx = tl.where(is_greedy, best_idx, chosen_idx)
    tl.store(out_ptr + pid * stride_out, final_idx.to(tl.int64))


# ------------------------------------------------------------------------------
#  Helper: vectorised top-p filtering (GPU, PyTorch)
# ------------------------------------------------------------------------------
def _filter_probs_top_p(probs: torch.Tensor, top_p: torch.Tensor) -> torch.Tensor:
    """
    Applies nucleus (top-p) filtering row-wise.
    The logic exactly matches the reference implementation.
    """
    vals_sorted, idx_sorted = torch.sort(probs, dim=1, descending=True)
    cdf = vals_sorted.cumsum(dim=1)

    # mask out everything AFTER the first value that makes CDF > p
    to_remove = cdf > top_p.unsqueeze(1)
    shifted = torch.zeros_like(to_remove)
    shifted[:, 1:] = to_remove[:, :-1]
    to_remove = shifted
    to_remove[:, 0] = False
    keep = ~to_remove

    filtered = torch.zeros_like(probs)
    filtered.scatter_(1, idx_sorted, vals_sorted * keep.float())

    row_sums = filtered.sum(dim=1, keepdim=True)
    row_sums = torch.where(row_sums == 0.0, torch.ones_like(row_sums), row_sums)
    return filtered / row_sums


# ------------------------------------------------------------------------------
#  Utility
# ------------------------------------------------------------------------------
def _to_gpu(t: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
    """
    Ensure `t` resides on a CUDA device and has the requested dtype.
    """
    if t.device.type == "cpu":
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is required but not available.")
        return t.cuda().to(dtype=dtype, copy=False)
    return t.to(dtype=dtype, copy=False)


# ------------------------------------------------------------------------------
#  Public entry point
# ------------------------------------------------------------------------------
def run(probs: torch.Tensor, top_p: torch.Tensor, *args: Any, **kwargs: Any) -> torch.Tensor:
    """
    Parameters
    ----------
    probs : (batch, VOCAB_SIZE) float32
        Probability distributions (already softmax-normalised).
    top_p : (batch,) float32
        Cumulative probability threshold per sequence.

    Returns
    -------
    (batch,) int64 tensor – sampled token indices (on same device as `probs`)
    """
    # ------------- sanity checks ------------------------------------------------
    if probs.ndim != 2:
        raise ValueError("`probs` must be 2-D [batch, vocab]")
    if top_p.ndim != 1:
        raise ValueError("`top_p` must be 1-D [batch]")
    batch_size, vocab_size = probs.shape
    if vocab_size != VOCAB_SIZE:
        raise ValueError(f"vocab_size must be {VOCAB_SIZE}, got {vocab_size}")
    if batch_size != top_p.shape[0]:
        raise ValueError("Batch size mismatch between `probs` and `top_p`")

    original_device = probs.device

    # ------------- move tensors to GPU -----------------------------------------
    probs_gpu = _to_gpu(probs, torch.float32)
    top_p_gpu = _to_gpu(top_p, torch.float32)

    # ------------- apply top-p filtering ---------------------------------------
    filtered_probs = probs_gpu.clone()
    mid_mask = (top_p_gpu > 0.0) & (top_p_gpu < 1.0)
    if mid_mask.any():
        filtered_probs_mid = _filter_probs_top_p(filtered_probs[mid_mask],
                                                 top_p_gpu[mid_mask])
        filtered_probs[mid_mask] = filtered_probs_mid

    # ------------- prepare RNG + output ----------------------------------------
    rand_vec = torch.rand(batch_size, dtype=torch.float32, device=filtered_probs.device)
    out_gpu  = torch.empty(batch_size, dtype=torch.int64, device=filtered_probs.device)

    # ------------- launch Triton kernel ----------------------------------------
    grid = (batch_size,)
    _top_p_sampling_kernel[grid](
        filtered_probs,
        top_p_gpu,
        rand_vec,
        out_gpu,
        filtered_probs.stride(0),
        top_p_gpu.stride(0),
        rand_vec.stride(0),
        out_gpu.stride(0),
        vocab_size=vocab_size,
        num_warps=1,                 # 1 warp is enough for the serial scan
    )

    # ------------- return result on original device ----------------------------
    return out_gpu.to(original_device, non_blocking=True)
scrolls · 185 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON