gpt-o3 / cudaf2ff2b
gpt-o3_cuda_f2ff2b · gpt-o3 · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 168 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-cuda-f2ff2b?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:43f3166b9276ab81e89339664a41364344833f2e317eab8025da91e143ca025f
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20
Kernel source
main.cpp168 lines
<{ return probs[a] > probs[b]; });
std::vector<float> tmp(VOCAB_SIZE, 0.f);
float sum_k = 0.f;
for (int i = 0; i < k; ++i)
{
int id = idx[i];
tmp[id] = probs[id];
sum_k += probs[id];
}
for (int i = 0; i < VOCAB_SIZE; ++i) tmp[i] /= sum_k;
probs.swap(tmp);
}
/* ---------------- greedy shortcut (p ≤ 0) --------------------- */
if (p <= 0.f)
{
return std::max_element(probs.begin(), probs.end()) - probs.begin();
}
/* ---------------- top-p (nucleus) filter ---------------------- */
if (p < 1.f)
{
std::vector<int> idx(VOCAB_SIZE);
std::iota(idx.begin(), idx.end(), 0);
std::sort(idx.begin(), idx.end(),
[&](int a, int b){ return probs[a] > probs[b]; });
std::vector<char> keep(VOCAB_SIZE, 0);
float cdf = 0.f;
for (int i = 0; i < VOCAB_SIZE; ++i)
{
cdf += probs[idx[i]];
keep[idx[i]] = 1;
if (cdf > p) break;
}
float norm = 0.f;
for (int i = 0; i < VOCAB_SIZE; ++i)
if (!keep[i]) probs[i] = 0.f;
else norm += probs[i];
for (float& v : probs) v /= norm;
}
/* ---------------- multinomial sample -------------------------- */
float r = static_cast<float>(std::rand()) / (RAND_MAX + 1.f); /* [0,1) */
float acc = 0.f;
int pick = VOCAB_SIZE - 1;
for (int i = 0; i < VOCAB_SIZE; ++i)
{
acc += probs[i];
if (r <= acc) { pick = i; break; }
}
return pick;
}
/* ================================================================= *
* 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 2-D (B, V)");
TORCH_CHECK(probs.scalar_type() == torch::kFloat32, "probs must be float32");
TORCH_CHECK(top_k.scalar_type() == torch::kInt32, "top_k must be int32");
TORCH_CHECK(top_p.scalar_type() == torch::kFloat32, "top_p must be float32");
const int64_t B = probs.size(0);
const int64_t V = probs.size(1);
TORCH_CHECK(V == VOCAB_SIZE,
"vocab dimension must be ", VOCAB_SIZE);
TORCH_CHECK(probs.is_cuda() && top_k.is_cuda() && top_p.is_cuda(),
"all inputs must reside on the same CUDA device");
/* make contiguous for reliable pointer arithmetic */
probs = probs.contiguous();
top_k = top_k.contiguous();
top_p = top_p.contiguous();
auto samples = torch::empty({B},
probs.options().dtype(torch::kInt64));
/* --------------- launch fast CUDA path ------------------------ */
cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
const std::uint64_t seed =
static_cast<std::uint64_t>(
std::chrono::high_resolution_clock::now()
.time_since_epoch().count());
launch_sampling_kernel(probs.data_ptr<float>(),
top_k.data_ptr<int>(),
top_p.data_ptr<float>(),
samples.data_ptr<std::int64_t>(),
static_cast<int>(B),
seed,
stream);
/* --------------- CPU fall-back (k ≤ 0 or k > 512) ------------- */
auto top_k_cpu = top_k.cpu();
auto top_p_cpu = top_p.cpu();
/* make sure GPU work is finished before we overwrite */
cudaStreamSynchronize(stream);
torch::Tensor probs_cpu; /* materialised lazily */
for (int64_t i = 0; i < B; ++i)
{
int k_val = top_k_cpu[i].item<int>();
float p_val = top_p_cpu[i].item<float>();
if (k_val > 0 && k_val <= FAST_TOP_K_MAX) continue; /* already done */
if (!probs_cpu.defined()) probs_cpu = probs.cpu();
std::int64_t token =
cpu_reference_one(probs_cpu[i].data_ptr<float>(), k_val, p_val);
samples.index_put_({i}, token);
}
return samples;
}
/* ================================================================= *
* pybind11 glue *
* ================================================================= */
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
{
m.def("run", &run,
"top_k_top_p_sampling_from_probs_v151936 (CUDA kernel + CPU fall-back)");
}
]]>scrolls · 168 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON