Skip to content
KernelIndex
Search⌘K

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