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