gemini-2.5-pro / cuda14d4de
gemini-2.5-pro_cuda_14d4de · gemini-2.5-pro · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 105 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-cuda-14d4de?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:732b34b88c02e9ad5b9ad837ff4d778c39ee0490171e04f8f0f77601e895d67a
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20
Kernel source
main.cpp105 lines
#include <torch/extension.h>
#include <pybind11/pybind11.h>
#include <stdexcept>
#include <vector>
#include <cmath>
#include <limits>
#include "kernel.h"
namespace py = pybind11;
// Helper to check 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");
// C++ entry point, called by Python
std::pair<torch::Tensor, torch::Tensor> gqa_ragged_prefill_causal_h32_kv8_d128(
const torch::Tensor& q,
const torch::Tensor& k,
const torch::Tensor& v,
const torch::Tensor& qo_indptr,
const torch::Tensor& kv_indptr,
py::object sm_scale_obj) {
// --- Input Validation ---
CHECK_TENSOR(q, "q");
CHECK_TENSOR(k, "k");
CHECK_TENSOR(v, "v");
CHECK_TENSOR(qo_indptr, "qo_indptr");
CHECK_TENSOR(kv_indptr, "kv_indptr");
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must have dtype bfloat16");
TORCH_CHECK(k.dtype() == torch::kBFloat16, "k must have dtype bfloat16");
TORCH_CHECK(v.dtype() == torch::kBFloat16, "v must have dtype bfloat16");
TORCH_CHECK(qo_indptr.dtype() == torch::kInt32, "qo_indptr must have dtype int32");
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must have dtype int32");
TORCH_CHECK(q.dim() == 3, "q must be 3-dimensional");
TORCH_CHECK(k.dim() == 3, "k must be 3-dimensional");
TORCH_CHECK(v.dim() == 3, "v must be 3-dimensional");
TORCH_CHECK(qo_indptr.dim() == 1, "qo_indptr must be 1-dimensional");
TORCH_CHECK(kv_indptr.dim() == 1, "kv_indptr must be 1-dimensional");
const int64_t total_q = q.size(0);
const int64_t num_qo_heads = q.size(1);
const int64_t head_dim = q.size(2);
const int64_t num_kv_heads = k.size(1);
const int64_t total_kv = k.size(0);
// Validate against fixed dimensions from spec
TORCH_CHECK(num_qo_heads == 32, "num_qo_heads must be 32");
TORCH_CHECK(num_kv_heads == 8, "num_kv_heads must be 8");
TORCH_CHECK(head_dim == 128, "head_dim must be 128");
// Validate constraints
if (qo_indptr.size(0) > 0) {
auto last_qo_indptr = qo_indptr.index({qo_indptr.size(0) - 1}).item<int>();
TORCH_CHECK(total_q == last_qo_indptr, "total_q must equal qo_indptr[-1]");
} else {
TORCH_CHECK(total_q == 0, "total_q must be 0 for empty qo_indptr");
}
if (kv_indptr.size(0) > 0) {
auto last_kv_indptr = kv_indptr.index({kv_indptr.size(0) - 1}).item<int>();
TORCH_CHECK(total_kv == last_kv_indptr, "total_kv must equal kv_indptr[-1]");
} else {
TORCH_CHECK(total_kv == 0, "total_kv must be 0 for empty kv_indptr");
}
// --- Softmax Scale ---
float sm_scale;
if (sm_scale_obj.is_none()) {
sm_scale = 1.0f / std::sqrt(static_cast<float>(head_dim));
} else {
sm_scale = py::cast<float>(sm_scale_obj);
}
// --- Output Allocation ---
auto output = torch::empty_like(q);
auto lse_options = torch::TensorOptions().device(q.device()).dtype(torch::kFloat32);
auto lse = torch::empty({total_q, num_qo_heads}, lse_options);
// --- Kernel Execution ---
run_gqa_ragged_prefill_causal_h32_kv8_d128(
q, k, v, qo_indptr, kv_indptr, sm_scale, output, lse
);
return {output, lse};
}
// --- Pybind11 Module Definition ---
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def(
"run",
&gqa_ragged_prefill_causal_h32_kv8_d128,
"gqa_ragged_prefill_causal_h32_kv8_d128 CUDA kernel",
py::arg("q"),
py::arg("k"),
py::arg("v"),
py::arg("qo_indptr"),
py::arg("kv_indptr"),
py::arg("sm_scale") = py::none()
);
}scrolls · 105 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON