Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro_cuda_292432

gemini-2.5-pro · cuda · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 93 lines, Apache-2.0, pinned at da91508.

main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-cuda-292432?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:60fbe2ed0c060a46c10573c8277b9436748d22bfb50cba978f06d6e34bec8252
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20

Kernel source

main.cpp93 lines
#include "kernel.h"
#include <torch/extension.h>
#include <pybind11/pybind11.h>
#include <pybind11/stl.h>

namespace py = pybind11;

// --- Helper for Tensor Validation ---
void validate_tensor(const torch::Tensor& t, const std::string& name, torch::ScalarType dtype, int dims) {
    TORCH_CHECK(t.is_cuda(), name, " must be a CUDA tensor");
    TORCH_CHECK(t.dtype() == dtype, name, " must have dtype ", dtype);
    TORCH_CHECK(t.dim() == dims, name, " must be ", dims, "D");
    TORCH_CHECK(t.is_contiguous(), name, " must be contiguous");
}

// --- Main 'run' function exposed to Python ---
py::dict run(
    torch::Tensor q_nope,
    torch::Tensor q_pe,
    torch::Tensor ckv_cache,
    torch::Tensor kpe_cache,
    torch::Tensor kv_indptr,
    torch::Tensor kv_indices,
    float sm_scale
) {
    // --- Input Validation ---
    validate_tensor(q_nope, "q_nope", torch::kBFloat16, 3);
    validate_tensor(q_pe, "q_pe", torch::kBFloat16, 3);
    validate_tensor(ckv_cache, "ckv_cache", torch::kBFloat16, 3);
    validate_tensor(kpe_cache, "kpe_cache", torch::kBFloat16, 3);
    validate_tensor(kv_indptr, "kv_indptr", torch::kInt32, 1);
    validate_tensor(kv_indices, "kv_indices", torch::kInt32, 1);

    const int batch_size = q_nope.size(0);
    const int num_qo_heads = q_nope.size(1);
    const int head_dim_ckv = q_nope.size(2);
    const int head_dim_kpe = q_pe.size(2);
    const int page_size = ckv_cache.size(1);

    // Check fixed dimensions from spec
    TORCH_CHECK(num_qo_heads == 16, "num_qo_heads must be 16, but got ", num_qo_heads);
    TORCH_CHECK(head_dim_ckv == 512, "head_dim_ckv must be 512, but got ", head_dim_ckv);
    TORCH_CHECK(head_dim_kpe == 64, "head_dim_kpe must be 64, but got ", head_dim_kpe);
    TORCH_CHECK(page_size == 1, "page_size must be 1, but got ", page_size);

    // Check constraints
    TORCH_CHECK(kv_indptr.size(0) == batch_size + 1, "len_indptr must be batch_size + 1");
    // The following check requires synchronization and is expensive.
    // It is assumed the user provides valid inputs as per spec.
    // torch::Tensor last_indptr = kv_indptr.index({-1}).to(torch::kCPU);
    // TORCH_CHECK(kv_indices.size(0) == last_indptr.item<int>(), "num_kv_indices must equal kv_indptr[-1]");

    // --- Create Output Tensors ---
    auto output = torch::empty_like(q_nope);
    auto lse = torch::empty({batch_size, num_qo_heads}, q_nope.options().dtype(torch::kFloat32));

    // --- Prepare Kernel Parameters ---
    MlaPagedDecodeParams params;
    params.q_nope_ptr = reinterpret_cast<const __nv_bfloat16*>(q_nope.data_ptr());
    params.q_pe_ptr = reinterpret_cast<const __nv_bfloat16*>(q_pe.data_ptr());
    params.ckv_cache_ptr = reinterpret_cast<const __nv_bfloat16*>(ckv_cache.data_ptr());
    params.kpe_cache_ptr = reinterpret_cast<const __nv_bfloat16*>(kpe_cache.data_ptr());
    params.kv_indptr_ptr = kv_indptr.data_ptr<int>();
    params.kv_indices_ptr = kv_indices.data_ptr<int>();
    params.sm_scale = sm_scale;

    params.output_ptr = reinterpret_cast<__nv_bfloat16*>(output.data_ptr());
    params.lse_ptr = lse.data_ptr<float>();

    params.batch_size = batch_size;

    // --- Launch Kernel ---
    cudaStream_t stream = at::cuda::getCurrentCUDAStream();
    mla_paged_decode_launch(params, stream);
    
    // --- Return Results ---
    py::dict result;
    result["output"] = output;
    result["lse"] = lse;
    return result;
}

// --- Pybind11 Module Definition ---
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run", &run, "MLA Paged Decode Kernel (BFLOAT16, Optimized)",
          py::arg("q_nope"),
          py::arg("q_pe"),
          py::arg("ckv_cache"),
          py::arg("kpe_cache"),
          py::arg("kv_indptr"),
          py::arg("kv_indices"),
          py::arg("sm_scale"));
}
scrolls · 93 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON