Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / tritonda906d

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-da906d?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:638e4458bb0f554ac090b6793ad99dc5013e5c0e419e14c4c40ee7035e1a65e3
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,
tile-n = 2048BLOCK_N = 2048

Kernel source

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


# Kernel: per-row sampling
# - If k <= 0 or k >= vocab_size: sample from full distribution (unnormalized probs)
# - Else: write sentinel -1 to signal host-side top-k sampling fallback
@triton.jit
def _sample_full_or_signal_topk(
    probs_ptr,           # float32*  [B, V]
    topk_ptr,            # int32*    [B]
    rand_ptr,            # float32*  [B], uniform in [0, 1)
    samples_ptr,         # int64*    [B]
    stride_row,          # int32     row stride in elements
    vocab_size: tl.constexpr,  # constexpr = 151936
    BLOCK_N: tl.constexpr,     # tile size along vocab dimension
):
    pid = tl.program_id(axis=0)
    row_ptr = probs_ptr + pid * stride_row

    # Load k and random u for this row
    k = tl.load(topk_ptr + pid)
    u = tl.load(rand_ptr + pid)

    # Decide path
    k_no_filter = (k <= 0) | (k >= vocab_size)

    # Default: signal fallback (sentinel -1)
    selected_idx = tl.full((), -1, dtype=tl.int32)

    if k_no_filter:
        # First pass: total sum across the row
        total_sum = tl.zeros((), dtype=tl.float32)
        for start in range(0, vocab_size, BLOCK_N):
            offs = start + tl.arange(0, BLOCK_N)
            mask = offs < vocab_size
            vals = tl.load(row_ptr + offs, mask=mask, other=0.0)
            total_sum += tl.sum(vals, axis=0)

        # Handle degenerate case: if total_sum <= 0, return argmax index (tie -> smallest idx)
        if total_sum <= 0:
            neg_inf = tl.full((), -float("inf"), dtype=tl.float32)
            best_val = neg_inf
            best_idx = tl.full((), 0, dtype=tl.int32)
            big_int = tl.full((), 2147483647, dtype=tl.int32)
            for start in range(0, vocab_size, BLOCK_N):
                offs = start + tl.arange(0, BLOCK_N)
                mask = offs < vocab_size
                vals = tl.load(row_ptr + offs, mask=mask, other=0.0)
                # invalid lanes -> -inf
                vals = tl.where(mask, vals, neg_inf)
                tile_max = tl.max(vals, axis=0)
                is_eq = vals == tile_max
                tie_idx = tl.min(tl.where(is_eq, offs, big_int), axis=0)
                take = tile_max > best_val
                best_val = tl.where(take, tile_max, best_val)
                best_idx = tl.where(take, tie_idx, best_idx)
            selected_idx = best_idx
        else:
            # Second pass: sample categorical by threshold t = u * total_sum
            t = u * total_sum
            prefix = tl.zeros((), dtype=tl.float32)
            found = tl.full((), 0, dtype=tl.int32)
            big_int = tl.full((), 2147483647, dtype=tl.int32)
            for start in range(0, vocab_size, BLOCK_N):
                offs = start + tl.arange(0, BLOCK_N)
                mask = offs < vocab_size
                vals = tl.load(row_ptr + offs, mask=mask, other=0.0)
                block_sum = tl.sum(vals, axis=0)
                need = (found == 0) & (prefix + block_sum >= t)
                if need:
                    # Find the first index within this block where cumsum crosses t
                    v = tl.where(mask, vals, 0.0)
                    csum = tl.cumsum(v, axis=0)
                    thr = t - prefix
                    cross = csum >= thr
                    cand = tl.where(cross, offs, big_int)
                    pick = tl.min(cand, axis=0)
                    selected_idx = pick
                    found = tl.full((), 1, dtype=tl.int32)
                else:
                    # update prefix only if not yet found
                    prefix = tl.where(found == 0, prefix + block_sum, prefix)
            # Fallback in case of numerical issues: last index
            selected_idx = tl.where(found == 1, selected_idx, vocab_size - 1)

    # Store result as int64
    tl.store(samples_ptr + pid, tl.cast(selected_idx, tl.int64))


def _ensure_cuda_device(t: torch.Tensor, name: str) -> torch.device:
    if t.is_cuda:
        return t.device
    if torch.cuda.is_available():
        return torch.device("cuda")
    raise RuntimeError(f"CUDA is required to run this kernel, but {name} is on CPU and no CUDA device is available.")


def _move_to_device(t: torch.Tensor, device: torch.device) -> torch.Tensor:
    if t.device != device:
        return t.to(device, non_blocking=True)
    return t


@torch.no_grad()
def run(*args, **kwargs):
    # Parse inputs
    if len(args) >= 2:
        probs, top_k = args[0], args[1]
    else:
        probs = kwargs.get("probs", None)
        top_k = kwargs.get("top_k", None)
    if probs is None or top_k is None:
        raise ValueError("Missing required arguments: probs and top_k")

    if probs.dim() != 2:
        raise ValueError("probs must be a 2D tensor of shape [batch_size, vocab_size]")

    batch_size, vocab_size = probs.shape
    if vocab_size != 151936:
        raise AssertionError(f"Expected vocab_size=151936, got {vocab_size}")

    # Ensure dtype
    probs = probs.to(dtype=torch.float32)

    # Device management
    device = _ensure_cuda_device(probs, "probs")
    _ = _ensure_cuda_device(top_k, "top_k")  # just to validate availability

    probs_gpu = _move_to_device(probs, device)
    top_k_gpu = _move_to_device(top_k, device).to(dtype=torch.int32)

    # Output and RNG
    samples_gpu = torch.empty((batch_size,), dtype=torch.int64, device=device)
    rand_gpu = torch.rand((batch_size,), dtype=torch.float32, device=device)

    # Kernel launch params
    stride_row = probs_gpu.stride(0)
    grid = (batch_size,)
    # BLOCK_N tuned for large V on B200; adjust if needed
    BLOCK_N = 2048

    _sample_full_or_signal_topk[grid](
        probs_gpu,
        top_k_gpu,
        rand_gpu,
        samples_gpu,
        stride_row,
        vocab_size=vocab_size,
        BLOCK_N=BLOCK_N,
        num_warps=8,
        num_stages=3,
    )

    # Host-side fallback for rows requiring top-k filtering (0 < k < vocab_size)
    # We detect those rows by sentinel -1 written by the kernel.
    with torch.no_grad():
        sentinel = (samples_gpu == -1)
        if sentinel.any():
            idx_rows = torch.nonzero(sentinel, as_tuple=False).squeeze(1)
            for r in idx_rows.tolist():
                row = probs_gpu[r]
                k = int(top_k_gpu[r].item())
                # Only do top-k filtering if 0 < k < vocab_size; otherwise, keep kernel result
                if 0 < k < vocab_size:
                    # Select top-k tokens and sample among them
                    top_vals, top_idx = torch.topk(row, k, largest=True, sorted=False)
                    denom = top_vals.sum()
                    if float(denom.item()) > 0.0:
                        weights = top_vals / denom
                        sel_local = torch.multinomial(weights, 1, replacement=True).item()
                        token = int(top_idx[sel_local].item())
                    else:
                        token = int(torch.argmax(row).item())
                    samples_gpu[r] = token
                else:
                    # Shouldn't happen because kernel handled this case, but in case it does:
                    # sample from full distribution via torch.multinomial
                    denom = row.sum()
                    if float(denom.item()) > 0.0:
                        weights = row / denom
                        samples_gpu[r] = int(torch.multinomial(weights, 1, replacement=True).item())
                    else:
                        samples_gpu[r] = int(torch.argmax(row).item())

    # Move to original device if needed
    if probs.device.type != "cuda":
        samples_out = samples_gpu.cpu()
    else:
        samples_out = samples_gpu

    return samples_out
scrolls · 194 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON