Skip to content
KernelIndex
Search⌘K

gpt-o3_cuda_037673

gpt-o3 · cuda · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-cuda-037673?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:2d89bc84d3006866b496c07f0424f8cfd9b24d4db668efe754c2b8959c53a848
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20

Kernel source

main.cpp94 lines
#include "kernel.h"

#include <ATen/cuda/CUDAGuard.h>
#include <torch/extension.h>
#include <cmath>
#include <limits>

/* ========================================================================= */
/*  Internal helper – wraps CUDA launcher                                    */
/* ========================================================================= */
static void run_kernel(
        const torch::Tensor& q,
        const torch::Tensor& k_cache,
        const torch::Tensor& v_cache,
        const torch::Tensor& qo_indptr,
        const torch::Tensor& kv_indptr,
        const torch::Tensor& kv_indices,
        double               sm_scale,
        torch::Tensor&       output,
        torch::Tensor&       lse)
{
    /* Basic sanity checks -------------------------------------------------- */
    TORCH_CHECK(q.is_cuda()          && q.scalar_type() == torch::kBFloat16, "q must be bf16 CUDA");
    TORCH_CHECK(k_cache.is_cuda()    && k_cache.scalar_type() == torch::kBFloat16, "k_cache must be bf16 CUDA");
    TORCH_CHECK(v_cache.is_cuda()    && v_cache.scalar_type() == torch::kBFloat16, "v_cache must be bf16 CUDA");
    TORCH_CHECK(qo_indptr.is_cuda()  && qo_indptr.scalar_type() == torch::kInt,     "qo_indptr must be int32 CUDA");
    TORCH_CHECK(kv_indptr.is_cuda()  && kv_indptr.scalar_type() == torch::kInt,     "kv_indptr must be int32 CUDA");
    TORCH_CHECK(kv_indices.is_cuda() && kv_indices.scalar_type() == torch::kInt,    "kv_indices must be int32 CUDA");

    /* Ensure we execute on the right device / stream ---------------------- */
    at::cuda::CUDAGuard device_guard(q.device());
    at::cuda::CUDAStream cuda_stream = at::cuda::getCurrentCUDAStream();

    launch_gqa_paged_prefill_causal_h32_kv4_d128_ps1(
        q, k_cache, v_cache,
        qo_indptr, kv_indptr, kv_indices,
        static_cast<float>(sm_scale),
        output, lse,
        cuda_stream.stream());
}

/* ========================================================================= */
/*  Python binding                                                           */
/* ========================================================================= */
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run",
          [](torch::Tensor  q,
             torch::Tensor  k_cache,
             torch::Tensor  v_cache,
             torch::Tensor  qo_indptr,
             torch::Tensor  kv_indptr,
             torch::Tensor  kv_indices,
             double         sm_scale = 1.0 / std::sqrt(static_cast<double>(HEAD_DIM)))
          {
              const int64_t total_q = q.size(0);

              /* Allocate outputs ----------------------------------------- */
              auto opts_out = torch::TensorOptions()
                                  .dtype(torch::kBFloat16)
                                  .device(q.device());
              auto opts_lse = torch::TensorOptions()
                                  .dtype(torch::kFloat32)
                                  .device(q.device());

              torch::Tensor output = torch::zeros({total_q,
                                                   NUM_QO_HEADS,
                                                   HEAD_DIM}, opts_out);

              torch::Tensor lse    = torch::full({total_q,
                                                  NUM_QO_HEADS},
                                                  -std::numeric_limits<float>::infinity(),
                                                  opts_lse);

              /* Invoke CUDA kernel --------------------------------------- */
              run_kernel(q, k_cache, v_cache,
                         qo_indptr, kv_indptr, kv_indices,
                         sm_scale,
                         output, lse);

              return std::make_tuple(output, lse);
          },
          pybind11::arg("q"),
          pybind11::arg("k_cache"),
          pybind11::arg("v_cache"),
          pybind11::arg("qo_indptr"),
          pybind11::arg("kv_indptr"),
          pybind11::arg("kv_indices"),
          pybind11::arg("sm_scale") = 1.0 / std::sqrt(static_cast<double>(HEAD_DIM)),
          R"doc(
B200-optimised GQA paged-prefill causal-attention kernel  
Hard-wired for 32 Qo-heads, 4 KV-heads, d = 128, page = 1.  
Returns  (output, lse)
)doc");
}
scrolls · 94 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON