Skip to content
KernelIndex
Search⌘K

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