gpt-5-2025-08-07 / cudaa6f41d
gpt-5-2025-08-07_cuda_a6f41d · gpt-5-2025-08-07 · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 184 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-cuda-a6f41d?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:2ace59b4d9fc02cc0fbef68fddf32f451aeb3e3fa61f73d850a9a2b0b15bd18f
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20
Kernel source
main.cpp184 lines
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <vector>
#include <random>
#include <chrono>
#include <stdexcept>
#include <cstring>
#include <limits>
#include <cstdint>
#include "kernel.h"
using namespace at::indexing;
// Helper to check tensor properties
static inline void check_inputs(const torch::Tensor& probs, const torch::Tensor& top_p) {
TORCH_CHECK(probs.is_cuda(), "probs must be a CUDA tensor");
TORCH_CHECK(top_p.is_cuda() || top_p.is_cpu(), "top_p must be on CPU or CUDA");
TORCH_CHECK(probs.dim() == 2, "probs must be 2D [batch, vocab]");
TORCH_CHECK(probs.scalar_type() == torch::kFloat32, "probs must be float32");
TORCH_CHECK(probs.size(1) == VOCAB_SIZE_QWEN3, "vocab_size must be 151936");
TORCH_CHECK(top_p.dim() == 1 && top_p.size(0) == probs.size(0), "top_p must be [batch]");
TORCH_CHECK(top_p.scalar_type() == torch::kFloat32, "top_p must be float32");
}
// Build host vectors of row indices for each category based on top_p
static inline void partition_rows(const torch::Tensor& top_p, std::vector<int32_t>& rows_argmax,
std::vector<int32_t>& rows_full, std::vector<int32_t>& rows_nucleus) {
auto top_p_cpu = top_p.to(torch::kCPU, /*non_blocking=*/false);
const float* tp = top_p_cpu.data_ptr<float>();
int64_t B = top_p_cpu.size(0);
rows_argmax.reserve(B);
rows_full.reserve(B);
rows_nucleus.reserve(B);
for (int64_t i = 0; i < B; ++i) {
float p = tp[i];
if (!(p > 0.0f)) {
rows_argmax.push_back(static_cast<int32_t>(i));
} else if (p >= 1.0f) {
rows_full.push_back(static_cast<int32_t>(i));
} else {
rows_nucleus.push_back(static_cast<int32_t>(i));
}
}
}
// Copy host vector to device tensor (int32)
static inline torch::Tensor vec_to_device_i32(const std::vector<int32_t>& v, c10::Device device, cudaStream_t stream) {
auto opts = torch::TensorOptions().dtype(torch::kInt32).device(device);
torch::Tensor t = torch::empty({static_cast<long long>(v.size())}, opts);
if (!v.empty()) {
CUDA_CHECK(cudaMemcpyAsync(t.data_ptr<int32_t>(), v.data(), v.size() * sizeof(int32_t),
cudaMemcpyHostToDevice, stream));
}
return t;
}
// Entry point
torch::Tensor run(torch::Tensor probs, torch::Tensor top_p) {
check_inputs(probs, top_p);
cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
const int64_t B = probs.size(0);
const int64_t V = probs.size(1);
TORCH_CHECK(V == VOCAB_SIZE_QWEN3, "vocab size mismatch");
// Ensure contiguous
probs = probs.contiguous();
// Make top_p available on device persistently
torch::Tensor top_p_dev = top_p.is_cuda()
? top_p.contiguous()
: top_p.to(probs.device(), /*non_blocking=*/false).contiguous();
// Output on GPU
auto out = torch::empty({B}, probs.options().dtype(torch::kInt64));
// Partition rows
std::vector<int32_t> rows_argmax, rows_full, rows_nucleus;
partition_rows(top_p, rows_argmax, rows_full, rows_nucleus);
// Device copies of row lists
torch::Tensor d_rows_argmax = vec_to_device_i32(rows_argmax, probs.device(), stream);
torch::Tensor d_rows_full = vec_to_device_i32(rows_full, probs.device(), stream);
torch::Tensor d_rows_nucleus= vec_to_device_i32(rows_nucleus,probs.device(), stream);
const float* d_probs = probs.data_ptr<float>();
const float* d_top_p = top_p_dev.data_ptr<float>();
int64_t* d_out = out.data_ptr<int64_t>();
// Seed for RNG used by kernels (full sampling path)
unsigned long long seed = static_cast<unsigned long long>(
(std::random_device{}()) ^
(static_cast<unsigned long long>(std::chrono::high_resolution_clock::now().time_since_epoch().count())));
// 1) Argmax rows (top_p <= 0)
if (!rows_argmax.empty()) {
launch_argmax_kernel(d_probs,
d_rows_argmax.data_ptr<int32_t>(),
static_cast<int>(rows_argmax.size()),
static_cast<int>(V),
d_out,
stream);
}
// 2) Full multinomial rows (top_p >= 1)
if (!rows_full.empty()) {
launch_sample_full_kernel(d_probs,
d_rows_full.data_ptr<int32_t>(),
static_cast<int>(rows_full.size()),
static_cast<int>(V),
seed,
d_out,
stream);
}
// 3) Nucleus rows (0 < top_p < 1): Use exact PyTorch semantics on GPU to guarantee correctness
if (!rows_nucleus.empty()) {
// Gather per-row probabilities and top_p for nucleus rows
auto sel_rows = d_rows_nucleus; // [R] int32 on device
auto sel_probs = probs.index_select(0, sel_rows.to(torch::kLong)); // [R, V]
auto sel_top_p = top_p_dev.index_select(0, sel_rows.to(torch::kLong)).unsqueeze(1); // [R,1]
// Sort descending like reference
auto sort_tuple = sel_probs.sort(1, /*descending=*/true);
auto vals_sorted = std::get<0>(sort_tuple); // [R, V], descending
auto idx_sorted = std::get<1>(sort_tuple).to(torch::kInt32); // [R, V], original indices
// CDF and keep mask with "shift" semantics
auto cdf = vals_sorted.cumsum(1); // [R, V]
auto to_remove = cdf.gt(sel_top_p); // [R, V], bool
// Shift mask right by 1, set first column to False
auto to_remove_shifted = torch::zeros_like(to_remove, to_remove.options().dtype(torch::kBool));
// to_remove_shifted[:, 1:] = to_remove[:, :-1]
to_remove_shifted.index_put_({Slice(), Slice(1, None)}, to_remove.index({Slice(), Slice(None, -1)}));
// to_remove_shifted[:, 0] already False
auto keep = (~to_remove_shifted); // [R, V] bool
// Build filtered distribution in sorted space and renormalize
auto keep_f = keep.to(vals_sorted.scalar_type()); // float
auto probs_kept_sorted = vals_sorted * keep_f; // [R, V]
auto sums = probs_kept_sorted.sum(1, true); // [R, 1]
// Handle degenerate rows defensively: if sum==0, fall back to picking the top token
auto safe_sums = torch::where(sums > 0.0, sums, torch::ones_like(sums));
auto normalized_sorted = probs_kept_sorted / safe_sums;
// For rows with sum==0, normalized will be zeros; set first position to 1
// Find rows where sums==0 and fix their normalized distributions
auto zero_rows_mask = (sums <= 0.0).squeeze(1); // [R] bool
if (zero_rows_mask.any().item<bool>()) {
// Build an index tensor for those rows and set normalized_sorted[row, 0] = 1
auto zero_rows = zero_rows_mask.nonzero().squeeze(1); // [Rz]
if (zero_rows.numel() > 0) {
normalized_sorted.index_put_({zero_rows, 0}, 1.0f);
}
}
// Sample indices in sorted space using PyTorch multinomial (GPU)
auto selected_sorted = torch::multinomial(normalized_sorted, /*num_samples=*/1, /*replacement=*/true); // [R,1]
// Map back to original vocabulary indices
auto final_idx = idx_sorted.gather(1, selected_sorted).squeeze(1).to(torch::kInt64); // [R]
// Scatter into output at the appropriate batch rows
out.index_copy_(0, sel_rows.to(torch::kLong), final_idx);
}
CUDA_CHECK(cudaGetLastError());
CUDA_CHECK(cudaStreamSynchronize(stream));
// Return results to CPU
return out.cpu();
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("run", &run, "top_p_sampling_from_probs_v151936 (CUDA)");
}scrolls · 184 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON