diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..5a05b1181fb1 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6ec6c068288c949c5900574cefee6957b29bac6535c14eaf5b6e6ab8b66fd974 +size 30306 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..21798a89556b --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7ab87005d5352fae201c6f70b410b200cbca30144fd492f0456432b1344abdc2 +size 29428 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..ef609577aa0d --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:562b4fdbb68a01d3db56687a43c6e2a48f7955254adcd26e0303b67364603dcb +size 27989 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..a6c26190f885 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0e317857e1c0dc09b4802a9cf0f0b742ade56768967b3174c04daaf024605e93 +size 33158 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..e5402932bd9d --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0becf6dd9d585b1e5f0a1bb97f7738e9fab95ea36f47690472339e5c941a5dbc +size 32633 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..72b9c5822daf --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:39103a67c0cf5683578cabc67c417ee8a842c474eb17a8da329d389e344e7ec3 +size 30977 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..5e5e09eaf1a6 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f3434840fedb9cabb4ee1c430dd213b96f8ab9e0ed8492544f64652baa9231d7 +size 36426 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..1a342b78ff60 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8d58835636c23d95678539364e24d5c6abb35d6e0f8675a32e35b78015346dbc +size 36458 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..37cf00701e67 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3f5c854ff6aee8c19584a764e87be9a39d1dceb90ca61a25093136a55cc9a0da +size 33953 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..7d292d9bbb50 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c944abb856061b8c7891d888ed612222e5e921161bf61ac2d00725f4826c457c +size 30371 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..1da12177452c --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a06a3fb1982d48890148f5877b7b13733b4d89017e3377f92ab89ddbc20d825f +size 29845 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..02f34895a72a --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8ff75170ab8fcf17aba56da5782f5af2bcf50ebdc19d2ff510165ce07d11378c +size 27949 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..0081155c5521 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:984ba62fcbfca61855344da887895a4414cd5ef5d212d99cd92c2931d7d52e08 +size 33171 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..17bf11947b95 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9e1b58fd80a68d58c2ed8b48dc9999dcc895b38296d3bff1872ab10c7eb25746 +size 33036 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..ca0befc2bff3 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4c8978d24ade1aeb361cbaab8e3675ae90c78ba245524e27b53b8dc51d2300ba +size 31177 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..b55935bac091 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:68288997abac9ca93827ce63ec52ef58905d217556d187647929ee55b471cdfa +size 36969 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..dedcb72637b4 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:00acef74a5051d112ceca982680bd951c6eaecee26105f800a77a4081525a8b1 +size 37040 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst new file mode 100644 index 000000000000..0aa34a7a8236 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:803976279b4b67641146db8765b1aaf31fc2b2e364492d2dbeb569cfd010d968 +size 35031 diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/kernelMetaInfo.h b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/kernelMetaInfo.h index 50e642011517..0cfa02672458 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/kernelMetaInfo.h +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/cubin/kernelMetaInfo.h @@ -909,6 +909,15 @@ extern unsigned char const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk192HV128Separ extern unsigned char const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk192HV128SeparateQkvDenseVarSeqQ128Kv128StaticContext_cubin[]; extern unsigned char const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk192HV128SeparateQkvDenseVarSeqSkipsSoftmaxQ128Kv128PersistentContext_cubin[]; extern unsigned char const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk192HV128SeparateQkvDenseVarSeqSkipsSoftmaxQ128Kv128StaticContext_cubin[]; +extern unsigned char const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin[]; extern unsigned char const FmhaSm103aKernel_QkvE4m3OBfloat16H128PackedQkvCausalVarSeqQ128Kv128PersistentContext_cubin[]; extern unsigned char const FmhaSm103aKernel_QkvE4m3OBfloat16H128PackedQkvCausalVarSeqQ128Kv128StaticContext_cubin[]; extern unsigned char const FmhaSm103aKernel_QkvE4m3OBfloat16H128PackedQkvCausalVarSeqSkipsSoftmaxQ128Kv128PersistentContext_cubin[]; @@ -1973,6 +1982,15 @@ extern unsigned int const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk192HV128Separa extern unsigned int const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk192HV128SeparateQkvDenseVarSeqQ128Kv128StaticContext_cubin_len; extern unsigned int const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk192HV128SeparateQkvDenseVarSeqSkipsSoftmaxQ128Kv128PersistentContext_cubin_len; extern unsigned int const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk192HV128SeparateQkvDenseVarSeqSkipsSoftmaxQ128Kv128StaticContext_cubin_len; +extern unsigned int const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len; extern unsigned int const FmhaSm103aKernel_QkvE4m3OBfloat16H128PackedQkvCausalVarSeqQ128Kv128PersistentContext_cubin_len; extern unsigned int const FmhaSm103aKernel_QkvE4m3OBfloat16H128PackedQkvCausalVarSeqQ128Kv128StaticContext_cubin_len; extern unsigned int const FmhaSm103aKernel_QkvE4m3OBfloat16H128PackedQkvCausalVarSeqSkipsSoftmaxQ128Kv128PersistentContext_cubin_len; @@ -3084,6 +3102,15 @@ extern unsigned char const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512Paged extern unsigned char const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvDenseStaticTokenSparseP1MultiCtasKvGmemSepVarSeqQ64Kv128StaticKeepsAbForGen_cubin[]; extern unsigned char const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvDenseStaticTokenSparseP1VarSeqQ64Kv128PersistentKeepsAbForGen_cubin[]; extern unsigned char const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvDenseStaticTokenSparseP1VarSeqQ64Kv128StaticKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin[]; +extern unsigned char const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin[]; extern unsigned char const FmhaSm100fKernel_QkvE4m3OBfloat16H128PackedQkvCausalVarSeqQ128Kv128PersistentContext_cubin[]; extern unsigned char const FmhaSm100fKernel_QkvE4m3OBfloat16H128PackedQkvCausalVarSeqQ128Kv128StaticContext_cubin[]; extern unsigned char const FmhaSm100fKernel_QkvE4m3OBfloat16H128PackedQkvCausalVarSeqSkipsSoftmaxQ128Kv128PersistentContext_cubin[]; @@ -4996,6 +5023,15 @@ extern unsigned int const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedK extern unsigned int const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvDenseStaticTokenSparseP1MultiCtasKvGmemSepVarSeqQ64Kv128StaticKeepsAbForGen_cubin_len; extern unsigned int const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvDenseStaticTokenSparseP1VarSeqQ64Kv128PersistentKeepsAbForGen_cubin_len; extern unsigned int const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvDenseStaticTokenSparseP1VarSeqQ64Kv128StaticKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin_len; +extern unsigned int const FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len; extern unsigned int const FmhaSm100fKernel_QkvE4m3OBfloat16H128PackedQkvCausalVarSeqQ128Kv128PersistentContext_cubin_len; extern unsigned int const FmhaSm100fKernel_QkvE4m3OBfloat16H128PackedQkvCausalVarSeqQ128Kv128StaticContext_cubin_len; extern unsigned int const FmhaSm100fKernel_QkvE4m3OBfloat16H128PackedQkvCausalVarSeqSkipsSoftmaxQ128Kv128PersistentContext_cubin_len; @@ -7883,6 +7919,15 @@ static const TllmGenFmhaKernelMetaInfo sTllmGenFmhaKernelMetaInfos[] = { { DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, 128, 128, 256, 128, 80, 80, 80, kSM_103, FmhaSm103aKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqQ128Kv128StaticContext_cubin, FmhaSm103aKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqQ128Kv128StaticContext_cubin_len, "fmhaSm103aKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqQ128Kv128StaticContext", 165136, 512, 2, 32, 2, 0, 0, 0, 0, 0, 0, 0, false, false, false, false, 0, false, false, false, false, false, false, "0639e63d7d9490f56cb006951bbd31a41fba256fd9f32af3fb14afe444497ac5"}, { DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, 128, 128, 256, 128, 80, 80, 80, kSM_103, FmhaSm103aKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqSkipsSoftmaxQ128Kv128PersistentContext_cubin, FmhaSm103aKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqSkipsSoftmaxQ128Kv128PersistentContext_cubin_len, "fmhaSm103aKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqSkipsSoftmaxQ128Kv128PersistentContext", 165264, 512, 2, 32, 2, 0, 1, 0, 0, 0, 0, 0, false, false, false, false, 0, true, false, false, false, false, false, "3cce3ec12b2f1d284380033b686b777116c856e3499050e5aa59855661d0335e"}, { DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, 128, 128, 256, 128, 80, 80, 80, kSM_103, FmhaSm103aKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqSkipsSoftmaxQ128Kv128StaticContext_cubin, FmhaSm103aKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqSkipsSoftmaxQ128Kv128StaticContext_cubin_len, "fmhaSm103aKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqSkipsSoftmaxQ128Kv128StaticContext", 165152, 512, 2, 32, 2, 0, 0, 0, 0, 0, 0, 0, false, false, false, false, 0, true, false, false, false, false, false, "89ce00fe1b2929fc8176c9d9bdfcdfefd969b9e0705b455ebed08dbe5cd165db"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 128, 576, 512, kSM_103, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len, "fmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen", 206112, 384, 2, 32, 1, 3, 0, 2, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "90a136349125fb199fabbaa783eacadb8daee31a3a6fbf0903cf14dfe0f5587d"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 128, 576, 512, kSM_103, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin_len, "fmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen", 206208, 384, 2, 32, 1, 3, 1, 0, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "3a563afb3f0026ab6875f3377932c13021755e15843916ba181c37485ad5e970"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 128, 576, 512, kSM_103, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len, "fmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen", 206096, 384, 2, 32, 1, 3, 0, 0, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "08703af6f770f746527495bfd659b91352613c080c0a6e8d8fba1dae621f3c3b"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 256, 576, 512, kSM_103, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len, "fmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen", 206112, 384, 2, 32, 1, 3, 0, 2, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "941db89494246620bc4e69a5d361b74d52641f0bb7569b5dc7f5576d76a57844"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 256, 576, 512, kSM_103, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin_len, "fmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen", 206208, 384, 2, 32, 1, 3, 1, 0, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "761489905145b668be3ce23b18fe24d3d81694e12f813f657d05ba56849f6751"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 256, 576, 512, kSM_103, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len, "fmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen", 206096, 384, 2, 32, 1, 3, 0, 0, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "8a8b34a686fd0c7ae67c0c33cdc12f9597bb264c597e66045a3339ba55ee7fd3"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 512, 576, 512, kSM_103, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len, "fmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen", 206112, 384, 2, 32, 1, 3, 0, 2, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "fe58788fab2346578d5e4b640dde86579be66bb24227ee8449c91324c75bd0b6"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 512, 576, 512, kSM_103, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin_len, "fmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen", 206208, 384, 2, 32, 1, 3, 1, 0, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "b7525dc38f5ebf238ffd9bf6ad8b373cb2785a9e661c931e206cf8bbbbcc541e"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 512, 576, 512, kSM_103, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin, FmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len, "fmhaSm103aKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen", 206096, 384, 2, 32, 1, 3, 0, 0, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "157aa91be677e53e43c03dca108eeabf9d7acb1fe4dca22bd9daadf4c5290254"}, #endif // EXCLUDE_SM_103 #ifndef EXCLUDE_SM_107 #endif // EXCLUDE_SM_107 @@ -9799,6 +9844,15 @@ static const TllmGenFmhaKernelMetaInfo sTllmGenFmhaKernelMetaInfos[] = { { DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, 128, 128, 256, 128, 80, 80, 80, kSM_100f, FmhaSm100fKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqQ128Kv128StaticContext_cubin, FmhaSm100fKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqQ128Kv128StaticContext_cubin_len, "fmhaSm100fKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqQ128Kv128StaticContext", 165136, 512, 2, 32, 2, 0, 0, 0, 0, 0, 0, 0, false, false, false, false, 0, false, false, false, false, false, false, "0e0ec6fc08482336c8079eaf23fbebe15d2e9b4e32d0f3738ad1913cee4b8894"}, { DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, 128, 128, 256, 128, 80, 80, 80, kSM_100f, FmhaSm100fKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqSkipsSoftmaxQ128Kv128PersistentContext_cubin, FmhaSm100fKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqSkipsSoftmaxQ128Kv128PersistentContext_cubin_len, "fmhaSm100fKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqSkipsSoftmaxQ128Kv128PersistentContext", 165264, 512, 2, 32, 2, 0, 1, 0, 0, 0, 0, 0, false, false, false, false, 0, true, false, false, false, false, false, "2b12fe51d2a0e03131b5d06cda96dea3d64a04ebc256d142b4103be0cfb00ec2"}, { DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, DATA_TYPE_FP16, 128, 128, 256, 128, 80, 80, 80, kSM_100f, FmhaSm100fKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqSkipsSoftmaxQ128Kv128StaticContext_cubin, FmhaSm100fKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqSkipsSoftmaxQ128Kv128StaticContext_cubin_len, "fmhaSm100fKernel_QkvFp16OFp16H80PagedKvSlidingOrChunkedCausalP32VarSeqSkipsSoftmaxQ128Kv128StaticContext", 165152, 512, 2, 32, 2, 0, 0, 0, 0, 0, 0, 0, false, false, false, false, 0, true, false, false, false, false, false, "dea1474d1fe7c4c8f2b8df2b1c686cb5f066dd4e8934aef6434139f247bb8fc4"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 128, 576, 512, kSM_100f, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len, "fmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen", 206112, 384, 2, 32, 1, 3, 0, 2, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "57cdaea5483cc9be1c3b99434cfff94652f63d44c6a8fd81e3a0d25b2096b0f7"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 128, 576, 512, kSM_100f, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin_len, "fmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen", 206208, 384, 2, 32, 1, 3, 1, 0, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "9f796f2c507f4e4a79b7f7e804b3ed43e8f297ab3e5f2b8eb0e6538e2d0a427c"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 128, 576, 512, kSM_100f, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len, "fmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen", 206096, 384, 2, 32, 1, 3, 0, 0, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "1c746daed8ec8d66f65913c2a4f4b40b4b0394f98241d087d1450d24f28d57c1"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 256, 576, 512, kSM_100f, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len, "fmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen", 206112, 384, 2, 32, 1, 3, 0, 2, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "a997c5796d1bce52b8f1063201405695148d90d0a22104b26809704b3ac18ca7"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 256, 576, 512, kSM_100f, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin_len, "fmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen", 206208, 384, 2, 32, 1, 3, 1, 0, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "aa2808ae7025857af459f9fd40e3713d2c105e32a2559fc916be048ac975414d"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 256, 576, 512, kSM_100f, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len, "fmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta256PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen", 206096, 384, 2, 32, 1, 3, 0, 0, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "f60f1d7ccb903a0e7c151ee99517d2ea31b8fb8144f0d695184daa940fc3604b"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 512, 576, 512, kSM_100f, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len, "fmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32MultiCtasKvGmemSepVarSeqQ64Kv128StaticGroupedKeepsAbForGen", 206112, 384, 2, 32, 1, 3, 0, 2, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "c3a1281fa43a7973f48cad4d9454f7849b34e93adb319e9fe72ebf8a68732ad4"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 512, 576, 512, kSM_100f, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen_cubin_len, "fmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128PersistentGroupedKeepsAbForGen", 206208, 384, 2, 32, 1, 3, 1, 0, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "9966a1261165f05bc1367c6cc828224f3e8e63adc44c346f869b3ede693b26e8"}, +{ DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, 64, 128, 64, 128, 512, 576, 512, kSM_100f, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin, FmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen_cubin_len, "fmhaSm100fKernel_QkvBfloat16OBfloat16HQk576HV512PagedKvCausalP32VarSeqQ64Kv128StaticGroupedKeepsAbForGen", 206096, 384, 2, 32, 1, 3, 0, 0, 0, 0, 0, 0, true, true, false, false, 0, false, false, false, false, false, false, "0c3cdbfd1158b558bf4a3429e215fe90a7c929ac3e8053b4862af057a8d5881f"}, #endif // EXCLUDE_SM_100F }; // clang-format on diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaKernels.h b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaKernels.h index 8bcc29e0dbdb..729b0738931c 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaKernels.h +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaKernels.h @@ -95,6 +95,21 @@ constexpr bool isSMCompatible(int gpuSM, int kernelSM) return gpuSM == kernelSM; } +#if defined(TLLM_FMHA_TEST_HOOKS) +// Test-only result of probeKernelSelectionForTesting: the kernel the autotuner would launch. +struct TllmGenFmhaSelectedKernel +{ + // Resolved cubin function name (empty when mFound is false). + std::string mFuncName; + bool mFound = false; + // True when the autotuner chose the NVRTC path instead of a precompiled cubin. + bool mUsedNvrtc = false; + // Grouping flags from the matched kernelMeta. + bool mGroupsHeadsQ = false; + bool mGroupsTokensHeadsQ = false; +}; +#endif // TLLM_FMHA_TEST_HOOKS + class TllmGenFmhaKernel { @@ -291,6 +306,58 @@ class TllmGenFmhaKernel return std::make_pair(true, info); } +#if defined(TLLM_FMHA_TEST_HOOKS) + // Test-only: report which kernelMeta the autotuner + hash lookup resolve to for these params, + // without launching the kernel. + TllmGenFmhaSelectedKernel probeKernelSelectionForTesting(RunnerParams const& params) const + { + TllmGenFmhaSelectedKernel result; + if (params.mHeadDimQk % 8 != 0 || params.mHeadDimV % 8 != 0) + { + return result; + } + if (params.mMaxSeqLenQ == 0 || params.mBatchSize == 0 + || (!isContextKernel(params.mKernelType) && params.mMaxSeqLenKv == 0)) + { + return result; + } + int32_t ctaDim = 512; + FmhaOptions options; + FmhaOptionsFromArgs optionsFromArgs; + parseOptionsFromRunnerParams(params, options); + options.mCudaArch = intToCudaArch(mSM); + + FmhaAutoTuner autoTuner(options, optionsFromArgs, params.mMultiProcessorCount); + std::tie(options, optionsFromArgs, ctaDim) = autoTuner.selectKernel(); + + checkFmhaOptions(options, optionsFromArgs); + updateFmhaOptions(options, optionsFromArgs); + + computeNumCtas(options, params.mMultiProcessorCount); + + if (shouldUseNvrtc(options)) + { + result.mUsedNvrtc = true; + return result; + } + + algoFilterForCubinPath(options); + auto [hashId, info] = hashFromFmhaOptions(options); + + auto const findIter = mFunctions.find(hashId); + if (findIter == mFunctions.end()) + { + return result; + } + auto const& kernelMeta = mKernelMeta[findIter->second.mMetaInfoIndex]; + result.mFound = true; + result.mFuncName = kernelMeta.mFuncName != nullptr ? std::string(kernelMeta.mFuncName) : std::string{}; + result.mGroupsHeadsQ = kernelMeta.mGroupsHeadsQ; + result.mGroupsTokensHeadsQ = kernelMeta.mGroupsTokensHeadsQ; + return result; + } +#endif // TLLM_FMHA_TEST_HOOKS + void algoFilterForCubinPath(FmhaOptions& options) const { if (!isContextKernel(options.mFmhaKernelType) && options.mMaskType == TrtllmGenAttentionMaskType::Dense @@ -999,6 +1066,9 @@ class TllmGenFmhaKernel options.mEnablesAutoTuner = true; options.mIsMlaGen = isMlaGenKernel(params); + // Let the autotuner pick the grouped-token Q64 MLA generation kernel when its + // capability predicate matches (causal spec decode, supported head ratios/dtypes). + options.mSelectsGroupedMla = options.mIsMlaGen; options.mDtypeQ = dataTypeToDtype(mDtypeQ); options.mDtypeKv = dataTypeToDtype(mDtypeK); options.mDtypeK = dataTypeToDtype(mDtypeK); diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaRunner.h b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaRunner.h index bb84c45e6ee2..4e300ad76431 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaRunner.h +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaRunner.h @@ -1,5 +1,5 @@ /* - * Copyright (c) 2020-2023, NVIDIA CORPORATION. All rights reserved. + * Copyright (c) 2020-2026, NVIDIA CORPORATION. All rights reserved. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -17,6 +17,7 @@ #pragma once #include +#include #include "fmhaKernels.h" #include "fmhaRunnerParams.h" @@ -50,6 +51,14 @@ class TllmGenFmhaRunner // Run the fmha kernel. void run(TllmGenFmhaRunnerParams const&); +#if defined(TLLM_FMHA_TEST_HOOKS) + // Test-only: probe which cubin the autotuner would select for these params, without launching. + TllmGenFmhaSelectedKernel probeKernelSelectionForTesting(TllmGenFmhaRunnerParams const& runnerParams) const + { + return mKernel->probeKernelSelectionForTesting(runnerParams); + } +#endif // TLLM_FMHA_TEST_HOOKS + private: // The input/output datatype. Data_type mDtypeQ, mDtypeK, mDtypeV, mDtypeOut; diff --git a/cpp/tests/unit_tests/kernels/CMakeLists.txt b/cpp/tests/unit_tests/kernels/CMakeLists.txt index 6b0e5a118211..71920950561f 100644 --- a/cpp/tests/unit_tests/kernels/CMakeLists.txt +++ b/cpp/tests/unit_tests/kernels/CMakeLists.txt @@ -113,3 +113,6 @@ endif() add_gtest(eaglePackDataTest eaglePackDataTest.cpp) add_gtest(sparseKvCacheTest sparseKvCacheTest.cu) add_gtest(prepareCustomMaskTest prepareCustomMaskTest.cpp) +add_gtest(kimiMlaGroupedSelectionTest kimiMlaGroupedSelectionTest.cpp) +# Enables the test-only kernel-selection probe in fmhaRunner.h / fmhaKernels.h. +target_compile_definitions(kimiMlaGroupedSelectionTest PRIVATE TLLM_FMHA_TEST_HOOKS=1) diff --git a/cpp/tests/unit_tests/kernels/kimiMlaGroupedSelectionTest.cpp b/cpp/tests/unit_tests/kernels/kimiMlaGroupedSelectionTest.cpp new file mode 100644 index 000000000000..93a7668ba7e7 --- /dev/null +++ b/cpp/tests/unit_tests/kernels/kimiMlaGroupedSelectionTest.cpp @@ -0,0 +1,1216 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT + * WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the + * License for the specific language governing permissions and limitations + * under the License. + */ + +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include + +#include "tensorrt_llm/common/cudaUtils.h" +#include "tensorrt_llm/kernels/multiHeadAttentionCommon.h" +#include "tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaRunner.h" +#include "tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaRunnerParams.h" + +namespace +{ + +using namespace tensorrt_llm::kernels; + +// Reinterpret a BF16 bit pattern (uint16_t) as the corresponding FP32 value +// by zero-extending into the high half of a 32-bit float. Lossless because +// BF16's exponent + mantissa share FP32's high 16 bits. +inline float bf16BitsToFloat(uint16_t bits) +{ + uint32_t f = static_cast(bits) << 16; + float result; + std::memcpy(&result, &f, sizeof(float)); + return result; +} + +// BF16 1.0 bit pattern (sign=0, exp=127=0x7F, mantissa=0): 0x3F80. +constexpr uint16_t kBf16One = 0x3F80; +// BF16 2.0 bit pattern (sign=0, exp=128=0x80, mantissa=0): 0x4000. +constexpr uint16_t kBf16Two = 0x4000; + +// BF16 NaN-like sentinel used to pre-fill the output buffer. cudaMemset(..., 0xFF) +// puts 0xFFFF in every BF16 slot (sign=1, exp=0xFF, mantissa nonzero -> NaN). +constexpr uint16_t kBf16NaNSentinel = 0xFFFF; + +// Multi-CTA-KV semaphore/counter buffer: sized generously so the test does not +// depend on the exact per-CTA semaphore layout of the current kernels. +constexpr size_t kCounterBytes = size_t{1} << 20; + +// Convert FP32 to BF16 with simple truncation (top 16 bits of the FP32 word). +// Matches what TensorRT-LLM's BF16 storage convention is for these tests. +inline uint16_t floatToBf16(float v) +{ + uint32_t bits; + std::memcpy(&bits, &v, sizeof(float)); + return static_cast(bits >> 16); +} + +// FP32 CPU reference for MLA generation with the Kimi shape (B=1, q_len=4, 16 heads, +// HQk=576, HV=512, paged KV). Causal spec-decode mask: query token qi sees kv indices +// [0, seqLenKv - q_len + 1 + qi). +// +// MLA softmax scale is 1/sqrt(QK_NOPE_HEAD_DIM + QK_ROPE_HEAD_DIM) = 1/sqrt(192), not +// 1/sqrt(headDimQk=576). The caller must plumb the same scale to the kernel: pass +// scaleSoftmaxLog2 = (1/sqrt(192)) * log2(e) and set params.mScaleQ = sqrt(192/576) so the +// host-side softmaxScale = (1 / (sqrt(headDimQk) * mScaleQ)) * log2(e) resolves identically. +void mlaReferenceCpu(int seqLenQ, int seqLenKv, int numHeadsQ, int headDimQk, int headDimV, int numTokensPerPage, + std::vector const& Q, std::vector const& KV, float softmaxScale, std::vector& Out) +{ + // Q shape: [seqLenQ * numHeadsQ * headDimQk] + // KV shape (paged, single page-table batch): [seqLenKv * headDimQk] when + // page table is identity. We assume the caller has built KV with pages + // in sequential order so logical KV index kv just indexes KV[kv*headDimQk + d]. + // Out shape: [seqLenQ * numHeadsQ * headDimV] + auto qAt = [&](int q, int h, int d) -> float + { return bf16BitsToFloat(Q[static_cast(q) * numHeadsQ * headDimQk + h * headDimQk + d]); }; + auto kvAt = [&](int kv, int d) -> float + { return bf16BitsToFloat(KV[static_cast(kv) * headDimQk + d]); }; + (void) numTokensPerPage; // page-table identity is captured by the caller's hPageIdx[i]=i. + + std::vector scores(seqLenKv); + std::vector weights(seqLenKv); + for (int q = 0; q < seqLenQ; ++q) + { + // Spec-decode causal upper bound for this query position. + int const validEnd = seqLenKv - seqLenQ + 1 + q; + for (int h = 0; h < numHeadsQ; ++h) + { + // Scores Q[q,h] dot K[kv]; K spans the full headDimQk (kv_lora_rank + qk_rope). + for (int kv = 0; kv < seqLenKv; ++kv) + { + float s = 0.f; + for (int d = 0; d < headDimQk; ++d) + { + s += qAt(q, h, d) * kvAt(kv, d); + } + scores[kv] = s * softmaxScale; + } + // Apply spec-decode causal mask: positions >= validEnd are -inf. + for (int kv = validEnd; kv < seqLenKv; ++kv) + { + scores[kv] = -std::numeric_limits::infinity(); + } + // Numerically stable softmax. + float maxScore = -std::numeric_limits::infinity(); + for (int kv = 0; kv < seqLenKv; ++kv) + { + maxScore = std::max(maxScore, scores[kv]); + } + float sumExp = 0.f; + for (int kv = 0; kv < seqLenKv; ++kv) + { + weights[kv] = (scores[kv] == -std::numeric_limits::infinity()) ? 0.f + : std::exp(scores[kv] - maxScore); + sumExp += weights[kv]; + } + float const invSum = (sumExp > 0.f) ? (1.f / sumExp) : 0.f; + for (int kv = 0; kv < seqLenKv; ++kv) + { + weights[kv] *= invSum; + } + // Output Out[q,h,d_v] = sum_kv weights[kv] * V[kv, d_v], V is first kv_lora_rank + // of headDimQk. + for (int d_v = 0; d_v < headDimV; ++d_v) + { + float o = 0.f; + for (int kv = 0; kv < seqLenKv; ++kv) + { + o += weights[kv] * kvAt(kv, d_v); + } + Out[static_cast(q) * numHeadsQ * headDimV + h * headDimV + d_v] = floatToBf16(o); + } + } + } +} + +// Selection + correctness tests for the SM103 grouped-token MLA generation cubin +// (tileSizeQ=64 groups tokensQ and headsQ into one CTA) for the Kimi K2.5/K2.6 EAGLE-3 +// decode shape. Selection is driven by the trtllm-gen static-library autotuner and the +// grouped kernelMetaInfo.h row; the TLLM_FMHA_TEST_HOOKS probe asserts the resolved cubin. + +class KimiMlaGroupedSelectionTest : public ::testing::Test +{ +protected: + void SetUp() override + { + int sm = tensorrt_llm::common::getSMVersion(); + if (sm != kSM_103) + { + GTEST_SKIP() << "Kimi grouped MLA cubin is registered only for SM103 (B300). " + << "Current SM=" << sm << ". Skipping selection test."; + } + // Establish a CUDA context so cuModuleLoadData inside the runner ctor + // can resolve. cudaFree(0) is the canonical way to lazy-init the + // primary context for the current device without allocating. Check + // both return values so environment failures surface clearly. + ASSERT_EQ(cudaSuccess, cudaSetDevice(0)) << "cudaSetDevice(0) failed"; + ASSERT_EQ(cudaSuccess, cudaFree(nullptr)) << "cudaFree(nullptr) failed (CUDA context init)"; + } + + // Build a RunnerParams that mirrors the Kimi K2.5/K2.6 EAGLE-3 MLA decode + // call shape (B=1, q_len=4, 16 heads, HQk=576, HV=512, paged KV P32, BF16). + // mIsMlaGen is auto-derived from mHeadDimQk==576 && mHeadDimV==512 inside + // parseOptionsFromRunnerParams (see fmhaKernels.h:1209). Same for + // mIsCausalSpecDecodingGen from mMaxSeqLenQ>1 && !mIsSpecDecTree. + static void buildKimiParams(TllmGenFmhaRunnerParams& p, int maxSeqLenQ = 4, int headDimQk = 576, int headDimV = 512, + int numHeadsQ = 16) + { + std::memset(&p, 0, sizeof(p)); + p.mQkvLayout = QkvLayout::PagedKv; + // Callers pass Dense for generation; the autotuner rewrites the mask for + // causal spec-decode generation kernels. + p.mMaskType = TrtllmGenAttentionMaskType::Dense; + p.mIsSpecDecTree = false; + p.mKernelType = FmhaKernelType::Generation; + p.mTileScheduler = TileScheduler::Static; + p.mMultiCtasKvMode = true; + p.mHeadDimQk = headDimQk; + p.mHeadDimV = headDimV; + p.mHeadDimQkNope = 512; + p.mNumHeadsQ = numHeadsQ; + p.mNumHeadsKv = 1; + p.mNumHeadsQPerKv = numHeadsQ; + p.mBatchSize = 1; + p.mMaxSeqLenQ = maxSeqLenQ; + p.mMaxSeqLenKv = 1024; + p.mNumTokensPerPage = 32; + p.mChunkedAttentionSize = INT_MAX; + p.mAttentionWindowSize = INT_MAX; + p.mScaleQ = 1.f; + p.mSparseAttention = SparseType::None; + p.mSparseTopK = 0; + // Query the actual device's SM count instead of hard-coding B300=148; + // the autotuner's CTA heuristics consume this value and a wrong count + // can perturb the selected tuple. + p.mMultiProcessorCount = tensorrt_llm::common::getMultiProcessorCount(); + p.mSumOfSeqLensQ = maxSeqLenQ * p.mBatchSize; + p.mSumOfSeqLensKv = p.mMaxSeqLenKv * p.mBatchSize; + p.mMaxNumPagesPerSeqKv = (p.mMaxSeqLenKv + p.mNumTokensPerPage - 1) / p.mNumTokensPerPage; + } +}; + +// The Kimi shape must resolve to the SM103 grouped cubin. +TEST_F(KimiMlaGroupedSelectionTest, KimiShape_SelectsGroupedCubin) +{ + TllmGenFmhaRunner runner(DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16); + TllmGenFmhaRunnerParams params; + buildKimiParams(params); + + auto [supported, info] = runner.isSupportedWithInfo(params); + EXPECT_TRUE(supported) << "Cubin lookup failed for Kimi MLA grouped shape: " << info; + + // Exact selected-cubin proof via the test-only probe: consults the matched + // mFunctions[hashId] -> mKernelMeta entry directly. + auto selected = runner.probeKernelSelectionForTesting(params); + EXPECT_TRUE(selected.mFound) << "probeKernelSelectionForTesting reported no match"; + EXPECT_FALSE(selected.mUsedNvrtc) << "probe took the NVRTC path instead of the grouped cubin"; + // The autotuner picks the headDimPerCtaV / reduction variant; assert the grouped MLA family. + EXPECT_NE(selected.mFuncName.find("HQk576HV512"), std::string::npos) << selected.mFuncName; + EXPECT_NE(selected.mFuncName.find("Q64Kv128"), std::string::npos) << selected.mFuncName; + EXPECT_NE(selected.mFuncName.find("GroupedKeepsAbForGen"), std::string::npos) << selected.mFuncName; + EXPECT_TRUE(selected.mGroupsHeadsQ); + EXPECT_TRUE(selected.mGroupsTokensHeadsQ); +} + +// q_len=1 (plain decode, M=16) must NOT select the grouped Q64 cubin. +TEST_F(KimiMlaGroupedSelectionTest, NonKimiShape_DoesNotSelectGroupedCubin) +{ + TllmGenFmhaRunner runner(DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16); + TllmGenFmhaRunnerParams params; + buildKimiParams(params, /*maxSeqLenQ=*/1); + + auto [supported, info] = runner.isSupportedWithInfo(params); + EXPECT_TRUE(supported) << "MLA decode path reported unsupported with q_len=1: " << info; + + auto selected = runner.probeKernelSelectionForTesting(params); + EXPECT_FALSE(selected.mGroupsTokensHeadsQ) << "Grouped cubin selected for a non-spec-decode shape: " + << selected.mFuncName; +} + +// Smoke run with Q=K=V=zero: both the FMHA and separate-reduction kernels must launch +// cleanly, every output slot must be written (no surviving 0xFFFF sentinel), and the +// output must be exactly zero. +TEST_F(KimiMlaGroupedSelectionTest, KimiShape_RunSmokeSucceeds) +{ + TllmGenFmhaRunner runner(DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16); + TllmGenFmhaRunnerParams params; + buildKimiParams(params); + + int const seqLenQ = params.mMaxSeqLenQ; // 4 + int const seqLenKv = params.mMaxSeqLenKv; // 1024 + int const batchSize = params.mBatchSize; // 1 + int const numHeadsQ = params.mNumHeadsQ; // 16 + int const headDimQk = params.mHeadDimQk; // 576 + int const headDimV = params.mHeadDimV; // 512 + int const numTokensPerPage = params.mNumTokensPerPage; // 32 + int const maxNumPagesPerSeqKv = params.mMaxNumPagesPerSeqKv; // ceilDiv(1024,32)=32 + int const numPages = maxNumPagesPerSeqKv * batchSize; // 32 + + constexpr size_t kBf16 = sizeof(uint16_t); + + // Allocate device buffers. Sizes mirror the Kimi MLA paged-KV layout: + // Q : [batch * seqLenQ, numHeadsQ, headDimQk] + // KV : [num_pages, page_size, headDimQk] - shared K/V cache, MLA layout + // O : [batch * seqLenQ, numHeadsQ, headDimV] + size_t const qBytes = static_cast(batchSize) * seqLenQ * numHeadsQ * headDimQk * kBf16; + size_t const kvBytes = static_cast(numPages) * numTokensPerPage * headDimQk * kBf16; + size_t const oBytes = static_cast(batchSize) * seqLenQ * numHeadsQ * headDimV * kBf16; + // multiCtasKv scratch holds partialStats (float2 per CTA) + partialO + // (float per Q tile slot * headDimV). 64 MB is comfortably above what + // the SM103 grouped grid requires (~4 MB) for this shape. + size_t const scratchBytes = static_cast(64) * 1024 * 1024; + + void* dQ = nullptr; + void* dKV = nullptr; + void* dO = nullptr; + void* dScratch = nullptr; + int32_t* dCounter = nullptr; + int32_t* dPageIdx = nullptr; + int32_t* dSeqLensKv = nullptr; + int32_t* dCumSeqLensQ = nullptr; + int32_t* dCumSeqLensKv = nullptr; + float* dScaleSoftmaxLog2 = nullptr; + float* dOutputScale = nullptr; + cudaStream_t stream = nullptr; + + ASSERT_EQ(cudaSuccess, cudaMalloc(&dQ, qBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dKV, kvBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dO, oBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dScratch, scratchBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCounter, kCounterBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dPageIdx, batchSize * maxNumPagesPerSeqKv * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dSeqLensKv, batchSize * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCumSeqLensQ, (batchSize + 1) * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCumSeqLensKv, (batchSize + 1) * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dScaleSoftmaxLog2, sizeof(float))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dOutputScale, sizeof(float))); + ASSERT_EQ(cudaSuccess, cudaStreamCreate(&stream)); + + // Zero the input + scratch buffers. With Q=K=V=zero, the expected output + // is zero (uniform softmax over zero scores -> uniform attention; uniform + // weights * V=zero -> zero per element). Pre-fill dO with a NaN sentinel + // (BF16 bits 0xFFFF = sign=1, exp=255, mantissa=0x7F -> NaN) so we can + // detect whether the kernel actually wrote every output position. Any + // 0xFFFF that survives after run() means the kernel skipped that slot. + ASSERT_EQ(cudaSuccess, cudaMemset(dQ, 0, qBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dKV, 0, kvBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dO, 0xFF, oBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dScratch, 0, scratchBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dCounter, 0, kCounterBytes)); + + // Page table: sequential page indices [0, 1, ..., numPages-1]. + std::vector hPageIdx(batchSize * maxNumPagesPerSeqKv); + for (int i = 0; i < static_cast(hPageIdx.size()); ++i) + { + hPageIdx[i] = i; + } + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dPageIdx, hPageIdx.data(), hPageIdx.size() * sizeof(int32_t), cudaMemcpyHostToDevice, stream)); + + // seqLensKv = [seqLenKv]; cumSeqLensQ = [0, seqLenQ]; cumSeqLensKv = [0, seqLenKv]. + std::vector const hSeqLensKv = {seqLenKv}; + std::vector const hCumSeqLensQ = {0, seqLenQ}; + std::vector const hCumSeqLensKv = {0, seqLenKv}; + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync( + dSeqLensKv, hSeqLensKv.data(), batchSize * sizeof(int32_t), cudaMemcpyHostToDevice, stream)); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dCumSeqLensQ, hCumSeqLensQ.data(), (batchSize + 1) * sizeof(int32_t), + cudaMemcpyHostToDevice, stream)); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dCumSeqLensKv, hCumSeqLensKv.data(), (batchSize + 1) * sizeof(int32_t), + cudaMemcpyHostToDevice, stream)); + + // softmaxScale_log2 = (1 / sqrt(headDimQk)) * log2(e). Mirrors the host + // value setFmhaData computes from params.mScaleQ=1 and mHeadDimQk. + float const hScaleSoftmaxLog2 = (1.f / std::sqrt(static_cast(headDimQk))) * static_cast(M_LOG2E); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dScaleSoftmaxLog2, &hScaleSoftmaxLog2, sizeof(float), cudaMemcpyHostToDevice, stream)); + float const hOutputScale = 1.0f; + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dOutputScale, &hOutputScale, sizeof(float), cudaMemcpyHostToDevice, stream)); + + // Wire up RunnerParams pointers. + params.qPtr = dQ; + params.kvPtr = dKV; + params.oPtr = dO; + params.kvPageIdxPtr = dPageIdx; + params.seqLensKvPtr = dSeqLensKv; + params.cumSeqLensQPtr = dCumSeqLensQ; + params.cumSeqLensKvPtr = dCumSeqLensKv; + params.scaleSoftmaxLog2Ptr = dScaleSoftmaxLog2; + params.outputScalePtr = dOutputScale; + params.multiCtasKvScratchPtr = dScratch; + params.multiCtasKvCounterPtr = dCounter; + params.stream = stream; + + // Sync any pending host-to-device copies so the kernel reads valid data. + ASSERT_EQ(cudaSuccess, cudaStreamSynchronize(stream)); + + ASSERT_NO_THROW(runner.run(params)); + + // Make sure both the FMHA kernel and the separate reduction kernel + // finished without runtime errors. + cudaError_t const syncErr = cudaStreamSynchronize(stream); + EXPECT_EQ(cudaSuccess, syncErr) << "cudaStreamSynchronize after run: " << cudaGetErrorString(syncErr); + cudaError_t const lastErr = cudaGetLastError(); + EXPECT_EQ(cudaSuccess, lastErr) << "cudaGetLastError after run: " << cudaGetErrorString(lastErr); + + // Output equivalence: pull dO back and verify (a) every BF16 element was + // overwritten (no 0xFFFF NaN sentinel survives -> the kernel actually + // wrote every output slot), (b) no NaN or Inf was produced, (c) every + // element equals exactly 0.0 because softmax(QK^T)=uniform and V=zero + // implies output=zero with no accumulation error. + size_t const numOutputElems = oBytes / sizeof(uint16_t); + std::vector hO(numOutputElems); + ASSERT_EQ(cudaSuccess, cudaMemcpy(hO.data(), dO, oBytes, cudaMemcpyDeviceToHost)); + + int numSentinelRemnants = 0; + int numNaN = 0; + int numInf = 0; + int numNonZero = 0; + for (uint16_t bits : hO) + { + if (bits == kBf16NaNSentinel) + { + ++numSentinelRemnants; + } + float const f = bf16BitsToFloat(bits); + if (std::isnan(f)) + { + ++numNaN; + } + if (std::isinf(f)) + { + ++numInf; + } + if (f != 0.0f) + { + ++numNonZero; + } + } + EXPECT_EQ(0, numSentinelRemnants) + << "0xFFFF NaN sentinels remain: " << numSentinelRemnants << "/" << numOutputElems + << " elements not written by the kernel (zero input should still cover every output position)"; + EXPECT_EQ(0, numNaN) << "Found " << numNaN << "/" << numOutputElems << " NaN elements in output"; + EXPECT_EQ(0, numInf) << "Found " << numInf << "/" << numOutputElems << " Inf elements in output"; + EXPECT_EQ(0, numNonZero) << "Zero input must produce zero output; found " << numNonZero << "/" + << numOutputElems << " non-zero elements"; + + // Cleanup. + cudaFree(dQ); + cudaFree(dKV); + cudaFree(dO); + cudaFree(dScratch); + cudaFree(dCounter); + cudaFree(dPageIdx); + cudaFree(dSeqLensKv); + cudaFree(dCumSeqLensQ); + cudaFree(dCumSeqLensKv); + cudaFree(dScaleSoftmaxLog2); + cudaFree(dOutputScale); + cudaStreamDestroy(stream); +} + +// Q=zero (uniform softmax) + KV=1.0 everywhere: output must be ~1.0 per element. +// Unlike the all-zero case this exercises V reads from the paged cache and cross-tile +// aggregation through the separate GMEM reduction kernel. +TEST_F(KimiMlaGroupedSelectionTest, KimiShape_RunSmokeConstantKVOutputsOne) +{ + TllmGenFmhaRunner runner(DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16); + TllmGenFmhaRunnerParams params; + buildKimiParams(params); + + int const seqLenQ = params.mMaxSeqLenQ; + int const seqLenKv = params.mMaxSeqLenKv; + int const batchSize = params.mBatchSize; + int const numHeadsQ = params.mNumHeadsQ; + int const headDimQk = params.mHeadDimQk; + int const headDimV = params.mHeadDimV; + int const numTokensPerPage = params.mNumTokensPerPage; + int const maxNumPagesPerSeqKv = params.mMaxNumPagesPerSeqKv; + int const numPages = maxNumPagesPerSeqKv * batchSize; + + constexpr size_t kBf16 = sizeof(uint16_t); + size_t const qBytes = static_cast(batchSize) * seqLenQ * numHeadsQ * headDimQk * kBf16; + size_t const kvBytes = static_cast(numPages) * numTokensPerPage * headDimQk * kBf16; + size_t const oBytes = static_cast(batchSize) * seqLenQ * numHeadsQ * headDimV * kBf16; + size_t const scratchBytes = static_cast(64) * 1024 * 1024; + + void* dQ = nullptr; + void* dKV = nullptr; + void* dO = nullptr; + void* dScratch = nullptr; + int32_t* dCounter = nullptr; + int32_t* dPageIdx = nullptr; + int32_t* dSeqLensKv = nullptr; + int32_t* dCumSeqLensQ = nullptr; + int32_t* dCumSeqLensKv = nullptr; + float* dScaleSoftmaxLog2 = nullptr; + float* dOutputScale = nullptr; + cudaStream_t stream = nullptr; + + ASSERT_EQ(cudaSuccess, cudaMalloc(&dQ, qBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dKV, kvBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dO, oBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dScratch, scratchBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCounter, kCounterBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dPageIdx, batchSize * maxNumPagesPerSeqKv * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dSeqLensKv, batchSize * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCumSeqLensQ, (batchSize + 1) * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCumSeqLensKv, (batchSize + 1) * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dScaleSoftmaxLog2, sizeof(float))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dOutputScale, sizeof(float))); + ASSERT_EQ(cudaSuccess, cudaStreamCreate(&stream)); + + // Q is all zeros. KV cache is BF16 1.0 across every entry (both the K + // slice the kernel reads for scores and the V slice it reads for the + // weighted sum). dO starts as NaN-sentinel so we can detect un-written + // slots after the kernel runs. + ASSERT_EQ(cudaSuccess, cudaMemset(dQ, 0, qBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dO, 0xFF, oBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dScratch, 0, scratchBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dCounter, 0, kCounterBytes)); + std::vector const hKV(kvBytes / sizeof(uint16_t), kBf16One); + ASSERT_EQ(cudaSuccess, cudaMemcpyAsync(dKV, hKV.data(), kvBytes, cudaMemcpyHostToDevice, stream)); + + std::vector hPageIdx(batchSize * maxNumPagesPerSeqKv); + for (int i = 0; i < static_cast(hPageIdx.size()); ++i) + { + hPageIdx[i] = i; + } + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dPageIdx, hPageIdx.data(), hPageIdx.size() * sizeof(int32_t), cudaMemcpyHostToDevice, stream)); + std::vector const hSeqLensKv = {seqLenKv}; + std::vector const hCumSeqLensQ = {0, seqLenQ}; + std::vector const hCumSeqLensKv = {0, seqLenKv}; + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dSeqLensKv, hSeqLensKv.data(), batchSize * sizeof(int32_t), cudaMemcpyHostToDevice, stream)); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dCumSeqLensQ, hCumSeqLensQ.data(), (batchSize + 1) * sizeof(int32_t), cudaMemcpyHostToDevice, + stream)); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dCumSeqLensKv, hCumSeqLensKv.data(), (batchSize + 1) * sizeof(int32_t), + cudaMemcpyHostToDevice, stream)); + float const hScaleSoftmaxLog2 = (1.f / std::sqrt(static_cast(headDimQk))) * static_cast(M_LOG2E); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dScaleSoftmaxLog2, &hScaleSoftmaxLog2, sizeof(float), cudaMemcpyHostToDevice, stream)); + float const hOutputScale = 1.0f; + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dOutputScale, &hOutputScale, sizeof(float), cudaMemcpyHostToDevice, stream)); + + params.qPtr = dQ; + params.kvPtr = dKV; + params.oPtr = dO; + params.kvPageIdxPtr = dPageIdx; + params.seqLensKvPtr = dSeqLensKv; + params.cumSeqLensQPtr = dCumSeqLensQ; + params.cumSeqLensKvPtr = dCumSeqLensKv; + params.scaleSoftmaxLog2Ptr = dScaleSoftmaxLog2; + params.outputScalePtr = dOutputScale; + params.multiCtasKvScratchPtr = dScratch; + params.multiCtasKvCounterPtr = dCounter; + params.stream = stream; + + ASSERT_EQ(cudaSuccess, cudaStreamSynchronize(stream)); + ASSERT_NO_THROW(runner.run(params)); + cudaError_t const syncErr = cudaStreamSynchronize(stream); + EXPECT_EQ(cudaSuccess, syncErr) << "cudaStreamSynchronize after run: " << cudaGetErrorString(syncErr); + cudaError_t const lastErr = cudaGetLastError(); + EXPECT_EQ(cudaSuccess, lastErr) << "cudaGetLastError after run: " << cudaGetErrorString(lastErr); + + size_t const numOutputElems = oBytes / sizeof(uint16_t); + std::vector hO(numOutputElems); + ASSERT_EQ(cudaSuccess, cudaMemcpy(hO.data(), dO, oBytes, cudaMemcpyDeviceToHost)); + + int numSentinelRemnants = 0; + int numNaN = 0; + int numInf = 0; + int numOutOfTolerance = 0; + float minVal = std::numeric_limits::infinity(); + float maxVal = -std::numeric_limits::infinity(); + constexpr float kExpectedValue = 1.0f; + constexpr float kAbsTolerance = 0.03f; // ~3% of |1.0| + for (uint16_t bits : hO) + { + if (bits == kBf16NaNSentinel) + { + ++numSentinelRemnants; + } + float const f = bf16BitsToFloat(bits); + if (std::isnan(f)) + { + ++numNaN; + } + if (std::isinf(f)) + { + ++numInf; + } + if (std::isfinite(f)) + { + minVal = std::min(minVal, f); + maxVal = std::max(maxVal, f); + if (std::abs(f - kExpectedValue) > kAbsTolerance) + { + ++numOutOfTolerance; + } + } + } + EXPECT_EQ(0, numSentinelRemnants) << "0xFFFF NaN sentinels remain: " << numSentinelRemnants << "/" << numOutputElems; + EXPECT_EQ(0, numNaN) << "Found " << numNaN << "/" << numOutputElems << " NaN elements"; + EXPECT_EQ(0, numInf) << "Found " << numInf << "/" << numOutputElems << " Inf elements"; + EXPECT_EQ(0, numOutOfTolerance) << "Output elements outside |x - 1.0| <= " << kAbsTolerance << ": " + << numOutOfTolerance << "/" << numOutputElems << "; min=" << minVal + << " max=" << maxVal; + + cudaFree(dQ); + cudaFree(dKV); + cudaFree(dO); + cudaFree(dScratch); + cudaFree(dCounter); + cudaFree(dPageIdx); + cudaFree(dSeqLensKv); + cudaFree(dCumSeqLensQ); + cudaFree(dCumSeqLensKv); + cudaFree(dScaleSoftmaxLog2); + cudaFree(dOutputScale); + cudaStreamDestroy(stream); +} + +// Q=zero + KV split into two halves (pages 0..15 -> 1.0, pages 16..31 -> 2.0): +// uniform attention over 1024 positions gives output ~1.5 per element. Catches +// dropped/double-counted halves of the multi-CTA reduction; per-position indexing +// is covered by the patterned-V and CPU-reference tests below. (In causal +// spec-decode the 4 packed Q tokens have slightly different visible lengths near +// the tail — far smaller than the tolerance, so no exact 1.5 check.) +TEST_F(KimiMlaGroupedSelectionTest, KimiShape_RunSmokeSplitKVOutputsOnePointFive) +{ + TllmGenFmhaRunner runner(DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16); + TllmGenFmhaRunnerParams params; + buildKimiParams(params); + + int const seqLenQ = params.mMaxSeqLenQ; + int const seqLenKv = params.mMaxSeqLenKv; + int const batchSize = params.mBatchSize; + int const numHeadsQ = params.mNumHeadsQ; + int const headDimQk = params.mHeadDimQk; + int const headDimV = params.mHeadDimV; + int const numTokensPerPage = params.mNumTokensPerPage; + int const maxNumPagesPerSeqKv = params.mMaxNumPagesPerSeqKv; + int const numPages = maxNumPagesPerSeqKv * batchSize; + ASSERT_EQ(0, numPages % 2) << "split-V test assumes an even number of pages"; + + constexpr size_t kBf16 = sizeof(uint16_t); + size_t const qBytes = static_cast(batchSize) * seqLenQ * numHeadsQ * headDimQk * kBf16; + size_t const kvBytes = static_cast(numPages) * numTokensPerPage * headDimQk * kBf16; + size_t const oBytes = static_cast(batchSize) * seqLenQ * numHeadsQ * headDimV * kBf16; + size_t const scratchBytes = static_cast(64) * 1024 * 1024; + + void* dQ = nullptr; + void* dKV = nullptr; + void* dO = nullptr; + void* dScratch = nullptr; + int32_t* dCounter = nullptr; + int32_t* dPageIdx = nullptr; + int32_t* dSeqLensKv = nullptr; + int32_t* dCumSeqLensQ = nullptr; + int32_t* dCumSeqLensKv = nullptr; + float* dScaleSoftmaxLog2 = nullptr; + float* dOutputScale = nullptr; + cudaStream_t stream = nullptr; + + ASSERT_EQ(cudaSuccess, cudaMalloc(&dQ, qBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dKV, kvBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dO, oBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dScratch, scratchBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCounter, kCounterBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dPageIdx, batchSize * maxNumPagesPerSeqKv * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dSeqLensKv, batchSize * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCumSeqLensQ, (batchSize + 1) * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCumSeqLensKv, (batchSize + 1) * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dScaleSoftmaxLog2, sizeof(float))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dOutputScale, sizeof(float))); + ASSERT_EQ(cudaSuccess, cudaStreamCreate(&stream)); + + ASSERT_EQ(cudaSuccess, cudaMemset(dQ, 0, qBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dO, 0xFF, oBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dScratch, 0, scratchBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dCounter, 0, kCounterBytes)); + + // Build the split KV cache on host. Layout: [page, slot_in_page, headDim]. + // Every BF16 element of a slot gets the same value (1.0 for the first half + // of pages, 2.0 for the second half) so both the K slice (used for scores) + // and the V slice (used for the weighted sum) are uniform within a page. + std::vector hKV(kvBytes / sizeof(uint16_t)); + size_t const elemsPerPage = static_cast(numTokensPerPage) * headDimQk; + int const halfPages = numPages / 2; + for (int page = 0; page < numPages; ++page) + { + uint16_t const value = page < halfPages ? kBf16One : kBf16Two; + size_t const pageStart = static_cast(page) * elemsPerPage; + std::fill_n(hKV.begin() + pageStart, elemsPerPage, value); + } + ASSERT_EQ(cudaSuccess, cudaMemcpyAsync(dKV, hKV.data(), kvBytes, cudaMemcpyHostToDevice, stream)); + + std::vector hPageIdx(batchSize * maxNumPagesPerSeqKv); + for (int i = 0; i < static_cast(hPageIdx.size()); ++i) + { + hPageIdx[i] = i; + } + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dPageIdx, hPageIdx.data(), hPageIdx.size() * sizeof(int32_t), cudaMemcpyHostToDevice, stream)); + std::vector const hSeqLensKv = {seqLenKv}; + std::vector const hCumSeqLensQ = {0, seqLenQ}; + std::vector const hCumSeqLensKv = {0, seqLenKv}; + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dSeqLensKv, hSeqLensKv.data(), batchSize * sizeof(int32_t), cudaMemcpyHostToDevice, stream)); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dCumSeqLensQ, hCumSeqLensQ.data(), (batchSize + 1) * sizeof(int32_t), cudaMemcpyHostToDevice, + stream)); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dCumSeqLensKv, hCumSeqLensKv.data(), (batchSize + 1) * sizeof(int32_t), + cudaMemcpyHostToDevice, stream)); + float const hScaleSoftmaxLog2 = (1.f / std::sqrt(static_cast(headDimQk))) * static_cast(M_LOG2E); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dScaleSoftmaxLog2, &hScaleSoftmaxLog2, sizeof(float), cudaMemcpyHostToDevice, stream)); + float const hOutputScale = 1.0f; + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dOutputScale, &hOutputScale, sizeof(float), cudaMemcpyHostToDevice, stream)); + + params.qPtr = dQ; + params.kvPtr = dKV; + params.oPtr = dO; + params.kvPageIdxPtr = dPageIdx; + params.seqLensKvPtr = dSeqLensKv; + params.cumSeqLensQPtr = dCumSeqLensQ; + params.cumSeqLensKvPtr = dCumSeqLensKv; + params.scaleSoftmaxLog2Ptr = dScaleSoftmaxLog2; + params.outputScalePtr = dOutputScale; + params.multiCtasKvScratchPtr = dScratch; + params.multiCtasKvCounterPtr = dCounter; + params.stream = stream; + + ASSERT_EQ(cudaSuccess, cudaStreamSynchronize(stream)); + ASSERT_NO_THROW(runner.run(params)); + cudaError_t const syncErr = cudaStreamSynchronize(stream); + EXPECT_EQ(cudaSuccess, syncErr) << "cudaStreamSynchronize after run: " << cudaGetErrorString(syncErr); + cudaError_t const lastErr = cudaGetLastError(); + EXPECT_EQ(cudaSuccess, lastErr) << "cudaGetLastError after run: " << cudaGetErrorString(lastErr); + + size_t const numOutputElems = oBytes / sizeof(uint16_t); + std::vector hO(numOutputElems); + ASSERT_EQ(cudaSuccess, cudaMemcpy(hO.data(), dO, oBytes, cudaMemcpyDeviceToHost)); + + int numSentinelRemnants = 0; + int numNaN = 0; + int numInf = 0; + int numOutOfTolerance = 0; + float minVal = std::numeric_limits::infinity(); + float maxVal = -std::numeric_limits::infinity(); + constexpr float kExpectedValue = 1.5f; + constexpr float kAbsTolerance = 0.03f; + for (uint16_t bits : hO) + { + if (bits == kBf16NaNSentinel) + { + ++numSentinelRemnants; + } + float const f = bf16BitsToFloat(bits); + if (std::isnan(f)) + { + ++numNaN; + } + if (std::isinf(f)) + { + ++numInf; + } + if (std::isfinite(f)) + { + minVal = std::min(minVal, f); + maxVal = std::max(maxVal, f); + if (std::abs(f - kExpectedValue) > kAbsTolerance) + { + ++numOutOfTolerance; + } + } + } + EXPECT_EQ(0, numSentinelRemnants) << "0xFFFF NaN sentinels remain: " << numSentinelRemnants << "/" << numOutputElems; + EXPECT_EQ(0, numNaN) << "Found " << numNaN << "/" << numOutputElems << " NaN elements"; + EXPECT_EQ(0, numInf) << "Found " << numInf << "/" << numOutputElems << " Inf elements"; + EXPECT_EQ(0, numOutOfTolerance) << "Output elements outside |x - 1.5| <= " << kAbsTolerance << ": " + << numOutOfTolerance << "/" << numOutputElems << "; min=" << minVal + << " max=" << maxVal; + + cudaFree(dQ); + cudaFree(dKV); + cudaFree(dO); + cudaFree(dScratch); + cudaFree(dCounter); + cudaFree(dPageIdx); + cudaFree(dSeqLensKv); + cudaFree(dCumSeqLensQ); + cudaFree(dCumSeqLensKv); + cudaFree(dScaleSoftmaxLog2); + cudaFree(dOutputScale); + cudaStreamDestroy(stream); +} + +// Q=zero + V[kv, d_v] = (d_v % 8) + 1 (identical across kv, exactly representable in +// BF16): uniform attention gives output[q, h, d_v] = (d_v % 8) + 1. Complements the +// split-half test by varying V along the d_v axis, so d_v mis-indexing (transposed or +// shifted strides) trips the per-element compare. +TEST_F(KimiMlaGroupedSelectionTest, KimiShape_RunSmokePatternedVOutputsPerDim) +{ + TllmGenFmhaRunner runner(DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16); + TllmGenFmhaRunnerParams params; + buildKimiParams(params); + + int const seqLenQ = params.mMaxSeqLenQ; + int const seqLenKv = params.mMaxSeqLenKv; + int const batchSize = params.mBatchSize; + int const numHeadsQ = params.mNumHeadsQ; + int const headDimQk = params.mHeadDimQk; + int const headDimV = params.mHeadDimV; + int const numTokensPerPage = params.mNumTokensPerPage; + int const maxNumPagesPerSeqKv = params.mMaxNumPagesPerSeqKv; + int const numPages = maxNumPagesPerSeqKv * batchSize; + + constexpr size_t kBf16 = sizeof(uint16_t); + size_t const qBytes = static_cast(batchSize) * seqLenQ * numHeadsQ * headDimQk * kBf16; + size_t const kvBytes = static_cast(numPages) * numTokensPerPage * headDimQk * kBf16; + size_t const oBytes = static_cast(batchSize) * seqLenQ * numHeadsQ * headDimV * kBf16; + size_t const scratchBytes = static_cast(64) * 1024 * 1024; + + void* dQ = nullptr; + void* dKV = nullptr; + void* dO = nullptr; + void* dScratch = nullptr; + int32_t* dCounter = nullptr; + int32_t* dPageIdx = nullptr; + int32_t* dSeqLensKv = nullptr; + int32_t* dCumSeqLensQ = nullptr; + int32_t* dCumSeqLensKv = nullptr; + float* dScaleSoftmaxLog2 = nullptr; + float* dOutputScale = nullptr; + cudaStream_t stream = nullptr; + + ASSERT_EQ(cudaSuccess, cudaMalloc(&dQ, qBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dKV, kvBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dO, oBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dScratch, scratchBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCounter, kCounterBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dPageIdx, batchSize * maxNumPagesPerSeqKv * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dSeqLensKv, batchSize * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCumSeqLensQ, (batchSize + 1) * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCumSeqLensKv, (batchSize + 1) * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dScaleSoftmaxLog2, sizeof(float))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dOutputScale, sizeof(float))); + ASSERT_EQ(cudaSuccess, cudaStreamCreate(&stream)); + + ASSERT_EQ(cudaSuccess, cudaMemset(dQ, 0, qBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dO, 0xFF, oBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dScratch, 0, scratchBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dCounter, 0, kCounterBytes)); + + // Build the patterned KV cache. Same value V_d = (d % 8) + 1 for every + // kv across the full head_dim_qk = 576. BF16 encoding for integer N in + // [1, 8] is float-to-BF16 of the integer (top 16 bits of the FP32 word). + auto floatToBf16 = [](float v) -> uint16_t { + uint32_t bits; + std::memcpy(&bits, &v, sizeof(float)); + return static_cast(bits >> 16); + }; + std::vector hKV(kvBytes / sizeof(uint16_t)); + size_t const totalKvSlots = static_cast(numPages) * numTokensPerPage; + for (size_t kv = 0; kv < totalKvSlots; ++kv) + { + for (int d = 0; d < headDimQk; ++d) + { + float const value = static_cast((d % 8) + 1); // 1..8 + hKV[kv * headDimQk + d] = floatToBf16(value); + } + } + ASSERT_EQ(cudaSuccess, cudaMemcpyAsync(dKV, hKV.data(), kvBytes, cudaMemcpyHostToDevice, stream)); + + std::vector hPageIdx(batchSize * maxNumPagesPerSeqKv); + for (int i = 0; i < static_cast(hPageIdx.size()); ++i) + { + hPageIdx[i] = i; + } + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dPageIdx, hPageIdx.data(), hPageIdx.size() * sizeof(int32_t), cudaMemcpyHostToDevice, stream)); + std::vector const hSeqLensKv = {seqLenKv}; + std::vector const hCumSeqLensQ = {0, seqLenQ}; + std::vector const hCumSeqLensKv = {0, seqLenKv}; + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dSeqLensKv, hSeqLensKv.data(), batchSize * sizeof(int32_t), cudaMemcpyHostToDevice, stream)); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dCumSeqLensQ, hCumSeqLensQ.data(), (batchSize + 1) * sizeof(int32_t), cudaMemcpyHostToDevice, + stream)); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dCumSeqLensKv, hCumSeqLensKv.data(), (batchSize + 1) * sizeof(int32_t), + cudaMemcpyHostToDevice, stream)); + float const hScaleSoftmaxLog2 = (1.f / std::sqrt(static_cast(headDimQk))) * static_cast(M_LOG2E); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dScaleSoftmaxLog2, &hScaleSoftmaxLog2, sizeof(float), cudaMemcpyHostToDevice, stream)); + float const hOutputScale = 1.0f; + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dOutputScale, &hOutputScale, sizeof(float), cudaMemcpyHostToDevice, stream)); + + params.qPtr = dQ; + params.kvPtr = dKV; + params.oPtr = dO; + params.kvPageIdxPtr = dPageIdx; + params.seqLensKvPtr = dSeqLensKv; + params.cumSeqLensQPtr = dCumSeqLensQ; + params.cumSeqLensKvPtr = dCumSeqLensKv; + params.scaleSoftmaxLog2Ptr = dScaleSoftmaxLog2; + params.outputScalePtr = dOutputScale; + params.multiCtasKvScratchPtr = dScratch; + params.multiCtasKvCounterPtr = dCounter; + params.stream = stream; + + ASSERT_EQ(cudaSuccess, cudaStreamSynchronize(stream)); + ASSERT_NO_THROW(runner.run(params)); + cudaError_t const syncErr = cudaStreamSynchronize(stream); + EXPECT_EQ(cudaSuccess, syncErr) << "cudaStreamSynchronize after run: " << cudaGetErrorString(syncErr); + cudaError_t const lastErr = cudaGetLastError(); + EXPECT_EQ(cudaSuccess, lastErr) << "cudaGetLastError after run: " << cudaGetErrorString(lastErr); + + size_t const numOutputElems = oBytes / sizeof(uint16_t); + std::vector hO(numOutputElems); + ASSERT_EQ(cudaSuccess, cudaMemcpy(hO.data(), dO, oBytes, cudaMemcpyDeviceToHost)); + + // Output layout: [batch, q_len, num_heads, head_dim_v] = [1, 4, 16, 512] + // contiguous; element (q, h, d_v) lives at index q*numHeadsQ*headDimV + + // h*headDimV + d_v. For each (q, h), the d_v dimension cycles 1..8 every + // 8 elements. + int numSentinelRemnants = 0; + int numNaN = 0; + int numInf = 0; + int numOutOfTolerance = 0; + float minVal = std::numeric_limits::infinity(); + float maxVal = -std::numeric_limits::infinity(); + constexpr float kAbsTolerance = 0.05f; + for (int q = 0; q < seqLenQ; ++q) + { + for (int h = 0; h < numHeadsQ; ++h) + { + for (int d_v = 0; d_v < headDimV; ++d_v) + { + size_t const idx + = static_cast(q) * numHeadsQ * headDimV + static_cast(h) * headDimV + d_v; + uint16_t const bits = hO[idx]; + if (bits == kBf16NaNSentinel) + { + ++numSentinelRemnants; + } + float const f = bf16BitsToFloat(bits); + if (std::isnan(f)) + { + ++numNaN; + } + if (std::isinf(f)) + { + ++numInf; + } + if (std::isfinite(f)) + { + minVal = std::min(minVal, f); + maxVal = std::max(maxVal, f); + float const expected = static_cast((d_v % 8) + 1); + if (std::abs(f - expected) > kAbsTolerance) + { + ++numOutOfTolerance; + } + } + } + } + } + EXPECT_EQ(0, numSentinelRemnants) << "0xFFFF NaN sentinels remain: " << numSentinelRemnants << "/" << numOutputElems; + EXPECT_EQ(0, numNaN) << "Found " << numNaN << "/" << numOutputElems << " NaN elements"; + EXPECT_EQ(0, numInf) << "Found " << numInf << "/" << numOutputElems << " Inf elements"; + EXPECT_EQ(0, numOutOfTolerance) << "Output elements outside |x - expected(d_v)| <= " << kAbsTolerance << ": " + << numOutOfTolerance << "/" << numOutputElems << "; min=" << minVal + << " max=" << maxVal; + + cudaFree(dQ); + cudaFree(dKV); + cudaFree(dO); + cudaFree(dScratch); + cudaFree(dCounter); + cudaFree(dPageIdx); + cudaFree(dSeqLensKv); + cudaFree(dCumSeqLensQ); + cudaFree(dCumSeqLensKv); + cudaFree(dScaleSoftmaxLog2); + cudaFree(dOutputScale); + cudaStreamDestroy(stream); +} + +// Random-input vs CPU FP32 reference (mlaReferenceCpu) diff — covers the +// non-uniform-softmax, causal-spec-decode-mask, and arbitrary-V cases the hand-derived +// tests cannot. Inputs are uniform [-1, 1] so post-scale logits have stddev ~0.58 and +// softmax weights are non-uniform enough (~33x max/min) for the wrong-scale negative +// control below to fire. Tolerance is 0.03 absolute (BF16 precision + 576-dim dot +// product + 1024-position softmax rounding). +// +// Negative control: re-run the CPU reference with the softmax scale doubled and expect +// >= 1% of elements to diverge from the cubin output. The cubin bakes the softmax scale +// in for this shape (perturbing mScaleQ or scaleSoftmaxLog2Ptr does not change its +// output), so the reference is the only scale lever; the control proves the inputs are +// non-uniform enough to catch a baked-in-scale drift via the positive case. +TEST_F(KimiMlaGroupedSelectionTest, KimiShape_RunMatchesCpuMlaReference) +{ + TllmGenFmhaRunner runner(DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16, DATA_TYPE_BF16); + TllmGenFmhaRunnerParams params; + buildKimiParams(params); + + int const seqLenQ = params.mMaxSeqLenQ; + int const seqLenKv = params.mMaxSeqLenKv; + int const batchSize = params.mBatchSize; + int const numHeadsQ = params.mNumHeadsQ; + int const headDimQk = params.mHeadDimQk; + int const headDimV = params.mHeadDimV; + int const numTokensPerPage = params.mNumTokensPerPage; + int const maxNumPagesPerSeqKv = params.mMaxNumPagesPerSeqKv; + int const numPages = maxNumPagesPerSeqKv * batchSize; + ASSERT_EQ(1, batchSize) << "reference assumes batch=1"; + + // MLA softmax scale: 1/sqrt(QK_NOPE + QK_ROPE) = 1/sqrt(128 + 64) = 1/sqrt(192). + constexpr int kQkNopeHeadDim = 128; + constexpr int kQkRopeHeadDim = 64; + float const softmaxScale = 1.f / std::sqrt(static_cast(kQkNopeHeadDim + kQkRopeHeadDim)); + + constexpr size_t kBf16 = sizeof(uint16_t); + size_t const qBytes = static_cast(batchSize) * seqLenQ * numHeadsQ * headDimQk * kBf16; + size_t const kvBytes = static_cast(numPages) * numTokensPerPage * headDimQk * kBf16; + size_t const oBytes = static_cast(batchSize) * seqLenQ * numHeadsQ * headDimV * kBf16; + size_t const scratchBytes = static_cast(64) * 1024 * 1024; + + // Build random Q and KV on host. Deterministic seed so a failure is + // reproducible. Uniform [-1.0, 1.0] makes scores have stddev ~8 and + // post-scale logits stddev ~0.58, so softmax weights are non-uniform + // (max/min ratio ~33x) and the wrong-scale negative control below + // actually fires. + std::mt19937 rng(/*seed=*/42); + std::uniform_real_distribution dist(-1.0f, 1.0f); + + std::vector hQ(qBytes / sizeof(uint16_t)); + for (auto& q : hQ) + { + q = floatToBf16(dist(rng)); + } + std::vector hKV(kvBytes / sizeof(uint16_t)); + for (auto& kv : hKV) + { + kv = floatToBf16(dist(rng)); + } + + // CPU reference: pre-compute expected output. Use the same SOFTMAX_SCALE + // we plumb into the cubin so internal-consistency holds. + size_t const numOutputElems = oBytes / sizeof(uint16_t); + std::vector hOutRef(numOutputElems); + mlaReferenceCpu(seqLenQ, seqLenKv, numHeadsQ, headDimQk, headDimV, numTokensPerPage, hQ, hKV, softmaxScale, hOutRef); + + void* dQ = nullptr; + void* dKV = nullptr; + void* dO = nullptr; + void* dScratch = nullptr; + int32_t* dCounter = nullptr; + int32_t* dPageIdx = nullptr; + int32_t* dSeqLensKv = nullptr; + int32_t* dCumSeqLensQ = nullptr; + int32_t* dCumSeqLensKv = nullptr; + float* dScaleSoftmaxLog2 = nullptr; + float* dOutputScale = nullptr; + cudaStream_t stream = nullptr; + + ASSERT_EQ(cudaSuccess, cudaMalloc(&dQ, qBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dKV, kvBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dO, oBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dScratch, scratchBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCounter, kCounterBytes)); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dPageIdx, batchSize * maxNumPagesPerSeqKv * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dSeqLensKv, batchSize * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCumSeqLensQ, (batchSize + 1) * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dCumSeqLensKv, (batchSize + 1) * sizeof(int32_t))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dScaleSoftmaxLog2, sizeof(float))); + ASSERT_EQ(cudaSuccess, cudaMalloc(&dOutputScale, sizeof(float))); + ASSERT_EQ(cudaSuccess, cudaStreamCreate(&stream)); + + ASSERT_EQ(cudaSuccess, cudaMemset(dO, 0xFF, oBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dScratch, 0, scratchBytes)); + ASSERT_EQ(cudaSuccess, cudaMemset(dCounter, 0, kCounterBytes)); + ASSERT_EQ(cudaSuccess, cudaMemcpyAsync(dQ, hQ.data(), qBytes, cudaMemcpyHostToDevice, stream)); + ASSERT_EQ(cudaSuccess, cudaMemcpyAsync(dKV, hKV.data(), kvBytes, cudaMemcpyHostToDevice, stream)); + + std::vector hPageIdx(batchSize * maxNumPagesPerSeqKv); + for (int i = 0; i < static_cast(hPageIdx.size()); ++i) + { + hPageIdx[i] = i; + } + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dPageIdx, hPageIdx.data(), hPageIdx.size() * sizeof(int32_t), cudaMemcpyHostToDevice, stream)); + std::vector const hSeqLensKv = {seqLenKv}; + std::vector const hCumSeqLensQ = {0, seqLenQ}; + std::vector const hCumSeqLensKv = {0, seqLenKv}; + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dSeqLensKv, hSeqLensKv.data(), batchSize * sizeof(int32_t), cudaMemcpyHostToDevice, stream)); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dCumSeqLensQ, hCumSeqLensQ.data(), (batchSize + 1) * sizeof(int32_t), cudaMemcpyHostToDevice, + stream)); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dCumSeqLensKv, hCumSeqLensKv.data(), (batchSize + 1) * sizeof(int32_t), + cudaMemcpyHostToDevice, stream)); + + // Plumb the same softmax scale to both the device pointer and via mScaleQ: + // setFmhaData computes softmaxScale = (1 / (sqrt(headDimQk) * mScaleQ)) * log2(e), + // so mScaleQ = sqrt(192/576) resolves it to (1/sqrt(192)) * log2(e). + params.mScaleQ = std::sqrt(static_cast(kQkNopeHeadDim + kQkRopeHeadDim) / static_cast(headDimQk)); + float const hScaleSoftmaxLog2 = softmaxScale * static_cast(M_LOG2E); + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dScaleSoftmaxLog2, &hScaleSoftmaxLog2, sizeof(float), cudaMemcpyHostToDevice, stream)); + float const hOutputScale = 1.0f; + ASSERT_EQ(cudaSuccess, + cudaMemcpyAsync(dOutputScale, &hOutputScale, sizeof(float), cudaMemcpyHostToDevice, stream)); + + params.qPtr = dQ; + params.kvPtr = dKV; + params.oPtr = dO; + params.kvPageIdxPtr = dPageIdx; + params.seqLensKvPtr = dSeqLensKv; + params.cumSeqLensQPtr = dCumSeqLensQ; + params.cumSeqLensKvPtr = dCumSeqLensKv; + params.scaleSoftmaxLog2Ptr = dScaleSoftmaxLog2; + params.outputScalePtr = dOutputScale; + params.multiCtasKvScratchPtr = dScratch; + params.multiCtasKvCounterPtr = dCounter; + params.stream = stream; + + ASSERT_EQ(cudaSuccess, cudaStreamSynchronize(stream)); + ASSERT_NO_THROW(runner.run(params)); + cudaError_t const syncErr = cudaStreamSynchronize(stream); + EXPECT_EQ(cudaSuccess, syncErr) << "cudaStreamSynchronize after run: " << cudaGetErrorString(syncErr); + cudaError_t const lastErr = cudaGetLastError(); + EXPECT_EQ(cudaSuccess, lastErr) << "cudaGetLastError after run: " << cudaGetErrorString(lastErr); + + std::vector hOutCubin(numOutputElems); + ASSERT_EQ(cudaSuccess, cudaMemcpy(hOutCubin.data(), dO, oBytes, cudaMemcpyDeviceToHost)); + + // Element-by-element diff. Track count + max abs diff for diagnosis on failure. + int numSentinelRemnants = 0; + int numNaN = 0; + int numInf = 0; + int numOutOfTolerance = 0; + float maxAbsDiff = 0.f; + float maxRefMag = 0.f; + constexpr float kAbsTolerance = 0.03f; + for (size_t i = 0; i < numOutputElems; ++i) + { + uint16_t const cubinBits = hOutCubin[i]; + if (cubinBits == kBf16NaNSentinel) + { + ++numSentinelRemnants; + } + float const fCubin = bf16BitsToFloat(cubinBits); + float const fRef = bf16BitsToFloat(hOutRef[i]); + if (std::isnan(fCubin)) + { + ++numNaN; + } + if (std::isinf(fCubin)) + { + ++numInf; + } + if (std::isfinite(fCubin) && std::isfinite(fRef)) + { + float const absDiff = std::abs(fCubin - fRef); + maxAbsDiff = std::max(maxAbsDiff, absDiff); + maxRefMag = std::max(maxRefMag, std::abs(fRef)); + if (absDiff > kAbsTolerance) + { + ++numOutOfTolerance; + } + } + } + EXPECT_EQ(0, numSentinelRemnants) << numSentinelRemnants << "/" << numOutputElems << " sentinels remain"; + EXPECT_EQ(0, numNaN) << numNaN << "/" << numOutputElems << " NaN elements in cubin output"; + EXPECT_EQ(0, numInf) << numInf << "/" << numOutputElems << " Inf elements in cubin output"; + EXPECT_EQ(0, numOutOfTolerance) + << "cubin vs CPU reference disagrees by > " << kAbsTolerance << " on " << numOutOfTolerance << "/" + << numOutputElems << " elements; maxAbsDiff=" << maxAbsDiff << " maxRefMag=" << maxRefMag; + + // Negative control (see test header comment): a 2x-wrong-scale reference must + // diverge from the cubin output, proving the positive case is scale-sensitive. + std::vector hOutRefWrong(numOutputElems); + mlaReferenceCpu(seqLenQ, seqLenKv, numHeadsQ, headDimQk, headDimV, numTokensPerPage, hQ, hKV, softmaxScale * 2.f, + hOutRefWrong); + int numWrongScaleDivergent = 0; + float maxWrongScaleAbsDiff = 0.f; + for (size_t i = 0; i < numOutputElems; ++i) + { + float const fCubin = bf16BitsToFloat(hOutCubin[i]); + float const fRefWrong = bf16BitsToFloat(hOutRefWrong[i]); + if (std::isfinite(fCubin) && std::isfinite(fRefWrong)) + { + float const absDiff = std::abs(fCubin - fRefWrong); + maxWrongScaleAbsDiff = std::max(maxWrongScaleAbsDiff, absDiff); + if (absDiff > kAbsTolerance) + { + ++numWrongScaleDivergent; + } + } + } + // Expect the 2x-wrong reference to disagree with the cubin on many + // elements. Threshold is conservative (>= 1% of elements) so this is + // robust to the exact random seed. + int const kMinDivergent = static_cast(numOutputElems / 100); + EXPECT_GE(numWrongScaleDivergent, kMinDivergent) + << "Negative control: 2x-wrong-scale CPU reference should diverge from cubin on >= " << kMinDivergent + << " elements (>" << kAbsTolerance << " abs); only " << numWrongScaleDivergent + << " did. Random inputs may be too uniform for scale sensitivity. maxWrongScaleAbsDiff=" + << maxWrongScaleAbsDiff; + + cudaFree(dQ); + cudaFree(dKV); + cudaFree(dO); + cudaFree(dScratch); + cudaFree(dCounter); + cudaFree(dPageIdx); + cudaFree(dSeqLensKv); + cudaFree(dCumSeqLensQ); + cudaFree(dCumSeqLensKv); + cudaFree(dScaleSoftmaxLog2); + cudaFree(dOutputScale); + cudaStreamDestroy(stream); +} + +} // namespace