claude-opus-4-1 / cudab3cd9c
claude-opus-4-1_cuda_b3cd9c · claude-opus-4-1-20250805 · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 139 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-cuda-b3cd9c?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:846b75cc0a71a24dd399103cba72ae2af5eb0eca5cfff182dd54e9311569aa1e
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.cpp139 lines
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <vector>
#include <stdexcept>
#include <cmath>
#include "kernel.h"
// 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_index) {
if (!t.is_cuda()) {
throw std::runtime_error(name + " must be a CUDA tensor");
}
if (t.device().index() != device_index) {
throw std::runtime_error(name + " must be on the same CUDA device");
}
if (t.scalar_type() != dtype) {
throw std::runtime_error(name + " has incorrect dtype");
}
if (!t.is_contiguous()) {
throw std::runtime_error(name + " must be contiguous");
}
}
std::tuple<torch::Tensor, torch::Tensor> run(
torch::Tensor q,
torch::Tensor k,
torch::Tensor v,
torch::Tensor qo_indptr,
torch::Tensor kv_indptr,
float sm_scale = -1.0f
) {
// Get device
int device_index = q.device().index();
cudaSetDevice(device_index);
// Check input tensors
check_tensor(q, "q", torch::kBFloat16, device_index);
check_tensor(k, "k", torch::kBFloat16, device_index);
check_tensor(v, "v", torch::kBFloat16, device_index);
check_tensor(qo_indptr, "qo_indptr", torch::kInt32, device_index);
check_tensor(kv_indptr, "kv_indptr", torch::kInt32, device_index);
// Get dimensions
int64_t total_q = q.size(0);
int64_t num_qo_heads = q.size(1);
int64_t head_dim = q.size(2);
int64_t total_kv = k.size(0);
int64_t num_kv_heads = k.size(1);
int64_t len_indptr = qo_indptr.size(0);
// Validate dimensions
if (num_qo_heads != NUM_QO_HEADS) {
throw std::runtime_error("num_qo_heads must be 32, got " + std::to_string(num_qo_heads));
}
if (num_kv_heads != NUM_KV_HEADS) {
throw std::runtime_error("num_kv_heads must be 8, got " + std::to_string(num_kv_heads));
}
if (head_dim != HEAD_DIM) {
throw std::runtime_error("head_dim must be 128, got " + std::to_string(head_dim));
}
// Validate K and V shapes
if (k.size(0) != total_kv || k.size(1) != num_kv_heads || k.size(2) != head_dim) {
throw std::runtime_error("K tensor has incorrect shape");
}
if (v.size(0) != total_kv || v.size(1) != num_kv_heads || v.size(2) != head_dim) {
throw std::runtime_error("V tensor has incorrect shape");
}
// Validate indptr shapes
if (kv_indptr.size(0) != len_indptr) {
throw std::runtime_error("kv_indptr and qo_indptr must have the same length");
}
// Set default sm_scale if not provided
if (sm_scale < 0.0f) {
sm_scale = 1.0f / std::sqrt(static_cast<float>(head_dim));
}
// Allocate output tensors
auto options_bf16 = torch::TensorOptions()
.dtype(torch::kBFloat16)
.device(q.device());
auto options_f32 = torch::TensorOptions()
.dtype(torch::kFloat32)
.device(q.device());
torch::Tensor output = torch::zeros({total_q, num_qo_heads, head_dim}, options_bf16);
torch::Tensor lse = torch::full({total_q, num_qo_heads},
-std::numeric_limits<float>::infinity(), options_f32);
// Get CUDA stream
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
// Launch kernel
launch_gqa_ragged_prefill(
reinterpret_cast<const __nv_bfloat16*>(q.data_ptr()),
reinterpret_cast<const __nv_bfloat16*>(k.data_ptr()),
reinterpret_cast<const __nv_bfloat16*>(v.data_ptr()),
qo_indptr.data_ptr<int32_t>(),
kv_indptr.data_ptr<int32_t>(),
reinterpret_cast<__nv_bfloat16*>(output.data_ptr()),
lse.data_ptr<float>(),
sm_scale,
static_cast<int>(len_indptr),
static_cast<int>(total_q),
static_cast<int>(total_kv),
stream
);
// Check for errors
CHECK_CUDA(cudaGetLastError());
return std::make_tuple(output, lse);
}
// Python bindings
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("run", &run, "GQA ragged prefill causal attention",
pybind11::arg("q"),
pybind11::arg("k"),
pybind11::arg("v"),
pybind11::arg("qo_indptr"),
pybind11::arg("kv_indptr"),
pybind11::arg("sm_scale") = -1.0f);
}scrolls · 139 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON