Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07_triton_657308

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-657308?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:c03e9d523983ade21c790201f164f3a02af2324a9046fbe226112145c89a703f
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 = 4num_stages=4

Kernel source

main.py440 lines
import math
from typing import Any, Dict, Tuple

import torch
import triton
import triton.language as tl

VOCAB_SIZE = 129280  # constant per spec


@triton.jit
def _argmax_kernel(
    probs_ptr,               # *const float32
    row_ids_ptr,             # *const int32 (indices into batch)
    out_ptr,                 # *mut int64 (write results at absolute row index)
    V: tl.constexpr,         # vocab size
    BLOCK_SIZE: tl.constexpr # tile size along vocab
):
    pid = tl.program_id(0)
    rid = tl.load(row_ids_ptr + pid, mask=True, other=0).to(tl.int32)
    base = rid * V

    best_val = tl.float32(-1.0e30)
    best_idx = tl.int32(0)

    for off in range(0, V, BLOCK_SIZE):
        idxs = off + tl.arange(0, BLOCK_SIZE)
        mask = idxs < V
        vals = tl.load(probs_ptr + base + idxs, mask=mask, other=tl.float32(-1.0e30))
        local_max = tl.max(vals, axis=0)
        local_arg = tl.argmax(vals, axis=0)  # index within the tile
        g_idx = off + local_arg
        better = local_max > best_val
        best_val = tl.where(better, local_max, best_val)
        best_idx = tl.where(better, g_idx, best_idx)

    tl.store(out_ptr + rid, best_idx.to(tl.int64))


@triton.jit
def _sample_full_kernel(
    probs_ptr,               # *const float32
    row_ids_ptr,             # *const int32
    rand_ptr,                # *const float32 (one uniform [0,1) per row)
    out_ptr,                 # *mut int64
    V: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
    SEG: tl.constexpr
):
    pid = tl.program_id(0)
    rid = tl.load(row_ids_ptr + pid, mask=True, other=0).to(tl.int32)
    base = rid * V

    # Pass 1: total sum
    total = tl.float32(0.0)
    for off in range(0, V, BLOCK_SIZE):
        idxs = off + tl.arange(0, BLOCK_SIZE)
        mask = idxs < V
        vals = tl.load(probs_ptr + base + idxs, mask=mask, other=tl.float32(0.0))
        total += tl.sum(vals, axis=0)

    u = tl.load(rand_ptr + rid)
    # Clamp to strictly less than total to avoid boundary issues
    eps = tl.float32(1e-7)
    target = u * total
    target = tl.where(target >= total, total - eps * total, target)

    prefix = tl.float32(0.0)
    res_idx = tl.int32(-1)

    for off in range(0, V, BLOCK_SIZE):
        idxs_block = off + tl.arange(0, BLOCK_SIZE)
        mask_block = idxs_block < V
        vals_block = tl.load(probs_ptr + base + idxs_block, mask=mask_block, other=tl.float32(0.0))
        block_sum = tl.sum(vals_block, axis=0)

        # If threshold is not in this block, skip
        in_block = (res_idx < 0) & (prefix + block_sum > target)

        # If in this block, find the exact index
        if in_block:
            # segmented search to limit inner unroll
            for so in range(0, BLOCK_SIZE, SEG):
                j = so + tl.arange(0, SEG)
                mask_seg = mask_block & (j < BLOCK_SIZE)
                vals_seg = tl.load(probs_ptr + base + off + j, mask=mask_seg, other=tl.float32(0.0))
                seg_sum = tl.sum(vals_seg, axis=0)

                in_seg = (res_idx < 0) & (prefix + seg_sum > target)
                # If the target is in this segment, do a linear search in the small segment
                if in_seg:
                    # Linear search within the small segment (SEG is small, e.g., 128)
                    for t in range(SEG):
                        v = vals_seg[t]
                        prefix = prefix + v
                        found_now = (res_idx < 0) & (prefix > target)
                        idx_found = off + so + t
                        res_idx = tl.where(found_now, idx_found.to(tl.int32), res_idx)
                else:
                    # target not in this segment
                    prefix = prefix + seg_sum

        else:
            prefix = prefix + block_sum

    # Safety: if due to numeric precision res_idx is still -1, choose last valid index
    res_idx = tl.where(res_idx < 0, (V - 1).to(tl.int32), res_idx)
    tl.store(out_ptr + rid, res_idx.to(tl.int64))


@triton.jit
def _topk_sample_kernel(
    probs_ptr,               # *const float32
    topk_ptr,                # *const int32
    topp_ptr,                # *const float32
    rand_ptr,                # *const float32
    row_ids_ptr,             # *const int32
    out_ptr,                 # *mut int64
    V: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
    K_MAX: tl.constexpr
):
    pid = tl.program_id(0)
    rid = tl.load(row_ids_ptr + pid, mask=True, other=0).to(tl.int32)
    base = rid * V

    k_val = tl.load(topk_ptr + rid).to(tl.int32)
    p_val = tl.load(topp_ptr + rid)
    u_val = tl.load(rand_ptr + rid)

    # Buffers for top-k
    sel_idx = tl.full([K_MAX], -1, dtype=tl.int32)
    sel_val = tl.zeros([K_MAX], dtype=tl.float32)

    # Iteratively select global top-k (k <= K_MAX) using masked argmax
    for t in range(K_MAX):
        active = t < k_val
        # initialize best for this iteration
        best_val = tl.float32(-1.0e30)
        best_gidx = tl.int32(0)

        for off in range(0, V, BLOCK_SIZE):
            idxs = off + tl.arange(0, BLOCK_SIZE)
            mask = idxs < V
            vals = tl.load(probs_ptr + base + idxs, mask=mask, other=tl.float32(-1.0e30))

            # Mask out previously selected indices
            # Build a blocked mask: True if idx equals any sel_idx[j] for j < t
            blocked = tl.zeros([BLOCK_SIZE], dtype=tl.int1)
            for j in range(K_MAX):
                if j < t:
                    sidx = sel_idx[j]
                    blocked = blocked | (idxs == sidx)

            masked_vals = tl.where(blocked, tl.float32(-1.0e30), vals)

            local_max = tl.max(masked_vals, axis=0)
            local_arg = tl.argmax(masked_vals, axis=0)
            g_idx = off + local_arg

            better = (local_max > best_val) & active
            best_val = tl.where(better, local_max, best_val)
            best_gidx = tl.where(better, g_idx, best_gidx)

        # Write selected
        if active:
            sel_idx[t] = best_gidx
            sel_val[t] = best_val

    # Sum of top-k values
    sum_k = tl.float32(0.0)
    for t in range(K_MAX):
        if t < k_val:
            sum_k += sel_val[t]

    # Determine allowed prefix count under top-p
    use_topp = p_val < 1.0
    threshold = tl.where(use_topp, p_val * sum_k, sum_k)
    # Compute minimal m such that prefix >= threshold; guarantee at least one
    m_count = tl.int32(0)
    cum = tl.float32(0.0)
    for t in range(K_MAX):
        if t < k_val:
            cum = cum + sel_val[t]
            set_now = use_topp & (m_count == 0) & (cum >= threshold)
            m_count = tl.where(set_now, (t + 1).to(tl.int32), m_count)

    allowed_count = tl.where(use_topp, m_count, k_val)
    # sum over allowed prefix for sampling
    sum_allowed = tl.float32(0.0)
    for t in range(K_MAX):
        if t < allowed_count:
            sum_allowed += sel_val[t]

    # Guard against degenerate sums
    eps = tl.float32(1e-7)
    sum_allowed = tl.where(sum_allowed <= tl.float32(0.0), eps, sum_allowed)
    r = u_val * sum_allowed
    # Clamp r to strictly less than sum_allowed
    r = tl.where(r >= sum_allowed, sum_allowed - eps * sum_allowed, r)

    # Sample within the allowed prefix
    pref = tl.float32(0.0)
    drawn_pos = tl.int32(0)
    taken = False
    for t in range(K_MAX):
        if t < allowed_count:
            pref = pref + sel_val[t]
            take_now = (not taken) & (pref > r)
            drawn_pos = tl.where(take_now, t.to(tl.int32), drawn_pos)
            # "taken" cannot be changed directly as Python bool; emulate
            taken = taken | (pref > r)

    out_idx = sel_idx[drawn_pos]
    tl.store(out_ptr + rid, out_idx.to(tl.int64))


def _to_cuda(t: torch.Tensor) -> torch.Tensor:
    if t.is_cuda:
        return t
    if torch.cuda.is_available():
        return t.cuda()
    raise RuntimeError("CUDA is required but not available; received a CPU tensor.")


def _ensure_dtype(t: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
    if t.dtype != dtype:
        return t.to(dtype)
    return t


def _call_argmax_kernel(probs: torch.Tensor, row_ids: torch.Tensor, out: torch.Tensor):
    assert probs.is_cuda and row_ids.is_cuda and out.is_cuda
    BLOCK_SIZE = 4096
    grid = (row_ids.numel(),)
    _argmax_kernel[grid](
        probs, row_ids, out,
        V=VOCAB_SIZE,
        BLOCK_SIZE=BLOCK_SIZE,
        num_warps=8,
        num_stages=4
    )


def _call_sample_full_kernel(probs: torch.Tensor, row_ids: torch.Tensor, rand: torch.Tensor, out: torch.Tensor):
    assert probs.is_cuda and row_ids.is_cuda and rand.is_cuda and out.is_cuda
    BLOCK_SIZE = 4096
    SEG = 128
    grid = (row_ids.numel(),)
    _sample_full_kernel[grid](
        probs, row_ids, rand, out,
        V=VOCAB_SIZE,
        BLOCK_SIZE=BLOCK_SIZE,
        SEG=SEG,
        num_warps=8,
        num_stages=4
    )


def _call_topk_sample_kernel(
    probs: torch.Tensor,
    top_k: torch.Tensor,
    top_p: torch.Tensor,
    row_ids: torch.Tensor,
    rand: torch.Tensor,
    out: torch.Tensor,
    k_max: int
):
    assert probs.is_cuda and top_k.is_cuda and top_p.is_cuda and row_ids.is_cuda and rand.is_cuda and out.is_cuda
    BLOCK_SIZE = 4096
    grid = (row_ids.numel(),)
    _topk_sample_kernel[grid](
        probs, top_k, top_p, rand, row_ids, out,
        V=VOCAB_SIZE,
        BLOCK_SIZE=BLOCK_SIZE,
        K_MAX=k_max,
        num_warps=8,
        num_stages=4
    )


@torch.no_grad()
def run(*args, **kwargs):
    """
    Entry point:
      run(probs, top_k, top_p) -> samples

    Implements top-k then top-p sampling as specified, optimized with Triton kernels on B200.
    """
    # Handle both positional and keyword forms
    if len(args) == 3 and not kwargs:
        probs, top_k, top_p = args
    else:
        probs = kwargs.get("probs", args[0] if len(args) > 0 else None)
        top_k = kwargs.get("top_k", args[1] if len(args) > 1 else None)
        top_p = kwargs.get("top_p", args[2] if len(args) > 2 else None)

    if probs is None or top_k is None or top_p is None:
        raise ValueError("Missing required arguments: probs, top_k, top_p")

    # Validate shapes and dtypes
    if probs.dim() != 2:
        raise ValueError(f"probs must be 2D [batch_size, vocab_size], got {tuple(probs.shape)}")
    batch_size, vocab_size = probs.shape
    if vocab_size != VOCAB_SIZE:
        raise AssertionError(f"vocab_size must be {VOCAB_SIZE}, got {vocab_size}")
    if top_k.shape != (batch_size,):
        raise ValueError(f"top_k must be shape [{batch_size}], got {tuple(top_k.shape)}")
    if top_p.shape != (batch_size,):
        raise ValueError(f"top_p must be shape [{batch_size}], got {tuple(top_p.shape)}")

    # Keep originals to restore device
    out_device = probs.device

    # Ensure on CUDA
    probs_dev = _to_cuda(_ensure_dtype(probs, torch.float32))
    top_k_dev = _to_cuda(_ensure_dtype(top_k, torch.int32))
    top_p_dev = _to_cuda(_ensure_dtype(top_p, torch.float32))

    # Allocate output on device
    samples_dev = torch.empty(batch_size, dtype=torch.int64, device=probs_dev.device)

    # Random uniforms per row for sampling kernels
    rand = torch.rand(batch_size, dtype=torch.float32, device=probs_dev.device)

    # Build row masks
    with torch.no_grad():
        k = top_k_dev
        p = top_p_dev

        # Masks
        mask_p_le_zero = p <= 0.0
        mask_use_topk_kernel = (p > 0.0) & (k > 0) & (k < VOCAB_SIZE)  # further restricted by K_MAX later
        mask_full_sampling = (p >= 1.0) & ((k <= 0) | (k >= VOCAB_SIZE))
        # The remainder will use a GPU PyTorch fallback for exactness

        # We will use a reasonable K_MAX for the Triton top-k kernel
        K_MAX = 128

        # Split the top-k mask based on K_MAX
        mask_topk_small = mask_use_topk_kernel & (k <= K_MAX)
        mask_topk_large = mask_use_topk_kernel & (k > K_MAX)

        # Category 1: p <= 0 -> argmax (no need to apply top-k since argmax is invariant)
        idxs = torch.nonzero(mask_p_le_zero, as_tuple=False).flatten()
        if idxs.numel() > 0:
            row_ids = idxs.to(torch.int32).contiguous()
            _call_argmax_kernel(probs_dev, row_ids, samples_dev)

        # Category 2: 0 < p, 0 < k < V, k <= K_MAX -> Triton top-k + top-p selection + sampling
        idxs = torch.nonzero(mask_topk_small, as_tuple=False).flatten()
        if idxs.numel() > 0:
            row_ids = idxs.to(torch.int32).contiguous()
            _call_topk_sample_kernel(
                probs_dev,
                top_k_dev,
                top_p_dev,
                row_ids=row_ids,
                rand=rand,
                out=samples_dev,
                k_max=K_MAX
            )

        # Category 3: p >= 1 and (k <= 0 or k >= V) -> sample from full distribution
        idxs = torch.nonzero(mask_full_sampling, as_tuple=False).flatten()
        if idxs.numel() > 0:
            row_ids = idxs.to(torch.int32).contiguous()
            _call_sample_full_kernel(probs_dev, row_ids, rand, samples_dev)

        # Category 4: Fallback exact GPU path using PyTorch ops for all remaining rows
        remaining_mask = ~(mask_p_le_zero | mask_topk_small | mask_full_sampling)
        idxs = torch.nonzero(remaining_mask, as_tuple=False).flatten()

        if idxs.numel() > 0:
            # Process each row independently for exactness, on GPU
            for rid in idxs.tolist():
                row = probs_dev[rid]
                ki = int(k[rid].item())
                pi = float(p[rid].item())

                # Apply top-k filtering if needed
                if 0 < ki < VOCAB_SIZE:
                    vals, idx_sorted = torch.sort(row, descending=True)
                    keep_idx_k = idx_sorted[:ki]
                    filtered_k = torch.zeros_like(row)
                    filtered_k[keep_idx_k] = row[keep_idx_k]
                    row_work = filtered_k
                else:
                    row_work = row

                # Apply top-p if needed
                if pi <= 0.0:
                    # This shouldn't happen due to mask, but keep for safety
                    samples_dev[rid] = torch.argmax(row_work).to(torch.int64)
                    continue

                if pi < 1.0:
                    vals, idx_sorted = torch.sort(row_work, descending=True)
                    cdf = torch.cumsum(vals, dim=0)
                    to_remove = cdf > pi
                    if VOCAB_SIZE > 1:
                        to_remove[1:] = to_remove[:-1].clone()
                        to_remove[0] = False
                    keep_idx_p = idx_sorted[~to_remove]
                    filtered_p = torch.zeros_like(row_work)
                    filtered_p[keep_idx_p] = row_work[keep_idx_p]
                    row_final = filtered_p
                else:
                    row_final = row_work

                # Renormalize and sample
                s = row_final.sum()
                if s.item() <= 0.0:
                    # Degenerate: pick argmax
                    samples_dev[rid] = torch.argmax(row).to(torch.int64)
                else:
                    probs_vec = row_final / s
                    draw = torch.multinomial(probs_vec, 1, replacement=True).squeeze(0)
                    samples_dev[rid] = draw.to(torch.int64)

    # Move to original device if needed
    if samples_dev.device != out_device:
        samples = samples_dev.to(out_device)
    else:
        samples = samples_dev
    return samples


if __name__ == "__main__":
    # Simple sanity check (will run on CUDA if available)
    bs = 4
    V = VOCAB_SIZE
    device = "cuda" if torch.cuda.is_available() else "cpu"
    torch.manual_seed(0)
    probs = torch.rand(bs, V, device=device, dtype=torch.float32)
    probs = probs / probs.sum(dim=1, keepdim=True)
    top_k = torch.tensor([50, 0, 100, 10], device=device, dtype=torch.int32)
    top_p = torch.tensor([0.9, 1.0, 0.95, 0.0], device=device, dtype=torch.float32)
    out = run(probs, top_k, top_p)
    print("Samples:", out)
scrolls · 440 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON