Skip to content
KernelIndex
Search⌘K

gpt-5 / cuda00b2dd

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

96 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=17 · num_kv_indices=2
NVIDIA B200
30.7µs
#3 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=17 · num_kv_indices=2
NVIDIA B200
31.1µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=8 · num_kv_indices=7
NVIDIA B200
33.2µs
#4 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=8 · num_kv_indices=7
NVIDIA B200
35.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=10 · num_kv_indices=9
NVIDIA B200
35.5µs
#3 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=10 · num_kv_indices=9
NVIDIA B200
38.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=12 · num_kv_indices=11
NVIDIA B200
41.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=12 · num_kv_indices=11
NVIDIA B200
42.3µs
#3 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
43.4µs
#4 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
46.5µs
#4 of 7
2025-10-16
Show all 96 measurements ›
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=57 · num_kv_indices=40
NVIDIA B200
67.6µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=57 · num_kv_indices=40
NVIDIA B200
68.6µs
#5 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
78.4µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
80.2µs
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=81 · num_kv_indices=64
NVIDIA B200
93.6µs
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=81 · num_kv_indices=64
NVIDIA B200
93.6µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9317 · num_kv_indices=72
NVIDIA B200
105.0µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9317 · num_kv_indices=72
NVIDIA B200
105.1µs
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9332 · num_kv_indices=87
NVIDIA B200
117.3µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9332 · num_kv_indices=87
NVIDIA B200
117.8µs
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9347 · num_kv_indices=102
NVIDIA B200
127.1µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9347 · num_kv_indices=102
NVIDIA B200
135.4µs
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=191 · num_kv_indices=141
NVIDIA B200
175.2µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=191 · num_kv_indices=141
NVIDIA B200
178.4µs
#6 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=302 · num_kv_indices=252
NVIDIA B200
291.5µs
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=302 · num_kv_indices=252
NVIDIA B200
292.1µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=1070 · num_kv_indices=1020
NVIDIA B200
351.3µs
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=1070 · num_kv_indices=1020
NVIDIA B200
352.0µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=412 · num_kv_indices=362
NVIDIA B200
375.0µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=412 · num_kv_indices=362
NVIDIA B200
410.3µs
#4 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=2158 · num_kv_indices=2108
NVIDIA B200
413.4µs
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=2158 · num_kv_indices=2108
NVIDIA B200
416.6µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=3246 · num_kv_indices=3196
NVIDIA B200
476.6µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=3246 · num_kv_indices=3196
NVIDIA B200
479.6µs
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=486 · num_kv_indices=436
NVIDIA B200
483.2µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=486 · num_kv_indices=436
NVIDIA B200
486.7µs
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=4334 · num_kv_indices=4284
NVIDIA B200
539.5µs
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=4334 · num_kv_indices=4284
NVIDIA B200
547.2µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=596 · num_kv_indices=546
NVIDIA B200
603.0µs
#5 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=596 · num_kv_indices=546
NVIDIA B200
603.3µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=5422 · num_kv_indices=5372
NVIDIA B200
608.6µs
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=5422 · num_kv_indices=5372
NVIDIA B200
608.6µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=6510 · num_kv_indices=6460
NVIDIA B200
676.9µs
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=6510 · num_kv_indices=6460
NVIDIA B200
686.6µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=7598 · num_kv_indices=7548
NVIDIA B200
732.9µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=7598 · num_kv_indices=7548
NVIDIA B200
736.8µs
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=8686 · num_kv_indices=8636
NVIDIA B200
793.4µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=8686 · num_kv_indices=8636
NVIDIA B200
798.4µs
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=24732 · num_kv_indices=15463
NVIDIA B200
2.83ms
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=28831 · num_kv_indices=28815
NVIDIA B200
2.89ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=25820 · num_kv_indices=16551
NVIDIA B200
2.89ms
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=35359 · num_kv_indices=35343
NVIDIA B200
2.95ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=31007 · num_kv_indices=30991
NVIDIA B200
3.03ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=28831 · num_kv_indices=28815
NVIDIA B200
3.06ms
#3 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=33183 · num_kv_indices=33167
NVIDIA B200
3.07ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=26908 · num_kv_indices=17639
NVIDIA B200
3.08ms
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=50655 · num_kv_indices=50639
NVIDIA B200
3.14ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=24732 · num_kv_indices=15463
NVIDIA B200
3.14ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=39711 · num_kv_indices=39695
NVIDIA B200
3.14ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=37535 · num_kv_indices=37519
NVIDIA B200
3.15ms
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=39711 · num_kv_indices=39695
NVIDIA B200
3.17ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=41887 · num_kv_indices=41871
NVIDIA B200
3.19ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=33183 · num_kv_indices=33167
NVIDIA B200
3.21ms
#4 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=44063 · num_kv_indices=44047
NVIDIA B200
3.22ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=35359 · num_kv_indices=35343
NVIDIA B200
3.22ms
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=29084 · num_kv_indices=19815
NVIDIA B200
3.23ms
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=57183 · num_kv_indices=57167
NVIDIA B200
3.25ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=37535 · num_kv_indices=37519
NVIDIA B200
3.25ms
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=48479 · num_kv_indices=48463
NVIDIA B200
3.26ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=46303 · num_kv_indices=46287
NVIDIA B200
3.26ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=44063 · num_kv_indices=44047
NVIDIA B200
3.28ms
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=41887 · num_kv_indices=41871
NVIDIA B200
3.28ms
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=25820 · num_kv_indices=16551
NVIDIA B200
3.30ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=31007 · num_kv_indices=30991
NVIDIA B200
3.30ms
#4 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=30172 · num_kv_indices=20903
NVIDIA B200
3.30ms
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=27996 · num_kv_indices=18727
NVIDIA B200
3.33ms
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=31260 · num_kv_indices=21991
NVIDIA B200
3.35ms
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=50655 · num_kv_indices=50639
NVIDIA B200
3.38ms
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=48479 · num_kv_indices=48463
NVIDIA B200
3.40ms
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=59359 · num_kv_indices=59343
NVIDIA B200
3.40ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=31260 · num_kv_indices=21991
NVIDIA B200
3.40ms
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=26908 · num_kv_indices=17639
NVIDIA B200
3.41ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=32348 · num_kv_indices=23079
NVIDIA B200
3.41ms
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=55007 · num_kv_indices=54991
NVIDIA B200
3.41ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=57183 · num_kv_indices=57167
NVIDIA B200
3.42ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=52831 · num_kv_indices=52815
NVIDIA B200
3.43ms
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=32348 · num_kv_indices=23079
NVIDIA B200
3.43ms
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=52831 · num_kv_indices=52815
NVIDIA B200
3.45ms
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=27996 · num_kv_indices=18727
NVIDIA B200
3.47ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=46303 · num_kv_indices=46287
NVIDIA B200
3.47ms
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=59359 · num_kv_indices=59343
NVIDIA B200
3.48ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=29084 · num_kv_indices=19815
NVIDIA B200
3.50ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=61535 · num_kv_indices=61519
NVIDIA B200
3.52ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=55007 · num_kv_indices=54991
NVIDIA B200
3.55ms
#3 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=61535 · num_kv_indices=61519
NVIDIA B200
3.56ms
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=30172 · num_kv_indices=20903
NVIDIA B200
3.59ms
#7 of 7
2025-10-16

Reproduction-ready · How evidence levels are derived →

Source and license

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

Kernel source

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

namespace py = pybind11;

static inline void check_inputs(
    const at::Tensor& q,
    const at::Tensor& k_cache,
    const at::Tensor& v_cache,
    const at::Tensor& kv_indptr,
    const at::Tensor& kv_indices) {

  TORCH_CHECK(q.is_cuda(), "q must be a CUDA tensor");
  TORCH_CHECK(k_cache.is_cuda(), "k_cache must be a CUDA tensor");
  TORCH_CHECK(v_cache.is_cuda(), "v_cache must be a CUDA tensor");
  TORCH_CHECK(kv_indptr.is_cuda(), "kv_indptr must be a CUDA tensor");
  TORCH_CHECK(kv_indices.is_cuda(), "kv_indices must be a CUDA tensor");

  TORCH_CHECK(q.dtype() == at::kBFloat16, "q must be bfloat16");
  TORCH_CHECK(k_cache.dtype() == at::kBFloat16, "k_cache must be bfloat16");
  TORCH_CHECK(v_cache.dtype() == at::kBFloat16, "v_cache must be bfloat16");
  TORCH_CHECK(kv_indptr.dtype() == at::kInt, "kv_indptr must be int32");
  TORCH_CHECK(kv_indices.dtype() == at::kInt, "kv_indices must be int32");

  TORCH_CHECK(q.dim() == 3, "q must have shape [B, 32, 128]");
  TORCH_CHECK(q.size(1) == GQA_NUM_QO_HEADS, "q.num_qo_heads must be 32");
  TORCH_CHECK(q.size(2) == GQA_HEAD_DIM, "q.head_dim must be 128");

  TORCH_CHECK(k_cache.dim() == 4, "k_cache must have shape [num_pages, 1, 4, 128]");
  TORCH_CHECK(v_cache.dim() == 4, "v_cache must have shape [num_pages, 1, 4, 128]");
  TORCH_CHECK(k_cache.size(1) == GQA_PAGE_SIZE, "page_size must be 1 for k_cache");
  TORCH_CHECK(v_cache.size(1) == GQA_PAGE_SIZE, "page_size must be 1 for v_cache");
  TORCH_CHECK(k_cache.size(2) == GQA_NUM_KV_HEADS, "k_cache.num_kv_heads must be 4");
  TORCH_CHECK(v_cache.size(2) == GQA_NUM_KV_HEADS, "v_cache.num_kv_heads must be 4");
  TORCH_CHECK(k_cache.size(3) == GQA_HEAD_DIM, "k_cache.head_dim must be 128");
  TORCH_CHECK(v_cache.size(3) == GQA_HEAD_DIM, "v_cache.head_dim must be 128");

  TORCH_CHECK(kv_indptr.dim() == 1, "kv_indptr must be 1D");
  TORCH_CHECK(kv_indices.dim() == 1, "kv_indices must be 1D");

  int64_t batch_size = q.size(0);
  TORCH_CHECK(kv_indptr.size(0) == batch_size + 1, "len_indptr must equal batch_size + 1");

  // Verify num_kv_indices == kv_indptr[-1].
  // Copy small vector to CPU to avoid tricky GPU indexing in C++ frontend.
  at::Tensor kv_indptr_cpu = kv_indptr.to(torch::kCPU);
  int32_t last_offset = kv_indptr_cpu.data_ptr<int32_t>()[kv_indptr_cpu.numel() - 1];
  TORCH_CHECK(kv_indices.size(0) == static_cast<int64_t>(last_offset),
              "num_kv_indices must equal kv_indptr[-1]");
}

std::tuple<at::Tensor, at::Tensor> run(
    at::Tensor q,             // [B, 32, 128] bfloat16
    at::Tensor k_cache,       // [P, 1, 4, 128] bfloat16
    at::Tensor v_cache,       // [P, 1, 4, 128] bfloat16
    at::Tensor kv_indptr,     // [B+1] int32
    at::Tensor kv_indices,    // [kv_indptr[-1]] int32
    c10::optional<double> sm_scale_opt // optional float
) {
  check_inputs(q, k_cache, v_cache, kv_indptr, kv_indices);

  // Make tensors contiguous for predictable addressing
  q = q.contiguous();
  k_cache = k_cache.contiguous();
  v_cache = v_cache.contiguous();
  kv_indptr = kv_indptr.contiguous();
  kv_indices = kv_indices.contiguous();

  int64_t batch_size = q.size(0);
  int64_t num_pages = k_cache.size(0);
  // Default sm_scale = 1/sqrt(128)
  float sm_scale = sm_scale_opt.has_value() ? static_cast<float>(sm_scale_opt.value())
                                            : static_cast<float>(1.0 / std::sqrt(static_cast<double>(GQA_HEAD_DIM)));

  // Allocate outputs
  auto out = torch::empty({batch_size, (int64_t)GQA_NUM_QO_HEADS, (int64_t)GQA_HEAD_DIM},
                          q.options()); // bfloat16
  auto lse = torch::empty({batch_size, (int64_t)GQA_NUM_QO_HEADS},
                          q.options().dtype(torch::kFloat32)); // float32

  // Launch kernel
  auto stream = at::cuda::getCurrentCUDAStream();

  const __nv_bfloat16* q_ptr = reinterpret_cast<const __nv_bfloat16*>(q.data_ptr<at::BFloat16>());
  const __nv_bfloat16* k_ptr = reinterpret_cast<const __nv_bfloat16*>(k_cache.data_ptr<at::BFloat16>());
  const __nv_bfloat16* v_ptr = reinterpret_cast<const __nv_bfloat16*>(v_cache.data_ptr<at::BFloat16>());
  const int32_t* indptr_ptr = kv_indptr.data_ptr<int32_t>();
  const int32_t* indices_ptr = kv_indices.data_ptr<int32_t>();
  __nv_bfloat16* out_ptr = reinterpret_cast<__nv_bfloat16*>(out.data_ptr<at::BFloat16>());
  float* lse_ptr = lse.data_ptr<float>();

  gqa_paged_decode_h32_kv4_d128_ps1_launcher(
      q_ptr,
      k_ptr,
      v_ptr,
      indptr_ptr,
      indices_ptr,
      static_cast<int>(batch_size),
      static_cast<int>(num_pages),
      sm_scale,
      out_ptr,
      lse_ptr,
      stream);

  // Ensure kernel launch was successful
  auto err = cudaGetLastError();
  TORCH_CHECK(err == cudaSuccess, "Kernel launch failed: ", cudaGetErrorString(err));

  return {out, lse};
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("run", &run,
        py::arg("q"),
        py::arg("k_cache"),
        py::arg("v_cache"),
        py::arg("kv_indptr"),
        py::arg("kv_indices"),
        py::arg("sm_scale") = py::none(),
        "GQA paged decode kernel for (h=32, kv=4, d=128, page_size=1).");
}
scrolls · 124 lines total

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

Best evidence level for this revision: reproducible

JSON