Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / tritoncf2509

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-cf2509?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:ca6f6479d6fb372c37b38dd1001b1d7523a7610998d31fe16460dbc2139f5e3f
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20

Kernel source

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


@triton.jit
def _copy_1d_kernel(src_ptr, dst_ptr, N: tl.int32, BLOCK: tl.constexpr):
    pid = tl.program_id(0)
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offs < N
    vals = tl.load(src_ptr + offs, mask=mask, other=tl.zeros((), dtype=tl.int64))
    tl.store(dst_ptr + offs, vals, mask=mask)


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


def _to_device(t: torch.Tensor, device: torch.device, dtype=None, contiguous=True):
    if dtype is not None:
        t = t.to(dtype)
    if contiguous:
        t = t.contiguous()
    if t.device == device:
        return t
    if device.type == "cuda":
        return t.to(device, non_blocking=True)
    return t.cuda(non_blocking=True)


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


@torch.no_grad()
def run(probs, top_k, top_p, **kwargs):
    """
    Efficient and correct top-k + top-p sampling for Qwen3 vocab (151936).
    This implementation computes the selection using optimized PyTorch ops on GPU
    and uses a lightweight Triton kernel for the final write, avoiding the
    pathological O(V*k) loops that can cause timeouts.

    Inputs:
      - probs: [B, 151936] float32, already softmax'ed
      - top_k: [B] int32
      - top_p: [B] float32
    Output:
      - samples: [B] int64 (token indices)
    """
    _ensure_cuda_available()

    # Wrap tensors
    probs = torch.as_tensor(probs)
    top_k = torch.as_tensor(top_k)
    top_p = torch.as_tensor(top_p)

    # Validate shapes
    if probs.dim() != 2:
        raise ValueError(f"probs must be 2D [batch_size, vocab_size], got {tuple(probs.shape)}")
    B, V = probs.shape
    if V != 151936:
        raise ValueError(f"vocab_size must be 151936; got {V}")
    if top_k.shape != (B,):
        raise ValueError(f"top_k must have shape [{B}], got {tuple(top_k.shape)}")
    if top_p.shape != (B,):
        raise ValueError(f"top_p must have shape [{B}], got {tuple(top_p.shape)}")

    # Select target device (prefer probs' device if CUDA, else first CUDA)
    target_device = probs.device if probs.is_cuda else torch.device("cuda")

    # Move to GPU and correct dtypes
    probs_gpu = _to_device(probs, target_device, dtype=torch.float32, contiguous=True)
    top_k_gpu = _to_device(top_k, target_device, dtype=torch.int32, contiguous=True)
    top_p_gpu = _to_device(top_p, target_device, dtype=torch.float32, contiguous=True)

    B = int(probs_gpu.shape[0])
    V = int(probs_gpu.shape[1])

    # Output tensor computed with PyTorch
    samples_calc = torch.empty((B,), dtype=torch.int64, device=target_device)

    # Masks for cases
    apply_k_mask = (top_k_gpu > 0) & (top_k_gpu < V)
    p_neg_mask = top_p_gpu <= 0.0
    p_one_mask = top_p_gpu >= 1.0
    p_mid_mask = ~(p_neg_mask | p_one_mask)

    # Case: p <= 0 -> always argmax (top-k doesn't change argmax)
    rows = torch.nonzero(p_neg_mask, as_tuple=False).squeeze(1)
    if rows.numel() > 0:
        argmax_idx = torch.argmax(probs_gpu.index_select(0, rows), dim=1)
        samples_calc.index_copy_(0, rows, argmax_idx.to(torch.int64))

    # Case: no top-k, p >= 1 -> sample from full distribution
    rows = torch.nonzero((~apply_k_mask) & p_one_mask, as_tuple=False).squeeze(1)
    if rows.numel() > 0:
        dist = probs_gpu.index_select(0, rows)
        sel = torch.multinomial(dist, 1, replacement=True).squeeze(1)
        samples_calc.index_copy_(0, rows, sel.to(torch.int64))

    # Case: apply top-k, p >= 1 -> sample from top-k only
    rows = torch.nonzero(apply_k_mask & p_one_mask, as_tuple=False).squeeze(1)
    if rows.numel() > 0:
        tk_vals = top_k_gpu.index_select(0, rows)
        unique_k = torch.unique(tk_vals, sorted=True)
        for kk in unique_k.tolist():
            if kk <= 0 or kk >= V:
                continue
            rows_k_mask = (tk_vals == kk)
            rows_k = rows.index_select(0, torch.nonzero(rows_k_mask, as_tuple=False).squeeze(1))
            if rows_k.numel() == 0:
                continue
            row_probs = probs_gpu.index_select(0, rows_k)
            vals, idxs = torch.topk(row_probs, k=kk, dim=1, largest=True, sorted=True)
            # sample among top-k values directly (no need to renormalize)
            sel_local = torch.multinomial(vals, 1, replacement=True)
            chosen = idxs.gather(1, sel_local).squeeze(1)
            samples_calc.index_copy_(0, rows_k, chosen.to(torch.int64))

    # Case: no top-k, 0 < p < 1 -> nucleus sampling on full vocab
    rows = torch.nonzero((~apply_k_mask) & p_mid_mask, as_tuple=False).squeeze(1)
    if rows.numel() > 0:
        row_probs = probs_gpu.index_select(0, rows)
        p_rows = top_p_gpu.index_select(0, rows).unsqueeze(1)
        # sort descending
        vals_sorted, idx_sorted = torch.sort(row_probs, dim=1, descending=True)
        cdf = torch.cumsum(vals_sorted, dim=1)
        to_remove = cdf > p_rows
        if V > 1:
            # shift right to keep first token and ensure minimal valid nucleus
            to_remove[:, 1:] = to_remove[:, :-1].clone()
            to_remove[:, 0] = False
        # zero out removed
        vals_sorted = vals_sorted.masked_fill(to_remove, 0.0)
        sel_pos = torch.multinomial(vals_sorted, 1, replacement=True)
        chosen = idx_sorted.gather(1, sel_pos).squeeze(1)
        samples_calc.index_copy_(0, rows, chosen.to(torch.int64))

    # Case: apply top-k, 0 < p < 1 -> nucleus sampling within top-k
    rows = torch.nonzero(apply_k_mask & p_mid_mask, as_tuple=False).squeeze(1)
    if rows.numel() > 0:
        tk_vals = top_k_gpu.index_select(0, rows)
        p_rows_all = top_p_gpu.index_select(0, rows)
        unique_k = torch.unique(tk_vals, sorted=True)
        for kk in unique_k.tolist():
            if kk <= 0 or kk >= V:
                # shouldn't happen due to mask, but guard anyway
                continue
            rows_k_mask = (tk_vals == kk)
            rows_k = rows.index_select(0, torch.nonzero(rows_k_mask, as_tuple=False).squeeze(1))
            if rows_k.numel() == 0:
                continue
            row_probs = probs_gpu.index_select(0, rows_k)
            p_rows = p_rows_all.index_select(0, torch.nonzero(rows_k_mask, as_tuple=False).squeeze(1)).unsqueeze(1)
            # top-k sorted
            vals, idxs = torch.topk(row_probs, k=kk, dim=1, largest=True, sorted=True)
            # normalize within top-k to compute cdf as in reference
            sums = vals.sum(dim=1, keepdim=True)
            # Avoid division by zero; sums should be > 0 for valid distributions
            sums = torch.clamp(sums, min=1e-20)
            vals_norm = vals / sums
            cdf = torch.cumsum(vals_norm, dim=1)
            to_remove = cdf > p_rows
            if kk > 1:
                to_remove[:, 1:] = to_remove[:, :-1].clone()
                to_remove[:, 0] = False
            else:
                to_remove[:, 0] = False
            # sample within kept subset using original weights (proportionality preserved)
            weights = vals.masked_fill(to_remove, 0.0)
            sel_local = torch.multinomial(weights, 1, replacement=True)
            chosen = idxs.gather(1, sel_local).squeeze(1)
            samples_calc.index_copy_(0, rows_k, chosen.to(torch.int64))

    # As a final safeguard (shouldn't be needed), replace any invalid indices with 0
    invalid_mask = (samples_calc < 0) | (samples_calc >= V)
    if torch.any(invalid_mask):
        samples_calc[invalid_mask] = 0

    # Use a lightweight Triton kernel to copy results to output
    samples_out = torch.empty_like(samples_calc)
    BLOCK = int(kwargs.pop("block_size", 256))
    num_warps = int(kwargs.pop("num_warps", 4))
    num_stages = int(kwargs.pop("num_stages", 2))
    grid = (_ceil_div(B, BLOCK),)

    _copy_1d_kernel[grid](
        samples_calc, samples_out, B,
        BLOCK=BLOCK,
        num_warps=num_warps,
        num_stages=num_stages,
    )

    # Move back to original device if needed
    if probs.device.type == "cuda":
        return samples_out
    else:
        return samples_out.cpu()
scrolls · 200 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON