Skip to content
KernelIndex
Search⌘K

gpt-5 / cuda0b3c0a

gpt-5_cuda_0b3c0a · gpt-5-2025-08-07 · cuda · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-cuda-0b3c0a?include=source"
interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32

Benchmark evidence

21 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
46.7µs
#5 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
46.7µs
#6 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #5fa6fa
NVIDIA B200
46.7µs
#3 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
58.7µs
#4 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
58.8µs
#5 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #cb4b00
NVIDIA B200
61.7µs
#3 of 5
2025-10-19
NVIDIA B200
103.9µs
#7 of 10
2025-10-19
NVIDIA B200
104.2µs
#8 of 10
2025-10-19
NVIDIA B200
190.1µs
#3 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
219.6µs
#17 of 20
2025-10-19
Show all 21 measurements ›
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
219.7µs
#18 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
219.9µs
#19 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
220.0µs
#20 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #f3d59b
NVIDIA B200
224.0µs
#5 of 5
2025-10-19
NVIDIA B200
536.3µs
#5 of 5
2025-10-19
NVIDIA B200
658.0µs
#9 of 10
2025-10-19
NVIDIA B200
658.2µs
#10 of 10
2025-10-19
NVIDIA B200
15.1ms
#5 of 5
2025-10-19
NVIDIA B200
1.12s
#5 of 5
2025-10-19
NVIDIA B200
1.12s
#5 of 5
2025-10-19
NVIDIA B200
1.20s
#5 of 5
2025-10-19

Reproduction-ready · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:18f1a868373ce1a1d5496de6f50719b62017fc936c073df735d10aff364a45b9
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20

Kernel source

main.cpp151 lines
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include "kernel.h"

#include <vector>
#include <cmath>
#include <limits>
#include <cstdint>

// Basic shape/dtype checks (device handling is done later)
static inline void check_shapes_and_types(
    const torch::Tensor& q,
    const torch::Tensor& k,
    const torch::Tensor& v,
    const torch::Tensor& qo_indptr,
    const torch::Tensor& kv_indptr) {

  TORCH_CHECK(q.scalar_type() == at::kBFloat16, "q must be bfloat16");
  TORCH_CHECK(k.scalar_type() == at::kBFloat16, "k must be bfloat16");
  TORCH_CHECK(v.scalar_type() == at::kBFloat16, "v must be bfloat16");
  TORCH_CHECK(qo_indptr.scalar_type() == at::kInt, "qo_indptr must be int32");
  TORCH_CHECK(kv_indptr.scalar_type() == at::kInt, "kv_indptr must be int32");

  TORCH_CHECK(q.dim() == 3, "q must have shape [total_q, 32, 128]");
  TORCH_CHECK(k.dim() == 3, "k must have shape [total_kv, 8, 128]");
  TORCH_CHECK(v.dim() == 3, "v must have shape [total_kv, 8, 128]");
  TORCH_CHECK(q.size(1) == NUM_QO_HEADS && q.size(2) == HEAD_DIM,
              "q must have shape [*, 32, 128]");
  TORCH_CHECK(k.size(1) == NUM_KV_HEADS && k.size(2) == HEAD_DIM,
              "k must have shape [*, 8, 128]");
  TORCH_CHECK(v.size(1) == NUM_KV_HEADS && v.size(2) == HEAD_DIM,
              "v must have shape [*, 8, 128]");

  TORCH_CHECK(qo_indptr.dim() == 1, "qo_indptr must be 1D");
  TORCH_CHECK(kv_indptr.dim() == 1, "kv_indptr must be 1D");
  TORCH_CHECK(qo_indptr.size(0) == kv_indptr.size(0),
              "qo_indptr and kv_indptr must have the same length");
}

static inline void check_constraints_on_device(
    const torch::Tensor& qo_indptr,
    const torch::Tensor& kv_indptr,
    int64_t total_q,
    int64_t total_kv) {
  // Fetch last elements from device using a synchronous memcpy (simple, robust)
  int32_t qo_last = 0;
  int32_t kv_last = 0;
  const int32_t* qo_ptr = qo_indptr.data_ptr<int32_t>();
  const int32_t* kv_ptr = kv_indptr.data_ptr<int32_t>();
  CUDA_CALL(cudaMemcpy(&qo_last, qo_ptr + (qo_indptr.size(0) - 1),
                       sizeof(int32_t), cudaMemcpyDeviceToHost));
  CUDA_CALL(cudaMemcpy(&kv_last, kv_ptr + (kv_indptr.size(0) - 1),
                       sizeof(int32_t), cudaMemcpyDeviceToHost));

  TORCH_CHECK(qo_last == total_q,
              "Constraint failed: total_q must equal qo_indptr[-1] (got ",
              qo_last, " vs ", total_q, ")");
  TORCH_CHECK(kv_last == total_kv,
              "Constraint failed: total_kv must equal kv_indptr[-1] (got ",
              kv_last, " vs ", total_kv, ")");
}

