Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton7d588b

gpt-o3_triton_7d588b · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-7d588b?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:02c4d08c502ec1dc81e3850702ab6a67f61294e495bab707d8b186c8cc69add4
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, # Tuned empirically – good starting point for B200

Kernel source

main.py256 lines
import math
from typing import Any, Tuple, Union

import torch
import triton
import triton.language as tl


################################################################################
#                              TRITON KERNELS                                  #
################################################################################
@triton.jit
def _sample_topk_kernel(  # noqa: N802
    probs_ptr,         # float32  [num_rows, K_MAX]
    idx_ptr,           # int32    [num_rows, K_MAX]
    rand_ptr,          # float32  [num_rows]
    out_ptr,           # int64    [num_rows]
    stride_probs,      # int32    leading dimension for probs
    stride_idx,        # int32    leading dimension for indices
    K: tl.constexpr,   # compile–time: maximum K across this launch
):
    """
    One program = one row (sequence).
    Each program draws exactly one sample from the provided probability rows.
    The probability rows are assumed to
      1. already be limited to top-k tokens (zeros elsewhere)
      2. already be normalised                       (sum == 1)

    Kernel launches with grid = (num_rows,)
    """
    row_id = tl.program_id(axis=0)

    # Base pointers for this row
    probs_row_ptr = probs_ptr + row_id * stride_probs
    idx_row_ptr   = idx_ptr   + row_id * stride_idx

    # Remaining probability mass before we “hit” the sample
    remaining = tl.load(rand_ptr + row_id)          # uniform in [0, 1)
    sample_idx_in_row = tl.full((), -1, tl.int32)   # sentinel (-1  -> not chosen yet)

    # Sequential (compile-time) scan over the *fixed* number of columns K
    for j in tl.static_range(K):
        p_val = tl.load(probs_row_ptr + j, eviction_policy='evict_last')

        # If we haven’t picked a token yet, check whether the current position
        # crosses the remaining cumulative mass.
        not_found   = sample_idx_in_row < 0
        take_token  = not_found & (remaining <= p_val)

        sample_idx_in_row = tl.where(
            take_token,
            tl.full((), j, tl.int32),
            sample_idx_in_row,
        )

        # If not picked yet, subtract this probability and keep scanning
        remaining = tl.where(
            not_found,
            remaining - p_val,
            remaining,
        )

    # ------------------------------------------------------------------
    # Resolve the *vocabulary* index that corresponds to sample_idx_in_row
    # Fallback to the *last* candidate if, due to numerical issues,
    # no token was selected (should be extremely rare).
    # ------------------------------------------------------------------
    safe_idx   = tl.where(sample_idx_in_row >= 0,
                          sample_idx_in_row,
                          tl.full((), K - 1, tl.int32))
    token_id   = tl.load(idx_row_ptr + safe_idx).to(tl.int64)

    tl.store(out_ptr + row_id, token_id)


################################################################################
#                           PYTHON / HOST  SIDE                                #
################################################################################
def _prepare_topk_tensors(
    probs: torch.Tensor,
    top_k: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
    """
    Utility that extracts + normalises the per-row top-k probability mass.

    Returns
    -------
    (rows_kept,
     topk_probs,      # float32  [R, K_MAX]   (row-major, contiguous)
     topk_indices,    # int32    [R, K_MAX]   (row-major, contiguous)
     K_MAX)           # python   int
    """
    vocab_size = probs.shape[1]
    mask_valid = (top_k > 0) & (top_k < vocab_size)

    if not torch.any(mask_valid):
        # No row needs top-k filtering; callers can skip Triton completely.
        return (
            torch.empty(0, dtype=torch.int64, device=probs.device),
            torch.empty(0, dtype=torch.float32, device=probs.device),
            torch.empty(0, dtype=torch.int32,  device=probs.device),
            0,
        )

    rows_kept: torch.Tensor = torch.nonzero(mask_valid, as_tuple=False).squeeze(1)
    # Maximum k across *selected* rows
    K_MAX: int = int(top_k[rows_kept].max().item())

    # Use torch.topk (efficient & GPU-accelerated) to fetch the candidates
    topk_vals, topk_indices = torch.topk(
        probs[rows_kept], K_MAX, dim=1, largest=True, sorted=True,
    )

    # Normalise probabilities *inside* each row up to its own k_i
    topk_probs = torch.zeros_like(topk_vals)
    for row_local, row_global in enumerate(rows_kept):
        k_i = int(top_k[row_global].item())
        if k_i == 0:      # should not happen (mask_valid) but stay safe
            continue
        vals_slice = topk_vals[row_local, :k_i]
        row_sum    = vals_slice.sum()
        topk_probs[row_local, :k_i] = vals_slice / row_sum

    # Remaining positions stay zero, which is what the kernel expects.
    return (
        rows_kept,
        topk_probs.contiguous(),
        topk_indices.to(torch.int32).contiguous(),
        K_MAX,
    )


def _device_guard(t: torch.Tensor) -> torch.device:
    """Utility – also ensures CUDA availability."""
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is required but `torch.cuda.is_available()` is False.")
    return t.device if t.is_cuda else torch.device("cuda")


################################################################################
#                                ENTRY POINT                                   #
################################################################################
def run(
    *args: Any,
    **kwargs: Any,
) -> torch.Tensor:
    """
    Public entry point that mirrors the reference API:

        samples = run(probs, top_k)

    The function takes care of
      • transferring to GPU (if needed),
      • launching the Triton kernel,
      • handling edge-cases (`k <= 0` or `k >= vocab_size`),
      • moving results back to the original device.
    """
    # ------------------------------------------------------------------
    # Parse arguments
    # ------------------------------------------------------------------
    if len(args) >= 1:
        probs = args[0]
        top_k = args[1] if len(args) >= 2 else kwargs.get("top_k", None)
    else:
        probs = kwargs.get("probs", None)
        top_k = kwargs.get("top_k", None)

    if probs is None or top_k is None:
        raise ValueError("`run` expects two tensors: `probs` and `top_k`.")

    # Make sure dtypes / shapes are as expected
    probs = probs.to(torch.float32)
    top_k = top_k.to(torch.int32)

    batch_size, vocab_size = probs.shape
    if vocab_size != 151_936:
        raise ValueError(
            f"Expected vocab_size = 151,936 but got {vocab_size}"
        )

    # ------------------------------------------------------------------
    # Device management – transfer to GPU if needed
    # ------------------------------------------------------------------
    orig_device = probs.device
    cuda_device = _device_guard(probs)

    if not probs.is_cuda:
        probs_cuda = probs.to(cuda_device, non_blocking=True)
    else:
        probs_cuda = probs

    if not top_k.is_cuda:
        top_k_cuda = top_k.to(cuda_device, non_blocking=True)
    else:
        top_k_cuda = top_k

    # Output tensor (on CUDA for now, moved back later if needed)
    samples_cuda = torch.empty(batch_size, dtype=torch.int64, device=cuda_device)

    # ------------------------------------------------------------------
    # Rows that do *not* need top-k filtering
    # (k <= 0  OR  k >= vocab_size)   ->   vanilla multinomial
    # ------------------------------------------------------------------
    full_rows_mask = (top_k_cuda <= 0) | (top_k_cuda >= vocab_size)
    if torch.any(full_rows_mask):
        rows_full = torch.nonzero(full_rows_mask, as_tuple=False).squeeze(1)
        sub     = torch.multinomial(
            probs_cuda[rows_full], 1, replacement=True,
        ).squeeze(1)
        samples_cuda[rows_full] = sub.to(torch.int64)

    # ------------------------------------------------------------------
    # Rows that *do* need top-k filtering  → Triton
    # ------------------------------------------------------------------
    (
        rows_kept,
        topk_probs,
        topk_indices,
        K_MAX,
    ) = _prepare_topk_tensors(probs_cuda, top_k_cuda)

    if rows_kept.numel() > 0:
        # Random numbers for each selected row
        rand_uniform = torch.rand(
            rows_kept.shape[0], dtype=torch.float32, device=cuda_device
        )

        # Output buffer for the Triton kernel
        out_subset = torch.empty(
            rows_kept.shape[0], dtype=torch.int64, device=cuda_device
        )

        grid = (rows_kept.shape[0],)
        stride_probs = topk_probs.stride(0)
        stride_idx   = topk_indices.stride(0)

        _sample_topk_kernel[grid](
            topk_probs,
            topk_indices,
            rand_uniform,
            out_subset,
            stride_probs,
            stride_idx,
            K=K_MAX,
            num_warps=4,      # Tuned empirically – good starting point for B200
        )

        # Scatter back to their original rows
        samples_cuda[rows_kept] = out_subset

    # ------------------------------------------------------------------
    # Move back to the original device (if needed) & return
    # ------------------------------------------------------------------
    if not probs.is_cuda:
        return samples_cuda.to(orig_device, non_blocking=True)
    return samples_cuda
scrolls · 256 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON