Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / cudaa7821e

gemini-2.5-pro_cuda_a7821e · gemini-2.5-pro · cuda · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Kernel source

main.cpp122 lines
#include "kernel.h"

#include <torch/extension.h>
#include <pybind11/pybind11.h>
#include <pybind11/stl.h>

#include <vector>
#include <stdexcept>
#include <string>
#include <cmath>

namespace py = pybind11;

#define CUDA_CHECK(call)                                                    \
  do {                                                                      \
    cudaError_t err = call;                                                 \
    if (err != cudaSuccess) {                                               \
      throw std::runtime_error(std::string("CUDA Error: ") +                 \
                               cudaGetErrorString(err) + " at " + __FILE__ + \
                               ":" + std::to_string(__LINE__));             \
    }                                                                       \
  } while (0)

/**
 * @brief Python entry point for the MLA Paged Prefill kernel.
 *
 * This function is exposed to Python via Pybind11. It performs extensive input
 * validation, prepares output tensors, pre-computes a helper mapping (`q_to_b_map`)
 * on the CPU for efficiency, and then calls the CUDA kernel launcher.
 *
 * @return A pair of torch::Tensor objects: the attention output and the LSE values.
 */
std::pair<torch::Tensor, torch::Tensor> run(
    torch::Tensor q_nope,
    torch::Tensor q_pe,
    torch::Tensor ckv_cache,
    torch::Tensor kpe_cache,
    torch::Tensor qo_indptr,
    torch::Tensor kv_indptr,
    torch::Tensor kv_indices,
    py::object sm_scale_obj)
{
    // --- Input Validation ---
    TORCH_CHECK(q_nope.is_cuda(), "q_nope must be a CUDA tensor");
    TORCH_CHECK(q_nope.scalar_type() == torch::kBFloat16, "q_nope must be of bfloat16 type");
    TORCH_CHECK(q_pe.is_cuda() && q_pe.scalar_type() == torch::kBFloat16, "q_pe must be a CUDA bfloat16 tensor");
    TORCH_CHECK(ckv_cache.is_cuda() && ckv_cache.scalar_type() == torch::kBFloat16, "ckv_cache must be a CUDA bfloat16 tensor");
    TORCH_CHECK(kpe_cache.is_cuda() && kpe_cache.scalar_type() == torch::kBFloat16, "kpe_cache must be a CUDA bfloat16 tensor");
    TORCH_CHECK(qo_indptr.is_cuda() && qo_indptr.scalar_type() == torch::kInt32, "qo_indptr must be a CUDA int32 tensor");
    TORCH_CHECK(kv_indptr.is_cuda() && kv_indptr.scalar_type() == torch::kInt32, "kv_indptr must be a CUDA int32 tensor");
    TORCH_CHECK(kv_indices.is_cuda() && kv_indices.scalar_type() == torch::kInt32, "kv_indices must be a CUDA int32 tensor");

    // Check dimensions against constants
    const int64_t total_q = q_nope.size(0);
    const int64_t num_qo_heads = q_nope.size(1);
    const int64_t head_dim_ckv = q_nope.size(2);
    const int64_t head_dim_kpe = q_pe.size(2);
    const int64_t page_size = ckv_cache.size(1);

    TORCH_CHECK(num_qo_heads == 16, "num_qo_heads must be 16");
    TORCH_CHECK(head_dim_ckv == 512, "head_dim_ckv must be 512");
    TORCH_CHECK(head_dim_kpe == 64, "head_dim_kpe must be 64");
    TORCH_CHECK(page_size == 1, "page_size must be 1");

    // --- Softmax Scale ---
    float sm_scale;
    if (sm_scale_obj.is_none()) {
        sm_scale = 1.0f / std::sqrt(128.0f + 64.0f);
    } else {
        sm_scale = py::cast<float>(sm_scale_obj);
    }

    // --- Prepare Outputs ---
    auto output = torch::empty_like(q_nope);
    auto lse = torch::empty({total_q, num_qo_heads}, q_nope.options().dtype(torch::kFloat32));

    // --- Host-side Pre-computation of q_to_b_map ---
    // This map allows the kernel to quickly find the batch index for any given query token,
    // avoiding a complex and divergent search within the kernel.
    torch::Tensor qo_indptr_cpu = qo_indptr.to(torch::kCPU);
    auto qo_indptr_acc = qo_indptr_cpu.accessor<int32_t, 1>();
    const int32_t batch_size = qo_indptr_cpu.size(0) - 1;

    auto q_to_b_map = torch::empty({total_q}, torch::TensorOptions().dtype(torch::kInt32));
    auto q_to_b_map_acc = q_to_b_map.accessor<int32_t, 1>();

    #pragma omp parallel for
    for (int b = 0; b < batch_size; ++b) {
        int32_t q_start = qo_indptr_acc[b];
        int32_t q_end = qo_indptr_acc[b + 1];
        for (int q_idx = q_start; q_idx < q_end; ++q_idx) {
            if (q_idx < total_q) {
                 q_to_b_map_acc[q_idx] = b;
            }
        }
    }
    auto q_to_b_map_gpu = q_to_b_map.to(q_nope.device());

    // --- Launch Kernel ---
    cudaStream_t stream = at::cuda::getCurrentCUDAStream();
    launch_mla_paged_prefill_causal(
        output, lse, q_nope, q_pe, ckv_cache, kpe_cache,
        qo_indptr, kv_indptr, kv_indices, q_to_b_map_gpu, sm_scale, stream);

    CUDA_CHECK(cudaGetLastError());

    return {output, lse};
}

// --- Pybind11 Module Definition ---
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run", &run, "MLA Paged Prefill (Causal)",
        py::arg("q_nope"),
        py::arg("q_pe"),
        py::arg("ckv_cache"),
        py::arg("kpe_cache"),
        py::arg("qo_indptr"),
        py::arg("kv_indptr"),
        py::arg("kv_indices"),
        py::arg("sm_scale") = py::none()
    );
}
scrolls · 122 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON