Kernels
drbh commited on
Commit
57c3a10
·
1 Parent(s): 6b27dde

feat: generate and vendor flashinfer kernels

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +1 -0
  2. .gitignore +3 -0
  3. .make_markers/patch_applied +0 -0
  4. .make_markers/submodule_initialized +0 -0
  5. README.md +10 -0
  6. build.toml +224 -0
  7. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode.cu +197 -0
  8. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_config.inc +71 -0
  9. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_jit_pybind.cu +40 -0
  10. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_kernel.cu +13 -0
  11. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode.cu +197 -0
  12. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_config.inc +71 -0
  13. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_jit_pybind.cu +40 -0
  14. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_kernel.cu +13 -0
  15. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode.cu +197 -0
  16. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_config.inc +71 -0
  17. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_jit_pybind.cu +40 -0
  18. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_kernel.cu +13 -0
  19. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode.cu +197 -0
  20. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_config.inc +71 -0
  21. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_jit_pybind.cu +40 -0
  22. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_kernel.cu +13 -0
  23. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode.cu +197 -0
  24. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_config.inc +71 -0
  25. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_jit_pybind.cu +40 -0
  26. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_kernel.cu +13 -0
  27. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode.cu +197 -0
  28. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_config.inc +71 -0
  29. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_jit_pybind.cu +40 -0
  30. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_kernel.cu +13 -0
  31. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode.cu +197 -0
  32. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_config.inc +71 -0
  33. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_jit_pybind.cu +40 -0
  34. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_kernel.cu +13 -0
  35. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode.cu +197 -0
  36. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_config.inc +71 -0
  37. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_jit_pybind.cu +40 -0
  38. csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_kernel.cu +13 -0
  39. csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode.cu +197 -0
  40. csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_config.inc +71 -0
  41. csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_jit_pybind.cu +40 -0
  42. csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_kernel.cu +13 -0
  43. csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode.cu +197 -0
  44. csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_config.inc +71 -0
  45. csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_jit_pybind.cu +40 -0
  46. csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_kernel.cu +13 -0
  47. csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode.cu +197 -0
  48. csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_config.inc +71 -0
  49. csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_jit_pybind.cu +40 -0
  50. csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_kernel.cu +13 -0
.gitattributes CHANGED
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ *.so filter=lfs diff=lfs merge=lfs -text
.gitignore ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ .bak
2
+ __pycache__
3
+ result
.make_markers/patch_applied ADDED
File without changes
.make_markers/submodule_initialized ADDED
File without changes
README.md ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ tags:
4
+ - kernel
5
+ ---
6
+
7
+
8
+ This kernel is a work in progress and requires more work to correctly add all of the FlashInfer kernels.
9
+
10
+ Please see the [generate-source.md](generate-source.md) for instructions on how to generate the source files that are contained in this kernel.
build.toml ADDED
@@ -0,0 +1,224 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [general]
2
+ name = "flashinfer"
3
+ universal = false
4
+
5
+ [torch]
6
+ src = [
7
+ "torch-ext/torch_binding.cpp",
8
+ "torch-ext/torch_binding.h"
9
+ ]
10
+
11
+ [kernel.flashinfer]
12
+ backend = "cuda"
13
+ cuda-capabilities = [
14
+ # "7.0",
15
+ # "7.2",
16
+ "7.5",
17
+ "8.0",
18
+ "8.6",
19
+ "8.7",
20
+ "8.9",
21
+ "9.0",
22
+ # "10.0",
23
+ # "10.1",
24
+ # "11.8",
25
+ # "12.0"
26
+ ]
27
+ cuda-flags = [
28
+ "-O3",
29
+ "-std=c++17",
30
+ "--use_fast_math",
31
+ "--expt-relaxed-constexpr",
32
+ "--expt-extended-lambda",
33
+ "-DFLASHINFER_ENABLE_F16",
34
+ "-DFLASHINFER_ENABLE_BF16",
35
+ "-DFLASHINFER_ENABLE_FP8_E4M3",
36
+ "-DFLASHINFER_ENABLE_FP8_E5M2",
37
+ "-DNDEBUG",
38
+ ]
39
+ cxx-flags = [
40
+ "-DFLASHINFER_ENABLE_F16",
41
+ "-DFLASHINFER_ENABLE_BF16",
42
+ "-DFLASHINFER_ENABLE_FP8_E4M3",
43
+ "-DFLASHINFER_ENABLE_FP8_E5M2",
44
+ ]
45
+ include = [ "include" ]
46
+ depends = ["torch", "cutlass_3_8"]
47
+ src = [
48
+ "include/flashinfer/page.cuh",
49
+ "include/flashinfer/arch_condition.h",
50
+ "include/flashinfer/cubin_loader.h",
51
+ "include/flashinfer/permuted_smem.cuh",
52
+ "include/flashinfer/exception.h",
53
+ "include/flashinfer/trtllm/common.h",
54
+ "include/flashinfer/trtllm/fmha/fmhaRunnerParams.h",
55
+ "include/flashinfer/trtllm/fmha/decoder_impl_common.h",
56
+ "include/flashinfer/trtllm/fmha/decoder_params.h",
57
+ "include/flashinfer/trtllm/fmha/kernelParams.h",
58
+ "include/flashinfer/trtllm/fmha/gen_kernel_launcher.cuh",
59
+ "include/flashinfer/trtllm/fmha/fmhaKernels.cuh",
60
+ "include/flashinfer/trtllm/fmha/cubin/kernelMetaInfo.h",
61
+ "include/flashinfer/trtllm/fmha/fmhaRunner.cuh",
62
+ "include/flashinfer/trtllm/common/cudaTypeUtils.cuh",
63
+ "include/flashinfer/trtllm/common/cudaUtils.h",
64
+ "include/flashinfer/trtllm/common/cudaBf16Wrapper.h",
65
+ "include/flashinfer/trtllm/common/cudaFp8Utils.h",
66
+ "include/flashinfer/trtllm/common/cudaBf16Fallbacks.cuh",
67
+ "include/flashinfer/trtllm/fused_moe/RoutingKernel.cuh",
68
+ "include/flashinfer/trtllm/fused_moe/DevKernel.h",
69
+ "include/flashinfer/trtllm/fused_moe/RoutingKernel.h",
70
+ "include/flashinfer/trtllm/fused_moe/runner.h",
71
+ "include/flashinfer/trtllm/fused_moe/IntFastDiv.h",
72
+ "include/flashinfer/trtllm/fused_moe/RoutingKernelTopK.cuh",
73
+ "include/flashinfer/trtllm/batched_gemm/KernelRunner.h",
74
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/GemmGatedActOptions.h",
75
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/trtllm/gen/MmaDecl.h",
76
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/trtllm/gen/SfLayoutDecl.h",
77
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/trtllm/gen/DtypeDecl.h",
78
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/trtllm/gen/CommonUtils.h",
79
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/trtllm/gen/CudaKernelLauncher.h",
80
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/BatchedGemmOptions.h",
81
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/TmaDescriptor.h",
82
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/BatchedGemmEnums.h",
83
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/KernelParamsDecl.h",
84
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/KernelTraits.h",
85
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/KernelParams.h",
86
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/Enums.h",
87
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/KernelMetaInfo.h",
88
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/BatchedGemmInterface.h",
89
+ "include/flashinfer/trtllm/batched_gemm/trtllmGen_bmm_export/GemmOptions.h",
90
+ "include/flashinfer/trtllm/gemm/trtllmGen_gemm_export/trtllm/gen/MmaDecl.h",
91
+ "include/flashinfer/trtllm/gemm/trtllmGen_gemm_export/trtllm/gen/SfLayoutDecl.h",
92
+ "include/flashinfer/trtllm/gemm/trtllmGen_gemm_export/trtllm/gen/DtypeDecl.h",
93
+ "include/flashinfer/trtllm/gemm/trtllmGen_gemm_export/trtllm/gen/CommonUtils.h",
94
+ "include/flashinfer/trtllm/gemm/trtllmGen_gemm_export/trtllm/gen/CudaKernelLauncher.h",
95
+ "include/flashinfer/trtllm/gemm/trtllmGen_gemm_export/TmaDescriptor.h",
96
+ "include/flashinfer/trtllm/gemm/trtllmGen_gemm_export/KernelTraits.h",
97
+ "include/flashinfer/trtllm/gemm/trtllmGen_gemm_export/GemmInterface.h",
98
+ "include/flashinfer/trtllm/gemm/trtllmGen_gemm_export/KernelParams.h",
99
+ "include/flashinfer/trtllm/gemm/trtllmGen_gemm_export/Enums.h",
100
+ "include/flashinfer/trtllm/gemm/trtllmGen_gemm_export/GemmOptions.h",
101
+ "include/flashinfer/page.cuh",
102
+ "include/flashinfer/vec_dtypes.cuh",
103
+ "include/flashinfer/sampling.cuh",
104
+ "include/flashinfer/logging.h",
105
+ "include/flashinfer/fp16.h",
106
+ "include/flashinfer/attention_impl.cuh",
107
+ "include/flashinfer/allocator.h",
108
+ "include/flashinfer/utils.cuh",
109
+ "include/flashinfer/fastdiv.cuh",
110
+ "include/flashinfer/norm.cuh",
111
+ "include/flashinfer/layout.cuh",
112
+ "include/flashinfer/math.cuh",
113
+ "include/flashinfer/gemm/group_gemm_sm90.cuh",
114
+ "include/flashinfer/gemm/group_gemv.cuh",
115
+ "include/flashinfer/gemm/fp8_gemm_cutlass_template.h",
116
+ "include/flashinfer/gemm/group_gemm_mxfp4_groupwise_sm100.cuh",
117
+ "include/flashinfer/gemm/fp4_gemm_cutlass.h",
118
+ "include/flashinfer/gemm/fp4_gemm_cutlass_template.h",
119
+ "include/flashinfer/gemm/cutlass_gemm_configs.h",
120
+ "include/flashinfer/gemm/fp4_gemm_template_sm100.h",
121
+ "include/flashinfer/gemm/group_gemm_fp8_groupwise_sm100.cuh",
122
+ "include/flashinfer/gemm/gemm_groupwise_sm100.cuh",
123
+ "include/flashinfer/gemm/group_gemm_lora.cuh",
124
+ "include/flashinfer/gemm/group_gemm.cuh",
125
+ "include/flashinfer/gemm/fp8_gemm_cutlass.h",
126
+ "include/flashinfer/gemm/fp8_gemm_template_sm100.h",
127
+ "include/flashinfer/gemm/bmm_fp8.cuh",
128
+ "include/flashinfer/activation.cuh",
129
+ "include/flashinfer/semaphore_utils.cuh",
130
+ "include/flashinfer/quantization.cuh",
131
+ "include/flashinfer/mma.cuh",
132
+ "include/flashinfer/attention/variant_helper.cuh",
133
+ "include/flashinfer/attention/hopper.cuh",
134
+ "include/flashinfer/attention/prefill.cuh",
135
+ "include/flashinfer/attention/persistent_template.cuh",
136
+ "include/flashinfer/attention/mla_hopper.cuh",
137
+ "include/flashinfer/attention/pod.cuh",
138
+ "include/flashinfer/attention/variants.cuh",
139
+ "include/flashinfer/attention/decode.cuh",
140
+ "include/flashinfer/attention/heap.h",
141
+ "include/flashinfer/attention/state.cuh",
142
+ "include/flashinfer/attention/mask.cuh",
143
+ "include/flashinfer/attention/default_prefill_params.cuh",
144
+ "include/flashinfer/attention/cutlass_mla.cuh",
145
+ "include/flashinfer/attention/scheduler.cuh",
146
+ "include/flashinfer/attention/blackwell/kernel/fmha_options.hpp",
147
+ "include/flashinfer/attention/blackwell/kernel/gather_tensor.hpp",
148
+ "include/flashinfer/attention/blackwell/kernel/fmha_tile_scheduler.hpp",
149
+ "include/flashinfer/attention/blackwell/kernel/sm100_fmha_fwd_kernel_tma_warpspecialized.hpp",
150
+ "include/flashinfer/attention/blackwell/kernel/sm100_fmha_mla_tma_warpspecialized.hpp",
151
+ "include/flashinfer/attention/blackwell/kernel/sm100_mla_tile_scheduler.hpp",
152
+ "include/flashinfer/attention/blackwell/kernel/sm100_fmha_gen_kernel_warpspecialized.hpp",
153
+ "include/flashinfer/attention/blackwell/kernel/sm100_fmha_mla_reduction.hpp",
154
+ "include/flashinfer/attention/blackwell/plan.cuh",
155
+ "include/flashinfer/attention/blackwell/common/pow_2.hpp",
156
+ "include/flashinfer/attention/blackwell/collective/sm100_fmha_load_tma_warpspecialized.hpp",
157
+ "include/flashinfer/attention/blackwell/collective/sm100_fmha_fwd_epilogue_tma_warpspecialized.hpp",
158
+ "include/flashinfer/attention/blackwell/collective/sm100_fmha_gen_mainloop_warpspecialized.hpp",
159
+ "include/flashinfer/attention/blackwell/collective/sm100_fmha_load_cpasync_warpspecialized.hpp",
160
+ "include/flashinfer/attention/blackwell/collective/sm100_fmha_gen_epilogue_warpspecialized.hpp",
161
+ "include/flashinfer/attention/blackwell/collective/sm100_fmha_fwd_mainloop_tma_warpspecialized.hpp",
162
+ "include/flashinfer/attention/blackwell/collective/fmha_fusion.hpp",
163
+ "include/flashinfer/attention/blackwell/collective/fmha_common.hpp",
164
+ "include/flashinfer/attention/blackwell/device/sm100_mla.hpp",
165
+ "include/flashinfer/attention/blackwell/device/fmha.hpp",
166
+ "include/flashinfer/attention/blackwell/fmha_cutlass_sm100.cuh",
167
+ "include/flashinfer/attention/mla.cuh",
168
+ "include/flashinfer/attention/decode_mla_cute_sm80.cuh",
169
+ "include/flashinfer/attention/cascade.cuh",
170
+ "include/flashinfer/attention/hopper/variant_helper.cuh",
171
+ "include/flashinfer/attention/hopper/default_params.cuh",
172
+ "include/flashinfer/attention/hopper/attention_updater.cuh",
173
+ "include/flashinfer/attention/hopper/epilogue.cuh",
174
+ "include/flashinfer/attention/hopper/variants.cuh",
175
+ "include/flashinfer/attention/hopper/mainloop.cuh",
176
+ "include/flashinfer/attention/hopper/block_sparse_gather.cuh",
177
+ "include/flashinfer/attention/hopper/utils.cuh",
178
+ "include/flashinfer/attention/hopper/prefill_sm90.cuh",
179
+ "include/flashinfer/attention/hopper/named_barrier.cuh",
180
+ "include/flashinfer/attention/hopper/tile_scheduler.cuh",
181
+ "include/flashinfer/attention/hopper/sparse_mainloop.cuh",
182
+ "include/flashinfer/attention/hopper/kernel_traits.cuh",
183
+ "include/flashinfer/attention/hopper/mainloop_mma.cuh",
184
+ "include/flashinfer/attention/hopper/quantization/epilogue.cuh",
185
+ "include/flashinfer/attention/hopper/quantization/prefill_sm90.cuh",
186
+ "include/flashinfer/attention/hopper/quantization/mainloop_load.cuh",
187
+ "include/flashinfer/attention/hopper/quantization/mainloop_sparse_load.cuh",
188
+ "include/flashinfer/attention/hopper/quantization/kernel_traits.cuh",
189
+ "include/flashinfer/attention/hopper/quantization/mainloop_mma.cuh",
190
+ "include/flashinfer/attention/default_decode_params.cuh",
191
+ "include/flashinfer/attention/mla_params.cuh",
192
+ "include/flashinfer/attention/persistent.cuh",
193
+ "include/flashinfer/pos_enc.cuh",
194
+ "include/flashinfer/cutlass_utils.cuh",
195
+ "include/flashinfer/comm/trtllm_moe_allreduce_fusion.cuh",
196
+ "include/flashinfer/comm/trtllm_mnnvl_allreduce.cuh",
197
+ "include/flashinfer/comm/vllm_custom_all_reduce.cuh",
198
+ "include/flashinfer/comm/trtllm_allreduce.cuh",
199
+ "include/flashinfer/comm/trtllm_allreduce_fusion.cuh",
200
+ "include/flashinfer/comm/trtllm_alltoall.cuh",
201
+ "include/flashinfer/frag_layout_swizzle.cuh",
202
+ "include/flashinfer/profiler.cuh",
203
+ "include/flashinfer/cp_async.cuh",
204
+
205
+ "include/pytorch_conversion_utils.h",
206
+ "include/pytorch_extension_utils.h",
207
+
208
+ # # batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False
209
+ # "csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_config.inc",
210
+ # "csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_jit_pybind.cu",
211
+ # "csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode.cu",
212
+ # "csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_kernel.cu",
213
+
214
+ # # batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False
215
+ # "csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_config.inc",
216
+ # "csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_jit_pybind.cu",
217
+ # "csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode.cu",
218
+ # "csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_kernel.cu",
219
+
220
+
221
+ "csrc/generated/gelu_and_mul.cu",
222
+ "csrc/generated/gelu_tanh_and_mul.cu",
223
+ "csrc/generated/silu_and_mul.cu",
224
+ ]
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode.cu ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include <flashinfer/attention/scheduler.cuh>
17
+ #include <flashinfer/pos_enc.cuh>
18
+ #include <flashinfer/utils.cuh>
19
+ #include <optional>
20
+
21
+ #include "batch_decode_config.inc"
22
+ #include "pytorch_conversion_utils.h"
23
+ #include "pytorch_extension_utils.h"
24
+
25
+ namespace flashinfer {
26
+
27
+ template <uint32_t HEAD_DIM, PosEncodingMode POS_ENCODING_MODE, typename AttentionVariant,
28
+ typename Params>
29
+ cudaError_t BatchDecodeWithPagedKVCacheDispatched(Params params, typename Params::DTypeO* tmp_v,
30
+ float* tmp_s, bool enable_pdl,
31
+ cudaStream_t stream);
32
+
33
+ } // namespace flashinfer
34
+
35
+ using namespace flashinfer;
36
+
37
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
38
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
39
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
40
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
41
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
42
+ at::Tensor empty_q_data, at::Tensor empty_kv_data) {
43
+ size_t float_workspace_size_in_bytes =
44
+ float_workspace_buffer.size(0) * float_workspace_buffer.element_size();
45
+ size_t int_workspace_size_in_bytes =
46
+ int_workspace_buffer.size(0) * int_workspace_buffer.element_size();
47
+
48
+ DecodePlanInfo plan_info;
49
+
50
+ auto q_scalar_type = empty_q_data.scalar_type();
51
+ auto kv_scalar_type = empty_kv_data.scalar_type();
52
+
53
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
54
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
55
+ "cores template for different head dim");
56
+
57
+ const c10::cuda::OptionalCUDAGuard device_guard(float_workspace_buffer.device());
58
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
59
+ DISPATCH_context(
60
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
61
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
62
+ DISPATCH_GQA_GROUP_SIZE(num_qo_heads / num_kv_heads, GROUP_SIZE, {
63
+ auto work_estimation_func = BatchDecodeWithPagedKVCacheWorkEstimationDispatched<
64
+ GROUP_SIZE, HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>;
65
+ cudaError_t status = DecodePlan<HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>(
66
+ static_cast<void*>(float_workspace_buffer.data_ptr()), float_workspace_size_in_bytes,
67
+ static_cast<void*>(int_workspace_buffer.data_ptr()),
68
+ static_cast<void*>(page_locked_int_workspace_buffer.data_ptr()),
69
+ int_workspace_size_in_bytes, plan_info, static_cast<IdType*>(indptr.data_ptr()),
70
+ batch_size, num_qo_heads, page_size, enable_cuda_graph,
71
+ /*stream=*/stream, work_estimation_func);
72
+
73
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
74
+ cudaGetErrorString(status));
75
+ return true;
76
+ });
77
+ });
78
+
79
+ return vec_to_tensor(plan_info.ToVector());
80
+ }
81
+
82
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
83
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
84
+ at::Tensor q, at::Tensor paged_k_cache,
85
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
86
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
87
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
88
+ int64_t kv_layout_code, int64_t window_left,
89
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS) {
90
+ DecodePlanInfo plan_info;
91
+ plan_info.FromVector(tensor_to_vec(plan_info_vec));
92
+ QKVLayout kv_layout = static_cast<QKVLayout>(kv_layout_code);
93
+ auto device = q.device();
94
+ int64_t batch_size = q.size(0);
95
+ int64_t num_qo_heads = q.size(1);
96
+ int64_t num_kv_heads, page_size;
97
+
98
+ if (kv_layout == QKVLayout::kHND) {
99
+ num_kv_heads = paged_k_cache.size(1);
100
+ page_size = paged_k_cache.size(2);
101
+ } else {
102
+ page_size = paged_k_cache.size(1);
103
+ num_kv_heads = paged_k_cache.size(2);
104
+ }
105
+ uint32_t head_dim_qk = q.size(2);
106
+ uint32_t head_dim_vo = paged_v_cache.size(3);
107
+
108
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
109
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
110
+ "cores template for different head dim");
111
+
112
+ if (maybe_lse) {
113
+ const auto& lse = *maybe_lse;
114
+ TORCH_CHECK(lse.size(0) == batch_size, lse.size(0), q.size(0));
115
+ TORCH_CHECK(lse.size(1) == num_qo_heads, lse.size(1), q.size(1));
116
+ }
117
+
118
+ void* float_buffer = static_cast<void*>(float_workspace_buffer.data_ptr());
119
+ void* int_buffer = static_cast<void*>(int_workspace_buffer.data_ptr());
120
+
121
+ // get q_scalar_type and kv_scalar_type
122
+ auto q_scalar_type = q.scalar_type();
123
+ auto kv_scalar_type = paged_k_cache.scalar_type();
124
+
125
+ // get q_stride_n and q_stride_h
126
+ const auto q_stride_n = q.stride(0);
127
+ const auto q_stride_h = q.stride(1);
128
+
129
+ // get kv_cache_strides
130
+ const int64_t* kv_cache_strides = nullptr;
131
+ auto k_strides = paged_k_cache.strides();
132
+ auto v_strides = paged_v_cache.strides();
133
+ TORCH_CHECK(k_strides == v_strides, "k/v strides must be identical");
134
+ kv_cache_strides = k_strides.data();
135
+
136
+ const c10::cuda::OptionalCUDAGuard device_guard(device);
137
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
138
+
139
+ DISPATCH_context(
140
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
141
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
142
+ paged_kv_t<DTypeKV, IdType> paged_kv(
143
+ num_kv_heads, page_size, HEAD_DIM_QK, batch_size, kv_layout,
144
+ static_cast<DTypeKV*>(paged_k_cache.data_ptr()),
145
+ static_cast<DTypeKV*>(paged_v_cache.data_ptr()), kv_cache_strides,
146
+ static_cast<IdType*>(paged_kv_indices.data_ptr()),
147
+ static_cast<IdType*>(paged_kv_indptr.data_ptr()),
148
+ static_cast<IdType*>(paged_kv_last_page_len.data_ptr()));
149
+
150
+ Params params;
151
+ params.q = static_cast<DTypeQ*>(q.data_ptr());
152
+ params.paged_kv = paged_kv;
153
+ params.o = static_cast<DTypeO*>(o.data_ptr());
154
+ params.lse = maybe_lse ? static_cast<float*>(maybe_lse->data_ptr()) : nullptr;
155
+ params.padded_batch_size = 0;
156
+ params.num_qo_heads = num_qo_heads;
157
+ params.q_stride_n = q_stride_n;
158
+ params.q_stride_h = q_stride_h;
159
+ params.window_left = window_left;
160
+ params.request_indices = nullptr;
161
+ params.kv_tile_indices = nullptr;
162
+ params.o_indptr = nullptr;
163
+ params.kv_chunk_size_ptr = nullptr;
164
+ params.block_valid_mask = nullptr;
165
+ params.partition_kv = false;
166
+
167
+ ADDITIONAL_PARAMS_SETTER
168
+
169
+ DTypeO* tmp_v = nullptr;
170
+ float* tmp_s = nullptr;
171
+ params.request_indices =
172
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.request_indices_offset);
173
+ params.kv_tile_indices =
174
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_tile_indices_offset);
175
+ params.o_indptr = GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.o_indptr_offset);
176
+ params.kv_chunk_size_ptr =
177
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_chunk_size_ptr_offset);
178
+ if (plan_info.split_kv) {
179
+ tmp_v = GetPtrFromBaseOffset<DTypeO>(float_buffer, plan_info.v_offset);
180
+ tmp_s = GetPtrFromBaseOffset<float>(float_buffer, plan_info.s_offset);
181
+ if (plan_info.enable_cuda_graph) {
182
+ params.block_valid_mask =
183
+ GetPtrFromBaseOffset<bool>(int_buffer, plan_info.block_valid_mask_offset);
184
+ }
185
+ }
186
+ params.padded_batch_size = plan_info.padded_batch_size;
187
+
188
+ cudaError_t status =
189
+ flashinfer::BatchDecodeWithPagedKVCacheDispatched<HEAD_DIM_QK, POS_ENCODING_MODE,
190
+ AttentionVariant>(params, tmp_v,
191
+ tmp_s, enable_pdl,
192
+ /*stream=*/stream);
193
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
194
+ cudaGetErrorString(status));
195
+ return true;
196
+ });
197
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_config.inc ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #pragma once
2
+ #include <flashinfer/page.cuh>
3
+ #include <flashinfer/math.cuh>
4
+ #include <flashinfer/layout.cuh>
5
+ #include <flashinfer/pos_enc.cuh>
6
+ #include <flashinfer/attention/variant_helper.cuh>
7
+
8
+ #define ADDITIONAL_FUNC_PARAMS , std::optional<at::Tensor> maybe_alibi_slopes, double logits_soft_cap, double sm_scale, double rope_rcp_scale, double rope_rcp_theta
9
+ #define ADDITIONAL_PARAMS_SETTER params.maybe_alibi_slopes = maybe_alibi_slopes ? static_cast<float*>(maybe_alibi_slopes->data_ptr()): nullptr; \
10
+ params.logits_soft_cap = logits_soft_cap; \
11
+ params.sm_scale = sm_scale; \
12
+ params.rope_rcp_scale = rope_rcp_scale; \
13
+ params.rope_rcp_theta = rope_rcp_theta;
14
+
15
+ #define DISPATCH_context(DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE, USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, ...) { \
16
+ using AttentionVariant = DefaultAttention<false, false, false, false>; \
17
+ __VA_ARGS__(); \
18
+ }
19
+
20
+ using namespace flashinfer;
21
+
22
+ using DTypeQ = nv_bfloat16;
23
+ using DTypeKV = nv_bfloat16;
24
+ using DTypeO = nv_bfloat16;
25
+ using IdType = int32_t;
26
+ constexpr int HEAD_DIM_QK = 128;
27
+ constexpr int HEAD_DIM_VO = 128;
28
+ constexpr auto USE_LOGITS_SOFT_CAP = false;
29
+ constexpr auto POS_ENCODING_MODE = PosEncodingMode::kNone;
30
+ constexpr auto USE_SLIDING_WINDOW = false;
31
+
32
+ struct Params {
33
+ using DTypeQ = DTypeQ;
34
+ using DTypeKV = DTypeKV;
35
+ using DTypeO = DTypeO;
36
+ using IdType = IdType;
37
+
38
+ DTypeQ* q;
39
+ paged_kv_t<DTypeKV, IdType> paged_kv;
40
+ DTypeO* o;
41
+ float* lse;
42
+
43
+ float* maybe_alibi_slopes;
44
+ double logits_soft_cap;
45
+ double sm_scale;
46
+ double rope_rcp_scale;
47
+ double rope_rcp_theta;
48
+
49
+
50
+ uint32_t padded_batch_size;
51
+ uint32_t num_qo_heads;
52
+ IdType q_stride_n;
53
+ IdType q_stride_h;
54
+ int32_t window_left;
55
+ bool enable_pdl;
56
+
57
+ IdType* request_indices;
58
+ IdType* kv_tile_indices;
59
+ IdType* o_indptr;
60
+ IdType* kv_chunk_size_ptr;
61
+ bool* block_valid_mask;
62
+ bool partition_kv;
63
+
64
+ __host__ __device__ __forceinline__ int32_t get_qo_len(int32_t batch_idx) const { return 1; }
65
+
66
+ __host__ __device__ __forceinline__ int32_t get_kv_len(int32_t batch_idx) const {
67
+ return paged_kv.get_length(batch_idx);
68
+ }
69
+ };
70
+
71
+ #include<flashinfer/attention/variants.cuh>
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_jit_pybind.cu ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023-2025 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include "batch_decode_config.inc"
17
+ #include "pytorch_extension_utils.h"
18
+
19
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
20
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
21
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
22
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
23
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
24
+ at::Tensor empty_q_data, at::Tensor empty_kv_data);
25
+
26
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
27
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
28
+ at::Tensor q, at::Tensor paged_k_cache,
29
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
30
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
31
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
32
+ int64_t kv_layout_code, int64_t window_left,
33
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS);
34
+
35
+ TORCH_LIBRARY_FRAGMENT(TORCH_EXTENSION_NAME, m) {
36
+ // Batched decode with paged KV-Cache plan
37
+ m.def("plan", BatchDecodeWithPagedKVCachePlan);
38
+ // Batched decode with paged KV-Cache run
39
+ m.def("run", BatchDecodeWithPagedKVCacheRun);
40
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_kernel.cu ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #include <flashinfer/attention/decode.cuh>
2
+ #include "batch_decode_config.inc"
3
+
4
+ using namespace flashinfer;
5
+
6
+ namespace flashinfer {
7
+
8
+ template cudaError_t
9
+ BatchDecodeWithPagedKVCacheDispatched<128, PosEncodingMode::kNone, DefaultAttention<false, false, false, false>, Params>(
10
+ Params params, nv_bfloat16* tmp_v,
11
+ float* tmp_s, bool enable_pdl, cudaStream_t stream);
12
+
13
+ };
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode.cu ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include <flashinfer/attention/scheduler.cuh>
17
+ #include <flashinfer/pos_enc.cuh>
18
+ #include <flashinfer/utils.cuh>
19
+ #include <optional>
20
+
21
+ #include "batch_decode_config.inc"
22
+ #include "pytorch_conversion_utils.h"
23
+ #include "pytorch_extension_utils.h"
24
+
25
+ namespace flashinfer {
26
+
27
+ template <uint32_t HEAD_DIM, PosEncodingMode POS_ENCODING_MODE, typename AttentionVariant,
28
+ typename Params>
29
+ cudaError_t BatchDecodeWithPagedKVCacheDispatched(Params params, typename Params::DTypeO* tmp_v,
30
+ float* tmp_s, bool enable_pdl,
31
+ cudaStream_t stream);
32
+
33
+ } // namespace flashinfer
34
+
35
+ using namespace flashinfer;
36
+
37
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
38
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
39
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
40
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
41
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
42
+ at::Tensor empty_q_data, at::Tensor empty_kv_data) {
43
+ size_t float_workspace_size_in_bytes =
44
+ float_workspace_buffer.size(0) * float_workspace_buffer.element_size();
45
+ size_t int_workspace_size_in_bytes =
46
+ int_workspace_buffer.size(0) * int_workspace_buffer.element_size();
47
+
48
+ DecodePlanInfo plan_info;
49
+
50
+ auto q_scalar_type = empty_q_data.scalar_type();
51
+ auto kv_scalar_type = empty_kv_data.scalar_type();
52
+
53
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
54
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
55
+ "cores template for different head dim");
56
+
57
+ const c10::cuda::OptionalCUDAGuard device_guard(float_workspace_buffer.device());
58
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
59
+ DISPATCH_context(
60
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
61
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
62
+ DISPATCH_GQA_GROUP_SIZE(num_qo_heads / num_kv_heads, GROUP_SIZE, {
63
+ auto work_estimation_func = BatchDecodeWithPagedKVCacheWorkEstimationDispatched<
64
+ GROUP_SIZE, HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>;
65
+ cudaError_t status = DecodePlan<HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>(
66
+ static_cast<void*>(float_workspace_buffer.data_ptr()), float_workspace_size_in_bytes,
67
+ static_cast<void*>(int_workspace_buffer.data_ptr()),
68
+ static_cast<void*>(page_locked_int_workspace_buffer.data_ptr()),
69
+ int_workspace_size_in_bytes, plan_info, static_cast<IdType*>(indptr.data_ptr()),
70
+ batch_size, num_qo_heads, page_size, enable_cuda_graph,
71
+ /*stream=*/stream, work_estimation_func);
72
+
73
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
74
+ cudaGetErrorString(status));
75
+ return true;
76
+ });
77
+ });
78
+
79
+ return vec_to_tensor(plan_info.ToVector());
80
+ }
81
+
82
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
83
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
84
+ at::Tensor q, at::Tensor paged_k_cache,
85
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
86
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
87
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
88
+ int64_t kv_layout_code, int64_t window_left,
89
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS) {
90
+ DecodePlanInfo plan_info;
91
+ plan_info.FromVector(tensor_to_vec(plan_info_vec));
92
+ QKVLayout kv_layout = static_cast<QKVLayout>(kv_layout_code);
93
+ auto device = q.device();
94
+ int64_t batch_size = q.size(0);
95
+ int64_t num_qo_heads = q.size(1);
96
+ int64_t num_kv_heads, page_size;
97
+
98
+ if (kv_layout == QKVLayout::kHND) {
99
+ num_kv_heads = paged_k_cache.size(1);
100
+ page_size = paged_k_cache.size(2);
101
+ } else {
102
+ page_size = paged_k_cache.size(1);
103
+ num_kv_heads = paged_k_cache.size(2);
104
+ }
105
+ uint32_t head_dim_qk = q.size(2);
106
+ uint32_t head_dim_vo = paged_v_cache.size(3);
107
+
108
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
109
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
110
+ "cores template for different head dim");
111
+
112
+ if (maybe_lse) {
113
+ const auto& lse = *maybe_lse;
114
+ TORCH_CHECK(lse.size(0) == batch_size, lse.size(0), q.size(0));
115
+ TORCH_CHECK(lse.size(1) == num_qo_heads, lse.size(1), q.size(1));
116
+ }
117
+
118
+ void* float_buffer = static_cast<void*>(float_workspace_buffer.data_ptr());
119
+ void* int_buffer = static_cast<void*>(int_workspace_buffer.data_ptr());
120
+
121
+ // get q_scalar_type and kv_scalar_type
122
+ auto q_scalar_type = q.scalar_type();
123
+ auto kv_scalar_type = paged_k_cache.scalar_type();
124
+
125
+ // get q_stride_n and q_stride_h
126
+ const auto q_stride_n = q.stride(0);
127
+ const auto q_stride_h = q.stride(1);
128
+
129
+ // get kv_cache_strides
130
+ const int64_t* kv_cache_strides = nullptr;
131
+ auto k_strides = paged_k_cache.strides();
132
+ auto v_strides = paged_v_cache.strides();
133
+ TORCH_CHECK(k_strides == v_strides, "k/v strides must be identical");
134
+ kv_cache_strides = k_strides.data();
135
+
136
+ const c10::cuda::OptionalCUDAGuard device_guard(device);
137
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
138
+
139
+ DISPATCH_context(
140
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
141
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
142
+ paged_kv_t<DTypeKV, IdType> paged_kv(
143
+ num_kv_heads, page_size, HEAD_DIM_QK, batch_size, kv_layout,
144
+ static_cast<DTypeKV*>(paged_k_cache.data_ptr()),
145
+ static_cast<DTypeKV*>(paged_v_cache.data_ptr()), kv_cache_strides,
146
+ static_cast<IdType*>(paged_kv_indices.data_ptr()),
147
+ static_cast<IdType*>(paged_kv_indptr.data_ptr()),
148
+ static_cast<IdType*>(paged_kv_last_page_len.data_ptr()));
149
+
150
+ Params params;
151
+ params.q = static_cast<DTypeQ*>(q.data_ptr());
152
+ params.paged_kv = paged_kv;
153
+ params.o = static_cast<DTypeO*>(o.data_ptr());
154
+ params.lse = maybe_lse ? static_cast<float*>(maybe_lse->data_ptr()) : nullptr;
155
+ params.padded_batch_size = 0;
156
+ params.num_qo_heads = num_qo_heads;
157
+ params.q_stride_n = q_stride_n;
158
+ params.q_stride_h = q_stride_h;
159
+ params.window_left = window_left;
160
+ params.request_indices = nullptr;
161
+ params.kv_tile_indices = nullptr;
162
+ params.o_indptr = nullptr;
163
+ params.kv_chunk_size_ptr = nullptr;
164
+ params.block_valid_mask = nullptr;
165
+ params.partition_kv = false;
166
+
167
+ ADDITIONAL_PARAMS_SETTER
168
+
169
+ DTypeO* tmp_v = nullptr;
170
+ float* tmp_s = nullptr;
171
+ params.request_indices =
172
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.request_indices_offset);
173
+ params.kv_tile_indices =
174
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_tile_indices_offset);
175
+ params.o_indptr = GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.o_indptr_offset);
176
+ params.kv_chunk_size_ptr =
177
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_chunk_size_ptr_offset);
178
+ if (plan_info.split_kv) {
179
+ tmp_v = GetPtrFromBaseOffset<DTypeO>(float_buffer, plan_info.v_offset);
180
+ tmp_s = GetPtrFromBaseOffset<float>(float_buffer, plan_info.s_offset);
181
+ if (plan_info.enable_cuda_graph) {
182
+ params.block_valid_mask =
183
+ GetPtrFromBaseOffset<bool>(int_buffer, plan_info.block_valid_mask_offset);
184
+ }
185
+ }
186
+ params.padded_batch_size = plan_info.padded_batch_size;
187
+
188
+ cudaError_t status =
189
+ flashinfer::BatchDecodeWithPagedKVCacheDispatched<HEAD_DIM_QK, POS_ENCODING_MODE,
190
+ AttentionVariant>(params, tmp_v,
191
+ tmp_s, enable_pdl,
192
+ /*stream=*/stream);
193
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
194
+ cudaGetErrorString(status));
195
+ return true;
196
+ });
197
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_config.inc ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #pragma once
2
+ #include <flashinfer/page.cuh>
3
+ #include <flashinfer/math.cuh>
4
+ #include <flashinfer/layout.cuh>
5
+ #include <flashinfer/pos_enc.cuh>
6
+ #include <flashinfer/attention/variant_helper.cuh>
7
+
8
+ #define ADDITIONAL_FUNC_PARAMS , std::optional<at::Tensor> maybe_alibi_slopes, double logits_soft_cap, double sm_scale, double rope_rcp_scale, double rope_rcp_theta
9
+ #define ADDITIONAL_PARAMS_SETTER params.maybe_alibi_slopes = maybe_alibi_slopes ? static_cast<float*>(maybe_alibi_slopes->data_ptr()): nullptr; \
10
+ params.logits_soft_cap = logits_soft_cap; \
11
+ params.sm_scale = sm_scale; \
12
+ params.rope_rcp_scale = rope_rcp_scale; \
13
+ params.rope_rcp_theta = rope_rcp_theta;
14
+
15
+ #define DISPATCH_context(DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE, USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, ...) { \
16
+ using AttentionVariant = DefaultAttention<false, true, true, false>; \
17
+ __VA_ARGS__(); \
18
+ }
19
+
20
+ using namespace flashinfer;
21
+
22
+ using DTypeQ = nv_bfloat16;
23
+ using DTypeKV = nv_bfloat16;
24
+ using DTypeO = nv_bfloat16;
25
+ using IdType = int32_t;
26
+ constexpr int HEAD_DIM_QK = 256;
27
+ constexpr int HEAD_DIM_VO = 256;
28
+ constexpr auto USE_LOGITS_SOFT_CAP = true;
29
+ constexpr auto POS_ENCODING_MODE = PosEncodingMode::kNone;
30
+ constexpr auto USE_SLIDING_WINDOW = true;
31
+
32
+ struct Params {
33
+ using DTypeQ = DTypeQ;
34
+ using DTypeKV = DTypeKV;
35
+ using DTypeO = DTypeO;
36
+ using IdType = IdType;
37
+
38
+ DTypeQ* q;
39
+ paged_kv_t<DTypeKV, IdType> paged_kv;
40
+ DTypeO* o;
41
+ float* lse;
42
+
43
+ float* maybe_alibi_slopes;
44
+ double logits_soft_cap;
45
+ double sm_scale;
46
+ double rope_rcp_scale;
47
+ double rope_rcp_theta;
48
+
49
+
50
+ uint32_t padded_batch_size;
51
+ uint32_t num_qo_heads;
52
+ IdType q_stride_n;
53
+ IdType q_stride_h;
54
+ int32_t window_left;
55
+ bool enable_pdl;
56
+
57
+ IdType* request_indices;
58
+ IdType* kv_tile_indices;
59
+ IdType* o_indptr;
60
+ IdType* kv_chunk_size_ptr;
61
+ bool* block_valid_mask;
62
+ bool partition_kv;
63
+
64
+ __host__ __device__ __forceinline__ int32_t get_qo_len(int32_t batch_idx) const { return 1; }
65
+
66
+ __host__ __device__ __forceinline__ int32_t get_kv_len(int32_t batch_idx) const {
67
+ return paged_kv.get_length(batch_idx);
68
+ }
69
+ };
70
+
71
+ #include<flashinfer/attention/variants.cuh>
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_jit_pybind.cu ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023-2025 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include "batch_decode_config.inc"
17
+ #include "pytorch_extension_utils.h"
18
+
19
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
20
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
21
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
22
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
23
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
24
+ at::Tensor empty_q_data, at::Tensor empty_kv_data);
25
+
26
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
27
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
28
+ at::Tensor q, at::Tensor paged_k_cache,
29
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
30
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
31
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
32
+ int64_t kv_layout_code, int64_t window_left,
33
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS);
34
+
35
+ TORCH_LIBRARY_FRAGMENT(TORCH_EXTENSION_NAME, m) {
36
+ // Batched decode with paged KV-Cache plan
37
+ m.def("plan", BatchDecodeWithPagedKVCachePlan);
38
+ // Batched decode with paged KV-Cache run
39
+ m.def("run", BatchDecodeWithPagedKVCacheRun);
40
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_kernel.cu ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #include <flashinfer/attention/decode.cuh>
2
+ #include "batch_decode_config.inc"
3
+
4
+ using namespace flashinfer;
5
+
6
+ namespace flashinfer {
7
+
8
+ template cudaError_t
9
+ BatchDecodeWithPagedKVCacheDispatched<256, PosEncodingMode::kNone, DefaultAttention<false, true, true, false>, Params>(
10
+ Params params, nv_bfloat16* tmp_v,
11
+ float* tmp_s, bool enable_pdl, cudaStream_t stream);
12
+
13
+ };
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode.cu ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include <flashinfer/attention/scheduler.cuh>
17
+ #include <flashinfer/pos_enc.cuh>
18
+ #include <flashinfer/utils.cuh>
19
+ #include <optional>
20
+
21
+ #include "batch_decode_config.inc"
22
+ #include "pytorch_conversion_utils.h"
23
+ #include "pytorch_extension_utils.h"
24
+
25
+ namespace flashinfer {
26
+
27
+ template <uint32_t HEAD_DIM, PosEncodingMode POS_ENCODING_MODE, typename AttentionVariant,
28
+ typename Params>
29
+ cudaError_t BatchDecodeWithPagedKVCacheDispatched(Params params, typename Params::DTypeO* tmp_v,
30
+ float* tmp_s, bool enable_pdl,
31
+ cudaStream_t stream);
32
+
33
+ } // namespace flashinfer
34
+
35
+ using namespace flashinfer;
36
+
37
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
38
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
39
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
40
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
41
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
42
+ at::Tensor empty_q_data, at::Tensor empty_kv_data) {
43
+ size_t float_workspace_size_in_bytes =
44
+ float_workspace_buffer.size(0) * float_workspace_buffer.element_size();
45
+ size_t int_workspace_size_in_bytes =
46
+ int_workspace_buffer.size(0) * int_workspace_buffer.element_size();
47
+
48
+ DecodePlanInfo plan_info;
49
+
50
+ auto q_scalar_type = empty_q_data.scalar_type();
51
+ auto kv_scalar_type = empty_kv_data.scalar_type();
52
+
53
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
54
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
55
+ "cores template for different head dim");
56
+
57
+ const c10::cuda::OptionalCUDAGuard device_guard(float_workspace_buffer.device());
58
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
59
+ DISPATCH_context(
60
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
61
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
62
+ DISPATCH_GQA_GROUP_SIZE(num_qo_heads / num_kv_heads, GROUP_SIZE, {
63
+ auto work_estimation_func = BatchDecodeWithPagedKVCacheWorkEstimationDispatched<
64
+ GROUP_SIZE, HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>;
65
+ cudaError_t status = DecodePlan<HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>(
66
+ static_cast<void*>(float_workspace_buffer.data_ptr()), float_workspace_size_in_bytes,
67
+ static_cast<void*>(int_workspace_buffer.data_ptr()),
68
+ static_cast<void*>(page_locked_int_workspace_buffer.data_ptr()),
69
+ int_workspace_size_in_bytes, plan_info, static_cast<IdType*>(indptr.data_ptr()),
70
+ batch_size, num_qo_heads, page_size, enable_cuda_graph,
71
+ /*stream=*/stream, work_estimation_func);
72
+
73
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
74
+ cudaGetErrorString(status));
75
+ return true;
76
+ });
77
+ });
78
+
79
+ return vec_to_tensor(plan_info.ToVector());
80
+ }
81
+
82
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
83
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
84
+ at::Tensor q, at::Tensor paged_k_cache,
85
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
86
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
87
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
88
+ int64_t kv_layout_code, int64_t window_left,
89
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS) {
90
+ DecodePlanInfo plan_info;
91
+ plan_info.FromVector(tensor_to_vec(plan_info_vec));
92
+ QKVLayout kv_layout = static_cast<QKVLayout>(kv_layout_code);
93
+ auto device = q.device();
94
+ int64_t batch_size = q.size(0);
95
+ int64_t num_qo_heads = q.size(1);
96
+ int64_t num_kv_heads, page_size;
97
+
98
+ if (kv_layout == QKVLayout::kHND) {
99
+ num_kv_heads = paged_k_cache.size(1);
100
+ page_size = paged_k_cache.size(2);
101
+ } else {
102
+ page_size = paged_k_cache.size(1);
103
+ num_kv_heads = paged_k_cache.size(2);
104
+ }
105
+ uint32_t head_dim_qk = q.size(2);
106
+ uint32_t head_dim_vo = paged_v_cache.size(3);
107
+
108
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
109
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
110
+ "cores template for different head dim");
111
+
112
+ if (maybe_lse) {
113
+ const auto& lse = *maybe_lse;
114
+ TORCH_CHECK(lse.size(0) == batch_size, lse.size(0), q.size(0));
115
+ TORCH_CHECK(lse.size(1) == num_qo_heads, lse.size(1), q.size(1));
116
+ }
117
+
118
+ void* float_buffer = static_cast<void*>(float_workspace_buffer.data_ptr());
119
+ void* int_buffer = static_cast<void*>(int_workspace_buffer.data_ptr());
120
+
121
+ // get q_scalar_type and kv_scalar_type
122
+ auto q_scalar_type = q.scalar_type();
123
+ auto kv_scalar_type = paged_k_cache.scalar_type();
124
+
125
+ // get q_stride_n and q_stride_h
126
+ const auto q_stride_n = q.stride(0);
127
+ const auto q_stride_h = q.stride(1);
128
+
129
+ // get kv_cache_strides
130
+ const int64_t* kv_cache_strides = nullptr;
131
+ auto k_strides = paged_k_cache.strides();
132
+ auto v_strides = paged_v_cache.strides();
133
+ TORCH_CHECK(k_strides == v_strides, "k/v strides must be identical");
134
+ kv_cache_strides = k_strides.data();
135
+
136
+ const c10::cuda::OptionalCUDAGuard device_guard(device);
137
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
138
+
139
+ DISPATCH_context(
140
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
141
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
142
+ paged_kv_t<DTypeKV, IdType> paged_kv(
143
+ num_kv_heads, page_size, HEAD_DIM_QK, batch_size, kv_layout,
144
+ static_cast<DTypeKV*>(paged_k_cache.data_ptr()),
145
+ static_cast<DTypeKV*>(paged_v_cache.data_ptr()), kv_cache_strides,
146
+ static_cast<IdType*>(paged_kv_indices.data_ptr()),
147
+ static_cast<IdType*>(paged_kv_indptr.data_ptr()),
148
+ static_cast<IdType*>(paged_kv_last_page_len.data_ptr()));
149
+
150
+ Params params;
151
+ params.q = static_cast<DTypeQ*>(q.data_ptr());
152
+ params.paged_kv = paged_kv;
153
+ params.o = static_cast<DTypeO*>(o.data_ptr());
154
+ params.lse = maybe_lse ? static_cast<float*>(maybe_lse->data_ptr()) : nullptr;
155
+ params.padded_batch_size = 0;
156
+ params.num_qo_heads = num_qo_heads;
157
+ params.q_stride_n = q_stride_n;
158
+ params.q_stride_h = q_stride_h;
159
+ params.window_left = window_left;
160
+ params.request_indices = nullptr;
161
+ params.kv_tile_indices = nullptr;
162
+ params.o_indptr = nullptr;
163
+ params.kv_chunk_size_ptr = nullptr;
164
+ params.block_valid_mask = nullptr;
165
+ params.partition_kv = false;
166
+
167
+ ADDITIONAL_PARAMS_SETTER
168
+
169
+ DTypeO* tmp_v = nullptr;
170
+ float* tmp_s = nullptr;
171
+ params.request_indices =
172
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.request_indices_offset);
173
+ params.kv_tile_indices =
174
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_tile_indices_offset);
175
+ params.o_indptr = GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.o_indptr_offset);
176
+ params.kv_chunk_size_ptr =
177
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_chunk_size_ptr_offset);
178
+ if (plan_info.split_kv) {
179
+ tmp_v = GetPtrFromBaseOffset<DTypeO>(float_buffer, plan_info.v_offset);
180
+ tmp_s = GetPtrFromBaseOffset<float>(float_buffer, plan_info.s_offset);
181
+ if (plan_info.enable_cuda_graph) {
182
+ params.block_valid_mask =
183
+ GetPtrFromBaseOffset<bool>(int_buffer, plan_info.block_valid_mask_offset);
184
+ }
185
+ }
186
+ params.padded_batch_size = plan_info.padded_batch_size;
187
+
188
+ cudaError_t status =
189
+ flashinfer::BatchDecodeWithPagedKVCacheDispatched<HEAD_DIM_QK, POS_ENCODING_MODE,
190
+ AttentionVariant>(params, tmp_v,
191
+ tmp_s, enable_pdl,
192
+ /*stream=*/stream);
193
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
194
+ cudaGetErrorString(status));
195
+ return true;
196
+ });
197
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_config.inc ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #pragma once
2
+ #include <flashinfer/page.cuh>
3
+ #include <flashinfer/math.cuh>
4
+ #include <flashinfer/layout.cuh>
5
+ #include <flashinfer/pos_enc.cuh>
6
+ #include <flashinfer/attention/variant_helper.cuh>
7
+
8
+ #define ADDITIONAL_FUNC_PARAMS , std::optional<at::Tensor> maybe_alibi_slopes, double logits_soft_cap, double sm_scale, double rope_rcp_scale, double rope_rcp_theta
9
+ #define ADDITIONAL_PARAMS_SETTER params.maybe_alibi_slopes = maybe_alibi_slopes ? static_cast<float*>(maybe_alibi_slopes->data_ptr()): nullptr; \
10
+ params.logits_soft_cap = logits_soft_cap; \
11
+ params.sm_scale = sm_scale; \
12
+ params.rope_rcp_scale = rope_rcp_scale; \
13
+ params.rope_rcp_theta = rope_rcp_theta;
14
+
15
+ #define DISPATCH_context(DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE, USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, ...) { \
16
+ using AttentionVariant = DefaultAttention<false, false, false, false>; \
17
+ __VA_ARGS__(); \
18
+ }
19
+
20
+ using namespace flashinfer;
21
+
22
+ using DTypeQ = nv_bfloat16;
23
+ using DTypeKV = nv_bfloat16;
24
+ using DTypeO = nv_bfloat16;
25
+ using IdType = int32_t;
26
+ constexpr int HEAD_DIM_QK = 64;
27
+ constexpr int HEAD_DIM_VO = 64;
28
+ constexpr auto USE_LOGITS_SOFT_CAP = false;
29
+ constexpr auto POS_ENCODING_MODE = PosEncodingMode::kNone;
30
+ constexpr auto USE_SLIDING_WINDOW = false;
31
+
32
+ struct Params {
33
+ using DTypeQ = DTypeQ;
34
+ using DTypeKV = DTypeKV;
35
+ using DTypeO = DTypeO;
36
+ using IdType = IdType;
37
+
38
+ DTypeQ* q;
39
+ paged_kv_t<DTypeKV, IdType> paged_kv;
40
+ DTypeO* o;
41
+ float* lse;
42
+
43
+ float* maybe_alibi_slopes;
44
+ double logits_soft_cap;
45
+ double sm_scale;
46
+ double rope_rcp_scale;
47
+ double rope_rcp_theta;
48
+
49
+
50
+ uint32_t padded_batch_size;
51
+ uint32_t num_qo_heads;
52
+ IdType q_stride_n;
53
+ IdType q_stride_h;
54
+ int32_t window_left;
55
+ bool enable_pdl;
56
+
57
+ IdType* request_indices;
58
+ IdType* kv_tile_indices;
59
+ IdType* o_indptr;
60
+ IdType* kv_chunk_size_ptr;
61
+ bool* block_valid_mask;
62
+ bool partition_kv;
63
+
64
+ __host__ __device__ __forceinline__ int32_t get_qo_len(int32_t batch_idx) const { return 1; }
65
+
66
+ __host__ __device__ __forceinline__ int32_t get_kv_len(int32_t batch_idx) const {
67
+ return paged_kv.get_length(batch_idx);
68
+ }
69
+ };
70
+
71
+ #include<flashinfer/attention/variants.cuh>
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_jit_pybind.cu ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023-2025 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include "batch_decode_config.inc"
17
+ #include "pytorch_extension_utils.h"
18
+
19
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
20
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
21
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
22
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
23
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
24
+ at::Tensor empty_q_data, at::Tensor empty_kv_data);
25
+
26
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
27
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
28
+ at::Tensor q, at::Tensor paged_k_cache,
29
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
30
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
31
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
32
+ int64_t kv_layout_code, int64_t window_left,
33
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS);
34
+
35
+ TORCH_LIBRARY_FRAGMENT(TORCH_EXTENSION_NAME, m) {
36
+ // Batched decode with paged KV-Cache plan
37
+ m.def("plan", BatchDecodeWithPagedKVCachePlan);
38
+ // Batched decode with paged KV-Cache run
39
+ m.def("run", BatchDecodeWithPagedKVCacheRun);
40
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_kernel.cu ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #include <flashinfer/attention/decode.cuh>
2
+ #include "batch_decode_config.inc"
3
+
4
+ using namespace flashinfer;
5
+
6
+ namespace flashinfer {
7
+
8
+ template cudaError_t
9
+ BatchDecodeWithPagedKVCacheDispatched<64, PosEncodingMode::kNone, DefaultAttention<false, false, false, false>, Params>(
10
+ Params params, nv_bfloat16* tmp_v,
11
+ float* tmp_s, bool enable_pdl, cudaStream_t stream);
12
+
13
+ };
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode.cu ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include <flashinfer/attention/scheduler.cuh>
17
+ #include <flashinfer/pos_enc.cuh>
18
+ #include <flashinfer/utils.cuh>
19
+ #include <optional>
20
+
21
+ #include "batch_decode_config.inc"
22
+ #include "pytorch_conversion_utils.h"
23
+ #include "pytorch_extension_utils.h"
24
+
25
+ namespace flashinfer {
26
+
27
+ template <uint32_t HEAD_DIM, PosEncodingMode POS_ENCODING_MODE, typename AttentionVariant,
28
+ typename Params>
29
+ cudaError_t BatchDecodeWithPagedKVCacheDispatched(Params params, typename Params::DTypeO* tmp_v,
30
+ float* tmp_s, bool enable_pdl,
31
+ cudaStream_t stream);
32
+
33
+ } // namespace flashinfer
34
+
35
+ using namespace flashinfer;
36
+
37
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
38
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
39
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
40
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
41
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
42
+ at::Tensor empty_q_data, at::Tensor empty_kv_data) {
43
+ size_t float_workspace_size_in_bytes =
44
+ float_workspace_buffer.size(0) * float_workspace_buffer.element_size();
45
+ size_t int_workspace_size_in_bytes =
46
+ int_workspace_buffer.size(0) * int_workspace_buffer.element_size();
47
+
48
+ DecodePlanInfo plan_info;
49
+
50
+ auto q_scalar_type = empty_q_data.scalar_type();
51
+ auto kv_scalar_type = empty_kv_data.scalar_type();
52
+
53
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
54
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
55
+ "cores template for different head dim");
56
+
57
+ const c10::cuda::OptionalCUDAGuard device_guard(float_workspace_buffer.device());
58
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
59
+ DISPATCH_context(
60
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
61
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
62
+ DISPATCH_GQA_GROUP_SIZE(num_qo_heads / num_kv_heads, GROUP_SIZE, {
63
+ auto work_estimation_func = BatchDecodeWithPagedKVCacheWorkEstimationDispatched<
64
+ GROUP_SIZE, HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>;
65
+ cudaError_t status = DecodePlan<HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>(
66
+ static_cast<void*>(float_workspace_buffer.data_ptr()), float_workspace_size_in_bytes,
67
+ static_cast<void*>(int_workspace_buffer.data_ptr()),
68
+ static_cast<void*>(page_locked_int_workspace_buffer.data_ptr()),
69
+ int_workspace_size_in_bytes, plan_info, static_cast<IdType*>(indptr.data_ptr()),
70
+ batch_size, num_qo_heads, page_size, enable_cuda_graph,
71
+ /*stream=*/stream, work_estimation_func);
72
+
73
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
74
+ cudaGetErrorString(status));
75
+ return true;
76
+ });
77
+ });
78
+
79
+ return vec_to_tensor(plan_info.ToVector());
80
+ }
81
+
82
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
83
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
84
+ at::Tensor q, at::Tensor paged_k_cache,
85
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
86
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
87
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
88
+ int64_t kv_layout_code, int64_t window_left,
89
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS) {
90
+ DecodePlanInfo plan_info;
91
+ plan_info.FromVector(tensor_to_vec(plan_info_vec));
92
+ QKVLayout kv_layout = static_cast<QKVLayout>(kv_layout_code);
93
+ auto device = q.device();
94
+ int64_t batch_size = q.size(0);
95
+ int64_t num_qo_heads = q.size(1);
96
+ int64_t num_kv_heads, page_size;
97
+
98
+ if (kv_layout == QKVLayout::kHND) {
99
+ num_kv_heads = paged_k_cache.size(1);
100
+ page_size = paged_k_cache.size(2);
101
+ } else {
102
+ page_size = paged_k_cache.size(1);
103
+ num_kv_heads = paged_k_cache.size(2);
104
+ }
105
+ uint32_t head_dim_qk = q.size(2);
106
+ uint32_t head_dim_vo = paged_v_cache.size(3);
107
+
108
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
109
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
110
+ "cores template for different head dim");
111
+
112
+ if (maybe_lse) {
113
+ const auto& lse = *maybe_lse;
114
+ TORCH_CHECK(lse.size(0) == batch_size, lse.size(0), q.size(0));
115
+ TORCH_CHECK(lse.size(1) == num_qo_heads, lse.size(1), q.size(1));
116
+ }
117
+
118
+ void* float_buffer = static_cast<void*>(float_workspace_buffer.data_ptr());
119
+ void* int_buffer = static_cast<void*>(int_workspace_buffer.data_ptr());
120
+
121
+ // get q_scalar_type and kv_scalar_type
122
+ auto q_scalar_type = q.scalar_type();
123
+ auto kv_scalar_type = paged_k_cache.scalar_type();
124
+
125
+ // get q_stride_n and q_stride_h
126
+ const auto q_stride_n = q.stride(0);
127
+ const auto q_stride_h = q.stride(1);
128
+
129
+ // get kv_cache_strides
130
+ const int64_t* kv_cache_strides = nullptr;
131
+ auto k_strides = paged_k_cache.strides();
132
+ auto v_strides = paged_v_cache.strides();
133
+ TORCH_CHECK(k_strides == v_strides, "k/v strides must be identical");
134
+ kv_cache_strides = k_strides.data();
135
+
136
+ const c10::cuda::OptionalCUDAGuard device_guard(device);
137
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
138
+
139
+ DISPATCH_context(
140
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
141
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
142
+ paged_kv_t<DTypeKV, IdType> paged_kv(
143
+ num_kv_heads, page_size, HEAD_DIM_QK, batch_size, kv_layout,
144
+ static_cast<DTypeKV*>(paged_k_cache.data_ptr()),
145
+ static_cast<DTypeKV*>(paged_v_cache.data_ptr()), kv_cache_strides,
146
+ static_cast<IdType*>(paged_kv_indices.data_ptr()),
147
+ static_cast<IdType*>(paged_kv_indptr.data_ptr()),
148
+ static_cast<IdType*>(paged_kv_last_page_len.data_ptr()));
149
+
150
+ Params params;
151
+ params.q = static_cast<DTypeQ*>(q.data_ptr());
152
+ params.paged_kv = paged_kv;
153
+ params.o = static_cast<DTypeO*>(o.data_ptr());
154
+ params.lse = maybe_lse ? static_cast<float*>(maybe_lse->data_ptr()) : nullptr;
155
+ params.padded_batch_size = 0;
156
+ params.num_qo_heads = num_qo_heads;
157
+ params.q_stride_n = q_stride_n;
158
+ params.q_stride_h = q_stride_h;
159
+ params.window_left = window_left;
160
+ params.request_indices = nullptr;
161
+ params.kv_tile_indices = nullptr;
162
+ params.o_indptr = nullptr;
163
+ params.kv_chunk_size_ptr = nullptr;
164
+ params.block_valid_mask = nullptr;
165
+ params.partition_kv = false;
166
+
167
+ ADDITIONAL_PARAMS_SETTER
168
+
169
+ DTypeO* tmp_v = nullptr;
170
+ float* tmp_s = nullptr;
171
+ params.request_indices =
172
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.request_indices_offset);
173
+ params.kv_tile_indices =
174
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_tile_indices_offset);
175
+ params.o_indptr = GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.o_indptr_offset);
176
+ params.kv_chunk_size_ptr =
177
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_chunk_size_ptr_offset);
178
+ if (plan_info.split_kv) {
179
+ tmp_v = GetPtrFromBaseOffset<DTypeO>(float_buffer, plan_info.v_offset);
180
+ tmp_s = GetPtrFromBaseOffset<float>(float_buffer, plan_info.s_offset);
181
+ if (plan_info.enable_cuda_graph) {
182
+ params.block_valid_mask =
183
+ GetPtrFromBaseOffset<bool>(int_buffer, plan_info.block_valid_mask_offset);
184
+ }
185
+ }
186
+ params.padded_batch_size = plan_info.padded_batch_size;
187
+
188
+ cudaError_t status =
189
+ flashinfer::BatchDecodeWithPagedKVCacheDispatched<HEAD_DIM_QK, POS_ENCODING_MODE,
190
+ AttentionVariant>(params, tmp_v,
191
+ tmp_s, enable_pdl,
192
+ /*stream=*/stream);
193
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
194
+ cudaGetErrorString(status));
195
+ return true;
196
+ });
197
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_config.inc ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #pragma once
2
+ #include <flashinfer/page.cuh>
3
+ #include <flashinfer/math.cuh>
4
+ #include <flashinfer/layout.cuh>
5
+ #include <flashinfer/pos_enc.cuh>
6
+ #include <flashinfer/attention/variant_helper.cuh>
7
+
8
+ #define ADDITIONAL_FUNC_PARAMS , std::optional<at::Tensor> maybe_alibi_slopes, double logits_soft_cap, double sm_scale, double rope_rcp_scale, double rope_rcp_theta
9
+ #define ADDITIONAL_PARAMS_SETTER params.maybe_alibi_slopes = maybe_alibi_slopes ? static_cast<float*>(maybe_alibi_slopes->data_ptr()): nullptr; \
10
+ params.logits_soft_cap = logits_soft_cap; \
11
+ params.sm_scale = sm_scale; \
12
+ params.rope_rcp_scale = rope_rcp_scale; \
13
+ params.rope_rcp_theta = rope_rcp_theta;
14
+
15
+ #define DISPATCH_context(DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE, USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, ...) { \
16
+ using AttentionVariant = DefaultAttention<false, true, false, false>; \
17
+ __VA_ARGS__(); \
18
+ }
19
+
20
+ using namespace flashinfer;
21
+
22
+ using DTypeQ = nv_bfloat16;
23
+ using DTypeKV = nv_bfloat16;
24
+ using DTypeO = nv_bfloat16;
25
+ using IdType = int32_t;
26
+ constexpr int HEAD_DIM_QK = 64;
27
+ constexpr int HEAD_DIM_VO = 64;
28
+ constexpr auto USE_LOGITS_SOFT_CAP = false;
29
+ constexpr auto POS_ENCODING_MODE = PosEncodingMode::kNone;
30
+ constexpr auto USE_SLIDING_WINDOW = true;
31
+
32
+ struct Params {
33
+ using DTypeQ = DTypeQ;
34
+ using DTypeKV = DTypeKV;
35
+ using DTypeO = DTypeO;
36
+ using IdType = IdType;
37
+
38
+ DTypeQ* q;
39
+ paged_kv_t<DTypeKV, IdType> paged_kv;
40
+ DTypeO* o;
41
+ float* lse;
42
+
43
+ float* maybe_alibi_slopes;
44
+ double logits_soft_cap;
45
+ double sm_scale;
46
+ double rope_rcp_scale;
47
+ double rope_rcp_theta;
48
+
49
+
50
+ uint32_t padded_batch_size;
51
+ uint32_t num_qo_heads;
52
+ IdType q_stride_n;
53
+ IdType q_stride_h;
54
+ int32_t window_left;
55
+ bool enable_pdl;
56
+
57
+ IdType* request_indices;
58
+ IdType* kv_tile_indices;
59
+ IdType* o_indptr;
60
+ IdType* kv_chunk_size_ptr;
61
+ bool* block_valid_mask;
62
+ bool partition_kv;
63
+
64
+ __host__ __device__ __forceinline__ int32_t get_qo_len(int32_t batch_idx) const { return 1; }
65
+
66
+ __host__ __device__ __forceinline__ int32_t get_kv_len(int32_t batch_idx) const {
67
+ return paged_kv.get_length(batch_idx);
68
+ }
69
+ };
70
+
71
+ #include<flashinfer/attention/variants.cuh>
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_jit_pybind.cu ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023-2025 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include "batch_decode_config.inc"
17
+ #include "pytorch_extension_utils.h"
18
+
19
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
20
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
21
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
22
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
23
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
24
+ at::Tensor empty_q_data, at::Tensor empty_kv_data);
25
+
26
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
27
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
28
+ at::Tensor q, at::Tensor paged_k_cache,
29
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
30
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
31
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
32
+ int64_t kv_layout_code, int64_t window_left,
33
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS);
34
+
35
+ TORCH_LIBRARY_FRAGMENT(TORCH_EXTENSION_NAME, m) {
36
+ // Batched decode with paged KV-Cache plan
37
+ m.def("plan", BatchDecodeWithPagedKVCachePlan);
38
+ // Batched decode with paged KV-Cache run
39
+ m.def("run", BatchDecodeWithPagedKVCacheRun);
40
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_bf16_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_kernel.cu ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #include <flashinfer/attention/decode.cuh>
2
+ #include "batch_decode_config.inc"
3
+
4
+ using namespace flashinfer;
5
+
6
+ namespace flashinfer {
7
+
8
+ template cudaError_t
9
+ BatchDecodeWithPagedKVCacheDispatched<64, PosEncodingMode::kNone, DefaultAttention<false, true, false, false>, Params>(
10
+ Params params, nv_bfloat16* tmp_v,
11
+ float* tmp_s, bool enable_pdl, cudaStream_t stream);
12
+
13
+ };
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode.cu ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include <flashinfer/attention/scheduler.cuh>
17
+ #include <flashinfer/pos_enc.cuh>
18
+ #include <flashinfer/utils.cuh>
19
+ #include <optional>
20
+
21
+ #include "batch_decode_config.inc"
22
+ #include "pytorch_conversion_utils.h"
23
+ #include "pytorch_extension_utils.h"
24
+
25
+ namespace flashinfer {
26
+
27
+ template <uint32_t HEAD_DIM, PosEncodingMode POS_ENCODING_MODE, typename AttentionVariant,
28
+ typename Params>
29
+ cudaError_t BatchDecodeWithPagedKVCacheDispatched(Params params, typename Params::DTypeO* tmp_v,
30
+ float* tmp_s, bool enable_pdl,
31
+ cudaStream_t stream);
32
+
33
+ } // namespace flashinfer
34
+
35
+ using namespace flashinfer;
36
+
37
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
38
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
39
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
40
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
41
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
42
+ at::Tensor empty_q_data, at::Tensor empty_kv_data) {
43
+ size_t float_workspace_size_in_bytes =
44
+ float_workspace_buffer.size(0) * float_workspace_buffer.element_size();
45
+ size_t int_workspace_size_in_bytes =
46
+ int_workspace_buffer.size(0) * int_workspace_buffer.element_size();
47
+
48
+ DecodePlanInfo plan_info;
49
+
50
+ auto q_scalar_type = empty_q_data.scalar_type();
51
+ auto kv_scalar_type = empty_kv_data.scalar_type();
52
+
53
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
54
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
55
+ "cores template for different head dim");
56
+
57
+ const c10::cuda::OptionalCUDAGuard device_guard(float_workspace_buffer.device());
58
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
59
+ DISPATCH_context(
60
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
61
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
62
+ DISPATCH_GQA_GROUP_SIZE(num_qo_heads / num_kv_heads, GROUP_SIZE, {
63
+ auto work_estimation_func = BatchDecodeWithPagedKVCacheWorkEstimationDispatched<
64
+ GROUP_SIZE, HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>;
65
+ cudaError_t status = DecodePlan<HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>(
66
+ static_cast<void*>(float_workspace_buffer.data_ptr()), float_workspace_size_in_bytes,
67
+ static_cast<void*>(int_workspace_buffer.data_ptr()),
68
+ static_cast<void*>(page_locked_int_workspace_buffer.data_ptr()),
69
+ int_workspace_size_in_bytes, plan_info, static_cast<IdType*>(indptr.data_ptr()),
70
+ batch_size, num_qo_heads, page_size, enable_cuda_graph,
71
+ /*stream=*/stream, work_estimation_func);
72
+
73
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
74
+ cudaGetErrorString(status));
75
+ return true;
76
+ });
77
+ });
78
+
79
+ return vec_to_tensor(plan_info.ToVector());
80
+ }
81
+
82
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
83
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
84
+ at::Tensor q, at::Tensor paged_k_cache,
85
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
86
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
87
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
88
+ int64_t kv_layout_code, int64_t window_left,
89
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS) {
90
+ DecodePlanInfo plan_info;
91
+ plan_info.FromVector(tensor_to_vec(plan_info_vec));
92
+ QKVLayout kv_layout = static_cast<QKVLayout>(kv_layout_code);
93
+ auto device = q.device();
94
+ int64_t batch_size = q.size(0);
95
+ int64_t num_qo_heads = q.size(1);
96
+ int64_t num_kv_heads, page_size;
97
+
98
+ if (kv_layout == QKVLayout::kHND) {
99
+ num_kv_heads = paged_k_cache.size(1);
100
+ page_size = paged_k_cache.size(2);
101
+ } else {
102
+ page_size = paged_k_cache.size(1);
103
+ num_kv_heads = paged_k_cache.size(2);
104
+ }
105
+ uint32_t head_dim_qk = q.size(2);
106
+ uint32_t head_dim_vo = paged_v_cache.size(3);
107
+
108
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
109
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
110
+ "cores template for different head dim");
111
+
112
+ if (maybe_lse) {
113
+ const auto& lse = *maybe_lse;
114
+ TORCH_CHECK(lse.size(0) == batch_size, lse.size(0), q.size(0));
115
+ TORCH_CHECK(lse.size(1) == num_qo_heads, lse.size(1), q.size(1));
116
+ }
117
+
118
+ void* float_buffer = static_cast<void*>(float_workspace_buffer.data_ptr());
119
+ void* int_buffer = static_cast<void*>(int_workspace_buffer.data_ptr());
120
+
121
+ // get q_scalar_type and kv_scalar_type
122
+ auto q_scalar_type = q.scalar_type();
123
+ auto kv_scalar_type = paged_k_cache.scalar_type();
124
+
125
+ // get q_stride_n and q_stride_h
126
+ const auto q_stride_n = q.stride(0);
127
+ const auto q_stride_h = q.stride(1);
128
+
129
+ // get kv_cache_strides
130
+ const int64_t* kv_cache_strides = nullptr;
131
+ auto k_strides = paged_k_cache.strides();
132
+ auto v_strides = paged_v_cache.strides();
133
+ TORCH_CHECK(k_strides == v_strides, "k/v strides must be identical");
134
+ kv_cache_strides = k_strides.data();
135
+
136
+ const c10::cuda::OptionalCUDAGuard device_guard(device);
137
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
138
+
139
+ DISPATCH_context(
140
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
141
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
142
+ paged_kv_t<DTypeKV, IdType> paged_kv(
143
+ num_kv_heads, page_size, HEAD_DIM_QK, batch_size, kv_layout,
144
+ static_cast<DTypeKV*>(paged_k_cache.data_ptr()),
145
+ static_cast<DTypeKV*>(paged_v_cache.data_ptr()), kv_cache_strides,
146
+ static_cast<IdType*>(paged_kv_indices.data_ptr()),
147
+ static_cast<IdType*>(paged_kv_indptr.data_ptr()),
148
+ static_cast<IdType*>(paged_kv_last_page_len.data_ptr()));
149
+
150
+ Params params;
151
+ params.q = static_cast<DTypeQ*>(q.data_ptr());
152
+ params.paged_kv = paged_kv;
153
+ params.o = static_cast<DTypeO*>(o.data_ptr());
154
+ params.lse = maybe_lse ? static_cast<float*>(maybe_lse->data_ptr()) : nullptr;
155
+ params.padded_batch_size = 0;
156
+ params.num_qo_heads = num_qo_heads;
157
+ params.q_stride_n = q_stride_n;
158
+ params.q_stride_h = q_stride_h;
159
+ params.window_left = window_left;
160
+ params.request_indices = nullptr;
161
+ params.kv_tile_indices = nullptr;
162
+ params.o_indptr = nullptr;
163
+ params.kv_chunk_size_ptr = nullptr;
164
+ params.block_valid_mask = nullptr;
165
+ params.partition_kv = false;
166
+
167
+ ADDITIONAL_PARAMS_SETTER
168
+
169
+ DTypeO* tmp_v = nullptr;
170
+ float* tmp_s = nullptr;
171
+ params.request_indices =
172
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.request_indices_offset);
173
+ params.kv_tile_indices =
174
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_tile_indices_offset);
175
+ params.o_indptr = GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.o_indptr_offset);
176
+ params.kv_chunk_size_ptr =
177
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_chunk_size_ptr_offset);
178
+ if (plan_info.split_kv) {
179
+ tmp_v = GetPtrFromBaseOffset<DTypeO>(float_buffer, plan_info.v_offset);
180
+ tmp_s = GetPtrFromBaseOffset<float>(float_buffer, plan_info.s_offset);
181
+ if (plan_info.enable_cuda_graph) {
182
+ params.block_valid_mask =
183
+ GetPtrFromBaseOffset<bool>(int_buffer, plan_info.block_valid_mask_offset);
184
+ }
185
+ }
186
+ params.padded_batch_size = plan_info.padded_batch_size;
187
+
188
+ cudaError_t status =
189
+ flashinfer::BatchDecodeWithPagedKVCacheDispatched<HEAD_DIM_QK, POS_ENCODING_MODE,
190
+ AttentionVariant>(params, tmp_v,
191
+ tmp_s, enable_pdl,
192
+ /*stream=*/stream);
193
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
194
+ cudaGetErrorString(status));
195
+ return true;
196
+ });
197
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_config.inc ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #pragma once
2
+ #include <flashinfer/page.cuh>
3
+ #include <flashinfer/math.cuh>
4
+ #include <flashinfer/layout.cuh>
5
+ #include <flashinfer/pos_enc.cuh>
6
+ #include <flashinfer/attention/variant_helper.cuh>
7
+
8
+ #define ADDITIONAL_FUNC_PARAMS , std::optional<at::Tensor> maybe_alibi_slopes, double logits_soft_cap, double sm_scale, double rope_rcp_scale, double rope_rcp_theta
9
+ #define ADDITIONAL_PARAMS_SETTER params.maybe_alibi_slopes = maybe_alibi_slopes ? static_cast<float*>(maybe_alibi_slopes->data_ptr()): nullptr; \
10
+ params.logits_soft_cap = logits_soft_cap; \
11
+ params.sm_scale = sm_scale; \
12
+ params.rope_rcp_scale = rope_rcp_scale; \
13
+ params.rope_rcp_theta = rope_rcp_theta;
14
+
15
+ #define DISPATCH_context(DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE, USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, ...) { \
16
+ using AttentionVariant = DefaultAttention<false, false, false, false>; \
17
+ __VA_ARGS__(); \
18
+ }
19
+
20
+ using namespace flashinfer;
21
+
22
+ using DTypeQ = nv_bfloat16;
23
+ using DTypeKV = __nv_fp8_e4m3;
24
+ using DTypeO = nv_bfloat16;
25
+ using IdType = int32_t;
26
+ constexpr int HEAD_DIM_QK = 128;
27
+ constexpr int HEAD_DIM_VO = 128;
28
+ constexpr auto USE_LOGITS_SOFT_CAP = false;
29
+ constexpr auto POS_ENCODING_MODE = PosEncodingMode::kNone;
30
+ constexpr auto USE_SLIDING_WINDOW = false;
31
+
32
+ struct Params {
33
+ using DTypeQ = DTypeQ;
34
+ using DTypeKV = DTypeKV;
35
+ using DTypeO = DTypeO;
36
+ using IdType = IdType;
37
+
38
+ DTypeQ* q;
39
+ paged_kv_t<DTypeKV, IdType> paged_kv;
40
+ DTypeO* o;
41
+ float* lse;
42
+
43
+ float* maybe_alibi_slopes;
44
+ double logits_soft_cap;
45
+ double sm_scale;
46
+ double rope_rcp_scale;
47
+ double rope_rcp_theta;
48
+
49
+
50
+ uint32_t padded_batch_size;
51
+ uint32_t num_qo_heads;
52
+ IdType q_stride_n;
53
+ IdType q_stride_h;
54
+ int32_t window_left;
55
+ bool enable_pdl;
56
+
57
+ IdType* request_indices;
58
+ IdType* kv_tile_indices;
59
+ IdType* o_indptr;
60
+ IdType* kv_chunk_size_ptr;
61
+ bool* block_valid_mask;
62
+ bool partition_kv;
63
+
64
+ __host__ __device__ __forceinline__ int32_t get_qo_len(int32_t batch_idx) const { return 1; }
65
+
66
+ __host__ __device__ __forceinline__ int32_t get_kv_len(int32_t batch_idx) const {
67
+ return paged_kv.get_length(batch_idx);
68
+ }
69
+ };
70
+
71
+ #include<flashinfer/attention/variants.cuh>
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_jit_pybind.cu ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023-2025 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include "batch_decode_config.inc"
17
+ #include "pytorch_extension_utils.h"
18
+
19
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
20
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
21
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
22
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
23
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
24
+ at::Tensor empty_q_data, at::Tensor empty_kv_data);
25
+
26
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
27
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
28
+ at::Tensor q, at::Tensor paged_k_cache,
29
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
30
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
31
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
32
+ int64_t kv_layout_code, int64_t window_left,
33
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS);
34
+
35
+ TORCH_LIBRARY_FRAGMENT(TORCH_EXTENSION_NAME, m) {
36
+ // Batched decode with paged KV-Cache plan
37
+ m.def("plan", BatchDecodeWithPagedKVCachePlan);
38
+ // Batched decode with paged KV-Cache run
39
+ m.def("run", BatchDecodeWithPagedKVCacheRun);
40
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_kernel.cu ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #include <flashinfer/attention/decode.cuh>
2
+ #include "batch_decode_config.inc"
3
+
4
+ using namespace flashinfer;
5
+
6
+ namespace flashinfer {
7
+
8
+ template cudaError_t
9
+ BatchDecodeWithPagedKVCacheDispatched<128, PosEncodingMode::kNone, DefaultAttention<false, false, false, false>, Params>(
10
+ Params params, nv_bfloat16* tmp_v,
11
+ float* tmp_s, bool enable_pdl, cudaStream_t stream);
12
+
13
+ };
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode.cu ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include <flashinfer/attention/scheduler.cuh>
17
+ #include <flashinfer/pos_enc.cuh>
18
+ #include <flashinfer/utils.cuh>
19
+ #include <optional>
20
+
21
+ #include "batch_decode_config.inc"
22
+ #include "pytorch_conversion_utils.h"
23
+ #include "pytorch_extension_utils.h"
24
+
25
+ namespace flashinfer {
26
+
27
+ template <uint32_t HEAD_DIM, PosEncodingMode POS_ENCODING_MODE, typename AttentionVariant,
28
+ typename Params>
29
+ cudaError_t BatchDecodeWithPagedKVCacheDispatched(Params params, typename Params::DTypeO* tmp_v,
30
+ float* tmp_s, bool enable_pdl,
31
+ cudaStream_t stream);
32
+
33
+ } // namespace flashinfer
34
+
35
+ using namespace flashinfer;
36
+
37
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
38
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
39
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
40
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
41
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
42
+ at::Tensor empty_q_data, at::Tensor empty_kv_data) {
43
+ size_t float_workspace_size_in_bytes =
44
+ float_workspace_buffer.size(0) * float_workspace_buffer.element_size();
45
+ size_t int_workspace_size_in_bytes =
46
+ int_workspace_buffer.size(0) * int_workspace_buffer.element_size();
47
+
48
+ DecodePlanInfo plan_info;
49
+
50
+ auto q_scalar_type = empty_q_data.scalar_type();
51
+ auto kv_scalar_type = empty_kv_data.scalar_type();
52
+
53
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
54
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
55
+ "cores template for different head dim");
56
+
57
+ const c10::cuda::OptionalCUDAGuard device_guard(float_workspace_buffer.device());
58
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
59
+ DISPATCH_context(
60
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
61
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
62
+ DISPATCH_GQA_GROUP_SIZE(num_qo_heads / num_kv_heads, GROUP_SIZE, {
63
+ auto work_estimation_func = BatchDecodeWithPagedKVCacheWorkEstimationDispatched<
64
+ GROUP_SIZE, HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>;
65
+ cudaError_t status = DecodePlan<HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>(
66
+ static_cast<void*>(float_workspace_buffer.data_ptr()), float_workspace_size_in_bytes,
67
+ static_cast<void*>(int_workspace_buffer.data_ptr()),
68
+ static_cast<void*>(page_locked_int_workspace_buffer.data_ptr()),
69
+ int_workspace_size_in_bytes, plan_info, static_cast<IdType*>(indptr.data_ptr()),
70
+ batch_size, num_qo_heads, page_size, enable_cuda_graph,
71
+ /*stream=*/stream, work_estimation_func);
72
+
73
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
74
+ cudaGetErrorString(status));
75
+ return true;
76
+ });
77
+ });
78
+
79
+ return vec_to_tensor(plan_info.ToVector());
80
+ }
81
+
82
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
83
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
84
+ at::Tensor q, at::Tensor paged_k_cache,
85
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
86
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
87
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
88
+ int64_t kv_layout_code, int64_t window_left,
89
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS) {
90
+ DecodePlanInfo plan_info;
91
+ plan_info.FromVector(tensor_to_vec(plan_info_vec));
92
+ QKVLayout kv_layout = static_cast<QKVLayout>(kv_layout_code);
93
+ auto device = q.device();
94
+ int64_t batch_size = q.size(0);
95
+ int64_t num_qo_heads = q.size(1);
96
+ int64_t num_kv_heads, page_size;
97
+
98
+ if (kv_layout == QKVLayout::kHND) {
99
+ num_kv_heads = paged_k_cache.size(1);
100
+ page_size = paged_k_cache.size(2);
101
+ } else {
102
+ page_size = paged_k_cache.size(1);
103
+ num_kv_heads = paged_k_cache.size(2);
104
+ }
105
+ uint32_t head_dim_qk = q.size(2);
106
+ uint32_t head_dim_vo = paged_v_cache.size(3);
107
+
108
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
109
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
110
+ "cores template for different head dim");
111
+
112
+ if (maybe_lse) {
113
+ const auto& lse = *maybe_lse;
114
+ TORCH_CHECK(lse.size(0) == batch_size, lse.size(0), q.size(0));
115
+ TORCH_CHECK(lse.size(1) == num_qo_heads, lse.size(1), q.size(1));
116
+ }
117
+
118
+ void* float_buffer = static_cast<void*>(float_workspace_buffer.data_ptr());
119
+ void* int_buffer = static_cast<void*>(int_workspace_buffer.data_ptr());
120
+
121
+ // get q_scalar_type and kv_scalar_type
122
+ auto q_scalar_type = q.scalar_type();
123
+ auto kv_scalar_type = paged_k_cache.scalar_type();
124
+
125
+ // get q_stride_n and q_stride_h
126
+ const auto q_stride_n = q.stride(0);
127
+ const auto q_stride_h = q.stride(1);
128
+
129
+ // get kv_cache_strides
130
+ const int64_t* kv_cache_strides = nullptr;
131
+ auto k_strides = paged_k_cache.strides();
132
+ auto v_strides = paged_v_cache.strides();
133
+ TORCH_CHECK(k_strides == v_strides, "k/v strides must be identical");
134
+ kv_cache_strides = k_strides.data();
135
+
136
+ const c10::cuda::OptionalCUDAGuard device_guard(device);
137
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
138
+
139
+ DISPATCH_context(
140
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
141
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
142
+ paged_kv_t<DTypeKV, IdType> paged_kv(
143
+ num_kv_heads, page_size, HEAD_DIM_QK, batch_size, kv_layout,
144
+ static_cast<DTypeKV*>(paged_k_cache.data_ptr()),
145
+ static_cast<DTypeKV*>(paged_v_cache.data_ptr()), kv_cache_strides,
146
+ static_cast<IdType*>(paged_kv_indices.data_ptr()),
147
+ static_cast<IdType*>(paged_kv_indptr.data_ptr()),
148
+ static_cast<IdType*>(paged_kv_last_page_len.data_ptr()));
149
+
150
+ Params params;
151
+ params.q = static_cast<DTypeQ*>(q.data_ptr());
152
+ params.paged_kv = paged_kv;
153
+ params.o = static_cast<DTypeO*>(o.data_ptr());
154
+ params.lse = maybe_lse ? static_cast<float*>(maybe_lse->data_ptr()) : nullptr;
155
+ params.padded_batch_size = 0;
156
+ params.num_qo_heads = num_qo_heads;
157
+ params.q_stride_n = q_stride_n;
158
+ params.q_stride_h = q_stride_h;
159
+ params.window_left = window_left;
160
+ params.request_indices = nullptr;
161
+ params.kv_tile_indices = nullptr;
162
+ params.o_indptr = nullptr;
163
+ params.kv_chunk_size_ptr = nullptr;
164
+ params.block_valid_mask = nullptr;
165
+ params.partition_kv = false;
166
+
167
+ ADDITIONAL_PARAMS_SETTER
168
+
169
+ DTypeO* tmp_v = nullptr;
170
+ float* tmp_s = nullptr;
171
+ params.request_indices =
172
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.request_indices_offset);
173
+ params.kv_tile_indices =
174
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_tile_indices_offset);
175
+ params.o_indptr = GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.o_indptr_offset);
176
+ params.kv_chunk_size_ptr =
177
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_chunk_size_ptr_offset);
178
+ if (plan_info.split_kv) {
179
+ tmp_v = GetPtrFromBaseOffset<DTypeO>(float_buffer, plan_info.v_offset);
180
+ tmp_s = GetPtrFromBaseOffset<float>(float_buffer, plan_info.s_offset);
181
+ if (plan_info.enable_cuda_graph) {
182
+ params.block_valid_mask =
183
+ GetPtrFromBaseOffset<bool>(int_buffer, plan_info.block_valid_mask_offset);
184
+ }
185
+ }
186
+ params.padded_batch_size = plan_info.padded_batch_size;
187
+
188
+ cudaError_t status =
189
+ flashinfer::BatchDecodeWithPagedKVCacheDispatched<HEAD_DIM_QK, POS_ENCODING_MODE,
190
+ AttentionVariant>(params, tmp_v,
191
+ tmp_s, enable_pdl,
192
+ /*stream=*/stream);
193
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
194
+ cudaGetErrorString(status));
195
+ return true;
196
+ });
197
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_config.inc ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #pragma once
2
+ #include <flashinfer/page.cuh>
3
+ #include <flashinfer/math.cuh>
4
+ #include <flashinfer/layout.cuh>
5
+ #include <flashinfer/pos_enc.cuh>
6
+ #include <flashinfer/attention/variant_helper.cuh>
7
+
8
+ #define ADDITIONAL_FUNC_PARAMS , std::optional<at::Tensor> maybe_alibi_slopes, double logits_soft_cap, double sm_scale, double rope_rcp_scale, double rope_rcp_theta
9
+ #define ADDITIONAL_PARAMS_SETTER params.maybe_alibi_slopes = maybe_alibi_slopes ? static_cast<float*>(maybe_alibi_slopes->data_ptr()): nullptr; \
10
+ params.logits_soft_cap = logits_soft_cap; \
11
+ params.sm_scale = sm_scale; \
12
+ params.rope_rcp_scale = rope_rcp_scale; \
13
+ params.rope_rcp_theta = rope_rcp_theta;
14
+
15
+ #define DISPATCH_context(DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE, USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, ...) { \
16
+ using AttentionVariant = DefaultAttention<false, true, true, false>; \
17
+ __VA_ARGS__(); \
18
+ }
19
+
20
+ using namespace flashinfer;
21
+
22
+ using DTypeQ = nv_bfloat16;
23
+ using DTypeKV = __nv_fp8_e4m3;
24
+ using DTypeO = nv_bfloat16;
25
+ using IdType = int32_t;
26
+ constexpr int HEAD_DIM_QK = 256;
27
+ constexpr int HEAD_DIM_VO = 256;
28
+ constexpr auto USE_LOGITS_SOFT_CAP = true;
29
+ constexpr auto POS_ENCODING_MODE = PosEncodingMode::kNone;
30
+ constexpr auto USE_SLIDING_WINDOW = true;
31
+
32
+ struct Params {
33
+ using DTypeQ = DTypeQ;
34
+ using DTypeKV = DTypeKV;
35
+ using DTypeO = DTypeO;
36
+ using IdType = IdType;
37
+
38
+ DTypeQ* q;
39
+ paged_kv_t<DTypeKV, IdType> paged_kv;
40
+ DTypeO* o;
41
+ float* lse;
42
+
43
+ float* maybe_alibi_slopes;
44
+ double logits_soft_cap;
45
+ double sm_scale;
46
+ double rope_rcp_scale;
47
+ double rope_rcp_theta;
48
+
49
+
50
+ uint32_t padded_batch_size;
51
+ uint32_t num_qo_heads;
52
+ IdType q_stride_n;
53
+ IdType q_stride_h;
54
+ int32_t window_left;
55
+ bool enable_pdl;
56
+
57
+ IdType* request_indices;
58
+ IdType* kv_tile_indices;
59
+ IdType* o_indptr;
60
+ IdType* kv_chunk_size_ptr;
61
+ bool* block_valid_mask;
62
+ bool partition_kv;
63
+
64
+ __host__ __device__ __forceinline__ int32_t get_qo_len(int32_t batch_idx) const { return 1; }
65
+
66
+ __host__ __device__ __forceinline__ int32_t get_kv_len(int32_t batch_idx) const {
67
+ return paged_kv.get_length(batch_idx);
68
+ }
69
+ };
70
+
71
+ #include<flashinfer/attention/variants.cuh>
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_jit_pybind.cu ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023-2025 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include "batch_decode_config.inc"
17
+ #include "pytorch_extension_utils.h"
18
+
19
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
20
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
21
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
22
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
23
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
24
+ at::Tensor empty_q_data, at::Tensor empty_kv_data);
25
+
26
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
27
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
28
+ at::Tensor q, at::Tensor paged_k_cache,
29
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
30
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
31
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
32
+ int64_t kv_layout_code, int64_t window_left,
33
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS);
34
+
35
+ TORCH_LIBRARY_FRAGMENT(TORCH_EXTENSION_NAME, m) {
36
+ // Batched decode with paged KV-Cache plan
37
+ m.def("plan", BatchDecodeWithPagedKVCachePlan);
38
+ // Batched decode with paged KV-Cache run
39
+ m.def("run", BatchDecodeWithPagedKVCacheRun);
40
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_kernel.cu ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #include <flashinfer/attention/decode.cuh>
2
+ #include "batch_decode_config.inc"
3
+
4
+ using namespace flashinfer;
5
+
6
+ namespace flashinfer {
7
+
8
+ template cudaError_t
9
+ BatchDecodeWithPagedKVCacheDispatched<256, PosEncodingMode::kNone, DefaultAttention<false, true, true, false>, Params>(
10
+ Params params, nv_bfloat16* tmp_v,
11
+ float* tmp_s, bool enable_pdl, cudaStream_t stream);
12
+
13
+ };
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode.cu ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include <flashinfer/attention/scheduler.cuh>
17
+ #include <flashinfer/pos_enc.cuh>
18
+ #include <flashinfer/utils.cuh>
19
+ #include <optional>
20
+
21
+ #include "batch_decode_config.inc"
22
+ #include "pytorch_conversion_utils.h"
23
+ #include "pytorch_extension_utils.h"
24
+
25
+ namespace flashinfer {
26
+
27
+ template <uint32_t HEAD_DIM, PosEncodingMode POS_ENCODING_MODE, typename AttentionVariant,
28
+ typename Params>
29
+ cudaError_t BatchDecodeWithPagedKVCacheDispatched(Params params, typename Params::DTypeO* tmp_v,
30
+ float* tmp_s, bool enable_pdl,
31
+ cudaStream_t stream);
32
+
33
+ } // namespace flashinfer
34
+
35
+ using namespace flashinfer;
36
+
37
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
38
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
39
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
40
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
41
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
42
+ at::Tensor empty_q_data, at::Tensor empty_kv_data) {
43
+ size_t float_workspace_size_in_bytes =
44
+ float_workspace_buffer.size(0) * float_workspace_buffer.element_size();
45
+ size_t int_workspace_size_in_bytes =
46
+ int_workspace_buffer.size(0) * int_workspace_buffer.element_size();
47
+
48
+ DecodePlanInfo plan_info;
49
+
50
+ auto q_scalar_type = empty_q_data.scalar_type();
51
+ auto kv_scalar_type = empty_kv_data.scalar_type();
52
+
53
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
54
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
55
+ "cores template for different head dim");
56
+
57
+ const c10::cuda::OptionalCUDAGuard device_guard(float_workspace_buffer.device());
58
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
59
+ DISPATCH_context(
60
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
61
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
62
+ DISPATCH_GQA_GROUP_SIZE(num_qo_heads / num_kv_heads, GROUP_SIZE, {
63
+ auto work_estimation_func = BatchDecodeWithPagedKVCacheWorkEstimationDispatched<
64
+ GROUP_SIZE, HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>;
65
+ cudaError_t status = DecodePlan<HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>(
66
+ static_cast<void*>(float_workspace_buffer.data_ptr()), float_workspace_size_in_bytes,
67
+ static_cast<void*>(int_workspace_buffer.data_ptr()),
68
+ static_cast<void*>(page_locked_int_workspace_buffer.data_ptr()),
69
+ int_workspace_size_in_bytes, plan_info, static_cast<IdType*>(indptr.data_ptr()),
70
+ batch_size, num_qo_heads, page_size, enable_cuda_graph,
71
+ /*stream=*/stream, work_estimation_func);
72
+
73
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
74
+ cudaGetErrorString(status));
75
+ return true;
76
+ });
77
+ });
78
+
79
+ return vec_to_tensor(plan_info.ToVector());
80
+ }
81
+
82
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
83
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
84
+ at::Tensor q, at::Tensor paged_k_cache,
85
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
86
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
87
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
88
+ int64_t kv_layout_code, int64_t window_left,
89
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS) {
90
+ DecodePlanInfo plan_info;
91
+ plan_info.FromVector(tensor_to_vec(plan_info_vec));
92
+ QKVLayout kv_layout = static_cast<QKVLayout>(kv_layout_code);
93
+ auto device = q.device();
94
+ int64_t batch_size = q.size(0);
95
+ int64_t num_qo_heads = q.size(1);
96
+ int64_t num_kv_heads, page_size;
97
+
98
+ if (kv_layout == QKVLayout::kHND) {
99
+ num_kv_heads = paged_k_cache.size(1);
100
+ page_size = paged_k_cache.size(2);
101
+ } else {
102
+ page_size = paged_k_cache.size(1);
103
+ num_kv_heads = paged_k_cache.size(2);
104
+ }
105
+ uint32_t head_dim_qk = q.size(2);
106
+ uint32_t head_dim_vo = paged_v_cache.size(3);
107
+
108
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
109
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
110
+ "cores template for different head dim");
111
+
112
+ if (maybe_lse) {
113
+ const auto& lse = *maybe_lse;
114
+ TORCH_CHECK(lse.size(0) == batch_size, lse.size(0), q.size(0));
115
+ TORCH_CHECK(lse.size(1) == num_qo_heads, lse.size(1), q.size(1));
116
+ }
117
+
118
+ void* float_buffer = static_cast<void*>(float_workspace_buffer.data_ptr());
119
+ void* int_buffer = static_cast<void*>(int_workspace_buffer.data_ptr());
120
+
121
+ // get q_scalar_type and kv_scalar_type
122
+ auto q_scalar_type = q.scalar_type();
123
+ auto kv_scalar_type = paged_k_cache.scalar_type();
124
+
125
+ // get q_stride_n and q_stride_h
126
+ const auto q_stride_n = q.stride(0);
127
+ const auto q_stride_h = q.stride(1);
128
+
129
+ // get kv_cache_strides
130
+ const int64_t* kv_cache_strides = nullptr;
131
+ auto k_strides = paged_k_cache.strides();
132
+ auto v_strides = paged_v_cache.strides();
133
+ TORCH_CHECK(k_strides == v_strides, "k/v strides must be identical");
134
+ kv_cache_strides = k_strides.data();
135
+
136
+ const c10::cuda::OptionalCUDAGuard device_guard(device);
137
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
138
+
139
+ DISPATCH_context(
140
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
141
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
142
+ paged_kv_t<DTypeKV, IdType> paged_kv(
143
+ num_kv_heads, page_size, HEAD_DIM_QK, batch_size, kv_layout,
144
+ static_cast<DTypeKV*>(paged_k_cache.data_ptr()),
145
+ static_cast<DTypeKV*>(paged_v_cache.data_ptr()), kv_cache_strides,
146
+ static_cast<IdType*>(paged_kv_indices.data_ptr()),
147
+ static_cast<IdType*>(paged_kv_indptr.data_ptr()),
148
+ static_cast<IdType*>(paged_kv_last_page_len.data_ptr()));
149
+
150
+ Params params;
151
+ params.q = static_cast<DTypeQ*>(q.data_ptr());
152
+ params.paged_kv = paged_kv;
153
+ params.o = static_cast<DTypeO*>(o.data_ptr());
154
+ params.lse = maybe_lse ? static_cast<float*>(maybe_lse->data_ptr()) : nullptr;
155
+ params.padded_batch_size = 0;
156
+ params.num_qo_heads = num_qo_heads;
157
+ params.q_stride_n = q_stride_n;
158
+ params.q_stride_h = q_stride_h;
159
+ params.window_left = window_left;
160
+ params.request_indices = nullptr;
161
+ params.kv_tile_indices = nullptr;
162
+ params.o_indptr = nullptr;
163
+ params.kv_chunk_size_ptr = nullptr;
164
+ params.block_valid_mask = nullptr;
165
+ params.partition_kv = false;
166
+
167
+ ADDITIONAL_PARAMS_SETTER
168
+
169
+ DTypeO* tmp_v = nullptr;
170
+ float* tmp_s = nullptr;
171
+ params.request_indices =
172
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.request_indices_offset);
173
+ params.kv_tile_indices =
174
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_tile_indices_offset);
175
+ params.o_indptr = GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.o_indptr_offset);
176
+ params.kv_chunk_size_ptr =
177
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_chunk_size_ptr_offset);
178
+ if (plan_info.split_kv) {
179
+ tmp_v = GetPtrFromBaseOffset<DTypeO>(float_buffer, plan_info.v_offset);
180
+ tmp_s = GetPtrFromBaseOffset<float>(float_buffer, plan_info.s_offset);
181
+ if (plan_info.enable_cuda_graph) {
182
+ params.block_valid_mask =
183
+ GetPtrFromBaseOffset<bool>(int_buffer, plan_info.block_valid_mask_offset);
184
+ }
185
+ }
186
+ params.padded_batch_size = plan_info.padded_batch_size;
187
+
188
+ cudaError_t status =
189
+ flashinfer::BatchDecodeWithPagedKVCacheDispatched<HEAD_DIM_QK, POS_ENCODING_MODE,
190
+ AttentionVariant>(params, tmp_v,
191
+ tmp_s, enable_pdl,
192
+ /*stream=*/stream);
193
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
194
+ cudaGetErrorString(status));
195
+ return true;
196
+ });
197
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_config.inc ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #pragma once
2
+ #include <flashinfer/page.cuh>
3
+ #include <flashinfer/math.cuh>
4
+ #include <flashinfer/layout.cuh>
5
+ #include <flashinfer/pos_enc.cuh>
6
+ #include <flashinfer/attention/variant_helper.cuh>
7
+
8
+ #define ADDITIONAL_FUNC_PARAMS , std::optional<at::Tensor> maybe_alibi_slopes, double logits_soft_cap, double sm_scale, double rope_rcp_scale, double rope_rcp_theta
9
+ #define ADDITIONAL_PARAMS_SETTER params.maybe_alibi_slopes = maybe_alibi_slopes ? static_cast<float*>(maybe_alibi_slopes->data_ptr()): nullptr; \
10
+ params.logits_soft_cap = logits_soft_cap; \
11
+ params.sm_scale = sm_scale; \
12
+ params.rope_rcp_scale = rope_rcp_scale; \
13
+ params.rope_rcp_theta = rope_rcp_theta;
14
+
15
+ #define DISPATCH_context(DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE, USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, ...) { \
16
+ using AttentionVariant = DefaultAttention<false, false, false, false>; \
17
+ __VA_ARGS__(); \
18
+ }
19
+
20
+ using namespace flashinfer;
21
+
22
+ using DTypeQ = nv_bfloat16;
23
+ using DTypeKV = __nv_fp8_e4m3;
24
+ using DTypeO = nv_bfloat16;
25
+ using IdType = int32_t;
26
+ constexpr int HEAD_DIM_QK = 64;
27
+ constexpr int HEAD_DIM_VO = 64;
28
+ constexpr auto USE_LOGITS_SOFT_CAP = false;
29
+ constexpr auto POS_ENCODING_MODE = PosEncodingMode::kNone;
30
+ constexpr auto USE_SLIDING_WINDOW = false;
31
+
32
+ struct Params {
33
+ using DTypeQ = DTypeQ;
34
+ using DTypeKV = DTypeKV;
35
+ using DTypeO = DTypeO;
36
+ using IdType = IdType;
37
+
38
+ DTypeQ* q;
39
+ paged_kv_t<DTypeKV, IdType> paged_kv;
40
+ DTypeO* o;
41
+ float* lse;
42
+
43
+ float* maybe_alibi_slopes;
44
+ double logits_soft_cap;
45
+ double sm_scale;
46
+ double rope_rcp_scale;
47
+ double rope_rcp_theta;
48
+
49
+
50
+ uint32_t padded_batch_size;
51
+ uint32_t num_qo_heads;
52
+ IdType q_stride_n;
53
+ IdType q_stride_h;
54
+ int32_t window_left;
55
+ bool enable_pdl;
56
+
57
+ IdType* request_indices;
58
+ IdType* kv_tile_indices;
59
+ IdType* o_indptr;
60
+ IdType* kv_chunk_size_ptr;
61
+ bool* block_valid_mask;
62
+ bool partition_kv;
63
+
64
+ __host__ __device__ __forceinline__ int32_t get_qo_len(int32_t batch_idx) const { return 1; }
65
+
66
+ __host__ __device__ __forceinline__ int32_t get_kv_len(int32_t batch_idx) const {
67
+ return paged_kv.get_length(batch_idx);
68
+ }
69
+ };
70
+
71
+ #include<flashinfer/attention/variants.cuh>
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_jit_pybind.cu ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023-2025 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include "batch_decode_config.inc"
17
+ #include "pytorch_extension_utils.h"
18
+
19
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
20
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
21
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
22
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
23
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
24
+ at::Tensor empty_q_data, at::Tensor empty_kv_data);
25
+
26
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
27
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
28
+ at::Tensor q, at::Tensor paged_k_cache,
29
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
30
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
31
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
32
+ int64_t kv_layout_code, int64_t window_left,
33
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS);
34
+
35
+ TORCH_LIBRARY_FRAGMENT(TORCH_EXTENSION_NAME, m) {
36
+ // Batched decode with paged KV-Cache plan
37
+ m.def("plan", BatchDecodeWithPagedKVCachePlan);
38
+ // Batched decode with paged KV-Cache run
39
+ m.def("run", BatchDecodeWithPagedKVCacheRun);
40
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_kernel.cu ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #include <flashinfer/attention/decode.cuh>
2
+ #include "batch_decode_config.inc"
3
+
4
+ using namespace flashinfer;
5
+
6
+ namespace flashinfer {
7
+
8
+ template cudaError_t
9
+ BatchDecodeWithPagedKVCacheDispatched<64, PosEncodingMode::kNone, DefaultAttention<false, false, false, false>, Params>(
10
+ Params params, nv_bfloat16* tmp_v,
11
+ float* tmp_s, bool enable_pdl, cudaStream_t stream);
12
+
13
+ };
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode.cu ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include <flashinfer/attention/scheduler.cuh>
17
+ #include <flashinfer/pos_enc.cuh>
18
+ #include <flashinfer/utils.cuh>
19
+ #include <optional>
20
+
21
+ #include "batch_decode_config.inc"
22
+ #include "pytorch_conversion_utils.h"
23
+ #include "pytorch_extension_utils.h"
24
+
25
+ namespace flashinfer {
26
+
27
+ template <uint32_t HEAD_DIM, PosEncodingMode POS_ENCODING_MODE, typename AttentionVariant,
28
+ typename Params>
29
+ cudaError_t BatchDecodeWithPagedKVCacheDispatched(Params params, typename Params::DTypeO* tmp_v,
30
+ float* tmp_s, bool enable_pdl,
31
+ cudaStream_t stream);
32
+
33
+ } // namespace flashinfer
34
+
35
+ using namespace flashinfer;
36
+
37
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
38
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
39
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
40
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
41
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
42
+ at::Tensor empty_q_data, at::Tensor empty_kv_data) {
43
+ size_t float_workspace_size_in_bytes =
44
+ float_workspace_buffer.size(0) * float_workspace_buffer.element_size();
45
+ size_t int_workspace_size_in_bytes =
46
+ int_workspace_buffer.size(0) * int_workspace_buffer.element_size();
47
+
48
+ DecodePlanInfo plan_info;
49
+
50
+ auto q_scalar_type = empty_q_data.scalar_type();
51
+ auto kv_scalar_type = empty_kv_data.scalar_type();
52
+
53
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
54
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
55
+ "cores template for different head dim");
56
+
57
+ const c10::cuda::OptionalCUDAGuard device_guard(float_workspace_buffer.device());
58
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
59
+ DISPATCH_context(
60
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
61
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
62
+ DISPATCH_GQA_GROUP_SIZE(num_qo_heads / num_kv_heads, GROUP_SIZE, {
63
+ auto work_estimation_func = BatchDecodeWithPagedKVCacheWorkEstimationDispatched<
64
+ GROUP_SIZE, HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>;
65
+ cudaError_t status = DecodePlan<HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>(
66
+ static_cast<void*>(float_workspace_buffer.data_ptr()), float_workspace_size_in_bytes,
67
+ static_cast<void*>(int_workspace_buffer.data_ptr()),
68
+ static_cast<void*>(page_locked_int_workspace_buffer.data_ptr()),
69
+ int_workspace_size_in_bytes, plan_info, static_cast<IdType*>(indptr.data_ptr()),
70
+ batch_size, num_qo_heads, page_size, enable_cuda_graph,
71
+ /*stream=*/stream, work_estimation_func);
72
+
73
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
74
+ cudaGetErrorString(status));
75
+ return true;
76
+ });
77
+ });
78
+
79
+ return vec_to_tensor(plan_info.ToVector());
80
+ }
81
+
82
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
83
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
84
+ at::Tensor q, at::Tensor paged_k_cache,
85
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
86
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
87
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
88
+ int64_t kv_layout_code, int64_t window_left,
89
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS) {
90
+ DecodePlanInfo plan_info;
91
+ plan_info.FromVector(tensor_to_vec(plan_info_vec));
92
+ QKVLayout kv_layout = static_cast<QKVLayout>(kv_layout_code);
93
+ auto device = q.device();
94
+ int64_t batch_size = q.size(0);
95
+ int64_t num_qo_heads = q.size(1);
96
+ int64_t num_kv_heads, page_size;
97
+
98
+ if (kv_layout == QKVLayout::kHND) {
99
+ num_kv_heads = paged_k_cache.size(1);
100
+ page_size = paged_k_cache.size(2);
101
+ } else {
102
+ page_size = paged_k_cache.size(1);
103
+ num_kv_heads = paged_k_cache.size(2);
104
+ }
105
+ uint32_t head_dim_qk = q.size(2);
106
+ uint32_t head_dim_vo = paged_v_cache.size(3);
107
+
108
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
109
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
110
+ "cores template for different head dim");
111
+
112
+ if (maybe_lse) {
113
+ const auto& lse = *maybe_lse;
114
+ TORCH_CHECK(lse.size(0) == batch_size, lse.size(0), q.size(0));
115
+ TORCH_CHECK(lse.size(1) == num_qo_heads, lse.size(1), q.size(1));
116
+ }
117
+
118
+ void* float_buffer = static_cast<void*>(float_workspace_buffer.data_ptr());
119
+ void* int_buffer = static_cast<void*>(int_workspace_buffer.data_ptr());
120
+
121
+ // get q_scalar_type and kv_scalar_type
122
+ auto q_scalar_type = q.scalar_type();
123
+ auto kv_scalar_type = paged_k_cache.scalar_type();
124
+
125
+ // get q_stride_n and q_stride_h
126
+ const auto q_stride_n = q.stride(0);
127
+ const auto q_stride_h = q.stride(1);
128
+
129
+ // get kv_cache_strides
130
+ const int64_t* kv_cache_strides = nullptr;
131
+ auto k_strides = paged_k_cache.strides();
132
+ auto v_strides = paged_v_cache.strides();
133
+ TORCH_CHECK(k_strides == v_strides, "k/v strides must be identical");
134
+ kv_cache_strides = k_strides.data();
135
+
136
+ const c10::cuda::OptionalCUDAGuard device_guard(device);
137
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
138
+
139
+ DISPATCH_context(
140
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
141
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
142
+ paged_kv_t<DTypeKV, IdType> paged_kv(
143
+ num_kv_heads, page_size, HEAD_DIM_QK, batch_size, kv_layout,
144
+ static_cast<DTypeKV*>(paged_k_cache.data_ptr()),
145
+ static_cast<DTypeKV*>(paged_v_cache.data_ptr()), kv_cache_strides,
146
+ static_cast<IdType*>(paged_kv_indices.data_ptr()),
147
+ static_cast<IdType*>(paged_kv_indptr.data_ptr()),
148
+ static_cast<IdType*>(paged_kv_last_page_len.data_ptr()));
149
+
150
+ Params params;
151
+ params.q = static_cast<DTypeQ*>(q.data_ptr());
152
+ params.paged_kv = paged_kv;
153
+ params.o = static_cast<DTypeO*>(o.data_ptr());
154
+ params.lse = maybe_lse ? static_cast<float*>(maybe_lse->data_ptr()) : nullptr;
155
+ params.padded_batch_size = 0;
156
+ params.num_qo_heads = num_qo_heads;
157
+ params.q_stride_n = q_stride_n;
158
+ params.q_stride_h = q_stride_h;
159
+ params.window_left = window_left;
160
+ params.request_indices = nullptr;
161
+ params.kv_tile_indices = nullptr;
162
+ params.o_indptr = nullptr;
163
+ params.kv_chunk_size_ptr = nullptr;
164
+ params.block_valid_mask = nullptr;
165
+ params.partition_kv = false;
166
+
167
+ ADDITIONAL_PARAMS_SETTER
168
+
169
+ DTypeO* tmp_v = nullptr;
170
+ float* tmp_s = nullptr;
171
+ params.request_indices =
172
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.request_indices_offset);
173
+ params.kv_tile_indices =
174
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_tile_indices_offset);
175
+ params.o_indptr = GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.o_indptr_offset);
176
+ params.kv_chunk_size_ptr =
177
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_chunk_size_ptr_offset);
178
+ if (plan_info.split_kv) {
179
+ tmp_v = GetPtrFromBaseOffset<DTypeO>(float_buffer, plan_info.v_offset);
180
+ tmp_s = GetPtrFromBaseOffset<float>(float_buffer, plan_info.s_offset);
181
+ if (plan_info.enable_cuda_graph) {
182
+ params.block_valid_mask =
183
+ GetPtrFromBaseOffset<bool>(int_buffer, plan_info.block_valid_mask_offset);
184
+ }
185
+ }
186
+ params.padded_batch_size = plan_info.padded_batch_size;
187
+
188
+ cudaError_t status =
189
+ flashinfer::BatchDecodeWithPagedKVCacheDispatched<HEAD_DIM_QK, POS_ENCODING_MODE,
190
+ AttentionVariant>(params, tmp_v,
191
+ tmp_s, enable_pdl,
192
+ /*stream=*/stream);
193
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
194
+ cudaGetErrorString(status));
195
+ return true;
196
+ });
197
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_config.inc ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #pragma once
2
+ #include <flashinfer/page.cuh>
3
+ #include <flashinfer/math.cuh>
4
+ #include <flashinfer/layout.cuh>
5
+ #include <flashinfer/pos_enc.cuh>
6
+ #include <flashinfer/attention/variant_helper.cuh>
7
+
8
+ #define ADDITIONAL_FUNC_PARAMS , std::optional<at::Tensor> maybe_alibi_slopes, double logits_soft_cap, double sm_scale, double rope_rcp_scale, double rope_rcp_theta
9
+ #define ADDITIONAL_PARAMS_SETTER params.maybe_alibi_slopes = maybe_alibi_slopes ? static_cast<float*>(maybe_alibi_slopes->data_ptr()): nullptr; \
10
+ params.logits_soft_cap = logits_soft_cap; \
11
+ params.sm_scale = sm_scale; \
12
+ params.rope_rcp_scale = rope_rcp_scale; \
13
+ params.rope_rcp_theta = rope_rcp_theta;
14
+
15
+ #define DISPATCH_context(DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE, USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, ...) { \
16
+ using AttentionVariant = DefaultAttention<false, true, false, false>; \
17
+ __VA_ARGS__(); \
18
+ }
19
+
20
+ using namespace flashinfer;
21
+
22
+ using DTypeQ = nv_bfloat16;
23
+ using DTypeKV = __nv_fp8_e4m3;
24
+ using DTypeO = nv_bfloat16;
25
+ using IdType = int32_t;
26
+ constexpr int HEAD_DIM_QK = 64;
27
+ constexpr int HEAD_DIM_VO = 64;
28
+ constexpr auto USE_LOGITS_SOFT_CAP = false;
29
+ constexpr auto POS_ENCODING_MODE = PosEncodingMode::kNone;
30
+ constexpr auto USE_SLIDING_WINDOW = true;
31
+
32
+ struct Params {
33
+ using DTypeQ = DTypeQ;
34
+ using DTypeKV = DTypeKV;
35
+ using DTypeO = DTypeO;
36
+ using IdType = IdType;
37
+
38
+ DTypeQ* q;
39
+ paged_kv_t<DTypeKV, IdType> paged_kv;
40
+ DTypeO* o;
41
+ float* lse;
42
+
43
+ float* maybe_alibi_slopes;
44
+ double logits_soft_cap;
45
+ double sm_scale;
46
+ double rope_rcp_scale;
47
+ double rope_rcp_theta;
48
+
49
+
50
+ uint32_t padded_batch_size;
51
+ uint32_t num_qo_heads;
52
+ IdType q_stride_n;
53
+ IdType q_stride_h;
54
+ int32_t window_left;
55
+ bool enable_pdl;
56
+
57
+ IdType* request_indices;
58
+ IdType* kv_tile_indices;
59
+ IdType* o_indptr;
60
+ IdType* kv_chunk_size_ptr;
61
+ bool* block_valid_mask;
62
+ bool partition_kv;
63
+
64
+ __host__ __device__ __forceinline__ int32_t get_qo_len(int32_t batch_idx) const { return 1; }
65
+
66
+ __host__ __device__ __forceinline__ int32_t get_kv_len(int32_t batch_idx) const {
67
+ return paged_kv.get_length(batch_idx);
68
+ }
69
+ };
70
+
71
+ #include<flashinfer/attention/variants.cuh>
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_jit_pybind.cu ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023-2025 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include "batch_decode_config.inc"
17
+ #include "pytorch_extension_utils.h"
18
+
19
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
20
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
21
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
22
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
23
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
24
+ at::Tensor empty_q_data, at::Tensor empty_kv_data);
25
+
26
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
27
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
28
+ at::Tensor q, at::Tensor paged_k_cache,
29
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
30
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
31
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
32
+ int64_t kv_layout_code, int64_t window_left,
33
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS);
34
+
35
+ TORCH_LIBRARY_FRAGMENT(TORCH_EXTENSION_NAME, m) {
36
+ // Batched decode with paged KV-Cache plan
37
+ m.def("plan", BatchDecodeWithPagedKVCachePlan);
38
+ // Batched decode with paged KV-Cache run
39
+ m.def("run", BatchDecodeWithPagedKVCacheRun);
40
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_bf16_dtype_kv_e4m3_dtype_o_bf16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_True_use_logits_cap_False/batch_decode_kernel.cu ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #include <flashinfer/attention/decode.cuh>
2
+ #include "batch_decode_config.inc"
3
+
4
+ using namespace flashinfer;
5
+
6
+ namespace flashinfer {
7
+
8
+ template cudaError_t
9
+ BatchDecodeWithPagedKVCacheDispatched<64, PosEncodingMode::kNone, DefaultAttention<false, true, false, false>, Params>(
10
+ Params params, nv_bfloat16* tmp_v,
11
+ float* tmp_s, bool enable_pdl, cudaStream_t stream);
12
+
13
+ };
csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode.cu ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include <flashinfer/attention/scheduler.cuh>
17
+ #include <flashinfer/pos_enc.cuh>
18
+ #include <flashinfer/utils.cuh>
19
+ #include <optional>
20
+
21
+ #include "batch_decode_config.inc"
22
+ #include "pytorch_conversion_utils.h"
23
+ #include "pytorch_extension_utils.h"
24
+
25
+ namespace flashinfer {
26
+
27
+ template <uint32_t HEAD_DIM, PosEncodingMode POS_ENCODING_MODE, typename AttentionVariant,
28
+ typename Params>
29
+ cudaError_t BatchDecodeWithPagedKVCacheDispatched(Params params, typename Params::DTypeO* tmp_v,
30
+ float* tmp_s, bool enable_pdl,
31
+ cudaStream_t stream);
32
+
33
+ } // namespace flashinfer
34
+
35
+ using namespace flashinfer;
36
+
37
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
38
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
39
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
40
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
41
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
42
+ at::Tensor empty_q_data, at::Tensor empty_kv_data) {
43
+ size_t float_workspace_size_in_bytes =
44
+ float_workspace_buffer.size(0) * float_workspace_buffer.element_size();
45
+ size_t int_workspace_size_in_bytes =
46
+ int_workspace_buffer.size(0) * int_workspace_buffer.element_size();
47
+
48
+ DecodePlanInfo plan_info;
49
+
50
+ auto q_scalar_type = empty_q_data.scalar_type();
51
+ auto kv_scalar_type = empty_kv_data.scalar_type();
52
+
53
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
54
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
55
+ "cores template for different head dim");
56
+
57
+ const c10::cuda::OptionalCUDAGuard device_guard(float_workspace_buffer.device());
58
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
59
+ DISPATCH_context(
60
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
61
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
62
+ DISPATCH_GQA_GROUP_SIZE(num_qo_heads / num_kv_heads, GROUP_SIZE, {
63
+ auto work_estimation_func = BatchDecodeWithPagedKVCacheWorkEstimationDispatched<
64
+ GROUP_SIZE, HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>;
65
+ cudaError_t status = DecodePlan<HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>(
66
+ static_cast<void*>(float_workspace_buffer.data_ptr()), float_workspace_size_in_bytes,
67
+ static_cast<void*>(int_workspace_buffer.data_ptr()),
68
+ static_cast<void*>(page_locked_int_workspace_buffer.data_ptr()),
69
+ int_workspace_size_in_bytes, plan_info, static_cast<IdType*>(indptr.data_ptr()),
70
+ batch_size, num_qo_heads, page_size, enable_cuda_graph,
71
+ /*stream=*/stream, work_estimation_func);
72
+
73
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
74
+ cudaGetErrorString(status));
75
+ return true;
76
+ });
77
+ });
78
+
79
+ return vec_to_tensor(plan_info.ToVector());
80
+ }
81
+
82
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
83
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
84
+ at::Tensor q, at::Tensor paged_k_cache,
85
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
86
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
87
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
88
+ int64_t kv_layout_code, int64_t window_left,
89
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS) {
90
+ DecodePlanInfo plan_info;
91
+ plan_info.FromVector(tensor_to_vec(plan_info_vec));
92
+ QKVLayout kv_layout = static_cast<QKVLayout>(kv_layout_code);
93
+ auto device = q.device();
94
+ int64_t batch_size = q.size(0);
95
+ int64_t num_qo_heads = q.size(1);
96
+ int64_t num_kv_heads, page_size;
97
+
98
+ if (kv_layout == QKVLayout::kHND) {
99
+ num_kv_heads = paged_k_cache.size(1);
100
+ page_size = paged_k_cache.size(2);
101
+ } else {
102
+ page_size = paged_k_cache.size(1);
103
+ num_kv_heads = paged_k_cache.size(2);
104
+ }
105
+ uint32_t head_dim_qk = q.size(2);
106
+ uint32_t head_dim_vo = paged_v_cache.size(3);
107
+
108
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
109
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
110
+ "cores template for different head dim");
111
+
112
+ if (maybe_lse) {
113
+ const auto& lse = *maybe_lse;
114
+ TORCH_CHECK(lse.size(0) == batch_size, lse.size(0), q.size(0));
115
+ TORCH_CHECK(lse.size(1) == num_qo_heads, lse.size(1), q.size(1));
116
+ }
117
+
118
+ void* float_buffer = static_cast<void*>(float_workspace_buffer.data_ptr());
119
+ void* int_buffer = static_cast<void*>(int_workspace_buffer.data_ptr());
120
+
121
+ // get q_scalar_type and kv_scalar_type
122
+ auto q_scalar_type = q.scalar_type();
123
+ auto kv_scalar_type = paged_k_cache.scalar_type();
124
+
125
+ // get q_stride_n and q_stride_h
126
+ const auto q_stride_n = q.stride(0);
127
+ const auto q_stride_h = q.stride(1);
128
+
129
+ // get kv_cache_strides
130
+ const int64_t* kv_cache_strides = nullptr;
131
+ auto k_strides = paged_k_cache.strides();
132
+ auto v_strides = paged_v_cache.strides();
133
+ TORCH_CHECK(k_strides == v_strides, "k/v strides must be identical");
134
+ kv_cache_strides = k_strides.data();
135
+
136
+ const c10::cuda::OptionalCUDAGuard device_guard(device);
137
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
138
+
139
+ DISPATCH_context(
140
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
141
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
142
+ paged_kv_t<DTypeKV, IdType> paged_kv(
143
+ num_kv_heads, page_size, HEAD_DIM_QK, batch_size, kv_layout,
144
+ static_cast<DTypeKV*>(paged_k_cache.data_ptr()),
145
+ static_cast<DTypeKV*>(paged_v_cache.data_ptr()), kv_cache_strides,
146
+ static_cast<IdType*>(paged_kv_indices.data_ptr()),
147
+ static_cast<IdType*>(paged_kv_indptr.data_ptr()),
148
+ static_cast<IdType*>(paged_kv_last_page_len.data_ptr()));
149
+
150
+ Params params;
151
+ params.q = static_cast<DTypeQ*>(q.data_ptr());
152
+ params.paged_kv = paged_kv;
153
+ params.o = static_cast<DTypeO*>(o.data_ptr());
154
+ params.lse = maybe_lse ? static_cast<float*>(maybe_lse->data_ptr()) : nullptr;
155
+ params.padded_batch_size = 0;
156
+ params.num_qo_heads = num_qo_heads;
157
+ params.q_stride_n = q_stride_n;
158
+ params.q_stride_h = q_stride_h;
159
+ params.window_left = window_left;
160
+ params.request_indices = nullptr;
161
+ params.kv_tile_indices = nullptr;
162
+ params.o_indptr = nullptr;
163
+ params.kv_chunk_size_ptr = nullptr;
164
+ params.block_valid_mask = nullptr;
165
+ params.partition_kv = false;
166
+
167
+ ADDITIONAL_PARAMS_SETTER
168
+
169
+ DTypeO* tmp_v = nullptr;
170
+ float* tmp_s = nullptr;
171
+ params.request_indices =
172
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.request_indices_offset);
173
+ params.kv_tile_indices =
174
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_tile_indices_offset);
175
+ params.o_indptr = GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.o_indptr_offset);
176
+ params.kv_chunk_size_ptr =
177
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_chunk_size_ptr_offset);
178
+ if (plan_info.split_kv) {
179
+ tmp_v = GetPtrFromBaseOffset<DTypeO>(float_buffer, plan_info.v_offset);
180
+ tmp_s = GetPtrFromBaseOffset<float>(float_buffer, plan_info.s_offset);
181
+ if (plan_info.enable_cuda_graph) {
182
+ params.block_valid_mask =
183
+ GetPtrFromBaseOffset<bool>(int_buffer, plan_info.block_valid_mask_offset);
184
+ }
185
+ }
186
+ params.padded_batch_size = plan_info.padded_batch_size;
187
+
188
+ cudaError_t status =
189
+ flashinfer::BatchDecodeWithPagedKVCacheDispatched<HEAD_DIM_QK, POS_ENCODING_MODE,
190
+ AttentionVariant>(params, tmp_v,
191
+ tmp_s, enable_pdl,
192
+ /*stream=*/stream);
193
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
194
+ cudaGetErrorString(status));
195
+ return true;
196
+ });
197
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_config.inc ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #pragma once
2
+ #include <flashinfer/page.cuh>
3
+ #include <flashinfer/math.cuh>
4
+ #include <flashinfer/layout.cuh>
5
+ #include <flashinfer/pos_enc.cuh>
6
+ #include <flashinfer/attention/variant_helper.cuh>
7
+
8
+ #define ADDITIONAL_FUNC_PARAMS , std::optional<at::Tensor> maybe_alibi_slopes, double logits_soft_cap, double sm_scale, double rope_rcp_scale, double rope_rcp_theta
9
+ #define ADDITIONAL_PARAMS_SETTER params.maybe_alibi_slopes = maybe_alibi_slopes ? static_cast<float*>(maybe_alibi_slopes->data_ptr()): nullptr; \
10
+ params.logits_soft_cap = logits_soft_cap; \
11
+ params.sm_scale = sm_scale; \
12
+ params.rope_rcp_scale = rope_rcp_scale; \
13
+ params.rope_rcp_theta = rope_rcp_theta;
14
+
15
+ #define DISPATCH_context(DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE, USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, ...) { \
16
+ using AttentionVariant = DefaultAttention<false, false, false, false>; \
17
+ __VA_ARGS__(); \
18
+ }
19
+
20
+ using namespace flashinfer;
21
+
22
+ using DTypeQ = half;
23
+ using DTypeKV = __nv_fp8_e4m3;
24
+ using DTypeO = half;
25
+ using IdType = int32_t;
26
+ constexpr int HEAD_DIM_QK = 128;
27
+ constexpr int HEAD_DIM_VO = 128;
28
+ constexpr auto USE_LOGITS_SOFT_CAP = false;
29
+ constexpr auto POS_ENCODING_MODE = PosEncodingMode::kNone;
30
+ constexpr auto USE_SLIDING_WINDOW = false;
31
+
32
+ struct Params {
33
+ using DTypeQ = DTypeQ;
34
+ using DTypeKV = DTypeKV;
35
+ using DTypeO = DTypeO;
36
+ using IdType = IdType;
37
+
38
+ DTypeQ* q;
39
+ paged_kv_t<DTypeKV, IdType> paged_kv;
40
+ DTypeO* o;
41
+ float* lse;
42
+
43
+ float* maybe_alibi_slopes;
44
+ double logits_soft_cap;
45
+ double sm_scale;
46
+ double rope_rcp_scale;
47
+ double rope_rcp_theta;
48
+
49
+
50
+ uint32_t padded_batch_size;
51
+ uint32_t num_qo_heads;
52
+ IdType q_stride_n;
53
+ IdType q_stride_h;
54
+ int32_t window_left;
55
+ bool enable_pdl;
56
+
57
+ IdType* request_indices;
58
+ IdType* kv_tile_indices;
59
+ IdType* o_indptr;
60
+ IdType* kv_chunk_size_ptr;
61
+ bool* block_valid_mask;
62
+ bool partition_kv;
63
+
64
+ __host__ __device__ __forceinline__ int32_t get_qo_len(int32_t batch_idx) const { return 1; }
65
+
66
+ __host__ __device__ __forceinline__ int32_t get_kv_len(int32_t batch_idx) const {
67
+ return paged_kv.get_length(batch_idx);
68
+ }
69
+ };
70
+
71
+ #include<flashinfer/attention/variants.cuh>
csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_jit_pybind.cu ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023-2025 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include "batch_decode_config.inc"
17
+ #include "pytorch_extension_utils.h"
18
+
19
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
20
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
21
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
22
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
23
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
24
+ at::Tensor empty_q_data, at::Tensor empty_kv_data);
25
+
26
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
27
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
28
+ at::Tensor q, at::Tensor paged_k_cache,
29
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
30
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
31
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
32
+ int64_t kv_layout_code, int64_t window_left,
33
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS);
34
+
35
+ TORCH_LIBRARY_FRAGMENT(TORCH_EXTENSION_NAME, m) {
36
+ // Batched decode with paged KV-Cache plan
37
+ m.def("plan", BatchDecodeWithPagedKVCachePlan);
38
+ // Batched decode with paged KV-Cache run
39
+ m.def("run", BatchDecodeWithPagedKVCacheRun);
40
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_128_head_dim_vo_128_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_kernel.cu ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #include <flashinfer/attention/decode.cuh>
2
+ #include "batch_decode_config.inc"
3
+
4
+ using namespace flashinfer;
5
+
6
+ namespace flashinfer {
7
+
8
+ template cudaError_t
9
+ BatchDecodeWithPagedKVCacheDispatched<128, PosEncodingMode::kNone, DefaultAttention<false, false, false, false>, Params>(
10
+ Params params, half* tmp_v,
11
+ float* tmp_s, bool enable_pdl, cudaStream_t stream);
12
+
13
+ };
csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode.cu ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include <flashinfer/attention/scheduler.cuh>
17
+ #include <flashinfer/pos_enc.cuh>
18
+ #include <flashinfer/utils.cuh>
19
+ #include <optional>
20
+
21
+ #include "batch_decode_config.inc"
22
+ #include "pytorch_conversion_utils.h"
23
+ #include "pytorch_extension_utils.h"
24
+
25
+ namespace flashinfer {
26
+
27
+ template <uint32_t HEAD_DIM, PosEncodingMode POS_ENCODING_MODE, typename AttentionVariant,
28
+ typename Params>
29
+ cudaError_t BatchDecodeWithPagedKVCacheDispatched(Params params, typename Params::DTypeO* tmp_v,
30
+ float* tmp_s, bool enable_pdl,
31
+ cudaStream_t stream);
32
+
33
+ } // namespace flashinfer
34
+
35
+ using namespace flashinfer;
36
+
37
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
38
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
39
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
40
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
41
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
42
+ at::Tensor empty_q_data, at::Tensor empty_kv_data) {
43
+ size_t float_workspace_size_in_bytes =
44
+ float_workspace_buffer.size(0) * float_workspace_buffer.element_size();
45
+ size_t int_workspace_size_in_bytes =
46
+ int_workspace_buffer.size(0) * int_workspace_buffer.element_size();
47
+
48
+ DecodePlanInfo plan_info;
49
+
50
+ auto q_scalar_type = empty_q_data.scalar_type();
51
+ auto kv_scalar_type = empty_kv_data.scalar_type();
52
+
53
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
54
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
55
+ "cores template for different head dim");
56
+
57
+ const c10::cuda::OptionalCUDAGuard device_guard(float_workspace_buffer.device());
58
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
59
+ DISPATCH_context(
60
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
61
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
62
+ DISPATCH_GQA_GROUP_SIZE(num_qo_heads / num_kv_heads, GROUP_SIZE, {
63
+ auto work_estimation_func = BatchDecodeWithPagedKVCacheWorkEstimationDispatched<
64
+ GROUP_SIZE, HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>;
65
+ cudaError_t status = DecodePlan<HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>(
66
+ static_cast<void*>(float_workspace_buffer.data_ptr()), float_workspace_size_in_bytes,
67
+ static_cast<void*>(int_workspace_buffer.data_ptr()),
68
+ static_cast<void*>(page_locked_int_workspace_buffer.data_ptr()),
69
+ int_workspace_size_in_bytes, plan_info, static_cast<IdType*>(indptr.data_ptr()),
70
+ batch_size, num_qo_heads, page_size, enable_cuda_graph,
71
+ /*stream=*/stream, work_estimation_func);
72
+
73
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
74
+ cudaGetErrorString(status));
75
+ return true;
76
+ });
77
+ });
78
+
79
+ return vec_to_tensor(plan_info.ToVector());
80
+ }
81
+
82
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
83
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
84
+ at::Tensor q, at::Tensor paged_k_cache,
85
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
86
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
87
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
88
+ int64_t kv_layout_code, int64_t window_left,
89
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS) {
90
+ DecodePlanInfo plan_info;
91
+ plan_info.FromVector(tensor_to_vec(plan_info_vec));
92
+ QKVLayout kv_layout = static_cast<QKVLayout>(kv_layout_code);
93
+ auto device = q.device();
94
+ int64_t batch_size = q.size(0);
95
+ int64_t num_qo_heads = q.size(1);
96
+ int64_t num_kv_heads, page_size;
97
+
98
+ if (kv_layout == QKVLayout::kHND) {
99
+ num_kv_heads = paged_k_cache.size(1);
100
+ page_size = paged_k_cache.size(2);
101
+ } else {
102
+ page_size = paged_k_cache.size(1);
103
+ num_kv_heads = paged_k_cache.size(2);
104
+ }
105
+ uint32_t head_dim_qk = q.size(2);
106
+ uint32_t head_dim_vo = paged_v_cache.size(3);
107
+
108
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
109
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
110
+ "cores template for different head dim");
111
+
112
+ if (maybe_lse) {
113
+ const auto& lse = *maybe_lse;
114
+ TORCH_CHECK(lse.size(0) == batch_size, lse.size(0), q.size(0));
115
+ TORCH_CHECK(lse.size(1) == num_qo_heads, lse.size(1), q.size(1));
116
+ }
117
+
118
+ void* float_buffer = static_cast<void*>(float_workspace_buffer.data_ptr());
119
+ void* int_buffer = static_cast<void*>(int_workspace_buffer.data_ptr());
120
+
121
+ // get q_scalar_type and kv_scalar_type
122
+ auto q_scalar_type = q.scalar_type();
123
+ auto kv_scalar_type = paged_k_cache.scalar_type();
124
+
125
+ // get q_stride_n and q_stride_h
126
+ const auto q_stride_n = q.stride(0);
127
+ const auto q_stride_h = q.stride(1);
128
+
129
+ // get kv_cache_strides
130
+ const int64_t* kv_cache_strides = nullptr;
131
+ auto k_strides = paged_k_cache.strides();
132
+ auto v_strides = paged_v_cache.strides();
133
+ TORCH_CHECK(k_strides == v_strides, "k/v strides must be identical");
134
+ kv_cache_strides = k_strides.data();
135
+
136
+ const c10::cuda::OptionalCUDAGuard device_guard(device);
137
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
138
+
139
+ DISPATCH_context(
140
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
141
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
142
+ paged_kv_t<DTypeKV, IdType> paged_kv(
143
+ num_kv_heads, page_size, HEAD_DIM_QK, batch_size, kv_layout,
144
+ static_cast<DTypeKV*>(paged_k_cache.data_ptr()),
145
+ static_cast<DTypeKV*>(paged_v_cache.data_ptr()), kv_cache_strides,
146
+ static_cast<IdType*>(paged_kv_indices.data_ptr()),
147
+ static_cast<IdType*>(paged_kv_indptr.data_ptr()),
148
+ static_cast<IdType*>(paged_kv_last_page_len.data_ptr()));
149
+
150
+ Params params;
151
+ params.q = static_cast<DTypeQ*>(q.data_ptr());
152
+ params.paged_kv = paged_kv;
153
+ params.o = static_cast<DTypeO*>(o.data_ptr());
154
+ params.lse = maybe_lse ? static_cast<float*>(maybe_lse->data_ptr()) : nullptr;
155
+ params.padded_batch_size = 0;
156
+ params.num_qo_heads = num_qo_heads;
157
+ params.q_stride_n = q_stride_n;
158
+ params.q_stride_h = q_stride_h;
159
+ params.window_left = window_left;
160
+ params.request_indices = nullptr;
161
+ params.kv_tile_indices = nullptr;
162
+ params.o_indptr = nullptr;
163
+ params.kv_chunk_size_ptr = nullptr;
164
+ params.block_valid_mask = nullptr;
165
+ params.partition_kv = false;
166
+
167
+ ADDITIONAL_PARAMS_SETTER
168
+
169
+ DTypeO* tmp_v = nullptr;
170
+ float* tmp_s = nullptr;
171
+ params.request_indices =
172
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.request_indices_offset);
173
+ params.kv_tile_indices =
174
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_tile_indices_offset);
175
+ params.o_indptr = GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.o_indptr_offset);
176
+ params.kv_chunk_size_ptr =
177
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_chunk_size_ptr_offset);
178
+ if (plan_info.split_kv) {
179
+ tmp_v = GetPtrFromBaseOffset<DTypeO>(float_buffer, plan_info.v_offset);
180
+ tmp_s = GetPtrFromBaseOffset<float>(float_buffer, plan_info.s_offset);
181
+ if (plan_info.enable_cuda_graph) {
182
+ params.block_valid_mask =
183
+ GetPtrFromBaseOffset<bool>(int_buffer, plan_info.block_valid_mask_offset);
184
+ }
185
+ }
186
+ params.padded_batch_size = plan_info.padded_batch_size;
187
+
188
+ cudaError_t status =
189
+ flashinfer::BatchDecodeWithPagedKVCacheDispatched<HEAD_DIM_QK, POS_ENCODING_MODE,
190
+ AttentionVariant>(params, tmp_v,
191
+ tmp_s, enable_pdl,
192
+ /*stream=*/stream);
193
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
194
+ cudaGetErrorString(status));
195
+ return true;
196
+ });
197
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_config.inc ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #pragma once
2
+ #include <flashinfer/page.cuh>
3
+ #include <flashinfer/math.cuh>
4
+ #include <flashinfer/layout.cuh>
5
+ #include <flashinfer/pos_enc.cuh>
6
+ #include <flashinfer/attention/variant_helper.cuh>
7
+
8
+ #define ADDITIONAL_FUNC_PARAMS , std::optional<at::Tensor> maybe_alibi_slopes, double logits_soft_cap, double sm_scale, double rope_rcp_scale, double rope_rcp_theta
9
+ #define ADDITIONAL_PARAMS_SETTER params.maybe_alibi_slopes = maybe_alibi_slopes ? static_cast<float*>(maybe_alibi_slopes->data_ptr()): nullptr; \
10
+ params.logits_soft_cap = logits_soft_cap; \
11
+ params.sm_scale = sm_scale; \
12
+ params.rope_rcp_scale = rope_rcp_scale; \
13
+ params.rope_rcp_theta = rope_rcp_theta;
14
+
15
+ #define DISPATCH_context(DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE, USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, ...) { \
16
+ using AttentionVariant = DefaultAttention<false, true, true, false>; \
17
+ __VA_ARGS__(); \
18
+ }
19
+
20
+ using namespace flashinfer;
21
+
22
+ using DTypeQ = half;
23
+ using DTypeKV = __nv_fp8_e4m3;
24
+ using DTypeO = half;
25
+ using IdType = int32_t;
26
+ constexpr int HEAD_DIM_QK = 256;
27
+ constexpr int HEAD_DIM_VO = 256;
28
+ constexpr auto USE_LOGITS_SOFT_CAP = true;
29
+ constexpr auto POS_ENCODING_MODE = PosEncodingMode::kNone;
30
+ constexpr auto USE_SLIDING_WINDOW = true;
31
+
32
+ struct Params {
33
+ using DTypeQ = DTypeQ;
34
+ using DTypeKV = DTypeKV;
35
+ using DTypeO = DTypeO;
36
+ using IdType = IdType;
37
+
38
+ DTypeQ* q;
39
+ paged_kv_t<DTypeKV, IdType> paged_kv;
40
+ DTypeO* o;
41
+ float* lse;
42
+
43
+ float* maybe_alibi_slopes;
44
+ double logits_soft_cap;
45
+ double sm_scale;
46
+ double rope_rcp_scale;
47
+ double rope_rcp_theta;
48
+
49
+
50
+ uint32_t padded_batch_size;
51
+ uint32_t num_qo_heads;
52
+ IdType q_stride_n;
53
+ IdType q_stride_h;
54
+ int32_t window_left;
55
+ bool enable_pdl;
56
+
57
+ IdType* request_indices;
58
+ IdType* kv_tile_indices;
59
+ IdType* o_indptr;
60
+ IdType* kv_chunk_size_ptr;
61
+ bool* block_valid_mask;
62
+ bool partition_kv;
63
+
64
+ __host__ __device__ __forceinline__ int32_t get_qo_len(int32_t batch_idx) const { return 1; }
65
+
66
+ __host__ __device__ __forceinline__ int32_t get_kv_len(int32_t batch_idx) const {
67
+ return paged_kv.get_length(batch_idx);
68
+ }
69
+ };
70
+
71
+ #include<flashinfer/attention/variants.cuh>
csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_jit_pybind.cu ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023-2025 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include "batch_decode_config.inc"
17
+ #include "pytorch_extension_utils.h"
18
+
19
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
20
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
21
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
22
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
23
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
24
+ at::Tensor empty_q_data, at::Tensor empty_kv_data);
25
+
26
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
27
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
28
+ at::Tensor q, at::Tensor paged_k_cache,
29
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
30
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
31
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
32
+ int64_t kv_layout_code, int64_t window_left,
33
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS);
34
+
35
+ TORCH_LIBRARY_FRAGMENT(TORCH_EXTENSION_NAME, m) {
36
+ // Batched decode with paged KV-Cache plan
37
+ m.def("plan", BatchDecodeWithPagedKVCachePlan);
38
+ // Batched decode with paged KV-Cache run
39
+ m.def("run", BatchDecodeWithPagedKVCacheRun);
40
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_256_head_dim_vo_256_posenc_0_use_swa_True_use_logits_cap_True/batch_decode_kernel.cu ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #include <flashinfer/attention/decode.cuh>
2
+ #include "batch_decode_config.inc"
3
+
4
+ using namespace flashinfer;
5
+
6
+ namespace flashinfer {
7
+
8
+ template cudaError_t
9
+ BatchDecodeWithPagedKVCacheDispatched<256, PosEncodingMode::kNone, DefaultAttention<false, true, true, false>, Params>(
10
+ Params params, half* tmp_v,
11
+ float* tmp_s, bool enable_pdl, cudaStream_t stream);
12
+
13
+ };
csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode.cu ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include <flashinfer/attention/scheduler.cuh>
17
+ #include <flashinfer/pos_enc.cuh>
18
+ #include <flashinfer/utils.cuh>
19
+ #include <optional>
20
+
21
+ #include "batch_decode_config.inc"
22
+ #include "pytorch_conversion_utils.h"
23
+ #include "pytorch_extension_utils.h"
24
+
25
+ namespace flashinfer {
26
+
27
+ template <uint32_t HEAD_DIM, PosEncodingMode POS_ENCODING_MODE, typename AttentionVariant,
28
+ typename Params>
29
+ cudaError_t BatchDecodeWithPagedKVCacheDispatched(Params params, typename Params::DTypeO* tmp_v,
30
+ float* tmp_s, bool enable_pdl,
31
+ cudaStream_t stream);
32
+
33
+ } // namespace flashinfer
34
+
35
+ using namespace flashinfer;
36
+
37
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
38
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
39
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
40
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
41
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
42
+ at::Tensor empty_q_data, at::Tensor empty_kv_data) {
43
+ size_t float_workspace_size_in_bytes =
44
+ float_workspace_buffer.size(0) * float_workspace_buffer.element_size();
45
+ size_t int_workspace_size_in_bytes =
46
+ int_workspace_buffer.size(0) * int_workspace_buffer.element_size();
47
+
48
+ DecodePlanInfo plan_info;
49
+
50
+ auto q_scalar_type = empty_q_data.scalar_type();
51
+ auto kv_scalar_type = empty_kv_data.scalar_type();
52
+
53
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
54
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
55
+ "cores template for different head dim");
56
+
57
+ const c10::cuda::OptionalCUDAGuard device_guard(float_workspace_buffer.device());
58
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
59
+ DISPATCH_context(
60
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
61
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
62
+ DISPATCH_GQA_GROUP_SIZE(num_qo_heads / num_kv_heads, GROUP_SIZE, {
63
+ auto work_estimation_func = BatchDecodeWithPagedKVCacheWorkEstimationDispatched<
64
+ GROUP_SIZE, HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>;
65
+ cudaError_t status = DecodePlan<HEAD_DIM_QK, POS_ENCODING_MODE, AttentionVariant, Params>(
66
+ static_cast<void*>(float_workspace_buffer.data_ptr()), float_workspace_size_in_bytes,
67
+ static_cast<void*>(int_workspace_buffer.data_ptr()),
68
+ static_cast<void*>(page_locked_int_workspace_buffer.data_ptr()),
69
+ int_workspace_size_in_bytes, plan_info, static_cast<IdType*>(indptr.data_ptr()),
70
+ batch_size, num_qo_heads, page_size, enable_cuda_graph,
71
+ /*stream=*/stream, work_estimation_func);
72
+
73
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
74
+ cudaGetErrorString(status));
75
+ return true;
76
+ });
77
+ });
78
+
79
+ return vec_to_tensor(plan_info.ToVector());
80
+ }
81
+
82
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
83
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
84
+ at::Tensor q, at::Tensor paged_k_cache,
85
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
86
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
87
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
88
+ int64_t kv_layout_code, int64_t window_left,
89
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS) {
90
+ DecodePlanInfo plan_info;
91
+ plan_info.FromVector(tensor_to_vec(plan_info_vec));
92
+ QKVLayout kv_layout = static_cast<QKVLayout>(kv_layout_code);
93
+ auto device = q.device();
94
+ int64_t batch_size = q.size(0);
95
+ int64_t num_qo_heads = q.size(1);
96
+ int64_t num_kv_heads, page_size;
97
+
98
+ if (kv_layout == QKVLayout::kHND) {
99
+ num_kv_heads = paged_k_cache.size(1);
100
+ page_size = paged_k_cache.size(2);
101
+ } else {
102
+ page_size = paged_k_cache.size(1);
103
+ num_kv_heads = paged_k_cache.size(2);
104
+ }
105
+ uint32_t head_dim_qk = q.size(2);
106
+ uint32_t head_dim_vo = paged_v_cache.size(3);
107
+
108
+ TORCH_CHECK(head_dim_qk == head_dim_vo,
109
+ "CUDA cores template only supports equal head dim for QK and VO, please use tensor "
110
+ "cores template for different head dim");
111
+
112
+ if (maybe_lse) {
113
+ const auto& lse = *maybe_lse;
114
+ TORCH_CHECK(lse.size(0) == batch_size, lse.size(0), q.size(0));
115
+ TORCH_CHECK(lse.size(1) == num_qo_heads, lse.size(1), q.size(1));
116
+ }
117
+
118
+ void* float_buffer = static_cast<void*>(float_workspace_buffer.data_ptr());
119
+ void* int_buffer = static_cast<void*>(int_workspace_buffer.data_ptr());
120
+
121
+ // get q_scalar_type and kv_scalar_type
122
+ auto q_scalar_type = q.scalar_type();
123
+ auto kv_scalar_type = paged_k_cache.scalar_type();
124
+
125
+ // get q_stride_n and q_stride_h
126
+ const auto q_stride_n = q.stride(0);
127
+ const auto q_stride_h = q.stride(1);
128
+
129
+ // get kv_cache_strides
130
+ const int64_t* kv_cache_strides = nullptr;
131
+ auto k_strides = paged_k_cache.strides();
132
+ auto v_strides = paged_v_cache.strides();
133
+ TORCH_CHECK(k_strides == v_strides, "k/v strides must be identical");
134
+ kv_cache_strides = k_strides.data();
135
+
136
+ const c10::cuda::OptionalCUDAGuard device_guard(device);
137
+ const cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
138
+
139
+ DISPATCH_context(
140
+ DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE,
141
+ USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, [&] {
142
+ paged_kv_t<DTypeKV, IdType> paged_kv(
143
+ num_kv_heads, page_size, HEAD_DIM_QK, batch_size, kv_layout,
144
+ static_cast<DTypeKV*>(paged_k_cache.data_ptr()),
145
+ static_cast<DTypeKV*>(paged_v_cache.data_ptr()), kv_cache_strides,
146
+ static_cast<IdType*>(paged_kv_indices.data_ptr()),
147
+ static_cast<IdType*>(paged_kv_indptr.data_ptr()),
148
+ static_cast<IdType*>(paged_kv_last_page_len.data_ptr()));
149
+
150
+ Params params;
151
+ params.q = static_cast<DTypeQ*>(q.data_ptr());
152
+ params.paged_kv = paged_kv;
153
+ params.o = static_cast<DTypeO*>(o.data_ptr());
154
+ params.lse = maybe_lse ? static_cast<float*>(maybe_lse->data_ptr()) : nullptr;
155
+ params.padded_batch_size = 0;
156
+ params.num_qo_heads = num_qo_heads;
157
+ params.q_stride_n = q_stride_n;
158
+ params.q_stride_h = q_stride_h;
159
+ params.window_left = window_left;
160
+ params.request_indices = nullptr;
161
+ params.kv_tile_indices = nullptr;
162
+ params.o_indptr = nullptr;
163
+ params.kv_chunk_size_ptr = nullptr;
164
+ params.block_valid_mask = nullptr;
165
+ params.partition_kv = false;
166
+
167
+ ADDITIONAL_PARAMS_SETTER
168
+
169
+ DTypeO* tmp_v = nullptr;
170
+ float* tmp_s = nullptr;
171
+ params.request_indices =
172
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.request_indices_offset);
173
+ params.kv_tile_indices =
174
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_tile_indices_offset);
175
+ params.o_indptr = GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.o_indptr_offset);
176
+ params.kv_chunk_size_ptr =
177
+ GetPtrFromBaseOffset<IdType>(int_buffer, plan_info.kv_chunk_size_ptr_offset);
178
+ if (plan_info.split_kv) {
179
+ tmp_v = GetPtrFromBaseOffset<DTypeO>(float_buffer, plan_info.v_offset);
180
+ tmp_s = GetPtrFromBaseOffset<float>(float_buffer, plan_info.s_offset);
181
+ if (plan_info.enable_cuda_graph) {
182
+ params.block_valid_mask =
183
+ GetPtrFromBaseOffset<bool>(int_buffer, plan_info.block_valid_mask_offset);
184
+ }
185
+ }
186
+ params.padded_batch_size = plan_info.padded_batch_size;
187
+
188
+ cudaError_t status =
189
+ flashinfer::BatchDecodeWithPagedKVCacheDispatched<HEAD_DIM_QK, POS_ENCODING_MODE,
190
+ AttentionVariant>(params, tmp_v,
191
+ tmp_s, enable_pdl,
192
+ /*stream=*/stream);
193
+ TORCH_CHECK(status == cudaSuccess, "BatchDecodeWithPagedKVCache failed with error ",
194
+ cudaGetErrorString(status));
195
+ return true;
196
+ });
197
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_config.inc ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #pragma once
2
+ #include <flashinfer/page.cuh>
3
+ #include <flashinfer/math.cuh>
4
+ #include <flashinfer/layout.cuh>
5
+ #include <flashinfer/pos_enc.cuh>
6
+ #include <flashinfer/attention/variant_helper.cuh>
7
+
8
+ #define ADDITIONAL_FUNC_PARAMS , std::optional<at::Tensor> maybe_alibi_slopes, double logits_soft_cap, double sm_scale, double rope_rcp_scale, double rope_rcp_theta
9
+ #define ADDITIONAL_PARAMS_SETTER params.maybe_alibi_slopes = maybe_alibi_slopes ? static_cast<float*>(maybe_alibi_slopes->data_ptr()): nullptr; \
10
+ params.logits_soft_cap = logits_soft_cap; \
11
+ params.sm_scale = sm_scale; \
12
+ params.rope_rcp_scale = rope_rcp_scale; \
13
+ params.rope_rcp_theta = rope_rcp_theta;
14
+
15
+ #define DISPATCH_context(DTypeQ, DTypeKV, DTypeO, IdType, HEAD_DIM_QK, HEAD_DIM_VO, POS_ENCODING_MODE, USE_SLIDING_WINDOW, USE_LOGITS_SOFT_CAP, AttentionVariant, Params, ...) { \
16
+ using AttentionVariant = DefaultAttention<false, false, false, false>; \
17
+ __VA_ARGS__(); \
18
+ }
19
+
20
+ using namespace flashinfer;
21
+
22
+ using DTypeQ = half;
23
+ using DTypeKV = __nv_fp8_e4m3;
24
+ using DTypeO = half;
25
+ using IdType = int32_t;
26
+ constexpr int HEAD_DIM_QK = 64;
27
+ constexpr int HEAD_DIM_VO = 64;
28
+ constexpr auto USE_LOGITS_SOFT_CAP = false;
29
+ constexpr auto POS_ENCODING_MODE = PosEncodingMode::kNone;
30
+ constexpr auto USE_SLIDING_WINDOW = false;
31
+
32
+ struct Params {
33
+ using DTypeQ = DTypeQ;
34
+ using DTypeKV = DTypeKV;
35
+ using DTypeO = DTypeO;
36
+ using IdType = IdType;
37
+
38
+ DTypeQ* q;
39
+ paged_kv_t<DTypeKV, IdType> paged_kv;
40
+ DTypeO* o;
41
+ float* lse;
42
+
43
+ float* maybe_alibi_slopes;
44
+ double logits_soft_cap;
45
+ double sm_scale;
46
+ double rope_rcp_scale;
47
+ double rope_rcp_theta;
48
+
49
+
50
+ uint32_t padded_batch_size;
51
+ uint32_t num_qo_heads;
52
+ IdType q_stride_n;
53
+ IdType q_stride_h;
54
+ int32_t window_left;
55
+ bool enable_pdl;
56
+
57
+ IdType* request_indices;
58
+ IdType* kv_tile_indices;
59
+ IdType* o_indptr;
60
+ IdType* kv_chunk_size_ptr;
61
+ bool* block_valid_mask;
62
+ bool partition_kv;
63
+
64
+ __host__ __device__ __forceinline__ int32_t get_qo_len(int32_t batch_idx) const { return 1; }
65
+
66
+ __host__ __device__ __forceinline__ int32_t get_kv_len(int32_t batch_idx) const {
67
+ return paged_kv.get_length(batch_idx);
68
+ }
69
+ };
70
+
71
+ #include<flashinfer/attention/variants.cuh>
csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_jit_pybind.cu ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /*
2
+ * Copyright (c) 2023-2025 by FlashInfer team.
3
+ *
4
+ * Licensed under the Apache License, Version 2.0 (the "License");
5
+ * you may not use this file except in compliance with the License.
6
+ * You may obtain a copy of the License at
7
+ *
8
+ * http://www.apache.org/licenses/LICENSE-2.0
9
+ *
10
+ * Unless required by applicable law or agreed to in writing, software
11
+ * distributed under the License is distributed on an "AS IS" BASIS,
12
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ * See the License for the specific language governing permissions and
14
+ * limitations under the License.
15
+ */
16
+ #include "batch_decode_config.inc"
17
+ #include "pytorch_extension_utils.h"
18
+
19
+ at::Tensor BatchDecodeWithPagedKVCachePlan(
20
+ at::Tensor float_workspace_buffer, at::Tensor int_workspace_buffer,
21
+ at::Tensor page_locked_int_workspace_buffer, at::Tensor indptr, int64_t batch_size,
22
+ int64_t num_qo_heads, int64_t num_kv_heads, int64_t page_size, bool enable_cuda_graph,
23
+ int64_t window_left, double logits_soft_cap, int64_t head_dim_qk, int64_t head_dim_vo,
24
+ at::Tensor empty_q_data, at::Tensor empty_kv_data);
25
+
26
+ void BatchDecodeWithPagedKVCacheRun(at::Tensor float_workspace_buffer,
27
+ at::Tensor int_workspace_buffer, at::Tensor plan_info_vec,
28
+ at::Tensor q, at::Tensor paged_k_cache,
29
+ at::Tensor paged_v_cache, at::Tensor paged_kv_indptr,
30
+ at::Tensor paged_kv_indices, at::Tensor paged_kv_last_page_len,
31
+ at::Tensor o, std::optional<at::Tensor> maybe_lse,
32
+ int64_t kv_layout_code, int64_t window_left,
33
+ bool enable_pdl ADDITIONAL_FUNC_PARAMS);
34
+
35
+ TORCH_LIBRARY_FRAGMENT(TORCH_EXTENSION_NAME, m) {
36
+ // Batched decode with paged KV-Cache plan
37
+ m.def("plan", BatchDecodeWithPagedKVCachePlan);
38
+ // Batched decode with paged KV-Cache run
39
+ m.def("run", BatchDecodeWithPagedKVCacheRun);
40
+ }
csrc/generated/batch_decode_with_kv_cache_dtype_q_f16_dtype_kv_e4m3_dtype_o_f16_dtype_idx_i32_head_dim_qk_64_head_dim_vo_64_posenc_0_use_swa_False_use_logits_cap_False/batch_decode_kernel.cu ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #include <flashinfer/attention/decode.cuh>
2
+ #include "batch_decode_config.inc"
3
+
4
+ using namespace flashinfer;
5
+
6
+ namespace flashinfer {
7
+
8
+ template cudaError_t
9
+ BatchDecodeWithPagedKVCacheDispatched<64, PosEncodingMode::kNone, DefaultAttention<false, false, false, false>, Params>(
10
+ Params params, half* tmp_v,
11
+ float* tmp_s, bool enable_pdl, cudaStream_t stream);
12
+
13
+ };