Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / triton8dfa99

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-8dfa99?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:ed6087065650deb7bcb452d8e1c601225d4b32a281576952903c3411108a8b2b
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 = 4_copy_i64_kernel[grid](samples_tmp, samples_out_dev, N, BLOCK=BLOCK, num_warps=4, num_stages=2)
stages = 2_copy_i64_kernel[grid](samples_tmp, samples_out_dev, N, BLOCK=BLOCK, num_warps=4, num_stages=2)

Kernel source

main.py193 lines
import math
from typing import Any, Dict

import torch
import triton
import triton.language as tl


def _ensure_cuda_device():
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is required to run the Triton kernel but is not available.")


def _prepare_device(t: torch.Tensor, device: torch.device) -> torch.Tensor:
    if t.device == device:
        return t
    return t.to(device)


def _ceil_div(a: int, b: int) -> int:
    return (a + b - 1) // b


@triton.jit
def _copy_i64_kernel(
    src_ptr,          # *i64 [N]
    dst_ptr,          # *i64 [N]
    n_elements,       # i32
    BLOCK: tl.constexpr,
):
    pid = tl.program_id(axis=0)
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offs < n_elements
    vals = tl.load(src_ptr + offs, mask=mask, other=tl.zeros((), dtype=tl.int64))
    tl.store(dst_ptr + offs, vals, mask=mask)


@torch.no_grad()
def run(probs: torch.Tensor, top_p: torch.Tensor) -> torch.Tensor:
    """
    top_p_sampling_from_probs_v129280
    Inputs:
      - probs: [batch_size, 129280] float32, probabilities (after softmax)
      - top_p: [batch_size] float32
    Outputs:
      - samples: [batch_size] int64

    Semantics match the provided reference exactly:
      - p <= 0.0: greedy argmax
      - 0.0 < p < 1.0: nucleus (top-p) sampling with "shift-keep" semantics
      - otherwise (p >= 1.0 or NaN): sample from full distribution
    """
    # Validate inputs
    if not isinstance(probs, torch.Tensor) or not isinstance(top_p, torch.Tensor):
        raise TypeError("probs and top_p must be torch.Tensor objects.")
    if probs.ndim != 2:
        raise ValueError(f"probs must be 2D [batch_size, vocab_size], got shape {tuple(probs.shape)}")
    if top_p.ndim != 1:
        raise ValueError(f"top_p must be 1D [batch_size], got shape {tuple(top_p.shape)}")
    if probs.shape[0] != top_p.shape[0]:
        raise ValueError("probs.shape[0] (batch_size) must match top_p.shape[0].")
    B, V = probs.shape
    if V != 129280:
        raise AssertionError(f"vocab_size must be 129280, got {V}")

    # Choose/prepare device
    if probs.is_cuda:
        device = probs.device
    elif top_p.is_cuda:
        device = top_p.device
    else:
        _ensure_cuda_device()
        device = torch.device("cuda", index=torch.cuda.current_device())

    orig_device = probs.device

    # Cast and move to GPU
    probs_gpu = _prepare_device(probs.to(dtype=torch.float32), device)
    top_p_gpu = _prepare_device(top_p.to(dtype=torch.float32), device)
    if not probs_gpu.is_contiguous():
        probs_gpu = probs_gpu.contiguous()
    if not top_p_gpu.is_contiguous():
        top_p_gpu = top_p_gpu.contiguous()

    # Output buffer on device
    samples_tmp = torch.empty(B, dtype=torch.int64, device=device)

    # Masks for cases - match reference control flow precisely, including NaN behavior
    # - p <= 0.0 -> argmax
    # - 0.0 < p < 1.0 -> top-p
    # - else (p >= 1.0 or NaN) -> full distribution
    mask_top_p = (top_p_gpu > 0.0) & (top_p_gpu < 1.0)
    mask_argmax = (top_p_gpu <= 0.0)
    mask_full = ~(mask_top_p | mask_argmax)

    # Case A: p <= 0 -> greedy argmax
    if mask_argmax.any():
        rows = mask_argmax.nonzero(as_tuple=False).squeeze(-1)
        rows_probs = probs_gpu.index_select(0, rows)
        argmax_idx = torch.argmax(rows_probs, dim=1)
        samples_tmp.index_copy_(0, rows, argmax_idx.to(torch.int64))

    # Case B: otherwise (p >= 1.0 or NaN) -> sample full distribution
    if mask_full.any():
        rows = mask_full.nonzero(as_tuple=False).squeeze(-1)
        full_rows = probs_gpu.index_select(0, rows)
        # Use torch.multinomial directly; assumes non-negative inputs (softmax outputs)
        # Degenerate rows (sum <= 0) fallback to argmax
        row_sums = full_rows.sum(dim=1)
        zero_sum_mask = row_sums <= 0.0
        if zero_sum_mask.any():
            zrows = rows[zero_sum_mask]
            zargmax = torch.argmax(probs_gpu.index_select(0, zrows), dim=1)
            samples_tmp.index_copy_(0, zrows, zargmax.to(torch.int64))
        nz_mask = ~zero_sum_mask
        if nz_mask.any():
            nz_rows = rows[nz_mask]
            nz_full = full_rows[nz_mask]
            picked = torch.multinomial(nz_full, num_samples=1, replacement=True).squeeze(1)
            samples_tmp.index_copy_(0, nz_rows, picked.to(torch.int64))

    # Case C: 0 < p < 1 -> nucleus (top-p) sampling with exact "shift-keep" semantics
    if mask_top_p.any():
        rows_all = mask_top_p.nonzero(as_tuple=False).squeeze(-1)
        # Process in row-chunks to control peak memory
        # With V=129280, ROWS_CHUNK=32 keeps working set modest
        ROWS_CHUNK = 32
        zeros_cache = torch.zeros((ROWS_CHUNK, V), dtype=torch.float32, device=device)
        for start in range(0, rows_all.numel(), ROWS_CHUNK):
            rows = rows_all[start : start + ROWS_CHUNK]
            sub = probs_gpu.index_select(0, rows)  # [R, V]
            R = sub.size(0)
            # Sort descending per row
            vals, idx = torch.sort(sub, dim=1, descending=True)  # both [R, V]
            # CDF
            cdf = torch.cumsum(vals, dim=1)
            # Build "to_remove" mask and shift as per reference to keep the first crossing token
            p_rows = top_p_gpu.index_select(0, rows).unsqueeze(1)  # [R, 1]
            to_remove = cdf > p_rows
            if V > 1:
                to_remove[:, 1:] = to_remove[:, :-1].clone()
            to_remove[:, 0] = False
            keep = ~to_remove

            # Keep values in original index space using scatter, matching reference implementation
            masked_vals = torch.where(keep, vals, torch.zeros_like(vals))
            # Allocate filtered distribution (re-use cached buffer when possible)
            if R != zeros_cache.size(0):
                filtered = torch.zeros_like(sub)
            else:
                filtered = zeros_cache[:R, :].zero_()
            filtered.scatter_(dim=1, index=idx, src=masked_vals)

            # Normalize the filtered distribution; handle degenerate rows
            sums = filtered.sum(dim=1, keepdim=True)  # [R, 1]
            deg_mask = (sums.squeeze(1) <= 0.0) | (~torch.isfinite(sums.squeeze(1)))
            picked_orig = torch.empty(R, dtype=torch.int64, device=device)

            if (~deg_mask).any():
                nz_rows_mask = ~deg_mask
                dist = filtered[nz_rows_mask] / sums[nz_rows_mask]
                pos = torch.multinomial(dist, num_samples=1, replacement=True).squeeze(1)
                picked_orig[nz_rows_mask] = pos.to(torch.int64)

            if deg_mask.any():
                # Fallback to argmax: idx[:, 0] maps to original index of top-1
                deg_idx0 = idx[deg_mask, 0]
                picked_orig[deg_mask] = deg_idx0.to(torch.int64)

            samples_tmp.index_copy_(0, rows, picked_orig)

    # Copy via Triton kernel (ensures Triton usage and allows future fusing)
    samples_out_dev = torch.empty_like(samples_tmp)
    N = samples_tmp.numel()
    BLOCK = 256
    grid = (_ceil_div(N, BLOCK),)
    _copy_i64_kernel[grid](samples_tmp, samples_out_dev, N, BLOCK=BLOCK, num_warps=4, num_stages=2)

    # Move back to original device if needed
    if orig_device != device:
        samples_out = samples_out_dev.to(orig_device)
    else:
        samples_out = samples_out_dev

    return samples_out


def entrypoint(*args: Any, **kwargs: Dict[str, Any]) -> torch.Tensor:
    if len(args) == 2 and not kwargs:
        return run(args[0], args[1])
    if "probs" in kwargs and "top_p" in kwargs:
        return run(kwargs["probs"], kwargs["top_p"])
    raise ValueError("Expected arguments: run(probs, top_p) or entrypoint(probs=..., top_p=...).")
scrolls · 193 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON