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-memory
size_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