Skip to content
KernelIndex
Search⌘K

claude-opus-4-1 / cudafeab9b

claude-opus-4-1_cuda_feab9b · claude-opus-4-1-20250805 · cuda · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-cuda-feab9b?include=source"
interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32

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:790df41678ff47667523bb19a6e547f0df751ce699422ffcdc09d0d2e4d4f32b
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20

Kernel source

main.cpp157 lines
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <vector>
#include <stdexcept>
#include <limits>
#include "kernel.h"

namespace py = pybind11;

// Helper function to check CUDA errors
#define CHECK_CUDA(x) \
    do { \
        cudaError_t err = x; \
        if (err != cudaSuccess) { \
            throw std::runtime_error(std::string("CUDA error: ") + cudaGetErrorString(err)); \
        } \
    } while(0)

// Helper function to check tensor properties
void check_tensor(const torch::Tensor& t, const std::string& name, 
                  torch::ScalarType dtype, int device_id) {
    if (!t.is_cuda()) {
        throw std::runtime_error(name + " must be a CUDA tensor");
    }
    if (t.device().index() != device_id) {
        throw std::runtime_error(name + " must be on the same CUDA device");
    }
    if (t.scalar_type() != dtype) {
        throw std::runtime_error(name + " must have the correct dtype");
    }
    if (!t.is_contiguous()) {
        throw std::runtime_error(name + " must be contiguous");
    }
}

// Main run function that returns a dictionary
py::dict run(
    torch::Tensor q_nope,
    torch::Tensor q_pe,
    torch::Tensor ckv_cache,
    torch::Tensor kpe_cache,
    torch::Tensor kv_indptr,
    torch::Tensor kv_indices,
    float sm_scale
) {
    // Get tensor dimensions
    const int batch_size = q_nope.size(0);
    const int num_qo_heads = q_nope.size(1);
    const int head_dim_ckv = q_nope.size(2);
    const int head_dim_kpe = q_pe.size(2);
    const int page_size = ckv_cache.size(1);
    const int len_indptr = kv_indptr.size(0);
    const int num_kv_indices = kv_indices.size(0);
    
    // Verify constants match specification
    if (num_qo_heads != NUM_QO_HEADS) {
        throw std::runtime_error("num_qo_heads must be 16");
    }
    if (head_dim_ckv != HEAD_DIM_CKV) {
        throw std::runtime_error("head_dim_ckv must be 512");
    }
    if (head_dim_kpe != HEAD_DIM_KPE) {
        throw std::runtime_error("head_dim_kpe must be 64");
    }
    if (page_size != PAGE_SIZE) {
        throw std::runtime_error("page_size must be 1");
    }
    
    // Verify constraints
    if (len_indptr != batch_size + 1) {
        throw std::runtime_error("len_indptr must equal batch_size + 1");
    }
    
    // Verify num_kv_indices constraint
    torch::Tensor last_indptr = kv_indptr.index({-1});
    int expected_num_indices = last_indptr.item<int>();
    if (num_kv_indices != expected_num_indices) {
        throw std::runtime_error("num_kv_indices must equal kv_indptr[-1]");
    }
    
    // Get device
    int device_id = q_nope.device().index();
    CHECK_CUDA(cudaSetDevice(device_id));
    
    // Check all tensors are on the same device and have correct properties
    check_tensor(q_nope, "q_nope", torch::kBFloat16, device_id);
    check_tensor(q_pe, "q_pe", torch::kBFloat16, device_id);
    check_tensor(ckv_cache, "ckv_cache", torch::kBFloat16, device_id);
    check_tensor(kpe_cache, "kpe_cache", torch::kBFloat16, device_id);
    check_tensor(kv_indptr, "kv_indptr", torch::kInt32, device_id);
    check_tensor(kv_indices, "kv_indices", torch::kInt32, device_id);
    
    // Ensure tensors are contiguous
    q_nope = q_nope.contiguous();
    q_pe = q_pe.contiguous();
    ckv_cache = ckv_cache.contiguous();
    kpe_cache = kpe_cache.contiguous();
    kv_indptr = kv_indptr.contiguous();
    kv_indices = kv_indices.contiguous();
    
    // Create output tensors
    auto options_bf16 = torch::TensorOptions()
        .dtype(torch::kBFloat16)
        .device(torch::kCUDA, device_id);
    auto options_f32 = torch::TensorOptions()
        .dtype(torch::kFloat32)
        .device(torch::kCUDA, device_id);
    
    torch::Tensor output = torch::zeros({batch_size, num_qo_heads, head_dim_ckv}, options_bf16);
    torch::Tensor lse = torch::full({batch_size, num_qo_heads}, 
                                   -std::numeric_limits<float>::infinity(), options_f32);
    
    // For page_size=1, we can treat the cache as 2D instead of 3D
    torch::Tensor ckv_cache_2d = ckv_cache.view({-1, head_dim_ckv});
    torch::Tensor kpe_cache_2d = kpe_cache.view({-1, head_dim_kpe});
    
    // Get data pointers
    const __nv_bfloat16* q_nope_ptr = reinterpret_cast<const __nv_bfloat16*>(q_nope.data_ptr());
    const __nv_bfloat16* q_pe_ptr = reinterpret_cast<const __nv_bfloat16*>(q_pe.data_ptr());
    const __nv_bfloat16* ckv_cache_ptr = reinterpret_cast<const __nv_bfloat16*>(ckv_cache_2d.data_ptr());
    const __nv_bfloat16* kpe_cache_ptr = reinterpret_cast<const __nv_bfloat16*>(kpe_cache_2d.data_ptr());
    const int32_t* kv_indptr_ptr = kv_indptr.data_ptr<int32_t>();
    const int32_t* kv_indices_ptr = kv_indices.data_ptr<int32_t>();
    __nv_bfloat16* output_ptr = reinterpret_cast<__nv_bfloat16*>(output.data_ptr());
    float* lse_ptr = lse.data_ptr<float>();
    
    // Get current CUDA stream
    cudaStream_t stream = at::cuda::getCurrentCUDAStream();
    
    // Launch kernel
    launch_mla_paged_decode(
        q_nope_ptr, q_pe_ptr, ckv_cache_ptr, kpe_cache_ptr,
        kv_indptr_ptr, kv_indices_ptr, sm_scale,
        output_ptr, lse_ptr, batch_size, stream
    );
    
    // Check for errors
    CHECK_CUDA(cudaGetLastError());
    
    // Return dictionary with both outputs
    py::dict result;
    result["output"] = output;
    result["lse"] = lse;
    return result;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run", &run, "MLA paged decode kernel",
          py::arg("q_nope"),
          py::arg("q_pe"),
          py::arg("ckv_cache"),
          py::arg("kpe_cache"),
          py::arg("kv_indptr"),
          py::arg("kv_indices"),
          py::arg("sm_scale"));
}
scrolls · 157 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON