Skip to content
KernelIndex
Search⌘K

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
<![CDATA[
/*
 *  PyTorch binding + CPU fall-back for extreme top-k values
 */

#include "kernel.h"

#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>

#include <algorithm>
#include <chrono>
#include <cstdlib>
#include <numeric>
#include <vector>

using namespace topk_topp_sampling;

/* ================================================================= *
 *                 plain C++ reference for CPU fall-back             *
 * ================================================================= */
static std::int64_t cpu_reference_one(const float* row,
                                      int          k,
                                      float        p)
{
  std::vector<float> probs(row, row + VOCAB_SIZE);

  /* ---------------- top-k filter -------------------------------- */
  if (k > 0 && k < VOCAB_SIZE)
  {
    std::vector<int> idx(VOCAB_SIZE);
    std::iota(idx.begin(), idx.end(), 0);
    std::partial_sort(idx.begin(), idx.begin() + k, idx.end(),
                      [&](int a, int b){ 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