Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07_triton_e65787

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-e65787?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:01b9edb886e80161050e8ab89f93ebbf954341eee5da6c910adee45652f54f91
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,
tile-k = 256BLOCK_K = 256

Kernel source

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


V_CONST = 129280  # DeepSeek V3 vocabulary size (constant)


@triton.jit
def sample_full_kernel(
    probs_ptr,            # *f32 [B, V_CONST]
    rand_ptr,             # *f32 [n_rows]
    rows_idx_ptr,         # *i32 [n_rows] mapping from local row-id -> global row-id
    out_ptr,              # *i64 [B]
    stride_probs,         # i64 stride between rows in probs (in elements)
    n_rows,               # i32 number of rows to process in this launch
    V: tl.constexpr,      # vocab size (constexpr)
    BLOCK: tl.constexpr   # tile width along vocab dimension
):
    pid = tl.program_id(axis=0)
    if pid >= n_rows:
        return

    # Load the mapped global row index and compute base pointer for that row
    row_global_i32 = tl.load(rows_idx_ptr + pid)
    row_global = row_global_i32.to(tl.int64)
    row_ptr = probs_ptr + row_global * stride_probs

    # Load the uniform random number in [0, 1)
    u = tl.load(rand_ptr + pid, eviction_policy='evict_last')

    # Pass 1: compute total mass (sum of probabilities/weights)
    total = tl.zeros((), dtype=tl.float32)
    for off in range(0, V, BLOCK):
        offs = off + tl.arange(0, BLOCK)
        mask = offs < V
        p = tl.load(row_ptr + offs, mask=mask, other=0.0)
        total += tl.sum(p, axis=0)

    # Threshold in [0, total]
    threshold = u * total

    # Pass 2: scan CDF and find first index where CDF >= threshold
    cdf = tl.zeros((), dtype=tl.float32)
    chosen = tl.full((), -1, dtype=tl.int64)
    large = tl.full((), V + 1, dtype=tl.int64)

    for off in range(0, V, BLOCK):
        offs = off + tl.arange(0, BLOCK)
        mask = offs < V
        p = tl.load(row_ptr + offs, mask=mask, other=0.0)
        pref = tl.cumsum(p, axis=0) + cdf
        hit = pref >= threshold
        idxs = (offs).to(tl.int64)
        hit_idxs = tl.where(hit & mask, idxs, large)
        first = tl.min(hit_idxs, axis=0)
        found = first < large
        chosen = tl.where((chosen < 0) & found, first, chosen)
        cdf += tl.sum(p, axis=0)

    # Fallback: if no element found due to numerical issues, select the last index
    chosen = tl.where(chosen < 0, tl.full((), V - 1, dtype=tl.int64), chosen)
    # Write result to the correct global row position
    tl.store(out_ptr + row_global, chosen)


@triton.jit
def sample_topk_kernel(
    vals_ptr,             # *f32 [G, K]
    inds_ptr,             # *i64 [G, K]
    rand_ptr,             # *f32 [G]
    rows_idx_ptr,         # *i32 [G] mapping to global rows
    out_ptr,              # *i64 [B]
    stride_vals,          # i64 stride between rows in vals (in elements)
    stride_inds,          # i64 stride between rows in inds (in elements)
    n_rows,               # i32 number of rows in this group
    K: tl.constexpr,      # number of columns (top-k) for this group (constexpr)
    BLOCK: tl.constexpr   # tile width along K
):
    pid = tl.program_id(axis=0)
    if pid >= n_rows:
        return

    # Pointers to this local row
    row_vals_ptr = vals_ptr + pid * stride_vals
    row_inds_ptr = inds_ptr + pid * stride_inds

    # Mapped global row id for storing final answer
    row_global_i32 = tl.load(rows_idx_ptr + pid)
    row_global = row_global_i32.to(tl.int64)

    # Random uniform in [0, 1)
    u = tl.load(rand_ptr + pid, eviction_policy='evict_last')

    # Pass 1: total mass
    total = tl.zeros((), dtype=tl.float32)
    for off in range(0, K, BLOCK):
        offs = off + tl.arange(0, BLOCK)
        mask = offs < K
        v = tl.load(row_vals_ptr + offs, mask=mask, other=0.0)
        total += tl.sum(v, axis=0)

    # Threshold in [0, total]
    threshold = u * total

    # Pass 2: scan CDF across K and select first where CDF >= threshold
    cdf = tl.zeros((), dtype=tl.float32)
    chosen_local = tl.full((), -1, dtype=tl.int64)
    large = tl.full((), K + 1, dtype=tl.int64)

    for off in range(0, K, BLOCK):
        offs = off + tl.arange(0, BLOCK)
        mask = offs < K
        v = tl.load(row_vals_ptr + offs, mask=mask, other=0.0)
        pref = tl.cumsum(v, axis=0) + cdf
        hit = pref >= threshold
        idxs = offs.to(tl.int64)
        hit_idxs = tl.where(hit & mask, idxs, large)
        first = tl.min(hit_idxs, axis=0)
        found = first < large
        chosen_local = tl.where((chosen_local < 0) & found, first, chosen_local)
        cdf += tl.sum(v, axis=0)

    # If not found (extreme numerical edge), choose last position
    chosen_local = tl.where(chosen_local < 0, tl.full((), K - 1, dtype=tl.int64), chosen_local)

    # Map to original vocab index using inds_ptr
    orig_idx = tl.load(row_inds_ptr + chosen_local)
    tl.store(out_ptr + row_global, orig_idx)


def _ensure_cuda_tensor(t: torch.Tensor, like: torch.device) -> torch.Tensor:
    if t.is_cuda:
        if t.device != like:
            return t.to(like)
        return t
    else:
        return t.to(like)


def run(probs, top_k):
    """
    Triton-accelerated top-k sampling from probability rows.

    Inputs:
      probs: [batch_size, 129280] float32 (probabilities after softmax)
      top_k: [batch_size] int32, per-row top-k to consider. If k <= 0 or k >= 129280, no filtering.

    Output:
      samples: [batch_size] int64 sampled indices per row
    """
    # Basic validation
    if not isinstance(probs, torch.Tensor) or not isinstance(top_k, torch.Tensor):
        raise TypeError("probs and top_k must be torch.Tensor")

    if probs.dim() != 2:
        raise ValueError(f"probs must be 2D [B, V], got shape {tuple(probs.shape)}")

    B, V = probs.shape
    if V != V_CONST:
        raise AssertionError(f"Expected vocab_size == {V_CONST}, got {V}")

    # DType checks/conversions
    if probs.dtype != torch.float32:
        probs = probs.to(torch.float32)

    if top_k.dtype != torch.int32:
        top_k = top_k.to(torch.int32)

    # Device management
    want_cuda = True  # We must run Triton; ensure we are on CUDA
    if want_cuda and not torch.cuda.is_available():
        raise RuntimeError("CUDA is required to run Triton kernels, but torch.cuda.is_not_available().")

    # Track original device for returning output
    orig_device = probs.device

    # Move inputs to CUDA if needed
    device = torch.device("cuda") if not probs.is_cuda else probs.device
    probs = _ensure_cuda_tensor(probs, device)
    top_k = _ensure_cuda_tensor(top_k, device)

    # Prepare output on device
    samples = torch.empty(B, dtype=torch.int64, device=device)

    # Common strides
    stride_probs = probs.stride(0)

    # Determine which rows use filtering
    valid_mask = (top_k > 0) & (top_k < V_CONST)
    invalid_mask = ~valid_mask

    # 1) Invalid k: sample directly from full distribution
    if invalid_mask.any():
        rows_invalid = torch.nonzero(invalid_mask, as_tuple=False).squeeze(1).to(torch.int32)
        n_invalid = rows_invalid.numel()
        if n_invalid > 0:
            rand = torch.rand(n_invalid, device=device, dtype=torch.float32)
            grid = (triton.cdiv(n_invalid, 1),)
            sample_full_kernel[grid](
                probs,
                rand,
                rows_invalid,
                samples,
                stride_probs,
                n_invalid,
                V=V_CONST,
                BLOCK=2048,
                num_warps=8,
                num_stages=4,
            )

    # 2) Valid k: group by unique k and process each group
    if valid_mask.any():
        unique_k = torch.unique(top_k[valid_mask], sorted=False)
        # Ensure unique_k on device
        unique_k = unique_k.to(device)
        for k_val in unique_k.tolist():
            k_int = int(k_val)
            group_mask = valid_mask & (top_k == k_int)
            rows_group = torch.nonzero(group_mask, as_tuple=False).squeeze(1)
            if rows_group.numel() == 0:
                continue
            # Gather rows and compute top-k per row using PyTorch (highly-optimized)
            sub_probs = probs.index_select(0, rows_group)
            # topk returns values and indices along dim=1; order within top-k doesn't affect sampling correctness
            vals, inds = torch.topk(sub_probs, k=k_int, dim=1, largest=True, sorted=False)
            # Normalize to probabilities (avoid division-by-zero by adding tiny eps)
            sums = vals.sum(dim=1, keepdim=True)
            # In case of extreme edge (row all zeros) - keep numeric safety
            eps = 0.0
            vals = vals / (sums + eps)

            G = rows_group.numel()
            rows_group_i32 = rows_group.to(torch.int32)
            rand = torch.rand(G, device=device, dtype=torch.float32)

            grid = (triton.cdiv(G, 1),)
            # Choose a practical block for K scanning; process in tiles if needed
            BLOCK_K = 256
            sample_topk_kernel[grid](
                vals,
                inds,
                rand,
                rows_group_i32,
                samples,
                vals.stride(0),
                inds.stride(0),
                G,
                K=k_int,
                BLOCK=BLOCK_K,
                num_warps=4,
                num_stages=3,
            )

    # Return samples on original device
    if samples.device != orig_device:
        return samples.to(orig_device)
    return samples


if __name__ == "__main__":
    # Minimal sanity check (not exhaustive)
    B = 4
    V = V_CONST
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    if device.type != "cuda":
        raise RuntimeError("This script requires CUDA to run.")

    torch.manual_seed(0)
    probs = torch.randn(B, V, device=device, dtype=torch.float32)
    probs = torch.softmax(probs, dim=1)
    top_k = torch.tensor([0, 1, 32, V_CONST], device=device, dtype=torch.int32)

    out = run(probs, top_k)
    print("Samples:", out)
scrolls · 277 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON