Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07_cuda_724008

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Kernel source

main.cpp163 lines
#include "kernel.h"

#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime.h>

#include <vector>
#include <chrono>
#include <limits>
#include <stdexcept>
#include <sstream>
#include <string>
#include <cstdint>
#include <climits>

namespace py = pybind11;

#ifndef CUDA_CHECK
#define CUDA_CHECK(ans) { gpuAssert((ans), __FILE__, __LINE__, true); }
inline void gpuAssert(cudaError_t code, const char *file, int line, bool abort) {
    if (code != cudaSuccess) {
        std::stringstream ss;
        ss << "CUDA error at " << file << ":" << line << " -> " << cudaGetErrorString(code);
        if (abort) throw std::runtime_error(ss.str());
    }
}
#endif

static inline void check_tensor_is_cuda_contig(const torch::Tensor& t, const std::string& name) {
    if (!t.is_cuda()) {
        throw std::invalid_argument(name + " must be a CUDA tensor");
    }
    if (!t.is_contiguous()) {
        throw std::invalid_argument(name + " must be contiguous");
    }
}

torch::Tensor run(torch::Tensor probs,
                  torch::Tensor top_k,
                  torch::Tensor top_p,
                  c10::optional<long long> seed_opt) {
    // Validate shapes and dtypes
    if (probs.dim() != 2) throw std::invalid_argument("probs must be 2D [batch_size, vocab_size]");
    const int64_t batch_size = probs.size(0);
    const int64_t vocab_size = probs.size(1);
    if (vocab_size != VOCAB_SIZE_129280) {
        throw std::invalid_argument("vocab_size must be 129280");
    }

    if (top_k.dim() != 1 || top_k.size(0) != batch_size) {
        throw std::invalid_argument("top_k must be 1D and match batch_size");
    }
    if (top_p.dim() != 1 || top_p.size(0) != batch_size) {
        throw std::invalid_argument("top_p must be 1D and match batch_size");
    }

    // DType and device checks
    if (probs.scalar_type() != at::kFloat) {
        throw std::invalid_argument("probs must be float32");
    }
    check_tensor_is_cuda_contig(probs, "probs");

    if (top_p.scalar_type() != at::kFloat) {
        throw std::invalid_argument("top_p must be float32");
    }
    // top_p may be CPU input; move to device
    torch::Tensor top_p_dev = top_p.is_cuda() ? top_p : top_p.to(probs.device(), /*non_blocking=*/true);
    torch::Tensor top_p_f32 = top_p_dev.contiguous();

    if (!(top_k.scalar_type() == at::kInt || top_k.scalar_type() == at::kLong)) {
        throw std::invalid_argument("top_k must be int32 or int64");
    }

    // Ensure top_k is on device and int32 contiguous
    torch::Tensor top_k_dev = top_k.is_cuda() ? top_k : top_k.to(probs.device(), /*non_blocking=*/true);
    torch::Tensor top_k_i32 = (top_k_dev.scalar_type() == at::kInt)
                                ? top_k_dev.contiguous()
                                : top_k_dev.to(at::kInt, /*non_blocking=*/true).contiguous();

    auto device = probs.device();
    auto opts_idx32 = torch::TensorOptions().dtype(at::kInt).device(device);
    auto opts_i64   = torch::TensorOptions().dtype(at::kLong).device(device);
    auto opts_f32   = torch::TensorOptions().dtype(at::kFloat).device(device);

    // Output samples
    torch::Tensor samples = torch::empty({batch_size}, opts_i64);

    // Allocate temporary buffers
    const int64_t total_items64 = batch_size * vocab_size;
    if (total_items64 > static_cast<int64_t>(std::numeric_limits<int>::max())) {
        throw std::invalid_argument("total_items exceeds int32 range required by CUB segmented sort");
    }
    const int total_items = static_cast<int>(total_items64);

    torch::Tensor d_begin = torch::empty({batch_size}, opts_idx32);
    torch::Tensor d_end   = torch::empty({batch_size}, opts_idx32);
    torch::Tensor d_indices_in = torch::empty({total_items64}, opts_idx32);
    torch::Tensor d_sorted_vals = torch::empty({total_items64}, opts_f32);
    torch::Tensor d_sorted_idx  = torch::empty({total_items64}, opts_idx32);

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

    // Fill segment offsets and iota indices
    launch_fill_segment_offsets(d_begin.data_ptr<int32_t>(), d_end.data_ptr<int32_t>(),
                                static_cast<int>(batch_size),
                                static_cast<int>(vocab_size),
                                stream.stream());
    CUDA_CHECK(cudaGetLastError());

    launch_iota_indices(d_indices_in.data_ptr<int32_t>(),
                        static_cast<int>(batch_size),
                        static_cast<int>(vocab_size),
                        stream.stream());
    CUDA_CHECK(cudaGetLastError());

    // Segmented sort per row (descending)
    launch_segmented_sort_desc_pairs(
        probs.data_ptr<float>(),
        d_sorted_vals.data_ptr<float>(),
        d_indices_in.data_ptr<int32_t>(),
        d_sorted_idx.data_ptr<int32_t>(),
        total_items,
        static_cast<int>(batch_size),
        d_begin.data_ptr<int32_t>(),
        d_end.data_ptr<int32_t>(),
        stream.stream());
    CUDA_CHECK(cudaGetLastError());

    // Seed handling
    uint64_t seed;
    if (seed_opt.has_value()) {
        seed = static_cast<uint64_t>(seed_opt.value());
    } else {
        auto now = static_cast<uint64_t>(std::chrono::high_resolution_clock::now().time_since_epoch().count());
        seed = now ^ (static_cast<uint64_t>(reinterpret_cast<uintptr_t>(stream.stream())) << 32);
    }

    // Launch sampling kernel
    launch_sampling_kernel(d_sorted_vals.data_ptr<float>(),
                           d_sorted_idx.data_ptr<int32_t>(),
                           top_k_i32.data_ptr<int32_t>(),
                           top_p_f32.data_ptr<float>(),
                           samples.data_ptr<int64_t>(),
                           static_cast<int>(batch_size),
                           seed,
                           stream.stream());
    CUDA_CHECK(cudaGetLastError());

    // Synchronize to ensure completion before returning
    CUDA_CHECK(cudaStreamSynchronize(stream.stream()));

    return samples;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run",
          &run,
          "top_k_top_p_sampling_from_probs_v129280 (B200-optimized segmented sort + block sampling)",
          py::arg("probs"),
          py::arg("top_k"),
          py::arg("top_p"),
          py::arg("seed") = c10::nullopt);
}
scrolls · 163 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON