Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / triton7230f5

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-7230f5?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:bdf3112b281d0909c51c10df0c5a799e42686adb033a18606b87af3c1f4561ab
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 = 4num_warps = 4
stages = 2num_stages = 2

Kernel source

main.py229 lines
import torch
import triton
import triton.language as tl


@triton.jit
def _sample_from_topk_kernel(
    topk_vals_ptr,       # float32* [N_valid, Kmax]
    topk_idx_ptr,        # int32*   [N_valid, Kmax]
    sizes_ptr,           # int32*   [N_valid]
    rand_ptr,            # float32* [N_valid], uniform in [0, sum(topk_vals_i))
    row_map_ptr,         # int32*   [N_valid], maps local row -> original batch row
    out_ptr,             # int64*   [batch_size], output indices
    Kmax: tl.constexpr,  # padded max-k across valid rows
    BLOCK: tl.constexpr, # tile size along K dimension, e.g., 256
):
    pid = tl.program_id(0)

    # Per-row metadata (0-D scalars)
    k_size = tl.load(sizes_ptr + pid)
    r = tl.load(rand_ptr + pid)
    row_out = tl.load(row_map_ptr + pid)

    # Running state (0-D scalars)
    acc = tl.zeros((), dtype=tl.float32)            # accumulated sum before current block
    found = tl.zeros((), dtype=tl.int32)            # 0/1 flag
    found_idx_global = tl.full((), -1, dtype=tl.int32)

    row_base = pid * Kmax
    arange = tl.arange(0, BLOCK)

    # Iterate blocks across K dimension with compile-time unrolling
    for start in tl.static_range(0, Kmax, BLOCK):
        offs = start + arange
        # Valid elements within this block for this row
        valid = offs < k_size
        vals = tl.load(topk_vals_ptr + row_base + offs, mask=valid, other=0.0)

        # Sum of this block
        block_sum = tl.sum(vals, axis=0)

        # Will the crossing happen within this block?
        cross_in_block = (found == 0) & (acc + block_sum >= r)

        # Sequential search within the block if needed. Avoid tensor indexing by scalar;
        # instead do masked scalar loads directly from memory.
        rem = k_size - start
        sel = tl.full((), -1, tl.int32)
        run_sum = acc
        for j in tl.static_range(0, BLOCK):
            j_mask = cross_in_block & (j < rem) & (sel < 0)
            v = tl.load(topk_vals_ptr + row_base + (start + j), mask=j_mask, other=0.0)
            run_sum = tl.where(j_mask, run_sum + v, run_sum)
            crossed = j_mask & (run_sum >= r) & (sel < 0)
            sel = tl.where(crossed, tl.full((), j, tl.int32), sel)

        block_found = sel >= 0
        found = tl.where(block_found & (found == 0), tl.full((), 1, tl.int32), found)
        found_idx_global = tl.where(
            block_found & (found_idx_global < 0),
            tl.full((), start, tl.int32) + sel,
            found_idx_global,
        )

        # If still not found, add this block's sum to acc
        acc = tl.where(found == 0, acc + block_sum, acc)

    # Fallback to last valid index if numerical corner-case prevented finding a crossing
    last_idx = tl.where(k_size > 0, k_size - tl.full((), 1, tl.int32), tl.full((), 0, tl.int32))
    final_pos = tl.where(found_idx_global < 0, last_idx, found_idx_global)

    # Gather original token index and store
    tok_i32 = tl.load(topk_idx_ptr + row_base + final_pos)
    tok_i64 = tok_i32.to(tl.int64)
    tl.store(out_ptr + row_out, tok_i64)


def _ensure_cuda_tensor(t: torch.Tensor, device: torch.device):
    if t.device.type == "cuda":
        if t.device != device:
            return t.to(device)
        return t
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is required but not available. Cannot move CPU tensors to GPU.")
    return t.to(device)


@torch.no_grad()
def run(probs, top_k):
    """
    Triton-accelerated top-k sampling from probability distributions.

    Args:
      probs: [batch_size, 128256] float32, probabilities after softmax.
      top_k: [batch_size] int32, per-row top-k values. If 0 < k < vocab_size, restrict to top-k tokens,
             renormalize implicitly via weighted sampling and sample. Otherwise sample from the full distribution.

    Returns:
      samples: [batch_size] int64, sampled token indices.
    """
    # Handle both args and kwargs robustly
    if isinstance(probs, dict):
        probs = probs.get("probs", None)
    if isinstance(top_k, dict):
        top_k = top_k.get("top_k", None)
    if probs is None or top_k is None:
        raise ValueError("Both 'probs' and 'top_k' must be provided.")

    # Basic validation and types
    if probs.ndim != 2:
        raise ValueError(f"probs must be 2D [batch_size, vocab_size], got shape {tuple(probs.shape)}")
    if top_k.ndim != 1:
        raise ValueError(f"top_k must be 1D [batch_size], got shape {tuple(top_k.shape)}")

    batch_size, vocab_size = probs.shape
    if vocab_size != 128256:
        raise AssertionError(f"vocab_size must be 128256, got {vocab_size}")

    # Convert dtypes exactly as in the reference
    probs = probs.to(torch.float32)
    top_k = top_k.to(torch.int32)

    # Device management
    orig_device = probs.device
    if torch.cuda.is_available():
        if probs.is_cuda:
            work_device = probs.device
        elif top_k.is_cuda:
            work_device = top_k.device
        else:
            work_device = torch.device("cuda")
    else:
        if probs.is_cuda or top_k.is_cuda:
            raise RuntimeError("CUDA is not available but a tensor is on GPU.")
        raise RuntimeError("CUDA is required for this Triton kernel but is not available.")

    # Move to working CUDA device if needed
    probs_gpu = _ensure_cuda_tensor(probs, work_device)
    top_k_gpu = _ensure_cuda_tensor(top_k, work_device)

    # Output buffer on GPU
    samples_gpu = torch.empty((batch_size,), dtype=torch.int64, device=work_device)

    # Mask rows by validity of k
    V = vocab_size
    k_valid_mask = (top_k_gpu > 0) & (top_k_gpu < V)
    valid_rows = torch.nonzero(k_valid_mask, as_tuple=False).flatten()
    invalid_rows = torch.nonzero(~k_valid_mask, as_tuple=False).flatten()

    # Handle invalid-k rows: sample from full distribution using torch.multinomial (GPU)
    if invalid_rows.numel() > 0:
        probs_invalid = probs_gpu.index_select(0, invalid_rows).contiguous()
        sampled_invalid = torch.multinomial(probs_invalid, num_samples=1, replacement=True).squeeze(1).to(torch.int64)
        samples_gpu.index_copy_(0, invalid_rows, sampled_invalid)

    # Handle valid-k rows with Triton kernel
    if valid_rows.numel() > 0:
        # Gather valid rows
        probs_valid = probs_gpu.index_select(0, valid_rows).contiguous()
        k_vals = top_k_gpu.index_select(0, valid_rows)  # [N_valid] int32
        Kmax = int(k_vals.max().item())
        N_valid = probs_valid.size(0)

        # Compute top-Kmax once for all valid rows (sorted desc)
        topk_vals_padded, topk_idx_padded = torch.topk(probs_valid, Kmax, dim=1, largest=True, sorted=True)
        topk_vals_padded = topk_vals_padded.contiguous()
        topk_idx_padded = topk_idx_padded.to(torch.int32).contiguous()  # Triton expects int32

        # Compute per-row sums across the first k_i entries only
        ar = torch.arange(Kmax, device=work_device, dtype=torch.int32).unsqueeze(0)  # [1, Kmax]
        sizes_broadcast = k_vals.unsqueeze(1)  # [N_valid, 1]
        mask2d = (ar < sizes_broadcast)  # [N_valid, Kmax], bool
        sums = (topk_vals_padded * mask2d.to(topk_vals_padded.dtype)).sum(dim=1)  # [N_valid]

        # Safety: if any sum is 0 (shouldn't happen), fall back to full distribution for those rows
        zero_sum_mask = sums <= 0
        if torch.any(zero_sum_mask):
            fix_rows_local = torch.nonzero(zero_sum_mask, as_tuple=False).flatten()
            if fix_rows_local.numel() > 0:
                fix_rows_global = valid_rows.index_select(0, fix_rows_local)
                probs_fix = probs_gpu.index_select(0, fix_rows_global).contiguous()
                sampled_fix = torch.multinomial(probs_fix, num_samples=1, replacement=True).squeeze(1).to(torch.int64)
                samples_gpu.index_copy_(0, fix_rows_global, sampled_fix)

            keep_mask = ~zero_sum_mask
            if torch.any(keep_mask):
                keep_idx = torch.nonzero(keep_mask, as_tuple=False).flatten()
                topk_vals_padded = topk_vals_padded.index_select(0, keep_idx).contiguous()
                topk_idx_padded = topk_idx_padded.index_select(0, keep_idx).contiguous()
                k_vals = k_vals.index_select(0, keep_idx).contiguous()
                valid_rows_kernel = valid_rows.index_select(0, keep_idx).contiguous()
                sums = sums.index_select(0, keep_idx).contiguous()
                N_valid_kernel = valid_rows_kernel.numel()
            else:
                N_valid_kernel = 0
        else:
            valid_rows_kernel = valid_rows
            N_valid_kernel = N_valid

        if N_valid_kernel > 0:
            # Prepare random thresholds in [0, sums)
            rands = torch.rand((N_valid_kernel,), dtype=torch.float32, device=work_device) * sums

            # Launch Triton kernel
            grid = (N_valid_kernel,)
            # Tuned params for B200
            num_warps = 4
            num_stages = 2
            BLOCK = 256  # good trade-off for memory coalescing vs. register pressure

            # Row mapping back to global batch indices
            row_map = valid_rows_kernel.to(torch.int32).contiguous()

            _sample_from_topk_kernel[grid](
                topk_vals_padded,
                topk_idx_padded,
                k_vals,
                rands,
                row_map,
                samples_gpu,
                Kmax=Kmax,
                BLOCK=BLOCK,
                num_warps=num_warps,
                num_stages=num_stages,
            )

    # Move result back to original device if needed
    samples = samples_gpu if orig_device.type == "cuda" else samples_gpu.to(orig_device)
    return samples
scrolls · 229 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON