Skip to content
KernelIndex
Search⌘K

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