gpt-o3_cuda_717406
gpt-o3 · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 110 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-cuda-717406?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:6adab27155212c0902b73d8d4b73e57529069e7ccba8346df2e2d518e69e053c
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20
Kernel source
main.cpp110 lines
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include "kernel.h"
/* ------------------------------------------------------------------------- */
/* Per-row nucleus (top-p) sampling */
/* ------------------------------------------------------------------------- */
static torch::Tensor top_p_row_sampling(const torch::Tensor& row_in,
float p_thresh)
{
/* `row_in` is 1-D float32 CUDA tensor of length VOCAB_SIZE_TP */
auto row = row_in.contiguous();
/* ---- degenerate cases ------------------------------------------------ */
if (p_thresh <= 0.0f) {
/* fall back to argmax (greedy) */
return row.argmax(0);
}
if (p_thresh >= 1.0f - 1e-8f) {
/* standard multinomial sampling */
return torch::multinomial(row, /*num_samples=*/1, /*replacement=*/true)
.squeeze(0);
}
/* ---- nucleus filtering ---------------------------------------------- */
auto sort_pair = torch::sort(row, /*dim=*/0, /*descending=*/true);
auto sorted_vals = std::get<0>(sort_pair); // probabilities, high→low
auto sorted_indices= std::get<1>(sort_pair); // original vocab indices
auto cdf = torch::cumsum(sorted_vals, 0);
auto to_remove = cdf > p_thresh; // bool mask CUDA
/* shift mask so we keep the first token that crosses the threshold */
if (to_remove.size(0) > 1) {
auto shifted = torch::cat(
{torch::zeros({1}, to_remove.options()),
to_remove.slice(/*dim=*/0, /*start=*/0,
/*end=*/to_remove.size(0) - 1)}, 0);
to_remove = shifted;
}
to_remove.index_put_({0}, false); // always keep top-1 token
auto keep_mask = torch::logical_not(to_remove);
auto keep_idx = sorted_indices.masked_select(keep_mask);
/* rebuild filtered distribution in original vocabulary order */
auto filtered = torch::zeros_like(row);
if (keep_idx.numel() == 0) {
/* numerical corner-case – should be almost impossible */
return row.argmax(0);
}
filtered.index_put_({keep_idx}, row.index_select(0, keep_idx));
const float norm = filtered.sum().item<float>();
if (norm <= 0.0f) { // extra safety
return row.argmax(0);
}
filtered /= norm;
return torch::multinomial(filtered, 1, true).squeeze(0);
}
/* ------------------------------------------------------------------------- */
/* Entry point */
/* ------------------------------------------------------------------------- */
torch::Tensor run(torch::Tensor probs, // [B, V] float32 CUDA
torch::Tensor top_p) // [B] float32 CUDA
{
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.scalar_type() == torch::kFloat32,
"`probs` must be float32");
TORCH_CHECK(top_p.scalar_type() == torch::kFloat32,
"`top_p` must be float32");
TORCH_CHECK(probs.dim() == 2,
"`probs` must have shape [batch, vocab]");
TORCH_CHECK(probs.size(1) == VOCAB_SIZE_TP,
"vocab size mismatch – expected ", VOCAB_SIZE_TP);
TORCH_CHECK(top_p.dim() == 1 && top_p.size(0) == probs.size(0),
"`top_p` must be 1-D and match batch size");
const int64_t B = probs.size(0);
auto samples = torch::empty(
{B},
probs.options().dtype(torch::kInt64)); // [B] int64 CUDA
/* bring `top_p` to host once – avoids B small GPU→CPU transfers */
auto top_p_host = top_p.to(torch::kCPU);
const float *p_host = top_p_host.data_ptr<float>();
for (int64_t i = 0; i < B; ++i) {
const float p_threshold = p_host[i];
auto row = probs[i]; // view [V]
auto id = top_p_row_sampling(row, p_threshold);
samples[i] = id.item<int64_t>();
}
return samples;
}
/* ------------------------------------------------------------------------- */
/* Python binding */
/* ------------------------------------------------------------------------- */
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
{
m.def("run",
&run,
"top_p_sampling_from_probs_v151936 (CUDA, ATen implementation)");
}scrolls · 110 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON