gemini-2.5-pro_cuda_292432
gemini-2.5-pro · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 93 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-cuda-292432?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:60fbe2ed0c060a46c10573c8277b9436748d22bfb50cba978f06d6e34bec8252
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20
Kernel source
main.cpp93 lines
#include "kernel.h"
#include <torch/extension.h>
#include <pybind11/pybind11.h>
#include <pybind11/stl.h>
namespace py = pybind11;
// --- Helper for Tensor Validation ---
void validate_tensor(const torch::Tensor& t, const std::string& name, torch::ScalarType dtype, int dims) {
TORCH_CHECK(t.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(t.dtype() == dtype, name, " must have dtype ", dtype);
TORCH_CHECK(t.dim() == dims, name, " must be ", dims, "D");
TORCH_CHECK(t.is_contiguous(), name, " must be contiguous");
}
// --- Main 'run' function exposed to Python ---
py::dict 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,
float sm_scale
) {
// --- Input Validation ---
validate_tensor(q_nope, "q_nope", torch::kBFloat16, 3);
validate_tensor(q_pe, "q_pe", torch::kBFloat16, 3);
validate_tensor(ckv_cache, "ckv_cache", torch::kBFloat16, 3);
validate_tensor(kpe_cache, "kpe_cache", torch::kBFloat16, 3);
validate_tensor(kv_indptr, "kv_indptr", torch::kInt32, 1);
validate_tensor(kv_indices, "kv_indices", torch::kInt32, 1);
const int batch_size = q_nope.size(0);
const int num_qo_heads = q_nope.size(1);
const int head_dim_ckv = q_nope.size(2);
const int head_dim_kpe = q_pe.size(2);
const int page_size = ckv_cache.size(1);
// Check fixed dimensions from spec
TORCH_CHECK(num_qo_heads == 16, "num_qo_heads must be 16, but got ", num_qo_heads);
TORCH_CHECK(head_dim_ckv == 512, "head_dim_ckv must be 512, but got ", head_dim_ckv);
TORCH_CHECK(head_dim_kpe == 64, "head_dim_kpe must be 64, but got ", head_dim_kpe);
TORCH_CHECK(page_size == 1, "page_size must be 1, but got ", page_size);
// Check constraints
TORCH_CHECK(kv_indptr.size(0) == batch_size + 1, "len_indptr must be batch_size + 1");
// The following check requires synchronization and is expensive.
// It is assumed the user provides valid inputs as per spec.
// torch::Tensor last_indptr = kv_indptr.index({-1}).to(torch::kCPU);
// TORCH_CHECK(kv_indices.size(0) == last_indptr.item<int>(), "num_kv_indices must equal kv_indptr[-1]");
// --- Create Output Tensors ---
auto output = torch::empty_like(q_nope);
auto lse = torch::empty({batch_size, num_qo_heads}, q_nope.options().dtype(torch::kFloat32));
// --- Prepare Kernel Parameters ---
MlaPagedDecodeParams params;
params.q_nope_ptr = reinterpret_cast<const __nv_bfloat16*>(q_nope.data_ptr());
params.q_pe_ptr = reinterpret_cast<const __nv_bfloat16*>(q_pe.data_ptr());
params.ckv_cache_ptr = reinterpret_cast<const __nv_bfloat16*>(ckv_cache.data_ptr());
params.kpe_cache_ptr = reinterpret_cast<const __nv_bfloat16*>(kpe_cache.data_ptr());
params.kv_indptr_ptr = kv_indptr.data_ptr<int>();
params.kv_indices_ptr = kv_indices.data_ptr<int>();
params.sm_scale = sm_scale;
params.output_ptr = reinterpret_cast<__nv_bfloat16*>(output.data_ptr());
params.lse_ptr = lse.data_ptr<float>();
params.batch_size = batch_size;
// --- Launch Kernel ---
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
mla_paged_decode_launch(params, stream);
// --- Return Results ---
py::dict result;
result["output"] = output;
result["lse"] = lse;
return result;
}
// --- Pybind11 Module Definition ---
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("run", &run, "MLA Paged Decode Kernel (BFLOAT16, Optimized)",
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 · 93 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON