Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / cudaed28aa

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-cuda-ed28aa?include=source"
interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16

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:c8990b6f0adfca24e9a17ff96d4aad795019e2ad480852414e5ace22237ece28
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20

Kernel source

main.cpp72 lines
#include <torch/extension.h>
#include <c10/cuda/CUDAStream.h>
#include "kernel.h"
#include <stdexcept>

// Helper to check tensor properties
void check_tensor(const torch::Tensor& tensor, const std::string& name) {
    if (!tensor.is_cuda()) {
        throw std::runtime_error(name + " tensor must be on a CUDA device.");
    }
    if (tensor.scalar_type() != torch::kFloat16) {
        throw std::runtime_error(name + " tensor must have dtype torch.float16.");
    }
    if (!tensor.is_contiguous()) {
        throw std::runtime_error(name + " tensor must be contiguous.");
    }
    if (tensor.dim() != 2) {
        throw std::runtime_error(name + " tensor must be 2-dimensional.");
    }
}

/**
 * @brief Python-bindable run function that executes the GEMM operation.
 *
 * This function acts as the interface between PyTorch and the custom CUDA kernel.
 * It performs input validation, allocates output tensor, and launches the kernel.
 *
 * @param A A torch.Tensor of shape [M, 2048] and dtype float16, on a CUDA device.
 * @param B A torch.Tensor of shape [128, 2048] and dtype float16, on a CUDA device.
 * @return A torch.Tensor of shape [M, 128] and dtype float16, on a CUDA device.
 */
torch::Tensor run(torch::Tensor A, torch::Tensor B) {
    // --- Input Validation ---
    check_tensor(A, "A");
    check_tensor(B, "B");

    const int M = A.size(0);
    const int K_A = A.size(1);
    const int N_B = B.size(0);
    const int K_B = B.size(1);
    
    // Fixed dimensions check
    if (K_A != 2048) {
        throw std::runtime_error("Dimension K of A must be 2048.");
    }
    if (N_B != 128) {
        throw std::runtime_error("Dimension N of B must be 128.");
    }
    if (K_B != 2048) {
        throw std::runtime_error("Dimension K of B must be 2048.");
    }
    
    // --- Output Allocation ---
    const int N = 128;
    auto C = torch::empty({M, N}, A.options());

    // --- Data Pointers ---
    const half* A_ptr = reinterpret_cast<const half*>(A.data_ptr<at::Half>());
    const half* B_ptr = reinterpret_cast<const half*>(B.data_ptr<at::Half>());
    half* C_ptr = reinterpret_cast<half*>(C.data_ptr<at::Half>());

    // --- Kernel Execution ---
    cudaStream_t stream = at::cuda::getCurrentCUDAStream();
    gemm_n128_k2048_launch(A_ptr, B_ptr, C_ptr, M, stream);

    return C;
}

// --- Pybind11 Module Definition ---
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run", &run, "GEMM N=128 K=2048 (FP16) CUDA kernel for B200");
}
scrolls · 72 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON