Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / tritonaf4b72

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-af4b72?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:7039f1a2e1cbda09006e854886751531ef63687d1b5dbee05c8fc1f463545b6b
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,

Kernel source

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


VOCAB_SIZE = 128256


@triton.jit
def sample_from_packed_kernel(
    probs_ptr,          # float32 [total_kept]
    idxs_ptr,           # int32   [total_kept]
    starts_ptr,         # int32   [n_rows]
    lens_ptr,           # int32   [n_rows]
    rand_ptr,           # float32 [n_rows]
    out_ptr,            # int64   [n_rows]
    BLOCK_SIZE: tl.constexpr,
    MAX_TILES: tl.constexpr,
):
    pid = tl.program_id(axis=0)

    # Load start, length, and per-row random
    start = tl.load(starts_ptr + pid)
    length = tl.load(lens_ptr + pid)
    u = tl.load(rand_ptr + pid)

    # Clamp u to [0, 1 - eps) to avoid edge-case where cdf never exceeds u
    eps = 1e-7
    one_minus_eps = 1.0 - eps
    u = tl.where(u < one_minus_eps, u, one_minus_eps)

    running = tl.zeros((), dtype=tl.float32)
    found = tl.full((), 0, tl.int32)
    found_tile = tl.full((), 0, tl.int32)
    carry = tl.zeros((), dtype=tl.float32)

    ar = tl.arange(0, BLOCK_SIZE)

    # Scan tiles to locate the tile containing u
    for t in tl.static_range(0, MAX_TILES):
        col_base = t * BLOCK_SIZE
        rem = length - col_base
        has = rem > 0
        offs = start + col_base + ar
        valid = has & (ar < rem)
        vals = tl.load(probs_ptr + offs, mask=valid, other=0.0)
        tile_sum = tl.sum(vals, axis=0)

        not_found = found == 0
        crosses = not_found & has & ((running + tile_sum) > u)

        carry = tl.where(crosses, running, carry)
        found_tile = tl.where(crosses, tl.full((), t, tl.int32), found_tile)
        found = tl.where(crosses, 1, found)

        running = running + tl.where(has, tile_sum, 0.0)

    # Now search inside the found tile sequentially
    col_base2 = found_tile * BLOCK_SIZE
    base2 = start + col_base2
    rem2 = length - col_base2
    target = u - carry

    acc = tl.zeros((), dtype=tl.float32)
    j = tl.full((), -1, tl.int32)
    for i in tl.static_range(0, BLOCK_SIZE):
        valid_i = i < rem2
        vi = tl.load(probs_ptr + base2 + i, mask=valid_i, other=0.0)
        acc = acc + tl.where(valid_i, vi, 0.0)
        take = (j < 0) & valid_i & (acc >= target)
        j = tl.where(take, tl.full((), i, tl.int32), j)

    last_idx = tl.where(rem2 > 0, rem2 - 1, 0)
    j = tl.where(j >= 0, j, last_idx)

    sel_off = base2 + j
    tok = tl.load(idxs_ptr + sel_off).to(tl.int64)
    tl.store(out_ptr + pid, tok)


@triton.jit
def sample_from_dense_kernel(
    probs_ptr,          # float32 [n_rows, VOCAB_SIZE] base pointer
    stride_row,         # int32 stride in elements between rows
    rand_ptr,           # float32 [n_rows]
    out_ptr,            # int64   [n_rows]
    N_COLS: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
    MAX_TILES: tl.constexpr,
):
    pid = tl.program_id(axis=0)

    # Pointer to the start of this row
    row_ptr = probs_ptr + pid * stride_row
    u = tl.load(rand_ptr + pid)

    # Clamp u to [0, 1 - eps)
    eps = 1e-7
    one_minus_eps = 1.0 - eps
    u = tl.where(u < one_minus_eps, u, one_minus_eps)

    running = tl.zeros((), dtype=tl.float32)
    found = tl.full((), 0, tl.int32)
    found_tile = tl.full((), 0, tl.int32)
    carry = tl.zeros((), dtype=tl.float32)

    ar = tl.arange(0, BLOCK_SIZE)

    # Scan tiles across the row
    for t in tl.static_range(0, MAX_TILES):
        col_base = t * BLOCK_SIZE
        rem = N_COLS - col_base
        has = rem > 0
        offs = col_base + ar
        valid = has & (ar < rem)
        vals = tl.load(row_ptr + offs, mask=valid, other=0.0)
        tile_sum = tl.sum(vals, axis=0)

        not_found = found == 0
        crosses = not_found & has & ((running + tile_sum) > u)

        carry = tl.where(crosses, running, carry)
        found_tile = tl.where(crosses, tl.full((), t, tl.int32), found_tile)
        found = tl.where(crosses, 1, found)

        running = running + tl.where(has, tile_sum, 0.0)

    # Search within found tile sequentially
    col_base2 = found_tile * BLOCK_SIZE
    rem2 = N_COLS - col_base2
    target = u - carry

    acc = tl.zeros((), dtype=tl.float32)
    j = tl.full((), -1, tl.int32)
    for i in tl.static_range(0, BLOCK_SIZE):
        valid_i = i < rem2
        vi = tl.load(row_ptr + col_base2 + i, mask=valid_i, other=0.0)
        acc = acc + tl.where(valid_i, vi, 0.0)
        take = (j < 0) & valid_i & (acc >= target)
        j = tl.where(take, tl.full((), i, tl.int32), j)

    last_idx = tl.where(rem2 > 0, rem2 - 1, 0)
    j = tl.where(j >= 0, j, last_idx)

    tok = (col_base2 + j).to(tl.int64)
    tl.store(out_ptr + pid, tok)


def run(*args, **kwargs):
    """
    Entry point: top_k_top_p_sampling_from_probs_v128256
    Inputs:
      probs: [batch, 128256] float32
      top_k: [batch] int32
      top_p: [batch] float32
    Output:
      samples: [batch] int64
    """
    # Accept args or kwargs
    if len(args) == 3 and not kwargs:
        probs, top_k, top_p = args
    else:
        probs = kwargs.get("probs", args[0] if len(args) > 0 else None)
        top_k = kwargs.get("top_k", args[1] if len(args) > 1 else None)
        top_p = kwargs.get("top_p", args[2] if len(args) > 2 else None)

    if probs is None or top_k is None or top_p is None:
        raise ValueError("Missing required arguments: probs, top_k, top_p")

    # Validate shapes and types
    if probs.ndim != 2:
        raise ValueError("probs must be 2D [batch, vocab_size]")
    if probs.shape[1] != VOCAB_SIZE:
        raise AssertionError(f"vocab_size must be {VOCAB_SIZE}, got {probs.shape[1]}")
    if top_k.ndim != 1 or top_k.shape[0] != probs.shape[0]:
        raise ValueError("top_k must be 1D with length equal to batch size")
    if top_p.ndim != 1 or top_p.shape[0] != probs.shape[0]:
        raise ValueError("top_p must be 1D with length equal to batch size")

    # Device management: ensure CUDA
    original_device = probs.device
    if original_device.type == "cuda":
        device = original_device
    else:
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is required for Triton kernels, but not available.")
        device = torch.device("cuda")

    # Move inputs to device
    probs_dev = probs.to(device=device, dtype=torch.float32, copy=False)
    top_k_dev = top_k.to(device=device, dtype=torch.int32, copy=False)
    top_p_dev = top_p.to(device=device, dtype=torch.float32, copy=False)

    batch = probs_dev.shape[0]
    samples_dev = torch.empty(batch, dtype=torch.int64, device=device)

    # Rows where p <= 0: select argmax (top-k filtering doesn't change argmax)
    mask_p_le_zero = (top_p_dev <= 0.0)
    if mask_p_le_zero.any():
        idx_rows = mask_p_le_zero.nonzero(as_tuple=False).squeeze(1)
        if idx_rows.numel() > 0:
            argmax_idx = torch.argmax(probs_dev.index_select(0, idx_rows), dim=1)
            samples_dev.index_copy_(0, idx_rows, argmax_idx.to(torch.int64))

    # Remaining rows (p > 0) need sampling
    mask_remaining = ~mask_p_le_zero
    if mask_remaining.any():
        rows_remaining = mask_remaining.nonzero(as_tuple=False).squeeze(1)
        k_all = top_k_dev.index_select(0, rows_remaining)
        p_all = top_p_dev.index_select(0, rows_remaining)

        # Case A: No top-k (k <= 0 or k >= vocab) and p >= 1 -> dense sampling from full distribution
        mask_no_topk = (k_all <= 0) | (k_all >= VOCAB_SIZE)
        mask_p_ge_one = p_all >= 1.0
        mask_dense = mask_no_topk & mask_p_ge_one
        if mask_dense.any():
            dense_rows_local = rows_remaining.index_select(0, mask_dense.nonzero(as_tuple=False).squeeze(1))
            if dense_rows_local.numel() > 0:
                dense_probs = probs_dev.index_select(0, dense_rows_local).contiguous()
                # Normalize to guard against small numerical drift
                row_sums = dense_probs.sum(dim=1, keepdim=True)
                dense_probs = dense_probs / torch.clamp(row_sums, min=1e-12)

                n_dense = dense_probs.shape[0]
                rand_u = torch.rand(n_dense, dtype=torch.float32, device=device)
                # Configure kernel
                BLOCK = 512
                MAX_TILES = max(1, triton.cdiv(VOCAB_SIZE, BLOCK))
                grid = (n_dense,)

                sample_out = torch.empty(n_dense, dtype=torch.int64, device=device)
                stride_row = dense_probs.stride(0)

                sample_from_dense_kernel[grid](
                    dense_probs,     # probs_ptr
                    stride_row,      # stride_row
                    rand_u,          # rand_ptr
                    sample_out,      # out_ptr
                    N_COLS=VOCAB_SIZE,
                    BLOCK_SIZE=BLOCK,
                    MAX_TILES=MAX_TILES,
                    num_warps=8,
                    num_stages=4,
                )
                samples_dev.index_copy_(0, dense_rows_local, sample_out)

        # Case B: Other rows -> build packed candidates via top-k and/or nucleus (top-p), then sample via packed kernel
        packed_probs_list = []
        packed_idxs_list = []
        packed_lens = []
        packed_row_ids = []

        # Helper for packing row-wise tensors with a boolean keep mask; expects sorted in descending probability
        def _pack_kept(vals_sorted: torch.Tensor, idx_sorted: torch.Tensor, keep_mask: torch.Tensor, row_ids: torch.Tensor):
            # vals_sorted, idx_sorted, keep_mask: [n, L]
            n = vals_sorted.shape[0]
            if n == 0:
                return
            # Per-row lengths
            lens = keep_mask.sum(dim=1)  # int64
            # Masked flatten
            flat_vals = vals_sorted[keep_mask]
            flat_idx = idx_sorted[keep_mask]
            # Renormalize per-row
            row_ids_expand = torch.repeat_interleave(torch.arange(n, device=device, dtype=torch.long), lens)
            sums = torch.zeros(n, dtype=torch.float32, device=device)
            sums.index_add_(0, row_ids_expand, flat_vals)
            scales = 1.0 / torch.clamp(sums, min=1e-12)
            flat_vals = flat_vals * scales.index_select(0, row_ids_expand)

            # Record
            packed_lens.extend(lens.to(torch.int32).tolist())
            packed_row_ids.extend(row_ids.to(torch.int32).tolist())
            packed_probs_list.append(flat_vals)
            packed_idxs_list.append(flat_idx.to(torch.int32))

        # Precompute masks within 'rows_remaining'
        mask_topk = (k_all > 0) & (k_all < VOCAB_SIZE)
        mask_no_topk_pcut = mask_no_topk & (p_all < 1.0)

        # Process top-k rows grouped by unique k
        if mask_topk.any():
            rows_topk_local = rows_remaining.index_select(0, mask_topk.nonzero(as_tuple=False).squeeze(1))
            k_topk_local = top_k_dev.index_select(0, rows_topk_local)
            p_topk_local = top_p_dev.index_select(0, rows_topk_local)

            # Split by p >= 1 and p < 1
            mask_topk_p_ge_one = (p_topk_local >= 1.0)
            mask_topk_p_lt_one = ~mask_topk_p_ge_one

            # Group by unique k for p >= 1.0
            if mask_topk_p_ge_one.any():
                rows_ge1 = rows_topk_local.index_select(0, mask_topk_p_ge_one.nonzero(as_tuple=False).squeeze(1))
                k_ge1 = k_topk_local.index_select(0, mask_topk_p_ge_one.nonzero(as_tuple=False).squeeze(1))
                uniq_k = torch.unique(k_ge1, sorted=True)
                for kv in uniq_k.tolist():
                    sel = (k_ge1 == kv)
                    if not sel.any():
                        continue
                    grp_rows = rows_ge1.index_select(0, sel.nonzero(as_tuple=False).squeeze(1))
                    grp_probs = probs_dev.index_select(0, grp_rows)
                    # topk sorted descending
                    vals, idx = torch.topk(grp_probs, kv, dim=1, largest=True, sorted=True)
                    # Normalize top-k distribution
                    sums = torch.clamp(vals.sum(dim=1, keepdim=True), min=1e-12)
                    vals = vals / sums
                    keep_mask = torch.ones_like(vals, dtype=torch.bool)
                    _pack_kept(vals, idx, keep_mask, grp_rows)

            # Group by unique k for p < 1.0
            if mask_topk_p_lt_one.any():
                rows_lt1 = rows_topk_local.index_select(0, mask_topk_p_lt_one.nonzero(as_tuple=False).squeeze(1))
                k_lt1 = k_topk_local.index_select(0, mask_topk_p_lt_one.nonzero(as_tuple=False).squeeze(1))
                p_lt1 = p_topk_local.index_select(0, mask_topk_p_lt_one.nonzero(as_tuple=False).squeeze(1))
                uniq_k = torch.unique(k_lt1, sorted=True)
                for kv in uniq_k.tolist():
                    sel = (k_lt1 == kv)
                    if not sel.any():
                        continue
                    grp_rows = rows_lt1.index_select(0, sel.nonzero(as_tuple=False).squeeze(1))
                    grp_probs = probs_dev.index_select(0, grp_rows)
                    grp_p = p_lt1.index_select(0, sel.nonzero(as_tuple=False).squeeze(1))
                    # topk sorted descending
                    vals, idx = torch.topk(grp_probs, kv, dim=1, largest=True, sorted=True)
                    # Normalize
                    sums = torch.clamp(vals.sum(dim=1, keepdim=True), min=1e-12)
                    vals = vals / sums
                    # Nucleus (top-p) selection
                    cdf = torch.cumsum(vals, dim=1)
                    pcol = grp_p.unsqueeze(1)
                    to_remove = cdf > pcol
                    if kv > 1:
                        to_remove[:, 1:] = to_remove[:, :-1].clone()
                        to_remove[:, 0] = False
                    elif kv == 1:
                        to_remove[:, 0] = False
                    keep_mask = ~to_remove
                    _pack_kept(vals, idx, keep_mask, grp_rows)

        # Process no-topk rows with p in (0,1) -> full sort + nucleus
        if mask_no_topk_pcut.any():
            rows_ntk = rows_remaining.index_select(0, mask_no_topk_pcut.nonzero(as_tuple=False).squeeze(1))
            if rows_ntk.numel() > 0:
                p_ntk = top_p_dev.index_select(0, rows_ntk)
                probs_ntk = probs_dev.index_select(0, rows_ntk)
                # Sort full vocab descending
                vals, idx = torch.sort(probs_ntk, dim=1, descending=True)
                # Nucleus selection per row
                cdf = torch.cumsum(vals, dim=1)
                pcol = p_ntk.unsqueeze(1)
                to_remove = cdf > pcol
                if VOCAB_SIZE > 1:
                    to_remove[:, 1:] = to_remove[:, :-1].clone()
                    to_remove[:, 0] = False
                else:
                    to_remove[:, 0] = False
                keep_mask = ~to_remove
                _pack_kept(vals, idx, keep_mask, rows_ntk)

        # Launch packed sampling kernel if we have any rows to sample
        if len(packed_lens) > 0:
            if len(packed_probs_list) == 1:
                packed_probs = packed_probs_list[0]
                packed_idxs = packed_idxs_list[0]
            else:
                packed_probs = torch.cat(packed_probs_list, dim=0)
                packed_idxs = torch.cat(packed_idxs_list, dim=0)

            lens_tensor = torch.tensor(packed_lens, dtype=torch.int32, device=device)
            row_ids_tensor = torch.tensor(packed_row_ids, dtype=torch.int32, device=device)

            # Build starts array
            starts_tensor = torch.zeros_like(lens_tensor)
            if lens_tensor.numel() > 1:
                starts_tensor[1:] = torch.cumsum(lens_tensor[:-1], dim=0)

            # Generate random uniforms per row
            rand_u = torch.rand(lens_tensor.shape[0], dtype=torch.float32, device=device)

            # Configure kernel
            BLOCK = 512
            max_len = int(lens_tensor.max().item())
            MAX_TILES = max(1, triton.cdiv(max_len, BLOCK))
            grid = (lens_tensor.shape[0],)
            sample_out = torch.empty(lens_tensor.shape[0], dtype=torch.int64, device=device)

            sample_from_packed_kernel[grid](
                packed_probs,   # probs_ptr
                packed_idxs,    # idxs_ptr
                starts_tensor,  # starts_ptr
                lens_tensor,    # lens_ptr
                rand_u,         # rand_ptr
                sample_out,     # out_ptr
                BLOCK_SIZE=BLOCK,
                MAX_TILES=MAX_TILES,
                num_warps=8,
                num_stages=4,
            )

            # Scatter back to global samples
            samples_dev.index_copy_(0, row_ids_tensor.to(torch.long), sample_out)

    # Move results to original device if needed
    if original_device.type != "cuda":
        samples = samples_dev.to(device=original_device)
    else:
        samples = samples_dev

    return samples
scrolls · 410 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON