claude-opus-4-1 / cuda53eadf
claude-opus-4-1_cuda_53eadf · claude-opus-4-1-20250805 · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 149 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-cuda-53eadf?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:d8d5a113b4c8ed2a19a59bcc75715d4125ca1836797606623a6bb427e57f6da2
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.cpp149 lines
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cublas_v2.h>
#include "kernel.h"
#include <iostream>
#include <stdexcept>
#include <memory>
#include <mutex>
// Helper macros for error checking
#define CUDA_CHECK(call) \
do { \
cudaError_t error = call; \
if (error != cudaSuccess) { \
throw std::runtime_error(std::string("CUDA error at ") + __FILE__ + ":" + \
std::to_string(__LINE__) + " - " + cudaGetErrorString(error)); \
} \
} while(0)
#define CUBLAS_CHECK(call) \
do { \
cublasStatus_t status = call; \
if (status != CUBLAS_STATUS_SUCCESS) { \
throw std::runtime_error(std::string("cuBLAS error at ") + __FILE__ + ":" + \
std::to_string(__LINE__) + " code: " + std::to_string(status)); \
} \
} while(0)
// Thread-safe cuBLAS handle management
class CublasHandleManager {
private:
cublasHandle_t handle;
static std::unique_ptr<CublasHandleManager> instance;
static std::mutex mutex;
CublasHandleManager() {
CUBLAS_CHECK(cublasCreate(&handle));
// Enable tensor cores
CUBLAS_CHECK(cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH));
}
public:
~CublasHandleManager() {
if (handle) {
cublasDestroy(handle);
}
}
static cublasHandle_t get() {
std::lock_guard<std::mutex> lock(mutex);
if (!instance) {
instance.reset(new CublasHandleManager());
}
return instance->handle;
}
CublasHandleManager(const CublasHandleManager&) = delete;
CublasHandleManager& operator=(const CublasHandleManager&) = delete;
};
std::unique_ptr<CublasHandleManager> CublasHandleManager::instance = nullptr;
std::mutex CublasHandleManager::mutex;
torch::Tensor run(torch::Tensor A, torch::Tensor B) {
// Input validation
TORCH_CHECK(A.dim() == 2, "Input A must be 2-dimensional, got ", A.dim());
TORCH_CHECK(B.dim() == 2, "Input B must be 2-dimensional, got ", B.dim());
TORCH_CHECK(A.size(1) == K_SIZE, "A must have ", K_SIZE, " columns, got ", A.size(1));
TORCH_CHECK(B.size(0) == N_SIZE, "B must have ", N_SIZE, " rows, got ", B.size(0));
TORCH_CHECK(B.size(1) == K_SIZE, "B must have ", K_SIZE, " columns, got ", B.size(1));
TORCH_CHECK(A.scalar_type() == torch::kFloat16, "A must be float16");
TORCH_CHECK(B.scalar_type() == torch::kFloat16, "B must be float16");
TORCH_CHECK(A.is_cuda(), "A must be on CUDA device");
TORCH_CHECK(B.is_cuda(), "B must be on CUDA device");
TORCH_CHECK(A.device() == B.device(), "A and B must be on the same device");
// Make tensors contiguous if needed
torch::Tensor A_contig = A.contiguous();
torch::Tensor B_contig = B.contiguous();
const int M = A_contig.size(0);
// Create output tensor
auto options = torch::TensorOptions()
.dtype(torch::kFloat16)
.device(A_contig.device())
.requires_grad(false);
torch::Tensor C = torch::empty({M, N_SIZE}, options);
// Get current CUDA stream
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
// Choose implementation based on matrix size
// For large matrices, cuBLAS is optimal on B200
if (M >= 256) {
// Use cuBLAS for optimal performance on large matrices
cublasHandle_t handle = CublasHandleManager::get();
CUBLAS_CHECK(cublasSetStream(handle, stream));
const __half alpha = __float2half(1.0f);
const __half beta = __float2half(0.0f);
// Compute C = A * B^T using cuBLAS
// We need to compute C = A * B^T
// In column-major view: C^T = B * A^T
// Since PyTorch uses row-major, we can directly compute:
// C(m,n) = A(m,:) * B(n,:)^T = A(m,:) * B^T(:,n)
// Using cublasGemmEx for better performance
CUBLAS_CHECK(cublasGemmEx(
handle,
CUBLAS_OP_T, // B needs to be transposed
CUBLAS_OP_N, // A is not transposed
N_SIZE, // m - rows of result
M, // n - cols of result
K_SIZE, // k - reduction dimension
&alpha,
B_contig.data_ptr<at::Half>(), // B
CUDA_R_16F, // B datatype
K_SIZE, // ldb - leading dimension of B
A_contig.data_ptr<at::Half>(), // A
CUDA_R_16F, // A datatype
K_SIZE, // lda - leading dimension of A
&beta,
C.data_ptr<at::Half>(), // C
CUDA_R_16F, // C datatype
N_SIZE, // ldc - leading dimension of C
CUBLAS_COMPUTE_16F, // compute type
CUBLAS_GEMM_DEFAULT_TENSOR_OP // algorithm
));
} else {
// Use custom kernel for smaller matrices
launch_gemm_kernel(
reinterpret_cast<const half*>(A_contig.data_ptr<at::Half>()),
reinterpret_cast<const half*>(B_contig.data_ptr<at::Half>()),
reinterpret_cast<half*>(C.data_ptr<at::Half>()),
M,
stream
);
}
return C;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("run", &run, "Optimized GEMM kernel for M x 4096 @ 28672 x 4096 -> M x 28672",
py::arg("A"), py::arg("B"));
}scrolls · 149 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON