Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / triton4a1fba

gpt-5-2025-08-07_triton_4a1fba · gpt-5-2025-08-07 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-4a1fba?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:6c942a0ebe33a55c2ebca69541d8c7b4a613104025e1fff7e14855a2e2f69148
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

num-warps = 8num_warps = 8
stages = 2num_stages = 2

Kernel source

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


VOCAB_SIZE = 151936


@triton.jit
def _top_p_sample_kernel(
    probs_ptr,            # *const float32 [B, V]
    work_ptr,             # *const float32 [B, V] - filtered probs (or same as probs)
    top_p_ptr,            # *const float32 [B]
    out_ptr,              # *mut int64   [B]
    B,                    # int32: batch size
    stride_p0, stride_p1, # strides for probs in elements
    stride_w0, stride_w1, # strides for work in elements
    seed,                 # uint32 seed for RNG
    V: tl.constexpr,      # vocab size (compile-time)
    BLOCK: tl.constexpr   # block size along vocab
):
    pid = tl.program_id(axis=0)
    if pid >= B:
        return

    # Load per-row top_p threshold
    p = tl.load(top_p_ptr + pid)

    # Compute base pointers for this row
    row_p_ptr = probs_ptr + pid * stride_p0
    row_w_ptr = work_ptr + pid * stride_w0

    # Constants
    eps = 1e-9
    big_pos = 1e20

    # Branch: p <= 0 => argmax over original probabilities
    if p <= 0.0:
        best_val = tl.full((), -float("inf"), dtype=tl.float32)  # scalar
        best_idx = tl.full((), 0, dtype=tl.int32)                # scalar
        for start in tl.static_range(0, V, BLOCK):
            offs = start + tl.arange(0, BLOCK)
            mask = offs < V
            vals = tl.load(row_p_ptr + offs * stride_p1, mask=mask, other=-float("inf"))
            # block max and first index achieving it
            block_max = tl.max(vals, axis=0)  # scalar
            eq = vals == block_max
            idxs = offs
            cand = tl.where(eq, idxs, tl.full([BLOCK], V, dtype=tl.int32))
            block_idx = tl.min(cand, axis=0)  # scalar
            # update global best (prefer smaller index on ties)
            better = block_max > best_val
            equal = block_max == best_val
            take_idx = tl.where(equal, block_idx < best_idx, better)
            best_idx = tl.where(take_idx, block_idx, best_idx)
            best_val = tl.where(take_idx, block_max, best_val)
        tl.store(out_ptr + pid, best_idx.to(tl.int64))
        return

    # Branch: p > 0 => sample from 'work' distribution using A-Res
    # If p >= 1, 'work' should be equal to 'probs'. If 0 < p < 1, 'work' is top-p filtered.
    best_r = tl.full((), big_pos, dtype=tl.float32)  # scalar
    best_i = tl.full((), 0, dtype=tl.int32)          # scalar
    sum_w = tl.full((), 0.0, dtype=tl.float32)       # scalar

    for start in tl.static_range(0, V, BLOCK):
        offs = start + tl.arange(0, BLOCK)
        mask = offs < V
        w = tl.load(row_w_ptr + offs * stride_w1, mask=mask, other=0.0)
        sum_w += tl.sum(w, axis=0)

        # RNG: unique per (row, token). Keep offsets in 32-bit domain.
        rng_offsets = (pid * V + offs).to(tl.int32)
        u = tl.rand(seed, rng_offsets)
        u = tl.maximum(u, eps)
        denom = tl.maximum(w, eps)
        r = -tl.log(u) / denom
        # for zero-weight elements, set r to big_pos so they never win
        r = tl.where(w > 0, r, big_pos)

        # block min of r and first index achieving it
        block_min_r = tl.min(r, axis=0)  # scalar
        eq = r == block_min_r
        idxs = offs
        cand = tl.where(eq, idxs, tl.full([BLOCK], V, dtype=tl.int32))
        block_min_idx = tl.min(cand, axis=0)  # scalar

        better = block_min_r < best_r
        best_r = tl.where(better, block_min_r, best_r)
        best_i = tl.where(better, block_min_idx, best_i)

    # If all weights were zero (degenerate), fallback to argmax over original probs
    has_weight = sum_w > 0.0
    if has_weight:
        tl.store(out_ptr + pid, best_i.to(tl.int64))
    else:
        best_val = tl.full((), -float("inf"), dtype=tl.float32)  # scalar
        best_idx = tl.full((), 0, dtype=tl.int32)                # scalar
        for start in tl.static_range(0, V, BLOCK):
            offs = start + tl.arange(0, BLOCK)
            mask = offs < V
            vals = tl.load(row_p_ptr + offs * stride_p1, mask=mask, other=-float("inf"))
            block_max = tl.max(vals, axis=0)  # scalar
            eq = vals == block_max
            idxs = offs
            cand = tl.where(eq, idxs, tl.full([BLOCK], V, dtype=tl.int32))
            block_idx = tl.min(cand, axis=0)  # scalar
            better = block_max > best_val
            equal = block_max == best_val
            take_idx = tl.where(equal, block_idx < best_idx, better)
            best_idx = tl.where(take_idx, block_idx, best_idx)
            best_val = tl.where(take_idx, block_max, best_val)
        tl.store(out_ptr + pid, best_idx.to(tl.int64))


def _as_cuda(t: torch.Tensor) -> torch.Tensor:
    if t.is_cuda:
        return t
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is required but not available. Cannot move CPU tensors to GPU.")
    return t.cuda(non_blocking=True)


def _check_inputs(probs: torch.Tensor, top_p: torch.Tensor):
    if probs.dtype != torch.float32:
        raise TypeError(f"probs must be float32, got {probs.dtype}")
    if probs.dim() != 2:
        raise ValueError(f"probs must be 2D [batch, vocab], got shape {tuple(probs.shape)}")
    if probs.shape[1] != VOCAB_SIZE:
        raise ValueError(f"Expected vocab_size={VOCAB_SIZE}, got {probs.shape[1]}")
    if top_p.dtype not in (torch.float32, torch.float64):
        raise TypeError(f"top_p must be float tensor, got {top_p.dtype}")
    if top_p.dim() != 1:
        raise ValueError(f"top_p must be 1D [batch], got shape {tuple(top_p.shape)}")
    if top_p.shape[0] != probs.shape[0]:
        raise ValueError(f"top_p length {top_p.shape[0]} must match batch_size {probs.shape[0]}")


def _build_top_p_filtered_work(probs: torch.Tensor, top_p: torch.Tensor) -> torch.Tensor:
    """
    Build a filtered copy of probs for top-p sampling:
      - For rows with 0 < p < 1: zero out tokens outside the nucleus (HuggingFace-style mask).
      - For rows with p <= 0 or p >= 1: leave as original probs.
    All operations are on the same device as probs (expected GPU).
    """
    B, V = probs.shape
    work = probs.clone()

    mask_rows = (top_p > 0.0) & (top_p < 1.0)
    if mask_rows.any():
        sel_probs = probs[mask_rows]  # [B_sel, V]
        sel_p = top_p[mask_rows].unsqueeze(1)  # [B_sel, 1]

        # Sort descending
        vals, idxs = torch.sort(sel_probs, dim=1, descending=True)  # [B_sel, V], [B_sel, V]
        cdf = torch.cumsum(vals, dim=1)  # [B_sel, V]

        # HF-style nucleus mask: remove tokens strictly after the first that crosses p
        to_remove = cdf > sel_p
        to_remove_shifted = torch.zeros_like(to_remove)
        to_remove_shifted[:, 1:] = to_remove[:, :-1]
        keep_sorted = ~to_remove_shifted

        filtered_sorted = vals * keep_sorted.to(vals.dtype)

        # Scatter back into original index space
        work_sel = torch.zeros_like(sel_probs)
        work_sel.scatter_(1, idxs, filtered_sorted)
        work[mask_rows] = work_sel

    return work


@torch.no_grad()
def run(*args, **kwargs):
    """
    Entry point: top_p_sampling_from_probs_v151936
    Inputs:
      - probs: [batch, 151936] float32 probabilities (after softmax)
      - top_p: [batch] float32 cumulative probability thresholds
    Output:
      - samples: [batch] int64 sampled token indices
    """
    # Extract args
    if len(args) >= 2:
        probs, top_p = args[0], args[1]
    else:
        probs = kwargs.get("probs", None)
        top_p = kwargs.get("top_p", None)
    if probs is None or top_p is None:
        raise ValueError("run expects arguments (probs, top_p) either as positional or keyword.")

    _check_inputs(probs, top_p)

    # Preserve original device of probs
    orig_device_probs = probs.device

    # Ensure CUDA tensors
    probs_cuda = _as_cuda(probs.contiguous())
    top_p_cuda = _as_cuda(top_p.to(torch.float32).contiguous())

    B, V = probs_cuda.shape
    assert V == VOCAB_SIZE, f"Expected vocab={VOCAB_SIZE}, got {V}"

    # Build filtered work tensor for 0 < p < 1 rows (GPU)
    work = _build_top_p_filtered_work(probs_cuda, top_p_cuda)

    # Output buffer on GPU
    out_cuda = torch.empty((B,), dtype=torch.int64, device=probs_cuda.device)

    # Strides in elements (not bytes)
    stride_p0, stride_p1 = probs_cuda.stride()
    stride_w0, stride_w1 = work.stride()

    # Random seed: uint32
    seed = torch.randint(0, 2**31 - 1, (1,), device=probs_cuda.device, dtype=torch.int64).item()
    seed = int(seed & 0xFFFFFFFF)

    # Kernel launch configuration - tuned for large V on B200
    BLOCK = 4096  # tile over vocab
    num_warps = 8
    num_stages = 2

    grid = lambda META: (B,)

    _top_p_sample_kernel[grid](
        probs_cuda,
        work,
        top_p_cuda,
        out_cuda,
        B,
        stride_p0, stride_p1,
        stride_w0, stride_w1,
        seed,
        V=VOCAB_SIZE,
        BLOCK=BLOCK,
        num_warps=num_warps,
        num_stages=num_stages,
    )

    # Move result back to original device of probs
    if orig_device_probs.type == "cpu":
        return out_cuda.cpu()
    return out_cuda
scrolls · 245 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON