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