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