gpt-5-2025-08-07 / cuda5fc7e3
gpt-5-2025-08-07_cuda_5fc7e3 · gpt-5-2025-08-07 · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 137 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-cuda-5fc7e3?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:ffacdb9d70668fc151cce7587ed7470f231c9b477d5d35a97c0867ecf74a31e0
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20
Kernel source
main.cpp137 lines
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cuda_runtime.h>
#include <vector>
#include <chrono>
#include <stdexcept>
#include "kernel.h"
namespace py = pybind11;
static inline void cuda_check_last(const char* msg) {
cudaError_t err = cudaGetLastError();
TORCH_CHECK(err == cudaSuccess, "CUDA error after ", msg, ": ", cudaGetErrorString(err));
}
static inline void cuda_check(cudaError_t err, const char* msg) {
TORCH_CHECK(err == cudaSuccess, "CUDA error in ", msg, ": ", cudaGetErrorString(err));
}
torch::Tensor run(torch::Tensor probs, torch::Tensor top_p) {
TORCH_CHECK(probs.is_cuda(), "probs 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_p.dtype() == torch::kFloat32, "top_p must be float32");
TORCH_CHECK(probs.dim() == 2, "probs must be 2D [batch, vocab]");
TORCH_CHECK(top_p.dim() == 1, "top_p must be 1D [batch]");
const int64_t B = probs.size(0);
const int64_t V = probs.size(1);
TORCH_CHECK(V == (int64_t)V_CONST, "vocab_size must be 128256");
TORCH_CHECK(top_p.size(0) == B, "top_p shape mismatch with batch size");
c10::cuda::CUDAGuard device_guard(probs.get_device());
probs = probs.contiguous();
top_p = top_p.contiguous();
auto options = torch::TensorOptions().dtype(torch::kInt64).device(probs.device());
auto samples = torch::empty({B}, options);
auto cu_stream = at::cuda::getCurrentCUDAStream();
// Compute status per row on device: 0 (argmax), 1 (top-p), 2 (full)
auto status = torch::empty({B}, probs.options().dtype(torch::kInt32));
compute_row_status_launcher(top_p.data_ptr<float>(), status.data_ptr<int32_t>(), B, cu_stream.stream());
cuda_check_last("compute_row_status");
cuda_check(cudaStreamSynchronize(cu_stream.stream()), "stream sync after compute_row_status");
// Copy status and top_p to host to build list of rows needing top-p
auto status_h = status.cpu();
auto top_p_h = top_p.cpu();
const int32_t* status_hp = status_h.data_ptr<int32_t>();
const float* top_p_hp = top_p_h.data_ptr<float>();
std::vector<int64_t> rows_top_p;
rows_top_p.reserve(static_cast<size_t>(B));
for (int64_t i = 0; i < B; ++i) {
if (status_hp[i] == 1) rows_top_p.push_back(i);
}
// Step 1: p <= 0: argmax
argmax_kernel_launcher(
probs.data_ptr<float>(),
status.data_ptr<int32_t>(),
B, V,
samples.data_ptr<int64_t>(),
cu_stream.stream());
cuda_check_last("argmax_kernel");
// Step 2: p >= 1: sample from full distribution
uint64_t base_seed =
static_cast<uint64_t>(
std::chrono::high_resolution_clock::now().time_since_epoch().count());
sample_full_kernel_launcher(
probs.data_ptr<float>(),
status.data_ptr<int32_t>(),
B, V,
base_seed,
samples.data_ptr<int64_t>(),
cu_stream.stream());
cuda_check_last("sample_full_kernel");
// Step 3: 0 < p < 1: top-p via per-row CUB radix sort + sampling
if (!rows_top_p.empty()) {
float* d_sorted_keys = nullptr;
int32_t* d_sorted_idx = nullptr;
int32_t* d_base_idx = nullptr;
void* d_cub_tmp = nullptr;
cuda_check(cudaMalloc(&d_sorted_keys, sizeof(float) * (size_t)V), "cudaMalloc d_sorted_keys");
cuda_check(cudaMalloc(&d_sorted_idx, sizeof(int32_t) * (size_t)V), "cudaMalloc d_sorted_idx");
cuda_check(cudaMalloc(&d_base_idx, sizeof(int32_t) * (size_t)V), "cudaMalloc d_base_idx");
arange_launcher(d_base_idx, V, cu_stream.stream());
cuda_check_last("arange_launcher");
size_t tmp_bytes = cub_sort_pairs_temp_bytes(V, cu_stream.stream());
cuda_check(cudaMalloc(&d_cub_tmp, tmp_bytes), "cudaMalloc d_cub_tmp");
const float* probs_ptr = probs.data_ptr<float>();
int64_t* out_ptr = samples.data_ptr<int64_t>();
for (size_t k = 0; k < rows_top_p.size(); ++k) {
int64_t row = rows_top_p[k];
const float* row_ptr = probs_ptr + row * V;
// Sort this row descending
sort_row_desc(row_ptr, d_sorted_keys, d_base_idx, d_sorted_idx, V, d_cub_tmp, tmp_bytes, cu_stream.stream());
cuda_check_last("cub::SortPairsDescending");
// Sample from top-p subset
float pval = top_p_hp[row];
uint64_t seed = base_seed + static_cast<uint64_t>(row * 1315423911ULL);
top_p_sample_row_kernel_launcher(
d_sorted_keys, d_sorted_idx,
pval, V, seed,
out_ptr + row,
cu_stream.stream());
cuda_check_last("top_p_sample_row_kernel");
}
cuda_check(cudaFree(d_cub_tmp), "cudaFree d_cub_tmp");
cuda_check(cudaFree(d_base_idx), "cudaFree d_base_idx");
cuda_check(cudaFree(d_sorted_idx), "cudaFree d_sorted_idx");
cuda_check(cudaFree(d_sorted_keys), "cudaFree d_sorted_keys");
}
cuda_check_last("final");
return samples;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("run", &run,
"top_p_sampling_from_probs_v128256 (B200-optimized)",
py::arg("probs"),
py::arg("top_p"));
}scrolls · 137 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON