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