gemini-2.5-pro / cuda4bc599
gemini-2.5-pro_cuda_4bc599 · gemini-2.5-pro · cuda · Apache-2.0
Kernel source · 89 lines ↓holds 21 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 89 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-cuda-4bc599?include=source"interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16
Benchmark evidence
43 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Show all 43 measurements ›Showing all 43 measurements ⌄
Reproduction-ready · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:5f7701c65c849769e2afde4758abeb07be8ea616b7ac1add7e78ff0988abdc06
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20
Kernel source
main.cpp89 lines
#include <torch/extension.h>
#include <c10/cuda/CUDAStream.h>
#include <cublas_v2.h>
#include <memory>
#include "kernel.h"
// Helper macros for PyTorch tensor validation
#define CHECK_CUDA(x) TORCH_CHECK(x.is_cuda(), #x " must be a CUDA tensor")
#define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
#define CHECK_DTYPE_FP16(x) TORCH_CHECK(x.scalar_type() == torch::kFloat16, #x " must be a Float16 tensor")
#define CHECK_INPUT(x) CHECK_CUDA(x); CHECK_CONTIGUOUS(x); CHECK_DTYPE_FP16(x)
// RAII wrapper for cublasHandle_t to ensure it's always destroyed.
struct CublasHandle {
cublasHandle_t handle;
CublasHandle() {
TORCH_CHECK(cublasCreate(&handle) == CUBLAS_STATUS_SUCCESS, "cuBLAS handle creation failed");
}
~CublasHandle() {
cublasDestroy(handle);
}
// Allow the struct to be passed directly to functions expecting a handle
operator cublasHandle_t() const { return handle; }
};
/**
* @brief Python-bindable function to execute the GEMM operation.
*
* This function takes two PyTorch tensors, A and B, validates them,
* and calls the custom CUDA/cuBLAS kernel to compute C = A * B.T.
*
* @param A A torch::Tensor with shape [M, 4096] and dtype float16, on a CUDA device.
* @param B A torch::Tensor with shape [6144, 4096] and dtype float16, on a CUDA device.
* @return A torch::Tensor with shape [M, 6144] and dtype float16, on the same CUDA device.
*/
torch::Tensor run(torch::Tensor A, torch::Tensor B) {
// ---- Input Validation ----
CHECK_INPUT(A);
CHECK_INPUT(B);
TORCH_CHECK(A.dim() == 2, "Input tensor A must be 2-dimensional");
TORCH_CHECK(B.dim() == 2, "Input tensor B must be 2-dimensional");
// Check against fixed dimensions from the specification
constexpr int N_dim = 6144;
constexpr int K_dim = 4096;
TORCH_CHECK(B.size(0) == N_dim, "B.shape[0] must be ", N_dim);
TORCH_CHECK(A.size(1) == K_dim, "A.shape[1] must be ", K_dim);
TORCH_CHECK(B.size(1) == K_dim, "B.shape[1] must be ", K_dim);
const int M_dim = A.size(0);
// ---- Output Tensor Preparation ----
auto C = torch::empty({M_dim, N_dim}, A.options());
// ---- Kernel Execution ----
try {
// Create a cuBLAS handle (RAII ensures cleanup)
static thread_local CublasHandle handle;
// Get the current CUDA stream from PyTorch
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
// Get raw data pointers from PyTorch tensors
const __half* ptr_A = reinterpret_cast<const __half*>(A.data_ptr<at::Half>());
const __half* ptr_B = reinterpret_cast<const __half*>(B.data_ptr<at::Half>());
__half* ptr_C = reinterpret_cast<__half*>(C.data_ptr<at::Half>());
// Launch the custom CUDA kernel
gemm_n6144_k4096_launcher(handle, M_dim, ptr_A, ptr_B, ptr_C, stream);
} catch (const std::exception& e) {
// Propagate exceptions from the CUDA code to PyTorch
TORCH_CHECK(false, "GEMM kernel execution failed: ", e.what());
}
// Check for any asynchronous CUDA errors from the kernel launch.
// Note: cuBLAS calls are also asynchronous.
cudaError_t err = cudaGetLastError();
TORCH_CHECK(err == cudaSuccess, "CUDA error after kernel launch: ", cudaGetErrorString(err));
return C;
}
// ---- Pybind11 Module Definition ----
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("run", &run, "GEMM(A[M, 4096], B[6144, 4096].T) implementation using cuBLAS, optimized for B200 GPU");
}scrolls · 89 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reproducible
JSON