Skip to content
KernelIndex
Search⌘K

gpt-o3 / cuda7cdf4b

gpt-o3_cuda_7cdf4b · gpt-o3 · cuda · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 111 lines, Apache-2.0, pinned at da91508.

main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-cuda-7cdf4b?include=source"
interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, 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:42af63ba39ee3450bd073933aa1e45be0d422879524cc565e82c7c71998d7e6f
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20

Kernel source

main.cpp111 lines
#include "kernel.h"

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

/* ------------------------------------------------------------------ */
/*                    basic tensor sanity checks                       */
/* ------------------------------------------------------------------ */
static void checkTensor(const torch::Tensor& t,
                        torch::ScalarType   expected_dtype,
                        const char*         name)
{
    TORCH_CHECK(t.is_cuda(),       name, " must reside on CUDA.");
    TORCH_CHECK(t.is_contiguous(), name, " must be contiguous.");
    TORCH_CHECK(t.scalar_type() == expected_dtype,
                name, " must have dtype ", expected_dtype, ".");
}

/* ------------------------------------------------------------------ */
/*                python-visible entry point (“run”)                   */
/* ------------------------------------------------------------------ */
torch::Dict<std::string, torch::Tensor>
run(torch::Tensor q_nope,
    torch::Tensor q_pe,
    torch::Tensor ckv_cache,
    torch::Tensor kpe_cache,
    torch::Tensor kv_indptr,
    torch::Tensor kv_indices,
    double        sm_scale_double)
{
    /* ---------------  sanity ------------------------------------- */
    checkTensor(q_nope,    torch::kBFloat16, "q_nope");
    checkTensor(q_pe,      torch::kBFloat16, "q_pe");
    checkTensor(ckv_cache, torch::kBFloat16, "ckv_cache");
    checkTensor(kpe_cache, torch::kBFloat16, "kpe_cache");

    TORCH_CHECK(kv_indptr.scalar_type()  == torch::kInt32,
                "kv_indptr must be int32.");
    TORCH_CHECK(kv_indices.scalar_type() == torch::kInt32,
                "kv_indices must be int32.");

    const int64_t B = q_nope.size(0);
    TORCH_CHECK(q_nope.sizes() == torch::IntArrayRef({B, 16, 512}),
                "q_nope shape must be [B,16,512].");
    TORCH_CHECK(q_pe.sizes()   == torch::IntArrayRef({B, 16,  64}),
                "q_pe shape must be [B,16,64].");

    /* ---------------  allocate outputs --------------------------- */
    const auto opts_bf16 = q_nope.options().dtype(torch::kBFloat16);
    const auto opts_f32  = q_nope.options().dtype(torch::kFloat32);

    torch::Tensor output = torch::empty({B, 16, 512}, opts_bf16);
    torch::Tensor lse    = torch::empty({B, 16},       opts_f32);

    /* ---------------  raw pointers ------------------------------- */
    const __nv_bfloat16* qn_ptr = reinterpret_cast<const __nv_bfloat16*>(
                                    q_nope.data_ptr<at::BFloat16>());
    const __nv_bfloat16* qp_ptr = reinterpret_cast<const __nv_bfloat16*>(
                                    q_pe.data_ptr<at::BFloat16>());
    const __nv_bfloat16* kc_ptr = reinterpret_cast<const __nv_bfloat16*>(
                                    ckv_cache.data_ptr<at::BFloat16>());
    const __nv_bfloat16* kp_ptr = reinterpret_cast<const __nv_bfloat16*>(
                                    kpe_cache.data_ptr<at::BFloat16>());

    const int32_t* kv_indptr_ptr  = kv_indptr.data_ptr<int32_t>();
    const int32_t* kv_indices_ptr = kv_indices.data_ptr<int32_t>();

    __nv_bfloat16* out_ptr = reinterpret_cast<__nv_bfloat16*>(
                                output.data_ptr<at::BFloat16>());
    float* lse_ptr = lse.data_ptr<float>();

    /* ---------------  launch kernel ------------------------------ */
    cudaStream_t stream = at::cuda::getCurrentCUDAStream();

    launch_mla_paged_decode_h16_ckv512_kpe64_ps1(
        qn_ptr, qp_ptr,
        kc_ptr, kp_ptr,
        kv_indptr_ptr, kv_indices_ptr,
        static_cast<int>(B),
        static_cast<float>(sm_scale_double),
        out_ptr, lse_ptr,
        stream);

    /* ---------------  sync & return ------------------------------ */
    CUDA_CHECK(cudaStreamSynchronize(stream));

    return {{"output", output},
            {"lse",    lse   }};
}

/* ------------------------------------------------------------------ */
/*          PyBind11 module definition (torch-extension)              */
/* ------------------------------------------------------------------ */
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run", &run,
          R"pbdoc(
              Paged-KV MLA single-query decode
              (num_qo_heads = 16, head_dim_ckv = 512,
               head_dim_kpe = 64, page_size = 1)
          )pbdoc",
          py::arg("q_nope"),
          py::arg("q_pe"),
          py::arg("ckv_cache"),
          py::arg("kpe_cache"),
          py::arg("kv_indptr"),
          py::arg("kv_indices"),
          py::arg("sm_scale"));
}
scrolls · 111 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON