Skip to content
KernelIndex
Search⌘K

gpt-o3 / cuda025ec1

gpt-o3_cuda_025ec1 · gpt-o3 · cuda · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

21 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
18.6µs
#3 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
18.9µs
#4 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #5fa6fa
NVIDIA B200
19.3µs
#2 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #cb4b00
NVIDIA B200
46.9µs
#2 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
47.7µs
#2 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
47.7µs
#3 of 10
2025-10-19
NVIDIA B200
96.9µs
#5 of 10
2025-10-19
NVIDIA B200
96.9µs
#6 of 10
2025-10-19
NVIDIA B200
160.0µs
#2 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
175.8µs
#13 of 20
2025-10-19
Show all 21 measurements ›
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
175.9µs
#14 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
175.9µs
#15 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #f3d59b
NVIDIA B200
175.9µs
#4 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
176.4µs
#16 of 20
2025-10-19
NVIDIA B200
322.1µs
#4 of 5
2025-10-19
NVIDIA B200
369.1µs
#7 of 10
2025-10-19
NVIDIA B200
369.4µs
#8 of 10
2025-10-19
NVIDIA B200
2.40ms
#4 of 5
2025-10-19
NVIDIA B200
122.2ms
#4 of 5
2025-10-19
NVIDIA B200
125.5ms
#4 of 5
2025-10-19
NVIDIA B200
126.6ms
#4 of 5
2025-10-19

Reproduction-ready · How evidence levels are derived →

Source and license

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

Kernel source

main.cpp95 lines
#include "kernel.h"

#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime.h>

#define CHECK_CUDA(x) TORCH_CHECK(x.is_cuda(), #x " must reside on CUDA device")
#define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
#define CHECK_INPUT(x)  \
    CHECK_CUDA(x);      \
    CHECK_CONTIGUOUS(x)

torch::Tensor gqa_run(torch::Tensor q,
                      torch::Tensor k,
                      torch::Tensor v,
                      torch::Tensor qo_indptr,
                      torch::Tensor kv_indptr,
                      float sm_scale,
                      torch::Tensor output,
                      torch::Tensor lse) {
    CHECK_INPUT(q);
    CHECK_INPUT(k);
    CHECK_INPUT(v);
    CHECK_INPUT(qo_indptr);
    CHECK_INPUT(kv_indptr);
    CHECK_INPUT(output);
    CHECK_INPUT(lse);

    TORCH_CHECK(q.scalar_type()  == at::kBFloat16, "q must be bfloat16");
    TORCH_CHECK(k.scalar_type()  == at::kBFloat16, "k must be bfloat16");
    TORCH_CHECK(v.scalar_type()  == at::kBFloat16, "v must be bfloat16");
    TORCH_CHECK(output.scalar_type() == at::kBFloat16, "output must be bfloat16");
    TORCH_CHECK(lse.scalar_type()    == at::kFloat,    "lse must be float32");

    int total_q  = q.size(0);
    int total_kv = k.size(0);
    int len_indptr = qo_indptr.numel();

    // pointer casts
    const __nv_bfloat16 *q_ptr  = reinterpret_cast<const __nv_bfloat16*>(q.data_ptr<at::BFloat16>());
    const __nv_bfloat16 *k_ptr  = reinterpret_cast<const __nv_bfloat16*>(k.data_ptr<at::BFloat16>());
    const __nv_bfloat16 *v_ptr  = reinterpret_cast<const __nv_bfloat16*>(v.data_ptr<at::BFloat16>());
    const int32_t *qo_ptr       = qo_indptr.data_ptr<int32_t>();
    const int32_t *kv_ptr       = kv_indptr.data_ptr<int32_t>();
    __nv_bfloat16 *out_ptr      = reinterpret_cast<__nv_bfloat16*>(output.data_ptr<at::BFloat16>());
    float *lse_ptr              = lse.data_ptr<float>();

    cudaStream_t stream = at::cuda::getCurrentCUDAStream();

    gqa_ragged_prefill_causal_h32_kv8_d128_launcher(
        q_ptr, k_ptr, v_ptr,
        qo_ptr, kv_ptr,
        sm_scale,
        len_indptr,
        total_q,
        total_kv,
        out_ptr,
        lse_ptr,
        stream);

    TORCH_CHECK(cudaGetLastError() == cudaSuccess,
                "Kernel launch failed");

    return output;   // we also changed lse in-place
}

// ------------------------------------------------------------------
// Python binding
// ------------------------------------------------------------------
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run",
          [](torch::Tensor q,
             torch::Tensor k,
             torch::Tensor v,
             torch::Tensor qo_indptr,
             torch::Tensor kv_indptr,
             float sm_scale) {
                auto options_bf16 = torch::dtype(torch::kBFloat16).device(q.device());
                auto options_f32  = torch::dtype(torch::kFloat32).device(q.device());

                int64_t total_q = q.size(0);

                auto output = torch::empty({total_q, 32, 128}, options_bf16);
                auto lse    = torch::empty({total_q, 32},       options_f32);

                gqa_run(q, k, v, qo_indptr, kv_indptr, sm_scale, output, lse);
                return std::make_tuple(output, lse);
          },
          py::arg("q"),
          py::arg("k"),
          py::arg("v"),
          py::arg("qo_indptr"),
          py::arg("kv_indptr"),
          py::arg("sm_scale") = 1.0f / 11.313708498984761f); // 1/sqrt(128)
}
scrolls · 95 lines total

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

Best evidence level for this revision: reproducible

JSON