Skip to content
KernelIndex
Search⌘K

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