Instructions to use replicate/flashinfer-draft with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Kernels
How to use replicate/flashinfer-draft with Kernels:
# !pip install kernels from kernels import get_kernel kernel = get_kernel("replicate/flashinfer-draft") - Notebooks
- Google Colab
- Kaggle
drbh commited on
Commit ·
57c3a10
1
Parent(s): 6b27dde
feat: generate and vendor flashinfer kernels
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +1 -0
- .gitignore +3 -0
- .make_markers/patch_applied +0 -0
- .make_markers/submodule_initialized +0 -0
- README.md +10 -0
- build.toml +224 -0
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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
- 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 |
+
};
|