gpt-5 / cuda99c7ef
gpt-5_cuda_99c7ef · gpt-5-2025-08-07 · cuda · Apache-2.0
Kernel source · 187 lines ↓holds 1 record
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 187 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-cuda-99c7ef?include=source"interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
2 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reproduction-ready · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:10679b43171eecbc036d217847c077bfe8ac015459c1e1c8404e090f00e62a36
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20
Kernel source
main.cpp187 lines
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime_api.h>
#include <vector>
#include <limits>
#include <stdexcept>
#include <cmath>
#include "kernel.h"
static inline void check_tensor(const torch::Tensor& t, at::ScalarType dtype, const std::string& name, bool on_cuda = true) {
if (on_cuda && !t.is_cuda()) {
throw std::runtime_error(name + " must be a CUDA tensor");
}
if (!on_cuda && t.is_cuda()) {
throw std::runtime_error(name + " must be a CPU tensor");
}
if (t.scalar_type() != dtype) {
throw std::runtime_error(name + " has wrong dtype");
}
if (!t.is_contiguous()) {
throw std::runtime_error(name + " must be contiguous");
}
}
static inline void check_shapes(
const torch::Tensor& q,
const torch::Tensor& k_cache,
const torch::Tensor& v_cache) {
if (q.dim() != 3) {
throw std::runtime_error("q must have shape [total_q, 32, 128]");
}
if (q.size(1) != HEADS_QO || q.size(2) != HEAD_DIM) {
throw std::runtime_error("q shape mismatch: expected [*, 32, 128]");
}
if (k_cache.dim() != 4 || v_cache.dim() != 4) {
throw std::runtime_error("k_cache/v_cache must have shape [num_pages, 1, 4, 128]");
}
if (k_cache.size(1) != 1 || v_cache.size(1) != 1 ||
k_cache.size(2) != HEADS_KV || v_cache.size(2) != HEADS_KV ||
k_cache.size(3) != HEAD_DIM || v_cache.size(3) != HEAD_DIM) {
throw std::runtime_error("k_cache/v_cache shape mismatch: expected [num_pages, 1, 4, 128]");
}
}
static inline void verify_constraints(
const torch::Tensor& qo_indptr,
const torch::Tensor& kv_indptr,
const torch::Tensor& kv_indices,
int64_t total_q) {
// Move to CPU for robust access
auto qo_cpu = qo_indptr.to(torch::kCPU, /*non_blocking=*/false);
auto kv_cpu = kv_indptr.to(torch::kCPU, /*non_blocking=*/false);
const int32_t qo_last = qo_cpu.data_ptr<int32_t>()[qo_cpu.size(0) - 1];
if (qo_last != total_q) {
throw std::runtime_error("Constraint failed: total_q == qo_indptr[-1]");
}
const int32_t kv_last = kv_cpu.data_ptr<int32_t>()[kv_cpu.size(0) - 1];
const auto num_kv_indices = static_cast<int32_t>(kv_indices.size(0));
if (kv_last != num_kv_indices) {
throw std::runtime_error("Constraint failed: num_kv_indices == kv_indptr[-1]");
}
}
// Prepare per-token sequence index and position vectors on CPU
static inline void build_q_token_maps_cpu(
const torch::Tensor& qo_indptr_cpu, // int32 CPU tensor [B+1]
std::vector<int32_t>& q_seq_index,
std::vector<int32_t>& q_seq_pos) {
const int32_t* ind = qo_indptr_cpu.data_ptr<int32_t>();
const int64_t len = qo_indptr_cpu.size(0);
const int32_t total_q = ind[len - 1];
q_seq_index.resize(total_q);
q_seq_pos.resize(total_q);
for (int64_t b = 0; b + 1 < len; ++b) {
const int32_t start = ind[b];
const int32_t end = ind[b + 1];
for (int32_t i = start; i < end; ++i) {
q_seq_index[i] = static_cast<int32_t>(b);
q_seq_pos[i] = i - start;
}
}
}
// Entry point called from Python
std::tuple<torch::Tensor, torch::Tensor> run(
torch::Tensor q, // [total_q, 32, 128] bfloat16
torch::Tensor k_cache, // [num_pages, 1, 4, 128] bfloat16
torch::Tensor v_cache, // [num_pages, 1, 4, 128] bfloat16
torch::Tensor qo_indptr, // [len_indptr] int32
torch::Tensor kv_indptr, // [len_indptr] int32
torch::Tensor kv_indices, // [num_kv_indices] int32
double sm_scale_opt = -1.0 // default -> 1/sqrt(128)
) {
// Type checks and contiguity
check_tensor(q, at::kBFloat16, "q");
check_tensor(k_cache, at::kBFloat16, "k_cache");
check_tensor(v_cache, at::kBFloat16, "v_cache");
check_tensor(qo_indptr, at::kInt, "qo_indptr");
check_tensor(kv_indptr, at::kInt, "kv_indptr");
check_tensor(kv_indices, at::kInt, "kv_indices");
check_shapes(q, k_cache, v_cache);
const int64_t total_q = q.size(0);
const int64_t num_pages = k_cache.size(0);
const int64_t len_indptr = qo_indptr.size(0);
verify_constraints(qo_indptr, kv_indptr, kv_indices, total_q);
// Default softmax scale if not provided
float sm_scale = static_cast<float>(sm_scale_opt);
if (sm_scale_opt < 0.0) {
sm_scale = 1.0f / std::sqrt(static_cast<float>(HEAD_DIM));
}
// Allocate outputs
auto opts_bf16 = torch::TensorOptions().dtype(torch::kBFloat16).device(q.device());
auto opts_f32 = torch::TensorOptions().dtype(torch::kFloat32).device(q.device());
torch::Tensor output = torch::zeros({total_q, HEADS_QO, HEAD_DIM}, opts_bf16);
torch::Tensor lse = torch::full({total_q, HEADS_QO}, -std::numeric_limits<float>::infinity(), opts_f32);
// Build q_seq_index and q_seq_pos on CPU
torch::Tensor qo_indptr_cpu = qo_indptr.to(torch::kCPU, /*non_blocking=*/false);
std::vector<int32_t> h_q_seq_index;
std::vector<int32_t> h_q_seq_pos;
build_q_token_maps_cpu(qo_indptr_cpu, h_q_seq_index, h_q_seq_pos);
// Move maps to device
auto opts_i32_cpu = torch::TensorOptions().dtype(torch::kInt32).device(torch::kCPU);
torch::Tensor q_seq_index_cpu = torch::from_blob(h_q_seq_index.data(), {total_q}, opts_i32_cpu).clone();
torch::Tensor q_seq_pos_cpu = torch::from_blob(h_q_seq_pos.data(), {total_q}, opts_i32_cpu).clone();
torch::Tensor q_seq_index_dev = q_seq_index_cpu.to(q.device(), /*non_blocking=*/false);
torch::Tensor q_seq_pos_dev = q_seq_pos_cpu.to(q.device(), /*non_blocking=*/false);
// Launch kernel
cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
const __nv_bfloat16* q_ptr = reinterpret_cast<const __nv_bfloat16*>(q.data_ptr<c10::BFloat16>());
const __nv_bfloat16* k_ptr = reinterpret_cast<const __nv_bfloat16*>(k_cache.data_ptr<c10::BFloat16>());
const __nv_bfloat16* v_ptr = reinterpret_cast<const __nv_bfloat16*>(v_cache.data_ptr<c10::BFloat16>());
const int32_t* qo_indptr_ptr = qo_indptr.data_ptr<int32_t>();
const int32_t* kv_indptr_ptr = kv_indptr.data_ptr<int32_t>();
const int32_t* kv_indices_ptr = kv_indices.data_ptr<int32_t>();
const int32_t* q_seq_index_ptr = q_seq_index_dev.data_ptr<int32_t>();
const int32_t* q_seq_pos_ptr = q_seq_pos_dev.data_ptr<int32_t>();
__nv_bfloat16* out_ptr = reinterpret_cast<__nv_bfloat16*>(output.data_ptr<c10::BFloat16>());
float* lse_ptr = lse.data_ptr<float>();
launch_gqa_paged_prefill_causal_h32_kv4_d128_ps1(
q_ptr, k_ptr, v_ptr,
qo_indptr_ptr, kv_indptr_ptr, kv_indices_ptr,
q_seq_index_ptr, q_seq_pos_ptr,
static_cast<int32_t>(len_indptr),
static_cast<int32_t>(total_q),
static_cast<int32_t>(num_pages),
sm_scale,
out_ptr, lse_ptr,
stream
);
return std::make_tuple(output, lse);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("run", &run,
py::arg("q"),
py::arg("k_cache"),
py::arg("v_cache"),
py::arg("qo_indptr"),
py::arg("kv_indptr"),
py::arg("kv_indices"),
py::arg("sm_scale") = -1.0,
"GQA paged prefill causal attention (h32, kv4, d128, ps1) optimized for B200.\n"
"Inputs:\n"
" q: [total_q, 32, 128] (bfloat16)\n"
" k_cache: [num_pages, 1, 4, 128] (bfloat16)\n"
" v_cache: [num_pages, 1, 4, 128] (bfloat16)\n"
" qo_indptr: [len_indptr] int32\n"
" kv_indptr: [len_indptr] int32\n"
" kv_indices: [num_kv_indices] int32\n"
" sm_scale: float (default 1/sqrt(128))\n"
"Outputs:\n"
" output: [total_q, 32, 128] (bfloat16)\n"
" lse: [total_q, 32] (float32), 2-based log-sum-exp");
}scrolls · 187 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reproducible
JSON