gemini-2.5-pro / cuda188d8f
gemini-2.5-pro_cuda_188d8f · gemini-2.5-pro · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 112 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-cuda-188d8f?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:af98170c948d4a3323e4bee129b48835b6cd80ca0146dc3f3d8dbb392a39f342
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20
Kernel source
main.cpp112 lines
#include "kernel.h"
#include <torch/extension.h>
#include <vector>
#include <stdexcept>
#include <limits>
// Helper macro for checking tensor properties
#define CHECK_TENSOR(T, NAME) \
TORCH_CHECK(T.is_cuda(), #NAME, " must be a CUDA tensor"); \
TORCH_CHECK(T.is_contiguous(), #NAME, " must be contiguous");
// Helper macro for CUDA error checking
#define CUDA_CHECK(call) \
do { \
cudaError_t err = call; \
if (err != cudaSuccess) { \
throw std::runtime_error(std::string("CUDA error in " #call ": ") + cudaGetErrorString(err)); \
} \
} while (0)
// C++ entry point, called by Python.
std::vector<torch::Tensor> run(
torch::Tensor q,
torch::Tensor k_cache,
torch::Tensor v_cache,
torch::Tensor qo_indptr,
torch::Tensor kv_indptr,
torch::Tensor kv_indices,
float sm_scale
) {
// --- 1. Input Validation ---
CHECK_TENSOR(q, "q");
CHECK_TENSOR(k_cache, "k_cache");
CHECK_TENSOR(v_cache, "v_cache");
CHECK_TENSOR(qo_indptr, "qo_indptr");
CHECK_TENSOR(kv_indptr, "kv_indptr");
CHECK_TENSOR(kv_indices, "kv_indices");
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bfloat16");
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bfloat16");
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bfloat16");
TORCH_CHECK(qo_indptr.dtype() == torch::kInt32, "qo_indptr must be int32");
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32");
TORCH_CHECK(kv_indices.dtype() == torch::kInt32, "kv_indices must be int32");
const int total_q = q.size(0);
const int num_qo_heads = q.size(1);
const int head_dim = q.size(2);
const int num_kv_heads = k_cache.size(2);
const int len_indptr = qo_indptr.size(0);
// Check against fixed dimensions from spec
TORCH_CHECK(num_qo_heads == 32, "num_qo_heads must be 32");
TORCH_CHECK(num_kv_heads == 4, "num_kv_heads must be 4");
TORCH_CHECK(head_dim == 128, "head_dim must be 128");
TORCH_CHECK(k_cache.size(1) == 1, "page_size must be 1");
// --- 2. Create Output Tensors ---
auto output = torch::empty_like(q);
auto lse = torch::full(
{total_q, num_qo_heads},
-std::numeric_limits<float>::infinity(),
torch::TensorOptions()
.dtype(torch::kFloat32)
.device(q.device())
);
// --- 3. Pack Parameters and Launch Kernel ---
if (total_q > 0) {
GqaPagedPrefillCausalParams params;
params.output_ptr = output.data_ptr();
params.lse_ptr = lse.data_ptr<float>();
params.q_ptr = q.data_ptr();
params.k_cache_ptr = k_cache.data_ptr();
params.v_cache_ptr = v_cache.data_ptr();
params.qo_indptr_ptr = qo_indptr.data_ptr<int>();
params.kv_indptr_ptr = kv_indptr.data_ptr<int>();
params.kv_indices_ptr = kv_indices.data_ptr<int>();
params.sm_scale = sm_scale;
params.total_q = total_q;
params.len_indptr = len_indptr;
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
try {
gqa_paged_prefill_causal_h32_kv4_d128_ps1_launch(params, stream);
CUDA_CHECK(cudaGetLastError());
} catch (const std::exception& e) {
// Make sure to rethrow C++ exceptions as Python exceptions
throw std::runtime_error(std::string("gqa_paged_prefill CUDA launch failed: ") + e.what());
}
}
return {output, lse};
}
// --- Pybind11 Module Definition ---
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def(
"run",
&run,
"GQA Paged Prefill Causal H32 KV4 D128 PS1 (CUDA)",
py::arg("q"),
py::arg("k_cache"),
py::arg("v_cache"),
py::arg("qo_indptr"),
py::arg("kv_indptr"),
py::arg("kv_indices"),
py::arg("sm_scale")
);
}scrolls · 112 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON