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