Skip to content
KernelIndex
Search⌘K

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