std::vector<torch::Tensor> run(
    torch::Tensor q,           // [total_q, 32, 128], bfloat16
    torch::Tensor k,           // [total_kv, 8, 128], bfloat16
    torch::Tensor v,           // [total_kv, 8, 128], bfloat16
    torch::Tensor qo_indptr,   // [len_indptr], int32
    torch::Tensor kv_indptr,   // [len_indptr], int32
    c10::optional<double> sm_scale_opt /* = None */) {

  // Validate shapes and types first
  check_shapes_and_types(q, k, v, qo_indptr, kv_indptr);

  // Move to CUDA if needed (support both CPU/CUDA inputs)
  if (!q.is_cuda()) q = q.to(torch::kCUDA);
  if (!k.is_cuda()) k = k.to(torch::kCUDA);
  if (!v.is_cuda()) v = v.to(torch::kCUDA);
  if (!qo_indptr.is_cuda()) qo_indptr = qo_indptr.to(torch::kCUDA);
  if (!kv_indptr.is_cuda()) kv_indptr = kv_indptr.to(torch::kCUDA);

  // Ensure contiguity on device
  q = q.contiguous();
  k = k.contiguous();
  v = v.contiguous();
  qo_indptr = qo_indptr.contiguous();
  kv_indptr = kv_indptr.contiguous();

  const int64_t total_q = q.size(0);
  const int64_t total_kv = k.size(0);
  const int64_t len_indptr = qo_indptr.size(0);

  // Check constraints (on device)
  check_constraints_on_device(qo_indptr, kv_indptr, total_q, total_kv);

  // Default softmax scale
  float sm_scale;
  if (sm_scale_opt.has_value()) {
    sm_scale = static_cast<float>(sm_scale_opt.value());
  } else {
    sm_scale = 1.0f / std::sqrt(static_cast<float>(HEAD_DIM)); // 1/sqrt(128)
  }

  // Allocate outputs on device: output zeros, lse -inf
  auto out_opts = q.options().dtype(at::kBFloat16);
  auto lse_opts = q.options().dtype(at::kFloat);
  torch::Tensor out = torch::zeros({total_q, NUM_QO_HEADS, HEAD_DIM}, out_opts);
  torch::Tensor lse = torch::full({total_q, NUM_QO_HEADS},
                                  -std::numeric_limits<float>::infinity(),
                                  lse_opts);

  // Raw pointers (16-bit storage)
  const uint16_t* q_ptr = reinterpret_cast<const uint16_t*>(q.data_ptr<at::BFloat16>());
  const uint16_t* k_ptr = reinterpret_cast<const uint16_t*>(k.data_ptr<at::BFloat16>());
  const uint16_t* v_ptr = reinterpret_cast<const uint16_t*>(v.data_ptr<at::BFloat16>());

  const int32_t* qo_ptr = qo_indptr.data_ptr<int32_t>();
  const int32_t* kv_ptr = kv_indptr.data_ptr<int32_t>();

  uint16_t* out_ptr = reinterpret_cast<uint16_t*>(out.data_ptr<at::BFloat16>());
  float* lse_ptr = lse.data_ptr<float>();

  // Current CUDA stream
  cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();

  // Launch kernel
  gqa_ragged_prefill_causal_h32_kv8_d128_launcher(
      q_ptr, k_ptr, v_ptr,
      qo_ptr, kv_ptr,
      static_cast<int32_t>(len_indptr),
      static_cast<int32_t>(total_q),
      static_cast<int32_t>(total_kv),
      sm_scale,
      out_ptr, lse_ptr,
      stream);

  return {out, lse};
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("run", &run,
        pybind11::arg("q"),
        pybind11::arg("k"),
        pybind11::arg("v"),
        pybind11::arg("qo_indptr"),
        pybind11::arg("kv_indptr"),
        pybind11::arg("sm_scale") = pybind11::none(),
        "gqa_ragged_prefill_causal_h32_kv8_d128 kernel: "
        "Computes ragged causal GQA prefill attention with 32 QO heads, 8 KV heads, head_dim=128. "
        "Inputs: q [total_q,32,128] bf16, k/v [total_kv,8,128] bf16, indptr arrays int32. "
        "Outputs: output bf16 [total_q,32,128], lse float [total_q,32] in base-2.");
}
scrolls · 151 lines total

Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0

Best evidence level for this revision: reproducible

JSON