gpt-5-2025-08-07 / cudaaec5f2
gpt-5-2025-08-07_cuda_aec5f2 · gpt-5-2025-08-07 · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 158 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-cuda-aec5f2?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:a9d02869ecc77d57789953be09df84ad744a38a0d0ed03d8b1f6da1ed3801b27
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20
Kernel source
main.cpp158 lines
#include "kernel.h"
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime_api.h>
#include <chrono>
#include <limits>
#include <pybind11/pybind11.h>
namespace {
inline void check_inputs(const torch::Tensor& probs,
const torch::Tensor& top_k,
const torch::Tensor& top_p) {
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(top_p.is_cuda(), "top_p must be a CUDA tensor");
TORCH_CHECK(probs.dtype() == torch::kFloat32, "probs must be float32");
TORCH_CHECK(top_k.dtype() == torch::kInt32, "top_k must be int32");
TORCH_CHECK(top_p.dtype() == torch::kFloat32, "top_p must be float32");
TORCH_CHECK(probs.dim() == 2, "probs must be 2D [batch_size, vocab_size]");
TORCH_CHECK(probs.size(1) == VOCAB_SIZE_V128256,
"vocab_size must be 128256, got ", probs.size(1));
TORCH_CHECK(top_k.dim() == 1 && top_k.size(0) == probs.size(0),
"top_k must be 1D with length equal to batch_size");
TORCH_CHECK(top_p.dim() == 1 && top_p.size(0) == probs.size(0),
"top_p must be 1D with length equal to batch_size");
TORCH_CHECK(probs.device().index() == top_k.device().index() &&
probs.device().index() == top_p.device().index(),
"All inputs must be on the same CUDA device");
}
} // namespace
torch::Tensor top_k_top_p_sampling_from_probs_v128256(torch::Tensor probs,
torch::Tensor top_k,
torch::Tensor top_p,
c10::optional<int64_t> seed_opt) {
check_inputs(probs, top_k, top_p);
at::cuda::CUDAGuard device_guard(probs.device());
probs = probs.contiguous();
top_k = top_k.contiguous();
top_p = top_p.contiguous();
const int batch_size = static_cast<int>(probs.size(0));
const int vocab_size = VOCAB_SIZE_V128256;
if (batch_size == 0) {
return torch::empty({0}, probs.options().dtype(torch::kInt64));
}
cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
const int64_t total_items_64 = static_cast<int64_t>(batch_size) * vocab_size;
TORCH_CHECK(total_items_64 <= static_cast<int64_t>(std::numeric_limits<int>::max()),
"Total items exceed supported range");
const int total_items = static_cast<int>(total_items_64);
// Allocate global sort buffers
auto keys_dtype = torch::kInt64; // store uint64_t keys in int64_t tensor
auto idx_dtype = torch::kInt32;
auto keys_in = torch::empty({total_items_64}, probs.options().dtype(keys_dtype));
auto keys_out = torch::empty({total_items_64}, probs.options().dtype(keys_dtype));
auto vals_in = torch::empty({total_items_64}, probs.options().dtype(idx_dtype));
auto vals_out = torch::empty({total_items_64}, probs.options().dtype(idx_dtype));
// Build composite keys and linear indices
build_composite_keys_and_linidx_launcher(
probs.data_ptr<float>(),
reinterpret_cast<uint64_t*>(keys_in.data_ptr<int64_t>()),
vals_in.data_ptr<int32_t>(),
total_items,
vocab_size,
stream);
// CUB global radix sort on (composite_key, linidx) pairs (ascending)
uint64_t* d_keys_in = reinterpret_cast<uint64_t*>(keys_in.data_ptr<int64_t>());
uint64_t* d_keys_out = reinterpret_cast<uint64_t*>(keys_out.data_ptr<int64_t>());
int32_t* d_vals_in = vals_in.data_ptr<int32_t>();
int32_t* d_vals_out = vals_out.data_ptr<int32_t>();
size_t temp_bytes = radix_sort_pairs_temp_bytes(
d_keys_in, d_keys_out, d_vals_in, d_vals_out, total_items, stream);
auto temp_storage = torch::empty({static_cast<int64_t>(temp_bytes)},
probs.options().dtype(torch::kUInt8));
void* temp_ptr = static_cast<void*>(temp_storage.data_ptr<uint8_t>());
radix_sort_pairs_launcher(
d_keys_in, d_keys_out, d_vals_in, d_vals_out,
total_items, temp_ptr, temp_bytes, stream);
// Gather sorted probs and token indices from linidx result
auto sorted_probs = torch::empty_like(probs); // [B, V]
auto sorted_indices = torch::empty({total_items_64}, probs.options().dtype(idx_dtype));
gather_sorted_from_linidx_launcher(
probs.data_ptr<float>(),
d_vals_out,
sorted_probs.data_ptr<float>(),
sorted_indices.data_ptr<int32_t>(),
total_items,
vocab_size,
stream);
// Output samples
auto samples = torch::empty({batch_size}, probs.options().dtype(torch::kInt64));
// Seed
uint64_t seed;
if (seed_opt.has_value() && seed_opt.value() != 0) {
seed = static_cast<uint64_t>(seed_opt.value());
} else {
seed = static_cast<uint64_t>(
std::chrono::high_resolution_clock::now().time_since_epoch().count());
}
// Launch sampling kernel: one block per row
sample_from_sorted_launcher(
sorted_probs.data_ptr<float>(),
sorted_indices.data_ptr<int32_t>(),
top_k.data_ptr<int32_t>(),
top_p.data_ptr<float>(),
samples.data_ptr<int64_t>(),
batch_size,
vocab_size,
seed,
stream);
return samples;
}
// Python binding entry point
torch::Tensor run(torch::Tensor probs,
torch::Tensor top_k,
torch::Tensor top_p,
c10::optional<int64_t> seed) {
return top_k_top_p_sampling_from_probs_v128256(std::move(probs),
std::move(top_k),
std::move(top_p),
seed);
}
namespace py = pybind11;
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Top-K + Top-P sampling from probability distributions (vocab=128256) optimized for B200";
m.def("run",
&run,
"Run top-k top-p sampling from probs (v128256)",
py::arg("probs"),
py::arg("top_k"),
py::arg("top_p"),
py::arg("seed") = py::none());
}scrolls · 158 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON