claude-opus-4-1 / cuda4bd468
claude-opus-4-1_cuda_4bd468 · claude-opus-4-1-20250805 · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 119 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-cuda-4bd468?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:c00bdf588f42af34e614395151a560327e273e9df85cb5255f4af5752065e5e6
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.cpp119 lines
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <vector>
#include <stdexcept>
#include <cmath>
#include "kernel.h"
// Helper macros for CUDA error checking
#define CHECK_CUDA(x) TORCH_CHECK(x == cudaSuccess, "CUDA error: ", cudaGetErrorString(x))
#define CHECK_INPUT(x) TORCH_CHECK(x.is_cuda(), #x " must be a CUDA tensor")
#define CHECK_DTYPE(x, dtype) TORCH_CHECK(x.scalar_type() == dtype, #x " must have dtype " #dtype)
std::tuple<torch::Tensor, torch::Tensor> run(
torch::Tensor q,
torch::Tensor k_cache,
torch::Tensor v_cache,
torch::Tensor kv_indptr,
torch::Tensor kv_indices,
float sm_scale
) {
// Input validation
CHECK_INPUT(q);
CHECK_INPUT(k_cache);
CHECK_INPUT(v_cache);
CHECK_INPUT(kv_indptr);
CHECK_INPUT(kv_indices);
CHECK_DTYPE(q, torch::kBFloat16);
CHECK_DTYPE(k_cache, torch::kBFloat16);
CHECK_DTYPE(v_cache, torch::kBFloat16);
CHECK_DTYPE(kv_indptr, torch::kInt32);
CHECK_DTYPE(kv_indices, torch::kInt32);
// Ensure contiguous memory layout
q = q.contiguous();
k_cache = k_cache.contiguous();
v_cache = v_cache.contiguous();
kv_indptr = kv_indptr.contiguous();
kv_indices = kv_indices.contiguous();
// Dimension checks
TORCH_CHECK(q.dim() == 3, "q must be 3D [batch_size, num_qo_heads, head_dim]");
TORCH_CHECK(k_cache.dim() == 4, "k_cache must be 4D [num_pages, page_size, num_kv_heads, head_dim]");
TORCH_CHECK(v_cache.dim() == 4, "v_cache must be 4D [num_pages, page_size, num_kv_heads, head_dim]");
TORCH_CHECK(kv_indptr.dim() == 1, "kv_indptr must be 1D");
TORCH_CHECK(kv_indices.dim() == 1, "kv_indices must be 1D");
// Extract dimensions
const int batch_size = q.size(0);
const int num_qo_heads = q.size(1);
const int head_dim = q.size(2);
const int page_size = k_cache.size(1);
const int num_kv_heads = k_cache.size(2);
const int len_indptr = kv_indptr.size(0);
// Verify constants
TORCH_CHECK(num_qo_heads == 32, "num_qo_heads must be 32, got ", num_qo_heads);
TORCH_CHECK(num_kv_heads == 4, "num_kv_heads must be 4, got ", num_kv_heads);
TORCH_CHECK(head_dim == 128, "head_dim must be 128, got ", head_dim);
TORCH_CHECK(page_size == 1, "page_size must be 1, got ", page_size);
// Verify shape consistency
TORCH_CHECK(k_cache.size(3) == head_dim, "k_cache head_dim mismatch");
TORCH_CHECK(v_cache.size(0) == k_cache.size(0), "v_cache num_pages mismatch");
TORCH_CHECK(v_cache.size(1) == page_size, "v_cache page_size mismatch");
TORCH_CHECK(v_cache.size(2) == num_kv_heads, "v_cache num_kv_heads mismatch");
TORCH_CHECK(v_cache.size(3) == head_dim, "v_cache head_dim mismatch");
// Verify constraints
TORCH_CHECK(len_indptr == batch_size + 1,
"len_indptr (", len_indptr, ") must equal batch_size + 1 (", batch_size + 1, ")");
// 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 output = torch::zeros({batch_size, num_qo_heads, head_dim},
q.options().dtype(torch::kBFloat16));
auto lse = torch::full({batch_size, num_qo_heads},
-std::numeric_limits<float>::infinity(),
q.options().dtype(torch::kFloat32));
// Get current CUDA stream
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
// Launch kernel
cudaError_t err = launch_gqa_paged_decode(
q.data_ptr(),
k_cache.data_ptr(),
v_cache.data_ptr(),
kv_indptr.data_ptr(),
kv_indices.data_ptr(),
output.data_ptr(),
lse.data_ptr(),
sm_scale,
batch_size,
stream
);
CHECK_CUDA(err);
CHECK_CUDA(cudaGetLastError());
return std::make_tuple(output, lse);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "GQA Paged Decode kernel optimized for B200 GPU";
m.def("run", &run,
"GQA Paged Decode with 32 query heads, 4 KV heads, 128 head dim, page size 1",
py::arg("q"),
py::arg("k_cache"),
py::arg("v_cache"),
py::arg("kv_indptr"),
py::arg("kv_indices"),
py::arg("sm_scale") = -1.0f);
}scrolls · 119 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON