Skip to content
KernelIndex
Search⌘K

gpt-o3 / cuda1dc3e4

gpt-o3_cuda_1dc3e4 · gpt-o3 · cuda · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-cuda-1dc3e4?include=source"
interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32

Benchmark evidence

36 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
189.1µs
#5 of 5
2025-10-16
NVIDIA B200
233.7µs
#5 of 5
2025-10-16
NVIDIA B200
359.6µs
#5 of 5
2025-10-16
NVIDIA B200
361.0µs
#5 of 5
2025-10-16
NVIDIA B200
404.6µs
#5 of 5
2025-10-16
NVIDIA B200
508.2µs
#5 of 5
2025-10-16
NVIDIA B200
533.1µs
#5 of 5
2025-10-16
NVIDIA B200
724.1µs
#5 of 5
2025-10-16
NVIDIA B200
838.7µs
#5 of 5
2025-10-16
NVIDIA B200
845.6µs
#5 of 5
2025-10-16
Show all 36 measurements ›
NVIDIA B200
845.7µs
#5 of 5
2025-10-16
NVIDIA B200
847.0µs
#5 of 5
2025-10-16
NVIDIA B200
1.41ms
#5 of 5
2025-10-16
NVIDIA B200
1.42ms
#5 of 5
2025-10-16
NVIDIA B200
1.85ms
#5 of 5
2025-10-16
NVIDIA B200
2.18ms
#5 of 5
2025-10-16
NVIDIA B200
2.42ms
#5 of 5
2025-10-16
NVIDIA B200
2.56ms
#5 of 5
2025-10-16
NVIDIA B200
5.03ms
#5 of 5
2025-10-16
NVIDIA B200
6.24ms
#5 of 5
2025-10-16
NVIDIA B200
6.54ms
#5 of 5
2025-10-16
NVIDIA B200
13.7ms
#5 of 5
2025-10-16
NVIDIA B200
24.6ms
#5 of 5
2025-10-16
NVIDIA B200
38.7ms
#5 of 5
2025-10-16
NVIDIA B200
39.1ms
#5 of 5
2025-10-16
NVIDIA B200
48.3ms
#5 of 5
2025-10-16
NVIDIA B200
128.5ms
#5 of 5
2025-10-16
NVIDIA B200
129.5ms
#5 of 5
2025-10-16
NVIDIA B200
248.1ms
#5 of 5
2025-10-16
NVIDIA B200
267.3ms
#5 of 5
2025-10-16
NVIDIA B200
625.5ms
#5 of 5
2025-10-16
NVIDIA B200
1.55s
#5 of 5
2025-10-16
NVIDIA B200
1.61s
#5 of 5
2025-10-16
NVIDIA B200
2.61s
#5 of 5
2025-10-16
NVIDIA B200
5.89s
#5 of 5
2025-10-16
NVIDIA B200
8.25s
#5 of 5
2025-10-16

Reproduction-ready · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:cfd188b3111ac75c36302be921a813023797585305dd7e81f334d0668f4d9fc0
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20

Kernel source

main.cpp102 lines
#include "kernel.h"
#include <torch/extension.h>
#include <vector>
#include <cmath>
#include <limits>

/* -------------- helpers ---------------- */
#define CHECK_CUDA(x)   TORCH_CHECK(x.is_cuda(),   #x " must be on CUDA")
#define CHECK_CONTIG(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
#define CHECK_INPUT(x)  CHECK_CUDA(x); CHECK_CONTIG(x)

/* ==================================================================== */
/*  public entry                                                        */
/* ==================================================================== */
std::vector<torch::Tensor> run(
        const torch::Tensor& q_nope,
        const torch::Tensor& q_pe,
        const torch::Tensor& ckv_cache,
        const torch::Tensor& kpe_cache,
        const torch::Tensor& qo_indptr,
        const torch::Tensor& kv_indptr,
        const torch::Tensor& kv_indices,
        const c10::optional<double> sm_scale_opt = c10::nullopt)
{
    /* ---------------- basic checks ---------------------------------- */
    CHECK_INPUT(q_nope);    CHECK_INPUT(q_pe);
    CHECK_INPUT(ckv_cache); CHECK_INPUT(kpe_cache);
    CHECK_INPUT(qo_indptr); CHECK_INPUT(kv_indptr); CHECK_INPUT(kv_indices);

    TORCH_CHECK(q_nope.dtype() == torch::kBFloat16 &&
                q_pe.dtype()   == torch::kBFloat16,
                "q_nope and q_pe must be bfloat16");
    TORCH_CHECK(ckv_cache.dtype() == torch::kBFloat16 &&
                kpe_cache.dtype() == torch::kBFloat16,
                "ckv_cache and kpe_cache must be bfloat16");

    TORCH_CHECK(q_nope.size(1) == NUM_QO_HEADS &&
                q_nope.size(2) == HEAD_DIM_CKV, "q_nope shape mismatch");
    TORCH_CHECK(q_pe.size(1) == NUM_QO_HEADS &&
                q_pe.size(2) == HEAD_DIM_KPE, "q_pe shape mismatch");

    TORCH_CHECK(ckv_cache.size(1) == PAGE_SIZE &&
                kpe_cache.size(1) == PAGE_SIZE,
                "page_size must be 1");

    const int total_q          = q_nope.size(0);
    const int num_kv_indices   = kv_indices.size(0);

    TORCH_CHECK(total_q ==
                qo_indptr[qo_indptr.size(0) - 1].item<int32_t>(),
                "total_q must equal qo_indptr[-1]");
    TORCH_CHECK(num_kv_indices ==
                kv_indptr[kv_indptr.size(0) - 1].item<int32_t>(),
                "num_kv_indices must equal kv_indptr[-1]");

    /* ---------------- outputs --------------------------------------- */
    auto output = torch::empty(
        {total_q, NUM_QO_HEADS, HEAD_DIM_CKV},
        torch::dtype(torch::kBFloat16).device(q_nope.device()));

    auto lse = torch::full(
        {total_q, NUM_QO_HEADS},
        -std::numeric_limits<float>::infinity(),
        torch::dtype(torch::kFloat32).device(q_nope.device()));

    /* ---------------- softmax scale --------------------------------- */
    float sm_scale = sm_scale_opt.has_value()
                     ? static_cast<float>(*sm_scale_opt)
                     : 1.f / std::sqrt(static_cast<float>(HEAD_DIM_CKV));

    /* ---------------- launch kernel --------------------------------- */
    launch_mla_paged_prefill_causal_h16_ckv512_kpe64_ps1(
        q_nope, q_pe,
        ckv_cache, kpe_cache,
        qo_indptr, kv_indptr, kv_indices,
        sm_scale,
        output, lse);

    return {output, lse};
}

/* ==================================================================== */
/*  pybind11 binding                                                    */
/* ==================================================================== */
#include <pybind11/pybind11.h>
namespace py = pybind11;

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
{
    m.doc() = "Optimised paged-prefill causal MLA "
              "(H=16, Dckv=512, Dkpe=64, PS=1)";
    m.def("run", &run,
          py::arg("q_nope"),
          py::arg("q_pe"),
          py::arg("ckv_cache"),
          py::arg("kpe_cache"),
          py::arg("qo_indptr"),
          py::arg("kv_indptr"),
          py::arg("kv_indices"),
          py::arg("sm_scale") = c10::nullopt,
          "Compute (output, lse) for the given queries / cache.");
}
scrolls · 102 lines total

Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0

Best evidence level for this revision: reproducible

JSON