gemini-2.5-pro / cudaa7821e
gemini-2.5-pro_cuda_a7821e · gemini-2.5-pro · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 122 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-cuda-a7821e?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:279c50d6ea269abe66bf63433c293fbcf36aae02ddb46d37160639cc016d1152
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20
Kernel source
main.cpp122 lines
#include "kernel.h"
#include <torch/extension.h>
#include <pybind11/pybind11.h>
#include <pybind11/stl.h>
#include <vector>
#include <stdexcept>
#include <string>
#include <cmath>
namespace py = pybind11;
#define CUDA_CHECK(call) \
do { \
cudaError_t err = call; \
if (err != cudaSuccess) { \
throw std::runtime_error(std::string("CUDA Error: ") + \
cudaGetErrorString(err) + " at " + __FILE__ + \
":" + std::to_string(__LINE__)); \
} \
} while (0)
/**
* @brief Python entry point for the MLA Paged Prefill kernel.
*
* This function is exposed to Python via Pybind11. It performs extensive input
* validation, prepares output tensors, pre-computes a helper mapping (`q_to_b_map`)
* on the CPU for efficiency, and then calls the CUDA kernel launcher.
*
* @return A pair of torch::Tensor objects: the attention output and the LSE values.
*/
std::pair<torch::Tensor, torch::Tensor> run(
torch::Tensor q_nope,
torch::Tensor q_pe,
torch::Tensor ckv_cache,
torch::Tensor kpe_cache,
torch::Tensor qo_indptr,
torch::Tensor kv_indptr,
torch::Tensor kv_indices,
py::object sm_scale_obj)
{
// --- Input Validation ---
TORCH_CHECK(q_nope.is_cuda(), "q_nope must be a CUDA tensor");
TORCH_CHECK(q_nope.scalar_type() == torch::kBFloat16, "q_nope must be of bfloat16 type");
TORCH_CHECK(q_pe.is_cuda() && q_pe.scalar_type() == torch::kBFloat16, "q_pe must be a CUDA bfloat16 tensor");
TORCH_CHECK(ckv_cache.is_cuda() && ckv_cache.scalar_type() == torch::kBFloat16, "ckv_cache must be a CUDA bfloat16 tensor");
TORCH_CHECK(kpe_cache.is_cuda() && kpe_cache.scalar_type() == torch::kBFloat16, "kpe_cache must be a CUDA bfloat16 tensor");
TORCH_CHECK(qo_indptr.is_cuda() && qo_indptr.scalar_type() == torch::kInt32, "qo_indptr must be a CUDA int32 tensor");
TORCH_CHECK(kv_indptr.is_cuda() && kv_indptr.scalar_type() == torch::kInt32, "kv_indptr must be a CUDA int32 tensor");
TORCH_CHECK(kv_indices.is_cuda() && kv_indices.scalar_type() == torch::kInt32, "kv_indices must be a CUDA int32 tensor");
// Check dimensions against constants
const int64_t total_q = q_nope.size(0);
const int64_t num_qo_heads = q_nope.size(1);
const int64_t head_dim_ckv = q_nope.size(2);
const int64_t head_dim_kpe = q_pe.size(2);
const int64_t page_size = ckv_cache.size(1);
TORCH_CHECK(num_qo_heads == 16, "num_qo_heads must be 16");
TORCH_CHECK(head_dim_ckv == 512, "head_dim_ckv must be 512");
TORCH_CHECK(head_dim_kpe == 64, "head_dim_kpe must be 64");
TORCH_CHECK(page_size == 1, "page_size must be 1");
// --- Softmax Scale ---
float sm_scale;
if (sm_scale_obj.is_none()) {
sm_scale = 1.0f / std::sqrt(128.0f + 64.0f);
} else {
sm_scale = py::cast<float>(sm_scale_obj);
}
// --- Prepare Outputs ---
auto output = torch::empty_like(q_nope);
auto lse = torch::empty({total_q, num_qo_heads}, q_nope.options().dtype(torch::kFloat32));
// --- Host-side Pre-computation of q_to_b_map ---
// This map allows the kernel to quickly find the batch index for any given query token,
// avoiding a complex and divergent search within the kernel.
torch::Tensor qo_indptr_cpu = qo_indptr.to(torch::kCPU);
auto qo_indptr_acc = qo_indptr_cpu.accessor<int32_t, 1>();
const int32_t batch_size = qo_indptr_cpu.size(0) - 1;
auto q_to_b_map = torch::empty({total_q}, torch::TensorOptions().dtype(torch::kInt32));
auto q_to_b_map_acc = q_to_b_map.accessor<int32_t, 1>();
#pragma omp parallel for
for (int b = 0; b < batch_size; ++b) {
int32_t q_start = qo_indptr_acc[b];
int32_t q_end = qo_indptr_acc[b + 1];
for (int q_idx = q_start; q_idx < q_end; ++q_idx) {
if (q_idx < total_q) {
q_to_b_map_acc[q_idx] = b;
}
}
}
auto q_to_b_map_gpu = q_to_b_map.to(q_nope.device());
// --- Launch Kernel ---
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
launch_mla_paged_prefill_causal(
output, lse, q_nope, q_pe, ckv_cache, kpe_cache,
qo_indptr, kv_indptr, kv_indices, q_to_b_map_gpu, sm_scale, stream);
CUDA_CHECK(cudaGetLastError());
return {output, lse};
}
// --- Pybind11 Module Definition ---
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("run", &run, "MLA Paged Prefill (Causal)",
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") = py::none()
);
}scrolls · 122 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON