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