gpt-o3 / cudad377ec
gpt-o3_cuda_d377ec · gpt-o3 · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 133 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-cuda-d377ec?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:e9ddc3067d3170de984ead5f883265183c97172fc8a1b322b5b4d5f24f9bed81
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20
Kernel source
main.cpp133 lines
#include "kernel.h"
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <thrust/device_ptr.h>
#include <thrust/sequence.h>
#include <thrust/sort.h>
#include <thrust/functional.h>
#include <chrono>
/* -------------------------------------------------------------------------- */
/* Helper : make sure tensor is on CUDA, correct dtype & contiguous */
/* -------------------------------------------------------------------------- */
static inline torch::Tensor ensure(torch::Tensor t, torch::ScalarType dt)
{
if (t.scalar_type() != dt) t = t.to(dt);
if (!t.is_cuda()) t = t.cuda();
if (!t.is_contiguous()) t = t.contiguous();
return t;
}
/* -------------------------------------------------------------------------- */
/* Device-side free (async when available) */
/* -------------------------------------------------------------------------- */
static void device_free(void* ptr, cudaStream_t stream)
{
#if defined(CUDART_VERSION) && CUDART_VERSION >= 11020
CUDA_CHECK(cudaFreeAsync(ptr, stream));
#else
cudaStreamSynchronize(stream);
CUDA_CHECK(cudaFree(ptr));
#endif
}
/* -------------------------------------------------------------------------- */
/* Python entry point */
/* -------------------------------------------------------------------------- */
torch::Tensor run(torch::Tensor probs,
torch::Tensor top_k,
torch::Tensor top_p)
{
TORCH_CHECK(probs.dim() == 2,
"probs must be [batch_size , vocab_size]");
TORCH_CHECK(probs.size(1) == VOCAB_SIZE,
"vocab_size must be ", VOCAB_SIZE);
const int64_t B = probs.size(0); /* batch size */
auto stream = c10::cuda::getCurrentCUDAStream();
/* Ensure correct dtype / device -------------------------------------- */
probs = ensure(probs, torch::kFloat32);
top_k = ensure(top_k, torch::kInt32);
top_p = ensure(top_p, torch::kFloat32);
float* p_ptr = probs.data_ptr<float>();
int32_t* k_ptr = top_k.data_ptr<int32_t>();
float* tp_ptr = top_p.data_ptr<float>();
/* -------------------------------------------------------------------- */
/* Build per-row descending sort (keys = probs , values = indices) */
/* -------------------------------------------------------------------- */
auto indices = torch::empty({B, VOCAB_SIZE},
probs.options().dtype(torch::kInt32));
int32_t* idx_ptr = indices.data_ptr<int32_t>();
auto exec = thrust::cuda::par.on(stream);
for (int64_t b = 0; b < B; ++b) {
float* row_probs = p_ptr + b * VOCAB_SIZE;
int32_t* row_idx = idx_ptr + b * VOCAB_SIZE;
thrust::device_ptr<float> d_probs(row_probs);
thrust::device_ptr<int32_t> d_idx (row_idx);
thrust::sequence(exec, d_idx, d_idx + VOCAB_SIZE, 0);
thrust::sort_by_key(exec,
d_probs,
d_probs + VOCAB_SIZE,
d_idx,
thrust::greater<float>());
}
/* -------------------------------------------------------------------- */
/* RNG states */
/* -------------------------------------------------------------------- */
curandStatePhilox4_32_10_t* d_states = nullptr;
CUDA_CHECK(cudaMalloc(&d_states, sizeof(*d_states) * B));
const unsigned long long seed =
static_cast<unsigned long long>(
std::chrono::high_resolution_clock::now()
.time_since_epoch().count());
initialize_random_states(d_states,
static_cast<int>(B),
seed,
stream);
/* -------------------------------------------------------------------- */
/* Output tensor */
/* -------------------------------------------------------------------- */
auto samples = torch::empty({B},
probs.options().dtype(torch::kInt64));
int64_t* s_ptr = samples.data_ptr<int64_t>();
/* -------------------------------------------------------------------- */
/* Launch sampling kernel */
/* -------------------------------------------------------------------- */
topk_topp_sample_kernel_launcher(p_ptr,
idx_ptr,
k_ptr,
tp_ptr,
s_ptr,
d_states,
static_cast<int>(B),
stream);
/* Free RNG states ----------------------------------------------------- */
device_free(d_states, stream);
return samples;
}
/* -------------------------------------------------------------------------- */
/* PyBind11 module */
/* -------------------------------------------------------------------------- */
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
{
m.def("run", &run,
"top_k_top_p_sampling_from_probs_v129280 (CUDA, B200-optimised)");
}scrolls · 133 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON