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
Show all 36 measurements ›Showing all 36 measurements ⌄
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [805, 16, 64]
NVIDIA B200
128.5ms
#5 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [1954, 16, 64]
NVIDIA B200
129.5ms
#5 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [1028, 16, 64]
NVIDIA B200
248.1ms
#5 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [1187, 16, 64]
NVIDIA B200
267.3ms
#5 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [3842, 16, 64]
NVIDIA B200
625.5ms
#5 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [15092, 16, 64]
NVIDIA B200
5.89s
#5 of 5
2025-10-16
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [15883, 16, 64]
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