Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / cuda1e905d

gemini-2.5-pro_cuda_1e905d · gemini-2.5-pro · cuda · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-cuda-1e905d?include=source"
interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32

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:44116acecc0baf731927c348e4650ce74c37a8c55ae8491e42d4e5a4d29852e5
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

shared-memorysize_t smem_size = sizeof(TempStorage);

Kernel source

main.cpp139 lines
#include "kernel.h"

#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cub/device/device_segmented_radix_sort.cuh>
#include <torch/extension.h>

#include <stdexcept>
#include <vector>
#include <ctime>

// Error handling macros
#define CUDA_CHECK(expr)                                                       \
  do {                                                                         \
    cudaError_t status = (expr);                                               \
    if (status != cudaSuccess) {                                               \
      throw std::runtime_error("CUDA error in " + std::string(__FILE__) + ":" + \
                               std::to_string(__LINE__) + " - " +              \
                               std::string(cudaGetErrorString(status)));       \
    }                                                                          \
  } while (0)

#define CHECK_CUDA(x) TORCH_CHECK(x.is_cuda(), #x " must be a CUDA tensor")
#define CHECK_CONTIGUOUS(x) \
  TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
#define CHECK_INPUT(x) \
  CHECK_CUDA(x);       \
  CHECK_CONTIGUOUS(x)

torch::Tensor top_p_sampling_from_probs_v129280(torch::Tensor probs,
                                                torch::Tensor top_p) {
    // --- Input Validation ---
    CHECK_INPUT(probs);
    CHECK_INPUT(top_p);
    TORCH_CHECK(probs.dim() == 2, "probs must be a 2D tensor");
    TORCH_CHECK(probs.size(1) == VOCAB_SIZE, "probs must have vocab_size of ", VOCAB_SIZE);
    TORCH_CHECK(probs.scalar_type() == torch::kFloat32, "probs must be a float32 tensor");
    TORCH_CHECK(top_p.dim() == 1, "top_p must be a 1D tensor");
    TORCH_CHECK(top_p.size(0) == probs.size(0), "top_p must have the same batch size as probs");
    TORCH_CHECK(top_p.scalar_type() == torch::kFloat32, "top_p must be a float32 tensor");

    const int batch_size = probs.size(0);
    if (batch_size == 0) {
        return torch::empty({0}, torch::dtype(torch::kInt64).device(probs.device()));
    }

    const at::cuda::CUDAGuard device_guard(probs.device());
    cudaStream_t stream = at::cuda::getCurrentCUDAStream();

    // --- Output Allocation ---
    auto samples = torch::empty({batch_size}, torch::dtype(torch::kInt64).device(probs.device()));

    // --- Temporary Storage Allocation ---
    // For cuRAND states
    auto d_curand_states = torch::empty(
        {(long long)batch_size * sizeof(curandState_t)},
        torch::dtype(torch::kUInt8).device(probs.device()));

    // For CUB segmented sort inputs and outputs
    auto d_keys_in = torch::empty_like(probs);
    auto d_values_in = torch::empty({(long long)batch_size, VOCAB_SIZE}, torch::dtype(torch::kInt32).device(probs.device()));
    auto d_keys_out = torch::empty_like(probs);
    auto d_values_out = torch::empty_like(d_values_in);

    // CUB requires an array of size `num_segments + 1` for segment offsets
    auto d_segment_offsets = torch::arange(0, (long long)batch_size * VOCAB_SIZE + 1, VOCAB_SIZE, torch::dtype(torch::kInt32).device(probs.device()));

    // Pointers for CUB API
    float* d_keys_in_ptr = d_keys_in.data_ptr<float>();
    int* d_values_in_ptr = d_values_in.data_ptr<int>();
    float* d_keys_out_ptr = d_keys_out.data_ptr<float>();
    int* d_values_out_ptr = d_values_out.data_ptr<int>();
    int* d_offsets_ptr = d_segment_offsets.data_ptr<int>();

    // Determine temporary storage size for CUB sort
    void* d_temp_storage = nullptr;
    size_t temp_storage_bytes = 0;
    
    // The CUB API with begin/end iterators for offsets requires the distance
    // between iterators to equal `num_segments + 1`. This was the source of the
    // COMPILE_ERROR.
    cub::DeviceSegmentedRadixSort::SortPairsDescending(
        d_temp_storage, temp_storage_bytes,
        d_keys_in_ptr, d_keys_out_ptr,
        d_values_in_ptr, d_values_out_ptr,
        (long long)batch_size * VOCAB_SIZE, batch_size,
        d_offsets_ptr, d_offsets_ptr + batch_size + 1, // Corrected end iterator
        0, 8 * sizeof(float), stream);

    auto d_temp_storage_tensor = torch::empty({(long)temp_storage_bytes}, torch::dtype(torch::kUInt8).device(probs.device()));
    d_temp_storage = d_temp_storage_tensor.data_ptr();

    // --- Kernel Launches ---

    // 1. Setup cuRAND states
    const int curand_threads = 256;
    const int curand_blocks = (batch_size + curand_threads - 1) / curand_threads;
    setup_curand_kernel<<<curand_blocks, curand_threads, 0, stream>>>(
        (curandState_t*)d_curand_states.data_ptr(),
        (unsigned long long)time(nullptr) + (unsigned long long)probs.data_ptr(),
        batch_size);
    CUDA_CHECK(cudaGetLastError());

    // 2. Prepare key-value pairs (prob, index) for sorting
    prepare_sort_kernel<<<batch_size, BLOCK_THREADS, 0, stream>>>(
        probs.data_ptr<float>(),
        d_keys_in_ptr,
        d_values_in_ptr,
        batch_size);
    CUDA_CHECK(cudaGetLastError());

    // 3. Perform sorting with CUB (actual call)
    cub::DeviceSegmentedRadixSort::SortPairsDescending(
        d_temp_storage, temp_storage_bytes,
        d_keys_in_ptr, d_keys_out_ptr,
        d_values_in_ptr, d_values_out_ptr,
        (long long)batch_size * VOCAB_SIZE, batch_size,
        d_offsets_ptr, d_offsets_ptr + batch_size + 1, // Corrected end iterator
        0, 8 * sizeof(float), stream);
    CUDA_CHECK(cudaGetLastError());

    // 4. Filter, renormalize, and sample from the sorted distributions
    size_t smem_size = sizeof(TempStorage);
    filter_and_sample_kernel<<<batch_size, BLOCK_THREADS, smem_size, stream>>>(
        d_keys_out.data_ptr<float>(),
        d_values_out.data_ptr<int>(),
        top_p.data_ptr<float>(),
        samples.data_ptr<int64_t>(),
        (curandState_t*)d_curand_states.data_ptr(),
        batch_size);
    CUDA_CHECK(cudaGetLastError());

    return samples;
}

// Pybind11 module definition
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run", &top_p_sampling_from_probs_v129280, "Top-P Sampling from Probabilities (CUDA)");
}
scrolls · 139 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON