Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / triton44f7ae

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-44f7ae?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:3c11b901ec5b15f8a93b72673b6125fdf63413ff7b4b055983827bb4e6369784
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 = 3num_stages=3,

Kernel source

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


VOCAB_SIZE = 128256


@triton.jit
def _top_p_sample_sorted_kernel(
    vals_ptr,          # float32 [B, V] sorted descending per row
    idx_ptr,           # int64   [B, V] corresponding original indices
    top_p_ptr,         # float32 [B]
    rand_ptr,          # float32 [B] uniform in [0, 1)
    out_ptr,           # int64   [B]
    batch_size: tl.constexpr,
    vocab_size: tl.constexpr,
    CHUNK_SIZE: tl.constexpr,
    CHUNKS: tl.constexpr,
):
    pid = tl.program_id(axis=0)
    if pid >= batch_size:
        return

    # Base pointers for this row
    row_vals_ptr = vals_ptr + pid * vocab_size
    row_idx_ptr = idx_ptr + pid * vocab_size

    # Per-row parameters
    p = tl.load(top_p_ptr + pid)
    u01 = tl.load(rand_ptr + pid)

    # Degenerate case: p <= 0.0 -> argmax (sorted -> first element)
    if p <= 0.0:
        best_idx = tl.load(row_idx_ptr + 0)
        tl.store(out_ptr + pid, best_idx)
        return

    # Compile-time arange for chunk indexing
    ar = tl.arange(0, CHUNK_SIZE)

    # First pass: find truncation boundary (if 0 < p < 1), and total mass
    total_mass = tl.full((), 0.0, dtype=tl.float32)
    cum_before = tl.full((), 0.0, dtype=tl.float32)
    found = tl.full((), False, dtype=tl.int1)
    bound_chunk = tl.full((), 0, dtype=tl.int32)      # chunk id where we cross p
    bound_i_local = tl.full((), 0, dtype=tl.int32)    # local index within bound_chunk
    t_mass = tl.full((), 0.0, dtype=tl.float32)       # truncated mass up to boundary (inclusive)

    big_i_vec = tl.full([CHUNK_SIZE], 2147483647, dtype=tl.int32)  # for reductions

    for j in range(CHUNKS):
        base = j * CHUNK_SIZE
        offs = base + ar
        valid = offs < vocab_size
        v = tl.load(row_vals_ptr + offs, mask=valid, other=0.0)
        s_chunk = tl.sum(v, axis=0)

        # Always accumulate total mass (for treat_full case)
        total_mass = total_mass + s_chunk

        # Check if we need to search the crossing in this chunk
        truncated = p < 1.0
        need = truncated & (~found)

        # Compute prefix sums within this chunk (masked by valid) plus cumulative before
        pref = tl.cumsum(v, axis=0) + cum_before

        # Determine if crossing happens within this chunk
        is_cross = pref > p
        # Replace tl.any with reduction to float and comparison
        any_cross = tl.max(tl.where(is_cross, 1.0, 0.0), axis=0) > 0.0
        any_cross = need & any_cross

        # First crossing index within this chunk (if any)
        idx_first = tl.min(tl.where(is_cross, ar, big_i_vec), axis=0)
        # Mass at crossing
        pref_selected = tl.sum(tl.where(ar == idx_first, pref, 0.0), axis=0)

        # Update boundary if we found crossing here
        bound_chunk = tl.where(any_cross, tl.full((), j, dtype=tl.int32), bound_chunk)
        bound_i_local = tl.where(any_cross, idx_first, bound_i_local)
        t_mass = tl.where(any_cross, pref_selected, t_mass)
        found = found | any_cross

        # If still not found and truncating, accumulate chunk mass into cum_before
        cum_before = tl.where(need & (~any_cross), cum_before + s_chunk, cum_before)

    # If not truncating or we never crossed p, treat as full distribution
    treat_full = (~(p < 1.0)) | (~found)
    t_mass = tl.where(treat_full, total_mass, t_mass)

    # If truncated mass is non-positive (degenerate), fall back to argmax
    if t_mass <= 0.0:
        best_idx = tl.load(row_idx_ptr + 0)
        tl.store(out_ptr + pid, best_idx)
        return

    # Sample u in [0, t_mass)
    u = u01 * t_mass

    # Second pass: locate sampled token by scanning allowed mass
    acc = tl.full((), 0.0, dtype=tl.float32)
    picked = tl.full((), False, dtype=tl.int1)
    sel_off = tl.full((), 0, dtype=tl.int32)

    for j in range(CHUNKS):
        base = j * CHUNK_SIZE
        offs = base + ar
        valid = offs < vocab_size
        v = tl.load(row_vals_ptr + offs, mask=valid, other=0.0)

        # Build allow mask depending on truncation:
        # allow_trunc = valid & ((j < bound_chunk) | ((j == bound_chunk) & (ar <= bound_i_local)))
        j_scalar = tl.full((), j, dtype=tl.int32)
        before = j_scalar < bound_chunk
        at = j_scalar == bound_chunk
        le_local = ar <= bound_i_local
        allow_trunc = valid & (before | (at & le_local))
        allow = tl.where(treat_full, valid, allow_trunc)

        w = tl.where(allow, v, 0.0)
        s_allow = tl.sum(w, axis=0)

        pref = tl.cumsum(w, axis=0) + acc
        cross = pref > u
        any_cross = tl.max(tl.where(cross, 1.0, 0.0), axis=0) > 0.0
        # Only consider first time we find the crossing
        want_pick = (~picked) & any_cross

        idx_first = tl.min(tl.where(cross, ar, big_i_vec), axis=0)
        pos_global = base + idx_first

        sel_off = tl.where(want_pick, pos_global, sel_off)
        picked = picked | want_pick

        # Update accumulator only if we still haven't picked
        acc = tl.where(~picked, acc + s_allow, acc)

    # If nothing picked due to numerical issues, default to first token
    sel_off = tl.where(picked, sel_off, tl.full((), 0, dtype=tl.int32))

    # Fetch original token index and store
    tok_idx = tl.load(row_idx_ptr + sel_off)
    tl.store(out_ptr + pid, tok_idx)


def _ensure_device(t: torch.Tensor, device: torch.device) -> torch.Tensor:
    if t.device == device:
        return t
    if device.type == "cuda":
        return t.to(device, non_blocking=True)
    return t.cpu()


def _validate_and_prepare_inputs(probs: torch.Tensor, top_p: torch.Tensor):
    if probs.dim() != 2:
        raise ValueError(f"probs must be 2D [batch_size, vocab_size], got {tuple(probs.shape)}")
    if probs.shape[1] != VOCAB_SIZE:
        raise ValueError(f"vocab_size must be {VOCAB_SIZE}, got {probs.shape[1]}")
    if top_p.dim() != 1:
        raise ValueError(f"top_p must be 1D [batch_size], got {tuple(top_p.shape)}")
    if top_p.shape[0] != probs.shape[0]:
        raise ValueError(f"top_p batch size {top_p.shape[0]} does not match probs batch size {probs.shape[0]}")
    if probs.dtype != torch.float32:
        probs = probs.to(torch.float32)
    if top_p.dtype != torch.float32:
        top_p = top_p.to(torch.float32)
    return probs, top_p


def _top_p_sample_impl(probs: torch.Tensor, top_p: torch.Tensor, generator: torch.Generator = None) -> torch.Tensor:
    # Validate and cast types
    probs, top_p = _validate_and_prepare_inputs(probs, top_p)
    batch_size, vocab_size = probs.shape

    # Device management
    if probs.is_cuda or top_p.is_cuda:
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available but GPU tensors were provided.")
        # Ensure both inputs on the same CUDA device
        target_device = probs.device if probs.is_cuda else top_p.device
        if probs.is_cuda and top_p.is_cuda and probs.device != top_p.device:
            target_device = probs.device
    else:
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is required to run the Triton kernel, but no CUDA device is available.")
        target_device = torch.device("cuda")

    probs_gpu = _ensure_device(probs.contiguous(), target_device)
    top_p_gpu = _ensure_device(top_p.contiguous(), target_device)

    # Sort probabilities descending per row; get sorted values and original indices
    sorted_vals, sorted_idx = torch.sort(probs_gpu, dim=1, descending=True, stable=False)

    # Random uniforms per row
    if generator is None:
        rand = torch.rand((batch_size,), dtype=torch.float32, device=target_device)
    else:
        try:
            rand = torch.rand((batch_size,), dtype=torch.float32, device=target_device, generator=generator)
        except Exception:
            rand = torch.rand((batch_size,), dtype=torch.float32, device=target_device)

    # Output buffer
    out = torch.empty((batch_size,), dtype=torch.int64, device=target_device)

    # Launch Triton kernel: one program per row
    CHUNK_SIZE = 1024  # tile size for vectorized scanning; good fit for B200
    CHUNKS = (VOCAB_SIZE + CHUNK_SIZE - 1) // CHUNK_SIZE
    grid = (batch_size,)

    _top_p_sample_sorted_kernel[grid](
        sorted_vals,
        sorted_idx,
        top_p_gpu,
        rand,
        out,
        batch_size=batch_size,
        vocab_size=vocab_size,
        CHUNK_SIZE=CHUNK_SIZE,
        CHUNKS=CHUNKS,
        num_warps=8,
        num_stages=3,
    )

    # Move result back to the original device of probs
    out_final = out.to(probs.device) if probs.device.type != "cuda" else out
    return out_final


def run(*args, **kwargs):
    """
    Entry point. Usage:
      samples = run(probs, top_p)
    """
    if len(args) < 2 and not ("probs" in kwargs and "top_p" in kwargs):
        raise ValueError("run requires 'probs' and 'top_p' arguments.")
    probs = args[0] if len(args) > 0 else kwargs["probs"]
    top_p = args[1] if len(args) > 1 else kwargs["top_p"]
    generator = kwargs.get("generator", None)
    return _top_p_sample_impl(probs, top_p, generator=generator)
scrolls · 243 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON