Skip to content
KernelIndex
Search⌘K

gpt-o3 / tritonfaf01e

gpt-o3_triton_faf01e · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Kernel source

main.py212 lines
import math
import torch
import triton
import triton.language as tl


###############################################################################
# Constants
###############################################################################
VOCAB_SIZE: int = 129_280       # fixed DeepSeek-V3 vocabulary size
BLOCK_SIZE: int = 1024          # tokens processed per kernel iteration


###############################################################################
# Triton kernel
###############################################################################
@triton.jit
def _sample_kernel(
    probs_ptr,           # float32  [batch, vocab]
    rand_ptr,            # float32  [batch]
    sample_ptr,          # int64    [batch]
    stride_row,          # stride between consecutive rows (vocab_size)
    vocab_size: tl.constexpr,      # 129_280 (compile–time constant)
    BLOCK: tl.constexpr            # 1 024   (compile–time constant)
):
    """
    Each Triton program samples ONE sequence (one distribution / row).
    The probabilities in `probs_ptr` MUST already be:
      • filtered (top-k / top-p) and
      • re-normalised so that they sum to 1.
    """
    pid = tl.program_id(axis=0)                    # sequence id
    row_offset = pid * stride_row                  # start of this row
    row_ptr = probs_ptr + row_offset               # pointer to first prob
    rng = tl.load(rand_ptr + pid)                  # U(0,1) for this row

    # Running state ----------------------------------------------------------
    cumsum_before = tl.zeros((), dtype=tl.float32) # cumulative mass processed
    found        = tl.zeros((), dtype=tl.int32)    # 0 -> still searching
    chosen_idx   = tl.full((), -1, dtype=tl.int32) # result placeholder

    # Utility: thread-local contiguous indices 0 … BLOCK-1
    idx_in_block = tl.arange(0, BLOCK)

    # Iterate over the vocabulary ------------------------------------------------
    for offs in range(0, vocab_size, BLOCK):
        global_idx = offs + idx_in_block
        block_mask = global_idx < vocab_size

        # load current chunk of probabilities
        probs = tl.load(row_ptr + global_idx, mask=block_mask, other=0.0)

        # sum of this BLOCK across all threads
        block_sum = tl.sum(probs, axis=0)

        # If we haven’t found the token yet and the running cumulative mass
        # crosses our random number *inside* this block, we must identify it.
        search_block = (found == 0) & (cumsum_before + block_sum > rng)

        # Prefix sums of probs within the block (only matters when searching)
        prefix = tl.cumsum(probs, axis=0)

        # Candidate positions: where prefix exceeds the residual mass
        residual    = rng - cumsum_before
        in_prefix   = prefix > residual
        candidate   = tl.where(search_block & in_prefix, idx_in_block,
                               BLOCK)                    # sentinel

        # First index in this block that satisfies the predicate
        first_in_blk = tl.min(candidate, axis=0)

        # If a valid index was found, record the global position
        is_valid     = first_in_blk < BLOCK
        chosen_idx   = tl.where(is_valid & (found == 0),
                                offs + first_in_blk, chosen_idx)
        found        = tl.where(is_valid, 1, found)

        # advance cumulative mass (only while still searching)
        cumsum_before += tl.where(found == 0, block_sum,
                                  tl.zeros((), dtype=tl.float32))

    # Fallback (numerical safety) – never happens in theory
    chosen_idx = tl.where(found == 0, vocab_size - 1, chosen_idx)

    # Write result as int64
    tl.store(sample_ptr + pid, chosen_idx.to(tl.int64))


###############################################################################
# Python wrapper
###############################################################################
def _ensure_cuda(t: torch.Tensor) -> torch.Tensor:
    """Move tensor to CUDA if it is on CPU.  Raises if CUDA is unavailable."""
    if t.device.type == "cuda":
        return t
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device is required for this kernel.")
    return t.cuda()


@torch.no_grad()
def run(probs: torch.Tensor,
        top_k: torch.Tensor,
        top_p: torch.Tensor) -> torch.Tensor:
    """
    Top-k / Top-p sampling implemented with a mix of high-level PyTorch
    primitives (for filtering) and a custom Triton kernel (for the final
    draw).  The output exactly matches the reference implementation.
    """
    # --------------------------------------------------------------------- #
    # 1. Device management & dtype normalisation
    # --------------------------------------------------------------------- #
    orig_device = probs.device
    probs  = probs.to(torch.float32)
    top_k  = top_k.to(torch.int32)
    top_p  = top_p.to(torch.float32)

    probs_gpu = _ensure_cuda(probs)
    k_gpu     = _ensure_cuda(top_k)
    p_gpu     = _ensure_cuda(top_p)

    batch, vocab = probs_gpu.shape
    if vocab != VOCAB_SIZE:
        raise ValueError(f"Expected vocab_size={VOCAB_SIZE}, got {vocab}")

    samples = torch.empty(batch, dtype=torch.int64, device=probs_gpu.device)

    # --------------------------------------------------------------------- #
    # 2. Per-row filtering (top-k, top-p)             — executed in PyTorch
    # --------------------------------------------------------------------- #
    rows_for_kernel = []
    for i in range(batch):
        row = probs_gpu[i]
        k   = int(k_gpu[i].item())
        p   = float(p_gpu[i].item())

        # -------- top-k --------------------------------------------------
        if 0 < k < VOCAB_SIZE:
            vals, idx = torch.topk(row, k, largest=True, sorted=False)
            mask = torch.zeros_like(row, dtype=torch.bool)
            mask[idx] = True
            row = row * mask.float()
            row /= row.sum()

        # deterministic maximum if nucleus threshold <= 0
        if p <= 0.0:
            samples[i] = torch.argmax(row).to(torch.int64)
            probs_gpu[i] = row           # store (normalised) for completeness
            continue

        # -------- top-p --------------------------------------------------
        if p < 1.0:
            vals, sidx = torch.sort(row, descending=True)
            cdf = torch.cumsum(vals, 0)
            remove = cdf > p
            if VOCAB_SIZE > 1:
                remove[1:] = remove[:-1].clone()
                remove[0]  = False
            keep = sidx[~remove]
            mask = torch.zeros_like(row, dtype=torch.bool)
            mask[keep] = True
            row = row * mask.float()
            row /= row.sum()

        # row now sums to 1  → store back
        probs_gpu[i] = row
        rows_for_kernel.append(i)

    # --------------------------------------------------------------------- #
    # 3. Sampling rows with stochastic nucleus          — Triton kernel
    # --------------------------------------------------------------------- #
    if rows_for_kernel:
        idx_tensor  = torch.tensor(rows_for_kernel,
                                   dtype=torch.int64,
                                   device=probs_gpu.device)
        sub_probs   = probs_gpu.index_select(0, idx_tensor).contiguous()
        rand_vec    = torch.rand(len(rows_for_kernel),
                                 dtype=torch.float32,
                                 device=probs_gpu.device)
        out_buf     = torch.empty(len(rows_for_kernel),
                                  dtype=torch.int64,
                                  device=probs_gpu.device)

        grid = (sub_probs.shape[0],)
        _sample_kernel[grid](
            sub_probs, rand_vec, out_buf,
            sub_probs.stride(0),          # stride between rows
            vocab_size=VOCAB_SIZE,
            BLOCK=BLOCK_SIZE,
        )
        samples.index_copy_(0, idx_tensor, out_buf)

    # --------------------------------------------------------------------- #
    # 4. Move result back to original device (if needed)
    # --------------------------------------------------------------------- #
    if orig_device.type == "cuda":
        return samples.to(orig_device)
    return samples.cpu()


# When this file is executed directly (not imported) -------------------------
if __name__ == "__main__":
    # Quick sanity check
    torch.manual_seed(0)
    bs = 4
    logits = torch.randn(bs, VOCAB_SIZE, dtype=torch.float32)
    probs  = torch.softmax(logits, dim=-1)
    top_k  = torch.tensor([50, 0,  10, VOCAB_SIZE], dtype=torch.int32)
    top_p  = torch.tensor([0.95, -1.0, 0.9, 0.0],   dtype=torch.float32)

    samples_out = run(probs, top_k, top_p)
    print("Sampled indices:", samples_out)
scrolls · 212 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON