Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / cuda52e243

gpt-5-2025-08-07_cuda_52e243 · gpt-5-2025-08-07 · cuda · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-cuda-52e243?include=source"
interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
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:204e09e8618383458da46f360db82019397aad8def274899c4932de2dc97a367
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20

Kernel source

main.cpp115 lines
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAStream.h>

#include <vector>
#include <random>
#include <limits>
#include "kernel.h"

namespace py = pybind11;

// Validate input tensors and return (batch_size, vocab_size)
static std::pair<int, int> validate_inputs(const torch::Tensor& probs, const torch::Tensor& top_k) {
    TORCH_CHECK(probs.is_cuda(), "probs must be a CUDA tensor");
    TORCH_CHECK(top_k.is_cuda(), "top_k must be a CUDA tensor");
    TORCH_CHECK(probs.dtype() == torch::kFloat32, "probs must be float32");
    TORCH_CHECK(probs.dim() == 2, "probs must be 2D [batch_size, vocab_size]");
    TORCH_CHECK(top_k.dim() == 1, "top_k must be 1D [batch_size]");
    TORCH_CHECK(probs.size(1) == VOCAB_SIZE_V151936, "vocab_size must be exactly 151936");
    TORCH_CHECK(probs.size(0) == top_k.size(0), "batch_size mismatch between probs and top_k");
    TORCH_CHECK(top_k.dtype() == torch::kInt32, "top_k must be int32");

    int batch_size = static_cast<int>(probs.size(0));
    int vocab_size = static_cast<int>(probs.size(1));
    return {batch_size, vocab_size};
}

// Entry point exposed to Python
torch::Tensor run(torch::Tensor probs, torch::Tensor top_k, c10::optional<uint64_t> seed_opt = c10::nullopt) {
    // Validate and make contiguous
    auto sizes = validate_inputs(probs, top_k);
    const int batch_size = sizes.first;
    const int vocab_size = sizes.second; // equals 151936

    auto probs_contig = probs.contiguous();
    auto topk_contig = top_k.contiguous();

    // Output tensor: [batch_size] int64
    auto options_i64 = torch::TensorOptions().dtype(torch::kInt64).device(probs.device());
    torch::Tensor samples = torch::empty({batch_size}, options_i64);

    // Workspace tensors
    auto options_f32 = torch::TensorOptions().dtype(torch::kFloat32).device(probs.device());
    auto options_i32 = torch::TensorOptions().dtype(torch::kInt32).device(probs.device());

    // Sorted outputs (full sort per row)
    torch::Tensor sorted_probs = torch::empty_like(probs_contig, options_f32);
    torch::Tensor indices_in = torch::empty_like(probs_contig, options_i32);
    torch::Tensor sorted_indices = torch::empty_like(indices_in, options_i32);

    // Stream
    auto stream = at::cuda::getCurrentCUDAStream();

    // Build indices_in = [0..V-1] per row
    launch_fill_indices(indices_in.data_ptr<int32_t>(), batch_size, vocab_size, stream.stream());

    // Build segmented offsets
    torch::Tensor begin_offsets = torch::empty({batch_size}, options_i32);
    torch::Tensor end_offsets   = torch::empty({batch_size}, options_i32);
    launch_build_segment_offsets(begin_offsets.data_ptr<int32_t>(),
                                 end_offsets.data_ptr<int32_t>(),
                                 batch_size, vocab_size, stream.stream());

    // Segmented radix sort by probability descending per row
    segmented_sort_pairs_descending(
        probs_contig.data_ptr<float>(),
        sorted_probs.data_ptr<float>(),
        indices_in.data_ptr<int32_t>(),
        sorted_indices.data_ptr<int32_t>(),
        begin_offsets.data_ptr<int32_t>(),
        end_offsets.data_ptr<int32_t>(),
        batch_size,
        static_cast<int64_t>(batch_size) * static_cast<int64_t>(vocab_size),
        stream.stream()
    );

    // RNG seed management
    uint64_t seed;
    if (seed_opt.has_value()) {
        seed = seed_opt.value();
    } else {
        // Non-deterministic seed per call
        std::random_device rd;
        seed = (static_cast<uint64_t>(rd()) << 32) ^ static_cast<uint64_t>(rd());
        if (seed == 0) seed = 0x9E3779B97F4A7C15ULL; // avoid zero seed
    }

    // Sample per row from sorted top-k
    launch_sample_from_sorted_topk(
        sorted_probs.data_ptr<float>(),
        sorted_indices.data_ptr<int32_t>(),
        topk_contig.data_ptr<int32_t>(),
        samples.data_ptr<int64_t>(),
        batch_size,
        vocab_size,
        seed,
        stream.stream()
    );

    // Ensure all work on current stream is queued correctly
    auto err = cudaGetLastError();
    TORCH_CHECK(err == cudaSuccess, "CUDA error prior to return: ", cudaGetErrorString(err));

    return samples;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run", &run,
          py::arg("probs"),
          py::arg("top_k"),
          py::arg("seed") = py::none(),
          "Top-k sampling from probability rows (Qwen3 vocab = 151936). "
          "If 0 < top_k[i] < V, restrict to top-k by probability; otherwise sample from full row. "
          "Returns int64 token indices. Optional seed for reproducibility.");
}
scrolls · 115 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON