From c4f603039f38a623118b6da36858491447a3f4dc Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Fri, 21 Aug 2026 06:59:13 -0500 Subject: [PATCH 01/50] feat(ops): port Kimi K3 AscendC operators Port the KDA and SiTU kernels already validated on the v0.26 release branch, update their bindings and build integration, and retain only repository nightly coverage outside csrc. Signed-off-by: maoxx241 --- csrc/attention/chunk_kda_fwd/CMakeLists.txt | 16 + .../chunk_kda_fwd/op_host/CMakeLists.txt | 33 + .../op_host/chunk_kda_fwd_def.cpp | 96 + .../op_host/chunk_kda_fwd_tiling.cpp | 193 ++ .../op_host/chunk_kda_fwd_tiling.h | 47 + .../op_host/op_api/aclnn_chunk_kda_fwd.cpp | 911 ++++++ .../op_host/op_api/aclnn_chunk_kda_fwd.h | 52 + .../op_host/op_api/chunk_kda_fwd.cpp | 156 + .../op_host/op_api/chunk_kda_fwd.h | 47 + .../chunk_kda_fwd/op_kernel/chunk_kda_fwd.cpp | 2584 +++++++++++++++++ csrc/attention/kda_gate_cumsum/CMakeLists.txt | 12 + .../kda_gate_cumsum/op_host/CMakeLists.txt | 19 + .../op_host/kda_gate_cumsum_def.cpp | 60 + .../op_host/kda_gate_cumsum_tiling.cpp | 136 + .../op_host/kda_gate_cumsum_tiling.h | 37 + .../op_host/op_api/aclnn_kda_gate_cumsum.cpp | 232 ++ .../op_host/op_api/aclnn_kda_gate_cumsum.h | 38 + .../op_host/op_api/kda_gate_cumsum.cpp | 54 + .../op_host/op_api/kda_gate_cumsum.h | 26 + .../op_kernel/kda_gate_cumsum.cpp | 330 +++ .../kda_layout_swap12/CMakeLists.txt | 12 + .../kda_layout_swap12/op_host/CMakeLists.txt | 19 + .../op_host/kda_layout_swap12_def.cpp | 48 + .../op_host/kda_layout_swap12_tiling.cpp | 99 + .../op_host/kda_layout_swap12_tiling.h | 27 + .../op_api/aclnn_kda_layout_swap12.cpp | 119 + .../op_host/op_api/aclnn_kda_layout_swap12.h | 31 + .../op_host/op_api/kda_layout_swap12.cpp | 46 + .../op_host/op_api/kda_layout_swap12.h | 24 + .../op_kernel/kda_layout_swap12.cpp | 193 ++ csrc/attention/recurrent_kda/CMakeLists.txt | 13 + .../recurrent_kda/op_host/CMakeLists.txt | 29 + .../op_host/op_api/aclnn_recurrent_kda.cpp | 473 +++ .../op_host/op_api/aclnn_recurrent_kda.h | 52 + .../op_host/op_api/recurrent_kda.cpp | 76 + .../op_host/op_api/recurrent_kda.h | 43 + .../op_host/recurrent_kda_def.cpp | 91 + .../op_host/recurrent_kda_infershape.cpp | 65 + .../op_host/recurrent_kda_tiling.cpp | 439 +++ .../op_host/recurrent_kda_tiling.h | 75 + .../op_host/recurrent_kda_tiling_processor.h | 678 +++++ .../op_kernel/arch35/recurrent_kda.h | 1004 +++++++ .../recurrent_kda/op_kernel/recurrent_kda.cpp | 39 + .../recurrent_kda/op_kernel/recurrent_kda.h | 968 ++++++ .../op_kernel/recurrent_kda_struct.h | 70 + .../op_kernel/recurrent_kda_tiling_data.h | 20 + .../recurrent_kda/recurrent_kda_torch_adpt.h | 157 + csrc/build_aclnn.sh | 15 + .../chunk_gated_delta_rule_fwd_h_def.cpp | 26 +- .../chunk_gated_delta_rule_fwd_h_tiling.cpp | 194 +- .../chunk_gated_delta_rule_fwd_h_tiling.h | 6 +- ..._gated_delta_rule_fwd_h_tiling_processor.h | 152 + .../aclnn_chunk_gated_delta_rule_fwd_h.cpp | 134 +- .../aclnn_chunk_gated_delta_rule_fwd_h.h | 6 +- .../op_api/chunk_gated_delta_rule_fwd_h.cpp | 10 +- .../op_api/chunk_gated_delta_rule_fwd_h.h | 4 +- .../gemm/block/block_scheduler_gdn_fwd_h.hpp | 3 +- .../arch20/gemm/kernel/gdn_fwd_h_kernel.hpp | 3 +- .../block/block_epilogue_gdn_fwdh_update.hpp | 432 ++- .../block/block_epilogue_gdn_fwdh_vnew.hpp | 444 ++- .../epilogue/gdn_fwd_h_epilogue_policies.hpp | 8 - .../gemm/block/block_scheduler_gdn_fwd_h.hpp | 256 +- .../arch22/gemm/kernel/gdn_fwd_h_kernel.hpp | 493 ++-- .../block/block_epilogue_gdn_fwdh_update.hpp | 399 +++ .../block/block_epilogue_gdn_fwdh_vnew.hpp | 508 ++++ .../epilogue/gdn_fwd_h_epilogue_policies.hpp | 27 + .../gemm/block/block_scheduler_gdn_fwd_h.hpp | 340 +++ .../arch35/gemm/kernel/gdn_fwd_h_kernel.hpp | 609 ++++ .../chunk_gated_delta_rule_fwd_h.cpp | 157 +- .../chunk_gated_delta_rule_fwd_h_struct.h | 53 + .../block_mmad_pingpong_tla_preloadA_l1B.hpp | 532 ++++ csrc/moe/dequant_situ_quant/CMakeLists.txt | 19 + .../dequant_situ_quant_torch_adpt.h | 74 + .../docs/aclnnDequantSituQuant.md | 222 ++ .../op_graph/dequant_situ_quant_proto.h | 75 + .../dequant_situ_quant/op_host/CMakeLists.txt | 39 + .../op_host/dequant_situ_quant_def.cpp | 81 + .../op_host/dequant_situ_quant_infershape.cpp | 93 + .../op_host/dequant_situ_quant_tiling.cpp | 701 +++++ .../op_host/dequant_situ_quant_tiling.h | 149 + .../op_kernel/dequant_situ_quant.cpp | 103 + .../op_kernel/dequant_situ_quant.h | 1060 +++++++ csrc/moe/situ_mx_quant/CMakeLists.txt | 19 + .../situ_mx_quant/docs/aclnnSituMxQuant.md | 87 + csrc/moe/situ_mx_quant/op_host/CMakeLists.txt | 32 + .../arch35/situ_mx_quant_tiling_arch35.cpp | 339 +++ .../arch35/situ_mx_quant_tiling_arch35.h | 110 + .../ascend950/situ_mx_quant_binary.json | 87 + .../op_host/situ_mx_quant_def.cpp | 64 + .../op_host/situ_mx_quant_infershape.cpp | 136 + .../arch35/situ_mx_quant_axis_last.h | 280 ++ .../op_kernel/arch35/situ_mx_quant_common.h | 500 ++++ .../arch35/situ_mx_quant_tiling_data.h | 52 + .../arch35/situ_mx_quant_tiling_key.h | 38 + .../op_kernel/inc/kernel_utils.h | 71 + .../situ_mx_quant/op_kernel/inc/platform.h | 81 + .../op_kernel/situ_mx_quant_apt.cpp | 53 + .../situ_mx_quant/situ_mx_quant_torch_adpt.h | 72 + csrc/torch_binding.cpp | 378 +++ csrc/torch_binding_meta.cpp | 271 ++ ...test_chunk_gated_delta_rule_fwd_h_aclnn.py | 334 +++ .../singlecard_ops/test_chunk_kda_aclnn.py | 426 +++ .../singlecard_ops/test_dequant_situ_quant.py | 201 ++ .../test_kimi_k3_situ_fusion_cases.py | 235 ++ .../test_kimi_kda_ascendc_npu.py | 113 + .../test_kimi_kda_recurrent_ascendc_npu.py | 235 ++ 106 files changed, 21002 insertions(+), 628 deletions(-) create mode 100644 csrc/attention/chunk_kda_fwd/CMakeLists.txt create mode 100644 csrc/attention/chunk_kda_fwd/op_host/CMakeLists.txt create mode 100644 csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_def.cpp create mode 100644 csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.cpp create mode 100644 csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.h create mode 100644 csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.cpp create mode 100644 csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.h create mode 100644 csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.cpp create mode 100644 csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.h create mode 100644 csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd.cpp create mode 100644 csrc/attention/kda_gate_cumsum/CMakeLists.txt create mode 100644 csrc/attention/kda_gate_cumsum/op_host/CMakeLists.txt create mode 100644 csrc/attention/kda_gate_cumsum/op_host/kda_gate_cumsum_def.cpp create mode 100644 csrc/attention/kda_gate_cumsum/op_host/kda_gate_cumsum_tiling.cpp create mode 100644 csrc/attention/kda_gate_cumsum/op_host/kda_gate_cumsum_tiling.h create mode 100644 csrc/attention/kda_gate_cumsum/op_host/op_api/aclnn_kda_gate_cumsum.cpp create mode 100644 csrc/attention/kda_gate_cumsum/op_host/op_api/aclnn_kda_gate_cumsum.h create mode 100644 csrc/attention/kda_gate_cumsum/op_host/op_api/kda_gate_cumsum.cpp create mode 100644 csrc/attention/kda_gate_cumsum/op_host/op_api/kda_gate_cumsum.h create mode 100644 csrc/attention/kda_gate_cumsum/op_kernel/kda_gate_cumsum.cpp create mode 100644 csrc/attention/kda_layout_swap12/CMakeLists.txt create mode 100644 csrc/attention/kda_layout_swap12/op_host/CMakeLists.txt create mode 100644 csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_def.cpp create mode 100644 csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_tiling.cpp create mode 100644 csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_tiling.h create mode 100644 csrc/attention/kda_layout_swap12/op_host/op_api/aclnn_kda_layout_swap12.cpp create mode 100644 csrc/attention/kda_layout_swap12/op_host/op_api/aclnn_kda_layout_swap12.h create mode 100644 csrc/attention/kda_layout_swap12/op_host/op_api/kda_layout_swap12.cpp create mode 100644 csrc/attention/kda_layout_swap12/op_host/op_api/kda_layout_swap12.h create mode 100644 csrc/attention/kda_layout_swap12/op_kernel/kda_layout_swap12.cpp create mode 100644 csrc/attention/recurrent_kda/CMakeLists.txt create mode 100644 csrc/attention/recurrent_kda/op_host/CMakeLists.txt create mode 100644 csrc/attention/recurrent_kda/op_host/op_api/aclnn_recurrent_kda.cpp create mode 100644 csrc/attention/recurrent_kda/op_host/op_api/aclnn_recurrent_kda.h create mode 100644 csrc/attention/recurrent_kda/op_host/op_api/recurrent_kda.cpp create mode 100644 csrc/attention/recurrent_kda/op_host/op_api/recurrent_kda.h create mode 100644 csrc/attention/recurrent_kda/op_host/recurrent_kda_def.cpp create mode 100644 csrc/attention/recurrent_kda/op_host/recurrent_kda_infershape.cpp create mode 100644 csrc/attention/recurrent_kda/op_host/recurrent_kda_tiling.cpp create mode 100644 csrc/attention/recurrent_kda/op_host/recurrent_kda_tiling.h create mode 100644 csrc/attention/recurrent_kda/op_host/recurrent_kda_tiling_processor.h create mode 100644 csrc/attention/recurrent_kda/op_kernel/arch35/recurrent_kda.h create mode 100644 csrc/attention/recurrent_kda/op_kernel/recurrent_kda.cpp create mode 100644 csrc/attention/recurrent_kda/op_kernel/recurrent_kda.h create mode 100644 csrc/attention/recurrent_kda/op_kernel/recurrent_kda_struct.h create mode 100644 csrc/attention/recurrent_kda/op_kernel/recurrent_kda_tiling_data.h create mode 100644 csrc/attention/recurrent_kda/recurrent_kda_torch_adpt.h create mode 100644 csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_tiling_processor.h create mode 100644 csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_update.hpp create mode 100644 csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp create mode 100644 csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/gdn_fwd_h_epilogue_policies.hpp create mode 100644 csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/block/block_scheduler_gdn_fwd_h.hpp create mode 100644 csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/kernel/gdn_fwd_h_kernel.hpp create mode 100644 csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/chunk_gated_delta_rule_fwd_h_struct.h create mode 100644 csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla_preloadA_l1B.hpp create mode 100644 csrc/moe/dequant_situ_quant/CMakeLists.txt create mode 100644 csrc/moe/dequant_situ_quant/dequant_situ_quant_torch_adpt.h create mode 100644 csrc/moe/dequant_situ_quant/docs/aclnnDequantSituQuant.md create mode 100644 csrc/moe/dequant_situ_quant/op_graph/dequant_situ_quant_proto.h create mode 100644 csrc/moe/dequant_situ_quant/op_host/CMakeLists.txt create mode 100644 csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_def.cpp create mode 100644 csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_infershape.cpp create mode 100644 csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_tiling.cpp create mode 100644 csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_tiling.h create mode 100644 csrc/moe/dequant_situ_quant/op_kernel/dequant_situ_quant.cpp create mode 100644 csrc/moe/dequant_situ_quant/op_kernel/dequant_situ_quant.h create mode 100644 csrc/moe/situ_mx_quant/CMakeLists.txt create mode 100644 csrc/moe/situ_mx_quant/docs/aclnnSituMxQuant.md create mode 100644 csrc/moe/situ_mx_quant/op_host/CMakeLists.txt create mode 100644 csrc/moe/situ_mx_quant/op_host/arch35/situ_mx_quant_tiling_arch35.cpp create mode 100644 csrc/moe/situ_mx_quant/op_host/arch35/situ_mx_quant_tiling_arch35.h create mode 100644 csrc/moe/situ_mx_quant/op_host/config/ascend950/situ_mx_quant_binary.json create mode 100644 csrc/moe/situ_mx_quant/op_host/situ_mx_quant_def.cpp create mode 100644 csrc/moe/situ_mx_quant/op_host/situ_mx_quant_infershape.cpp create mode 100644 csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_axis_last.h create mode 100644 csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_common.h create mode 100644 csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_tiling_data.h create mode 100644 csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_tiling_key.h create mode 100644 csrc/moe/situ_mx_quant/op_kernel/inc/kernel_utils.h create mode 100644 csrc/moe/situ_mx_quant/op_kernel/inc/platform.h create mode 100644 csrc/moe/situ_mx_quant/op_kernel/situ_mx_quant_apt.cpp create mode 100644 csrc/moe/situ_mx_quant/situ_mx_quant_torch_adpt.h create mode 100644 tests/e2e/nightly/single_node/ops/singlecard_ops/test_chunk_gated_delta_rule_fwd_h_aclnn.py create mode 100644 tests/e2e/nightly/single_node/ops/singlecard_ops/test_chunk_kda_aclnn.py create mode 100644 tests/e2e/nightly/single_node/ops/singlecard_ops/test_dequant_situ_quant.py create mode 100644 tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_k3_situ_fusion_cases.py create mode 100644 tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_kda_ascendc_npu.py create mode 100644 tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_kda_recurrent_ascendc_npu.py diff --git a/csrc/attention/chunk_kda_fwd/CMakeLists.txt b/csrc/attention/chunk_kda_fwd/CMakeLists.txt new file mode 100644 index 000000000000..89dc500ecef7 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/CMakeLists.txt @@ -0,0 +1,16 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Tianjin University, Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# ----------------------------------------------------------------------------------------------------------- +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +if(NOT ENABLE_TEST AND NOT BENCHMARK) + list(REMOVE_ITEM CURRENT_DIRS tests) +endif() +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/csrc/attention/chunk_kda_fwd/op_host/CMakeLists.txt b/csrc/attention/chunk_kda_fwd/op_host/CMakeLists.txt new file mode 100644 index 000000000000..e6a6a541264d --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_host/CMakeLists.txt @@ -0,0 +1,33 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Tianjin University, Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# ----------------------------------------------------------------------------------------------------------- +set(CURRENT_CMAKE_DIR ${CMAKE_CURRENT_SOURCE_DIR}) +set(CATLASS_INCLUDE_DIR "${CMAKE_SOURCE_DIR}/third_party/catlass/include") +get_filename_component(CATLASS_INCLUDE_DIR_ABS ${CATLASS_INCLUDE_DIR} ABSOLUTE) + +add_op_to_compiled_list() +if (BUILD_OPEN_PROJECT) + target_sources(op_host_aclnnExc PRIVATE + chunk_kda_fwd_def.cpp + ) +endif() + +add_ops_compile_options( + OP_NAME ChunkKdaFwd + OPTIONS + --cce-auto-sync=off + -Wno-deprecated-declarations + -I${CATLASS_INCLUDE_DIR_ABS} +) + +if (NOT BUILD_OPS_RTY_KERNEL) + add_modules_sources(OPTYPE chunk_kda_fwd ACLNNTYPE aclnn_exclude) + target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE + ${CMAKE_CURRENT_SOURCE_DIR} + ${CATLASS_INCLUDE_DIR_ABS} + ) +endif() diff --git a/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_def.cpp b/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_def.cpp new file mode 100644 index 000000000000..90f6080f2c39 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_def.cpp @@ -0,0 +1,96 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#include "register/op_def_registry.h" + +namespace ops { +class ChunkKdaFwd : public OpDef { +public: + explicit ChunkKdaFwd(const char *name) : OpDef(name) + { + const std::initializer_list dataTypes = { + ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, + ge::DT_FLOAT16, ge::DT_BF16 + }; + const std::initializer_list stateTypes = { + ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, + ge::DT_FLOAT, ge::DT_FLOAT + }; + const std::initializer_list akkTypes = { + ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, + ge::DT_FLOAT16, ge::DT_BF16 + }; + const std::initializer_list outputDataTypes = { + ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT, + ge::DT_FLOAT16, ge::DT_BF16 + }; + const std::initializer_list formats = { + ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, + ge::FORMAT_ND, ge::FORMAT_ND + }; + + this->Input("q").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Input("k").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Input("v").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Input("gk").ParamType(REQUIRED) + .DataType(stateTypes) + .Format(formats).UnknownShapeFormat(formats); + this->Input("beta").ParamType(REQUIRED) + .DataType(stateTypes) + .Format(formats).UnknownShapeFormat(formats); + this->Input("initial_state").ParamType(OPTIONAL).DataType(stateTypes).Format(formats).UnknownShapeFormat(formats); + this->Input("cu_seqlens").ParamType(OPTIONAL).ValueDepend(OPTIONAL) + .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, + ge::DT_INT64, ge::DT_INT64}) + .Format(formats).UnknownShapeFormat(formats); + this->Input("chunk_indices").ParamType(OPTIONAL).ValueDepend(OPTIONAL) + .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, + ge::DT_INT64, ge::DT_INT64}) + .Format(formats).UnknownShapeFormat(formats); + this->Input("stage_qg").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Input("stage_aqk").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Input("stage_v_new").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Input("stage_h").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + + this->Output("o").ParamType(REQUIRED).DataType(outputDataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("final_state").ParamType(REQUIRED).DataType(stateTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("Aqk").ParamType(REQUIRED).DataType(stateTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("Akk").ParamType(REQUIRED).DataType(akkTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("w").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("u").ParamType(REQUIRED).DataType(outputDataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("qg").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("kg").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("v_new").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("h").ParamType(REQUIRED).DataType(stateTypes).Format(formats).UnknownShapeFormat(formats); + + this->Attr("scale").AttrType(REQUIRED).Float(1.0); + this->Attr("chunk_size").AttrType(REQUIRED).Int(64); + this->Attr("output_final_state").AttrType(REQUIRED).Bool(false); + this->Attr("total_chunks").AttrType(REQUIRED).Int(1); + this->Attr("stage").AttrType(OPTIONAL).Int(0); + + OpAICoreConfig aicoreConfig; + aicoreConfig.DynamicCompileStaticFlag(true) + .DynamicFormatFlag(true) + .DynamicRankSupportFlag(true) + .DynamicShapeSupportFlag(true) + .NeedCheckSupportFlag(false) + .PrecisionReduceFlag(true) + .ExtendCfgInfo("prebuildPattern.value", "Opaque") + .ExtendCfgInfo("coreType.value", "AiCore") + .ExtendCfgInfo("aclnnSupport.value", "support_aclnn"); + + this->AICore().AddConfig("ascend910b", aicoreConfig); + this->AICore().AddConfig("ascend910_93", aicoreConfig); + this->AICore().AddConfig("ascend950", aicoreConfig); + } +}; + +OP_ADD(ChunkKdaFwd); +} // namespace ops diff --git a/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.cpp b/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.cpp new file mode 100644 index 000000000000..43d1b8e9c633 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.cpp @@ -0,0 +1,193 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#include "chunk_kda_fwd_tiling.h" +#include +#include +#include +#include "tiling/platform/platform_ascendc.h" + +namespace optiling { +namespace { +constexpr size_t INPUT_Q_IDX = 0; +constexpr size_t INPUT_V_IDX = 2; +constexpr size_t INPUT_GK_IDX = 3; +constexpr size_t INPUT_INITIAL_IDX = 5; +constexpr size_t INPUT_CU_SEQLENS_IDX = 6; +constexpr size_t INPUT_CHUNK_INDICES_IDX = 7; +constexpr size_t ATTR_SCALE_IDX = 0; +constexpr size_t ATTR_CHUNK_SIZE_IDX = 1; +constexpr size_t ATTR_OUTPUT_FINAL_STATE_IDX = 2; +constexpr size_t ATTR_TOTAL_CHUNKS_IDX = 3; +constexpr size_t ATTR_STAGE_IDX = 4; +constexpr uint64_t KDA_SOLVE_SCRATCH_SLOTS = 5; +constexpr uint64_t KDA_SCORE_QUEUE_SLOTS = 2; +constexpr uint64_t KDA_SCORE_SCRATCH_PLANES = 3; +constexpr uint64_t KDA_FP32_BYTES = sizeof(float); +constexpr uint64_t KDA_WORKSPACE_ALIGN = 512; + +constexpr size_t DIM_B = 0; +constexpr size_t DIM_H = 1; +constexpr size_t DIM_T = 2; +constexpr size_t DIM_D = 3; + +int64_t DTypeCode(ge::DataType dtype) +{ + if (dtype == ge::DT_BF16) { + return 1; + } + if (dtype == ge::DT_FLOAT) { + return 2; + } + return 0; +} +} // namespace + +ge::graphStatus Tiling4ChunkKdaFwd(gert::TilingContext *context) +{ + ChunkKdaFwdTilingData tiling; + + auto qShape = context->GetOptionalInputShape(INPUT_Q_IDX)->GetStorageShape(); + auto vShape = context->GetOptionalInputShape(INPUT_V_IDX)->GetStorageShape(); + auto qDesc = context->GetInputDesc(INPUT_Q_IDX); + auto gDesc = context->GetInputDesc(INPUT_GK_IDX); + if (qDesc == nullptr || gDesc == nullptr) { + return ge::GRAPH_FAILED; + } + + auto attrPtr = context->GetAttrs(); + if (attrPtr == nullptr) { + return ge::GRAPH_FAILED; + } + float scale = static_cast(*(attrPtr->GetAttrPointer(ATTR_SCALE_IDX))); + int64_t chunkSize = *(attrPtr->GetAttrPointer(ATTR_CHUNK_SIZE_IDX)); + bool outputFinalState = *(attrPtr->GetAttrPointer(ATTR_OUTPUT_FINAL_STATE_IDX)); + int64_t totalChunks = *(attrPtr->GetAttrPointer(ATTR_TOTAL_CHUNKS_IDX)); + int64_t stage = 0; + const int64_t *stagePtr = attrPtr->GetAttrPointer(ATTR_STAGE_IDX); + if (stagePtr != nullptr) { + stage = *stagePtr; + } + + bool isVarLen = context->GetOptionalInputTensor(INPUT_CU_SEQLENS_IDX) != nullptr; + int64_t batch = qShape.GetDim(DIM_B); + int64_t seqNum = batch; + std::array seqStart{}; + std::array seqEnd{}; + std::array seqChunkOffset{}; + if (isVarLen) { + auto cuTensor = context->GetOptionalInputTensor(INPUT_CU_SEQLENS_IDX); + seqNum = cuTensor->GetStorageShape().GetDim(0) - 1; + auto chunkMetadata = context->GetOptionalInputTensor(INPUT_CHUNK_INDICES_IDX); + if (seqNum <= 0 || seqNum > KDA_MAX_TILING_SEQUENCES || chunkMetadata == nullptr || + chunkMetadata->GetStorageShape().GetShapeSize() != totalChunks * 4) { + return ge::GRAPH_FAILED; + } + const int64_t *cu = cuTensor->GetData(); + if (cu == nullptr) { + return ge::GRAPH_FAILED; + } + int64_t chunkOffset = 0; + for (int64_t seq = 0; seq < seqNum; ++seq) { + if (cu[seq] < 0 || cu[seq + 1] < cu[seq]) { + return ge::GRAPH_FAILED; + } + seqStart[seq] = cu[seq]; + seqEnd[seq] = cu[seq + 1]; + seqChunkOffset[seq] = chunkOffset; + const int64_t seqLength = cu[seq + 1] - cu[seq]; + chunkOffset += (seqLength + chunkSize - 1) / chunkSize; + } + seqChunkOffset[seqNum] = chunkOffset; + if (chunkOffset != totalChunks) { + return ge::GRAPH_FAILED; + } + } + bool hasInitialState = context->GetOptionalInputTensor(INPUT_INITIAL_IDX) != nullptr; + + const auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo()); + uint32_t coreNum = ascendcPlatform.GetCoreNumAic(); + int64_t taskNum = seqNum * vShape.GetDim(DIM_H); + if (stage == 1 || stage == 2 || stage == 3) { + taskNum = (isVarLen ? totalChunks : batch * totalChunks) * vShape.GetDim(DIM_H); + } + uint32_t blockDim = static_cast(std::min(taskNum, coreNum)); + if (stage == 1 || stage == 2 || stage == 3 || + (qDesc->GetDataType() != ge::DT_FLOAT && qShape.GetDim(DIM_D) >= 16)) { + blockDim = coreNum; + } + context->SetBlockDim(blockDim == 0 ? 1 : blockDim); + size_t *workspace = context->GetWorkspaceSizes(1); + uint64_t kernelScratch = 0; + if (stage == 1) { + const uint64_t usedCoreNum = static_cast(blockDim == 0 ? 1 : blockDim); + const uint64_t solveScratch = usedCoreNum * KDA_SOLVE_SCRATCH_SLOTS * + static_cast(chunkSize) * static_cast(chunkSize) * + KDA_FP32_BYTES; + const uint64_t alignedSolveScratch = + (solveScratch + KDA_WORKSPACE_ALIGN - 1) / KDA_WORKSPACE_ALIGN * KDA_WORKSPACE_ALIGN; + const uint64_t scoreElementBytes = qDesc->GetDataType() == ge::DT_FLOAT ? sizeof(float) : sizeof(uint16_t); + const uint64_t scoreScratch = usedCoreNum * KDA_SCORE_QUEUE_SLOTS * KDA_SCORE_SCRATCH_PLANES * + static_cast(chunkSize) * + static_cast(qShape.GetDim(DIM_D)) * scoreElementBytes; + kernelScratch = alignedSolveScratch + scoreScratch; + } else if (stage == 2) { + const uint64_t outputElements = static_cast(batch) * + static_cast(vShape.GetDim(DIM_H)) * + static_cast(qShape.GetDim(DIM_T)) * + static_cast(vShape.GetDim(DIM_D)); + kernelScratch = 2 * outputElements * KDA_FP32_BYTES; + } + kernelScratch = (kernelScratch + KDA_WORKSPACE_ALIGN - 1) / KDA_WORKSPACE_ALIGN * KDA_WORKSPACE_ALIGN; + workspace[0] = ascendcPlatform.GetLibApiWorkSpaceSize() + kernelScratch; + + tiling.set_batch(batch); + tiling.set_seqNum(seqNum); + tiling.set_qHeadNum(qShape.GetDim(DIM_H)); + tiling.set_vHeadNum(vShape.GetDim(DIM_H)); + tiling.set_seqlen(qShape.GetDim(DIM_T)); + tiling.set_kHeadDim(qShape.GetDim(DIM_D)); + tiling.set_vHeadDim(vShape.GetDim(DIM_D)); + tiling.set_chunkSize(chunkSize); + tiling.set_totalChunks(totalChunks); + tiling.set_scale(scale); + tiling.set_hasInitialState(hasInitialState); + tiling.set_outputFinalState(outputFinalState); + tiling.set_isVarLen(isVarLen); + tiling.set_dataType(DTypeCode(qDesc->GetDataType())); + tiling.set_gateDataType(DTypeCode(gDesc->GetDataType())); + tiling.set_usedCoreNum(blockDim == 0 ? 1 : blockDim); + tiling.set_stage(stage); + tiling.set_seqStart(seqStart.data()); + tiling.set_seqEnd(seqEnd.data()); + tiling.set_seqChunkOffset(seqChunkOffset.data()); + + if (qDesc->GetDataType() == ge::DT_FLOAT) { + context->SetTilingKey(0); + } else if (qShape.GetDim(DIM_D) < 16) { + context->SetTilingKey(2); + } else { + context->SetTilingKey(1); + } + tiling.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity()); + context->GetRawTilingData()->SetDataSize(tiling.GetDataSize()); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus TilingPrepare4ChunkKdaFwd(gert::TilingParseContext *context) +{ + (void)context; + return ge::GRAPH_SUCCESS; +} + +IMPL_OP_OPTILING(ChunkKdaFwd) + .Tiling(Tiling4ChunkKdaFwd) + .TilingParse(TilingPrepare4ChunkKdaFwd); + +} // namespace optiling diff --git a/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.h b/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.h new file mode 100644 index 000000000000..b22d2de71464 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.h @@ -0,0 +1,47 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#pragma once + +#include +#include +#include + +namespace optiling { + +constexpr int64_t KDA_MAX_TILING_SEQUENCES = 1024; +constexpr int64_t KDA_MAX_TILING_SEQUENCE_OFFSETS = KDA_MAX_TILING_SEQUENCES + 1; + +BEGIN_TILING_DATA_DEF(ChunkKdaFwdTilingData) +TILING_DATA_FIELD_DEF(int64_t, batch); +TILING_DATA_FIELD_DEF(int64_t, seqNum); +TILING_DATA_FIELD_DEF(int64_t, qHeadNum); +TILING_DATA_FIELD_DEF(int64_t, vHeadNum); +TILING_DATA_FIELD_DEF(int64_t, seqlen); +TILING_DATA_FIELD_DEF(int64_t, kHeadDim); +TILING_DATA_FIELD_DEF(int64_t, vHeadDim); +TILING_DATA_FIELD_DEF(int64_t, chunkSize); +TILING_DATA_FIELD_DEF(int64_t, totalChunks); +TILING_DATA_FIELD_DEF(float, scale); +TILING_DATA_FIELD_DEF(bool, hasInitialState); +TILING_DATA_FIELD_DEF(bool, outputFinalState); +TILING_DATA_FIELD_DEF(bool, isVarLen); +TILING_DATA_FIELD_DEF(int64_t, dataType); +TILING_DATA_FIELD_DEF(int64_t, gateDataType); +TILING_DATA_FIELD_DEF(int64_t, usedCoreNum); +TILING_DATA_FIELD_DEF(int64_t, stage); +TILING_DATA_FIELD_DEF_ARR(int64_t, KDA_MAX_TILING_SEQUENCES, seqStart); +TILING_DATA_FIELD_DEF_ARR(int64_t, KDA_MAX_TILING_SEQUENCES, seqEnd); +TILING_DATA_FIELD_DEF_ARR(int64_t, KDA_MAX_TILING_SEQUENCE_OFFSETS, seqChunkOffset); +END_TILING_DATA_DEF; + +REGISTER_TILING_DATA_CLASS(ChunkKdaFwd, ChunkKdaFwdTilingData) + +struct ChunkKdaFwdCompileInfo {}; +} // namespace optiling diff --git a/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.cpp b/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.cpp new file mode 100644 index 000000000000..7b387cb0a4eb --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.cpp @@ -0,0 +1,911 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#include "aclnn_chunk_kda_fwd.h" +#include "chunk_kda_fwd.h" +#include "../../../kda_layout_swap12/op_host/op_api/kda_layout_swap12.h" +#include "moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/chunk_gated_delta_rule_fwd_h.h" + +#include + +#include "acl/acl.h" +#include "aclnn/aclnn_base.h" +#include "aclnn_kernels/cast.h" +#include "aclnn_kernels/common/op_error_check.h" +#include "aclnn_kernels/contiguous.h" +#include "aclnn_kernels/reshape.h" +#include "opdev/make_op_executor.h" +#include "opdev/op_dfx.h" +#include "opdev/op_executor.h" +#include "opdev/op_log.h" +#include "opdev/tensor_view_utils.h" + +using namespace op; + +namespace l0op { +const aclTensor *Muls(const aclTensor *self, float alpha, aclOpExecutor *executor); +const aclTensor *ZerosLike(const aclTensor *self, aclOpExecutor *executor); +} + +#ifdef __cplusplus +extern "C" { +#endif + +namespace { +constexpr int64_t MAX_KDA_K_DIM = 256; +constexpr int64_t MAX_KDA_HEAD_NUM = 128; +constexpr int64_t MAX_KDA_VARLEN_SEQUENCES = 1024; + +struct ChunkKdaFwdParams { + const aclTensor *q = nullptr; + const aclTensor *k = nullptr; + const aclTensor *v = nullptr; + const aclTensor *gk = nullptr; + const aclTensor *beta = nullptr; + const aclTensor *initialStateOptional = nullptr; + const aclIntArray *cuSeqlensOptional = nullptr; + const aclIntArray *chunkIndicesOptional = nullptr; + const char *layout = "BSND"; + double scale = 1.0; + int64_t chunkSize = 64; + bool outputFinalState = false; + int64_t totalChunks = 1; + const aclTensor *oOut = nullptr; + const aclTensor *finalStateOut = nullptr; + const aclTensor *aqkOut = nullptr; + const aclTensor *akkOut = nullptr; + const aclTensor *wOut = nullptr; + const aclTensor *uOut = nullptr; + const aclTensor *qgOut = nullptr; + const aclTensor *kgOut = nullptr; + const aclTensor *vNewOut = nullptr; + const aclTensor *hOut = nullptr; +}; + +aclnnStatus KdaFwdDataContiguous(const aclTensor *&tensor, aclOpExecutor *executor) +{ + if (tensor == nullptr) { + return ACLNN_SUCCESS; + } + tensor = l0op::Contiguous(tensor, executor); + CHECK_RET(tensor != nullptr, ACLNN_ERR_INNER_NULLPTR); + return ACLNN_SUCCESS; +} + +op::Shape KdaFwdMakeShape(std::initializer_list dims) +{ + op::Shape shape; + for (int64_t dim : dims) { + shape.AppendDim(dim); + } + return shape; +} + +int64_t KdaFwdDim(const aclTensor *tensor, size_t idx) +{ + return tensor->GetViewShape().GetDim(idx); +} + +const aclTensor *KdaFwdMaybeCast(const aclTensor *tensor, DataType dataType, aclOpExecutor *executor) +{ + if (tensor == nullptr || tensor->GetDataType() == dataType) { + return tensor; + } + return l0op::Cast(tensor, dataType, executor); +} + +aclnnStatus KdaFwdViewCopyMaybeCast(const aclTensor *src, const aclTensor *dst, aclOpExecutor *executor) +{ + const aclTensor *castSrc = KdaFwdMaybeCast(src, dst->GetDataType(), executor); + CHECK_RET(castSrc != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::ViewCopy(castSrc, dst, executor) != nullptr, ACLNN_ERR_INNER_NULLPTR); + return ACLNN_SUCCESS; +} + +int64_t KdaFwdNumel(const aclTensor *tensor) +{ + const auto shape = tensor->GetViewShape(); + int64_t numel = 1; + for (size_t idx = 0; idx < shape.GetDimNum(); ++idx) { + numel *= shape.GetDim(idx); + } + return numel; +} + +aclnnStatus KdaFwdCopyMaybeCastAfter(const aclTensor *src, const aclTensor *dependency, + const aclTensor *dst, aclOpExecutor *executor) +{ + const aclTensor *castSrc = KdaFwdMaybeCast(src, dst->GetDataType(), executor); + CHECK_RET(castSrc != nullptr, ACLNN_ERR_INNER_NULLPTR); + // The split-forward intermediates already use the destination layout. Reuse + // the internal swap kernel as a dependency-ordered device copy, making its + // dim-1/dim-2 swap an identity by flattening both swapped dimensions to 1. + // This deliberately calls the l0op directly; the public aclnn swap shape + // contract applies to layout conversion, not to this internal copy barrier. + const aclTensor *linearSrc = l0op::Reshape(castSrc, KdaFwdMakeShape({1, 1, 1, KdaFwdNumel(castSrc)}), executor); + CHECK_RET(linearSrc != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(linearSrc, dependency, dst, executor)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + return ACLNN_SUCCESS; +} + +size_t KdaFwdRank(const aclTensor *tensor) +{ + return tensor->GetViewShape().GetDimNum(); +} + +aclnnStatus KdaFwdCheckCuSeqlens(const aclIntArray *cuSeqlensOptional, int64_t seqlen) +{ + if (cuSeqlensOptional == nullptr) { + return ACLNN_SUCCESS; + } + const aclIntArray &cu = *cuSeqlensOptional; + CHECK_COND(cu.Size() >= 2, ACLNN_ERR_PARAM_INVALID, + "cuSeqlensOptional must contain at least [0, total_tokens]."); + CHECK_COND(cu[0] == 0, ACLNN_ERR_PARAM_INVALID, "cuSeqlensOptional[0] must be 0."); + CHECK_COND(cu[cu.Size() - 1] == seqlen, ACLNN_ERR_PARAM_INVALID, + "cuSeqlensOptional last element must equal the sequence length."); + for (size_t idx = 0; idx + 1 < cu.Size(); ++idx) { + CHECK_COND(cu[idx] <= cu[idx + 1], ACLNN_ERR_PARAM_INVALID, + "cuSeqlensOptional must be nondecreasing."); + } + return ACLNN_SUCCESS; +} + +int64_t KdaFwdExpectedChunks(const aclIntArray *cuSeqlensOptional, int64_t seqlen, int64_t chunkSize) +{ + if (cuSeqlensOptional == nullptr) { + return (seqlen + chunkSize - 1) / chunkSize; + } + int64_t total = 0; + const aclIntArray &cu = *cuSeqlensOptional; + for (size_t idx = 0; idx + 1 < cu.Size(); ++idx) { + int64_t length = cu[idx + 1] - cu[idx]; + total += (length + chunkSize - 1) / chunkSize; + } + return total; +} + +aclnnStatus KdaFwdCheckChunkIndices(const aclIntArray *chunkIndicesOptional, + const aclIntArray *cuSeqlensOptional, + int64_t totalChunks, + int64_t expectedChunks, + int64_t chunkSize) +{ + CHECK_COND(totalChunks == expectedChunks, ACLNN_ERR_PARAM_INVALID, + "totalChunks must equal the number of chunks derived from sequence lengths and chunkSize."); + if (chunkIndicesOptional == nullptr) { + return ACLNN_SUCCESS; + } + CHECK_COND(cuSeqlensOptional != nullptr, ACLNN_ERR_PARAM_INVALID, + "chunkIndicesOptional is only valid when cuSeqlensOptional is provided."); + CHECK_COND(chunkIndicesOptional->Size() % 2 == 0, ACLNN_ERR_PARAM_INVALID, + "chunkIndicesOptional must contain (seq_id, chunk_id) pairs."); + CHECK_COND(static_cast(chunkIndicesOptional->Size() / 2) == expectedChunks, + ACLNN_ERR_PARAM_INVALID, + "chunkIndicesOptional must contain exactly totalChunks (seq_id, chunk_id) pairs."); + const aclIntArray &indices = *chunkIndicesOptional; + const aclIntArray &cu = *cuSeqlensOptional; + int64_t seqNum = static_cast(cu.Size()) - 1; + for (size_t idx = 0; idx < indices.Size(); idx += 2) { + int64_t seq = indices[idx]; + int64_t localChunk = indices[idx + 1]; + CHECK_COND(seq >= 0 && seq < seqNum, ACLNN_ERR_PARAM_INVALID, + "chunkIndicesOptional seq_id must be in [0, seq_num)."); + int64_t seqLength = cu[seq + 1] - cu[seq]; + int64_t seqChunks = (seqLength + chunkSize - 1) / chunkSize; + CHECK_COND(localChunk >= 0 && localChunk < seqChunks, ACLNN_ERR_PARAM_INVALID, + "chunkIndicesOptional chunk_id is outside the selected sequence."); + } + size_t expectedIdx = 0; + for (int64_t seq = 0; seq < seqNum; ++seq) { + int64_t seqLength = cu[seq + 1] - cu[seq]; + int64_t seqChunks = (seqLength + chunkSize - 1) / chunkSize; + for (int64_t localChunk = 0; localChunk < seqChunks; ++localChunk) { + CHECK_COND(indices[expectedIdx] == seq && indices[expectedIdx + 1] == localChunk, + ACLNN_ERR_PARAM_INVALID, + "chunkIndicesOptional must use canonical sequence-major chunk order."); + expectedIdx += 2; + } + } + return ACLNN_SUCCESS; +} + +int64_t KdaFwdSeqNum(int64_t batch, const aclIntArray *cuSeqlensOptional) +{ + if (cuSeqlensOptional == nullptr) { + return batch; + } + return static_cast(cuSeqlensOptional->Size()) - 1; +} + +aclnnStatus KdaFwdCheckStateShape(const aclTensor *state, const char *name, int64_t seqNum, int64_t hvNum, + int64_t kDim, int64_t vDim) +{ + if (state == nullptr) { + return ACLNN_SUCCESS; + } + const auto shape = state->GetViewShape(); + CHECK_COND(shape.GetDimNum() == 4 && shape.GetDim(0) == seqNum && shape.GetDim(1) == hvNum && + shape.GetDim(2) == kDim && shape.GetDim(3) == vDim, + ACLNN_ERR_PARAM_INVALID, + "%s must be [seq_num, HV, K, V], where seq_num is batch for dense input or " + "len(cuSeqlensOptional)-1 for varlen input.", + name); + return ACLNN_SUCCESS; +} + +enum class KdaFwdLayout { + BSND, + BNSD, + TND, + NTD, +}; + +bool KdaFwdSameShape(const aclTensor *lhs, const aclTensor *rhs) +{ + if (KdaFwdRank(lhs) != KdaFwdRank(rhs)) { + return false; + } + for (size_t idx = 0; idx < KdaFwdRank(lhs); ++idx) { + if (KdaFwdDim(lhs, idx) != KdaFwdDim(rhs, idx)) { + return false; + } + } + return true; +} + +aclnnStatus KdaFwdParseLayout(const char *layout, KdaFwdLayout &parsed) +{ + CHECK_COND(layout != nullptr, ACLNN_ERR_PARAM_INVALID, + "layout must not be nullptr and must be one of BSND, BNSD, TND, NTD."); + if (std::strcmp(layout, "BSND") == 0) { + parsed = KdaFwdLayout::BSND; + return ACLNN_SUCCESS; + } + if (std::strcmp(layout, "BNSD") == 0) { + parsed = KdaFwdLayout::BNSD; + return ACLNN_SUCCESS; + } + if (std::strcmp(layout, "TND") == 0) { + parsed = KdaFwdLayout::TND; + return ACLNN_SUCCESS; + } + if (std::strcmp(layout, "NTD") == 0) { + parsed = KdaFwdLayout::NTD; + return ACLNN_SUCCESS; + } + CHECK_COND(false, ACLNN_ERR_PARAM_INVALID, + "layout must be one of BSND, BNSD, TND, NTD and must be uppercase."); + return ACLNN_ERR_PARAM_INVALID; +} + +aclnnStatus KdaFwdCheckLayoutShape(const ChunkKdaFwdParams ¶ms, KdaFwdLayout layout) +{ + CHECK_COND(KdaFwdSameShape(params.q, params.k), ACLNN_ERR_PARAM_INVALID, + "q and k must have identical shape."); + if (layout == KdaFwdLayout::TND) { + CHECK_COND(KdaFwdRank(params.q) == 3 && KdaFwdRank(params.v) == 3 && + KdaFwdRank(params.gk) == 3 && KdaFwdRank(params.beta) == 2, + ACLNN_ERR_PARAM_INVALID, + "layout TND expects q/k [T,H,K], v [T,HV,V], gk [T,HV,K], beta [T,HV]."); + CHECK_COND(KdaFwdDim(params.v, 0) == KdaFwdDim(params.q, 0) && + KdaFwdDim(params.gk, 0) == KdaFwdDim(params.q, 0) && + KdaFwdDim(params.beta, 0) == KdaFwdDim(params.q, 0) && + KdaFwdDim(params.gk, 1) == KdaFwdDim(params.v, 1) && + KdaFwdDim(params.beta, 1) == KdaFwdDim(params.v, 1) && + KdaFwdDim(params.gk, 2) == KdaFwdDim(params.q, 2), + ACLNN_ERR_PARAM_INVALID, + "layout TND shape mismatch."); + } else if (layout == KdaFwdLayout::NTD) { + CHECK_COND(KdaFwdRank(params.q) == 3 && KdaFwdRank(params.v) == 3 && + KdaFwdRank(params.gk) == 3 && KdaFwdRank(params.beta) == 2, + ACLNN_ERR_PARAM_INVALID, + "layout NTD expects q/k [H,T,K], v [HV,T,V], gk [HV,T,K], beta [HV,T]."); + CHECK_COND(KdaFwdDim(params.v, 1) == KdaFwdDim(params.q, 1) && + KdaFwdDim(params.gk, 0) == KdaFwdDim(params.v, 0) && + KdaFwdDim(params.beta, 0) == KdaFwdDim(params.v, 0) && + KdaFwdDim(params.gk, 1) == KdaFwdDim(params.q, 1) && + KdaFwdDim(params.beta, 1) == KdaFwdDim(params.q, 1) && + KdaFwdDim(params.gk, 2) == KdaFwdDim(params.q, 2), + ACLNN_ERR_PARAM_INVALID, + "layout NTD shape mismatch."); + } else if (layout == KdaFwdLayout::BSND) { + CHECK_COND(KdaFwdRank(params.q) == 4 && KdaFwdRank(params.v) == 4 && + KdaFwdRank(params.gk) == 4 && KdaFwdRank(params.beta) == 3, + ACLNN_ERR_PARAM_INVALID, + "layout BSND expects q/k [B,T,H,K], v [B,T,HV,V], gk [B,T,HV,K], beta [B,T,HV]."); + CHECK_COND(KdaFwdDim(params.v, 0) == KdaFwdDim(params.q, 0) && + KdaFwdDim(params.v, 1) == KdaFwdDim(params.q, 1) && + KdaFwdDim(params.gk, 0) == KdaFwdDim(params.q, 0) && + KdaFwdDim(params.gk, 1) == KdaFwdDim(params.q, 1) && + KdaFwdDim(params.beta, 0) == KdaFwdDim(params.q, 0) && + KdaFwdDim(params.beta, 1) == KdaFwdDim(params.q, 1) && + KdaFwdDim(params.gk, 2) == KdaFwdDim(params.v, 2) && + KdaFwdDim(params.beta, 2) == KdaFwdDim(params.v, 2) && + KdaFwdDim(params.gk, 3) == KdaFwdDim(params.q, 3), + ACLNN_ERR_PARAM_INVALID, + "layout BSND shape mismatch."); + } else { + CHECK_COND(KdaFwdRank(params.q) == 4 && KdaFwdRank(params.v) == 4 && + KdaFwdRank(params.gk) == 4 && KdaFwdRank(params.beta) == 3, + ACLNN_ERR_PARAM_INVALID, + "layout BNSD expects q/k [B,H,T,K], v [B,HV,T,V], gk [B,HV,T,K], beta [B,HV,T]."); + CHECK_COND(KdaFwdDim(params.v, 0) == KdaFwdDim(params.q, 0) && + KdaFwdDim(params.v, 2) == KdaFwdDim(params.q, 2) && + KdaFwdDim(params.gk, 0) == KdaFwdDim(params.q, 0) && + KdaFwdDim(params.gk, 1) == KdaFwdDim(params.v, 1) && + KdaFwdDim(params.beta, 0) == KdaFwdDim(params.q, 0) && + KdaFwdDim(params.beta, 1) == KdaFwdDim(params.v, 1) && + KdaFwdDim(params.gk, 2) == KdaFwdDim(params.q, 2) && + KdaFwdDim(params.beta, 2) == KdaFwdDim(params.q, 2) && + KdaFwdDim(params.gk, 3) == KdaFwdDim(params.q, 3), + ACLNN_ERR_PARAM_INVALID, + "layout BNSD shape mismatch."); + } + return ACLNN_SUCCESS; +} + +aclnnStatus KdaFwdCheckParams(const ChunkKdaFwdParams ¶ms) +{ + CHECK_COND(params.q != nullptr, ACLNN_ERR_PARAM_NULLPTR, "q must not be nullptr."); + CHECK_COND(params.k != nullptr, ACLNN_ERR_PARAM_NULLPTR, "k must not be nullptr."); + CHECK_COND(params.v != nullptr, ACLNN_ERR_PARAM_NULLPTR, "v must not be nullptr."); + CHECK_COND(params.gk != nullptr, ACLNN_ERR_PARAM_NULLPTR, "gk must not be nullptr."); + CHECK_COND(params.beta != nullptr, ACLNN_ERR_PARAM_NULLPTR, "beta must not be nullptr."); + CHECK_COND(params.oOut != nullptr && params.finalStateOut != nullptr && params.aqkOut != nullptr && + params.akkOut != nullptr && params.wOut != nullptr && params.uOut != nullptr && + params.qgOut != nullptr && params.kgOut != nullptr && params.vNewOut != nullptr && + params.hOut != nullptr, + ACLNN_ERR_PARAM_NULLPTR, "ChunkKdaFwd outputs must not be nullptr."); + CHECK_COND(params.chunkSize > 0, ACLNN_ERR_PARAM_INVALID, "chunkSize must be positive."); + CHECK_COND(params.totalChunks > 0, ACLNN_ERR_PARAM_INVALID, "totalChunks must be positive."); + size_t qRank = KdaFwdRank(params.q); + size_t betaRank = KdaFwdRank(params.beta); + CHECK_COND((qRank == 4 && betaRank == 3) || (qRank == 3 && betaRank == 2), ACLNN_ERR_PARAM_INVALID, + "q/k/v/gk must be BSND/BNSD rank4 with beta rank3, or TND/NTD rank3 with beta rank2."); + size_t kDimIdx = (qRank == 4) ? 3 : 2; + CHECK_COND(params.q->GetViewShape().GetDim(kDimIdx) <= MAX_KDA_K_DIM, ACLNN_ERR_PARAM_INVALID, + "k head dimension must be less than or equal to 256."); + return ACLNN_SUCCESS; +} + +bool KdaFwdSplitCubePathSupported(const ChunkKdaFwdParams ¶ms, int64_t kDim, int64_t vDim) +{ + auto qDtype = params.q->GetDataType(); + auto kDtype = params.k->GetDataType(); + auto vDtype = params.v->GetDataType(); + bool dataDtypeSupported = (qDtype == DataType::DT_FLOAT16 || qDtype == DataType::DT_BF16) && + kDtype == qDtype && vDtype == qDtype; + return dataDtypeSupported && + (params.chunkSize == 64 || params.chunkSize == 128) && kDim >= 16 && vDim >= 16 && + kDim % 16 == 0 && vDim % 16 == 0 && vDim <= 256; +} + +aclnnStatus KdaFwdParamsDataContiguous(ChunkKdaFwdParams ¶ms, aclOpExecutor *executor) +{ + CHECK_RET(KdaFwdDataContiguous(params.q, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(KdaFwdDataContiguous(params.k, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(KdaFwdDataContiguous(params.v, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(KdaFwdDataContiguous(params.gk, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(KdaFwdDataContiguous(params.beta, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(KdaFwdDataContiguous(params.initialStateOptional, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + return ACLNN_SUCCESS; +} +} // namespace + +aclnnStatus aclnnChunkKdaFwdGetWorkspaceSize( + const aclTensor *q, + const aclTensor *k, + const aclTensor *v, + const aclTensor *gk, + const aclTensor *beta, + const aclTensor *initialStateOptional, + const aclIntArray *cuSeqlensOptional, + const aclIntArray *chunkIndicesOptional, + const char *layout, + double scale, + int64_t chunkSize, + bool outputFinalState, + int64_t totalChunks, + const aclTensor *oOut, + const aclTensor *finalStateOut, + const aclTensor *aqkOut, + const aclTensor *akkOut, + const aclTensor *wOut, + const aclTensor *uOut, + const aclTensor *qgOut, + const aclTensor *kgOut, + const aclTensor *vNewOut, + const aclTensor *hOut, + uint64_t *workspaceSize, + aclOpExecutor **executor) +{ + ChunkKdaFwdParams params{q, k, v, gk, beta, initialStateOptional, cuSeqlensOptional, chunkIndicesOptional, layout, + scale, chunkSize, outputFinalState, totalChunks, oOut, finalStateOut, aqkOut, + akkOut, wOut, uOut, qgOut, kgOut, vNewOut, hOut}; + L2_DFX_PHASE_1(aclnnChunkKdaFwd, + DFX_IN(q, k, v, gk, beta, initialStateOptional, cuSeqlensOptional, chunkIndicesOptional), + DFX_OUT(oOut, finalStateOut, aqkOut, akkOut, wOut, uOut, qgOut, kgOut, vNewOut, hOut)); + auto uniqueExecutor = CREATE_EXECUTOR(); + CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); + auto executorPtr = uniqueExecutor.get(); + CHECK_RET(KdaFwdCheckParams(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(KdaFwdParamsDataContiguous(params, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + + KdaFwdLayout parsedLayout = KdaFwdLayout::BSND; + CHECK_RET(KdaFwdParseLayout(params.layout, parsedLayout) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(KdaFwdCheckLayoutShape(params, parsedLayout) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + bool isTnd = parsedLayout == KdaFwdLayout::TND || parsedLayout == KdaFwdLayout::NTD; + bool isInternalLayout = parsedLayout == KdaFwdLayout::BNSD || parsedLayout == KdaFwdLayout::NTD; + int64_t batch = isTnd ? 1 : KdaFwdDim(params.q, 0); + int64_t seqlen = parsedLayout == KdaFwdLayout::TND ? KdaFwdDim(params.q, 0) : + (parsedLayout == KdaFwdLayout::NTD ? KdaFwdDim(params.q, 1) : + (parsedLayout == KdaFwdLayout::BNSD ? KdaFwdDim(params.q, 2) : KdaFwdDim(params.q, 1))); + int64_t hNum = parsedLayout == KdaFwdLayout::TND ? KdaFwdDim(params.q, 1) : + (parsedLayout == KdaFwdLayout::NTD ? KdaFwdDim(params.q, 0) : + (parsedLayout == KdaFwdLayout::BNSD ? KdaFwdDim(params.q, 1) : KdaFwdDim(params.q, 2))); + int64_t kDim = isTnd ? KdaFwdDim(params.q, 2) : KdaFwdDim(params.q, 3); + int64_t hvNum = parsedLayout == KdaFwdLayout::TND ? KdaFwdDim(params.v, 1) : + (parsedLayout == KdaFwdLayout::NTD ? KdaFwdDim(params.v, 0) : + (parsedLayout == KdaFwdLayout::BNSD ? KdaFwdDim(params.v, 1) : KdaFwdDim(params.v, 2))); + int64_t vDim = isTnd ? KdaFwdDim(params.v, 2) : KdaFwdDim(params.v, 3); + int64_t seqNum = KdaFwdSeqNum(batch, params.cuSeqlensOptional); + CHECK_COND(hNum <= MAX_KDA_HEAD_NUM && hvNum <= MAX_KDA_HEAD_NUM, ACLNN_ERR_PARAM_INVALID, + "H and HV must be less than or equal to 128."); + CHECK_COND(hNum > 0 && hvNum >= hNum && hvNum % hNum == 0, ACLNN_ERR_PARAM_INVALID, + "H and HV must be positive, HV must be greater than or equal to H, and HV must be divisible by H."); + CHECK_COND(parsedLayout != KdaFwdLayout::TND || hNum == 1, ACLNN_ERR_PARAM_INVALID, + "TND layout with H > 1 is not supported by npu_chunk_kda_fwd; use NTD [H,T,D] layout " + "for multi-head rank3 input."); + CHECK_RET(KdaFwdCheckCuSeqlens(params.cuSeqlensOptional, seqlen) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + int64_t expectedChunks = KdaFwdExpectedChunks(params.cuSeqlensOptional, seqlen, params.chunkSize); + CHECK_COND(params.cuSeqlensOptional == nullptr || seqNum <= MAX_KDA_VARLEN_SEQUENCES, + ACLNN_ERR_PARAM_INVALID, + "varlen input supports at most 1024 sequences in one call; split a larger request at sequence " + "boundaries."); + CHECK_RET(KdaFwdCheckChunkIndices(params.chunkIndicesOptional, params.cuSeqlensOptional, params.totalChunks, + expectedChunks, params.chunkSize) == ACLNN_SUCCESS, + ACLNN_ERR_PARAM_INVALID); + CHECK_COND(params.cuSeqlensOptional == nullptr || isTnd || batch == 1, ACLNN_ERR_PARAM_INVALID, + "rank4 varlen input with cuSeqlensOptional currently requires B=1."); + CHECK_RET(KdaFwdCheckStateShape(params.initialStateOptional, "initialStateOptional", seqNum, hvNum, kDim, vDim) == + ACLNN_SUCCESS, + ACLNN_ERR_PARAM_INVALID); + CHECK_RET(KdaFwdCheckStateShape(params.finalStateOut, "finalStateOut", seqNum, hvNum, kDim, vDim) == + ACLNN_SUCCESS, + ACLNN_ERR_PARAM_INVALID); + CHECK_COND(KdaFwdSplitCubePathSupported(params, kDim, vDim), ACLNN_ERR_PARAM_INVALID, + "npu_chunk_kda_fwd only supports the AscendC split cube/vector path: q/k/v dtype must be the same " + "fp16/bf16 type, chunkSize must be 64 or 128, K/V must be multiples of 16, and V must be <= 256."); + + const aclTensor *qBsnd = params.q; + const aclTensor *kBsnd = params.k; + const aclTensor *vBsnd = params.v; + const aclTensor *gkBsnd = params.gk; + const aclTensor *betaBsn = params.beta; + if (parsedLayout == KdaFwdLayout::TND) { + qBsnd = l0op::Reshape(params.q, KdaFwdMakeShape({1, seqlen, hNum, kDim}), executorPtr); + kBsnd = l0op::Reshape(params.k, KdaFwdMakeShape({1, seqlen, hNum, kDim}), executorPtr); + vBsnd = l0op::Reshape(params.v, KdaFwdMakeShape({1, seqlen, hvNum, vDim}), executorPtr); + gkBsnd = l0op::Reshape(params.gk, KdaFwdMakeShape({1, seqlen, hvNum, kDim}), executorPtr); + betaBsn = l0op::Reshape(params.beta, KdaFwdMakeShape({1, seqlen, hvNum}), executorPtr); + CHECK_RET(qBsnd != nullptr && kBsnd != nullptr && vBsnd != nullptr && gkBsnd != nullptr && betaBsn != nullptr, + ACLNN_ERR_INNER_NULLPTR); + } else if (parsedLayout == KdaFwdLayout::NTD) { + qBsnd = l0op::Reshape(params.q, KdaFwdMakeShape({1, hNum, seqlen, kDim}), executorPtr); + kBsnd = l0op::Reshape(params.k, KdaFwdMakeShape({1, hNum, seqlen, kDim}), executorPtr); + vBsnd = l0op::Reshape(params.v, KdaFwdMakeShape({1, hvNum, seqlen, vDim}), executorPtr); + gkBsnd = l0op::Reshape(params.gk, KdaFwdMakeShape({1, hvNum, seqlen, kDim}), executorPtr); + betaBsn = l0op::Reshape(params.beta, KdaFwdMakeShape({1, hvNum, seqlen}), executorPtr); + CHECK_RET(qBsnd != nullptr && kBsnd != nullptr && vBsnd != nullptr && gkBsnd != nullptr && betaBsn != nullptr, + ACLNN_ERR_INNER_NULLPTR); + } + + const aclTensor *qBnsd = isInternalLayout ? qBsnd : + executorPtr->AllocTensor(KdaFwdMakeShape({batch, hNum, seqlen, kDim}), + params.q->GetDataType(), Format::FORMAT_ND); + const aclTensor *kBnsd = isInternalLayout ? kBsnd : + executorPtr->AllocTensor(KdaFwdMakeShape({batch, hNum, seqlen, kDim}), + params.k->GetDataType(), Format::FORMAT_ND); + const aclTensor *vBnsd = isInternalLayout ? vBsnd : + executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, vDim}), + params.v->GetDataType(), Format::FORMAT_ND); + const aclTensor *gkBnsdRaw = isInternalLayout ? gkBsnd : + executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), + params.gk->GetDataType(), Format::FORMAT_ND); + const aclTensor *betaBnsRaw = isInternalLayout ? betaBsn : + executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen}), + params.beta->GetDataType(), Format::FORMAT_ND); + bool returnIntermediates = KdaFwdNumel(params.aqkOut) != 0; + const aclTensor *oBnsd = nullptr; + const aclTensor *aqkBnst = nullptr; + const aclTensor *akkBnst = nullptr; + const aclTensor *wBnsd = nullptr; + const aclTensor *uBnsd = nullptr; + const aclTensor *qgBnsd = nullptr; + const aclTensor *kgBnsd = nullptr; + const aclTensor *vNewBnsd = nullptr; + const aclTensor *hBnst = nullptr; + if (isInternalLayout) { + if (parsedLayout == KdaFwdLayout::NTD) { + oBnsd = l0op::Reshape(params.oOut, KdaFwdMakeShape({1, hvNum, seqlen, vDim}), executorPtr); + if (returnIntermediates) { + aqkBnst = l0op::Reshape(params.aqkOut, KdaFwdMakeShape({1, hvNum, seqlen, params.chunkSize}), + executorPtr); + akkBnst = l0op::Reshape(params.akkOut, KdaFwdMakeShape({1, hvNum, seqlen, params.chunkSize}), + executorPtr); + wBnsd = l0op::Reshape(params.wOut, KdaFwdMakeShape({1, hvNum, seqlen, kDim}), executorPtr); + uBnsd = l0op::Reshape(params.uOut, KdaFwdMakeShape({1, hvNum, seqlen, vDim}), executorPtr); + qgBnsd = l0op::Reshape(params.qgOut, KdaFwdMakeShape({1, hvNum, seqlen, kDim}), executorPtr); + kgBnsd = l0op::Reshape(params.kgOut, KdaFwdMakeShape({1, hvNum, seqlen, kDim}), executorPtr); + vNewBnsd = l0op::Reshape(params.vNewOut, KdaFwdMakeShape({1, hvNum, seqlen, vDim}), executorPtr); + hBnst = l0op::Reshape(params.hOut, KdaFwdMakeShape({1, hvNum, params.totalChunks, kDim, vDim}), + executorPtr); + } + } else { + oBnsd = params.oOut; + if (returnIntermediates) { + aqkBnst = params.aqkOut; + akkBnst = params.akkOut; + wBnsd = params.wOut; + uBnsd = params.uOut; + qgBnsd = params.qgOut; + kgBnsd = params.kgOut; + vNewBnsd = params.vNewOut; + hBnst = params.hOut; + } + } + } else { + oBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, vDim}), + params.oOut->GetDataType(), Format::FORMAT_ND); + } + if (!isInternalLayout) { + wBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), + params.wOut->GetDataType(), Format::FORMAT_ND); + uBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, vDim}), + params.uOut->GetDataType(), Format::FORMAT_ND); + qgBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), + params.qgOut->GetDataType(), Format::FORMAT_ND); + kgBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), + params.kgOut->GetDataType(), Format::FORMAT_ND); + vNewBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, vDim}), + params.vNewOut->GetDataType(), Format::FORMAT_ND); + hBnst = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, params.totalChunks, kDim, vDim}), + params.hOut->GetDataType(), Format::FORMAT_ND); + } + const bool internalIntermediateOutputsReady = !isInternalLayout || !returnIntermediates || + (aqkBnst != nullptr && akkBnst != nullptr && wBnsd != nullptr && uBnsd != nullptr && + qgBnsd != nullptr && kgBnsd != nullptr && vNewBnsd != nullptr && hBnst != nullptr); + const bool externalComputeBuffersReady = isInternalLayout || + (wBnsd != nullptr && uBnsd != nullptr && qgBnsd != nullptr && kgBnsd != nullptr && + vNewBnsd != nullptr && hBnst != nullptr); + CHECK_RET(qBnsd != nullptr && kBnsd != nullptr && vBnsd != nullptr && gkBnsdRaw != nullptr && + betaBnsRaw != nullptr && oBnsd != nullptr && internalIntermediateOutputsReady && + externalComputeBuffersReady, + ACLNN_ERR_INNER_NULLPTR); + + if (!isInternalLayout) { + CHECK_RET(l0op::KdaLayoutSwap12(qBsnd, qBnsd, executorPtr)[0] != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(kBsnd, kBnsd, executorPtr)[0] != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(vBsnd, vBnsd, executorPtr)[0] != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(gkBsnd, gkBnsdRaw, executorPtr)[0] != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(betaBsn, betaBnsRaw, executorPtr)[0] != nullptr, ACLNN_ERR_INNER_NULLPTR); + } + + const aclTensor *gkBnsd = gkBnsdRaw; + const aclTensor *betaBns = betaBnsRaw; + if (gkBnsd->GetDataType() != DataType::DT_FLOAT) { + gkBnsd = l0op::Cast(gkBnsd, DataType::DT_FLOAT, executorPtr); + CHECK_RET(gkBnsd != nullptr, ACLNN_ERR_INNER_NULLPTR); + } + if (betaBns->GetDataType() != DataType::DT_FLOAT) { + betaBns = l0op::Cast(betaBns, DataType::DT_FLOAT, executorPtr); + CHECK_RET(betaBns != nullptr, ACLNN_ERR_INNER_NULLPTR); + } + + std::array result; + bool useSplitForward = true; + const aclTensor *aqkComputeBnst = aqkBnst; + const aclTensor *akkComputeBnst = akkBnst; + const aclTensor *wComputeBnsd = wBnsd; + const aclTensor *uComputeBnsd = uBnsd; + const aclTensor *qgComputeBnsd = qgBnsd; + const aclTensor *kgComputeBnsd = kgBnsd; + const aclTensor *vNewComputeBnsd = vNewBnsd; + const aclTensor *hComputeBnst = hBnst; + const aclTensor *oOutComputeBnsd = nullptr; + const aclTensor *wPreComputeBnsd = nullptr; + const aclTensor *kgScratchComputeBnsd = nullptr; + if (useSplitForward && isInternalLayout) { + aqkComputeBnst = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, params.chunkSize}), + DataType::DT_FLOAT, Format::FORMAT_ND); + akkComputeBnst = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, params.chunkSize}), + DataType::DT_FLOAT, Format::FORMAT_ND); + wComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), + params.wOut->GetDataType(), Format::FORMAT_ND); + uComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, vDim}), + params.uOut->GetDataType(), Format::FORMAT_ND); + qgComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), + params.qgOut->GetDataType(), Format::FORMAT_ND); + kgComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), + params.kgOut->GetDataType(), Format::FORMAT_ND); + vNewComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, vDim}), + params.vNewOut->GetDataType(), Format::FORMAT_ND); + hComputeBnst = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, params.totalChunks, kDim, vDim}), + params.hOut->GetDataType(), Format::FORMAT_ND); + CHECK_RET(aqkComputeBnst != nullptr && akkComputeBnst != nullptr && wComputeBnsd != nullptr && + uComputeBnsd != nullptr && qgComputeBnsd != nullptr && + kgComputeBnsd != nullptr && vNewComputeBnsd != nullptr && hComputeBnst != nullptr, + ACLNN_ERR_INNER_NULLPTR); + } + if (useSplitForward) { + if (!isInternalLayout) { + aqkComputeBnst = executorPtr->AllocTensor( + KdaFwdMakeShape({batch, hvNum, seqlen, params.chunkSize}), + DataType::DT_FLOAT, Format::FORMAT_ND); + akkComputeBnst = executorPtr->AllocTensor( + KdaFwdMakeShape({batch, hvNum, seqlen, params.chunkSize}), + DataType::DT_FLOAT, Format::FORMAT_ND); + CHECK_RET(aqkComputeBnst != nullptr && akkComputeBnst != nullptr, ACLNN_ERR_INNER_NULLPTR); + } + oOutComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, Format::FORMAT_ND); + CHECK_RET(oOutComputeBnsd != nullptr, ACLNN_ERR_INNER_NULLPTR); + wPreComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), + params.wOut->GetDataType(), Format::FORMAT_ND); + kgScratchComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), + params.kgOut->GetDataType(), Format::FORMAT_ND); + auto stage1ODummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.oOut->GetDataType(), + Format::FORMAT_ND); + auto stage1FinalStateDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, + Format::FORMAT_ND); + auto stage1UDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.uOut->GetDataType(), + Format::FORMAT_ND); + auto stage1HDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, Format::FORMAT_ND); + CHECK_RET(wPreComputeBnsd != nullptr && kgScratchComputeBnsd != nullptr && stage1ODummy != nullptr && + stage1FinalStateDummy != nullptr && stage1UDummy != nullptr && stage1HDummy != nullptr, + ACLNN_ERR_INNER_NULLPTR); + auto prepResult = l0op::ChunkKdaFwd(qBnsd, kBnsd, vBnsd, gkBnsd, betaBns, params.initialStateOptional, + params.cuSeqlensOptional, params.chunkIndicesOptional, nullptr, nullptr, + nullptr, nullptr, params.scale, params.chunkSize, params.outputFinalState, + params.totalChunks, 1, stage1ODummy, stage1FinalStateDummy, + aqkComputeBnst, akkComputeBnst, wPreComputeBnsd, stage1UDummy, + qgComputeBnsd, kgScratchComputeBnsd, + vNewComputeBnsd, stage1HDummy, executorPtr); + for (auto tensor : prepResult) { + CHECK_RET(tensor != nullptr, ACLNN_ERR_PARAM_NULLPTR); + } + const aclTensor *aqkScaledBnst = + l0op::Muls(aqkComputeBnst, static_cast(params.scale), executorPtr); + CHECK_RET(aqkScaledBnst != nullptr, ACLNN_ERR_INNER_NULLPTR); + const aclTensor *aqkForOutBnst = KdaFwdMaybeCast(aqkScaledBnst, qBnsd->GetDataType(), executorPtr); + CHECK_RET(aqkForOutBnst != nullptr, ACLNN_ERR_INNER_NULLPTR); + const aclTensor *qgScaledBnsd = + l0op::Muls(qgComputeBnsd, static_cast(params.scale), executorPtr); + CHECK_RET(qgScaledBnsd != nullptr, ACLNN_ERR_INNER_NULLPTR); + + auto wScratchBntd = executorPtr->AllocTensor( + KdaFwdMakeShape({batch, hvNum, params.totalChunks, params.chunkSize, kDim}), + DataType::DT_FLOAT, Format::FORMAT_ND); + const aclTensor *akkPostBnst = KdaFwdMaybeCast(akkComputeBnst, qBnsd->GetDataType(), executorPtr); + CHECK_RET(wScratchBntd != nullptr && akkPostBnst != nullptr, ACLNN_ERR_INNER_NULLPTR); + + auto stage3ODummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.oOut->GetDataType(), + Format::FORMAT_ND); + auto stage3FinalStateDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, + Format::FORMAT_ND); + auto stage3AqkDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, + Format::FORMAT_ND); + auto stage3AkkDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), qBnsd->GetDataType(), + Format::FORMAT_ND); + auto stage3QGDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.qgOut->GetDataType(), + Format::FORMAT_ND); + auto stage3VNewDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.vNewOut->GetDataType(), + Format::FORMAT_ND); + CHECK_RET(stage3ODummy != nullptr && stage3FinalStateDummy != nullptr && stage3AqkDummy != nullptr && + stage3AkkDummy != nullptr && stage3QGDummy != nullptr && stage3VNewDummy != nullptr, + ACLNN_ERR_INNER_NULLPTR); + + auto postResult = l0op::ChunkKdaFwd( + qBnsd, kBnsd, vBnsd, gkBnsd, betaBns, params.initialStateOptional, + params.cuSeqlensOptional, params.chunkIndicesOptional, wPreComputeBnsd, akkPostBnst, + vNewComputeBnsd, nullptr, params.scale, params.chunkSize, params.outputFinalState, + params.totalChunks, 3, stage3ODummy, stage3FinalStateDummy, stage3AqkDummy, stage3AkkDummy, + wComputeBnsd, uComputeBnsd, stage3QGDummy, kgComputeBnsd, stage3VNewDummy, wScratchBntd, + executorPtr); + for (auto tensor : postResult) { + CHECK_RET(tensor != nullptr, ACLNN_ERR_PARAM_NULLPTR); + } + + const aclTensor *neutralGForH = l0op::ZerosLike(betaBns, executorPtr); + CHECK_RET(neutralGForH != nullptr, ACLNN_ERR_INNER_NULLPTR); + auto hResult = l0op::ChunkGatedDeltaRuleFwdH( + kgComputeBnsd, wComputeBnsd, uComputeBnsd, neutralGForH, gkBnsd, + params.initialStateOptional, params.cuSeqlensOptional, params.chunkIndicesOptional, + params.outputFinalState, params.chunkSize, hComputeBnst, vNewComputeBnsd, + params.finalStateOut, executorPtr); + for (auto tensor : hResult) { + CHECK_RET(tensor != nullptr, ACLNN_ERR_PARAM_NULLPTR); + } + + auto oLocalDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, + Format::FORMAT_ND); + auto stage2FinalStateDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, + Format::FORMAT_ND); + auto stage2AqkDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, + Format::FORMAT_ND); + auto stage2AkkDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, + Format::FORMAT_ND); + auto stage2WDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.wOut->GetDataType(), + Format::FORMAT_ND); + auto stage2QGDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.qgOut->GetDataType(), + Format::FORMAT_ND); + auto stage2KGDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.kgOut->GetDataType(), + Format::FORMAT_ND); + auto stage2HDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, + Format::FORMAT_ND); + CHECK_RET(oLocalDummy != nullptr && stage2FinalStateDummy != nullptr && + stage2AqkDummy != nullptr && stage2AkkDummy != nullptr && stage2WDummy != nullptr && + stage2QGDummy != nullptr && stage2KGDummy != nullptr && stage2HDummy != nullptr, + ACLNN_ERR_INNER_NULLPTR); + auto outResult = l0op::ChunkKdaFwd( + qBnsd, kBnsd, vBnsd, gkBnsd, betaBns, params.initialStateOptional, + params.cuSeqlensOptional, params.chunkIndicesOptional, qgScaledBnsd, aqkForOutBnst, + vNewComputeBnsd, hComputeBnst, params.scale, params.chunkSize, false, params.totalChunks, 2, + oOutComputeBnsd, stage2FinalStateDummy, stage2AqkDummy, stage2AkkDummy, stage2WDummy, + oLocalDummy, stage2QGDummy, stage2KGDummy, oBnsd, stage2HDummy, executorPtr); + for (auto tensor : outResult) { + CHECK_RET(tensor != nullptr, ACLNN_ERR_PARAM_NULLPTR); + } + result = {oBnsd, params.finalStateOut, aqkScaledBnst, akkComputeBnst, wComputeBnsd, + uComputeBnsd, qgComputeBnsd, kgComputeBnsd, vNewComputeBnsd, hComputeBnst}; + } + for (auto tensor : result) { + CHECK_RET(tensor != nullptr, ACLNN_ERR_PARAM_NULLPTR); + } + if (isInternalLayout) { + if (useSplitForward) { + if (result[0] != oBnsd) { + CHECK_RET(KdaFwdViewCopyMaybeCast(result[0], oBnsd, executorPtr) == ACLNN_SUCCESS, + ACLNN_ERR_INNER_NULLPTR); + } + if (returnIntermediates) { + CHECK_RET(KdaFwdCopyMaybeCastAfter(result[2], oBnsd, aqkBnst, executorPtr) == ACLNN_SUCCESS, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(KdaFwdCopyMaybeCastAfter(result[3], aqkBnst, akkBnst, executorPtr) == ACLNN_SUCCESS, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(KdaFwdCopyMaybeCastAfter(result[4], akkBnst, wBnsd, executorPtr) == ACLNN_SUCCESS, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(KdaFwdCopyMaybeCastAfter(result[5], wBnsd, uBnsd, executorPtr) == ACLNN_SUCCESS, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(KdaFwdCopyMaybeCastAfter(result[6], uBnsd, qgBnsd, executorPtr) == ACLNN_SUCCESS, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(KdaFwdCopyMaybeCastAfter(result[7], qgBnsd, kgBnsd, executorPtr) == ACLNN_SUCCESS, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(KdaFwdCopyMaybeCastAfter(result[8], kgBnsd, vNewBnsd, executorPtr) == ACLNN_SUCCESS, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(KdaFwdCopyMaybeCastAfter(result[9], vNewBnsd, hBnst, executorPtr) == ACLNN_SUCCESS, + ACLNN_ERR_INNER_NULLPTR); + } + } + } else if (isTnd) { + auto oBsnd = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, vDim}), + params.oOut->GetDataType(), Format::FORMAT_ND); + CHECK_RET(oBsnd != nullptr, ACLNN_ERR_INNER_NULLPTR); + const aclTensor *oForLayout = KdaFwdMaybeCast(result[0], params.oOut->GetDataType(), executorPtr); + CHECK_RET(oForLayout != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(oForLayout, nullptr, oBsnd, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::ViewCopy(l0op::Reshape(oBsnd, KdaFwdMakeShape({seqlen, hvNum, vDim}), executorPtr), + params.oOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); + if (returnIntermediates) { + const aclTensor *aqkForLayout = KdaFwdMaybeCast(result[2], params.aqkOut->GetDataType(), executorPtr); + CHECK_RET(aqkForLayout != nullptr, ACLNN_ERR_INNER_NULLPTR); + const aclTensor *akkForLayout = KdaFwdMaybeCast(result[3], params.akkOut->GetDataType(), executorPtr); + CHECK_RET(akkForLayout != nullptr, ACLNN_ERR_INNER_NULLPTR); + auto aqkBsnt = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, params.chunkSize}), + params.aqkOut->GetDataType(), Format::FORMAT_ND); + auto akkBsnt = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, params.chunkSize}), + params.akkOut->GetDataType(), Format::FORMAT_ND); + auto wBsnd = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, kDim}), + params.wOut->GetDataType(), Format::FORMAT_ND); + auto uBsnd = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, vDim}), + params.uOut->GetDataType(), Format::FORMAT_ND); + auto qgBsnd = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, kDim}), + params.qgOut->GetDataType(), Format::FORMAT_ND); + auto kgBsnd = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, kDim}), + params.kgOut->GetDataType(), Format::FORMAT_ND); + auto vNewBsnd = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, vDim}), + params.vNewOut->GetDataType(), Format::FORMAT_ND); + auto hBsnt = executorPtr->AllocTensor(KdaFwdMakeShape({1, params.totalChunks, hvNum, kDim, vDim}), + params.hOut->GetDataType(), Format::FORMAT_ND); + CHECK_RET(aqkBsnt != nullptr && akkBsnt != nullptr && wBsnd != nullptr && uBsnd != nullptr && + qgBsnd != nullptr && kgBsnd != nullptr && vNewBsnd != nullptr && hBsnt != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(aqkForLayout, oBsnd, aqkBsnt, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(akkForLayout, aqkBsnt, akkBsnt, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(result[4], akkBsnt, wBsnd, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(result[5], wBsnd, uBsnd, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(result[6], uBsnd, qgBsnd, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(result[7], qgBsnd, kgBsnd, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(result[8], kgBsnd, vNewBsnd, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(result[9], vNewBsnd, hBsnt, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::ViewCopy(l0op::Reshape(aqkBsnt, KdaFwdMakeShape({seqlen, hvNum, params.chunkSize}), + executorPtr), + params.aqkOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::ViewCopy(l0op::Reshape(akkBsnt, KdaFwdMakeShape({seqlen, hvNum, params.chunkSize}), + executorPtr), + params.akkOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::ViewCopy(l0op::Reshape(wBsnd, KdaFwdMakeShape({seqlen, hvNum, kDim}), executorPtr), + params.wOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::ViewCopy(l0op::Reshape(uBsnd, KdaFwdMakeShape({seqlen, hvNum, vDim}), executorPtr), + params.uOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::ViewCopy(l0op::Reshape(qgBsnd, KdaFwdMakeShape({seqlen, hvNum, kDim}), executorPtr), + params.qgOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::ViewCopy(l0op::Reshape(kgBsnd, KdaFwdMakeShape({seqlen, hvNum, kDim}), executorPtr), + params.kgOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::ViewCopy(l0op::Reshape(vNewBsnd, KdaFwdMakeShape({seqlen, hvNum, vDim}), executorPtr), + params.vNewOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::ViewCopy(l0op::Reshape(hBsnt, KdaFwdMakeShape({params.totalChunks, hvNum, kDim, vDim}), + executorPtr), + params.hOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); + } + } else { + const aclTensor *oBsnd = params.oOut; + const aclTensor *oForLayout = KdaFwdMaybeCast(result[0], params.oOut->GetDataType(), executorPtr); + CHECK_RET(oForLayout != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(oForLayout, nullptr, oBsnd, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + if (returnIntermediates) { + const aclTensor *aqkForLayout = KdaFwdMaybeCast(result[2], params.aqkOut->GetDataType(), executorPtr); + CHECK_RET(aqkForLayout != nullptr, ACLNN_ERR_INNER_NULLPTR); + const aclTensor *akkForLayout = KdaFwdMaybeCast(result[3], params.akkOut->GetDataType(), executorPtr); + CHECK_RET(akkForLayout != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(aqkForLayout, oBsnd, params.aqkOut, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(akkForLayout, params.aqkOut, params.akkOut, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(result[4], params.akkOut, params.wOut, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(result[5], params.wOut, params.uOut, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(result[6], params.uOut, params.qgOut, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(result[7], params.qgOut, params.kgOut, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(result[8], params.kgOut, params.vNewOut, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::KdaLayoutSwap12(result[9], params.vNewOut, params.hOut, executorPtr)[0] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + } + } + *workspaceSize = uniqueExecutor->GetWorkspaceSize(); + uniqueExecutor.ReleaseTo(executor); + return ACLNN_SUCCESS; +} + +aclnnStatus aclnnChunkKdaFwd(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream) +{ + L2_DFX_PHASE_2(aclnnChunkKdaFwd); + CHECK_COND(CommonOpExecutorRun(workspace, workspaceSize, executor, stream) == ACLNN_SUCCESS, ACLNN_ERR_INNER, + "ChunkKdaFwd launch failed."); + return ACLNN_SUCCESS; +} + +#ifdef __cplusplus +} +#endif diff --git a/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.h b/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.h new file mode 100644 index 000000000000..63301d16d227 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.h @@ -0,0 +1,52 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ +#ifndef OP_API_INC_ACLNN_CHUNK_KDA_FWD_H +#define OP_API_INC_ACLNN_CHUNK_KDA_FWD_H + +#include "aclnn/aclnn_base.h" +#include "aclnn_util.h" + +#ifdef __cplusplus +extern "C" { +#endif + +aclnnStatus aclnnChunkKdaFwdGetWorkspaceSize( + const aclTensor *q, + const aclTensor *k, + const aclTensor *v, + const aclTensor *gk, + const aclTensor *beta, + const aclTensor *initialStateOptional, + const aclIntArray *cuSeqlensOptional, + const aclIntArray *chunkIndicesOptional, + const char *layout, + double scale, + int64_t chunkSize, + bool outputFinalState, + int64_t totalChunks, + const aclTensor *oOut, + const aclTensor *finalStateOut, + const aclTensor *aqkOut, + const aclTensor *akkOut, + const aclTensor *wOut, + const aclTensor *uOut, + const aclTensor *qgOut, + const aclTensor *kgOut, + const aclTensor *vNewOut, + const aclTensor *hOut, + uint64_t *workspaceSize, + aclOpExecutor **executor); + +aclnnStatus aclnnChunkKdaFwd(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream); + +#ifdef __cplusplus +} +#endif + +#endif diff --git a/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.cpp b/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.cpp new file mode 100644 index 000000000000..999c8e2b367b --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.cpp @@ -0,0 +1,156 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#include "chunk_kda_fwd.h" + +#include "opdev/make_op_executor.h" +#include "opdev/op_dfx.h" +#include "opdev/op_log.h" + +#include + +using namespace op; + +namespace l0op { +OP_TYPE_REGISTER(ChunkKdaFwd); + +namespace { +const aclIntArray *BuildPackedChunkMetadata(const aclIntArray *cuSeqlens, + const aclIntArray *chunkIndices, + int64_t chunkSize, + int64_t totalChunks, + aclOpExecutor *executor) +{ + if (cuSeqlens == nullptr || cuSeqlens->Size() < 2 || chunkSize <= 0 || totalChunks <= 0) { + return nullptr; + } + + const aclIntArray &cu = *cuSeqlens; + std::vector packed; + packed.reserve(static_cast(totalChunks) * 4); + auto appendChunk = [&](int64_t seq, int64_t localChunk) -> bool { + if (seq < 0 || static_cast(seq + 1) >= cu.Size() || localChunk < 0) { + return false; + } + int64_t seqStart = cu[static_cast(seq)]; + int64_t seqEnd = cu[static_cast(seq + 1)]; + int64_t start = seqStart + localChunk * chunkSize; + if (start < seqStart || start >= seqEnd) { + return false; + } + int64_t end = start + chunkSize; + if (end > seqEnd) { + end = seqEnd; + } + packed.insert(packed.end(), {seq, start, end, 0}); + return true; + }; + + if (chunkIndices != nullptr) { + if (chunkIndices->Size() != static_cast(totalChunks) * 2) { + return nullptr; + } + for (size_t idx = 0; idx < chunkIndices->Size(); idx += 2) { + if (!appendChunk((*chunkIndices)[idx], (*chunkIndices)[idx + 1])) { + return nullptr; + } + } + } else { + for (size_t seq = 0; seq + 1 < cu.Size(); ++seq) { + int64_t seqLength = cu[seq + 1] - cu[seq]; + int64_t chunkCount = (seqLength + chunkSize - 1) / chunkSize; + for (int64_t localChunk = 0; localChunk < chunkCount; ++localChunk) { + if (!appendChunk(static_cast(seq), localChunk)) { + return nullptr; + } + } + } + } + if (packed.size() != static_cast(totalChunks) * 4) { + return nullptr; + } + return executor->AllocIntArray(packed.data(), packed.size()); +} +} // namespace + +const std::array ChunkKdaFwd( + const aclTensor *q, + const aclTensor *k, + const aclTensor *v, + const aclTensor *gk, + const aclTensor *beta, + const aclTensor *initialStateOptional, + const aclIntArray *cuSeqlensOptional, + const aclIntArray *chunkIndicesOptional, + const aclTensor *stageQGInputOptional, + const aclTensor *stageAqkInputOptional, + const aclTensor *stageVNewInputOptional, + const aclTensor *stageHInputOptional, + double scale, + int64_t chunkSize, + bool outputFinalState, + int64_t totalChunks, + int64_t stage, + const aclTensor *oOut, + const aclTensor *finalStateOut, + const aclTensor *aqkOut, + const aclTensor *akkOut, + const aclTensor *wOut, + const aclTensor *uOut, + const aclTensor *qgOut, + const aclTensor *kgOut, + const aclTensor *vNewOut, + const aclTensor *hOut, + aclOpExecutor *executor) +{ + L0_DFX(ChunkKdaFwd, q, k, v, gk, beta, initialStateOptional, cuSeqlensOptional, chunkIndicesOptional, + stageQGInputOptional, stageAqkInputOptional, stageVNewInputOptional, stageHInputOptional, scale, chunkSize, + outputFinalState, totalChunks, stage, oOut, finalStateOut, aqkOut, akkOut, wOut, uOut, qgOut, kgOut, + vNewOut, hOut); + + const aclTensor *actualCuSeqlens = nullptr; + if (cuSeqlensOptional != nullptr) { + actualCuSeqlens = executor->ConvertToTensor(cuSeqlensOptional, DataType::DT_INT64); + const_cast(actualCuSeqlens)->SetStorageFormat(Format::FORMAT_ND); + const_cast(actualCuSeqlens)->SetViewFormat(Format::FORMAT_ND); + const_cast(actualCuSeqlens)->SetOriginalFormat(Format::FORMAT_ND); + } + + const aclTensor *actualChunkIndices = nullptr; + if (cuSeqlensOptional != nullptr) { + const aclIntArray *packedChunkMetadata = BuildPackedChunkMetadata( + cuSeqlensOptional, chunkIndicesOptional, chunkSize, totalChunks, executor); + if (packedChunkMetadata == nullptr) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "failed to build packed chunk metadata."); + return {nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr}; + } + actualChunkIndices = executor->ConvertToTensor(packedChunkMetadata, DataType::DT_INT64); + if (actualChunkIndices == nullptr) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "failed to convert packed chunk metadata to tensor."); + return {nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr}; + } + const_cast(actualChunkIndices)->SetStorageFormat(Format::FORMAT_ND); + const_cast(actualChunkIndices)->SetViewFormat(Format::FORMAT_ND); + const_cast(actualChunkIndices)->SetOriginalFormat(Format::FORMAT_ND); + } + + auto ret = ADD_TO_LAUNCHER_LIST_AICORE( + ChunkKdaFwd, + OP_INPUT(q, k, v, gk, beta, initialStateOptional, actualCuSeqlens, actualChunkIndices, + stageQGInputOptional, stageAqkInputOptional, stageVNewInputOptional, stageHInputOptional), + OP_OUTPUT(oOut, finalStateOut, aqkOut, akkOut, wOut, uOut, qgOut, kgOut, vNewOut, hOut), + OP_ATTR(scale, chunkSize, outputFinalState, totalChunks, stage)); + if (ret != ACLNN_SUCCESS) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "ADD_TO_LAUNCHER_LIST_AICORE ChunkKdaFwd failed."); + return {nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr}; + } + return {oOut, finalStateOut, aqkOut, akkOut, wOut, uOut, qgOut, kgOut, vNewOut, hOut}; +} + +} // namespace l0op diff --git a/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.h b/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.h new file mode 100644 index 000000000000..b606715c96a8 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.h @@ -0,0 +1,47 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ +#ifndef OP_API_INC_LEVEL0_OP_CHUNK_KDA_FWD_OP_H +#define OP_API_INC_LEVEL0_OP_CHUNK_KDA_FWD_OP_H + +#include +#include "opdev/op_executor.h" + +namespace l0op { +const std::array ChunkKdaFwd( + const aclTensor *q, + const aclTensor *k, + const aclTensor *v, + const aclTensor *gk, + const aclTensor *beta, + const aclTensor *initialStateOptional, + const aclIntArray *cuSeqlensOptional, + const aclIntArray *chunkIndicesOptional, + const aclTensor *stageQGInputOptional, + const aclTensor *stageAqkInputOptional, + const aclTensor *stageVNewInputOptional, + const aclTensor *stageHInputOptional, + double scale, + int64_t chunkSize, + bool outputFinalState, + int64_t totalChunks, + int64_t stage, + const aclTensor *oOut, + const aclTensor *finalStateOut, + const aclTensor *aqkOut, + const aclTensor *akkOut, + const aclTensor *wOut, + const aclTensor *uOut, + const aclTensor *qgOut, + const aclTensor *kgOut, + const aclTensor *vNewOut, + const aclTensor *hOut, + aclOpExecutor *executor); +} + +#endif diff --git a/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd.cpp b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd.cpp new file mode 100644 index 000000000000..05ff5f4e607f --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd.cpp @@ -0,0 +1,2584 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#include "kernel_operator.h" + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +#define CATLASS_ARCH 3510 +#include "catlass/arch/arch.hpp" +#include "catlass/catlass.hpp" +#include "catlass/gemm/block/block_mmad.hpp" +#include "catlass/gemm/dispatch_policy.hpp" +#include "catlass/gemm/gemm_type.hpp" +#include "catlass/gemm_coord.hpp" +#include "catlass/layout/layout.hpp" +#include "catlass/arch/cross_core_sync.hpp" +#include "tla/layout.hpp" +#include "tla/tensor.hpp" +using _128 = tla::Int<128>; +#else +#define CATLASS_ARCH 2201 +#include "catlass/arch/arch.hpp" +#include "catlass/catlass.hpp" +#include "catlass/gemm/block/block_mmad.hpp" +#include "catlass/gemm/dispatch_policy.hpp" +#include "catlass/gemm/gemm_type.hpp" +#include "catlass/gemm_coord.hpp" +#include "catlass/layout/layout.hpp" +#include "catlass/arch/cross_core_sync.hpp" +#include "tla/layout.hpp" +#include "tla/tensor.hpp" +#endif + +#ifndef TORCH_MODE +#include "lib/matmul_intf.h" +#endif + +using namespace AscendC; +using _64 = tla::Int<64>; + +namespace { +constexpr float LN2 = 0.69314718055994530942f; +constexpr float KDA_EXP2_CLAMP = 80.0f; +constexpr float KDA_EXP_INPUT_MAX = KDA_EXP2_CLAMP * LN2; +constexpr float KDA_EXP_INPUT_MIN = -KDA_EXP2_CLAMP * LN2; +constexpr float KDA_FP16_MAX = 65504.0f; +constexpr uint32_t EXP2_UB_ELEMENTS = 256; +constexpr uint32_t EXP2_EVENT_ID = 0; +constexpr uint32_t KDA_MTE2_V_EVENT_ID = 1; +constexpr uint32_t KDA_SCALAR_V_MTE3_EVENT_ID = 4; +constexpr uint32_t KDA_SCALAR_MTE3_V_EVENT_ID = 5; +constexpr uint32_t KDA_MTE2_MTE3_EVENT_ID = 6; +constexpr uint32_t KDA_MTE3_MTE2_EVENT_ID = 7; +constexpr uint32_t KDA_VEC_BUFFER_NUM = 2; +constexpr uint32_t KDA_SOLVE_BT = 64; +constexpr uint32_t KDA_SOLVE_MATRIX_ELEMENTS = KDA_SOLVE_BT * KDA_SOLVE_BT; +constexpr uint32_t KDA_SOLVE_SCRATCH_X = 0; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y0 = 1; +constexpr uint32_t KDA_SOLVE_SCRATCH_TMP = 2; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y1 = 3; +constexpr uint32_t KDA_SOLVE_SCRATCH_IDENTITY = 4; +constexpr uint32_t KDA_SOLVE_SCRATCH_SLOTS = 5; +constexpr uint32_t KDA_SOLVE_DIAG_BT = 16; +constexpr uint32_t KDA_SOLVE_DIAG_BLOCKS = KDA_SOLVE_BT / KDA_SOLVE_DIAG_BT; +constexpr uint32_t KDA_SOLVE_DIAG_MCH_ITERS = 3; +constexpr uint32_t KDA_SCORE_REF_BC = 16; +constexpr uint32_t KDA_VEC_ARENA_ELEMENTS = 32768; +constexpr uint32_t KDA_BITS_PER_MASK_BYTE = 8; +constexpr uint32_t KDA_SELECT_COL_BLOCKS = 2; +constexpr uint32_t KDA_SELECT_COL_MASK_BYTES = KDA_SOLVE_MATRIX_ELEMENTS / KDA_BITS_PER_MASK_BYTE; +constexpr uint32_t KDA_SELECT_MASK_BYTES = KDA_SELECT_COL_BLOCKS * KDA_SELECT_COL_MASK_BYTES; +constexpr uint32_t KDA_SELECT_AQK_MASK_BYTE_OFFSET = 120 * 1024; +constexpr uint32_t KDA_SELECT_AKK_MASK_BYTE_OFFSET = KDA_SELECT_AQK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_BYTE_OFFSET = KDA_SELECT_AKK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_FLOAT_OFFSET = KDA_SELECT_ZERO_BYTE_OFFSET / sizeof(float); +constexpr uint8_t KDA_SCORE_DONE_FLAG0 = 2; +constexpr uint8_t KDA_SCORE_DONE_FLAG1 = 3; +constexpr uint8_t KDA_SCORE_READY_FLAG0 = 4; +constexpr uint8_t KDA_SCORE_READY_FLAG1 = 5; +constexpr uint32_t KDA_SCORE_QUEUE_DEPTH = 2; +constexpr uint32_t KDA_SYNC_REVERSE_DEPTH = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_PLANES = 3; +constexpr uint32_t KDA_SCORE_SCRATCH_QG = 0; +constexpr uint32_t KDA_SCORE_SCRATCH_W = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_KG = 2; +constexpr uint64_t KDA_WORKSPACE_ALIGN = 512; + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +using KdaArchTag = Catlass::Arch::Ascend950; +#else +using KdaArchTag = Catlass::Arch::AtlasA2; +#endif +using KdaDispatchPolicy = Catlass::Gemm::MmadPingpong; +using KdaSolveDispatchPolicy = Catlass::Gemm::MmadPingpong; +static_assert(!KdaSolveDispatchPolicy::USE_HF32_MODE, "KDA triangular solve must use IEEE FP32 Cube mode"); +using KdaL1TileShape = tla::Shape<_64, _128, _128>; +using KdaL0TileShape = KdaL1TileShape; +using KdaSolveL1TileShape = tla::Shape<_64, _64, _64>; +using KdaSolveL0TileShape = KdaSolveL1TileShape; + +__aicore__ inline uint32_t FloatToBits(float value) +{ + union Bits { + __aicore__ Bits() {} + float f; + uint32_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline float BitsToFloat(uint32_t value) +{ + union Bits { + __aicore__ Bits() {} + uint32_t u; + float f; + } bits; + bits.u = value; + return bits.f; +} + +__aicore__ inline uint16_t Bf16ToBits(bfloat16_t value) +{ + union Bits { + __aicore__ Bits() {} + bfloat16_t f; + uint16_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline bfloat16_t BitsToBf16(uint16_t value) +{ + union Bits { + __aicore__ Bits() {} + uint16_t u; + bfloat16_t f; + } bits; + bits.u = value; + return bits.f; +} + +template +__aicore__ inline T FloatToType(float value) +{ + if constexpr (IsSameType::value) { + uint32_t bits = FloatToBits(value); + uint32_t bias = 0x7FFFu + ((bits >> 16) & 1u); + return BitsToBf16(static_cast((bits + bias) >> 16)); + } + return static_cast(value); +} + +template +class ChunkKdaFwdKernel { +public: + __aicore__ inline void Init(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR stageQG, GM_ADDR stageAqk, + GM_ADDR stageVNew, GM_ADDR stageH, GM_ADDR o, GM_ADDR finalState, GM_ADDR aqk, + GM_ADDR akk, GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR workspace, const ChunkKdaFwdTilingData &tiling, TPipe *pipe, + bool initVecBuffers = true) + { + pipe_ = pipe; + q_.SetGlobalBuffer((__gm__ T *)q); + k_.SetGlobalBuffer((__gm__ T *)k); + v_.SetGlobalBuffer((__gm__ T *)v); + gk_.SetGlobalBuffer((__gm__ float *)gk); + beta_.SetGlobalBuffer((__gm__ float *)beta); + if (initialState != nullptr) { + initialState_.SetGlobalBuffer((__gm__ float *)initialState); + } + if (cuSeqlens != nullptr) { + cuSeqlens_.SetGlobalBuffer((__gm__ int64_t *)cuSeqlens); + } + if (stageQG != nullptr) { + stageQG_.SetGlobalBuffer((__gm__ T *)stageQG); + } + if (stageAqk != nullptr) { + stageAqk_.SetGlobalBuffer((__gm__ T *)stageAqk); + } + if (stageVNew != nullptr) { + stageVNew_.SetGlobalBuffer((__gm__ T *)stageVNew); + } + if (stageH != nullptr) { + stageH_.SetGlobalBuffer((__gm__ T *)stageH); + } + hasChunkIndices_ = chunkIndices != nullptr; + o_.SetGlobalBuffer((__gm__ OUT_T *)o); + finalState_.SetGlobalBuffer((__gm__ float *)finalState); + aqk_.SetGlobalBuffer((__gm__ float *)aqk); + akk_.SetGlobalBuffer((__gm__ AKK_T *)akk); + w_.SetGlobalBuffer((__gm__ T *)w); + u_.SetGlobalBuffer((__gm__ OUT_T *)u); + qg_.SetGlobalBuffer((__gm__ T *)qg); + kg_.SetGlobalBuffer((__gm__ T *)kg); + vNew_.SetGlobalBuffer((__gm__ T *)vNew); + h_.SetGlobalBuffer((__gm__ float *)h); + solveWorkspace_.SetGlobalBuffer((__gm__ float *)workspace); + + B_ = tiling.batch; + N_ = tiling.seqNum; + H_ = tiling.qHeadNum; + HV_ = tiling.vHeadNum; + T_ = tiling.seqlen; + K_ = tiling.kHeadDim; + V_ = tiling.vHeadDim; + BT_ = tiling.chunkSize; + NT_ = tiling.totalChunks; + scale_ = tiling.scale; + hasInitial_ = tiling.hasInitialState; + isVarLen_ = tiling.isVarLen; + usedCoreNum_ = tiling.usedCoreNum; + stage_ = tiling.stage; + if (stage_ == 1) { + const uint64_t solveBytes = usedCoreNum_ * KDA_SOLVE_SCRATCH_SLOTS * BT_ * BT_ * sizeof(float); + const uint64_t alignedSolveBytes = + (solveBytes + KDA_WORKSPACE_ALIGN - 1) / KDA_WORKSPACE_ALIGN * KDA_WORKSPACE_ALIGN; + scoreWorkspace_.SetGlobalBuffer((__gm__ T *)(workspace + alignedSolveBytes)); + } + if (stage_ == 2) { + const uint64_t outputElements = B_ * HV_ * T_ * V_; + o_.SetGlobalBuffer((__gm__ OUT_T *)workspace); + u_.SetGlobalBuffer((__gm__ OUT_T *)workspace + outputElements); + } + if ASCEND_IS_AIV { + uint64_t subBlockNum = static_cast(GetSubBlockNum()); + solveCoreIdx_ = subBlockNum == 0 ? 0 : static_cast(GetBlockIdx()) / subBlockNum; + } else { + solveCoreIdx_ = static_cast(GetBlockIdx()); + } + seqStart_ = tiling.seqStart; + seqEnd_ = tiling.seqEnd; + seqChunkOffset_ = tiling.seqChunkOffset; + + if (pipe_ != nullptr && initVecBuffers) { + pipe_->InitBuffer(exp2Buf_, EXP2_UB_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(vecBuf_, KDA_VEC_ARENA_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(qInQue_, KDA_VEC_BUFFER_NUM, EXP2_UB_ELEMENTS * sizeof(T)); + pipe_->InitBuffer(kInQue_, KDA_VEC_BUFFER_NUM, EXP2_UB_ELEMENTS * sizeof(T)); + pipe_->InitBuffer(gInQue_, KDA_VEC_BUFFER_NUM, EXP2_UB_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(qgOutQue_, KDA_VEC_BUFFER_NUM, EXP2_UB_ELEMENTS * sizeof(T)); + pipe_->InitBuffer(wOutQue_, KDA_VEC_BUFFER_NUM, EXP2_UB_ELEMENTS * sizeof(T)); + pipe_->InitBuffer(kgOutQue_, KDA_VEC_BUFFER_NUM, EXP2_UB_ELEMENTS * sizeof(T)); + } + } + + __aicore__ inline void ProcessAivOnly() + { + if (stage_ == 1) { + isAivOnly_ = true; + ProcessPreAiv(); + return; + } + if (stage_ == 2) { + isAivOnly_ = true; + ProcessOutAiv(); + return; + } + if (stage_ == 3) { + isAivOnly_ = true; + ProcessPostAiv(); + return; + } + return; + } + + __aicore__ inline void ProcessAiv() + { + if (stage_ == 1) { + ProcessPreAiv(); + return; + } + if (stage_ == 2) { + ProcessOutAiv(); + return; + } + if (stage_ == 3) { + ProcessPostAiv(); + return; + } + return; + } + + __aicore__ inline void ProcessAic() + { + if (stage_ == 1) { + ProcessPreAic(); + return; + } + if (stage_ == 2) { + ProcessOutAic(); + return; + } + if (stage_ == 3) { + ProcessPostAic(); + return; + } + return; + } + +private: + __aicore__ inline uint64_t QOffset(uint64_t b, uint64_t h, uint64_t t, uint64_t d) const + { + return ((b * H_ + h) * T_ + t) * K_ + d; + } + + __aicore__ inline uint64_t KVOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t d, uint64_t dim) const + { + return ((b * HV_ + hv) * T_ + t) * dim + d; + } + + __aicore__ inline uint64_t BetaOffset(uint64_t b, uint64_t hv, uint64_t t) const + { + return (b * HV_ + hv) * T_ + t; + } + + __aicore__ inline uint64_t AOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t j) const + { + return ((b * HV_ + hv) * T_ + t) * BT_ + j; + } + + __aicore__ inline uint64_t HOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t d, uint64_t r) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * K_ + d) * V_ + r; + } + + __aicore__ inline uint64_t WScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t t, uint64_t d) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * BT_ + t) * K_ + d; + } + + __aicore__ inline uint64_t SolveScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t slot) const + { + (void)b; + (void)hv; + (void)chunkIdx; + uint64_t matrixElements = BT_ * BT_; + return solveCoreIdx_ * KDA_SOLVE_SCRATCH_SLOTS * matrixElements + slot * matrixElements; + } + + __aicore__ inline uint64_t ScoreScratchOffset(uint64_t slot, uint64_t plane, uint64_t t = 0, + uint64_t d = 0) const + { + return (((solveCoreIdx_ * KDA_SCORE_QUEUE_DEPTH + slot) * KDA_SCORE_SCRATCH_PLANES + plane) * BT_ + t) * + K_ + + d; + } + + + + __aicore__ inline uint64_t ScoreRefBlockSize() const + { + if constexpr (IsSameType::value) { + return 2; + } + return KDA_SCORE_REF_BC; + } + + __aicore__ inline uint64_t ScoreRowBlockCount(uint64_t curT, uint64_t rowBegin) const + { + uint64_t blockSize = ScoreRefBlockSize(); + uint64_t rowCount = curT - rowBegin; + if (rowCount > blockSize) { + rowCount = blockSize; + } + return rowCount; + } + + __aicore__ inline uint64_t ScoreRefToken(uint64_t start, uint64_t curT, uint64_t rowBegin, + uint64_t rowCount) const + { + uint64_t ref = rowBegin + rowCount / 2; + if (ref >= curT) { + ref = curT - 1; + } + return start + ref; + } + + __aicore__ inline void RunExp2(LocalTensor &tensor, uint32_t count) + { + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + ClampExpInput(tensor, count); + Exp(tensor, tensor, count); + PipeBarrier(); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + } + + __aicore__ inline void ClampExpInput(LocalTensor &tensor, uint32_t count) + { + Mins(tensor, tensor, KDA_EXP_INPUT_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, KDA_EXP_INPUT_MIN, count); + PipeBarrier(); + } + + __aicore__ inline void ClampFp32ToOutputType(LocalTensor &tensor, uint32_t count) + { + if constexpr (IsSameType::value) { + Mins(tensor, tensor, KDA_FP16_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, -KDA_FP16_MAX, count); + PipeBarrier(); + } + } + + template + __aicore__ inline void CopyVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst, src[offset], static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + + template + __aicore__ inline void CopyVectorOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst[offset], src, static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPad(dst[offset], src, params); + } + + template + __aicore__ inline void CopyRowIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset) + { + CopyVectorIn(dst, src, offset, K_); + } + + template + __aicore__ inline void CopyRowOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src) + { + CopyVectorOut(dst, offset, src, K_); + } + + __aicore__ inline LocalTensor VecScratch(uint64_t slot) + { + return vecBuf_.Get()[slot * EXP2_UB_ELEMENTS]; + } + + template + __aicore__ inline void LoadAsFloatRow(GlobalTensor &src, uint64_t srcOffset, LocalTensor &dst, + uint64_t count) + { + if constexpr (IsSameType::value) { + CopyVectorIn(dst, src, srcOffset, count); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + Adds(dst, dst, 0.0f, static_cast(count)); + PipeBarrier(); + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + } else { + LocalTensor rowLocal = exp2Buf_.Get(); + CopyVectorIn(rowLocal, src, srcOffset, count); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + Cast(dst, rowLocal, RoundMode::CAST_NONE, static_cast(count)); + PipeBarrier(); + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + } + PipeBarrier(); + } + + template + __aicore__ inline void StoreFloatRow(GlobalTensor &dst, uint64_t dstOffset, LocalTensor &src, + uint64_t count) + { + if constexpr (IsSameType::value) { + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + CopyVectorOut(dst, dstOffset, src, count); + } else { + LocalTensor rowLocal = exp2Buf_.Get(); + Cast(rowLocal, src, RoundMode::CAST_RINT, static_cast(count)); + PipeBarrier(); + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + CopyVectorOut(dst, dstOffset, rowLocal, count); + } + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + } + + + + + + __aicore__ inline LocalTensor Exp2NegG(uint64_t b, uint64_t hv, uint64_t t) + { + LocalTensor exp2Local = exp2Buf_.Get(); + CopyRowIn(exp2Local, gk_, KVOffset(b, hv, t, 0, K_)); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + Muls(exp2Local, exp2Local, -LN2, static_cast(K_)); + PipeBarrier(); + RunExp2(exp2Local, static_cast(K_)); + return exp2Local; + } + + + __aicore__ inline uint64_t GateProductToken(uint64_t start, uint64_t logicalIdx, uint64_t subBlockIdx, + uint64_t subBlockNum) const + { + return start + subBlockIdx + logicalIdx * subBlockNum; + } + + __aicore__ inline void LoadGateProductRow(uint64_t b, uint64_t h, uint64_t hv, uint64_t ti) + { + LocalTensor qLocal = qInQue_.AllocTensor(); + LocalTensor kLocal = kInQue_.AllocTensor(); + LocalTensor gLocal = gInQue_.AllocTensor(); + CopyRowIn(qLocal, q_, QOffset(b, h, ti, 0)); + CopyRowIn(kLocal, k_, QOffset(b, h, ti, 0)); + CopyRowIn(gLocal, gk_, KVOffset(b, hv, ti, 0, K_)); + qInQue_.EnQue(qLocal); + kInQue_.EnQue(kLocal); + gInQue_.EnQue(gLocal); + } + + __aicore__ inline void StoreGateProductRow(uint64_t b, uint64_t hv, uint64_t ti) + { + LocalTensor qPosLocal = qgOutQue_.DeQue(); + LocalTensor kPosLocal = wOutQue_.DeQue(); + LocalTensor kNegLocal = kgOutQue_.DeQue(); + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + CopyRowOut(qg_, KVOffset(b, hv, ti, 0, K_), qPosLocal); + CopyRowOut(w_, KVOffset(b, hv, ti, 0, K_), kPosLocal); + CopyRowOut(kg_, KVOffset(b, hv, ti, 0, K_), kNegLocal); + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + qgOutQue_.FreeTensor(qPosLocal); + wOutQue_.FreeTensor(kPosLocal); + kgOutQue_.FreeTensor(kNegLocal); + } + + __aicore__ inline void ComputeGateProductRow(LocalTensor &qFp32, LocalTensor &kFp32, + LocalTensor &gFp32, LocalTensor &refFp32, + LocalTensor &expFp32, LocalTensor &outFp32, + bool useRef, bool zeroKg) + { + LocalTensor qPosLocal = qgOutQue_.AllocTensor(); + LocalTensor kPosLocal = wOutQue_.AllocTensor(); + LocalTensor kNegLocal = kgOutQue_.AllocTensor(); + + if (useRef) { + Sub(expFp32, gFp32, refFp32, static_cast(K_)); + } else { + Adds(expFp32, gFp32, 0.0f, static_cast(K_)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(K_)); + PipeBarrier(); + ClampExpInput(expFp32, static_cast(K_)); + Exp(expFp32, expFp32, static_cast(K_)); + PipeBarrier(); + + Mul(outFp32, qFp32, expFp32, static_cast(K_)); + PipeBarrier(); + ClampFp32ToOutputType(outFp32, static_cast(K_)); + if constexpr (IsSameType::value) { + DataCopy(qPosLocal, outFp32, static_cast(K_)); + } else { + Cast(qPosLocal, outFp32, RoundMode::CAST_RINT, static_cast(K_)); + } + PipeBarrier(); + + Mul(outFp32, kFp32, expFp32, static_cast(K_)); + PipeBarrier(); + ClampFp32ToOutputType(outFp32, static_cast(K_)); + if constexpr (IsSameType::value) { + DataCopy(kPosLocal, outFp32, static_cast(K_)); + } else { + Cast(kPosLocal, outFp32, RoundMode::CAST_RINT, static_cast(K_)); + } + PipeBarrier(); + + if (zeroKg) { + Duplicate(outFp32, 0.0f, static_cast(K_)); + PipeBarrier(); + } else { + if (useRef) { + Sub(expFp32, refFp32, gFp32, static_cast(K_)); + } else { + Muls(expFp32, gFp32, -1.0f, static_cast(K_)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(K_)); + PipeBarrier(); + ClampExpInput(expFp32, static_cast(K_)); + Exp(expFp32, expFp32, static_cast(K_)); + PipeBarrier(); + Mul(outFp32, kFp32, expFp32, static_cast(K_)); + PipeBarrier(); + } + ClampFp32ToOutputType(outFp32, static_cast(K_)); + if constexpr (IsSameType::value) { + DataCopy(kNegLocal, outFp32, static_cast(K_)); + } else { + Cast(kNegLocal, outFp32, RoundMode::CAST_RINT, static_cast(K_)); + } + + qgOutQue_.EnQue(qPosLocal); + wOutQue_.EnQue(kPosLocal); + kgOutQue_.EnQue(kNegLocal); + } + + __aicore__ inline uint64_t ScoreVectorMaxRows(uint64_t bytesPerElem) const + { + constexpr uint64_t arenaBytes = static_cast(KDA_VEC_ARENA_ELEMENTS) * sizeof(float); + uint64_t maxRows = (arenaBytes / bytesPerElem) / K_; + if (K_ >= 128 && maxRows > 32) { + maxRows = 32; + } + return maxRows; + } + + __aicore__ inline void PrepareScoreFactorsBulk(uint64_t b, uint64_t h, uint64_t hv, uint64_t start, + uint64_t subBlockIdx, uint64_t subBlockNum, + uint64_t refToken, uint64_t scoreRowBegin, + uint64_t scoreRowCount, uint64_t validColEnd, + uint64_t scoreSlot) + { + LocalTensor refFp32 = exp2Buf_.Get(); + LoadAsFloatRow(gk_, KVOffset(b, hv, refToken, 0, K_), refFp32, K_); + + uint64_t qwBegin = scoreRowBegin + (scoreRowCount * subBlockIdx) / subBlockNum; + uint64_t qwEnd = scoreRowBegin + (scoreRowCount * (subBlockIdx + 1)) / subBlockNum; + uint64_t qwMaxRows = ScoreVectorMaxRows(5 * sizeof(float) + 2 * sizeof(T)); + for (uint64_t tileRow = qwBegin; tileRow < qwEnd; tileRow += qwMaxRows) { + uint64_t tileRows = qwEnd - tileRow; + if (tileRows > qwMaxRows) { + tileRows = qwMaxRows; + } + uint64_t elems = tileRows * K_; + LocalTensor arena = vecBuf_.Get(); + LocalTensor qFp32 = arena; + LocalTensor kFp32 = arena[elems]; + LocalTensor gFp32 = arena[2 * elems]; + LocalTensor expFp32 = arena[3 * elems]; + LocalTensor outFp32 = arena[4 * elems]; + uint64_t typedOffset = (5 * elems * sizeof(float) + sizeof(T) - 1) / sizeof(T); + LocalTensor typedBase = vecBuf_.Get()[typedOffset]; + LocalTensor qTyped = typedBase; + LocalTensor kTyped = typedBase[elems]; + + uint64_t token = start + tileRow; + CopyVectorIn(qTyped, q_, QOffset(b, h, token, 0), elems); + CopyVectorIn(kTyped, k_, QOffset(b, h, token, 0), elems); + CopyVectorIn(gFp32, gk_, KVOffset(b, hv, token, 0, K_), elems); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + Cast(qFp32, qTyped, RoundMode::CAST_NONE, static_cast(elems)); + Cast(kFp32, kTyped, RoundMode::CAST_NONE, static_cast(elems)); + PipeBarrier(); + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expFp32[row * K_], gFp32[row * K_], refFp32, static_cast(K_)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(elems)); + PipeBarrier(); + ClampExpInput(expFp32, static_cast(elems)); + Exp(expFp32, expFp32, static_cast(elems)); + PipeBarrier(); + + Mul(outFp32, qFp32, expFp32, static_cast(elems)); + PipeBarrier(); + ClampFp32ToOutputType(outFp32, static_cast(elems)); + Cast(qTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); + Mul(outFp32, kFp32, expFp32, static_cast(elems)); + PipeBarrier(); + ClampFp32ToOutputType(outFp32, static_cast(elems)); + Cast(kTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); + + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG, tileRow), + qTyped, elems); + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W, tileRow), + kTyped, elems); + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + } + + uint64_t kgBegin = (validColEnd * subBlockIdx) / subBlockNum; + uint64_t kgEnd = (validColEnd * (subBlockIdx + 1)) / subBlockNum; + uint64_t kgMaxRows = ScoreVectorMaxRows(4 * sizeof(float) + sizeof(T)); + for (uint64_t tileRow = kgBegin; tileRow < kgEnd; tileRow += kgMaxRows) { + uint64_t tileRows = kgEnd - tileRow; + if (tileRows > kgMaxRows) { + tileRows = kgMaxRows; + } + uint64_t elems = tileRows * K_; + LocalTensor arena = vecBuf_.Get(); + LocalTensor kFp32 = arena; + LocalTensor gFp32 = arena[elems]; + LocalTensor expFp32 = arena[2 * elems]; + LocalTensor outFp32 = arena[3 * elems]; + uint64_t typedOffset = (4 * elems * sizeof(float) + sizeof(T) - 1) / sizeof(T); + LocalTensor kTyped = vecBuf_.Get()[typedOffset]; + + uint64_t token = start + tileRow; + CopyVectorIn(kTyped, k_, QOffset(b, h, token, 0), elems); + CopyVectorIn(gFp32, gk_, KVOffset(b, hv, token, 0, K_), elems); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + Cast(kFp32, kTyped, RoundMode::CAST_NONE, static_cast(elems)); + PipeBarrier(); + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expFp32[row * K_], refFp32, gFp32[row * K_], static_cast(K_)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(elems)); + PipeBarrier(); + ClampExpInput(expFp32, static_cast(elems)); + Exp(expFp32, expFp32, static_cast(elems)); + PipeBarrier(); + Mul(outFp32, kFp32, expFp32, static_cast(elems)); + PipeBarrier(); + ClampFp32ToOutputType(outFp32, static_cast(elems)); + Cast(kTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); + + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG, tileRow), + kTyped, elems); + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + } + } + + __aicore__ inline bool PrepareGateProductsBulk(uint64_t b, uint64_t h, uint64_t hv, uint64_t start, + uint64_t curT, uint64_t subBlockIdx, uint64_t subBlockNum, + bool useRef, uint64_t refToken, uint64_t validColEnd, + bool writeScoreScratch, uint64_t scoreSlot) + { + if constexpr (IsSameType::value) { + return false; + } + if (subBlockNum == 0 || subBlockIdx >= subBlockNum || K_ == 0) { + return false; + } + uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; + uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + if (rowBegin >= rowEnd) { + return true; + } + + constexpr uint64_t arenaBytes = static_cast(KDA_VEC_ARENA_ELEMENTS) * sizeof(float); + constexpr uint64_t bytesPerElem = 5 * sizeof(float) + 3 * sizeof(T); + uint64_t maxElems = arenaBytes / bytesPerElem; + uint64_t maxRows = maxElems / K_; + // Keep the multi-row SIMD tile below the 192 KiB per-core UB budget. + // K=128 uses five FP32 work planes plus three typed planes; 32 rows + // leaves headroom for alignment and the surrounding pipeline buffers. + if (K_ >= 128 && maxRows > 32) { + maxRows = 32; + } + if (maxRows == 0) { + return false; + } + LocalTensor refFp32 = exp2Buf_.Get(); + if (useRef) { + LoadAsFloatRow(gk_, KVOffset(b, hv, refToken, 0, K_), refFp32, K_); + } + + for (uint64_t tileRow = rowBegin; tileRow < rowEnd; tileRow += maxRows) { + uint64_t tileRows = rowEnd - tileRow; + if (tileRows > maxRows) { + tileRows = maxRows; + } + uint64_t elems = tileRows * K_; + LocalTensor arena = vecBuf_.Get(); + LocalTensor qFp32 = arena; + LocalTensor kFp32 = arena[elems]; + LocalTensor gFp32 = arena[2 * elems]; + LocalTensor expFp32 = arena[3 * elems]; + LocalTensor outFp32 = arena[4 * elems]; + + uint64_t typedOffset = (5 * elems * sizeof(float) + sizeof(T) - 1) / sizeof(T); + uint64_t typedCapacity = arenaBytes / sizeof(T); + if (typedOffset + 3 * elems > typedCapacity) { + return false; + } + LocalTensor typedBase = vecBuf_.Get()[typedOffset]; + LocalTensor qTyped = typedBase; + LocalTensor kTyped = typedBase[elems]; + LocalTensor kgTyped = typedBase[2 * elems]; + + uint64_t token = start + tileRow; + CopyVectorIn(qTyped, q_, QOffset(b, h, token, 0), elems); + CopyVectorIn(kTyped, k_, QOffset(b, h, token, 0), elems); + CopyVectorIn(gFp32, gk_, KVOffset(b, hv, token, 0, K_), elems); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + + Cast(qFp32, qTyped, RoundMode::CAST_NONE, static_cast(elems)); + Cast(kFp32, kTyped, RoundMode::CAST_NONE, static_cast(elems)); + PipeBarrier(); + + if (useRef) { + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expFp32[row * K_], gFp32[row * K_], refFp32, static_cast(K_)); + } + } else { + Adds(expFp32, gFp32, 0.0f, static_cast(elems)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(elems)); + PipeBarrier(); + ClampExpInput(expFp32, static_cast(elems)); + Exp(expFp32, expFp32, static_cast(elems)); + PipeBarrier(); + + Mul(outFp32, qFp32, expFp32, static_cast(elems)); + PipeBarrier(); + ClampFp32ToOutputType(outFp32, static_cast(elems)); + Cast(qTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); + + Mul(outFp32, kFp32, expFp32, static_cast(elems)); + PipeBarrier(); + ClampFp32ToOutputType(outFp32, static_cast(elems)); + Cast(kTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); + + if (useRef) { + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expFp32[row * K_], refFp32, gFp32[row * K_], static_cast(K_)); + } + } else { + Muls(expFp32, gFp32, -1.0f, static_cast(elems)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(elems)); + PipeBarrier(); + ClampExpInput(expFp32, static_cast(elems)); + Exp(expFp32, expFp32, static_cast(elems)); + PipeBarrier(); + Mul(outFp32, kFp32, expFp32, static_cast(elems)); + PipeBarrier(); + if (useRef && tileRow + tileRows > validColEnd) { + for (uint64_t row = 0; row < tileRows; ++row) { + if (tileRow + row >= validColEnd) { + Duplicate(outFp32[row * K_], 0.0f, static_cast(K_)); + } + } + PipeBarrier(); + } + ClampFp32ToOutputType(outFp32, static_cast(elems)); + Cast(kgTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); + + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + if (writeScoreScratch) { + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG, tileRow), + qTyped, elems); + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W, tileRow), + kTyped, elems); + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG, tileRow), + kgTyped, elems); + } else { + CopyVectorOut(qg_, KVOffset(b, hv, token, 0, K_), qTyped, elems); + CopyVectorOut(w_, KVOffset(b, hv, token, 0, K_), kTyped, elems); + CopyVectorOut(kg_, KVOffset(b, hv, token, 0, K_), kgTyped, elems); + } + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + } + return true; + } + + __aicore__ inline void PrepareGateProducts(uint64_t b, uint64_t h, uint64_t hv, uint64_t start, uint64_t curT, + uint64_t subBlockIdx, uint64_t subBlockNum, bool useRef = false, + uint64_t refToken = 0, uint64_t validColEnd = 0, + bool writeScoreScratch = false, uint64_t scoreSlot = 0, + uint64_t scoreRowBegin = 0, uint64_t scoreRowCount = 0) + { + if (subBlockNum == 0 || subBlockIdx >= subBlockNum) { + return; + } + if (validColEnd == 0 || validColEnd > curT) { + validColEnd = curT; + } + if (writeScoreScratch) { + PrepareScoreFactorsBulk(b, h, hv, start, subBlockIdx, subBlockNum, refToken, scoreRowBegin, + scoreRowCount, validColEnd, scoreSlot); + return; + } + if (PrepareGateProductsBulk(b, h, hv, start, curT, subBlockIdx, subBlockNum, useRef, refToken, + validColEnd, writeScoreScratch, scoreSlot)) { + return; + } + + if (subBlockIdx >= curT) { + return; + } + + LocalTensor vecLocal = vecBuf_.Get(); + LocalTensor qFp32 = vecLocal; + LocalTensor kFp32 = vecLocal[EXP2_UB_ELEMENTS]; + LocalTensor gFp32 = vecLocal[2 * EXP2_UB_ELEMENTS]; + LocalTensor expFp32 = vecLocal[3 * EXP2_UB_ELEMENTS]; + LocalTensor outFp32 = vecLocal[4 * EXP2_UB_ELEMENTS]; + LocalTensor refFp32 = vecLocal[5 * EXP2_UB_ELEMENTS]; + if (useRef) { + LoadAsFloatRow(gk_, KVOffset(b, hv, refToken, 0, K_), refFp32, K_); + } + + uint64_t rowCount = 0; + for (uint64_t i = subBlockIdx; i < curT; i += subBlockNum) { + ++rowCount; + } + + LoadGateProductRow(b, h, hv, GateProductToken(start, 0, subBlockIdx, subBlockNum)); + for (uint64_t logicalIdx = 0; logicalIdx < rowCount; ++logicalIdx) { + uint64_t ti = GateProductToken(start, logicalIdx, subBlockIdx, subBlockNum); + LocalTensor qLocal = qInQue_.DeQue(); + LocalTensor kLocal = kInQue_.DeQue(); + LocalTensor gLocal = gInQue_.DeQue(); + + if (logicalIdx + 1 < rowCount) { + LoadGateProductRow(b, h, hv, GateProductToken(start, logicalIdx + 1, subBlockIdx, subBlockNum)); + } + + if constexpr (IsSameType::value) { + DataCopy(qFp32, qLocal, static_cast(K_)); + DataCopy(kFp32, kLocal, static_cast(K_)); + } else { + Cast(qFp32, qLocal, RoundMode::CAST_NONE, static_cast(K_)); + Cast(kFp32, kLocal, RoundMode::CAST_NONE, static_cast(K_)); + } + DataCopy(gFp32, gLocal, static_cast(K_)); + qInQue_.FreeTensor(qLocal); + kInQue_.FreeTensor(kLocal); + gInQue_.FreeTensor(gLocal); + PipeBarrier(); + + if (logicalIdx > 0) { + StoreGateProductRow(b, hv, GateProductToken(start, logicalIdx - 1, subBlockIdx, subBlockNum)); + } + bool zeroKg = useRef && (ti - start >= validColEnd); + ComputeGateProductRow(qFp32, kFp32, gFp32, refFp32, expFp32, outFp32, useRef, zeroKg); + } + StoreGateProductRow(b, hv, GateProductToken(start, rowCount - 1, subBlockIdx, subBlockNum)); + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + } + + __aicore__ inline void ComputeRawAqkAkkCube(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT) + { + ComputeRawAqkAkkCubeBlock(b, hv, start, curT, 0, curT); + } + + __aicore__ inline void ComputeRawAqkAkkCubeBlock(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT, + uint64_t rowBegin, uint64_t rowCount, + bool readScoreScratch = false, uint64_t scoreSlot = 0, + uint64_t colCount = 0) + { + using ElementA = T; + using ElementB = T; + using ElementC = float; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::ColumnMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using BlockMmad = Catlass::Gemm::Block::BlockMmadTla; + + Catlass::Arch::Resource resource; + BlockMmad blockMmad(resource); + auto layoutA = tla::MakeLayout(BT_, K_); + auto layoutB = tla::MakeLayout(K_, BT_); + auto layoutC = tla::MakeLayout(BT_, BT_); + if (colCount == 0 || colCount > curT) { + colCount = curT; + } + Catlass::GemmCoord shape{static_cast(rowCount), static_cast(colCount), + static_cast(K_)}; + + auto tensorQPos = readScoreScratch ? + tla::MakeTensor(scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG)], + layoutA, Catlass::Arch::PositionGM{}) : + tla::MakeTensor(qg_[KVOffset(b, hv, start, 0, K_)], layoutA, + Catlass::Arch::PositionGM{}); + auto tensorKPos = readScoreScratch ? + tla::MakeTensor(scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W)], + layoutA, Catlass::Arch::PositionGM{}) : + tla::MakeTensor(w_[KVOffset(b, hv, start, 0, K_)], layoutA, + Catlass::Arch::PositionGM{}); + auto tensorKNeg = readScoreScratch ? + tla::MakeTensor(scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG)], + layoutB, Catlass::Arch::PositionGM{}) : + tla::MakeTensor(kg_[KVOffset(b, hv, start, 0, K_)], layoutB, + Catlass::Arch::PositionGM{}); + auto tensorAqk = tla::MakeTensor(aqk_[AOffset(b, hv, start, 0)], layoutC, + Catlass::Arch::PositionGM{}); + auto tensorAkk = tla::MakeTensor(akk_[AOffset(b, hv, start, 0)], layoutC, + Catlass::Arch::PositionGM{}); + + auto blockQPos = GetTile(tensorQPos, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.k())); + auto blockKPos = GetTile(tensorKPos, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.k())); + auto blockKNeg = GetTile(tensorKNeg, tla::MakeCoord(0, 0), tla::MakeShape(shape.k(), shape.n())); + auto blockAqk = GetTile(tensorAqk, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.n())); + auto blockAkk = GetTile(tensorAkk, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.n())); + + blockMmad(blockQPos, blockKNeg, blockAqk, shape); + PipeBarrier(); + blockMmad(blockKPos, blockKNeg, blockAkk, shape); + PipeBarrier(); + } + + __aicore__ inline bool UseAkkCubeSolve(uint64_t curT) const + { + return curT > 0 && curT <= BT_ && (BT_ == 64 || BT_ == 128) && K_ >= 16 && V_ >= 16 && + V_ <= 256 && K_ % 16 == 0 && V_ % 16 == 0; + } + + __aicore__ inline bool UsePostWuCube(uint64_t curT) const + { + return curT > 0 && curT <= BT_ && (BT_ == 64 || BT_ == 128) && K_ >= 16 && V_ >= 16 && + V_ <= 256 && K_ % 16 == 0 && V_ % 16 == 0; + } + + __aicore__ inline void CopyLocalFloat(LocalTensor dst, LocalTensor src, uint64_t count) + { + if (count == 0) { + return; + } + Adds(dst, src, 0.0f, static_cast(count)); + PipeBarrier(); + } + + __aicore__ inline void FillLocalFloat(LocalTensor dst, float value, uint64_t count) + { + if (count == 0) { + return; + } + Duplicate(dst, value, static_cast(count)); + PipeBarrier(); + } + + __aicore__ inline void BuildPrefixMask(LocalTensor dst, uint64_t prefix, uint64_t count) + { + if (prefix > count) { + prefix = count; + } + Duplicate(dst, 0.0f, static_cast(count)); + if (prefix > 0) { + Duplicate(dst, 1.0f, static_cast(prefix)); + } + PipeBarrier(); + } + + __aicore__ inline uint64_t BuildCausalMask(uint64_t threshold, uint64_t colBegin) const + { + if (threshold <= colBegin) { + return ~0ULL; + } + if (threshold >= colBegin + KDA_SOLVE_BT) { + return 0ULL; + } + return ~0ULL << (threshold - colBegin); + } + + __aicore__ inline void BuildCausalSelectMasks(LocalTensor aqkMask, LocalTensor akkMask, + uint64_t rowBegin, uint64_t rowCount, uint64_t colBegin) + { + __ubuf__ uint64_t *aqkMaskPtr = reinterpret_cast<__ubuf__ uint64_t *>(aqkMask.GetPhyAddr()); + __ubuf__ uint64_t *akkMaskPtr = reinterpret_cast<__ubuf__ uint64_t *>(akkMask.GetPhyAddr()); + for (uint32_t localRow = 0; localRow < rowCount; ++localRow) { + uint32_t row = static_cast(rowBegin + localRow); + aqkMaskPtr[localRow] = BuildCausalMask(static_cast(row) + 1, colBegin); + akkMaskPtr[localRow] = BuildCausalMask(static_cast(row), colBegin); + } + } + + __aicore__ inline void SelectCausalRows(LocalTensor aqkMat, LocalTensor akkMat, + uint64_t rowBegin, uint64_t rowCount) + { + LocalTensor aqkMask = vecBuf_.Get()[KDA_SELECT_AQK_MASK_BYTE_OFFSET]; + LocalTensor akkMask = vecBuf_.Get()[KDA_SELECT_AKK_MASK_BYTE_OFFSET]; + LocalTensor zeroLocal = vecBuf_.Get()[KDA_SELECT_ZERO_FLOAT_OFFSET]; + Duplicate(zeroLocal, 0.0f, 8); + PipeBarrier(); + + uint64_t colBlockCount = (BT_ + KDA_SOLVE_BT - 1) / KDA_SOLVE_BT; + for (uint64_t colBlock = 0; colBlock < colBlockCount; ++colBlock) { + uint64_t maskOffset = colBlock * KDA_SELECT_COL_MASK_BYTES; + uint64_t colBegin = colBlock * KDA_SOLVE_BT; + BuildCausalSelectMasks(aqkMask[maskOffset], akkMask[maskOffset], rowBegin, rowCount, colBegin); + } + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + + uint8_t rowStride = static_cast(BT_ * sizeof(float) / 32); + BinaryRepeatParams repeatParams = {1, 0, 1, rowStride, 0, rowStride}; + for (uint64_t colBlock = 0; colBlock < colBlockCount; ++colBlock) { + uint64_t maskOffset = colBlock * KDA_SELECT_COL_MASK_BYTES; + uint64_t colBegin = colBlock * KDA_SOLVE_BT; + Select(aqkMat[colBegin], aqkMask[maskOffset], zeroLocal, aqkMat[colBegin], + SELMODE::VSEL_TENSOR_TENSOR_MODE, KDA_SOLVE_BT, static_cast(rowCount), repeatParams); + Select(akkMat[colBegin], akkMask[maskOffset], zeroLocal, akkMat[colBegin], + SELMODE::VSEL_TENSOR_TENSOR_MODE, KDA_SOLVE_BT, static_cast(rowCount), repeatParams); + } + PipeBarrier(); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + } + + __aicore__ inline void PrepareAqkAkkSolveInput64(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start) + { + LocalTensor arena = vecBuf_.Get(); + LocalTensor aqkMat = arena; + LocalTensor akkMat = arena[KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor xMat = arena[2 * KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor betaLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor betaBrcb = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT]; + LocalTensor maskLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512]; + LocalTensor oneHotLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512 + KDA_SOLVE_BT]; + + LoadAsFloatRow(beta_, BetaOffset(b, hv, start), betaLocal, KDA_SOLVE_BT); + Brcb(betaBrcb, betaLocal, 8, {1, 8}); + PipeBarrier(); + + DataCopy(aqkMat, aqk_[AOffset(b, hv, start, 0)], KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(akkMat, akk_[AOffset(b, hv, start, 0)], KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + + for (uint64_t col = 0; col < KDA_SOLVE_BT; col += 8) { + Mul(akkMat[col], akkMat[col], betaBrcb, 8, KDA_SOLVE_BT, {1, 1, 1, 8, 8, 1}); + PipeBarrier(); + } + SelectCausalRows(aqkMat, akkMat, 0, KDA_SOLVE_BT); + + Muls(xMat, akkMat, -1.0f, KDA_SOLVE_MATRIX_ELEMENTS); + PipeBarrier(); + for (uint64_t row = 0; row < KDA_SOLVE_BT; ++row) { + BuildPrefixMask(maskLocal, row + 1, KDA_SOLVE_BT); + BuildPrefixMask(oneHotLocal, row, KDA_SOLVE_BT); + Sub(maskLocal, maskLocal, oneHotLocal, KDA_SOLVE_BT); + PipeBarrier(); + Add(xMat[row * KDA_SOLVE_BT], xMat[row * KDA_SOLVE_BT], maskLocal, KDA_SOLVE_BT); + PipeBarrier(); + } + + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + DataCopy(aqk_[AOffset(b, hv, start, 0)], aqkMat, KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(akk_[AOffset(b, hv, start, 0)], akkMat, KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(h_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X)], xMat, + KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + } + + __aicore__ inline void PrepareAqkAkkSolveInputTail(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT) + { + uint64_t elemCount = curT * KDA_SOLVE_BT; + DataCopyParams aqkValidParams{1, static_cast(elemCount * sizeof(float)), 0, 0}; + DataCopyParams akkValidParams{1, static_cast(elemCount * sizeof(float)), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + LocalTensor arena = vecBuf_.Get(); + LocalTensor aqkMat = arena; + LocalTensor akkMat = arena[KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor xMat = arena[2 * KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor betaLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor betaBrcb = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT]; + LocalTensor maskLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512]; + LocalTensor oneHotLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512 + KDA_SOLVE_BT]; + + FillLocalFloat(betaLocal, 0.0f, KDA_SOLVE_BT); + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + LoadAsFloatRow(beta_, BetaOffset(b, hv, start), betaLocal, curT); + Brcb(betaBrcb, betaLocal, 8, {1, 8}); + PipeBarrier(); + + DataCopyPad(aqkMat, aqk_[AOffset(b, hv, start, 0)], aqkValidParams, padParams); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + if (elemCount < KDA_SOLVE_MATRIX_ELEMENTS) { + FillLocalFloat(aqkMat[elemCount], 0.0f, KDA_SOLVE_MATRIX_ELEMENTS - elemCount); + } + DataCopyPad(akkMat, akk_[AOffset(b, hv, start, 0)], akkValidParams, padParams); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + if (elemCount < KDA_SOLVE_MATRIX_ELEMENTS) { + FillLocalFloat(akkMat[elemCount], 0.0f, KDA_SOLVE_MATRIX_ELEMENTS - elemCount); + } + + for (uint64_t col = 0; col < KDA_SOLVE_BT; col += 8) { + Mul(akkMat[col], akkMat[col], betaBrcb, 8, KDA_SOLVE_BT, {1, 1, 1, 8, 8, 1}); + PipeBarrier(); + } + SelectCausalRows(aqkMat, akkMat, 0, KDA_SOLVE_BT); + + Muls(xMat, akkMat, -1.0f, KDA_SOLVE_MATRIX_ELEMENTS); + PipeBarrier(); + for (uint64_t row = 0; row < KDA_SOLVE_BT; ++row) { + BuildPrefixMask(maskLocal, row + 1, KDA_SOLVE_BT); + BuildPrefixMask(oneHotLocal, row, KDA_SOLVE_BT); + Sub(maskLocal, maskLocal, oneHotLocal, KDA_SOLVE_BT); + PipeBarrier(); + Add(xMat[row * KDA_SOLVE_BT], xMat[row * KDA_SOLVE_BT], maskLocal, KDA_SOLVE_BT); + PipeBarrier(); + } + + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + DataCopyPad(aqk_[AOffset(b, hv, start, 0)], aqkMat, aqkValidParams); + DataCopy(h_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X)], xMat, + KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(h_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0)], akkMat, + KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + } + + __aicore__ inline void GetSolveRowRange(uint64_t curT, uint64_t subBlockIdx, uint64_t subBlockNum, + uint64_t &rowBegin, uint64_t &rowEnd) const + { + if (subBlockNum == 0 || subBlockIdx >= subBlockNum) { + rowBegin = 0; + rowEnd = 0; + return; + } + rowBegin = (curT * subBlockIdx) / subBlockNum; + rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + } + + __aicore__ inline void PrepareAqkAkkSolveInputRows(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT, uint64_t rowBegin, + uint64_t rowEnd, bool storeLToAkk, bool storeLToScratch) + { + uint64_t rowCount = rowEnd - rowBegin; + if (rowCount == 0) { + return; + } + uint64_t validRowCount = rowBegin < curT ? curT - rowBegin : 0; + if (validRowCount > rowCount) { + validRowCount = rowCount; + } + uint64_t elemCount = rowCount * BT_; + uint64_t validElemCount = validRowCount * BT_; + DataCopyParams aqkValidParams{1, static_cast(validElemCount * sizeof(float)), 0, 0}; + DataCopyParams akkValidParams{1, static_cast(validElemCount * sizeof(float)), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + LocalTensor arena = vecBuf_.Get(); + LocalTensor aqkMat = arena; + LocalTensor akkMat = arena[elemCount]; + LocalTensor xMat = arena[2 * elemCount]; + LocalTensor betaLocal = arena[3 * elemCount]; + LocalTensor betaBrcb = arena[3 * elemCount + BT_]; + LocalTensor maskLocal = arena[3 * elemCount + BT_ + 512]; + LocalTensor oneHotLocal = arena[3 * elemCount + BT_ + 512 + BT_]; + + uint64_t token = start + rowBegin; + + FillLocalFloat(aqkMat, 0.0f, elemCount); + FillLocalFloat(akkMat, 0.0f, elemCount); + FillLocalFloat(betaLocal, 0.0f, rowCount); + SetFlag(KDA_MTE2_MTE3_EVENT_ID); + WaitFlag(KDA_MTE2_MTE3_EVENT_ID); + if (validRowCount > 0) { + LoadAsFloatRow(beta_, BetaOffset(b, hv, token), betaLocal, validRowCount); + DataCopyPad(aqkMat, aqk_[AOffset(b, hv, token, 0)], aqkValidParams, padParams); + DataCopyPad(akkMat, akk_[AOffset(b, hv, token, 0)], akkValidParams, padParams); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + } + Brcb(betaBrcb, betaLocal, static_cast((rowCount + 7) / 8), {1, 8}); + PipeBarrier(); + + uint8_t rowStride = static_cast(BT_ * sizeof(float) / 32); + for (uint64_t col = 0; col < BT_; col += 8) { + Mul(akkMat[col], akkMat[col], betaBrcb, 8, static_cast(rowCount), + {1, 1, 0, rowStride, rowStride, 1}); + PipeBarrier(); + } + if (validRowCount > 0) { + SelectCausalRows(aqkMat, akkMat, rowBegin, validRowCount); + } + + Muls(xMat, akkMat, -1.0f, static_cast(elemCount)); + PipeBarrier(); + for (uint64_t localRow = 0; localRow < rowCount; ++localRow) { + uint64_t row = rowBegin + localRow; + BuildPrefixMask(maskLocal, row + 1, BT_); + BuildPrefixMask(oneHotLocal, row, BT_); + Sub(maskLocal, maskLocal, oneHotLocal, static_cast(BT_)); + PipeBarrier(); + Add(xMat[localRow * BT_], xMat[localRow * BT_], maskLocal, static_cast(BT_)); + PipeBarrier(); + } + + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; + uint64_t lBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0) + rowBegin * BT_; + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + if (validRowCount > 0) { + DataCopyPad(aqk_[AOffset(b, hv, token, 0)], aqkMat, aqkValidParams); + if (storeLToAkk) { + DataCopyPad(akk_[AOffset(b, hv, token, 0)], akkMat, akkValidParams); + } + } + DataCopy(solveWorkspace_[xBase], xMat, static_cast(elemCount)); + if (storeLToScratch) { + DataCopy(solveWorkspace_[lBase], akkMat, static_cast(elemCount)); + } + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + } + + __aicore__ inline void CubeGemmSolveSub(GlobalTensor &tensorA, uint64_t baseA, uint64_t rowA, uint64_t colA, + GlobalTensor &tensorB, uint64_t baseB, uint64_t rowB, uint64_t colB, + GlobalTensor &tensorC, uint64_t baseC, uint64_t rowC, uint64_t colC, + uint32_t m, uint32_t n, uint32_t k) + { + using ElementA = float; + using ElementB = float; + using ElementC = float; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using BlockMmad = Catlass::Gemm::Block::BlockMmadTla; + Catlass::Arch::Resource resource; + auto layoutA = tla::MakeLayout(BT_, BT_); + auto layoutB = tla::MakeLayout(BT_, BT_); + auto layoutC = tla::MakeLayout(BT_, BT_); + auto tensorLayoutA = tla::MakeTensor(tensorA[baseA], layoutA, Catlass::Arch::PositionGM{}); + auto tensorLayoutB = tla::MakeTensor(tensorB[baseB], layoutB, Catlass::Arch::PositionGM{}); + auto tensorLayoutC = tla::MakeTensor(tensorC[baseC], layoutC, Catlass::Arch::PositionGM{}); + Catlass::GemmCoord shape{m, n, k}; + auto blockA = GetTile(tensorLayoutA, tla::MakeCoord(rowA, colA), tla::MakeShape(shape.m(), shape.k())); + auto blockB = GetTile(tensorLayoutB, tla::MakeCoord(rowB, colB), tla::MakeShape(shape.k(), shape.n())); + auto blockC = GetTile(tensorLayoutC, tla::MakeCoord(rowC, colC), tla::MakeShape(shape.m(), shape.n())); + BlockMmad blockMmad(resource); + blockMmad(blockA, blockB, blockC, shape); + PipeBarrier(); + } + + __aicore__ inline void AddSolveTmpToX(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + bool storeAkk) + { + LocalTensor arena = vecBuf_.Get(); + LocalTensor xLocal = arena; + LocalTensor tmpLocal = arena[KDA_SOLVE_MATRIX_ELEMENTS]; + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + DataCopy(xLocal, h_[xBase], KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(tmpLocal, h_[tmpBase], KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + + Add(xLocal, xLocal, tmpLocal, KDA_SOLVE_MATRIX_ELEMENTS); + PipeBarrier(); + + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + DataCopy(h_[xBase], xLocal, KDA_SOLVE_MATRIX_ELEMENTS); + if (storeAkk) { + DataCopy(akk_[AOffset(b, hv, start, 0)], xLocal, KDA_SOLVE_MATRIX_ELEMENTS); + } + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + } + + __aicore__ inline void AddSolveTmpToXTail(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT, bool storeAkk) + { + uint64_t elemCount = curT * KDA_SOLVE_BT; + DataCopyParams validParams{1, static_cast(elemCount * sizeof(float)), 0, 0}; + LocalTensor arena = vecBuf_.Get(); + LocalTensor xLocal = arena; + LocalTensor tmpLocal = arena[KDA_SOLVE_MATRIX_ELEMENTS]; + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + DataCopy(xLocal, h_[xBase], KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(tmpLocal, h_[tmpBase], KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + + Add(xLocal, xLocal, tmpLocal, KDA_SOLVE_MATRIX_ELEMENTS); + PipeBarrier(); + + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + DataCopy(h_[xBase], xLocal, KDA_SOLVE_MATRIX_ELEMENTS); + if (storeAkk) { + DataCopyPad(akk_[AOffset(b, hv, start, 0)], xLocal, validParams); + } + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + } + + __aicore__ inline void AddSolveTmpToXRows(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT, uint64_t rowBegin, uint64_t rowEnd, bool storeAkk) + { + uint64_t rowCount = rowEnd - rowBegin; + if (rowCount == 0) { + return; + } + uint64_t validRowCount = rowBegin < curT ? curT - rowBegin : 0; + if (validRowCount > rowCount) { + validRowCount = rowCount; + } + uint64_t elemCount = rowCount * BT_; + uint64_t validElemCount = validRowCount * BT_; + DataCopyParams validParams{1, static_cast(validElemCount * sizeof(float)), 0, 0}; + LocalTensor arena = vecBuf_.Get(); + LocalTensor xLocal = arena; + LocalTensor tmpLocal = arena[elemCount]; + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP) + rowBegin * BT_; + uint64_t token = start + rowBegin; + + DataCopy(xLocal, solveWorkspace_[xBase], static_cast(elemCount)); + DataCopy(tmpLocal, solveWorkspace_[tmpBase], static_cast(elemCount)); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + + Add(xLocal, xLocal, tmpLocal, static_cast(elemCount)); + PipeBarrier(); + + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + DataCopy(solveWorkspace_[xBase], xLocal, static_cast(elemCount)); + if (storeAkk && validRowCount > 0) { + DataCopyPad(akk_[AOffset(b, hv, token, 0)], xLocal, validParams); + } + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + } + + __aicore__ inline void AddSolveTmpToXDiagRows(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t rowBegin, uint64_t rowEnd, bool storeAkk) + { + uint64_t rowCount = rowEnd - rowBegin; + if (rowCount == 0) { + return; + } + uint64_t elemCount = rowCount * BT_; + DataCopyParams validParams{1, static_cast(elemCount * sizeof(float)), 0, 0}; + LocalTensor arena = vecBuf_.Get(); + LocalTensor xLocal = arena; + LocalTensor tmpLocal = arena[elemCount]; + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP) + rowBegin * BT_; + uint64_t token = start + rowBegin; + + DataCopy(xLocal, solveWorkspace_[xBase], static_cast(elemCount)); + DataCopy(tmpLocal, solveWorkspace_[tmpBase], static_cast(elemCount)); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + + for (uint64_t localRow = 0; localRow < rowCount; ++localRow) { + uint64_t row = rowBegin + localRow; + uint64_t col = (row / KDA_SOLVE_DIAG_BT) * KDA_SOLVE_DIAG_BT; + uint64_t offset = localRow * BT_ + col; + Add(xLocal[offset], xLocal[offset], tmpLocal[offset], KDA_SOLVE_DIAG_BT); + PipeBarrier(); + } + + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + DataCopy(solveWorkspace_[xBase], xLocal, static_cast(elemCount)); + if (storeAkk) { + DataCopyPad(akk_[AOffset(b, hv, token, 0)], xLocal, validParams); + } + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + } + + __aicore__ inline void StoreSolveXRowsToAkk(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT, uint64_t rowBegin, uint64_t rowEnd) + { + uint64_t validRowCount = rowBegin < curT ? curT - rowBegin : 0; + uint64_t rowCount = rowEnd - rowBegin; + if (validRowCount > rowCount) { + validRowCount = rowCount; + } + if (validRowCount == 0) { + return; + } + uint64_t elemCount = validRowCount * BT_; + DataCopyParams validParams{1, static_cast(elemCount * sizeof(float)), 0, 0}; + LocalTensor xLocal = vecBuf_.Get(); + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; + + DataCopy(xLocal, solveWorkspace_[xBase], static_cast(elemCount)); + SetFlag(KDA_MTE2_MTE3_EVENT_ID); + WaitFlag(KDA_MTE2_MTE3_EVENT_ID); + DataCopyPad(akk_[AOffset(b, hv, start + rowBegin, 0)], xLocal, validParams); + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + } + + __aicore__ inline void ComputeAkkMergeCube(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start) + { + uint64_t aiBase = AOffset(b, hv, start, 0); + uint64_t negABase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + for (uint32_t mergeSize = 2 * KDA_SOLVE_DIAG_BT; mergeSize <= BT_; mergeSize *= 2) { + uint32_t half = mergeSize / 2; + for (uint32_t block = 0; block < BT_; block += mergeSize) { + uint32_t lower = block + half; + CubeGemmSolveSub(akk_, aiBase, lower, lower, solveWorkspace_, negABase, lower, block, + solveWorkspace_, tmpBase, 0, 0, half, half, half); + CubeGemmSolveSub(solveWorkspace_, tmpBase, 0, 0, akk_, aiBase, block, block, + akk_, aiBase, lower, block, half, half, half); + } + } + } + + __aicore__ inline void ComputeAkkMergeCubeWorkspace(uint64_t b, uint64_t hv, uint64_t chunkIdx) + { + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + for (uint32_t mergeSize = 2 * KDA_SOLVE_DIAG_BT; mergeSize <= BT_; mergeSize *= 2) { + uint32_t half = mergeSize / 2; + for (uint32_t block = 0; block < BT_; block += mergeSize) { + uint32_t lower = block + half; + CubeGemmSolveSub(solveWorkspace_, xBase, lower, lower, solveWorkspace_, xBase, lower, block, + solveWorkspace_, tmpBase, 0, 0, half, half, half); + CubeGemmSolveSub(solveWorkspace_, tmpBase, 0, 0, solveWorkspace_, xBase, block, block, + solveWorkspace_, xBase, lower, block, half, half, half); + } + } + } + + __aicore__ inline void ComputeAkkInverseMchFull(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start) + { + uint64_t aBase = AOffset(b, hv, start, 0); + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t yBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0); + uint64_t yNextBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y1); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + uint32_t diagBlocks = static_cast(BT_ / KDA_SOLVE_DIAG_BT); + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(akk_, aBase, off, off, akk_, aBase, off, off, solveWorkspace_, yBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + for (uint32_t iter = 0; iter < KDA_SOLVE_DIAG_MCH_ITERS; ++iter) { + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(solveWorkspace_, xBase, off, off, solveWorkspace_, yBase, off, off, + solveWorkspace_, tmpBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); + if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(solveWorkspace_, yBase, off, off, solveWorkspace_, yBase, off, off, + solveWorkspace_, yNextBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + } + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(syncReadyFlag_); + if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { + uint64_t oldYBase = yBase; + yBase = yNextBase; + yNextBase = oldYBase; + } + } + ComputeAkkMergeCube(b, hv, chunkIdx, start); + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); + } + + __aicore__ inline void ComputeAkkInverseMchTail(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT) + { + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t lBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0); + uint64_t yBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y1); + uint64_t yNextBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + (void)start; + (void)curT; + + uint32_t diagBlocks = static_cast(BT_ / KDA_SOLVE_DIAG_BT); + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(solveWorkspace_, lBase, off, off, solveWorkspace_, lBase, off, off, + solveWorkspace_, yBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + for (uint32_t iter = 0; iter < KDA_SOLVE_DIAG_MCH_ITERS; ++iter) { + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(solveWorkspace_, xBase, off, off, solveWorkspace_, yBase, off, off, + solveWorkspace_, tmpBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); + if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(solveWorkspace_, yBase, off, off, solveWorkspace_, yBase, off, off, + solveWorkspace_, yNextBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + } + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(syncReadyFlag_); + if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { + uint64_t oldYBase = yBase; + yBase = yNextBase; + yNextBase = oldYBase; + } + } + ComputeAkkMergeCubeWorkspace(b, hv, chunkIdx); + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); + } + + + + __aicore__ inline void ScaleRowsByBeta(GlobalTensor &src, GlobalTensor &dst, uint64_t b, uint64_t hv, + uint64_t start, uint64_t rowBegin, uint64_t rowCount, uint64_t dim, + LocalTensor &betaBrcb, LocalTensor &matrixLocal) + { + constexpr uint64_t vecElemsPerRepeat = 64; + constexpr uint64_t typedOffsetFloats = 20480; + constexpr uint64_t typedOffset = typedOffsetFloats * sizeof(float) / sizeof(T); + uint64_t elemCount = rowCount * dim; + uint64_t baseOffset = KVOffset(b, hv, start + rowBegin, 0, dim); + + if constexpr (IsSameType::value) { + DataCopy(matrixLocal, src[baseOffset], static_cast(elemCount)); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + } else { + LocalTensor matrixTyped = vecBuf_.Get()[typedOffset]; + DataCopy(matrixTyped, src[baseOffset], static_cast(elemCount)); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + Cast(matrixLocal, matrixTyped, RoundMode::CAST_NONE, static_cast(elemCount)); + PipeBarrier(); + } + + uint8_t repeatStride = static_cast(dim * sizeof(float) / 32); + for (uint64_t col = 0; col < dim; col += vecElemsPerRepeat) { + uint64_t mask = dim - col; + if (mask > vecElemsPerRepeat) { + mask = vecElemsPerRepeat; + } + Mul(matrixLocal[col], matrixLocal[col], betaBrcb, mask, static_cast(rowCount), + {1, 1, 0, repeatStride, repeatStride, 1}); + PipeBarrier(); + } + + if constexpr (IsSameType::value) { + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + DataCopy(dst[baseOffset], matrixLocal, static_cast(elemCount)); + } else { + LocalTensor matrixTyped = vecBuf_.Get()[typedOffset]; + Cast(matrixTyped, matrixLocal, RoundMode::CAST_RINT, static_cast(elemCount)); + PipeBarrier(); + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + DataCopy(dst[baseOffset], matrixTyped, static_cast(elemCount)); + } + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + } + + __aicore__ inline void PrepareWuCubeInputs(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT, + uint64_t subBlockIdx, uint64_t subBlockNum) + { + uint64_t rowsPerSubBlock = (curT + subBlockNum - 1) / subBlockNum; + uint64_t rowBegin = subBlockIdx * rowsPerSubBlock; + if (rowBegin >= curT) { + return; + } + uint64_t rowCount = curT - rowBegin; + if (rowCount > rowsPerSubBlock) { + rowCount = rowsPerSubBlock; + } + LocalTensor arena = vecBuf_.Get(); + LocalTensor betaLocal = arena; + LocalTensor betaBrcb = arena[KDA_SOLVE_BT]; + LocalTensor matrixLocal = arena[KDA_SOLVE_BT + 512]; + LoadAsFloatRow(beta_, BetaOffset(b, hv, start + rowBegin), betaLocal, rowCount); + Brcb(betaBrcb, betaLocal, static_cast((rowCount + 7) / 8), {1, 8}); + PipeBarrier(); + ScaleRowsByBeta(w_, w_, b, hv, start, rowBegin, rowCount, K_, betaBrcb, matrixLocal); + ScaleRowsByBeta(v_, vNew_, b, hv, start, rowBegin, rowCount, V_, betaBrcb, matrixLocal); + } + + __aicore__ inline void ComputePostWuCube(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT) + { + using ElementA = AKK_T; + using ElementB = T; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using WTileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using UTileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using PostL1TileShape128 = tla::Shape<_128, _128, tla::_256>; + using PostL0TileShape128 = tla::Shape<_128, _128, _128>; + using PostL1TileShape256 = tla::Shape<_128, tla::_256, tla::_256>; + using PostL0TileShape256 = tla::Shape<_128, tla::_256, _64>; + using WBlockMmad = Catlass::Gemm::Block::BlockMmadTla; + using UBlockMmad128 = Catlass::Gemm::Block::BlockMmadTla; + using UBlockMmad256 = Catlass::Gemm::Block::BlockMmadTla; + + LayoutTagA tagA = LayoutTagA::template MakeLayout(BT_, BT_); + auto layoutA = tla::MakeLayoutFromTag(tagA); + auto tensorA = tla::MakeTensor(stageAqk_[AOffset(b, hv, start, 0)], layoutA, + Catlass::Arch::PositionGM{}); + + { + LayoutTagB tagB = LayoutTagB::template MakeLayout(BT_, K_); + LayoutTagC tagC = LayoutTagC::template MakeLayout(BT_, K_); + auto layoutB = tla::MakeLayoutFromTag(tagB); + auto layoutC = tla::MakeLayoutFromTag(tagC); + Catlass::GemmCoord shape{static_cast(curT), static_cast(K_), + static_cast(curT)}; + auto tensorB = tla::MakeTensor(stageQG_[KVOffset(b, hv, start, 0, K_)], layoutB, + Catlass::Arch::PositionGM{}); + auto tensorC = tla::MakeTensor(h_[WScratchOffset(b, hv, chunkIdx, 0, 0)], layoutC, + Catlass::Arch::PositionGM{}); + auto blockA = GetTile(tensorA, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.k())); + auto blockB = GetTile(tensorB, tla::MakeCoord(0, 0), tla::MakeShape(shape.k(), shape.n())); + auto blockC = GetTile(tensorC, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.n())); + Catlass::Arch::Resource wResource; + WBlockMmad wBlockMmad(wResource); + wBlockMmad(blockA, blockB, blockC, shape); + PipeBarrier(); + } + + { + LayoutTagB tagB = LayoutTagB::template MakeLayout(BT_, V_); + LayoutTagC tagC = LayoutTagC::template MakeLayout(BT_, V_); + auto layoutB = tla::MakeLayoutFromTag(tagB); + auto layoutC = tla::MakeLayoutFromTag(tagC); + Catlass::GemmCoord shape{static_cast(curT), static_cast(V_), + static_cast(curT)}; + auto tensorB = tla::MakeTensor(stageVNew_[KVOffset(b, hv, start, 0, V_)], layoutB, + Catlass::Arch::PositionGM{}); + auto tensorC = tla::MakeTensor(u_[KVOffset(b, hv, start, 0, V_)], layoutC, + Catlass::Arch::PositionGM{}); + auto blockA = GetTile(tensorA, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.k())); + auto blockB = GetTile(tensorB, tla::MakeCoord(0, 0), tla::MakeShape(shape.k(), shape.n())); + auto blockC = GetTile(tensorC, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.n())); + Catlass::Arch::Resource uResource; + if (V_ <= 128) { + UBlockMmad128 uBlockMmad(uResource); + uBlockMmad(blockA, blockB, blockC, shape); + } else { + UBlockMmad256 uBlockMmad(uResource); + uBlockMmad(blockA, blockB, blockC, shape); + } + PipeBarrier(); + } + + } + + __aicore__ inline void CopyScratchWAndFinalizeKg(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT, uint64_t subBlockIdx, + uint64_t subBlockNum) + { + constexpr uint64_t typedOffsetFloats = 20480; + constexpr uint64_t typedOffset = typedOffsetFloats * sizeof(float) / sizeof(T); + constexpr uint64_t kgFp32Planes = 4; + uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; + uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + if (rowBegin >= rowEnd) { + return; + } + uint64_t maxRows = (typedOffsetFloats / kgFp32Planes) / K_; + if (maxRows > 32) { + maxRows = 32; + } + if (maxRows == 0) { + return; + } + + uint64_t last = start + curT - 1; + LocalTensor arena = vecBuf_.Get(); + LocalTensor gateLast = exp2Buf_.Get(); + LocalTensor typedLocal = vecBuf_.Get()[typedOffset]; + LoadAsFloatRow(gk_, KVOffset(b, hv, last, 0, K_), gateLast, K_); + + for (uint64_t tileRow = rowBegin; tileRow < rowEnd; tileRow += maxRows) { + uint64_t tileRows = rowEnd - tileRow; + if (tileRows > maxRows) { + tileRows = maxRows; + } + uint64_t elemCount = tileRows * K_; + uint64_t scratchBase = WScratchOffset(b, hv, chunkIdx, tileRow, 0); + uint64_t token = start + tileRow; + + DataCopy(arena, h_[scratchBase], static_cast(elemCount)); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + Cast(typedLocal, arena, RoundMode::CAST_RINT, static_cast(elemCount)); + PipeBarrier(); + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + DataCopy(w_[KVOffset(b, hv, token, 0, K_)], typedLocal, static_cast(elemCount)); + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + + LocalTensor kLocal = arena; + LocalTensor gLocal = arena[elemCount]; + LocalTensor expLocal = arena[2 * elemCount]; + LocalTensor outLocal = arena[3 * elemCount]; + CopyVectorIn(typedLocal, k_, QOffset(b, h, token, 0), elemCount); + CopyVectorIn(gLocal, gk_, KVOffset(b, hv, token, 0, K_), elemCount); + SetFlag(KDA_MTE2_V_EVENT_ID); + WaitFlag(KDA_MTE2_V_EVENT_ID); + Cast(kLocal, typedLocal, RoundMode::CAST_NONE, static_cast(elemCount)); + PipeBarrier(); + + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expLocal[row * K_], gateLast, gLocal[row * K_], static_cast(K_)); + } + PipeBarrier(); + Muls(expLocal, expLocal, LN2, static_cast(elemCount)); + PipeBarrier(); + ClampExpInput(expLocal, static_cast(elemCount)); + Exp(expLocal, expLocal, static_cast(elemCount)); + PipeBarrier(); + Mul(outLocal, kLocal, expLocal, static_cast(elemCount)); + PipeBarrier(); + ClampFp32ToOutputType(outLocal, static_cast(elemCount)); + Cast(typedLocal, outLocal, RoundMode::CAST_RINT, static_cast(elemCount)); + PipeBarrier(); + SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); + CopyVectorOut(kg_, KVOffset(b, hv, token, 0, K_), typedLocal, elemCount); + SetFlag(KDA_MTE3_MTE2_EVENT_ID); + WaitFlag(KDA_MTE3_MTE2_EVENT_ID); + } + SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); + } + + + + + + __aicore__ inline void ComputeOutputCube(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT) + { + using ElementA = T; + using ElementB = T; + using ElementC = OUT_T; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using BlockMmad = Catlass::Gemm::Block::BlockMmadTla; + + Catlass::Arch::Resource resource; + BlockMmad blockMmad(resource); + + auto layoutQ = tla::MakeLayout(BT_, K_); + auto layoutH = tla::MakeLayout(K_, V_); + auto layoutO = tla::MakeLayout(BT_, V_); + for (uint64_t nOffset = 0; nOffset < V_; nOffset += 128) { + uint32_t curN = static_cast((V_ - nOffset) > 128 ? 128 : (V_ - nOffset)); + auto tensorH = tla::MakeTensor(stageH_[HOffset(b, hv, chunkIdx, 0, nOffset)], layoutH, + Catlass::Arch::PositionGM{}); + for (uint64_t mOffset = 0; mOffset < curT; mOffset += 64) { + uint32_t curM = static_cast((curT - mOffset) > 64 ? 64 : (curT - mOffset)); + Catlass::GemmCoord shapeQH{curM, curN, static_cast(K_)}; + auto tensorQ = tla::MakeTensor(stageQG_[KVOffset(b, hv, start + mOffset, 0, K_)], layoutQ, + Catlass::Arch::PositionGM{}); + auto tensorO = tla::MakeTensor(o_[KVOffset(b, hv, start + mOffset, nOffset, V_)], layoutO, + Catlass::Arch::PositionGM{}); + auto blockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(shapeQH.m(), shapeQH.k())); + auto blockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(shapeQH.k(), shapeQH.n())); + auto blockO = GetTile(tensorO, tla::MakeCoord(0, 0), tla::MakeShape(shapeQH.m(), shapeQH.n())); + blockMmad(blockQ, blockH, blockO, shapeQH); + PipeBarrier(); + } + } + + auto layoutAqk = tla::MakeLayout(BT_, BT_); + auto layoutV = tla::MakeLayout(BT_, V_); + for (uint64_t nOffset = 0; nOffset < V_; nOffset += 128) { + uint32_t curN = static_cast((V_ - nOffset) > 128 ? 128 : (V_ - nOffset)); + auto tensorVNew = tla::MakeTensor(stageVNew_[KVOffset(b, hv, start, nOffset, V_)], layoutV, + Catlass::Arch::PositionGM{}); + for (uint64_t mOffset = 0; mOffset < curT; mOffset += 64) { + uint32_t curM = static_cast((curT - mOffset) > 64 ? 64 : (curT - mOffset)); + Catlass::GemmCoord shapeAV{curM, curN, static_cast(curT)}; + auto tensorAqk = tla::MakeTensor(stageAqk_[AOffset(b, hv, start + mOffset, 0)], layoutAqk, + Catlass::Arch::PositionGM{}); + auto tensorLocal = tla::MakeTensor(u_[KVOffset(b, hv, start + mOffset, nOffset, V_)], layoutO, + Catlass::Arch::PositionGM{}); + auto blockAqk = GetTile(tensorAqk, tla::MakeCoord(0, 0), tla::MakeShape(shapeAV.m(), shapeAV.k())); + auto blockVNew = GetTile(tensorVNew, tla::MakeCoord(0, 0), tla::MakeShape(shapeAV.k(), shapeAV.n())); + auto blockLocal = GetTile(tensorLocal, tla::MakeCoord(0, 0), tla::MakeShape(shapeAV.m(), shapeAV.n())); + blockMmad(blockAqk, blockVNew, blockLocal, shapeAV); + PipeBarrier(); + } + } + } + + __aicore__ inline void FinalizeOutputRows(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT, + uint64_t subBlockIdx, uint64_t subBlockNum) + { + LocalTensor stateLocal = VecScratch(0); + LocalTensor localLocal = VecScratch(1); + LocalTensor outLocal = VecScratch(2); + for (uint64_t i = subBlockIdx; i < curT; i += subBlockNum) { + uint64_t ti = start + i; + LoadAsFloatRow(o_, KVOffset(b, hv, ti, 0, V_), stateLocal, V_); + LoadAsFloatRow(u_, KVOffset(b, hv, ti, 0, V_), localLocal, V_); + Add(outLocal, stateLocal, localLocal, static_cast(V_)); + PipeBarrier(); + ClampFp32ToOutputType(outLocal, static_cast(V_)); + StoreFloatRow(vNew_, KVOffset(b, hv, ti, 0, V_), outLocal, V_); + } + } + + + __aicore__ inline bool ResolveFlatChunk(uint64_t task, uint64_t &seq, uint64_t &b, uint64_t &h, uint64_t &hv, + uint64_t &chunkIdx, uint64_t &start, uint64_t &end) + { + hv = task % HV_; + uint64_t flatChunk = task / HV_; + if (!isVarLen_) { + seq = flatChunk / NT_; + b = seq; + chunkIdx = flatChunk % NT_; + start = chunkIdx * BT_; + end = start + BT_; + if (end > T_) { + end = T_; + } + } else { + if (hasChunkIndices_) { + uint64_t low = 0; + uint64_t high = N_; + while (low + 1 < high) { + uint64_t mid = (low + high) >> 1; + if (flatChunk < static_cast(seqChunkOffset_[mid])) { + high = mid; + } else { + low = mid; + } + } + seq = low; + uint64_t localChunk = flatChunk - static_cast(seqChunkOffset_[seq]); + start = static_cast(seqStart_[seq]) + localChunk * BT_; + end = start + BT_; + uint64_t seqEnd = static_cast(seqEnd_[seq]); + if (end > seqEnd) { + end = seqEnd; + } + b = 0; + chunkIdx = flatChunk; + h = hv / (HV_ / H_); + return start < end; + } + return false; + } + h = hv / (HV_ / H_); + return start < end; + } + + __aicore__ inline void ProcessChunkPreAiv(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t end, uint64_t subBlockIdx, + uint64_t subBlockNum) + { + if constexpr (IsSameType::value) { + ProcessChunkPreAivFp32(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); + } + } + + template + __aicore__ inline void RunAicAfterBothAivReady(uint64_t subBlockIdx, uint64_t subBlockNum) + { + if constexpr (CORE_TYPE == AscendC::AIV) { + (void)subBlockIdx; + (void)subBlockNum; + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_MTE3>(syncReadyFlag_); + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(syncDoneFlag_); + } + } + + __aicore__ inline void ProcessChunkPreAivFp32(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t end, uint64_t subBlockIdx, + uint64_t subBlockNum) + { + uint64_t curT = end - start; + if (curT == 0) { + return; + } + if constexpr (IsSameType::value) { + return; + } + + if (K_ < 16) { + return; + } + bool usePostWuCube = UsePostWuCube(curT); + bool useAkkCubeSolve = UseAkkCubeSolve(curT); + uint64_t solveRowBegin = 0; + uint64_t solveRowEnd = 0; + GetSolveRowRange(BT_, subBlockIdx, subBlockNum, solveRowBegin, solveRowEnd); + uint64_t scoreBlockSize = ScoreRefBlockSize(); + uint64_t scoreBlockCount = (curT + scoreBlockSize - 1) / scoreBlockSize; + uint64_t pipelineBlockCount = + (scoreBlockCount + KDA_SCORE_QUEUE_DEPTH - 1) / KDA_SCORE_QUEUE_DEPTH * KDA_SCORE_QUEUE_DEPTH; + for (uint64_t block = 0; block < pipelineBlockCount; ++block) { + if (block < scoreBlockCount) { + uint64_t rowBegin = block * scoreBlockSize; + uint64_t rowCount = ScoreRowBlockCount(curT, rowBegin); + uint64_t refToken = ScoreRefToken(start, curT, rowBegin, rowCount); + PrepareGateProducts(b, h, hv, start, curT, subBlockIdx, subBlockNum, true, refToken, + rowBegin + rowCount, true, block % KDA_SCORE_QUEUE_DEPTH, + rowBegin, rowCount); + } + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_MTE3>(scoreReadyFlag_); + if (block > 0) { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(scoreDoneFlag_); + } + } + if (pipelineBlockCount > 0) { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(scoreDoneFlag_); + } + PrepareGateProducts(b, h, hv, start, curT, subBlockIdx, subBlockNum); + if (useAkkCubeSolve) { + bool fullChunk = curT == BT_; + PrepareAqkAkkSolveInputRows(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd, + fullChunk, !fullChunk); + } + if (useAkkCubeSolve) { + bool fullChunk = curT == BT_; + uint32_t solveIters = KDA_SOLVE_DIAG_MCH_ITERS; + RunAicAfterBothAivReady(subBlockIdx, subBlockNum); + for (uint32_t iter = 0; iter < solveIters; ++iter) { + AddSolveTmpToXDiagRows(b, hv, chunkIdx, start, solveRowBegin, solveRowEnd, + fullChunk && iter + 1 == solveIters); + RunAicAfterBothAivReady(subBlockIdx, subBlockNum); + } + if (!fullChunk) { + StoreSolveXRowsToAkk(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd); + } + } + // Host validation guarantees every accepted shape has enough workspace for this cube path. + PrepareWuCubeInputs(b, hv, start, curT, subBlockIdx, subBlockNum); + } + + __aicore__ inline void ProcessChunkPreAic(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + if constexpr (IsSameType::value) { + ProcessChunkPreAicFp32(b, hv, chunkIdx, start, end); + } + } + + __aicore__ inline void ProcessChunkPreAicFp32(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + uint64_t curT = end - start; + if (curT == 0 || K_ < 16) { + return; + } + uint64_t scoreBlockSize = ScoreRefBlockSize(); + uint64_t scoreBlockCount = (curT + scoreBlockSize - 1) / scoreBlockSize; + uint64_t pipelineBlockCount = + (scoreBlockCount + KDA_SCORE_QUEUE_DEPTH - 1) / KDA_SCORE_QUEUE_DEPTH * KDA_SCORE_QUEUE_DEPTH; + for (uint64_t block = 0; block < pipelineBlockCount; ++block) { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(scoreReadyFlag_); + if (block < scoreBlockCount) { + uint64_t rowBegin = block * scoreBlockSize; + uint64_t rowCount = ScoreRowBlockCount(curT, rowBegin); + ComputeRawAqkAkkCubeBlock(b, hv, start, curT, rowBegin, rowCount, true, + block % KDA_SCORE_QUEUE_DEPTH, rowBegin + rowCount); + } + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(scoreDoneFlag_); + } + bool usePostWuCube = UsePostWuCube(curT); + bool useAkkCubeSolve = UseAkkCubeSolve(curT); + if (useAkkCubeSolve) { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(syncReadyFlag_); + if (curT == BT_) { + ComputeAkkInverseMchFull(b, hv, chunkIdx, start); + } else { + ComputeAkkInverseMchTail(b, hv, chunkIdx, start, curT); + } + } + (void)usePostWuCube; + (void)chunkIdx; + } + + __aicore__ inline void ProcessChunkPostAiv(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t end, uint64_t subBlockIdx, + uint64_t subBlockNum) + { + uint64_t curT = end - start; + if (curT == 0 || !UsePostWuCube(curT)) { + return; + } + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(syncDoneFlag_); + CopyScratchWAndFinalizeKg(b, h, hv, chunkIdx, start, curT, subBlockIdx, subBlockNum); + } + + __aicore__ inline void ProcessChunkPostAic(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + if constexpr (IsSameType::value) { + ProcessChunkPostAicTyped(b, hv, chunkIdx, start, end); + } + } + + __aicore__ inline void ProcessChunkPostAicTyped(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + uint64_t curT = end - start; + if (curT == 0 || !UsePostWuCube(curT)) { + return; + } + ComputePostWuCube(b, hv, chunkIdx, start, curT); + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); + } + + __aicore__ inline void ProcessChunkOutAiv(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end, uint64_t subBlockIdx, uint64_t subBlockNum) + { + uint64_t curT = end - start; + if (curT == 0) { + return; + } + if constexpr (IsSameType::value) { + return; + } + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(syncDoneFlag_); + FinalizeOutputRows(b, hv, start, curT, subBlockIdx, subBlockNum); + } + + __aicore__ inline void ProcessChunkOutAic(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + uint64_t curT = end - start; + if (curT == 0) { + return; + } + ComputeOutputCube(b, hv, chunkIdx, start, curT); + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); + } + + __aicore__ inline void ProcessPreAiv() + { + if constexpr (IsSameType::value) { + isAivOnly_ = true; + } + uint64_t subBlockNum = isAivOnly_ ? 1 : static_cast(GetSubBlockNum()); + if (subBlockNum == 0) { + return; + } + uint64_t subBlockIdx = isAivOnly_ ? 0 : static_cast(GetSubBlockIdx()); + uint64_t coreNum = isAivOnly_ ? static_cast(GetBlockNum()) : usedCoreNum_; + uint64_t coreIdx = isAivOnly_ ? static_cast(GetBlockIdx()) : + static_cast(GetBlockIdx()) / subBlockNum; + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + for (uint64_t task = coreIdx; task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + ProcessChunkPreAiv(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); + } + } + } + + __aicore__ inline void ProcessPreAic() + { + if constexpr (IsSameType::value) { + return; + } + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + (void)h; + ProcessChunkPreAic(b, hv, chunkIdx, start, end); + } + } + } + + __aicore__ inline void ProcessPostAiv() + { + if constexpr (IsSameType::value) { + return; + } + uint64_t subBlockNum = static_cast(GetSubBlockNum()); + if (subBlockNum == 0) { + return; + } + uint64_t subBlockIdx = static_cast(GetSubBlockIdx()); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + uint64_t coreIdx = static_cast(GetBlockIdx()) / subBlockNum; + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + for (uint64_t task = coreIdx; task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + ProcessChunkPostAiv(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); + } + } + } + + __aicore__ inline void ProcessPostAic() + { + if constexpr (IsSameType::value) { + return; + } + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + (void)h; + ProcessChunkPostAic(b, hv, chunkIdx, start, end); + } + } + } + + __aicore__ inline void ProcessOutAiv() + { + if constexpr (IsSameType::value) { + return; + } + uint64_t subBlockNum = static_cast(GetSubBlockNum()); + if (subBlockNum == 0) { + return; + } + uint64_t subBlockIdx = static_cast(GetSubBlockIdx()); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + uint64_t coreIdx = static_cast(GetBlockIdx()) / subBlockNum; + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + for (uint64_t task = coreIdx; task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + (void)h; + (void)chunkIdx; + ProcessChunkOutAiv(b, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); + } + } + } + + __aicore__ inline void ProcessOutAic() + { + if constexpr (IsSameType::value) { + return; + } + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + (void)h; + ProcessChunkOutAic(b, hv, chunkIdx, start, end); + } + } + } + + + + +private: + GlobalTensor q_; + GlobalTensor k_; + GlobalTensor v_; + GlobalTensor gk_; + GlobalTensor beta_; + GlobalTensor initialState_; + GlobalTensor cuSeqlens_; + GlobalTensor o_; + GlobalTensor finalState_; + GlobalTensor aqk_; + GlobalTensor akk_; + GlobalTensor w_; + GlobalTensor u_; + GlobalTensor qg_; + GlobalTensor kg_; + GlobalTensor vNew_; + GlobalTensor h_; + GlobalTensor stageQG_; + GlobalTensor stageAqk_; + GlobalTensor stageVNew_; + GlobalTensor stageH_; + GlobalTensor solveWorkspace_; + GlobalTensor scoreWorkspace_; + TPipe *pipe_ = nullptr; + TBuf exp2Buf_; + TBuf vecBuf_; + TQue qInQue_; + TQue kInQue_; + TQue gInQue_; + TQue qgOutQue_; + TQue wOutQue_; + TQue kgOutQue_; + Catlass::Arch::CrossCoreFlagWithReverse scoreReadyFlag_{KDA_SCORE_READY_FLAG0, + KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse scoreDoneFlag_{KDA_SCORE_DONE_FLAG0, + KDA_SCORE_DONE_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse syncReadyFlag_{KDA_SCORE_READY_FLAG0, + KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse syncDoneFlag_{KDA_SCORE_DONE_FLAG0, + KDA_SCORE_DONE_FLAG1}; + uint64_t B_ = 0; + uint64_t N_ = 0; + uint64_t H_ = 0; + uint64_t HV_ = 0; + uint64_t T_ = 0; + uint64_t K_ = 0; + uint64_t V_ = 0; + uint64_t BT_ = 0; + uint64_t NT_ = 0; + float scale_ = 1.0f; + bool hasInitial_ = false; + bool isVarLen_ = false; + bool hasChunkIndices_ = false; + bool isAivOnly_ = false; + uint64_t usedCoreNum_ = 1; + uint64_t solveCoreIdx_ = 0; + int64_t stage_ = 0; + const int64_t *seqStart_ = nullptr; + const int64_t *seqEnd_ = nullptr; + const int64_t *seqChunkOffset_ = nullptr; +}; +} // namespace + +extern "C" __global__ __aicore__ void chunk_kda_fwd(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, + GM_ADDR initial_state, GM_ADDR cu_seqlens, + GM_ADDR chunk_indices, GM_ADDR stage_qg, GM_ADDR stage_aqk, + GM_ADDR stage_v_new, GM_ADDR stage_h, GM_ADDR o, + GM_ADDR final_state, GM_ADDR aqk, GM_ADDR akk, GM_ADDR w, + GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR v_new, GM_ADDR h, + GM_ADDR workspace, GM_ADDR tiling) +{ + GM_ADDR userWS = AscendC::GetUserWorkspace(workspace); + (void)userWS; + GET_TILING_DATA(tilingData, tiling); + TPipe pipe; + if (TILING_KEY_IS(0)) { + KERNEL_TASK_TYPE(0, KERNEL_TYPE_AIV_ONLY); + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, stage_v_new, + stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, &pipe); + op.ProcessAivOnly(); + } else if (TILING_KEY_IS(1)) { + KERNEL_TASK_TYPE(1, KERNEL_TYPE_MIX_AIC_1_2); + if (tilingData.dataType == 1) { + if ASCEND_IS_AIC { + if (tilingData.stage == 2) { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, + stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, + &pipe, false); + op.ProcessAic(); + } else if (tilingData.stage == 3) { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, + stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, + &pipe, false); + op.ProcessAic(); + } else { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, + stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, + &pipe, false); + op.ProcessAic(); + } + } + if ASCEND_IS_AIV { + if (tilingData.stage == 2) { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, + stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, + &pipe); + op.ProcessAiv(); + } else if (tilingData.stage == 3) { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, + stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, + &pipe); + op.ProcessAiv(); + } else { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, + stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, + &pipe); + op.ProcessAiv(); + } + } + } else { + if ASCEND_IS_AIC { + if (tilingData.stage == 2) { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, + stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, + &pipe, false); + op.ProcessAic(); + } else if (tilingData.stage == 3) { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, + stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, + &pipe, false); + op.ProcessAic(); + } else { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, + stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, + &pipe, false); + op.ProcessAic(); + } + } + if ASCEND_IS_AIV { + if (tilingData.stage == 2) { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, + stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, + &pipe); + op.ProcessAiv(); + } else if (tilingData.stage == 3) { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, + stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, + &pipe); + op.ProcessAiv(); + } else { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, + stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, + &pipe); + op.ProcessAiv(); + } + } + } + } else if (TILING_KEY_IS(2)) { + KERNEL_TASK_TYPE(2, KERNEL_TYPE_AIV_ONLY); + if (tilingData.dataType == 1) { + if (tilingData.stage == 3) { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, stage_v_new, + stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, &pipe); + op.ProcessAivOnly(); + } else { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, stage_v_new, + stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, &pipe); + op.ProcessAivOnly(); + } + } else { + if (tilingData.stage == 3) { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, stage_v_new, + stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, &pipe); + op.ProcessAivOnly(); + } else { + ChunkKdaFwdKernel op; + op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, stage_v_new, + stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, &pipe); + op.ProcessAivOnly(); + } + } + } +} diff --git a/csrc/attention/kda_gate_cumsum/CMakeLists.txt b/csrc/attention/kda_gate_cumsum/CMakeLists.txt new file mode 100644 index 000000000000..24deaa2180c5 --- /dev/null +++ b/csrc/attention/kda_gate_cumsum/CMakeLists.txt @@ -0,0 +1,12 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Tianjin University, Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# the BSD 3-Clause License (the "License"). Please refer to the License for details. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. +# ----------------------------------------------------------------------------------------------------------- +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/csrc/attention/kda_gate_cumsum/op_host/CMakeLists.txt b/csrc/attention/kda_gate_cumsum/op_host/CMakeLists.txt new file mode 100644 index 000000000000..c186d4731cc7 --- /dev/null +++ b/csrc/attention/kda_gate_cumsum/op_host/CMakeLists.txt @@ -0,0 +1,19 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Tianjin University, Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# the BSD 3-Clause License (the "License"). Please refer to the License for details. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. +# ----------------------------------------------------------------------------------------------------------- +add_op_to_compiled_list() +if (BUILD_OPEN_PROJECT) + target_sources(op_host_aclnnExc PRIVATE + kda_gate_cumsum_def.cpp + ) +endif() + +add_modules_sources(OPTYPE kda_gate_cumsum ACLNNTYPE aclnn_exclude) +add_ops_compile_options( + OP_NAME KdaGateCumsum + OPTIONS --cce-auto-sync=off + -Wno-deprecated-declarations +) diff --git a/csrc/attention/kda_gate_cumsum/op_host/kda_gate_cumsum_def.cpp b/csrc/attention/kda_gate_cumsum/op_host/kda_gate_cumsum_def.cpp new file mode 100644 index 000000000000..e1f922f8231f --- /dev/null +++ b/csrc/attention/kda_gate_cumsum/op_host/kda_gate_cumsum_def.cpp @@ -0,0 +1,60 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#include "register/op_def_registry.h" + +namespace ops { +class KdaGateCumsum : public OpDef { +public: + explicit KdaGateCumsum(const char *name) : OpDef(name) + { + const std::initializer_list gateTypes = { + ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16 + }; + const std::initializer_list floatTypes = { + ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT + }; + const std::initializer_list intTypes = { + ge::DT_INT64, ge::DT_INT64, ge::DT_INT64 + }; + const std::initializer_list formats = { + ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND + }; + + this->Input("g").ParamType(REQUIRED).DataType(gateTypes).Format(formats).UnknownShapeFormat(formats); + this->Input("a_log").ParamType(OPTIONAL).DataType(floatTypes).Format(formats).UnknownShapeFormat(formats); + this->Input("dt_bias").ParamType(OPTIONAL).DataType(floatTypes).Format(formats).UnknownShapeFormat(formats); + this->Input("cu_seqlens").ParamType(OPTIONAL).ValueDepend(OPTIONAL) + .DataType(intTypes).Format(formats).UnknownShapeFormat(formats); + + this->Output("gk").ParamType(REQUIRED).DataType(floatTypes).Format(formats).UnknownShapeFormat(formats); + + this->Attr("chunk_size").AttrType(REQUIRED).Int(64); + this->Attr("use_gate_in_kernel").AttrType(REQUIRED).Bool(false); + this->Attr("safe_gate").AttrType(REQUIRED).Bool(false); + this->Attr("lower_bound").AttrType(REQUIRED).Float(-5.0); + this->Attr("layout").AttrType(OPTIONAL).String("BSND"); + + OpAICoreConfig aicoreConfig; + aicoreConfig.DynamicCompileStaticFlag(true) + .DynamicFormatFlag(true) + .DynamicRankSupportFlag(true) + .DynamicShapeSupportFlag(true) + .NeedCheckSupportFlag(false) + .PrecisionReduceFlag(true) + .ExtendCfgInfo("prebuildPattern.value", "Opaque") + .ExtendCfgInfo("coreType.value", "AiCore") + .ExtendCfgInfo("aclnnSupport.value", "support_aclnn"); + + this->AICore().AddConfig("ascend910b", aicoreConfig); + this->AICore().AddConfig("ascend910_93", aicoreConfig); + this->AICore().AddConfig("ascend950", aicoreConfig); + } +}; + +OP_ADD(KdaGateCumsum); +} // namespace ops diff --git a/csrc/attention/kda_gate_cumsum/op_host/kda_gate_cumsum_tiling.cpp b/csrc/attention/kda_gate_cumsum/op_host/kda_gate_cumsum_tiling.cpp new file mode 100644 index 000000000000..2afcd9565431 --- /dev/null +++ b/csrc/attention/kda_gate_cumsum/op_host/kda_gate_cumsum_tiling.cpp @@ -0,0 +1,136 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#include "kda_gate_cumsum_tiling.h" +#include +#include +#include +#include "tiling/platform/platform_ascendc.h" + +namespace optiling { +namespace { +constexpr size_t INPUT_G_IDX = 0; +constexpr size_t INPUT_A_LOG_IDX = 1; +constexpr size_t INPUT_DT_BIAS_IDX = 2; +constexpr size_t INPUT_CU_SEQLENS_IDX = 3; +constexpr size_t ATTR_CHUNK_SIZE_IDX = 0; +constexpr size_t ATTR_USE_GATE_IDX = 1; +constexpr size_t ATTR_SAFE_GATE_IDX = 2; +constexpr size_t ATTR_LOWER_BOUND_IDX = 3; +constexpr size_t ATTR_LAYOUT_IDX = 4; +constexpr int64_t MAX_K_DIM = 256; + +enum class KdaGateLayout : int64_t { + BSND = 0, + BNSD = 1, + TND = 2, + NTD = 3, +}; +} // namespace + +ge::graphStatus Tiling4KdaGateCumsum(gert::TilingContext *context) +{ + KdaGateCumsumTilingData tiling; + auto gShape = context->GetOptionalInputShape(INPUT_G_IDX)->GetStorageShape(); + auto gDesc = context->GetInputDesc(INPUT_G_IDX); + if (gDesc == nullptr || (gShape.GetDimNum() != 3 && gShape.GetDimNum() != 4)) { + return ge::GRAPH_FAILED; + } + + int64_t rank = static_cast(gShape.GetDimNum()); + const char *layoutAttr = context->GetAttrs()->GetAttrPointer(ATTR_LAYOUT_IDX); + if (layoutAttr == nullptr) { + return ge::GRAPH_FAILED; + } + KdaGateLayout layout; + if (std::strcmp(layoutAttr, "BSND") == 0) { + layout = KdaGateLayout::BSND; + } else if (std::strcmp(layoutAttr, "BNSD") == 0) { + layout = KdaGateLayout::BNSD; + } else if (std::strcmp(layoutAttr, "TND") == 0) { + layout = KdaGateLayout::TND; + } else if (std::strcmp(layoutAttr, "NTD") == 0) { + layout = KdaGateLayout::NTD; + } else { + return ge::GRAPH_FAILED; + } + if ((rank == 4 && (layout == KdaGateLayout::TND || layout == KdaGateLayout::NTD)) || + (rank == 3 && (layout == KdaGateLayout::BSND || layout == KdaGateLayout::BNSD))) { + return ge::GRAPH_FAILED; + } + int64_t batch = (rank == 4) ? gShape.GetDim(0) : 1; + int64_t t = (layout == KdaGateLayout::BNSD) ? gShape.GetDim(2) : + ((layout == KdaGateLayout::NTD) ? gShape.GetDim(1) : + ((rank == 4) ? gShape.GetDim(1) : gShape.GetDim(0))); + int64_t hv = (layout == KdaGateLayout::BNSD) ? gShape.GetDim(1) : + ((layout == KdaGateLayout::NTD) ? gShape.GetDim(0) : + ((rank == 4) ? gShape.GetDim(2) : gShape.GetDim(1))); + int64_t k = (rank == 4) ? gShape.GetDim(3) : gShape.GetDim(2); + if (k > MAX_K_DIM) { + return ge::GRAPH_FAILED; + } + + int64_t chunkSize = *context->GetAttrs()->GetAttrPointer(ATTR_CHUNK_SIZE_IDX); + bool useGate = *context->GetAttrs()->GetAttrPointer(ATTR_USE_GATE_IDX); + bool safeGate = *context->GetAttrs()->GetAttrPointer(ATTR_SAFE_GATE_IDX); + float lowerBound = *context->GetAttrs()->GetAttrPointer(ATTR_LOWER_BOUND_IDX); + + const auto cuShape = context->GetOptionalInputShape(INPUT_CU_SEQLENS_IDX); + int64_t hasCuSeqlens = (cuShape != nullptr) ? 1 : 0; + int64_t seqNum = hasCuSeqlens ? (cuShape->GetStorageShape().GetDim(0) - 1) : batch; + int64_t maxChunks = (t + chunkSize - 1) / chunkSize; + // Dense input keeps chunk-level parallelism. Varlen owns one (sequence, head) pair and + // iterates only that sequence's real chunks, avoiding a rectangular grid of empty tasks. + int64_t taskCount = hasCuSeqlens ? seqNum * hv : batch * hv * maxChunks; + + const auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo()); + uint32_t coreNum = ascendcPlatform.GetCoreNumAiv(); + uint32_t blockDim = static_cast(std::min(taskCount, coreNum)); + context->SetBlockDim(blockDim == 0 ? 1 : blockDim); + + size_t *workspace = context->GetWorkspaceSizes(1); + workspace[0] = ascendcPlatform.GetLibApiWorkSpaceSize(); + + tiling.set_batch(batch); + tiling.set_t(t); + tiling.set_hv(hv); + tiling.set_k(k); + tiling.set_rank(rank); + tiling.set_layout(static_cast(layout)); + tiling.set_chunkSize(chunkSize); + tiling.set_seqNum(seqNum); + tiling.set_hasCuSeqlens(hasCuSeqlens); + tiling.set_hasALog(context->GetOptionalInputDesc(INPUT_A_LOG_IDX) != nullptr ? 1 : 0); + tiling.set_hasDtBias(context->GetOptionalInputDesc(INPUT_DT_BIAS_IDX) != nullptr ? 1 : 0); + int64_t dataType = 0; + if (gDesc->GetDataType() == ge::DT_FLOAT) { + dataType = 2; + } else if (gDesc->GetDataType() == ge::DT_BF16) { + dataType = 1; + } + tiling.set_dataType(dataType); + tiling.set_useGateInKernel(useGate ? 1 : 0); + tiling.set_safeGate(safeGate ? 1 : 0); + tiling.set_lowerBound(lowerBound); + tiling.set_usedCoreNum(blockDim == 0 ? 1 : blockDim); + + tiling.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity()); + context->GetRawTilingData()->SetDataSize(tiling.GetDataSize()); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus TilingPrepare4KdaGateCumsum(gert::TilingParseContext *context) +{ + (void)context; + return ge::GRAPH_SUCCESS; +} + +IMPL_OP_OPTILING(KdaGateCumsum) + .Tiling(Tiling4KdaGateCumsum) + .TilingParse(TilingPrepare4KdaGateCumsum); + +} // namespace optiling diff --git a/csrc/attention/kda_gate_cumsum/op_host/kda_gate_cumsum_tiling.h b/csrc/attention/kda_gate_cumsum/op_host/kda_gate_cumsum_tiling.h new file mode 100644 index 000000000000..c8e9be4964fd --- /dev/null +++ b/csrc/attention/kda_gate_cumsum/op_host/kda_gate_cumsum_tiling.h @@ -0,0 +1,37 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#pragma once + +#include +#include + +namespace optiling { + +BEGIN_TILING_DATA_DEF(KdaGateCumsumTilingData) +TILING_DATA_FIELD_DEF(int64_t, batch); +TILING_DATA_FIELD_DEF(int64_t, t); +TILING_DATA_FIELD_DEF(int64_t, hv); +TILING_DATA_FIELD_DEF(int64_t, k); +TILING_DATA_FIELD_DEF(int64_t, rank); +TILING_DATA_FIELD_DEF(int64_t, layout); +TILING_DATA_FIELD_DEF(int64_t, chunkSize); +TILING_DATA_FIELD_DEF(int64_t, seqNum); +TILING_DATA_FIELD_DEF(int64_t, hasCuSeqlens); +TILING_DATA_FIELD_DEF(int64_t, hasALog); +TILING_DATA_FIELD_DEF(int64_t, hasDtBias); +TILING_DATA_FIELD_DEF(int64_t, dataType); +TILING_DATA_FIELD_DEF(int64_t, useGateInKernel); +TILING_DATA_FIELD_DEF(int64_t, safeGate); +TILING_DATA_FIELD_DEF(float, lowerBound); +TILING_DATA_FIELD_DEF(int64_t, usedCoreNum); +END_TILING_DATA_DEF; + +REGISTER_TILING_DATA_CLASS(KdaGateCumsum, KdaGateCumsumTilingData) + +struct KdaGateCumsumCompileInfo {}; +} // namespace optiling diff --git a/csrc/attention/kda_gate_cumsum/op_host/op_api/aclnn_kda_gate_cumsum.cpp b/csrc/attention/kda_gate_cumsum/op_host/op_api/aclnn_kda_gate_cumsum.cpp new file mode 100644 index 000000000000..67749edfaf00 --- /dev/null +++ b/csrc/attention/kda_gate_cumsum/op_host/op_api/aclnn_kda_gate_cumsum.cpp @@ -0,0 +1,232 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#include "aclnn_kda_gate_cumsum.h" +#include "kda_gate_cumsum.h" + +#include "acl/acl.h" +#include "aclnn/aclnn_base.h" +#include "aclnn_kernels/common/op_error_check.h" +#include "aclnn_kernels/contiguous.h" +#include "opdev/make_op_executor.h" +#include "opdev/op_dfx.h" +#include "opdev/op_executor.h" +#include "opdev/op_log.h" +#include "opdev/tensor_view_utils.h" + +#include + +using namespace op; + +#ifdef __cplusplus +extern "C" { +#endif + +namespace { +constexpr int64_t MAX_KDA_K_DIM = 256; + +enum class KdaGateLayout : int64_t { + BSND = 0, + BNSD = 1, + TND = 2, + NTD = 3, +}; + +aclnnStatus KdaGateDataContiguous(const aclTensor *&tensor, aclOpExecutor *executor) +{ + if (tensor == nullptr) { + return ACLNN_SUCCESS; + } + tensor = l0op::Contiguous(tensor, executor); + CHECK_RET(tensor != nullptr, ACLNN_ERR_INNER_NULLPTR); + return ACLNN_SUCCESS; +} + +int64_t KdaGateDim(const aclTensor *tensor, size_t idx) +{ + return tensor->GetViewShape().GetDim(idx); +} + +size_t KdaGateRank(const aclTensor *tensor) +{ + return tensor->GetViewShape().GetDimNum(); +} + +bool KdaGateSameShape(const aclTensor *lhs, const aclTensor *rhs) +{ + size_t rank = KdaGateRank(lhs); + if (rank != KdaGateRank(rhs)) { + return false; + } + for (size_t idx = 0; idx < rank; ++idx) { + if (KdaGateDim(lhs, idx) != KdaGateDim(rhs, idx)) { + return false; + } + } + return true; +} + +aclnnStatus ParseKdaGateLayout(const char *layout, KdaGateLayout &parsed) +{ + CHECK_COND(layout != nullptr, ACLNN_ERR_PARAM_INVALID, + "layout must be one of BSND, BNSD, TND or NTD."); + if (std::strcmp(layout, "BSND") == 0) { + parsed = KdaGateLayout::BSND; + return ACLNN_SUCCESS; + } + if (std::strcmp(layout, "BNSD") == 0) { + parsed = KdaGateLayout::BNSD; + return ACLNN_SUCCESS; + } + if (std::strcmp(layout, "TND") == 0) { + parsed = KdaGateLayout::TND; + return ACLNN_SUCCESS; + } + if (std::strcmp(layout, "NTD") == 0) { + parsed = KdaGateLayout::NTD; + return ACLNN_SUCCESS; + } + CHECK_COND(false, ACLNN_ERR_PARAM_INVALID, + "layout must be uppercase and one of BSND, BNSD, TND or NTD."); + return ACLNN_ERR_PARAM_INVALID; +} + +int64_t KdaGateSeqLen(const aclTensor *g, KdaGateLayout layout) +{ + if (layout == KdaGateLayout::TND) { + return KdaGateDim(g, 0); + } + if (layout == KdaGateLayout::NTD) { + return KdaGateDim(g, 1); + } + return layout == KdaGateLayout::BNSD ? KdaGateDim(g, 2) : KdaGateDim(g, 1); +} + +aclnnStatus KdaGateCheckCuSeqlens(const aclIntArray *cuSeqlensOptional, int64_t seqlen) +{ + if (cuSeqlensOptional == nullptr) { + return ACLNN_SUCCESS; + } + const aclIntArray &cu = *cuSeqlensOptional; + CHECK_COND(cu.Size() >= 2, ACLNN_ERR_PARAM_INVALID, + "cuSeqlensOptional must contain at least [0, total_tokens]."); + CHECK_COND(cu[0] == 0, ACLNN_ERR_PARAM_INVALID, "cuSeqlensOptional[0] must be 0."); + CHECK_COND(cu[cu.Size() - 1] == seqlen, ACLNN_ERR_PARAM_INVALID, + "cuSeqlensOptional last element must equal the sequence length."); + for (size_t idx = 0; idx + 1 < cu.Size(); ++idx) { + CHECK_COND(cu[idx] <= cu[idx + 1], ACLNN_ERR_PARAM_INVALID, + "cuSeqlensOptional must be nondecreasing."); + } + return ACLNN_SUCCESS; +} + +aclnnStatus KdaGateCheckParams( + const aclTensor *g, + const aclTensor *aLogOptional, + const aclTensor *dtBiasOptional, + const aclIntArray *cuSeqlensOptional, + int64_t chunkSize, + bool useGateInKernel, + bool safeGate, + double lowerBound, + const char *layoutText, + const aclTensor *gkOut) +{ + CHECK_COND(g != nullptr, ACLNN_ERR_PARAM_NULLPTR, "g must not be nullptr."); + CHECK_COND(gkOut != nullptr, ACLNN_ERR_PARAM_NULLPTR, "gkOut must not be nullptr."); + CHECK_COND(chunkSize == 32 || chunkSize == 64 || chunkSize == 128, ACLNN_ERR_PARAM_INVALID, + "chunkSize must be 32, 64 or 128."); + size_t rank = KdaGateRank(g); + CHECK_COND(rank == 3 || rank == 4, ACLNN_ERR_PARAM_INVALID, + "g must be BSND/BNSD rank4 or TND/NTD rank3."); + size_t kDimIdx = rank == 4 ? 3 : 2; + CHECK_COND(KdaGateDim(g, kDimIdx) <= MAX_KDA_K_DIM, ACLNN_ERR_PARAM_INVALID, + "K must be less than or equal to 256."); + CHECK_COND(gkOut->GetDataType() == DataType::DT_FLOAT, ACLNN_ERR_PARAM_INVALID, + "gkOut must be float32."); + CHECK_COND(KdaGateSameShape(gkOut, g), ACLNN_ERR_PARAM_INVALID, "gkOut shape must match g shape."); + KdaGateLayout layout = KdaGateLayout::BSND; + CHECK_RET(ParseKdaGateLayout(layoutText, layout) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + bool rankMatchesLayout = (rank == 4 && (layout == KdaGateLayout::BSND || layout == KdaGateLayout::BNSD)) || + (rank == 3 && (layout == KdaGateLayout::TND || layout == KdaGateLayout::NTD)); + CHECK_COND(rankMatchesLayout, ACLNN_ERR_PARAM_INVALID, + "layout rank does not match g rank: BSND/BNSD require rank 4 and TND/NTD require rank 3."); + CHECK_RET(KdaGateCheckCuSeqlens(cuSeqlensOptional, KdaGateSeqLen(g, layout)) == ACLNN_SUCCESS, + ACLNN_ERR_PARAM_INVALID); + CHECK_COND(cuSeqlensOptional == nullptr || rank == 3 || KdaGateDim(g, 0) == 1, ACLNN_ERR_PARAM_INVALID, + "rank4 varlen input with cuSeqlensOptional currently requires B=1."); + int64_t hv = (layout == KdaGateLayout::BNSD) ? KdaGateDim(g, 1) : + ((layout == KdaGateLayout::NTD) ? KdaGateDim(g, 0) : + ((rank == 4) ? KdaGateDim(g, 2) : KdaGateDim(g, 1))); + if (useGateInKernel) { + CHECK_COND(aLogOptional != nullptr, ACLNN_ERR_PARAM_NULLPTR, + "aLogOptional must be provided when useGateInKernel is true."); + CHECK_COND(safeGate, ACLNN_ERR_PARAM_INVALID, + "Only safe_gate raw-gate path is supported; set safeGate=true with lowerBound."); + CHECK_COND(lowerBound >= -5.0 && lowerBound < 0.0, ACLNN_ERR_PARAM_INVALID, + "lowerBound must be in [-5, 0)."); + int64_t k = KdaGateDim(g, kDimIdx); + CHECK_COND(KdaGateRank(aLogOptional) == 1 && KdaGateDim(aLogOptional, 0) == hv, + ACLNN_ERR_PARAM_INVALID, "aLogOptional shape must be [HV]."); + if (dtBiasOptional != nullptr) { + size_t biasRank = KdaGateRank(dtBiasOptional); + bool validBias = (biasRank == 1 && KdaGateDim(dtBiasOptional, 0) == hv * k) || + (biasRank == 2 && KdaGateDim(dtBiasOptional, 0) == hv && + KdaGateDim(dtBiasOptional, 1) == k); + CHECK_COND(validBias, ACLNN_ERR_PARAM_INVALID, "dtBiasOptional shape must be [HV*K] or [HV, K]."); + } + } else { + CHECK_COND(!safeGate, ACLNN_ERR_PARAM_INVALID, + "safeGate only takes effect when useGateInKernel is true."); + } + return ACLNN_SUCCESS; +} +} // namespace + +aclnnStatus aclnnKdaGateCumsumGetWorkspaceSize( + const aclTensor *g, + const aclTensor *aLogOptional, + const aclTensor *dtBiasOptional, + const aclIntArray *cuSeqlensOptional, + int64_t chunkSize, + bool useGateInKernel, + bool safeGate, + double lowerBound, + const char *layout, + const aclTensor *gkOut, + uint64_t *workspaceSize, + aclOpExecutor **executor) +{ + L2_DFX_PHASE_1(aclnnKdaGateCumsum, DFX_IN(g, aLogOptional, dtBiasOptional, cuSeqlensOptional), DFX_OUT(gkOut)); + auto uniqueExecutor = CREATE_EXECUTOR(); + CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); + auto executorPtr = uniqueExecutor.get(); + CHECK_RET(KdaGateCheckParams(g, aLogOptional, dtBiasOptional, cuSeqlensOptional, chunkSize, useGateInKernel, + safeGate, lowerBound, layout, gkOut) == ACLNN_SUCCESS, + ACLNN_ERR_PARAM_INVALID); + CHECK_RET(KdaGateDataContiguous(g, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(KdaGateDataContiguous(aLogOptional, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(KdaGateDataContiguous(dtBiasOptional, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + + auto result = l0op::KdaGateCumsum(g, aLogOptional, dtBiasOptional, cuSeqlensOptional, chunkSize, + useGateInKernel, safeGate, lowerBound, layout, gkOut, executorPtr); + CHECK_RET(result[0] != nullptr, ACLNN_ERR_INNER_NULLPTR); + + *workspaceSize = uniqueExecutor->GetWorkspaceSize(); + uniqueExecutor.ReleaseTo(executor); + return ACLNN_SUCCESS; +} + +aclnnStatus aclnnKdaGateCumsum(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream) +{ + L2_DFX_PHASE_2(aclnnKdaGateCumsum); + return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); +} + +#ifdef __cplusplus +} +#endif diff --git a/csrc/attention/kda_gate_cumsum/op_host/op_api/aclnn_kda_gate_cumsum.h b/csrc/attention/kda_gate_cumsum/op_host/op_api/aclnn_kda_gate_cumsum.h new file mode 100644 index 000000000000..6a4386a306b7 --- /dev/null +++ b/csrc/attention/kda_gate_cumsum/op_host/op_api/aclnn_kda_gate_cumsum.h @@ -0,0 +1,38 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#ifndef OP_API_INC_ACLNN_KDA_GATE_CUMSUM_H +#define OP_API_INC_ACLNN_KDA_GATE_CUMSUM_H + +#include "aclnn/aclnn_base.h" +#include "aclnn_util.h" + +#ifdef __cplusplus +extern "C" { +#endif + +aclnnStatus aclnnKdaGateCumsumGetWorkspaceSize( + const aclTensor *g, + const aclTensor *aLogOptional, + const aclTensor *dtBiasOptional, + const aclIntArray *cuSeqlensOptional, + int64_t chunkSize, + bool useGateInKernel, + bool safeGate, + double lowerBound, + const char *layout, + const aclTensor *gkOut, + uint64_t *workspaceSize, + aclOpExecutor **executor); + +aclnnStatus aclnnKdaGateCumsum(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream); + +#ifdef __cplusplus +} +#endif + +#endif diff --git a/csrc/attention/kda_gate_cumsum/op_host/op_api/kda_gate_cumsum.cpp b/csrc/attention/kda_gate_cumsum/op_host/op_api/kda_gate_cumsum.cpp new file mode 100644 index 000000000000..bc39bfc909fd --- /dev/null +++ b/csrc/attention/kda_gate_cumsum/op_host/op_api/kda_gate_cumsum.cpp @@ -0,0 +1,54 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#include "kda_gate_cumsum.h" + +#include "opdev/make_op_executor.h" +#include "opdev/op_dfx.h" +#include "opdev/op_log.h" + +using namespace op; + +namespace l0op { +OP_TYPE_REGISTER(KdaGateCumsum); + +const std::array KdaGateCumsum( + const aclTensor *g, + const aclTensor *aLogOptional, + const aclTensor *dtBiasOptional, + const aclIntArray *cuSeqlensOptional, + int64_t chunkSize, + bool useGateInKernel, + bool safeGate, + double lowerBound, + const char *layout, + const aclTensor *gkOut, + aclOpExecutor *executor) +{ + L0_DFX(KdaGateCumsum, g, aLogOptional, dtBiasOptional, cuSeqlensOptional, chunkSize, useGateInKernel, + safeGate, lowerBound, layout, gkOut); + + const aclTensor *actualCuSeqlens = nullptr; + if (cuSeqlensOptional != nullptr) { + actualCuSeqlens = executor->ConvertToTensor(cuSeqlensOptional, DataType::DT_INT64); + const_cast(actualCuSeqlens)->SetStorageFormat(Format::FORMAT_ND); + const_cast(actualCuSeqlens)->SetViewFormat(Format::FORMAT_ND); + const_cast(actualCuSeqlens)->SetOriginalFormat(Format::FORMAT_ND); + } + + auto ret = ADD_TO_LAUNCHER_LIST_AICORE( + KdaGateCumsum, + OP_INPUT(g, aLogOptional, dtBiasOptional, actualCuSeqlens), + OP_OUTPUT(gkOut), + OP_ATTR(chunkSize, useGateInKernel, safeGate, static_cast(lowerBound), layout)); + if (ret != ACLNN_SUCCESS) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "ADD_TO_LAUNCHER_LIST_AICORE KdaGateCumsum failed."); + return {nullptr}; + } + return {gkOut}; +} +} // namespace l0op diff --git a/csrc/attention/kda_gate_cumsum/op_host/op_api/kda_gate_cumsum.h b/csrc/attention/kda_gate_cumsum/op_host/op_api/kda_gate_cumsum.h new file mode 100644 index 000000000000..bf4d53200a68 --- /dev/null +++ b/csrc/attention/kda_gate_cumsum/op_host/op_api/kda_gate_cumsum.h @@ -0,0 +1,26 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#pragma once + +#include "aclnn/aclnn_base.h" +#include + +namespace l0op { +const std::array KdaGateCumsum( + const aclTensor *g, + const aclTensor *aLogOptional, + const aclTensor *dtBiasOptional, + const aclIntArray *cuSeqlensOptional, + int64_t chunkSize, + bool useGateInKernel, + bool safeGate, + double lowerBound, + const char *layout, + const aclTensor *gkOut, + aclOpExecutor *executor); +} // namespace l0op diff --git a/csrc/attention/kda_gate_cumsum/op_kernel/kda_gate_cumsum.cpp b/csrc/attention/kda_gate_cumsum/op_kernel/kda_gate_cumsum.cpp new file mode 100644 index 000000000000..23bce8fd1c4d --- /dev/null +++ b/csrc/attention/kda_gate_cumsum/op_kernel/kda_gate_cumsum.cpp @@ -0,0 +1,330 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#include "kernel_operator.h" + +using namespace AscendC; + +namespace { +constexpr float RCP_LN2 = 1.4426950408889634f; +constexpr uint32_t GATE_MTE2_V_EVENT_ID = 0; +constexpr uint32_t GATE_V_MTE3_EVENT_ID = 1; +constexpr uint32_t GATE_MTE3_MTE2_EVENT_ID = 2; +constexpr uint32_t GATE_SCALAR_MTE2_V_EVENT_ID = 3; +constexpr uint32_t GATE_SCALAR_V_S_EVENT_ID = 4; +constexpr uint32_t GATE_MTE3_V_EVENT_ID = 5; +constexpr uint32_t GATE_ROW_ELEMENTS = 256; + +template +class KdaGateCumsumKernel { +public: + __aicore__ inline void Init(GM_ADDR g, GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR cuSeqlens, GM_ADDR gk, + const KdaGateCumsumTilingData &tiling, TPipe *pipe) + { + g_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(g)); + aLog_.SetGlobalBuffer(reinterpret_cast<__gm__ float *>(aLog)); + dtBias_.SetGlobalBuffer(reinterpret_cast<__gm__ float *>(dtBias)); + cuSeqlens_.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(cuSeqlens)); + gk_.SetGlobalBuffer(reinterpret_cast<__gm__ float *>(gk)); + pipe_ = pipe; + batch_ = static_cast(tiling.batch); + t_ = static_cast(tiling.t); + hv_ = static_cast(tiling.hv); + k_ = static_cast(tiling.k); + rank_ = static_cast(tiling.rank); + layout_ = static_cast(tiling.layout); + chunkSize_ = static_cast(tiling.chunkSize); + seqNum_ = static_cast(tiling.seqNum); + hasCuSeqlens_ = tiling.hasCuSeqlens != 0; + hasALog_ = tiling.hasALog != 0; + hasDtBias_ = tiling.hasDtBias != 0; + lowerBound_ = tiling.lowerBound; + usedCoreNum_ = static_cast(tiling.usedCoreNum); + maxChunks_ = (t_ + chunkSize_ - 1) / chunkSize_; + pipe_->InitBuffer(rowBuf_, GATE_ROW_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(accBuf_, GATE_ROW_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(tmpBuf_, GATE_ROW_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(oneBuf_, GATE_ROW_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(inBuf_, GATE_ROW_ELEMENTS * sizeof(T)); + pipe_->InitBuffer(scalarBuf_, 32); + pipe_->InitBuffer(scalarI64Buf_, 32); + } + + __aicore__ inline void Process() + { + uint64_t taskCount = hasCuSeqlens_ ? seqNum_ * hv_ : batch_ * hv_ * maxChunks_; + uint64_t coreIdx = static_cast(GetBlockIdx()); + for (uint64_t task = coreIdx; task < taskCount; task += usedCoreNum_) { + ProcessTask(task); + } + } + +private: + __aicore__ inline uint64_t Offset(uint64_t b, uint64_t t, uint64_t hv, uint64_t k) const + { + if (layout_ == 1) { + return ((b * hv_ + hv) * t_ + t) * k_ + k; + } + if (layout_ == 3) { + return (hv * t_ + t) * k_ + k; + } + if (rank_ == 4) { + return ((b * t_ + t) * hv_ + hv) * k_ + k; + } + return (t * hv_ + hv) * k_ + k; + } + + __aicore__ inline void CopyVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(T)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst, src[offset], static_cast(count)); + } else { + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + } + + __aicore__ inline void CopyFloatVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, + uint64_t count) + { + uint64_t rowBytes = count * sizeof(float); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst, src[offset], static_cast(count)); + } else { + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + } + + __aicore__ inline void CopyFloatVectorOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src, + uint64_t count) + { + uint64_t rowBytes = count * sizeof(float); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst[offset], src, static_cast(count)); + } else { + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPad(dst[offset], src, params); + } + } + + __aicore__ inline void LoadGateRow(uint64_t offset, LocalTensor &row) + { + if constexpr (IsSameType::value) { + CopyVectorIn(row, g_, offset, k_); + } else { + LocalTensor inLocal = inBuf_.Get(); + CopyVectorIn(inLocal, g_, offset, k_); + SetFlag(GATE_MTE2_V_EVENT_ID); + WaitFlag(GATE_MTE2_V_EVENT_ID); + Cast(row, inLocal, RoundMode::CAST_NONE, static_cast(k_)); + PipeBarrier(); + return; + } + SetFlag(GATE_MTE2_V_EVENT_ID); + WaitFlag(GATE_MTE2_V_EVENT_ID); + Adds(row, row, 0.0f, static_cast(k_)); + PipeBarrier(); + } + + __aicore__ inline float ReadFloat(GlobalTensor &tensor, uint64_t offset) + { + LocalTensor scalar = scalarBuf_.Get(); + DataCopyParams params{1, static_cast(sizeof(float)), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(scalar, tensor[offset], params, padParams); + SetFlag(GATE_SCALAR_MTE2_V_EVENT_ID); + WaitFlag(GATE_SCALAR_MTE2_V_EVENT_ID); + Adds(scalar, scalar, 0.0f, 1); + PipeBarrier(); + SetFlag(GATE_SCALAR_V_S_EVENT_ID); + WaitFlag(GATE_SCALAR_V_S_EVENT_ID); + __ubuf__ float *ptr = (__ubuf__ float *)scalar.GetPhyAddr(); + return ptr[0]; + } + + __aicore__ inline int64_t ReadInt64(GlobalTensor &tensor, uint64_t offset) + { + LocalTensor scalar = scalarI64Buf_.Get(); + DataCopyParams params{1, static_cast(sizeof(int64_t)), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(scalar, tensor[offset], params, padParams); + SetFlag(GATE_SCALAR_MTE2_V_EVENT_ID); + WaitFlag(GATE_SCALAR_MTE2_V_EVENT_ID); + SetFlag(GATE_SCALAR_V_S_EVENT_ID); + WaitFlag(GATE_SCALAR_V_S_EVENT_ID); + __ubuf__ int64_t *ptr = (__ubuf__ int64_t *)scalar.GetPhyAddr(); + return ptr[0]; + } + + __aicore__ inline float ExpScalar(float x) + { + LocalTensor scalar = scalarBuf_.Get(); + Duplicate(scalar, x, 1); + PipeBarrier(); + Exp(scalar, scalar, 1); + PipeBarrier(); + SetFlag(GATE_SCALAR_V_S_EVENT_ID); + WaitFlag(GATE_SCALAR_V_S_EVENT_ID); + __ubuf__ float *ptr = (__ubuf__ float *)scalar.GetPhyAddr(); + return ptr[0]; + } + + __aicore__ inline void ApplyGate(uint64_t hv, LocalTensor &row) + { + if constexpr (SAFE_GATE) { + if (hasDtBias_) { + LocalTensor tmp = tmpBuf_.Get(); + CopyFloatVectorIn(tmp, dtBias_, hv * k_, k_); + SetFlag(GATE_MTE2_V_EVENT_ID); + WaitFlag(GATE_MTE2_V_EVENT_ID); + Add(row, row, tmp, static_cast(k_)); + PipeBarrier(); + } + + float expA = hasALog_ ? ExpScalar(ReadFloat(aLog_, hv)) : 1.0f; + Muls(row, row, expA, static_cast(k_)); + PipeBarrier(); + + LocalTensor tmp = tmpBuf_.Get(); + Muls(tmp, row, -1.0f, static_cast(k_)); + PipeBarrier(); + Exp(tmp, tmp, static_cast(k_)); + PipeBarrier(); + Adds(tmp, tmp, 1.0f, static_cast(k_)); + PipeBarrier(); + + LocalTensor one = oneBuf_.Get(); + Duplicate(one, 1.0f, static_cast(k_)); + PipeBarrier(); + Div(row, one, tmp, static_cast(k_)); + PipeBarrier(); + Muls(row, row, lowerBound_, static_cast(k_)); + PipeBarrier(); + } + } + + __aicore__ inline void ProcessTask(uint64_t task) + { + if (!hasCuSeqlens_) { + uint64_t chunk = task % maxChunks_; + uint64_t hv = (task / maxChunks_) % hv_; + uint64_t b = task / (maxChunks_ * hv_); + uint64_t start = chunk * chunkSize_; + uint64_t end = start + chunkSize_; + if (end > t_) { + end = t_; + } + ProcessChunk(b, hv, start, end); + return; + } + uint64_t hv = task % hv_; + uint64_t seq = task / hv_; + uint64_t seqStart = static_cast(ReadInt64(cuSeqlens_, seq)); + uint64_t seqEnd = static_cast(ReadInt64(cuSeqlens_, seq + 1)); + for (uint64_t start = seqStart; start < seqEnd; start += chunkSize_) { + uint64_t end = start + chunkSize_; + if (end > seqEnd) { + end = seqEnd; + } + ProcessChunk(0, hv, start, end); + } + } + + __aicore__ inline void ProcessChunk(uint64_t b, uint64_t hv, uint64_t start, uint64_t end) + { + LocalTensor acc = accBuf_.Get(); + LocalTensor row = rowBuf_.Get(); + Duplicate(acc, 0.0f, static_cast(k_)); + PipeBarrier(); + for (uint64_t t = start; t < end; ++t) { + LoadGateRow(Offset(b, t, hv, 0), row); + ApplyGate(hv, row); + Muls(row, row, RCP_LN2, static_cast(k_)); + PipeBarrier(); + Add(acc, acc, row, static_cast(k_)); + PipeBarrier(); + SetFlag(GATE_V_MTE3_EVENT_ID); + WaitFlag(GATE_V_MTE3_EVENT_ID); + CopyFloatVectorOut(gk_, Offset(b, t, hv, 0), acc, k_); + SetFlag(GATE_MTE3_MTE2_EVENT_ID); + WaitFlag(GATE_MTE3_MTE2_EVENT_ID); + SetFlag(GATE_MTE3_V_EVENT_ID); + WaitFlag(GATE_MTE3_V_EVENT_ID); + } + } + + GlobalTensor g_; + GlobalTensor aLog_; + GlobalTensor dtBias_; + GlobalTensor cuSeqlens_; + GlobalTensor gk_; + TPipe *pipe_ = nullptr; + TBuf rowBuf_; + TBuf accBuf_; + TBuf tmpBuf_; + TBuf oneBuf_; + TBuf inBuf_; + TBuf scalarBuf_; + TBuf scalarI64Buf_; + uint64_t batch_ = 0; + uint64_t t_ = 0; + uint64_t hv_ = 0; + uint64_t k_ = 0; + uint64_t rank_ = 0; + uint64_t layout_ = 0; + uint64_t chunkSize_ = 0; + uint64_t seqNum_ = 0; + uint64_t maxChunks_ = 0; + bool hasCuSeqlens_ = false; + bool hasALog_ = false; + bool hasDtBias_ = false; + float lowerBound_ = -5.0f; + uint64_t usedCoreNum_ = 1; +}; + +template +__aicore__ inline void RunKdaGateCumsum(GM_ADDR g, GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR cuSeqlens, GM_ADDR gk, + const KdaGateCumsumTilingData &tilingData, TPipe *pipe) +{ + KdaGateCumsumKernel op; + op.Init(g, aLog, dtBias, cuSeqlens, gk, tilingData, pipe); + op.Process(); +} + +template +__aicore__ inline void DispatchKdaGateCumsumBySafeGate(GM_ADDR g, GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR cuSeqlens, + GM_ADDR gk, const KdaGateCumsumTilingData &tilingData, + TPipe *pipe) +{ + if (tilingData.safeGate != 0) { + RunKdaGateCumsum(g, aLog, dtBias, cuSeqlens, gk, tilingData, pipe); + } else { + RunKdaGateCumsum(g, aLog, dtBias, cuSeqlens, gk, tilingData, pipe); + } +} +} // namespace + +extern "C" __global__ __aicore__ void kda_gate_cumsum(GM_ADDR g, GM_ADDR aLog, GM_ADDR dtBias, + GM_ADDR cuSeqlens, GM_ADDR gk, GM_ADDR workspace, + GM_ADDR tiling) +{ + (void)workspace; + KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY); + GET_TILING_DATA(tilingData, tiling); + TPipe pipe; + if (tilingData.dataType == 2) { + DispatchKdaGateCumsumBySafeGate(g, aLog, dtBias, cuSeqlens, gk, tilingData, &pipe); + } else if (tilingData.dataType == 1) { + DispatchKdaGateCumsumBySafeGate(g, aLog, dtBias, cuSeqlens, gk, tilingData, &pipe); + } else { + DispatchKdaGateCumsumBySafeGate(g, aLog, dtBias, cuSeqlens, gk, tilingData, &pipe); + } +} diff --git a/csrc/attention/kda_layout_swap12/CMakeLists.txt b/csrc/attention/kda_layout_swap12/CMakeLists.txt new file mode 100644 index 000000000000..24deaa2180c5 --- /dev/null +++ b/csrc/attention/kda_layout_swap12/CMakeLists.txt @@ -0,0 +1,12 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Tianjin University, Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# the BSD 3-Clause License (the "License"). Please refer to the License for details. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. +# ----------------------------------------------------------------------------------------------------------- +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/csrc/attention/kda_layout_swap12/op_host/CMakeLists.txt b/csrc/attention/kda_layout_swap12/op_host/CMakeLists.txt new file mode 100644 index 000000000000..0095d7a4e283 --- /dev/null +++ b/csrc/attention/kda_layout_swap12/op_host/CMakeLists.txt @@ -0,0 +1,19 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2026 Tianjin University, Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# the BSD 3-Clause License (the "License"). Please refer to the License for details. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. +# ----------------------------------------------------------------------------------------------------------- +add_op_to_compiled_list() +if (BUILD_OPEN_PROJECT) + target_sources(op_host_aclnnExc PRIVATE + kda_layout_swap12_def.cpp + ) +endif() + +add_modules_sources(OPTYPE kda_layout_swap12 ACLNNTYPE aclnn_exclude) +add_ops_compile_options( + OP_NAME KdaLayoutSwap12 + OPTIONS --cce-auto-sync=off + -Wno-deprecated-declarations +) diff --git a/csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_def.cpp b/csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_def.cpp new file mode 100644 index 000000000000..1e6c7cd6236a --- /dev/null +++ b/csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_def.cpp @@ -0,0 +1,48 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#include "register/op_def_registry.h" + +namespace ops { +class KdaLayoutSwap12 : public OpDef { +public: + explicit KdaLayoutSwap12(const char *name) : OpDef(name) + { + const std::initializer_list dataTypes = { + ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16 + }; + const std::initializer_list dependencyTypes = { + ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT + }; + const std::initializer_list formats = { + ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND + }; + + this->Input("x").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Input("dependency").ParamType(OPTIONAL) + .DataType(dependencyTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("y").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + + OpAICoreConfig aicoreConfig; + aicoreConfig.DynamicCompileStaticFlag(true) + .DynamicFormatFlag(true) + .DynamicRankSupportFlag(true) + .DynamicShapeSupportFlag(true) + .NeedCheckSupportFlag(false) + .PrecisionReduceFlag(true) + .ExtendCfgInfo("prebuildPattern.value", "Opaque") + .ExtendCfgInfo("coreType.value", "AiCore") + .ExtendCfgInfo("aclnnSupport.value", "support_aclnn"); + + this->AICore().AddConfig("ascend910b", aicoreConfig); + this->AICore().AddConfig("ascend910_93", aicoreConfig); + this->AICore().AddConfig("ascend950", aicoreConfig); + } +}; + +OP_ADD(KdaLayoutSwap12); +} // namespace ops diff --git a/csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_tiling.cpp b/csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_tiling.cpp new file mode 100644 index 000000000000..141a68374154 --- /dev/null +++ b/csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_tiling.cpp @@ -0,0 +1,99 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#include "kda_layout_swap12_tiling.h" +#include +#include +#include "tiling/platform/platform_ascendc.h" + +namespace optiling { +namespace { +constexpr size_t INPUT_X_IDX = 0; +constexpr size_t OUTPUT_Y_IDX = 0; +constexpr size_t DIM_B = 0; +constexpr size_t DIM_FIRST = 1; +constexpr size_t DIM_SECOND = 2; + +int64_t DTypeCode(ge::DataType dtype) +{ + if (dtype == ge::DT_BF16) { + return 1; + } + if (dtype == ge::DT_FLOAT) { + return 2; + } + return 0; +} +} // namespace + +ge::graphStatus Tiling4KdaLayoutSwap12(gert::TilingContext *context) +{ + KdaLayoutSwap12TilingData tiling; + auto xShape = context->GetOptionalInputShape(INPUT_X_IDX)->GetStorageShape(); + auto yShape = context->GetOutputShape(OUTPUT_Y_IDX)->GetStorageShape(); + auto xDesc = context->GetInputDesc(INPUT_X_IDX); + if (xDesc == nullptr || xShape.GetDimNum() < 3) { + return ge::GRAPH_FAILED; + } + + int64_t batch = xShape.GetDim(DIM_B); + int64_t firstDim = xShape.GetDim(DIM_FIRST); + int64_t secondDim = xShape.GetDim(DIM_SECOND); + int64_t tailDim = 1; + for (size_t idx = 3; idx < xShape.GetDimNum(); ++idx) { + tailDim *= xShape.GetDim(idx); + } + + if (xShape.GetDimNum() == 3 && yShape.GetDimNum() == 3 && + yShape.GetDim(0) == xShape.GetDim(1) && + yShape.GetDim(1) == xShape.GetDim(0) && + yShape.GetDim(2) == xShape.GetDim(2)) { + batch = 1; + firstDim = xShape.GetDim(0); + secondDim = xShape.GetDim(1); + tailDim = xShape.GetDim(2); + } + + const auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo()); + uint32_t coreNum = ascendcPlatform.GetCoreNumAiv(); + int64_t rowCount = batch * firstDim * secondDim; + uint32_t blockDim = static_cast(std::min(rowCount, coreNum)); + context->SetBlockDim(blockDim == 0 ? 1 : blockDim); + + size_t *workspace = context->GetWorkspaceSizes(1); + workspace[0] = ascendcPlatform.GetLibApiWorkSpaceSize(); + + tiling.set_batch(batch); + tiling.set_firstDim(firstDim); + tiling.set_secondDim(secondDim); + tiling.set_tailDim(tailDim); + tiling.set_dataType(DTypeCode(xDesc->GetDataType())); + tiling.set_usedCoreNum(blockDim == 0 ? 1 : blockDim); + + if (xDesc->GetDataType() == ge::DT_FLOAT) { + context->SetTilingKey(0); + } else if (xDesc->GetDataType() == ge::DT_BF16) { + context->SetTilingKey(1); + } else { + context->SetTilingKey(2); + } + tiling.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity()); + context->GetRawTilingData()->SetDataSize(tiling.GetDataSize()); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus TilingPrepare4KdaLayoutSwap12(gert::TilingParseContext *context) +{ + (void)context; + return ge::GRAPH_SUCCESS; +} + +IMPL_OP_OPTILING(KdaLayoutSwap12) + .Tiling(Tiling4KdaLayoutSwap12) + .TilingParse(TilingPrepare4KdaLayoutSwap12); + +} // namespace optiling diff --git a/csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_tiling.h b/csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_tiling.h new file mode 100644 index 000000000000..ace67dc088bc --- /dev/null +++ b/csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_tiling.h @@ -0,0 +1,27 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#pragma once + +#include +#include + +namespace optiling { + +BEGIN_TILING_DATA_DEF(KdaLayoutSwap12TilingData) +TILING_DATA_FIELD_DEF(int64_t, batch); +TILING_DATA_FIELD_DEF(int64_t, firstDim); +TILING_DATA_FIELD_DEF(int64_t, secondDim); +TILING_DATA_FIELD_DEF(int64_t, tailDim); +TILING_DATA_FIELD_DEF(int64_t, dataType); +TILING_DATA_FIELD_DEF(int64_t, usedCoreNum); +END_TILING_DATA_DEF; + +REGISTER_TILING_DATA_CLASS(KdaLayoutSwap12, KdaLayoutSwap12TilingData) + +struct KdaLayoutSwap12CompileInfo {}; +} // namespace optiling diff --git a/csrc/attention/kda_layout_swap12/op_host/op_api/aclnn_kda_layout_swap12.cpp b/csrc/attention/kda_layout_swap12/op_host/op_api/aclnn_kda_layout_swap12.cpp new file mode 100644 index 000000000000..8e013dfdc229 --- /dev/null +++ b/csrc/attention/kda_layout_swap12/op_host/op_api/aclnn_kda_layout_swap12.cpp @@ -0,0 +1,119 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#include "aclnn_kda_layout_swap12.h" +#include "kda_layout_swap12.h" + +#include "acl/acl.h" +#include "aclnn/aclnn_base.h" +#include "aclnn_kernels/common/op_error_check.h" +#include "aclnn_kernels/contiguous.h" +#include "opdev/make_op_executor.h" +#include "opdev/op_dfx.h" +#include "opdev/op_executor.h" +#include "opdev/op_log.h" + +using namespace op; + +#ifdef __cplusplus +extern "C" { +#endif + +namespace { +aclnnStatus KdaLayoutSwapDataContiguous(const aclTensor *&tensor, aclOpExecutor *executor) +{ + if (tensor == nullptr) { + return ACLNN_SUCCESS; + } + tensor = l0op::Contiguous(tensor, executor); + CHECK_RET(tensor != nullptr, ACLNN_ERR_INNER_NULLPTR); + return ACLNN_SUCCESS; +} + +bool KdaLayoutSwapSameShape(const aclTensor *lhs, const aclTensor *rhs) +{ + auto lhsShape = lhs->GetViewShape(); + auto rhsShape = rhs->GetViewShape(); + if (lhsShape.GetDimNum() != rhsShape.GetDimNum()) { + return false; + } + for (size_t idx = 0; idx < lhsShape.GetDimNum(); ++idx) { + if (lhsShape.GetDim(idx) != rhsShape.GetDim(idx)) { + return false; + } + } + return true; +} + +aclnnStatus KdaLayoutSwapCheckParams( + const aclTensor *x, + const aclTensor *dependencyOptional, + const aclTensor *yOut) +{ + CHECK_COND(x != nullptr, ACLNN_ERR_PARAM_NULLPTR, "x must not be nullptr."); + CHECK_COND(yOut != nullptr, ACLNN_ERR_PARAM_NULLPTR, "yOut must not be nullptr."); + auto xShape = x->GetViewShape(); + auto yShape = yOut->GetViewShape(); + CHECK_COND(xShape.GetDimNum() >= 3, ACLNN_ERR_PARAM_INVALID, "x must have rank >= 3."); + CHECK_COND(yShape.GetDimNum() == xShape.GetDimNum(), ACLNN_ERR_PARAM_INVALID, + "yOut rank must match x rank."); + if (xShape.GetDimNum() == 3) { + CHECK_COND(yShape.GetDim(0) == xShape.GetDim(1) && yShape.GetDim(1) == xShape.GetDim(0) && + yShape.GetDim(2) == xShape.GetDim(2), + ACLNN_ERR_PARAM_INVALID, "rank3 yOut shape must be [x.dim1, x.dim0, x.dim2]."); + } else { + CHECK_COND(yShape.GetDim(0) == xShape.GetDim(0), ACLNN_ERR_PARAM_INVALID, + "yOut dim 0 must match x dim 0."); + CHECK_COND(yShape.GetDim(1) == xShape.GetDim(2) && yShape.GetDim(2) == xShape.GetDim(1), + ACLNN_ERR_PARAM_INVALID, "yOut dims 1 and 2 must swap x dims 1 and 2."); + for (size_t idx = 3; idx < xShape.GetDimNum(); ++idx) { + CHECK_COND(yShape.GetDim(idx) == xShape.GetDim(idx), ACLNN_ERR_PARAM_INVALID, + "yOut tail dims must match x tail dims."); + } + } + CHECK_COND(x->GetDataType() == yOut->GetDataType(), ACLNN_ERR_PARAM_INVALID, + "x and yOut dtype must match."); + if (dependencyOptional != nullptr) { + CHECK_COND(KdaLayoutSwapSameShape(dependencyOptional, yOut), ACLNN_ERR_PARAM_INVALID, + "dependencyOptional shape must match yOut shape."); + } + return ACLNN_SUCCESS; +} +} // namespace + +aclnnStatus aclnnKdaLayoutSwap12GetWorkspaceSize( + const aclTensor *x, + const aclTensor *dependencyOptional, + const aclTensor *yOut, + uint64_t *workspaceSize, + aclOpExecutor **executor) +{ + L2_DFX_PHASE_1(aclnnKdaLayoutSwap12, DFX_IN(x, dependencyOptional), DFX_OUT(yOut)); + auto uniqueExecutor = CREATE_EXECUTOR(); + CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); + auto executorPtr = uniqueExecutor.get(); + CHECK_RET(KdaLayoutSwapCheckParams(x, dependencyOptional, yOut) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(KdaLayoutSwapDataContiguous(x, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(KdaLayoutSwapDataContiguous(dependencyOptional, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + + auto result = l0op::KdaLayoutSwap12(x, dependencyOptional, yOut, executorPtr); + CHECK_RET(result[0] != nullptr, ACLNN_ERR_INNER_NULLPTR); + + *workspaceSize = uniqueExecutor->GetWorkspaceSize(); + uniqueExecutor.ReleaseTo(executor); + return ACLNN_SUCCESS; +} + +aclnnStatus aclnnKdaLayoutSwap12(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream) +{ + L2_DFX_PHASE_2(aclnnKdaLayoutSwap12); + return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); +} + +#ifdef __cplusplus +} +#endif diff --git a/csrc/attention/kda_layout_swap12/op_host/op_api/aclnn_kda_layout_swap12.h b/csrc/attention/kda_layout_swap12/op_host/op_api/aclnn_kda_layout_swap12.h new file mode 100644 index 000000000000..2df484d98c40 --- /dev/null +++ b/csrc/attention/kda_layout_swap12/op_host/op_api/aclnn_kda_layout_swap12.h @@ -0,0 +1,31 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#ifndef OP_API_INC_ACLNN_KDA_LAYOUT_SWAP12_H +#define OP_API_INC_ACLNN_KDA_LAYOUT_SWAP12_H + +#include "aclnn/aclnn_base.h" +#include "aclnn_util.h" + +#ifdef __cplusplus +extern "C" { +#endif + +aclnnStatus aclnnKdaLayoutSwap12GetWorkspaceSize( + const aclTensor *x, + const aclTensor *dependencyOptional, + const aclTensor *yOut, + uint64_t *workspaceSize, + aclOpExecutor **executor); + +aclnnStatus aclnnKdaLayoutSwap12(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream); + +#ifdef __cplusplus +} +#endif + +#endif diff --git a/csrc/attention/kda_layout_swap12/op_host/op_api/kda_layout_swap12.cpp b/csrc/attention/kda_layout_swap12/op_host/op_api/kda_layout_swap12.cpp new file mode 100644 index 000000000000..d7dc782492b2 --- /dev/null +++ b/csrc/attention/kda_layout_swap12/op_host/op_api/kda_layout_swap12.cpp @@ -0,0 +1,46 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#include "kda_layout_swap12.h" + +#include "opdev/make_op_executor.h" +#include "opdev/op_dfx.h" +#include "opdev/op_log.h" + +using namespace op; + +namespace l0op { +OP_TYPE_REGISTER(KdaLayoutSwap12); + +const std::array KdaLayoutSwap12( + const aclTensor *x, + const aclTensor *dependency, + const aclTensor *y, + aclOpExecutor *executor) +{ + L0_DFX(KdaLayoutSwap12, x, dependency, y); + auto ret = ADD_TO_LAUNCHER_LIST_AICORE( + KdaLayoutSwap12, + OP_INPUT(x, dependency), + OP_OUTPUT(y), + OP_ATTR()); + if (ret != ACLNN_SUCCESS) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "ADD_TO_LAUNCHER_LIST_AICORE KdaLayoutSwap12 failed."); + return {nullptr}; + } + (void)executor; + return {y}; +} + +const std::array KdaLayoutSwap12( + const aclTensor *x, + const aclTensor *y, + aclOpExecutor *executor) +{ + return KdaLayoutSwap12(x, nullptr, y, executor); +} +} // namespace l0op diff --git a/csrc/attention/kda_layout_swap12/op_host/op_api/kda_layout_swap12.h b/csrc/attention/kda_layout_swap12/op_host/op_api/kda_layout_swap12.h new file mode 100644 index 000000000000..09009729f7eb --- /dev/null +++ b/csrc/attention/kda_layout_swap12/op_host/op_api/kda_layout_swap12.h @@ -0,0 +1,24 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#pragma once + +#include "aclnn/aclnn_base.h" +#include + +namespace l0op { +const std::array KdaLayoutSwap12( + const aclTensor *x, + const aclTensor *dependency, + const aclTensor *y, + aclOpExecutor *executor); + +const std::array KdaLayoutSwap12( + const aclTensor *x, + const aclTensor *y, + aclOpExecutor *executor); +} // namespace l0op diff --git a/csrc/attention/kda_layout_swap12/op_kernel/kda_layout_swap12.cpp b/csrc/attention/kda_layout_swap12/op_kernel/kda_layout_swap12.cpp new file mode 100644 index 000000000000..551e4ba905cd --- /dev/null +++ b/csrc/attention/kda_layout_swap12/op_kernel/kda_layout_swap12.cpp @@ -0,0 +1,193 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#include "kernel_operator.h" + +using namespace AscendC; + +namespace { +constexpr uint32_t SWAP12_MTE2_MTE3_EVENT_ID = 0; +constexpr uint32_t SWAP12_MTE3_MTE2_EVENT_ID = 1; +constexpr uint32_t SWAP12_UB_ELEMENTS = 8192; +constexpr uint32_t SWAP12_MAX_GROUP_ROWS = 64; + +template +class KdaLayoutSwap12Kernel { +public: + __aicore__ inline void Init(GM_ADDR x, GM_ADDR y, const KdaLayoutSwap12TilingData &tiling, TPipe *pipe) + { + x_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(x)); + y_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(y)); + pipe_ = pipe; + batch_ = static_cast(tiling.batch); + firstDim_ = static_cast(tiling.firstDim); + secondDim_ = static_cast(tiling.secondDim); + tailDim_ = static_cast(tiling.tailDim); + usedCoreNum_ = static_cast(tiling.usedCoreNum); + pipe_->InitBuffer(copyBuf_, SWAP12_UB_ELEMENTS * sizeof(T)); + } + + __aicore__ inline void Process() + { + if (CanUseGroupedCopy()) { + ProcessGroupedRows(); + return; + } + ProcessSingleRows(); + } + +private: + __aicore__ inline bool CanUseGroupedCopy() const + { + uint64_t rowBytes = tailDim_ * static_cast(sizeof(T)); + uint64_t blockLen = rowBytes / 32; + uint64_t srcStride = secondDim_ > 0 ? (secondDim_ - 1) * blockLen : 0; + return rowBytes >= 32 && rowBytes % 32 == 0 && blockLen > 0 && blockLen <= 65535 && + srcStride <= 65535 && tailDim_ <= SWAP12_UB_ELEMENTS; + } + + __aicore__ inline uint64_t GroupRows() const + { + uint64_t rows = SWAP12_UB_ELEMENTS / tailDim_; + if (rows == 0) { + rows = 1; + } + if (rows > SWAP12_MAX_GROUP_ROWS) { + rows = SWAP12_MAX_GROUP_ROWS; + } + return rows; + } + + __aicore__ inline void ProcessSingleRows() + { + uint64_t rowCount = batch_ * firstDim_ * secondDim_; + uint64_t coreIdx = static_cast(GetBlockIdx()); + for (uint64_t row = coreIdx; row < rowCount; row += usedCoreNum_) { + CopySwappedRow(row); + } + } + + __aicore__ inline void ProcessGroupedRows() + { + uint64_t rowsPerTile = GroupRows(); + uint64_t tileCount = (firstDim_ + rowsPerTile - 1) / rowsPerTile; + uint64_t taskCount = batch_ * secondDim_ * tileCount; + uint64_t coreIdx = static_cast(GetBlockIdx()); + for (uint64_t task = coreIdx; task < taskCount; task += usedCoreNum_) { + uint64_t tileIdx = task % tileCount; + uint64_t rem = task / tileCount; + uint64_t j = rem % secondDim_; + uint64_t b = rem / secondDim_; + uint64_t firstStart = tileIdx * rowsPerTile; + uint64_t rows = firstDim_ - firstStart; + if (rows > rowsPerTile) { + rows = rowsPerTile; + } + CopySwappedRows(b, firstStart, j, rows); + } + } + + __aicore__ inline void CopySwappedRows(uint64_t b, uint64_t firstStart, uint64_t j, uint64_t rows) + { + LocalTensor local = copyBuf_.Get(); + uint64_t rowBytes = tailDim_ * static_cast(sizeof(T)); + uint64_t blockLen = rowBytes / 32; + uint64_t srcBase = ((b * firstDim_ + firstStart) * secondDim_ + j) * tailDim_; + uint64_t dstBase = ((b * secondDim_ + j) * firstDim_ + firstStart) * tailDim_; + + DataCopyParams copyInParams; + copyInParams.blockCount = static_cast(rows); + copyInParams.blockLen = static_cast(blockLen); + copyInParams.srcStride = static_cast((secondDim_ - 1) * blockLen); + copyInParams.dstStride = 0; + DataCopy(local, x_[srcBase], copyInParams); + SetFlag(SWAP12_MTE2_MTE3_EVENT_ID); + WaitFlag(SWAP12_MTE2_MTE3_EVENT_ID); + + DataCopy(y_[dstBase], local, static_cast(rows * tailDim_)); + SetFlag(SWAP12_MTE3_MTE2_EVENT_ID); + WaitFlag(SWAP12_MTE3_MTE2_EVENT_ID); + } + + __aicore__ inline void CopyTile(uint64_t srcOffset, uint64_t dstOffset, uint64_t elems) + { + LocalTensor local = copyBuf_.Get(); + uint64_t rowBytes = elems * static_cast(sizeof(T)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(local, x_[srcOffset], static_cast(elems)); + } else { + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(local, x_[srcOffset], params, padParams); + } + SetFlag(SWAP12_MTE2_MTE3_EVENT_ID); + WaitFlag(SWAP12_MTE2_MTE3_EVENT_ID); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(y_[dstOffset], local, static_cast(elems)); + } else { + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPad(y_[dstOffset], local, params); + } + SetFlag(SWAP12_MTE3_MTE2_EVENT_ID); + WaitFlag(SWAP12_MTE3_MTE2_EVENT_ID); + } + + __aicore__ inline void CopySwappedRow(uint64_t row) + { + uint64_t perBatchRows = firstDim_ * secondDim_; + uint64_t b = row / perBatchRows; + uint64_t rem = row - b * perBatchRows; + uint64_t i = rem / secondDim_; + uint64_t j = rem - i * secondDim_; + + uint64_t srcBase = ((b * firstDim_ + i) * secondDim_ + j) * tailDim_; + uint64_t dstBase = ((b * secondDim_ + j) * firstDim_ + i) * tailDim_; + for (uint64_t off = 0; off < tailDim_; off += SWAP12_UB_ELEMENTS) { + uint64_t elems = tailDim_ - off; + if (elems > SWAP12_UB_ELEMENTS) { + elems = SWAP12_UB_ELEMENTS; + } + CopyTile(srcBase + off, dstBase + off, elems); + } + } + + GlobalTensor x_; + GlobalTensor y_; + TBuf copyBuf_; + TPipe *pipe_ = nullptr; + uint64_t batch_ = 0; + uint64_t firstDim_ = 0; + uint64_t secondDim_ = 0; + uint64_t tailDim_ = 0; + uint64_t usedCoreNum_ = 1; +}; +} // namespace + +extern "C" __global__ __aicore__ void kda_layout_swap12(GM_ADDR x, GM_ADDR dependency, GM_ADDR y, + GM_ADDR workspace, GM_ADDR tiling) +{ + (void)dependency; + (void)workspace; + GET_TILING_DATA(tilingData, tiling); + TPipe pipe; + if (TILING_KEY_IS(0)) { + KERNEL_TASK_TYPE(0, KERNEL_TYPE_AIV_ONLY); + KdaLayoutSwap12Kernel op; + op.Init(x, y, tilingData, &pipe); + op.Process(); + } else if (TILING_KEY_IS(1)) { + KERNEL_TASK_TYPE(1, KERNEL_TYPE_AIV_ONLY); + KdaLayoutSwap12Kernel op; + op.Init(x, y, tilingData, &pipe); + op.Process(); + } else if (TILING_KEY_IS(2)) { + KERNEL_TASK_TYPE(2, KERNEL_TYPE_AIV_ONLY); + KdaLayoutSwap12Kernel op; + op.Init(x, y, tilingData, &pipe); + op.Process(); + } +} diff --git a/csrc/attention/recurrent_kda/CMakeLists.txt b/csrc/attention/recurrent_kda/CMakeLists.txt new file mode 100644 index 000000000000..2f5ee687b137 --- /dev/null +++ b/csrc/attention/recurrent_kda/CMakeLists.txt @@ -0,0 +1,13 @@ +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +if(NOT ENABLE_TEST AND NOT BENCHMARK) + list(REMOVE_ITEM CURRENT_DIRS tests) +endif() +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/csrc/attention/recurrent_kda/op_host/CMakeLists.txt b/csrc/attention/recurrent_kda/op_host/CMakeLists.txt new file mode 100644 index 000000000000..40edee4363a5 --- /dev/null +++ b/csrc/attention/recurrent_kda/op_host/CMakeLists.txt @@ -0,0 +1,29 @@ +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + +add_op_to_compiled_list() + +if (BUILD_OPEN_PROJECT) + target_sources(op_host_aclnnExc PRIVATE + recurrent_kda_def.cpp + ) +endif() + +if (BUILD_OPS_RTY_KERNEL) # 回黄kernel + add_ops_compile_options( + OP_NAME RecurrentKda + OPTIONS --cce-auto-sync=off + -Wno-deprecated-declarations + -Werror + ) + if("${ASCEND_COMPUTE_UNIT}" STREQUAL "ascend950") + add_ops_compile_options( + OP_NAME RecurrentKda + COMPUTE_UNIT Ascend950PR_9599 + OPTIONS -mllvm -cce-aicore-dcci-before-kernel-end=false + ) + endif() +else() + add_modules_sources(OPTYPE recurrent_kda ACLNNTYPE aclnn_exclude) +endif() diff --git a/csrc/attention/recurrent_kda/op_host/op_api/aclnn_recurrent_kda.cpp b/csrc/attention/recurrent_kda/op_host/op_api/aclnn_recurrent_kda.cpp new file mode 100644 index 000000000000..ee57e6a9ddd0 --- /dev/null +++ b/csrc/attention/recurrent_kda/op_host/op_api/aclnn_recurrent_kda.cpp @@ -0,0 +1,473 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +/*! + * \file aclnn_recurrent_kda.cpp + * \brief + */ +#include "aclnn_recurrent_kda.h" +#include "recurrent_kda.h" + +#include "aclnn_kernels/common/op_error_check.h" +#include "aclnn_kernels/contiguous.h" +#include "opdev/common_types.h" +#include "opdev/op_dfx.h" +#include "opdev/op_executor.h" +#include "opdev/op_log.h" +#include "opdev/tensor_view_utils.h" + +#include + +using namespace op; + +#ifdef __cplusplus +extern "C" { +#endif + +namespace { +constexpr size_t DIM0 = 0; +constexpr size_t DIM1 = 1; +constexpr size_t DIM2 = 2; +constexpr size_t DIM3 = 3; + +enum class RecurrentKdaLayout { + BSND, + TND, +}; + +struct RecurrentKdaParams { + const aclTensor *query = nullptr; + const aclTensor *key = nullptr; + const aclTensor *value = nullptr; + const aclTensor *gate = nullptr; + const aclTensor *beta = nullptr; + aclTensor *initialStateRef = nullptr; + const aclTensor *cuSeqlensOptional = nullptr; + const aclTensor *ssmStateIndicesOptional = nullptr; + const aclTensor *aLogOptional = nullptr; + const aclTensor *dtBiasOptional = nullptr; + const aclTensor *numAcceptedTokensOptional = nullptr; + const char *layout = "BSND"; + double scale = 1.0; + bool outputFinalState = false; + bool inplaceFinalState = true; + bool useQkL2normInKernel = false; + bool useGateInKernel = false; + bool useBetaSigmoidInKernel = false; + bool allowNegEigval = false; + bool safeGate = false; + double lowerBound = -5.0; + bool stateVFirst = false; + const aclTensor *attnOut = nullptr; + const aclTensor *finalState = nullptr; +}; + +static const std::initializer_list QKV_TYPE_SUPPORT_LIST = {op::DataType::DT_BF16}; +static const std::initializer_list STATE_TYPE_SUPPORT_LIST = {op::DataType::DT_BF16, + op::DataType::DT_FLOAT}; +static const std::initializer_list GATE_TYPE_SUPPORT_LIST = {op::DataType::DT_FLOAT, + op::DataType::DT_BF16, + op::DataType::DT_FLOAT16}; +static const std::initializer_list F32_TYPE_SUPPORT_LIST = {op::DataType::DT_FLOAT}; +static const std::initializer_list INT_TYPE_SUPPORT_LIST = {op::DataType::DT_INT32, + op::DataType::DT_INT64}; + +static size_t Rank(const aclTensor *tensor) +{ + return tensor->GetViewShape().GetDimNum(); +} + +static int64_t Dim(const aclTensor *tensor, size_t idx) +{ + return tensor->GetViewShape().GetDim(idx); +} + +static bool SameShape(const aclTensor *lhs, const aclTensor *rhs) +{ + if (Rank(lhs) != Rank(rhs)) { + return false; + } + for (size_t i = 0; i < Rank(lhs); ++i) { + if (Dim(lhs, i) != Dim(rhs, i)) { + return false; + } + } + return true; +} + +static bool ParseLayout(const char *layout, RecurrentKdaLayout &parsed) +{ + if (layout == nullptr || std::strcmp(layout, "BSND") == 0) { + parsed = RecurrentKdaLayout::BSND; + return true; + } + if (std::strcmp(layout, "TND") == 0) { + parsed = RecurrentKdaLayout::TND; + return true; + } + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: layout must be BSND or TND."); + return false; +} + +bool CheckCuSeqlensShape(const aclTensor *cuSeqlens, const char *opName) +{ + if (cuSeqlens == nullptr) { + return true; + } + if (Rank(cuSeqlens) != 1 || Dim(cuSeqlens, DIM0) < 2) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, + "%s: cuSeqlensOptional must be a 1D tensor with at least two elements.", opName); + return false; + } + return true; +} + +bool CheckShape(const RecurrentKdaParams ¶ms, RecurrentKdaLayout layout) +{ + if (!SameShape(params.query, params.key)) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: query and key must have identical shape."); + return false; + } + int64_t totalTokens = 0; + int64_t denseSeqLen = 0; + int64_t batch = 1; + int64_t h = 0; + int64_t hv = 0; + int64_t kDim = 0; + int64_t vDim = 0; + if (layout == RecurrentKdaLayout::TND) { + if (Rank(params.query) != 3 || Rank(params.value) != 3 || Rank(params.gate) != 3 || Rank(params.beta) != 2) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, + "npu_recurrent_kda: TND expects q/k [T,H,K], v [T,HV,V], g [T,HV,K], beta [T,HV]."); + return false; + } + totalTokens = Dim(params.query, DIM0); + denseSeqLen = totalTokens; + h = Dim(params.query, DIM1); + kDim = Dim(params.query, DIM2); + hv = Dim(params.value, DIM1); + vDim = Dim(params.value, DIM2); + if (Dim(params.value, DIM0) != totalTokens || Dim(params.gate, DIM0) != totalTokens || + Dim(params.beta, DIM0) != totalTokens || Dim(params.gate, DIM1) != hv || + Dim(params.beta, DIM1) != hv || Dim(params.gate, DIM2) != kDim) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: TND shape mismatch."); + return false; + } + } else { + if (Rank(params.query) != 4 || Rank(params.value) != 4 || Rank(params.gate) != 4 || Rank(params.beta) != 3) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, + "npu_recurrent_kda: BSND expects q/k [B,T,H,K], v [B,T,HV,V], g [B,T,HV,K], beta [B,T,HV]."); + return false; + } + batch = Dim(params.query, DIM0); + denseSeqLen = Dim(params.query, DIM1); + totalTokens = batch * denseSeqLen; + h = Dim(params.query, DIM2); + kDim = Dim(params.query, DIM3); + hv = Dim(params.value, DIM2); + vDim = Dim(params.value, DIM3); + if (Dim(params.value, DIM0) != batch || Dim(params.value, DIM1) != denseSeqLen || + Dim(params.gate, DIM0) != batch || Dim(params.gate, DIM1) != denseSeqLen || + Dim(params.gate, DIM2) != hv || Dim(params.gate, DIM3) != kDim || + Dim(params.beta, DIM0) != batch || Dim(params.beta, DIM1) != denseSeqLen || + Dim(params.beta, DIM2) != hv) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: BSND shape mismatch."); + return false; + } + } + if (h <= 0 || hv <= 0 || kDim <= 0 || vDim <= 0 || totalTokens <= 0 || denseSeqLen <= 0) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: all shape dimensions must be positive."); + return false; + } + if (hv % h != 0) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: HV must be divisible by H."); + return false; + } + if (kDim != 128 || (vDim != 128 && vDim != 256)) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, + "npu_recurrent_kda: K/V currently support only K=128,V=128 or K=128,V=256, but K=%ld,V=%ld.", + kDim, vDim); + return false; + } + if (!CheckCuSeqlensShape(params.cuSeqlensOptional, "npu_recurrent_kda")) { + return false; + } + int64_t seqNum = params.cuSeqlensOptional == nullptr ? + ((layout == RecurrentKdaLayout::BSND) ? batch : 1) : + Dim(params.cuSeqlensOptional, DIM0) - 1; + if (Rank(params.initialStateRef) != 4) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: initialStateRef must be rank 4."); + return false; + } + int64_t stateCapacity = Dim(params.initialStateRef, DIM0); + bool stateTailMatches = params.stateVFirst ? + (Dim(params.initialStateRef, DIM2) == vDim && Dim(params.initialStateRef, DIM3) == kDim) : + (Dim(params.initialStateRef, DIM2) == kDim && Dim(params.initialStateRef, DIM3) == vDim); + if (stateCapacity <= 0 || Dim(params.initialStateRef, DIM1) != hv || !stateTailMatches || + (params.ssmStateIndicesOptional == nullptr && stateCapacity != seqNum)) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, + "npu_recurrent_kda: state must be [state_capacity,HV,V,K] when stateVFirst=true or " + "[state_capacity,HV,K,V] otherwise; without ssmStateIndicesOptional, state_capacity must equal seq_num."); + return false; + } + if (!SameShape(params.attnOut, params.value)) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: attnOut shape must match value."); + return false; + } + if (!SameShape(params.finalState, params.initialStateRef)) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: finalState shape must match initialStateRef."); + return false; + } + if (params.ssmStateIndicesOptional != nullptr) { + size_t rank = Rank(params.ssmStateIndicesOptional); + bool packed1d = rank == 1 && Dim(params.ssmStateIndicesOptional, DIM0) >= totalTokens; + bool speculative2d = rank == 2 && Dim(params.ssmStateIndicesOptional, DIM0) == seqNum && + Dim(params.ssmStateIndicesOptional, DIM1) > 0; + if (!packed1d && !speculative2d) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, + "npu_recurrent_kda: ssm_state_indices must be packed [T] or speculative [seq_num,max_step]."); + return false; + } + } + if (params.numAcceptedTokensOptional != nullptr && + (Rank(params.numAcceptedTokensOptional) != 1 || Dim(params.numAcceptedTokensOptional, DIM0) != seqNum)) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: num_accepted_tokens length must equal sequence number."); + return false; + } + if (params.useGateInKernel && params.aLogOptional == nullptr) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: A_log is required when use_gate_in_kernel=True."); + return false; + } + if (!params.useGateInKernel && (params.safeGate || params.aLogOptional != nullptr || params.dtBiasOptional != nullptr)) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, + "npu_recurrent_kda: A_log, dt_bias and safe_gate require use_gate_in_kernel=True."); + return false; + } + if (params.aLogOptional != nullptr && (Rank(params.aLogOptional) != 1 || Dim(params.aLogOptional, DIM0) != hv)) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: A_log must be float32 with shape [HV]."); + return false; + } + if (params.dtBiasOptional != nullptr) { + bool dtBiasOk = (Rank(params.dtBiasOptional) == 1 && Dim(params.dtBiasOptional, DIM0) == hv * kDim) || + (Rank(params.dtBiasOptional) == 2 && Dim(params.dtBiasOptional, DIM0) == hv && + Dim(params.dtBiasOptional, DIM1) == kDim); + if (!dtBiasOk) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: dt_bias must be float32 with shape [HV*K] or [HV,K]."); + return false; + } + } + if (params.safeGate && (params.lowerBound < -5.0 || params.lowerBound >= 0.0)) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: lower_bound must be in [-5, 0) when safe_gate=True."); + return false; + } + return true; +} + +bool CheckNotNull(const RecurrentKdaParams ¶ms) +{ + OP_CHECK_NULL(params.query, return false); + OP_CHECK_NULL(params.key, return false); + OP_CHECK_NULL(params.value, return false); + OP_CHECK_NULL(params.gate, return false); + OP_CHECK_NULL(params.beta, return false); + OP_CHECK_NULL(params.initialStateRef, return false); + OP_CHECK_NULL(params.attnOut, return false); + OP_CHECK_NULL(params.finalState, return false); + return true; +} + +bool CheckDtypeValid(const RecurrentKdaParams ¶ms) +{ + OP_CHECK_DTYPE_NOT_SUPPORT(params.query, QKV_TYPE_SUPPORT_LIST, return false); + OP_CHECK_DTYPE_NOT_SUPPORT(params.key, QKV_TYPE_SUPPORT_LIST, return false); + OP_CHECK_DTYPE_NOT_SUPPORT(params.value, QKV_TYPE_SUPPORT_LIST, return false); + OP_CHECK_DTYPE_NOT_SUPPORT(params.gate, GATE_TYPE_SUPPORT_LIST, return false); + OP_CHECK_DTYPE_NOT_SUPPORT(params.beta, GATE_TYPE_SUPPORT_LIST, return false); + OP_CHECK_DTYPE_NOT_SUPPORT(params.initialStateRef, STATE_TYPE_SUPPORT_LIST, return false); + OP_CHECK_DTYPE_NOT_SUPPORT(params.attnOut, QKV_TYPE_SUPPORT_LIST, return false); + OP_CHECK_DTYPE_NOT_SUPPORT(params.finalState, STATE_TYPE_SUPPORT_LIST, return false); + if (params.finalState->GetDataType() != params.initialStateRef->GetDataType()) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: finalState dtype must match initialStateRef."); + return false; + } + if (params.cuSeqlensOptional != nullptr) { + OP_CHECK_DTYPE_NOT_SUPPORT(params.cuSeqlensOptional, INT_TYPE_SUPPORT_LIST, return false); + } + if (params.ssmStateIndicesOptional != nullptr) { + OP_CHECK_DTYPE_NOT_SUPPORT(params.ssmStateIndicesOptional, INT_TYPE_SUPPORT_LIST, return false); + } + if (params.aLogOptional != nullptr) { + OP_CHECK_DTYPE_NOT_SUPPORT(params.aLogOptional, F32_TYPE_SUPPORT_LIST, return false); + } + if (params.dtBiasOptional != nullptr) { + OP_CHECK_DTYPE_NOT_SUPPORT(params.dtBiasOptional, F32_TYPE_SUPPORT_LIST, return false); + } + if (params.numAcceptedTokensOptional != nullptr) { + OP_CHECK_DTYPE_NOT_SUPPORT(params.numAcceptedTokensOptional, INT_TYPE_SUPPORT_LIST, return false); + } + return true; +} + +aclnnStatus DataContiguous(const aclTensor *&tensor, aclOpExecutor *executor) +{ + if (tensor == nullptr) { + return ACLNN_SUCCESS; + } + tensor = l0op::Contiguous(tensor, executor); + CHECK_RET(tensor != nullptr, ACLNN_ERR_INNER_NULLPTR); + return ACLNN_SUCCESS; +} + +void SetTensorOriginalShape(const aclTensor *tensor) +{ + if (tensor != nullptr) { + tensor->SetOriginalShape(tensor->GetViewShape()); + } +} + +void SetInputOriginalShape(RecurrentKdaParams ¶ms) +{ + SetTensorOriginalShape(params.query); + SetTensorOriginalShape(params.key); + SetTensorOriginalShape(params.value); + SetTensorOriginalShape(params.gate); + SetTensorOriginalShape(params.beta); + SetTensorOriginalShape(params.initialStateRef); + SetTensorOriginalShape(params.cuSeqlensOptional); + SetTensorOriginalShape(params.ssmStateIndicesOptional); + SetTensorOriginalShape(params.aLogOptional); + SetTensorOriginalShape(params.dtBiasOptional); + SetTensorOriginalShape(params.numAcceptedTokensOptional); +} + +aclnnStatus PreProcess(RecurrentKdaParams ¶ms, aclOpExecutor *executor) +{ + SetInputOriginalShape(params); + CHECK_RET(DataContiguous(params.query, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(DataContiguous(params.key, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(DataContiguous(params.value, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(DataContiguous(params.gate, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(DataContiguous(params.beta, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(DataContiguous(params.cuSeqlensOptional, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(DataContiguous(params.ssmStateIndicesOptional, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(DataContiguous(params.aLogOptional, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(DataContiguous(params.dtBiasOptional, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(DataContiguous(params.numAcceptedTokensOptional, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR); + if (params.ssmStateIndicesOptional == nullptr && params.numAcceptedTokensOptional != nullptr) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, + "npu_recurrent_kda: num_accepted_tokens requires ssm_state_indices."); + return ACLNN_ERR_PARAM_INVALID; + } + return ACLNN_SUCCESS; +} +} // namespace + +aclnnStatus aclnnRecurrentKdaGetWorkspaceSize( + const aclTensor *query, + const aclTensor *key, + const aclTensor *value, + const aclTensor *gate, + const aclTensor *beta, + aclTensor *initialStateRef, + const aclTensor *cuSeqlensOptional, + const aclTensor *ssmStateIndicesOptional, + const aclTensor *aLogOptional, + const aclTensor *dtBiasOptional, + const aclTensor *numAcceptedTokensOptional, + const char *layout, + double scale, + bool outputFinalState, + bool inplaceFinalState, + bool useQkL2normInKernel, + bool useGateInKernel, + bool useBetaSigmoidInKernel, + bool allowNegEigval, + bool safeGate, + double lowerBound, + bool stateVFirst, + const aclTensor *attnOut, + const aclTensor *finalState, + uint64_t *workspaceSize, + aclOpExecutor **executor) +{ + L2_DFX_PHASE_1(aclnnRecurrentKda, + DFX_IN(query, key, value, gate, beta, initialStateRef, cuSeqlensOptional, + ssmStateIndicesOptional, aLogOptional, dtBiasOptional, numAcceptedTokensOptional, + layout, scale, outputFinalState, inplaceFinalState, useQkL2normInKernel, + useGateInKernel, useBetaSigmoidInKernel, allowNegEigval, safeGate, lowerBound, + stateVFirst), + DFX_OUT(attnOut, initialStateRef, finalState)); + + auto uniqueExecutor = CREATE_EXECUTOR(); + CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); + auto executorPtr = uniqueExecutor.get(); + + RecurrentKdaParams params{query, key, value, gate, beta, initialStateRef, cuSeqlensOptional, + ssmStateIndicesOptional, aLogOptional, dtBiasOptional, + numAcceptedTokensOptional, layout, scale, outputFinalState, inplaceFinalState, + useQkL2normInKernel, useGateInKernel, useBetaSigmoidInKernel, + allowNegEigval, safeGate, lowerBound, stateVFirst, attnOut, finalState}; + + CHECK_RET(CheckNotNull(params), ACLNN_ERR_PARAM_INVALID); + CHECK_RET(CheckDtypeValid(params), ACLNN_ERR_PARAM_INVALID); + RecurrentKdaLayout parsedLayout = RecurrentKdaLayout::BSND; + CHECK_RET(ParseLayout(params.layout, parsedLayout), ACLNN_ERR_PARAM_INVALID); + CHECK_RET(CheckShape(params, parsedLayout), ACLNN_ERR_PARAM_INVALID); + CHECK_RET(PreProcess(params, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + + aclTensor *initialStateForKernel = params.initialStateRef; + if (!IsContiguous(initialStateForKernel)) { + initialStateForKernel = executorPtr->CreateView( + initialStateForKernel, + initialStateForKernel->GetViewShape(), + initialStateForKernel->GetStorageShape(), + initialStateForKernel->GetViewStrides(), + initialStateForKernel->GetViewOffset()); + CHECK_RET(initialStateForKernel != nullptr, ACLNN_ERR_INNER_NULLPTR); + } + const aclTensor *finalStateForKernel = params.finalState; + if (!params.inplaceFinalState || params.outputFinalState) { + if (!IsContiguous(finalStateForKernel)) { + finalStateForKernel = executorPtr->CreateView( + finalStateForKernel, + finalStateForKernel->GetViewShape(), + finalStateForKernel->GetStorageShape(), + finalStateForKernel->GetViewStrides(), + finalStateForKernel->GetViewOffset()); + CHECK_RET(finalStateForKernel != nullptr, ACLNN_ERR_INNER_NULLPTR); + } + } + + auto result = l0op::RecurrentKda( + params.query, params.key, params.value, params.gate, params.beta, initialStateForKernel, + params.cuSeqlensOptional, params.ssmStateIndicesOptional, params.aLogOptional, + params.dtBiasOptional, params.numAcceptedTokensOptional, params.layout, params.scale, + params.outputFinalState, params.inplaceFinalState, params.useQkL2normInKernel, + params.useGateInKernel, params.useBetaSigmoidInKernel, params.allowNegEigval, + params.safeGate, params.lowerBound, params.stateVFirst, params.attnOut, + finalStateForKernel, executorPtr); + CHECK_RET(result[0] != nullptr && result[1] != nullptr && result[2] != nullptr, + ACLNN_ERR_INNER_NULLPTR); + if (params.inplaceFinalState && params.outputFinalState && + params.finalState != params.initialStateRef) { + CHECK_RET(l0op::ViewCopy(result[1], params.finalState, executorPtr) != nullptr, + ACLNN_ERR_INNER_NULLPTR); + } + + *workspaceSize = uniqueExecutor->GetWorkspaceSize(); + uniqueExecutor.ReleaseTo(executor); + return ACLNN_SUCCESS; +} + +aclnnStatus aclnnRecurrentKda(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream) +{ + L2_DFX_PHASE_2(aclnnRecurrentKda); + return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); +} + +#ifdef __cplusplus +} +#endif diff --git a/csrc/attention/recurrent_kda/op_host/op_api/aclnn_recurrent_kda.h b/csrc/attention/recurrent_kda/op_host/op_api/aclnn_recurrent_kda.h new file mode 100644 index 000000000000..bdcb890a9480 --- /dev/null +++ b/csrc/attention/recurrent_kda/op_host/op_api/aclnn_recurrent_kda.h @@ -0,0 +1,52 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +#ifndef OP_API_ACLNN_RECURRENT_KDA_H +#define OP_API_ACLNN_RECURRENT_KDA_H + +#include "aclnn/aclnn_base.h" +#include "aclnn_util.h" + +#ifdef __cplusplus +extern "C" { +#endif + +ACLNN_API aclnnStatus aclnnRecurrentKdaGetWorkspaceSize( + const aclTensor *query, + const aclTensor *key, + const aclTensor *value, + const aclTensor *gate, + const aclTensor *beta, + aclTensor *initialStateRef, + const aclTensor *cuSeqlensOptional, + const aclTensor *ssmStateIndicesOptional, + const aclTensor *aLogOptional, + const aclTensor *dtBiasOptional, + const aclTensor *numAcceptedTokensOptional, + const char *layout, + double scale, + bool outputFinalState, + bool inplaceFinalState, + bool useQkL2normInKernel, + bool useGateInKernel, + bool useBetaSigmoidInKernel, + bool allowNegEigval, + bool safeGate, + double lowerBound, + bool stateVFirst, + const aclTensor *attnOut, + const aclTensor *finalState, + uint64_t *workspaceSize, + aclOpExecutor **executor); + +ACLNN_API aclnnStatus aclnnRecurrentKda(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, + aclrtStream stream); + +#ifdef __cplusplus +} +#endif + +#endif // OP_API_ACLNN_RECURRENT_KDA_H diff --git a/csrc/attention/recurrent_kda/op_host/op_api/recurrent_kda.cpp b/csrc/attention/recurrent_kda/op_host/op_api/recurrent_kda.cpp new file mode 100644 index 000000000000..a575c9dbf133 --- /dev/null +++ b/csrc/attention/recurrent_kda/op_host/op_api/recurrent_kda.cpp @@ -0,0 +1,76 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +/*! + * \file recurrent_kda.cpp + * \brief + */ +#include "recurrent_kda.h" + +#include "aclnn_kernels/common/op_error_check.h" +#include "opdev/make_op_executor.h" +#include "opdev/op_def.h" +#include "opdev/op_dfx.h" +#include "opdev/op_executor.h" +#include "opdev/op_log.h" + +using namespace op; + +namespace l0op { + +OP_TYPE_REGISTER(RecurrentKda); + +const std::array RecurrentKda( + const aclTensor *query, + const aclTensor *key, + const aclTensor *value, + const aclTensor *gate, + const aclTensor *beta, + aclTensor *initialStateRef, + const aclTensor *cuSeqlensOptional, + const aclTensor *ssmStateIndicesOptional, + const aclTensor *aLogOptional, + const aclTensor *dtBiasOptional, + const aclTensor *numAcceptedTokensOptional, + const char *layout, + double scale, + bool outputFinalState, + bool inplaceFinalState, + bool useQkL2normInKernel, + bool useGateInKernel, + bool useBetaSigmoidInKernel, + bool allowNegEigval, + bool safeGate, + double lowerBound, + bool stateVFirst, + const aclTensor *attnOut, + const aclTensor *finalState, + aclOpExecutor *executor) +{ + L0_DFX(RecurrentKda, query, key, value, gate, beta, initialStateRef, cuSeqlensOptional, + ssmStateIndicesOptional, aLogOptional, dtBiasOptional, numAcceptedTokensOptional, layout, scale, + outputFinalState, inplaceFinalState, useQkL2normInKernel, useGateInKernel, + useBetaSigmoidInKernel, allowNegEigval, safeGate, lowerBound, stateVFirst, attnOut, + initialStateRef, finalState); + + float scaleAttr = static_cast(scale); + float lowerBoundAttr = static_cast(lowerBound); + auto ret = ADD_TO_LAUNCHER_LIST_AICORE( + RecurrentKda, + OP_INPUT(query, key, value, gate, beta, initialStateRef, cuSeqlensOptional, ssmStateIndicesOptional, + aLogOptional, dtBiasOptional, numAcceptedTokensOptional), + OP_OUTPUT(attnOut, initialStateRef, finalState), + OP_ATTR(layout, scaleAttr, outputFinalState, inplaceFinalState, useQkL2normInKernel, + useGateInKernel, useBetaSigmoidInKernel, allowNegEigval, safeGate, lowerBoundAttr, + stateVFirst)); + if (ret != ACLNN_SUCCESS) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "RecurrentKda ADD_TO_LAUNCHER_LIST_AICORE failed."); + return {nullptr, nullptr, nullptr}; + } + + return {attnOut, initialStateRef, finalState}; +} +} // namespace l0op diff --git a/csrc/attention/recurrent_kda/op_host/op_api/recurrent_kda.h b/csrc/attention/recurrent_kda/op_host/op_api/recurrent_kda.h new file mode 100644 index 000000000000..34696cb7b353 --- /dev/null +++ b/csrc/attention/recurrent_kda/op_host/op_api/recurrent_kda.h @@ -0,0 +1,43 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +#ifndef PTA_NPU_OP_API_COMMON_INC_LEVEL0_OP_RECURRENT_KDA +#define PTA_NPU_OP_API_COMMON_INC_LEVEL0_OP_RECURRENT_KDA + +#include "opdev/make_op_executor.h" +#include "opdev/op_executor.h" +#include + +namespace l0op { +const std::array RecurrentKda( + const aclTensor *query, + const aclTensor *key, + const aclTensor *value, + const aclTensor *gate, + const aclTensor *beta, + aclTensor *initialStateRef, + const aclTensor *cuSeqlensOptional, + const aclTensor *ssmStateIndicesOptional, + const aclTensor *aLogOptional, + const aclTensor *dtBiasOptional, + const aclTensor *numAcceptedTokensOptional, + const char *layout, + double scale, + bool outputFinalState, + bool inplaceFinalState, + bool useQkL2normInKernel, + bool useGateInKernel, + bool useBetaSigmoidInKernel, + bool allowNegEigval, + bool safeGate, + double lowerBound, + bool stateVFirst, + const aclTensor *attnOut, + const aclTensor *finalState, + aclOpExecutor *executor); +} + +#endif // PTA_NPU_OP_API_COMMON_INC_LEVEL0_OP_RECURRENT_KDA diff --git a/csrc/attention/recurrent_kda/op_host/recurrent_kda_def.cpp b/csrc/attention/recurrent_kda/op_host/recurrent_kda_def.cpp new file mode 100644 index 000000000000..5a147c4bb78e --- /dev/null +++ b/csrc/attention/recurrent_kda/op_host/recurrent_kda_def.cpp @@ -0,0 +1,91 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +/*! + * \file recurrent_kda_def.cpp + * \brief Recurrent KDA operator definition. + */ +#include "register/op_def_registry.h" + +namespace ops { +class RecurrentKda : public OpDef { +public: + explicit RecurrentKda(const char *name) : OpDef(name) + { + const std::initializer_list qkvTypes = {ge::DT_BF16, ge::DT_BF16}; + const std::initializer_list floatTypes = {ge::DT_FLOAT, ge::DT_FLOAT}; + const std::initializer_list stateTypes = {ge::DT_BF16, ge::DT_FLOAT}; + + this->Input("query").ParamType(REQUIRED).DataType(qkvTypes).FormatList({ge::FORMAT_ND}); + this->Input("key").ParamType(REQUIRED).DataType(qkvTypes).FormatList({ge::FORMAT_ND}); + this->Input("value").ParamType(REQUIRED).DataType(qkvTypes).FormatList({ge::FORMAT_ND}); + this->Input("gate").ParamType(REQUIRED) + .DataTypeList({ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16}).FormatList({ge::FORMAT_ND}); + this->Input("beta").ParamType(REQUIRED) + .DataTypeList({ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16}).FormatList({ge::FORMAT_ND}); + this->Input("initial_state") + .ParamType(REQUIRED) + .DataType(stateTypes) + .FormatList({ge::FORMAT_ND}) + .IgnoreContiguous(); + this->Input("cu_seqlens") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32, ge::DT_INT64}) + .FormatList({ge::FORMAT_ND}); + this->Input("ssm_state_indices") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32, ge::DT_INT64}) + .FormatList({ge::FORMAT_ND}); + this->Input("A_log").ParamType(OPTIONAL).DataType(floatTypes).FormatList({ge::FORMAT_ND}); + this->Input("dt_bias").ParamType(OPTIONAL).DataType(floatTypes).FormatList({ge::FORMAT_ND}); + this->Input("num_accepted_tokens") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32, ge::DT_INT64}) + .FormatList({ge::FORMAT_ND}); + this->Output("attn_out").ParamType(REQUIRED).DataType(qkvTypes).FormatList({ge::FORMAT_ND}); + this->Output("initial_state") + .ParamType(REQUIRED) + .DataType(stateTypes) + .FormatList({ge::FORMAT_ND}) + .IgnoreContiguous(); + this->Output("final_state") + .ParamType(REQUIRED) + .DataType(stateTypes) + .FormatList({ge::FORMAT_ND}) + .IgnoreContiguous(); + + this->Attr("layout").AttrType(OPTIONAL).String("BSND"); + this->Attr("scale").AttrType(OPTIONAL).Float(1.0); + this->Attr("output_final_state").AttrType(OPTIONAL).Bool(false); + this->Attr("inplace_final_state").AttrType(OPTIONAL).Bool(true); + this->Attr("use_qk_l2norm_in_kernel").AttrType(OPTIONAL).Bool(false); + this->Attr("use_gate_in_kernel").AttrType(OPTIONAL).Bool(false); + this->Attr("use_beta_sigmoid_in_kernel").AttrType(OPTIONAL).Bool(false); + this->Attr("allow_neg_eigval").AttrType(OPTIONAL).Bool(false); + this->Attr("safe_gate").AttrType(OPTIONAL).Bool(false); + this->Attr("lower_bound").AttrType(OPTIONAL).Float(-5.0); + this->Attr("state_v_first").AttrType(OPTIONAL).Bool(false); + + OpAICoreConfig aicConfig; + aicConfig.DynamicCompileStaticFlag(true) + .DynamicFormatFlag(true) + .DynamicRankSupportFlag(true) + .DynamicShapeSupportFlag(true) + .NeedCheckSupportFlag(false) + .PrecisionReduceFlag(true) + .ExtendCfgInfo("prebuildPattern.value", "Opaque") + .ExtendCfgInfo("coreType.value", "AiCore") + .ExtendCfgInfo("aclnnSupport.value", "support_aclnn") + .ExtendCfgInfo("softsync.flag", "true"); + this->AICore().AddConfig("ascend910b", aicConfig); + this->AICore().AddConfig("ascend910_93", aicConfig); + this->AICore().AddConfig("ascend950", aicConfig); + } +}; + +OP_ADD(RecurrentKda); + +} // namespace ops diff --git a/csrc/attention/recurrent_kda/op_host/recurrent_kda_infershape.cpp b/csrc/attention/recurrent_kda/op_host/recurrent_kda_infershape.cpp new file mode 100644 index 000000000000..35e832ae2b06 --- /dev/null +++ b/csrc/attention/recurrent_kda/op_host/recurrent_kda_infershape.cpp @@ -0,0 +1,65 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +/* ! + * \file recurrent_kda_infershape.cpp + * \brief + */ +#include "exe_graph/runtime/infer_shape_context.h" +#include "exe_graph/runtime/shape.h" +#include "exe_graph/runtime/storage_shape.h" +#include "register/op_impl_registry.h" +#include "log/log.h" +#include "err/ops_err.h" + +using namespace gert; +namespace ops { + +const size_t VALUE_INDEX = 2; +const size_t STATE_INDEX = 5; + +const size_t DIM_0 = 0; +const size_t DIM_1 = 1; +const size_t DIM_2 = 2; + +static ge::graphStatus InferShapeRecurrentKda(InferShapeContext *context) +{ + if (context == nullptr) { + OP_LOGE("RecurrentKda", "inference context is null"); + return ge::GRAPH_FAILED; + } + + auto opName = context->GetNodeName(); + auto shapeValue = context->GetInputShape(VALUE_INDEX); + auto shapeInitialState = context->GetInputShape(STATE_INDEX); + auto shapeOut = context->GetOutputShape(DIM_0); + auto shapeInitialStateOut = context->GetOutputShape(DIM_1); + auto shapeFinalState = context->GetOutputShape(DIM_2); + if (shapeValue == nullptr || shapeInitialState == nullptr || shapeOut == nullptr || + shapeInitialStateOut == nullptr || shapeFinalState == nullptr) { + OP_LOGE(opName, "[InferShape] shape is null"); + return ge::GRAPH_FAILED; + } + + *shapeOut = *shapeValue; + *shapeInitialStateOut = *shapeInitialState; + *shapeFinalState = *shapeInitialState; + + return ge::GRAPH_SUCCESS; +} + +static ge::graphStatus InferDataTypeRecurrentKda(gert::InferDataTypeContext *context) +{ + context->SetOutputDataType(0, ge::DT_BF16); + context->SetOutputDataType(1, context->GetInputDataType(STATE_INDEX)); + context->SetOutputDataType(2, context->GetInputDataType(STATE_INDEX)); + return ge::GRAPH_SUCCESS; +} + +IMPL_OP_INFERSHAPE(RecurrentKda) + .InferShape(InferShapeRecurrentKda) + .InferDataType(InferDataTypeRecurrentKda); +} // namespace ops diff --git a/csrc/attention/recurrent_kda/op_host/recurrent_kda_tiling.cpp b/csrc/attention/recurrent_kda/op_host/recurrent_kda_tiling.cpp new file mode 100644 index 000000000000..4391a708cb34 --- /dev/null +++ b/csrc/attention/recurrent_kda/op_host/recurrent_kda_tiling.cpp @@ -0,0 +1,439 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +/*! + * \file recurrent_kda_tiling.cpp + * \brief + */ +#include "recurrent_kda_tiling.h" + +#include "err/ops_err.h" +#include "log/log.h" +#include "platform/platform_infos_def.h" +#include "register/op_def_registry.h" +#include "tiling/platform/platform_ascendc.h" +#include "tiling_base/tiling_templates_registry.h" +#include + +namespace optiling { + +REGISTER_OPS_TILING_TEMPLATE(RecurrentKda, RecurrentKdaTiling, 0); + +const size_t QUERY_INDEX = 0; +const size_t KEY_INDEX = 1; +const size_t VALUE_INDEX = 2; +const size_t GATE_INDEX = 3; +const size_t BETA_INDEX = 4; +const size_t STATE_INDEX = 5; +const size_t CU_SEQLENS_INDEX = 6; +const size_t SSM_STATE_INDICES_INDEX = 7; +const size_t A_LOG_INDEX = 8; +const size_t DT_BIAS_INDEX = 9; +const size_t ACC_TOKEN_INDEX = 10; + +const size_t INITIAL_STATE_OUTPUT_INDEX = 1; +const size_t FINAL_STATE_OUTPUT_INDEX = 2; + +const size_t ATTR_LAYOUT_INDEX = 0; +const size_t ATTR_SCALE_INDEX = 1; +const size_t ATTR_OUTPUT_FINAL_STATE_INDEX = 2; +const size_t ATTR_INPLACE_FINAL_STATE_INDEX = 3; +const size_t ATTR_USE_QK_L2NORM_INDEX = 4; +const size_t ATTR_USE_GATE_INDEX = 5; +const size_t ATTR_USE_BETA_SIGMOID_INDEX = 6; +const size_t ATTR_ALLOW_NEG_EIGVAL_INDEX = 7; +const size_t ATTR_SAFE_GATE_INDEX = 8; +const size_t ATTR_LOWER_BOUND_INDEX = 9; +const size_t ATTR_STATE_V_FIRST_INDEX = 10; + +void RecurrentKdaTiling::InitCompileInfo() +{ + auto platformInfoPtr = context_->GetPlatformInfo(); + if (platformInfoPtr == nullptr) { + OP_LOGE(context_->GetNodeName(), "platformInfoPtr is null"); + return; + } + const auto &ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr); + ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, compileInfo_.ubSize); + compileInfo_.aivNum = ascendcPlatform.GetCoreNumAiv(); + + if (compileInfo_.aivNum <= 0) { + OP_LOGE(context_->GetNodeName(), "aivNum <= 0"); + return; + } + tilingData_.vectorCoreNum = static_cast(compileInfo_.aivNum); +} + +namespace { +void CopyOptionalOriginShape(gert::TilingContext *context, size_t index, gert::Shape &dst) +{ + auto shape = context->GetOptionalInputShape(index); + if (shape != nullptr) { + dst = shape->GetOriginShape(); + } +} + +template +void CopyStateStrides(const StrideType *src, std::array &dst, bool &hasStrides) +{ + if (src == nullptr || src->GetDimNum() != RKDA_STATE_DIM_NUM) { + return; + } + for (size_t i = 0; i < RKDA_STATE_DIM_NUM; ++i) { + dst[i] = src->GetStride(i); + } + hasStrides = true; +} +} // namespace + +RecurrentKdaTilingContext RecurrentKdaTiling::BuildProcessorContext() const +{ + RecurrentKdaTilingContext ctx; + ctx.nodeName = context_->GetNodeName(); + ctx.queryShape = context_->GetInputShape(QUERY_INDEX)->GetOriginShape(); + ctx.keyShape = context_->GetInputShape(KEY_INDEX)->GetOriginShape(); + ctx.valueShape = context_->GetInputShape(VALUE_INDEX)->GetOriginShape(); + ctx.gateShape = context_->GetInputShape(GATE_INDEX)->GetOriginShape(); + ctx.betaShape = context_->GetInputShape(BETA_INDEX)->GetOriginShape(); + ctx.stateShape = context_->GetInputShape(STATE_INDEX)->GetOriginShape(); + CopyStateStrides(context_->GetInputStride(STATE_INDEX), ctx.stateInStrides, ctx.hasStateInStrides); + const size_t stateOutputIndex = tilingData_.inplaceFinalState == 1 ? + INITIAL_STATE_OUTPUT_INDEX : FINAL_STATE_OUTPUT_INDEX; + CopyStateStrides(context_->GetOutputStride(stateOutputIndex), ctx.stateOutStrides, ctx.hasStateOutStrides); + if (!ctx.hasStateOutStrides && tilingData_.inplaceFinalState == 1 && ctx.hasStateInStrides) { + ctx.stateOutStrides = ctx.stateInStrides; + ctx.hasStateOutStrides = true; + } + CopyOptionalOriginShape(context_, CU_SEQLENS_INDEX, ctx.cuSeqlensShape); + CopyOptionalOriginShape(context_, SSM_STATE_INDICES_INDEX, ctx.ssmStateShape); + CopyOptionalOriginShape(context_, A_LOG_INDEX, ctx.aLogShape); + CopyOptionalOriginShape(context_, DT_BIAS_INDEX, ctx.dtBiasShape); + CopyOptionalOriginShape(context_, ACC_TOKEN_INDEX, ctx.acceptedTokensShape); + ctx.aivNum = compileInfo_.aivNum; + ctx.ubSize = compileInfo_.ubSize; + ctx.stateDtype = context_->GetInputDesc(STATE_INDEX)->GetDataType(); + ctx.gateDtype = context_->GetInputDesc(GATE_INDEX)->GetDataType(); + ctx.betaDtype = context_->GetInputDesc(BETA_INDEX)->GetDataType(); + if (context_->GetOptionalInputDesc(CU_SEQLENS_INDEX) != nullptr) { + ctx.cuSeqlensDtype = context_->GetOptionalInputDesc(CU_SEQLENS_INDEX)->GetDataType(); + } + if (context_->GetOptionalInputDesc(SSM_STATE_INDICES_INDEX) != nullptr) { + ctx.ssmStateIndicesDtype = context_->GetOptionalInputDesc(SSM_STATE_INDICES_INDEX)->GetDataType(); + } + if (context_->GetOptionalInputDesc(ACC_TOKEN_INDEX) != nullptr) { + ctx.acceptedTokensDtype = context_->GetOptionalInputDesc(ACC_TOKEN_INDEX)->GetDataType(); + } + ctx.scale = tilingData_.scale; + ctx.lowerBound = tilingData_.lowerBound; + ctx.layout = tilingData_.layout; + ctx.hasCuSeqlens = tilingData_.hasCuSeqlens; + ctx.hasSsmStateIndices = tilingData_.hasSsmStateIndices; + ctx.hasALog = tilingData_.hasALog; + ctx.hasDtBias = tilingData_.hasDtBias; + ctx.hasAcceptedTokens = tilingData_.hasAcceptedTokens; + ctx.useQkL2norm = tilingData_.useQkL2norm; + ctx.useGateInKernel = tilingData_.useGateInKernel; + ctx.useBetaSigmoid = tilingData_.useBetaSigmoid; + ctx.allowNegEigval = tilingData_.allowNegEigval; + ctx.safeGate = tilingData_.safeGate; + ctx.stateVFirst = tilingData_.stateVFirst; + ctx.outputFinalState = tilingData_.outputFinalState; + ctx.inplaceFinalState = tilingData_.inplaceFinalState; + return ctx; +} + +ge::graphStatus RecurrentKdaTiling::GetPlatformInfo() +{ + return ge::GRAPH_SUCCESS; +}; + +ge::graphStatus RecurrentKdaTiling::GetShapeAttrsInfo() +{ + OP_CHECK_IF(CheckContext() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "Invalid context."), + return ge::GRAPH_FAILED); + + OP_CHECK_IF(GetOptionalInput() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "Invalid optional input."), + return ge::GRAPH_FAILED); + + OP_CHECK_IF(GetAttrsInfo() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "Invalid attrs."), + return ge::GRAPH_FAILED); + + OP_CHECK_IF(AnalyzeDtype() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "Invalid dtypes."), + return ge::GRAPH_FAILED); + + OP_CHECK_IF(AnalyzeShapes() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "Invalid shapes."), + return ge::GRAPH_FAILED); + + OP_CHECK_IF(AnalyzeFormat() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "Invalid format."), + return ge::GRAPH_FAILED); + + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus RecurrentKdaTiling::DoOpTiling() +{ + OP_CHECK_IF(CalUbSize() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "CalUbSize failed."), + return ge::GRAPH_FAILED); + + PrintTilingData(); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus RecurrentKdaTiling::DoLibApiTiling() +{ + tilingKey_ = 0; + return ge::GRAPH_SUCCESS; +}; + +uint64_t RecurrentKdaTiling::GetTilingKey() const +{ + return tilingKey_; +}; + +ge::graphStatus RecurrentKdaTiling::GetWorkspaceSize() +{ + workspaceSize_ = static_cast(RKDA_SYS_WORKSPACE_SIZE); + return ge::GRAPH_SUCCESS; +}; + +ge::graphStatus RecurrentKdaTiling::PostTiling() +{ + context_->SetBlockDim(tilingData_.vectorCoreNum); + auto tilingDataSize = sizeof(RecurrentKdaTilingData); + errno_t ret = memcpy_s(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity(), + reinterpret_cast(&tilingData_), tilingDataSize); + if (ret != EOK) { + OP_LOGE(context_->GetNodeName(), "memcpy_s failed, ret=%d", ret); + return ge::GRAPH_FAILED; + } + context_->GetRawTilingData()->SetDataSize(tilingDataSize); + + size_t *workspaces = context_->GetWorkspaceSizes(1); + OP_CHECK_IF(workspaces == nullptr, OPS_REPORT_CUBE_INNER_ERR(context_->GetNodeName(), "workspaces is null"), + return ge::GRAPH_FAILED); + workspaces[0] = workspaceSize_; + + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus RecurrentKdaTiling::CheckContext() +{ + OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputShape(QUERY_INDEX)); + OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputDesc(QUERY_INDEX)); + OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputShape(KEY_INDEX)); + OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputDesc(KEY_INDEX)); + OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputShape(VALUE_INDEX)); + OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputDesc(VALUE_INDEX)); + OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputShape(GATE_INDEX)); + OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputDesc(GATE_INDEX)); + OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputShape(BETA_INDEX)); + OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputDesc(BETA_INDEX)); + OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputShape(STATE_INDEX)); + OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputDesc(STATE_INDEX)); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus RecurrentKdaTiling::AnalyzeDtype() +{ + auto queryDtype = context_->GetInputDesc(QUERY_INDEX)->GetDataType(); + auto keyDtype = context_->GetInputDesc(KEY_INDEX)->GetDataType(); + auto valueDtype = context_->GetInputDesc(VALUE_INDEX)->GetDataType(); + OP_CHECK_IF(queryDtype != ge::DT_BF16 || keyDtype != ge::DT_BF16 || valueDtype != ge::DT_BF16, + OP_LOGE(context_->GetNodeName(), "query, key and value dtype should be bfloat16"), + return ge::GRAPH_FAILED); + + auto gateDtype = context_->GetInputDesc(GATE_INDEX)->GetDataType(); + auto betaDtype = context_->GetInputDesc(BETA_INDEX)->GetDataType(); + auto stateDtype = context_->GetInputDesc(STATE_INDEX)->GetDataType(); + OP_CHECK_IF((gateDtype != ge::DT_FLOAT && gateDtype != ge::DT_BF16 && gateDtype != ge::DT_FLOAT16) || + (betaDtype != ge::DT_FLOAT && betaDtype != ge::DT_BF16 && betaDtype != ge::DT_FLOAT16), + OP_LOGE(context_->GetNodeName(), "gate and beta dtype should be float32, bfloat16 or float16"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(stateDtype != ge::DT_FLOAT && stateDtype != ge::DT_BF16, + OP_LOGE(context_->GetNodeName(), "initial_state dtype should be bfloat16 or float32"), + return ge::GRAPH_FAILED); + if (context_->GetOptionalInputDesc(CU_SEQLENS_INDEX) != nullptr) { + auto dtype = context_->GetOptionalInputDesc(CU_SEQLENS_INDEX)->GetDataType(); + OP_CHECK_IF(dtype != ge::DT_INT32 && dtype != ge::DT_INT64, + OP_LOGE(context_->GetNodeName(), "cu_seqlens dtype should be int32 or int64"), + return ge::GRAPH_FAILED); + } + if (context_->GetOptionalInputDesc(SSM_STATE_INDICES_INDEX) != nullptr) { + auto dtype = context_->GetOptionalInputDesc(SSM_STATE_INDICES_INDEX)->GetDataType(); + OP_CHECK_IF(dtype != ge::DT_INT32 && dtype != ge::DT_INT64, + OP_LOGE(context_->GetNodeName(), "ssm_state_indices dtype should be int32 or int64"), + return ge::GRAPH_FAILED); + } + if (context_->GetOptionalInputDesc(A_LOG_INDEX) != nullptr) { + OP_CHECK_IF(context_->GetOptionalInputDesc(A_LOG_INDEX)->GetDataType() != ge::DT_FLOAT, + OP_LOGE(context_->GetNodeName(), "A_log dtype should be float32"), + return ge::GRAPH_FAILED); + } + if (context_->GetOptionalInputDesc(DT_BIAS_INDEX) != nullptr) { + OP_CHECK_IF(context_->GetOptionalInputDesc(DT_BIAS_INDEX)->GetDataType() != ge::DT_FLOAT, + OP_LOGE(context_->GetNodeName(), "dt_bias dtype should be float32"), + return ge::GRAPH_FAILED); + } + if (context_->GetOptionalInputDesc(ACC_TOKEN_INDEX) != nullptr) { + auto dtype = context_->GetOptionalInputDesc(ACC_TOKEN_INDEX)->GetDataType(); + OP_CHECK_IF(dtype != ge::DT_INT32 && dtype != ge::DT_INT64, + OP_LOGE(context_->GetNodeName(), "num_accepted_tokens dtype should be int32 or int64"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus RecurrentKdaTiling::AnalyzeShapes() +{ + RecurrentKdaTilingProcessor processor(BuildProcessorContext()); + return processor.ProcessShapes(tilingData_); +} + +bool RecurrentKdaTiling::CheckFormat(ge::Format format, const std::string &desc) +{ + if (format == ge::FORMAT_FRACTAL_NZ) { + OP_LOGE(context_->GetNodeName(), "%s format does not support NZ", desc.c_str()); + return false; + } + return true; +} + +ge::graphStatus RecurrentKdaTiling::AnalyzeFormat() +{ + if (!CheckFormat(context_->GetInputDesc(QUERY_INDEX)->GetStorageFormat(), "query") || + !CheckFormat(context_->GetInputDesc(KEY_INDEX)->GetStorageFormat(), "key") || + !CheckFormat(context_->GetInputDesc(VALUE_INDEX)->GetStorageFormat(), "value") || + !CheckFormat(context_->GetInputDesc(GATE_INDEX)->GetStorageFormat(), "gate") || + !CheckFormat(context_->GetInputDesc(BETA_INDEX)->GetStorageFormat(), "beta") || + !CheckFormat(context_->GetInputDesc(STATE_INDEX)->GetStorageFormat(), "initial_state")) { + return ge::GRAPH_FAILED; + } + + const std::array, 5> optionalInputs = {{ + {CU_SEQLENS_INDEX, "cu_seqlens"}, + {SSM_STATE_INDICES_INDEX, "ssm_state_indices"}, + {A_LOG_INDEX, "A_log"}, + {DT_BIAS_INDEX, "dt_bias"}, + {ACC_TOKEN_INDEX, "num_accepted_tokens"}, + }}; + for (const auto &item : optionalInputs) { + auto desc = context_->GetOptionalInputDesc(item.first); + if (desc != nullptr && !CheckFormat(desc->GetStorageFormat(), item.second)) { + return ge::GRAPH_FAILED; + } + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus RecurrentKdaTiling::GetAttrsInfo() +{ + auto attrs = context_->GetAttrs(); + OP_CHECK_IF(attrs == nullptr, OP_LOGE(context_->GetNodeName(), "attrs is null"), return ge::GRAPH_FAILED); + const char *layout = attrs->GetAttrPointer(ATTR_LAYOUT_INDEX); + if (layout == nullptr || std::strcmp(layout, "BSND") == 0) { + tilingData_.layout = RKDA_LAYOUT_BSND; + } else if (std::strcmp(layout, "TND") == 0) { + tilingData_.layout = RKDA_LAYOUT_TND; + } else { + OP_LOGE(context_->GetNodeName(), "layout must be BSND or TND for RecurrentKda, got %s", layout); + return ge::GRAPH_FAILED; + } + tilingData_.scale = *attrs->GetAttrPointer(ATTR_SCALE_INDEX); + tilingData_.outputFinalState = *attrs->GetAttrPointer(ATTR_OUTPUT_FINAL_STATE_INDEX) ? 1 : 0; + tilingData_.inplaceFinalState = *attrs->GetAttrPointer(ATTR_INPLACE_FINAL_STATE_INDEX) ? 1 : 0; + tilingData_.useQkL2norm = *attrs->GetAttrPointer(ATTR_USE_QK_L2NORM_INDEX) ? 1 : 0; + tilingData_.useGateInKernel = *attrs->GetAttrPointer(ATTR_USE_GATE_INDEX) ? 1 : 0; + tilingData_.useBetaSigmoid = *attrs->GetAttrPointer(ATTR_USE_BETA_SIGMOID_INDEX) ? 1 : 0; + tilingData_.allowNegEigval = *attrs->GetAttrPointer(ATTR_ALLOW_NEG_EIGVAL_INDEX) ? 1 : 0; + tilingData_.safeGate = *attrs->GetAttrPointer(ATTR_SAFE_GATE_INDEX) ? 1 : 0; + tilingData_.lowerBound = *attrs->GetAttrPointer(ATTR_LOWER_BOUND_INDEX); + tilingData_.stateVFirst = *attrs->GetAttrPointer(ATTR_STATE_V_FIRST_INDEX) ? 1 : 0; + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus RecurrentKdaTiling::GetOptionalInput() +{ + tilingData_.hasCuSeqlens = (context_->GetOptionalInputDesc(CU_SEQLENS_INDEX) == nullptr) ? 0 : 1; + tilingData_.hasSsmStateIndices = (context_->GetOptionalInputDesc(SSM_STATE_INDICES_INDEX) == nullptr) ? 0 : 1; + tilingData_.hasALog = (context_->GetOptionalInputDesc(A_LOG_INDEX) == nullptr) ? 0 : 1; + tilingData_.hasDtBias = (context_->GetOptionalInputDesc(DT_BIAS_INDEX) == nullptr) ? 0 : 1; + tilingData_.hasAcceptedTokens = (context_->GetOptionalInputDesc(ACC_TOKEN_INDEX) == nullptr) ? 0 : 1; + return ge::GRAPH_SUCCESS; +} + +void RecurrentKdaTiling::PrintTilingData() +{ + OP_LOGD(context_->GetNodeName(), "vectorCoreNum: [%u]", tilingData_.vectorCoreNum); + OP_LOGD(context_->GetNodeName(), "ubCalSize: [%u]", tilingData_.ubCalSize); + OP_LOGD(context_->GetNodeName(), "ubRestBytes: [%u]", tilingData_.ubRestBytes); + OP_LOGD(context_->GetNodeName(), "t: [%u]", tilingData_.t); + OP_LOGD(context_->GetNodeName(), "seqLen: [%u]", tilingData_.seqLen); + OP_LOGD(context_->GetNodeName(), "nk: [%u]", tilingData_.nk); + OP_LOGD(context_->GetNodeName(), "dk: [%u]", tilingData_.dk); + OP_LOGD(context_->GetNodeName(), "nv: [%u]", tilingData_.nv); + OP_LOGD(context_->GetNodeName(), "dv: [%u]", tilingData_.dv); + OP_LOGD(context_->GetNodeName(), "sBlockNum: [%u]", tilingData_.sBlockNum); + OP_LOGD(context_->GetNodeName(), "ssmStateStride: [%u]", tilingData_.ssmStateStride); + OP_LOGD(context_->GetNodeName(), "b: [%u]", tilingData_.b); + OP_LOGD(context_->GetNodeName(), "vStep: [%u]", tilingData_.vStep); + OP_LOGD(context_->GetNodeName(), "stateOutBufferNum: [%u]", tilingData_.stateOutBufferNum); + OP_LOGD(context_->GetNodeName(), "attnOutBufferNum: [%u]", tilingData_.attnOutBufferNum); + OP_LOGD(context_->GetNodeName(), "scale: [%f]", tilingData_.scale); + OP_LOGD(context_->GetNodeName(), "lowerBound: [%f]", tilingData_.lowerBound); + OP_LOGD(context_->GetNodeName(), "layout: [%u]", tilingData_.layout); + OP_LOGD(context_->GetNodeName(), "hasCuSeqlens: [%u]", tilingData_.hasCuSeqlens); + OP_LOGD(context_->GetNodeName(), "hasSsmStateIndices: [%u]", tilingData_.hasSsmStateIndices); + OP_LOGD(context_->GetNodeName(), "hasALog: [%u]", tilingData_.hasALog); + OP_LOGD(context_->GetNodeName(), "hasDtBias: [%u]", tilingData_.hasDtBias); + OP_LOGD(context_->GetNodeName(), "hasAcceptedTokens: [%u]", tilingData_.hasAcceptedTokens); + OP_LOGD(context_->GetNodeName(), "useQkL2norm: [%u]", tilingData_.useQkL2norm); + OP_LOGD(context_->GetNodeName(), "useGateInKernel: [%u]", tilingData_.useGateInKernel); + OP_LOGD(context_->GetNodeName(), "useBetaSigmoid: [%u]", tilingData_.useBetaSigmoid); + OP_LOGD(context_->GetNodeName(), "allowNegEigval: [%u]", tilingData_.allowNegEigval); + OP_LOGD(context_->GetNodeName(), "safeGate: [%u]", tilingData_.safeGate); + OP_LOGD(context_->GetNodeName(), "stateVFirst: [%u]", tilingData_.stateVFirst); + OP_LOGD(context_->GetNodeName(), "outputFinalState: [%u]", tilingData_.outputFinalState); + OP_LOGD(context_->GetNodeName(), "inplaceFinalState: [%u]", tilingData_.inplaceFinalState); + OP_LOGD(context_->GetNodeName(), "gateDtype: [%u]", tilingData_.gateDtype); + OP_LOGD(context_->GetNodeName(), "betaDtype: [%u]", tilingData_.betaDtype); + OP_LOGD(context_->GetNodeName(), "cuSeqlensDtype: [%u]", tilingData_.cuSeqlensDtype); + OP_LOGD(context_->GetNodeName(), "ssmStateIndicesDtype: [%u]", tilingData_.ssmStateIndicesDtype); + OP_LOGD(context_->GetNodeName(), "acceptedTokensDtype: [%u]", tilingData_.acceptedTokensDtype); +} + +ge::graphStatus RecurrentKdaTiling::CalUbSize() +{ + RecurrentKdaTilingProcessor processor(BuildProcessorContext()); + return processor.ProcessUb(tilingData_); +} + +static ge::graphStatus RecurrentKdaTilingFunc(gert::TilingContext *context) +{ + OP_CHECK_IF(context == nullptr, OPS_REPORT_CUBE_INNER_ERR("RecurrentKda", "context is null"), + return ge::GRAPH_FAILED); + return Ops::Transformer::OpTiling::TilingRegistry::GetInstance().DoTilingImpl(context); +} + +static ge::graphStatus TilingPrepareForRecurrentKda(gert::TilingParseContext *context) +{ + OP_CHECK_IF(context == nullptr, OPS_REPORT_CUBE_INNER_ERR("RecurrentKda", "context is null"), + return ge::GRAPH_FAILED); + + fe::PlatFormInfos *platformInfo = context->GetPlatformInfo(); + OP_CHECK_IF(platformInfo == nullptr, OPS_REPORT_CUBE_INNER_ERR(context->GetNodeName(), "platformInfoPtr is null"), + return ge::GRAPH_FAILED); + + auto compileInfoPtr = context->GetCompiledInfo(); + OP_CHECK_IF(compileInfoPtr == nullptr, OPS_REPORT_CUBE_INNER_ERR(context->GetNodeName(), "compileInfoPtr is null"), + return ge::GRAPH_FAILED); + + return ge::GRAPH_SUCCESS; +} + +IMPL_OP_OPTILING(RecurrentKda) + .Tiling(RecurrentKdaTilingFunc) + .TilingParse(TilingPrepareForRecurrentKda); +} // namespace optiling diff --git a/csrc/attention/recurrent_kda/op_host/recurrent_kda_tiling.h b/csrc/attention/recurrent_kda/op_host/recurrent_kda_tiling.h new file mode 100644 index 000000000000..b69371cabe4d --- /dev/null +++ b/csrc/attention/recurrent_kda/op_host/recurrent_kda_tiling.h @@ -0,0 +1,75 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +/*! + * \file recurrent_kda_tiling.h + * \brief + */ +#ifndef __OP_HOST_RECURRENT_KDA_TILING_H__ +#define __OP_HOST_RECURRENT_KDA_TILING_H__ +#include +#include "register/tilingdata_base.h" +#include "tiling_base/tiling_base.h" +#include "err/ops_err.h" +#include "../op_kernel/recurrent_kda_tiling_data.h" +#include "recurrent_kda_tiling_processor.h" + +namespace optiling { +using namespace RecurrentKda; + +struct RecurrentKdaCompileInfo { + uint64_t aivNum{0UL}; + uint64_t ubSize{0UL}; +}; + +struct RecurrentKdaInfo { +public: + int64_t usedCoreNum = 0; + const char *opName = "RecurrentKda"; +}; + +class RecurrentKdaTiling : public Ops::Transformer::OpTiling::TilingBaseClass { +public: + explicit RecurrentKdaTiling(gert::TilingContext *context) : Ops::Transformer::OpTiling::TilingBaseClass(context) + { + InitCompileInfo(); + }; + ~RecurrentKdaTiling() override = default; + +protected: + bool IsCapable() override + { + return true; + } + ge::graphStatus GetPlatformInfo() override; + ge::graphStatus GetShapeAttrsInfo() override; + ge::graphStatus DoOpTiling() override; + ge::graphStatus DoLibApiTiling() override; + uint64_t GetTilingKey() const override; + ge::graphStatus GetWorkspaceSize() override; + ge::graphStatus PostTiling() override; + +protected: + void InitCompileInfo(); + void PrintTilingData(); + RecurrentKdaTilingContext BuildProcessorContext() const; + + ge::graphStatus CheckContext(); + ge::graphStatus AnalyzeDtype(); + ge::graphStatus AnalyzeShapes(); + ge::graphStatus CalUbSize(); + ge::graphStatus GetAttrsInfo(); + ge::graphStatus GetOptionalInput(); + ge::graphStatus AnalyzeFormat(); + bool CheckFormat(ge::Format format, const std::string &Desc); + + RecurrentKdaCompileInfo compileInfo_; + RecurrentKdaTilingData tilingData_; + RecurrentKdaInfo inputParams_; +}; + +} // namespace optiling +#endif // __OP_HOST_RECURRENT_KDA_TILING_H__ diff --git a/csrc/attention/recurrent_kda/op_host/recurrent_kda_tiling_processor.h b/csrc/attention/recurrent_kda/op_host/recurrent_kda_tiling_processor.h new file mode 100644 index 000000000000..4d716a9b22f0 --- /dev/null +++ b/csrc/attention/recurrent_kda/op_host/recurrent_kda_tiling_processor.h @@ -0,0 +1,678 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +/*! + * \file recurrent_kda_tiling_processor.h + * \brief Tiling processor shared by aclnn tiling and fast kernel launch. + */ + +#ifndef RECURRENT_KDA_TILING_PROCESSOR_H +#define RECURRENT_KDA_TILING_PROCESSOR_H + +#include "../op_kernel/recurrent_kda_struct.h" +#include "err/ops_err.h" +#include "log/log.h" +#include "register/op_impl_registry.h" +#include "tiling/tiling_api.h" +#include "util/math_util.h" +#include +#include +#include +#include + +using RecurrentKdaTilingData = RecurrentKda::RecurrentKdaTilingData; + +namespace optiling { + +static constexpr size_t RKDA_RANK3_QKV_DIM_NUM = 3; +static constexpr size_t RKDA_RANK4_QKV_DIM_NUM = 4; +static constexpr size_t RKDA_RANK2_BETA_DIM_NUM = 2; +static constexpr size_t RKDA_RANK3_BETA_DIM_NUM = 3; +static constexpr size_t RKDA_STATE_DIM_NUM = 4; +static constexpr size_t RKDA_METADATA_RANK1 = 1; +static constexpr size_t RKDA_METADATA_RANK2 = 2; + +static constexpr size_t RKDA_DIM_0 = 0; +static constexpr size_t RKDA_DIM_1 = 1; +static constexpr size_t RKDA_DIM_2 = 2; +static constexpr size_t RKDA_DIM_3 = 3; + +static constexpr uint32_t RKDA_LAYOUT_BSND = 0; +static constexpr uint32_t RKDA_LAYOUT_TND = 1; +static constexpr size_t RKDA_MAX_MTP = 8; +static constexpr int64_t RKDA_UB_GUARD_BYTES = 2048; +static constexpr size_t RKDA_SYS_WORKSPACE_SIZE = 16U * 1024U * 1024U; + +struct RecurrentKdaTilingContext { + const char *nodeName = "RecurrentKda"; + gert::Shape queryShape; + gert::Shape keyShape; + gert::Shape valueShape; + gert::Shape gateShape; + gert::Shape betaShape; + gert::Shape stateShape; + gert::Shape cuSeqlensShape; + gert::Shape ssmStateShape; + gert::Shape aLogShape; + gert::Shape dtBiasShape; + gert::Shape acceptedTokensShape; + float scale = 1.0f; + float lowerBound = -5.0f; + uint32_t layout = RKDA_LAYOUT_BSND; + uint32_t hasCuSeqlens = 0; + uint32_t hasSsmStateIndices = 0; + uint32_t hasALog = 0; + uint32_t hasDtBias = 0; + uint32_t hasAcceptedTokens = 0; + uint32_t useQkL2norm = 0; + uint32_t useGateInKernel = 0; + uint32_t useBetaSigmoid = 0; + uint32_t allowNegEigval = 0; + uint32_t safeGate = 0; + uint32_t stateVFirst = 0; + uint32_t outputFinalState = 0; + uint32_t inplaceFinalState = 1; + ge::DataType stateDtype = ge::DT_BF16; + ge::DataType gateDtype = ge::DT_FLOAT; + ge::DataType betaDtype = ge::DT_FLOAT; + ge::DataType cuSeqlensDtype = ge::DT_INT64; + ge::DataType ssmStateIndicesDtype = ge::DT_INT64; + ge::DataType acceptedTokensDtype = ge::DT_INT64; + std::array stateInStrides = {}; + std::array stateOutStrides = {}; + bool hasStateInStrides = false; + bool hasStateOutStrides = false; + uint64_t aivNum = 0; + uint64_t ubSize = 0; +}; + +class RecurrentKdaTilingProcessor { +public: + explicit RecurrentKdaTilingProcessor(const RecurrentKdaTilingContext &ctx) : ctx_(ctx) {} + + ge::graphStatus ProcessShapes(RecurrentKdaTilingData &tiling) const + { + tiling.vectorCoreNum = static_cast(ctx_.aivNum); + + struct RuleItem { + const char *name; + ge::graphStatus (RecurrentKdaTilingProcessor::*fn)(RecurrentKdaTilingData &) const; + }; + const std::array shapeRules = {{ + {"RuleCheckShapeDimAndRelation", &RecurrentKdaTilingProcessor::RuleCheckShapeDimAndRelation}, + {"RuleFillTilingShapeData", &RecurrentKdaTilingProcessor::RuleFillTilingShapeData}, + {"RuleCheckShapeValueRangeAndRule", &RecurrentKdaTilingProcessor::RuleCheckShapeValueRangeAndRule}, + {"RuleUpdateDynamicBlockDimByTaskUnits", &RecurrentKdaTilingProcessor::RuleUpdateDynamicBlockDimByTaskUnits}, + }}; + for (const auto &rule : shapeRules) { + OP_CHECK_IF((this->*(rule.fn))(tiling) != ge::GRAPH_SUCCESS, + OP_LOGE(ctx_.nodeName, "ProcessShapes rule failed: %s", rule.name), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; + } + + ge::graphStatus ProcessUb(RecurrentKdaTilingData &tiling) const + { + struct RuleItem { + const char *name; + ge::graphStatus (RecurrentKdaTilingProcessor::*fn)(RecurrentKdaTilingData &, UbCalcContext &) const; + }; + UbCalcContext ubCalcCtx; + const std::array ubRules = {{ + {"RuleInitUbCalcContext", &RecurrentKdaTilingProcessor::RuleInitUbCalcContext}, + {"RuleCalcFixedUbBytes", &RecurrentKdaTilingProcessor::RuleCalcFixedUbBytes}, + {"RuleCalcWorkingUbBytes", &RecurrentKdaTilingProcessor::RuleCalcWorkingUbBytes}, + {"RuleCalcVStepCoeff", &RecurrentKdaTilingProcessor::RuleCalcVStepCoeff}, + {"RuleFinalizeVStepFromUb", &RecurrentKdaTilingProcessor::RuleFinalizeVStepFromUb}, + }}; + for (const auto &rule : ubRules) { + OP_CHECK_IF((this->*(rule.fn))(tiling, ubCalcCtx) != ge::GRAPH_SUCCESS, + OP_LOGE(ctx_.nodeName, "ProcessUb rule failed: %s", rule.name), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; + } + +private: + struct UbCalcContext { + int64_t ubSize = 0; + int64_t aNv = 0; + int64_t aDv = 0; + int64_t aDk = 0; + int64_t fixedUbBytes = 0; + int64_t workingUbBytes = 0; + int64_t coeff = 0; + }; + + struct BufferProfile { + uint32_t stateOutBufferNum = 1; + uint32_t attnOutBufferNum = 1; + uint32_t vStep = 0; + uint32_t repeatTime = 0; + bool valid = false; + }; + + RecurrentKdaTilingContext ctx_; + + std::array ResolveStateStrides(bool output) const + { + const bool hasStrides = output ? ctx_.hasStateOutStrides : ctx_.hasStateInStrides; + if (hasStrides) { + return output ? ctx_.stateOutStrides : ctx_.stateInStrides; + } + const auto &shape = ctx_.stateShape; + std::array strides = {}; + strides[RKDA_DIM_3] = 1; + strides[RKDA_DIM_2] = shape.GetDim(RKDA_DIM_3); + strides[RKDA_DIM_1] = shape.GetDim(RKDA_DIM_2) * strides[RKDA_DIM_2]; + strides[RKDA_DIM_0] = shape.GetDim(RKDA_DIM_1) * strides[RKDA_DIM_1]; + return strides; + } + + ge::graphStatus ValidateStateStrides( + const std::array &strides, const char *name) const + { + for (size_t i = 0; i < RKDA_STATE_DIM_NUM; ++i) { + OP_CHECK_IF(strides[i] <= 0, + OP_LOGE(ctx_.nodeName, "%s stride[%zu] must be positive, but it is %ld.", + name, i, strides[i]), + return ge::GRAPH_FAILED); + } + const auto &shape = ctx_.stateShape; + const int64_t denseRow = shape.GetDim(RKDA_DIM_3); + const int64_t densePlane = shape.GetDim(RKDA_DIM_2) * denseRow; + OP_CHECK_IF(strides[RKDA_DIM_3] != 1 || strides[RKDA_DIM_2] != denseRow, + OP_LOGE(ctx_.nodeName, + "%s must keep its inner state matrix dense: stride[3]=1 and stride[2]=%ld, " + "but got [%ld, %ld, %ld, %ld].", + name, denseRow, strides[RKDA_DIM_0], strides[RKDA_DIM_1], + strides[RKDA_DIM_2], strides[RKDA_DIM_3]), + return ge::GRAPH_FAILED); + const int64_t minSlotStride = + (shape.GetDim(RKDA_DIM_1) - 1) * strides[RKDA_DIM_1] + densePlane; + OP_CHECK_IF(strides[RKDA_DIM_1] < densePlane || + strides[RKDA_DIM_0] < minSlotStride, + OP_LOGE(ctx_.nodeName, + "%s outer strides overlap state rows or heads: [%ld, %ld, %ld, %ld].", + name, strides[RKDA_DIM_0], strides[RKDA_DIM_1], + strides[RKDA_DIM_2], strides[RKDA_DIM_3]), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; + } + + ge::graphStatus CheckStateStrides() const + { + if (ValidateStateStrides(ResolveStateStrides(false), "initial_state") != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + if (ValidateStateStrides(ResolveStateStrides(true), "state output") != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; + } + + bool CheckDim(const gert::Shape &shape, const size_t dim, const std::string &dimDesc) const + { + if (shape.GetDimNum() != dim) { + OP_LOGE(ctx_.nodeName, "The number of dimensions of %s should be %zu, but it is %zu", + dimDesc.c_str(), dim, shape.GetDimNum()); + return false; + } + return true; + } + + bool CheckDimEqual(const gert::Shape &a, const int64_t dimA, const gert::Shape &b, const int64_t dimB, + const std::string &nameA, const std::string &nameB, const std::string &dimDesc) const + { + if (a.GetDim(dimA) != b.GetDim(dimB)) { + OP_LOGE(ctx_.nodeName, "The %s of %s and %s should be the same, but %s is %ld while %s is %ld", + dimDesc.c_str(), nameA.c_str(), nameB.c_str(), nameA.c_str(), a.GetDim(dimA), nameB.c_str(), + b.GetDim(dimB)); + return false; + } + return true; + } + + int64_t SeqNum(const gert::Shape &queryShape, const gert::Shape &cuSeqlensShape) const + { + if (ctx_.hasCuSeqlens) { + return cuSeqlensShape.GetDim(RKDA_DIM_0) - 1; + } + return ctx_.layout == RKDA_LAYOUT_BSND ? queryShape.GetDim(RKDA_DIM_0) : 1; + } + + ge::graphStatus CheckMetadataShapes(const gert::Shape &queryShape, const gert::Shape &cuSeqlensShape) const + { + if (ctx_.hasCuSeqlens) { + if (!CheckDim(cuSeqlensShape, RKDA_METADATA_RANK1, "cu_seqlens")) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF(cuSeqlensShape.GetDim(RKDA_DIM_0) < 2, + OP_LOGE(ctx_.nodeName, "cu_seqlens must contain at least 2 elements."), + return ge::GRAPH_FAILED); + } + int64_t totalTokens = + (ctx_.layout == RKDA_LAYOUT_TND) ? queryShape.GetDim(RKDA_DIM_0) : + queryShape.GetDim(RKDA_DIM_0) * queryShape.GetDim(RKDA_DIM_1); + int64_t seqNum = SeqNum(queryShape, cuSeqlensShape); + + if (ctx_.hasSsmStateIndices) { + size_t rank = ctx_.ssmStateShape.GetDimNum(); + bool packed1d = rank == RKDA_METADATA_RANK1 && + ctx_.ssmStateShape.GetDim(RKDA_DIM_0) >= totalTokens; + bool speculative2d = rank == RKDA_METADATA_RANK2 && + ctx_.ssmStateShape.GetDim(RKDA_DIM_0) == seqNum && + ctx_.ssmStateShape.GetDim(RKDA_DIM_1) > 0; + OP_CHECK_IF(!packed1d && !speculative2d, + OP_LOGE(ctx_.nodeName, + "ssm_state_indices must be packed [T] or speculative [seq_num,max_step]."), + return ge::GRAPH_FAILED); + } + if (ctx_.hasAcceptedTokens) { + if (!CheckDim(ctx_.acceptedTokensShape, RKDA_METADATA_RANK1, "num_accepted_tokens")) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF(ctx_.acceptedTokensShape.GetDim(RKDA_DIM_0) != seqNum, + OP_LOGE(ctx_.nodeName, "num_accepted_tokens length must equal sequence number."), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; + } + + ge::graphStatus CheckOptionalGateShapes(int64_t hvNum, int64_t kDim) const + { + if (ctx_.useGateInKernel && !ctx_.hasALog) { + OP_LOGE(ctx_.nodeName, "A_log is required when use_gate_in_kernel is true."); + return ge::GRAPH_FAILED; + } + if (ctx_.hasALog) { + if (!CheckDim(ctx_.aLogShape, RKDA_METADATA_RANK1, "A_log")) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF(ctx_.aLogShape.GetDim(RKDA_DIM_0) != hvNum, + OP_LOGE(ctx_.nodeName, "A_log shape must be [HV]."), + return ge::GRAPH_FAILED); + } + if (ctx_.hasDtBias) { + size_t rank = ctx_.dtBiasShape.GetDimNum(); + bool valid = (rank == 1 && ctx_.dtBiasShape.GetDim(RKDA_DIM_0) == hvNum * kDim) || + (rank == 2 && ctx_.dtBiasShape.GetDim(RKDA_DIM_0) == hvNum && + ctx_.dtBiasShape.GetDim(RKDA_DIM_1) == kDim); + OP_CHECK_IF(!valid, OP_LOGE(ctx_.nodeName, "dt_bias must be [HV*K] or [HV, K]."), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; + } + + ge::graphStatus CheckShapeDimAndRelation(const gert::Shape &queryShape, const gert::Shape &keyShape, + const gert::Shape &valueShape, const gert::Shape &gateShape, + const gert::Shape &betaShape, const gert::Shape &stateShape, + const gert::Shape &cuSeqlensShape) const + { + int64_t totalTokens = 0; + int64_t denseSeqLen = 0; + int64_t hNum = 0; + int64_t hvNum = 0; + int64_t kDim = 0; + int64_t vDim = 0; + if (ctx_.layout == RKDA_LAYOUT_TND) { + if (!CheckDim(queryShape, RKDA_RANK3_QKV_DIM_NUM, "query") || + !CheckDim(keyShape, RKDA_RANK3_QKV_DIM_NUM, "key") || + !CheckDim(valueShape, RKDA_RANK3_QKV_DIM_NUM, "value") || + !CheckDim(gateShape, RKDA_RANK3_QKV_DIM_NUM, "gate") || + !CheckDim(betaShape, RKDA_RANK2_BETA_DIM_NUM, "beta")) { + return ge::GRAPH_FAILED; + } + totalTokens = queryShape.GetDim(RKDA_DIM_0); + denseSeqLen = totalTokens; + hNum = queryShape.GetDim(RKDA_DIM_1); + kDim = queryShape.GetDim(RKDA_DIM_2); + hvNum = valueShape.GetDim(RKDA_DIM_1); + vDim = valueShape.GetDim(RKDA_DIM_2); + OP_CHECK_IF(valueShape.GetDim(RKDA_DIM_0) != totalTokens || + gateShape.GetDim(RKDA_DIM_0) != totalTokens || + betaShape.GetDim(RKDA_DIM_0) != totalTokens || + gateShape.GetDim(RKDA_DIM_1) != hvNum || + betaShape.GetDim(RKDA_DIM_1) != hvNum || + gateShape.GetDim(RKDA_DIM_2) != kDim, + OP_LOGE(ctx_.nodeName, + "TND expects q/k [T,H,K], value [T,HV,V], gate [T,HV,K], beta [T,HV]."), + return ge::GRAPH_FAILED); + } else { + if (!CheckDim(queryShape, RKDA_RANK4_QKV_DIM_NUM, "query") || + !CheckDim(keyShape, RKDA_RANK4_QKV_DIM_NUM, "key") || + !CheckDim(valueShape, RKDA_RANK4_QKV_DIM_NUM, "value") || + !CheckDim(gateShape, RKDA_RANK4_QKV_DIM_NUM, "gate") || + !CheckDim(betaShape, RKDA_RANK3_BETA_DIM_NUM, "beta")) { + return ge::GRAPH_FAILED; + } + int64_t batch = queryShape.GetDim(RKDA_DIM_0); + denseSeqLen = queryShape.GetDim(RKDA_DIM_1); + totalTokens = batch * denseSeqLen; + hNum = queryShape.GetDim(RKDA_DIM_2); + kDim = queryShape.GetDim(RKDA_DIM_3); + hvNum = valueShape.GetDim(RKDA_DIM_2); + vDim = valueShape.GetDim(RKDA_DIM_3); + OP_CHECK_IF(valueShape.GetDim(RKDA_DIM_0) != batch || + valueShape.GetDim(RKDA_DIM_1) != denseSeqLen || + gateShape.GetDim(RKDA_DIM_0) != batch || + gateShape.GetDim(RKDA_DIM_1) != denseSeqLen || + gateShape.GetDim(RKDA_DIM_2) != hvNum || + gateShape.GetDim(RKDA_DIM_3) != kDim || + betaShape.GetDim(RKDA_DIM_0) != batch || + betaShape.GetDim(RKDA_DIM_1) != denseSeqLen || + betaShape.GetDim(RKDA_DIM_2) != hvNum, + OP_LOGE(ctx_.nodeName, + "BSND expects q/k [B,T,H,K], value [B,T,HV,V], gate [B,T,HV,K], beta [B,T,HV]."), + return ge::GRAPH_FAILED); + } + + if (!CheckDimEqual(queryShape, RKDA_DIM_0, keyShape, RKDA_DIM_0, "query", "key", + "leading token dimension") || + !CheckDimEqual(queryShape, queryShape.GetDimNum() - 2, keyShape, keyShape.GetDimNum() - 2, + "query", "key", "H dimension") || + !CheckDimEqual(queryShape, queryShape.GetDimNum() - 1, keyShape, keyShape.GetDimNum() - 1, + "query", "key", "K dimension")) { + return ge::GRAPH_FAILED; + } + + OP_CHECK_IF(hNum <= 0 || hvNum <= 0 || kDim <= 0 || vDim <= 0 || totalTokens <= 0 || denseSeqLen <= 0, + OP_LOGE(ctx_.nodeName, "input shape dimensions must be positive."), return ge::GRAPH_FAILED); + OP_CHECK_IF(hvNum % hNum != 0, + OP_LOGE(ctx_.nodeName, "HV must be an integer multiple of H, but HV is %ld and H is %ld.", + hvNum, hNum), + return ge::GRAPH_FAILED); + if (!CheckDim(stateShape, RKDA_STATE_DIM_NUM, "initial_state")) { + return ge::GRAPH_FAILED; + } + int64_t seqNum = SeqNum(queryShape, cuSeqlensShape); + bool stateTailMatches = ctx_.stateVFirst ? + (stateShape.GetDim(RKDA_DIM_2) == vDim && stateShape.GetDim(RKDA_DIM_3) == kDim) : + (stateShape.GetDim(RKDA_DIM_2) == kDim && stateShape.GetDim(RKDA_DIM_3) == vDim); + OP_CHECK_IF(stateShape.GetDim(RKDA_DIM_0) <= 0 || + (!ctx_.hasSsmStateIndices && stateShape.GetDim(RKDA_DIM_0) != seqNum) || + stateShape.GetDim(RKDA_DIM_1) != hvNum || !stateTailMatches, + OP_LOGE(ctx_.nodeName, + "state must be [state_capacity, HV, V, K] when state_v_first=true or " + "[state_capacity, HV, K, V] otherwise; without ssm_state_indices, " + "state_capacity must equal seq_num."), + return ge::GRAPH_FAILED); + + if (CheckStateStrides() != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + + if (CheckMetadataShapes(queryShape, cuSeqlensShape) != ge::GRAPH_SUCCESS || + CheckOptionalGateShapes(hvNum, kDim) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; + } + + void FillTilingShapeData(const gert::Shape &queryShape, const gert::Shape &valueShape, + const gert::Shape &stateShape, const gert::Shape &cuSeqlensShape, + RecurrentKdaTilingData &tiling) const + { + if (ctx_.layout == RKDA_LAYOUT_TND) { + tiling.t = static_cast(queryShape.GetDim(RKDA_DIM_0)); + tiling.seqLen = static_cast(queryShape.GetDim(RKDA_DIM_0)); + tiling.nk = static_cast(queryShape.GetDim(RKDA_DIM_1)); + tiling.dk = static_cast(queryShape.GetDim(RKDA_DIM_2)); + tiling.nv = static_cast(valueShape.GetDim(RKDA_DIM_1)); + tiling.dv = static_cast(valueShape.GetDim(RKDA_DIM_2)); + tiling.b = static_cast(SeqNum(queryShape, cuSeqlensShape)); + } else { + tiling.seqLen = static_cast(queryShape.GetDim(RKDA_DIM_1)); + tiling.t = static_cast(queryShape.GetDim(RKDA_DIM_0) * queryShape.GetDim(RKDA_DIM_1)); + tiling.nk = static_cast(queryShape.GetDim(RKDA_DIM_2)); + tiling.dk = static_cast(queryShape.GetDim(RKDA_DIM_3)); + tiling.nv = static_cast(valueShape.GetDim(RKDA_DIM_2)); + tiling.dv = static_cast(valueShape.GetDim(RKDA_DIM_3)); + tiling.b = static_cast(SeqNum(queryShape, cuSeqlensShape)); + } + tiling.sBlockNum = static_cast(stateShape.GetDim(RKDA_DIM_0)); + tiling.ssmStateStride = (ctx_.hasSsmStateIndices && ctx_.ssmStateShape.GetDimNum() == RKDA_METADATA_RANK2) ? + static_cast(ctx_.ssmStateShape.GetDim(RKDA_DIM_1)) : 0; + const auto stateInStrides = ResolveStateStrides(false); + const auto stateOutStrides = ResolveStateStrides(true); + tiling.stateInStride0 = static_cast(stateInStrides[RKDA_DIM_0]); + tiling.stateInStride1 = static_cast(stateInStrides[RKDA_DIM_1]); + tiling.stateInStride2 = static_cast(stateInStrides[RKDA_DIM_2]); + tiling.stateInStride3 = static_cast(stateInStrides[RKDA_DIM_3]); + tiling.stateOutStride0 = static_cast(stateOutStrides[RKDA_DIM_0]); + tiling.stateOutStride1 = static_cast(stateOutStrides[RKDA_DIM_1]); + tiling.stateOutStride2 = static_cast(stateOutStrides[RKDA_DIM_2]); + tiling.stateOutStride3 = static_cast(stateOutStrides[RKDA_DIM_3]); + tiling.scale = ctx_.scale; + tiling.lowerBound = ctx_.lowerBound; + tiling.layout = ctx_.layout; + tiling.hasCuSeqlens = ctx_.hasCuSeqlens; + tiling.hasSsmStateIndices = ctx_.hasSsmStateIndices; + tiling.hasALog = ctx_.hasALog; + tiling.hasDtBias = ctx_.hasDtBias; + tiling.hasAcceptedTokens = ctx_.hasAcceptedTokens; + tiling.useQkL2norm = ctx_.useQkL2norm; + tiling.useGateInKernel = ctx_.useGateInKernel; + tiling.useBetaSigmoid = ctx_.useBetaSigmoid; + tiling.allowNegEigval = ctx_.allowNegEigval; + tiling.safeGate = ctx_.safeGate; + tiling.stateVFirst = ctx_.stateVFirst; + tiling.outputFinalState = ctx_.outputFinalState; + tiling.inplaceFinalState = ctx_.inplaceFinalState; + tiling.gateDtype = ctx_.gateDtype == ge::DT_FLOAT ? 0 : (ctx_.gateDtype == ge::DT_BF16 ? 1 : 2); + tiling.betaDtype = ctx_.betaDtype == ge::DT_FLOAT ? 0 : (ctx_.betaDtype == ge::DT_BF16 ? 1 : 2); + tiling.cuSeqlensDtype = ctx_.cuSeqlensDtype == ge::DT_INT32 ? 0 : 1; + tiling.ssmStateIndicesDtype = ctx_.ssmStateIndicesDtype == ge::DT_INT32 ? 0 : 1; + tiling.acceptedTokensDtype = ctx_.acceptedTokensDtype == ge::DT_INT32 ? 0 : 1; + } + + ge::graphStatus CheckShapeValueRangeAndRule(const RecurrentKdaTilingData &tiling) const + { + OP_CHECK_IF(tiling.nk > 256 || tiling.nv > 256, + OP_LOGE(ctx_.nodeName, + "H/HV must be <= 256, but H=%u, HV=%u.", tiling.nk, tiling.nv), + return ge::GRAPH_FAILED); + OP_CHECK_IF(tiling.dk != 128 || (tiling.dv != 128 && tiling.dv != 256), + OP_LOGE(ctx_.nodeName, + "K/V currently support only K=128,V=128 or K=128,V=256, but K=%u, V=%u.", + tiling.dk, tiling.dv), + return ge::GRAPH_FAILED); + OP_CHECK_IF(tiling.nv % tiling.nk != 0, + OP_LOGE(ctx_.nodeName, "HV must be divisible by H, but HV=%u and H=%u.", + tiling.nv, tiling.nk), + return ge::GRAPH_FAILED); + OP_CHECK_IF(tiling.safeGate && (tiling.lowerBound < -5.0f || tiling.lowerBound >= 0.0f), + OP_LOGE(ctx_.nodeName, "lower_bound must be in [-5, 0) when safe_gate is true."), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; + } + + void UpdateDynamicBlockDimByTaskUnits(RecurrentKdaTilingData &tiling) const + { + uint64_t taskUnits = static_cast(tiling.b) * static_cast(tiling.nv); + if (taskUnits == 0) { + taskUnits = 1; + } + uint64_t maxCoreNum = (ctx_.aivNum > 0) ? ctx_.aivNum : 1; + uint64_t selectedCoreNum = (taskUnits < maxCoreNum) ? taskUnits : maxCoreNum; + tiling.vectorCoreNum = static_cast(selectedCoreNum); + OP_LOGD(ctx_.nodeName, "taskUnits: [%llu], selected vectorCoreNum: [%u]", + static_cast(taskUnits), tiling.vectorCoreNum); + } + + ge::graphStatus RuleCheckShapeDimAndRelation(RecurrentKdaTilingData &tiling) const + { + (void)tiling; + return CheckShapeDimAndRelation(ctx_.queryShape, ctx_.keyShape, ctx_.valueShape, ctx_.gateShape, + ctx_.betaShape, ctx_.stateShape, ctx_.cuSeqlensShape); + } + + ge::graphStatus RuleFillTilingShapeData(RecurrentKdaTilingData &tiling) const + { + FillTilingShapeData(ctx_.queryShape, ctx_.valueShape, ctx_.stateShape, ctx_.cuSeqlensShape, tiling); + return ge::GRAPH_SUCCESS; + } + + ge::graphStatus RuleCheckShapeValueRangeAndRule(RecurrentKdaTilingData &tiling) const + { + return CheckShapeValueRangeAndRule(tiling); + } + + ge::graphStatus RuleUpdateDynamicBlockDimByTaskUnits(RecurrentKdaTilingData &tiling) const + { + UpdateDynamicBlockDimByTaskUnits(tiling); + return ge::GRAPH_SUCCESS; + } + + int64_t CalcFixedUbBytes(int64_t aNv, int64_t aDv, int64_t aDk) const + { + int64_t usedUbBytes = 0; + usedUbBytes += RKDA_MAX_MTP * (4 * aDk + 2 * aDv); // q/k/v input queues, bf16. + usedUbBytes += RKDA_MAX_MTP * 4 * aDk; // gate input queue, fp32. + usedUbBytes += RKDA_MAX_MTP * 4 * aNv; // beta input queue, fp32. + usedUbBytes += 64 + RKDA_UB_GUARD_BYTES; // scalar buffer and guard. + return usedUbBytes; + } + + int64_t CalcWorkingUbBytes(int64_t aNv, int64_t aDv, int64_t aDk) const + { + int64_t usedUbBytes = CalcFixedUbBytes(aNv, aDv, aDk); + usedUbBytes += RKDA_MAX_MTP * (4 * aDv + 12 * aDk + 4 * aNv); + return usedUbBytes; + } + + int64_t CalcVStepCoeff(int64_t aDk, uint32_t stateOutBufferNum, uint32_t attnOutBufferNum) const + { + int64_t stateDtypeSize = (ctx_.stateDtype == ge::DT_FLOAT) ? 4 : 2; + int64_t coeff = stateDtypeSize * aDk; // state input queue. + coeff += static_cast(stateOutBufferNum) * stateDtypeSize * aDk; + coeff += static_cast(attnOutBufferNum) * 2; + coeff += 8 * aDk + 8; + return coeff; + } + + bool EvaluateBufferProfile(int64_t ubSize, int64_t usedUbBytes, int64_t aDk, uint32_t stateOutBufferNum, + uint32_t attnOutBufferNum, const RecurrentKdaTilingData &tiling, + BufferProfile &profile) const + { + int64_t coeff = CalcVStepCoeff(aDk, stateOutBufferNum, attnOutBufferNum); + int64_t vStep = (ubSize - usedUbBytes) / coeff / 8 * 8; + if (vStep < static_cast(RKDA_MAX_MTP)) { + return false; + } + int64_t repeatTime = Ops::Base::CeilDiv(tiling.dv, static_cast(vStep)); + vStep = Ops::Base::CeilAlign(Ops::Base::CeilDiv(tiling.dv, static_cast(repeatTime)), + static_cast(8)); + if (vStep < static_cast(RKDA_MAX_MTP)) { + return false; + } + profile.stateOutBufferNum = stateOutBufferNum; + profile.attnOutBufferNum = attnOutBufferNum; + profile.vStep = static_cast(vStep); + profile.repeatTime = static_cast(repeatTime); + profile.valid = true; + return true; + } + + bool IsBetterProfile(const BufferProfile &candidate, const BufferProfile ¤t) const + { + if (!current.valid) { + return true; + } + if (candidate.repeatTime != current.repeatTime) { + return candidate.repeatTime < current.repeatTime; + } + uint32_t candidateDepth = candidate.stateOutBufferNum + candidate.attnOutBufferNum; + uint32_t currentDepth = current.stateOutBufferNum + current.attnOutBufferNum; + if (candidateDepth != currentDepth) { + return candidateDepth > currentDepth; + } + return candidate.vStep > current.vStep; + } + + ge::graphStatus FinalizeVStepFromUb(int64_t ubSize, int64_t usedUbBytes, int64_t coeff, + RecurrentKdaTilingData &tiling, UbCalcContext &ubCalcCtx) const + { + (void)coeff; + int64_t aDk = Ops::Base::CeilAlign(tiling.dk, static_cast(16)); + BufferProfile selected; + const std::array candidates = {{ + {1, 1, 0, 0, false}, + {1, 2, 0, 0, false}, + {2, 2, 0, 0, false}, + }}; + for (const auto &candidate : candidates) { + BufferProfile profile; + if (!EvaluateBufferProfile(ubSize, usedUbBytes, aDk, candidate.stateOutBufferNum, + candidate.attnOutBufferNum, tiling, profile)) { + continue; + } + if (IsBetterProfile(profile, selected)) { + selected = profile; + } + } + + if (!selected.valid) { + OP_LOGE(ctx_.nodeName, "vStep should be at least %zu, shape is too big", RKDA_MAX_MTP); + return ge::GRAPH_FAILED; + } + + int64_t queueCoeff = CalcVStepCoeff(aDk, selected.stateOutBufferNum, selected.attnOutBufferNum) - + (8 * aDk + 8); + int64_t ubRestBytes = ubSize - ubCalcCtx.fixedUbBytes - + queueCoeff * static_cast(selected.vStep); + if (ubRestBytes < 0) { + OP_LOGE(ctx_.nodeName, "ubRestBytes should be non-negative, but got %ld", ubRestBytes); + return ge::GRAPH_FAILED; + } + tiling.ubCalSize = static_cast(ctx_.ubSize); + tiling.vStep = selected.vStep; + tiling.stateOutBufferNum = selected.stateOutBufferNum; + tiling.attnOutBufferNum = selected.attnOutBufferNum; + tiling.ubRestBytes = static_cast(ubRestBytes); + return ge::GRAPH_SUCCESS; + } + + ge::graphStatus RuleInitUbCalcContext(RecurrentKdaTilingData &tiling, UbCalcContext &ubCalcCtx) const + { + ubCalcCtx.ubSize = static_cast(ctx_.ubSize); + ubCalcCtx.aNv = Ops::Base::CeilAlign(tiling.nv, static_cast(16)); + ubCalcCtx.aDv = Ops::Base::CeilAlign(tiling.dv, static_cast(16)); + ubCalcCtx.aDk = Ops::Base::CeilAlign(tiling.dk, static_cast(16)); + return ge::GRAPH_SUCCESS; + } + + ge::graphStatus RuleCalcFixedUbBytes(RecurrentKdaTilingData &tiling, UbCalcContext &ubCalcCtx) const + { + ubCalcCtx.fixedUbBytes = CalcFixedUbBytes(ubCalcCtx.aNv, ubCalcCtx.aDv, ubCalcCtx.aDk); + tiling.ubRestBytes = static_cast(ubCalcCtx.ubSize - ubCalcCtx.fixedUbBytes); + return ge::GRAPH_SUCCESS; + } + + ge::graphStatus RuleCalcWorkingUbBytes(RecurrentKdaTilingData &tiling, UbCalcContext &ubCalcCtx) const + { + (void)tiling; + ubCalcCtx.workingUbBytes = CalcWorkingUbBytes(ubCalcCtx.aNv, ubCalcCtx.aDv, ubCalcCtx.aDk); + return ge::GRAPH_SUCCESS; + } + + ge::graphStatus RuleCalcVStepCoeff(RecurrentKdaTilingData &tiling, UbCalcContext &ubCalcCtx) const + { + (void)tiling; + ubCalcCtx.coeff = CalcVStepCoeff(ubCalcCtx.aDk, 1, 1); + return ge::GRAPH_SUCCESS; + } + + ge::graphStatus RuleFinalizeVStepFromUb(RecurrentKdaTilingData &tiling, UbCalcContext &ubCalcCtx) const + { + return FinalizeVStepFromUb(ubCalcCtx.ubSize, ubCalcCtx.workingUbBytes, ubCalcCtx.coeff, tiling, ubCalcCtx); + } +}; + +} // namespace optiling + +#endif // RECURRENT_KDA_TILING_PROCESSOR_H diff --git a/csrc/attention/recurrent_kda/op_kernel/arch35/recurrent_kda.h b/csrc/attention/recurrent_kda/op_kernel/arch35/recurrent_kda.h new file mode 100644 index 000000000000..7dc6ac58b0b8 --- /dev/null +++ b/csrc/attention/recurrent_kda/op_kernel/arch35/recurrent_kda.h @@ -0,0 +1,1004 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +/*! + * \file recurrent_kda.h + * \brief Single-kernel fused recurrent KDA implementation. + */ + +#ifndef __RECURRENT_KDA_KERNEL_H_ +#define __RECURRENT_KDA_KERNEL_H_ + +#include "kernel_operator.h" +#include "lib/matmul_intf.h" +#include "../recurrent_kda_tiling_data.h" + +namespace RecurrentKda { + +using namespace matmul; +using namespace AscendC; +using namespace AscendC::MicroAPI; +constexpr uint64_t BUFFER_NUM = 1; +constexpr uint32_t MAX_OUT_BUFFER_NUM = 2; +constexpr uint64_t MAX_MTP = 8; +constexpr uint64_t BF16_NUM_PER_BLOCK = 16; +constexpr uint64_t FP32_NUM_PER_BLOCK = 8; +constexpr uint32_t REPEAT_LENTH = 64; // 256 bytes for float. +constexpr uint32_t MAX_REPEAT_TIME = 255; +constexpr uint32_t ADD_FOLD_REDUCE_MIN_K = 128; +constexpr uint16_t V_LENGTH = VECTOR_REG_WIDTH / sizeof(float); +constexpr uint16_t TWO_V_LENGTH = 2 * V_LENGTH; +constexpr uint64_t INVALID_STATE_SLOT = static_cast(-1); + +#ifndef RKDA_ENABLE_ADD_FOLD_REDUCE +#define RKDA_ENABLE_ADD_FOLD_REDUCE 1 +#endif + +struct RKDAInitParams { + GM_ADDR query; + GM_ADDR key; + GM_ADDR value; + GM_ADDR gate; + GM_ADDR beta; + GM_ADDR initState; + GM_ADDR cuSeqlens; + GM_ADDR ssmStateIndices; + GM_ADDR aLog; + GM_ADDR dtBias; + GM_ADDR numAcceptedTokens; + GM_ADDR attnOut; + GM_ADDR finalState; +}; + +template +class RKDA { +public: + __aicore__ inline explicit RKDA(const RecurrentKdaTilingData *tilingData) + { + B_ = tilingData->b; + T_ = tilingData->t; + seqLen_ = tilingData->seqLen; + NK_ = tilingData->nk; + realK_ = tilingData->dk; + NV_ = tilingData->nv; + realV_ = tilingData->dv; + stateCapacity_ = tilingData->sBlockNum; + ssmStateStride_ = tilingData->ssmStateStride; + stateInStride0_ = tilingData->stateInStride0; + stateInStride1_ = tilingData->stateInStride1; + stateInStride2_ = tilingData->stateInStride2; + stateInStride3_ = tilingData->stateInStride3; + stateOutStride0_ = tilingData->stateOutStride0; + stateOutStride1_ = tilingData->stateOutStride1; + stateOutStride2_ = tilingData->stateOutStride2; + stateOutStride3_ = tilingData->stateOutStride3; + scale_ = tilingData->scale; + lowerBound_ = tilingData->lowerBound; + hasCuSeqlens_ = (tilingData->hasCuSeqlens == 1); + hasSsmStateIndices_ = (tilingData->hasSsmStateIndices == 1); + hasAcceptedTokens_ = (tilingData->hasAcceptedTokens == 1); + hasALog_ = (tilingData->hasALog == 1); + hasDtBias_ = (tilingData->hasDtBias == 1); + useQkL2norm_ = (tilingData->useQkL2norm == 1); + useGateInKernel_ = (tilingData->useGateInKernel == 1); + useBetaSigmoid_ = (tilingData->useBetaSigmoid == 1); + allowNegEigval_ = (tilingData->allowNegEigval == 1); + safeGate_ = (tilingData->safeGate == 1); + stateVFirst_ = (tilingData->stateVFirst == 1); + shouldStoreState_ = (tilingData->inplaceFinalState == 1 || tilingData->outputFinalState == 1); + gateDtype_ = tilingData->gateDtype; + betaDtype_ = tilingData->betaDtype; + cuSeqlensDtype_ = tilingData->cuSeqlensDtype; + ssmStateIndicesDtype_ = tilingData->ssmStateIndicesDtype; + acceptedTokensDtype_ = tilingData->acceptedTokensDtype; + useAddFoldReduce_ = (RKDA_ENABLE_ADD_FOLD_REDUCE != 0); + vStep_ = tilingData->vStep; + stateOutBufferNum_ = (tilingData->stateOutBufferNum == MAX_OUT_BUFFER_NUM) ? MAX_OUT_BUFFER_NUM : BUFFER_NUM; + attnOutBufferNum_ = (tilingData->attnOutBufferNum == MAX_OUT_BUFFER_NUM) ? MAX_OUT_BUFFER_NUM : BUFFER_NUM; + restUbSize_ = tilingData->ubRestBytes; + alignK_ = Ceil(tilingData->dk, BF16_NUM_PER_BLOCK) * BF16_NUM_PER_BLOCK; + alignV_ = Ceil(tilingData->dv, BF16_NUM_PER_BLOCK) * BF16_NUM_PER_BLOCK; + eventMte2ToVInitialized_ = false; + eventVToMte2Initialized_ = false; + eventVToSInitialized_ = false; + } + + __aicore__ inline void Init(const RKDAInitParams &initParams, TPipe *pipe) + { + uint64_t blockDim = GetBlockNum(); + blockIdx = GetBlockIdx(); + if (blockIdx >= blockDim) { + return; + } + pipe_ = pipe; + SetGlobalTensors(initParams); + InitLocalBuffers(); + } + + __aicore__ inline void SetGlobalTensors(const RKDAInitParams &initParams) + { + queryGm_.SetGlobalBuffer((__gm__ inType *)initParams.query); + keyGm_.SetGlobalBuffer((__gm__ inType *)initParams.key); + valueGm_.SetGlobalBuffer((__gm__ inType *)initParams.value); + gateFloatGm_.SetGlobalBuffer((__gm__ float *)initParams.gate); + gateBf16Gm_.SetGlobalBuffer((__gm__ bfloat16_t *)initParams.gate); + gateFp16Gm_.SetGlobalBuffer((__gm__ half *)initParams.gate); + betaFloatGm_.SetGlobalBuffer((__gm__ float *)initParams.beta); + betaBf16Gm_.SetGlobalBuffer((__gm__ bfloat16_t *)initParams.beta); + betaFp16Gm_.SetGlobalBuffer((__gm__ half *)initParams.beta); + initStateGm_.SetGlobalBuffer((__gm__ stateType *)initParams.initState); + cuSeqlensInt32Gm_.SetGlobalBuffer((__gm__ int32_t *)initParams.cuSeqlens); + cuSeqlensInt64Gm_.SetGlobalBuffer((__gm__ int64_t *)initParams.cuSeqlens); + ssmStateIndicesInt32Gm_.SetGlobalBuffer((__gm__ int32_t *)initParams.ssmStateIndices); + ssmStateIndicesInt64Gm_.SetGlobalBuffer((__gm__ int64_t *)initParams.ssmStateIndices); + aLogGm_.SetGlobalBuffer((__gm__ float *)initParams.aLog); + dtBiasGm_.SetGlobalBuffer((__gm__ float *)initParams.dtBias); + numAcceptedTokensInt32Gm_.SetGlobalBuffer((__gm__ int32_t *)initParams.numAcceptedTokens); + numAcceptedTokensInt64Gm_.SetGlobalBuffer((__gm__ int64_t *)initParams.numAcceptedTokens); + finalStateGm_.SetGlobalBuffer((__gm__ stateType *)initParams.finalState); + attnOutGm_.SetGlobalBuffer((__gm__ outType *)initParams.attnOut); + } + + __aicore__ inline void InitLocalBuffers() + { + uint32_t cubeSize = alignK_ * vStep_ * sizeof(float); + uint32_t singleVSize = vStep_ * sizeof(float); + uint32_t vSize = MAX_MTP * alignV_ * sizeof(float); + uint32_t kSize = MAX_MTP * alignK_ * sizeof(float); + uint32_t betaUbSize = + Ceil(MAX_MTP * NV_, FP32_NUM_PER_BLOCK) * FP32_NUM_PER_BLOCK * sizeof(float); + pipe_->InitBuffer(qInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(inType)); + pipe_->InitBuffer(kInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(inType)); + pipe_->InitBuffer(vInQueue_, BUFFER_NUM, MAX_MTP * alignV_ * sizeof(inType)); + pipe_->InitBuffer(gateInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(float)); + pipe_->InitBuffer(betaInQueue_, BUFFER_NUM, betaUbSize); + pipe_->InitBuffer(stateInQueue_, BUFFER_NUM, alignK_ * vStep_ * sizeof(stateType)); + pipe_->InitBuffer(stateOutQueue_, stateOutBufferNum_, alignK_ * vStep_ * sizeof(stateType)); + pipe_->InitBuffer(attnOutQueue_, attnOutBufferNum_, vStep_ * sizeof(outType)); + pipe_->InitBuffer(tmpBuff, restUbSize_); + pipe_->InitBuffer(scalarBuf_, 64); + + uint32_t buffOffset = 0; + deltaInUb = tmpBuff.GetWithOffset(static_cast(vStep_), buffOffset); + buffOffset += singleVSize; + attnInUb = tmpBuff.GetWithOffset(static_cast(vStep_), buffOffset); + buffOffset += singleVSize; + vInUb = tmpBuff.GetWithOffset(static_cast(MAX_MTP * alignV_), buffOffset); + buffOffset += vSize; + qInUb = tmpBuff.GetWithOffset(static_cast(MAX_MTP * alignK_), buffOffset); + buffOffset += kSize; + kInUb = tmpBuff.GetWithOffset(static_cast(MAX_MTP * alignK_), buffOffset); + buffOffset += kSize; + stateInUb = tmpBuff.GetWithOffset(static_cast(alignK_ * vStep_), buffOffset); + buffOffset += cubeSize; + broadTmpInUb = tmpBuff.GetWithOffset(static_cast(alignK_ * vStep_), buffOffset); + buffOffset += cubeSize; + betaInUb = tmpBuff.GetWithOffset(static_cast(betaUbSize / sizeof(float)), buffOffset); + buffOffset += betaUbSize; + gateInUb = tmpBuff.GetWithOffset(static_cast(MAX_MTP * alignK_), buffOffset); + } + + __aicore__ inline void SyncMte2ToV() + { + if (!eventMte2ToVInitialized_) { + eventIdMte2ToV_ = GetTPipePtr()->FetchEventID(HardEvent::MTE2_V); + eventMte2ToVInitialized_ = true; + } + SetFlag(eventIdMte2ToV_); + WaitFlag(eventIdMte2ToV_); + } + + __aicore__ inline void SyncVToMte2() + { + if (!eventVToMte2Initialized_) { + eventIdVToMte2_ = GetTPipePtr()->FetchEventID(HardEvent::V_MTE2); + eventVToMte2Initialized_ = true; + } + SetFlag(eventIdVToMte2_); + WaitFlag(eventIdVToMte2_); + } + + __aicore__ inline void SyncVToS() + { + if (!eventVToSInitialized_) { + eventIdVToS_ = GetTPipePtr()->FetchEventID(HardEvent::V_S); + eventVToSInitialized_ = true; + } + SetFlag(eventIdVToS_); + WaitFlag(eventIdVToS_); + } + + __aicore__ inline void ReleaseEvents() + { + if (eventMte2ToVInitialized_) { + GetTPipePtr()->ReleaseEventID(eventIdMte2ToV_); + eventMte2ToVInitialized_ = false; + } + if (eventVToMte2Initialized_) { + GetTPipePtr()->ReleaseEventID(eventIdVToMte2_); + eventVToMte2Initialized_ = false; + } + if (eventVToSInitialized_) { + GetTPipePtr()->ReleaseEventID(eventIdVToS_); + eventVToSInitialized_ = false; + } + } + + __aicore__ inline void Process() + { + if (!ValidateCuSeqlens()) { + ReleaseEvents(); + return; + } + uint64_t vectorCoreNum = GetBlockNum(); + uint64_t taskNum = B_ * NV_; + for (uint64_t taskIdx = blockIdx; taskIdx < taskNum; taskIdx += vectorCoreNum) { + uint64_t batch_i = taskIdx / NV_; + uint64_t head_i = taskIdx % NV_; + int64_t seq0 = SequenceStart(batch_i); + int64_t seq1 = SequenceEnd(batch_i); + int64_t seqLen64 = seq1 - seq0; + if (seqLen64 == 0) { + continue; + } + int32_t seqLen = static_cast(seqLen64); + if (!ValidateStateSlots(batch_i, seq0, seqLen)) { + ReleaseEvents(); + return; + } + + uint64_t stateSlot = ResolveInitialStateSlot(batch_i, seq0, seqLen); + if (stateSlot == INVALID_STATE_SLOT) { + ReleaseEvents(); + return; + } + CopyInBeta(seq0, seq1); + ProcessHead(batch_i, seq0, seq1, head_i, stateSlot); + } + ReleaseEvents(); + } + +private: + __aicore__ inline bool ValidateCuSeqlens() const + { + if (!hasCuSeqlens_) { + return seqLen_ <= MAX_MTP; + } + int64_t seq0 = LoadCuSeqlens(0); + if (seq0 != 0) { + return false; + } + for (uint64_t i = 0; i < B_; i++) { + int64_t seq1 = LoadCuSeqlens(i + 1); + int64_t length = seq1 - seq0; + if (seq1 < seq0 || seq1 > static_cast(T_) || + length > static_cast(MAX_MTP) || + (hasSsmStateIndices_ && ssmStateStride_ > 0 && length > ssmStateStride_)) { + return false; + } + seq0 = seq1; + } + return seq0 <= static_cast(T_); + } + + __aicore__ inline int64_t LoadCuSeqlens(uint64_t index) const + { + return cuSeqlensDtype_ == 0 ? static_cast(cuSeqlensInt32Gm_.GetValue(index)) : + cuSeqlensInt64Gm_.GetValue(index); + } + + __aicore__ inline int64_t SequenceStart(uint64_t batchIdx) const + { + return hasCuSeqlens_ ? LoadCuSeqlens(batchIdx) : static_cast(batchIdx * seqLen_); + } + + __aicore__ inline int64_t SequenceEnd(uint64_t batchIdx) const + { + return hasCuSeqlens_ ? LoadCuSeqlens(batchIdx + 1) : + static_cast((batchIdx + 1) * seqLen_); + } + + __aicore__ inline int64_t LoadSsmStateIndex(uint64_t index) const + { + return ssmStateIndicesDtype_ == 0 ? static_cast(ssmStateIndicesInt32Gm_.GetValue(index)) : + ssmStateIndicesInt64Gm_.GetValue(index); + } + + __aicore__ inline int64_t LoadAcceptedTokens(uint64_t index) const + { + return acceptedTokensDtype_ == 0 ? static_cast(numAcceptedTokensInt32Gm_.GetValue(index)) : + numAcceptedTokensInt64Gm_.GetValue(index); + } + + __aicore__ inline uint64_t StateMetadataOffset(uint64_t batchIdx, int64_t seq0, int64_t tokenIdx) const + { + if (ssmStateStride_ == 0) { + return static_cast(tokenIdx); + } + return batchIdx * ssmStateStride_ + static_cast(tokenIdx - seq0); + } + + __aicore__ inline uint64_t LoadStateSlot(uint64_t batchIdx, int64_t seq0, int64_t tokenIdx) const + { + int64_t stateSlot = LoadSsmStateIndex(StateMetadataOffset(batchIdx, seq0, tokenIdx)); + if (stateSlot < 0 || stateSlot >= static_cast(stateCapacity_)) { + return INVALID_STATE_SLOT; + } + return static_cast(stateSlot); + } + + __aicore__ inline bool ValidateStateSlots(uint64_t batchIdx, int64_t seq0, int32_t seqLen) const + { + if (!hasSsmStateIndices_) { + return batchIdx < stateCapacity_; + } + if (hasAcceptedTokens_) { + int64_t acceptedTokenNum = LoadAcceptedTokens(batchIdx); + if (acceptedTokenNum <= 0 || acceptedTokenNum > seqLen) { + return false; + } + } + for (int32_t step = 0; step < seqLen; ++step) { + if (LoadStateSlot(batchIdx, seq0, seq0 + step) == INVALID_STATE_SLOT) { + return false; + } + } + return true; + } + + __aicore__ inline uint64_t ResolveInitialStateSlot(uint64_t batchIdx, int64_t seq0, int32_t seqLen) const + { + if (!hasSsmStateIndices_) { + return batchIdx; + } + int64_t tokenIdx = seq0; + if (hasAcceptedTokens_) { + tokenIdx = seq0 + LoadAcceptedTokens(batchIdx) - 1; + } + return LoadStateSlot(batchIdx, seq0, tokenIdx); + } + + template + __aicore__ inline void CopyVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, + uint64_t count) + { + uint64_t rowBytes = count * sizeof(dataType); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst, src[offset], static_cast(count)); + } else { + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + } + + __aicore__ inline float ReadFloat(GlobalTensor &tensor, uint64_t offset) + { + LocalTensor scalar = scalarBuf_.Get(); + DataCopyParams params{1, static_cast(sizeof(float)), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(scalar, tensor[offset], params, padParams); + SyncMte2ToV(); + Adds(scalar, scalar, 0.0f, 1); + PipeBarrier(); + SyncVToS(); + __ubuf__ float *ptr = (__ubuf__ float *)scalar.GetPhyAddr(); + return ptr[0]; + } + + __aicore__ inline float ExpScalar(float x) + { + LocalTensor scalar = scalarBuf_.Get(); + Duplicate(scalar, x, 1); + PipeBarrier(); + Exp(scalar, scalar, 1); + PipeBarrier(); + SyncVToS(); + __ubuf__ float *ptr = (__ubuf__ float *)scalar.GetPhyAddr(); + return ptr[0]; + } + + __aicore__ inline float SigmoidScalar(float x) + { + float denom = 1.0f + ExpScalar(-x); + return 1.0f / denom; + } + + __aicore__ inline void NormalizeRows(LocalTensor &tensor, int32_t seqLen) + { + for (int32_t row = 0; row < seqLen; ++row) { + uint32_t rowOffset = static_cast(row) * alignK_; + Mul(broadTmpInUb, tensor[rowOffset], tensor[rowOffset], alignK_); + PipeBarrier(); + ReduceSumDispatch(deltaInUb, broadTmpInUb, 1); + PipeBarrier(); + Sqrt(deltaInUb, deltaInUb, 1); + PipeBarrier(); + SyncVToS(); + float norm = deltaInUb.GetValue(0); + if (norm > 0.0f) { + Muls(tensor[rowOffset], tensor[rowOffset], 1.0f / norm, alignK_); + PipeBarrier(); + } + } + } + + __aicore__ inline void ApplyGateInKernel(uint64_t head, int32_t seqLen) + { + uint32_t total = static_cast(seqLen) * alignK_; + float expA = hasALog_ ? ExpScalar(ReadFloat(aLogGm_, head)) : 1.0f; + + if (hasDtBias_) { + CopyVectorIn(broadTmpInUb, dtBiasGm_, head * realK_, realK_); + SyncMte2ToV(); + for (int32_t row = 0; row < seqLen; ++row) { + Add(gateInUb[row * alignK_], gateInUb[row * alignK_], broadTmpInUb, alignK_); + PipeBarrier(); + } + } + + if (safeGate_) { + Muls(gateInUb, gateInUb, expA, total); + PipeBarrier(); + Muls(broadTmpInUb, gateInUb, -1.0f, total); + PipeBarrier(); + Exp(broadTmpInUb, broadTmpInUb, total); + PipeBarrier(); + Adds(broadTmpInUb, broadTmpInUb, 1.0f, total); + PipeBarrier(); + Duplicate(gateInUb, 1.0f, total); + PipeBarrier(); + Div(gateInUb, gateInUb, broadTmpInUb, total); + PipeBarrier(); + Muls(gateInUb, gateInUb, lowerBound_, total); + PipeBarrier(); + } else { + Exp(gateInUb, gateInUb, total); + PipeBarrier(); + Adds(gateInUb, gateInUb, 1.0f, total); + PipeBarrier(); + Ln(gateInUb, gateInUb, total); + PipeBarrier(); + Muls(gateInUb, gateInUb, -expA, total); + PipeBarrier(); + } + } + + template + __aicore__ inline void CopyInGate(uint64_t gateOffset, int32_t seqLen) + { + LocalTensor gateLocal = gateInQueue_.AllocTensor(); + Duplicate(gateLocal, static_cast(0), alignK_ * static_cast(seqLen)); + SyncVToMte2(); + DataCopyExtParams gateInParams{static_cast(seqLen), + static_cast(realK_ * sizeof(gateType)), + static_cast((NV_ - 1) * realK_ * sizeof(gateType)), 0, 0}; + DataCopyPadExtParams gatePadParams{ + true, 0, static_cast(alignK_ - realK_), static_cast(0)}; + if constexpr (std::is_same()) { + DataCopyPad(gateLocal, gateFloatGm_[gateOffset], gateInParams, gatePadParams); + } else if constexpr (std::is_same()) { + DataCopyPad(gateLocal, gateBf16Gm_[gateOffset], gateInParams, gatePadParams); + } else { + DataCopyPad(gateLocal, gateFp16Gm_[gateOffset], gateInParams, gatePadParams); + } + gateInQueue_.EnQue(gateLocal); + gateLocal = gateInQueue_.DeQue(); + if constexpr (std::is_same()) { + Adds(gateInUb, gateLocal, 0.0f, alignK_ * static_cast(seqLen)); + } else { + Cast(gateInUb, gateLocal, AscendC::RoundMode::CAST_NONE, + alignK_ * static_cast(seqLen)); + } + gateInQueue_.FreeTensor(gateLocal); + PipeBarrier(); + } + + __aicore__ inline void CopyInQKVGate(uint64_t vOffset, uint64_t qkOffset, uint64_t gateOffset, int32_t seqLen, + uint64_t head) + { + LocalTensor qLocal = qInQueue_.AllocTensor(); + LocalTensor kLocal = kInQueue_.AllocTensor(); + LocalTensor vLocal = vInQueue_.AllocTensor(); + + DataCopyExtParams qkInParams{static_cast(seqLen), static_cast(realK_ * sizeof(inType)), + static_cast((NK_ - 1) * realK_ * sizeof(inType)), 0, 0}; + DataCopyExtParams vInParams{static_cast(seqLen), static_cast(realV_ * sizeof(inType)), + static_cast((NV_ - 1) * realV_ * sizeof(inType)), 0, 0}; + DataCopyPadExtParams qkPadParams{true, 0, static_cast(alignK_ - realK_), 0}; + DataCopyPadExtParams vPadParams{true, 0, static_cast(alignV_ - realV_), 0}; + + DataCopyPad(qLocal, queryGm_[qkOffset], qkInParams, qkPadParams); + DataCopyPad(kLocal, keyGm_[qkOffset], qkInParams, qkPadParams); + DataCopyPad(vLocal, valueGm_[vOffset], vInParams, vPadParams); + qInQueue_.EnQue(qLocal); + kInQueue_.EnQue(kLocal); + vInQueue_.EnQue(vLocal); + if (gateDtype_ == 0) { + CopyInGate(gateOffset, seqLen); + } else if (gateDtype_ == 1) { + CopyInGate(gateOffset, seqLen); + } else { + CopyInGate(gateOffset, seqLen); + } + + qLocal = qInQueue_.DeQue(); + kLocal = kInQueue_.DeQue(); + vLocal = vInQueue_.DeQue(); + Cast(qInUb, qLocal, AscendC::RoundMode::CAST_NONE, alignK_ * seqLen); + Cast(kInUb, kLocal, AscendC::RoundMode::CAST_NONE, alignK_ * seqLen); + Cast(vInUb, vLocal, AscendC::RoundMode::CAST_NONE, alignV_ * seqLen); + AscendC::PipeBarrier(); + if (useQkL2norm_) { + NormalizeRows(qInUb, seqLen); + NormalizeRows(kInUb, seqLen); + } + Muls(qInUb, qInUb, scale_, seqLen * alignK_); + AscendC::PipeBarrier(); + if (useGateInKernel_) { + ApplyGateInKernel(head, seqLen); + } + Exp(gateInUb, gateInUb, alignK_ * seqLen); + AscendC::PipeBarrier(); + + qInQueue_.FreeTensor(qLocal); + kInQueue_.FreeTensor(kLocal); + vInQueue_.FreeTensor(vLocal); + } + + __aicore__ inline void PrefetchState(uint64_t stateSlot, uint64_t head, uint64_t vOffset, + uint32_t curSingleV) + { + LocalTensor stateLocal = stateInQueue_.AllocTensor(); + if (stateVFirst_) { + uint64_t stateOffset = + stateInStride0_ * stateSlot + stateInStride1_ * head + stateInStride2_ * vOffset; + DataCopyExtParams stateInParams{static_cast(curSingleV), + static_cast(realK_ * sizeof(stateType)), 0, 0, 0}; + DataCopyPadExtParams padParams{true, 0, static_cast(alignK_ - realK_), 0}; + DataCopyPad(stateLocal, initStateGm_[stateOffset], stateInParams, padParams); + } else { + for (uint32_t v = 0; v < curSingleV; ++v) { + for (uint32_t k = 0; k < realK_; ++k) { + uint64_t stateOffset = stateInStride0_ * stateSlot + stateInStride1_ * head + + stateInStride2_ * k + + stateInStride3_ * (vOffset + v); + stateLocal.SetValue(v * alignK_ + k, initStateGm_.GetValue(stateOffset)); + } + } + } + stateInQueue_.EnQue(stateLocal); + } + + __aicore__ inline void LoadPrefetchedState(uint32_t curSingleV) + { + LocalTensor stateLocal = stateInQueue_.DeQue(); + if constexpr (std::is_same()) { + DataCopy(stateInUb, stateLocal, alignK_ * curSingleV); + } else { + Cast(stateInUb, stateLocal, AscendC::RoundMode::CAST_NONE, alignK_ * curSingleV); + } + stateInQueue_.FreeTensor(stateLocal); + } + + __aicore__ inline void MatVecMul(const LocalTensor &cubeTensor, const LocalTensor &vecTensor, + LocalTensor &dstTensor, uint32_t rows) + { + __ubuf__ float* cubeAddr = (__ubuf__ float*)cubeTensor.GetPhyAddr(); + __ubuf__ float* vecAddr = (__ubuf__ float*)vecTensor.GetPhyAddr(); + __ubuf__ float* dstAddr = (__ubuf__ float*)dstTensor.GetPhyAddr(); + + uint16_t rowNum = static_cast(rows); + uint16_t colLoopTimes = static_cast(Ceil(alignK_, V_LENGTH)); + uint32_t colLength = alignK_; + __VEC_SCOPE__ + { + RegTensor cube; + RegTensor vec; + RegTensor dst; + MaskReg pregLoop; + for (uint16_t j = 0; j < colLoopTimes; j++) { + pregLoop = UpdateMask(colLength); + DataCopy(vec, vecAddr + j * V_LENGTH); + for (uint16_t i = 0; i < rowNum; i ++) { + DataCopy(cube, cubeAddr + i * alignK_ + j * V_LENGTH); + Mul(dst, cube, vec, pregLoop); + DataCopy(dstAddr + i * alignK_ + j * V_LENGTH, dst, pregLoop); + } + } + } + } + + __aicore__ inline void ProcessKQ(const LocalTensor &cubeTensor, const LocalTensor &vec1Tensor, + LocalTensor &dst1Tensor, const LocalTensor &vec2Tensor, + LocalTensor &dst2Tensor, uint32_t rows) + { + __ubuf__ float* cubeAddr = (__ubuf__ float*)cubeTensor.GetPhyAddr(); + __ubuf__ float* vec1Addr = (__ubuf__ float*)vec1Tensor.GetPhyAddr(); + __ubuf__ float* vec2Addr = (__ubuf__ float*)vec2Tensor.GetPhyAddr(); + __ubuf__ float* dst1Addr = (__ubuf__ float*)dst1Tensor.GetPhyAddr(); + __ubuf__ float* dst2Addr = (__ubuf__ float*)dst2Tensor.GetPhyAddr(); + + uint16_t rowNum = static_cast(rows); + uint16_t colLoopTimes = static_cast(Ceil(alignK_, V_LENGTH)); + uint32_t colLength = alignK_; + __VEC_SCOPE__ + { + RegTensor cube; + RegTensor vec1; + RegTensor vec2; + RegTensor dst1; + RegTensor dst2; + MaskReg pregLoop; + for (uint16_t j = 0; j < colLoopTimes; j++) { + pregLoop = UpdateMask(colLength); + DataCopy(vec1, vec1Addr + j * V_LENGTH); + DataCopy(vec2, vec2Addr + j * V_LENGTH); + for (uint16_t i = 0; i < rowNum; i ++) { + DataCopy(cube, cubeAddr + i); + DataCopy(dst1, dst1Addr + i * alignK_ + j * V_LENGTH); + Mul(cube, cube, vec1, pregLoop); + Add(dst1, dst1, cube, pregLoop); + Mul(dst2, dst1, vec2, pregLoop); + DataCopy(dst1Addr + i * alignK_ + j * V_LENGTH, dst1, pregLoop); + DataCopy(dst2Addr + i * alignK_ + j * V_LENGTH, dst2, pregLoop); + } + } + } + } + + __aicore__ inline void ReduceSum64(__ubuf__ float* dstAddr, __ubuf__ float* srcAddr, uint16_t rowNum) + { + uint32_t colLength = alignK_; + __VEC_SCOPE__ + { + RegTensor src; + RegTensor sum; + MaskReg pregLoop = UpdateMask(colLength); + for (uint16_t i = 0;i < rowNum;i ++) { + DataCopy(src, srcAddr + i * alignK_); + ReduceSum(sum, src, pregLoop); + DataCopy(dstAddr + i, sum, pregLoop); + } + } + } + + __aicore__ inline void ReduceSum128(__ubuf__ float* dstAddr, __ubuf__ float* srcAddr, uint16_t rowNum) + { + uint32_t colLength = alignK_ - V_LENGTH; + __VEC_SCOPE__ + { + RegTensor src1; + RegTensor src2; + RegTensor sum; + MaskReg pregFull = CreateMask(); + MaskReg pregLoop = UpdateMask(colLength); + for (uint16_t i = 0;i < rowNum;i ++) { + DataCopy(src1, srcAddr + i * alignK_); + DataCopy(src2, srcAddr + i * alignK_ + V_LENGTH); + Add(src1, src1, src2, pregLoop); + ReduceSum(sum, src1, pregFull); + DataCopy(dstAddr + i, sum, pregFull); + } + } + } + + __aicore__ inline void ReduceSumVF(__ubuf__ float* dstAddr, __ubuf__ float* srcAddr, uint16_t rowNum) + { + uint16_t colLoopTimes = static_cast(Ceil(alignK_, V_LENGTH)); + __VEC_SCOPE__ + { + RegTensor src; + RegTensor tmp; + RegTensor sum; + MaskReg pregFull = CreateMask(); + MaskReg pregLoop; + for (uint16_t i = 0;i < rowNum;i ++) { + uint32_t colLength = alignK_; + Duplicate(tmp, 0.0f); + for (uint16_t j = 0; j < colLoopTimes; j++) { + pregLoop = UpdateMask(colLength); + DataCopy(src, srcAddr + i * alignK_ + j * V_LENGTH); + Add(tmp, tmp, src, pregLoop); + } + ReduceSum(sum, tmp, pregFull); + DataCopy(dstAddr + i, sum, pregFull); + } + } + } + + __aicore__ inline void ReduceSumDispatch(LocalTensor &dstTensor, LocalTensor &srcTensor, + uint32_t rows) + { + __ubuf__ float* srcAddr = (__ubuf__ float*)srcTensor.GetPhyAddr(); + __ubuf__ float* dstAddr = (__ubuf__ float*)dstTensor.GetPhyAddr(); + uint16_t rowNum = static_cast(rows); + if (alignK_ <= V_LENGTH) { + ReduceSum64(dstAddr, srcAddr, rowNum); + } else if (alignK_ <= TWO_V_LENGTH) { + ReduceSum128(dstAddr, srcAddr, rowNum); + } else { + ReduceSumVF(dstAddr, srcAddr, rowNum); + } + } + + __aicore__ inline void Compute(uint32_t curSingleV, uint64_t curQKOffset, uint64_t curVOffset) + { + MatVecMul(stateInUb, gateInUb[curQKOffset], stateInUb, curSingleV); + AscendC::PipeBarrier(); + MatVecMul(stateInUb, kInUb[curQKOffset], broadTmpInUb, curSingleV); + AscendC::PipeBarrier(); + ReduceSumDispatch(deltaInUb, broadTmpInUb, curSingleV); + AscendC::PipeBarrier(); + Sub(deltaInUb, vInUb[curVOffset], deltaInUb, curSingleV); + AscendC::PipeBarrier(); + Muls(deltaInUb, deltaInUb, beta_, curSingleV); + AscendC::PipeBarrier(); + ProcessKQ(deltaInUb, kInUb[curQKOffset], stateInUb, qInUb[curQKOffset], broadTmpInUb, curSingleV); + AscendC::PipeBarrier(); + ReduceSumDispatch(attnInUb, broadTmpInUb, curSingleV); + LocalTensor attnOutLocal = attnOutQueue_.AllocTensor(); + if (shouldStoreState_) { + LocalTensor stateOutLocal = stateOutQueue_.AllocTensor(); + if constexpr (std::is_same()) { + DataCopy(stateOutLocal, stateInUb, alignK_ * curSingleV); + } else { + Cast(stateOutLocal, stateInUb, AscendC::RoundMode::CAST_RINT, alignK_ * curSingleV); + } + stateOutQueue_.EnQue(stateOutLocal); + } + Cast(attnOutLocal, attnInUb, AscendC::RoundMode::CAST_RINT, curSingleV); + attnOutQueue_.EnQue(attnOutLocal); + } + + __aicore__ inline void CopyOutAttn(uint64_t attnOffset, uint32_t curSingleV) + { + LocalTensor attnLocal = attnOutQueue_.DeQue(); + DataCopyParams attnOutParams{1, static_cast(curSingleV * sizeof(outType)), 0, 0}; + DataCopyPad(attnOutGm_[attnOffset], attnLocal, attnOutParams); + attnOutQueue_.FreeTensor(attnLocal); + } + + __aicore__ inline void CopyOutState(uint64_t stateSlot, uint64_t head, uint64_t vOffset, + uint32_t curSingleV) + { + LocalTensor stateOutLocal = stateOutQueue_.DeQue(); + if (stateVFirst_) { + uint64_t stateOffset = + stateOutStride0_ * stateSlot + stateOutStride1_ * head + stateOutStride2_ * vOffset; + DataCopyParams stateOutParams{static_cast(curSingleV), + static_cast(realK_ * sizeof(stateType)), 0, 0}; + DataCopyPad(finalStateGm_[stateOffset], stateOutLocal, stateOutParams); + } else { + SyncVToS(); + for (uint32_t v = 0; v < curSingleV; ++v) { + for (uint32_t k = 0; k < realK_; ++k) { + uint64_t stateOffset = stateOutStride0_ * stateSlot + stateOutStride1_ * head + + stateOutStride2_ * k + + stateOutStride3_ * (vOffset + v); + finalStateGm_.SetValue(stateOffset, stateOutLocal.GetValue(v * alignK_ + k)); + } + } + } + stateOutQueue_.FreeTensor(stateOutLocal); + } + + template + __aicore__ inline void CopyInBetaTyped(int64_t seq0, int64_t seq1) + { + int64_t seqLen = seq1 - seq0; + uint64_t betaCount = static_cast(seqLen) * NV_; + uint64_t betaBatchSize = Ceil(betaCount, FP32_NUM_PER_BLOCK) * FP32_NUM_PER_BLOCK; + LocalTensor betaLocal = betaInQueue_.AllocTensor(); + if constexpr (std::is_same()) { + CopyVectorIn(betaLocal, betaFloatGm_, static_cast(seq0) * NV_, betaCount); + } else if constexpr (std::is_same()) { + CopyVectorIn(betaLocal, betaBf16Gm_, static_cast(seq0) * NV_, betaCount); + } else { + CopyVectorIn(betaLocal, betaFp16Gm_, static_cast(seq0) * NV_, betaCount); + } + betaInQueue_.EnQue(betaLocal); + betaLocal = betaInQueue_.DeQue(); + if constexpr (std::is_same()) { + Adds(betaInUb, betaLocal, 0.0f, static_cast(betaBatchSize)); + } else { + Cast(betaInUb, betaLocal, AscendC::RoundMode::CAST_NONE, static_cast(betaBatchSize)); + } + betaInQueue_.FreeTensor(betaLocal); + PipeBarrier(); + SyncVToS(); + } + + __aicore__ inline void CopyInBeta(int64_t seq0, int64_t seq1) + { + if (betaDtype_ == 0) { + CopyInBetaTyped(seq0, seq1); + } else if (betaDtype_ == 1) { + CopyInBetaTyped(seq0, seq1); + } else { + CopyInBetaTyped(seq0, seq1); + } + } + + __aicore__ inline uint64_t StateSlotForToken(uint64_t batchIdx, int64_t seq0, int64_t tokenIdx) const + { + if (hasSsmStateIndices_) { + return LoadStateSlot(batchIdx, seq0, tokenIdx); + } + return batchIdx; + } + + __aicore__ inline float LoadBeta(uint64_t gbOffset) + { + float beta = betaInUb.GetValue(gbOffset); + if (useBetaSigmoid_) { + beta = SigmoidScalar(beta); + if (allowNegEigval_) { + beta *= 2.0f; + } + } + return beta; + } + + __aicore__ inline void ProcessHead(uint64_t batchIdx, int64_t seq0, int64_t seq1, + uint64_t head_i, uint64_t stateSlot) + { + uint64_t vOffset = (static_cast(seq0) * NV_ + head_i) * realV_; + uint64_t qkOffset = (static_cast(seq0) * NK_ + head_i / (NV_ / NK_)) * realK_; + uint64_t gateOffset = (static_cast(seq0) * NV_ + head_i) * realK_; + CopyInQKVGate(vOffset, qkOffset, gateOffset, static_cast(seq1 - seq0), head_i); + if (realV_ == 0) { + return; + } + uint64_t nextVOffset = 0; + uint32_t nextSingleV = realV_ > vStep_ ? vStep_ : realV_; + PrefetchState(stateSlot, head_i, 0, nextSingleV); + for (uint64_t v_i = 0; v_i < realV_; v_i += vStep_) { + uint32_t curSingleV = v_i + vStep_ > realV_ ? realV_ - v_i : vStep_; + LoadPrefetchedState(curSingleV); + nextVOffset = v_i + vStep_; + if (nextVOffset < realV_) { + nextSingleV = nextVOffset + vStep_ > realV_ ? realV_ - nextVOffset : vStep_; + PrefetchState(stateSlot, head_i, nextVOffset, nextSingleV); + } + uint64_t pendingAttnOffset = 0; + uint64_t pendingStateSlot = 0; + bool hasPendingAttn = false; + bool hasPendingState = false; + for (int64_t seq_i = seq0; seq_i < seq1; seq_i++) { + uint64_t gbOffset = head_i + static_cast(seq_i - seq0) * NV_; + uint64_t curQKOffset = static_cast(seq_i - seq0) * alignK_; + uint64_t curVOffset = static_cast(seq_i - seq0) * alignV_ + v_i; + uint64_t attnOffset = (static_cast(seq_i) * NV_ + head_i) * realV_ + v_i; + uint64_t curStateSlot = StateSlotForToken(batchIdx, seq0, seq_i); + uint64_t curStateOutSlot = curStateSlot; + beta_ = LoadBeta(gbOffset); + Compute(curSingleV, curQKOffset, curVOffset); + if (attnOutBufferNum_ == BUFFER_NUM) { + CopyOutAttn(attnOffset, curSingleV); + } else { + if (hasPendingAttn) { + CopyOutAttn(pendingAttnOffset, curSingleV); + } + pendingAttnOffset = attnOffset; + hasPendingAttn = true; + } + if (shouldStoreState_) { + if (stateOutBufferNum_ == BUFFER_NUM) { + CopyOutState(curStateOutSlot, head_i, v_i, curSingleV); + } else { + if (hasPendingState) { + CopyOutState(pendingStateSlot, head_i, v_i, curSingleV); + } + pendingStateSlot = curStateOutSlot; + hasPendingState = true; + } + } + } + if (hasPendingAttn) { + CopyOutAttn(pendingAttnOffset, curSingleV); + } + if (hasPendingState) { + CopyOutState(pendingStateSlot, head_i, v_i, curSingleV); + } + } + } + +private: + GlobalTensor queryGm_; + GlobalTensor keyGm_; + GlobalTensor valueGm_; + GlobalTensor gateFloatGm_; + GlobalTensor gateBf16Gm_; + GlobalTensor gateFp16Gm_; + GlobalTensor betaFloatGm_; + GlobalTensor betaBf16Gm_; + GlobalTensor betaFp16Gm_; + GlobalTensor initStateGm_; + GlobalTensor cuSeqlensInt32Gm_; + GlobalTensor cuSeqlensInt64Gm_; + GlobalTensor ssmStateIndicesInt32Gm_; + GlobalTensor ssmStateIndicesInt64Gm_; + GlobalTensor aLogGm_; + GlobalTensor dtBiasGm_; + GlobalTensor numAcceptedTokensInt32Gm_; + GlobalTensor numAcceptedTokensInt64Gm_; + GlobalTensor finalStateGm_; + GlobalTensor attnOutGm_; + TPipe *pipe_; + TQue qInQueue_; + TQue kInQueue_; + TQue vInQueue_; + TQue gateInQueue_; + TQue betaInQueue_; + TQue stateInQueue_; + TQue attnOutQueue_; + TQue stateOutQueue_; + TBuf tmpBuff; + TBuf scalarBuf_; + LocalTensor qInUb; + LocalTensor kInUb; + LocalTensor vInUb; + LocalTensor gateInUb; + LocalTensor betaInUb; + LocalTensor deltaInUb; + LocalTensor broadTmpInUb; + LocalTensor attnInUb; + LocalTensor stateInUb; + TEventID eventIdMte2ToV_; + TEventID eventIdVToMte2_; + TEventID eventIdVToS_; + bool eventMte2ToVInitialized_; + bool eventVToMte2Initialized_; + bool eventVToSInitialized_; + uint32_t B_; + uint32_t T_; + uint32_t seqLen_; + uint32_t NK_; + uint32_t alignK_; + uint32_t realK_; + uint32_t NV_; + uint32_t alignV_; + uint32_t realV_; + uint32_t stateCapacity_; + uint32_t ssmStateStride_; + uint64_t stateInStride0_; + uint64_t stateInStride1_; + uint64_t stateInStride2_; + uint64_t stateInStride3_; + uint64_t stateOutStride0_; + uint64_t stateOutStride1_; + uint64_t stateOutStride2_; + uint64_t stateOutStride3_; + uint32_t vStep_; + uint32_t stateOutBufferNum_; + uint32_t attnOutBufferNum_; + uint32_t restUbSize_; + uint32_t gateDtype_; + uint32_t betaDtype_; + uint32_t cuSeqlensDtype_; + uint32_t ssmStateIndicesDtype_; + uint32_t acceptedTokensDtype_; + bool hasCuSeqlens_; + bool hasSsmStateIndices_; + bool hasAcceptedTokens_; + bool hasALog_; + bool hasDtBias_; + bool useQkL2norm_; + bool useGateInKernel_; + bool useBetaSigmoid_; + bool allowNegEigval_; + bool safeGate_; + bool stateVFirst_; + bool shouldStoreState_; + bool useAddFoldReduce_; + float beta_; + float scale_; + float lowerBound_; + uint64_t blockIdx; +}; +} // namespace RecurrentKda +#endif diff --git a/csrc/attention/recurrent_kda/op_kernel/recurrent_kda.cpp b/csrc/attention/recurrent_kda/op_kernel/recurrent_kda.cpp new file mode 100644 index 000000000000..b577a7f88921 --- /dev/null +++ b/csrc/attention/recurrent_kda/op_kernel/recurrent_kda.cpp @@ -0,0 +1,39 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +/*! + * \file recurrent_kda.cpp + * \brief + */ +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +#include "arch35/recurrent_kda.h" +#else +#include "recurrent_kda.h" +#endif +#include "recurrent_kda_tiling_data.h" + + +using namespace AscendC; +using namespace matmul; +using namespace RecurrentKda; + + +extern "C" __global__ __aicore__ void +recurrent_kda(GM_ADDR query, GM_ADDR key, GM_ADDR value, GM_ADDR gate, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR ssmStateIndices, GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR numAcceptedTokens, + GM_ADDR out, GM_ADDR initialStateOut, GM_ADDR finalState, GM_ADDR workspaceGM, GM_ADDR tilingGM) +{ + REGISTER_TILING_DEFAULT(RecurrentKdaTilingData); + GET_TILING_DATA(tilingData, tilingGM); + KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY); + TPipe pipe; + RKDA op(&tilingData); + GM_ADDR stateOutput = tilingData.inplaceFinalState == 1 ? initialStateOut : finalState; + RKDAInitParams initParams{query, key, value, gate, beta, initialState, cuSeqlens, ssmStateIndices, + aLog, dtBias, numAcceptedTokens, out, stateOutput}; + op.Init(initParams, &pipe); + op.Process(); +} diff --git a/csrc/attention/recurrent_kda/op_kernel/recurrent_kda.h b/csrc/attention/recurrent_kda/op_kernel/recurrent_kda.h new file mode 100644 index 000000000000..e28923fa931a --- /dev/null +++ b/csrc/attention/recurrent_kda/op_kernel/recurrent_kda.h @@ -0,0 +1,968 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +/*! + * \file recurrent_kda.h + * \brief Single-kernel fused recurrent KDA implementation. + */ + +#ifndef __RECURRENT_KDA_KERNEL_H_ +#define __RECURRENT_KDA_KERNEL_H_ + +#include "kernel_operator.h" +#include "lib/matmul_intf.h" +#include "recurrent_kda_tiling_data.h" + +namespace RecurrentKda { + +using namespace matmul; +using namespace AscendC; +constexpr uint64_t BUFFER_NUM = 1; +constexpr uint32_t MAX_OUT_BUFFER_NUM = 2; +constexpr uint64_t MAX_MTP = 8; +constexpr uint64_t BF16_NUM_PER_BLOCK = 16; +constexpr uint64_t FP32_NUM_PER_BLOCK = 8; +constexpr uint32_t REPEAT_LENTH = 64; // 256 bytes for float. +constexpr uint32_t MAX_REPEAT_TIME = 255; +constexpr uint32_t ADD_FOLD_REDUCE_MIN_K = 128; +constexpr uint64_t INVALID_STATE_SLOT = static_cast(-1); + +#ifndef RKDA_ENABLE_ADD_FOLD_REDUCE +#define RKDA_ENABLE_ADD_FOLD_REDUCE 1 +#endif + +struct RKDAInitParams { + GM_ADDR query; + GM_ADDR key; + GM_ADDR value; + GM_ADDR gate; + GM_ADDR beta; + GM_ADDR initState; + GM_ADDR cuSeqlens; + GM_ADDR ssmStateIndices; + GM_ADDR aLog; + GM_ADDR dtBias; + GM_ADDR numAcceptedTokens; + GM_ADDR attnOut; + GM_ADDR finalState; +}; + +template +class RKDA { +public: + __aicore__ inline explicit RKDA(const RecurrentKdaTilingData *tilingData) + { + B_ = tilingData->b; + T_ = tilingData->t; + seqLen_ = tilingData->seqLen; + NK_ = tilingData->nk; + realK_ = tilingData->dk; + NV_ = tilingData->nv; + realV_ = tilingData->dv; + stateCapacity_ = tilingData->sBlockNum; + ssmStateStride_ = tilingData->ssmStateStride; + stateInStride0_ = tilingData->stateInStride0; + stateInStride1_ = tilingData->stateInStride1; + stateInStride2_ = tilingData->stateInStride2; + stateInStride3_ = tilingData->stateInStride3; + stateOutStride0_ = tilingData->stateOutStride0; + stateOutStride1_ = tilingData->stateOutStride1; + stateOutStride2_ = tilingData->stateOutStride2; + stateOutStride3_ = tilingData->stateOutStride3; + scale_ = tilingData->scale; + lowerBound_ = tilingData->lowerBound; + hasCuSeqlens_ = (tilingData->hasCuSeqlens == 1); + hasSsmStateIndices_ = (tilingData->hasSsmStateIndices == 1); + hasAcceptedTokens_ = (tilingData->hasAcceptedTokens == 1); + hasALog_ = (tilingData->hasALog == 1); + hasDtBias_ = (tilingData->hasDtBias == 1); + useQkL2norm_ = (tilingData->useQkL2norm == 1); + useGateInKernel_ = (tilingData->useGateInKernel == 1); + useBetaSigmoid_ = (tilingData->useBetaSigmoid == 1); + allowNegEigval_ = (tilingData->allowNegEigval == 1); + safeGate_ = (tilingData->safeGate == 1); + stateVFirst_ = (tilingData->stateVFirst == 1); + shouldStoreState_ = (tilingData->inplaceFinalState == 1 || tilingData->outputFinalState == 1); + gateDtype_ = tilingData->gateDtype; + betaDtype_ = tilingData->betaDtype; + cuSeqlensDtype_ = tilingData->cuSeqlensDtype; + ssmStateIndicesDtype_ = tilingData->ssmStateIndicesDtype; + acceptedTokensDtype_ = tilingData->acceptedTokensDtype; + useAddFoldReduce_ = (RKDA_ENABLE_ADD_FOLD_REDUCE != 0); + vStep_ = tilingData->vStep; + stateOutBufferNum_ = (tilingData->stateOutBufferNum == MAX_OUT_BUFFER_NUM) ? MAX_OUT_BUFFER_NUM : BUFFER_NUM; + attnOutBufferNum_ = (tilingData->attnOutBufferNum == MAX_OUT_BUFFER_NUM) ? MAX_OUT_BUFFER_NUM : BUFFER_NUM; + restUbSize_ = tilingData->ubRestBytes; + alignK_ = Ceil(tilingData->dk, BF16_NUM_PER_BLOCK) * BF16_NUM_PER_BLOCK; + alignV_ = Ceil(tilingData->dv, BF16_NUM_PER_BLOCK) * BF16_NUM_PER_BLOCK; + eventMte2ToVInitialized_ = false; + eventVToMte2Initialized_ = false; + eventVToSInitialized_ = false; + } + + __aicore__ inline void Init(const RKDAInitParams &initParams, TPipe *pipe) + { + uint64_t blockDim = GetBlockNum(); + blockIdx = GetBlockIdx(); + if (blockIdx >= blockDim) { + return; + } + pipe_ = pipe; + SetGlobalTensors(initParams); + InitLocalBuffers(); + } + + __aicore__ inline void SetGlobalTensors(const RKDAInitParams &initParams) + { + queryGm_.SetGlobalBuffer((__gm__ inType *)initParams.query); + keyGm_.SetGlobalBuffer((__gm__ inType *)initParams.key); + valueGm_.SetGlobalBuffer((__gm__ inType *)initParams.value); + gateFloatGm_.SetGlobalBuffer((__gm__ float *)initParams.gate); + gateBf16Gm_.SetGlobalBuffer((__gm__ bfloat16_t *)initParams.gate); + gateFp16Gm_.SetGlobalBuffer((__gm__ half *)initParams.gate); + betaFloatGm_.SetGlobalBuffer((__gm__ float *)initParams.beta); + betaBf16Gm_.SetGlobalBuffer((__gm__ bfloat16_t *)initParams.beta); + betaFp16Gm_.SetGlobalBuffer((__gm__ half *)initParams.beta); + initStateGm_.SetGlobalBuffer((__gm__ stateType *)initParams.initState); + cuSeqlensInt32Gm_.SetGlobalBuffer((__gm__ int32_t *)initParams.cuSeqlens); + cuSeqlensInt64Gm_.SetGlobalBuffer((__gm__ int64_t *)initParams.cuSeqlens); + ssmStateIndicesInt32Gm_.SetGlobalBuffer((__gm__ int32_t *)initParams.ssmStateIndices); + ssmStateIndicesInt64Gm_.SetGlobalBuffer((__gm__ int64_t *)initParams.ssmStateIndices); + aLogGm_.SetGlobalBuffer((__gm__ float *)initParams.aLog); + dtBiasGm_.SetGlobalBuffer((__gm__ float *)initParams.dtBias); + numAcceptedTokensInt32Gm_.SetGlobalBuffer((__gm__ int32_t *)initParams.numAcceptedTokens); + numAcceptedTokensInt64Gm_.SetGlobalBuffer((__gm__ int64_t *)initParams.numAcceptedTokens); + finalStateGm_.SetGlobalBuffer((__gm__ stateType *)initParams.finalState); + attnOutGm_.SetGlobalBuffer((__gm__ outType *)initParams.attnOut); + } + + __aicore__ inline void InitLocalBuffers() + { + uint32_t cubeSize = alignK_ * vStep_ * sizeof(float); + uint32_t singleVSize = vStep_ * sizeof(float); + uint32_t vSize = MAX_MTP * alignV_ * sizeof(float); + uint32_t kSize = MAX_MTP * alignK_ * sizeof(float); + uint32_t betaUbSize = + Ceil(MAX_MTP * NV_, FP32_NUM_PER_BLOCK) * FP32_NUM_PER_BLOCK * sizeof(float); + pipe_->InitBuffer(qInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(inType)); + pipe_->InitBuffer(kInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(inType)); + pipe_->InitBuffer(vInQueue_, BUFFER_NUM, MAX_MTP * alignV_ * sizeof(inType)); + pipe_->InitBuffer(gateInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(float)); + pipe_->InitBuffer(betaInQueue_, BUFFER_NUM, betaUbSize); + pipe_->InitBuffer(stateInQueue_, BUFFER_NUM, alignK_ * vStep_ * sizeof(stateType)); + pipe_->InitBuffer(stateOutQueue_, stateOutBufferNum_, alignK_ * vStep_ * sizeof(stateType)); + pipe_->InitBuffer(attnOutQueue_, attnOutBufferNum_, vStep_ * sizeof(outType)); + pipe_->InitBuffer(tmpBuff, restUbSize_); + pipe_->InitBuffer(scalarBuf_, 64); + + uint32_t buffOffset = 0; + deltaInUb = tmpBuff.GetWithOffset(static_cast(vStep_), buffOffset); + buffOffset += singleVSize; + attnInUb = tmpBuff.GetWithOffset(static_cast(vStep_), buffOffset); + buffOffset += singleVSize; + vInUb = tmpBuff.GetWithOffset(static_cast(MAX_MTP * alignV_), buffOffset); + buffOffset += vSize; + qInUb = tmpBuff.GetWithOffset(static_cast(MAX_MTP * alignK_), buffOffset); + buffOffset += kSize; + kInUb = tmpBuff.GetWithOffset(static_cast(MAX_MTP * alignK_), buffOffset); + buffOffset += kSize; + stateInUb = tmpBuff.GetWithOffset(static_cast(alignK_ * vStep_), buffOffset); + buffOffset += cubeSize; + broadTmpInUb = tmpBuff.GetWithOffset(static_cast(alignK_ * vStep_), buffOffset); + buffOffset += cubeSize; + betaInUb = tmpBuff.GetWithOffset(static_cast(betaUbSize / sizeof(float)), buffOffset); + buffOffset += betaUbSize; + gateInUb = tmpBuff.GetWithOffset(static_cast(MAX_MTP * alignK_), buffOffset); + } + + __aicore__ inline void SyncMte2ToV() + { + if (!eventMte2ToVInitialized_) { + eventIdMte2ToV_ = GetTPipePtr()->FetchEventID(HardEvent::MTE2_V); + eventMte2ToVInitialized_ = true; + } + SetFlag(eventIdMte2ToV_); + WaitFlag(eventIdMte2ToV_); + } + + __aicore__ inline void SyncVToMte2() + { + if (!eventVToMte2Initialized_) { + eventIdVToMte2_ = GetTPipePtr()->FetchEventID(HardEvent::V_MTE2); + eventVToMte2Initialized_ = true; + } + SetFlag(eventIdVToMte2_); + WaitFlag(eventIdVToMte2_); + } + + __aicore__ inline void SyncVToS() + { + if (!eventVToSInitialized_) { + eventIdVToS_ = GetTPipePtr()->FetchEventID(HardEvent::V_S); + eventVToSInitialized_ = true; + } + SetFlag(eventIdVToS_); + WaitFlag(eventIdVToS_); + } + + __aicore__ inline void ReleaseEvents() + { + if (eventMte2ToVInitialized_) { + GetTPipePtr()->ReleaseEventID(eventIdMte2ToV_); + eventMte2ToVInitialized_ = false; + } + if (eventVToMte2Initialized_) { + GetTPipePtr()->ReleaseEventID(eventIdVToMte2_); + eventVToMte2Initialized_ = false; + } + if (eventVToSInitialized_) { + GetTPipePtr()->ReleaseEventID(eventIdVToS_); + eventVToSInitialized_ = false; + } + } + + __aicore__ inline void Process() + { + if (!ValidateCuSeqlens()) { + ReleaseEvents(); + return; + } + for (uint64_t batch_i = 0; batch_i < B_; batch_i++) { + int64_t seq0 = SequenceStart(batch_i); + int64_t seq1 = SequenceEnd(batch_i); + int64_t seqLen64 = seq1 - seq0; + if (seqLen64 == 0) { + continue; + } + int32_t seqLen = static_cast(seqLen64); + if (!ValidateStateSlots(batch_i, seq0, seqLen)) { + ReleaseEvents(); + return; + } + + uint32_t copyFlag = 0; + uint64_t stateSlot = batch_i; + for (uint64_t head_i = 0; head_i < NV_; head_i++) { + if (!IsCurrentTask(batch_i, head_i)) { + continue; + } + copyFlag++; + if (copyFlag == 1) { + stateSlot = ResolveInitialStateSlot(batch_i, seq0, seqLen); + if (stateSlot == INVALID_STATE_SLOT) { + ReleaseEvents(); + return; + } + CopyInBeta(seq0, seq1); + } + ProcessHead(batch_i, seq0, seq1, head_i, stateSlot); + } + } + ReleaseEvents(); + } + +private: + __aicore__ inline bool ValidateCuSeqlens() const + { + if (!hasCuSeqlens_) { + return seqLen_ <= MAX_MTP; + } + int64_t seq0 = LoadCuSeqlens(0); + if (seq0 != 0) { + return false; + } + for (uint64_t i = 0; i < B_; i++) { + int64_t seq1 = LoadCuSeqlens(i + 1); + int64_t length = seq1 - seq0; + if (seq1 < seq0 || seq1 > static_cast(T_) || + length > static_cast(MAX_MTP) || + (hasSsmStateIndices_ && ssmStateStride_ > 0 && length > ssmStateStride_)) { + return false; + } + seq0 = seq1; + } + return seq0 <= static_cast(T_); + } + + __aicore__ inline int64_t LoadCuSeqlens(uint64_t index) const + { + return cuSeqlensDtype_ == 0 ? static_cast(cuSeqlensInt32Gm_.GetValue(index)) : + cuSeqlensInt64Gm_.GetValue(index); + } + + __aicore__ inline int64_t SequenceStart(uint64_t batchIdx) const + { + return hasCuSeqlens_ ? LoadCuSeqlens(batchIdx) : static_cast(batchIdx * seqLen_); + } + + __aicore__ inline int64_t SequenceEnd(uint64_t batchIdx) const + { + return hasCuSeqlens_ ? LoadCuSeqlens(batchIdx + 1) : + static_cast((batchIdx + 1) * seqLen_); + } + + __aicore__ inline int64_t LoadSsmStateIndex(uint64_t index) const + { + return ssmStateIndicesDtype_ == 0 ? static_cast(ssmStateIndicesInt32Gm_.GetValue(index)) : + ssmStateIndicesInt64Gm_.GetValue(index); + } + + __aicore__ inline int64_t LoadAcceptedTokens(uint64_t index) const + { + return acceptedTokensDtype_ == 0 ? static_cast(numAcceptedTokensInt32Gm_.GetValue(index)) : + numAcceptedTokensInt64Gm_.GetValue(index); + } + + __aicore__ inline uint64_t StateMetadataOffset(uint64_t batchIdx, int64_t seq0, int64_t tokenIdx) const + { + if (ssmStateStride_ == 0) { + return static_cast(tokenIdx); + } + return batchIdx * ssmStateStride_ + static_cast(tokenIdx - seq0); + } + + __aicore__ inline uint64_t LoadStateSlot(uint64_t batchIdx, int64_t seq0, int64_t tokenIdx) const + { + int64_t stateSlot = LoadSsmStateIndex(StateMetadataOffset(batchIdx, seq0, tokenIdx)); + if (stateSlot < 0 || stateSlot >= static_cast(stateCapacity_)) { + return INVALID_STATE_SLOT; + } + return static_cast(stateSlot); + } + + __aicore__ inline bool ValidateStateSlots(uint64_t batchIdx, int64_t seq0, int32_t seqLen) const + { + if (!hasSsmStateIndices_) { + return batchIdx < stateCapacity_; + } + if (hasAcceptedTokens_) { + int64_t acceptedTokenNum = LoadAcceptedTokens(batchIdx); + if (acceptedTokenNum <= 0 || acceptedTokenNum > seqLen) { + return false; + } + } + for (int32_t step = 0; step < seqLen; ++step) { + if (LoadStateSlot(batchIdx, seq0, seq0 + step) == INVALID_STATE_SLOT) { + return false; + } + } + return true; + } + + __aicore__ inline uint64_t ResolveInitialStateSlot(uint64_t batchIdx, int64_t seq0, int32_t seqLen) const + { + if (!hasSsmStateIndices_) { + return batchIdx; + } + int64_t tokenIdx = seq0; + if (hasAcceptedTokens_) { + tokenIdx = seq0 + LoadAcceptedTokens(batchIdx) - 1; + } + return LoadStateSlot(batchIdx, seq0, tokenIdx); + } + + template + __aicore__ inline void CopyVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, + uint64_t count) + { + uint64_t rowBytes = count * sizeof(dataType); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst, src[offset], static_cast(count)); + } else { + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + } + + __aicore__ inline float ReadFloat(GlobalTensor &tensor, uint64_t offset) + { + LocalTensor scalar = scalarBuf_.Get(); + DataCopyParams params{1, static_cast(sizeof(float)), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(scalar, tensor[offset], params, padParams); + SyncMte2ToV(); + Adds(scalar, scalar, 0.0f, 1); + PipeBarrier(); + SyncVToS(); + __ubuf__ float *ptr = (__ubuf__ float *)scalar.GetPhyAddr(); + return ptr[0]; + } + + __aicore__ inline float ExpScalar(float x) + { + LocalTensor scalar = scalarBuf_.Get(); + Duplicate(scalar, x, 1); + PipeBarrier(); + Exp(scalar, scalar, 1); + PipeBarrier(); + SyncVToS(); + __ubuf__ float *ptr = (__ubuf__ float *)scalar.GetPhyAddr(); + return ptr[0]; + } + + __aicore__ inline float SigmoidScalar(float x) + { + float denom = 1.0f + ExpScalar(-x); + return 1.0f / denom; + } + + __aicore__ inline void NormalizeRows(LocalTensor &tensor, int32_t seqLen) + { + for (int32_t row = 0; row < seqLen; ++row) { + uint32_t rowOffset = static_cast(row) * alignK_; + Mul(broadTmpInUb, tensor[rowOffset], tensor[rowOffset], alignK_); + PipeBarrier(); + ReduceSumDispatch(deltaInUb, broadTmpInUb, 1); + PipeBarrier(); + Sqrt(deltaInUb, deltaInUb, 1); + PipeBarrier(); + SyncVToS(); + float norm = deltaInUb.GetValue(0); + if (norm > 0.0f) { + Muls(tensor[rowOffset], tensor[rowOffset], 1.0f / norm, alignK_); + PipeBarrier(); + } + } + } + + __aicore__ inline void ApplyGateInKernel(uint64_t head, int32_t seqLen) + { + uint32_t total = static_cast(seqLen) * alignK_; + float expA = hasALog_ ? ExpScalar(ReadFloat(aLogGm_, head)) : 1.0f; + + if (hasDtBias_) { + CopyVectorIn(broadTmpInUb, dtBiasGm_, head * realK_, realK_); + SyncMte2ToV(); + for (int32_t row = 0; row < seqLen; ++row) { + Add(gateInUb[row * alignK_], gateInUb[row * alignK_], broadTmpInUb, alignK_); + PipeBarrier(); + } + } + + if (safeGate_) { + Muls(gateInUb, gateInUb, expA, total); + PipeBarrier(); + Muls(broadTmpInUb, gateInUb, -1.0f, total); + PipeBarrier(); + Exp(broadTmpInUb, broadTmpInUb, total); + PipeBarrier(); + Adds(broadTmpInUb, broadTmpInUb, 1.0f, total); + PipeBarrier(); + Duplicate(gateInUb, 1.0f, total); + PipeBarrier(); + Div(gateInUb, gateInUb, broadTmpInUb, total); + PipeBarrier(); + Muls(gateInUb, gateInUb, lowerBound_, total); + PipeBarrier(); + } else { + Exp(gateInUb, gateInUb, total); + PipeBarrier(); + Adds(gateInUb, gateInUb, 1.0f, total); + PipeBarrier(); + Ln(gateInUb, gateInUb, total); + PipeBarrier(); + Muls(gateInUb, gateInUb, -expA, total); + PipeBarrier(); + } + } + + template + __aicore__ inline void CopyInGate(uint64_t gateOffset, int32_t seqLen) + { + LocalTensor gateLocal = gateInQueue_.AllocTensor(); + Duplicate(gateLocal, static_cast(0), alignK_ * static_cast(seqLen)); + SyncVToMte2(); + DataCopyExtParams gateInParams{static_cast(seqLen), + static_cast(realK_ * sizeof(gateType)), + static_cast((NV_ - 1) * realK_ * sizeof(gateType)), 0, 0}; + DataCopyPadExtParams gatePadParams{ + true, 0, static_cast(alignK_ - realK_), static_cast(0)}; + if constexpr (std::is_same()) { + DataCopyPad(gateLocal, gateFloatGm_[gateOffset], gateInParams, gatePadParams); + } else if constexpr (std::is_same()) { + DataCopyPad(gateLocal, gateBf16Gm_[gateOffset], gateInParams, gatePadParams); + } else { + DataCopyPad(gateLocal, gateFp16Gm_[gateOffset], gateInParams, gatePadParams); + } + gateInQueue_.EnQue(gateLocal); + gateLocal = gateInQueue_.DeQue(); + if constexpr (std::is_same()) { + Adds(gateInUb, gateLocal, 0.0f, alignK_ * static_cast(seqLen)); + } else { + Cast(gateInUb, gateLocal, AscendC::RoundMode::CAST_NONE, + alignK_ * static_cast(seqLen)); + } + gateInQueue_.FreeTensor(gateLocal); + PipeBarrier(); + } + + __aicore__ inline void CopyInQKVGate(uint64_t vOffset, uint64_t qkOffset, uint64_t gateOffset, int32_t seqLen, + uint64_t head) + { + LocalTensor qLocal = qInQueue_.AllocTensor(); + LocalTensor kLocal = kInQueue_.AllocTensor(); + LocalTensor vLocal = vInQueue_.AllocTensor(); + + DataCopyExtParams qkInParams{static_cast(seqLen), static_cast(realK_ * sizeof(inType)), + static_cast((NK_ - 1) * realK_ * sizeof(inType)), 0, 0}; + DataCopyExtParams vInParams{static_cast(seqLen), static_cast(realV_ * sizeof(inType)), + static_cast((NV_ - 1) * realV_ * sizeof(inType)), 0, 0}; + DataCopyPadExtParams qkPadParams{true, 0, static_cast(alignK_ - realK_), 0}; + DataCopyPadExtParams vPadParams{true, 0, static_cast(alignV_ - realV_), 0}; + + DataCopyPad(qLocal, queryGm_[qkOffset], qkInParams, qkPadParams); + DataCopyPad(kLocal, keyGm_[qkOffset], qkInParams, qkPadParams); + DataCopyPad(vLocal, valueGm_[vOffset], vInParams, vPadParams); + qInQueue_.EnQue(qLocal); + kInQueue_.EnQue(kLocal); + vInQueue_.EnQue(vLocal); + if (gateDtype_ == 0) { + CopyInGate(gateOffset, seqLen); + } else if (gateDtype_ == 1) { + CopyInGate(gateOffset, seqLen); + } else { + CopyInGate(gateOffset, seqLen); + } + + qLocal = qInQueue_.DeQue(); + kLocal = kInQueue_.DeQue(); + vLocal = vInQueue_.DeQue(); + Cast(qInUb, qLocal, AscendC::RoundMode::CAST_NONE, alignK_ * seqLen); + Cast(kInUb, kLocal, AscendC::RoundMode::CAST_NONE, alignK_ * seqLen); + Cast(vInUb, vLocal, AscendC::RoundMode::CAST_NONE, alignV_ * seqLen); + AscendC::PipeBarrier(); + if (useQkL2norm_) { + NormalizeRows(qInUb, seqLen); + NormalizeRows(kInUb, seqLen); + } + Muls(qInUb, qInUb, scale_, seqLen * alignK_); + AscendC::PipeBarrier(); + if (useGateInKernel_) { + ApplyGateInKernel(head, seqLen); + } + Exp(gateInUb, gateInUb, alignK_ * seqLen); + AscendC::PipeBarrier(); + + qInQueue_.FreeTensor(qLocal); + kInQueue_.FreeTensor(kLocal); + vInQueue_.FreeTensor(vLocal); + } + + __aicore__ inline void PrefetchState(uint64_t stateSlot, uint64_t head, uint64_t vOffset, + uint32_t curSingleV) + { + LocalTensor stateLocal = stateInQueue_.AllocTensor(); + if (stateVFirst_) { + uint64_t stateOffset = + stateInStride0_ * stateSlot + stateInStride1_ * head + stateInStride2_ * vOffset; + DataCopyExtParams stateInParams{static_cast(curSingleV), + static_cast(realK_ * sizeof(stateType)), 0, 0, 0}; + DataCopyPadExtParams padParams{true, 0, static_cast(alignK_ - realK_), 0}; + DataCopyPad(stateLocal, initStateGm_[stateOffset], stateInParams, padParams); + } else { + for (uint32_t v = 0; v < curSingleV; ++v) { + for (uint32_t k = 0; k < realK_; ++k) { + uint64_t stateOffset = stateInStride0_ * stateSlot + stateInStride1_ * head + + stateInStride2_ * k + + stateInStride3_ * (vOffset + v); + stateLocal.SetValue(v * alignK_ + k, initStateGm_.GetValue(stateOffset)); + } + } + } + stateInQueue_.EnQue(stateLocal); + } + + __aicore__ inline void LoadPrefetchedState(uint32_t curSingleV) + { + LocalTensor stateLocal = stateInQueue_.DeQue(); + if constexpr (std::is_same()) { + DataCopy(stateInUb, stateLocal, alignK_ * curSingleV); + } else { + Cast(stateInUb, stateLocal, AscendC::RoundMode::CAST_NONE, alignK_ * curSingleV); + } + stateInQueue_.FreeTensor(stateLocal); + } + + __aicore__ inline void MatVecMul(const LocalTensor &cubeTensor, const LocalTensor &vecTensor, + LocalTensor &dstTensor, uint32_t cols, bool isAdd) + { + uint8_t repeatStride = alignK_ / FP32_NUM_PER_BLOCK; + for (uint32_t i = 0; i < alignK_; i += REPEAT_LENTH) { + uint64_t mask = Std::min(REPEAT_LENTH, alignK_ - i); + for (uint32_t j = 0; j < cols; j += MAX_REPEAT_TIME) { + uint64_t repeatTime = Std::min(MAX_REPEAT_TIME, cols - j); + if (isAdd) { + MulAddDst(dstTensor[j * alignK_ + i], cubeTensor[j * alignK_ + i], vecTensor[i], mask, repeatTime, + {1, 1, 1, repeatStride, repeatStride, 0}); + } else { + Mul(dstTensor[j * alignK_ + i], cubeTensor[j * alignK_ + i], vecTensor[i], mask, repeatTime, + {1, 1, 1, repeatStride, repeatStride, 0}); + } + } + } + } + + __aicore__ inline void ReduceSumBaseline(LocalTensor &dstTensor, const LocalTensor &srcTensor, + uint32_t rows) + { + uint32_t stateShape[2] = {rows, alignK_}; + ReduceSum(dstTensor, srcTensor, stateShape, true); + } + + __aicore__ inline bool CanUseK128AddFoldFastPath(uint32_t rows) const + { + if (alignK_ != ADD_FOLD_REDUCE_MIN_K) { + return false; + } + if (rows == 0 || rows > MAX_REPEAT_TIME) { + return false; + } + return true; + } + + __aicore__ inline void ReduceSumAddFoldK128(LocalTensor &dstTensor, LocalTensor &srcTensor, + uint32_t rows) + { + const uint8_t repeatTime = static_cast(rows); + const uint8_t rowRepStride = static_cast(alignK_ / FP32_NUM_PER_BLOCK); + + Add(srcTensor[REPEAT_LENTH], srcTensor, srcTensor[REPEAT_LENTH], REPEAT_LENTH, repeatTime, + {1, 1, 1, rowRepStride, rowRepStride, rowRepStride}); + AscendC::PipeBarrier(); + WholeReduceSum(dstTensor, srcTensor[REPEAT_LENTH], REPEAT_LENTH, repeatTime, 1, 1, rowRepStride); + } + + __aicore__ inline void ReduceSumAddFold(LocalTensor &dstTensor, LocalTensor &srcTensor, + uint32_t rows) + { + if (alignK_ < REPEAT_LENTH) { + ReduceSumBaseline(dstTensor, srcTensor, rows); + return; + } + + if ((alignK_ & (alignK_ - 1)) != 0) { + ReduceSumBaseline(dstTensor, srcTensor, rows); + return; + } + + if (CanUseK128AddFoldFastPath(rows)) { + ReduceSumAddFoldK128(dstTensor, srcTensor, rows); + return; + } + + for (uint32_t row = 0; row < rows; ++row) { + uint32_t rowOffset = row * alignK_; + uint32_t activeLen = alignK_; + while (activeLen > REPEAT_LENTH) { + uint32_t half = activeLen >> 1; + Add(srcTensor[rowOffset], srcTensor[rowOffset], srcTensor[rowOffset + half], half); + AscendC::PipeBarrier(); + activeLen = half; + } + + WholeReduceSum(dstTensor[row], srcTensor[rowOffset], REPEAT_LENTH, 1, 1, 1, FP32_NUM_PER_BLOCK); + } + } + + __aicore__ inline void ReduceSumDispatch(LocalTensor &dstTensor, LocalTensor &srcTensor, + uint32_t rows) + { + if (useAddFoldReduce_ && alignK_ >= ADD_FOLD_REDUCE_MIN_K) { + ReduceSumAddFold(dstTensor, srcTensor, rows); + return; + } + ReduceSumBaseline(dstTensor, srcTensor, rows); + } + + __aicore__ inline void Compute(uint32_t curSingleV, uint64_t curQKOffset, uint64_t curVOffset) + { + uint32_t stateShape[2] = {curSingleV, alignK_}; + uint32_t deltaShape[2] = {curSingleV, 1}; + MatVecMul(stateInUb, gateInUb[curQKOffset], stateInUb, curSingleV, false); + AscendC::PipeBarrier(); + MatVecMul(stateInUb, kInUb[curQKOffset], broadTmpInUb, curSingleV, false); + AscendC::PipeBarrier(); + ReduceSumDispatch(deltaInUb, broadTmpInUb, curSingleV); + AscendC::PipeBarrier(); + deltaInUb = vInUb[curVOffset] - deltaInUb; + AscendC::PipeBarrier(); + Muls(deltaInUb, deltaInUb, beta_, curSingleV); + AscendC::PipeBarrier(); + Broadcast(broadTmpInUb, deltaInUb, stateShape, deltaShape); + AscendC::PipeBarrier(); + MatVecMul(broadTmpInUb, kInUb[curQKOffset], stateInUb, curSingleV, true); + AscendC::PipeBarrier(); + MatVecMul(stateInUb, qInUb[curQKOffset], broadTmpInUb, curSingleV, false); + AscendC::PipeBarrier(); + ReduceSumDispatch(attnInUb, broadTmpInUb, curSingleV); + LocalTensor attnOutLocal = attnOutQueue_.AllocTensor(); + if (shouldStoreState_) { + LocalTensor stateOutLocal = stateOutQueue_.AllocTensor(); + if constexpr (std::is_same()) { + DataCopy(stateOutLocal, stateInUb, alignK_ * curSingleV); + } else { + Cast(stateOutLocal, stateInUb, AscendC::RoundMode::CAST_RINT, alignK_ * curSingleV); + } + stateOutQueue_.EnQue(stateOutLocal); + } + Cast(attnOutLocal, attnInUb, AscendC::RoundMode::CAST_RINT, curSingleV); + attnOutQueue_.EnQue(attnOutLocal); + } + + __aicore__ inline void CopyOutAttn(uint64_t attnOffset, uint32_t curSingleV) + { + LocalTensor attnLocal = attnOutQueue_.DeQue(); + DataCopyParams attnOutParams{1, static_cast(curSingleV * sizeof(outType)), 0, 0}; + DataCopyPad(attnOutGm_[attnOffset], attnLocal, attnOutParams); + attnOutQueue_.FreeTensor(attnLocal); + } + + __aicore__ inline void CopyOutState(uint64_t stateSlot, uint64_t head, uint64_t vOffset, + uint32_t curSingleV) + { + LocalTensor stateOutLocal = stateOutQueue_.DeQue(); + if (stateVFirst_) { + uint64_t stateOffset = + stateOutStride0_ * stateSlot + stateOutStride1_ * head + stateOutStride2_ * vOffset; + DataCopyParams stateOutParams{static_cast(curSingleV), + static_cast(realK_ * sizeof(stateType)), 0, 0}; + DataCopyPad(finalStateGm_[stateOffset], stateOutLocal, stateOutParams); + } else { + SyncVToS(); + for (uint32_t v = 0; v < curSingleV; ++v) { + for (uint32_t k = 0; k < realK_; ++k) { + uint64_t stateOffset = stateOutStride0_ * stateSlot + stateOutStride1_ * head + + stateOutStride2_ * k + + stateOutStride3_ * (vOffset + v); + finalStateGm_.SetValue(stateOffset, stateOutLocal.GetValue(v * alignK_ + k)); + } + } + } + stateOutQueue_.FreeTensor(stateOutLocal); + } + + template + __aicore__ inline void CopyInBetaTyped(int64_t seq0, int64_t seq1) + { + int64_t seqLen = seq1 - seq0; + uint64_t betaCount = static_cast(seqLen) * NV_; + uint64_t betaBatchSize = Ceil(betaCount, FP32_NUM_PER_BLOCK) * FP32_NUM_PER_BLOCK; + LocalTensor betaLocal = betaInQueue_.AllocTensor(); + if constexpr (std::is_same()) { + CopyVectorIn(betaLocal, betaFloatGm_, static_cast(seq0) * NV_, betaCount); + } else if constexpr (std::is_same()) { + CopyVectorIn(betaLocal, betaBf16Gm_, static_cast(seq0) * NV_, betaCount); + } else { + CopyVectorIn(betaLocal, betaFp16Gm_, static_cast(seq0) * NV_, betaCount); + } + betaInQueue_.EnQue(betaLocal); + betaLocal = betaInQueue_.DeQue(); + if constexpr (std::is_same()) { + Adds(betaInUb, betaLocal, 0.0f, static_cast(betaBatchSize)); + } else { + Cast(betaInUb, betaLocal, AscendC::RoundMode::CAST_NONE, static_cast(betaBatchSize)); + } + betaInQueue_.FreeTensor(betaLocal); + PipeBarrier(); + SyncVToS(); + } + + __aicore__ inline void CopyInBeta(int64_t seq0, int64_t seq1) + { + if (betaDtype_ == 0) { + CopyInBetaTyped(seq0, seq1); + } else if (betaDtype_ == 1) { + CopyInBetaTyped(seq0, seq1); + } else { + CopyInBetaTyped(seq0, seq1); + } + } + + __aicore__ inline uint64_t StateSlotForToken(uint64_t batchIdx, int64_t seq0, int64_t tokenIdx) const + { + if (hasSsmStateIndices_) { + return LoadStateSlot(batchIdx, seq0, tokenIdx); + } + return batchIdx; + } + + __aicore__ inline float LoadBeta(uint64_t gbOffset) + { + float beta = betaInUb.GetValue(gbOffset); + if (useBetaSigmoid_) { + beta = SigmoidScalar(beta); + if (allowNegEigval_) { + beta *= 2.0f; + } + } + return beta; + } + + __aicore__ inline void ProcessHead(uint64_t batchIdx, int64_t seq0, int64_t seq1, + uint64_t head_i, uint64_t stateSlot) + { + uint64_t vOffset = (static_cast(seq0) * NV_ + head_i) * realV_; + uint64_t qkOffset = (static_cast(seq0) * NK_ + head_i / (NV_ / NK_)) * realK_; + uint64_t gateOffset = (static_cast(seq0) * NV_ + head_i) * realK_; + CopyInQKVGate(vOffset, qkOffset, gateOffset, static_cast(seq1 - seq0), head_i); + if (realV_ == 0) { + return; + } + uint64_t nextVOffset = 0; + uint32_t nextSingleV = realV_ > vStep_ ? vStep_ : realV_; + PrefetchState(stateSlot, head_i, 0, nextSingleV); + for (uint64_t v_i = 0; v_i < realV_; v_i += vStep_) { + uint32_t curSingleV = v_i + vStep_ > realV_ ? realV_ - v_i : vStep_; + LoadPrefetchedState(curSingleV); + nextVOffset = v_i + vStep_; + if (nextVOffset < realV_) { + nextSingleV = nextVOffset + vStep_ > realV_ ? realV_ - nextVOffset : vStep_; + PrefetchState(stateSlot, head_i, nextVOffset, nextSingleV); + } + uint64_t pendingAttnOffset = 0; + uint64_t pendingStateSlot = 0; + bool hasPendingAttn = false; + bool hasPendingState = false; + for (int64_t seq_i = seq0; seq_i < seq1; seq_i++) { + uint64_t gbOffset = head_i + static_cast(seq_i - seq0) * NV_; + uint64_t curQKOffset = static_cast(seq_i - seq0) * alignK_; + uint64_t curVOffset = static_cast(seq_i - seq0) * alignV_ + v_i; + uint64_t attnOffset = (static_cast(seq_i) * NV_ + head_i) * realV_ + v_i; + uint64_t curStateSlot = StateSlotForToken(batchIdx, seq0, seq_i); + uint64_t curStateOutSlot = curStateSlot; + beta_ = LoadBeta(gbOffset); + Compute(curSingleV, curQKOffset, curVOffset); + if (attnOutBufferNum_ == BUFFER_NUM) { + CopyOutAttn(attnOffset, curSingleV); + } else { + if (hasPendingAttn) { + CopyOutAttn(pendingAttnOffset, curSingleV); + } + pendingAttnOffset = attnOffset; + hasPendingAttn = true; + } + if (shouldStoreState_) { + if (stateOutBufferNum_ == BUFFER_NUM) { + CopyOutState(curStateOutSlot, head_i, v_i, curSingleV); + } else { + if (hasPendingState) { + CopyOutState(pendingStateSlot, head_i, v_i, curSingleV); + } + pendingStateSlot = curStateOutSlot; + hasPendingState = true; + } + } + } + if (hasPendingAttn) { + CopyOutAttn(pendingAttnOffset, curSingleV); + } + if (hasPendingState) { + CopyOutState(pendingStateSlot, head_i, v_i, curSingleV); + } + } + } + + __aicore__ inline bool IsCurrentTask(uint64_t batchIdx, uint64_t headIdx) const + { + return ((batchIdx * NV_ + headIdx) % GetBlockNum()) == blockIdx; + } + +private: + GlobalTensor queryGm_; + GlobalTensor keyGm_; + GlobalTensor valueGm_; + GlobalTensor gateFloatGm_; + GlobalTensor gateBf16Gm_; + GlobalTensor gateFp16Gm_; + GlobalTensor betaFloatGm_; + GlobalTensor betaBf16Gm_; + GlobalTensor betaFp16Gm_; + GlobalTensor initStateGm_; + GlobalTensor cuSeqlensInt32Gm_; + GlobalTensor cuSeqlensInt64Gm_; + GlobalTensor ssmStateIndicesInt32Gm_; + GlobalTensor ssmStateIndicesInt64Gm_; + GlobalTensor aLogGm_; + GlobalTensor dtBiasGm_; + GlobalTensor numAcceptedTokensInt32Gm_; + GlobalTensor numAcceptedTokensInt64Gm_; + GlobalTensor finalStateGm_; + GlobalTensor attnOutGm_; + TPipe *pipe_; + TQue qInQueue_; + TQue kInQueue_; + TQue vInQueue_; + TQue gateInQueue_; + TQue betaInQueue_; + TQue stateInQueue_; + TQue attnOutQueue_; + TQue stateOutQueue_; + TBuf tmpBuff; + TBuf scalarBuf_; + LocalTensor qInUb; + LocalTensor kInUb; + LocalTensor vInUb; + LocalTensor gateInUb; + LocalTensor betaInUb; + LocalTensor deltaInUb; + LocalTensor broadTmpInUb; + LocalTensor attnInUb; + LocalTensor stateInUb; + TEventID eventIdMte2ToV_; + TEventID eventIdVToMte2_; + TEventID eventIdVToS_; + bool eventMte2ToVInitialized_; + bool eventVToMte2Initialized_; + bool eventVToSInitialized_; + uint32_t B_; + uint32_t T_; + uint32_t seqLen_; + uint32_t NK_; + uint32_t alignK_; + uint32_t realK_; + uint32_t NV_; + uint32_t alignV_; + uint32_t realV_; + uint32_t stateCapacity_; + uint32_t ssmStateStride_; + uint64_t stateInStride0_; + uint64_t stateInStride1_; + uint64_t stateInStride2_; + uint64_t stateInStride3_; + uint64_t stateOutStride0_; + uint64_t stateOutStride1_; + uint64_t stateOutStride2_; + uint64_t stateOutStride3_; + uint32_t vStep_; + uint32_t stateOutBufferNum_; + uint32_t attnOutBufferNum_; + uint32_t restUbSize_; + uint32_t gateDtype_; + uint32_t betaDtype_; + uint32_t cuSeqlensDtype_; + uint32_t ssmStateIndicesDtype_; + uint32_t acceptedTokensDtype_; + bool hasCuSeqlens_; + bool hasSsmStateIndices_; + bool hasAcceptedTokens_; + bool hasALog_; + bool hasDtBias_; + bool useQkL2norm_; + bool useGateInKernel_; + bool useBetaSigmoid_; + bool allowNegEigval_; + bool safeGate_; + bool stateVFirst_; + bool shouldStoreState_; + bool useAddFoldReduce_; + float beta_; + float scale_; + float lowerBound_; + uint64_t blockIdx; +}; +} // namespace RecurrentKda +#endif diff --git a/csrc/attention/recurrent_kda/op_kernel/recurrent_kda_struct.h b/csrc/attention/recurrent_kda/op_kernel/recurrent_kda_struct.h new file mode 100644 index 000000000000..a7d714ca5f11 --- /dev/null +++ b/csrc/attention/recurrent_kda/op_kernel/recurrent_kda_struct.h @@ -0,0 +1,70 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +/*! + * \file recurrent_kda_struct.h + * \brief Plain tiling struct shared by aclnn tiling and fast kernel launch. + */ + +#ifndef RECURRENT_KDA_STRUCT_H +#define RECURRENT_KDA_STRUCT_H + +#include + +namespace RecurrentKda { + +#pragma pack(push, 8) +struct alignas(8) RecurrentKdaTilingData { + uint32_t vectorCoreNum; + uint32_t ubCalSize; + uint32_t ubRestBytes; + uint32_t t; + uint32_t seqLen; + uint32_t nk; + uint32_t dk; + uint32_t nv; + uint32_t dv; + uint32_t sBlockNum; + uint32_t ssmStateStride; + uint32_t b; + uint32_t vStep; + uint32_t stateOutBufferNum; + uint32_t attnOutBufferNum; + float scale; + float lowerBound; + uint32_t layout; + uint32_t hasSsmStateIndices; + uint32_t hasALog; + uint32_t hasDtBias; + uint32_t hasAcceptedTokens; + uint32_t useQkL2norm; + uint32_t useGateInKernel; + uint32_t useBetaSigmoid; + uint32_t allowNegEigval; + uint32_t safeGate; + uint32_t stateVFirst; + uint32_t outputFinalState; + uint32_t inplaceFinalState; + uint32_t hasCuSeqlens; + uint32_t gateDtype; + uint32_t betaDtype; + uint32_t cuSeqlensDtype; + uint32_t ssmStateIndicesDtype; + uint32_t acceptedTokensDtype; + uint64_t stateInStride0; + uint64_t stateInStride1; + uint64_t stateInStride2; + uint64_t stateInStride3; + uint64_t stateOutStride0; + uint64_t stateOutStride1; + uint64_t stateOutStride2; + uint64_t stateOutStride3; +}; +#pragma pack(pop) + +} // namespace RecurrentKda + +#endif // RECURRENT_KDA_STRUCT_H diff --git a/csrc/attention/recurrent_kda/op_kernel/recurrent_kda_tiling_data.h b/csrc/attention/recurrent_kda/op_kernel/recurrent_kda_tiling_data.h new file mode 100644 index 000000000000..86a32b27c790 --- /dev/null +++ b/csrc/attention/recurrent_kda/op_kernel/recurrent_kda_tiling_data.h @@ -0,0 +1,20 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ + +/*! + * \file recurrent_kda_tiling_data.h + * \brief + */ +#ifndef RECURRENT_KDA_TILING_DATA_H +#define RECURRENT_KDA_TILING_DATA_H + +#include "recurrent_kda_struct.h" + +#ifndef TORCH_MODE +#include "kernel_tiling/kernel_tiling.h" +#endif + +#endif // RECURRENT_KDA_TILING_DATA_H diff --git a/csrc/attention/recurrent_kda/recurrent_kda_torch_adpt.h b/csrc/attention/recurrent_kda/recurrent_kda_torch_adpt.h new file mode 100644 index 000000000000..6560de89a1a7 --- /dev/null +++ b/csrc/attention/recurrent_kda/recurrent_kda_torch_adpt.h @@ -0,0 +1,157 @@ +/* + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ +#ifndef RECURRENT_KDA_TORCH_ADPT_H +#define RECURRENT_KDA_TORCH_ADPT_H + +namespace vllm_ascend { + +at::Tensor recurrent_kda( + const at::Tensor& query, + const at::Tensor& key, + const at::Tensor& value, + const at::Tensor& gate, + const at::Tensor& beta, + at::Tensor& initial_state, + const at::Tensor& cu_seqlens, + const at::Tensor& ssm_state_indices, + const at::Tensor& a_log, + const at::Tensor& dt_bias, + const c10::optional& num_accepted_tokens, + double scale, + bool use_qk_l2norm_in_kernel, + bool use_gate_in_kernel, + bool use_beta_sigmoid_in_kernel, + bool allow_neg_eigval, + bool safe_gate, + double lower_bound) +{ + const bool is_tnd = query.dim() == 3; + TORCH_CHECK((is_tnd && key.dim() == 3 && value.dim() == 3 && gate.dim() == 3 && beta.dim() == 2) || + (!is_tnd && query.dim() == 4 && key.dim() == 4 && value.dim() == 4 && + gate.dim() == 4 && beta.dim() == 3), + "recurrent_kda: TND expects q/k [T,H,K], v [T,HV,V], gate [T,HV,K], beta [T,HV]; " + "BSND expects q/k [B,T,H,K], v [B,T,HV,V], gate [B,T,HV,K], beta [B,T,HV]."); + TORCH_CHECK(query.sizes() == key.sizes(), + "recurrent_kda: query and key must have identical shapes."); + TORCH_CHECK(query.scalar_type() == at::kBFloat16 && + key.scalar_type() == at::kBFloat16 && + value.scalar_type() == at::kBFloat16, + "recurrent_kda: query/key/value must be bfloat16."); + TORCH_CHECK((gate.scalar_type() == at::kFloat || gate.scalar_type() == at::kBFloat16 || + gate.scalar_type() == at::kHalf) && + (beta.scalar_type() == at::kFloat || beta.scalar_type() == at::kBFloat16 || + beta.scalar_type() == at::kHalf), + "recurrent_kda: gate and beta must be float32, bfloat16 or float16."); + TORCH_CHECK(key.device() == query.device() && value.device() == query.device() && + gate.device() == query.device() && beta.device() == query.device() && + initial_state.device() == query.device(), + "recurrent_kda: query/key/value/gate/beta/state must be on the same device."); + TORCH_CHECK(cu_seqlens.dim() == 1 && cu_seqlens.numel() >= 2, + "recurrent_kda: cu_seqlens must be a 1D device tensor with at least two elements."); + TORCH_CHECK(cu_seqlens.scalar_type() == at::kInt || cu_seqlens.scalar_type() == at::kLong, + "recurrent_kda: cu_seqlens must be int32 or int64."); + TORCH_CHECK(cu_seqlens.device() == query.device(), + "recurrent_kda: cu_seqlens must be on the same device as query."); + + const int64_t batch = is_tnd ? 1 : query.size(0); + const int64_t total_tokens = is_tnd ? query.size(0) : query.size(0) * query.size(1); + const int64_t seq_num = cu_seqlens.size(0) - 1; + const int64_t h = is_tnd ? query.size(1) : query.size(2); + const int64_t k_dim = is_tnd ? query.size(2) : query.size(3); + const int64_t hv = is_tnd ? value.size(1) : value.size(2); + const int64_t v_dim = is_tnd ? value.size(2) : value.size(3); + TORCH_CHECK(total_tokens > 0 && h > 0 && hv > 0, + "recurrent_kda: token and head dimensions must be positive."); + TORCH_CHECK(hv % h == 0, + "recurrent_kda: HV must be divisible by H."); + TORCH_CHECK(k_dim == 128 && (v_dim == 128 || v_dim == 256), + "recurrent_kda: the Kimi K3 integration requires K=128 and V=128 or 256."); + TORCH_CHECK((is_tnd && value.size(0) == total_tokens && gate.size(0) == total_tokens && + beta.size(0) == total_tokens && gate.size(1) == hv && gate.size(2) == k_dim && + beta.size(1) == hv) || + (!is_tnd && value.size(0) == batch && value.size(1) == query.size(1) && + gate.size(0) == batch && gate.size(1) == query.size(1) && gate.size(2) == hv && + gate.size(3) == k_dim && beta.size(0) == batch && beta.size(1) == query.size(1) && + beta.size(2) == hv), + "recurrent_kda: value/gate/beta shapes do not match the selected layout."); + const bool packed_indices = ssm_state_indices.dim() == 1 && + ssm_state_indices.numel() >= total_tokens; + const bool speculative_indices = ssm_state_indices.dim() == 2 && + ssm_state_indices.size(0) == seq_num && + ssm_state_indices.size(1) > 0; + TORCH_CHECK((ssm_state_indices.scalar_type() == at::kInt || + ssm_state_indices.scalar_type() == at::kLong) && + (packed_indices || speculative_indices), + "recurrent_kda: ssm_state_indices must be int32/int64 packed [T] or " + "speculative [seq_num,max_step]."); + TORCH_CHECK(ssm_state_indices.device() == query.device(), + "recurrent_kda: ssm_state_indices must be on the same device as query."); + TORCH_CHECK(initial_state.dim() == 4 && initial_state.size(0) >= 1 && + initial_state.size(1) == hv && initial_state.size(2) == v_dim && + initial_state.size(3) == k_dim, + "recurrent_kda: initial_state must be a non-empty [state_capacity,HV,V,K] pool."); + TORCH_CHECK(initial_state.scalar_type() == at::kFloat || initial_state.scalar_type() == at::kBFloat16, + "recurrent_kda: initial_state must be float32 or bfloat16."); + TORCH_CHECK(a_log.scalar_type() == at::kFloat && a_log.dim() == 1 && a_log.numel() == hv, + "recurrent_kda: A_log must be float32 [HV]."); + TORCH_CHECK(dt_bias.scalar_type() == at::kFloat && + ((dt_bias.dim() == 1 && dt_bias.numel() == hv * k_dim) || + (dt_bias.dim() == 2 && dt_bias.size(0) == hv && dt_bias.size(1) == k_dim)), + "recurrent_kda: dt_bias must be float32 [HV*K] or [HV,K]."); + TORCH_CHECK(a_log.device() == query.device() && dt_bias.device() == query.device(), + "recurrent_kda: A_log and dt_bias must be on the same device as query."); + if (num_accepted_tokens.has_value() && num_accepted_tokens->defined()) { + TORCH_CHECK(num_accepted_tokens->dim() == 1 && num_accepted_tokens->size(0) == seq_num && + (num_accepted_tokens->scalar_type() == at::kInt || + num_accepted_tokens->scalar_type() == at::kLong), + "recurrent_kda: num_accepted_tokens must be int32/int64 [seq_num]."); + TORCH_CHECK(num_accepted_tokens->device() == query.device(), + "recurrent_kda: num_accepted_tokens must be on the same device as query."); + } + TORCH_CHECK(!safe_gate || (lower_bound >= -5.0 && lower_bound < 0.0), + "recurrent_kda: lower_bound must be in [-5,0) for safe gate."); + + at::Tensor output = at::empty_like(value); + at::Tensor final_state = initial_state; + const at::Tensor& accepted = c10::value_or_else( + num_accepted_tokens, [] { return at::Tensor(); }); + const char* layout = is_tnd ? "TND" : "BSND"; + // vLLM consumes the cache mutation through initial_state and returns only + // the attention output. Avoid materializing a second full state tensor. + bool output_final_state = false; + bool inplace_final_state = true; + bool state_v_first = true; + EXEC_NPU_CMD( + aclnnRecurrentKda, + query, + key, + value, + gate, + beta, + initial_state, + cu_seqlens, + ssm_state_indices, + a_log, + dt_bias, + accepted, + layout, + scale, + output_final_state, + inplace_final_state, + use_qk_l2norm_in_kernel, + use_gate_in_kernel, + use_beta_sigmoid_in_kernel, + allow_neg_eigval, + safe_gate, + lower_bound, + state_v_first, + output, + final_state); + return output; +} + +} // namespace vllm_ascend +#endif diff --git a/csrc/build_aclnn.sh b/csrc/build_aclnn.sh index 80451f120a98..b3e7d3aace90 100755 --- a/csrc/build_aclnn.sh +++ b/csrc/build_aclnn.sh @@ -119,12 +119,17 @@ elif [[ "$SOC_VERSION" =~ ^ascend910b ]]; then "hc_post" "inplace_partial_rotary_mul" "rms_norm_dynamic_quant" + "dequant_situ_quant" "dequant_swiglu_quant" "grouped_matmul_swiglu_quant" "grouped_matmul_swiglu_quant_v2" "recurrent_gated_delta_rule" + "recurrent_kda" "chunk_fwd_o" "chunk_gated_delta_rule_fwd_h" + "chunk_kda_fwd" + "kda_gate_cumsum" + "kda_layout_swap12" "store_kv_block" "store_kv_block_metadata" "sparse_attention_score" @@ -167,12 +172,17 @@ elif [[ "$SOC_VERSION" =~ ^ascend910_93 ]]; then "hc_post" "inplace_partial_rotary_mul" "rms_norm_dynamic_quant" + "dequant_situ_quant" "dequant_swiglu_quant" "grouped_matmul_swiglu_quant" "grouped_matmul_swiglu_quant_v2" "recurrent_gated_delta_rule" + "recurrent_kda" "chunk_fwd_o" "chunk_gated_delta_rule_fwd_h" + "chunk_kda_fwd" + "kda_gate_cumsum" + "kda_layout_swap12" "store_kv_block" "store_kv_block_metadata" "sparse_attention_score" @@ -200,11 +210,16 @@ elif [[ "$SOC_VERSION" =~ ^ascend950 ]]; then "hc_post" "hc_pre" "swiglu_group_quant" + "situ_mx_quant" "indexer_compress_epilog_v2" "causal_conv1d" "recurrent_gated_delta_rule" + "recurrent_kda" "chunk_fwd_o" "chunk_gated_delta_rule_fwd_h" + "chunk_kda_fwd" + "kda_gate_cumsum" + "kda_layout_swap12" "store_kv_block" "store_kv_block_metadata" "k2q_csr" diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_def.cpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_def.cpp index 630c87da9212..4b1841bbb757 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_def.cpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_def.cpp @@ -1,11 +1,11 @@ /** - * Copyright (c) 2025 Tianjin University, Ltd. - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of - * the BSD 3-Clause License (the "License"). - * Please refer to the License for details. You may not use this file except in compliance with the License. - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. - */ + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ /*! * \file chunk_gated_delta_rule_fwd_h_def.cpp @@ -42,7 +42,14 @@ class ChunkGatedDeltaRuleFwdH : public OpDef { .AutoContiguous(); this->Input("g") - .ParamType(REQUIRED) + .ParamType(OPTIONAL) + .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .AutoContiguous(); + + this->Input("gk") + .ParamType(OPTIONAL) .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16}) .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) @@ -53,7 +60,7 @@ class ChunkGatedDeltaRuleFwdH : public OpDef { .DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_FLOAT}) .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) - .IgnoreContiguous(); + .AutoContiguous(); this->Input("cu_seqlens") .ParamType(OPTIONAL) @@ -91,7 +98,6 @@ class ChunkGatedDeltaRuleFwdH : public OpDef { this->Attr("output_final_state").AttrType(REQUIRED).Bool(false); this->Attr("chunk_size").AttrType(REQUIRED).Int(64); - this->Attr("inital_state_stride0").AttrType(REQUIRED).Int(0); OpAICoreConfig aicore_config; aicore_config.DynamicCompileStaticFlag(true) diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_tiling.cpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_tiling.cpp index 27172db94514..673ff7d8e865 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_tiling.cpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_tiling.cpp @@ -1,5 +1,5 @@ /** - * Copyright (c) 2025 Tianjin University, Ltd. + * Copyright (c) 2026 Tianjin University, Ltd.  * This program is free software, you can redistribute it and/or modify it under the terms and conditions of  * the BSD 3-Clause License (the "License").  * Please refer to the License for details. You may not use this file except in compliance with the License. @@ -14,27 +14,45 @@ #include "chunk_gated_delta_rule_fwd_h_tiling.h" #include -#include "../tiling_base/data_copy_transpose_tiling.h" -#include "../tiling_base/tiling_templates_registry.h" +#include "tiling_base/data_copy_transpose_tiling.h" +#include "tiling_base/tiling_templates_registry.h" +#include "chunk_gated_delta_rule_fwd_h_tiling_processor.h" namespace optiling { + +// Maps a ge::DataType to the {fp16:0, bf16:1, fp32:2} convention shared with the kernel. +static int64_t GdnFwdHDtypeToEnum(ge::DataType dtype) +{ + if (dtype == ge::DT_BF16) { + return GDN_FWD_H_DTYPE_BF16; + } + if (dtype == ge::DT_FLOAT16) { + return GDN_FWD_H_DTYPE_FP16; + } + return GDN_FWD_H_DTYPE_FP32; +} static constexpr size_t INPUT_K_IDX = 0; static constexpr size_t INPUT_W_IDX = 1; static constexpr size_t INPUT_U_IDX = 2; static constexpr size_t INPUT_G_IDX = 3; -static constexpr size_t INPUT_INITIAL_STATE_IDX = 4; -static constexpr size_t INPUT_SEQLENS_IDX = 5; -static constexpr size_t INPUT_CHUNK_INDICES_IDX = 6; +static constexpr size_t INPUT_GK_IDX = 4; +static constexpr size_t INPUT_INITIAL_STATE_IDX = 5; +static constexpr size_t INPUT_SEQLENS_IDX = 6; +static constexpr size_t INPUT_CHUNK_INDICES_IDX = 7; static constexpr size_t ATTR_STORE_FINAL_STATE_IDX = 0; static constexpr size_t ATTR_CHUNK_SIZE_IDX = 1; -static constexpr size_t ATTR_INITIAL_STATE_STRIDE_IDX = 2; static constexpr size_t DIM_BATCH = 0; static constexpr size_t DIM_HEAD_NUM = 1; static constexpr size_t DIM_SEQLEN = 2; static constexpr size_t DIM_HEAD_DIM = 3; +static constexpr uint32_t TILING_KEY_DEFAULT = 0; +static constexpr uint32_t TILING_KEY_V128 = 1; +static constexpr uint32_t TILING_KEY_V256 = 2; +static constexpr int64_t V_DIM_128 = 128; +static constexpr int64_t V_DIM_256 = 256; static void ChunkGatedDeltaRuleFwdHTilingDataPrint(gert::TilingContext *context, ChunkGatedDeltaRuleFwdHTilingData &tiling) { @@ -49,6 +67,7 @@ static void ChunkGatedDeltaRuleFwdHTilingDataPrint(gert::TilingContext *context, OP_LOGD(nodeName, "=== chunkSize: %ld", tiling.get_chunkSize()); OP_LOGD(nodeName, "=== useInitialState: %ld", tiling.get_useInitialState()); OP_LOGD(nodeName, "=== storeFinalState: %ld", tiling.get_storeFinalState()); + OP_LOGD(nodeName, "=== useGk: %d", tiling.get_useGk()); OP_LOGD(nodeName, "=== dataType: %ld", tiling.get_dataType()); OP_LOGD(nodeName, "=== isVariedLen: %ld", tiling.get_isVariedLen()); OP_LOGD(nodeName, "=== shapeBatch: %ld", tiling.get_shapeBatch()); @@ -60,104 +79,99 @@ ge::graphStatus Tiling4ChunkGatedDeltaRuleFwdH(gert::TilingContext *context) { OP_LOGD(context->GetNodeName(), "Tiling4ChunkGatedDeltaRuleFwdH start."); ChunkGatedDeltaRuleFwdHTilingData tiling; - + gert::Shape kStorageShape = context->GetOptionalInputShape(INPUT_K_IDX)->GetStorageShape(); gert::Shape uStorageShape = context->GetOptionalInputShape(INPUT_U_IDX)->GetStorageShape(); - int64_t seqlen = kStorageShape.GetDim(DIM_SEQLEN); - int64_t kNumHead = kStorageShape.GetDim(DIM_HEAD_NUM); - int64_t vNumHead = uStorageShape.GetDim(DIM_HEAD_NUM); - int64_t kHeadDim = kStorageShape.GetDim(DIM_HEAD_DIM); - int64_t vHeadDim = uStorageShape.GetDim(DIM_HEAD_DIM); - int64_t batch, isVariedLen, shapeBatch, tokenBatch; - auto cuSeqlensTensor = context->GetOptionalInputTensor(INPUT_SEQLENS_IDX); - if (cuSeqlensTensor == nullptr) { - isVariedLen = false; - shapeBatch = kStorageShape.GetDim(DIM_BATCH); - tokenBatch = 1; - batch = shapeBatch; - } else { - isVariedLen = true; - shapeBatch = 1; - tokenBatch = cuSeqlensTensor->GetStorageShape().GetDim(DIM_BATCH) - 1; - batch = tokenBatch; - } - auto initialStateTensor = context->GetOptionalInputTensor(INPUT_INITIAL_STATE_IDX); bool useInitialState = initialStateTensor != nullptr; - int64_t stateDataType = 2; - if (useInitialState) { - auto stateDType = initialStateTensor->GetDataType(); - if (stateDType == ge::DT_BF16) { - stateDataType = 1; - } else if (stateDType == ge::DT_FLOAT16) { - stateDataType = 0; - } - } - - auto gDType = context->GetOptionalInputTensor(INPUT_G_IDX)->GetDataType(); - int64_t gDataType = 2; - if (gDType == ge::DT_BF16) { - gDataType = 1; - } else if (gDType == ge::DT_FLOAT16) { - gDataType = 0; - } - + auto gTensor = context->GetOptionalInputTensor(INPUT_G_IDX); + auto gkTensor = context->GetOptionalInputTensor(INPUT_GK_IDX); + bool useGk = gkTensor != nullptr; + OP_CHECK_IF(gTensor == nullptr && gkTensor == nullptr, + OP_LOGE(context->GetNodeName(), "Either g or gk must be provided."), + return ge::GRAPH_FAILED); + auto gateTensor = gTensor != nullptr ? gTensor : gkTensor; + auto attrPtr = context->GetAttrs(); bool storeFinalState = *(attrPtr->GetAttrPointer(ATTR_STORE_FINAL_STATE_IDX)); int64_t chunkSize = *(attrPtr->GetAttrPointer(ATTR_CHUNK_SIZE_IDX)); - int64_t initalStateStride0 = *(attrPtr->GetAttrPointer(ATTR_INITIAL_STATE_STRIDE_IDX)); - - auto dtype = context->GetInputTensor(0)->GetDataType(); - uint64_t dataType = dtype == ge::DT_BF16 ? 1 : 0; const auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo()); - uint32_t aicCoreNum = ascendcPlatform.GetCoreNumAic(); - context->SetBlockDim(aicCoreNum); - - constexpr size_t WORKSPACE_RSV_BYTE = 16 * 1024 * 1024; - constexpr size_t GM_ALIGN = 512; - constexpr int64_t PING_PONG_STAGES = 2; - - size_t workspaceOffset = ascendcPlatform.GetLibApiWorkSpaceSize(); - workspaceOffset += WORKSPACE_RSV_BYTE; - - tiling.set_vWorkspaceOffset(workspaceOffset); - workspaceOffset += (aicCoreNum * chunkSize * vHeadDim * sizeof(float) * PING_PONG_STAGES + GM_ALIGN) / GM_ALIGN * GM_ALIGN; - - tiling.set_vUpdateWorkspaceOffset(workspaceOffset); - workspaceOffset += (aicCoreNum * chunkSize * vHeadDim * sizeof(float) * PING_PONG_STAGES + GM_ALIGN) / GM_ALIGN * GM_ALIGN; - tiling.set_hWorkspaceOffset(workspaceOffset); - workspaceOffset += (aicCoreNum * kHeadDim * vHeadDim * sizeof(float) * PING_PONG_STAGES + GM_ALIGN) / GM_ALIGN * GM_ALIGN; - - tiling.set_numSeqWorkspaceOffset(workspaceOffset); - workspaceOffset += ((tokenBatch + 1) * sizeof(int64_t) + GM_ALIGN) / GM_ALIGN * GM_ALIGN; + ChunkGatedDeltaRuleFwdHTilingContext tilingCtx{}; + tilingCtx.seqlen = kStorageShape.GetDim(DIM_SEQLEN); + tilingCtx.kNumHead = kStorageShape.GetDim(DIM_HEAD_NUM); + tilingCtx.kHeadDim = kStorageShape.GetDim(DIM_HEAD_DIM); + tilingCtx.vNumHead = uStorageShape.GetDim(DIM_HEAD_NUM); + tilingCtx.vHeadDim = uStorageShape.GetDim(DIM_HEAD_DIM); + tilingCtx.shapeBatchDim = kStorageShape.GetDim(DIM_BATCH); + tilingCtx.hasCuSeqlens = cuSeqlensTensor != nullptr; + tilingCtx.cuSeqlensDim0 = + cuSeqlensTensor != nullptr ? cuSeqlensTensor->GetStorageShape().GetDim(DIM_BATCH) : 0; + tilingCtx.dataType = GdnFwdHDtypeToEnum(context->GetInputTensor(0)->GetDataType()); + tilingCtx.gDataType = GdnFwdHDtypeToEnum(gateTensor->GetDataType()); + tilingCtx.useInitialState = useInitialState; + tilingCtx.stateDataType = + useInitialState ? GdnFwdHDtypeToEnum(initialStateTensor->GetDataType()) : GDN_FWD_H_DTYPE_FP32; + tilingCtx.storeFinalState = storeFinalState; + tilingCtx.chunkSize = chunkSize; + tilingCtx.useGk = useGk; + tilingCtx.aicCoreNum = ascendcPlatform.GetCoreNumAic(); + tilingCtx.libApiWorkSpaceSize = ascendcPlatform.GetLibApiWorkSpaceSize(); + + if (tilingCtx.vNumHead % tilingCtx.kNumHead != 0) { + OP_LOGE(context->GetNodeName(), "Check head num failed, vNumHead should be divisible by kNumHead."); + return ge::GRAPH_FAILED; + } + if (tilingCtx.vHeadDim > V_DIM_256) { + OP_LOGE(context->GetNodeName(), "Check u shape failed, vHeadDim should be <= %ld, but get %ld.", + V_DIM_256, tilingCtx.vHeadDim); + return ge::GRAPH_FAILED; + } - tiling.set_numChunksWorkspaceOffset(workspaceOffset); - workspaceOffset += ((tokenBatch + 1) * sizeof(int64_t) + GM_ALIGN) / GM_ALIGN * GM_ALIGN; + ::ChunkGatedDeltaRuleFwdHTilingData plainTiling{}; + uint32_t blockDim = 0; + size_t workspaceSize = 0; + ChunkGatedDeltaRuleFwdHTilingProcessor processor(tilingCtx); + processor.Process(plainTiling, blockDim, workspaceSize); + + // 310P only ships KERNEL_TASK_TYPE_DEFAULT; keep key 0 so BinaryGetFunctionByEntry succeeds. + // Newer arches use keys 1/2 to select TileShapes128/256. + uint32_t tilingKey = TILING_KEY_DEFAULT; + if (ascendcPlatform.GetSocVersion() != platform_ascendc::SocVersion::ASCEND310P) { + tilingKey = plainTiling.vHeadDim > V_DIM_128 ? TILING_KEY_V256 : TILING_KEY_V128; // gitleaks:allow + } + context->SetTilingKey(tilingKey); + OP_LOGD(context->GetNodeName(), "tilingKey: %u (vHeadDim=%ld)", tilingKey, plainTiling.vHeadDim); - workspaceOffset += WORKSPACE_RSV_BYTE; + context->SetBlockDim(blockDim); size_t *currentWorkspace = context->GetWorkspaceSizes(1); - currentWorkspace[0] = (workspaceOffset - 0); - - tiling.set_batch(batch); - tiling.set_seqlen(seqlen); - tiling.set_kNumHead(kNumHead); - tiling.set_vNumHead(vNumHead); - tiling.set_kHeadDim(kHeadDim); - tiling.set_vHeadDim(vHeadDim); - tiling.set_chunkSize(chunkSize); - tiling.set_initalStateStride0(initalStateStride0); - tiling.set_useInitialState(useInitialState); - tiling.set_storeFinalState(storeFinalState); - tiling.set_dataType(dataType); - tiling.set_stateDataType(stateDataType); - tiling.set_gDataType(gDataType); - tiling.set_isVariedLen(isVariedLen); - tiling.set_shapeBatch(shapeBatch); - tiling.set_tokenBatch(tokenBatch); + currentWorkspace[0] = workspaceSize; + + tiling.set_batch(plainTiling.batch); + tiling.set_seqlen(plainTiling.seqlen); + tiling.set_kNumHead(plainTiling.kNumHead); + tiling.set_vNumHead(plainTiling.vNumHead); + tiling.set_kHeadDim(plainTiling.kHeadDim); + tiling.set_vHeadDim(plainTiling.vHeadDim); + tiling.set_chunkSize(plainTiling.chunkSize); + tiling.set_useInitialState(plainTiling.useInitialState); + tiling.set_storeFinalState(plainTiling.storeFinalState); + tiling.set_dataType(plainTiling.dataType); + tiling.set_stateDataType(plainTiling.stateDataType); + tiling.set_gDataType(plainTiling.gDataType); + tiling.set_isVariedLen(plainTiling.isVariedLen); + tiling.set_shapeBatch(plainTiling.shapeBatch); + tiling.set_tokenBatch(plainTiling.tokenBatch); + tiling.set_useGk(plainTiling.useGk); + tiling.set_vWorkspaceOffset(plainTiling.vWorkspaceOffset); + tiling.set_vUpdateWorkspaceOffset(plainTiling.vUpdateWorkspaceOffset); + tiling.set_kDecayWorkspaceOffset(plainTiling.kDecayWorkspaceOffset); + tiling.set_hWorkspaceOffset(plainTiling.hWorkspaceOffset); + tiling.set_numSeqWorkspaceOffset(plainTiling.numSeqWorkspaceOffset); + tiling.set_numChunksWorkspaceOffset(plainTiling.numChunksWorkspaceOffset); tiling.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity()); context->GetRawTilingData()->SetDataSize(tiling.GetDataSize()); diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_tiling.h b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_tiling.h index 44f3a01a947e..c3123404b875 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_tiling.h +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_tiling.h @@ -1,5 +1,5 @@ /** - * Copyright (c) 2025 Tianjin University, Ltd. + * Copyright (c) 2026 Tianjin University, Ltd.  * This program is free software, you can redistribute it and/or modify it under the terms and conditions of  * the BSD 3-Clause License (the "License").  * Please refer to the License for details. You may not use this file except in compliance with the License. @@ -28,17 +28,19 @@ TILING_DATA_FIELD_DEF(int64_t, vNumHead); TILING_DATA_FIELD_DEF(int64_t, kHeadDim); TILING_DATA_FIELD_DEF(int64_t, vHeadDim); TILING_DATA_FIELD_DEF(int64_t, chunkSize); -TILING_DATA_FIELD_DEF(int64_t, initalStateStride0); TILING_DATA_FIELD_DEF(bool, useInitialState); TILING_DATA_FIELD_DEF(bool, storeFinalState); TILING_DATA_FIELD_DEF(int64_t, dataType); TILING_DATA_FIELD_DEF(int64_t, gDataType); TILING_DATA_FIELD_DEF(int64_t, stateDataType); +TILING_DATA_FIELD_DEF(bool, hasGk); TILING_DATA_FIELD_DEF(int64_t, isVariedLen); TILING_DATA_FIELD_DEF(int64_t, shapeBatch); TILING_DATA_FIELD_DEF(int64_t, tokenBatch); +TILING_DATA_FIELD_DEF(bool, useGk); TILING_DATA_FIELD_DEF(int64_t, vWorkspaceOffset); TILING_DATA_FIELD_DEF(int64_t, vUpdateWorkspaceOffset); +TILING_DATA_FIELD_DEF(int64_t, kDecayWorkspaceOffset); TILING_DATA_FIELD_DEF(int64_t, hWorkspaceOffset); TILING_DATA_FIELD_DEF(int64_t, numSeqWorkspaceOffset); TILING_DATA_FIELD_DEF(int64_t, numChunksWorkspaceOffset); diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_tiling_processor.h b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_tiling_processor.h new file mode 100644 index 000000000000..1853175d170a --- /dev/null +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/chunk_gated_delta_rule_fwd_h_tiling_processor.h @@ -0,0 +1,152 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +/*! + * \file chunk_gated_delta_rule_fwd_h_tiling_processor.h + * \brief Tiling computation decoupled from gert::TilingContext, reusable by both the + * aclnn tiling entry and the fast kernel launch C++ extension. + * + * The caller is responsible for resolving framework-specific information (shapes, dtypes, + * platform core number, lib-api workspace size) into the plain context struct below. The + * processor then fills the plain ChunkGatedDeltaRuleFwdHTilingData together with the block + * dim and the total workspace size, mirroring exactly the original Tiling4ChunkGatedDeltaRuleFwdH. + */ + +#ifndef CHUNK_GATED_DELTA_RULE_FWD_H_TILING_PROCESSOR_H +#define CHUNK_GATED_DELTA_RULE_FWD_H_TILING_PROCESSOR_H + +#include +#include + +#include "../op_kernel/chunk_gated_delta_rule_fwd_h_struct.h" + +namespace optiling { + +// dtype enum convention shared with the kernel: 0 - fp16, 1 - bf16, 2 - fp32 +static constexpr int64_t GDN_FWD_H_DTYPE_FP16 = 0; +static constexpr int64_t GDN_FWD_H_DTYPE_BF16 = 1; +static constexpr int64_t GDN_FWD_H_DTYPE_FP32 = 2; + +static constexpr size_t GDN_FWD_H_WORKSPACE_RSV_BYTE = 16 * 1024 * 1024; +static constexpr size_t GDN_FWD_H_GM_ALIGN = 512; +static constexpr int64_t GDN_FWD_H_PING_PONG_STAGES = 2; + +// Plain, framework-agnostic inputs needed to compute the tiling. +struct ChunkGatedDeltaRuleFwdHTilingContext { + // shapes + int64_t seqlen; // k.dim(2) + int64_t kNumHead; // k.dim(1) + int64_t kHeadDim; // k.dim(3) + int64_t vNumHead; // u.dim(1) + int64_t vHeadDim; // u.dim(3) + int64_t shapeBatchDim; // k.dim(0) + // variable length + bool hasCuSeqlens; + int64_t cuSeqlensDim0; // length of cu_seqlens (only used when hasCuSeqlens) + // dtypes (use GDN_FWD_H_DTYPE_*) + int64_t dataType; // input (k/w/u) dtype: fp16 or bf16 + int64_t gDataType; // g dtype + bool useInitialState; + int64_t stateDataType; // initial/final state dtype + bool useGk; + // attrs + bool storeFinalState; + int64_t chunkSize; + // platform + uint32_t aicCoreNum; + size_t libApiWorkSpaceSize; +}; + +class ChunkGatedDeltaRuleFwdHTilingProcessor { +public: + explicit ChunkGatedDeltaRuleFwdHTilingProcessor(const ChunkGatedDeltaRuleFwdHTilingContext &ctx) : ctx_(ctx) {} + + // Fills the plain tiling struct, the block dim and the total workspace size. + void Process(::ChunkGatedDeltaRuleFwdHTilingData &tiling, uint32_t &blockDim, size_t &workspaceSize) const + { + int64_t isVariedLen; + int64_t shapeBatch; + int64_t tokenBatch; + int64_t batch; + + if (!ctx_.hasCuSeqlens) { + isVariedLen = 0; + shapeBatch = ctx_.shapeBatchDim; + tokenBatch = 1; + batch = shapeBatch; + } else { + isVariedLen = 1; + shapeBatch = 1; + tokenBatch = ctx_.cuSeqlensDim0 - 1; // gitleaks:allow + batch = tokenBatch; + } + + blockDim = ctx_.aicCoreNum; + const int64_t aicCoreNum = static_cast(ctx_.aicCoreNum); + const int64_t chunkSize = ctx_.chunkSize; + const int64_t kHeadDim = ctx_.kHeadDim; + const int64_t vHeadDim = ctx_.vHeadDim; + + size_t workspaceOffset = ctx_.libApiWorkSpaceSize; + workspaceOffset += GDN_FWD_H_WORKSPACE_RSV_BYTE; + + tiling.vWorkspaceOffset = static_cast(workspaceOffset); + workspaceOffset += AlignUp(static_cast(aicCoreNum * chunkSize * vHeadDim * static_cast(sizeof(float)) * GDN_FWD_H_PING_PONG_STAGES)); + + tiling.vUpdateWorkspaceOffset = static_cast(workspaceOffset); + workspaceOffset += AlignUp(static_cast(aicCoreNum * chunkSize * vHeadDim * static_cast(sizeof(float)) * GDN_FWD_H_PING_PONG_STAGES)); + + tiling.kDecayWorkspaceOffset = static_cast(workspaceOffset); + if (ctx_.useGk) { + workspaceOffset += AlignUp(static_cast(aicCoreNum * chunkSize * kHeadDim * static_cast(sizeof(float)) * GDN_FWD_H_PING_PONG_STAGES)); + } + + tiling.hWorkspaceOffset = static_cast(workspaceOffset); + workspaceOffset += AlignUp(static_cast(aicCoreNum * kHeadDim * vHeadDim * static_cast(sizeof(float)) * GDN_FWD_H_PING_PONG_STAGES)); + + tiling.numSeqWorkspaceOffset = static_cast(workspaceOffset); + workspaceOffset += AlignUp(static_cast((tokenBatch + 1) * static_cast(sizeof(int64_t)))); + + tiling.numChunksWorkspaceOffset = static_cast(workspaceOffset); + workspaceOffset += AlignUp(static_cast((tokenBatch + 1) * static_cast(sizeof(int64_t)))); + + workspaceOffset += GDN_FWD_H_WORKSPACE_RSV_BYTE; + workspaceSize = workspaceOffset; + + tiling.batch = batch; + tiling.seqlen = ctx_.seqlen; + tiling.kNumHead = ctx_.kNumHead; + tiling.vNumHead = ctx_.vNumHead; + tiling.kHeadDim = ctx_.kHeadDim; + tiling.vHeadDim = ctx_.vHeadDim; + tiling.chunkSize = chunkSize; + tiling.useInitialState = ctx_.useInitialState; + tiling.storeFinalState = ctx_.storeFinalState; + tiling.dataType = ctx_.dataType; + tiling.gDataType = ctx_.gDataType; + tiling.stateDataType = ctx_.stateDataType; + tiling.isVariedLen = isVariedLen; + tiling.shapeBatch = shapeBatch; + tiling.tokenBatch = tokenBatch; + tiling.useGk = ctx_.useGk; + } + +private: + // Mirrors the original "(x + GM_ALIGN) / GM_ALIGN * GM_ALIGN" alignment. + static size_t AlignUp(size_t x) + { + return (x + GDN_FWD_H_GM_ALIGN) / GDN_FWD_H_GM_ALIGN * GDN_FWD_H_GM_ALIGN; + } + + const ChunkGatedDeltaRuleFwdHTilingContext &ctx_; +}; + +} // namespace optiling + +#endif // CHUNK_GATED_DELTA_RULE_FWD_H_TILING_PROCESSOR_H diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/aclnn_chunk_gated_delta_rule_fwd_h.cpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/aclnn_chunk_gated_delta_rule_fwd_h.cpp index 55386a4c8634..5d00daece806 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/aclnn_chunk_gated_delta_rule_fwd_h.cpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/aclnn_chunk_gated_delta_rule_fwd_h.cpp @@ -1,5 +1,5 @@ /** - * Copyright (c) 2025 Tianjin University, Ltd. + * Copyright (c) 2026 Tianjin University, Ltd. * This program is free software, you can redistribute it and/or modify it under the terms and conditions of * CANN Open Software License Agreement Version 2.0 (the "License"). * Please refer to the License for details. You may not use this file except in compliance with the License. @@ -11,10 +11,11 @@ #include "chunk_gated_delta_rule_fwd_h.h" #include #include -#include #include "aclnn_kernels/transdata.h" #include "aclnn_kernels/contiguous.h" +#include "aclnn_kernels/reshape.h" +#include "aclnn_kernels/slice.h" #include "acl/acl.h" #include "aclnn/aclnn_base.h" #include "aclnn_kernels/common/op_error_check.h" @@ -32,6 +33,10 @@ using namespace op; +namespace l0op { +const aclTensor *ZerosLike(const aclTensor *self, aclOpExecutor *executor); +} + #ifdef __cplusplus extern "C" { #endif @@ -73,11 +78,77 @@ static aclnnStatus CheckFormat(ChunkGatedDeltaRuleFwdHParams params) static aclnnStatus CheckShape(ChunkGatedDeltaRuleFwdHParams params) { + auto kShape = params.k->GetViewShape(); + auto wShape = params.w->GetViewShape(); + auto uShape = params.u->GetViewShape(); + CHECK_COND(kShape.GetDimNum() == 4 && wShape.GetDimNum() == 4 && uShape.GetDimNum() == 4, + ACLNN_ERR_PARAM_INVALID, "k, w and u must be rank-4 BNSD tensors."); + CHECK_COND(kShape.GetDim(0) == wShape.GetDim(0) && kShape.GetDim(0) == uShape.GetDim(0) && + wShape.GetDim(1) == uShape.GetDim(1) && kShape.GetDim(2) == wShape.GetDim(2) && + kShape.GetDim(2) == uShape.GetDim(2) && kShape.GetDim(3) == wShape.GetDim(3), + ACLNN_ERR_PARAM_INVALID, + "k, w and u must match in B/T, w and u must match in HV, and k and w must match in K."); + CHECK_COND(uShape.GetDim(1) >= kShape.GetDim(1) && uShape.GetDim(1) % kShape.GetDim(1) == 0, + ACLNN_ERR_PARAM_INVALID, "u HV must be greater than or equal to k H and divisible by H."); + if (params.gOptional != nullptr) { + auto gShape = params.gOptional->GetViewShape(); + CHECK_COND(gShape.GetDimNum() == 3 && gShape.GetDim(0) == uShape.GetDim(0) && + gShape.GetDim(1) == uShape.GetDim(1) && gShape.GetDim(2) == uShape.GetDim(2), + ACLNN_ERR_PARAM_INVALID, "g must have shape [B, HV, T]."); + } return ACLNN_SUCCESS; } +static const aclTensor *MakeNeutralGate(const ChunkGatedDeltaRuleFwdHParams ¶ms, aclOpExecutor *executor) +{ + auto gkShape = params.gkOptional->GetViewShape(); + int64_t offsetsData[] = {0, 0, 0, 0}; + int64_t sizesData[] = {gkShape.GetDim(0), gkShape.GetDim(1), gkShape.GetDim(2), 1}; + auto offsets = executor->AllocIntArray(offsetsData, 4); + auto sizes = executor->AllocIntArray(sizesData, 4); + if (offsets == nullptr || sizes == nullptr) { + return nullptr; + } + auto gateLane = l0op::Slice(params.gkOptional, offsets, sizes, executor); + if (gateLane == nullptr) { + return nullptr; + } + gateLane = l0op::Contiguous(gateLane, executor); + if (gateLane == nullptr) { + return nullptr; + } + op::Shape gateShape; + gateShape.AppendDim(gkShape.GetDim(0)); + gateShape.AppendDim(gkShape.GetDim(1)); + gateShape.AppendDim(gkShape.GetDim(2)); + gateLane = l0op::Reshape(gateLane, gateShape, executor); + return gateLane == nullptr ? nullptr : l0op::ZerosLike(gateLane, executor); +} + static aclnnStatus CheckDtype(ChunkGatedDeltaRuleFwdHParams params) { + auto inputDtype = params.k->GetDataType(); + CHECK_COND(inputDtype == DataType::DT_FLOAT16 || inputDtype == DataType::DT_BF16, + ACLNN_ERR_PARAM_INVALID, "k dtype must be float16 or bfloat16."); + CHECK_COND(params.w->GetDataType() == inputDtype && params.u->GetDataType() == inputDtype, + ACLNN_ERR_PARAM_INVALID, "k, w and u must have the same dtype."); + CHECK_COND(params.hOut->GetDataType() == inputDtype && params.vNewOut->GetDataType() == inputDtype, + ACLNN_ERR_PARAM_INVALID, "hOut and vNewOut dtype must match k, w and u."); + auto gateDtype = params.gOptional != nullptr ? params.gOptional->GetDataType() : params.gkOptional->GetDataType(); + CHECK_COND(gateDtype == DataType::DT_FLOAT || gateDtype == inputDtype, + ACLNN_ERR_PARAM_INVALID, "g/gk dtype must be float32 or match k dtype."); + if (params.gOptional != nullptr && params.gkOptional != nullptr) { + CHECK_COND(params.gOptional->GetDataType() == params.gkOptional->GetDataType(), + ACLNN_ERR_PARAM_INVALID, "g and gk must have the same dtype when both are provided."); + } + if (params.outputFinalState) { + CHECK_COND(params.finalStateOut != nullptr, ACLNN_ERR_PARAM_NULLPTR, + "finalStateOut must be provided when outputFinalState is true."); + auto stateDtype = params.initalStateOptional != nullptr ? params.initalStateOptional->GetDataType() + : DataType::DT_FLOAT; + CHECK_COND(params.finalStateOut->GetDataType() == stateDtype, ACLNN_ERR_PARAM_INVALID, + "finalStateOut dtype must match initial state, or be float32 when initial state is absent."); + } return ACLNN_SUCCESS; } @@ -96,8 +167,14 @@ static aclnnStatus ParamsDataContiguous(ChunkGatedDeltaRuleFwdHParams ¶ms, a "Contiguous w failed."); CHECK_COND(DataContiguous(params.u, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID, "Contiguous u failed."); - CHECK_COND(DataContiguous(params.gOptional, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID, - "Contiguous gOptional failed."); + if (params.gOptional != nullptr) { + CHECK_COND(DataContiguous(params.gOptional, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID, + "Contiguous gOptional failed."); + } + if (params.gkOptional != nullptr) { + CHECK_COND(DataContiguous(params.gkOptional, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID, + "Contiguous gkOptional failed."); + } if (params.initalStateOptional != nullptr) { CHECK_COND(DataContiguous(params.initalStateOptional, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID, "Contiguous initalStateOptional failed."); @@ -106,17 +183,15 @@ static aclnnStatus ParamsDataContiguous(ChunkGatedDeltaRuleFwdHParams ¶ms, a return ACLNN_SUCCESS; } -static aclnnStatus CheckGOptionalNonNull(const ChunkGatedDeltaRuleFwdHParams ¶ms) +static aclnnStatus CheckGateOptionalNonNull(const ChunkGatedDeltaRuleFwdHParams ¶ms) { - CHECK_COND(params.gOptional != nullptr, ACLNN_ERR_PARAM_INVALID, - "g is an optional-parameter slot in the API but only a non-null aclTensor is supported; nullptr is not allowed until g=None is implemented."); + CHECK_COND(params.gOptional != nullptr || params.gkOptional != nullptr, ACLNN_ERR_PARAM_INVALID, + "Either g or gk must be provided."); return ACLNN_SUCCESS; } static aclnnStatus CheckReservedOptions(const ChunkGatedDeltaRuleFwdHParams ¶ms) { - CHECK_COND(params.gkOptional == nullptr, ACLNN_ERR_PARAM_INVALID, - "gk is reserved for ChunkGatedDeltaRuleFwdH and must be nullptr."); CHECK_COND(params.saveNewValue, ACLNN_ERR_PARAM_INVALID, "save_new_value is reserved and only true is supported."); CHECK_COND(!params.useExp2, ACLNN_ERR_PARAM_INVALID, @@ -126,11 +201,34 @@ static aclnnStatus CheckReservedOptions(const ChunkGatedDeltaRuleFwdHParams &par return ACLNN_SUCCESS; } +static aclnnStatus CheckGkParams(const ChunkGatedDeltaRuleFwdHParams ¶ms) +{ + if (params.gkOptional != nullptr) { + auto gkShape = params.gkOptional->GetViewShape(); + CHECK_COND(gkShape.GetDimNum() == 4, ACLNN_ERR_PARAM_INVALID, + "gk must have rank 4 when provided, got rank %ld.", gkShape.GetDimNum()); + CHECK_COND(gkShape.GetDim(3) == params.k->GetViewShape().GetDim(3), ACLNN_ERR_PARAM_INVALID, + "gk.shape[3] (K) must match k.shape[3] (K)."); + CHECK_COND(gkShape.GetDim(2) == params.k->GetViewShape().GetDim(2), ACLNN_ERR_PARAM_INVALID, + "gk.shape[2] (T) must match k.shape[2] (T)."); + CHECK_COND(gkShape.GetDim(1) == params.u->GetViewShape().GetDim(1), ACLNN_ERR_PARAM_INVALID, + "gk.shape[1] (HV) must match u.shape[1] (HV)."); + CHECK_COND(gkShape.GetDim(0) == params.k->GetViewShape().GetDim(0), ACLNN_ERR_PARAM_INVALID, + "gk.shape[0] (B) must match k.shape[0] (B)."); + if (params.gOptional != nullptr) { + CHECK_COND(params.gkOptional->GetDataType() == params.gOptional->GetDataType(), ACLNN_ERR_PARAM_INVALID, + "gk.dtype must match g.dtype when both are provided."); + } + } + return ACLNN_SUCCESS; +} + static aclnnStatus CheckParams(ChunkGatedDeltaRuleFwdHParams params) { CHECK_RET(CheckNotNull(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); - CHECK_RET(CheckGOptionalNonNull(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(CheckGateOptionalNonNull(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); CHECK_RET(CheckReservedOptions(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(CheckGkParams(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); CHECK_RET(CheckFormat(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); CHECK_RET(CheckShape(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); CHECK_RET(CheckDtype(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); @@ -184,19 +282,11 @@ aclnnStatus aclnnChunkGatedDeltaRuleFwdHGetWorkspaceSize( CHECK_RET(ret == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); CHECK_COND(ParamsDataContiguous(params, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID, "ParamsDataContiguous failed."); - - // aclGetViewStrides obtains the strides and the number of strides corresponding to aclTensor - int64_t *initialStateStridesValuePtr = nullptr; - int64_t initialStateStridesValue = 0; - uint64_t initialStateStridesNum = 0; - - if (initalStateOptional != nullptr) { - ret = aclGetViewStrides(initalStateOptional, &initialStateStridesValuePtr, &initialStateStridesNum); - CHECK_RET(ret == ACLNN_SUCCESS, ret); - initialStateStridesValue = initialStateStridesValuePtr[initialStateStridesNum - 2]; + if (params.gOptional == nullptr) { + params.gOptional = MakeNeutralGate(params, executorPtr); + CHECK_RET(params.gOptional != nullptr, ACLNN_ERR_INNER_NULLPTR); } - - auto result = l0op::ChunkGatedDeltaRuleFwdH(params.k, params.w, params.u, params.gOptional, params.initalStateOptional, params.cuSeqlensOptional, params.chunkIndicesOptional, params.outputFinalState, params.chunkSize, initialStateStridesValue, params.hOut, params.vNewOut, params.finalStateOut, executorPtr); + auto result = l0op::ChunkGatedDeltaRuleFwdH(params.k, params.w, params.u, params.gOptional, params.gkOptional, params.initalStateOptional, params.cuSeqlensOptional, params.chunkIndicesOptional, params.outputFinalState, params.chunkSize, params.hOut, params.vNewOut, params.finalStateOut, executorPtr); CHECK_RET(result[0] != nullptr, ACLNN_ERR_PARAM_NULLPTR); // If the output tensor is non-contiguous, convert the calculated contiguous tensor to non-contiguous. diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/aclnn_chunk_gated_delta_rule_fwd_h.h b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/aclnn_chunk_gated_delta_rule_fwd_h.h index 956c9f63f4c6..7e77615a5d10 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/aclnn_chunk_gated_delta_rule_fwd_h.h +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/aclnn_chunk_gated_delta_rule_fwd_h.h @@ -1,5 +1,5 @@ /** - * Copyright (c) 2025 Tianjin University, Ltd. + * Copyright (c) 2026 Tianjin University, Ltd. * This program is free software, you can redistribute it and/or modify it under the terms and conditions of * CANN Open Software License Agreement Version 2.0 (the "License"). * Please refer to the License for details. You may not use this file except in compliance with the License. @@ -20,8 +20,8 @@ extern "C" { * k : required * w : required * u : required - * gOptional : optional, only non-null aclTensor is supported - * gkOptional : optional, reserved (must be nullptr) + * gOptional : optional, scalar gate tensor; either gOptional or gkOptional must be non-null + * gkOptional : optional, key-wise gate tensor; either gOptional or gkOptional must be non-null * initalStateOptional : optional * outputFinalState : required * chunkSize : required diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/chunk_gated_delta_rule_fwd_h.cpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/chunk_gated_delta_rule_fwd_h.cpp index eb05d74eb4d0..de4f9b783dbe 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/chunk_gated_delta_rule_fwd_h.cpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/chunk_gated_delta_rule_fwd_h.cpp @@ -1,5 +1,5 @@ /** - * Copyright (c) 2025 Tianjin University, Ltd. + * Copyright (c) 2026 Tianjin University, Ltd. * This program is free software, you can redistribute it and/or modify it under the terms and conditions of * CANN Open Software License Agreement Version 2.0 (the "License"). * Please refer to the License for details. You may not use this file except in compliance with the License. @@ -23,18 +23,18 @@ const std::array ChunkGatedDeltaRuleFwdH( const aclTensor *w, const aclTensor *u, const aclTensor *g, + const aclTensor *gkOptional, const aclTensor *initalStateOptional, const aclIntArray *cuSeqlensOptional, const aclIntArray *chunkIndicesOptional, bool outputFinalState, int64_t chunkSize, - int64_t initialStateStridesValue, const aclTensor *hOut, const aclTensor *vNewOut, const aclTensor *finalStateOut, aclOpExecutor *executor) { - L0_DFX(ChunkGatedDeltaRuleFwdH, k, w, u, g, initalStateOptional, cuSeqlensOptional, chunkIndicesOptional, outputFinalState, chunkSize, initialStateStridesValue, hOut, vNewOut, finalStateOut); + L0_DFX(ChunkGatedDeltaRuleFwdH, k, w, u, g, gkOptional, initalStateOptional, cuSeqlensOptional, chunkIndicesOptional, outputFinalState, chunkSize, hOut, vNewOut, finalStateOut); const aclTensor *actualCuSeqlens = nullptr; if (cuSeqlensOptional) { @@ -57,9 +57,9 @@ const std::array ChunkGatedDeltaRuleFwdH( } auto ret = ADD_TO_LAUNCHER_LIST_AICORE(ChunkGatedDeltaRuleFwdH, - OP_INPUT(k, w, u, g, initalStateOptional, actualCuSeqlens, actualChunkIndices), + OP_INPUT(k, w, u, g, gkOptional, initalStateOptional, actualCuSeqlens, actualChunkIndices), OP_OUTPUT(hOut, vNewOut, finalStateOut), - OP_ATTR(outputFinalState, chunkSize, initialStateStridesValue)); + OP_ATTR(outputFinalState, chunkSize)); if (ret != ACLNN_SUCCESS) { OP_LOGE(ACLNN_ERR_PARAM_INVALID, "ADD_TO_LAUNCHER_LIST_AICORE failed."); return {nullptr, nullptr, nullptr}; diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/chunk_gated_delta_rule_fwd_h.h b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/chunk_gated_delta_rule_fwd_h.h index 98817016cd73..d5ee31ef0bad 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/chunk_gated_delta_rule_fwd_h.h +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/chunk_gated_delta_rule_fwd_h.h @@ -1,5 +1,5 @@ /** - * Copyright (c) 2025 Tianjin University, Ltd. + * Copyright (c) 2026 Tianjin University, Ltd. * This program is free software, you can redistribute it and/or modify it under the terms and conditions of * CANN Open Software License Agreement Version 2.0 (the "License"). * Please refer to the License for details. You may not use this file except in compliance with the License. @@ -18,12 +18,12 @@ const std::array ChunkGatedDeltaRuleFwdH( const aclTensor *w, const aclTensor *u, const aclTensor *g, + const aclTensor *gkOptional, const aclTensor *initalStateOptional, const aclIntArray *cuSeqlensOptional, const aclIntArray *chunkIndicesOptional, bool outputFinalState, int64_t chunkSize, - int64_t initialStateStridesValue, const aclTensor *hOut, const aclTensor *vNewOut, const aclTensor *finalStateOut, diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch20/gemm/block/block_scheduler_gdn_fwd_h.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch20/gemm/block/block_scheduler_gdn_fwd_h.hpp index e24d8283bf4c..bf6737d5b552 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch20/gemm/block/block_scheduler_gdn_fwd_h.hpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch20/gemm/block/block_scheduler_gdn_fwd_h.hpp @@ -130,7 +130,8 @@ struct BlockSchedulerGdnFwdH { kHeadDim = gdnFwdHTilingData->kHeadDim; vHeadDim = gdnFwdHTilingData->vHeadDim; chunkSize = gdnFwdHTilingData->chunkSize; - initalStateStride0 = gdnFwdHTilingData->initalStateStride0; + // Contiguous state layout: stride equals vHeadDim (initalStateStride0 attr removed). + initalStateStride0 = gdnFwdHTilingData->vHeadDim; isVariedLen = gdnFwdHTilingData->isVariedLen; shapeBatch = gdnFwdHTilingData->shapeBatch; tokenBatch = gdnFwdHTilingData->tokenBatch; diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch20/gemm/kernel/gdn_fwd_h_kernel.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch20/gemm/kernel/gdn_fwd_h_kernel.hpp index 4288e533dce0..8120b757fd91 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch20/gemm/kernel/gdn_fwd_h_kernel.hpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch20/gemm/kernel/gdn_fwd_h_kernel.hpp @@ -174,7 +174,8 @@ class GDNFwdHKernel { kHeadDim = gdnFwdHTilingData->kHeadDim; vHeadDim = gdnFwdHTilingData->vHeadDim; chunkSize = gdnFwdHTilingData->chunkSize; - initalStateStride0 = gdnFwdHTilingData->initalStateStride0; + // Contiguous state layout: stride equals vHeadDim (initalStateStride0 attr removed). + initalStateStride0 = gdnFwdHTilingData->vHeadDim; useInitialState = gdnFwdHTilingData->useInitialState; storeFinalState = gdnFwdHTilingData->storeFinalState; isVariedLen = gdnFwdHTilingData->isVariedLen; diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/epilogue/block/block_epilogue_gdn_fwdh_update.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/epilogue/block/block_epilogue_gdn_fwdh_update.hpp index 0fb901d7d9a6..f4325a02e8a8 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/epilogue/block/block_epilogue_gdn_fwdh_update.hpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/epilogue/block/block_epilogue_gdn_fwdh_update.hpp @@ -4,7 +4,7 @@ * the BSD 3-Clause License (the "License"). * Please refer to the License for details. You may not use this file except in compliance with the License. * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY OR FITNESS FOR A PARTICULAR PURPOSE. */ #ifndef CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_UPDATE_HPP @@ -23,7 +23,8 @@ template < class GInputType_, class HInputType_, class HUpdateInputType_, - class FinalStateType_ + class FinalStateType_, + class KGatedTag > class BlockEpilogue < EpilogueAtlasGDNFwdHUpdate, @@ -31,10 +32,11 @@ class BlockEpilogue < GInputType_, HInputType_, HUpdateInputType_, - FinalStateType_ + FinalStateType_, + KGatedTag > { + static constexpr bool kGated = KGatedTag::value; public: - // Type aliases using DispatchPolicy = EpilogueAtlasGDNFwdHUpdate; using ArchTag = typename DispatchPolicy::ArchTag; @@ -50,26 +52,93 @@ class BlockEpilogue < constexpr uint32_t CALC_BUF_OFFSET = 0; constexpr uint32_t PING_BUF_0_OFFSET = 32 * 1024; - constexpr uint32_t PING_BUF_1_OFFSET = 64 * 1024; - constexpr uint32_t PING_BUF_2_OFFSET = 80 * 1024; + constexpr uint32_t PING_BUF_1_OFFSET = 48 * 1024; + constexpr uint32_t PING_BUF_2_OFFSET = 64 * 1024; + constexpr uint32_t PING_BUF_3_OFFSET = 80 * 1024; + constexpr uint32_t PONG_BUF_0_OFFSET = 96 * 1024; + constexpr uint32_t PONG_BUF_1_OFFSET = 112 * 1024; + constexpr uint32_t PONG_BUF_2_OFFSET = 128 * 1024; + constexpr uint32_t PONG_BUF_3_OFFSET = 144 * 1024; constexpr uint32_t PING_G_BUF_OFFSET = 160 * 1024; + constexpr uint32_t PONG_G_BUF_OFFSET = 161 * 1024; + constexpr uint32_t PING_G_SUB_BUF_OFFSET = 162 * 1024; + constexpr uint32_t PONG_G_SUB_BUF_OFFSET = 163 * 1024; + constexpr uint32_t PING_G_INPUT_BUF_OFFSET = 164 * 1024; + constexpr uint32_t PONG_G_INPUT_BUF_OFFSET = 165 * 1024; + constexpr uint32_t SHARE_BUF_OFFSET = 166 * 1024; calcUbTensor = resource.ubBuf.template GetBufferByByte(CALC_BUF_OFFSET); - hUpdateUbTensor = resource.ubBuf.template GetBufferByByte(PING_BUF_1_OFFSET); - hUbTensor = resource.ubBuf.template GetBufferByByte(PING_BUF_0_OFFSET); + hUpdateUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_0_OFFSET); + gkBroadcastUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_1_OFFSET); + hUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_3_OFFSET); + finalOutputUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_3_OFFSET); + glastUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_INPUT_BUF_OFFSET); - hOutputUbTensor = resource.ubBuf.template GetBufferByByte(PING_BUF_1_OFFSET); - finalOutputUbTensor = resource.ubBuf.template GetBufferByByte(PING_BUF_1_OFFSET); - - glastUbTensor = resource.ubBuf.template GetBufferByByte(PING_G_BUF_OFFSET); + hUpdateUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_0_OFFSET); + gkBroadcastUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_1_OFFSET); + hUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_3_OFFSET); + finalOutputUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_3_OFFSET); + glastUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_INPUT_BUF_OFFSET); + if constexpr (kGated) { + gkLastUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_SUB_BUF_OFFSET); + gkLastUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_SUB_BUF_OFFSET); + gkInputUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_INPUT_BUF_OFFSET); + gkInputUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_INPUT_BUF_OFFSET); + shareBufferGk_ = resource.ubBuf.template GetBufferByByte(SHARE_BUF_OFFSET); + } } CATLASS_DEVICE ~BlockEpilogue() {} + template + CATLASS_DEVICE + void CopyGmToUb( + AscendC::LocalTensor dst, + AscendC::GlobalTensor src, + uint32_t rows, + uint32_t cols, + uint32_t srcStride) + { + if (cols == srcStride) { + AscendC::DataCopy(dst, src, rows * cols); + return; + } + AscendC::DataCopyExtParams copyParams{ + static_cast(rows), + static_cast(cols * sizeof(Element)), + static_cast((srcStride - cols) * sizeof(Element)), + 0, + 0}; + AscendC::DataCopyPadExtParams padParams{false, 0, 0, 0}; + AscendC::DataCopyPad(dst, src, copyParams, padParams); + } + + template + CATLASS_DEVICE + void CopyUbToGm( + AscendC::GlobalTensor dst, + AscendC::LocalTensor src, + uint32_t rows, + uint32_t cols, + uint32_t dstStride) + { + if (cols == dstStride) { + AscendC::DataCopy(dst, src, rows * cols); + return; + } + AscendC::DataCopyExtParams copyParams{ + static_cast(rows), + static_cast(cols * sizeof(Element)), + 0, + static_cast((dstStride - cols) * sizeof(Element)), + 0}; + AscendC::DataCopyPad(dst, src, copyParams); + } + CATLASS_DEVICE void operator()( AscendC::GlobalTensor hOutput, @@ -77,39 +146,47 @@ class BlockEpilogue < AscendC::GlobalTensor gInput, AscendC::GlobalTensor hInput, AscendC::GlobalTensor hUpdateInput, + AscendC::GlobalTensor gkInput, uint32_t chunkSize, uint32_t kHeadDim, + uint32_t vBlockDim, uint32_t vHeadDim, Arch::CrossCoreFlag cube2Done, - bool isFinalState + bool isInitialState, + bool isFinalState, + bool storeFinalState, + bool isPing ) { + static constexpr uint32_t ROW_TILE = 16; uint32_t mActual = kHeadDim; - uint32_t nActual = vHeadDim; + uint32_t nActual = vBlockDim; + uint32_t outputStride = vHeadDim; uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); uint32_t subBlockNum = AscendC::GetSubBlockNum(); - uint32_t mActualPerSubBlock = CeilDiv(mActual, subBlockNum); - uint32_t mActualThisSubBlock = (subBlockIdx == 0) ? mActualPerSubBlock : (mActual - mActualPerSubBlock); - uint32_t mOffset = subBlockIdx * mActualPerSubBlock; - uint32_t nOffset = 0; - int64_t offsetH = mOffset * nActual + nOffset; + uint32_t rowsPerSubBlock = CeilDiv(mActual, subBlockNum); + uint32_t rowBegin = subBlockIdx * rowsPerSubBlock; + uint32_t rowEnd = rowBegin + rowsPerSubBlock; + if (rowEnd > mActual) { + rowEnd = mActual; + } + if (rowBegin >= mActual) { + Arch::CrossCoreWaitFlag(cube2Done); + return; + } AscendC::ResetMask(); - AscendC::GlobalTensor hOutputThisSubBlock = hOutput[offsetH]; AscendC::GlobalTensor gInputThisSubBlock = gInput; - AscendC::GlobalTensor hInputThisSubBlock = hInput[offsetH]; - AscendC::GlobalTensor hUpdateInputThisSubBlock = hUpdateInput[offsetH]; - AscendC::GlobalTensor finalStateThisSubBlock = finalState[offsetH]; - - AscendC::SetFlag(EVENT_ID0); - AscendC::WaitFlag(EVENT_ID0); - AscendC::DataCopy(hUbTensor, hInputThisSubBlock, mActualThisSubBlock * nActual); - AscendC::SetFlag(EVENT_ID0); - AscendC::WaitFlag(EVENT_ID0); - AscendC::Cast(calcUbTensor, hUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nActual); - AscendC::PipeBarrier(); - + + uint32_t pingpongFlag = isPing ? 0 : pongBaseEvent; + AscendC::LocalTensor hUpdateUbTensor = isPing ? hUpdateUbTensor_ping : hUpdateUbTensor_pong; + AscendC::LocalTensor gkBroadcastUbTensor = + isPing ? gkBroadcastUbTensor_ping : gkBroadcastUbTensor_pong; + AscendC::LocalTensor hUbTensor = isPing ? hUbTensor_ping : hUbTensor_pong; + AscendC::LocalTensor finalOutputUbTensor = isPing ? finalOutputUbTensor_ping : finalOutputUbTensor_pong; + AscendC::LocalTensor glastUbTensor = isPing ? glastUbTensor_ping : glastUbTensor_pong; + GElementInput gLastVal = gInputThisSubBlock.GetValue(chunkSize-1); float gLastFloat = 0.0f; if constexpr(std::is_same::value) { @@ -121,60 +198,277 @@ class BlockEpilogue < } glastUbTensor.SetValue(0, gLastFloat); - AscendC::SetFlag(EVENT_ID0); - AscendC::WaitFlag(EVENT_ID0); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); AscendC::Exp(glastUbTensor, glastUbTensor, 1); - AscendC::SetFlag(EVENT_ID0); - AscendC::WaitFlag(EVENT_ID0); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); float muls = glastUbTensor.GetValue(0); - AscendC::SetFlag(EVENT_ID0); - AscendC::WaitFlag(EVENT_ID0); - AscendC::Muls(calcUbTensor, calcUbTensor, muls, mActualThisSubBlock * nActual); + if constexpr (kGated) { + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + + if (nActual <= 128 && nActual == outputStride) { + uint32_t mActualThisSubBlock = rowEnd - rowBegin; + AscendC::GlobalTensor hOutputThisSubBlock = hOutput[rowBegin * outputStride]; + AscendC::GlobalTensor hInputThisSubBlock = hInput[rowBegin * outputStride]; + AscendC::GlobalTensor hUpdateInputThisSubBlock = hUpdateInput[rowBegin * nActual]; + AscendC::GlobalTensor finalStateThisSubBlock = finalState[rowBegin * outputStride]; + + if (storeFinalState && isInitialState && std::is_same::value) { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + } + AscendC::DataCopy(hUbTensor, hInputThisSubBlock, mActualThisSubBlock * nActual); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + + AscendC::Cast(calcUbTensor, hUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + if (storeFinalState && isFinalState && std::is_same::value) { + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } + + AscendC::Muls(calcUbTensor, calcUbTensor, muls, mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + + if constexpr (kGated) { + AscendC::GlobalTensor gkLastInput = gkInput[(chunkSize - 1) * kHeadDim + rowBegin]; + AscendC::LocalTensor gkLastUbTensor = isPing ? gkLastUbTensor_ping : gkLastUbTensor_pong; + AscendC::LocalTensor gkInputUbTensor = isPing ? gkInputUbTensor_ping : gkInputUbTensor_pong; + + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + if constexpr(std::is_same::value) { + AscendC::DataCopy(gkLastUbTensor, gkLastInput, mActualThisSubBlock); + } else { + AscendC::DataCopy(gkInputUbTensor, gkLastInput, mActualThisSubBlock); + } + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + if constexpr(!std::is_same::value) { + AscendC::Cast(gkLastUbTensor, gkInputUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock); + } + AscendC::PipeBarrier(); + AscendC::Muls(gkLastUbTensor, gkLastUbTensor, 0.6931471805599453f, + mActualThisSubBlock); + AscendC::PipeBarrier(); + AscendC::Exp(gkLastUbTensor, gkLastUbTensor, mActualThisSubBlock); + AscendC::PipeBarrier(); + + uint32_t gkBrcReptime = (mActualThisSubBlock + 8 - 1) / 8; + uint32_t dstShapeGk[2] = {gkBrcReptime * 8, nActual}; + uint32_t srcShapeGk[2] = {gkBrcReptime * 8, 1}; + AscendC::Broadcast(hUpdateUbTensor, gkLastUbTensor, dstShapeGk, srcShapeGk, shareBufferGk_); + AscendC::PipeBarrier(); + AscendC::Mul(calcUbTensor, calcUbTensor, hUpdateUbTensor, mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + } + + Arch::CrossCoreWaitFlag(cube2Done); + + if constexpr (kGated) { + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + } + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::DataCopy(hUpdateUbTensor, hUpdateInputThisSubBlock, mActualThisSubBlock * nActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::Add(hUpdateUbTensor, calcUbTensor, hUpdateUbTensor, mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + + if constexpr(std::is_same::value) { + if (storeFinalState && isFinalState) { + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::DataCopy(finalStateThisSubBlock, hUpdateUbTensor, mActualThisSubBlock * nActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + } else { + AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + AscendC::DataCopy(hOutputThisSubBlock, hUbTensor, mActualThisSubBlock * nActual); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } + } else { + if (storeFinalState && isFinalState) { + AscendC::Cast(finalOutputUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + AscendC::DataCopy(finalStateThisSubBlock, finalOutputUbTensor, mActualThisSubBlock * nActual); + } else { + AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + AscendC::DataCopy(hOutputThisSubBlock, hUbTensor, mActualThisSubBlock * nActual); + } + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } + return; + } Arch::CrossCoreWaitFlag(cube2Done); - AscendC::SetFlag(EVENT_ID0); - AscendC::WaitFlag(EVENT_ID0); - AscendC::SetFlag(EVENT_ID1); - AscendC::WaitFlag(EVENT_ID1); - AscendC::DataCopy(hUpdateUbTensor, hUpdateInputThisSubBlock, mActualThisSubBlock * nActual); - AscendC::SetFlag(EVENT_ID1); - AscendC::WaitFlag(EVENT_ID1); - AscendC::Add(hUpdateUbTensor, calcUbTensor, hUpdateUbTensor, mActualThisSubBlock * nActual); - - if (isFinalState) { - if constexpr(!std::is_same::value) { + bool waitHFromV = storeFinalState && isInitialState && std::is_same::value; + bool waitUpdateFromMte3 = false; + bool waitGkFromScalar = true; + for (uint32_t rowStart = rowBegin; rowStart < rowEnd; rowStart += ROW_TILE) { + uint32_t rowsThisTile = rowEnd - rowStart; + if (rowsThisTile > ROW_TILE) { + rowsThisTile = ROW_TILE; + } + + AscendC::GlobalTensor hOutputThisTile = hOutput[rowStart * outputStride]; + AscendC::GlobalTensor hInputThisTile = hInput[rowStart * outputStride]; + AscendC::GlobalTensor hUpdateInputThisTile = hUpdateInput[rowStart * nActual]; + AscendC::GlobalTensor finalStateThisTile = finalState[rowStart * outputStride]; + + if (waitHFromV) { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + } + CopyGmToUb(hUbTensor, hInputThisTile, rowsThisTile, nActual, outputStride); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + + AscendC::Cast(calcUbTensor, hUbTensor, AscendC::RoundMode::CAST_NONE, rowsThisTile * nActual); + AscendC::PipeBarrier(); + if (storeFinalState && isFinalState && std::is_same::value) { + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + waitHFromV = true; + } else { + waitHFromV = false; + } + + AscendC::Muls(calcUbTensor, calcUbTensor, muls, rowsThisTile * nActual); + AscendC::PipeBarrier(); + + if constexpr (kGated) { + AscendC::GlobalTensor gkLastInput = gkInput[(chunkSize - 1) * kHeadDim + rowStart]; + AscendC::LocalTensor gkLastUbTensor = isPing ? gkLastUbTensor_ping : gkLastUbTensor_pong; + AscendC::LocalTensor gkInputUbTensor = isPing ? gkInputUbTensor_ping : gkInputUbTensor_pong; + + if (waitGkFromScalar) { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + waitGkFromScalar = false; + } else { + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + } + if constexpr(std::is_same::value) { + AscendC::DataCopy(gkLastUbTensor, gkLastInput, rowsThisTile); + } else { + AscendC::DataCopy(gkInputUbTensor, gkLastInput, rowsThisTile); + } + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + if constexpr(!std::is_same::value) { + AscendC::Cast(gkLastUbTensor, gkInputUbTensor, AscendC::RoundMode::CAST_NONE, rowsThisTile); + } + AscendC::PipeBarrier(); + AscendC::Muls(gkLastUbTensor, gkLastUbTensor, 0.6931471805599453f, + rowsThisTile); + AscendC::PipeBarrier(); + AscendC::Exp(gkLastUbTensor, gkLastUbTensor, rowsThisTile); AscendC::PipeBarrier(); - AscendC::Cast(finalOutputUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nActual); - AscendC::SetFlag(EVENT_ID0); - AscendC::WaitFlag(EVENT_ID0); - AscendC::DataCopy(finalStateThisSubBlock, finalOutputUbTensor, mActualThisSubBlock * nActual); + + uint32_t gkBrcReptime = (rowsThisTile + 8 - 1) / 8; + uint32_t dstShapeGk[2] = {gkBrcReptime * 8, nActual}; + uint32_t srcShapeGk[2] = {gkBrcReptime * 8, 1}; + AscendC::Broadcast(gkBroadcastUbTensor, gkLastUbTensor, dstShapeGk, srcShapeGk, + shareBufferGk_); + AscendC::PipeBarrier(); + AscendC::Mul(calcUbTensor, calcUbTensor, gkBroadcastUbTensor, rowsThisTile * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + } + + if (waitUpdateFromMte3) { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); } else { - AscendC::SetFlag(EVENT_ID0); - AscendC::WaitFlag(EVENT_ID0); - AscendC::DataCopy(finalStateThisSubBlock, hUpdateUbTensor, mActualThisSubBlock * nActual); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); } - } else { + AscendC::DataCopy(hUpdateUbTensor, hUpdateInputThisTile, rowsThisTile * nActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::Add(hUpdateUbTensor, calcUbTensor, hUpdateUbTensor, rowsThisTile * nActual); AscendC::PipeBarrier(); - AscendC::Cast(hOutputUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nActual); - AscendC::SetFlag(EVENT_ID0); - AscendC::WaitFlag(EVENT_ID0); - AscendC::DataCopy(hOutputThisSubBlock, hOutputUbTensor, mActualThisSubBlock * nActual); + + if constexpr(std::is_same::value) { + if (storeFinalState && isFinalState) { + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + CopyUbToGm(finalStateThisTile, hUpdateUbTensor, rowsThisTile, nActual, outputStride); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + waitUpdateFromMte3 = true; + } else { + AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + CopyUbToGm(hOutputThisTile, hUbTensor, rowsThisTile, nActual, outputStride); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + waitUpdateFromMte3 = false; + } + } else { + if (storeFinalState && isFinalState) { + AscendC::Cast(finalOutputUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + CopyUbToGm(finalStateThisTile, finalOutputUbTensor, rowsThisTile, nActual, outputStride); + } else { + AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + CopyUbToGm(hOutputThisTile, hUbTensor, rowsThisTile, nActual, outputStride); + } + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + waitUpdateFromMte3 = false; + } } + if constexpr (kGated) { + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + } + } private: - AscendC::LocalTensor calcUbTensor; + uint32_t pongBaseEvent = 4; - AscendC::LocalTensor hUbTensor; - AscendC::LocalTensor hUpdateUbTensor; + AscendC::LocalTensor calcUbTensor; - AscendC::LocalTensor hOutputUbTensor; - AscendC::LocalTensor finalOutputUbTensor; + AscendC::LocalTensor hUpdateUbTensor_ping; + AscendC::LocalTensor gkBroadcastUbTensor_ping; + AscendC::LocalTensor hUbTensor_ping; + AscendC::LocalTensor finalOutputUbTensor_ping; + AscendC::LocalTensor glastUbTensor_ping; - AscendC::LocalTensor glastUbTensor; + AscendC::LocalTensor hUpdateUbTensor_pong; + AscendC::LocalTensor gkBroadcastUbTensor_pong; + AscendC::LocalTensor hUbTensor_pong; + AscendC::LocalTensor finalOutputUbTensor_pong; + AscendC::LocalTensor glastUbTensor_pong; + AscendC::LocalTensor gkLastUbTensor_ping; + AscendC::LocalTensor gkLastUbTensor_pong; + AscendC::LocalTensor gkInputUbTensor_ping; + AscendC::LocalTensor gkInputUbTensor_pong; + AscendC::LocalTensor shareBufferGk_; }; } -#endif \ No newline at end of file +#endif diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp index 9f9349225d07..c3d889fa2d2b 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp @@ -24,15 +24,20 @@ template < class VOutputType_, class GInputType_, class UInputType_, - class WSInputType_ + class WSInputType_, + class FinalStateType_, + class KGatedTag > class BlockEpilogue < EpilogueAtlasGDNFwdHVnew, VOutputType_, GInputType_, UInputType_, - WSInputType_ + WSInputType_, + FinalStateType_, + KGatedTag > { + static constexpr bool kGated = KGatedTag::value; public: // Type aliases using DispatchPolicy = EpilogueAtlasGDNFwdHVnew; @@ -42,6 +47,7 @@ class BlockEpilogue < using GElementInput = typename GInputType_::Element; using UElementInput = typename UInputType_::Element; using WSElementInput = typename WSInputType_::Element; + using FinalStateElement = typename FinalStateType_::Element; CATLASS_DEVICE BlockEpilogue(Arch::Resource &resource) @@ -49,9 +55,13 @@ class BlockEpilogue < constexpr uint32_t CALC_BUF_OFFSET = 0; constexpr uint32_t PING_BUF_0_OFFSET = 32 * 1024; - constexpr uint32_t PING_BUF_1_OFFSET = 64 * 1024; + constexpr uint32_t PING_BUF_1_OFFSET = 48 * 1024; + constexpr uint32_t PING_BUF_2_OFFSET = 64 * 1024; + constexpr uint32_t PING_BUF_3_OFFSET = 80 * 1024; constexpr uint32_t PONG_BUF_0_OFFSET = 96 * 1024; - constexpr uint32_t PONG_BUF_1_OFFSET = 128 * 1024; + constexpr uint32_t PONG_BUF_1_OFFSET = 112 * 1024; + constexpr uint32_t PONG_BUF_2_OFFSET = 128 * 1024; + constexpr uint32_t PONG_BUF_3_OFFSET = 144 * 1024; constexpr uint32_t PING_G_BUF_OFFSET = 160 * 1024; constexpr uint32_t PONG_G_BUF_OFFSET = 161 * 1024; constexpr uint32_t PING_G_SUB_BUF_OFFSET = 162 * 1024; @@ -62,23 +72,21 @@ class BlockEpilogue < calcUbTensor = resource.ubBuf.template GetBufferByByte(CALC_BUF_OFFSET); - uUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_1_OFFSET); - uUbFloatTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_0_OFFSET); - wsUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_1_OFFSET); + uUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_2_OFFSET); + wsUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_0_OFFSET); gUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_BUF_OFFSET); gLastUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_SUB_BUF_OFFSET); - gInputUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_INPUT_BUF_OFFSET); - vNewOutputUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_1_OFFSET); - vNewDecayUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_0_OFFSET); + gInputUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_SUB_BUF_OFFSET); + vNewOutputUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_2_OFFSET); + vNewDecayUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_2_OFFSET); - uUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_1_OFFSET); - uUbFloatTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_0_OFFSET); - wsUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_1_OFFSET); + uUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_2_OFFSET); + wsUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_0_OFFSET); gUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_BUF_OFFSET); gLastUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_SUB_BUF_OFFSET); - gInputUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_INPUT_BUF_OFFSET); - vNewOutputUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_1_OFFSET); - vNewDecayUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_0_OFFSET); + gInputUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_SUB_BUF_OFFSET); + vNewOutputUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_2_OFFSET); + vNewDecayUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_2_OFFSET); shareBuffer_ = resource.ubBuf.template GetBufferByByte(SHARE_BUF_OFFSET); @@ -87,98 +95,86 @@ class BlockEpilogue < CATLASS_DEVICE ~BlockEpilogue() {} + template CATLASS_DEVICE - void operator()( - AscendC::GlobalTensor vnewOutput, - AscendC::GlobalTensor vnewdecayOutput, - AscendC::GlobalTensor gInput, - AscendC::GlobalTensor uInput, - AscendC::GlobalTensor wsInput, - uint32_t chunkSize, - uint32_t kHeadDim, - uint32_t vHeadDim, - Arch::CrossCoreFlag cube1Done - // const LayoutOutput &layoutOutput, - // const LayoutInput &LayoutInput - ) + void CopyGmToUb( + AscendC::LocalTensor dst, + AscendC::GlobalTensor src, + uint32_t rows, + uint32_t cols, + uint32_t srcStride) { - uint32_t mActual = chunkSize; - uint32_t nkActual = kHeadDim; - uint32_t nvActual = vHeadDim; - - uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); - uint32_t subBlockNum = AscendC::GetSubBlockNum(); - uint32_t mActualPerSubBlock = CeilDiv(mActual, subBlockNum); - uint32_t mActualThisSubBlock = (subBlockIdx == 0) ? mActualPerSubBlock : (mActual - mActualPerSubBlock); - uint32_t mOffset = subBlockIdx * mActualPerSubBlock; - uint32_t nOffset = 0; - // 当前场景内部一定连续 - // k [B, H, T, D] - // g [B, H, T] - // 在外部offset的基础上进一步offset - // 当前asset kdim == vHeadDim - int64_t offsetK = mOffset * nvActual + nOffset; - int64_t offsetD = 0; // 因为要用最后一个数减去之前所有,所以全部读入 - - uint32_t gbrcStart, gbrcRealStart, gbrcReptime, gbrcEffStart, gbrcEffEnd; - if(subBlockIdx==0) - { - gbrcStart = 0; - gbrcRealStart = 0; - gbrcReptime = (mActualThisSubBlock + 8 - 1) / 8; - - } - else - { - gbrcStart = mActualPerSubBlock; - gbrcRealStart = gbrcStart & ~15; - gbrcReptime = (mActual - gbrcRealStart + 8 - 1) / 8; + if (cols == srcStride) { + AscendC::DataCopy(dst, src, rows * cols); + return; } - gbrcEffStart = gbrcStart-gbrcRealStart; - gbrcEffEnd = gbrcEffStart + mActualThisSubBlock; + AscendC::DataCopyExtParams copyParams{ + static_cast(rows), + static_cast(cols * sizeof(Element)), + static_cast((srcStride - cols) * sizeof(Element)), + 0, + 0}; + AscendC::DataCopyPadExtParams padParams{false, 0, 0, 0}; + AscendC::DataCopyPad(dst, src, copyParams, padParams); + } - AscendC::ResetMask(); + template + CATLASS_DEVICE + void CopyUbToGm( + AscendC::GlobalTensor dst, + AscendC::LocalTensor src, + uint32_t rows, + uint32_t cols, + uint32_t dstStride) + { + if (cols == dstStride) { + AscendC::DataCopy(dst, src, rows * cols); + return; + } + AscendC::DataCopyExtParams copyParams{ + static_cast(rows), + static_cast(cols * sizeof(Element)), + 0, + static_cast((dstStride - cols) * sizeof(Element)), + 0}; + AscendC::DataCopyPad(dst, src, copyParams); + } - AscendC::GlobalTensor vnewOutputThisSubBlock = vnewOutput[offsetK]; - AscendC::GlobalTensor vnewdecayOutputThisSubBlock = vnewdecayOutput[offsetK]; - AscendC::GlobalTensor gInputThisSubBlock = gInput; - AscendC::GlobalTensor uInputThisSubBlock = uInput[offsetK]; - AscendC::GlobalTensor wsInputThisSubBlock = wsInput[offsetK]; - - pingpongFlag = isFirst ? 0 : 4; - AscendC::LocalTensor uUbTensor = isFirst ? uUbTensor_ping : uUbTensor_pong; - AscendC::LocalTensor uUbFloatTensor = isFirst ? uUbFloatTensor_ping : uUbFloatTensor_pong; - AscendC::LocalTensor wsUbTensor = isFirst ? wsUbTensor_ping : wsUbTensor_pong; - AscendC::LocalTensor gUbTensor = isFirst ? gUbTensor_ping : gUbTensor_pong; - AscendC::LocalTensor gLastUbTensor = isFirst ? gLastUbTensor_ping : gLastUbTensor_pong; - AscendC::LocalTensor gInputUbTensor = isFirst ? gInputUbTensor_ping : gInputUbTensor_pong; - AscendC::LocalTensor vNewOutputUbTensor = isFirst ? vNewOutputUbTensor_ping : vNewOutputUbTensor_pong; - AscendC::LocalTensor vNewDecayUbTensor = isFirst ? vNewDecayUbTensor_ping : vNewDecayUbTensor_pong; - - AscendC::SetFlag(EVENT_ID0 + pingpongFlag); - AscendC::SetFlag(EVENT_ID2 + pingpongFlag); - - AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + CATLASS_DEVICE + void PrepareG( + AscendC::LocalTensor gUbTensor, + AscendC::LocalTensor gLastUbTensor, + AscendC::LocalTensor gInputUbTensor, + AscendC::GlobalTensor gInputThisSubBlock, + uint32_t mActual, + uint32_t pingpongFlag) + { + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + if (mActual == 1) { + AscendC::Duplicate(gUbTensor, 1.0f, 1); + AscendC::PipeBarrier(); + return; + } if constexpr(std::is_same::value) { AscendC::DataCopyParams gUbParams{1, (uint16_t)(mActual * sizeof(float)), 0, 0}; AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0}; AscendC::DataCopyPad(gUbTensor, gInputThisSubBlock, gUbParams, gUbPadParams); } else { - AscendC::DataCopyParams gUbParams{1, (uint16_t)(mActual * sizeof(half)), 0, 0}; + AscendC::DataCopyParams gUbParams{1, (uint16_t)(mActual * sizeof(GElementInput)), 0, 0}; AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0}; AscendC::DataCopyPad(gInputUbTensor, gInputThisSubBlock, gUbParams, gUbPadParams); } - AscendC::SetFlag(EVENT_ID2 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); if constexpr(!std::is_same::value) { AscendC::Cast(gUbTensor, gInputUbTensor, AscendC::RoundMode::CAST_NONE, mActual); } - AscendC::SetFlag(EVENT_ID2 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); - float inputVal = gUbTensor.GetValue(mActual-1); - AscendC::SetFlag(EVENT_ID2 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + float inputVal = gUbTensor.GetValue(mActual - 1); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); AscendC::PipeBarrier(); AscendC::Duplicate(gLastUbTensor, inputVal, mActual); @@ -186,64 +182,243 @@ class BlockEpilogue < AscendC::Sub(gUbTensor, gLastUbTensor, gUbTensor, mActual); AscendC::PipeBarrier(); - AscendC::Exp(gUbTensor, gUbTensor, mActual); AscendC::PipeBarrier(); + } - uint32_t dstShape_[2] = {gbrcReptime*8, nvActual}; - uint32_t srcShape_[2] = {gbrcReptime*8, 1}; - AscendC::Broadcast(calcUbTensor, gUbTensor[gbrcRealStart], dstShape_, srcShape_, shareBuffer_); - AscendC::PipeBarrier(); + CATLASS_DEVICE + void operator()( + AscendC::GlobalTensor vnewOutput, + AscendC::GlobalTensor vnewdecayOutput, + AscendC::GlobalTensor gInput, + AscendC::GlobalTensor uInput, + AscendC::GlobalTensor wsInput, + AscendC::GlobalTensor gkInput, + AscendC::GlobalTensor kInput, + AscendC::GlobalTensor kDecayWorkspace, + uint32_t chunkSize, + uint32_t kHeadDim, + uint32_t vBlockDim, + uint32_t vHeadDim, + Arch::CrossCoreFlag cube1Done, + Arch::CrossCoreFlag vec1Done, + bool isInitialState, + bool isFinalState, + bool storeFinalState, + bool waitWsFromMte3, + bool isPing + ) + { + static constexpr uint32_t ROW_TILE = 16; + uint32_t mActual = chunkSize; + uint32_t nvActual = vBlockDim; + uint32_t nkActual = kHeadDim; + uint32_t inputStride = vHeadDim; - Arch::CrossCoreWaitFlag(cube1Done); + uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); + uint32_t subBlockNum = AscendC::GetSubBlockNum(); + uint32_t rowsPerSubBlock = CeilDiv(mActual, subBlockNum); + uint32_t rowBegin = subBlockIdx * rowsPerSubBlock; + uint32_t rowEnd = rowBegin + rowsPerSubBlock; + if (rowEnd > mActual) { + rowEnd = mActual; + } + uint32_t pingpongFlag = isPing ? 0 : pongBaseEvent; + if (rowBegin >= mActual) { + Arch::CrossCoreWaitFlag(cube1Done); + // Even an empty vector subblock owns this stream's workspace + // event. Normalize it for V2 before the stream can be reused. + if (waitWsFromMte3) { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + return; + } + AscendC::ResetMask(); - AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); - AscendC::SetFlag(EVENT_ID0 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); - - AscendC::DataCopy(uUbTensor, uInputThisSubBlock, mActualThisSubBlock * nvActual); - AscendC::SetFlag(EVENT_ID0 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); - AscendC::Cast(uUbFloatTensor, uUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nvActual); - AscendC::PipeBarrier(); + AscendC::GlobalTensor gInputThisSubBlock = gInput; - AscendC::SetFlag(EVENT_ID1 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); - AscendC::DataCopy(wsUbTensor, wsInputThisSubBlock, mActualThisSubBlock * nvActual); - AscendC::SetFlag(EVENT_ID1 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); - - AscendC::Sub(uUbFloatTensor, uUbFloatTensor, wsUbTensor, mActualThisSubBlock * nvActual); - AscendC::PipeBarrier(); - AscendC::Cast(vNewOutputUbTensor, uUbFloatTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nvActual); - AscendC::SetFlag(EVENT_ID0 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); - AscendC::DataCopy(vnewOutputThisSubBlock, vNewOutputUbTensor, mActualThisSubBlock * nvActual); + AscendC::LocalTensor uUbTensor = isPing ? uUbTensor_ping : uUbTensor_pong; + AscendC::LocalTensor wsUbTensor = isPing ? wsUbTensor_ping : wsUbTensor_pong; + AscendC::LocalTensor gUbTensor = isPing ? gUbTensor_ping : gUbTensor_pong; + AscendC::LocalTensor gLastUbTensor = isPing ? gLastUbTensor_ping : gLastUbTensor_pong; + AscendC::LocalTensor gInputUbTensor = isPing ? gInputUbTensor_ping : gInputUbTensor_pong; + AscendC::LocalTensor vNewOutputUbTensor = isPing ? vNewOutputUbTensor_ping : vNewOutputUbTensor_pong; + AscendC::LocalTensor vNewDecayUbTensor = isPing ? vNewDecayUbTensor_ping : vNewDecayUbTensor_pong; + + if (nvActual <= 128 && nvActual == inputStride) { + uint32_t mActualThisSubBlock = rowEnd - rowBegin; + uint32_t gbrcRealStart = rowBegin & ~7; + uint32_t gbrcEffStart = rowBegin - gbrcRealStart; + uint32_t gbrcRealProcess = gbrcEffStart + mActualThisSubBlock; + uint32_t dstShape_[2] = {gbrcRealProcess, nvActual}; + uint32_t srcShape_[2] = {gbrcRealProcess, 1}; + + AscendC::GlobalTensor vnewOutputThisSubBlock = vnewOutput[rowBegin * inputStride]; + AscendC::GlobalTensor vnewdecayOutputThisSubBlock = vnewdecayOutput[rowBegin * nvActual]; + AscendC::GlobalTensor uInputThisSubBlock = uInput[rowBegin * inputStride]; + AscendC::GlobalTensor wsInputThisSubBlock = wsInput[rowBegin * nvActual]; + + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(uUbTensor, uInputThisSubBlock, mActualThisSubBlock * nvActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::Cast(calcUbTensor, uUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nvActual); + + PrepareG(gUbTensor, gLastUbTensor, gInputUbTensor, gInputThisSubBlock, mActual, pingpongFlag); + Arch::CrossCoreWaitFlag(cube1Done); + + if (waitWsFromMte3) { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } + AscendC::DataCopy(wsUbTensor, wsInputThisSubBlock, mActualThisSubBlock * nvActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + + AscendC::Sub(wsUbTensor, calcUbTensor, wsUbTensor, mActualThisSubBlock * nvActual); + AscendC::PipeBarrier(); + + AscendC::Broadcast(calcUbTensor, gUbTensor[gbrcRealStart], dstShape_, srcShape_, shareBuffer_); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + + AscendC::Mul(calcUbTensor[gbrcEffStart * nvActual], wsUbTensor, calcUbTensor[gbrcEffStart * nvActual], mActualThisSubBlock * nvActual); + AscendC::PipeBarrier(); + + AscendC::Cast(vNewDecayUbTensor, calcUbTensor[gbrcEffStart * nvActual], AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nvActual); + AscendC::PipeBarrier(); + + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(vnewdecayOutputThisSubBlock, vNewDecayUbTensor, mActualThisSubBlock * nvActual); + if constexpr (!kGated) { + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + } + + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::Cast(vNewOutputUbTensor, wsUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nvActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(vnewOutputThisSubBlock, vNewOutputUbTensor, mActualThisSubBlock * nvActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + + if constexpr (kGated) { + AscendC::GlobalTensor kInputThisSubBlock = kInput[rowBegin * nkActual]; + AscendC::GlobalTensor kDecayWorkspaceThisSubBlock = kDecayWorkspace[rowBegin * nkActual]; + // KDA passes kg = k * exp2(g_last - gk). Keep that decay exactly once. + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(vNewOutputUbTensor, kInputThisSubBlock, mActualThisSubBlock * nkActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(kDecayWorkspaceThisSubBlock, vNewOutputUbTensor, + mActualThisSubBlock * nkActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + } + return; + } - AscendC::PipeBarrier(); - AscendC::Mul(calcUbTensor[gbrcEffStart*nvActual], uUbFloatTensor, calcUbTensor[gbrcEffStart*nvActual], mActualThisSubBlock * nvActual); - AscendC::PipeBarrier(); - AscendC::Cast(vNewDecayUbTensor, calcUbTensor[gbrcEffStart*nvActual], AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nvActual); - - AscendC::SetFlag(EVENT_ID0 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); - AscendC::DataCopy(vnewdecayOutputThisSubBlock, vNewDecayUbTensor, mActualThisSubBlock * nvActual); - - if (isFirst) { - AscendC::PipeBarrier(); + PrepareG(gUbTensor, gLastUbTensor, gInputUbTensor, gInputThisSubBlock, mActual, pingpongFlag); + Arch::CrossCoreWaitFlag(cube1Done); + + bool waitWsThisTileFromMte3 = waitWsFromMte3; + for (uint32_t rowStart = rowBegin; rowStart < rowEnd;) { + uint32_t alignExtra = rowStart & 7; + uint32_t maxRowsThisTile = ROW_TILE - alignExtra; + uint32_t rowsThisTile = rowEnd - rowStart; + if (rowsThisTile > maxRowsThisTile) { + rowsThisTile = maxRowsThisTile; + } + uint32_t gbrcRealStart = rowStart & ~7; + uint32_t gbrcRealProcess = alignExtra + rowsThisTile; + uint32_t dstShape_[2] = {gbrcRealProcess, nvActual}; + uint32_t srcShape_[2] = {gbrcRealProcess, 1}; + + AscendC::GlobalTensor vnewOutputThisTile = vnewOutput[rowStart * inputStride]; + AscendC::GlobalTensor vnewdecayOutputThisTile = vnewdecayOutput[rowStart * nvActual]; + AscendC::GlobalTensor uInputThisTile = uInput[rowStart * inputStride]; + AscendC::GlobalTensor wsInputThisTile = wsInput[rowStart * nvActual]; + + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + CopyGmToUb(uUbTensor, uInputThisTile, rowsThisTile, nvActual, inputStride); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::Cast(calcUbTensor, uUbTensor, AscendC::RoundMode::CAST_NONE, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + + if (waitWsThisTileFromMte3) { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } + AscendC::DataCopy(wsUbTensor, wsInputThisTile, rowsThisTile * nvActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + waitWsThisTileFromMte3 = false; + + AscendC::Sub(wsUbTensor, calcUbTensor, wsUbTensor, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + + AscendC::Broadcast(calcUbTensor, gUbTensor[gbrcRealStart], dstShape_, srcShape_, shareBuffer_); + AscendC::PipeBarrier(); + AscendC::Mul(calcUbTensor[alignExtra * nvActual], wsUbTensor, calcUbTensor[alignExtra * nvActual], rowsThisTile * nvActual); + AscendC::PipeBarrier(); + + AscendC::Cast(vNewDecayUbTensor, calcUbTensor[alignExtra * nvActual], AscendC::RoundMode::CAST_RINT, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(vnewdecayOutputThisTile, vNewDecayUbTensor, rowsThisTile * nvActual); + + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + if (rowStart + rowsThisTile >= rowEnd) { + if constexpr (!kGated) { + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + } + } + AscendC::Cast(vNewOutputUbTensor, wsUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + CopyUbToGm(vnewOutputThisTile, vNewOutputUbTensor, rowsThisTile, nvActual, inputStride); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + rowStart += rowsThisTile; } - isFirst = false; + if constexpr (kGated) { + uint32_t mActualThisSubBlock = rowEnd - rowBegin; + AscendC::GlobalTensor kInputThisSubBlock = kInput[rowBegin * nkActual]; + AscendC::GlobalTensor kDecayWorkspaceThisSubBlock = kDecayWorkspace[rowBegin * nkActual]; + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + CopyGmToUb(vNewOutputUbTensor, kInputThisSubBlock, mActualThisSubBlock, nkActual, nkActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + CopyUbToGm(kDecayWorkspaceThisSubBlock, vNewOutputUbTensor, mActualThisSubBlock, nkActual, nkActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + } + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); } private: - uint32_t pingpongFlag = 0; - bool isFirst = true; + uint32_t pongBaseEvent = 4; AscendC::LocalTensor calcUbTensor; AscendC::LocalTensor uUbTensor_ping; - AscendC::LocalTensor uUbFloatTensor_ping; AscendC::LocalTensor wsUbTensor_ping; AscendC::LocalTensor gUbTensor_ping; AscendC::LocalTensor gLastUbTensor_ping; @@ -252,7 +427,6 @@ class BlockEpilogue < AscendC::LocalTensor vNewDecayUbTensor_ping; AscendC::LocalTensor uUbTensor_pong; - AscendC::LocalTensor uUbFloatTensor_pong; AscendC::LocalTensor wsUbTensor_pong; AscendC::LocalTensor gUbTensor_pong; AscendC::LocalTensor gLastUbTensor_pong; @@ -265,4 +439,4 @@ class BlockEpilogue < }; } -#endif \ No newline at end of file +#endif diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/epilogue/gdn_fwd_h_epilogue_policies.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/epilogue/gdn_fwd_h_epilogue_policies.hpp index eb5cc74c0fa2..d5287d55defd 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/epilogue/gdn_fwd_h_epilogue_policies.hpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/epilogue/gdn_fwd_h_epilogue_policies.hpp @@ -15,19 +15,11 @@ namespace Catlass::Epilogue { struct EpilogueAtlasGDNFwdHVnew { -#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 - using ArchTag = Arch::Ascend950; -#else using ArchTag = Arch::AtlasA2; -#endif }; struct EpilogueAtlasGDNFwdHUpdate { -#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 - using ArchTag = Arch::Ascend950; -#else using ArchTag = Arch::AtlasA2; -#endif }; } // namespace Catlass::Epilogue diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/gemm/block/block_scheduler_gdn_fwd_h.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/gemm/block/block_scheduler_gdn_fwd_h.hpp index a83b08ef26f9..e6e21fd1b355 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/gemm/block/block_scheduler_gdn_fwd_h.hpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/gemm/block/block_scheduler_gdn_fwd_h.hpp @@ -15,6 +15,12 @@ using namespace Catlass; // constexpr uint32_t PING_PONG_STAGES = 1; constexpr uint32_t PING_PONG_STAGES = 2; +constexpr uint32_t BYTE_SIZE_16_BIT = 2; +constexpr uint32_t BYTES_PER_C0 = 32; +constexpr uint32_t BYTE_SIZE_PER_REPEAT = 256; +constexpr uint32_t SIZE_16_NUM_PER_C0 = BYTES_PER_C0 / BYTE_SIZE_16_BIT; +constexpr uint32_t FLOAT_NUM_PER_REPEAT = BYTE_SIZE_PER_REPEAT / sizeof(float); +constexpr uint32_t NZ_BLOCK_SIZE = 16; template CATLASS_DEVICE T AlignUp(T a, T b) { @@ -40,14 +46,18 @@ struct GDNFwdHOffsets { uint32_t wkOffset; uint32_t wOffset; uint32_t gOffset; + uint32_t gkOffset; uint32_t hWorkOffset; uint32_t vWorkOffset; + uint32_t kDecayWorkOffset; + uint32_t vBlockOffset; + uint32_t vBlockDim; uint32_t initialStateOffset; uint32_t finalStateOffset; bool isInitialState; bool isFinalState; uint32_t blockTokens; - bool isDummyHead; + uint32_t streamId; // for debug uint32_t batchIdx; uint32_t headIdx; @@ -55,6 +65,28 @@ struct GDNFwdHOffsets { }; +struct GDNFwdHStream { + uint32_t batchIdx; + uint32_t chunkIdx{0}; + uint32_t vHeadIdx; + uint32_t kHeadIdx; + uint32_t shapeBatchIdx; + uint32_t tokenBatchIdx; + + uint32_t chunkOffset; + uint32_t tokenOffset; + uint32_t batchChunks{0}; + uint32_t batchTokens; + uint32_t nextTaskIdx{0}; + bool active{false}; + + GDNFwdHOffsets offset; +}; + +struct GDNFwdHRunningQ { + GDNFwdHStream streams[PING_PONG_STAGES]; +}; + struct BlockSchedulerGdnFwdH { uint32_t batch; uint32_t seqlen; @@ -63,7 +95,6 @@ struct BlockSchedulerGdnFwdH { uint32_t kHeadDim; uint32_t vHeadDim; uint32_t chunkSize; - uint32_t initalStateStride0; uint32_t vBlockSize{128}; uint32_t isVariedLen; uint32_t shapeBatch; @@ -74,48 +105,26 @@ struct BlockSchedulerGdnFwdH { uint32_t numChunksWorkspaceOffset; uint32_t taskIdx; - uint32_t taskLoops; + uint32_t taskStride; uint32_t cubeCoreIdx; uint32_t cubeCoreNum; - uint32_t vLoops; uint32_t taskNum; uint32_t headGroups; uint32_t totalChunks; uint32_t totalTokens; - uint32_t headInnerLoop; - uint32_t iterId {0}; - bool hasDummyHead; - bool isRunning; - bool processNewTask {true}; - bool firstLoop {true}; - bool lastLoop {false}; - GDNFwdHOffsets offsets[PING_PONG_STAGES]; - int32_t currStage{PING_PONG_STAGES - 1}; + GDNFwdHRunningQ runningQ; - uint32_t vIdx; - uint32_t batchIdx; - uint32_t baseHeadIdx; - uint32_t chunkIdx; - uint32_t headInnerIdx; - uint32_t vHeadIdx; - uint32_t kHeadIdx; - uint32_t shapeBatchIdx; - uint32_t tokenBatchIdx; - - uint32_t chunkOffset; - uint32_t tokenOffset; - uint32_t batchChunks; - uint32_t batchTokens; + bool isRunning; AscendC::GlobalTensor gmSeqlen; AscendC::GlobalTensor gmNumSeq; AscendC::GlobalTensor gmNumChunks; - Arch::CrossCoreFlag cube1Done{0}; - Arch::CrossCoreFlag vec1Done{1}; - Arch::CrossCoreFlag cube2Done{2}; - Arch::CrossCoreFlag vec2Done{3}; + Arch::CrossCoreFlag cube1Done[PING_PONG_STAGES] = {0, 1}; + Arch::CrossCoreFlag vec1Done[PING_PONG_STAGES] = {2, 3}; + Arch::CrossCoreFlag cube2Done[PING_PONG_STAGES] = {4, 5}; + Arch::CrossCoreFlag vec2Done[PING_PONG_STAGES] = {6, 7}; CATLASS_DEVICE BlockSchedulerGdnFwdH() {} @@ -131,7 +140,6 @@ struct BlockSchedulerGdnFwdH { kHeadDim = gdnFwdHTilingData->kHeadDim; vHeadDim = gdnFwdHTilingData->vHeadDim; chunkSize = gdnFwdHTilingData->chunkSize; - initalStateStride0 = gdnFwdHTilingData->initalStateStride0; isVariedLen = gdnFwdHTilingData->isVariedLen; shapeBatch = gdnFwdHTilingData->shapeBatch; tokenBatch = gdnFwdHTilingData->tokenBatch; @@ -171,91 +179,134 @@ struct BlockSchedulerGdnFwdH { cubeCoreIdx = coreIdx; cubeCoreNum = coreNum; - vLoops = vHeadDim / vBlockSize; - taskNum = vLoops * batch * vNumHead; + vBlockSize = vHeadDim; + taskNum = batch * vNumHead; headGroups = vNumHead / kNumHead; - hasDummyHead = (taskNum % (PING_PONG_STAGES * cubeCoreNum) <= cubeCoreNum) && (taskNum % (PING_PONG_STAGES * cubeCoreNum) > 0); - taskLoops = (taskNum + cubeCoreNum * PING_PONG_STAGES - 1) / (cubeCoreNum * PING_PONG_STAGES); - headInnerLoop = taskNum > cubeCoreNum ? PING_PONG_STAGES : 1; - taskIdx = cubeCoreIdx * headInnerLoop; - isRunning = taskIdx < taskNum; + uint32_t maxTaskCntPerLoop = taskNum > cubeCoreNum ? PING_PONG_STAGES : 1; + taskStride = cubeCoreNum * maxTaskCntPerLoop; + for (uint32_t streamId = 0; streamId < PING_PONG_STAGES; ++streamId) { + auto& stream = runningQ.streams[streamId]; + stream.nextTaskIdx = (streamId < maxTaskCntPerLoop) ? (cubeCoreIdx * maxTaskCntPerLoop + streamId) : taskNum; + stream.chunkIdx = 0; + stream.batchChunks = 0; + stream.active = false; + } + isRunning = cubeCoreIdx * maxTaskCntPerLoop < taskNum; } + CATLASS_DEVICE - void InitTask() { - iterId++; - currStage = (currStage + 1) % PING_PONG_STAGES; - if (processNewTask) { - if (taskIdx >= taskNum) { - lastLoop = true; - isRunning = false; - return; - } - vIdx = taskIdx / (batch * vNumHead); - batchIdx = (taskIdx - vIdx * batch * vNumHead) / vNumHead; - baseHeadIdx = taskIdx % vNumHead; - shapeBatchIdx = isVariedLen ? 0 : batchIdx; - tokenBatchIdx = isVariedLen ? batchIdx : 0; - chunkOffset = isVariedLen ? gmNumChunks.GetValue(tokenBatchIdx) : 0; - batchChunks = isVariedLen ? (gmNumChunks.GetValue(tokenBatchIdx + 1) - chunkOffset) : totalChunks; - tokenOffset = isVariedLen ? gmNumSeq.GetValue(tokenBatchIdx) : 0; - batchTokens = isVariedLen ? (gmNumSeq.GetValue(tokenBatchIdx + 1) - tokenOffset) : totalTokens; - chunkIdx = 0; - headInnerIdx = 0; - } else { - chunkIdx = headInnerIdx == PING_PONG_STAGES - 1 ? chunkIdx + 1 : chunkIdx; - headInnerIdx = (headInnerIdx + 1) % PING_PONG_STAGES; + void InitNewStream(GDNFwdHStream& newStream) { + newStream.batchIdx = taskIdx / vNumHead; + newStream.vHeadIdx = taskIdx % vNumHead; + newStream.kHeadIdx = newStream.vHeadIdx / headGroups; + newStream.shapeBatchIdx = isVariedLen ? 0 : newStream.batchIdx; + newStream.tokenBatchIdx = isVariedLen ? newStream.batchIdx : 0; + newStream.chunkOffset = isVariedLen ? gmNumChunks.GetValue(newStream.tokenBatchIdx) : 0; + newStream.batchChunks = isVariedLen ? (gmNumChunks.GetValue(newStream.tokenBatchIdx + 1) - newStream.chunkOffset) : totalChunks; + newStream.tokenOffset = isVariedLen ? gmNumSeq.GetValue(newStream.tokenBatchIdx) : 0; + newStream.batchTokens = isVariedLen ? (gmNumSeq.GetValue(newStream.tokenBatchIdx + 1) - newStream.tokenOffset) : totalTokens; + newStream.chunkIdx = 0; + } + + CATLASS_DEVICE + void AssignNextStream(uint32_t streamId) { + auto& stream = runningQ.streams[streamId]; + taskIdx = stream.nextTaskIdx; + if (taskIdx >= taskNum) { + stream.active = false; + stream.batchChunks = 0; + return; + } + + stream.nextTaskIdx += taskStride; + InitNewStream(stream); + stream.active = stream.batchChunks > 0; + if (stream.active) { + UpdateTask(streamId); + } + } + + CATLASS_DEVICE + void UpdateTask(uint32_t streamId) { + auto& stream = runningQ.streams[streamId]; + auto& offset = stream.offset; + + offset.isInitialState = stream.chunkIdx == 0; + offset.isFinalState = stream.chunkIdx == (stream.batchChunks - 1); + uint32_t vBlockOffset = 0; + uint32_t vBlockDim = vBlockSize; + offset.initialStateOffset = (stream.batchIdx * vNumHead + stream.vHeadIdx) * kHeadDim * vHeadDim + vBlockOffset; + offset.finalStateOffset = (stream.batchIdx * vNumHead + stream.vHeadIdx) * kHeadDim * vHeadDim + vBlockOffset; + offset.hSrcOffset = (stream.shapeBatchIdx * vNumHead * totalChunks + stream.vHeadIdx * totalChunks + stream.chunkOffset + stream.chunkIdx) * kHeadDim * vHeadDim + vBlockOffset; + offset.hDstOffset = offset.hSrcOffset + kHeadDim * vHeadDim; + if (storeFinalState && offset.isFinalState) { + offset.hDstOffset = offset.hSrcOffset; } - - vHeadIdx = baseHeadIdx + headInnerIdx; - kHeadIdx = vHeadIdx / headGroups; - offsets[currStage].isInitialState = chunkIdx == 0; - offsets[currStage].isFinalState = chunkIdx == (batchChunks - 1); - offsets[currStage].initialStateOffset = (batchIdx * vNumHead + vHeadIdx) * kHeadDim * initalStateStride0; - offsets[currStage].finalStateOffset = (batchIdx * vNumHead + vHeadIdx) * kHeadDim * vHeadDim; - offsets[currStage].hSrcOffset = (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset + chunkIdx) * kHeadDim * vHeadDim; - offsets[currStage].hDstOffset = offsets[currStage].hSrcOffset + kHeadDim * vHeadDim; - offsets[currStage].uvOffset = (shapeBatchIdx * vNumHead * totalTokens + vHeadIdx * totalTokens + tokenOffset + chunkIdx * chunkSize) * vHeadDim; - offsets[currStage].wkOffset = (shapeBatchIdx * kNumHead * totalTokens + kHeadIdx * totalTokens + tokenOffset + chunkIdx * chunkSize) * kHeadDim; - offsets[currStage].wOffset = (shapeBatchIdx * vNumHead * totalTokens + vHeadIdx * totalTokens + tokenOffset + chunkIdx * chunkSize) * kHeadDim; - offsets[currStage].gOffset = shapeBatchIdx * vNumHead * totalTokens + vHeadIdx * totalTokens + tokenOffset + chunkIdx * chunkSize; - offsets[currStage].hWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + currStage) * kHeadDim * vHeadDim; - offsets[currStage].vWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + currStage) * chunkSize * vHeadDim; - offsets[currStage].blockTokens = offsets[currStage].isFinalState ? (batchTokens - chunkIdx * chunkSize) : chunkSize; - offsets[currStage].isDummyHead = headInnerLoop < PING_PONG_STAGES && headInnerIdx >= headInnerLoop; - offsets[currStage].batchIdx = batchIdx; - offsets[currStage].headIdx = vHeadIdx; - offsets[currStage].chunkIdx = chunkIdx; - - processNewTask = chunkIdx == batchChunks - 1 && headInnerIdx == PING_PONG_STAGES - 1; - if (processNewTask) { - uint32_t currLoopIdx = taskIdx / (PING_PONG_STAGES * cubeCoreNum); - headInnerLoop = ((currLoopIdx + 2 == taskLoops) && hasDummyHead) ? 1 : PING_PONG_STAGES; - taskIdx = (currLoopIdx + 1) * PING_PONG_STAGES * cubeCoreNum + headInnerLoop * cubeCoreIdx; + offset.uvOffset = (stream.shapeBatchIdx * vNumHead * totalTokens + stream.vHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize) * vHeadDim + vBlockOffset; + offset.wkOffset = (stream.shapeBatchIdx * kNumHead * totalTokens + stream.kHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize) * kHeadDim; + offset.wOffset = (stream.shapeBatchIdx * vNumHead * totalTokens + stream.vHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize) * kHeadDim; + offset.gOffset = stream.shapeBatchIdx * vNumHead * totalTokens + stream.vHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize; + offset.gkOffset = (stream.shapeBatchIdx * vNumHead * totalTokens + stream.vHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize) * kHeadDim; + offset.hWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + streamId) * kHeadDim * vBlockSize; + offset.vWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + streamId) * chunkSize * vBlockSize; + offset.kDecayWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + streamId) * chunkSize * kHeadDim; + offset.vBlockOffset = vBlockOffset; + offset.vBlockDim = vBlockDim; + offset.blockTokens = offset.isFinalState ? (stream.batchTokens - stream.chunkIdx * chunkSize) : chunkSize; + offset.streamId = streamId; + offset.batchIdx = stream.batchIdx; + offset.headIdx = stream.vHeadIdx; + offset.chunkIdx = stream.chunkIdx; + } + + CATLASS_DEVICE + void InitTasks() { + isRunning = false; + for (uint32_t streamId = 0; streamId < PING_PONG_STAGES; ++streamId) { + auto& stream = runningQ.streams[streamId]; + if (stream.active) { + stream.chunkIdx += 1; + if (stream.chunkIdx >= stream.batchChunks) { + stream.active = false; + stream.batchChunks = 0; + } + } + if (!stream.active) { + AssignNextStream(streamId); + } else { + UpdateTask(streamId); + } + if (stream.active) { + isRunning = true; + } } } CATLASS_DEVICE - GDNFwdHOffsets& GetStage1Offsets() { - return offsets[currStage]; + const GDNFwdHStream& GetStream(uint32_t i) const { + return runningQ.streams[i]; + } + + CATLASS_DEVICE + uint32_t GetStreamId(uint32_t i) const { + return i; } - + CATLASS_DEVICE - bool NeedProcessStage1() { - GDNFwdHOffsets& stage1Offsets = GetStage1Offsets(); - return !(lastLoop || stage1Offsets.isDummyHead); + const GDNFwdHOffsets& GetCurTaskOffsets(const GDNFwdHStream& stream) const { + return stream.offset; } CATLASS_DEVICE - GDNFwdHOffsets& GetStage2Offsets() { - return offsets[(currStage - 1) % PING_PONG_STAGES]; + bool StreamIsDone(const GDNFwdHStream& stream) const { + return !stream.active; } CATLASS_DEVICE - bool NeedProcessStage2() { - GDNFwdHOffsets& stage2Offsets = GetStage2Offsets(); - return !(iterId == 1 || (!storeFinalState && stage2Offsets.isFinalState) || stage2Offsets.isDummyHead); + bool NeedProcessStage2(const GDNFwdHStream& stream) { + return storeFinalState || !stream.offset.isFinalState; } }; @@ -276,11 +327,14 @@ struct BlockSchedulerGdnFwdHVec : public BlockSchedulerGdnFwdH { CATLASS_DEVICE void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR tiling, GM_ADDR user) { - BlockSchedulerGdnFwdH::Init(cu_seqlens, chunk_indices, tiling, user, AscendC::GetBlockIdx() / AscendC::GetSubBlockNum(), AscendC::GetBlockNum()); + BlockSchedulerGdnFwdH::Init( + cu_seqlens, chunk_indices, tiling, user, + AscendC::GetBlockIdx() / AscendC::GetSubBlockNum(), + AscendC::GetBlockNum()); } }; } // namespace Catlass::Gemm::Block -#endif // CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP \ No newline at end of file +#endif // CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/gemm/kernel/gdn_fwd_h_kernel.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/gemm/kernel/gdn_fwd_h_kernel.hpp index 4a97a98882e4..6924d5d91d46 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/gemm/kernel/gdn_fwd_h_kernel.hpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch22/gemm/kernel/gdn_fwd_h_kernel.hpp @@ -7,49 +7,6 @@ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. */ -#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 -#define CATLASS_ARCH 3510 - -#include "catlass/arch/arch.hpp" -#include "catlass/arch/cross_core_sync.hpp" -#include "catlass/arch/resource.hpp" -#include "catlass/catlass.hpp" -#include "catlass/debug.hpp" -#include "catlass/epilogue/block/block_epilogue.hpp" -#include "../../epilogue/block/block_epilogue_gdn_fwdh_update.hpp" -#include "../../epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp" -#include "catlass/gemm/block/block_mmad.hpp" -#include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp" -#include "catlass/gemm/block/block_swizzle.hpp" -#include "../block/block_scheduler_gdn_fwd_h.hpp" -#include "catlass/gemm/dispatch_policy.hpp" -#include "catlass/gemm/gemm_type.hpp" -#include "catlass/layout/layout.hpp" -#include "catlass/gemm_coord.hpp" -#include "tla/tensor.hpp" -#include "tla/layout.hpp" -#include "tla/tensor.hpp" - -using _0 = tla::Int<0>; -using _1 = tla::Int<1>; -using _2 = tla::Int<2>; -using _4 = tla::Int<4>; -using _8 = tla::Int<8>; -using _16 = tla::Int<16>; -using _32 = tla::Int<32>; -using _64 = tla::Int<64>; -using _128 = tla::Int<128>; -using _256 = tla::Int<256>; -using _512 = tla::Int<512>; -using _1024 = tla::Int<1024>; -using _2048 = tla::Int<2048>; -using _4096 = tla::Int<4096>; -using _8192 = tla::Int<8192>; -using _16384 = tla::Int<16384>; -using _32768 = tla::Int<32768>; -using _65536 = tla::Int<65536>; - -#else #define CATLASS_ARCH 2201 #include "catlass/arch/arch.hpp" @@ -71,7 +28,6 @@ using _65536 = tla::Int<65536>; #include "tla/tensor.hpp" #include "tla/layout.hpp" #include "tla/tensor.hpp" -#endif @@ -81,26 +37,34 @@ using namespace tla; namespace Catlass::Gemm::Kernel { +struct GDNFwdHTileShapes128 { + using L1TileShape = Shape<_128, _128, _128>; + using L0TileShape = L1TileShape; +}; + +struct GDNFwdHTileShapes256 { + using L1TileShape = Shape<_128, _256, _128>; + using L0TileShape = Shape<_128, _256, _64>; +}; + template< typename INPUT_TYPE, typename G_TYPE, typename STATE_TYPE, - typename WORKSPACE_TYPE + typename WORKSPACE_TYPE, + typename TileShapes = GDNFwdHTileShapes128, + bool kGated = false > class GDNFwdHKernel { public: - -#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 - using ArchTag = Arch::Ascend950; -#else + using ArchTag = Arch::AtlasA2; -#endif using CubeScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdHCube; using VecScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdHVec; using DispatchPolicyTla = Gemm::MmadPingpongTlaMulti; - using L1TileShapeTla = Shape<_128, _128, _128>; - using L0TileShapeTla = L1TileShapeTla; + using L1TileShapeVTla = typename TileShapes::L1TileShape; + using L0TileShapeVTla = typename TileShapes::L0TileShape; using WType = Gemm::GemmType; using HType = Gemm::GemmType; @@ -114,19 +78,19 @@ class GDNFwdHKernel { // cube 1 using TileCopyWH = Catlass::Gemm::Tile::PackedTileCopyTla; - using BlockMmadWH = Gemm::Block::BlockMmadTla; + using BlockMmadWH = Gemm::Block::BlockMmadTla; // cube 2 using TileCopyKV = Catlass::Gemm::Tile::PackedTileCopyTla; - using BlockMmadKV = Gemm::Block::BlockMmadTla; + using BlockMmadKV = Gemm::Block::BlockMmadTla; // vec 1 using DispatchPolicyGDNFwdHVnew = Epilogue::EpilogueAtlasGDNFwdHVnew; - using EpilogueGDNFwdHVnew = Epilogue::Block::BlockEpilogue; + using EpilogueGDNFwdHVnew = Epilogue::Block::BlockEpilogue>; // vec 2 using DispatchPolicyGDNFwdHUpdate = Epilogue::EpilogueAtlasGDNFwdHUpdate; - using EpilogueGDNFwdHUpdate = Epilogue::Block::BlockEpilogue; + using EpilogueGDNFwdHUpdate = Epilogue::Block::BlockEpilogue>; using GDNFwdHOffsets = Catlass::Gemm::Block::GDNFwdHOffsets; @@ -140,13 +104,13 @@ class GDNFwdHKernel { using ElementHWork = WORKSPACE_TYPE; using ElementInitialState = STATE_TYPE; using ElementFinalState = STATE_TYPE; - + using LayoutW = Catlass::layout::RowMajor; using LayoutH = Catlass::layout::RowMajor; using LayoutV = Catlass::layout::RowMajor; using LayoutK = Catlass::layout::ColumnMajor; - + uint32_t batch; uint32_t seqlen; uint32_t kNumHead; @@ -154,7 +118,6 @@ class GDNFwdHKernel { uint32_t kHeadDim; uint32_t vHeadDim; uint32_t chunkSize; - uint32_t initalStateStride0; bool useInitialState; bool storeFinalState; uint32_t isVariedLen; @@ -165,7 +128,8 @@ class GDNFwdHKernel { uint32_t hWorkspaceOffset; uint32_t numSeqWorkspaceOffset; uint32_t numChunksWorkspaceOffset; - + uint32_t kDecayWorkspaceOffset; + AscendC::GlobalTensor gmK; AscendC::GlobalTensor gmW; AscendC::GlobalTensor gmU; @@ -177,7 +141,10 @@ class GDNFwdHKernel { AscendC::GlobalTensor gmVWorkspace; AscendC::GlobalTensor gmVUpdateWorkspace; AscendC::GlobalTensor gmHWorkspace; - + + AscendC::GlobalTensor gmGk; + AscendC::GlobalTensor gmKDecayWorkspace; + AscendC::GlobalTensor gmSeqlen; AscendC::GlobalTensor gmNumSeq; AscendC::GlobalTensor gmNumChunks; @@ -190,9 +157,9 @@ class GDNFwdHKernel { __aicore__ inline GDNFwdHKernel() {} - __aicore__ inline void Init(GM_ADDR k, GM_ADDR w, GM_ADDR u, GM_ADDR g, GM_ADDR inital_state, GM_ADDR cu_seqlens, GM_ADDR chunk_indices, + __aicore__ inline void Init(GM_ADDR k, GM_ADDR w, GM_ADDR u, GM_ADDR g, GM_ADDR gk, GM_ADDR inital_state, GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR h, GM_ADDR v_new, GM_ADDR final_state, GM_ADDR tiling, GM_ADDR user) { - + __gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict gdnFwdHTilingData = reinterpret_cast<__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict>(tiling); batch = gdnFwdHTilingData->batch; @@ -202,7 +169,6 @@ class GDNFwdHKernel { kHeadDim = gdnFwdHTilingData->kHeadDim; vHeadDim = gdnFwdHTilingData->vHeadDim; chunkSize = gdnFwdHTilingData->chunkSize; - initalStateStride0 = gdnFwdHTilingData->initalStateStride0; useInitialState = gdnFwdHTilingData->useInitialState; storeFinalState = gdnFwdHTilingData->storeFinalState; isVariedLen = gdnFwdHTilingData->isVariedLen; @@ -213,7 +179,8 @@ class GDNFwdHKernel { hWorkspaceOffset = gdnFwdHTilingData->hWorkspaceOffset; numSeqWorkspaceOffset = gdnFwdHTilingData->numSeqWorkspaceOffset; numChunksWorkspaceOffset = gdnFwdHTilingData->numChunksWorkspaceOffset; - + kDecayWorkspaceOffset = gdnFwdHTilingData->kDecayWorkspaceOffset; + gmK.SetGlobalBuffer((__gm__ ElementK *)k); gmW.SetGlobalBuffer((__gm__ ElementW *)w); gmU.SetGlobalBuffer((__gm__ ElementU *)u); @@ -225,6 +192,8 @@ class GDNFwdHKernel { gmVWorkspace.SetGlobalBuffer((__gm__ ElementVWork *)(user + vWorkspaceOffset)); gmVUpdateWorkspace.SetGlobalBuffer((__gm__ ElementV *)(user + vUpdateWorkspaceOffset)); gmHWorkspace.SetGlobalBuffer((__gm__ ElementHWork *)(user + hWorkspaceOffset)); + gmGk.SetGlobalBuffer((__gm__ ElementG *)gk); + gmKDecayWorkspace.SetGlobalBuffer((__gm__ ElementK *)(user + kDecayWorkspaceOffset)); gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens); gmNumSeq.SetGlobalBuffer((__gm__ int64_t *)(user + numSeqWorkspaceOffset)); @@ -238,8 +207,11 @@ class GDNFwdHKernel { vecBlockScheduler.Init(cu_seqlens, chunk_indices, tiling, user); } } - + __aicore__ inline void Process() { + if (isVariedLen) { + AscendC::SyncAll(); + } if ASCEND_IS_AIC { uint32_t coreIdx = AscendC::GetBlockIdx(); @@ -250,67 +222,83 @@ class GDNFwdHKernel { auto wLayout = tla::MakeLayout(shapeBatch * kNumHead * cubeBlockScheduler.totalTokens, kHeadDim); auto hLayout = tla::MakeLayout(shapeBatch * vNumHead * cubeBlockScheduler.totalChunks * kHeadDim, vHeadDim); - auto vLayout = tla::MakeLayout(coreNum * chunkSize * PING_PONG_STAGES, vHeadDim); - - auto kLayout = tla::MakeLayout(kHeadDim, shapeBatch * kNumHead * cubeBlockScheduler.totalTokens); - auto vworkLayout = tla::MakeLayout(coreNum * chunkSize * PING_PONG_STAGES, vHeadDim); - auto hworkLayout = tla::MakeLayout(coreNum * kHeadDim * PING_PONG_STAGES, vHeadDim); + auto vLayout = tla::MakeLayout(coreNum * chunkSize * PING_PONG_STAGES, cubeBlockScheduler.vBlockSize); + auto kLayout = tla::MakeLayout(kHeadDim, shapeBatch * kNumHead * cubeBlockScheduler.totalTokens); + auto vworkLayout = tla::MakeLayout(coreNum * chunkSize * PING_PONG_STAGES, cubeBlockScheduler.vBlockSize); + auto hworkLayout = tla::MakeLayout(coreNum * kHeadDim * PING_PONG_STAGES, cubeBlockScheduler.vBlockSize); + AscendC::SyncAll(); + uint32_t currStage = 0; // 0: C1, 1: C2 while (cubeBlockScheduler.isRunning) { - cubeBlockScheduler.InitTask(); - // step 1: v_work = w @ h[i] - GDNFwdHOffsets& cube1Offsets = cubeBlockScheduler.GetStage1Offsets(); - Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done); - if (cubeBlockScheduler.NeedProcessStage1()) { - int64_t cube1OffsetW = cube1Offsets.wOffset; - int64_t cube1OffsetH = cube1Offsets.hSrcOffset; - int64_t cube1OffsetVwork = cube1Offsets.vWorkOffset; - auto tensorW = tla::MakeTensor(gmW[cube1OffsetW], wLayout, Catlass::Arch::PositionGM{}); - auto tensorH = tla::MakeTensor(gmH[cube1OffsetH], hLayout, Catlass::Arch::PositionGM{}); - auto tensorV = tla::MakeTensor(gmVWorkspace[cube1OffsetVwork], vLayout, Catlass::Arch::PositionGM{}); - GemmCoord cube1Shape {cube1Offsets.blockTokens, vHeadDim, kHeadDim}; - auto tensorBlockW = GetTile(tensorW, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k())); - auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n())); - auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n())); - blockMmadWH.preSetFlags(); - blockMmadWH(tensorBlockW, tensorBlockH, tensorBlockV, cube1Shape); - blockMmadWH.finalWaitFlags(); - } - Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube1Done); - - if (cubeBlockScheduler.iterId > 1) { - Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done); - GDNFwdHOffsets& cube2Offsets = cubeBlockScheduler.GetStage2Offsets(); - if (cubeBlockScheduler.NeedProcessStage2()) { - // step 3: h[i+1] = k.T @ v_work - int64_t cube2OffsetK = cube2Offsets.wkOffset; - int64_t cube2OffsetVwork = cube2Offsets.vWorkOffset; - int64_t cube2OffsetH = cube2Offsets.hWorkOffset; - auto tensorK = tla::MakeTensor(gmK[cube2OffsetK], kLayout, Catlass::Arch::PositionGM{}); - auto tensorVwork = tla::MakeTensor(gmVUpdateWorkspace[cube2OffsetVwork], vworkLayout, Catlass::Arch::PositionGM{}); - auto tensorHwork = tla::MakeTensor(gmHWorkspace[cube2OffsetH], hworkLayout, Catlass::Arch::PositionGM{}); - GemmCoord cube2Shape{kHeadDim, vHeadDim, cube2Offsets.blockTokens}; - auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k())); - auto tensorBlockVwork = GetTile(tensorVwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n())); - auto tensorBlockHwork = GetTile(tensorHwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n())); - blockMmadKV.preSetFlags(); - blockMmadKV(tensorBlockK, tensorBlockVwork, tensorBlockHwork, cube2Shape); - blockMmadKV.finalWaitFlags(); + if (currStage == 0) { + /* C1: v_work = w @ h[i] */ + cubeBlockScheduler.InitTasks(); + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = cubeBlockScheduler.GetStreamId(i); + const auto& stream = cubeBlockScheduler.GetStream(i); + if (cubeBlockScheduler.StreamIsDone(stream)) { + continue; + } + + const GDNFwdHOffsets& cube1Offsets = cubeBlockScheduler.GetCurTaskOffsets(stream); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[streamId]); + int64_t cube1OffsetW = cube1Offsets.wOffset; + int64_t cube1OffsetH = cube1Offsets.hSrcOffset; + int64_t cube1OffsetVwork = cube1Offsets.vWorkOffset; + auto tensorW = tla::MakeTensor(gmW[cube1OffsetW], wLayout, Catlass::Arch::PositionGM{}); + auto tensorH = tla::MakeTensor(gmH[cube1OffsetH], hLayout, Catlass::Arch::PositionGM{}); + auto tensorV = tla::MakeTensor(gmVWorkspace[cube1OffsetVwork], vLayout, Catlass::Arch::PositionGM{}); + GemmCoord cube1Shape {cube1Offsets.blockTokens, cube1Offsets.vBlockDim, kHeadDim}; + auto tensorBlockW = GetTile(tensorW, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k())); + auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n())); + auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n())); + blockMmadWH.preSetFlags(); + blockMmadWH(tensorBlockW, tensorBlockH, tensorBlockV, cube1Shape); + blockMmadWH.finalWaitFlags(); + Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube1Done[streamId]); + } + } else { + /* C2: h[i+1] = k.T @ v_work */ + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = cubeBlockScheduler.GetStreamId(i); + const auto& stream = cubeBlockScheduler.GetStream(i); + if (cubeBlockScheduler.StreamIsDone(stream)) { + continue; + } + const GDNFwdHOffsets& cube2Offsets = cubeBlockScheduler.GetCurTaskOffsets(stream); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done[streamId]); + + if (cubeBlockScheduler.NeedProcessStage2(stream)) { + // step 3: h[i+1] = k.T @ v_work + int64_t cube2OffsetKwork = kGated ? cube2Offsets.kDecayWorkOffset : cube2Offsets.wkOffset; + int64_t cube2OffsetVwork = cube2Offsets.vWorkOffset; + int64_t cube2OffsetH = cube2Offsets.hWorkOffset; + auto tensorK = kGated + ? tla::MakeTensor(gmKDecayWorkspace[cube2OffsetKwork], kLayout, Catlass::Arch::PositionGM{}) + : tla::MakeTensor(gmK[cube2OffsetKwork], kLayout, Catlass::Arch::PositionGM{}); + auto tensorVwork = tla::MakeTensor(gmVUpdateWorkspace[cube2OffsetVwork], vworkLayout, Catlass::Arch::PositionGM{}); + auto tensorHwork = tla::MakeTensor(gmHWorkspace[cube2OffsetH], hworkLayout, Catlass::Arch::PositionGM{}); + GemmCoord cube2Shape{kHeadDim, cube2Offsets.vBlockDim, cube2Offsets.blockTokens}; + auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k())); + auto tensorBlockVwork = GetTile(tensorVwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n())); + auto tensorBlockHwork = GetTile(tensorHwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n())); + blockMmadKV.preSetFlags(); + blockMmadKV(tensorBlockK, tensorBlockVwork, tensorBlockHwork, cube2Shape); + blockMmadKV.finalWaitFlags(); + } + Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube2Done[streamId]); } - Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube2Done); } + currStage ^= 0x01; } - Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[0]); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[1]); } if ASCEND_IS_AIV { uint32_t coreIdx = AscendC::GetBlockIdx(); uint32_t coreNum = AscendC::GetBlockNum(); - uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); - uint32_t subBlockNum = AscendC::GetSubBlockNum(); - - EpilogueGDNFwdHVnew epilogueGDNFwdHVnew(resource); if (useInitialState) { AscendC::LocalTensor stateUbTensorPing = resource.ubBuf.template GetBufferByByte(0); @@ -318,91 +306,262 @@ class GDNFwdHKernel { AscendC::LocalTensor hUbTensorPing = resource.ubBuf.template GetBufferByByte(64 * 1024); AscendC::LocalTensor hUbTensorPong = resource.ubBuf.template GetBufferByByte(160 * 1024); uint32_t totalChunks = isVariedLen ? vecBlockScheduler.totalChunks : ((seqlen + chunkSize - 1) / chunkSize); + uint32_t transferCount = isVariedLen ? (vecBlockScheduler.tokenBatch * vNumHead / coreNum) : (shapeBatch * vNumHead / coreNum); + uint32_t remainderFlag = isVariedLen ? (((vecBlockScheduler.tokenBatch * vNumHead) % coreNum) != 0): (((shapeBatch * vNumHead) % coreNum) != 0); + uint32_t step = transferCount + remainderFlag; uint32_t stateBlockSize = kHeadDim * vHeadDim; uint32_t pingpongFlag = 1; + uint32_t start = coreIdx * step; + uint32_t end = start + step; + uint32_t maxLimit = isVariedLen ? vecBlockScheduler.tokenBatch * vNumHead : shapeBatch * vNumHead; + uint32_t realEnd = min(end, maxLimit); AscendC::SetFlag(EVENT_ID0); AscendC::SetFlag(EVENT_ID1); - AscendC::DataCopyParams repeatParams = {static_cast(kHeadDim), static_cast(vHeadDim * sizeof(ElementInitialState) / 32), - static_cast((initalStateStride0 - vHeadDim)* sizeof(ElementInitialState) / 32), static_cast(0)}; - for (uint32_t shapeBatchIdx = 0; shapeBatchIdx < shapeBatch; shapeBatchIdx++) { - for (uint32_t vHeadIdx = 0; vHeadIdx < vNumHead; vHeadIdx++) { - for (uint32_t tokenBatchIdx = 0; tokenBatchIdx < vecBlockScheduler.tokenBatch; tokenBatchIdx++) { - uint32_t batchIdx = isVariedLen ? tokenBatchIdx : shapeBatchIdx; - uint32_t chunkOffset = isVariedLen ? gmNumChunks.GetValue(tokenBatchIdx) : 0; - uint32_t initialStateSrcOffset = (batchIdx * vNumHead + vHeadIdx) * kHeadDim * initalStateStride0; - uint32_t hOffset = (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset) * stateBlockSize; + for (uint32_t initialStateBlockOffset = start; initialStateBlockOffset >= start && initialStateBlockOffset < realEnd; initialStateBlockOffset++) { + uint32_t batchIdx = initialStateBlockOffset / vNumHead; + uint32_t vHeadIdx = initialStateBlockOffset % vNumHead; + uint32_t chunkOffset = isVariedLen ? gmNumChunks.GetValue(batchIdx) : 0; + uint32_t initialStateBaseOffset = initialStateBlockOffset * stateBlockSize; + uint32_t shapeBatchIdx = isVariedLen ? 0 : batchIdx; + uint32_t hBaseOffset = (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset) * stateBlockSize; + if (vHeadDim <= 128) { + AscendC::LocalTensor stateUbTensor = pingpongFlag ? stateUbTensorPing : stateUbTensorPong; + AscendC::LocalTensor hUbTensor = pingpongFlag ? hUbTensorPing : hUbTensorPong; + auto event_id = pingpongFlag ? EVENT_ID1 : EVENT_ID0; + AscendC::WaitFlag(event_id); + if constexpr(!std::is_same::value) { + AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateBaseOffset], stateBlockSize); + AscendC::SetFlag(event_id); + AscendC::WaitFlag(event_id); + AscendC::Cast(hUbTensor, stateUbTensor, AscendC::RoundMode::CAST_RINT, stateBlockSize); + AscendC::SetFlag(event_id); + AscendC::WaitFlag(event_id); + AscendC::DataCopy(gmH[hBaseOffset], hUbTensor, stateBlockSize); + } else { + AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateBaseOffset], stateBlockSize); + AscendC::SetFlag(event_id); + AscendC::WaitFlag(event_id); + AscendC::DataCopy(gmH[hBaseOffset], stateUbTensor, stateBlockSize); + } + AscendC::SetFlag(event_id); + pingpongFlag = 1 - pingpongFlag; + } else { + uint32_t stateRowTile = (32 * 1024) / (vHeadDim * sizeof(ElementH)); + for (uint32_t rowOffset = 0; rowOffset < kHeadDim; rowOffset += stateRowTile) { + uint32_t rowsThisTile = Min(stateRowTile, kHeadDim - rowOffset); + uint32_t stateTileElems = rowsThisTile * vHeadDim; + uint32_t initialStateOffset = initialStateBaseOffset + rowOffset * vHeadDim; + uint32_t hOffset = hBaseOffset + rowOffset * vHeadDim; AscendC::LocalTensor stateUbTensor = pingpongFlag ? stateUbTensorPing : stateUbTensorPong; AscendC::LocalTensor hUbTensor = pingpongFlag ? hUbTensorPing : hUbTensorPong; auto event_id = pingpongFlag ? EVENT_ID1 : EVENT_ID0; AscendC::WaitFlag(event_id); if constexpr(!std::is_same::value) { - AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateSrcOffset], repeatParams); + AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateOffset], stateTileElems); AscendC::SetFlag(event_id); AscendC::WaitFlag(event_id); - AscendC::Cast(hUbTensor, stateUbTensor, AscendC::RoundMode::CAST_RINT, stateBlockSize); + AscendC::Cast(hUbTensor, stateUbTensor, AscendC::RoundMode::CAST_RINT, stateTileElems); AscendC::SetFlag(event_id); AscendC::WaitFlag(event_id); - AscendC::DataCopy(gmH[hOffset], hUbTensor, stateBlockSize); + AscendC::DataCopy(gmH[hOffset], hUbTensor, stateTileElems); } else { - AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateSrcOffset], repeatParams); + AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateOffset], stateTileElems); AscendC::SetFlag(event_id); AscendC::WaitFlag(event_id); - AscendC::DataCopy(gmH[hOffset], stateUbTensor, stateBlockSize); + AscendC::DataCopy(gmH[hOffset], stateUbTensor, stateTileElems); } AscendC::SetFlag(event_id); pingpongFlag = 1 - pingpongFlag; } - + } + } + + AscendC::WaitFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID1); + } else { + uint32_t stateRowTile = (32 * 1024) / (vHeadDim * sizeof(ElementH)); + AscendC::LocalTensor chunkOffsetsUb = + resource.ubBuf.template GetBufferByByte(0); + if (isVariedLen) { + uint32_t chunkOffsetBytes = (vecBlockScheduler.tokenBatch + 1) * sizeof(int64_t); + AscendC::DataCopyParams copyParams{ + 1, static_cast(chunkOffsetBytes), 0, 0}; + AscendC::DataCopyPadParams padParams{false, 0, 0, 0}; + AscendC::DataCopyPad(chunkOffsetsUb, gmNumChunks[0], copyParams, padParams); + AscendC::SetFlag(EVENT_ID3); + AscendC::WaitFlag(EVENT_ID3); + AscendC::SetFlag(EVENT_ID3); + AscendC::WaitFlag(EVENT_ID3); + } + auto chunkOffsets = reinterpret_cast<__ubuf__ int64_t *>(chunkOffsetsUb.GetPhyAddr()); + AscendC::LocalTensor hUbTensorPing = + resource.ubBuf.template GetBufferByByte(64 * 1024); + AscendC::LocalTensor hUbTensorPong = + resource.ubBuf.template GetBufferByByte(160 * 1024); + uint32_t totalChunks = + isVariedLen ? vecBlockScheduler.totalChunks : ((seqlen + chunkSize - 1) / chunkSize); + uint32_t taskCount = + (isVariedLen ? vecBlockScheduler.tokenBatch : shapeBatch) * vNumHead; + uint32_t step = taskCount / coreNum + ((taskCount % coreNum) != 0); + uint32_t start = coreIdx * step; + uint32_t realEnd = Min(start + step, taskCount); + uint32_t pingpongFlag = 1; + AscendC::SetFlag(EVENT_ID0); + AscendC::SetFlag(EVENT_ID1); + for (uint32_t taskIdx = start; taskIdx < realEnd; ++taskIdx) { + uint32_t batchIdx = taskIdx / vNumHead; + uint32_t vHeadIdx = taskIdx % vNumHead; + uint32_t chunkOffset = isVariedLen ? static_cast(chunkOffsets[batchIdx]) : 0; + uint32_t shapeBatchIdx = isVariedLen ? 0 : batchIdx; + uint32_t hBaseOffset = + (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset) * + kHeadDim * vHeadDim; + for (uint32_t rowOffset = 0; rowOffset < kHeadDim; rowOffset += stateRowTile) { + uint32_t rowsThisTile = Min(stateRowTile, kHeadDim - rowOffset); + uint32_t stateTileElems = rowsThisTile * vHeadDim; + AscendC::LocalTensor hUbTensor = + pingpongFlag ? hUbTensorPing : hUbTensorPong; + auto eventId = pingpongFlag ? EVENT_ID1 : EVENT_ID0; + AscendC::WaitFlag(eventId); + AscendC::Duplicate(hUbTensor, static_cast(0), stateTileElems); + AscendC::SetFlag(eventId); + AscendC::WaitFlag(eventId); + AscendC::DataCopy(gmH[hBaseOffset + rowOffset * vHeadDim], hUbTensor, stateTileElems); + AscendC::SetFlag(eventId); + pingpongFlag = 1 - pingpongFlag; } } AscendC::WaitFlag(EVENT_ID0); AscendC::WaitFlag(EVENT_ID1); } - Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done); - Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done); + AscendC::SyncAll(); + + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done[0]); + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done[1]); + + EpilogueGDNFwdHVnew epilogueGDNFwdHVnew(resource); + EpilogueGDNFwdHUpdate epilogueGDNFwdHUpdate(resource); + uint32_t pongBaseEvent = 4; + + if (storeFinalState && std::is_same::value) { + AscendC::SetFlag(EVENT_ID0); // preset v + AscendC::SetFlag(EVENT_ID0 + pongBaseEvent); + AscendC::SetFlag(EVENT_ID2); // preset h + AscendC::SetFlag(EVENT_ID2 + pongBaseEvent); + } else { + AscendC::SetFlag(EVENT_ID0); // preset v + AscendC::SetFlag(EVENT_ID0 + pongBaseEvent); + AscendC::SetFlag(EVENT_ID2); // preset h + AscendC::SetFlag(EVENT_ID2 + pongBaseEvent); + } + AscendC::SetFlag(EVENT_ID1); // preset u + AscendC::SetFlag(EVENT_ID1 + pongBaseEvent); + AscendC::SetFlag(EVENT_ID3); // preset g + AscendC::SetFlag(EVENT_ID3 + pongBaseEvent); + uint32_t currStage = 0; // 0: V1, 1: V2 + bool event0FromMte3[PING_PONG_STAGES] = {false, false}; + bool event2FromMte3[PING_PONG_STAGES] = {!(storeFinalState && std::is_same::value), + !(storeFinalState && std::is_same::value)}; while (vecBlockScheduler.isRunning) { - vecBlockScheduler.InitTask(); - // step 2: - GDNFwdHOffsets& vec1Offsets = vecBlockScheduler.GetStage1Offsets(); - // gmV = gmU - gmVWorkspace - // g_buf = gmG[-1] - gmG - // g_buf = exp(g_buf) - // gmVWorkspace = g_buf * gmV - if (vecBlockScheduler.NeedProcessStage1()) { - epilogueGDNFwdHVnew( - gmV[vec1Offsets.uvOffset], gmVUpdateWorkspace[vec1Offsets.vWorkOffset], - gmG[vec1Offsets.gOffset], gmU[vec1Offsets.uvOffset], gmVWorkspace[vec1Offsets.vWorkOffset], - vec1Offsets.blockTokens, kHeadDim, vHeadDim, vecBlockScheduler.cube1Done - ); - } else { - Arch::CrossCoreWaitFlag(vecBlockScheduler.cube1Done); - } - Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec1Done); - - if (vecBlockScheduler.iterId > 1) { - GDNFwdHOffsets& vec2Offsets = vecBlockScheduler.GetStage2Offsets(); - if (vecBlockScheduler.NeedProcessStage2()) { - // step 4: h[i+1] += h_work if i < num_chunks - 1 else None - EpilogueGDNFwdHUpdate epilogueGDNFwdHUpdate(resource); - epilogueGDNFwdHUpdate( - gmH[vec2Offsets.hDstOffset], gmFinalState[vec2Offsets.finalStateOffset], - gmG[vec2Offsets.gOffset], - gmH[vec2Offsets.hSrcOffset], - gmHWorkspace[vec2Offsets.hWorkOffset], - vec2Offsets.blockTokens, kHeadDim, vHeadDim, vecBlockScheduler.cube2Done, - (vec2Offsets.isFinalState && storeFinalState) + if (currStage == 0) { + /* V1: + * gmV = gmU - gmVWorkspace + * g_buf = gmG[-1] - gmG + * g_buf = exp(g_buf) + * gmVWorkspace = g_buf * gmV + */ + vecBlockScheduler.InitTasks(); + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = vecBlockScheduler.GetStreamId(i); + const auto& stream = vecBlockScheduler.GetStream(i); + if (vecBlockScheduler.StreamIsDone(stream)) { + continue; + } + const GDNFwdHOffsets& vec1Offsets = vecBlockScheduler.GetCurTaskOffsets(stream); + bool waitWsFromMte3 = storeFinalState && std::is_same::value && + event0FromMte3[streamId]; + epilogueGDNFwdHVnew( + gmV[vec1Offsets.uvOffset], gmVUpdateWorkspace[vec1Offsets.vWorkOffset], + gmG[vec1Offsets.gOffset], gmU[vec1Offsets.uvOffset], gmVWorkspace[vec1Offsets.vWorkOffset], + gmGk[vec1Offsets.gkOffset], gmK[vec1Offsets.wkOffset], gmKDecayWorkspace[vec1Offsets.kDecayWorkOffset], + vec1Offsets.blockTokens, kHeadDim, vec1Offsets.vBlockDim, vHeadDim, + vecBlockScheduler.cube1Done[streamId], vecBlockScheduler.vec1Done[streamId], + vec1Offsets.isInitialState, vec1Offsets.isFinalState, storeFinalState, + waitWsFromMte3, (streamId == 0) ); - } else { - Arch::CrossCoreWaitFlag(vecBlockScheduler.cube2Done); + if (storeFinalState && std::is_same::value) { + event0FromMte3[streamId] = false; + } + } + } else { + /* V2: h[i+1] += h_work if i < num_chunks - 1 else None */ + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = vecBlockScheduler.GetStreamId(i); + const auto& stream = vecBlockScheduler.GetStream(i); + if (vecBlockScheduler.StreamIsDone(stream)) { + continue; + } + const GDNFwdHOffsets& vec2Offsets = vecBlockScheduler.GetCurTaskOffsets(stream); + if (vecBlockScheduler.NeedProcessStage2(stream)) { + if (storeFinalState && std::is_same::value) { + event0FromMte3[streamId] = vec2Offsets.isFinalState; + event2FromMte3[streamId] = !vec2Offsets.isFinalState; + } + // step 4: h[i+1] += h_work if i < num_chunks - 1 else None + epilogueGDNFwdHUpdate( + gmH[vec2Offsets.hDstOffset], gmFinalState[vec2Offsets.finalStateOffset], + gmG[vec2Offsets.gOffset], + gmH[vec2Offsets.hSrcOffset], + gmHWorkspace[vec2Offsets.hWorkOffset], + gmGk[vec2Offsets.gkOffset], + vec2Offsets.blockTokens, kHeadDim, vec2Offsets.vBlockDim, vHeadDim, vecBlockScheduler.cube2Done[streamId], + vec2Offsets.isInitialState, vec2Offsets.isFinalState, storeFinalState, (streamId == 0) + ); + } else { + Arch::CrossCoreWaitFlag(vecBlockScheduler.cube2Done[streamId]); + } + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done[streamId]); } - Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done); } + currStage ^= 0x01; + } + + if (storeFinalState && std::is_same::value) { + if (event0FromMte3[0]) { + AscendC::WaitFlag(EVENT_ID0); + } else { + AscendC::WaitFlag(EVENT_ID0); + } + if (event0FromMte3[1]) { + AscendC::WaitFlag(EVENT_ID0 + pongBaseEvent); + } else { + AscendC::WaitFlag(EVENT_ID0 + pongBaseEvent); + } + if (event2FromMte3[0]) { + AscendC::WaitFlag(EVENT_ID2); + } else { + AscendC::WaitFlag(EVENT_ID2); + } + if (event2FromMte3[1]) { + AscendC::WaitFlag(EVENT_ID2 + pongBaseEvent); + } else { + AscendC::WaitFlag(EVENT_ID2 + pongBaseEvent); + } + } else { + AscendC::WaitFlag(EVENT_ID0); // preset v + AscendC::WaitFlag(EVENT_ID0 + pongBaseEvent); + AscendC::WaitFlag(EVENT_ID2); // preset h + AscendC::WaitFlag(EVENT_ID2 + pongBaseEvent); } + AscendC::WaitFlag(EVENT_ID1); // preset u + AscendC::WaitFlag(EVENT_ID1 + pongBaseEvent); + AscendC::WaitFlag(EVENT_ID3); // preset g + AscendC::WaitFlag(EVENT_ID3 + pongBaseEvent); } } }; -} \ No newline at end of file +} diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_update.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_update.hpp new file mode 100644 index 000000000000..9d141633659e --- /dev/null +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_update.hpp @@ -0,0 +1,399 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#ifndef CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_UPDATE_HPP +#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_UPDATE_HPP +#include "catlass/catlass.hpp" +#include "catlass/arch/resource.hpp" +#include "../gdn_fwd_h_epilogue_policies.hpp" +#include "catlass/gemm_coord.hpp" +#include "catlass/matrix_coord.hpp" +#include "catlass/epilogue/tile/tile_copy.hpp" + +namespace Catlass::Epilogue::Block { + +template < + class HOutputType_, + class GInputType_, + class HInputType_, + class HUpdateInputType_, + class FinalStateType_, + class KGatedTag +> +class BlockEpilogue < + EpilogueAtlasGDNFwdHUpdate, + HOutputType_, + GInputType_, + HInputType_, + HUpdateInputType_, + FinalStateType_, + KGatedTag +> { + static constexpr bool kGated = KGatedTag::value; +public: + // Type aliases + using DispatchPolicy = EpilogueAtlasGDNFwdHUpdate; + using ArchTag = typename DispatchPolicy::ArchTag; + + using HElementOutput = typename HOutputType_::Element; + using GElementInput = typename GInputType_::Element; + using HElementInput = typename HInputType_::Element; + using HUpdateElementInput = typename HUpdateInputType_::Element; + using FinalStateElement = typename FinalStateType_::Element; + + CATLASS_DEVICE + BlockEpilogue(Arch::Resource &resource) + { + + constexpr uint32_t CALC_BUF_OFFSET = 0; + constexpr uint32_t PING_BUF_0_OFFSET = 32 * 1024; + constexpr uint32_t PING_BUF_1_OFFSET = 48 * 1024; + constexpr uint32_t PING_BUF_2_OFFSET = 64 * 1024; + constexpr uint32_t PING_BUF_3_OFFSET = 80 * 1024; + constexpr uint32_t PONG_BUF_0_OFFSET = 96 * 1024; + constexpr uint32_t PONG_BUF_1_OFFSET = 112 * 1024; + constexpr uint32_t PONG_BUF_2_OFFSET = 128 * 1024; + constexpr uint32_t PONG_BUF_3_OFFSET = 144 * 1024; + constexpr uint32_t PING_G_BUF_OFFSET = 168 * 1024; + constexpr uint32_t PONG_G_BUF_OFFSET = 169 * 1024; + constexpr uint32_t PING_G_SUB_BUF_OFFSET = 170 * 1024; + constexpr uint32_t PONG_G_SUB_BUF_OFFSET = 171 * 1024; + constexpr uint32_t PING_G_INPUT_BUF_OFFSET = 172 * 1024; + constexpr uint32_t PONG_G_INPUT_BUF_OFFSET = 173 * 1024; + constexpr uint32_t UPDATE_SCRATCH_BUF_OFFSET = 160 * 1024; + constexpr uint32_t UPDATE_G_BUF_OFFSET = 176 * 1024; + + + calcUbTensor = resource.ubBuf.template GetBufferByByte(CALC_BUF_OFFSET); + + hUpdateUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_0_OFFSET); + hUbTensor_ping = resource.ubBuf.template GetBufferByByte(UPDATE_SCRATCH_BUF_OFFSET); + finalOutputUbTensor_ping = resource.ubBuf.template GetBufferByByte(UPDATE_SCRATCH_BUF_OFFSET); + glastUbTensor_ping = resource.ubBuf.template GetBufferByByte(UPDATE_G_BUF_OFFSET); + + hUpdateUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_0_OFFSET); + hUbTensor_pong = resource.ubBuf.template GetBufferByByte(UPDATE_SCRATCH_BUF_OFFSET); + finalOutputUbTensor_pong = resource.ubBuf.template GetBufferByByte(UPDATE_SCRATCH_BUF_OFFSET); + glastUbTensor_pong = resource.ubBuf.template GetBufferByByte(UPDATE_G_BUF_OFFSET); + + if constexpr (kGated) { + gkLastUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_SUB_BUF_OFFSET); + gkLastUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_SUB_BUF_OFFSET); + gkInputUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_INPUT_BUF_OFFSET); + gkInputUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_INPUT_BUF_OFFSET); + gkBrcbUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_INPUT_BUF_OFFSET); + gkBrcbUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_INPUT_BUF_OFFSET); + } + + } + + CATLASS_DEVICE + ~BlockEpilogue() {} + + template + CATLASS_DEVICE + void CopyGmToUb( + AscendC::LocalTensor dst, + AscendC::GlobalTensor src, + uint32_t rows, + uint32_t cols, + uint32_t srcStride) + { + if (cols == srcStride) { + AscendC::DataCopy(dst, src, rows * cols); + return; + } + AscendC::DataCopyExtParams copyParams{ + static_cast(rows), + static_cast(cols * sizeof(Element)), + static_cast((srcStride - cols) * sizeof(Element)), + 0, + 0}; + AscendC::DataCopyPadExtParams padParams{false, 0, 0, 0}; + AscendC::DataCopyPad(dst, src, copyParams, padParams); + } + + template + CATLASS_DEVICE + void CopyUbToGm( + AscendC::GlobalTensor dst, + AscendC::LocalTensor src, + uint32_t rows, + uint32_t cols, + uint32_t dstStride) + { + if (cols == dstStride) { + AscendC::DataCopy(dst, src, rows * cols); + return; + } + AscendC::DataCopyExtParams copyParams{ + static_cast(rows), + static_cast(cols * sizeof(Element)), + 0, + static_cast((dstStride - cols) * sizeof(Element)), + 0}; + AscendC::DataCopyPad(dst, src, copyParams); + } + + CATLASS_DEVICE + void ApplyRowScale( + AscendC::LocalTensor matrix, + AscendC::LocalTensor rowScale, + AscendC::LocalTensor rowScaleBrcb, + uint32_t rows, + uint32_t cols) + { + constexpr uint32_t FP32_PER_BLOCK = 8; + constexpr uint32_t FP32_PER_REPEAT = 64; + uint8_t rowStride = static_cast(cols / FP32_PER_BLOCK); + AscendC::BinaryRepeatParams params(1, 1, 0, rowStride, rowStride, 1); + for (uint32_t row = 0; row < rows; row += FP32_PER_BLOCK) { + uint32_t rowsThisBlock = Min(FP32_PER_BLOCK, rows - row); + AscendC::Brcb(rowScaleBrcb, rowScale[row], 1, {1, FP32_PER_BLOCK}); + AscendC::PipeBarrier(); + for (uint32_t col = 0; col < cols; col += FP32_PER_REPEAT) { + uint32_t count = Min(FP32_PER_REPEAT, cols - col); + uint32_t offset = row * cols + col; + AscendC::Mul(matrix[offset], matrix[offset], rowScaleBrcb, + count, rowsThisBlock, params); + } + AscendC::PipeBarrier(); + } + } + + CATLASS_DEVICE + void operator()( + AscendC::GlobalTensor hOutput, + AscendC::GlobalTensor finalState, + AscendC::GlobalTensor gInput, + AscendC::GlobalTensor hInput, + AscendC::GlobalTensor hUpdateInput, + AscendC::GlobalTensor gkInput, + uint32_t chunkSize, + uint32_t kHeadDim, + uint32_t vBlockDim, + uint32_t vHeadDim, + Arch::CrossCoreFlag cube2Done, + bool isInitialState, + bool isFinalState, + bool storeFinalState, + bool isPing + ) + { + static constexpr uint32_t ROW_TILE = 16; + uint32_t mActual = kHeadDim; + uint32_t nActual = vBlockDim; + uint32_t outputStride = vHeadDim; + uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); + uint32_t subBlockNum = AscendC::GetSubBlockNum(); + uint32_t rowsPerSubBlock = CeilDiv(mActual, subBlockNum); + uint32_t rowBegin = subBlockIdx * rowsPerSubBlock; + uint32_t rowEnd = rowBegin + rowsPerSubBlock; + if (rowEnd > mActual) { + rowEnd = mActual; + } + if (rowBegin >= mActual) { + Arch::CrossCoreWaitFlag(cube2Done); + return; + } + + AscendC::ResetMask(); + + AscendC::GlobalTensor gInputThisSubBlock = gInput; + + uint32_t pingpongFlag = isPing ? 0 : pongBaseEvent; + AscendC::LocalTensor hUpdateUbTensor = isPing ? hUpdateUbTensor_ping : hUpdateUbTensor_pong; + AscendC::LocalTensor hUbTensor = isPing ? hUbTensor_ping : hUbTensor_pong; + AscendC::LocalTensor finalOutputUbTensor = isPing ? finalOutputUbTensor_ping : finalOutputUbTensor_pong; + AscendC::LocalTensor glastUbTensor = isPing ? glastUbTensor_ping : glastUbTensor_pong; + + GElementInput gLastVal = gInputThisSubBlock.GetValue(chunkSize-1); + float gLastFloat = 0.0f; + if constexpr(std::is_same::value) { + gLastFloat = gLastVal; + } else if constexpr(std::is_same::value) { + gLastFloat = (float)gLastVal; + } else if constexpr(std::is_same::value) { + gLastFloat = AscendC::ToFloat(gLastVal); + } + glastUbTensor.SetValue(0, gLastFloat); + + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + AscendC::Exp(glastUbTensor, glastUbTensor, 1); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + float muls = glastUbTensor.GetValue(0); + if constexpr (kGated) { + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + } + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + + Arch::CrossCoreWaitFlag(cube2Done); + // fix: need to adapt kGated. issue: A5 do not have vdim128 branch. + bool waitHFromV = storeFinalState && isInitialState && std::is_same::value; + bool waitUpdateFromMte3 = false; + for (uint32_t rowStart = rowBegin; rowStart < rowEnd; rowStart += ROW_TILE) { + uint32_t rowsThisTile = rowEnd - rowStart; + if (rowsThisTile > ROW_TILE) { + rowsThisTile = ROW_TILE; + } + + AscendC::GlobalTensor hOutputThisTile = hOutput[rowStart * outputStride]; + AscendC::GlobalTensor hInputThisTile = hInput[rowStart * outputStride]; + AscendC::GlobalTensor hUpdateInputThisTile = hUpdateInput[rowStart * nActual]; + AscendC::GlobalTensor finalStateThisTile = finalState[rowStart * outputStride]; + + if (waitHFromV) { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + } + CopyGmToUb(hUbTensor, hInputThisTile, rowsThisTile, nActual, outputStride); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + + AscendC::Cast(calcUbTensor, hUbTensor, AscendC::RoundMode::CAST_NONE, rowsThisTile * nActual); + AscendC::PipeBarrier(); + if (storeFinalState && isFinalState && std::is_same::value) { + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + waitHFromV = true; + } else { + waitHFromV = false; + } + + AscendC::Muls(calcUbTensor, calcUbTensor, muls, rowsThisTile * nActual); + AscendC::PipeBarrier(); + + if constexpr (kGated) { + AscendC::GlobalTensor gkLastInput = + gkInput[(chunkSize - 1) * kHeadDim + rowStart]; + AscendC::LocalTensor gkLastUbTensor = + isPing ? gkLastUbTensor_ping : gkLastUbTensor_pong; + AscendC::LocalTensor gkInputUbTensor = + isPing ? gkInputUbTensor_ping : gkInputUbTensor_pong; + AscendC::LocalTensor gkBrcbUbTensor = + isPing ? gkBrcbUbTensor_ping : gkBrcbUbTensor_pong; + + if (rowStart == rowBegin) { + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + } + if constexpr(std::is_same::value) { + AscendC::DataCopy(gkLastUbTensor, gkLastInput, rowsThisTile); + } else { + AscendC::DataCopy(gkInputUbTensor, gkLastInput, rowsThisTile); + } + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + if constexpr(!std::is_same::value) { + AscendC::Cast(gkLastUbTensor, gkInputUbTensor, + AscendC::RoundMode::CAST_NONE, rowsThisTile); + } + AscendC::PipeBarrier(); + AscendC::Muls(gkLastUbTensor, gkLastUbTensor, 0.6931471805599453f, + rowsThisTile); + AscendC::PipeBarrier(); + AscendC::Exp(gkLastUbTensor, gkLastUbTensor, rowsThisTile); + AscendC::PipeBarrier(); + + ApplyRowScale(calcUbTensor, gkLastUbTensor, gkBrcbUbTensor, + rowsThisTile, nActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + } + + if (waitUpdateFromMte3) { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } + CopyGmToUb(hUpdateUbTensor, hUpdateInputThisTile, rowsThisTile, nActual, nActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::Add(hUpdateUbTensor, calcUbTensor, hUpdateUbTensor, rowsThisTile * nActual); + AscendC::PipeBarrier(); + + if constexpr(std::is_same::value) { + if (storeFinalState && isFinalState) { + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + CopyUbToGm(finalStateThisTile, hUpdateUbTensor, rowsThisTile, nActual, outputStride); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + waitUpdateFromMte3 = true; + } else { + AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + CopyUbToGm(hOutputThisTile, hUbTensor, rowsThisTile, nActual, outputStride); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + waitUpdateFromMte3 = false; + } + } else { + if (storeFinalState && isFinalState) { + AscendC::Cast(finalOutputUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + CopyUbToGm(finalStateThisTile, finalOutputUbTensor, rowsThisTile, nActual, outputStride); + } else { + AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + CopyUbToGm(hOutputThisTile, hUbTensor, rowsThisTile, nActual, outputStride); + } + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + waitUpdateFromMte3 = false; + } + } + + if (storeFinalState && isFinalState && std::is_same::value) { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } + if constexpr (kGated) { + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + } + + } + +private: + uint32_t pongBaseEvent = 4; + + AscendC::LocalTensor calcUbTensor; + + AscendC::LocalTensor hUpdateUbTensor_ping; + AscendC::LocalTensor hUbTensor_ping; + AscendC::LocalTensor finalOutputUbTensor_ping; + AscendC::LocalTensor glastUbTensor_ping; + + AscendC::LocalTensor hUpdateUbTensor_pong; + AscendC::LocalTensor hUbTensor_pong; + AscendC::LocalTensor finalOutputUbTensor_pong; + AscendC::LocalTensor glastUbTensor_pong; + + AscendC::LocalTensor gkLastUbTensor_ping; + AscendC::LocalTensor gkLastUbTensor_pong; + AscendC::LocalTensor gkInputUbTensor_ping; + AscendC::LocalTensor gkInputUbTensor_pong; + AscendC::LocalTensor gkBrcbUbTensor_ping; + AscendC::LocalTensor gkBrcbUbTensor_pong; +}; +} + +#endif diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp new file mode 100644 index 000000000000..696af333ae59 --- /dev/null +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp @@ -0,0 +1,508 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#ifndef CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_VNEW_HPP +#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_VNEW_HPP +#include "catlass/catlass.hpp" +#include "catlass/arch/resource.hpp" +#include "../gdn_fwd_h_epilogue_policies.hpp" +#include "catlass/gemm_coord.hpp" +#include "catlass/matrix_coord.hpp" +#include "catlass/epilogue/tile/tile_copy.hpp" + + + +namespace Catlass::Epilogue::Block { + +template < + class VOutputType_, + class GInputType_, + class UInputType_, + class WSInputType_, + class VUpdateType_, + class FinalStateType_, + class KGatedTag +> +class BlockEpilogue < + EpilogueAtlasGDNFwdHVnew, + VOutputType_, + GInputType_, + UInputType_, + WSInputType_, + VUpdateType_, + FinalStateType_, + KGatedTag +> { + static constexpr bool kGated = KGatedTag::value; +public: + using DispatchPolicy = EpilogueAtlasGDNFwdHVnew; + using ArchTag = typename DispatchPolicy::ArchTag; + + using VElementOutput = typename VOutputType_::Element; + using GElementInput = typename GInputType_::Element; + using UElementInput = typename UInputType_::Element; + using WSElementInput = typename WSInputType_::Element; + using VElementUpdate = typename VUpdateType_::Element; + using VLayoutUpdate = typename VUpdateType_::Layout; + using FinalStateElement = typename FinalStateType_::Element; + + CATLASS_DEVICE + BlockEpilogue(Arch::Resource &resource) + { + + constexpr uint32_t CALC_BUF_OFFSET = 0; + constexpr uint32_t PING_BUF_0_OFFSET = 32 * 1024; + constexpr uint32_t PING_BUF_1_OFFSET = 48 * 1024; + constexpr uint32_t PING_BUF_2_OFFSET = 64 * 1024; + constexpr uint32_t PING_BUF_3_OFFSET = 80 * 1024; + constexpr uint32_t PONG_BUF_0_OFFSET = 96 * 1024; + constexpr uint32_t PONG_BUF_1_OFFSET = 112 * 1024; + constexpr uint32_t PONG_BUF_2_OFFSET = 128 * 1024; + constexpr uint32_t PONG_BUF_3_OFFSET = 144 * 1024; + constexpr uint32_t PING_WIDE_IO_BUF_OFFSET = 16 * 1024; + constexpr uint32_t PONG_WIDE_IO_BUF_OFFSET = 24 * 1024; + constexpr uint32_t PING_G_BUF_OFFSET = 160 * 1024; + constexpr uint32_t PONG_G_BUF_OFFSET = 161 * 1024; + constexpr uint32_t PING_G_SUB_BUF_OFFSET = 162 * 1024; + constexpr uint32_t PONG_G_SUB_BUF_OFFSET = 163 * 1024; + constexpr uint32_t PING_G_INPUT_BUF_OFFSET = 164 * 1024; + constexpr uint32_t PONG_G_INPUT_BUF_OFFSET = 165 * 1024; + constexpr uint32_t SHARE_BUF_OFFSET = 166 * 1024; + + calcUbTensor = resource.ubBuf.template GetBufferByByte(CALC_BUF_OFFSET); + + uUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_2_OFFSET); + wsUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_0_OFFSET); + gUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_BUF_OFFSET); + gLastUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_SUB_BUF_OFFSET); + gInputUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_SUB_BUF_OFFSET); + vNewOutputUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_2_OFFSET); + vNewDecayUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_2_OFFSET); + wideIoUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_WIDE_IO_BUF_OFFSET); + + uUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_2_OFFSET); + wsUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_0_OFFSET); + gUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_BUF_OFFSET); + gLastUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_SUB_BUF_OFFSET); + gInputUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_SUB_BUF_OFFSET); + vNewOutputUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_2_OFFSET); + vNewDecayUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_2_OFFSET); + wideIoUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_WIDE_IO_BUF_OFFSET); + + gBrcbUbTensor_ = resource.ubBuf.template GetBufferByByte(SHARE_BUF_OFFSET); + + } + + CATLASS_DEVICE + ~BlockEpilogue() {} + + template + CATLASS_DEVICE + void CopyGmToUb( + AscendC::LocalTensor dst, + AscendC::GlobalTensor src, + uint32_t rows, + uint32_t cols, + uint32_t srcStride) + { + if (cols == srcStride) { + AscendC::DataCopy(dst, src, rows * cols); + return; + } + AscendC::DataCopyExtParams copyParams{ + static_cast(rows), + static_cast(cols * sizeof(Element)), + static_cast((srcStride - cols) * sizeof(Element)), + 0, + 0}; + AscendC::DataCopyPadExtParams padParams{false, 0, 0, 0}; + AscendC::DataCopyPad(dst, src, copyParams, padParams); + } + + template + CATLASS_DEVICE + void CopyUbToGm( + AscendC::GlobalTensor dst, + AscendC::LocalTensor src, + uint32_t rows, + uint32_t cols, + uint32_t dstStride) + { + if (cols == dstStride) { + AscendC::DataCopy(dst, src, rows * cols); + return; + } + AscendC::DataCopyExtParams copyParams{ + static_cast(rows), + static_cast(cols * sizeof(Element)), + 0, + static_cast((dstStride - cols) * sizeof(Element)), + 0}; + AscendC::DataCopyPad(dst, src, copyParams); + } + + CATLASS_DEVICE + void PrepareG( + AscendC::LocalTensor gUbTensor, + AscendC::LocalTensor gLastUbTensor, + AscendC::LocalTensor gInputUbTensor, + AscendC::GlobalTensor gInputThisSubBlock, + uint32_t mActual, + uint32_t pingpongFlag) + { + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + if (mActual == 1) { + AscendC::Duplicate(gUbTensor, 1.0f, 1); + AscendC::PipeBarrier(); + return; + } + if constexpr(std::is_same::value) { + AscendC::DataCopyParams gUbParams{1, (uint16_t)(mActual * sizeof(float)), 0, 0}; + AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0}; + AscendC::DataCopyPad(gUbTensor, gInputThisSubBlock, gUbParams, gUbPadParams); + } else { + AscendC::DataCopyParams gUbParams{1, (uint16_t)(mActual * sizeof(GElementInput)), 0, 0}; + AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0}; + AscendC::DataCopyPad(gInputUbTensor, gInputThisSubBlock, gUbParams, gUbPadParams); + } + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + if constexpr(!std::is_same::value) { + AscendC::Cast(gUbTensor, gInputUbTensor, AscendC::RoundMode::CAST_NONE, mActual); + AscendC::PipeBarrier(); + } + + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + float inputVal = gUbTensor.GetValue(mActual - 1); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + + AscendC::PipeBarrier(); + AscendC::Duplicate(gLastUbTensor, inputVal, mActual); + AscendC::PipeBarrier(); + + AscendC::Sub(gUbTensor, gLastUbTensor, gUbTensor, mActual); + AscendC::PipeBarrier(); + AscendC::Exp(gUbTensor, gUbTensor, mActual); + AscendC::PipeBarrier(); + } + + CATLASS_DEVICE + void ApplyRowScale( + AscendC::LocalTensor matrix, + AscendC::LocalTensor rowScale, + uint32_t rowScaleOffset, + uint32_t rows, + uint32_t cols) + { + constexpr uint32_t FP32_PER_BLOCK = 8; + constexpr uint32_t FP32_PER_REPEAT = 64; + uint8_t rowStride = static_cast(cols / FP32_PER_BLOCK); + AscendC::BinaryRepeatParams params(1, 1, 0, rowStride, rowStride, 1); + uint32_t localRow = 0; + while (localRow < rows) { + uint32_t scaleRow = rowScaleOffset + localRow; + uint32_t alignedScaleRow = scaleRow & ~(FP32_PER_BLOCK - 1); + uint32_t firstScaleLane = scaleRow - alignedScaleRow; + uint32_t rowsThisBlock = Min(FP32_PER_BLOCK - firstScaleLane, rows - localRow); + AscendC::Brcb(gBrcbUbTensor_, rowScale[alignedScaleRow], 1, {1, FP32_PER_BLOCK}); + AscendC::PipeBarrier(); + for (uint32_t col = 0; col < cols; col += FP32_PER_REPEAT) { + uint32_t count = Min(FP32_PER_REPEAT, cols - col); + uint32_t matrixOffset = localRow * cols + col; + AscendC::Mul(matrix[matrixOffset], matrix[matrixOffset], + gBrcbUbTensor_[firstScaleLane * FP32_PER_BLOCK], count, + rowsThisBlock, params); + } + AscendC::PipeBarrier(); + localRow += rowsThisBlock; + } + } + + CATLASS_DEVICE + void operator()( + AscendC::GlobalTensor vnewOutput, + AscendC::GlobalTensor vnewdecayOutput, + AscendC::LocalTensor l1VUpdate, + AscendC::GlobalTensor gInput, + AscendC::GlobalTensor uInput, + AscendC::GlobalTensor wsInput, + AscendC::GlobalTensor gkInput, + AscendC::GlobalTensor kInput, + AscendC::GlobalTensor kDecayWorkspace, + uint32_t chunkSize, + uint32_t kHeadDim, + uint32_t vBlockDim, + uint32_t vHeadDim, + Arch::CrossCoreFlag cube1Done, + Arch::CrossCoreFlag vec1Done, + bool isInitialState, + bool isFinalState, + bool storeFinalState, + bool waitWsFromMte3, + bool isPing + ) + { + static constexpr uint32_t ROW_TILE = 16; + uint32_t mActual = chunkSize; + uint32_t nvActual = vBlockDim; + uint32_t nkActual = kHeadDim; + uint32_t inputStride = vHeadDim; + + uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); + uint32_t subBlockNum = AscendC::GetSubBlockNum(); + uint32_t rowsPerSubBlock = CeilDiv(mActual, subBlockNum); + uint32_t rowBegin = subBlockIdx * rowsPerSubBlock; + uint32_t rowEnd = rowBegin + rowsPerSubBlock; + if (rowEnd > mActual) { + rowEnd = mActual; + } + if (rowBegin >= mActual) { + Arch::CrossCoreWaitFlag(cube1Done); + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + return; + } + AscendC::ResetMask(); + + AscendC::GlobalTensor gInputThisSubBlock = gInput; + + uint32_t pingpongFlag = isPing ? 0 : pongBaseEvent; + AscendC::LocalTensor uUbTensor = isPing ? uUbTensor_ping : uUbTensor_pong; + AscendC::LocalTensor wsUbTensor = isPing ? wsUbTensor_ping : wsUbTensor_pong; + AscendC::LocalTensor gUbTensor = isPing ? gUbTensor_ping : gUbTensor_pong; + AscendC::LocalTensor gLastUbTensor = isPing ? gLastUbTensor_ping : gLastUbTensor_pong; + AscendC::LocalTensor gInputUbTensor = isPing ? gInputUbTensor_ping : gInputUbTensor_pong; + AscendC::LocalTensor vNewOutputUbTensor = isPing ? vNewOutputUbTensor_ping : vNewOutputUbTensor_pong; + AscendC::LocalTensor vNewDecayUbTensor = isPing ? vNewDecayUbTensor_ping : vNewDecayUbTensor_pong; + if (nvActual > 128) { + uUbTensor = isPing ? wideIoUbTensor_ping : wideIoUbTensor_pong; + vNewOutputUbTensor = isPing ? wideIoUbTensor_ping : wideIoUbTensor_pong; + vNewDecayUbTensor = isPing ? wideIoUbTensor_ping : wideIoUbTensor_pong; + } + + if (rowBegin < rowEnd && nvActual <= 128 && nvActual == inputStride) { + uint32_t mActualThisSubBlock = rowEnd - rowBegin; + AscendC::GlobalTensor vnewOutputThisSubBlock = vnewOutput[rowBegin * inputStride]; + AscendC::GlobalTensor uInputThisSubBlock = uInput[rowBegin * inputStride]; + AscendC::GlobalTensor wsInputThisSubBlock = wsInput[rowBegin * nvActual]; + + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(uUbTensor, uInputThisSubBlock, mActualThisSubBlock * nvActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::Cast(calcUbTensor, uUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nvActual); + + PrepareG(gUbTensor, gLastUbTensor, gInputUbTensor, gInputThisSubBlock, mActual, pingpongFlag); + Arch::CrossCoreWaitFlag(cube1Done); + + if (waitWsFromMte3) { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } + AscendC::DataCopy(wsUbTensor, wsInputThisSubBlock, mActualThisSubBlock * nvActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + + AscendC::Sub(wsUbTensor, calcUbTensor, wsUbTensor, mActualThisSubBlock * nvActual); + AscendC::PipeBarrier(); + + AscendC::Copy(calcUbTensor, wsUbTensor, mActualThisSubBlock * nvActual); + AscendC::PipeBarrier(); + ApplyRowScale(calcUbTensor, gUbTensor, rowBegin, mActualThisSubBlock, nvActual); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + + uint32_t nvLoops = nvActual / FLOAT_NUM_PER_REPEAT; + for (uint32_t nLoop = 0; nLoop < nvLoops; nLoop++) { + uint32_t castSrcOffset = nLoop * FLOAT_NUM_PER_REPEAT; + uint32_t castDstOffset = nLoop * mActualThisSubBlock * FLOAT_NUM_PER_REPEAT; + AscendC::Cast(vNewDecayUbTensor[castDstOffset], calcUbTensor[castSrcOffset], AscendC::RoundMode::CAST_RINT, FLOAT_NUM_PER_REPEAT, mActualThisSubBlock, {(uint16_t)mActualThisSubBlock, 1, 1, (uint8_t)(nvLoops * 8)}); + } + + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + + uint32_t ubToL1Loops = nvActual / SIZE_16_NUM_PER_C0; + uint32_t mActualPadded = (mActual + NZ_BLOCK_SIZE - 1) / NZ_BLOCK_SIZE * NZ_BLOCK_SIZE; + AscendC::DataCopyParams intriParams; + intriParams.blockCount = ubToL1Loops; + intriParams.blockLen = mActualThisSubBlock; + intriParams.srcGap = 0; + intriParams.dstGap = mActualPadded - mActualThisSubBlock; + uint32_t l1Addr = rowBegin * SIZE_16_NUM_PER_C0; + AscendC::DataCopy(vnewdecayOutput[l1Addr], vNewDecayUbTensor, intriParams); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + if constexpr (!kGated) { + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + } + + AscendC::Cast(vNewOutputUbTensor, wsUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nvActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(vnewOutputThisSubBlock, vNewOutputUbTensor, mActualThisSubBlock * nvActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + + if constexpr (kGated) { + AscendC::GlobalTensor kInputThisSubBlock = kInput[rowBegin * nkActual]; + AscendC::GlobalTensor kDecayWorkspaceThisSubBlock = kDecayWorkspace[rowBegin * nkActual]; + // KDA passes kg = k * exp2(g_last - gk). Keep that decay exactly once. + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(vNewOutputUbTensor, kInputThisSubBlock, + mActualThisSubBlock * nkActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(kDecayWorkspaceThisSubBlock, vNewOutputUbTensor, + mActualThisSubBlock * nkActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + } + return; + } + + PrepareG(gUbTensor, gLastUbTensor, gInputUbTensor, gInputThisSubBlock, mActual, pingpongFlag); + Arch::CrossCoreWaitFlag(cube1Done); + + uint32_t mActualPadded = (mActual + NZ_BLOCK_SIZE - 1) / NZ_BLOCK_SIZE * NZ_BLOCK_SIZE; + bool waitWsThisTileFromMte3 = waitWsFromMte3; + for (uint32_t rowStart = rowBegin; rowStart < rowEnd;) { + uint32_t alignExtra = rowStart & 7; + uint32_t maxRowsThisTile = ROW_TILE - alignExtra; + uint32_t rowsThisTile = rowEnd - rowStart; + if (rowsThisTile > maxRowsThisTile) { + rowsThisTile = maxRowsThisTile; + } + uint32_t localRowStart = rowStart - rowBegin; + + AscendC::GlobalTensor vnewOutputThisTile = vnewOutput[rowStart * inputStride]; + AscendC::GlobalTensor uInputThisTile = uInput[rowStart * inputStride]; + AscendC::GlobalTensor wsInputThisTile = wsInput[rowStart * nvActual]; + AscendC::LocalTensor wsUbTensorThisTile = wsUbTensor[localRowStart * nvActual]; + + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + CopyGmToUb(uUbTensor, uInputThisTile, rowsThisTile, nvActual, inputStride); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::Cast(calcUbTensor, uUbTensor, AscendC::RoundMode::CAST_NONE, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + + if (waitWsThisTileFromMte3) { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + waitWsThisTileFromMte3 = false; + } else { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } + CopyGmToUb(wsUbTensorThisTile, wsInputThisTile, rowsThisTile, nvActual, nvActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::Sub(wsUbTensorThisTile, calcUbTensor, wsUbTensorThisTile, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + + AscendC::Copy(calcUbTensor, wsUbTensorThisTile, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + ApplyRowScale(calcUbTensor, gUbTensor, rowStart, rowsThisTile, nvActual); + + uint32_t nvLoops = nvActual / FLOAT_NUM_PER_REPEAT; + for (uint32_t nLoop = 0; nLoop < nvLoops; nLoop++) { + uint32_t castSrcOffset = nLoop * FLOAT_NUM_PER_REPEAT; + uint32_t castDstOffset = nLoop * rowsThisTile * FLOAT_NUM_PER_REPEAT; + AscendC::Cast(vNewDecayUbTensor[castDstOffset], calcUbTensor[castSrcOffset], AscendC::RoundMode::CAST_RINT, FLOAT_NUM_PER_REPEAT, rowsThisTile, {(uint16_t)rowsThisTile, 1, 1, (uint8_t)(nvLoops * 8)}); + } + + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + + uint32_t ubToL1Loops = nvActual / SIZE_16_NUM_PER_C0; + AscendC::DataCopyParams intriParams; + intriParams.blockCount = ubToL1Loops; + intriParams.blockLen = rowsThisTile; + intriParams.srcGap = 0; + intriParams.dstGap = mActualPadded - rowsThisTile; + uint32_t l1Addr = rowStart * SIZE_16_NUM_PER_C0; + AscendC::DataCopy(vnewdecayOutput[l1Addr], vNewDecayUbTensor, intriParams); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + if constexpr (!kGated) { + if (rowStart + rowsThisTile >= rowEnd) { + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + } + } + + AscendC::Cast(vNewOutputUbTensor, wsUbTensorThisTile, AscendC::RoundMode::CAST_RINT, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + CopyUbToGm(vnewOutputThisTile, vNewOutputUbTensor, rowsThisTile, nvActual, inputStride); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + rowStart += rowsThisTile; + } + + if constexpr (kGated) { + uint32_t mActualThisSubBlock = rowEnd - rowBegin; + for (uint32_t rowOffset = 0; rowOffset < mActualThisSubBlock; rowOffset += ROW_TILE) { + uint32_t rowsThisTile = Min(ROW_TILE, mActualThisSubBlock - rowOffset); + uint32_t row = rowBegin + rowOffset; + AscendC::GlobalTensor kInputThisTile = kInput[row * nkActual]; + AscendC::GlobalTensor kDecayWorkspaceThisTile = + kDecayWorkspace[row * nkActual]; + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + CopyGmToUb(vNewOutputUbTensor, kInputThisTile, rowsThisTile, nkActual, nkActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + CopyUbToGm(kDecayWorkspaceThisTile, vNewOutputUbTensor, + rowsThisTile, nkActual, nkActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + } + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + } + + if (rowBegin < rowEnd) { + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + } + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + + } + +private: + uint32_t pongBaseEvent = 4; + + AscendC::LocalTensor calcUbTensor; + + AscendC::LocalTensor uUbTensor_ping; + AscendC::LocalTensor wsUbTensor_ping; + AscendC::LocalTensor gUbTensor_ping; + AscendC::LocalTensor gLastUbTensor_ping; + AscendC::LocalTensor gInputUbTensor_ping; + AscendC::LocalTensor vNewOutputUbTensor_ping; + AscendC::LocalTensor vNewDecayUbTensor_ping; + + AscendC::LocalTensor uUbTensor_pong; + AscendC::LocalTensor wsUbTensor_pong; + AscendC::LocalTensor gUbTensor_pong; + AscendC::LocalTensor gLastUbTensor_pong; + AscendC::LocalTensor gInputUbTensor_pong; + AscendC::LocalTensor vNewOutputUbTensor_pong; + AscendC::LocalTensor vNewDecayUbTensor_pong; + AscendC::LocalTensor wideIoUbTensor_ping; + AscendC::LocalTensor wideIoUbTensor_pong; + + AscendC::LocalTensor gBrcbUbTensor_; + +}; +} + +#endif diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/gdn_fwd_h_epilogue_policies.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/gdn_fwd_h_epilogue_policies.hpp new file mode 100644 index 000000000000..640b2fbccbd3 --- /dev/null +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/gdn_fwd_h_epilogue_policies.hpp @@ -0,0 +1,27 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#ifndef CATLASS_EPILOGUE_GDN_FWD_H_EPILOGUE_POLICIES_HPP +#define CATLASS_EPILOGUE_GDN_FWD_H_EPILOGUE_POLICIES_HPP + +#include "catlass/catlass.hpp" + +namespace Catlass::Epilogue { + +struct EpilogueAtlasGDNFwdHVnew { + using ArchTag = Arch::Ascend950; +}; + +struct EpilogueAtlasGDNFwdHUpdate { + using ArchTag = Arch::Ascend950; +}; + +} // namespace Catlass::Epilogue + +#endif // CATLASS_EPILOGUE_GDN_FWD_H_EPILOGUE_POLICIES_HPP diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/block/block_scheduler_gdn_fwd_h.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/block/block_scheduler_gdn_fwd_h.hpp new file mode 100644 index 000000000000..e6e21fd1b355 --- /dev/null +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/block/block_scheduler_gdn_fwd_h.hpp @@ -0,0 +1,340 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#include "catlass/gemm_coord.hpp" +using namespace Catlass; + +#ifndef CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP +#define CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP + +// constexpr uint32_t PING_PONG_STAGES = 1; +constexpr uint32_t PING_PONG_STAGES = 2; +constexpr uint32_t BYTE_SIZE_16_BIT = 2; +constexpr uint32_t BYTES_PER_C0 = 32; +constexpr uint32_t BYTE_SIZE_PER_REPEAT = 256; +constexpr uint32_t SIZE_16_NUM_PER_C0 = BYTES_PER_C0 / BYTE_SIZE_16_BIT; +constexpr uint32_t FLOAT_NUM_PER_REPEAT = BYTE_SIZE_PER_REPEAT / sizeof(float); +constexpr uint32_t NZ_BLOCK_SIZE = 16; + +template +CATLASS_DEVICE T AlignUp(T a, T b) { + return (b == 0) ? 0 : (a + b - 1) / b * b; +} + +template +CATLASS_DEVICE T Min(T a, T b) { + return (a > b) ? b : a; +} + +template +CATLASS_DEVICE T Max(T a, T b) { + return (a > b) ? a : b; +} + +namespace Catlass::Gemm::Block { + +struct GDNFwdHOffsets { + uint32_t hSrcOffset; + uint32_t hDstOffset; + uint32_t uvOffset; + uint32_t wkOffset; + uint32_t wOffset; + uint32_t gOffset; + uint32_t gkOffset; + uint32_t hWorkOffset; + uint32_t vWorkOffset; + uint32_t kDecayWorkOffset; + uint32_t vBlockOffset; + uint32_t vBlockDim; + uint32_t initialStateOffset; + uint32_t finalStateOffset; + bool isInitialState; + bool isFinalState; + uint32_t blockTokens; + uint32_t streamId; + // for debug + uint32_t batchIdx; + uint32_t headIdx; + uint32_t chunkIdx; + +}; + +struct GDNFwdHStream { + uint32_t batchIdx; + uint32_t chunkIdx{0}; + uint32_t vHeadIdx; + uint32_t kHeadIdx; + uint32_t shapeBatchIdx; + uint32_t tokenBatchIdx; + + uint32_t chunkOffset; + uint32_t tokenOffset; + uint32_t batchChunks{0}; + uint32_t batchTokens; + uint32_t nextTaskIdx{0}; + bool active{false}; + + GDNFwdHOffsets offset; +}; + +struct GDNFwdHRunningQ { + GDNFwdHStream streams[PING_PONG_STAGES]; +}; + +struct BlockSchedulerGdnFwdH { + uint32_t batch; + uint32_t seqlen; + uint32_t kNumHead; + uint32_t vNumHead; + uint32_t kHeadDim; + uint32_t vHeadDim; + uint32_t chunkSize; + uint32_t vBlockSize{128}; + uint32_t isVariedLen; + uint32_t shapeBatch; + uint32_t tokenBatch; + bool useInitialState; + bool storeFinalState; + uint32_t numSeqWorkspaceOffset; + uint32_t numChunksWorkspaceOffset; + + uint32_t taskIdx; + uint32_t taskStride; + uint32_t cubeCoreIdx; + uint32_t cubeCoreNum; + uint32_t taskNum; + uint32_t headGroups; + uint32_t totalChunks; + uint32_t totalTokens; + + GDNFwdHRunningQ runningQ; + + bool isRunning; + + AscendC::GlobalTensor gmSeqlen; + AscendC::GlobalTensor gmNumSeq; + AscendC::GlobalTensor gmNumChunks; + + Arch::CrossCoreFlag cube1Done[PING_PONG_STAGES] = {0, 1}; + Arch::CrossCoreFlag vec1Done[PING_PONG_STAGES] = {2, 3}; + Arch::CrossCoreFlag cube2Done[PING_PONG_STAGES] = {4, 5}; + Arch::CrossCoreFlag vec2Done[PING_PONG_STAGES] = {6, 7}; + + CATLASS_DEVICE + BlockSchedulerGdnFwdH() {} + + CATLASS_DEVICE + void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR tiling, GM_ADDR user, uint32_t coreIdx, uint32_t coreNum) { + __gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict gdnFwdHTilingData = reinterpret_cast<__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict>(tiling); + + batch = gdnFwdHTilingData->batch; + seqlen = gdnFwdHTilingData->seqlen; + kNumHead = gdnFwdHTilingData->kNumHead; + vNumHead = gdnFwdHTilingData->vNumHead; + kHeadDim = gdnFwdHTilingData->kHeadDim; + vHeadDim = gdnFwdHTilingData->vHeadDim; + chunkSize = gdnFwdHTilingData->chunkSize; + isVariedLen = gdnFwdHTilingData->isVariedLen; + shapeBatch = gdnFwdHTilingData->shapeBatch; + tokenBatch = gdnFwdHTilingData->tokenBatch; + useInitialState = gdnFwdHTilingData->useInitialState; + storeFinalState = gdnFwdHTilingData->storeFinalState; + numSeqWorkspaceOffset = gdnFwdHTilingData->numSeqWorkspaceOffset; + numChunksWorkspaceOffset = gdnFwdHTilingData->numChunksWorkspaceOffset; + + gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens); + gmNumSeq.SetGlobalBuffer((__gm__ int64_t *)(user + numSeqWorkspaceOffset)); + gmNumChunks.SetGlobalBuffer((__gm__ int64_t *)(user + numChunksWorkspaceOffset)); + + if (isVariedLen) { + gmNumChunks.SetValue(0, 0); + gmNumSeq.SetValue(0, 0); + uint32_t actualBatch = 0; + int64_t prevSeq = 0, currSeq; + for (uint32_t b = 1; b <= tokenBatch; b++) { + currSeq = gmSeqlen.GetValue(b); + int64_t batchSeqLen = currSeq - prevSeq; + if (batchSeqLen > 0) { + actualBatch++; + gmNumSeq.SetValue(actualBatch, currSeq); + int64_t batchChunk = (batchSeqLen + chunkSize - 1) / chunkSize; + gmNumChunks.SetValue(actualBatch, gmNumChunks.GetValue(actualBatch - 1) + batchChunk); + } + prevSeq = currSeq; + } + tokenBatch = actualBatch; + batch = actualBatch; + totalChunks = gmNumChunks.GetValue(tokenBatch); + totalTokens = gmNumSeq.GetValue(tokenBatch); + } else { + totalChunks = (seqlen + chunkSize - 1) / chunkSize; + totalTokens = seqlen; + } + + cubeCoreIdx = coreIdx; + cubeCoreNum = coreNum; + vBlockSize = vHeadDim; + taskNum = batch * vNumHead; + headGroups = vNumHead / kNumHead; + uint32_t maxTaskCntPerLoop = taskNum > cubeCoreNum ? PING_PONG_STAGES : 1; + taskStride = cubeCoreNum * maxTaskCntPerLoop; + for (uint32_t streamId = 0; streamId < PING_PONG_STAGES; ++streamId) { + auto& stream = runningQ.streams[streamId]; + stream.nextTaskIdx = (streamId < maxTaskCntPerLoop) ? (cubeCoreIdx * maxTaskCntPerLoop + streamId) : taskNum; + stream.chunkIdx = 0; + stream.batchChunks = 0; + stream.active = false; + } + isRunning = cubeCoreIdx * maxTaskCntPerLoop < taskNum; + + } + + + CATLASS_DEVICE + void InitNewStream(GDNFwdHStream& newStream) { + newStream.batchIdx = taskIdx / vNumHead; + newStream.vHeadIdx = taskIdx % vNumHead; + newStream.kHeadIdx = newStream.vHeadIdx / headGroups; + newStream.shapeBatchIdx = isVariedLen ? 0 : newStream.batchIdx; + newStream.tokenBatchIdx = isVariedLen ? newStream.batchIdx : 0; + newStream.chunkOffset = isVariedLen ? gmNumChunks.GetValue(newStream.tokenBatchIdx) : 0; + newStream.batchChunks = isVariedLen ? (gmNumChunks.GetValue(newStream.tokenBatchIdx + 1) - newStream.chunkOffset) : totalChunks; + newStream.tokenOffset = isVariedLen ? gmNumSeq.GetValue(newStream.tokenBatchIdx) : 0; + newStream.batchTokens = isVariedLen ? (gmNumSeq.GetValue(newStream.tokenBatchIdx + 1) - newStream.tokenOffset) : totalTokens; + newStream.chunkIdx = 0; + } + + CATLASS_DEVICE + void AssignNextStream(uint32_t streamId) { + auto& stream = runningQ.streams[streamId]; + taskIdx = stream.nextTaskIdx; + if (taskIdx >= taskNum) { + stream.active = false; + stream.batchChunks = 0; + return; + } + + stream.nextTaskIdx += taskStride; + InitNewStream(stream); + stream.active = stream.batchChunks > 0; + if (stream.active) { + UpdateTask(streamId); + } + } + + CATLASS_DEVICE + void UpdateTask(uint32_t streamId) { + auto& stream = runningQ.streams[streamId]; + auto& offset = stream.offset; + + offset.isInitialState = stream.chunkIdx == 0; + offset.isFinalState = stream.chunkIdx == (stream.batchChunks - 1); + uint32_t vBlockOffset = 0; + uint32_t vBlockDim = vBlockSize; + offset.initialStateOffset = (stream.batchIdx * vNumHead + stream.vHeadIdx) * kHeadDim * vHeadDim + vBlockOffset; + offset.finalStateOffset = (stream.batchIdx * vNumHead + stream.vHeadIdx) * kHeadDim * vHeadDim + vBlockOffset; + offset.hSrcOffset = (stream.shapeBatchIdx * vNumHead * totalChunks + stream.vHeadIdx * totalChunks + stream.chunkOffset + stream.chunkIdx) * kHeadDim * vHeadDim + vBlockOffset; + offset.hDstOffset = offset.hSrcOffset + kHeadDim * vHeadDim; + if (storeFinalState && offset.isFinalState) { + offset.hDstOffset = offset.hSrcOffset; + } + offset.uvOffset = (stream.shapeBatchIdx * vNumHead * totalTokens + stream.vHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize) * vHeadDim + vBlockOffset; + offset.wkOffset = (stream.shapeBatchIdx * kNumHead * totalTokens + stream.kHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize) * kHeadDim; + offset.wOffset = (stream.shapeBatchIdx * vNumHead * totalTokens + stream.vHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize) * kHeadDim; + offset.gOffset = stream.shapeBatchIdx * vNumHead * totalTokens + stream.vHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize; + offset.gkOffset = (stream.shapeBatchIdx * vNumHead * totalTokens + stream.vHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize) * kHeadDim; + offset.hWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + streamId) * kHeadDim * vBlockSize; + offset.vWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + streamId) * chunkSize * vBlockSize; + offset.kDecayWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + streamId) * chunkSize * kHeadDim; + offset.vBlockOffset = vBlockOffset; + offset.vBlockDim = vBlockDim; + offset.blockTokens = offset.isFinalState ? (stream.batchTokens - stream.chunkIdx * chunkSize) : chunkSize; + offset.streamId = streamId; + offset.batchIdx = stream.batchIdx; + offset.headIdx = stream.vHeadIdx; + offset.chunkIdx = stream.chunkIdx; + } + + CATLASS_DEVICE + void InitTasks() { + isRunning = false; + for (uint32_t streamId = 0; streamId < PING_PONG_STAGES; ++streamId) { + auto& stream = runningQ.streams[streamId]; + if (stream.active) { + stream.chunkIdx += 1; + if (stream.chunkIdx >= stream.batchChunks) { + stream.active = false; + stream.batchChunks = 0; + } + } + if (!stream.active) { + AssignNextStream(streamId); + } else { + UpdateTask(streamId); + } + if (stream.active) { + isRunning = true; + } + } + } + + CATLASS_DEVICE + const GDNFwdHStream& GetStream(uint32_t i) const { + return runningQ.streams[i]; + } + + CATLASS_DEVICE + uint32_t GetStreamId(uint32_t i) const { + return i; + } + + CATLASS_DEVICE + const GDNFwdHOffsets& GetCurTaskOffsets(const GDNFwdHStream& stream) const { + return stream.offset; + } + + CATLASS_DEVICE + bool StreamIsDone(const GDNFwdHStream& stream) const { + return !stream.active; + } + + CATLASS_DEVICE + bool NeedProcessStage2(const GDNFwdHStream& stream) { + return storeFinalState || !stream.offset.isFinalState; + } +}; + +struct BlockSchedulerGdnFwdHCube : public BlockSchedulerGdnFwdH { + CATLASS_DEVICE + BlockSchedulerGdnFwdHCube() {} + + CATLASS_DEVICE + void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR tiling, GM_ADDR user) { + BlockSchedulerGdnFwdH::Init(cu_seqlens, chunk_indices, tiling, user, AscendC::GetBlockIdx(), AscendC::GetBlockNum()); + } + +}; + +struct BlockSchedulerGdnFwdHVec : public BlockSchedulerGdnFwdH { + CATLASS_DEVICE + BlockSchedulerGdnFwdHVec() {} + + CATLASS_DEVICE + void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR tiling, GM_ADDR user) { + BlockSchedulerGdnFwdH::Init( + cu_seqlens, chunk_indices, tiling, user, + AscendC::GetBlockIdx() / AscendC::GetSubBlockNum(), + AscendC::GetBlockNum()); + } + +}; + +} // namespace Catlass::Gemm::Block + +#endif // CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/kernel/gdn_fwd_h_kernel.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/kernel/gdn_fwd_h_kernel.hpp new file mode 100644 index 000000000000..03919be1d948 --- /dev/null +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/kernel/gdn_fwd_h_kernel.hpp @@ -0,0 +1,609 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#define CATLASS_ARCH 3510 + +#include "catlass/arch/arch.hpp" +#include "catlass/arch/cross_core_sync.hpp" +#include "catlass/arch/resource.hpp" +#include "catlass/catlass.hpp" +#include "catlass/debug.hpp" +#include "../block/block_scheduler_gdn_fwd_h.hpp" +#include "catlass/epilogue/block/block_epilogue.hpp" +#include "../../epilogue/block/block_epilogue_gdn_fwdh_update.hpp" +#include "../../epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp" +#include "catlass/gemm/block/block_mmad.hpp" +#include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp" +#include "kernel_utils/block/block_mmad_pingpong_tla_preloadA_l1B.hpp" +#include "catlass/gemm/block/block_swizzle.hpp" +#include "catlass/gemm/dispatch_policy.hpp" +#include "catlass/gemm/gemm_type.hpp" +#include "catlass/layout/layout.hpp" +#include "catlass/gemm_coord.hpp" +#include "tla/tensor.hpp" +#include "tla/layout.hpp" +#include "tla/tensor.hpp" + +using _0 = tla::Int<0>; +using _1 = tla::Int<1>; +using _2 = tla::Int<2>; +using _4 = tla::Int<4>; +using _8 = tla::Int<8>; +using _16 = tla::Int<16>; +using _32 = tla::Int<32>; +using _64 = tla::Int<64>; +using _128 = tla::Int<128>; +using _256 = tla::Int<256>; +using _512 = tla::Int<512>; +using _1024 = tla::Int<1024>; +using _2048 = tla::Int<2048>; +using _4096 = tla::Int<4096>; +using _8192 = tla::Int<8192>; +using _16384 = tla::Int<16384>; +using _32768 = tla::Int<32768>; +using _65536 = tla::Int<65536>; + + +#include "kernel_operator.h" +using namespace Catlass; +using namespace tla; + +namespace Catlass::Gemm::Kernel { + +struct GDNFwdHTileShapes128 { + using L1TileShape = Shape<_128, _128, _128>; + using L0TileShape = L1TileShape; +}; + +struct GDNFwdHTileShapes256 { + using L1TileShape = Shape<_128, _256, _128>; + using L0TileShape = Shape<_128, _256, _64>; +}; + +template< + typename INPUT_TYPE, + typename G_TYPE, + typename STATE_TYPE, + typename WORKSPACE_TYPE, + typename TileShapes = GDNFwdHTileShapes128, + bool kGated = false +> +class GDNFwdHKernel { +public: + + using ArchTag = Arch::Ascend950; + using CubeScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdHCube; + using VecScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdHVec; + + using DispatchPolicyTlaMulti = Gemm::MmadPingpongTlaMulti; + using DispatchPolicyTlaPreloadAL1B = Gemm::MmadPingpongTlaPreloadAL1B; + using L1TileShapeVTla = typename TileShapes::L1TileShape; + using L0TileShapeVTla = typename TileShapes::L0TileShape; + + using WType = Gemm::GemmType; + using HType = Gemm::GemmType; + using VworkType = Gemm::GemmType; + using KType = Gemm::GemmType; + using HworkType = Gemm::GemmType; + using VType = Gemm::GemmType; + using GType = Gemm::GemmType; + using UType = Gemm::GemmType; + using FinalStateType = Gemm::GemmType; + using VUpdateType = Gemm::GemmType; + + // cube 1 + using TileCopyWH = Catlass::Gemm::Tile::PackedTileCopyTla; + using BlockMmadWH = Gemm::Block::BlockMmadTla; + + // cube 2 + using TileCopyKV = Catlass::Gemm::Tile::PackedTileCopyTla; + using TileMmadKV = Gemm::Tile::TileMmadTla; + using BlockMmadKV = Gemm::Block::BlockMmadTla; + + // vec 1 + using DispatchPolicyGDNFwdHVnew = Epilogue::EpilogueAtlasGDNFwdHVnew; + using EpilogueGDNFwdHVnew = Epilogue::Block::BlockEpilogue>; + + // vec 2 + using DispatchPolicyGDNFwdHUpdate = Epilogue::EpilogueAtlasGDNFwdHUpdate; + using EpilogueGDNFwdHUpdate = Epilogue::Block::BlockEpilogue>; + + using GDNFwdHOffsets = Catlass::Gemm::Block::GDNFwdHOffsets; + + using ElementK = INPUT_TYPE; + using ElementW = INPUT_TYPE; + using ElementU = INPUT_TYPE; + using ElementG = G_TYPE; + using ElementH = INPUT_TYPE; + using ElementV = INPUT_TYPE; + using ElementVUpdate = INPUT_TYPE; + using ElementVWork = WORKSPACE_TYPE; + using ElementHWork = WORKSPACE_TYPE; + using ElementInitialState = STATE_TYPE; + using ElementFinalState = STATE_TYPE; + + using LayoutW = Catlass::layout::RowMajor; + using LayoutH = Catlass::layout::RowMajor; + using LayoutV = Catlass::layout::RowMajor; + using LayoutK = Catlass::layout::ColumnMajor; + using LayoutVUpdate = typename VUpdateType::Layout; + + + uint32_t batch; + uint32_t seqlen; + uint32_t kNumHead; + uint32_t vNumHead; + uint32_t kHeadDim; + uint32_t vHeadDim; + uint32_t chunkSize; + bool useInitialState; + bool storeFinalState; + uint32_t isVariedLen; + uint32_t shapeBatch; + uint32_t tokenBatch; + uint32_t vWorkspaceOffset; + uint32_t vUpdateWorkspaceOffset; + uint32_t hWorkspaceOffset; + uint32_t numSeqWorkspaceOffset; + uint32_t numChunksWorkspaceOffset; + uint32_t kDecayWorkspaceOffset; + + AscendC::GlobalTensor gmK; + AscendC::GlobalTensor gmW; + AscendC::GlobalTensor gmU; + AscendC::GlobalTensor gmG; + AscendC::GlobalTensor gmGk; + AscendC::GlobalTensor gmInitialState; + AscendC::GlobalTensor gmH; + AscendC::GlobalTensor gmV; + AscendC::GlobalTensor gmFinalState; + AscendC::GlobalTensor gmVWorkspace; + AscendC::GlobalTensor gmVUpdateWorkspace; + AscendC::GlobalTensor gmHWorkspace; + AscendC::GlobalTensor gmKDecayWorkspace; + + AscendC::GlobalTensor gmSeqlen; + AscendC::GlobalTensor gmNumSeq; + AscendC::GlobalTensor gmNumChunks; + + AscendC::LocalTensor ubHUpdatePing; + AscendC::LocalTensor ubHUpdatePong; + AscendC::LocalTensor ubVWorkPing; + AscendC::LocalTensor ubVWorkPong; + + AscendC::LocalTensor l1VUpdatePing; + AscendC::LocalTensor l1VUpdatePong; + + CubeScheduler cubeBlockScheduler; + VecScheduler vecBlockScheduler; + + Arch::Resource resource; + + + __aicore__ inline GDNFwdHKernel() {} + + __aicore__ inline void Init(GM_ADDR k, GM_ADDR w, GM_ADDR u, GM_ADDR g, GM_ADDR gk, GM_ADDR inital_state, GM_ADDR cu_seqlens, GM_ADDR chunk_indices, + GM_ADDR h, GM_ADDR v_new, GM_ADDR final_state, GM_ADDR tiling, GM_ADDR user) { + + __gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict gdnFwdHTilingData = reinterpret_cast<__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict>(tiling); + + batch = gdnFwdHTilingData->batch; + seqlen = gdnFwdHTilingData->seqlen; + kNumHead = gdnFwdHTilingData->kNumHead; + vNumHead = gdnFwdHTilingData->vNumHead; + kHeadDim = gdnFwdHTilingData->kHeadDim; + vHeadDim = gdnFwdHTilingData->vHeadDim; + chunkSize = gdnFwdHTilingData->chunkSize; + useInitialState = gdnFwdHTilingData->useInitialState; + storeFinalState = gdnFwdHTilingData->storeFinalState; + isVariedLen = gdnFwdHTilingData->isVariedLen; + shapeBatch = gdnFwdHTilingData->shapeBatch; + tokenBatch = gdnFwdHTilingData->tokenBatch; + vWorkspaceOffset = gdnFwdHTilingData->vWorkspaceOffset; + vUpdateWorkspaceOffset = gdnFwdHTilingData->vUpdateWorkspaceOffset; + hWorkspaceOffset = gdnFwdHTilingData->hWorkspaceOffset; + numSeqWorkspaceOffset = gdnFwdHTilingData->numSeqWorkspaceOffset; + numChunksWorkspaceOffset = gdnFwdHTilingData->numChunksWorkspaceOffset; + kDecayWorkspaceOffset = gdnFwdHTilingData->kDecayWorkspaceOffset; + + gmK.SetGlobalBuffer((__gm__ ElementK *)k); + gmW.SetGlobalBuffer((__gm__ ElementW *)w); + gmU.SetGlobalBuffer((__gm__ ElementU *)u); + gmG.SetGlobalBuffer((__gm__ ElementG *)g); + gmGk.SetGlobalBuffer((__gm__ ElementG *)gk); + gmInitialState.SetGlobalBuffer((__gm__ ElementInitialState *)inital_state); + gmH.SetGlobalBuffer((__gm__ ElementH *)h); + gmV.SetGlobalBuffer((__gm__ ElementV *)v_new); + gmFinalState.SetGlobalBuffer((__gm__ ElementFinalState *)final_state); + gmVWorkspace.SetGlobalBuffer((__gm__ ElementVWork *)(user + vWorkspaceOffset)); + gmVUpdateWorkspace.SetGlobalBuffer((__gm__ ElementV *)(user + vUpdateWorkspaceOffset)); + gmHWorkspace.SetGlobalBuffer((__gm__ ElementHWork *)(user + hWorkspaceOffset)); + gmKDecayWorkspace.SetGlobalBuffer((__gm__ ElementK *)(user + kDecayWorkspaceOffset)); + + gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens); + gmNumSeq.SetGlobalBuffer((__gm__ int64_t *)(user + numSeqWorkspaceOffset)); + gmNumChunks.SetGlobalBuffer((__gm__ int64_t *)(user + numChunksWorkspaceOffset)); + + ubHUpdatePing = resource.ubBuf.template GetBufferByByte(32 * 1024); + ubHUpdatePong = resource.ubBuf.template GetBufferByByte(96 * 1024); + ubVWorkPing = resource.ubBuf.template GetBufferByByte(32 * 1024); + ubVWorkPong = resource.ubBuf.template GetBufferByByte(96 * 1024); + + l1VUpdatePing = resource.l1Buf.template GetBufferByByte(0); + l1VUpdatePong = resource.l1Buf.template GetBufferByByte(chunkSize * vHeadDim * sizeof(ElementV)); + + if ASCEND_IS_AIC { + cubeBlockScheduler.Init(cu_seqlens, chunk_indices, tiling, user); + } + + if ASCEND_IS_AIV { + vecBlockScheduler.Init(cu_seqlens, chunk_indices, tiling, user); + } + } + + __aicore__ inline void Process() { + if (isVariedLen) { + AscendC::SyncAll(); + } + + if ASCEND_IS_AIC { + uint32_t coreIdx = AscendC::GetBlockIdx(); + uint32_t coreNum = AscendC::GetBlockNum(); + + BlockMmadWH blockMmadWH(resource, chunkSize * cubeBlockScheduler.vBlockSize * sizeof(ElementV) * PING_PONG_STAGES); + BlockMmadKV blockMmadKV(resource, chunkSize * cubeBlockScheduler.vBlockSize * sizeof(ElementV) * PING_PONG_STAGES); + + auto wLayout = tla::MakeLayout(shapeBatch * kNumHead * cubeBlockScheduler.totalTokens, kHeadDim); + auto hLayout = tla::MakeLayout(shapeBatch * vNumHead * cubeBlockScheduler.totalChunks * kHeadDim, vHeadDim); + + auto kLayout = tla::MakeLayout(kHeadDim, shapeBatch * kNumHead * cubeBlockScheduler.totalTokens); + auto hworkLayout = tla::MakeLayout(kHeadDim, cubeBlockScheduler.vBlockSize); + + AscendC::SyncAll(); + uint32_t currStage = 0; // 0: C1, 1: C2 + while (cubeBlockScheduler.isRunning) { + if (currStage == 0) { + /* C1: v_work = w @ h[i] */ + cubeBlockScheduler.InitTasks(); + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = cubeBlockScheduler.GetStreamId(i); + const auto& stream = cubeBlockScheduler.GetStream(i); + if (cubeBlockScheduler.StreamIsDone(stream)) { + continue; + } + + const GDNFwdHOffsets& cube1Offsets = cubeBlockScheduler.GetCurTaskOffsets(stream); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[streamId]); + auto vLayout = tla::MakeLayout(cube1Offsets.blockTokens, cube1Offsets.vBlockDim); + int64_t cube1OffsetW = cube1Offsets.wOffset; + int64_t cube1OffsetH = cube1Offsets.hSrcOffset; + int64_t cube1OffsetVwork = cube1Offsets.vWorkOffset; + auto tensorW = tla::MakeTensor(gmW[cube1OffsetW], wLayout, Catlass::Arch::PositionGM{}); + auto tensorH = tla::MakeTensor(gmH[cube1OffsetH], hLayout, Catlass::Arch::PositionGM{}); + auto tensorV = tla::MakeTensor(gmVWorkspace[cube1OffsetVwork], vLayout, Catlass::Arch::PositionGM{}); + GemmCoord cube1Shape {cube1Offsets.blockTokens, cube1Offsets.vBlockDim, kHeadDim}; + auto tensorBlockW = GetTile(tensorW, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k())); + auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n())); + auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n())); + + blockMmadWH.preSetFlags(); + blockMmadWH(tensorBlockW, tensorBlockH, tensorBlockV, cube1Shape); + blockMmadWH.finalWaitFlags(); + Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube1Done[streamId]); + } + } else { + /* C2: h[i+1] = k.T @ v_work */ + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = cubeBlockScheduler.GetStreamId(i); + const auto& stream = cubeBlockScheduler.GetStream(i); + if (cubeBlockScheduler.StreamIsDone(stream)) { + continue; + } + const GDNFwdHOffsets& cube2Offsets = cubeBlockScheduler.GetCurTaskOffsets(stream); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done[streamId]); + + if (cubeBlockScheduler.NeedProcessStage2(stream)) { + // step 3: h[i+1] = k.T @ v_work + int64_t cube2OffsetK = kGated ? cube2Offsets.kDecayWorkOffset : cube2Offsets.wkOffset; + int64_t cube2OffsetVwork = cube2Offsets.vWorkOffset; + auto tensorK = kGated + ? tla::MakeTensor(gmKDecayWorkspace[cube2OffsetK], kLayout, Catlass::Arch::PositionGM{}) + : tla::MakeTensor(gmK[cube2OffsetK], kLayout, Catlass::Arch::PositionGM{}); + auto vUpdateLayout = tla::MakeLayout(cube2Offsets.blockTokens, cube2Offsets.vBlockDim); + auto tensorVwork = tla::MakeTensor(gmVUpdateWorkspace[cube2OffsetVwork], vUpdateLayout, Catlass::Arch::PositionGM{}); + auto tensorHwork = tla::MakeTensor(gmHWorkspace[cube2Offsets.hWorkOffset], hworkLayout, Catlass::Arch::PositionGM{}); + GemmCoord cube2Shape{kHeadDim, cube2Offsets.vBlockDim, cube2Offsets.blockTokens}; + auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k())); + auto tensorBlockVwork = GetTile(tensorVwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n())); + auto tensorBlockHwork = GetTile(tensorHwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n())); + + blockMmadKV.preSetFlags(); + blockMmadKV(tensorBlockK, tensorBlockVwork, tensorBlockHwork, cube2Shape); + blockMmadKV.finalWaitFlags(); + } + Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube2Done[streamId]); + } + } + currStage ^= 0x01; + } + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[0]); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[1]); + + } + + if ASCEND_IS_AIV { + uint32_t coreIdx = AscendC::GetBlockIdx(); + uint32_t coreNum = AscendC::GetBlockNum(); + + if (useInitialState) { + AscendC::LocalTensor stateUbTensorPing = resource.ubBuf.template GetBufferByByte(0); + AscendC::LocalTensor stateUbTensorPong = resource.ubBuf.template GetBufferByByte(96 * 1024); + AscendC::LocalTensor hUbTensorPing = resource.ubBuf.template GetBufferByByte(64 * 1024); + AscendC::LocalTensor hUbTensorPong = resource.ubBuf.template GetBufferByByte(160 * 1024); + uint32_t totalChunks = isVariedLen ? vecBlockScheduler.totalChunks : ((seqlen + chunkSize - 1) / chunkSize); + uint32_t transferCount = isVariedLen ? (vecBlockScheduler.tokenBatch * vNumHead / coreNum) : (shapeBatch * vNumHead / coreNum); + uint32_t remainderFlag = isVariedLen ? (((vecBlockScheduler.tokenBatch * vNumHead) % coreNum) != 0): (((shapeBatch * vNumHead) % coreNum) != 0); + uint32_t step = transferCount + remainderFlag; + uint32_t stateBlockSize = kHeadDim * vHeadDim; + uint32_t pingpongFlag = 1; + uint32_t start = coreIdx * step; + uint32_t end = start + step; + uint32_t maxLimit = isVariedLen ? vecBlockScheduler.tokenBatch * vNumHead : shapeBatch * vNumHead; + uint32_t realEnd = min(end, maxLimit); + AscendC::SetFlag(EVENT_ID0); + AscendC::SetFlag(EVENT_ID1); + for (uint32_t initialStateBlockOffset = start; initialStateBlockOffset >= start && initialStateBlockOffset < realEnd; initialStateBlockOffset++) { + uint32_t batchIdx = initialStateBlockOffset / vNumHead; + uint32_t vHeadIdx = initialStateBlockOffset % vNumHead; + uint32_t chunkOffset = isVariedLen ? gmNumChunks.GetValue(batchIdx) : 0; + uint32_t initialStateBaseOffset = initialStateBlockOffset * stateBlockSize; + uint32_t shapeBatchIdx = isVariedLen ? 0 : batchIdx; + uint32_t hBaseOffset = (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset) * stateBlockSize; + if (vHeadDim <= 128) { + AscendC::LocalTensor stateUbTensor = pingpongFlag ? stateUbTensorPing : stateUbTensorPong; + AscendC::LocalTensor hUbTensor = pingpongFlag ? hUbTensorPing : hUbTensorPong; + auto event_id = pingpongFlag ? EVENT_ID1 : EVENT_ID0; + AscendC::WaitFlag(event_id); + if constexpr(!std::is_same::value) { + AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateBaseOffset], stateBlockSize); + AscendC::SetFlag(event_id); + AscendC::WaitFlag(event_id); + AscendC::Cast(hUbTensor, stateUbTensor, AscendC::RoundMode::CAST_RINT, stateBlockSize); + AscendC::SetFlag(event_id); + AscendC::WaitFlag(event_id); + AscendC::DataCopy(gmH[hBaseOffset], hUbTensor, stateBlockSize); + } else { + AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateBaseOffset], stateBlockSize); + AscendC::SetFlag(event_id); + AscendC::WaitFlag(event_id); + AscendC::DataCopy(gmH[hBaseOffset], stateUbTensor, stateBlockSize); + } + AscendC::SetFlag(event_id); + pingpongFlag = 1 - pingpongFlag; + } else { + uint32_t stateRowTile = (32 * 1024) / (vHeadDim * sizeof(ElementH)); + for (uint32_t rowOffset = 0; rowOffset < kHeadDim; rowOffset += stateRowTile) { + uint32_t rowsThisTile = Min(stateRowTile, kHeadDim - rowOffset); + uint32_t stateTileElems = rowsThisTile * vHeadDim; + uint32_t initialStateOffset = initialStateBaseOffset + rowOffset * vHeadDim; + uint32_t hOffset = hBaseOffset + rowOffset * vHeadDim; + AscendC::LocalTensor stateUbTensor = pingpongFlag ? stateUbTensorPing : stateUbTensorPong; + AscendC::LocalTensor hUbTensor = pingpongFlag ? hUbTensorPing : hUbTensorPong; + auto event_id = pingpongFlag ? EVENT_ID1 : EVENT_ID0; + AscendC::WaitFlag(event_id); + if constexpr(!std::is_same::value) { + AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateOffset], stateTileElems); + AscendC::SetFlag(event_id); + AscendC::WaitFlag(event_id); + AscendC::Cast(hUbTensor, stateUbTensor, AscendC::RoundMode::CAST_RINT, stateTileElems); + AscendC::SetFlag(event_id); + AscendC::WaitFlag(event_id); + AscendC::DataCopy(gmH[hOffset], hUbTensor, stateTileElems); + } else { + AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateOffset], stateTileElems); + AscendC::SetFlag(event_id); + AscendC::WaitFlag(event_id); + AscendC::DataCopy(gmH[hOffset], stateUbTensor, stateTileElems); + } + AscendC::SetFlag(event_id); + pingpongFlag = 1 - pingpongFlag; + } + } + } + + AscendC::WaitFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID1); + } else { + uint32_t stateRowTile = (32 * 1024) / (vHeadDim * sizeof(ElementH)); + AscendC::LocalTensor chunkOffsetsUb = + resource.ubBuf.template GetBufferByByte(0); + if (isVariedLen) { + uint32_t chunkOffsetBytes = (vecBlockScheduler.tokenBatch + 1) * sizeof(int64_t); + AscendC::DataCopyParams copyParams{ + 1, static_cast(chunkOffsetBytes), 0, 0}; + AscendC::DataCopyPadParams padParams{false, 0, 0, 0}; + AscendC::DataCopyPad(chunkOffsetsUb, gmNumChunks[0], copyParams, padParams); + AscendC::SetFlag(EVENT_ID3); + AscendC::WaitFlag(EVENT_ID3); + AscendC::SetFlag(EVENT_ID3); + AscendC::WaitFlag(EVENT_ID3); + } + auto chunkOffsets = reinterpret_cast<__ubuf__ int64_t *>(chunkOffsetsUb.GetPhyAddr()); + AscendC::LocalTensor hUbTensorPing = + resource.ubBuf.template GetBufferByByte(64 * 1024); + AscendC::LocalTensor hUbTensorPong = + resource.ubBuf.template GetBufferByByte(160 * 1024); + uint32_t totalChunks = + isVariedLen ? vecBlockScheduler.totalChunks : ((seqlen + chunkSize - 1) / chunkSize); + uint32_t taskCount = + (isVariedLen ? vecBlockScheduler.tokenBatch : shapeBatch) * vNumHead; + uint32_t step = taskCount / coreNum + ((taskCount % coreNum) != 0); + uint32_t start = coreIdx * step; + uint32_t realEnd = Min(start + step, taskCount); + uint32_t pingpongFlag = 1; + AscendC::SetFlag(EVENT_ID0); + AscendC::SetFlag(EVENT_ID1); + for (uint32_t taskIdx = start; taskIdx < realEnd; ++taskIdx) { + uint32_t batchIdx = taskIdx / vNumHead; + uint32_t vHeadIdx = taskIdx % vNumHead; + uint32_t chunkOffset = isVariedLen ? static_cast(chunkOffsets[batchIdx]) : 0; + uint32_t shapeBatchIdx = isVariedLen ? 0 : batchIdx; + uint32_t hBaseOffset = + (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset) * + kHeadDim * vHeadDim; + for (uint32_t rowOffset = 0; rowOffset < kHeadDim; rowOffset += stateRowTile) { + uint32_t rowsThisTile = Min(stateRowTile, kHeadDim - rowOffset); + uint32_t stateTileElems = rowsThisTile * vHeadDim; + AscendC::LocalTensor hUbTensor = + pingpongFlag ? hUbTensorPing : hUbTensorPong; + auto eventId = pingpongFlag ? EVENT_ID1 : EVENT_ID0; + AscendC::WaitFlag(eventId); + AscendC::Duplicate(hUbTensor, static_cast(0), stateTileElems); + AscendC::SetFlag(eventId); + AscendC::WaitFlag(eventId); + AscendC::DataCopy(gmH[hBaseOffset + rowOffset * vHeadDim], hUbTensor, stateTileElems); + AscendC::SetFlag(eventId); + pingpongFlag = 1 - pingpongFlag; + } + } + AscendC::WaitFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID1); + } + + AscendC::SyncAll(); + + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done[0]); + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done[1]); + + EpilogueGDNFwdHVnew epilogueGDNFwdHVnew(resource); + EpilogueGDNFwdHUpdate epilogueGDNFwdHUpdate(resource); + uint32_t pongBaseEvent = 4; + + if (storeFinalState && std::is_same::value) { + AscendC::SetFlag(EVENT_ID0); // preset final_state + AscendC::SetFlag(EVENT_ID0 + pongBaseEvent); + AscendC::SetFlag(EVENT_ID2); // preset h + AscendC::SetFlag(EVENT_ID2 + pongBaseEvent); + } else { + AscendC::SetFlag(EVENT_ID0); // preset h_update + AscendC::SetFlag(EVENT_ID0 + pongBaseEvent); + AscendC::SetFlag(EVENT_ID2); // preset h + AscendC::SetFlag(EVENT_ID2 + pongBaseEvent); + } + AscendC::SetFlag(EVENT_ID1); // preset u + AscendC::SetFlag(EVENT_ID1 + pongBaseEvent); + AscendC::SetFlag(EVENT_ID3); // preset g + AscendC::SetFlag(EVENT_ID3 + pongBaseEvent); + uint32_t currStage = 0; // 0: V1, 1: V2 + bool event0FromMte3[PING_PONG_STAGES] = {false, false}; + bool event2FromMte3[PING_PONG_STAGES] = {!(storeFinalState && std::is_same::value), + !(storeFinalState && std::is_same::value)}; + while (vecBlockScheduler.isRunning) { + if (currStage == 0) { + /* V1: + * gmV = gmU - gmVWorkspace + * g_buf = gmG[-1] - gmG + * g_buf = exp(g_buf) + * gmVWorkspace = g_buf * gmV + */ + vecBlockScheduler.InitTasks(); + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = vecBlockScheduler.GetStreamId(i); + const auto& stream = vecBlockScheduler.GetStream(i); + if (vecBlockScheduler.StreamIsDone(stream)) { + continue; + } + const GDNFwdHOffsets& vec1Offsets = vecBlockScheduler.GetCurTaskOffsets(stream); + AscendC::LocalTensor l1VUpdate = (i == 0) ? l1VUpdatePing : l1VUpdatePong; + bool waitWsFromMte3 = storeFinalState && std::is_same::value && + event0FromMte3[streamId]; + epilogueGDNFwdHVnew( + gmV[vec1Offsets.uvOffset], gmVUpdateWorkspace[vec1Offsets.vWorkOffset], l1VUpdate, + gmG[vec1Offsets.gOffset], gmU[vec1Offsets.uvOffset], gmVWorkspace[vec1Offsets.vWorkOffset], + gmGk[vec1Offsets.gkOffset], gmK[vec1Offsets.wkOffset], gmKDecayWorkspace[vec1Offsets.kDecayWorkOffset], + vec1Offsets.blockTokens, kHeadDim, vec1Offsets.vBlockDim, vHeadDim, + vecBlockScheduler.cube1Done[streamId], vecBlockScheduler.vec1Done[streamId], + vec1Offsets.isInitialState, vec1Offsets.isFinalState, storeFinalState, + waitWsFromMte3, (i == 0) + ); + if (storeFinalState && std::is_same::value) { + event0FromMte3[streamId] = false; + } + } + } else { + /* V2: h[i+1] += h_work if i < num_chunks - 1 else None */ + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = vecBlockScheduler.GetStreamId(i); + const auto& stream = vecBlockScheduler.GetStream(i); + if (vecBlockScheduler.StreamIsDone(stream)) { + continue; + } + const GDNFwdHOffsets& vec2Offsets = vecBlockScheduler.GetCurTaskOffsets(stream); + if (vecBlockScheduler.NeedProcessStage2(stream)) { + if (storeFinalState && std::is_same::value) { + event0FromMte3[streamId] = vec2Offsets.isFinalState; + event2FromMte3[streamId] = !vec2Offsets.isFinalState; + } + // step 4: h[i+1] += h_work if i < num_chunks - 1 else None + epilogueGDNFwdHUpdate( + gmH[vec2Offsets.hDstOffset], gmFinalState[vec2Offsets.finalStateOffset], + gmG[vec2Offsets.gOffset], + gmH[vec2Offsets.hSrcOffset], + gmHWorkspace[vec2Offsets.hWorkOffset], + gmGk[vec2Offsets.gkOffset], + vec2Offsets.blockTokens, kHeadDim, vec2Offsets.vBlockDim, vHeadDim, vecBlockScheduler.cube2Done[streamId], + vec2Offsets.isInitialState, vec2Offsets.isFinalState, storeFinalState, (i == 0) + ); + } else { + Arch::CrossCoreWaitFlag(vecBlockScheduler.cube2Done[streamId]); + } + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done[streamId]); + } + } + currStage ^= 0x01; + } + + if (storeFinalState && std::is_same::value) { + if (event0FromMte3[0]) { + AscendC::WaitFlag(EVENT_ID0); + } else { + AscendC::WaitFlag(EVENT_ID0); + } + if (event0FromMte3[1]) { + AscendC::WaitFlag(EVENT_ID0 + pongBaseEvent); + } else { + AscendC::WaitFlag(EVENT_ID0 + pongBaseEvent); + } + if (event2FromMte3[0]) { + AscendC::WaitFlag(EVENT_ID2); + } else { + AscendC::WaitFlag(EVENT_ID2); + } + if (event2FromMte3[1]) { + AscendC::WaitFlag(EVENT_ID2 + pongBaseEvent); + } else { + AscendC::WaitFlag(EVENT_ID2 + pongBaseEvent); + } + } else { + AscendC::WaitFlag(EVENT_ID0); // preset h_update + AscendC::WaitFlag(EVENT_ID0 + pongBaseEvent); + AscendC::WaitFlag(EVENT_ID2); // preset h + AscendC::WaitFlag(EVENT_ID2 + pongBaseEvent); + } + AscendC::WaitFlag(EVENT_ID1); // preset u + AscendC::WaitFlag(EVENT_ID1 + pongBaseEvent); + AscendC::WaitFlag(EVENT_ID3); // preset g + AscendC::WaitFlag(EVENT_ID3 + pongBaseEvent); + + } + } + +}; + +} diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/chunk_gated_delta_rule_fwd_h.cpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/chunk_gated_delta_rule_fwd_h.cpp index 9167f2bcd458..d40bf3bbbb24 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/chunk_gated_delta_rule_fwd_h.cpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/chunk_gated_delta_rule_fwd_h.cpp @@ -1,38 +1,46 @@ /** - * Copyright (c) 2026 Tianjin University, Ltd. - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of - * the BSD 3-Clause License (the "License"). - * Please refer to the License for details. You may not use this file except in compliance with the License. - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. - */ + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ /*! * \file chunk_gated_delta_rule_fwd_h.cpp * \brief */ -// #include "chunk_gated_delta_rule_fwd_h.h" #if defined(__CCE_AICORE__) && (__CCE_AICORE__ == 200) #include "arch20/compat_310p.h" #include "arch20/gemm/kernel/gdn_fwd_h_kernel.hpp" +#elif defined(__CCE_AICORE__) && (__CCE_AICORE__ == 310) +#include "arch35/gemm/kernel/gdn_fwd_h_kernel.hpp" #else #include "arch22/gemm/kernel/gdn_fwd_h_kernel.hpp" #endif + #include "lib/matmul_intf.h" using namespace Catlass; +#if defined(__CCE_AICORE__) && (__CCE_AICORE__ == 200) +// 310P keeps the arch20 kernel ABI: no TileShapes/kGated templates and no gk path. +// The host still passes the unified gk argument; arch20 ignores it. extern "C" __global__ __aicore__ void chunk_gated_delta_rule_fwd_h(GM_ADDR k, GM_ADDR w, GM_ADDR u, GM_ADDR g, - GM_ADDR inital_state, GM_ADDR cu_seqlens, GM_ADDR chunk_indices, - GM_ADDR h, GM_ADDR v_new, GM_ADDR final_state, - GM_ADDR workspace, GM_ADDR tiling) + GM_ADDR gk, GM_ADDR inital_state, GM_ADDR cu_seqlens, + GM_ADDR chunk_indices, GM_ADDR h, GM_ADDR v_new, + GM_ADDR final_state, GM_ADDR workspace, GM_ADDR tiling) { + (void)gk; + // 310P binary codegen cannot register non-default tiling keys here; host forces key 0. KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2); GM_ADDR user = AscendC::GetUserWorkspace(workspace); - __gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict gdnFwdHTilingData = reinterpret_cast<__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict>(tiling); + __gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict gdnFwdHTilingData = + reinterpret_cast<__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict>(tiling); using workspaceType = float; // dtype: 0 - fp16, 1 - bf16, 2 - fp32 @@ -93,3 +101,128 @@ extern "C" __global__ __aicore__ void chunk_gated_delta_rule_fwd_h(GM_ADDR k, GM } } } +#else +namespace GDN { + +template +__aicore__ inline void ChunkGatedDeltaRuleFwdHKernelImpl(GM_ADDR k, GM_ADDR w, GM_ADDR u, GM_ADDR g, GM_ADDR gk, + GM_ADDR inital_state, GM_ADDR cu_seqlens, + GM_ADDR chunk_indices, GM_ADDR h, GM_ADDR v_new, + GM_ADDR final_state, GM_ADDR tiling, GM_ADDR user) +{ + using GDNFwdHKernel = Catlass::Gemm::Kernel::GDNFwdHKernel; + GDNFwdHKernel gdnFwdH; + gdnFwdH.Init(k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + gdnFwdH.Process(); +} + +template +__aicore__ inline void ChunkGatedDeltaRuleFwdHDispatch(GM_ADDR k, GM_ADDR w, GM_ADDR u, GM_ADDR g, GM_ADDR gk, + GM_ADDR inital_state, GM_ADDR cu_seqlens, + GM_ADDR chunk_indices, GM_ADDR h, GM_ADDR v_new, + GM_ADDR final_state, GM_ADDR tiling, GM_ADDR user) +{ + __gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict gdnFwdHTilingData = + reinterpret_cast<__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict>(tiling); + using WorkspaceT = float; + bool useGk = gdnFwdHTilingData->useGk; + // dtype: 0 - fp16, 1 - bf16, 2 - fp32 + if (gdnFwdHTilingData->dataType == 1) { + if (gdnFwdHTilingData->stateDataType == 2) { + if (gdnFwdHTilingData->gDataType == 2) { + if (useGk) { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } else { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } + } else { + if (useGk) { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } else { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } + } + } else { + if (gdnFwdHTilingData->gDataType == 2) { + if (useGk) { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } else { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } + } else { + if (useGk) { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } else { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } + } + } + } else { + if (gdnFwdHTilingData->stateDataType == 2) { + if (gdnFwdHTilingData->gDataType == 2) { + if (useGk) { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } else { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } + } else { + if (useGk) { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } else { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } + } + } else { + if (gdnFwdHTilingData->gDataType == 2) { + if (useGk) { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } else { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } + } else { + if (useGk) { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } else { + ChunkGatedDeltaRuleFwdHKernelImpl( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } + } + } + } +} + +} // namespace GDN + +extern "C" __global__ __aicore__ void chunk_gated_delta_rule_fwd_h(GM_ADDR k, GM_ADDR w, GM_ADDR u, GM_ADDR g, + GM_ADDR gk, GM_ADDR inital_state, GM_ADDR cu_seqlens, GM_ADDR chunk_indices, + GM_ADDR h, GM_ADDR v_new, GM_ADDR final_state, + GM_ADDR workspace, GM_ADDR tiling) +{ + GM_ADDR user = AscendC::GetUserWorkspace(workspace); + + if (TILING_KEY_IS(1)) { + KERNEL_TASK_TYPE(1, KERNEL_TYPE_MIX_AIC_1_2); + GDN::ChunkGatedDeltaRuleFwdHDispatch( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } else if (TILING_KEY_IS(2)) { + KERNEL_TASK_TYPE(2, KERNEL_TYPE_MIX_AIC_1_2); + GDN::ChunkGatedDeltaRuleFwdHDispatch( + k, w, u, g, gk, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user); + } +} +#endif diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/chunk_gated_delta_rule_fwd_h_struct.h b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/chunk_gated_delta_rule_fwd_h_struct.h new file mode 100644 index 000000000000..42e56289355c --- /dev/null +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/chunk_gated_delta_rule_fwd_h_struct.h @@ -0,0 +1,53 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +/*! + * \file chunk_gated_delta_rule_fwd_h_struct.h + * \brief Plain tiling data struct for chunk_gated_delta_rule_fwd_h. + * + * The aclnn/ascendc framework auto-generates a kernel-side tiling struct named + * ChunkGatedDeltaRuleFwdHTilingData (global scope) from the BEGIN_TILING_DATA_DEF + * macro in chunk_gated_delta_rule_fwd_h_tiling.h. The fast kernel launch extension + * compiles the kernel standalone (without that auto-generated header), so it provides + * the same plain struct here. The field order/types mirror the macro definition so the + * kernel/scheduler can read it the same way, and so the struct can be passed by value + * to the kernel via the <<<>>> launch (its address is a valid GM_ADDR on Atlas A2). + */ + +#ifndef CHUNK_GATED_DELTA_RULE_FWD_H_STRUCT_H +#define CHUNK_GATED_DELTA_RULE_FWD_H_STRUCT_H + +#include + +struct ChunkGatedDeltaRuleFwdHTilingData { + int64_t batch; + int64_t seqlen; + int64_t kNumHead; + int64_t vNumHead; + int64_t kHeadDim; + int64_t vHeadDim; + int64_t chunkSize; + bool useInitialState; + bool storeFinalState; + int64_t dataType; + int64_t gDataType; + int64_t stateDataType; + int64_t isVariedLen; + int64_t shapeBatch; + int64_t tokenBatch; + bool useGk; + int64_t vWorkspaceOffset; + int64_t vUpdateWorkspaceOffset; + int64_t kDecayWorkspaceOffset; + int64_t hWorkspaceOffset; + int64_t numSeqWorkspaceOffset; + int64_t numChunksWorkspaceOffset; +}; + +#endif // CHUNK_GATED_DELTA_RULE_FWD_H_STRUCT_H diff --git a/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla_preloadA_l1B.hpp b/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla_preloadA_l1B.hpp new file mode 100644 index 000000000000..4be3bfac91d6 --- /dev/null +++ b/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla_preloadA_l1B.hpp @@ -0,0 +1,532 @@ +/** + * Copyright (c) 2025 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#ifndef CATLASS_GEMM_BLOCK_BLOCK_MMAD_PINGPONG_TLA_PRELOADA_L1B_HPP +#define CATLASS_GEMM_BLOCK_BLOCK_MMAD_PINGPONG_TLA_PRELOADA_L1B_HPP + + +#include "catlass/catlass.hpp" +#include "catlass/arch/resource.hpp" +#include "catlass/coord.hpp" +#include "catlass/gemm_coord.hpp" +#include "catlass/gemm/dispatch_policy.hpp" +#include "catlass/gemm/helper.hpp" +#include "catlass/gemm/tile/tile_copy.hpp" +#include "catlass/gemm/tile/tile_mmad.hpp" +#include "tla/layout.hpp" +#include "tla/tensor.hpp" + +namespace Catlass::Gemm { + +template +struct MmadPingpongTlaPreloadAL1B : public MmadBase { + static constexpr uint32_t L1A_STAGES = L1A_STAGES_; + static constexpr uint32_t L1B_STAGES = L1B_STAGES_; + static constexpr uint32_t L0A_STAGES = L0A_STAGES_; + static constexpr uint32_t L0B_STAGES = L0B_STAGES_; + static constexpr uint32_t L0C_STAGES = L0C_STAGES_; + static constexpr bool ENABLE_UNIT_FLAG = ENABLE_UNIT_FLAG_; +}; + +} // namespace Catlass::Gemm + +namespace Catlass::Gemm::Block { + +template < + class ArchTag_, + bool ENABLE_UNIT_FLAG_, + uint32_t L0C_STAGES_, + uint32_t L1A_STAGES_, + uint32_t L1B_STAGES_, + uint32_t L0A_STAGES_, + uint32_t L0B_STAGES_, + class L1TileShape_, + class L0TileShape_, + class ElementA_, + class ElementB_, + class ElementC_, + class ElementBias_, + class TileCopy_, + class TileMmad_ +> +struct BlockMmadTla < + MmadPingpongTlaPreloadAL1B, + L1TileShape_, + L0TileShape_, + ElementA_, + ElementB_, + ElementC_, + ElementBias_, + TileCopy_, + TileMmad_ +> { +public: + // Type Aliases + using DispatchPolicy = MmadPingpongTlaPreloadAL1B; + using ArchTag = typename DispatchPolicy::ArchTag; + using TileCopy = TileCopy_; + using L1TileShape = L1TileShape_; + using L0TileShape = L0TileShape_; + using ElementA = ElementA_; + using LayoutA = typename TileCopy::LayoutA; + using ElementB = ElementB_; + using LayoutB = typename TileCopy::LayoutB; + using ElementC = ElementC_; + using LayoutC = typename TileCopy::LayoutC; + + using TileMmad = TileMmad_; + + using CopyL1ToL0A = typename TileCopy::CopyL1ToL0A; + using CopyL1ToL0B = typename TileCopy::CopyL1ToL0B; + using CopyL1ToBT = typename TileCopy::CopyL1ToBT; + + using ElementAccumulator = typename TileCopy::ElementAccumulator; + + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + using LayoutTagL0A = typename TileCopy::LayoutTagL0A; + using LayoutTagL0B = typename TileCopy::LayoutTagL0B; + + using L1AAlignHelper = typename TileCopy_::L1AAlignHelper; + using L1BAlignHelper = typename TileCopy_::L1BAlignHelper; + + static_assert(tla::is_tuple::value && tla::is_static::value, + "L1TileShape must be tla::tuple and static!"); + static_assert(tla::is_tuple::value && tla::is_static::value, + "L0TileShape must be tla::tuple and static!"); + + static constexpr bool ENABLE_UNIT_FLAG = DispatchPolicy::ENABLE_UNIT_FLAG; + static constexpr uint32_t L1A_STAGES = DispatchPolicy::L1A_STAGES; + static constexpr uint32_t L1B_STAGES = DispatchPolicy::L1B_STAGES; + static constexpr uint32_t L0A_STAGES = DispatchPolicy::L0A_STAGES; + static constexpr uint32_t L0B_STAGES = DispatchPolicy::L0B_STAGES; + static constexpr uint32_t L0C_STAGES = DispatchPolicy::L0C_STAGES; + static constexpr uint32_t L1_TILE_M = tla::get<0>(L1TileShape{}); + static constexpr uint32_t L1_TILE_N = tla::get<1>(L1TileShape{}); + static constexpr uint32_t L1_TILE_K = tla::get<2>(L1TileShape{}); + static constexpr uint32_t L0_TILE_M = tla::get<0>(L0TileShape{}); + static constexpr uint32_t L0_TILE_N = tla::get<1>(L0TileShape{}); + static constexpr uint32_t L0_TILE_K = tla::get<2>(L0TileShape{}); + + // L1 tile size + static constexpr uint32_t L1A_TILE_SIZE = L1_TILE_M * L1_TILE_K * sizeof(ElementA); + static constexpr uint32_t L1B_TILE_SIZE = L1_TILE_N * L1_TILE_K * sizeof(ElementB); + // L0 tile size + static constexpr uint32_t L0A_TILE_SIZE = L0_TILE_M * L0_TILE_K * sizeof(ElementA); + static constexpr uint32_t L0B_TILE_SIZE = L0_TILE_K * L0_TILE_N * sizeof(ElementB); + static constexpr uint32_t L0C_TILE_SIZE = L1_TILE_M * L1_TILE_N * sizeof(ElementAccumulator); + + // Check L0C_STAGES + static_assert(!(ENABLE_UNIT_FLAG && L0C_STAGES != 1), "L0C_STAGES must be 1 when UnitFlag is true!"); + + // Check LayoutC + static_assert(tla::detail::isRowMajor::value || + ((std::is_same_v || std::is_same_v || + std::is_same_v) && tla::detail::iszN::value), + "LayoutC only supports zN in half or bfloat16 or float, RowMajor in all dtype yet!"); + + // Check L1TileShape + static_assert(L1A_TILE_SIZE * L1A_STAGES + L1B_TILE_SIZE * L1B_STAGES <= ArchTag::L1_SIZE, + "L1TileShape exceeding the L1 space!"); + + // Check L0TileShape + static_assert(L0A_TILE_SIZE * L0A_STAGES <= ArchTag::L0A_SIZE, "L0TileShape exceeding the L0A space!"); + static_assert(L0B_TILE_SIZE * L0B_STAGES <= ArchTag::L0B_SIZE, "L0TileShape exceeding the L0B space!"); + static_assert(L0C_TILE_SIZE * L0C_STAGES <= ArchTag::L0C_SIZE, "L0TileShape exceeding the L0C space!"); + + static constexpr uint32_t _32B = 32*8; // in bits + static_assert(L1_TILE_M == L0_TILE_M && L1_TILE_N == L0_TILE_N, + "The situation where the basic blocks of L1 and L0 differ on the m and n axes is not supported yet"); + static_assert(L0_TILE_K <= L1_TILE_K, "L0TileShape::K cannot exceed L1TileShape::K"); +#if (defined (CATLASS_ARCH) && CATLASS_ARCH == 2201) + static_assert(L1_TILE_M * SizeOfBits::value % _32B == 0, "L1TileShape::M must be 32B aligned."); + static_assert(L1_TILE_K * SizeOfBits::value % _32B == 0, "L1TileShape::K must be 32B aligned."); + static_assert(L1_TILE_K * SizeOfBits::value % _32B == 0, "L1TileShape::K must be 32B aligned."); + static_assert(L1_TILE_N * SizeOfBits::value % _32B == 0, "L1TileShape::N must be 32B aligned."); + static_assert(L0_TILE_K * SizeOfBits::value % _32B == 0, "L0TileShape::K must be 32B aligned."); +#endif + + static constexpr auto L1A_LAYOUT = + tla::MakeLayout(tla::Int{}, tla::Int{}); + static constexpr auto L1B_LAYOUT = + tla::MakeLayout(tla::Int{}, tla::Int{}); + static constexpr auto L1BIAS_LAYOUT = tla::MakeLayout(tla::Int{}); + static constexpr auto L0BIAS_LAYOUT = tla::MakeLayout(tla::Int{}); + + // When enabling L1 resident mode, restore the pointer and coordinates that record the last state + // to the initial state. if two blockmmad instances need to be consecutively invoked at the kernel layer, + // RestoreStatus() must be inserted between them. + CATLASS_DEVICE + void RestoreStatus() + { + for (int i = 0; i < L1A_STAGES; ++i) { + lastAddrA[i] = nullptr; + lastCoordA[i] = MatrixCoord{0U, 0U}; + } + for (int i = 0; i < L1B_STAGES; ++i) { + lastAddrB[i] = nullptr; + lastCoordB[i] = MatrixCoord{0U, 0U}; + } + } + + /// Construct + CATLASS_DEVICE + BlockMmadTla(Arch::Resource &resource, uint32_t l1BufAddrStart = 0) + { + if ASCEND_IS_AIC { + uint32_t l1AOffset = l1BufAddrStart; + uint32_t l1BOffset = l1BufAddrStart + L1A_TILE_SIZE * L1A_STAGES; + // Init buffers + for (uint32_t i = 0; i < L1A_STAGES; i++) { + // Assign L1/L0A/L0B space for each stages + l1ATensorList[i] = resource.l1Buf.template GetBufferByByte(l1AOffset + L1A_TILE_SIZE * i); + // Assign event ID for each stages + l1AEventList[i] = i; + } + for (uint32_t i = 0; i < L1B_STAGES; i++) { + // Assign L1/L0A/L0B space for each stages + l1BTensorList[i] = resource.l1Buf.template GetBufferByByte(l1BOffset + L1B_TILE_SIZE * i); + // Assign event ID for each stages + l1BEventList[i] = i + L1A_STAGES; + } + for (uint32_t i = 0; i < L0A_STAGES; i++) { + // Assign L1/L0A/L0B space for each stages + l0ATensorList[i] = resource.l0ABuf.template GetBufferByByte(L0A_TILE_SIZE * i); + // Assign event ID for each stages + l0AEventList[i] = i; + } + for (uint32_t i = 0; i < L0B_STAGES; i++) { + // Assign L1/L0A/L0B space for each stages + l0BTensorList[i] = resource.l0BBuf.template GetBufferByByte(L0B_TILE_SIZE * i); + // Assign event ID for each stages + l0BEventList[i] = i + L0A_STAGES; + } + if constexpr(!ENABLE_UNIT_FLAG) { + for (uint32_t i = 0; i < L0C_STAGES; i++) { + l0CTensorList[i] = resource.l0CBuf.template GetBufferByByte(L0C_TILE_SIZE * i); + l0CEventList[i] = i; + } + } else { + l0CTensorList[0] = resource.l0CBuf.template GetBufferByByte(0); + } + } + } + + /// Destructor + CATLASS_DEVICE + ~BlockMmadTla() {} + + CATLASS_DEVICE + void preSetFlags() { + + if ASCEND_IS_AIC { + if constexpr (ENABLE_UNIT_FLAG && tla::detail::isRowMajor::value) { + AscendC::SetMMLayoutTransform(true); + } + for (uint32_t i = 0; i < L1A_STAGES; i++) { + AscendC::SetFlag(l1AEventList[i]); + } + for (uint32_t i = 0; i < L1B_STAGES; i++) { + AscendC::SetFlag(l1BEventList[i]); + } + for (uint32_t i = 0; i < L0A_STAGES; i++) { + AscendC::SetFlag(l0AEventList[i]); + } + for (uint32_t i = 0; i < L0B_STAGES; i++) { + AscendC::SetFlag(l0BEventList[i]); + } + if constexpr(!ENABLE_UNIT_FLAG) { + for (uint32_t i = 0; i < L0C_STAGES; i++) { + AscendC::SetFlag(l0CEventList[i]); + } + } + } + } + + CATLASS_DEVICE + void finalWaitFlags() { + if ASCEND_IS_AIC { + if constexpr (ENABLE_UNIT_FLAG && tla::detail::isRowMajor::value) { + AscendC::SetMMLayoutTransform(false); + } + for (uint32_t i = 0; i < L1A_STAGES; i++) { + AscendC::WaitFlag(l1AEventList[i]); + } + for (uint32_t i = 0; i < L1B_STAGES; i++) { + AscendC::WaitFlag(l1BEventList[i]); + } + for (uint32_t i = 0; i < L0A_STAGES; i++) { + AscendC::WaitFlag(l0AEventList[i]); + } + for (uint32_t i = 0; i < L0B_STAGES; i++) { + AscendC::WaitFlag(l0BEventList[i]); + } + if constexpr(!ENABLE_UNIT_FLAG) { + for (uint32_t i = 0; i < L0C_STAGES; i++) { + AscendC::WaitFlag(l0CEventList[i]); + } + } + } + } + + /// Perform a block-scoped matrix multiply-accumulate + template + CATLASS_DEVICE void operator()(TensorA &tensorA, TensorB &tensorB, TensorC &tensorC, GemmCoord const &actualShape) + { + using CopyGmToL1A = typename TileCopy_::template CopyGmToL1A; + // using CopyGmToL1B = typename TileCopy_::template CopyGmToL1B; + CopyGmToL1A copyGmToL1A; + // CopyGmToL1B copyGmToL1B; +#if (defined (CATLASS_ARCH) && CATLASS_ARCH == 2201) + using CopyL0CToGm = typename TileCopy_::template CopyL0CToGm; + CopyL0CToGm copyL0CToDst; +#endif +#if (defined (CATLASS_ARCH) && CATLASS_ARCH == 3510) + using CopyL0CToDst = typename TileCopy_::template CopyL0CToDst; + CopyL0CToDst copyL0CToDst; +#endif + + uint32_t mBlockActual = actualShape.m(); + uint32_t kBlockActual = actualShape.k(); + uint32_t nBlockActual = actualShape.n(); + + uint32_t mL1Actual = mBlockActual; + if constexpr (std::is_same_v) { + // Avoid using the gemv mode in mmad + if (mL1Actual == 1) { + mL1Actual = 16; + } + } + uint32_t nL1Actual = nBlockActual; + + auto layoutInL0C = tla::MakeLayoutL0C(mL1Actual, nL1Actual); + auto tensorL0C = tla::MakeTensor(l0CTensorList[l0CListId], layoutInL0C, Arch::PositionL0C{}); + auto tensorL0Bias = tla::MakeTensor(l0BiasTensor, L0BIAS_LAYOUT, Arch::PositionBias{}); + + uint32_t kL1Actual = min(kBlockActual, L1_TILE_K); + // load first matrix A tile from GM to L1 + AscendC::WaitFlag(l1AEventList[l1AListId]); + auto tensorL1A = tla::MakeTensor(l1ATensorList[l1AListId], L1A_LAYOUT, Arch::PositionL1{}); + auto tensorTileA = GetTileA(tensorA, 0, 0, mBlockActual, kL1Actual); + copyGmToL1A(tensorL1A, tensorTileA); + AscendC::SetFlag(l1AEventList[l1AListId]); + + // load first matrix B tile from GM to L1 + AscendC::WaitFlag(l1BEventList[l1BListId]); + // auto tensorL1B = tla::MakeTensor(l1BTensorList[l1BListId], L1B_LAYOUT, Arch::PositionL1{}); + // auto tensorTileB = GetTile(tensorB, tla::MakeCoord(0, 0), tla::MakeShape(kL1Actual, nBlockActual)); + // copyGmToL1B(tensorL1B, tensorTileB); + AscendC::SetFlag(l1BEventList[l1BListId]); + + if constexpr (!ENABLE_UNIT_FLAG) { + AscendC::WaitFlag(l0CEventList[l0CListId]); + } + + uint32_t mL0Loop = CeilDiv(mL1Actual); + uint32_t nL0Loop = CeilDiv(nL1Actual); + + // main loop + uint32_t kL1Loop = CeilDiv(kBlockActual); + for (uint32_t kL1Idx = 0; kL1Idx < kL1Loop; kL1Idx++) { + uint32_t l1AListIdNext = (l1AListId + 1 < L1A_STAGES) ? (l1AListId + 1) : 0; + uint32_t l1BListIdNext = (l1BListId + 1 < L1B_STAGES) ? (l1BListId + 1) : 0; + uint32_t kL1ActualNext{0}; + // preload next tile from GM to L1 + if (kL1Idx < kL1Loop - 1) { + uint32_t kL1IdxNext = kL1Idx + 1; + kL1ActualNext = (kL1IdxNext < kL1Loop - 1) ? L1_TILE_K : (kBlockActual - kL1IdxNext * L1_TILE_K); + + // Get L1 tensor for next stage + auto l1ATensor = l1ATensorList[l1AListIdNext]; + auto l1BTensor = l1BTensorList[l1BListIdNext]; + auto tensorL1A = tla::MakeTensor(l1ATensor, L1A_LAYOUT, Arch::PositionL1{}); + auto tensorL1B = tla::MakeTensor(l1BTensor, L1B_LAYOUT, Arch::PositionL1{}); + // Get GM tile for next stage + auto tensorTileA = GetTileA(tensorA, 0, kL1IdxNext * L1_TILE_K, mBlockActual, kL1ActualNext); + auto tensorTileB = GetTile(tensorB, tla::MakeCoord(kL1IdxNext * L1_TILE_K, 0), + tla::MakeShape(kL1ActualNext, nBlockActual)); + + // load next matrix A tile from GM to L1 + AscendC::WaitFlag(l1AEventList[l1AListIdNext]); + copyGmToL1A(tensorL1A, tensorTileA); + AscendC::SetFlag(l1AEventList[l1AListIdNext]); + + // load next matrix B tile from GM to L1 + AscendC::WaitFlag(l1BEventList[l1BListIdNext]); + // copyGmToL1B(tensorL1B, tensorTileB); + AscendC::SetFlag(l1BEventList[l1BListIdNext]); + } + + // Get L1 tensor for current stage + auto l1ATensor = l1ATensorList[l1AListId]; + // auto l1BTensor = l1BTensorList[l1BListId]; + tensorL1A = tla::MakeTensor(l1ATensor, L1A_LAYOUT, Arch::PositionL1{}); + // tensorL1B = tla::MakeTensor(l1BTensor, L1B_LAYOUT, Arch::PositionL1{}); + // Get the loop nums on L0 + uint32_t kL0Loop = CeilDiv(kL1Actual); + + for (int mL0Idx = 0; mL0Idx < mL0Loop; mL0Idx++) { + uint32_t mL0Actual = (mL0Idx < mL0Loop - 1) ? L0_TILE_M : (mL1Actual - mL0Idx * L0_TILE_M); + + for (int kL0Idx = 0; kL0Idx < kL0Loop; kL0Idx++) { + uint32_t kL0Actual = (kL0Idx < kL0Loop - 1) ? L0_TILE_K : (kL1Actual - kL0Idx * L0_TILE_K); + + // Locate the current tile on L0A + auto l0ATile = l0ATensorList[l0AListId]; + auto layoutAInL0 = tla::MakeLayout(mL0Actual, kL0Actual); + auto tensorL0A = tla::MakeTensor(l0ATile, layoutAInL0, Arch::PositionL0A{}); + // Locate the current tile of matrix A on L1 + auto tensorTileL1A = GetTileA(tensorL1A, mL0Idx * L0_TILE_M, kL0Idx * L0_TILE_K, mL0Actual, kL0Actual); + + AscendC::WaitFlag(l0AEventList[l0AListId]); + if ((mL0Idx == 0) && (kL0Idx == 0)) { + AscendC::WaitFlag(l1AEventList[l1AListId]); + } + + // Load current tile from L1 to L0A + copyL1ToL0A(tensorL0A, tensorTileL1A); + + if ((mL0Idx == mL0Loop - 1) && (kL0Idx == kL0Loop - 1)) { + AscendC::SetFlag(l1AEventList[l1AListId]); + } + + bool initC = ((kL1Idx == 0) && (kL0Idx == 0)); + for (int nL0Idx = 0; nL0Idx < nL0Loop; nL0Idx++) { + uint32_t nL0Actual = (nL0Idx < nL0Loop - 1) ? L0_TILE_N : (nL1Actual - nL0Idx * L0_TILE_N); + + // Locate the current tile on L0B + auto l0BTile = l0BTensorList[l0BListId]; + auto layoutBInL0 = tla::MakeLayout(kL0Actual, nL0Actual); + auto tensorL0B = tla::MakeTensor(l0BTile, layoutBInL0, Arch::PositionL0B{}); + // Locate the current tile of matrix B on L1 + // auto tensorTileL1B = GetTile(tensorL1B, + // tla::MakeCoord(kL0Idx * L0_TILE_K, nL0Idx * L0_TILE_N), + // tla::MakeShape(kL0Actual, nL0Actual)); + auto tensorTileL1B = GetTile(tensorB, + tla::MakeCoord(kL0Idx * L0_TILE_K, nL0Idx * L0_TILE_N), + tla::MakeShape(kL0Actual, nL0Actual)); + + // Wait for mmad finished + AscendC::WaitFlag(l0BEventList[l0BListId]); + // If the current tile is the first one on the k&n axis, wait for loading matrix B from GM to L1 + if ((mL0Idx == 0) && (kL0Idx == 0) && (nL0Idx == 0)) { + AscendC::WaitFlag(l1BEventList[l1BListId]); + } + + // Load current tile from L1 to L0B + copyL1ToL0B(tensorL0B, tensorTileL1B); + + // If the current tile is the last one on the k&n axis, notify to load matrix B from GM to L1 + if ((mL0Idx == mL0Loop - 1) && (kL0Idx == kL0Loop - 1) && (nL0Idx == nL0Loop - 1)) { + AscendC::SetFlag(l1BEventList[l1BListId]); + } + + // Notify to do mmad + AscendC::SetFlag(l0CEventList[l0CListId]); + + // Locate the current tile on L0C + auto tensorTileL0C = GetTile(tensorL0C, + tla::MakeCoord(mL0Idx * L0_TILE_M, nL0Idx * L0_TILE_N), + tla::MakeShape(mL0Actual, nL0Actual)); + + // Compute the matrix multiplication on L0A and L0B and write the result to the accumulator + // Wait for loading L0B + AscendC::WaitFlag(l0CEventList[l0CListId]); + + // If the unit flag is enabled, the unit flag is set according to the calculation progress + uint8_t unitFlag = 0b00; + if constexpr (ENABLE_UNIT_FLAG) { + if ((kL1Idx == kL1Loop - 1) && (mL0Idx == mL0Loop - 1) && + (kL0Idx == kL0Loop - 1) && (nL0Idx == nL0Loop - 1)) { + unitFlag = 0b11; + } else { + unitFlag = 0b10; + } + } + + tileMmad(tensorTileL0C, tensorL0A, tensorL0B, + mL0Actual, nL0Actual, kL0Actual, initC, unitFlag); + + // Notify to move the next L0B tile + AscendC::SetFlag(l0BEventList[l0BListId]); + l0BListId = (l0BListId + 1 < L0B_STAGES) ? (l0BListId + 1) : 0; + } + AscendC::SetFlag(l0AEventList[l0AListId]); + l0AListId = (l0AListId + 1 < L0A_STAGES) ? (l0AListId + 1) : 0; + } + } + l1AListId = l1AListIdNext; + l1BListId = l1BListIdNext; + kL1Actual = kL1ActualNext; + } + + // copy block out + if constexpr (!ENABLE_UNIT_FLAG) { + AscendC::SetFlag(l0CEventList[l0CListId]); + AscendC::WaitFlag(l0CEventList[l0CListId]); + copyL0CToDst(tensorC, tensorL0C); + AscendC::SetFlag(l0CEventList[l0CListId]); + l0CListId = (l0CListId + 1 < L0C_STAGES) ? (l0CListId + 1) : 0; + } else { + copyL0CToDst(tensorC, tensorL0C, 0b11); + } + } + +protected: + template + CATLASS_DEVICE auto GetTileA(TensorA &tensorA, uint32_t mIndex, uint32_t kIndex, uint32_t mSize, uint32_t kSize) + { + if constexpr(tla::detail::isVector::value) { + return GetTile(tensorA, tla::MakeCoord(kIndex), tla::MakeShape(kSize)); + } else { + return GetTile(tensorA, tla::MakeCoord(mIndex, kIndex), tla::MakeShape(mSize, kSize)); + } + } + + // Multi-stage tensors list + AscendC::LocalTensor l1ATensorList[L1A_STAGES]; + AscendC::LocalTensor l1BTensorList[L1B_STAGES]; + AscendC::LocalTensor l0ATensorList[L0A_STAGES]; + AscendC::LocalTensor l0BTensorList[L0B_STAGES]; + AscendC::LocalTensor l0CTensorList[L0C_STAGES]; + AscendC::LocalTensor l1BiasTensor; + AscendC::LocalTensor l0BiasTensor; + + // Multi-stage event id list + int32_t l1AEventList[L1A_STAGES]; + int32_t l1BEventList[L1B_STAGES]; + int32_t l0AEventList[L0A_STAGES]; + int32_t l0BEventList[L0B_STAGES]; + int32_t l0CEventList[L0C_STAGES]; + + __gm__ typename AscendC::GlobalTensor::PrimType* lastAddrA[L1A_STAGES]; + __gm__ typename AscendC::GlobalTensor::PrimType* lastAddrB[L1B_STAGES]; + MatrixCoord lastCoordA[L1A_STAGES]; + MatrixCoord lastCoordB[L1B_STAGES]; + + // The id of current stage + uint32_t l1AListId{0}; + uint32_t l1BListId{0}; + uint32_t l0AListId{0}; + uint32_t l0BListId{0}; + uint32_t l0CListId{0}; + + TileMmad tileMmad; + CopyL1ToL0A copyL1ToL0A; + CopyL1ToL0B copyL1ToL0B; + CopyL1ToBT copyL1ToBT; +}; + +} // namespace Catlass::Gemm::Block + +#endif // CATLASS_GEMM_BLOCK_BLOCK_MMAD_PINGPONG_TLA_PRELOADA_L1B_HPP diff --git a/csrc/moe/dequant_situ_quant/CMakeLists.txt b/csrc/moe/dequant_situ_quant/CMakeLists.txt new file mode 100644 index 000000000000..dcb22d20cbed --- /dev/null +++ b/csrc/moe/dequant_situ_quant/CMakeLists.txt @@ -0,0 +1,19 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2025 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ---------------------------------------------------------------------------- + +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +if(NOT ENABLE_TEST AND NOT BENCHMARK) + list(REMOVE_ITEM CURRENT_DIRS tests) +endif() +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/csrc/moe/dequant_situ_quant/dequant_situ_quant_torch_adpt.h b/csrc/moe/dequant_situ_quant/dequant_situ_quant_torch_adpt.h new file mode 100644 index 000000000000..8c9d279a034e --- /dev/null +++ b/csrc/moe/dequant_situ_quant/dequant_situ_quant_torch_adpt.h @@ -0,0 +1,74 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2026. 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. + * 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. + */ + +#ifndef DEQUANT_SITU_QUANT_TORCH_ADPT_H +#define DEQUANT_SITU_QUANT_TORCH_ADPT_H + +namespace vllm_ascend { + +std::tuple dequant_situ_quant( + const at::Tensor& x, + const c10::optional& weight_scale, + const c10::optional& activation_scale, + const c10::optional& bias, + const c10::optional& quant_scale, + const c10::optional& quant_offset, + const c10::optional& group_index, + double beta, + double linear_beta, + bool activate_left, + c10::string_view quant_mode) +{ + TORCH_CHECK(x.dim() == 2, + "dequant_situ_quant: x must be 2-dimensional [rows, width], but got rank ", + x.dim()); + const int64_t input_width = x.size(1); + + at::Tensor y = at::empty({x.size(0), input_width / 2}, x.options().dtype(at::kChar)); + at::Tensor scale = at::empty({x.size(0)}, x.options().dtype(at::kFloat)); + if (x.size(0) == 0) { + return {y, scale}; + } + + const at::Tensor weight_scale_value = weight_scale.value_or(at::Tensor()); + const at::Tensor activation_scale_value = activation_scale.value_or(at::Tensor()); + const at::Tensor bias_value = bias.value_or(at::Tensor()); + const at::Tensor quant_scale_value = quant_scale.value_or(at::Tensor()); + const at::Tensor quant_offset_value = quant_offset.value_or(at::Tensor()); + const at::Tensor group_index_value = group_index.value_or(at::Tensor()); + std::string quant_mode_string(quant_mode); + char* quant_mode_ptr = quant_mode_string.data(); + + EXEC_NPU_CMD(aclnnDequantSituQuant, + x, + weight_scale_value, + activation_scale_value, + bias_value, + quant_scale_value, + quant_offset_value, + group_index_value, + beta, + linear_beta, + activate_left, + quant_mode_ptr, + y, + scale); + return {y, scale}; +} + +} // namespace vllm_ascend + +#endif // DEQUANT_SITU_QUANT_TORCH_ADPT_H diff --git a/csrc/moe/dequant_situ_quant/docs/aclnnDequantSituQuant.md b/csrc/moe/dequant_situ_quant/docs/aclnnDequantSituQuant.md new file mode 100644 index 000000000000..2eb393db4e92 --- /dev/null +++ b/csrc/moe/dequant_situ_quant/docs/aclnnDequantSituQuant.md @@ -0,0 +1,222 @@ +# aclnnDequantSituQuant + +## 产品支持情况 + +|产品 | 是否支持 | +|:-------------------------|:----------:| +| Ascend 950PR/Ascend 950DT | × | +| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ | +| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ | +| Atlas 200I/500 A2 推理产品 | × | +| Atlas 推理系列产品 | × | +| Atlas 训练系列产品 | × | + +## 功能说明 + +- 接口功能:在Situ激活函数前后添加dequant和quant操作,实现x的DequantSituQuant计算。 +- 计算公式: + + $$ + dequantOut = cast\_to\_float(x) \times dequantScale + dequantBias + $$ + + $$ + situOut = Situ(dequantOut) = \beta \times \tanh(gate / \beta) \times sigmoid(gate) \times up + $$ + + $$ + out = Quant(situOut, quantScale, quantOffset) + $$ + + 其中,当activateLeft为true时,gate取dequantOut的前半部分,up取后半部分;当activateLeft为false时,gate取dequantOut的后半部分,up取前半部分。当linearBeta > 0时,up会被进一步变换为 $linear\_beta \times \tanh(up / linear\_beta)$。 + +## 函数原型 + +每个算子分为[两段式接口](../../../docs/zh/context/两段式接口.md),必须先调用"aclnnDequantSituQuantGetWorkspaceSize"接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用"aclnnDequantSituQuant"接口执行计算。 + +```Cpp +aclnnStatus aclnnDequantSituQuantGetWorkspaceSize( + const aclTensor *x, + const aclTensor *dequantScale, + const aclTensor *dequantBiasOptional, + const aclTensor *quantScaleOptional, + const aclTensor *quantOffsetOptional, + float beta, + float linearBeta, + bool activateLeft, + char *quantModeOptional, + const aclTensor *yOut, + const aclTensor *scaleOut, + uint64_t *workspaceSize, + aclOpExecutor **executor) +``` + +```Cpp +aclnnStatus aclnnDequantSituQuant( + void *workspace, + uint64_t workspaceSize, + aclOpExecutor *executor, + aclrtStream stream) +``` + +## aclnnDequantSituQuantGetWorkspaceSize + +- **参数说明:** + + | 参数名 | 输入/输出 | 描述 | 使用说明 | 数据类型 | 维度(shape) | + |--------|-----------|------|----------|----------|-------------| + | x | 输入 | 输入待处理的数据 | shape为(N...,H),最后一维需要是2的倍数,且x的维度必须大于1维。不支持空Tensor。 | INT8 | 2-8 | + | dequantScale | 输入 | 反量化scale | shape为(H,)或(1,)。当shape为(H,)时,取值H和x最后一维保持一致。 | FLOAT32 | 1 | + | dequantBiasOptional | 输入 | 反量化bias | shape为(H,)或(1,)。可选参数,支持传空指针。 | FLOAT32 | 1 | + | quantScaleOptional | 输入 | 量化的scale | 当quantModeOptional为static时,shape为(H/2,)或(1,);当quantModeOptional为dynamic时,shape为(H/2,),作为smoothScale使用。可选参数,支持传空指针。 | FLOAT32 | 1 | + | quantOffsetOptional | 输入 | 量化的offset | shape为(H/2,)或(1,)。仅当quantModeOptional为static时有效。可选参数,支持传空指针。 | FLOAT32 | 1 | + | beta | 输入 | Situ激活的beta参数 | 不能为0。 | Float | - | + | linearBeta | 输入 | Situ激活的linear_beta参数 | 当值≤0时不启用linear_beta变换。 | Float | - | + | activateLeft | 输入 | 是否对输入的左半部分做Situ激活 | 当值为false时,对输入的右半部分做激活。 | Bool | - | + | quantModeOptional | 输入 | 量化模式 | 支持"static"和"dynamic"。 | String | - | + | yOut | 输出 | 量化后的输出 | shape为(N...,H/2)。 | INT8 | - | + | scaleOut | 输出 | 动态量化的scale | shape为(N,...),与yOut去除尾轴后的shape一致。 | FLOAT32 | - | + +- **返回值:** + + aclnnStatus:返回状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn返回码.md)。 + +## 约束说明 + +- x的最后一维需要是2的倍数,且x的维数必须大于1维。 +- beta参数不能为0。 +- 当quantModeOptional为static时,quantScaleOptional必须提供。 +- 当quantModeOptional为dynamic时,quantScaleOptional可选(作为smoothScale使用)。 + +## 调用示例 + +示例代码如下,仅供参考,具体编译和执行过程请参考[编译与运行样例](../../../docs/zh/context/编译与运行样例.md)。 + +```C++ +#include +#include +#include "acl/acl.h" +#include "aclnnop/aclnn_dequant_situ_quant.h" + +int64_t GetShapeSize(const std::vector& shape) +{ + int64_t shapeSize = 1; + for (auto dim : shape) { + shapeSize *= dim; + } + return shapeSize; +} + +template +int CreateAclTensor(const std::vector& hostData, const std::vector& shape, void** deviceAddr, + aclDataType dataType, aclTensor** tensor) +{ + auto size = GetShapeSize(shape) * sizeof(T); + auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); + if (ret != ACL_SUCCESS) return ret; + ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); + if (ret != ACL_SUCCESS) return ret; + + std::vector strides(shape.size(), 1); + for (int64_t i = shape.size() - 2; i >= 0; i--) { + strides[i] = shape[i + 1] * strides[i + 1]; + } + *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, + aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), *deviceAddr); + return 0; +} + +int main() { + // 1. Initialize device and stream + int32_t deviceId = 0; + aclrtStream stream; + auto ret = aclInit(nullptr); + ret = aclrtSetDevice(deviceId); + ret = aclrtCreateStream(&stream); + + // 2. Construct inputs: x=[16, 64], static quant mode + int64_t rowLen = 16; + int64_t inDimy = 64; + int64_t outDimy = 32; + double beta = 1.0; + double linearBeta = 0.0; + bool activateLeft = false; + + std::vector xShape = {rowLen, inDimy}; + std::vector dequantScaleShape = {inDimy}; + std::vector quantScaleShape = {1}; + std::vector quantOffsetShape = {1}; + std::vector yShape = {rowLen, outDimy}; + std::vector scaleOutShape = {rowLen}; + + auto xSize = GetShapeSize(xShape); + std::vector xHostData(xSize); + for (int64_t i = 0; i < xSize; i++) { + xHostData[i] = static_cast((i * 7 + 3) % 100); + } + std::vector dequantScaleHostData(inDimy, 0.1f); + std::vector quantScaleHostData(1, 1.0f); + std::vector quantOffsetHostData(1, 0.0f); + std::vector yHostData(GetShapeSize(yShape), 0); + std::vector scaleOutHostData(GetShapeSize(scaleOutShape), 0.0f); + + void* xDeviceAddr = nullptr; + void* dsDeviceAddr = nullptr; + void* qsDeviceAddr = nullptr; + void* qoDeviceAddr = nullptr; + void* yDeviceAddr = nullptr; + void* scaleDeviceAddr = nullptr; + aclTensor* x = nullptr; + aclTensor* dequantScale = nullptr; + aclTensor* quantScale = nullptr; + aclTensor* quantOffset = nullptr; + aclTensor* y = nullptr; + aclTensor* scaleOut = nullptr; + + CreateAclTensor(xHostData, xShape, &xDeviceAddr, aclDataType::ACL_INT8, &x); + CreateAclTensor(dequantScaleHostData, dequantScaleShape, &dsDeviceAddr, aclDataType::ACL_FLOAT, &dequantScale); + CreateAclTensor(quantScaleHostData, quantScaleShape, &qsDeviceAddr, aclDataType::ACL_FLOAT, &quantScale); + CreateAclTensor(quantOffsetHostData, quantOffsetShape, &qoDeviceAddr, aclDataType::ACL_FLOAT, &quantOffset); + CreateAclTensor(yHostData, yShape, &yDeviceAddr, aclDataType::ACL_INT8, &y); + CreateAclTensor(scaleOutHostData, scaleOutShape, &scaleDeviceAddr, aclDataType::ACL_FLOAT, &scaleOut); + + // 3. Call aclnnDequantSituQuantGetWorkspaceSize + uint64_t workspaceSize = 0; + aclOpExecutor* executor; + ret = aclnnDequantSituQuantGetWorkspaceSize(x, dequantScale, nullptr, quantScale, quantOffset, + beta, linearBeta, activateLeft, const_cast("static"), + y, scaleOut, &workspaceSize, &executor); + + // 4. Allocate workspace and call aclnnDequantSituQuant + void* workspaceAddr = nullptr; + if (workspaceSize > 0) { + aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); + } + ret = aclnnDequantSituQuant(workspaceAddr, workspaceSize, executor, stream); + ret = aclrtSynchronizeStream(stream); + + // 5. Copy output and cleanup + std::vector npuYResult(GetShapeSize(yShape), 0); + aclrtMemcpy(npuYResult.data(), npuYResult.size() * sizeof(int8_t), yDeviceAddr, + GetShapeSize(yShape) * sizeof(int8_t), ACL_MEMCPY_DEVICE_TO_HOST); + + aclDestroyTensor(x); + aclDestroyTensor(dequantScale); + aclDestroyTensor(quantScale); + aclDestroyTensor(quantOffset); + aclDestroyTensor(y); + aclDestroyTensor(scaleOut); + aclrtFree(xDeviceAddr); + aclrtFree(dsDeviceAddr); + aclrtFree(qsDeviceAddr); + aclrtFree(qoDeviceAddr); + aclrtFree(yDeviceAddr); + aclrtFree(scaleDeviceAddr); + if (workspaceSize > 0) aclrtFree(workspaceAddr); + + aclrtDestroyStream(stream); + aclrtResetDevice(deviceId); + aclFinalize(); + return 0; +} +``` diff --git a/csrc/moe/dequant_situ_quant/op_graph/dequant_situ_quant_proto.h b/csrc/moe/dequant_situ_quant/op_graph/dequant_situ_quant_proto.h new file mode 100644 index 000000000000..fa1e414ee42b --- /dev/null +++ b/csrc/moe/dequant_situ_quant/op_graph/dequant_situ_quant_proto.h @@ -0,0 +1,75 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file dequant_situ_quant_proto.h + * \brief Kimi K3 routed-MoE Dequant + SiTU + dynamic quantization. + */ +#ifndef OPS_QUANT_DEQUANT_SITU_QUANT_PROTO_H_ +#define OPS_QUANT_DEQUANT_SITU_QUANT_PROTO_H_ +#include "graph/operator_reg.h" + +namespace ge { + +/** +* @brief Dequantize an INT32 grouped-matmul accumulator or consume an already +* dequantized BF16 WeightNz GMM output, apply Kimi K3 SiTU, and dynamically +* quantize each row. + +* @par Inputs: +* Seven inputs, matching the grouped DequantSwigluQuant parameter order: +* @li x: Required INT32/BF16 tensor. Routed shape is [rows, 6144]. The shared-expert +* local gate-up width is 12288/6144/3072/1536/768 for TP1/2/4/8/16. +* @li weight_scale: FP32 tensor. Routed shape is [experts, 6144]; shared shape +* is [local_width] or [1, local_width]. Required when x is INT32. +* @li activation_scale: FP32 tensor containing one value per row. Required +* when x is INT32. +* @li bias: Optional FP32 tensor with the same grouped shape as weight_scale. +* @li quant_scale: Reserved optional input. It must be absent. +* @li quant_offset: Reserved optional input. It must be absent. +* @li group_index: Optional INT64 1-D tensor with shape [experts]. Each value +* is the number of consecutive routed rows belonging to that expert. None +* selects the single-group shared-expert path. +* For BF16 x, GMM has already applied dequantization; weight_scale, +* activation_scale, bias, and group_index must all be absent. + +* @par Outputs: +* @li y: INT8 tensor with shape [rows, x.shape[-1] / 2]. +* @li scale: FP32 tensor with shape [rows]. + +* @par Attributes: +* @li beta: Must be 4.0. +* @li linear_beta: Must be 25.0. +* @li activate_left: Must be true. +* @li quant_mode: Must be "dynamic". + +* @attention Constraints: +* @li Only the Kimi K3 A2/A3 W4A8 routed and TP-sharded shared contracts are supported. +* @li Non-empty group counts must cover x rows in expert order. Empty experts +* (zero counts) are supported. +*/ +REG_OP(DequantSituQuant) + .INPUT(x, TensorType({DT_INT32, DT_BF16})) + .OPTIONAL_INPUT(weight_scale, TensorType({DT_FLOAT})) + .OPTIONAL_INPUT(activation_scale, TensorType({DT_FLOAT})) + .OPTIONAL_INPUT(bias, TensorType({DT_FLOAT})) + .OPTIONAL_INPUT(quant_scale, TensorType({DT_FLOAT})) + .OPTIONAL_INPUT(quant_offset, TensorType({DT_FLOAT})) + .OPTIONAL_INPUT(group_index, TensorType({DT_INT64})) + .OUTPUT(y, TensorType({DT_INT8})) + .OUTPUT(scale, TensorType({DT_FLOAT})) + .ATTR(beta, Float, 4.0) + .ATTR(linear_beta, Float, 25.0) + .ATTR(activate_left, Bool, true) + .ATTR(quant_mode, String, "dynamic") + .OP_END_FACTORY_REG(DequantSituQuant) +} // namespace ge + +#endif // OPS_QUANT_DEQUANT_SITU_QUANT_PROTO_H_ diff --git a/csrc/moe/dequant_situ_quant/op_host/CMakeLists.txt b/csrc/moe/dequant_situ_quant/op_host/CMakeLists.txt new file mode 100644 index 000000000000..21a2bcadb7cf --- /dev/null +++ b/csrc/moe/dequant_situ_quant/op_host/CMakeLists.txt @@ -0,0 +1,39 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2025 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ---------------------------------------------------------------------------- + +set(DEQUANT_SITU_QUANT_SUPPORTED_SOCS ascend910b ascend910_93) +set(DEQUANT_SITU_QUANT_ENABLED FALSE) +foreach(SOC_VERSION ${ASCEND_COMPUTE_UNIT}) + if(SOC_VERSION IN_LIST DEQUANT_SITU_QUANT_SUPPORTED_SOCS) + set(DEQUANT_SITU_QUANT_ENABLED TRUE) + endif() +endforeach() +if(NOT DEQUANT_SITU_QUANT_ENABLED) + return() +endif() + +add_op_to_compiled_list() + +if(BUILD_OPEN_PROJECT) + target_sources(op_host_aclnn PRIVATE + dequant_situ_quant_def.cpp + ) +endif() + +add_ops_compile_options( + OP_NAME DequantSituQuant + OPTIONS --cce-auto-sync=on + -Wno-deprecated-declarations + -Werror +) + +if(NOT BUILD_OPS_RTY_KERNEL) + add_modules_sources(OPTYPE dequant_situ_quant ACLNNTYPE aclnn) +endif() diff --git a/csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_def.cpp b/csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_def.cpp new file mode 100644 index 000000000000..a0a01a773fc9 --- /dev/null +++ b/csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_def.cpp @@ -0,0 +1,81 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file dequant_situ_quant_def.cpp + * \brief + */ + +#include +#include "register/op_def_registry.h" + +namespace ops { +class DequantSituQuant : public OpDef { +public: + explicit DequantSituQuant(const char* name) : OpDef(name) + { + this->Input("x") + .ParamType(REQUIRED) + .DataType({ge::DT_INT32, ge::DT_BF16}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND}); + this->Input("weight_scale") + .ParamType(OPTIONAL) + .DataType({ge::DT_FLOAT, ge::DT_FLOAT}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND}); + this->Input("activation_scale") + .ParamType(OPTIONAL) + .DataType({ge::DT_FLOAT, ge::DT_FLOAT}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND}); + this->Input("bias") + .ParamType(OPTIONAL) + .DataType({ge::DT_FLOAT, ge::DT_FLOAT}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND}); + this->Input("quant_scale") + .ParamType(OPTIONAL) + .DataType({ge::DT_FLOAT, ge::DT_FLOAT}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND}); + this->Input("quant_offset") + .ParamType(OPTIONAL) + .DataType({ge::DT_FLOAT, ge::DT_FLOAT}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND}); + this->Input("group_index") + .ParamType(OPTIONAL) + .DataType({ge::DT_INT64, ge::DT_INT64}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND}); + this->Output("y") + .ParamType(REQUIRED) + .DataType({ge::DT_INT8, ge::DT_INT8}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND}); + this->Output("scale") + .ParamType(REQUIRED) + .DataType({ge::DT_FLOAT, ge::DT_FLOAT}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND}); + + this->Attr("beta").AttrType(OPTIONAL).Float(4.0); + this->Attr("linear_beta").AttrType(OPTIONAL).Float(25.0); + this->Attr("activate_left").AttrType(OPTIONAL).Bool(true); + this->Attr("quant_mode").AttrType(OPTIONAL).String("dynamic"); + + this->AICore().AddConfig("ascend910b"); + this->AICore().AddConfig("ascend910_93"); + } +}; + +OP_ADD(DequantSituQuant); +} // namespace ops diff --git a/csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_infershape.cpp b/csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_infershape.cpp new file mode 100644 index 000000000000..94b2b8261d1d --- /dev/null +++ b/csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_infershape.cpp @@ -0,0 +1,93 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file dequant_situ_quant_infershape.cpp + * \brief + */ +#include "register/op_impl_registry.h" +#include "graph/utils/type_utils.h" +#include "util/shape_util.h" +#include "log/log.h" +#include "util/math_util.h" + +using namespace ge; +namespace ops { +constexpr size_t INPUT_IDX_X = 0; +constexpr size_t OUTPUT_IDX_Y = 0; +constexpr size_t OUTPUT_IDX_SCALE = 1; +constexpr int64_t CONST_UNKNOW_SHAPE = -1; +constexpr int64_t NUM_TWO = 2; + +graphStatus InferShape4DequantSituQuant(gert::InferShapeContext* context) +{ + OP_LOGD(context, "Begin to do InferShape4DequantSituQuant."); + + const gert::Shape* xShape = context->GetInputShape(INPUT_IDX_X); + OP_CHECK_NULL_WITH_CONTEXT(context, xShape); + gert::Shape* yShape = context->GetOutputShape(OUTPUT_IDX_Y); + OP_CHECK_NULL_WITH_CONTEXT(context, yShape); + gert::Shape* scaleShape = context->GetOutputShape(OUTPUT_IDX_SCALE); + OP_CHECK_NULL_WITH_CONTEXT(context, scaleShape); + + *yShape = *xShape; + OP_CHECK_IF(Ops::Base::IsUnknownRank(*xShape), + OP_LOGD(context, "End to do InferShape4DequantSituQuant, inputx is [-2]."), return GRAPH_SUCCESS); + + int64_t xShapeRank = static_cast(xShape->GetDimNum()); + + auto inputDesc = context->GetInputDesc(INPUT_IDX_X); + OP_CHECK_NULL_WITH_CONTEXT(context, inputDesc); + ge::DataType xDataType = inputDesc->GetDataType(); + + bool isInt8 = (xDataType == ge::DT_INT8); + if (isInt8) { + OP_CHECK_IF(xShapeRank <= 1, + OP_LOGE(context, "x shape rank must > 1 for int8, but is %ld", xShapeRank), + return ge::GRAPH_FAILED); + } else { + OP_CHECK_IF(xShapeRank != 2, + OP_LOGE(context, "x shape rank must be 2 for int32/bfloat16, but is %ld", xShapeRank), + return ge::GRAPH_FAILED); + } + + int64_t lastDim = xShape->GetDim(xShapeRank - 1); + int64_t outLastDim = lastDim == CONST_UNKNOW_SHAPE ? CONST_UNKNOW_SHAPE : lastDim / NUM_TWO; + OP_CHECK_IF((lastDim != CONST_UNKNOW_SHAPE) && (lastDim % NUM_TWO != 0), + OP_LOGE(context, "The last dim of x must be even number, but is %ld", lastDim), + return ge::GRAPH_FAILED); + + yShape->SetDim(xShapeRank - 1, outLastDim); + + *scaleShape = *yShape; + if (isInt8) { + scaleShape->SetDimNum(xShapeRank - 1); + } else { + scaleShape->SetDimNum(1); + scaleShape->SetDim(0, xShape->GetDim(0)); + } + + OP_LOGD(context, "End to do InferShape4DequantSituQuant"); + return GRAPH_SUCCESS; +} + +graphStatus InferDtype4DequantSituQuant(gert::InferDataTypeContext* context) +{ + OP_LOGD(context, "InferDtype4DequantSituQuant enter"); + context->SetOutputDataType(OUTPUT_IDX_Y, DT_INT8); + context->SetOutputDataType(OUTPUT_IDX_SCALE, DT_FLOAT); + OP_LOGD(context, "InferDtype4DequantSituQuant end"); + return GRAPH_SUCCESS; +} + +IMPL_OP_INFERSHAPE(DequantSituQuant) + .InferShape(InferShape4DequantSituQuant) + .InferDataType(InferDtype4DequantSituQuant); +} // namespace ops diff --git a/csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_tiling.cpp b/csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_tiling.cpp new file mode 100644 index 000000000000..c2e309440052 --- /dev/null +++ b/csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_tiling.cpp @@ -0,0 +1,701 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file dequant_situ_quant_tiling.cpp + * \brief + */ + +#include +#include +#include +#include +#include +#include "tiling/tiling_api.h" +#include "dequant_situ_quant_tiling.h" +#include "../../dequant_swiglu_quant/tiling_base/tiling_util.h" +#include "../../dequant_swiglu_quant/tiling_base/tiling_templates_registry.h" + +#define CHECK_FAIL(cont, cond, ...) \ + do { \ + if (cond) { \ + OP_LOGE(cont->GetNodeName(), ##__VA_ARGS__); \ + return ge::GRAPH_FAILED; \ + } \ + } while (0) + +namespace optiling { +using Ops::NN::Optiling::TilingBaseClass; +using Ops::NN::Optiling::TilingRegistry; + +constexpr uint32_t UB_RESERVED_BUFF = 1024; +constexpr uint32_t ALIGN_UINT_IN_CACHE_32B = 32; +constexpr uint32_t PACK_UINT_IN_CACHE_512B = 512; +constexpr uint32_t DEFAULT_BUFFER_NUM = 2; +constexpr uint32_t MAX_CORE_NUMBER = 64; +constexpr uint32_t PERFORMANCE_COL_LEN = 1536; +constexpr uint32_t PERFORMANCE_ROW_LEN = 128; +constexpr uint32_t MIN_CORE = 12; +constexpr uint64_t USER_WORKSPACE = 16777216; // 16 * 1024 * 1024 +constexpr uint64_t BLOCK_BYTES = 32; + +constexpr size_t INDEX_IN_X = 0; +constexpr size_t INDEX_IN_WEIGHT_SCALE = 1; +constexpr size_t INDEX_IN_ACTIVATION_SCALE = 2; +constexpr size_t INDEX_IN_BIAS = 3; +constexpr size_t INDEX_IN_QUANT_SCALE = 4; +constexpr size_t INDEX_IN_QUANT_OFFSET = 5; +constexpr size_t INDEX_IN_GROUP_INDEX = 6; +constexpr int64_t NUM_TWO = 2; +constexpr int64_t QUANT_MODE_STATIC = 0; +constexpr int64_t QUANT_MODE_DYNAMIC = 1; + +void DequantSituQuantTiling::Reset() +{ + opName = nullptr; + totalCore = 0; + totalUsedCoreNum = 0; + inputDTypeLen = 1; + ubMinBlockLen = 32; + cacheLineLen = 512; + maxTileLen = 0; + optBaseRowLen = 0; + optBaseColLen = 0; + workspaceSize_ = 0; + hasDequantBias = false; + hasQuantScale = false; + hasQuantOffset = false; + quantIsOne = false; + activateLeft = 0; + quantMode = 0; + beta = 4.0f; + linearBeta = 25.0f; + quantScaleShapeSize = 0; + inDimx = 0; + inDimy = 0; + outDimy = 0; + xDtype_ = ge::DT_INT8; + isPreDequantized_ = false; + hasWeightScale_ = false; + hasActivationScale_ = false; + hasGroupIndex_ = false; + expertNum_ = 1; +} + +bool DequantSituQuantTiling::IsCapable() +{ + if (Ops::NN::OpTiling::IsRegbaseSocVersion(context_)) { + return false; + } + return true; +} + +ge::graphStatus DequantSituQuantTiling::GetPlatformInfo() +{ + auto platformInfo = context_->GetPlatformInfo(); + OP_CHECK_IF(platformInfo == nullptr, OP_LOGE(opName, "fail to get platform info"), return ge::GRAPH_FAILED); + auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo); + totalCore = ascendcPlatform.GetCoreNumAiv(); + aicoreParams_.numBlocks = totalCore; + uint64_t ubSizePlatForm; + ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm); + aicoreParams_.ubSize = ubSizePlatForm; + socVersion = ascendcPlatform.GetSocVersion(); + return ge::GRAPH_SUCCESS; +} + +bool DequantSituQuantTiling::SetAttrs(const gert::RuntimeAttrs* attrs) +{ + auto betaPtr = attrs->GetAttrPointer(0); + beta = (betaPtr == nullptr) ? 4.0f : *betaPtr; + OP_CHECK_IF(beta == 0.0f, OP_LOGE(context_->GetNodeName(), "beta must not be 0"), return false); + + auto linearBetaPtr = attrs->GetAttrPointer(1); + linearBeta = (linearBetaPtr == nullptr) ? 25.0f : *linearBetaPtr; + + auto isActivateLeftPtr = attrs->GetBool(2); + bool isActivateLeft = isActivateLeftPtr == nullptr ? true : *isActivateLeftPtr; + activateLeft = (isActivateLeft ? 1 : 0); + + auto str = attrs->GetStr(3); + OP_CHECK_IF(str == nullptr, OP_LOGE(context_->GetNodeName(), "quant_mode attr is null"), return false); + std::string quantModeAttr{str}; + std::transform(quantModeAttr.begin(), quantModeAttr.end(), quantModeAttr.begin(), ::tolower); + OP_CHECK_IF((quantModeAttr != "static") && (quantModeAttr != "dynamic"), + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context_->GetNodeName(), "quant_mode", quantModeAttr.c_str(), + "quant_mode should be static or dynamic with case insensitive"), + return false); + quantMode = ((quantModeAttr == "static") ? QUANT_MODE_STATIC : QUANT_MODE_DYNAMIC); + + tilingData.set_activateLeft(activateLeft); + tilingData.set_beta(beta); + tilingData.set_linearBeta(linearBeta); + tilingData.set_quantMode(quantMode); + return true; +} + +ge::graphStatus DequantSituQuantTiling::CheckInputShapesInt8(int64_t dimNum, int64_t inDimy, int64_t outDimy) +{ + CHECK_FAIL(context_, dimNum <= 1, "The shape dim of x can not be less than 2"); + + // weight_scale: required, shape [inDimy] or [1] + auto weightScaleShapePtr = context_->GetOptionalInputShape(INDEX_IN_WEIGHT_SCALE); + OP_CHECK_IF(weightScaleShapePtr == nullptr, + OP_LOGE(context_->GetNodeName(), "weight_scale is required for int8 x"), + return ge::GRAPH_FAILED); + uint64_t weightScaleSize = weightScaleShapePtr->GetStorageShape().GetShapeSize(); + OP_CHECK_IF(weightScaleSize != static_cast(inDimy) && weightScaleSize != 1, + OP_LOGE_FOR_INVALID_SHAPESIZE(context_->GetNodeName(), "weight_scale", + std::to_string(weightScaleSize).c_str(), + (std::to_string(inDimy) + " or 1").c_str()), + return ge::GRAPH_FAILED); + + // activation_scale: must be absent for int8 + OP_CHECK_IF(context_->GetOptionalInputShape(INDEX_IN_ACTIVATION_SCALE) != nullptr, + OP_LOGE(context_->GetNodeName(), "activation_scale must be absent for int8 x"), + return ge::GRAPH_FAILED); + + // bias: optional, shape [inDimy] or [1] + auto biasShapePtr = context_->GetOptionalInputShape(INDEX_IN_BIAS); + hasDequantBias = (biasShapePtr != nullptr); + tilingData.set_dequantBiasIsEmpty(!hasDequantBias); + if (hasDequantBias) { + uint64_t biasSize = biasShapePtr->GetStorageShape().GetShapeSize(); + OP_CHECK_IF(biasSize != static_cast(inDimy) && biasSize != 1, + OP_LOGE_FOR_INVALID_SHAPESIZE(context_->GetNodeName(), "bias", + std::to_string(biasSize).c_str(), + (std::to_string(inDimy) + " or 1").c_str()), + return ge::GRAPH_FAILED); + } + + // quant_scale: optional, shape [outDimy] or [1] + auto quantScaleShapePtr = context_->GetOptionalInputShape(INDEX_IN_QUANT_SCALE); + hasQuantScale = (quantScaleShapePtr != nullptr); + tilingData.set_quantScaleIsEmpty(!hasQuantScale); + if (hasQuantScale) { + quantScaleShapeSize = quantScaleShapePtr->GetStorageShape().GetShapeSize(); + OP_CHECK_IF(quantScaleShapeSize != static_cast(outDimy) && quantScaleShapeSize != 1, + OP_LOGE_FOR_INVALID_SHAPESIZE(context_->GetNodeName(), "quant_scale", + std::to_string(quantScaleShapeSize).c_str(), + (std::to_string(outDimy) + " or 1").c_str()), + return ge::GRAPH_FAILED); + quantIsOne = (quantScaleShapeSize == 1); + } else { + quantIsOne = false; + } + tilingData.set_quantIsOne(quantIsOne); + + // quant_offset: optional (only static), shape [outDimy] or [1] + auto quantOffsetShapePtr = context_->GetOptionalInputShape(INDEX_IN_QUANT_OFFSET); + hasQuantOffset = (quantOffsetShapePtr != nullptr); + tilingData.set_quantOffsetIsEmpty(!hasQuantOffset); + if (quantMode == QUANT_MODE_STATIC) { + OP_CHECK_IF(!hasQuantScale, + OP_LOGE(context_->GetNodeName(), "quant_scale must be provided when quant_mode is static"), + return ge::GRAPH_FAILED); + if (hasQuantOffset) { + uint64_t quantOffsetSize = quantOffsetShapePtr->GetStorageShape().GetShapeSize(); + OP_CHECK_IF(quantOffsetSize != static_cast(outDimy) && quantOffsetSize != 1, + OP_LOGE_FOR_INVALID_SHAPESIZE(context_->GetNodeName(), "quant_offset", + std::to_string(quantOffsetSize).c_str(), + (std::to_string(outDimy) + " or 1").c_str()), + return ge::GRAPH_FAILED); + } + } + + // group_index: must be absent for int8 + OP_CHECK_IF(context_->GetOptionalInputShape(INDEX_IN_GROUP_INDEX) != nullptr, + OP_LOGE(context_->GetNodeName(), "group_index must be absent for int8 x"), + return ge::GRAPH_FAILED); + + // check output shapes + auto yShapePtr = context_->GetOutputShape(0); + OP_CHECK_NULL_WITH_CONTEXT(context_, yShapePtr); + const gert::Shape yShape = yShapePtr->GetStorageShape(); + OP_CHECK_IF(yShape.GetDimNum() != static_cast(dimNum), + OP_LOGE(context_->GetNodeName(), "y shape dim must equal to x shape dim"), return ge::GRAPH_FAILED); + OP_CHECK_IF(yShape.GetDim(dimNum - 1) != outDimy, + OP_LOGE(context_->GetNodeName(), "y last dim must be x last dim / 2"), return ge::GRAPH_FAILED); + + auto scaleShapePtr = context_->GetOutputShape(1); + OP_CHECK_NULL_WITH_CONTEXT(context_, scaleShapePtr); + const gert::Shape scaleShape = scaleShapePtr->GetStorageShape(); + OP_CHECK_IF(static_cast(scaleShape.GetShapeSize()) != static_cast(inDimx), + OP_LOGE(context_->GetNodeName(), "scale shape size must equal to rowLen"), return ge::GRAPH_FAILED); + + tilingData.set_expertNum(1); + tilingData.set_hasGroupIndex(0); + tilingData.set_hasActivationScale(0); + tilingData.set_isPreDequantized(0); + tilingData.set_inputWidth(0); + tilingData.set_outputWidth(0); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus DequantSituQuantTiling::ValidateInt32Contract() +{ + // quant_scale/quant_offset must be absent + OP_CHECK_IF(context_->GetOptionalInputShape(INDEX_IN_QUANT_SCALE) != nullptr, + OP_LOGE(context_->GetNodeName(), "quant_scale is not supported for int32 x"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context_->GetOptionalInputShape(INDEX_IN_QUANT_OFFSET) != nullptr, + OP_LOGE(context_->GetNodeName(), "quant_offset is not supported for int32 x"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(quantMode != QUANT_MODE_DYNAMIC, + OP_LOGE(context_->GetNodeName(), "quant_mode must be dynamic for int32 x"), + return ge::GRAPH_FAILED); + + // weight_scale: required + auto weightScaleShapePtr = context_->GetOptionalInputShape(INDEX_IN_WEIGHT_SCALE); + OP_CHECK_IF(weightScaleShapePtr == nullptr, + OP_LOGE(context_->GetNodeName(), "weight_scale is required for int32 x"), + return ge::GRAPH_FAILED); + + // activation_scale: required + auto activationScaleShapePtr = context_->GetOptionalInputShape(INDEX_IN_ACTIVATION_SCALE); + OP_CHECK_IF(activationScaleShapePtr == nullptr, + OP_LOGE(context_->GetNodeName(), "activation_scale is required for int32 x"), + return ge::GRAPH_FAILED); + + // group_index: optional + auto groupIndexShapePtr = context_->GetOptionalInputShape(INDEX_IN_GROUP_INDEX); + hasGroupIndex_ = (groupIndexShapePtr != nullptr); + + expertNum_ = 1; + if (hasGroupIndex_) { + const gert::Shape groupIndexShape = groupIndexShapePtr->GetStorageShape(); + OP_CHECK_IF(groupIndexShape.GetDimNum() != 1 || groupIndexShape.GetDim(0) <= 0, + OP_LOGE(context_->GetNodeName(), "group_index must be a non-empty 1-D tensor"), + return ge::GRAPH_FAILED); + expertNum_ = static_cast(groupIndexShape.GetDim(0)); + } + + // weight_scale shape validation + const gert::Shape weightScaleShape = weightScaleShapePtr->GetStorageShape(); + const bool weightShapeValid = hasGroupIndex_ ? + (weightScaleShape.GetDimNum() == 2 && + weightScaleShape.GetDim(0) == static_cast(expertNum_) && + weightScaleShape.GetDim(1) == inDimy) : + ((weightScaleShape.GetDimNum() == 1 && weightScaleShape.GetDim(0) == inDimy) || + (weightScaleShape.GetDimNum() == 2 && weightScaleShape.GetDim(0) == 1 && + weightScaleShape.GetDim(1) == inDimy)); + OP_CHECK_IF(!weightShapeValid, + OP_LOGE(context_->GetNodeName(), "weight_scale shape does not match the group contract"), + return ge::GRAPH_FAILED); + + // activation_scale shape: [rows] or [rows, 1] + const gert::Shape activationScaleShape = activationScaleShapePtr->GetStorageShape(); + const bool actScaleValid = (activationScaleShape.GetDimNum() == 1 || + (activationScaleShape.GetDimNum() == 2 && activationScaleShape.GetDim(1) == 1)) && + activationScaleShape.GetShapeSize() == inDimx; + OP_CHECK_IF(!actScaleValid, + OP_LOGE(context_->GetNodeName(), "activation_scale must contain one FP32 value per row"), + return ge::GRAPH_FAILED); + + // bias: optional, same shape as weight_scale + auto biasShapePtr = context_->GetOptionalInputShape(INDEX_IN_BIAS); + hasDequantBias = (biasShapePtr != nullptr); + tilingData.set_dequantBiasIsEmpty(!hasDequantBias); + if (hasDequantBias) { + const gert::Shape biasShape = biasShapePtr->GetStorageShape(); + const bool biasShapeValid = hasGroupIndex_ ? + (biasShape.GetDimNum() == 2 && biasShape.GetDim(0) == static_cast(expertNum_) && + biasShape.GetDim(1) == inDimy) : + ((biasShape.GetDimNum() == 1 && biasShape.GetDim(0) == inDimy) || + (biasShape.GetDimNum() == 2 && biasShape.GetDim(0) == 1 && biasShape.GetDim(1) == inDimy)); + OP_CHECK_IF(!biasShapeValid, + OP_LOGE(context_->GetNodeName(), "bias shape must match weight_scale"), + return ge::GRAPH_FAILED); + } + + // check output shapes + auto yShapePtr = context_->GetOutputShape(0); + OP_CHECK_NULL_WITH_CONTEXT(context_, yShapePtr); + const gert::Shape yShape = yShapePtr->GetStorageShape(); + OP_CHECK_IF(yShape.GetDimNum() != 2 || yShape.GetDim(0) != inDimx || yShape.GetDim(1) != outDimy, + OP_LOGE(context_->GetNodeName(), "y shape must be [rows, input_width / 2]"), + return ge::GRAPH_FAILED); + + auto scaleShapePtr = context_->GetOutputShape(1); + OP_CHECK_NULL_WITH_CONTEXT(context_, scaleShapePtr); + const gert::Shape scaleShape = scaleShapePtr->GetStorageShape(); + OP_CHECK_IF(scaleShape.GetDimNum() != 1 || scaleShape.GetDim(0) != inDimx, + OP_LOGE(context_->GetNodeName(), "scale shape must be [rows]"), + return ge::GRAPH_FAILED); + + // UB capacity check + const uint64_t inputBytes = static_cast(inDimy) * sizeof(int32_t); + const uint64_t paramBytes = static_cast(inDimy) * sizeof(float); + const uint64_t tempBytes = static_cast(inDimy) * sizeof(float); + const uint64_t outputBytes = static_cast(outDimy) * sizeof(int8_t) + BLOCK_BYTES; + const uint64_t requiredUb = inputBytes + paramBytes + (hasDequantBias ? paramBytes : 0) + + tempBytes + outputBytes + UB_RESERVED_BUFF; + OP_CHECK_IF(aicoreParams_.ubSize < requiredUb, + OP_LOGE(context_->GetNodeName(), "UB size %lu is smaller than required %lu bytes", + aicoreParams_.ubSize, requiredUb), + return ge::GRAPH_FAILED); + + tilingData.set_expertNum(expertNum_); + tilingData.set_hasGroupIndex(hasGroupIndex_ ? 1 : 0); + tilingData.set_hasActivationScale(1); + tilingData.set_isPreDequantized(0); + tilingData.set_inputWidth(static_cast(inDimy)); + tilingData.set_outputWidth(static_cast(outDimy)); + tilingData.set_quantScaleIsEmpty(1); + tilingData.set_quantOffsetIsEmpty(1); + tilingData.set_quantIsOne(0); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus DequantSituQuantTiling::CheckInputShapesBF16() +{ + OP_CHECK_IF(quantMode != QUANT_MODE_DYNAMIC, + OP_LOGE(context_->GetNodeName(), "quant_mode must be dynamic for bfloat16 x"), + return ge::GRAPH_FAILED); + + OP_CHECK_IF(context_->GetOptionalInputShape(INDEX_IN_WEIGHT_SCALE) != nullptr, + OP_LOGE(context_->GetNodeName(), "weight_scale must be absent for bfloat16 x"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context_->GetOptionalInputShape(INDEX_IN_ACTIVATION_SCALE) != nullptr, + OP_LOGE(context_->GetNodeName(), "activation_scale must be absent for bfloat16 x"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context_->GetOptionalInputShape(INDEX_IN_BIAS) != nullptr, + OP_LOGE(context_->GetNodeName(), "bias must be absent for bfloat16 x"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context_->GetOptionalInputShape(INDEX_IN_QUANT_SCALE) != nullptr, + OP_LOGE(context_->GetNodeName(), "quant_scale must be absent for bfloat16 x"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context_->GetOptionalInputShape(INDEX_IN_QUANT_OFFSET) != nullptr, + OP_LOGE(context_->GetNodeName(), "quant_offset must be absent for bfloat16 x"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(context_->GetOptionalInputShape(INDEX_IN_GROUP_INDEX) != nullptr, + OP_LOGE(context_->GetNodeName(), "group_index must be absent for bfloat16 x"), + return ge::GRAPH_FAILED); + + // check output shapes + auto yShapePtr = context_->GetOutputShape(0); + OP_CHECK_NULL_WITH_CONTEXT(context_, yShapePtr); + const gert::Shape yShape = yShapePtr->GetStorageShape(); + OP_CHECK_IF(yShape.GetDimNum() != 2 || yShape.GetDim(0) != inDimx || yShape.GetDim(1) != outDimy, + OP_LOGE(context_->GetNodeName(), "y shape must be [rows, input_width / 2]"), + return ge::GRAPH_FAILED); + + auto scaleShapePtr = context_->GetOutputShape(1); + OP_CHECK_NULL_WITH_CONTEXT(context_, scaleShapePtr); + const gert::Shape scaleShape = scaleShapePtr->GetStorageShape(); + OP_CHECK_IF(scaleShape.GetDimNum() != 1 || scaleShape.GetDim(0) != inDimx, + OP_LOGE(context_->GetNodeName(), "scale shape must be [rows]"), + return ge::GRAPH_FAILED); + + // UB capacity check + const uint64_t inputBytes = static_cast(inDimy) * 2; // BF16 = 2 bytes + const uint64_t tempBytes = static_cast(inDimy) * sizeof(float); + const uint64_t dequantBytes = tempBytes; + const uint64_t outputBytes = static_cast(outDimy) * sizeof(int8_t) + BLOCK_BYTES; + const uint64_t requiredUb = inputBytes + tempBytes + dequantBytes + outputBytes + UB_RESERVED_BUFF; + OP_CHECK_IF(aicoreParams_.ubSize < requiredUb, + OP_LOGE(context_->GetNodeName(), "UB size %lu is smaller than required %lu bytes", + aicoreParams_.ubSize, requiredUb), + return ge::GRAPH_FAILED); + + hasDequantBias = false; + tilingData.set_dequantBiasIsEmpty(1); + tilingData.set_quantScaleIsEmpty(1); + tilingData.set_quantOffsetIsEmpty(1); + tilingData.set_quantIsOne(0); + tilingData.set_expertNum(1); + tilingData.set_hasGroupIndex(0); + tilingData.set_hasActivationScale(0); + tilingData.set_isPreDequantized(1); + tilingData.set_inputWidth(static_cast(inDimy)); + tilingData.set_outputWidth(static_cast(outDimy)); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus DequantSituQuantTiling::CheckInputShapes() +{ + auto xShapePtr = context_->GetInputShape(INDEX_IN_X); + OP_CHECK_NULL_WITH_CONTEXT(context_, xShapePtr); + const gert::Shape xShape = xShapePtr->GetStorageShape(); + int64_t dimNum = xShape.GetDimNum(); + + int64_t shapeBefore = 1; + for (int64_t i = 0; i < dimNum - 1; i++) { + shapeBefore *= xShape.GetDim(i); + } + inDimy = xShape.GetDim(dimNum - 1); + CHECK_FAIL(context_, inDimy % NUM_TWO != 0, "The last dim of x must be even number"); + + inDimx = shapeBefore; + outDimy = inDimy / NUM_TWO; + + tilingData.set_rowLen(inDimx); + tilingData.set_colLen(outDimy); + + if (xDtype_ == ge::DT_INT8) { + return CheckInputShapesInt8(dimNum, inDimy, outDimy); + } else if (xDtype_ == ge::DT_INT32) { + CHECK_FAIL(context_, dimNum != 2, "x shape rank must be 2 for int32, but is %ld", dimNum); + return ValidateInt32Contract(); + } else if (xDtype_ == ge::DT_BF16) { + CHECK_FAIL(context_, dimNum != 2, "x shape rank must be 2 for bfloat16, but is %ld", dimNum); + return CheckInputShapesBF16(); + } + return ge::GRAPH_FAILED; +} + +ge::graphStatus DequantSituQuantTiling::GetShapeAttrsInfo() +{ + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus DequantSituQuantTiling::GetShapeAttrsInfoInner() +{ + opName = context_->GetNodeName(); + + auto inputDesc = context_->GetInputDesc(INDEX_IN_X); + OP_CHECK_NULL_WITH_CONTEXT(context_, inputDesc); + xDtype_ = inputDesc->GetDataType(); + OP_CHECK_IF(xDtype_ != ge::DT_INT8 && xDtype_ != ge::DT_INT32 && xDtype_ != ge::DT_BF16, + OP_LOGE_FOR_INVALID_DTYPE(context_->GetNodeName(), "x", + ge::TypeUtils::DataTypeToSerialString(xDtype_).c_str(), + "int8/int32/bfloat16"), + return ge::GRAPH_FAILED); + + isPreDequantized_ = (xDtype_ == ge::DT_BF16); + + if (xDtype_ == ge::DT_INT8) { + inputDTypeLen = 1; + } else if (xDtype_ == ge::DT_INT32) { + inputDTypeLen = 4; + } else { + inputDTypeLen = 2; + } + + const gert::RuntimeAttrs* attrs = context_->GetAttrs(); + OP_CHECK_NULL_WITH_CONTEXT(context_, attrs); + if (!SetAttrs(attrs)) { + return ge::GRAPH_FAILED; + } + + if (CheckInputShapes() == ge::GRAPH_FAILED) { + return ge::GRAPH_FAILED; + } + + return ge::GRAPH_SUCCESS; +} + +bool DequantSituQuantTiling::CalcUbMaxTileLen(const uint64_t ubSize, uint32_t& maxTileLen) +{ + uint64_t bytesPerElement = 8 + 2 + 16 + 16 + 4; // 46 (weightScale + inQueueX + tmpBuf + castBuf + outQueue) + if (hasDequantBias) { + bytesPerElement += 8; + } + if (hasQuantScale && !quantIsOne) { + bytesPerElement += (quantMode == QUANT_MODE_DYNAMIC) ? 4 : 8; + } + + uint64_t availableUb = ubSize - UB_RESERVED_BUFF - ALIGN_UINT_IN_CACHE_32B; + uint64_t maxElements = availableUb / bytesPerElement; + maxTileLen = static_cast(maxElements / ALIGN_UINT_IN_CACHE_32B * ALIGN_UINT_IN_CACHE_32B); + OP_LOGI(opName, "CalcUbMaxTileLen ubSize:%lu, maxTileLen:%u, bytesPerElement:%lu", ubSize, maxTileLen, + bytesPerElement); + return true; +} + +bool DequantSituQuantTiling::CalcOptBaseShape(uint32_t maxTileLen) +{ + uint32_t colLen = static_cast(outDimy); + uint32_t rowLen = static_cast(inDimx); + + uint32_t baseColLen = std::min(colLen, maxTileLen); + if (baseColLen < colLen && baseColLen > cacheLineLen) { + baseColLen = baseColLen / cacheLineLen * cacheLineLen; + } + if (baseColLen == 0) { + baseColLen = std::min(colLen, static_cast(ALIGN_UINT_IN_CACHE_32B)); + } + + uint32_t baseRowLen = 1; + + optBaseRowLen = baseRowLen; + optBaseColLen = baseColLen; + + totalUsedCoreNum = std::min(rowLen, static_cast(totalCore)); + if (colLen < PERFORMANCE_COL_LEN && rowLen < PERFORMANCE_ROW_LEN) { + totalUsedCoreNum = std::min(totalUsedCoreNum, static_cast(MIN_CORE)); + } + totalUsedCoreNum = std::min(totalUsedCoreNum, static_cast(MAX_CORE_NUMBER)); + + return true; +} + +bool DequantSituQuantTiling::CalcTiling(const uint32_t totalCores, const uint64_t ubSize) +{ + ubMinBlockLen = ALIGN_UINT_IN_CACHE_32B / inputDTypeLen; + cacheLineLen = PACK_UINT_IN_CACHE_512B / inputDTypeLen; + + if (xDtype_ == ge::DT_INT32 || xDtype_ == ge::DT_BF16) { + // INT32/BF16 path: no column tiling, whole row in UB + tilingData.set_is32BAligned(1); + tilingData.set_isDoubleBuffer(1); + tilingData.set_baseRowLen(1); + tilingData.set_baseColLen(static_cast(outDimy)); + totalUsedCoreNum = inDimx == 0 ? 0 : std::min(static_cast(inDimx), totalCore); + tilingData.set_usedCoreNum(totalUsedCoreNum); + return true; + } + + // INT8 path + tilingData.set_is32BAligned(static_cast(outDimy % ubMinBlockLen == 0)); + tilingData.set_isDoubleBuffer(1); + + if (!CalcUbMaxTileLen(ubSize, maxTileLen)) { + return false; + } + if (!CalcOptBaseShape(maxTileLen)) { + return false; + } + + tilingData.set_baseRowLen(optBaseRowLen); + tilingData.set_baseColLen(optBaseColLen); + tilingData.set_usedCoreNum(totalUsedCoreNum); + return true; +} + +ge::graphStatus DequantSituQuantTiling::DoOpTiling() +{ + if (GetShapeAttrsInfoInner() == ge::GRAPH_FAILED) { + return ge::GRAPH_FAILED; + } + if (!CalcTiling(totalCore, aicoreParams_.ubSize)) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus DequantSituQuantTiling::DoLibApiTiling() +{ + return ge::GRAPH_SUCCESS; +} + +uint64_t DequantSituQuantTiling::GetTilingKey() const +{ + if (xDtype_ == ge::DT_INT32) { + return DSQ_INT32_DYNAMIC; + } + if (xDtype_ == ge::DT_BF16) { + return DSQ_BF16_DYNAMIC; + } + // INT8 path + if (quantMode == QUANT_MODE_STATIC) { + if (quantIsOne) { + return hasDequantBias ? DSQ_STATIC_QUANT_ONE_BIAS : DSQ_STATIC_QUANT_ONE; + } else { + return hasDequantBias ? DSQ_STATIC_QUANT_VEC_BIAS : DSQ_STATIC_QUANT_VEC; + } + } else { + if (hasQuantScale) { + return hasDequantBias ? DSQ_DYNAMIC_QUANT_SMOOTH_BIAS : DSQ_DYNAMIC_QUANT_SMOOTH; + } else { + return hasDequantBias ? DSQ_DYNAMIC_QUANT_NO_SMOOTH_BIAS : DSQ_DYNAMIC_QUANT_NO_SMOOTH; + } + } +} + +ge::graphStatus DequantSituQuantTiling::GetWorkspaceSize() +{ + if (xDtype_ == ge::DT_INT32 || xDtype_ == ge::DT_BF16) { + workspaceSize_ = 0; + return ge::GRAPH_SUCCESS; + } + workspaceSize_ = USER_WORKSPACE; + if (quantMode == QUANT_MODE_DYNAMIC && (static_cast(outDimy) > optBaseColLen)) { + workspaceSize_ += (totalUsedCoreNum * static_cast(outDimy) * sizeof(float)); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus DequantSituQuantTiling::PostTiling() +{ + context_->SetTilingKey(GetTilingKey()); + if (xDtype_ == ge::DT_INT32 || xDtype_ == ge::DT_BF16) { + context_->SetBlockDim(totalUsedCoreNum == 0 ? 1 : totalUsedCoreNum); + } else { + context_->SetBlockDim(totalUsedCoreNum); + } + size_t* workspaces = context_->GetWorkspaceSizes(1); + workspaces[0] = workspaceSize_; + OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetRawTilingData()); + tilingData.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity()); + context_->GetRawTilingData()->SetDataSize(tilingData.GetDataSize()); + return ge::GRAPH_SUCCESS; +} + +void DequantSituQuantTiling::ShowTilingData() +{ + std::ostringstream info; + info << "rowLen: " << tilingData.get_rowLen(); + info << ", colLen: " << tilingData.get_colLen(); + info << ", baseRowLen: " << tilingData.get_baseRowLen(); + info << ", baseColLen: " << tilingData.get_baseColLen(); + info << ", usedCoreNum: " << tilingData.get_usedCoreNum(); + info << ", quantMode: " << tilingData.get_quantMode(); + info << ", quantIsOne: " << tilingData.get_quantIsOne(); + info << ", beta: " << tilingData.get_beta(); + info << ", linearBeta: " << tilingData.get_linearBeta(); + info << ", dequantBiasIsEmpty: " << tilingData.get_dequantBiasIsEmpty(); + info << ", expertNum: " << tilingData.get_expertNum(); + info << ", hasGroupIndex: " << tilingData.get_hasGroupIndex(); + info << ", hasActivationScale: " << tilingData.get_hasActivationScale(); + info << ", isPreDequantized: " << tilingData.get_isPreDequantized(); + info << ", inputWidth: " << tilingData.get_inputWidth(); + info << ", outputWidth: " << tilingData.get_outputWidth(); + OP_LOGI(opName, "%s", info.str().c_str()); +} + +REGISTER_TILING_TEMPLATE("DequantSituQuant", DequantSituQuantTiling, 0); + +ge::graphStatus TilingForDequantSituQuant(gert::TilingContext* context) +{ + return TilingRegistry::GetInstance().DoTilingImpl(context); +} + +ge::graphStatus TilingPrepareForDequantSituQuant(gert::TilingParseContext* context) +{ + OP_LOGD(context, "TilingPrepare4DequantSituQuant enter."); + auto compileInfo = context->GetCompiledInfo(); + OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo); + auto platformInfo = context->GetPlatformInfo(); + OP_CHECK_NULL_WITH_CONTEXT(context, platformInfo); + auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo); + compileInfo->coreNum = ascendcPlatform.GetCoreNumAiv(); + OP_CHECK_IF((compileInfo->coreNum <= 0), + OP_LOGE(context->GetNodeName(), "Get core num failed, core num: %u", + static_cast(compileInfo->coreNum)), + return ge::GRAPH_FAILED); + + uint64_t ubSize; + ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize); + compileInfo->ubSize = ubSize; + OP_CHECK_IF((compileInfo->ubSize <= 0), + OP_LOGE(context->GetNodeName(), "Get ub size failed, ub size: %u", + static_cast(compileInfo->ubSize)), + return ge::GRAPH_FAILED); + + OP_LOGD(context, "TilingPrepare4DequantSituQuant exit."); + return ge::GRAPH_SUCCESS; +} + +IMPL_OP_OPTILING(DequantSituQuant) + .Tiling(TilingForDequantSituQuant) + .TilingParse(TilingPrepareForDequantSituQuant); + +} // namespace optiling diff --git a/csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_tiling.h b/csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_tiling.h new file mode 100644 index 000000000000..d28c3e06aca1 --- /dev/null +++ b/csrc/moe/dequant_situ_quant/op_host/dequant_situ_quant_tiling.h @@ -0,0 +1,149 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file dequant_situ_quant_tiling.h + * \brief + */ + +#ifndef DEQUANT_SITU_QUANT_TILING_H +#define DEQUANT_SITU_QUANT_TILING_H + +#include +#include +#include "register/op_impl_registry.h" +#include "util/math_util.h" +#include "log/log.h" +#include "tiling/platform/platform_ascendc.h" +#include "platform/platform_infos_def.h" +#include "register/tilingdata_base.h" +#include "tiling/tiling_api.h" +#include "../op_graph/dequant_situ_quant_proto.h" +#include "../../dequant_swiglu_quant/tiling_base/tiling_base.h" +#include "../../dequant_swiglu_quant/tiling_base/tiling_templates_registry.h" + +using namespace Ops::NN::Optiling; +namespace optiling { + +BEGIN_TILING_DATA_DEF(DequantSituQuantTilingData) +TILING_DATA_FIELD_DEF(uint32_t, is32BAligned); +TILING_DATA_FIELD_DEF(uint32_t, isDoubleBuffer); +TILING_DATA_FIELD_DEF(uint64_t, rowLen); +TILING_DATA_FIELD_DEF(uint64_t, colLen); +TILING_DATA_FIELD_DEF(uint32_t, baseRowLen); +TILING_DATA_FIELD_DEF(uint32_t, baseColLen); +TILING_DATA_FIELD_DEF(uint32_t, activateLeft); +TILING_DATA_FIELD_DEF(uint32_t, dequantBiasIsEmpty); +TILING_DATA_FIELD_DEF(uint32_t, quantScaleIsEmpty); +TILING_DATA_FIELD_DEF(uint32_t, quantOffsetIsEmpty); +TILING_DATA_FIELD_DEF(uint32_t, quantIsOne); +TILING_DATA_FIELD_DEF(uint32_t, usedCoreNum); +TILING_DATA_FIELD_DEF(int64_t, quantMode); +TILING_DATA_FIELD_DEF(float, beta); +TILING_DATA_FIELD_DEF(float, linearBeta); +TILING_DATA_FIELD_DEF(uint32_t, expertNum); +TILING_DATA_FIELD_DEF(uint32_t, hasGroupIndex); +TILING_DATA_FIELD_DEF(uint32_t, hasActivationScale); +TILING_DATA_FIELD_DEF(uint32_t, isPreDequantized); +TILING_DATA_FIELD_DEF(uint32_t, inputWidth); +TILING_DATA_FIELD_DEF(uint32_t, outputWidth); +END_TILING_DATA_DEF; + +REGISTER_TILING_DATA_CLASS(DequantSituQuant, DequantSituQuantTilingData) + +struct DequantSituQuantCompileInfo { + uint64_t coreNum = 0; + uint64_t ubSize = 0; +}; + +constexpr int64_t DSQ_STATIC_QUANT_ONE = 10000; +constexpr int64_t DSQ_STATIC_QUANT_VEC = 10001; +constexpr int64_t DSQ_STATIC_QUANT_ONE_BIAS = 10002; +constexpr int64_t DSQ_STATIC_QUANT_VEC_BIAS = 10003; +constexpr int64_t DSQ_DYNAMIC_QUANT_NO_SMOOTH = 20000; +constexpr int64_t DSQ_DYNAMIC_QUANT_SMOOTH = 20001; +constexpr int64_t DSQ_DYNAMIC_QUANT_NO_SMOOTH_BIAS = 20002; +constexpr int64_t DSQ_DYNAMIC_QUANT_SMOOTH_BIAS = 20003; +constexpr int64_t DSQ_INT32_DYNAMIC = 30000; +constexpr int64_t DSQ_BF16_DYNAMIC = 40000; + +class DequantSituQuantTiling : public Ops::NN::Optiling::TilingBaseClass { +public: + explicit DequantSituQuantTiling(gert::TilingContext* cont) : TilingBaseClass(cont) { Reset(); } + ~DequantSituQuantTiling() override = default; + + void Reset(gert::TilingContext* cont) override + { + TilingBaseClass::Reset(cont); + Reset(); + } + +protected: + bool IsCapable() override; + ge::graphStatus GetPlatformInfo() override; + ge::graphStatus GetShapeAttrsInfo() override; + ge::graphStatus DoOpTiling() override; + ge::graphStatus DoLibApiTiling() override; + uint64_t GetTilingKey() const override; + ge::graphStatus GetWorkspaceSize() override; + ge::graphStatus PostTiling() override; + void Reset(); + +private: + void ShowTilingData(); + ge::graphStatus GetShapeAttrsInfoInner(); + ge::graphStatus CheckInputShapes(); + ge::graphStatus CheckInputShapesInt8(int64_t dimNum, int64_t inDimy, int64_t outDimy); + ge::graphStatus CheckInputShapesInt32(int64_t inDimy, int64_t outDimy); + ge::graphStatus CheckInputShapesBF16(); + bool SetAttrs(const gert::RuntimeAttrs* attrs); + bool CalcTiling(const uint32_t totalCores, const uint64_t ubSize); + bool CalcUbMaxTileLen(const uint64_t ubSize, uint32_t& maxTileLen); + bool CalcOptBaseShape(uint32_t maxTileLen); + ge::graphStatus ValidateInt32Contract(); + + const char* opName = ""; + uint32_t totalCore = 0; + uint32_t totalUsedCoreNum = 0; + uint32_t inputDTypeLen = 1; + uint32_t ubMinBlockLen = 32; + uint32_t cacheLineLen = 512; + uint32_t maxTileLen = 0; + uint32_t optBaseRowLen = 0; + uint32_t optBaseColLen = 0; + uint64_t workspaceSize_ = 0; + + bool hasDequantBias = false; + bool hasQuantScale = false; + bool hasQuantOffset = false; + bool quantIsOne = false; + uint32_t activateLeft = 0; + int64_t quantMode = 0; + float beta = 4.0f; + float linearBeta = 25.0f; + uint64_t quantScaleShapeSize = 0; + + int64_t inDimx = 0; + int64_t inDimy = 0; + int64_t outDimy = 0; + + ge::DataType xDtype_ = ge::DT_INT8; + bool isPreDequantized_ = false; + bool hasWeightScale_ = false; + bool hasActivationScale_ = false; + bool hasGroupIndex_ = false; + uint32_t expertNum_ = 1; + + DequantSituQuantTilingData tilingData; + platform_ascendc::SocVersion socVersion = platform_ascendc::SocVersion::ASCEND910B; +}; + +} // namespace optiling +#endif // DEQUANT_SITU_QUANT_TILING_H diff --git a/csrc/moe/dequant_situ_quant/op_kernel/dequant_situ_quant.cpp b/csrc/moe/dequant_situ_quant/op_kernel/dequant_situ_quant.cpp new file mode 100644 index 000000000000..ea15eb835ed2 --- /dev/null +++ b/csrc/moe/dequant_situ_quant/op_kernel/dequant_situ_quant.cpp @@ -0,0 +1,103 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file dequant_situ_quant.cpp + * \brief Kernel entry point for DequantSituQuant operator. + */ + +#include "kernel_operator.h" +#include "kernel_tiling/kernel_tiling.h" +#include "dequant_situ_quant.h" + +using namespace AscendC; + +#define DSQ_STATIC_QUANT_ONE 10000 +#define DSQ_STATIC_QUANT_VEC 10001 +#define DSQ_STATIC_QUANT_ONE_BIAS 10002 +#define DSQ_STATIC_QUANT_VEC_BIAS 10003 +#define DSQ_DYNAMIC_QUANT_NO_SMOOTH 20000 +#define DSQ_DYNAMIC_QUANT_SMOOTH 20001 +#define DSQ_DYNAMIC_QUANT_NO_SMOOTH_BIAS 20002 +#define DSQ_DYNAMIC_QUANT_SMOOTH_BIAS 20003 +#define DSQ_INT32_DYNAMIC 30000 +#define DSQ_BF16_DYNAMIC 40000 + +extern "C" __global__ __aicore__ void dequant_situ_quant( + GM_ADDR xGM, GM_ADDR weightScaleGM, GM_ADDR activationScaleGM, GM_ADDR biasGM, + GM_ADDR quantScaleGM, GM_ADDR quantOffsetGM, GM_ADDR groupIndexGM, + GM_ADDR yGM, GM_ADDR scaleGM, GM_ADDR workspace, GM_ADDR tiling) +{ +#if (ORIG_DTYPE_X == DT_INT8) + if (workspace == nullptr) { + return; + } + GM_ADDR userspace = GetUserWorkspace(workspace); + if (userspace == nullptr) { + return; + } + + TPipe pipe; + GET_TILING_DATA_WITH_STRUCT(DequantSituQuantTilingData, tilingDataIn, tiling); + const DequantSituQuantTilingData* __restrict__ tilingData = &tilingDataIn; + + if (TILING_KEY_IS(DSQ_STATIC_QUANT_ONE)) { + DequantSituQuantOps::DequantSituQuantKernel op(&pipe); + op.Init(xGM, weightScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace, tilingData); + op.Process(); + } else if (TILING_KEY_IS(DSQ_STATIC_QUANT_VEC)) { + DequantSituQuantOps::DequantSituQuantKernel op(&pipe); + op.Init(xGM, weightScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace, tilingData); + op.Process(); + } else if (TILING_KEY_IS(DSQ_STATIC_QUANT_ONE_BIAS)) { + DequantSituQuantOps::DequantSituQuantKernel op(&pipe); + op.Init(xGM, weightScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace, tilingData); + op.Process(); + } else if (TILING_KEY_IS(DSQ_STATIC_QUANT_VEC_BIAS)) { + DequantSituQuantOps::DequantSituQuantKernel op(&pipe); + op.Init(xGM, weightScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace, tilingData); + op.Process(); + } else if (TILING_KEY_IS(DSQ_DYNAMIC_QUANT_NO_SMOOTH)) { + DequantSituQuantOps::DequantSituQuantKernel op(&pipe); + op.Init(xGM, weightScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace, tilingData); + op.Process(); + } else if (TILING_KEY_IS(DSQ_DYNAMIC_QUANT_SMOOTH)) { + DequantSituQuantOps::DequantSituQuantKernel op(&pipe); + op.Init(xGM, weightScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace, tilingData); + op.Process(); + } else if (TILING_KEY_IS(DSQ_DYNAMIC_QUANT_NO_SMOOTH_BIAS)) { + DequantSituQuantOps::DequantSituQuantKernel op(&pipe); + op.Init(xGM, weightScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace, tilingData); + op.Process(); + } else if (TILING_KEY_IS(DSQ_DYNAMIC_QUANT_SMOOTH_BIAS)) { + DequantSituQuantOps::DequantSituQuantKernel op(&pipe); + op.Init(xGM, weightScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace, tilingData); + op.Process(); + } +#elif (ORIG_DTYPE_X == DT_INT32) + if (TILING_KEY_IS(DSQ_INT32_DYNAMIC)) { + TPipe pipe; + GET_TILING_DATA_WITH_STRUCT(DequantSituQuantTilingData, tilingDataIn, tiling); + const DequantSituQuantTilingData* __restrict__ tilingData = &tilingDataIn; + DequantSituQuantOps::DequantSituQuantK3Kernel op(&pipe); + op.Init(xGM, weightScaleGM, activationScaleGM, biasGM, groupIndexGM, yGM, scaleGM, tilingData); + op.Process(); + } +#elif (ORIG_DTYPE_X == DT_BF16) + if (TILING_KEY_IS(DSQ_BF16_DYNAMIC)) { + TPipe pipe; + GET_TILING_DATA_WITH_STRUCT(DequantSituQuantTilingData, tilingDataIn, tiling); + const DequantSituQuantTilingData* __restrict__ tilingData = &tilingDataIn; + DequantSituQuantOps::DequantSituQuantK3Kernel op(&pipe); + op.Init(xGM, weightScaleGM, activationScaleGM, biasGM, groupIndexGM, yGM, scaleGM, tilingData); + op.Process(); + } +#endif +} diff --git a/csrc/moe/dequant_situ_quant/op_kernel/dequant_situ_quant.h b/csrc/moe/dequant_situ_quant/op_kernel/dequant_situ_quant.h new file mode 100644 index 000000000000..f2e55729bd53 --- /dev/null +++ b/csrc/moe/dequant_situ_quant/op_kernel/dequant_situ_quant.h @@ -0,0 +1,1060 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file dequant_situ_quant.h + * \brief DequantSituQuant kernel: Dequant -> Situ -> Quant + */ + +#ifndef DEQUANT_SITU_QUANT_H +#define DEQUANT_SITU_QUANT_H + +#include "kernel_tiling/kernel_tiling.h" +#include "kernel_operator.h" +#include + +#define TEMPLATE_DSQ_DECLARE template +#define TEMPLATE_DSQ_ARGS hasDequantBias + +namespace DequantSituQuantOps { +using namespace AscendC; + +constexpr static int64_t DB_BUFFER = 1; +constexpr static int64_t BLOCK_SIZE = 32; +constexpr static int64_t BLOCK_ELEM = BLOCK_SIZE / sizeof(float); +constexpr static int64_t MASK_NUM_T32 = 256 / sizeof(float); +constexpr static int64_t MASK_BLK_STRIDE = 8; +constexpr static int64_t ELEM_PER_REP_FP32 = 64; +constexpr static int64_t MAX_REPEAT = 255; +constexpr static int64_t SWI_FACTOR = 2; +constexpr static float DYNAMIC_QUANT_FACTOR = 1.0 / 127.0; + +TEMPLATE_DSQ_DECLARE +class DequantSituQuantKernel { +public: + __aicore__ inline DequantSituQuantKernel(TPipe* pipe) { pipe_ = pipe; } + + __aicore__ inline void Init(GM_ADDR x, GM_ADDR dequantScale, GM_ADDR dequantBias, GM_ADDR quantScale, + GM_ADDR quantOffset, GM_ADDR y, GM_ADDR scale, GM_ADDR workspace, + const DequantSituQuantTilingData* tilingData) + { + tl_ = tilingData; + blockIdx_ = GetBlockIdx(); + + rowLen_ = tl_->rowLen; + colLen_ = tl_->colLen; + inDimy_ = colLen_ * SWI_FACTOR; + outDimy_ = colLen_; + baseRowLen_ = tl_->baseRowLen; + baseColLen_ = tl_->baseColLen < colLen_ ? tl_->baseColLen : colLen_; + usedCoreNum_ = tl_->usedCoreNum; + activateLeft_ = tl_->activateLeft; + quantMode_ = tl_->quantMode; + quantIsOne_ = tl_->quantIsOne; + quantScaleIsEmpty_ = tl_->quantScaleIsEmpty; + quantOffsetIsEmpty_ = tl_->quantOffsetIsEmpty; + beta_ = tl_->beta; + linearBeta_ = tl_->linearBeta; + + if (rowLen_ < usedCoreNum_) { + usedCoreNum_ = rowLen_; + } + int64_t perRoundCnt = usedCoreNum_ == 0 ? 0 : rowLen_ / usedCoreNum_; + int64_t remainCnt = rowLen_ - usedCoreNum_ * perRoundCnt; + curCoreRowNum_ = perRoundCnt; + if (blockIdx_ < remainCnt) { + curCoreRowNum_ = perRoundCnt + 1; + inputCopyOffset_ = blockIdx_ * curCoreRowNum_; + } else { + inputCopyOffset_ = remainCnt * (perRoundCnt + 1) + (blockIdx_ - remainCnt) * perRoundCnt; + } + + xGm_.SetGlobalBuffer((__gm__ int8_t*)x + inputCopyOffset_ * inDimy_, curCoreRowNum_ * inDimy_); + dequantScaleGm_.SetGlobalBuffer((__gm__ float*)dequantScale); + if constexpr (hasDequantBias) { + dequantBiasGm_.SetGlobalBuffer((__gm__ float*)dequantBias); + } + if (quantScaleIsEmpty_ == 0) { + quantScaleGm_.SetGlobalBuffer((__gm__ float*)quantScale); + } + if (quantOffsetIsEmpty_ == 0) { + quantOffsetGm_.SetGlobalBuffer((__gm__ float*)quantOffset); + } + yGm_.SetGlobalBuffer((__gm__ int8_t*)y + inputCopyOffset_ * outDimy_, curCoreRowNum_ * outDimy_); + scaleGm_.SetGlobalBuffer((__gm__ float*)scale + inputCopyOffset_, curCoreRowNum_); + + if (quantScaleIsEmpty_ == 0 && quantIsOne_) { + quantScaleVal_ = quantScaleGm_.GetValue(0); + if (quantScaleVal_ == 0.0f) { + quantScaleVal_ = 1.0f; + } else { + quantScaleVal_ = 1.0f / quantScaleVal_; + } + } + if (quantOffsetIsEmpty_ == 0 && quantIsOne_) { + quantOffsetVal_ = quantOffsetGm_.GetValue(0); + } + + // Check if dequant_scale is scalar (shape [1]) + dequantScaleIsOne_ = (dequantScaleGm_.GetSize() == 1); + if (dequantScaleIsOne_) { + dequantScaleVal_ = dequantScaleGm_.GetValue(0); + } + if constexpr (hasDequantBias) { + dequantBiasIsOne_ = (dequantBiasGm_.GetSize() == 1); + if (dequantBiasIsOne_) { + dequantBiasVal_ = dequantBiasGm_.GetValue(0); + } + } + + curColNum_ = baseColLen_; + InitUbBuffer(); + } + + __aicore__ inline void Process() + { + if (blockIdx_ >= usedCoreNum_) { + return; + } + processCompute(); + } + +protected: + __aicore__ inline void InitUbBuffer() + { + int64_t alignColNum = curColNum_ == Align(curColNum_, sizeof(int8_t)) ? + curColNum_ : + Align(curColNum_, sizeof(int8_t)); + int64_t alignInDimy = alignColNum * SWI_FACTOR; + + pipe_->InitBuffer(inQueueX_, DB_BUFFER, alignInDimy * sizeof(int8_t)); + pipe_->InitBuffer(dequantScaleBuf_, alignInDimy * sizeof(float)); + if constexpr (hasDequantBias) { + pipe_->InitBuffer(dequantBiasBuf_, alignInDimy * sizeof(float)); + } + if (quantScaleIsEmpty_ == 0 && !quantIsOne_) { + int64_t quantBufElems = (quantMode_ == 1) ? alignColNum : alignInDimy; + pipe_->InitBuffer(quantBuf_, quantBufElems * sizeof(float)); + } + // outQueue_ must hold max(int8 output, float situ output) + scale + padding + // Situ computation uses outQueue as float buffer [H] floats = H*4 bytes + // Final output is [H] int8 = H bytes + scale [1] float = 4 bytes + pipe_->InitBuffer(outQueue_, 1, alignColNum * sizeof(float) + sizeof(float) + BLOCK_SIZE); + + // temp buffers for compute: dequantOut[2H] + situTemp[2H] = 4H floats + pipe_->InitBuffer(tmpBuf_, alignInDimy * SWI_FACTOR * sizeof(float)); + // cast buffer for int8<->float conversion intermediates + pipe_->InitBuffer(castBuf_, alignInDimy * SWI_FACTOR * sizeof(float)); + } + + __aicore__ inline void CopyInDequantParams(int64_t colOffset) + { + // Sync V→MTE2: ensure previous tile's V operations finish before + // overwriting dequantScaleBuf_ (TBuf has no automatic pipeline sync) + event_t eventV2MTE2 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_MTE2)); + SetFlag(eventV2MTE2); + WaitFlag(eventV2MTE2); + + if (!dequantScaleIsOne_) { + // x layout: [up(0:H), gate(H:2H)] — load up and gate separately + DataCopyExtParams params = {1, static_cast(curColNum_ * sizeof(float)), 0, 0, 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + LocalTensor dequantScaleLocal = dequantScaleBuf_.template Get(); + // Load dequant_scale for up: ds[colOffset : colOffset + curColNum] + DataCopyPad(dequantScaleLocal, dequantScaleGm_[colOffset], params, padParams); + // Load dequant_scale for gate: ds[outDimy_ + colOffset : outDimy_ + colOffset + curColNum] + DataCopyPad(dequantScaleLocal[curColNum_], dequantScaleGm_[outDimy_ + colOffset], params, padParams); + dequantScaleLocal_ = dequantScaleLocal; + } + + if constexpr (hasDequantBias) { + if (!dequantBiasIsOne_) { + DataCopyExtParams params = {1, static_cast(curColNum_ * sizeof(float)), 0, 0, 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + LocalTensor dequantBiasLocal = dequantBiasBuf_.template Get(); + DataCopyPad(dequantBiasLocal, dequantBiasGm_[colOffset], params, padParams); + DataCopyPad(dequantBiasLocal[curColNum_], dequantBiasGm_[outDimy_ + colOffset], params, padParams); + dequantBiasLocal_ = dequantBiasLocal; + } + } + + // Sync MTE2→V: ensure TBuf DataCopyPad completes before vector compute reads it + event_t eventMTE2ToV = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V)); + SetFlag(eventMTE2ToV); + WaitFlag(eventMTE2ToV); + } + + __aicore__ inline void CopyInQuantParams(int64_t colOffset) + { + if (quantScaleIsEmpty_ == 0 && !quantIsOne_) { + DataCopyExtParams params = {1, static_cast(curColNum_ * sizeof(float)), 0, 0, 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + LocalTensor quantLocal = quantBuf_.template Get(); + DataCopyPad(quantLocal, quantScaleGm_[colOffset], params, padParams); + if (quantOffsetIsEmpty_ == 0) { + DataCopyPad(quantLocal[curColNum_], quantOffsetGm_[colOffset], params, padParams); + } + quantLocal_ = quantLocal; + + // Sync MTE2→V: ensure TBuf DataCopyPad completes before vector compute reads it + event_t eventMTE2ToV = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V)); + SetFlag(eventMTE2ToV); + WaitFlag(eventMTE2ToV); + } + } + + __aicore__ inline void CopyIn(int64_t rowIdx, int64_t colOffset) + { + // x layout: [up(0:H), gate(H:2H)] — load up and gate separately + DataCopyExtParams params = {1, static_cast(curColNum_ * sizeof(int8_t)), 0, 0, 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + + LocalTensor xLocal = inQueueX_.template AllocTensor(); + // Load up: x[rowIdx * inDimy + colOffset : ... + curColNum] + DataCopyPad(xLocal, xGm_[rowIdx * inDimy_ + colOffset], params, padParams); + // Load gate: x[rowIdx * inDimy + outDimy_ + colOffset : ... + curColNum] + DataCopyPad(xLocal[curColNum_], xGm_[rowIdx * inDimy_ + outDimy_ + colOffset], params, padParams); + inQueueX_.EnQue(xLocal); + } + + __aicore__ inline void ComputeDequant(int64_t rowIdx) + { + LocalTensor xLocalI8 = inQueueX_.template DeQue(); + + LocalTensor tmpF32 = tmpBuf_.template Get(); + int64_t tileLen = curColNum_ * SWI_FACTOR; + LocalTensor dequantOut = tmpF32; + LocalTensor situTemp = tmpF32[tileLen]; + + // Step 1: Cast int8 -> half -> fp32 + LocalTensor tmpHalf = castBuf_.template Get(); + Cast(tmpHalf, xLocalI8, RoundMode::CAST_NONE, tileLen); + PipeBarrier(); + Cast(dequantOut, tmpHalf, RoundMode::CAST_NONE, tileLen); + PipeBarrier(); + inQueueX_.FreeTensor(xLocalI8); + + // Step 2: Mul dequant_scale + if (dequantScaleIsOne_) { + Muls(dequantOut, dequantOut, dequantScaleVal_, tileLen); + } else { + Mul(dequantOut, dequantOut, dequantScaleLocal_, tileLen); + } + PipeBarrier(); + + // Step 3: Add dequant_bias (if exists) + if constexpr (hasDequantBias) { + if (dequantBiasIsOne_) { + Adds(dequantOut, dequantOut, dequantBiasVal_, tileLen); + PipeBarrier(); + } else { + Add(dequantOut, dequantOut, dequantBiasLocal_, tileLen); + PipeBarrier(); + } + } + + // Store dequantOut for Situ computation + dequantOut_ = dequantOut; + situTemp_ = situTemp; + } + + __aicore__ inline void ComputeSitu() + { + int64_t H = curColNum_; + LocalTensor dequantOut = dequantOut_; + LocalTensor tmp = situTemp_; + + // gate and up: activateLeft=0 means gate=right half, up=left half + // activateLeft=1 means gate=left half, up=right half + int64_t gateOffset = (activateLeft_ == 1) ? 0 : H; + int64_t upOffset = (activateLeft_ == 1) ? H : 0; + + LocalTensor gate = dequantOut[gateOffset]; + LocalTensor up = dequantOut[upOffset]; + + // tmpBuf_ layout: [0:2H] = dequantOut (no longer needed), [2H:4H] = situTemp + // Reuse situTemp for Situ computation: + // tmp[0:H] = tanh result (beta * tanh(gate/beta)) + // tmp[H:2H] = sigmoid result + sigmoid denom + LocalTensor tanhResult = tmp; + LocalTensor sigmoidResult = tmp[H]; + + // Step 1: tanh(gate / beta) * beta + float invBeta = 1.0f / beta_; + Muls(tanhResult, gate, invBeta, H); + PipeBarrier(); + Tanh(tanhResult, tanhResult, H); + PipeBarrier(); + + Muls(tanhResult, tanhResult, beta_, H); + PipeBarrier(); + + // Step 2: sigmoid(gate) = 1 / (1 + exp(-gate)) + // Numerically stable: avoids positive-input exp overflow. + LocalTensor denomTmp = dequantOut[gateOffset]; + Muls(sigmoidResult, gate, -1.0f, H); + PipeBarrier(); + Exp(sigmoidResult, sigmoidResult, H); + PipeBarrier(); + + Adds(denomTmp, sigmoidResult, 1.0f, H); // 1 + exp(-gate) + PipeBarrier(); + + // sigmoid = 1 / (1 + exp(-gate)) + // Use Level 0 Div instead of Reciprocal for better precision. + // src0 is a single datablock of 1.0f, reused across all repeats via + // src0BlkStride=0 and src0RepStride=0. + LocalTensor onesBlock = castBuf_.template Get(); + Duplicate(onesBlock, 1.0f, 8); + PipeBarrier(); + + constexpr uint64_t maskFp32 = static_cast(ELEM_PER_REP_FP32); + uint32_t fullReps = static_cast(H / maskFp32); + uint32_t remainder = static_cast(H % maskFp32); + BinaryRepeatParams divParams(1, 0, 1, 8, 0, 8); + + if (fullReps > 0) { + Div(sigmoidResult, onesBlock, denomTmp, maskFp32, + static_cast(fullReps), divParams); + PipeBarrier(); + } + if (remainder > 0) { + Div(sigmoidResult[fullReps * maskFp32], onesBlock, + denomTmp[fullReps * maskFp32], remainder, 1, divParams); + PipeBarrier(); + } + + // Step 3: situ_a = tanhResult * sigmoidResult + Mul(tanhResult, tanhResult, sigmoidResult, H); + PipeBarrier(); + + // Step 4: if linear_beta > 0: up = linear_beta * tanh(up / linear_beta) + if (linearBeta_ > 0.0f) { + float invLinearBeta = 1.0f / linearBeta_; + Muls(up, up, invLinearBeta, H); + PipeBarrier(); + Tanh(up, up, H); + PipeBarrier(); + Muls(up, up, linearBeta_, H); + PipeBarrier(); + } + + // Step 5: output = situ_a * up = tanhResult * up + // Write to gate buffer (no longer needed, avoids aliasing with up) + LocalTensor situOut = dequantOut[gateOffset]; + Mul(situOut, tanhResult, up, H); + PipeBarrier(); + + // situOut now holds the Situ output [H] in fp32, stored in gate buffer region + situOut_ = situOut; + } + + __aicore__ inline void ComputeQuant() + { + int64_t H = curColNum_; + LocalTensor situOut = situOut_; + + if (quantMode_ == 1) { + // Dynamic quant + DynamicQuant(situOut); + } else { + // Static quant + StaticQuant(situOut); + } + } + + __aicore__ inline void StaticQuant(LocalTensor& situOut) + { + int64_t H = curColNum_; + + if (quantScaleIsEmpty_ == 0) { + if (quantIsOne_) { + Muls(situOut, situOut, quantScaleVal_, H); + PipeBarrier(); + if (quantOffsetIsEmpty_ == 0) { + Adds(situOut, situOut, quantOffsetVal_, H); + PipeBarrier(); + } + } else { + Div(situOut, situOut, quantLocal_, H); + PipeBarrier(); + if (quantOffsetIsEmpty_ == 0) { + Add(situOut, situOut, quantLocal_[H], H); + PipeBarrier(); + } + } + } + + // Allocate outQueue and cast fp32 -> int8 + LocalTensor outLocal = outQueue_.template AllocTensor(); + LocalTensor yOut = outLocal.template ReinterpretCast(); + CastFloatToInt8(situOut, yOut, H); + outQueue_.EnQue(outLocal); + } + + __aicore__ inline void DynamicQuant(LocalTensor& situOut) + { + int64_t H = curColNum_; + + if (quantScaleIsEmpty_ == 0) { + Mul(situOut, situOut, quantLocal_, H); + PipeBarrier(); + } + + // Compute per-row abs max using situTemp_ (no longer needed after Situ) + LocalTensor absBuf = situTemp_; + Abs(absBuf, situOut, H); + PipeBarrier(); + + // Allocate outQueue for final output: [H int8][1 float scale] + LocalTensor outLocal = outQueue_.template AllocTensor(); + LocalTensor scaleOut = outLocal[H]; + LocalTensor yOut = outLocal.template ReinterpretCast(); + + ComputeReduceMax(absBuf, H); + PipeBarrier(); + + Muls(scaleOut, absBuf, DYNAMIC_QUANT_FACTOR, 1); + PipeBarrier(); + + event_t eventV2S = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_S)); + SetFlag(eventV2S); + WaitFlag(eventV2S); + float scaleVal = scaleOut.GetValue(0); + if (scaleVal == 0.0f) { + scaleVal = 1.0f; + } + float invScale = 1.0f / scaleVal; + Muls(situOut, situOut, invScale, H); + PipeBarrier(); + + CastFloatToInt8(situOut, yOut, H); + outQueue_.EnQue(outLocal); + } + + __aicore__ inline void CastFloatToInt8(const LocalTensor& src, LocalTensor& dst, int64_t count) + { + // FP32 -> INT32 (rint) + LocalTensor tmpI32 = castBuf_.template Get(); + Cast(tmpI32, src, RoundMode::CAST_RINT, count); + PipeBarrier(); + SetDeqScale((half)1.000000e+00f); + + // INT32 -> FP16 (round) + LocalTensor tmpF32 = castBuf_.template Get(); + LocalTensor tmpF16 = tmpF32.ReinterpretCast(); + Cast(tmpF16, tmpI32, RoundMode::CAST_ROUND, count); + PipeBarrier(); + + // FP16 -> INT8 (trunc) + Cast(dst, tmpF16, RoundMode::CAST_TRUNC, count); + PipeBarrier(); + } + + __aicore__ inline void ComputeReduceMax(const LocalTensor& tempRes, int32_t calCount) + { + uint32_t repsFp32 = static_cast(calCount >> 6); + uint32_t offsetsFp32 = repsFp32 << 6; + uint32_t remsFp32 = static_cast(calCount & 0x3f); + + if (likely(repsFp32 > 1)) { + if (repsFp32 - 1 > MAX_REPEAT) { + Max(tempRes, tempRes[ELEM_PER_REP_FP32], tempRes, ELEM_PER_REP_FP32, MAX_REPEAT, + {1, 1, 1, 0, 8, 0}); + PipeBarrier(); + Max(tempRes, tempRes[ELEM_PER_REP_FP32 * MAX_REPEAT], tempRes, ELEM_PER_REP_FP32, + repsFp32 - MAX_REPEAT - 1, {1, 1, 1, 0, 8, 0}); + } else { + Max(tempRes, tempRes[ELEM_PER_REP_FP32], tempRes, ELEM_PER_REP_FP32, repsFp32 - 1, + {1, 1, 1, 0, 8, 0}); + } + PipeBarrier(); + } + if (unlikely(remsFp32 > 0) && unlikely(offsetsFp32 > 0)) { + Max(tempRes, tempRes[offsetsFp32], tempRes, remsFp32, 1, {1, 1, 1, 0, 8, 0}); + PipeBarrier(); + } + uint32_t mask = repsFp32 > 0 ? ELEM_PER_REP_FP32 : calCount; + WholeReduceMax(tempRes, tempRes, mask, 1, 8, 1, 8); + } + + __aicore__ inline void CopyOut(int64_t rowIdx, int64_t colOffset) + { + LocalTensor outLocal = outQueue_.template DeQue(); + LocalTensor yOut = outLocal.template ReinterpretCast(); + + DataCopyExtParams dataCopyYParams{1, static_cast(curColNum_ * sizeof(int8_t)), 0, 0, 0}; + DataCopyPad(yGm_[rowIdx * outDimy_ + colOffset], yOut, dataCopyYParams); + + if (quantMode_ == 1) { + LocalTensor scaleOut = outLocal[curColNum_]; + DataCopyExtParams dataCopyScaleParams{1, static_cast(sizeof(float)), 0, 0, 0}; + DataCopyPad(scaleGm_[rowIdx], scaleOut, dataCopyScaleParams); + } + + outQueue_.FreeTensor(outLocal); + } + + __aicore__ inline void processCompute() + { + int64_t lastColNum = baseColLen_; + int64_t colLoops = 1; + if (baseColLen_ < colLen_) { + colLoops = (colLen_ + baseColLen_ - 1) / baseColLen_; + lastColNum = colLen_ - (colLoops - 1) * baseColLen_; + } + + if (quantMode_ == 1 && colLoops > 1) { + // Dynamic mode with column tiling: two-pass (recompute) approach + // Pass 1: compute per-row absmax across all column tiles + // Pass 2: re-compute Situ, quantize with global scale, output + for (int64_t i = 0; i < curCoreRowNum_; i++) { + float scaleVal = DynamicComputeRowScale(i, colLoops, lastColNum); + // Sync MTE3→V between Pass 1 and Pass 2 + event_t eventMTE3ToV = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE3_V)); + SetFlag(eventMTE3ToV); + WaitFlag(eventMTE3ToV); + DynamicQuantizeAndOutput(i, scaleVal, colLoops, lastColNum); + } + } else { + // Single-pass approach (no column tiling or static mode) + // Row-major order: process all tiles for each row before moving to next row + for (int64_t i = 0; i < curCoreRowNum_; i++) { + for (int64_t colLoop = 0; colLoop < colLoops; colLoop++) { + curColNum_ = (colLoop == colLoops - 1) ? lastColNum : baseColLen_; + curColNum_ = (curColNum_ == 0) ? baseColLen_ : curColNum_; + + CopyInDequantParams(colLoop * baseColLen_); + if (quantScaleIsEmpty_ == 0 && !quantIsOne_) { + CopyInQuantParams(colLoop * baseColLen_); + } + CopyIn(i, colLoop * baseColLen_); + ComputeDequant(i); + ComputeSitu(); + ComputeQuant(); + CopyOut(i, colLoop * baseColLen_); + } + } + } + } + + __aicore__ inline void ApplySmoothScale(LocalTensor& situOut) + { + if (quantScaleIsEmpty_ == 0) { + if (quantIsOne_) { + Muls(situOut, situOut, quantScaleVal_, curColNum_); + } else { + Mul(situOut, situOut, quantLocal_, curColNum_); + } + PipeBarrier(); + } + } + + __aicore__ inline float DynamicComputeRowScale(int64_t rowIdx, int64_t colLoops, int64_t lastColNum) + { + float rowAbsMax = 0.0f; + for (int64_t colLoop = 0; colLoop < colLoops; colLoop++) { + curColNum_ = (colLoop == colLoops - 1) ? lastColNum : baseColLen_; + curColNum_ = (curColNum_ == 0) ? baseColLen_ : curColNum_; + + CopyInDequantParams(colLoop * baseColLen_); + if (quantScaleIsEmpty_ == 0 && !quantIsOne_) { + CopyInQuantParams(colLoop * baseColLen_); + } + CopyIn(rowIdx, colLoop * baseColLen_); + ComputeDequant(rowIdx); + ComputeSitu(); + ApplySmoothScale(situOut_); + + // Compute absmax for this tile + LocalTensor absBuf = situTemp_; + Abs(absBuf, situOut_, curColNum_); + PipeBarrier(); + ComputeReduceMax(absBuf, curColNum_); + PipeBarrier(); + + event_t eventV2S = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_S)); + SetFlag(eventV2S); + WaitFlag(eventV2S); + float tileMax = absBuf.GetValue(0); + if (tileMax > rowAbsMax) { + rowAbsMax = tileMax; + } + } + + float scaleVal = rowAbsMax * DYNAMIC_QUANT_FACTOR; + if (scaleVal == 0.0f) { + scaleVal = 1.0f; + } + return scaleVal; + } + + __aicore__ inline void DynamicQuantizeAndOutput(int64_t rowIdx, float scaleVal, int64_t colLoops, int64_t lastColNum) + { + float invScale = 1.0f / scaleVal; + + for (int64_t colLoop = 0; colLoop < colLoops; colLoop++) { + curColNum_ = (colLoop == colLoops - 1) ? lastColNum : baseColLen_; + curColNum_ = (curColNum_ == 0) ? baseColLen_ : curColNum_; + + CopyInDequantParams(colLoop * baseColLen_); + if (quantScaleIsEmpty_ == 0 && !quantIsOne_) { + CopyInQuantParams(colLoop * baseColLen_); + } + CopyIn(rowIdx, colLoop * baseColLen_); + ComputeDequant(rowIdx); + ComputeSitu(); + ApplySmoothScale(situOut_); + + // Quantize with global scale + Muls(situOut_, situOut_, invScale, curColNum_); + PipeBarrier(); + + // Cast to int8 and output + LocalTensor outLocal = outQueue_.template AllocTensor(); + LocalTensor yOut = outLocal.template ReinterpretCast(); + CastFloatToInt8(situOut_, yOut, curColNum_); + + // Write scale for first tile only + if (colLoop == 0) { + LocalTensor scaleOut = outLocal[curColNum_]; + Duplicate(scaleOut, scaleVal, 1); + PipeBarrier(); + } + + outQueue_.EnQue(outLocal); + + // CopyOut y (always) and scale (first tile only) + LocalTensor outLocalDeq = outQueue_.template DeQue(); + LocalTensor yOutDeq = outLocalDeq.template ReinterpretCast(); + DataCopyExtParams dataCopyYParams{1, static_cast(curColNum_ * sizeof(int8_t)), 0, 0, 0}; + DataCopyPad(yGm_[rowIdx * outDimy_ + colLoop * baseColLen_], yOutDeq, dataCopyYParams); + + if (colLoop == 0) { + LocalTensor scaleOutDeq = outLocalDeq[curColNum_]; + DataCopyExtParams dataCopyScaleParams{1, static_cast(sizeof(float)), 0, 0, 0}; + DataCopyPad(scaleGm_[rowIdx], scaleOutDeq, dataCopyScaleParams); + } + + outQueue_.FreeTensor(outLocalDeq); + + // Sync MTE3→MTE2 between tiles: ensure copy-out completes before + // next tile's CopyInDequantParams overwrites TBuf buffers + if (colLoop + 1 < colLoops) { + event_t eventMTE3ToMTE2 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE3_MTE2)); + SetFlag(eventMTE3ToMTE2); + WaitFlag(eventMTE3ToMTE2); + } + } + } + + __aicore__ inline int64_t Align(int64_t elementNum, int64_t bytes) + { + constexpr int64_t BLOCK_BYTES = 32; + if (bytes == 0) { + return 0; + } + return (elementNum * bytes + BLOCK_BYTES - 1) / BLOCK_BYTES * BLOCK_BYTES / bytes; + } + +protected: + TPipe* pipe_ = nullptr; + const DequantSituQuantTilingData* tl_ = nullptr; + + int64_t blockIdx_ = 0; + int64_t rowLen_ = 0; + int64_t colLen_ = 0; + int64_t inDimy_ = 0; + int64_t outDimy_ = 0; + int64_t baseRowLen_ = 0; + int64_t baseColLen_ = 0; + int64_t curColNum_ = 0; + int64_t usedCoreNum_ = 0; + int64_t curCoreRowNum_ = 0; + int64_t inputCopyOffset_ = 0; + int64_t activateLeft_ = 0; + int64_t quantMode_ = 0; + bool quantIsOne_ = false; + int64_t quantScaleIsEmpty_ = 1; + int64_t quantOffsetIsEmpty_ = 1; + float beta_ = 1.0f; + float linearBeta_ = 0.0f; + float quantScaleVal_ = 1.0f; + float quantOffsetVal_ = 0.0f; + bool dequantScaleIsOne_ = false; + float dequantScaleVal_ = 1.0f; + bool dequantBiasIsOne_ = false; + float dequantBiasVal_ = 0.0f; + + GlobalTensor xGm_; + GlobalTensor dequantScaleGm_; + GlobalTensor dequantBiasGm_; + GlobalTensor quantScaleGm_; + GlobalTensor quantOffsetGm_; + GlobalTensor yGm_; + GlobalTensor scaleGm_; + + TQue inQueueX_; + TBuf dequantScaleBuf_; + TBuf dequantBiasBuf_; + TBuf quantBuf_; + TQue outQueue_; + TBuf tmpBuf_; + TBuf castBuf_; + + // Intermediate results passed between compute stages + LocalTensor dequantScaleLocal_; + LocalTensor dequantBiasLocal_; + LocalTensor quantLocal_; + LocalTensor dequantOut_; + LocalTensor situTemp_; + LocalTensor situOut_; +}; + +// --------------------------------------------------------------------------- +// K3 Kernel: INT32/BF16 path with MoE routing and per-row dynamic quant +// --------------------------------------------------------------------------- + +constexpr int64_t K3_MASK_FP32 = 256 / sizeof(float); +constexpr int64_t K3_MASK_BLK_STRIDE = 8; +constexpr float K3_DYNAMIC_QUANT_FACTOR = 1.0f / 127.0f; + +template +class DequantSituQuantK3Kernel { +public: + __aicore__ inline explicit DequantSituQuantK3Kernel(TPipe* pipe) : pipe_(pipe) {} + + __aicore__ inline void Init( + GM_ADDR x, GM_ADDR weightScale, GM_ADDR activationScale, GM_ADDR bias, GM_ADDR groupIndex, + GM_ADDR y, GM_ADDR scale, const DequantSituQuantTilingData* tilingData) + { + tilingData_ = tilingData; + blockIdx_ = GetBlockIdx(); + rowLen_ = static_cast(tilingData_->rowLen); + inputWidth_ = static_cast(tilingData_->inputWidth); + outputWidth_ = static_cast(tilingData_->outputWidth); + expertNum_ = static_cast(tilingData_->expertNum); + usedCoreNum_ = static_cast(tilingData_->usedCoreNum); + hasBias_ = tilingData_->dequantBiasIsEmpty == 0; + hasGroupIndex_ = tilingData_->hasGroupIndex != 0; + activateLeft_ = tilingData_->activateLeft; + beta_ = tilingData_->beta; + linearBeta_ = tilingData_->linearBeta; + + xGm_.SetGlobalBuffer((__gm__ XType*)x, rowLen_ * inputWidth_); + if constexpr (std::is_same_v) { + weightScaleGm_.SetGlobalBuffer((__gm__ float*)weightScale, expertNum_ * inputWidth_); + activationScaleGm_.SetGlobalBuffer((__gm__ float*)activationScale, rowLen_); + if (hasBias_) { + biasGm_.SetGlobalBuffer((__gm__ float*)bias, expertNum_ * inputWidth_); + } + if (hasGroupIndex_) { + groupIndexGm_.SetGlobalBuffer((__gm__ int64_t*)groupIndex, expertNum_); + } + } + yGm_.SetGlobalBuffer((__gm__ int8_t*)y, rowLen_ * outputWidth_); + scaleGm_.SetGlobalBuffer((__gm__ float*)scale, rowLen_); + + const int64_t inputBytes = inputWidth_ * static_cast(sizeof(XType)); + const int64_t paramBytes = inputWidth_ * static_cast(sizeof(float)); + const int64_t outputBytes = outputWidth_ * static_cast(sizeof(int8_t)) + 32; + pipe_->InitBuffer(xQueue_, 1, inputBytes); + if constexpr (std::is_same_v) { + pipe_->InitBuffer(weightScaleQueue_, 1, paramBytes); + if (hasBias_) { + pipe_->InitBuffer(biasQueue_, 1, paramBytes); + } + } else { + pipe_->InitBuffer(dequantBuf_, inputWidth_ * static_cast(sizeof(float))); + } + pipe_->InitBuffer(outQueue_, 1, outputBytes); + pipe_->InitBuffer(tmpBuf_, inputWidth_ * static_cast(sizeof(float))); + } + + __aicore__ inline void Process() + { + if (usedCoreNum_ <= 0 || blockIdx_ >= usedCoreNum_) { + return; + } + + if constexpr (!std::is_same_v) { + ProcessGroup(0, rowLen_, 0); + return; + } + + if (!hasGroupIndex_) { + ProcessGroup(0, rowLen_, 0); + return; + } + + int64_t groupOffset = 0; + for (int64_t expertIdx = 0; expertIdx < expertNum_ && groupOffset < rowLen_; ++expertIdx) { + const int64_t requestedRows = groupIndexGm_.GetValue(expertIdx); + const int64_t remainingRows = rowLen_ - groupOffset; + const int64_t groupRows = requestedRows <= 0 ? 0 : + (requestedRows > remainingRows ? remainingRows : requestedRows); + if (groupRows > 0) { + ProcessGroup(expertIdx, groupRows, groupOffset); + groupOffset += groupRows; + } + } + } + +private: + __aicore__ inline void ProcessGroup(int64_t expertIdx, int64_t groupRows, int64_t groupOffset) + { + const int64_t rowsPerCore = (groupRows + usedCoreNum_ - 1) / usedCoreNum_; + const int64_t localGroupOffset = blockIdx_ * rowsPerCore; + if (localGroupOffset >= groupRows) { + return; + } + const int64_t localRows = + groupRows - localGroupOffset < rowsPerCore ? groupRows - localGroupOffset : rowsPerCore; + const int64_t firstRow = groupOffset + localGroupOffset; + + if constexpr (std::is_same_v) { + CopyInExpertParams(expertIdx); + weightScaleLocal_ = weightScaleQueue_.DeQue(); + if (hasBias_) { + biasLocal_ = biasQueue_.DeQue(); + } + } + + for (int64_t localRow = 0; localRow < localRows; ++localRow) { + const int64_t rowIdx = firstRow + localRow; + CopyInRow(rowIdx); + ComputeRow(rowIdx); + CopyOutRow(rowIdx); + } + + if constexpr (std::is_same_v) { + weightScaleQueue_.FreeTensor(weightScaleLocal_); + if (hasBias_) { + biasQueue_.FreeTensor(biasLocal_); + } + } + } + + __aicore__ inline void CopyInExpertParams(int64_t expertIdx) + { + const uint32_t paramBytes = static_cast(inputWidth_ * sizeof(float)); + DataCopyExtParams params{1, paramBytes, 0, 0, 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + const int64_t paramOffset = expertIdx * inputWidth_; + + LocalTensor weightScaleLocal = weightScaleQueue_.AllocTensor(); + DataCopyPad(weightScaleLocal, weightScaleGm_[paramOffset], params, padParams); + weightScaleQueue_.EnQue(weightScaleLocal); + + if (hasBias_) { + LocalTensor biasLocal = biasQueue_.AllocTensor(); + DataCopyPad(biasLocal, biasGm_[paramOffset], params, padParams); + biasQueue_.EnQue(biasLocal); + } + } + + __aicore__ inline void CopyInRow(int64_t rowIdx) + { + const uint32_t inputBytes = static_cast(inputWidth_ * sizeof(XType)); + DataCopyExtParams params{1, inputBytes, 0, 0, 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + LocalTensor xLocal = xQueue_.AllocTensor(); + DataCopyPad(xLocal, xGm_[rowIdx * inputWidth_], params, padParams); + xQueue_.EnQue(xLocal); + } + + __aicore__ inline void ComputeRow(int64_t rowIdx) + { + LocalTensor xLocal = xQueue_.DeQue(); + LocalTensor xLocalF32; + if constexpr (std::is_same_v) { + xLocalF32 = xLocal.template ReinterpretCast(); + } else { + xLocalF32 = dequantBuf_.Get(); + } + Cast(xLocalF32, xLocal, RoundMode::CAST_NONE, inputWidth_); + PipeBarrier(); + if constexpr (std::is_same_v) { + const float activationScale = activationScaleGm_.GetValue(rowIdx); + Mul(xLocalF32, xLocalF32, weightScaleLocal_, inputWidth_); + PipeBarrier(); + Muls(xLocalF32, xLocalF32, activationScale, inputWidth_); + PipeBarrier(); + if (hasBias_) { + Add(xLocalF32, xLocalF32, biasLocal_, inputWidth_); + PipeBarrier(); + } + } + + LocalTensor temp = tmpBuf_.Get(); + int64_t gateOffset = (activateLeft_ == 1) ? 0 : outputWidth_; + int64_t upOffset = (activateLeft_ == 1) ? outputWidth_ : 0; + + LocalTensor gate = xLocalF32[gateOffset]; + LocalTensor up = xLocalF32[upOffset]; + LocalTensor sigmoid = temp; + LocalTensor ones = temp[outputWidth_]; + + Adds(sigmoid, gate, 0.0f, outputWidth_); + PipeBarrier(); + Muls(gate, gate, 1.0f / beta_, outputWidth_); + PipeBarrier(); + Tanh(gate, gate, outputWidth_); + PipeBarrier(); + Muls(gate, gate, beta_, outputWidth_); + PipeBarrier(); + + Muls(sigmoid, sigmoid, -1.0f, outputWidth_); + PipeBarrier(); + Exp(sigmoid, sigmoid, outputWidth_); + PipeBarrier(); + Adds(sigmoid, sigmoid, 1.0f, outputWidth_); + PipeBarrier(); + Duplicate(ones, 1.0f, outputWidth_); + PipeBarrier(); + Div(sigmoid, ones, sigmoid, outputWidth_); + PipeBarrier(); + Mul(gate, gate, sigmoid, outputWidth_); + PipeBarrier(); + + if (linearBeta_ > 0.0f) { + Muls(up, up, 1.0f / linearBeta_, outputWidth_); + PipeBarrier(); + Tanh(up, up, outputWidth_); + PipeBarrier(); + Muls(up, up, linearBeta_, outputWidth_); + PipeBarrier(); + } + Mul(gate, gate, up, outputWidth_); + PipeBarrier(); + + DynamicQuant(gate, temp); + xQueue_.FreeTensor(xLocal); + } + + __aicore__ inline void DynamicQuant(LocalTensor& situ, LocalTensor& temp) + { + Abs(temp, situ, outputWidth_); + PipeBarrier(); + ComputeReduceMax(temp); + PipeBarrier(); + WholeReduceMax(temp, temp, K3_MASK_FP32, 1, K3_MASK_BLK_STRIDE, 1, K3_MASK_BLK_STRIDE, + ReduceOrder::ORDER_ONLY_VALUE); + PipeBarrier(); + + event_t eventVToS = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_S)); + SetFlag(eventVToS); + WaitFlag(eventVToS); + float scaleValue = temp.GetValue(0) * K3_DYNAMIC_QUANT_FACTOR; + if (scaleValue <= 0.0f) { + scaleValue = 1.0f; + } + Muls(situ, situ, 1.0f / scaleValue, outputWidth_); + PipeBarrier(); + + LocalTensor outLocal = outQueue_.AllocTensor(); + LocalTensor yLocal = outLocal.ReinterpretCast(); + // Scale is packed after int8 data, aligned to float boundary + int64_t scaleIdx = (outputWidth_ + static_cast(sizeof(float)) - 1) / sizeof(float); + LocalTensor scaleLocal = outLocal[scaleIdx]; + Duplicate(scaleLocal, scaleValue, 1); + PipeBarrier(); + CastFloatToInt8(situ, temp, yLocal); + outQueue_.EnQue(outLocal); + } + + __aicore__ inline void ComputeReduceMax(const LocalTensor& temp) + { + const uint32_t vectorCycles = static_cast(outputWidth_ / K3_MASK_FP32); + const uint32_t remainder = static_cast(outputWidth_ % K3_MASK_FP32); + + if (vectorCycles > 1) { + BinaryRepeatParams repeatParams; + repeatParams.dstBlkStride = 1; + repeatParams.src0BlkStride = 1; + repeatParams.src1BlkStride = 1; + repeatParams.dstRepStride = 0; + repeatParams.src0RepStride = K3_MASK_BLK_STRIDE; + repeatParams.src1RepStride = 0; + Max(temp, temp[K3_MASK_FP32], temp, K3_MASK_FP32, static_cast(vectorCycles - 1), repeatParams); + PipeBarrier(); + } + if (remainder > 0 && vectorCycles > 0) { + Max(temp, temp[vectorCycles * K3_MASK_FP32], temp, remainder, 1, {1, 1, 1, 0, 8, 0}); + PipeBarrier(); + } + uint32_t mask = vectorCycles > 0 ? K3_MASK_FP32 : outputWidth_; + WholeReduceMax(temp, temp, mask, 1, K3_MASK_BLK_STRIDE, 1, K3_MASK_BLK_STRIDE, + ReduceOrder::ORDER_ONLY_VALUE); + } + + __aicore__ inline void CastFloatToInt8( + const LocalTensor& src, LocalTensor& temp, LocalTensor& dst) + { + LocalTensor tempInt32 = temp[outputWidth_].ReinterpretCast(); + Cast(tempInt32, src, RoundMode::CAST_RINT, outputWidth_); + PipeBarrier(); + SetDeqScale((half)1.0f); + + LocalTensor tempHalf = temp.ReinterpretCast(); + Cast(tempHalf, tempInt32, RoundMode::CAST_ROUND, outputWidth_); + PipeBarrier(); + Cast(dst, tempHalf, RoundMode::CAST_TRUNC, outputWidth_); + PipeBarrier(); + } + + __aicore__ inline void CopyOutRow(int64_t rowIdx) + { + LocalTensor outLocal = outQueue_.DeQue(); + LocalTensor yLocal = outLocal.ReinterpretCast(); + int64_t scaleIdx = (outputWidth_ + static_cast(sizeof(float)) - 1) / sizeof(float); + LocalTensor scaleLocal = outLocal[scaleIdx]; + + DataCopyExtParams yParams{1, static_cast(outputWidth_ * sizeof(int8_t)), 0, 0, 0}; + DataCopyPad(yGm_[rowIdx * outputWidth_], yLocal, yParams); + DataCopyExtParams scaleParams{1, static_cast(sizeof(float)), 0, 0, 0}; + DataCopyPad(scaleGm_[rowIdx], scaleLocal, scaleParams); + outQueue_.FreeTensor(outLocal); + } + + TPipe* pipe_ = nullptr; + const DequantSituQuantTilingData* tilingData_ = nullptr; + int64_t blockIdx_ = 0; + int64_t rowLen_ = 0; + int64_t inputWidth_ = 0; + int64_t outputWidth_ = 0; + int64_t expertNum_ = 0; + int64_t usedCoreNum_ = 0; + uint32_t activateLeft_ = 1; + bool hasBias_ = false; + bool hasGroupIndex_ = false; + float beta_ = 4.0f; + float linearBeta_ = 25.0f; + + GlobalTensor xGm_; + GlobalTensor weightScaleGm_; + GlobalTensor activationScaleGm_; + GlobalTensor biasGm_; + GlobalTensor groupIndexGm_; + GlobalTensor yGm_; + GlobalTensor scaleGm_; + + TQue xQueue_; + TQue weightScaleQueue_; + TQue biasQueue_; + TQue outQueue_; + TBuf tmpBuf_; + TBuf dequantBuf_; + LocalTensor weightScaleLocal_; + LocalTensor biasLocal_; +}; + +} // namespace DequantSituQuantOps +#endif // DEQUANT_SITU_QUANT_H diff --git a/csrc/moe/situ_mx_quant/CMakeLists.txt b/csrc/moe/situ_mx_quant/CMakeLists.txt new file mode 100644 index 000000000000..a67c7eb0977f --- /dev/null +++ b/csrc/moe/situ_mx_quant/CMakeLists.txt @@ -0,0 +1,19 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ---------------------------------------------------------------------------- + +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +if(NOT ENABLE_TEST AND NOT BENCHMARK) + list(REMOVE_ITEM CURRENT_DIRS tests) +endif() +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/csrc/moe/situ_mx_quant/docs/aclnnSituMxQuant.md b/csrc/moe/situ_mx_quant/docs/aclnnSituMxQuant.md new file mode 100644 index 000000000000..83c2c9c29f57 --- /dev/null +++ b/csrc/moe/situ_mx_quant/docs/aclnnSituMxQuant.md @@ -0,0 +1,87 @@ +# aclnnSituMxQuant + +## 功能说明 + +SituMxQuant 算子执行 Situ 激活,随后进行动态 MX (Microscaling) 量化。 + +计算公式: + +```text +situ_a = beta * tanh(gate / beta) * sigmoid(gate) +situOut = situ_a * up (+ optional linear_beta * tanh(up / linear_beta) on up) +shared_exp = floor(log2(max(|situOut_i|))) - emax +mxscale = 2^shared_exp (E8M0) +y = cast_to_fp8(situOut / mxscale) +``` + +## 接口定义 + +```cpp +aclnnStatus aclnnSituMxQuant( + void* workspace, + uint64_t workspaceSize, + aclOpExecutor* executor, + aclrtStream stream) +``` + +## 算子原型 + +```cpp +REG_OP(SituMxQuant) + .INPUT(x, TensorType({DT_BF16})) + .OUTPUT(y, TensorType({DT_FLOAT8_E4M3FN, DT_FLOAT8_E5M2})) + .OUTPUT(mxscale, TensorType({DT_FLOAT8_E8M0})) + .ATTR(beta, Float, 1.0f) + .ATTR(linear_beta, Float, 0.0f) + .ATTR(activate_left, Bool, false) + .ATTR(axis, Int, -1) + .ATTR(dst_type, Int, 36) + .OP_END_FACTORY_REG(SituMxQuant) +``` + +## 参数说明 + +| 参数 | 输入/输出 | 类型 | 说明 | +|------|-----------|------|------| +| x | 输入 | bfloat16 | 输入张量,shape 为 [N..., 2H],最后一维必须为偶数 | +| y | 输出 | float8_e4m3fn / float8_e5m2 | 量化输出,shape 为 [N..., H] | +| mxscale | 输出 | float8_e8m0 | MX scale,shape 为 [N..., ceil(H/64), 2] | + +## 属性说明 + +| 属性 | 类型 | 默认值 | 说明 | +|------|------|--------|------| +| beta | Float | 1.0 | Situ 激活的 beta 参数,必须 > 0 | +| linear_beta | Float | 0.0 | Situ 激活的 linear_beta 参数,≤0 时不启用 up 的 tanh 变换 | +| activate_left | Bool | false | true: gate 在前半部分;false: gate 在后半部分 | +| axis | Int | -1 | 量化轴,当前仅支持 -1 | +| dst_type | Int | 36 | 输出数据类型: 36=FP8_E4M3FN, 35=FP8_E5M2 | + +## 约束条件 + +- 输入 x 的最后一维必须能被 2 整除 +- 输入 x 支持 1-7 维张量 +- axis 必须为 -1(尾轴量化) +- beta 必须 > 0 +- dst_type 必须为 36 (FP8_E4M3FN) 或 35 (FP8_E5M2) +- 仅支持 Ascend950 平台 + +## mxscale Shape 计算 + +- `H = x.shape[-1] / 2` +- `scaleNum = ceil(H / 64)` (偶对齐的 32 元素 block 数) +- `mxscale.shape = x.shape[:-1] + [scaleNum, 2]` + +## 调用示例 + +```cpp +// 1. 创建 op executor +aclOpExecutor* executor; +aclnnSituMxQuantGetWorkspaceSize(x, y, mxscale, beta, linearBeta, activateLeft, axis, dstType, &workspaceSize, &executor); +void* workspace = nullptr; +if (workspaceSize > 0) { + aclrtMalloc(&workspace, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); +} +// 2. 执行算子 +aclnnSituMxQuant(workspace, workspaceSize, executor, stream); +``` diff --git a/csrc/moe/situ_mx_quant/op_host/CMakeLists.txt b/csrc/moe/situ_mx_quant/op_host/CMakeLists.txt new file mode 100644 index 000000000000..85f66563a4ca --- /dev/null +++ b/csrc/moe/situ_mx_quant/op_host/CMakeLists.txt @@ -0,0 +1,32 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2026 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ---------------------------------------------------------------------------- + +if(NOT "ascend950" IN_LIST ASCEND_COMPUTE_UNIT) + return() +endif() + +add_op_to_compiled_list() + +if(BUILD_OPEN_PROJECT) + target_sources(op_host_aclnn PRIVATE + situ_mx_quant_def.cpp + ) +endif() + +add_ops_compile_options( + OP_NAME SituMxQuant + OPTIONS --cce-auto-sync=on + -Wno-deprecated-declarations + -Werror +) + +if(NOT BUILD_OPS_RTY_KERNEL) + add_modules_sources_with_soc(OPTYPE situ_mx_quant ACLNNTYPE aclnn) +endif() diff --git a/csrc/moe/situ_mx_quant/op_host/arch35/situ_mx_quant_tiling_arch35.cpp b/csrc/moe/situ_mx_quant/op_host/arch35/situ_mx_quant_tiling_arch35.cpp new file mode 100644 index 000000000000..f4d3399b6f06 --- /dev/null +++ b/csrc/moe/situ_mx_quant/op_host/arch35/situ_mx_quant_tiling_arch35.cpp @@ -0,0 +1,339 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file situ_mx_quant_tiling_arch35.cpp + * \brief Tiling implementation for Situ + MX quantization + */ + +#include "situ_mx_quant_tiling_arch35.h" +#include "../../op_kernel/arch35/situ_mx_quant_tiling_data.h" +#include "../../op_kernel/arch35/situ_mx_quant_tiling_key.h" + +#include +#include +#include "log/log.h" +#include "platform/platform_info.h" +#include "util/math_util.h" + +using namespace std; +using namespace ge; +using namespace AscendC; +using namespace SituMxQuantOp; + +namespace optiling { +// ==================== Constants ==================== +constexpr int64_t INDEX_ATTR_BETA = 0; +constexpr int64_t INDEX_ATTR_LINEAR_BETA = 1; +constexpr int64_t INDEX_ATTR_ACTIVATE_LEFT = 2; +constexpr int64_t INDEX_ATTR_AXIS = 3; +constexpr int64_t INDEX_ATTR_DST_TYPE = 4; + +constexpr int64_t BYTES_OF_BF16 = 2; +constexpr int64_t BYTES_OF_FP8 = 1; +constexpr int64_t BYTES_OF_INT16 = 2; +constexpr int64_t RESERVED_UB_SIZE = 32; +constexpr int64_t RESERVED_UB_FOR_ALIGN = 128; +constexpr int64_t BLOCK_SIZE = 32; +constexpr int64_t DOUBLE_BUFFER = 2; +constexpr int64_t CONST_TWO = 2; +constexpr int64_t DTYPE_35 = 35; // FP8_E5M2 +constexpr int64_t DTYPE_36 = 36; // FP8_E4M3FN +constexpr int64_t BASE_DIM1 = 256; // basic block size for last axis + +const set INPUT_SUPPORT_DTYPE_SET = {ge::DT_BF16}; +const set Y_SUPPORT_DTYPE_SET = {ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2}; +const set SCALE_SUPPORT_DTYPE_SET = {ge::DT_FLOAT8_E8M0}; + +// ==================== Helper Functions ==================== +template +static string Shape2String(const T& shape) +{ + ostringstream oss; + oss << "["; + if (shape.GetDimNum() > 0) { + for (size_t i = 0; i < shape.GetDimNum() - 1; ++i) { + oss << shape.GetDim(i) << ", "; + } + oss << shape.GetDim(shape.GetDimNum() - 1); + } + oss << "]"; + return oss.str(); +} + +// ==================== Class Methods ==================== +ge::graphStatus SituMxQuantRegbaseTiling::GetNpuInfo() +{ + OP_LOGD(context_->GetNodeName(), "GetNpuInfo begin."); + auto platformInfo = context_->GetPlatformInfo(); + OP_CHECK_NULL_WITH_CONTEXT(context_, platformInfo); + auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo); + compileInfo_.totalCoreNum = ascendcPlatform.GetCoreNumAiv(); + OP_CHECK_IF((compileInfo_.totalCoreNum <= 0), OP_LOGE(context_->GetNodeName(), "Failed to get core num."), + return ge::GRAPH_FAILED); + uint64_t ubSize = 0; + ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize); + compileInfo_.ubSize = static_cast(ubSize); + OP_CHECK_IF((compileInfo_.ubSize <= 0), OP_LOGE(context_->GetNodeName(), "Failed to get UB size."), + return ge::GRAPH_FAILED); + OP_LOGI(context_->GetNodeName(), "CompileInfo: totalCoreNum=%ld, ubSize=%ld", compileInfo_.totalCoreNum, + compileInfo_.ubSize); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SituMxQuantRegbaseTiling::ParseAttrs() +{ + auto* attrs = context_->GetAttrs(); + OP_CHECK_NULL_WITH_CONTEXT(context_, attrs); + + auto* attrBeta = attrs->GetAttrPointer(INDEX_ATTR_BETA); + attrParam_.beta = (attrBeta != nullptr) ? *attrBeta : 1.0f; + OP_CHECK_IF((attrParam_.beta <= 0.0f), + OP_LOGE(context_->GetNodeName(), "beta must be greater than 0, but got %f.", attrParam_.beta), + return ge::GRAPH_FAILED); + + auto* attrLinearBeta = attrs->GetAttrPointer(INDEX_ATTR_LINEAR_BETA); + attrParam_.linearBeta = (attrLinearBeta != nullptr) ? *attrLinearBeta : 0.0f; + attrParam_.hasLinearBeta = (attrParam_.linearBeta > 0.0f); + + auto* attrActivateLeft = attrs->GetAttrPointer(INDEX_ATTR_ACTIVATE_LEFT); + attrParam_.activateLeft = (attrActivateLeft != nullptr) ? *attrActivateLeft : false; + + auto* attrAxis = attrs->GetAttrPointer(INDEX_ATTR_AXIS); + attrParam_.axis = (attrAxis != nullptr) ? static_cast(*attrAxis) : -1; + OP_CHECK_IF((attrParam_.axis != -1), + OP_LOGE(context_->GetNodeName(), "Only axis=-1 is supported currently, but got %ld.", attrParam_.axis), + return ge::GRAPH_FAILED); + + auto* attrDstType = attrs->GetAttrPointer(INDEX_ATTR_DST_TYPE); + attrParam_.dstType = (attrDstType != nullptr) ? static_cast(*attrDstType) : DTYPE_36; + OP_CHECK_IF((attrParam_.dstType != DTYPE_35) && (attrParam_.dstType != DTYPE_36), + OP_LOGE(context_->GetNodeName(), "Invalid dstType: %ld, only 35(E5M2) or 36(E4M3FN) supported.", + attrParam_.dstType), + return ge::GRAPH_FAILED); + + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SituMxQuantRegbaseTiling::ValidateInput() +{ + auto xDtype = context_->GetInputDesc(0)->GetDataType(); + OP_CHECK_IF((INPUT_SUPPORT_DTYPE_SET.find(xDtype) == INPUT_SUPPORT_DTYPE_SET.end()), + OP_LOGE(context_->GetNodeName(), "Input x dtype %d is not supported. Only BF16 is supported.", + static_cast(xDtype)), + return ge::GRAPH_FAILED); + inputInfo_.xDtype = xDtype; + + auto xShape = context_->GetInputShape(0); + OP_CHECK_NULL_WITH_CONTEXT(context_, xShape); + int64_t dimNum = static_cast(xShape->GetStorageShape().GetDimNum()); + OP_LOGI(context_->GetNodeName(), "Input x shape = %s", Shape2String(xShape->GetStorageShape()).c_str()); + int64_t xSize = xShape->GetStorageShape().GetShapeSize(); + OP_CHECK_IF((dimNum < 1 || xSize == 0), + OP_LOGE(context_->GetNodeName(), "rank of x must >= 1, but is %ld, and not support empty tensor", + dimNum), + return ge::GRAPH_FAILED); + inputInfo_.dimNum = dimNum; + inputInfo_.inputDim2 = xShape->GetStorageShape().GetDim(dimNum - 1); + OP_CHECK_IF((inputInfo_.inputDim2 % CONST_TWO != 0), + OP_LOGE(context_->GetNodeName(), "Last dimension must be divisible by 2, but got %ld.", + inputInfo_.inputDim2), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SituMxQuantRegbaseTiling::ValidateOutput() +{ + auto yDtype = context_->GetOutputDesc(0)->GetDataType(); + OP_CHECK_IF((Y_SUPPORT_DTYPE_SET.find(yDtype) == Y_SUPPORT_DTYPE_SET.end()), + OP_LOGE(context_->GetNodeName(), "Output y dtype %d is not supported.", static_cast(yDtype)), + return ge::GRAPH_FAILED); + outputInfo_.yDtype = yDtype; + OP_CHECK_IF((static_cast(outputInfo_.yDtype) != attrParam_.dstType), + OP_LOGE(context_->GetNodeName(), + "attr dst_type(%ld) does not match output y dtype(%ld)", + attrParam_.dstType, static_cast(outputInfo_.yDtype)), + return ge::GRAPH_FAILED); + + auto mxscaleDtype = context_->GetOutputDesc(1)->GetDataType(); + OP_CHECK_IF((SCALE_SUPPORT_DTYPE_SET.find(mxscaleDtype) == SCALE_SUPPORT_DTYPE_SET.end()), + OP_LOGE(context_->GetNodeName(), "Output mxscale dtype %d is not supported.", + static_cast(mxscaleDtype)), + return ge::GRAPH_FAILED); + outputInfo_.mxscaleDtype = mxscaleDtype; + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SituMxQuantRegbaseTiling::PreProcess() +{ + const gert::StorageShape* xShape = context_->GetInputShape(0); + int64_t dimNum = inputInfo_.dimNum; + const gert::StorageShape* yShape = context_->GetOutputShape(0); + OP_CHECK_NULL_WITH_CONTEXT(context_, yShape); + int64_t yDimNum = static_cast(yShape->GetStorageShape().GetDimNum()); + OP_CHECK_IF((yDimNum != dimNum), + OP_LOGE(context_->GetNodeName(), "rank of yShape(%ld) must equal rank of xShape(%ld)", yDimNum, dimNum), + return ge::GRAPH_FAILED); + + outputInfo_.outputDim2 = yShape->GetStorageShape().GetDim(yDimNum - 1); + // Collapse leading dims into inputDim1 (batch * rows) + int64_t inDim1 = 1; + for (int64_t i = 0; i < dimNum - 1; i++) { + inDim1 *= xShape->GetStorageShape().GetDim(i); + } + inputInfo_.inputDim1 = inDim1; + outputInfo_.outputDim1 = inDim1; // same as input since only last dim is halved + + OP_LOGI(context_->GetNodeName(), "3D view: dim1=%ld, dim2=%ld(2H), outputDim2=%ld(H)", + inputInfo_.inputDim1, inputInfo_.inputDim2, outputInfo_.outputDim2); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SituMxQuantRegbaseTiling::CalculateTiling() +{ + tilingResult_.basicDim1 = 1; + tilingResult_.basicDim2 = BASE_DIM1; // 256 + tilingResult_.dimNBlockNum = Ops::Base::CeilDiv(outputInfo_.outputDim2, tilingResult_.basicDim2); + + // UB capacity calculation + int64_t availableUB = compileInfo_.ubSize - RESERVED_UB_SIZE - RESERVED_UB_FOR_ALIGN; + int64_t bytesPerIteration = 0; + // Input x: 2 halves (gate + up), each basicDim1 * basicDim2 * sizeof(BF16) + bytesPerIteration += tilingResult_.basicDim1 * tilingResult_.basicDim2 * CONST_TWO * BYTES_OF_BF16; + // Output y: FP8, basicDim1 * basicDim2 * sizeof(uint8_t) + bytesPerIteration += tilingResult_.basicDim1 * tilingResult_.basicDim2 * BYTES_OF_FP8; + // Output mxscale: E8M0 + int64_t scaleCount = tilingResult_.basicDim1 * tilingResult_.basicDim2 / BLOCK_SIZE; + bytesPerIteration += scaleCount * BYTES_OF_FP8; + // Double buffer + bytesPerIteration *= DOUBLE_BUFFER; + // Situ output buffer (BF16) + bytesPerIteration += tilingResult_.basicDim1 * tilingResult_.basicDim2 * BYTES_OF_BF16; + // maxExp + halfScale (uint16_t each) + bytesPerIteration += scaleCount * BYTES_OF_INT16 * CONST_TWO; + + int64_t ubTotalBasicBlock = availableUB / bytesPerIteration; + OP_LOGI(context_->GetNodeName(), "ubTotalBasicBlock is %ld", ubTotalBasicBlock); + + if (ubTotalBasicBlock >= tilingResult_.dimNBlockNum) { + tilingResult_.maxBasicNumUbDim2 = tilingResult_.dimNBlockNum; + tilingResult_.maxBasicNumUbDim1 = Ops::Base::FloorDiv(ubTotalBasicBlock, tilingResult_.dimNBlockNum); + } else { + tilingResult_.maxBasicNumUbDim2 = (ubTotalBasicBlock > 0) ? ubTotalBasicBlock : 1; + tilingResult_.maxBasicNumUbDim1 = 1; + } + + // Core distribution: M × N grid + int64_t dimM = outputInfo_.outputDim1; + int64_t bmCores = std::min(dimM, compileInfo_.totalCoreNum); + int64_t nCores = 1; + if (bmCores < compileInfo_.totalCoreNum) { + nCores = compileInfo_.totalCoreNum / bmCores; + nCores = std::min(nCores, tilingResult_.dimNBlockNum); + } + tilingResult_.mCorePerB = bmCores; + tilingResult_.nCoreNum = nCores; + tilingResult_.usedCoreNum = bmCores * nCores; + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SituMxQuantRegbaseTiling::FillTilingData() +{ + tilingData_ = context_->GetTilingData(); + OP_CHECK_IF(tilingData_ == nullptr, OP_LOGE(context_->GetNodeName(), "get tilingdata ptr failed"), + return ge::GRAPH_FAILED); + OP_CHECK_IF((memset_s(tilingData_, sizeof(SituMxQuantTilingData), 0, sizeof(SituMxQuantTilingData)) != EOK), + OP_LOGE(context_->GetNodeName(), "memset tilingData failed"), return ge::GRAPH_FAILED); + + tilingData_->usedCoreNum = tilingResult_.usedCoreNum; + tilingData_->inputDim0 = 1; + tilingData_->inputDim1 = outputInfo_.outputDim1; + tilingData_->inputDim2 = outputInfo_.outputDim2; + tilingData_->dimNBlockNum = tilingResult_.dimNBlockNum; + tilingData_->maxBasicNumUbDim2 = tilingResult_.maxBasicNumUbDim2; + tilingData_->maxBasicNumUbDim1 = tilingResult_.maxBasicNumUbDim1; + tilingData_->nCoreNum = tilingResult_.nCoreNum; + tilingData_->mCorePerB = tilingResult_.mCorePerB; + tilingData_->frontCoreNum = 0; + tilingData_->tailCoreBasicNumDim1 = 0; + tilingData_->activateLeft = attrParam_.activateLeft ? 1 : 0; + tilingData_->beta = attrParam_.beta; + tilingData_->linearBeta = attrParam_.linearBeta; + tilingData_->hasLinearBeta = attrParam_.hasLinearBeta ? 1 : 0; + return ge::GRAPH_SUCCESS; +} + +void SituMxQuantRegbaseTiling::SetTilingKeyAndCore() +{ + hasLinearBeta_ = attrParam_.hasLinearBeta ? TPL_HAS_LINEAR_BETA : TPL_NO_LINEAR_BETA; + dstTypeIndex_ = (attrParam_.dstType == DTYPE_36) ? TPL_DST_E4M3FN : TPL_DST_E5M2; + + int64_t tilingKey = GET_TPL_TILING_KEY(hasLinearBeta_, dstTypeIndex_); + OP_LOGI(context_->GetNodeName(), "hasLinearBeta=%lu, dstTypeIndex=%lu, tilingKey=%ld", + hasLinearBeta_, dstTypeIndex_, tilingKey); + context_->SetTilingKey(tilingKey); + context_->SetBlockDim(tilingData_->usedCoreNum); +} + +void SituMxQuantRegbaseTiling::PrintTilingData() const +{ + OP_LOGI(context_->GetNodeName(), + "TilingData: usedCoreNum=%ld, inputDim1=%ld, inputDim2=%ld, dimNBlockNum=%ld, " + "maxBasicNumUbDim2=%ld, maxBasicNumUbDim1=%ld, nCoreNum=%ld, mCorePerB=%ld, " + "beta=%f, linearBeta=%f, hasLinearBeta=%ld", + tilingData_->usedCoreNum, tilingData_->inputDim1, tilingData_->inputDim2, + tilingData_->dimNBlockNum, tilingData_->maxBasicNumUbDim2, tilingData_->maxBasicNumUbDim1, + tilingData_->nCoreNum, tilingData_->mCorePerB, tilingData_->beta, tilingData_->linearBeta, + tilingData_->hasLinearBeta); +} + +// ==================== Entry Functions ==================== +ge::graphStatus Tiling4SituMxQuant(gert::TilingContext* context) +{ + OP_LOGD(context->GetNodeName(), "Begin to do Tiling4SituMxQuant"); + SituMxQuantRegbaseTiling tiling(context); + + OP_CHECK_IF(tiling.GetNpuInfo() != ge::GRAPH_SUCCESS, + OP_LOGE(context->GetNodeName(), "GetNpuInfo failed"), return ge::GRAPH_FAILED); + OP_CHECK_IF(tiling.ParseAttrs() != ge::GRAPH_SUCCESS, + OP_LOGE(context->GetNodeName(), "ParseAttrs failed"), return ge::GRAPH_FAILED); + OP_CHECK_IF(tiling.ValidateInput() != ge::GRAPH_SUCCESS, + OP_LOGE(context->GetNodeName(), "ValidateInput failed"), return ge::GRAPH_FAILED); + OP_CHECK_IF(tiling.ValidateOutput() != ge::GRAPH_SUCCESS, + OP_LOGE(context->GetNodeName(), "ValidateOutput failed"), return ge::GRAPH_FAILED); + OP_CHECK_IF(tiling.PreProcess() != ge::GRAPH_SUCCESS, + OP_LOGE(context->GetNodeName(), "PreProcess failed"), return ge::GRAPH_FAILED); + OP_CHECK_IF(tiling.CalculateTiling() != ge::GRAPH_SUCCESS, + OP_LOGE(context->GetNodeName(), "CalculateTiling failed"), return ge::GRAPH_FAILED); + OP_CHECK_IF(tiling.FillTilingData() != ge::GRAPH_SUCCESS, + OP_LOGE(context->GetNodeName(), "FillTilingData failed"), return ge::GRAPH_FAILED); + tiling.SetTilingKeyAndCore(); + tiling.PrintTilingData(); + + // Set workspace size + size_t workspaceSize = 0; + size_t* currentWorkspace = context->GetWorkspaceSizes(1); + *currentWorkspace = workspaceSize; + + OP_LOGI(context->GetNodeName(), "End to do Tiling4SituMxQuant"); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus TilingPrepare4SituMxQuant(gert::TilingParseContext* context) +{ + return ge::GRAPH_SUCCESS; +} + +// ==================== Registration ==================== +IMPL_OP_OPTILING(SituMxQuant) + .Tiling(Tiling4SituMxQuant) + .TilingParse(TilingPrepare4SituMxQuant); + +} // namespace optiling diff --git a/csrc/moe/situ_mx_quant/op_host/arch35/situ_mx_quant_tiling_arch35.h b/csrc/moe/situ_mx_quant/op_host/arch35/situ_mx_quant_tiling_arch35.h new file mode 100644 index 000000000000..026e96bab470 --- /dev/null +++ b/csrc/moe/situ_mx_quant/op_host/arch35/situ_mx_quant_tiling_arch35.h @@ -0,0 +1,110 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file situ_mx_quant_tiling_arch35.h + * \brief Tiling data structure and parameters for Situ + MX quantization + */ + +#ifndef QUANT_SITU_MX_QUANT_TILING_ARCH35_H +#define QUANT_SITU_MX_QUANT_TILING_ARCH35_H + +#include +#include +#include +#include +#include "register/op_def_registry.h" +#include "register/tilingdata_base.h" +#include "tiling/tiling_api.h" +#include "util/math_util.h" +#include "../../op_kernel/arch35/situ_mx_quant_tiling_data.h" + +namespace optiling { + +// ==================== CompileInfo ==================== +struct SituMxQuantCompileInfo { + int64_t totalCoreNum{0}; + int64_t ubSize{0}; +}; + +// ==================== Input Info ==================== +struct SituMxQuantInputInfo { + ge::DataType xDtype{ge::DT_UNDEFINED}; + int64_t dimNum{0}; + int64_t inputDim0{1}; // batch dim (collapsed) + int64_t inputDim1{1}; // row dim (collapsed) + int64_t inputDim2{0}; // 2H (input last dim) +}; + +// ==================== Output Info ==================== +struct SituMxQuantOutputInfo { + ge::DataType yDtype{ge::DT_UNDEFINED}; + ge::DataType mxscaleDtype{ge::DT_UNDEFINED}; + int64_t outputDim2{0}; // H (output last dim) + int64_t outputDim1{1}; // row dim (collapsed) +}; + +// ==================== Attr Params ==================== +struct SituMxQuantAttrParam { + float beta{1.0f}; + float linearBeta{0.0f}; + bool activateLeft{false}; + int64_t axis{-1}; + int64_t dstType{36}; + bool hasLinearBeta{false}; +}; + +// ==================== Tiling Result ==================== +struct SituMxQuantTilingResult { + int64_t basicDim2{256}; + int64_t basicDim1{1}; + int64_t dimNBlockNum{0}; + int64_t maxBasicNumUbDim2{0}; + int64_t maxBasicNumUbDim1{0}; + int64_t usedCoreNum{0}; + int64_t nCoreNum{1}; + int64_t mCorePerB{1}; +}; + +// ==================== Tiling Class ==================== +class SituMxQuantRegbaseTiling { +public: + explicit SituMxQuantRegbaseTiling(gert::TilingContext* context) : context_(context){}; + + ge::graphStatus GetNpuInfo(); + ge::graphStatus ParseAttrs(); + ge::graphStatus ValidateInput(); + ge::graphStatus ValidateOutput(); + ge::graphStatus PreProcess(); + ge::graphStatus CalculateTiling(); + ge::graphStatus FillTilingData(); + void SetTilingKeyAndCore(); + void PrintTilingData() const; + +private: + SituMxQuantCompileInfo compileInfo_; + SituMxQuantInputInfo inputInfo_; + SituMxQuantOutputInfo outputInfo_; + SituMxQuantAttrParam attrParam_; + SituMxQuantTilingResult tilingResult_; + + uint64_t hasLinearBeta_ = 0; + uint64_t dstTypeIndex_ = 0; + + gert::TilingContext* context_ = nullptr; + SituMxQuantTilingData* tilingData_ = nullptr; +}; + +// ==================== Main function declarations ==================== +ge::graphStatus Tiling4SituMxQuant(gert::TilingContext* context); +ge::graphStatus TilingPrepare4SituMxQuant(gert::TilingParseContext* context); + +} // namespace optiling +#endif // QUANT_SITU_MX_QUANT_TILING_ARCH35_H diff --git a/csrc/moe/situ_mx_quant/op_host/config/ascend950/situ_mx_quant_binary.json b/csrc/moe/situ_mx_quant/op_host/config/ascend950/situ_mx_quant_binary.json new file mode 100644 index 000000000000..c69bbaf455e6 --- /dev/null +++ b/csrc/moe/situ_mx_quant/op_host/config/ascend950/situ_mx_quant_binary.json @@ -0,0 +1,87 @@ +{ + "op_type": "SituMxQuant", + "op_list": [ + { + "bin_filename": "SituMxQuant_bf16_e4m3fn", + "inputs": [ + { + "name": "x", + "index": 0, + "dtype": "bfloat16", + "format": "ND", + "paramType": "required", + "shape": [-2], + "format_match_mode": "FormatAgnostic" + } + ], + "outputs": [ + { + "name": "y", + "index": 0, + "dtype": "float8_e4m3fn", + "format": "ND", + "paramType": "required", + "shape": [-2], + "format_match_mode": "FormatAgnostic" + }, + { + "name": "mxscale", + "index": 1, + "dtype": "float8_e8m0", + "format": "ND", + "paramType": "required", + "shape": [-2], + "format_match_mode": "FormatAgnostic" + } + ], + "attrs": [ + {"name": "beta", "dtype": "float", "value": null}, + {"name": "linear_beta", "dtype": "float", "value": null}, + {"name": "activate_left", "dtype": "bool", "value": null}, + {"name": "axis", "dtype": "int", "value": null}, + {"name": "dst_type", "dtype": "int", "value": null} + ] + }, + { + "bin_filename": "SituMxQuant_bf16_e5m2", + "inputs": [ + { + "name": "x", + "index": 0, + "dtype": "bfloat16", + "format": "ND", + "paramType": "required", + "shape": [-2], + "format_match_mode": "FormatAgnostic" + } + ], + "outputs": [ + { + "name": "y", + "index": 0, + "dtype": "float8_e5m2", + "format": "ND", + "paramType": "required", + "shape": [-2], + "format_match_mode": "FormatAgnostic" + }, + { + "name": "mxscale", + "index": 1, + "dtype": "float8_e8m0", + "format": "ND", + "paramType": "required", + "shape": [-2], + "format_match_mode": "FormatAgnostic" + } + ], + "attrs": [ + {"name": "beta", "dtype": "float", "value": null}, + {"name": "linear_beta", "dtype": "float", "value": null}, + {"name": "activate_left", "dtype": "bool", "value": null}, + {"name": "axis", "dtype": "int", "value": null}, + {"name": "dst_type", "dtype": "int", "value": null} + ] + } + ] +} diff --git a/csrc/moe/situ_mx_quant/op_host/situ_mx_quant_def.cpp b/csrc/moe/situ_mx_quant/op_host/situ_mx_quant_def.cpp new file mode 100644 index 000000000000..7aecf520b52c --- /dev/null +++ b/csrc/moe/situ_mx_quant/op_host/situ_mx_quant_def.cpp @@ -0,0 +1,64 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file situ_mx_quant_def.cpp + * \brief Situ activation combined with dynamic MX quantization operator definition + */ + +#include + +namespace ops { + +class SituMxQuant : public OpDef { +public: + explicit SituMxQuant(const char* name) : OpDef(name) + { + // dtype 组合(2种): + // x=BF16 → y=FP8_E4M3FN, mxscale=E8M0 + // x=BF16 → y=FP8_E5M2, mxscale=E8M0 + this->Input("x") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND}) + .AutoContiguous(); + + this->Output("y") + .ParamType(REQUIRED) + .DataType({ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND}); + + this->Output("mxscale") + .ParamType(REQUIRED) + .DataType({ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND}); + + this->Attr("beta").AttrType(OPTIONAL).Float(1.0f); + this->Attr("linear_beta").AttrType(OPTIONAL).Float(0.0f); + this->Attr("activate_left").AttrType(OPTIONAL).Bool(false); + this->Attr("axis").AttrType(OPTIONAL).Int(-1); + this->Attr("dst_type").AttrType(OPTIONAL).Int(36); + + // Ascend 950 (arch35) configuration using Regbase + OpAICoreConfig regbaseCfg; + regbaseCfg.DynamicCompileStaticFlag(true) + .DynamicRankSupportFlag(true) + .DynamicShapeSupportFlag(true) + .ExtendCfgInfo("opFile.value", "situ_mx_quant_apt"); + this->AICore().AddConfig("ascend950", regbaseCfg); + } +}; + +OP_ADD(SituMxQuant); + +} // namespace ops diff --git a/csrc/moe/situ_mx_quant/op_host/situ_mx_quant_infershape.cpp b/csrc/moe/situ_mx_quant/op_host/situ_mx_quant_infershape.cpp new file mode 100644 index 000000000000..5f8c0c89b3cf --- /dev/null +++ b/csrc/moe/situ_mx_quant/op_host/situ_mx_quant_infershape.cpp @@ -0,0 +1,136 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file situ_mx_quant_infershape.cpp + * \brief Shape inference for Situ + MX quantization operator + */ + +#include "graph/utils/type_utils.h" +#include "runtime/infer_shape_context.h" +#include "register/op_impl_registry.h" +#include "log/log.h" +#include "util/shape_util.h" +#include "util/math_util.h" + +using namespace ge; + +namespace ops { +constexpr int64_t UNKNOWN_DIM_VALUE_ = -1; +constexpr int64_t UNKNOWN_RANK_DIM = -2; +constexpr size_t INDEX_INPUT_X = 0; +constexpr size_t INDEX_OUTPUT_Y = 0; +constexpr size_t INDEX_OUTPUT_MXSCALE = 1; + +constexpr size_t INDEX_ATTR_BETA = 0; +constexpr size_t INDEX_ATTR_LINEAR_BETA = 1; +constexpr size_t INDEX_ATTR_ACTIVATE_LEFT = 2; +constexpr size_t INDEX_ATTR_AXIS = 3; +constexpr size_t INDEX_ATTR_DST_TYPE = 4; + +constexpr int64_t SPLIT_NUM = 2; +constexpr int64_t BLOCK_SIZE = 32; +constexpr int64_t ALIGN_NUM = 2; +constexpr size_t MAX_DIM_NUM = 7; + +static const std::initializer_list Y_SUPPORT_DTYPE_SET = {ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2}; + +graphStatus InferShapeForSituMxQuant(gert::InferShapeContext* context) +{ + OP_LOGD(context->GetNodeName(), "Begin to do InferShapeForSituMxQuant"); + const gert::Shape* xShape = context->GetInputShape(INDEX_INPUT_X); + OP_CHECK_NULL_WITH_CONTEXT(context, xShape); + + gert::Shape* yShape = context->GetOutputShape(INDEX_OUTPUT_Y); + OP_CHECK_NULL_WITH_CONTEXT(context, yShape); + + gert::Shape* mxscaleShape = context->GetOutputShape(INDEX_OUTPUT_MXSCALE); + OP_CHECK_NULL_WITH_CONTEXT(context, mxscaleShape); + + OP_CHECK_IF(xShape->GetDimNum() < 1 || xShape->GetDimNum() > MAX_DIM_NUM, + OP_LOGE(context->GetNodeName(), "Input x rank[%lu] should be in [1, 7].", xShape->GetDimNum()), + return ge::GRAPH_FAILED); + + if (xShape->GetDimNum() == 1 && xShape->GetDim(0) == UNKNOWN_RANK_DIM) { + OP_LOGD(context->GetNodeName(), "x shape is UnknownRank, set y, mxscale shape to (-2, )"); + *yShape = *xShape; + *mxscaleShape = *xShape; + return ge::GRAPH_SUCCESS; + } + + auto attrsPtr = context->GetAttrs(); + OP_CHECK_NULL_WITH_CONTEXT(context, attrsPtr); + + const int64_t* axis = attrsPtr->GetAttrPointer(INDEX_ATTR_AXIS); + OP_CHECK_NULL_WITH_CONTEXT(context, axis); + + int64_t xRank = xShape->GetDimNum(); + int64_t lastDimIdx = xRank - 1; + int64_t axisNorm = (*axis >= 0) ? static_cast(*axis) : static_cast(*axis + xRank); + OP_CHECK_IF(axisNorm != lastDimIdx, + OP_LOGE(context->GetNodeName(), "axis must be -1 (last axis), but got axis=%ld (normalized=%ld)", + *axis, axisNorm), + return ge::GRAPH_FAILED); + + // Validate last dimension is divisible by 2 + if (xShape->GetDim(lastDimIdx) != UNKNOWN_DIM_VALUE_ && xShape->GetDim(lastDimIdx) % SPLIT_NUM != 0) { + OP_LOGE(context->GetNodeName(), "The last dimension must be divisible by 2, but got [%ld].", + xShape->GetDim(lastDimIdx)); + return ge::GRAPH_FAILED; + } + + // Step 1: Compute y shape (Situ output) + // y 的 shape 与 x 的 shape 维度一致,在最后一维上是 x 的一半 + *yShape = *xShape; + if (xShape->GetDim(lastDimIdx) != UNKNOWN_DIM_VALUE_) { + yShape->SetDim(lastDimIdx, xShape->GetDim(lastDimIdx) / SPLIT_NUM); + } + + // Step 2: Compute mxscale shape + // mxscale.shape = y.shape + // mxscale.shape[axis] = CeilDiv(y.shape[axis], 2 * 32) (even-aligned block count) + // mxscale.shape += [2] + *mxscaleShape = *yShape; + int64_t yAxisSize = 0; + if (yShape->GetDim(axisNorm) == UNKNOWN_DIM_VALUE_) { + yAxisSize = UNKNOWN_DIM_VALUE_; + } else { + int64_t yDim = yShape->GetDim(axisNorm); + yAxisSize = Ops::Base::CeilDiv(yDim, ALIGN_NUM * BLOCK_SIZE); + } + mxscaleShape->SetDim(axisNorm, yAxisSize); + mxscaleShape->AppendDim(ALIGN_NUM); + + OP_LOGI(context->GetNodeName(), "End to do InferShapeForSituMxQuant"); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus InferDataTypeForSituMxQuant(gert::InferDataTypeContext* context) +{ + OP_LOGD(context->GetNodeName(), "Begin to do InferDataTypeForSituMxQuant"); + auto attrsPtr = context->GetAttrs(); + OP_CHECK_NULL_WITH_CONTEXT(context, attrsPtr); + const int64_t* dstDtype = attrsPtr->GetAttrPointer(INDEX_ATTR_DST_TYPE); + OP_CHECK_NULL_WITH_CONTEXT(context, dstDtype); + ge::DataType outDtype = static_cast(*dstDtype); + OP_CHECK_IF( + std::find(Y_SUPPORT_DTYPE_SET.begin(), Y_SUPPORT_DTYPE_SET.end(), outDtype) == Y_SUPPORT_DTYPE_SET.end(), + OP_LOGE(context->GetNodeName(), + "dst_type is illegal, only supports 36(FLOAT8_E4M3FN) or 35(FLOAT8_E5M2). but got %d(%s) please check.", + *dstDtype, ge::TypeUtils::DataTypeToAscendString(outDtype).GetString()), + return ge::GRAPH_FAILED); + context->SetOutputDataType(INDEX_OUTPUT_Y, outDtype); + context->SetOutputDataType(INDEX_OUTPUT_MXSCALE, ge::DT_FLOAT8_E8M0); + OP_LOGI(context->GetNodeName(), "End to do InferDataTypeForSituMxQuant"); + return ge::GRAPH_SUCCESS; +} + +IMPL_OP_INFERSHAPE(SituMxQuant).InferShape(InferShapeForSituMxQuant).InferDataType(InferDataTypeForSituMxQuant); +} // namespace ops diff --git a/csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_axis_last.h b/csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_axis_last.h new file mode 100644 index 000000000000..cabacd4c5394 --- /dev/null +++ b/csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_axis_last.h @@ -0,0 +1,280 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/* ! + * \file situ_mx_quant_axis_last.h + * \brief Regbase implementation for Situ + MX quantization (activate_dim=-1, axis=-1) + */ + +#ifndef SITU_MX_QUANT_AXIS_LAST_H +#define SITU_MX_QUANT_AXIS_LAST_H + +#include "situ_mx_quant_common.h" +#include "kernel_operator.h" +#include "kernel_tiling/kernel_tiling.h" + +namespace SituMxQuant { +using namespace AscendC; + +template +class SituMxQuantAxisLast { +public: + __aicore__ inline SituMxQuantAxisLast(){}; + + __aicore__ inline void Init(GM_ADDR x, GM_ADDR y, GM_ADDR mxscale, GM_ADDR workspace, + const SituMxQuantTilingData* __restrict tilingData, AscendC::TPipe* pipe); + __aicore__ inline void Process(); + +private: + __aicore__ inline void Compute(int64_t dim0Size, int64_t dim1Size, int64_t dim1AlignSize); + __aicore__ inline void CopyIn(int64_t rowOffset, int64_t colBlockStart, int64_t dim0OnceSize, int64_t dim1OnceSize); + __aicore__ inline void CopyOut(int64_t rowOffset, int64_t colBlockStart, int64_t dim0OnceSize, int64_t dim1OnceSize, + int64_t dim1OnceSizeAlgin); + +private: + GlobalTensor xGm_; + GlobalTensor yGm_; + GlobalTensor scaleGm_; + const SituMxQuantTilingData* tiling_; + AscendC::TPipe* pipe_; + int32_t blockIdx_ = 0; + + AscendC::TQue inQuex_; + AscendC::TQue outQuey_; + AscendC::TQue outQueScale_; + + TBuf situBuffer_; + TBuf maxExpBuffer_; + TBuf halfScaleBuffer_; + + int64_t realCoreNum_ = 0; + int64_t activateLeft_ = 0; + float beta_ = 1.0f; + float invBeta_ = 1.0f; + float linearBeta_ = 0.0f; + float invLinearBeta_ = 0.0f; + + int64_t dimM_ = 0; + int64_t dim2N_ = 0; + int64_t dimN_ = 0; + int64_t factorDim0Size_ = 0; + int64_t factorDim1Size_ = 0; + + int64_t mStart_ = 0; + int64_t nStart_ = 0; + int64_t loopTimesPerBatch_ = 0; + int64_t tailPerBatch_ = 0; + int64_t loopTimesN_ = 0; + int64_t tailN_ = 0; + uint16_t f8Emax_ = 0; + int64_t outputScaleRowBytes_ = 0; +}; + +template +__aicore__ inline void SituMxQuantAxisLast::Init( + GM_ADDR x, GM_ADDR y, GM_ADDR mxscale, GM_ADDR workspace, + const SituMxQuantTilingData* __restrict tilingData, AscendC::TPipe* pipe) +{ +#if (__NPU_ARCH__ == 3510) + AscendC::SetCtrlSpr(0); +#endif + tiling_ = tilingData; + pipe_ = pipe; + blockIdx_ = GetBlockIdx(); + xGm_.SetGlobalBuffer((__gm__ T*)x); + yGm_.SetGlobalBuffer((__gm__ uint8_t*)y); + scaleGm_.SetGlobalBuffer((__gm__ uint8_t*)mxscale); + + dimM_ = tiling_->inputDim1; + dimN_ = tiling_->inputDim2; + dim2N_ = dimN_ * CONST_2; + + int64_t dimNBlockNum = tiling_->dimNBlockNum; + realCoreNum_ = tiling_->usedCoreNum; + factorDim0Size_ = tiling_->maxBasicNumUbDim1; + factorDim1Size_ = tiling_->maxBasicNumUbDim2; + + activateLeft_ = tiling_->activateLeft; + beta_ = tiling_->beta; + invBeta_ = 1.0f / beta_; + linearBeta_ = tiling_->linearBeta; + if constexpr (hasLinearBeta) { + invLinearBeta_ = 1.0f / linearBeta_; + } + + // Initialize pipe buffers + int32_t factorSize = factorDim0Size_ * factorDim1Size_; + pipe_->InitBuffer(inQuex_, CONST_2, factorSize * X_ONCE_NUM * sizeof(T)); + pipe_->InitBuffer(outQuey_, CONST_2, (factorSize * QUANT_ONCE_NUM) * sizeof(uint8_t)); + int32_t scaleUbSize = factorSize * SCALE_ONCE_NUM; + scaleUbSize = ((scaleUbSize + CONST_64 - 1) / CONST_64) * CONST_64; + pipe_->InitBuffer(outQueScale_, CONST_2, scaleUbSize); + + pipe_->InitBuffer(situBuffer_, factorSize * QUANT_ONCE_NUM * sizeof(T)); + int32_t maxExpUbSize = factorSize * SCALE_ONCE_NUM * sizeof(uint16_t); + maxExpUbSize = ((maxExpUbSize + ONE_BLOCK_UB - 1) / ONE_BLOCK_UB) * ONE_BLOCK_UB; + pipe_->InitBuffer(maxExpBuffer_, maxExpUbSize); + pipe_->InitBuffer(halfScaleBuffer_, maxExpUbSize); + + // Core grid distribution (axis=-1 only) + int64_t mCorePerB = tiling_->mCorePerB; + int64_t nCoreNum = tiling_->nCoreNum; + + int64_t nIdx = blockIdx_ % nCoreNum; + int64_t mIdx = blockIdx_ / nCoreNum; + + int64_t mHeadCores = dimM_ % mCorePerB; + int64_t mBase = dimM_ / mCorePerB; + mStart_ = (mIdx < mHeadCores) ? mIdx * (mBase + 1) : mHeadCores * (mBase + 1) + (mIdx - mHeadCores) * mBase; + int64_t mRows = (mIdx < mHeadCores) ? mBase + 1 : mBase; + + loopTimesPerBatch_ = ops::CeilDiv(mRows, factorDim0Size_); + tailPerBatch_ = mRows - (loopTimesPerBatch_ - 1) * factorDim0Size_; + + int64_t nHeadCores = dimNBlockNum % nCoreNum; + int64_t blockPerNCore = dimNBlockNum / nCoreNum; + int64_t nFrontCore = blockPerNCore + 1; + nStart_ = (nIdx < nHeadCores) ? nIdx * nFrontCore : nHeadCores * nFrontCore + (nIdx - nHeadCores) * blockPerNCore; + int64_t loopPerCoreN = (nIdx < nHeadCores) ? nFrontCore : blockPerNCore; + loopTimesN_ = ops::CeilDiv(loopPerCoreN, factorDim1Size_); + if (nIdx < nCoreNum - 1) { + tailN_ = loopPerCoreN * 256 - (loopTimesN_ - 1) * factorDim1Size_ * 256; + } else { + tailN_ = dimN_ - nStart_ * 256 - (loopTimesN_ - 1) * factorDim1Size_ * 256; + } + + outputScaleRowBytes_ = ((dimN_ + 64 - 1) / 64) * 2; + if constexpr (ops::IsSame::value) { + f8Emax_ = FP8_E4M3_MAX_EXP; + } + if constexpr (ops::IsSame::value) { + f8Emax_ = FP8_E5M2_MAX_EXP; + } +} + +template +__aicore__ inline void SituMxQuantAxisLast::Process() +{ + if (blockIdx_ >= realCoreNum_) { + return; + } + int64_t dim1Size = factorDim1Size_ * QUANT_ONCE_NUM; + int64_t dim1AlignSize = ((tailN_ + CONST_64 - 1) / CONST_64) * CONST_64; + for (int64_t mGroup = 0; mGroup < loopTimesPerBatch_; mGroup++) { + int64_t dim0Size = (mGroup == loopTimesPerBatch_ - 1) ? tailPerBatch_ : factorDim0Size_; + int64_t rowOffset = mStart_ + mGroup * factorDim0Size_; + for (int64_t nLoop = 0; nLoop < loopTimesN_; nLoop++) { + int64_t colOffset = nStart_ + nLoop * factorDim1Size_; + bool isTailDim1 = (nLoop == loopTimesN_ - 1); + int64_t dim1SizeNow = isTailDim1 ? tailN_ : dim1Size; + int64_t dim1AlignSizeNow = isTailDim1 ? dim1AlignSize : dim1Size; + CopyIn(rowOffset, colOffset, dim0Size, dim1SizeNow); + Compute(dim0Size, dim1SizeNow, dim1AlignSizeNow); + CopyOut(rowOffset, colOffset, dim0Size, dim1SizeNow, dim1AlignSizeNow); + } + } +} + +template +__aicore__ inline void SituMxQuantAxisLast::Compute( + int64_t dim0OnceSize, int64_t dim1OnceSize, int64_t dim1AlignSize) +{ + LocalTensor xlocal = inQuex_.DeQue(); + auto x1UbAddr = (__ubuf__ T*)xlocal.GetPhyAddr(); + auto x2UbAddr = (__ubuf__ T*)xlocal[factorDim0Size_ * factorDim1Size_ * QUANT_ONCE_NUM].GetPhyAddr(); + + // Determine gate and up based on activateLeft + // activateLeft=true: gate=first half (x1), up=second half (x2) + // activateLeft=false: gate=second half (x2), up=first half (x1) + __ubuf__ T* gateUbAddr = x1UbAddr; + __ubuf__ T* upUbAddr = x2UbAddr; + if (activateLeft_ == 0) { + gateUbAddr = x2UbAddr; + upUbAddr = x1UbAddr; + } + + LocalTensor situUb = situBuffer_.Get(); + auto situUbAddr = (__ubuf__ T*)situUb.GetPhyAddr(); + + // Step 1: Situ activation + ComputeVfSitu(gateUbAddr, upUbAddr, situUbAddr, dim0OnceSize, dim1OnceSize, dim1AlignSize, + beta_, invBeta_, linearBeta_, invLinearBeta_); + inQuex_.FreeTensor(xlocal); + + // Step 2: MxQuant - extract max exponent per 32-element block + LocalTensor maxExpUb = maxExpBuffer_.Get(); + auto maxExpUbAddr = (__ubuf__ uint16_t*)maxExpUb.GetPhyAddr(); + ComputeVfMaxExpVfLast(situUbAddr, maxExpUbAddr, dim0OnceSize, dim1AlignSize); + + // Step 3: MxQuant - compute E8M0 scale and reciprocal scale + LocalTensor mxScaleLocal = outQueScale_.AllocTensor(); + auto mxScaleLocalAddr = (__ubuf__ uint16_t*)mxScaleLocal.GetPhyAddr(); + LocalTensor halfScaleLocal = halfScaleBuffer_.Get(); + auto halfScaleLocalAddr = reinterpret_cast<__ubuf__ uint16_t*>(halfScaleLocal.GetPhyAddr()); + ComputeScaleLast(f8Emax_, maxExpUbAddr, mxScaleLocalAddr, halfScaleLocalAddr, dim0OnceSize, dim1AlignSize); + outQueScale_.EnQue(mxScaleLocal); + + // Step 4: MxQuant - quantize to FP8 + LocalTensor outLocal = outQuey_.AllocTensor(); + auto outLocalAddr = (__ubuf__ int8_t*)outLocal.GetPhyAddr(); + ComputeDataF8Last(situUbAddr, halfScaleLocalAddr, outLocalAddr, dim0OnceSize, dim1AlignSize); + outQuey_.EnQue(outLocal); +} + +template +__aicore__ inline void SituMxQuantAxisLast::CopyIn( + int64_t rowOffset, int64_t colBlockStart, int64_t dim0OnceSize, int64_t dim1OnceSize) +{ + LocalTensor xlocal = inQuex_.AllocTensor(); + DataCopyExtParams copyInParam = {0, 0, 0, 0, 0}; + DataCopyPadExtParams copyPadParams = {false, 0, 0, 0}; + // Load two halves of input: gate (first H) and up (second H) + // Input x shape: [..., 2H], first half = gate, second half = up + int64_t offset = rowOffset * dim2N_ + colBlockStart * QUANT_ONCE_NUM; + copyInParam.blockCount = dim0OnceSize; + copyInParam.blockLen = dim1OnceSize * sizeof(T); + copyInParam.srcStride = (dim2N_ - dim1OnceSize) * sizeof(T); + DataCopyPad(xlocal, xGm_[offset], copyInParam, copyPadParams); + DataCopyPad(xlocal[factorDim0Size_ * factorDim1Size_ * QUANT_ONCE_NUM], xGm_[offset + dimN_], copyInParam, + copyPadParams); + inQuex_.EnQue(xlocal); +} + +template +__aicore__ inline void SituMxQuantAxisLast::CopyOut( + int64_t rowOffset, int64_t colBlockStart, int64_t dim0OnceSize, int64_t dim1OnceSize, int64_t dim1OnceSizeAlgin) +{ + LocalTensor mxScaleLocal = outQueScale_.DeQue(); + LocalTensor outLocal = outQuey_.DeQue(); + + // Copy FP8 output + DataCopyExtParams copyOutParamData = {0, 0, 0, 0, 0}; + copyOutParamData.blockCount = dim0OnceSize; + copyOutParamData.blockLen = dim1OnceSize; + copyOutParamData.srcStride = (dim1OnceSizeAlgin - copyOutParamData.blockLen) / ONE_BLOCK_UB; + copyOutParamData.dstStride = dimN_ - copyOutParamData.blockLen; + int64_t offset = rowOffset * dimN_ + colBlockStart * 256; + DataCopyPad(yGm_[offset], outLocal, copyOutParamData); + + // Copy E8M0 scale output + DataCopyExtParams copyOutParamScale = {0, 0, 0, 0, 0}; + uint32_t usedFactorDim1 = dim1OnceSizeAlgin / ONE_BLOCK_UB; + copyOutParamScale.blockCount = dim0OnceSize; + copyOutParamScale.blockLen = usedFactorDim1; + copyOutParamScale.srcStride = 0; + copyOutParamScale.dstStride = outputScaleRowBytes_ - copyOutParamScale.blockLen; + int64_t offsetScale = rowOffset * outputScaleRowBytes_ + colBlockStart * SCALE_ONCE_NUM; + DataCopyPad(scaleGm_[offsetScale], mxScaleLocal, copyOutParamScale); + + outQuey_.FreeTensor(outLocal); + outQueScale_.FreeTensor(mxScaleLocal); +} +} // namespace SituMxQuant +#endif // SITU_MX_QUANT_AXIS_LAST_H diff --git a/csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_common.h b/csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_common.h new file mode 100644 index 000000000000..bd95a128bfda --- /dev/null +++ b/csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_common.h @@ -0,0 +1,500 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/* ! + * \file situ_mx_quant_common.h + * \brief Common definitions and shared regbase impl for Situ + MX quantization + */ + +#ifndef SITU_MX_QUANT_COMMON_H +#define SITU_MX_QUANT_COMMON_H + +#define FLOAT_OVERFLOW_MODE_CTRL 60 + +#include "kernel_operator.h" +#include "kernel_tiling/kernel_tiling.h" +#include "../inc/platform.h" +#include "../inc/kernel_utils.h" + +namespace SituMxQuant { +// ==================== Constants ==================== +constexpr uint16_t NAN_CUSTOMIZATION = 0x7f81; +constexpr uint16_t MAX_EXP_FOR_BF16 = 0x7f80; +constexpr uint16_t MAX_EXP_FOR_FP8 = 0x00ff; +constexpr uint16_t SPECIAL_EXP_THRESHOLD = 0x0040; +constexpr int16_t SHR_NUM_FOR_BF16 = 7; +constexpr uint16_t FP8_E4M3_MAX_EXP = 0x0400; +constexpr uint16_t FP8_E5M2_MAX_EXP = 0x0780; +constexpr uint16_t BF16_EXP_BIAS = 0x7f00; +constexpr int64_t OUT_ELE_NUM_ONE_BLK = 64; +constexpr int64_t QUANT_ONCE_NUM = 256; +constexpr int64_t X_ONCE_NUM = 512; +constexpr int64_t QUANT_ONCE_NUM_FP4 = 128; +constexpr int64_t SCALE_ONCE_NUM = 8; +constexpr int64_t CONST_64 = 64; +constexpr int64_t CONST_32 = 32; +constexpr int64_t CONST_2 = 2; +constexpr int64_t CONST_4 = 4; +constexpr uint32_t VF_LEN_T = platform::GetVRegSize() / sizeof(half); // 128 +constexpr uint32_t VF_LEN_FP32 = platform::GetVRegSize() / sizeof(float); // 64 +constexpr uint32_t ONE_BLOCK_UB = platform::GetUbBlockSize(); +constexpr uint32_t ONE_BLOCK_NUM = ONE_BLOCK_UB / sizeof(half); // 16 + +// ==================== Cast Traits ==================== +static constexpr AscendC::MicroAPI::CastTrait CAST_ZERO = { + AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::UNKNOWN, AscendC::MicroAPI::MaskMergeMode::ZEROING, + AscendC::RoundMode::UNKNOWN}; +static constexpr AscendC::MicroAPI::CastTrait CAST_ONE = { + AscendC::MicroAPI::RegLayout::ONE, AscendC::MicroAPI::SatMode::UNKNOWN, AscendC::MicroAPI::MaskMergeMode::ZEROING, + AscendC::RoundMode::UNKNOWN}; +static constexpr AscendC::MicroAPI::CastTrait CAST_FP32_TO_BF16 = { + AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::NO_SAT, AscendC::MicroAPI::MaskMergeMode::ZEROING, + AscendC::RoundMode::CAST_RINT}; +static constexpr AscendC::MicroAPI::CastTrait CAST_32_TO_80 = { + AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::SAT, AscendC::MicroAPI::MaskMergeMode::ZEROING, + AscendC::RoundMode::CAST_RINT}; +static constexpr AscendC::MicroAPI::CastTrait CAST_32_TO_81 = { + AscendC::MicroAPI::RegLayout::ONE, AscendC::MicroAPI::SatMode::SAT, AscendC::MicroAPI::MaskMergeMode::ZEROING, + AscendC::RoundMode::CAST_RINT}; +static constexpr AscendC::MicroAPI::CastTrait CAST_32_TO_82 = { + AscendC::MicroAPI::RegLayout::TWO, AscendC::MicroAPI::SatMode::SAT, AscendC::MicroAPI::MaskMergeMode::ZEROING, + AscendC::RoundMode::CAST_RINT}; +static constexpr AscendC::MicroAPI::CastTrait CAST_32_TO_83 = { + AscendC::MicroAPI::RegLayout::THREE, AscendC::MicroAPI::SatMode::SAT, AscendC::MicroAPI::MaskMergeMode::ZEROING, + AscendC::RoundMode::CAST_RINT}; + +using namespace AscendC; + +// =================================================================== +// MxQuant helper: Extract BF16 exponent and compute per-32-block max +// Adapted from swiglu_mx_quant_common.h (BF16-only path) +// =================================================================== +template +__aicore__ inline void ComputeVfMaxExpVfLast(__ubuf__ T* srcAddr, __ubuf__ uint16_t* maxExpAddr, int64_t dim0OnceSize, + int64_t alignDim1Size) +{ + uint32_t totalCountInUB = dim0OnceSize * alignDim1Size; + uint16_t loopNum = CeilDivision(totalCountInUB, QUANT_ONCE_NUM); + uint16_t maxExpbf16 = MAX_EXP_FOR_BF16; + int64_t onceNum = QUANT_ONCE_NUM; + int64_t scaleNum = SCALE_ONCE_NUM; + __VEC_SCOPE__ + { + AscendC::MicroAPI::RegTensor vdExp0, vdExp1; + AscendC::MicroAPI::RegTensor vdExpExtract0, vdExpExtract1; + AscendC::MicroAPI::RegTensor expMaskBF16, vdMaxExp; + AscendC::MicroAPI::Duplicate(expMaskBF16, maxExpbf16); + AscendC::MicroAPI::MaskReg scaleMask1; + AscendC::MicroAPI::UnalignReg u1; + for (uint16_t i = 0; i < loopNum; i++) { + scaleMask1 = AscendC::MicroAPI::UpdateMask(totalCountInUB); + AscendC::MicroAPI::DataCopy(vdExp0, vdExp1, srcAddr, onceNum); + AscendC::MicroAPI::And(vdExpExtract0, (AscendC::MicroAPI::RegTensor&)vdExp0, expMaskBF16, + scaleMask1); + AscendC::MicroAPI::And(vdExpExtract1, (AscendC::MicroAPI::RegTensor&)vdExp1, expMaskBF16, + scaleMask1); + AscendC::MicroAPI::Max(vdMaxExp, vdExpExtract0, vdExpExtract1, scaleMask1); + AscendC::MicroAPI::ReduceMaxWithDataBlock(vdMaxExp, vdMaxExp, scaleMask1); + AscendC::MicroAPI::DataCopyUnAlign( + maxExpAddr, vdMaxExp, u1, scaleNum); + } + AscendC::MicroAPI::DataCopyUnAlignPost(maxExpAddr, u1, 0); + } +} + +// =================================================================== +// MxQuant helper: Compute E8M0 scale and reciprocal scale (OCP algorithm) +// Adapted from swiglu_mx_quant_common.h +// =================================================================== +template +__aicore__ inline void ComputeScaleLast(uint16_t fEmax, __ubuf__ uint16_t* maxExpAddr, + __ubuf__ uint16_t* mxScaleLocalAddr, __ubuf__ uint16_t* halfScaleLocalAddr, + int64_t dim0OnceSize, int64_t alignDim1Size) +{ + uint32_t totalScaleInUB = dim0OnceSize * (alignDim1Size / CONST_32); + uint16_t loopNumScale = CeilDivision(totalScaleInUB, QUANT_ONCE_NUM_FP4); + uint16_t maxExpBf16 = MAX_EXP_FOR_BF16; + int64_t onceNum = QUANT_ONCE_NUM_FP4; + int64_t onceNumMxScale = CONST_64; + uint16_t bf16ExpBias = BF16_EXP_BIAS; + uint16_t maxExpFp8 = MAX_EXP_FOR_FP8; + uint16_t nanCustomZation = NAN_CUSTOMIZATION; + uint16_t specailExpThreshold = SPECIAL_EXP_THRESHOLD; + __VEC_SCOPE__ + { + AscendC::MicroAPI::RegTensor expMask, vdMaxExp; + AscendC::MicroAPI::Duplicate(expMask, maxExpBf16); + AscendC::MicroAPI::MaskReg cmpResult, zeroMask, cmpResultSub, preMaskScale; + AscendC::MicroAPI::RegTensor maxExpValue, sharedExp, scaleValue, scaleBias, halfScale; + AscendC::MicroAPI::Duplicate(maxExpValue, fEmax); + AscendC::MicroAPI::Duplicate(scaleBias, bf16ExpBias); + AscendC::MicroAPI::RegTensor fp8NanRegTensor, zeroRegTensor, nanRegTensor; + AscendC::MicroAPI::Duplicate(fp8NanRegTensor, maxExpFp8); + AscendC::MicroAPI::Duplicate(zeroRegTensor, 0); + AscendC::MicroAPI::Duplicate(nanRegTensor, nanCustomZation); + AscendC::MicroAPI::MaskReg invalidDataMask, specialDataMask; + AscendC::MicroAPI::RegTensor specialExpRegTensor; + AscendC::MicroAPI::Duplicate(specialExpRegTensor, specailExpThreshold); + for (uint16_t i = 0; i < loopNumScale; i++) { + preMaskScale = AscendC::MicroAPI::UpdateMask(totalScaleInUB); + AscendC::MicroAPI::DataCopy( + vdMaxExp, maxExpAddr, onceNum); + AscendC::MicroAPI::Compare(cmpResult, vdMaxExp, expMask, preMaskScale); + AscendC::MicroAPI::Compare(zeroMask, vdMaxExp, zeroRegTensor, preMaskScale); + AscendC::MicroAPI::Compare(invalidDataMask, vdMaxExp, maxExpValue, preMaskScale); + AscendC::MicroAPI::Select(vdMaxExp, maxExpValue, vdMaxExp, invalidDataMask); + AscendC::MicroAPI::Sub(sharedExp, vdMaxExp, maxExpValue, preMaskScale); + AscendC::MicroAPI::ShiftRights(scaleValue, sharedExp, SHR_NUM_FOR_BF16, preMaskScale); + AscendC::MicroAPI::Select(scaleValue, scaleValue, fp8NanRegTensor, cmpResult); + AscendC::MicroAPI::Select(scaleValue, scaleValue, zeroRegTensor, zeroMask); + AscendC::MicroAPI::DataCopy(mxScaleLocalAddr, scaleValue, + onceNumMxScale, preMaskScale); + AscendC::MicroAPI::Compare(specialDataMask, sharedExp, scaleBias, preMaskScale); + AscendC::MicroAPI::Sub(halfScale, scaleBias, sharedExp, preMaskScale); + AscendC::MicroAPI::Select(halfScale, halfScale, nanRegTensor, cmpResult); + AscendC::MicroAPI::Select(halfScale, halfScale, zeroRegTensor, zeroMask); + AscendC::MicroAPI::Select(halfScale, specialExpRegTensor, halfScale, specialDataMask); + AscendC::MicroAPI::DataCopy( + halfScaleLocalAddr, halfScale, onceNum, preMaskScale); + } + } +} + +// =================================================================== +// MxQuant helper: Quantize BF16 data to FP8 (multiply by reciprocal scale, then cast) +// Adapted from swiglu_mx_quant_common.h (BF16-only path) +// =================================================================== +template +__aicore__ inline void ComputeDataF8Last(__ubuf__ T* srcAddr, __ubuf__ uint16_t* halfScaleLocalAddr, + __ubuf__ int8_t* outLocalAddr, int64_t dim0OnceSize, int64_t dim1AlignSize) +{ + uint32_t totalCountInUB = dim0OnceSize * dim1AlignSize; + uint16_t loopNum = CeilDivision(totalCountInUB, QUANT_ONCE_NUM); + int64_t elementAfterReduce = SCALE_ONCE_NUM; + int64_t onceXNum = QUANT_ONCE_NUM; + __VEC_SCOPE__ + { + AscendC::MicroAPI::RegTensor halfScaleForMul; + AscendC::MicroAPI::RegTensor vdExp0, vdExp1; + AscendC::MicroAPI::RegTensor vdExp0FP32Zero, vdExp0FP32One; + AscendC::MicroAPI::RegTensor vdExp1FP32Zero, vdExp1FP32One; + AscendC::MicroAPI::RegTensor vdExp0FP8Zero, vdExp0FP8One; + AscendC::MicroAPI::RegTensor vdExp1FP8Zero, vdExp1FP8One; + AscendC::MicroAPI::MaskReg + maskAll = AscendC::MicroAPI::CreateMask(); + AscendC::MicroAPI::MaskReg + maskAllB8 = AscendC::MicroAPI::CreateMask(); + for (uint16_t i = 0; i < loopNum; i++) { + AscendC::MicroAPI::DataCopy(vdExp0, vdExp1, srcAddr, + onceXNum); + AscendC::MicroAPI::DataCopy(halfScaleForMul, halfScaleLocalAddr, + elementAfterReduce); + // BF16 path: multiply in BF16 domain, then cast to FP32, then to FP8 + AscendC::MicroAPI::Mul(vdExp0, vdExp0, (AscendC::MicroAPI::RegTensor&)halfScaleForMul, maskAll); + AscendC::MicroAPI::Mul(vdExp1, vdExp1, (AscendC::MicroAPI::RegTensor&)halfScaleForMul, maskAll); + AscendC::MicroAPI::Cast(vdExp0FP32Zero, vdExp0, maskAll); + AscendC::MicroAPI::Cast(vdExp0FP32One, vdExp0, maskAll); + AscendC::MicroAPI::Cast(vdExp1FP32Zero, vdExp1, maskAll); + AscendC::MicroAPI::Cast(vdExp1FP32One, vdExp1, maskAll); + AscendC::MicroAPI::Cast(vdExp0FP8Zero, vdExp0FP32Zero, maskAll); + AscendC::MicroAPI::Cast(vdExp0FP8One, vdExp0FP32One, maskAll); + AscendC::MicroAPI::Cast(vdExp1FP8Zero, vdExp1FP32Zero, maskAll); + AscendC::MicroAPI::Cast(vdExp1FP8One, vdExp1FP32One, maskAll); + AscendC::MicroAPI::Add((AscendC::MicroAPI::RegTensor&)vdExp0FP8Zero, + (AscendC::MicroAPI::RegTensor&)vdExp0FP8Zero, + (AscendC::MicroAPI::RegTensor&)vdExp0FP8One, maskAllB8); + AscendC::MicroAPI::Add((AscendC::MicroAPI::RegTensor&)vdExp0FP8Zero, + (AscendC::MicroAPI::RegTensor&)vdExp0FP8Zero, + (AscendC::MicroAPI::RegTensor&)vdExp1FP8Zero, maskAllB8); + AscendC::MicroAPI::Add((AscendC::MicroAPI::RegTensor&)vdExp0FP8Zero, + (AscendC::MicroAPI::RegTensor&)vdExp0FP8Zero, + (AscendC::MicroAPI::RegTensor&)vdExp1FP8One, maskAllB8); + AscendC::MicroAPI::DataCopy( + outLocalAddr, (AscendC::MicroAPI::RegTensor&)vdExp0FP8Zero, onceXNum, maskAllB8); + } + } +} + +// =================================================================== +// Situ activation: beta * tanh(gate / beta) * sigmoid(gate) * up +// (+ optional linear_beta * tanh(up / linear_beta) on up) +// Replaces ComputeVfSwigluV1 from swiglu_mx_quant +// =================================================================== +template +__aicore__ inline void ComputeVfSitu(__local_mem__ T* gateUbAddr, __local_mem__ T* upUbAddr, + __local_mem__ T* situUbAddr, int64_t dim0OnceSize, int64_t dim1OnceSize, + int64_t dim1AlignSize, float beta, float invBeta, float linearBeta, + float invLinearBeta) +{ + uint16_t dim0VfTimes = dim0OnceSize; + uint16_t dim1VfTimes = dim1OnceSize / VF_LEN_FP32; + uint32_t dim1Tail = dim1OnceSize % VF_LEN_FP32; + uint16_t dim1TailTimes = 0; + uint16_t dim1Tail2 = 0; + uint32_t mask1Num = 0; + uint32_t mask2Num = 0; + uint32_t mask3Num = 0; + uint32_t alignDim1In = ((dim1OnceSize + ONE_BLOCK_NUM - 1) / ONE_BLOCK_NUM) * ONE_BLOCK_NUM; + uint32_t alignDim1Out = dim1AlignSize; + auto gateUbAddr1 = gateUbAddr; + auto upUbAddr1 = upUbAddr; + auto situUbAddr1 = situUbAddr; + auto situUbAddr2 = situUbAddr; + T numZero = 0; + if (dim1Tail > 0) { + mask1Num = dim1Tail; + dim1TailTimes = 1; + uint32_t padNum = alignDim1Out - dim1VfTimes * VF_LEN_FP32; + if (padNum <= VF_LEN_FP32) { + mask2Num = padNum; + } else { + dim1Tail2 = 1; + mask2Num = VF_LEN_FP32; + mask3Num = padNum - VF_LEN_FP32; + } + int32_t offsetAlgin = dim1VfTimes * VF_LEN_FP32; + gateUbAddr1 = gateUbAddr + offsetAlgin; + upUbAddr1 = upUbAddr + offsetAlgin; + situUbAddr1 = situUbAddr + offsetAlgin; + situUbAddr2 = situUbAddr + offsetAlgin + dim1TailTimes * VF_LEN_FP32; + } + float scalarOne = 1.0f; + float negScalarOne = -1.0f; + float scalarTwo = 2.0f; + float negTwo = -2.0f; + // Two-path tanh (adapted from tanh.h reference): + // |x| < 0.6: degree-9 polynomial, FMA Horner (matches tanh.h exactly) + // |x| >= 0.6: sigmoid decomposition, sign naturally preserved + float tanhC1 = -0.333327681f; + float tanhC2 = 0.133152977f; + float tanhC3 = -0.0523039624f; + float tanhC4 = 0.0157396831f; + float tanhThreshold = 0.6f; + __VEC_SCOPE__ + { + AscendC::MicroAPI::RegTensor vregGate; + AscendC::MicroAPI::RegTensor vregUp; + AscendC::MicroAPI::RegTensor gateF; + AscendC::MicroAPI::RegTensor upF; + AscendC::MicroAPI::RegTensor gateDivBeta; + AscendC::MicroAPI::RegTensor polyReg; // sigmoid path result + AscendC::MicroAPI::RegTensor x2; // x² for Horner / temp + AscendC::MicroAPI::MaskReg cmpMask; // comparison result for Select + AscendC::MicroAPI::RegTensor negGate; + AscendC::MicroAPI::RegTensor expReg; // sigmoid Exp + AscendC::MicroAPI::RegTensor oneReg; + AscendC::MicroAPI::RegTensor sigmoidReg; // sigmoid result / linear_beta work reg + AscendC::MicroAPI::RegTensor c1Reg; // tanh polynomial coeff c1 (preloaded) + AscendC::MicroAPI::RegTensor c2Reg; // tanh polynomial coeff c2 (preloaded) + AscendC::MicroAPI::RegTensor outFReg; + AscendC::MicroAPI::RegTensor outTReg; + AscendC::MicroAPI::MaskReg mask = AscendC::MicroAPI::CreateMask(); + AscendC::MicroAPI::MaskReg mask1 = AscendC::MicroAPI::UpdateMask(mask1Num); + AscendC::MicroAPI::MaskReg mask2 = AscendC::MicroAPI::UpdateMask(mask2Num); + AscendC::MicroAPI::MaskReg mask3 = AscendC::MicroAPI::UpdateMask(mask3Num); + AscendC::MicroAPI::Duplicate(oneReg, scalarOne); + AscendC::MicroAPI::Duplicate(c1Reg, tanhC1); + AscendC::MicroAPI::Duplicate(c2Reg, tanhC2); + for (uint16_t dim0vfLoopIdx = 0; dim0vfLoopIdx < dim0VfTimes; dim0vfLoopIdx++) { + for (uint16_t dim1vfLoopIdx = 0; dim1vfLoopIdx < dim1VfTimes; dim1vfLoopIdx++) { + AscendC::MicroAPI::AddrReg srcIdxOffset = AscendC::MicroAPI::CreateAddrReg( + dim0vfLoopIdx, alignDim1In, dim1vfLoopIdx, 64); + AscendC::MicroAPI::DataCopy(vregGate, gateUbAddr, + srcIdxOffset); + AscendC::MicroAPI::DataCopy(vregUp, upUbAddr, + srcIdxOffset); + AscendC::MicroAPI::Cast(gateF, vregGate, mask); + AscendC::MicroAPI::Cast(upF, vregUp, mask); + + // Two-path tanh(gate/beta) — adapted from tanh.h reference: + // small |x|: degree-7 polynomial (Horner) + // large |x|: sigmoid decomposition on |x|, sign restore + AscendC::MicroAPI::Muls(gateDivBeta, gateF, invBeta, mask); // x = gate/beta + + // --- Polynomial path (all x, used for |x| < 0.6) --- + // tanh(x) ≈ x * (1 + c1*x² + c2*x⁴ + c3*x⁶ + c4*x⁸) + // FMA Horner (matches tanh.h reference: 7 ops, 7 roundings) + AscendC::MicroAPI::Mul(x2, gateDivBeta, gateDivBeta, mask); // x² + AscendC::MicroAPI::Muls(sigmoidReg, x2, tanhC4, mask); // c4*x² + AscendC::MicroAPI::Adds(sigmoidReg, sigmoidReg, tanhC3, mask); // +c3 + AscendC::MicroAPI::FusedMulDstAdd(sigmoidReg, x2, c2Reg, mask); // *x²+c2 + AscendC::MicroAPI::FusedMulDstAdd(sigmoidReg, x2, c1Reg, mask); // *x²+c1 + AscendC::MicroAPI::Mul(sigmoidReg, sigmoidReg, x2, mask); // *x² + AscendC::MicroAPI::FusedMulDstAdd(sigmoidReg, gateDivBeta, gateDivBeta, mask); // *x+x = x*(p+1) + + // --- Sigmoid path (used for |x| >= 0.6) --- + // tanh(x) = 2*sigmoid(2x) - 1 = 2/(1+exp(-2x)) - 1, sign naturally preserved + AscendC::MicroAPI::Muls(negGate, gateDivBeta, negTwo, mask); // -2x + AscendC::MicroAPI::Exp(expReg, negGate, mask); + AscendC::MicroAPI::Adds(expReg, expReg, scalarOne, mask); // 1+exp(-2x) + AscendC::MicroAPI::Div(polyReg, oneReg, expReg, mask); // sigmoid = 1/(1+exp(-2x)) + AscendC::MicroAPI::Muls(polyReg, polyReg, scalarTwo, mask); // 2*sigmoid + AscendC::MicroAPI::Adds(polyReg, polyReg, negScalarOne, mask); // 2*sigmoid - 1 + + // --- Path selection: sigmoid if |x| >= 0.6, else polynomial --- + AscendC::MicroAPI::Muls(x2, gateDivBeta, negScalarOne, mask); + AscendC::MicroAPI::Max(x2, gateDivBeta, x2, mask); // |x| + AscendC::MicroAPI::Duplicate(expReg, tanhThreshold); // 0.6 + AscendC::MicroAPI::Compare(cmpMask, x2, expReg, mask); + AscendC::MicroAPI::Select(sigmoidReg, polyReg, sigmoidReg, cmpMask); + // sigmoidReg = tanh(gate/beta) — save to negGate (free after |x|) + AscendC::MicroAPI::Mul(negGate, sigmoidReg, oneReg, mask); // negGate = tanh + + // sigmoid(gate) = 1 / (1 + exp(-gate)) → result in polyReg + AscendC::MicroAPI::Muls(polyReg, gateF, negScalarOne, mask); // -gate + AscendC::MicroAPI::Exp(expReg, polyReg, mask); + AscendC::MicroAPI::Adds(expReg, expReg, scalarOne, mask); + AscendC::MicroAPI::Div(polyReg, oneReg, expReg, mask); // sigmoid(gate) + + // situ_a = beta * tanh * sigmoid + AscendC::MicroAPI::Mul(polyReg, negGate, polyReg, mask); // tanh * sigmoid + AscendC::MicroAPI::Muls(polyReg, polyReg, beta, mask); // * beta + + // Optional: up = linear_beta * tanh(up / linear_beta) + // Uses sigmoidReg/negGate as work registers to preserve polyReg (situ_a) + if constexpr (hasLinearBeta) { + AscendC::MicroAPI::Muls(upF, upF, invLinearBeta, mask); // x = up/lb + + // Poly path → sigmoidReg (FMA Horner) + AscendC::MicroAPI::Mul(x2, upF, upF, mask); + AscendC::MicroAPI::Muls(sigmoidReg, x2, tanhC4, mask); + AscendC::MicroAPI::Adds(sigmoidReg, sigmoidReg, tanhC3, mask); + AscendC::MicroAPI::FusedMulDstAdd(sigmoidReg, x2, c2Reg, mask); + AscendC::MicroAPI::FusedMulDstAdd(sigmoidReg, x2, c1Reg, mask); + AscendC::MicroAPI::Mul(sigmoidReg, sigmoidReg, x2, mask); + AscendC::MicroAPI::FusedMulDstAdd(sigmoidReg, upF, upF, mask); + + // Sigmoid path on x (sign naturally preserved) → negGate + AscendC::MicroAPI::Muls(expReg, upF, negTwo, mask); // -2x + AscendC::MicroAPI::Exp(expReg, expReg, mask); + AscendC::MicroAPI::Adds(expReg, expReg, scalarOne, mask); // 1+exp(-2x) + AscendC::MicroAPI::Div(negGate, oneReg, expReg, mask); + AscendC::MicroAPI::Muls(negGate, negGate, scalarTwo, mask); + AscendC::MicroAPI::Adds(negGate, negGate, negScalarOne, mask); // 2*sig-1 + + // Path selection → sigmoidReg = tanh(up/lb) + AscendC::MicroAPI::Muls(x2, upF, negScalarOne, mask); + AscendC::MicroAPI::Max(x2, upF, x2, mask); // |x| + AscendC::MicroAPI::Duplicate(expReg, tanhThreshold); + AscendC::MicroAPI::Compare(cmpMask, x2, expReg, mask); + AscendC::MicroAPI::Select(sigmoidReg, negGate, sigmoidReg, cmpMask); + + AscendC::MicroAPI::Muls(upF, sigmoidReg, linearBeta, mask); + } + + // situOut = situ_a * up + AscendC::MicroAPI::Mul(outFReg, polyReg, upF, mask); + + AscendC::MicroAPI::Cast(outTReg, outFReg, mask); + AscendC::MicroAPI::AddrReg outOffset = AscendC::MicroAPI::CreateAddrReg(dim0vfLoopIdx, alignDim1Out, + dim1vfLoopIdx, 64); + DataCopy(situUbAddr, outTReg, outOffset, mask); + } + // Handle tail elements + AscendC::MicroAPI::AddrReg srcIdxOffset1 = AscendC::MicroAPI::CreateAddrReg(dim0vfLoopIdx, alignDim1In); + AscendC::MicroAPI::AddrReg outOffset1 = AscendC::MicroAPI::CreateAddrReg(dim0vfLoopIdx, alignDim1Out); + for (uint16_t aa = 0; aa < dim1TailTimes; aa++) { + AscendC::MicroAPI::DataCopy(vregGate, gateUbAddr1, + srcIdxOffset1); + AscendC::MicroAPI::DataCopy(vregUp, upUbAddr1, + srcIdxOffset1); + AscendC::MicroAPI::Cast(gateF, vregGate, mask1); + AscendC::MicroAPI::Cast(upF, vregUp, mask1); + + // Two-path tanh(gate/beta) — tail path + AscendC::MicroAPI::Muls(gateDivBeta, gateF, invBeta, mask1); + + // Poly path → sigmoidReg (FMA Horner) + AscendC::MicroAPI::Mul(x2, gateDivBeta, gateDivBeta, mask1); + AscendC::MicroAPI::Muls(sigmoidReg, x2, tanhC4, mask1); + AscendC::MicroAPI::Adds(sigmoidReg, sigmoidReg, tanhC3, mask1); + AscendC::MicroAPI::FusedMulDstAdd(sigmoidReg, x2, c2Reg, mask1); + AscendC::MicroAPI::FusedMulDstAdd(sigmoidReg, x2, c1Reg, mask1); + AscendC::MicroAPI::Mul(sigmoidReg, sigmoidReg, x2, mask1); + AscendC::MicroAPI::FusedMulDstAdd(sigmoidReg, gateDivBeta, gateDivBeta, mask1); + + // Sigmoid path on x → polyReg: tanh(x) = 2/(1+exp(-2x)) - 1 + AscendC::MicroAPI::Muls(negGate, gateDivBeta, negTwo, mask1); + AscendC::MicroAPI::Exp(expReg, negGate, mask1); + AscendC::MicroAPI::Adds(expReg, expReg, scalarOne, mask1); + AscendC::MicroAPI::Div(polyReg, oneReg, expReg, mask1); + AscendC::MicroAPI::Muls(polyReg, polyReg, scalarTwo, mask1); + AscendC::MicroAPI::Adds(polyReg, polyReg, negScalarOne, mask1); + + // Path selection → sigmoidReg = tanh + AscendC::MicroAPI::Muls(x2, gateDivBeta, negScalarOne, mask1); + AscendC::MicroAPI::Max(x2, gateDivBeta, x2, mask1); // |x| + AscendC::MicroAPI::Duplicate(expReg, tanhThreshold); + AscendC::MicroAPI::Compare(cmpMask, x2, expReg, mask1); + AscendC::MicroAPI::Select(sigmoidReg, polyReg, sigmoidReg, cmpMask); + AscendC::MicroAPI::Mul(negGate, sigmoidReg, oneReg, mask1); // save tanh + + // sigmoid(gate) → polyReg + AscendC::MicroAPI::Muls(polyReg, gateF, negScalarOne, mask1); + AscendC::MicroAPI::Exp(expReg, polyReg, mask1); + AscendC::MicroAPI::Adds(expReg, expReg, scalarOne, mask1); + AscendC::MicroAPI::Div(polyReg, oneReg, expReg, mask1); + + // situ_a = beta * tanh * sigmoid + AscendC::MicroAPI::Mul(polyReg, negGate, polyReg, mask1); + AscendC::MicroAPI::Muls(polyReg, polyReg, beta, mask1); + + // Optional: up = linear_beta * tanh(up / linear_beta) + if constexpr (hasLinearBeta) { + AscendC::MicroAPI::Muls(upF, upF, invLinearBeta, mask1); + + AscendC::MicroAPI::Mul(x2, upF, upF, mask1); + AscendC::MicroAPI::Muls(sigmoidReg, x2, tanhC4, mask1); + AscendC::MicroAPI::Adds(sigmoidReg, sigmoidReg, tanhC3, mask1); + AscendC::MicroAPI::FusedMulDstAdd(sigmoidReg, x2, c2Reg, mask1); + AscendC::MicroAPI::FusedMulDstAdd(sigmoidReg, x2, c1Reg, mask1); + AscendC::MicroAPI::Mul(sigmoidReg, sigmoidReg, x2, mask1); + AscendC::MicroAPI::FusedMulDstAdd(sigmoidReg, upF, upF, mask1); + + // Sigmoid path on x → negGate: tanh(x) = 2/(1+exp(-2x)) - 1 + AscendC::MicroAPI::Muls(expReg, upF, negTwo, mask1); + AscendC::MicroAPI::Exp(expReg, expReg, mask1); + AscendC::MicroAPI::Adds(expReg, expReg, scalarOne, mask1); + AscendC::MicroAPI::Div(negGate, oneReg, expReg, mask1); + AscendC::MicroAPI::Muls(negGate, negGate, scalarTwo, mask1); + AscendC::MicroAPI::Adds(negGate, negGate, negScalarOne, mask1); + + // Path selection → sigmoidReg = tanh(up/lb) + AscendC::MicroAPI::Muls(x2, upF, negScalarOne, mask1); + AscendC::MicroAPI::Max(x2, upF, x2, mask1); // |x| + AscendC::MicroAPI::Duplicate(expReg, tanhThreshold); + AscendC::MicroAPI::Compare(cmpMask, x2, expReg, mask1); + AscendC::MicroAPI::Select(sigmoidReg, negGate, sigmoidReg, cmpMask); + + AscendC::MicroAPI::Muls(upF, sigmoidReg, linearBeta, mask1); + } + + // situOut = situ_a * up + AscendC::MicroAPI::Mul(outFReg, polyReg, upF, mask1); + + AscendC::MicroAPI::Cast(outTReg, outFReg, mask1); + DataCopy(situUbAddr1, outTReg, outOffset1, mask2); + } + for (uint16_t cc = 0; cc < dim1Tail2; cc++) { + AscendC::MicroAPI::Duplicate(vregGate, numZero); + DataCopy(situUbAddr2, vregGate, outOffset1, mask3); + } + } + } +} + +} // namespace SituMxQuant + +#endif // SITU_MX_QUANT_COMMON_H diff --git a/csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_tiling_data.h b/csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_tiling_data.h new file mode 100644 index 000000000000..586f7267c236 --- /dev/null +++ b/csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_tiling_data.h @@ -0,0 +1,52 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file situ_mx_quant_tiling_data.h + * \brief Tiling data structure for Situ + MX quantization + */ + +#ifndef SITU_MX_QUANT_TILING_DATA_H +#define SITU_MX_QUANT_TILING_DATA_H + +struct SituMxQuantTilingData { + // Basic parameters + int64_t usedCoreNum; + + // 3D data shape: [inputDim0, inputDim1, inputDim2] + // inputDim0 = batch dim (=1 for 2D, product of leading dims for >2D) + // inputDim1 = row dim (M) + // inputDim2 = Situ output dim = last_dim / 2 (N) + int64_t inputDim0; + int64_t inputDim1; + int64_t inputDim2; + + // Block distribution + int64_t dimNBlockNum; // CeilDiv(N, 256) + + // Memory allocation parameters + int64_t maxBasicNumUbDim2; // UB 内最大列方向 block 数 + int64_t maxBasicNumUbDim1; // UB 内最大行数 + + // Core grid distribution + int64_t nCoreNum; // cores in N direction + int64_t mCorePerB; // M-cores per batch + + // Inter-core split parameters + int64_t frontCoreNum; + int64_t tailCoreBasicNumDim1; + + // Attributes + int64_t activateLeft; + float beta; + float linearBeta; + int64_t hasLinearBeta; // 0 or 1 +}; +#endif // SITU_MX_QUANT_TILING_DATA_H diff --git a/csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_tiling_key.h b/csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_tiling_key.h new file mode 100644 index 000000000000..2249b9145186 --- /dev/null +++ b/csrc/moe/situ_mx_quant/op_kernel/arch35/situ_mx_quant_tiling_key.h @@ -0,0 +1,38 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file situ_mx_quant_tiling_key.h + * \brief TPL tiling key template argument declarations for SituMxQuant + */ + +#ifndef SITU_MX_QUANT_TILING_KEY_H +#define SITU_MX_QUANT_TILING_KEY_H + +#include "ascendc/host_api/tiling/template_argument.h" + +#define TPL_NO_LINEAR_BETA 0 +#define TPL_HAS_LINEAR_BETA 1 + +#define TPL_DST_E4M3FN 0 +#define TPL_DST_E5M2 1 + +namespace SituMxQuantOp { +ASCENDC_TPL_ARGS_DECL(SituMxQuant, + ASCENDC_TPL_UINT_DECL(hasLinearBeta, 2, ASCENDC_TPL_UI_LIST, TPL_NO_LINEAR_BETA, + TPL_HAS_LINEAR_BETA), + ASCENDC_TPL_UINT_DECL(dstTypeIndex, 2, ASCENDC_TPL_UI_LIST, TPL_DST_E4M3FN, TPL_DST_E5M2)); + +ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL( + ASCENDC_TPL_UINT_SEL(hasLinearBeta, ASCENDC_TPL_UI_LIST, TPL_NO_LINEAR_BETA, TPL_HAS_LINEAR_BETA), + ASCENDC_TPL_UINT_SEL(dstTypeIndex, ASCENDC_TPL_UI_LIST, TPL_DST_E4M3FN, TPL_DST_E5M2))); +} // namespace SituMxQuantOp + +#endif // SITU_MX_QUANT_TILING_KEY_H diff --git a/csrc/moe/situ_mx_quant/op_kernel/inc/kernel_utils.h b/csrc/moe/situ_mx_quant/op_kernel/inc/kernel_utils.h new file mode 100644 index 000000000000..32089593cde3 --- /dev/null +++ b/csrc/moe/situ_mx_quant/op_kernel/inc/kernel_utils.h @@ -0,0 +1,71 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file kernel_utils.h + * \brief + */ +#ifndef OPS_BUILT_IN_OP_ASCENDC_KERNEL_UTILS_H_ +#define OPS_BUILT_IN_OP_ASCENDC_KERNEL_UTILS_H_ + +namespace ops { + +template +struct IntegralConstant { + static constexpr Tp value = v; +}; +using trueType = IntegralConstant; +using falseType = IntegralConstant; +template +struct IsSame : public falseType {}; +template +struct IsSame : public trueType {}; + +template +__aicore__ inline T Ceil(T a, T b) +{ + return (a + b - 1) / b; +} + +template +__aicore__ inline T CeilAlign(T a, T b) +{ + return (a + b - 1) / b * b; +} + +template +__aicore__ inline T CeilDiv(T a, T b) +{ + if (b == 0) { + return a; + } + return (a + b - 1) / b; +} + +template +__aicore__ inline T FloorDiv(T a, T b) +{ + if (b == 0) { + return a; + } + return a / b; +} + +template +__aicore__ inline T Aligned(T value, T alignment) +{ + if (alignment == 0) { + return value; + } + return (value + alignment - 1) / alignment * alignment; +} + +} // namespace ops +#endif // OPS_BUILT_IN_OP_ASCENDC_KERNEL_UTILS_H_ \ No newline at end of file diff --git a/csrc/moe/situ_mx_quant/op_kernel/inc/platform.h b/csrc/moe/situ_mx_quant/op_kernel/inc/platform.h new file mode 100644 index 000000000000..19d59e8c04a0 --- /dev/null +++ b/csrc/moe/situ_mx_quant/op_kernel/inc/platform.h @@ -0,0 +1,81 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file platform.h + * \brief platform apator + */ +#ifndef OPS_BUILT_IN_OP_ASCENDC_PLATFORM_INFO_H_ +#define OPS_BUILT_IN_OP_ASCENDC_PLATFORM_INFO_H_ + +#if ASC_DEVKIT_MAJOR >= 9 +#include "kernel_basic_intf.h" +#else +#include "kernel_operator.h" +#endif +#include "kernel_tiling/kernel_tiling.h" +#include "kernel_utils.h" + +#ifndef KERNEL_API +#define KERNEL_API extern "C" __global__ __aicore__ +#endif + +namespace platform { + +#define MID_THREAD_NUM 1024 + +__aicore__ inline constexpr bool IsDataCopyPadSupport() +{ +#if __CCE_AICORE__ == 220 + return true; +#else + return false; +#endif +} + +/** + * Get the block size of unified buffer in bytes + */ +__aicore__ inline constexpr uint32_t GetUbBlockSize() { return 32U; } + +/** + * Get the size of vector registers in bytes + */ +__aicore__ inline constexpr uint32_t GetVRegSize() +{ +#if __CCE_AICORE__ == 310 + return AscendC::VECTOR_REG_WIDTH; +#else + return 256U; +#endif +} + +/** + * Check whether the type is supported by atomic add for simd + */ +template +__aicore__ inline constexpr bool IsSupportAtomicAddTypeSIMD() +{ +#if __CCE_AICORE__ == 310 + return ops::IsSame::value || ops::IsSame::value || ops::IsSame::value || + ops::IsSame::value || ops::IsSame::value || ops::IsSame::value; +#else + return false; +#endif +} + +} // namespace platform + +namespace PlatformSocInfo { +__aicore__ inline constexpr bool IsDataCopyPadSupport() { return platform::IsDataCopyPadSupport(); } + +} // namespace PlatformSocInfo + +#endif // OPS_BUILT_IN_OP_ASCENDC_PLATFORM_INFO_H_ \ No newline at end of file diff --git a/csrc/moe/situ_mx_quant/op_kernel/situ_mx_quant_apt.cpp b/csrc/moe/situ_mx_quant/op_kernel/situ_mx_quant_apt.cpp new file mode 100644 index 000000000000..ac3b79e9a321 --- /dev/null +++ b/csrc/moe/situ_mx_quant/op_kernel/situ_mx_quant_apt.cpp @@ -0,0 +1,53 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file situ_mx_quant_apt.cpp + * \brief Kernel entry point for Situ + MX quantization operator + */ + +#include "kernel_tiling/kernel_tiling.h" +#include "kernel_operator.h" +#include "arch35/situ_mx_quant_tiling_key.h" +#include "arch35/situ_mx_quant_tiling_data.h" +#include "arch35/situ_mx_quant_common.h" +#include "arch35/situ_mx_quant_axis_last.h" + +using namespace AscendC; +using namespace SituMxQuantOp; + +template +__global__ __aicore__ void situ_mx_quant(GM_ADDR x, GM_ADDR y, GM_ADDR mxscale, GM_ADDR workspace, GM_ADDR tiling) +{ + KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY); + REGISTER_TILING_DEFAULT(SituMxQuantTilingData); + GET_TILING_DATA_WITH_STRUCT(SituMxQuantTilingData, tilingData, tiling); + GM_ADDR usrWorkspace = AscendC::GetUserWorkspace(workspace); + TPipe pipe; +#if (__NPU_ARCH__ == 3510) + int64_t oriOverflowMode = AscendC::GetCtrlSpr(); +#endif + + if constexpr (dstTypeIndex == TPL_DST_E4M3FN) { + constexpr bool useLinearBeta = (hasLinearBeta == TPL_HAS_LINEAR_BETA); + SituMxQuant::SituMxQuantAxisLast op; + op.Init(x, y, mxscale, usrWorkspace, &tilingData, &pipe); + op.Process(); + } else { + constexpr bool useLinearBeta = (hasLinearBeta == TPL_HAS_LINEAR_BETA); + SituMxQuant::SituMxQuantAxisLast op; + op.Init(x, y, mxscale, usrWorkspace, &tilingData, &pipe); + op.Process(); + } + +#if (__NPU_ARCH__ == 3510) + AscendC::SetCtrlSpr(oriOverflowMode); +#endif +} diff --git a/csrc/moe/situ_mx_quant/situ_mx_quant_torch_adpt.h b/csrc/moe/situ_mx_quant/situ_mx_quant_torch_adpt.h new file mode 100644 index 000000000000..b800b523cb4b --- /dev/null +++ b/csrc/moe/situ_mx_quant/situ_mx_quant_torch_adpt.h @@ -0,0 +1,72 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2026. 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. + * 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. + */ + +#ifndef SITU_MX_QUANT_TORCH_ADPT_H +#define SITU_MX_QUANT_TORCH_ADPT_H + +namespace vllm_ascend { + +std::tuple situ_mx_quant( + const at::Tensor& x, + double beta, + double linear_beta, + bool activate_left, + int64_t dst_type) +{ + constexpr int64_t DST_TYPE_E5M2 = 35; + constexpr int64_t DST_TYPE_E4M3FN = 36; + constexpr int64_t MX_BLOCK_SPAN = 64; + constexpr int64_t MX_SCALE_ALIGN = 2; + + TORCH_CHECK(x.dim() >= 1, + "situ_mx_quant: x must be at least 1-dimensional, but got ", + x.dim()); + TORCH_CHECK(x.size(-1) % 2 == 0, + "situ_mx_quant: x last dim must be even, but got ", x.size(-1)); + TORCH_CHECK(x.scalar_type() == at::kBFloat16, + "situ_mx_quant: x must be bfloat16, but got ", x.scalar_type()); + TORCH_CHECK(beta > 0.0, + "situ_mx_quant: beta must be greater than 0, but got ", beta); + TORCH_CHECK(dst_type == DST_TYPE_E4M3FN || dst_type == DST_TYPE_E5M2, + "situ_mx_quant: dst_type must be 36 (E4M3FN) or 35 (E5M2), but got ", + dst_type); + + std::vector y_shape(x.sizes().begin(), x.sizes().end()); + y_shape.back() /= 2; + std::vector mxscale_shape(y_shape.begin(), y_shape.end()); + mxscale_shape.back() = (y_shape.back() + MX_BLOCK_SPAN - 1) / MX_BLOCK_SPAN; + mxscale_shape.push_back(MX_SCALE_ALIGN); + + auto y_dtype = dst_type == DST_TYPE_E5M2 ? at::kFloat8_e5m2 : at::kFloat8_e4m3fn; + at::Tensor y = at::empty(y_shape, x.options().dtype(y_dtype)); + at::Tensor mxscale = at::empty(mxscale_shape, x.options().dtype(at::kFloat8_e8m0fnu)); + + constexpr int64_t AXIS = -1; + EXEC_NPU_CMD(aclnnSituMxQuant, + x, + beta, + linear_beta, + activate_left, + AXIS, + dst_type, + y, + mxscale); + return {y, mxscale}; +} + +} // namespace vllm_ascend + +#endif // SITU_MX_QUANT_TORCH_ADPT_H diff --git a/csrc/torch_binding.cpp b/csrc/torch_binding.cpp index 88cf47e3aeae..21d747ae0eb3 100644 --- a/csrc/torch_binding.cpp +++ b/csrc/torch_binding.cpp @@ -44,12 +44,15 @@ #include "attention/lightning_indexer_quant/lightning_indexer_quant_torch_adpt.h" #include "moe/causal_conv1d_v310/causal_conv1d_310_torch_adpt.h" #include "attention/recurrent_gated_delta_rule/recurrent_gated_delta_rule_torch_adpt.h" +#include "attention/recurrent_kda/recurrent_kda_torch_adpt.h" #include "attention/recurrent_gated_delta_rule_v310/recurrent_gated_delta_rule_310_torch_adpt.h" #include "attention/k2q_csr/k2q_csr_torch_adpt.h" #include "attention/msa_index_score/msa_index_score_torch_adpt.h" #include "attention/sparse_attention_score/sparse_attention_score_torch_adpt.h" #include "attention/store_kv_block/store_kv_block_torch_adpt.h" #include "attention/store_kv_block_metadata/store_kv_block_metadata_torch_adpt.cpp" +#include "moe/dequant_situ_quant/dequant_situ_quant_torch_adpt.h" +#include "moe/situ_mx_quant/situ_mx_quant_torch_adpt.h" #include #include #include @@ -1920,6 +1923,320 @@ at::Tensor npu_sparse_attention_score_prefill( return output; } +namespace { + +int64_t kda_ceil_div(int64_t x, int64_t y) +{ + return (x + y - 1) / y; +} + +int64_t get_kda_seq_num(int64_t batch, const c10::optional &cu_seqlens) +{ + if (!cu_seqlens.has_value()) { + return batch; + } + return static_cast(cu_seqlens.value().size()) - 1; +} + +void check_kda_cu_seqlens(const c10::optional &cu_seqlens, + int64_t total_tokens, + const char *op_name) +{ + if (!cu_seqlens.has_value()) { + return; + } + auto cu = cu_seqlens.value(); + TORCH_CHECK(cu.size() >= 2, op_name, ": cu_seqlens must contain at least [0, total_tokens]."); + TORCH_CHECK(cu[0] == 0, op_name, ": cu_seqlens[0] must be 0, but got ", cu[0], "."); + TORCH_CHECK(cu[cu.size() - 1] == total_tokens, + op_name, ": cu_seqlens[-1] must equal sequence length ", + total_tokens, ", but got ", cu[cu.size() - 1], "."); + for (size_t i = 0; i + 1 < cu.size(); ++i) { + TORCH_CHECK(cu[i] <= cu[i + 1], + op_name, ": cu_seqlens must be nondecreasing, but cu_seqlens[", + i, "]=", cu[i], " > cu_seqlens[", i + 1, "]=", cu[i + 1], "."); + } +} + +void check_kda_chunk_indices(const c10::optional &chunk_indices, + const c10::optional &cu_seqlens, + int64_t chunk_size, + const char *op_name) +{ + if (!chunk_indices.has_value()) { + return; + } + auto indices = chunk_indices.value(); + TORCH_CHECK(indices.size() % 2 == 0, + op_name, ": chunk_indices must contain (seq_id, chunk_id) pairs, but got ", + indices.size(), " elements."); + TORCH_CHECK(cu_seqlens.has_value(), op_name, ": chunk_indices requires cu_seqlens."); + auto cu = cu_seqlens.value(); + int64_t expected_chunks = 0; + for (size_t seq = 0; seq + 1 < cu.size(); ++seq) { + expected_chunks += kda_ceil_div(cu[seq + 1] - cu[seq], chunk_size); + } + TORCH_CHECK(static_cast(indices.size() / 2) == expected_chunks, + op_name, ": chunk_indices must contain exactly one pair per chunk."); + for (size_t idx = 0; idx < indices.size(); idx += 2) { + int64_t seq = indices[idx]; + int64_t chunk = indices[idx + 1]; + TORCH_CHECK(seq >= 0 && seq + 1 < static_cast(cu.size()), + op_name, ": chunk_indices seq_id is out of range."); + int64_t chunks = kda_ceil_div(cu[seq + 1] - cu[seq], chunk_size); + TORCH_CHECK(chunk >= 0 && chunk < chunks, + op_name, ": chunk_indices chunk_id is out of range."); + } +} + +int64_t get_kda_total_chunks(int64_t batch, + int64_t seqlen, + int64_t chunk_size, + const c10::optional &cu_seqlens, + const c10::optional &chunk_indices) +{ + if (chunk_indices.has_value()) { + return static_cast(chunk_indices.value().size()) / 2; + } + if (!cu_seqlens.has_value()) { + return kda_ceil_div(seqlen, chunk_size); + } + (void)batch; + int64_t total = 0; + auto cu = cu_seqlens.value(); + for (size_t i = 0; i + 1 < cu.size(); ++i) { + total += kda_ceil_div(cu[i + 1] - cu[i], chunk_size); + } + return total; +} + +std::vector build_kda_chunk_indices(at::IntArrayRef cu_seqlens, int64_t chunk_size) +{ + std::vector indices; + int64_t total_chunks = 0; + for (size_t i = 0; i + 1 < cu_seqlens.size(); ++i) { + total_chunks += kda_ceil_div(cu_seqlens[i + 1] - cu_seqlens[i], chunk_size); + } + indices.reserve(static_cast(total_chunks * 2)); + for (size_t seq = 0; seq + 1 < cu_seqlens.size(); ++seq) { + int64_t seq_len = cu_seqlens[seq + 1] - cu_seqlens[seq]; + int64_t chunks = kda_ceil_div(seq_len, chunk_size); + for (int64_t chunk = 0; chunk < chunks; ++chunk) { + indices.push_back(static_cast(seq)); + indices.push_back(chunk); + } + } + return indices; +} + +} // namespace + +std::tuple +chunk_kda_fwd( + const at::Tensor &q, + const at::Tensor &k, + const at::Tensor &v, + const at::Tensor &gk, + const at::Tensor &beta, + double scale, + int64_t chunk_size, + c10::string_view layout, + const c10::optional &initial_state, + c10::optional output_final_state, + c10::optional cu_seqlens, + c10::optional chunk_indices, + c10::optional return_intermediate, + c10::optional safe_gate, + c10::optional transpose_state_layout) +{ + std::string layout_str(layout.data(), layout.size()); + TORCH_CHECK(layout_str == "BSND" || layout_str == "BNSD" || layout_str == "TND" || layout_str == "NTD", + "chunk_kda_fwd: layout must be one of BSND, BNSD, TND, NTD and must be uppercase."); + TORCH_CHECK(!safe_gate.value_or(false), "chunk_kda_fwd: safe_gate=True is not supported."); + TORCH_CHECK(!transpose_state_layout.value_or(false), + "chunk_kda_fwd: transpose_state_layout=True is not supported."); + TORCH_CHECK(chunk_size == 32 || chunk_size == 64 || chunk_size == 128, + "chunk_kda_fwd: chunk_size must be 32, 64 or 128."); + + bool is_tnd = layout_str == "TND"; + bool is_ntd = layout_str == "NTD"; + bool is_bsnd = layout_str == "BSND"; + bool is_bnsd = layout_str == "BNSD"; + bool is_rank3 = is_tnd || is_ntd; + TORCH_CHECK((is_rank3 && q.dim() == 3 && k.dim() == 3 && v.dim() == 3 && gk.dim() == 3 && beta.dim() == 2) || + (!is_rank3 && q.dim() == 4 && k.dim() == 4 && v.dim() == 4 && gk.dim() == 4 && beta.dim() == 3), + "chunk_kda_fwd: layout/rank mismatch."); + TORCH_CHECK(q.sizes() == k.sizes(), "chunk_kda_fwd: q and k must have identical shape."); + + auto q_sizes = q.sizes(); + auto v_sizes = v.sizes(); + bool is_internal_layout = is_bnsd || is_ntd; + int64_t B = is_rank3 ? 1 : q_sizes[0]; + int64_t T = is_tnd ? q_sizes[0] : (is_ntd ? q_sizes[1] : (is_bnsd ? q_sizes[2] : q_sizes[1])); + int64_t H = is_tnd ? q_sizes[1] : (is_ntd ? q_sizes[0] : (is_bnsd ? q_sizes[1] : q_sizes[2])); + int64_t K = is_rank3 ? q_sizes[2] : q_sizes[3]; + int64_t HV = is_tnd ? v_sizes[1] : (is_ntd ? v_sizes[0] : (is_bnsd ? v_sizes[1] : v_sizes[2])); + int64_t V = is_rank3 ? v_sizes[2] : v_sizes[3]; + TORCH_CHECK(H > 0 && HV >= H, "chunk_kda_fwd: H and HV must be positive and H must be <= HV."); + TORCH_CHECK(H <= 128 && HV <= 128, "chunk_kda_fwd: H and HV must be <= 128."); + TORCH_CHECK(!is_tnd || H == 1, + "chunk_kda_fwd: TND layout with H > 1 is not supported; use NTD for multi-head rank3 input."); + check_kda_cu_seqlens(cu_seqlens, T, "chunk_kda_fwd"); + check_kda_chunk_indices(chunk_indices, cu_seqlens, chunk_size, "chunk_kda_fwd"); + TORCH_CHECK(!cu_seqlens.has_value() || is_rank3 || B == 1, + "chunk_kda_fwd: rank4 varlen input with cu_seqlens currently requires B=1."); + TORCH_CHECK(HV % H == 0, "chunk_kda_fwd: HV must be divisible by H."); + TORCH_CHECK(q.scalar_type() == at::kHalf || q.scalar_type() == at::kBFloat16, + "chunk_kda_fwd: q/k/v must use float16 or bfloat16."); + TORCH_CHECK(k.scalar_type() == q.scalar_type() && v.scalar_type() == q.scalar_type(), + "chunk_kda_fwd: q/k/v dtype must match."); + TORCH_CHECK(chunk_size == 64 && K >= 16 && V >= 16 && K % 16 == 0 && V % 16 == 0 && V <= 256 && + K * V >= 4 * 64 * 64 && K * V >= chunk_size * (K + V), + "chunk_kda_fwd: shape is outside the supported split cube/vector template."); + + int64_t seq_num = get_kda_seq_num(B, cu_seqlens); + at::Tensor initial_state_tensor = initial_state.value_or(at::Tensor()); + if (initial_state_tensor.defined()) { + TORCH_CHECK(initial_state_tensor.scalar_type() == at::kFloat, + "chunk_kda_fwd: initial_state must be float32 when provided."); + TORCH_CHECK(initial_state_tensor.dim() == 4 && initial_state_tensor.size(0) == seq_num && + initial_state_tensor.size(1) == HV && initial_state_tensor.size(2) == K && + initial_state_tensor.size(3) == V, + "chunk_kda_fwd: initial_state must be [seq_num,Hv,K,V]."); + } + + std::vector generated_chunk_indices; + c10::optional chunk_indices_for_call; + if (chunk_indices.has_value()) { + chunk_indices_for_call = chunk_indices.value(); + } else if (cu_seqlens.has_value()) { + generated_chunk_indices = build_kda_chunk_indices(cu_seqlens.value(), chunk_size); + chunk_indices_for_call = at::IntArrayRef(generated_chunk_indices); + } else { + chunk_indices_for_call = c10::nullopt; + } + + int64_t total_chunks = get_kda_total_chunks(B, T, chunk_size, cu_seqlens, chunk_indices_for_call); + at::Tensor o = at::empty_like(v); + at::Tensor final_state_work = at::empty({seq_num, HV, K, V}, q.options().dtype(at::kFloat)); + at::Tensor aqk = is_rank3 ? (is_internal_layout ? at::empty({HV, T, chunk_size}, q.options()) : + at::empty({T, HV, chunk_size}, q.options())) : (is_internal_layout ? + at::empty({B, HV, T, chunk_size}, q.options()) : at::empty({B, T, HV, chunk_size}, q.options())); + at::Tensor akk = at::empty_like(aqk); + at::Tensor w = is_rank3 ? (is_internal_layout ? at::empty({HV, T, K}, q.options()) : + at::empty({T, HV, K}, q.options())) : (is_internal_layout ? + at::empty({B, HV, T, K}, q.options()) : at::empty({B, T, HV, K}, q.options())); + at::Tensor u = at::empty_like(v); + at::Tensor qg = at::empty_like(w); + at::Tensor kg = at::empty_like(w); + at::Tensor v_new = at::empty_like(v); + at::Tensor h = is_rank3 ? (is_internal_layout ? at::empty({HV, total_chunks, K, V}, q.options()) : + at::empty({total_chunks, HV, K, V}, q.options())) : (is_internal_layout ? + at::empty({B, HV, total_chunks, K, V}, q.options()) : + at::empty({B, total_chunks, HV, K, V}, q.options())); + + bool recompute_output_final_state = true; + char *layout_cstr = const_cast(layout_str.c_str()); + EXEC_NPU_CMD( + aclnnChunkKdaFwd, + q, k, v, gk, beta, initial_state_tensor, + cu_seqlens, chunk_indices_for_call, + layout_cstr, scale, chunk_size, recompute_output_final_state, total_chunks, + o, final_state_work, aqk, akk, w, u, qg, kg, v_new, h + ); + + at::Tensor final_state = output_final_state.value_or(false) ? + final_state_work : at::empty({0}, q.options().dtype(at::kFloat)); + at::Tensor empty = at::empty({0}, q.options()); + at::Tensor g = gk.scalar_type() == at::kFloat ? gk : gk.to(at::kFloat); + at::Tensor initial_state_out = initial_state_tensor.defined() ? initial_state_tensor : empty; + (void)return_intermediate; + return std::make_tuple(o, final_state, g, aqk, akk, w, u, qg, kg, v_new, h, initial_state_out); +} + +at::Tensor kda_gate_cumsum( + const at::Tensor &g, + int64_t chunk_size, + const c10::optional &A_log, + const c10::optional &dt_bias, + c10::optional cu_seqlens, + c10::optional use_gate_in_kernel, + c10::optional safe_gate, + c10::optional lower_bound, + c10::string_view layout) +{ + TORCH_CHECK(g.dim() == 3 || g.dim() == 4, + "kda_gate_cumsum: g must be BSND/BNSD rank4 or TND/NTD rank3."); + TORCH_CHECK(chunk_size == 32 || chunk_size == 64 || chunk_size == 128, + "kda_gate_cumsum: chunk_size must be 32, 64 or 128."); + auto gate_dtype = g.scalar_type(); + TORCH_CHECK(gate_dtype == at::kFloat || gate_dtype == at::kBFloat16 || gate_dtype == at::kHalf, + "kda_gate_cumsum: g must be float32, bfloat16 or float16."); + std::string layout_str(layout.data(), layout.size()); + TORCH_CHECK(layout_str == "BSND" || layout_str == "BNSD" || layout_str == "TND" || layout_str == "NTD", + "kda_gate_cumsum: layout must be uppercase and one of BSND, BNSD, TND or NTD."); + bool is_bnsd = layout_str == "BNSD"; + bool is_ntd = layout_str == "NTD"; + int64_t T = is_bnsd ? g.sizes()[2] : (is_ntd ? g.sizes()[1] : (g.dim() == 4 ? g.sizes()[1] : g.sizes()[0])); + int64_t K = g.dim() == 4 ? g.sizes()[3] : g.sizes()[2]; + int64_t HV = is_bnsd ? g.sizes()[1] : (is_ntd ? g.sizes()[0] : (g.dim() == 4 ? g.sizes()[2] : g.sizes()[1])); + TORCH_CHECK(K <= 256, "kda_gate_cumsum: K must be <= 256."); + check_kda_cu_seqlens(cu_seqlens, T, "kda_gate_cumsum"); + TORCH_CHECK(!cu_seqlens.has_value() || g.dim() == 3 || g.sizes()[0] == 1, + "kda_gate_cumsum: rank4 varlen input with cu_seqlens currently requires B=1."); + + bool use_gate = use_gate_in_kernel.value_or(false); + bool safe = safe_gate.value_or(false); + double lower = lower_bound.value_or(-5.0); + at::Tensor A_log_tensor = A_log.value_or(at::Tensor()); + at::Tensor dt_bias_tensor = dt_bias.value_or(at::Tensor()); + if (use_gate) { + TORCH_CHECK(A_log_tensor.defined(), "kda_gate_cumsum: A_log is required when use_gate_in_kernel=True."); + TORCH_CHECK(A_log_tensor.scalar_type() == at::kFloat && A_log_tensor.dim() == 1 && A_log_tensor.sizes()[0] == HV, + "kda_gate_cumsum: A_log must be float32 with shape [HV]."); + TORCH_CHECK(safe, "kda_gate_cumsum: raw gate path currently requires safe_gate=True."); + TORCH_CHECK(lower >= -5.0 && lower < 0.0, "kda_gate_cumsum: lower_bound must be in [-5, 0)."); + } else { + TORCH_CHECK(!safe, "kda_gate_cumsum: safe_gate only applies when use_gate_in_kernel=True."); + } + + at::Tensor gk = at::empty(g.sizes(), g.options().dtype(at::kFloat)); + char *layout_cstr = const_cast(layout_str.c_str()); + EXEC_NPU_CMD( + aclnnKdaGateCumsum, + g, A_log_tensor, dt_bias_tensor, cu_seqlens, + chunk_size, use_gate, safe, lower, layout_cstr, gk + ); + return gk; +} + +at::Tensor kda_layout_swap12( + const at::Tensor &x, + const c10::optional &dependency) +{ + TORCH_CHECK(x.dim() >= 3, "kda_layout_swap12: x must have rank >= 3."); + auto dtype = x.scalar_type(); + TORCH_CHECK(dtype == at::kFloat || dtype == at::kHalf || dtype == at::kBFloat16, + "kda_layout_swap12: x must be float32, float16 or bfloat16."); + + std::vector y_sizes(x.sizes().begin(), x.sizes().end()); + if (x.dim() == 3) { + std::swap(y_sizes[0], y_sizes[1]); + } else { + std::swap(y_sizes[1], y_sizes[2]); + } + at::Tensor y = at::empty(y_sizes, x.options()); + at::Tensor dependency_tensor = dependency.value_or(at::Tensor()); + if (dependency_tensor.defined()) { + TORCH_CHECK(dependency_tensor.sizes() == y.sizes(), + "kda_layout_swap12: dependency must have the same shape as output."); + } + + EXEC_NPU_CMD(aclnnKdaLayoutSwap12, x, dependency_tensor, y); + return y; +} + std::vector get_npu_storage_shape(const at::Tensor& tensor) { TORCH_CHECK( @@ -1974,6 +2291,21 @@ TORCH_LIBRARY_EXPAND(CONCAT(_C, _ascend), ops) "chunk_fwd_o(Tensor q, Tensor k, Tensor v, Tensor h, float scale, *, Tensor? g=None, Tensor? g_gamma=None, int[]? cu_seqlens=None, int[]? chunk_indices=None, int? chunk_size=None, bool? transpose_state_layout=False) -> Tensor" ); ops.impl("chunk_fwd_o", torch::kPrivateUse1, &vllm_ascend::chunk_fwd_o); + + ops.def( + "chunk_kda_fwd(Tensor q, Tensor k, Tensor v, Tensor gk, Tensor beta, float scale, int chunk_size, str layout=\"BSND\", *, Tensor? initial_state=None, bool? output_final_state=False, int[]? cu_seqlens=None, int[]? chunk_indices=None, bool? return_intermediate=False, bool? safe_gate=False, bool? transpose_state_layout=False) -> (Tensor o, Tensor final_state, Tensor g, Tensor aqk, Tensor akk, Tensor w, Tensor u, Tensor qg, Tensor kg, Tensor v_new, Tensor h, Tensor initial_state_out)" + ); + ops.impl("chunk_kda_fwd", torch::kPrivateUse1, &vllm_ascend::chunk_kda_fwd); + + ops.def( + "kda_gate_cumsum(Tensor g, int chunk_size, *, Tensor? A_log=None, Tensor? dt_bias=None, int[]? cu_seqlens=None, bool? use_gate_in_kernel=False, bool? safe_gate=False, float? lower_bound=-5.0, str layout=\"BSND\") -> Tensor" + ); + ops.impl("kda_gate_cumsum", torch::kPrivateUse1, &vllm_ascend::kda_gate_cumsum); + + ops.def( + "kda_layout_swap12(Tensor x, *, Tensor? dependency=None) -> Tensor" + ); + ops.impl("kda_layout_swap12", torch::kPrivateUse1, &vllm_ascend::kda_layout_swap12); } #else // Pybind on other platform @@ -2005,6 +2337,37 @@ TORCH_LIBRARY_EXPAND(CONCAT(_C, _ascend), ops) " Tensor? gk=None) -> Tensor"); ops.impl("npu_recurrent_gated_delta_rule", torch::kPrivateUse1, &vllm_ascend::npu_recurrent_gated_delta_rule); + ops.def( + "recurrent_kda(Tensor query, Tensor key, Tensor value, Tensor gate, Tensor beta, " + "Tensor(a!) initial_state, Tensor cu_seqlens, Tensor ssm_state_indices, Tensor A_log, Tensor dt_bias, *, " + "Tensor? num_accepted_tokens=None, float scale=0.08838834764831845, " + "bool use_qk_l2norm_in_kernel=True, bool use_gate_in_kernel=True, " + "bool use_beta_sigmoid_in_kernel=False, bool allow_neg_eigval=False, " + "bool safe_gate=True, float lower_bound=-5.0) -> Tensor output"); + ops.impl("recurrent_kda", torch::kPrivateUse1, &vllm_ascend::recurrent_kda); + + ops.def( + "dequant_situ_quant(Tensor x, " + " *, Tensor? weight_scale=None, " + " Tensor? activation_scale=None, " + " Tensor? bias=None, " + " Tensor? quant_scale=None, " + " Tensor? quant_offset=None, " + " Tensor? group_index=None, " + " float beta=4.0, " + " float linear_beta=25.0, " + " bool activate_left=True, " + " str quant_mode=\"dynamic\") -> (Tensor y, Tensor scale)"); + ops.impl("dequant_situ_quant", torch::kPrivateUse1, &vllm_ascend::dequant_situ_quant); + + ops.def( + "situ_mx_quant(Tensor x, " + " float beta=1.0, " + " float linear_beta=0.0, " + " bool activate_left=False, " + " int dst_type=36) -> (Tensor y, Tensor mxscale)"); + ops.impl("situ_mx_quant", torch::kPrivateUse1, &vllm_ascend::situ_mx_quant); + #ifdef VLLM_ENABLE_ATB_AND_DIRECT_KERNELS // Direct kernel custom ops ops.def("bgmv_shrink(Tensor! x, Tensor! weight, Tensor! indices, Tensor! y, float scale) -> ()"); @@ -2580,6 +2943,21 @@ TORCH_LIBRARY_EXPAND(CONCAT(_C, _ascend), ops) ); ops.impl("chunk_fwd_o", torch::kPrivateUse1, &vllm_ascend::chunk_fwd_o); + ops.def( + "chunk_kda_fwd(Tensor q, Tensor k, Tensor v, Tensor gk, Tensor beta, float scale, int chunk_size, str layout=\"BSND\", *, Tensor? initial_state=None, bool? output_final_state=False, int[]? cu_seqlens=None, int[]? chunk_indices=None, bool? return_intermediate=False, bool? safe_gate=False, bool? transpose_state_layout=False) -> (Tensor o, Tensor final_state, Tensor g, Tensor aqk, Tensor akk, Tensor w, Tensor u, Tensor qg, Tensor kg, Tensor v_new, Tensor h, Tensor initial_state_out)" + ); + ops.impl("chunk_kda_fwd", torch::kPrivateUse1, &vllm_ascend::chunk_kda_fwd); + + ops.def( + "kda_gate_cumsum(Tensor g, int chunk_size, *, Tensor? A_log=None, Tensor? dt_bias=None, int[]? cu_seqlens=None, bool? use_gate_in_kernel=False, bool? safe_gate=False, float? lower_bound=-5.0, str layout=\"BSND\") -> Tensor" + ); + ops.impl("kda_gate_cumsum", torch::kPrivateUse1, &vllm_ascend::kda_gate_cumsum); + + ops.def( + "kda_layout_swap12(Tensor x, *, Tensor? dependency=None) -> Tensor" + ); + ops.impl("kda_layout_swap12", torch::kPrivateUse1, &vllm_ascend::kda_layout_swap12); + //store_kv_block ops.def( "store_kv_block_metadata(Tensor slot_mapping_npu, Tensor group_len, Tensor group_key_idx, Tensor group_key_cache_idx, int block_size=0)" diff --git a/csrc/torch_binding_meta.cpp b/csrc/torch_binding_meta.cpp index f5d5007e75c5..6c6b88c025dd 100644 --- a/csrc/torch_binding_meta.cpp +++ b/csrc/torch_binding_meta.cpp @@ -640,6 +640,46 @@ at::Tensor npu_recurrent_gated_delta_rule_meta( return output; } +at::Tensor recurrent_kda_meta( + const at::Tensor& query, + const at::Tensor& key, + const at::Tensor& value, + const at::Tensor& gate, + const at::Tensor& beta, + at::Tensor& initial_state, + const at::Tensor& actual_seq_lengths, + const at::Tensor& ssm_state_indices, + const at::Tensor& a_log, + const at::Tensor& dt_bias, + const c10::optional& num_accepted_tokens, + double scale, + bool use_qk_l2norm_in_kernel, + bool use_gate_in_kernel, + bool use_beta_sigmoid_in_kernel, + bool allow_neg_eigval, + bool safe_gate, + double lower_bound) +{ + (void)query; + (void)key; + (void)gate; + (void)beta; + (void)actual_seq_lengths; + (void)ssm_state_indices; + (void)a_log; + (void)dt_bias; + (void)num_accepted_tokens; + (void)scale; + (void)use_qk_l2norm_in_kernel; + (void)use_gate_in_kernel; + (void)use_beta_sigmoid_in_kernel; + (void)allow_neg_eigval; + (void)safe_gate; + (void)lower_bound; + (void)initial_state; + return at::empty_symint(value.sym_sizes(), value.options()); +} + std::vector moe_grouped_matmul_meta( at::Tensor x, at::Tensor weight, @@ -1464,6 +1504,143 @@ at::Tensor chunk_fwd_o_meta( return o; } +std::tuple +chunk_kda_fwd_meta( + const at::Tensor &q, + const at::Tensor &k, + const at::Tensor &v, + const at::Tensor &gk, + const at::Tensor &beta, + double scale, + int64_t chunk_size, + c10::string_view layout, + const c10::optional &initial_state, + c10::optional output_final_state, + c10::optional cu_seqlens, + c10::optional chunk_indices, + c10::optional return_intermediate, + c10::optional safe_gate, + c10::optional transpose_state_layout) +{ + std::string layout_str = std::string(layout); + bool is_tnd = layout_str == "TND"; + bool is_ntd = layout_str == "NTD"; + bool is_bnsd = layout_str == "BNSD"; + bool is_rank3 = is_tnd || is_ntd; + bool is_internal_layout = is_bnsd || is_ntd; + + c10::SymInt B = is_rank3 ? c10::SymInt(1) : q.sym_size(0); + c10::SymInt T = is_tnd ? q.sym_size(0) : + (is_ntd ? q.sym_size(1) : (is_bnsd ? q.sym_size(2) : q.sym_size(1))); + c10::SymInt K = is_rank3 ? q.sym_size(2) : q.sym_size(3); + c10::SymInt HV = is_tnd ? v.sym_size(1) : + (is_ntd ? v.sym_size(0) : (is_bnsd ? v.sym_size(1) : v.sym_size(2))); + c10::SymInt V = is_rank3 ? v.sym_size(2) : v.sym_size(3); + // symbolic-meta-ok: cu_seqlens is an IntArrayRef schema argument, not a Tensor shape. + c10::SymInt seq_num = cu_seqlens.has_value() ? + c10::SymInt(static_cast(cu_seqlens->size()) - 1) : B; + c10::SymInt total_chunks(0); + if (chunk_indices.has_value()) { + // symbolic-meta-ok: chunk_indices is an IntArrayRef schema argument, not a Tensor shape. + total_chunks = c10::SymInt(static_cast(chunk_indices->size()) / 2); + } else if (cu_seqlens.has_value()) { + int64_t concrete_total_chunks = 0; + // symbolic-meta-ok: cu_seqlens is an IntArrayRef schema argument, not a Tensor shape. + for (size_t i = 0; i + 1 < cu_seqlens->size(); ++i) { + concrete_total_chunks += ((*cu_seqlens)[i + 1] - (*cu_seqlens)[i] + chunk_size - 1) / chunk_size; + } + total_chunks = c10::SymInt(concrete_total_chunks); + } else { + total_chunks = (T + c10::SymInt(chunk_size - 1)) / c10::SymInt(chunk_size); + } + + at::Tensor o = at::empty_like(v); + at::Tensor final_state_work = at::empty_symint( + c10::SymDimVector{seq_num, HV, K, V}, q.options().dtype(at::kFloat)); + at::Tensor final_state = output_final_state.value_or(false) ? + final_state_work : at::empty_symint(c10::SymDimVector{c10::SymInt(0)}, q.options().dtype(at::kFloat)); + at::Tensor g = gk.scalar_type() == at::kFloat ? + gk : at::empty_symint(gk.sym_sizes(), gk.options().dtype(at::kFloat)); + c10::SymInt chunk_size_sym(chunk_size); + c10::SymDimVector aqk_shape; + if (is_rank3) { + aqk_shape = is_internal_layout ? c10::SymDimVector{HV, T, chunk_size_sym} : + c10::SymDimVector{T, HV, chunk_size_sym}; + } else { + aqk_shape = is_internal_layout ? c10::SymDimVector{B, HV, T, chunk_size_sym} : + c10::SymDimVector{B, T, HV, chunk_size_sym}; + } + at::Tensor aqk = at::empty_symint(aqk_shape, q.options()); + at::Tensor akk = at::empty_like(aqk); + c10::SymDimVector w_shape; + if (is_rank3) { + w_shape = is_internal_layout ? c10::SymDimVector{HV, T, K} : c10::SymDimVector{T, HV, K}; + } else { + w_shape = is_internal_layout ? c10::SymDimVector{B, HV, T, K} : c10::SymDimVector{B, T, HV, K}; + } + at::Tensor w = at::empty_symint(w_shape, q.options()); + at::Tensor u = at::empty_like(v); + at::Tensor qg = at::empty_like(w); + at::Tensor kg = at::empty_like(w); + at::Tensor v_new = at::empty_like(v); + c10::SymDimVector h_shape; + if (is_rank3) { + h_shape = is_internal_layout ? c10::SymDimVector{HV, total_chunks, K, V} : + c10::SymDimVector{total_chunks, HV, K, V}; + } else { + h_shape = is_internal_layout ? c10::SymDimVector{B, HV, total_chunks, K, V} : + c10::SymDimVector{B, total_chunks, HV, K, V}; + } + at::Tensor h = at::empty_symint(h_shape, q.options()); + at::Tensor initial_state_tensor = initial_state.value_or(at::Tensor()); + at::Tensor initial_state_out = initial_state_tensor.defined() ? + initial_state_tensor : at::empty_symint(c10::SymDimVector{c10::SymInt(0)}, q.options()); + (void)k; + (void)beta; + (void)scale; + (void)return_intermediate; + (void)safe_gate; + (void)transpose_state_layout; + return std::make_tuple(o, final_state, g, aqk, akk, w, u, qg, kg, v_new, h, initial_state_out); +} + +at::Tensor kda_gate_cumsum_meta( + const at::Tensor &g, + int64_t chunk_size, + const c10::optional &A_log, + const c10::optional &dt_bias, + c10::optional cu_seqlens, + c10::optional use_gate_in_kernel, + c10::optional safe_gate, + c10::optional lower_bound, + c10::string_view layout) +{ + (void)chunk_size; + (void)A_log; + (void)dt_bias; + (void)cu_seqlens; + (void)use_gate_in_kernel; + (void)safe_gate; + (void)lower_bound; + (void)layout; + return at::empty_symint(g.sym_sizes(), g.options().dtype(at::kFloat)); +} + +at::Tensor kda_layout_swap12_meta( + const at::Tensor &x, + const c10::optional &dependency) +{ + c10::SymDimVector y_sizes(x.sym_sizes().begin(), x.sym_sizes().end()); + if (x.dim() == 3) { + std::swap(y_sizes[0], y_sizes[1]); + } else { + std::swap(y_sizes[1], y_sizes[2]); + } + (void)dependency; + return at::empty_symint(y_sizes, x.options()); +} + void store_kv_block_metadata( const at::Tensor &slot_mapping_npu, const at::Tensor &group_len, @@ -1485,6 +1662,85 @@ void store_kv_block( return; } +std::tuple dequant_situ_quant_meta( + const at::Tensor& x, + const c10::optional& weight_scale, + const c10::optional& activation_scale, + const c10::optional& bias, + const c10::optional& quant_scale, + const c10::optional& quant_offset, + const c10::optional& group_index, + double beta, + double linear_beta, + bool activate_left, + c10::string_view quant_mode) +{ + (void)weight_scale; + (void)activation_scale; + (void)bias; + (void)quant_scale; + (void)quant_offset; + (void)group_index; + (void)beta; + (void)linear_beta; + (void)activate_left; + (void)quant_mode; + + TORCH_CHECK(x.dim() == 2, + "dequant_situ_quant: x must be 2-dimensional [rows, width], but got rank ", + x.dim()); + TORCH_CHECK(x.scalar_type() == at::kInt || x.scalar_type() == at::kBFloat16, + "dequant_situ_quant: x must be int32 or bfloat16, but got ", x.scalar_type()); + const c10::SymInt input_width = x.sym_size(1); + TORCH_CHECK(input_width % 2 == 0, + "dequant_situ_quant: x last dimension must be even"); + + c10::SymDimVector y_shape(x.sym_sizes().begin(), x.sym_sizes().end()); + y_shape.back() = input_width / 2; + c10::SymDimVector scale_shape; + scale_shape.push_back(x.sym_size(0)); + at::Tensor y = at::empty_symint(y_shape, x.options().dtype(at::kChar)); + at::Tensor scale = at::empty_symint(scale_shape, x.options().dtype(at::kFloat)); + return {y, scale}; +} + +std::tuple situ_mx_quant_meta( + const at::Tensor& x, + double beta, + double linear_beta, + bool activate_left, + int64_t dst_type) +{ + constexpr int64_t DST_TYPE_E5M2 = 35; + constexpr int64_t DST_TYPE_E4M3FN = 36; + constexpr int64_t MX_BLOCK_SPAN = 64; + constexpr int64_t MX_SCALE_ALIGN = 2; + + TORCH_CHECK(x.dim() >= 1, + "situ_mx_quant: x must be at least 1-dimensional, but got ", + x.dim()); + TORCH_CHECK(x.scalar_type() == at::kBFloat16, + "situ_mx_quant: x must be bfloat16, but got ", x.scalar_type()); + TORCH_CHECK(beta > 0.0, + "situ_mx_quant: beta must be greater than 0, but got ", beta); + TORCH_CHECK(dst_type == DST_TYPE_E4M3FN || dst_type == DST_TYPE_E5M2, + "situ_mx_quant: dst_type must be 36 (E4M3FN) or 35 (E5M2), but got ", + dst_type); + + (void)linear_beta; + (void)activate_left; + + c10::SymDimVector y_shape(x.sym_sizes().begin(), x.sym_sizes().end()); + y_shape.back() = y_shape.back() / 2; + c10::SymDimVector mxscale_shape(y_shape.begin(), y_shape.end()); + mxscale_shape.back() = (mxscale_shape.back() + MX_BLOCK_SPAN - 1) / MX_BLOCK_SPAN; + mxscale_shape.emplace_back(MX_SCALE_ALIGN); + + auto y_dtype = dst_type == DST_TYPE_E5M2 ? at::kFloat8_e5m2 : at::kFloat8_e4m3fn; + at::Tensor y = at::empty_symint(y_shape, x.options().dtype(y_dtype)); + at::Tensor mxscale = at::empty_symint(mxscale_shape, x.options().dtype(at::kFloat8_e8m0fnu)); + return {y, mxscale}; +} } // namespace meta } // namespace vllm_ascend @@ -1503,6 +1759,12 @@ TORCH_LIBRARY_IMPL_EXPAND(CONCAT(_C, _ascend), Meta, ops) { ops.impl("chunk_gated_delta_rule_fwd_h", &vllm_ascend::meta::chunk_gated_delta_rule_fwd_h_meta); // chunk_fwd_o ops.impl("chunk_fwd_o", &vllm_ascend::meta::chunk_fwd_o_meta); + // chunk_kda_fwd + ops.impl("chunk_kda_fwd", &vllm_ascend::meta::chunk_kda_fwd_meta); + // kda_gate_cumsum + ops.impl("kda_gate_cumsum", &vllm_ascend::meta::kda_gate_cumsum_meta); + // kda_layout_swap12 + ops.impl("kda_layout_swap12", &vllm_ascend::meta::kda_layout_swap12_meta); } } #else @@ -1513,6 +1775,9 @@ TORCH_LIBRARY_IMPL_EXPAND(CONCAT(_C, _ascend), Meta, ops) { ops.impl("npu_gemma_rms_norm", &vllm_ascend::meta::npu_gemma_rms_norm_meta); // recurrent_gated_delta_rule meta implementation ops.impl("npu_recurrent_gated_delta_rule", &vllm_ascend::meta::npu_recurrent_gated_delta_rule_meta); + ops.impl("recurrent_kda", &vllm_ascend::meta::recurrent_kda_meta); + ops.impl("dequant_situ_quant", &vllm_ascend::meta::dequant_situ_quant_meta); + ops.impl("situ_mx_quant", &vllm_ascend::meta::situ_mx_quant_meta); // Launch host print from device ops.impl("device_print", &vllm_ascend::meta::device_print_meta); // launch host print from device for tensors @@ -1588,6 +1853,12 @@ TORCH_LIBRARY_IMPL_EXPAND(CONCAT(_C, _ascend), Meta, ops) { ops.impl("chunk_gated_delta_rule_fwd_h", &vllm_ascend::meta::chunk_gated_delta_rule_fwd_h_meta); // chunk_fwd_o ops.impl("chunk_fwd_o", &vllm_ascend::meta::chunk_fwd_o_meta); + // chunk_kda_fwd + ops.impl("chunk_kda_fwd", &vllm_ascend::meta::chunk_kda_fwd_meta); + // kda_gate_cumsum + ops.impl("kda_gate_cumsum", &vllm_ascend::meta::kda_gate_cumsum_meta); + // kda_layout_swap12 + ops.impl("kda_layout_swap12", &vllm_ascend::meta::kda_layout_swap12_meta); // store_kv_block ops.impl("store_kv_block_pre", &vllm_ascend::meta::store_kv_block_metadata); ops.impl("store_kv_block", &vllm_ascend::meta::store_kv_block); diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_chunk_gated_delta_rule_fwd_h_aclnn.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_chunk_gated_delta_rule_fwd_h_aclnn.py new file mode 100644 index 000000000000..22187dd85ab6 --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_chunk_gated_delta_rule_fwd_h_aclnn.py @@ -0,0 +1,334 @@ +# +# Copyright (c) 2026 Huawei Technologies Co., Ltd. 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. +# 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. +# This file is a part of the vllm-ascend project. + +import gc +import math + +import pytest +import torch +import torch_npu + +from vllm_ascend.utils import enable_custom_op + +torch_npu.npu.config.allow_internal_format = True +enable_custom_op() + +CHUNK_SIZE = 64 +DETERMINISM_REPEATS = 20 +FWD_H_OUTPUT_NAMES = ("h", "v_new", "final_state") + + +def _cleanup_npu(): + gc.collect() + torch.npu.empty_cache() + torch.npu.reset_peak_memory_stats() + + +def _prepare_chunk_offsets(cu_seqlens, chunk_size): + if cu_seqlens is None: + return None + + num_chunks = 0 + for seq_start, seq_end in zip(cu_seqlens[:-1], cu_seqlens[1:]): + num_chunks += math.ceil((seq_end - seq_start) / chunk_size) + return [0, 0] * num_chunks + + +def _make_cumulative_gate(shape_batch, v_num_head, seqlen, chunk_size, cu_seqlens): + torch.manual_seed(20260720 + seqlen + v_num_head) + + g = -torch.rand(shape_batch, v_num_head, seqlen, dtype=torch.float32) * 0.05 + if cu_seqlens is None: + cu_seqlens = [0, seqlen] + + for batch_idx in range(shape_batch): + for head_idx in range(v_num_head): + for seq_start, seq_end in zip(cu_seqlens[:-1], cu_seqlens[1:]): + for chunk_start in range(seq_start, seq_end, chunk_size): + chunk_end = min(chunk_start + chunk_size, seq_end) + g[batch_idx, head_idx, chunk_start:chunk_end] = torch.cumsum( + g[batch_idx, head_idx, chunk_start:chunk_end], + dim=0, + ) + return g + + +def _make_inputs(batch, seqlen, k_num_head, v_num_head, k_dim, v_dim, dtype, is_varlen): + torch.manual_seed(20260720 + batch + seqlen + k_num_head + v_num_head + v_dim) + + shape_batch = 1 if is_varlen else batch + cu_seqlens = [0, seqlen // 3, seqlen] if is_varlen else None + k = torch.randn(shape_batch, k_num_head, seqlen, k_dim, dtype=dtype) * 0.04 + w = torch.randn(shape_batch, v_num_head, seqlen, k_dim, dtype=dtype) * 0.04 + u = torch.randn(shape_batch, v_num_head, seqlen, v_dim, dtype=dtype) * 0.04 + g = _make_cumulative_gate(shape_batch, v_num_head, seqlen, CHUNK_SIZE, cu_seqlens) + chunk_indices = _prepare_chunk_offsets(cu_seqlens, CHUNK_SIZE) + return k, w, u, g, cu_seqlens, chunk_indices + + +def _chunk_gated_delta_rule_fwd_h_reference(k, w, u, g, chunk_size, cu_seqlens): + dtype = k.dtype + k = k.float() + w = w.float() + u = u.float() + g = g.float() + + shape_batch, k_num_head, seqlen, k_dim = k.shape + v_num_head, v_dim = u.shape[1], u.shape[3] + head_ratio = v_num_head // k_num_head + if cu_seqlens is None: + cu_seqlens = [0, seqlen] + num_sequences = shape_batch + num_chunks = (seqlen + chunk_size - 1) // chunk_size + else: + num_sequences = len(cu_seqlens) - 1 + num_chunks = sum( + math.ceil((seq_end - seq_start) / chunk_size) for seq_start, seq_end in zip(cu_seqlens[:-1], cu_seqlens[1:]) + ) + + h = torch.zeros(shape_batch, v_num_head, num_chunks, k_dim, v_dim, dtype=torch.float32) + v_new = torch.zeros(shape_batch, v_num_head, seqlen, v_dim, dtype=torch.float32) + + for seq_idx in range(num_sequences): + shape_batch_idx = 0 if len(cu_seqlens) > 2 else seq_idx + seq_start = 0 if len(cu_seqlens) == 2 and shape_batch > 1 else cu_seqlens[seq_idx] + seq_end = seqlen if len(cu_seqlens) == 2 and shape_batch > 1 else cu_seqlens[seq_idx + 1] + chunk_base = 0 + if len(cu_seqlens) > 2: + chunk_base = sum( + math.ceil((end - start) / chunk_size) + for start, end in zip(cu_seqlens[:seq_idx], cu_seqlens[1 : seq_idx + 1]) + ) + seq_chunks = math.ceil((seq_end - seq_start) / chunk_size) + + for v_head_idx in range(v_num_head): + k_head_idx = v_head_idx // head_ratio + for chunk_idx in range(seq_chunks): + token_start = seq_start + chunk_idx * chunk_size + actual_len = min(chunk_size, seq_end - token_start) + h_idx = chunk_base + chunk_idx + + k_sel = torch.zeros(chunk_size, k_dim, dtype=torch.float32) + w_sel = torch.zeros(chunk_size, k_dim, dtype=torch.float32) + u_sel = torch.zeros(chunk_size, v_dim, dtype=torch.float32) + g_sel = torch.zeros(chunk_size, dtype=torch.float32) + token_slice = slice(token_start, token_start + actual_len) + k_sel[:actual_len] = k[shape_batch_idx, k_head_idx, token_slice] + w_sel[:actual_len] = w[shape_batch_idx, v_head_idx, token_slice] + u_sel[:actual_len] = u[shape_batch_idx, v_head_idx, token_slice] + g_sel[:actual_len] = g[shape_batch_idx, v_head_idx, token_slice] + + current_h = h[shape_batch_idx, v_head_idx, h_idx] + v_work = u_sel - w_sel @ current_h + if chunk_idx != seq_chunks - 1: + gate = (g_sel[actual_len - 1] - g_sel).exp().unsqueeze(-1) + h[shape_batch_idx, v_head_idx, h_idx + 1] = current_h * g_sel[ + actual_len - 1 + ].exp() + k_sel.transpose(-1, -2) @ (v_work * gate) + v_new[shape_batch_idx, v_head_idx, token_slice] = v_work[:actual_len] + + return h.to(dtype), v_new.to(dtype) + + +def _chunk_gated_delta_rule_fwd_h_kda_reference(k, w, u, gk, initial_state, chunk_size): + dtype = k.dtype + k = k.float() + w = w.float() + u = u.float() + gk = gk.float() + state = initial_state.float().clone() + + batch, k_num_head, seqlen, k_dim = k.shape + v_num_head, v_dim = u.shape[1], u.shape[3] + head_ratio = v_num_head // k_num_head + num_chunks = (seqlen + chunk_size - 1) // chunk_size + h = torch.zeros(batch, v_num_head, num_chunks, k_dim, v_dim, dtype=torch.float32) + v_new = torch.zeros(batch, v_num_head, seqlen, v_dim, dtype=torch.float32) + + for batch_idx in range(batch): + for v_head_idx in range(v_num_head): + k_head_idx = v_head_idx // head_ratio + for chunk_idx, token_start in enumerate(range(0, seqlen, chunk_size)): + token_end = min(token_start + chunk_size, seqlen) + token_slice = slice(token_start, token_end) + current_h = state[batch_idx, v_head_idx] + h[batch_idx, v_head_idx, chunk_idx] = current_h + v_work = u[batch_idx, v_head_idx, token_slice] - w[batch_idx, v_head_idx, token_slice] @ current_h + v_new[batch_idx, v_head_idx, token_slice] = v_work + last_gk = gk[batch_idx, v_head_idx, token_end - 1] + state[batch_idx, v_head_idx] = ( + torch.exp2(last_gk)[:, None] * current_h + + k[batch_idx, k_head_idx, token_slice].transpose(-1, -2) @ v_work + ) + + return h.to(dtype), v_new.to(dtype), state.to(initial_state.dtype) + + +def _assert_cosine_close(name, actual, expected, threshold=0.99): + actual = actual.detach().cpu().float().flatten() + expected = expected.detach().cpu().float().flatten() + if actual.norm() == 0 and expected.norm() == 0: + return + + cosine = torch.nn.functional.cosine_similarity(actual.unsqueeze(0), expected.unsqueeze(0)).item() + assert cosine >= threshold, f"{name} cosine={cosine:.6f}, expected >= {threshold}" + + +def _snapshot_outputs(outputs): + torch.npu.synchronize() + return tuple(output.detach().cpu().contiguous() for output in outputs) + + +def _assert_outputs_bitwise_equal(reference, actual, repeat): + assert len(reference) == len(actual) == len(FWD_H_OUTPUT_NAMES) + for name, expected, current in zip(FWD_H_OUTPUT_NAMES, reference, actual): + same_metadata = expected.shape == current.shape and expected.dtype == current.dtype + same_bits = same_metadata and torch.equal(expected.view(torch.uint8), current.view(torch.uint8)) + if same_bits: + continue + + expected_float = expected.float() + current_float = current.float() + finite = torch.isfinite(expected_float) & torch.isfinite(current_float) + max_abs_diff = ( + (expected_float[finite] - current_float[finite]).abs().max().item() if finite.any() else float("nan") + ) + expected_nonfinite = (~torch.isfinite(expected_float)).sum().item() + current_nonfinite = (~torch.isfinite(current_float)).sum().item() + raise AssertionError( + f"repeat={repeat} output={name} is not bitwise deterministic: " + f"expected_nonfinite={expected_nonfinite}, current_nonfinite={current_nonfinite}, " + f"max_abs_diff={max_abs_diff}" + ) + + +@pytest.mark.parametrize( + ("batch", "seqlen", "k_num_head", "v_num_head", "k_dim", "v_dim", "dtype", "is_varlen"), + [ + (1, 128, 1, 1, 128, 128, torch.float16, False), + (1, 128, 1, 2, 128, 256, torch.float16, False), + (1, 128, 2, 2, 128, 256, torch.bfloat16, False), + (2, 96, 1, 2, 128, 256, torch.float16, True), + ], +) +@torch.inference_mode() +def test_chunk_gated_delta_rule_fwd_h_matches_reference( + batch, + seqlen, + k_num_head, + v_num_head, + k_dim, + v_dim, + dtype, + is_varlen, +): + k, w, u, g, cu_seqlens, chunk_indices = _make_inputs( + batch, + seqlen, + k_num_head, + v_num_head, + k_dim, + v_dim, + dtype, + is_varlen, + ) + expected_h, expected_v = _chunk_gated_delta_rule_fwd_h_reference( + k, + w, + u, + g, + CHUNK_SIZE, + cu_seqlens, + ) + + h_out, v_new, final_state = torch.ops._C_ascend.chunk_gated_delta_rule_fwd_h( + k.npu(), + w.npu(), + u.npu(), + g=g.npu(), + output_final_state=False, + chunk_size=CHUNK_SIZE, + save_new_value=True, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + ) + + assert final_state is None + assert torch.isfinite(h_out).all().item() + assert torch.isfinite(v_new).all().item() + _assert_cosine_close("h_out", h_out, expected_h) + _assert_cosine_close("v_new", v_new, expected_v) + _cleanup_npu() + + +@torch.inference_mode() +def test_chunk_gated_delta_rule_fwd_h_kda_is_bitwise_deterministic(): + batch, seqlen, k_num_head, v_num_head, k_dim, v_dim = 1, 128, 2, 2, 128, 256 + dtype = torch.bfloat16 + torch.manual_seed(20260720 + seqlen + k_num_head + v_num_head + v_dim) + + k = torch.randn(batch, k_num_head, seqlen, k_dim, dtype=dtype) * 0.04 + w = torch.randn(batch, v_num_head, seqlen, k_dim, dtype=dtype) * 0.04 + u = torch.randn(batch, v_num_head, seqlen, v_dim, dtype=dtype) * 0.04 + raw_gate = -torch.rand(batch, v_num_head, seqlen, k_dim, dtype=torch.float32) * 0.04 + gk = torch.empty_like(raw_gate) + rcp_ln2 = 1.4426950408889634 + for token_start in range(0, seqlen, CHUNK_SIZE): + token_end = min(token_start + CHUNK_SIZE, seqlen) + gk[:, :, token_start:token_end] = torch.cumsum( + raw_gate[:, :, token_start:token_end] * rcp_ln2, + dim=2, + ) + initial_state = torch.randn(batch, v_num_head, k_dim, v_dim, dtype=torch.float32) * 0.01 + expected = _chunk_gated_delta_rule_fwd_h_kda_reference( + k, + w, + u, + gk, + initial_state, + CHUNK_SIZE, + ) + + k_npu = k.npu() + w_npu = w.npu() + u_npu = u.npu() + gk_npu = gk.npu() + initial_state_npu = initial_state.npu() + + def run_fwd_h(): + return torch.ops._C_ascend.chunk_gated_delta_rule_fwd_h( + k_npu, + w_npu, + u_npu, + gk=gk_npu, + initial_state=initial_state_npu, + output_final_state=True, + chunk_size=CHUNK_SIZE, + save_new_value=True, + ) + + run_fwd_h() + torch.npu.synchronize() + got = run_fwd_h() + reference_outputs = _snapshot_outputs(got) + for repeat in range(1, DETERMINISM_REPEATS): + current_outputs = _snapshot_outputs(run_fwd_h()) + _assert_outputs_bitwise_equal(reference_outputs, current_outputs, repeat) + + for name, actual, expected_output in zip(FWD_H_OUTPUT_NAMES, got, expected): + assert torch.isfinite(actual).all().item() + _assert_cosine_close(name, actual, expected_output) + _cleanup_npu() diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_chunk_kda_aclnn.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_chunk_kda_aclnn.py new file mode 100644 index 000000000000..f91b59188812 --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_chunk_kda_aclnn.py @@ -0,0 +1,426 @@ +# +# Copyright (c) 2026 Huawei Technologies Co., Ltd. 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. +# 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. +# This file is a part of the vllm-ascend project. + +import gc +from dataclasses import dataclass + +import pytest +import torch +import torch_npu + +from vllm_ascend.utils import enable_custom_op + +torch_npu.npu.config.allow_internal_format = True +enable_custom_op() + +DETERMINISM_REPEATS = 20 +CHUNK_KDA_OUTPUT_NAMES = ( + "o", + "final_state", + "g", + "aqk", + "akk", + "w", + "u", + "qg", + "kg", + "v_new", + "h", + "initial_state_out", +) + + +@dataclass +class ChunkKdaReferenceResult: + o: torch.Tensor + final_state: torch.Tensor | None + + +def _lower_inverse(mat: torch.Tensor) -> torch.Tensor: + eye = torch.eye(mat.shape[-1], device=mat.device, dtype=torch.float32) + lhs = eye + torch.tril(mat.to(torch.float32), diagonal=-1) + return torch.linalg.solve_triangular(lhs, eye, upper=False) + + +def chunk_kda_forward_reference( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + gk: torch.Tensor, + beta: torch.Tensor, + scale: float, + chunk_size: int = 64, + initial_state: torch.Tensor | None = None, + output_final_state: bool = False, +) -> ChunkKdaReferenceResult: + bsz, total_t, hq, kdim = q.shape + _, _, hv, vdim = v.shape + group = hv // hq + out_dtype = v.dtype + device = q.device + o = torch.zeros((bsz, total_t, hv, vdim), device=device, dtype=out_dtype) + state = ( + torch.zeros((bsz, hv, kdim, vdim), device=device, dtype=torch.float32) + if initial_state is None + else initial_state.to(torch.float32).clone() + ) + + for b in range(bsz): + for start in range(0, total_t, chunk_size): + end = min(start + chunk_size, total_t) + cur_t = end - start + for ihv in range(hv): + ih = ihv // group + q_blk = q[b, start:end, ih].to(torch.float32) + k_blk = k[b, start:end, ih].to(torch.float32) + v_blk = v[b, start:end, ihv].to(torch.float32) + g_blk = gk[b, start:end, ihv].to(torch.float32) + beta_blk = beta[b, start:end, ihv].to(torch.float32) + + causal = torch.ones((cur_t, cur_t), device=device, dtype=torch.bool).tril() + strict_causal = torch.ones((cur_t, cur_t), device=device, dtype=torch.bool).tril(diagonal=-1) + rel = g_blk[:, None, :] - g_blk[None, :, :] + rel = rel.masked_fill(~causal[:, :, None], 0.0) + gate = torch.exp2(rel) + qk = torch.einsum("ik,jk,ijk->ij", q_blk, k_blk, gate) * float(scale) + kk = torch.einsum("ik,jk,ijk->ij", k_blk, k_blk, gate) + tril_qk = torch.where(causal, qk, torch.zeros_like(qk)) + tril_kk = torch.where(strict_causal, kk * beta_blk[:, None], torch.zeros_like(kk)) + inv_akk = _lower_inverse(tril_kk) + + k_beta_g = k_blk * beta_blk[:, None] * torch.exp2(g_blk) + v_beta = v_blk * beta_blk[:, None] + w_blk = inv_akk @ k_beta_g + u_blk = inv_akk @ v_beta + + last_g = g_blk[cur_t - 1] + qg_blk = q_blk * torch.exp2(g_blk) + kg_blk = k_blk * torch.exp2(last_g[None, :] - g_blk) + h_prev = state[b, ihv].clone() + v_new_blk = u_blk - w_blk @ h_prev + state[b, ihv] = torch.exp2(last_g)[:, None] * h_prev + kg_blk.T @ v_new_blk + + o_inter = qg_blk @ h_prev * float(scale) + o_local = tril_qk @ v_new_blk + o[b, start:end, ihv] = (o_inter + o_local).to(out_dtype) + + return ChunkKdaReferenceResult(o=o, final_state=state if output_final_state else None) + + +def _cleanup_npu(): + gc.collect() + torch.npu.empty_cache() + torch.npu.reset_peak_memory_stats() + + +def _assert_close(name, actual, expected, rtol=5e-2, atol=5e-2): + torch.testing.assert_close(actual.detach().cpu(), expected.detach().cpu(), rtol=rtol, atol=atol, msg=name) + + +def _snapshot_outputs(outputs): + torch.npu.synchronize() + return tuple(output.detach().cpu().contiguous() for output in outputs) + + +def _assert_outputs_bitwise_equal(reference, actual, repeat): + assert len(reference) == len(actual) == len(CHUNK_KDA_OUTPUT_NAMES) + for name, expected, current in zip(CHUNK_KDA_OUTPUT_NAMES, reference, actual): + same_metadata = expected.shape == current.shape and expected.dtype == current.dtype + same_bits = same_metadata and torch.equal(expected.view(torch.uint8), current.view(torch.uint8)) + if same_bits: + continue + + expected_float = expected.float() + current_float = current.float() + finite = torch.isfinite(expected_float) & torch.isfinite(current_float) + max_abs_diff = ( + (expected_float[finite] - current_float[finite]).abs().max().item() if finite.any() else float("nan") + ) + expected_nonfinite = (~torch.isfinite(expected_float)).sum().item() + current_nonfinite = (~torch.isfinite(current_float)).sum().item() + raise AssertionError( + f"repeat={repeat} output={name} is not bitwise deterministic: " + f"expected_nonfinite={expected_nonfinite}, current_nonfinite={current_nonfinite}, " + f"max_abs_diff={max_abs_diff}" + ) + + +def _gate_cumsum_reference(g, chunk_size, cu_seqlens=None): + g_cpu = g.detach().cpu().to(torch.float32) + ref = torch.empty_like(g_cpu) + rcp_ln2 = 1.4426950408889634 + if cu_seqlens is None: + cu_seqlens = [0, g_cpu.shape[1]] + for start, end in zip(cu_seqlens[:-1], cu_seqlens[1:]): + for chunk_start in range(start, end, chunk_size): + chunk_end = min(chunk_start + chunk_size, end) + ref[:, chunk_start:chunk_end] = torch.cumsum(g_cpu[:, chunk_start:chunk_end] * rcp_ln2, dim=1) + return ref + + +def _layout_swap12_reference(x): + return x.transpose(1, 2).contiguous() + + +def test_kda_torch_bindings_have_shape_correct_meta_kernels(): + q = torch.empty((1, 64, 1, 128), device="meta", dtype=torch.bfloat16) + k = torch.empty_like(q) + v = torch.empty((1, 64, 2, 256), device="meta", dtype=torch.bfloat16) + raw_gate = torch.empty((1, 64, 2, 128), device="meta", dtype=torch.bfloat16) + beta = torch.empty((1, 64, 2), device="meta", dtype=torch.float32) + + gk = torch.ops._C_ascend.kda_gate_cumsum(raw_gate, 64, layout="BSND") + outputs = torch.ops._C_ascend.chunk_kda_fwd( + q, + k, + v, + gk, + beta, + 128**-0.5, + 64, + layout="BSND", + output_final_state=True, + return_intermediate=True, + ) + swapped = torch.ops._C_ascend.kda_layout_swap12(raw_gate) + + assert gk.shape == raw_gate.shape + assert gk.dtype == torch.float32 + assert [tuple(output.shape) for output in outputs] == [ + (1, 64, 2, 256), + (1, 2, 128, 256), + (1, 64, 2, 128), + (1, 64, 2, 64), + (1, 64, 2, 64), + (1, 64, 2, 128), + (1, 64, 2, 256), + (1, 64, 2, 128), + (1, 64, 2, 128), + (1, 64, 2, 256), + (1, 1, 2, 128, 256), + (0,), + ] + assert outputs[0].dtype == torch.bfloat16 + assert outputs[1].dtype == torch.float32 + assert swapped.shape == (1, 2, 64, 128) + assert swapped.dtype == torch.bfloat16 + + +@torch.inference_mode() +def test_chunk_kda_fwd_matches_reference_bsnd(): + torch.manual_seed(20260720) + + bsz, total_t, hq, hv, kdim, vdim = 1, 64, 1, 1, 128, 128 + dtype = torch.float16 + q = (torch.randn(bsz, total_t, hq, kdim, dtype=dtype) * 0.05).npu() + k = (torch.randn(bsz, total_t, hq, kdim, dtype=dtype) * 0.05).npu() + v = (torch.randn(bsz, total_t, hv, vdim, dtype=dtype) * 0.05).npu() + g = (-torch.rand(bsz, total_t, hv, kdim, dtype=torch.float32) * 0.05).npu() + beta = torch.sigmoid(torch.randn(bsz, total_t, hv, dtype=torch.float32)).npu() + gk = torch.ops._C_ascend.kda_gate_cumsum(g, 64, layout="BSND") + initial_state = (torch.randn(bsz, hv, kdim, vdim, dtype=torch.float32) * 0.01).npu() + scale = kdim**-0.5 + + got = torch.ops._C_ascend.chunk_kda_fwd( + q, + k, + v, + gk, + beta, + scale, + 64, + layout="BSND", + initial_state=initial_state, + output_final_state=True, + return_intermediate=True, + ) + ref = chunk_kda_forward_reference( + q.cpu(), + k.cpu(), + v.cpu(), + gk.cpu(), + beta.cpu(), + scale=scale, + chunk_size=64, + initial_state=initial_state.cpu(), + output_final_state=True, + ) + + assert torch.isfinite(got[0]).all().item() + assert torch.isfinite(got[1]).all().item() + _assert_close("o", got[0], ref.o, rtol=2e-2, atol=2e-2) + _assert_close("final_state", got[1], ref.final_state, rtol=2e-2, atol=2e-2) + _cleanup_npu() + + +@pytest.mark.parametrize( + ("shape", "dtype", "with_dependency"), + [ + ((1, 64, 2, 128), torch.float32, False), + ((1, 64, 2, 128), torch.float16, True), + ((1, 64, 2, 128), torch.bfloat16, False), + ], +) +@torch.inference_mode() +def test_kda_layout_swap12_matches_reference(shape, dtype, with_dependency): + torch.manual_seed(20260720 + len(shape) + shape[-1]) + + x = (torch.randn(*shape, dtype=dtype) * 0.04).npu() + expected = _layout_swap12_reference(x.cpu()) + dependency = torch.empty_like(expected).npu() if with_dependency else None + + got = torch.ops._C_ascend.kda_layout_swap12(x, dependency=dependency) + + assert torch.isfinite(got).all().item() + _assert_close("layout_swap12", got, expected, rtol=0, atol=0) + _cleanup_npu() + + +@pytest.mark.parametrize( + ("total_t", "hq", "hv", "kdim", "vdim", "dtype"), + [ + (64, 1, 1, 128, 128, torch.float16), + (128, 1, 2, 128, 256, torch.float16), + (128, 2, 2, 128, 256, torch.bfloat16), + ], +) +@torch.inference_mode() +def test_chunk_kda_fwd_c128_v256_path(total_t, hq, hv, kdim, vdim, dtype): + torch.manual_seed(20260720 + total_t + hq + hv + vdim) + + q = (torch.randn(1, total_t, hq, kdim, dtype=dtype) * 0.04).npu() + k = (torch.randn(1, total_t, hq, kdim, dtype=dtype) * 0.04).npu() + v = (torch.randn(1, total_t, hv, vdim, dtype=dtype) * 0.04).npu() + g = (-torch.rand(1, total_t, hv, kdim, dtype=torch.float32) * 0.04).npu() + beta = torch.sigmoid(torch.randn(1, total_t, hv, dtype=torch.float32)).npu() + initial_state = (torch.randn(1, hv, kdim, vdim, dtype=torch.float32) * 0.01).npu() + scale = kdim**-0.5 + + gk = torch.ops._C_ascend.kda_gate_cumsum(g, 64, layout="BSND") + + def run_chunk_kda_fwd(): + return torch.ops._C_ascend.chunk_kda_fwd( + q, + k, + v, + gk, + beta, + scale, + 64, + layout="BSND", + initial_state=initial_state, + output_final_state=True, + return_intermediate=True, + ) + + is_a5_determinism_case = ( + total_t == 128 and hq == 2 and hv == 2 and kdim == 128 and vdim == 256 and dtype == torch.bfloat16 + ) + if is_a5_determinism_case: + run_chunk_kda_fwd() + torch.npu.synchronize() + got = run_chunk_kda_fwd() + reference_outputs = _snapshot_outputs(got) + for repeat in range(1, DETERMINISM_REPEATS): + current_outputs = _snapshot_outputs(run_chunk_kda_fwd()) + _assert_outputs_bitwise_equal(reference_outputs, current_outputs, repeat) + else: + got = run_chunk_kda_fwd() + ref = chunk_kda_forward_reference( + q.cpu(), + k.cpu(), + v.cpu(), + gk.cpu(), + beta.cpu(), + scale=scale, + chunk_size=64, + initial_state=initial_state.cpu(), + output_final_state=True, + ) + + assert torch.isfinite(got[0]).all().item() + assert torch.isfinite(got[1]).all().item() + _assert_close("o", got[0], ref.o) + _assert_close("final_state", got[1], ref.final_state) + _cleanup_npu() + + +@torch.inference_mode() +def test_kda_gate_cumsum_matches_reference(): + torch.manual_seed(20260720) + + g = (-torch.rand(1, 96, 2, 128, dtype=torch.float32) * 0.05).npu() + cu_seqlens = [0, 31, 96] + out = torch.ops._C_ascend.kda_gate_cumsum(g, 64, cu_seqlens=cu_seqlens, layout="BSND") + ref = _gate_cumsum_reference(g, 64, cu_seqlens) + + assert torch.isfinite(out).all().item() + _assert_close("gk", out, ref, rtol=2e-3, atol=2e-3) + _cleanup_npu() + + +@torch.inference_mode() +def test_chunk_kda_fwd_bnsd_layout_matches_reference(): + torch.manual_seed(20260720) + + total_t, hq, hv, kdim, vdim = 64, 1, 1, 128, 128 + q_bsnd = (torch.randn(1, total_t, hq, kdim, dtype=torch.float16) * 0.04).npu() + k_bsnd = (torch.randn(1, total_t, hq, kdim, dtype=torch.float16) * 0.04).npu() + v_bsnd = (torch.randn(1, total_t, hv, vdim, dtype=torch.float16) * 0.04).npu() + g_bsnd = (-torch.rand(1, total_t, hv, kdim, dtype=torch.float32) * 0.04).npu() + beta_bsn = torch.sigmoid(torch.randn(1, total_t, hv, dtype=torch.float32)).npu() + initial_state = (torch.randn(1, hv, kdim, vdim, dtype=torch.float32) * 0.01).npu() + scale = kdim**-0.5 + + q_bnsd = q_bsnd.transpose(1, 2).contiguous() + k_bnsd = k_bsnd.transpose(1, 2).contiguous() + v_bnsd = v_bsnd.transpose(1, 2).contiguous() + g_bnsd = g_bsnd.transpose(1, 2).contiguous() + beta_bns = beta_bsn.transpose(1, 2).contiguous() + + gk_bnsd = torch.ops._C_ascend.kda_gate_cumsum(g_bnsd, 64, layout="BNSD") + got = torch.ops._C_ascend.chunk_kda_fwd( + q_bnsd, + k_bnsd, + v_bnsd, + gk_bnsd, + beta_bns, + scale, + 64, + layout="BNSD", + initial_state=initial_state, + output_final_state=True, + return_intermediate=True, + ) + gk_bsnd = gk_bnsd.transpose(1, 2).contiguous() + ref = chunk_kda_forward_reference( + q_bsnd.cpu(), + k_bsnd.cpu(), + v_bsnd.cpu(), + gk_bsnd.cpu(), + beta_bsn.cpu(), + scale=scale, + chunk_size=64, + initial_state=initial_state.cpu(), + output_final_state=True, + ) + + out_bsnd = got[0].transpose(1, 2).contiguous() + assert torch.isfinite(out_bsnd).all().item() + assert torch.isfinite(got[1]).all().item() + _assert_close("o", out_bsnd, ref.o) + _assert_close("final_state", got[1], ref.final_state) + _cleanup_npu() diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_dequant_situ_quant.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_dequant_situ_quant.py new file mode 100644 index 000000000000..d59180d67afd --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_dequant_situ_quant.py @@ -0,0 +1,201 @@ +# +# Copyright (c) 2026 Huawei Technologies Co., Ltd. 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. +# 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. +# This file is a part of the vllm-ascend project. + +"""Kimi K3 W4A8 numerical coverage for DequantSituQuant.""" + +import pytest +import torch +import torch_npu # noqa: F401 + +from vllm_ascend.utils import enable_custom_op + +enable_custom_op() + +K3_ROUTED_INPUT_WIDTH = 6144 +K3_SHARED_TP_CASES = ( + (1, 12288), + (2, 6144), + (4, 3072), + (8, 1536), + (16, 768), +) +K3_SITU_BETA = 4.0 +K3_SITU_LINEAR_BETA = 25.0 +K3_LOCAL_EXPERTS = 14 +K3_TOP_K = 16 + + +def _kimi_k3_reference( + x: torch.Tensor, + weight_scale: torch.Tensor, + activation_scale: torch.Tensor, + bias: torch.Tensor | None, + group_index: torch.Tensor | None, +) -> tuple[torch.Tensor, torch.Tensor]: + if group_index is None: + row_expert = torch.zeros(x.shape[0], dtype=torch.long) + weight_scale = weight_scale.reshape(1, -1) + bias = None if bias is None else bias.reshape(1, -1) + else: + row_expert = torch.repeat_interleave(torch.arange(group_index.numel()), group_index) + assert row_expert.numel() == x.shape[0] + + dequant = x.float() * weight_scale[row_expert] * activation_scale.reshape(-1, 1) + if bias is not None: + dequant = dequant + bias[row_expert] + gate, up = dequant.chunk(2, dim=-1) + gate = K3_SITU_BETA * torch.tanh(gate / K3_SITU_BETA) * torch.sigmoid(gate) + up = K3_SITU_LINEAR_BETA * torch.tanh(up / K3_SITU_LINEAR_BETA) + situ = gate * up + scale = situ.abs().amax(dim=-1) / 127.0 + scale = torch.where(scale == 0, torch.ones_like(scale), scale) + y = torch.round(situ / scale.unsqueeze(-1)).clamp(-128, 127).to(torch.int8) + return y, scale + + +def _run_dequant_situ_quant( + x: torch.Tensor, + weight_scale: torch.Tensor | None, + activation_scale: torch.Tensor | None, + bias: torch.Tensor | None, + group_index: torch.Tensor | None, +) -> tuple[torch.Tensor, torch.Tensor]: + return torch.ops._C_ascend.dequant_situ_quant( + x=x.npu(), + weight_scale=None if weight_scale is None else weight_scale.npu(), + activation_scale=None if activation_scale is None else activation_scale.npu(), + bias=None if bias is None else bias.npu(), + quant_scale=None, + quant_offset=None, + group_index=None if group_index is None else group_index.npu(), + beta=K3_SITU_BETA, + linear_beta=K3_SITU_LINEAR_BETA, + activate_left=True, + quant_mode="dynamic", + ) + + +def _kimi_k3_predequantized_reference(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + gate, up = x.float().chunk(2, dim=-1) + gate = K3_SITU_BETA * torch.tanh(gate / K3_SITU_BETA) * torch.sigmoid(gate) + up = K3_SITU_LINEAR_BETA * torch.tanh(up / K3_SITU_LINEAR_BETA) + situ = gate * up + scale = situ.abs().amax(dim=-1) / 127.0 + scale = torch.where(scale == 0, torch.ones_like(scale), scale) + y = torch.round(situ / scale.unsqueeze(-1)).clamp(-128, 127).to(torch.int8) + return y, scale + + +@pytest.mark.skip_global_cleanup +@pytest.mark.parametrize( + "group_counts", + ( + # One decode token expands to top-k=16 routed rows. Zero-count + # experts are expected on an EP rank and must not shift later scales. + (2, 0, 1, 1, 0, 2, 1, 1, 1, 0, 2, 1, 2, 2), + # A 65-token prefill expands to 65 * top-k rows. Keep all 14 local + # experts non-empty and uneven so every expert boundary is exercised. + (75, 75, 75, 75, 74, 74, 74, 74, 74, 74, 74, 74, 74, 74), + ), + ids=("decode_t1_topk16", "prefill_t65_topk16"), +) +@torch.inference_mode() +def test_kimi_k3_routed_multi_expert_dequant_situ_quant(group_counts: tuple[int, ...]): + if not hasattr(torch.ops._C_ascend, "dequant_situ_quant"): + pytest.skip("requires the DequantSituQuant custom operator") + + assert len(group_counts) == K3_LOCAL_EXPERTS + assert sum(group_counts) in (K3_TOP_K, 65 * K3_TOP_K) + group_index = torch.tensor(group_counts, dtype=torch.int64) + rows = int(group_index.sum()) + x = (torch.arange(rows * K3_ROUTED_INPUT_WIDTH, dtype=torch.int64) * 37) % 40001 - 20000 + x = x.to(torch.int32).reshape(rows, K3_ROUTED_INPUT_WIDTH) + + phase = torch.linspace(-torch.pi, torch.pi, K3_ROUTED_INPUT_WIDTH, dtype=torch.float32) + expert_offset = torch.arange(K3_LOCAL_EXPERTS, dtype=torch.float32).reshape(-1, 1) + weight_scale = 0.0008 + 0.00005 * expert_offset + 0.0002 * phase.cos() + bias = -0.75 + 0.125 * expert_offset + 0.35 * phase.sin() + activation_scale = torch.linspace(0.020, 0.080, rows, dtype=torch.float32) + + expected_y, expected_scale = _kimi_k3_reference(x, weight_scale, activation_scale, bias, group_index) + actual_y, actual_scale = _run_dequant_situ_quant(x, weight_scale, activation_scale, bias, group_index) + + assert tuple(actual_y.shape) == (rows, K3_ROUTED_INPUT_WIDTH // 2) + assert tuple(actual_scale.shape) == (rows,) + torch.testing.assert_close(actual_y.cpu(), expected_y, rtol=0, atol=1) + torch.testing.assert_close(actual_scale.cpu(), expected_scale, rtol=5e-3, atol=1e-5) + + +@pytest.mark.skip_global_cleanup +@pytest.mark.parametrize("rows", (1, 16, 1040), ids=("decode_m1", "decode_topk16", "prefill_t65_topk16")) +@torch.inference_mode() +def test_kimi_k3_predequantized_bf16_situ_quant(rows: int): + if not hasattr(torch.ops._C_ascend, "dequant_situ_quant"): + pytest.skip("requires the DequantSituQuant custom operator") + + values = torch.linspace(-32.0, 32.0, rows * K3_ROUTED_INPUT_WIDTH, dtype=torch.float32) + x = (values + 0.125 * torch.sin(values)).to(torch.bfloat16).reshape(rows, K3_ROUTED_INPUT_WIDTH) + expected_y, expected_scale = _kimi_k3_predequantized_reference(x) + actual_y, actual_scale = _run_dequant_situ_quant(x, None, None, None, None) + + assert tuple(actual_y.shape) == (rows, K3_ROUTED_INPUT_WIDTH // 2) + assert tuple(actual_scale.shape) == (rows,) + torch.testing.assert_close(actual_y.cpu(), expected_y, rtol=0, atol=1) + torch.testing.assert_close(actual_scale.cpu(), expected_scale, rtol=5e-3, atol=1e-5) + + +@pytest.mark.skip_global_cleanup +@pytest.mark.parametrize( + ("tp_size", "input_width"), + K3_SHARED_TP_CASES, + ids=("tp1", "tp2", "tp4", "tp8", "tp16"), +) +@pytest.mark.parametrize("rows", (1, 65), ids=("decode_m1", "prefill_m65")) +@torch.inference_mode() +def test_kimi_k3_shared_expert_without_bias_or_group_index(tp_size: int, input_width: int, rows: int): + if not hasattr(torch.ops._C_ascend, "dequant_situ_quant"): + pytest.skip("requires the DequantSituQuant custom operator") + + assert input_width == 12288 // tp_size + x = (torch.arange(rows * input_width, dtype=torch.int64) * 19) % 30001 - 15000 + x = x.to(torch.int32).reshape(rows, input_width) + weight_scale = torch.linspace(0.0007, 0.0017, input_width, dtype=torch.float32) + activation_scale = torch.linspace(0.025, 0.075, rows, dtype=torch.float32).reshape(rows, 1) + + expected_y, expected_scale = _kimi_k3_reference(x, weight_scale, activation_scale, None, None) + actual_y, actual_scale = _run_dequant_situ_quant(x, weight_scale, activation_scale, None, None) + + assert tuple(actual_y.shape) == (rows, input_width // 2) + assert tuple(actual_scale.shape) == (rows,) + torch.testing.assert_close(actual_y.cpu(), expected_y, rtol=0, atol=1) + torch.testing.assert_close(actual_scale.cpu(), expected_scale, rtol=5e-3, atol=1e-5) + + +@pytest.mark.skip_global_cleanup +@torch.inference_mode() +def test_kimi_k3_zero_routed_rows_return_empty_outputs(): + if not hasattr(torch.ops._C_ascend, "dequant_situ_quant"): + pytest.skip("requires the DequantSituQuant custom operator") + + group_index = torch.zeros(K3_LOCAL_EXPERTS, dtype=torch.int64) + x = torch.empty((0, K3_ROUTED_INPUT_WIDTH), dtype=torch.int32) + weight_scale = torch.ones((K3_LOCAL_EXPERTS, K3_ROUTED_INPUT_WIDTH), dtype=torch.float32) + activation_scale = torch.empty((0,), dtype=torch.float32) + bias = torch.zeros_like(weight_scale) + actual_y, actual_scale = _run_dequant_situ_quant(x, weight_scale, activation_scale, bias, group_index) + + assert tuple(actual_y.shape) == (0, K3_ROUTED_INPUT_WIDTH // 2) + assert tuple(actual_scale.shape) == (0,) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_k3_situ_fusion_cases.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_k3_situ_fusion_cases.py new file mode 100644 index 000000000000..7befee5cbcf7 --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_k3_situ_fusion_cases.py @@ -0,0 +1,235 @@ +# +# Copyright (c) 2026 Huawei Technologies Co., Ltd. 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. +# 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. +# This file is a part of the vllm-ascend project. + +"""Single-operator Kimi K3 SiTU fusion fixtures. + +The A3 fixtures intentionally call DequantSituQuant directly. Shared experts +exercise its INT32 dequant mode; routed experts exercise the same operator's +pre-dequantized BF16 mode with every dequant input absent. +""" + +import math + +import pytest +import torch +import torch_npu + +from vllm_ascend.utils import enable_custom_op + +enable_custom_op() + +K3_BETA = 4.0 +K3_LINEAR_BETA = 25.0 +K3_ROUTED_INPUT_WIDTH = 6144 +K3_LOCAL_EXPERTS = 14 +K3_SHARED_TP_CASES = ( + (1, 12288), + (2, 6144), + (4, 3072), + (8, 1536), + (16, 768), +) + + +def _situ(values: torch.Tensor) -> torch.Tensor: + gate, up = values.float().chunk(2, dim=-1) + gate = K3_BETA * torch.tanh(gate / K3_BETA) * torch.sigmoid(gate) + up = K3_LINEAR_BETA * torch.tanh(up / K3_LINEAR_BETA) + return gate * up + + +def _dynamic_int8_quant(values: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + if values.shape[0] == 0: + return ( + torch.empty((0, values.shape[-1]), dtype=torch.int8), + torch.empty((0,), dtype=torch.float32), + ) + scale = values.abs().amax(dim=-1) / 127.0 + scale = torch.where(scale == 0, torch.ones_like(scale), scale) + output = torch.round(values / scale[:, None]).clamp(-128, 127).to(torch.int8) + return output, scale + + +def _shared_inputs(rows: int, width: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + accumulator = (torch.arange(rows * width, dtype=torch.int64) * 19) % 30001 - 15000 + accumulator = accumulator.to(torch.int32).reshape(rows, width) + weight_scale = torch.linspace(0.0007, 0.0017, width, dtype=torch.float32) + activation_scale = torch.linspace(0.025, 0.075, rows, dtype=torch.float32).reshape(rows, 1) + return accumulator, weight_scale, activation_scale + + +def _bf16_input(rows: int, width: int) -> torch.Tensor: + if rows == 0: + return torch.empty((0, width), dtype=torch.bfloat16) + values = torch.arange(rows * width, dtype=torch.int64) + values = ((values * 37) % 4097).float() / 64.0 - 32.0 + return values.to(torch.bfloat16).reshape(rows, width) + + +def _routed_input(rows: int) -> torch.Tensor: + return _bf16_input(rows, K3_ROUTED_INPUT_WIDTH) + + +def _run_a3( + x: torch.Tensor, + weight_scale: torch.Tensor | None, + activation_scale: torch.Tensor | None, +) -> tuple[torch.Tensor, torch.Tensor]: + assert hasattr(torch.ops._C_ascend, "dequant_situ_quant") + return torch.ops._C_ascend.dequant_situ_quant( + x=x.npu(), + weight_scale=None if weight_scale is None else weight_scale.npu(), + activation_scale=None if activation_scale is None else activation_scale.npu(), + bias=None, + quant_scale=None, + quant_offset=None, + group_index=None, + beta=K3_BETA, + linear_beta=K3_LINEAR_BETA, + activate_left=True, + quant_mode="dynamic", + ) + + +@pytest.mark.skip_global_cleanup +@pytest.mark.parametrize( + ("tp_size", "input_width"), + K3_SHARED_TP_CASES, + ids=("tp1", "tp2", "tp4", "tp8", "tp16"), +) +@pytest.mark.parametrize(("phase", "rows"), (("decode", 1), ("prefill", 65)), ids=("decode", "prefill")) +@torch.inference_mode() +def test_a3_shared_dequant_situ_quant_single_op( + tp_size: int, + input_width: int, + phase: str, + rows: int, +): + """Shared experts: INT32 accumulator plus FP32 dequant scales.""" + if not hasattr(torch.ops._C_ascend, "dequant_situ_quant"): + pytest.skip("requires the DequantSituQuant custom operator") + + assert phase in {"decode", "prefill"} + assert input_width == 12288 // tp_size + x, weight_scale, activation_scale = _shared_inputs(rows, input_width) + dequantized = x.float() * weight_scale[None, :] * activation_scale + expected_y, expected_scale = _dynamic_int8_quant(_situ(dequantized)) + + actual_y, actual_scale = _run_a3(x, weight_scale, activation_scale) + + assert tuple(actual_y.shape) == (rows, input_width // 2) + assert actual_y.dtype == torch.int8 + assert tuple(actual_scale.shape) == (rows,) + assert actual_scale.dtype == torch.float32 + torch.testing.assert_close(actual_y.cpu(), expected_y, rtol=0, atol=1) + torch.testing.assert_close(actual_scale.cpu(), expected_scale, rtol=5e-3, atol=1e-5) + + +@pytest.mark.skip_global_cleanup +@pytest.mark.parametrize( + ("phase", "rows"), + ( + ("decode_empty", 0), + ("decode_max_local", K3_LOCAL_EXPERTS), + ("prefill_max_local", 65 * K3_LOCAL_EXPERTS), + ), + ids=("decode_empty_m0", "decode_max_m14", "prefill_max_m910"), +) +@torch.inference_mode() +def test_a3_routed_bf16_dequant_situ_quant_single_op(phase: str, rows: int): + """Routed experts: the same A3 op, BF16 input, no dequant stage.""" + if not hasattr(torch.ops._C_ascend, "dequant_situ_quant"): + pytest.skip("requires the DequantSituQuant custom operator") + + assert phase in {"decode_empty", "decode_max_local", "prefill_max_local"} + x = _routed_input(rows) + expected_y, expected_scale = _dynamic_int8_quant(_situ(x)) + + actual_y, actual_scale = _run_a3(x, None, None) + + assert tuple(actual_y.shape) == (rows, K3_ROUTED_INPUT_WIDTH // 2) + assert actual_y.dtype == torch.int8 + assert tuple(actual_scale.shape) == (rows,) + assert actual_scale.dtype == torch.float32 + torch.testing.assert_close(actual_y.cpu(), expected_y, rtol=0, atol=1) + torch.testing.assert_close(actual_scale.cpu(), expected_scale, rtol=5e-3, atol=1e-5) + + +def _is_ascend_950() -> bool: + try: + return "950" in torch.npu.get_device_name(0) + except Exception: + return False + + +@pytest.mark.skip_global_cleanup +@pytest.mark.parametrize( + ("tp_size", "input_width"), + K3_SHARED_TP_CASES, + ids=("tp1", "tp2", "tp4", "tp8", "tp16"), +) +@pytest.mark.parametrize(("phase", "rows"), (("decode", 1), ("prefill", 65)), ids=("decode", "prefill")) +@torch.inference_mode() +def test_a5_shared_situ_mx_quant_shapes_single_op( + tp_size: int, + input_width: int, + phase: str, + rows: int, +): + """A5 shared-expert MXFP output and scale layout.""" + if not _is_ascend_950() or not hasattr(torch.ops._C_ascend, "situ_mx_quant"): + pytest.skip("requires an Ascend 950 device and SituMxQuant") + + assert phase in {"decode", "prefill"} + assert input_width == 12288 // tp_size + x = _bf16_input(rows, input_width) + y, mxscale = torch.ops._C_ascend.situ_mx_quant( + x.npu(), + beta=K3_BETA, + linear_beta=K3_LINEAR_BETA, + activate_left=True, + dst_type=36, + ) + + output_width = input_width // 2 + assert tuple(y.shape) == (rows, output_width) + assert y.dtype == torch.float8_e4m3fn + assert tuple(mxscale.shape) == (rows, math.ceil(output_width / 64), 2) + assert mxscale.dtype == torch_npu.float8_e8m0fnu + + +@pytest.mark.skip_global_cleanup +@pytest.mark.parametrize(("phase", "rows"), (("decode_max_local", 14), ("prefill_max_local", 910))) +@torch.inference_mode() +def test_a5_routed_situ_mx_quant_shapes_single_op(phase: str, rows: int): + """A5 routed-expert shape is TP-invariant.""" + if not _is_ascend_950() or not hasattr(torch.ops._C_ascend, "situ_mx_quant"): + pytest.skip("requires an Ascend 950 device and SituMxQuant") + + assert phase in {"decode_max_local", "prefill_max_local"} + x = _routed_input(rows) + y, mxscale = torch.ops._C_ascend.situ_mx_quant( + x.npu(), + beta=K3_BETA, + linear_beta=K3_LINEAR_BETA, + activate_left=True, + dst_type=36, + ) + + assert tuple(y.shape) == (rows, 3072) + assert y.dtype == torch.float8_e4m3fn + assert tuple(mxscale.shape) == (rows, 48, 2) + assert mxscale.dtype == torch_npu.float8_e8m0fnu diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_kda_ascendc_npu.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_kda_ascendc_npu.py new file mode 100644 index 000000000000..04486cd1ec27 --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_kda_ascendc_npu.py @@ -0,0 +1,113 @@ +# +# Copyright (c) 2026 Huawei Technologies Co., Ltd. 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. +# 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. +# This file is a part of the vllm-ascend project. + +"""Kimi K3 integration coverage for the AscendC prefill operators.""" + +import pytest +import torch +import torch_npu # noqa: F401 + +from vllm_ascend.utils import enable_custom_op + +enable_custom_op() + + +def _l2norm(x: torch.Tensor) -> torch.Tensor: + dtype = x.dtype + x = x.float() + return (x * torch.rsqrt((x * x).sum(dim=-1, keepdim=True) + 1e-6)).to(dtype) + + +def _naive_kda( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + gate: torch.Tensor, + beta: torch.Tensor, + initial_state_kv: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + dtype = v.dtype + q, k, v, gate, beta = (x.float() for x in (q, k, v, gate, beta)) + state = initial_state_kv.float().clone() + out = torch.empty_like(v) + scale = q.shape[-1] ** -0.5 + for token in range(q.shape[1]): + state *= gate[:, token].exp().unsqueeze(-1) + residual = v[:, token] - torch.einsum("bhk,bhkv->bhv", k[:, token], state) + state += torch.einsum("bhk,bhv->bhkv", beta[:, token].unsqueeze(-1) * k[:, token], residual) + out[:, token] = torch.einsum("bhk,bhkv->bhv", q[:, token] * scale, state) + return out.to(dtype), state + + +@pytest.mark.skip_global_cleanup +@torch.inference_mode() +def test_kimi_k3_safe_gate_prefill_and_transposed_state_layout(): + if not hasattr(torch.ops._C_ascend, "kda_gate_cumsum") or not hasattr(torch.ops._C_ascend, "chunk_kda_fwd"): + pytest.skip("requires the KDA AscendC operators") + + torch.manual_seed(20260720) + tokens, heads, head_dim = 64, 1, 128 + dtype = torch.float16 + q = _l2norm(torch.randn(1, tokens, heads, head_dim, dtype=dtype, device="npu")) + k = _l2norm(torch.randn_like(q)) + v = torch.randn_like(q) * 0.05 + raw_gate = torch.randn_like(q) * 0.1 + beta = torch.rand(1, tokens, heads, dtype=torch.float32, device="npu").sigmoid() + a_log = torch.randn(heads, dtype=torch.float32, device="npu") * 0.05 + dt_bias = torch.randn(heads * head_dim, dtype=torch.float32, device="npu") * 0.05 + cache_vk = torch.randn(1, heads, head_dim, head_dim, dtype=torch.float32, device="npu") * 0.01 + lower_bound = -5.0 + cu_seqlens = (0, tokens) + chunk_indices = (0, 0) + + gate_cumsum = torch.ops._C_ascend.kda_gate_cumsum( + raw_gate, + 64, + A_log=a_log, + dt_bias=dt_bias, + cu_seqlens=cu_seqlens, + use_gate_in_kernel=True, + safe_gate=True, + lower_bound=lower_bound, + layout="BSND", + ) + initial_state_kv = cache_vk.transpose(-1, -2).contiguous() + got = torch.ops._C_ascend.chunk_kda_fwd( + q, + k, + v, + gate_cumsum, + beta, + head_dim**-0.5, + 64, + layout="BSND", + initial_state=initial_state_kv, + output_final_state=True, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + return_intermediate=False, + ) + + safe_gate = lower_bound * torch.sigmoid( + (raw_gate.float() + dt_bias.view(1, 1, heads, head_dim)) * a_log.exp().view(1, 1, heads, 1) + ) + expected_out, expected_state_kv = _naive_kda(q, k, v, safe_gate, beta, initial_state_kv) + + torch.testing.assert_close(got[0], expected_out, rtol=3e-2, atol=3e-2) + torch.testing.assert_close(got[1], expected_state_kv, rtol=3e-2, atol=3e-2) + # The vLLM decode cache remains [H,V,K] after crossing the AscendC boundary. + cache_vk.copy_(got[1].transpose(-1, -2)) + torch.testing.assert_close(cache_vk.transpose(-1, -2), expected_state_kv, rtol=3e-2, atol=3e-2) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_kda_recurrent_ascendc_npu.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_kda_recurrent_ascendc_npu.py new file mode 100644 index 000000000000..4b0eebad34ad --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_kda_recurrent_ascendc_npu.py @@ -0,0 +1,235 @@ +# +# Copyright (c) 2026 Huawei Technologies Co., Ltd. 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. +# 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. +# This file is a part of the vllm-ascend project. + +"""Kimi K3 recurrent-KDA reference and direct AscendC accuracy coverage.""" + +from __future__ import annotations + +from collections.abc import Sequence + +import torch +import torch.nn.functional as F +import torch_npu # noqa: F401 + + +def _flatten_bsnd(x: torch.Tensor, layout: str) -> torch.Tensor: + if layout == "TND": + return x + if layout != "BSND": + raise ValueError("layout must be BSND or TND") + return x.reshape(x.shape[0] * x.shape[1], *x.shape[2:]) + + +def _restore_layout(x: torch.Tensor, ref: torch.Tensor, layout: str) -> torch.Tensor: + return x if layout == "TND" else x.reshape(ref.shape) + + +def _seq_ranges(total_tokens: int, cu_seqlens: Sequence[int]) -> list[tuple[int, int]]: + offsets = [int(offset) for offset in cu_seqlens] + if len(offsets) < 2: + raise ValueError("cu_seqlens must contain at least two cumulative offsets") + if offsets[0] != 0: + raise ValueError("cu_seqlens must start at zero") + if any(end < start for start, end in zip(offsets, offsets[1:])): + raise ValueError("cu_seqlens must be nondecreasing") + if offsets[-1] != total_tokens: + raise ValueError("the last cu_seqlens offset must equal the packed token count") + return list(zip(offsets, offsets[1:])) + + +def _state_slot(ssm_state_indices: torch.Tensor, seq_idx: int, start: int, token: int) -> int: + if ssm_state_indices.ndim == 1: + return int(ssm_state_indices[token].item()) + if ssm_state_indices.ndim == 2: + return int(ssm_state_indices[seq_idx, token - start].item()) + raise ValueError("ssm_state_indices must be packed [T] or speculative [seq_num,max_step]") + + +def recurrent_kda_reference( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + initial_state: torch.Tensor | None = None, + *, + cu_seqlens: Sequence[int], + ssm_state_indices: torch.Tensor | None = None, + A_log: torch.Tensor | None = None, + dt_bias: torch.Tensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + layout: str = "BSND", + scale: float | None = None, + output_final_state: bool = True, + use_qk_l2norm_in_kernel: bool = False, + use_gate_in_kernel: bool = False, + use_beta_sigmoid_in_kernel: bool = False, + allow_neg_eigval: bool = False, + safe_gate: bool = False, + lower_bound: float = -5.0, + state_v_first: bool = True, +) -> tuple[torch.Tensor, torch.Tensor]: + del output_final_state + if not state_v_first: + raise ValueError("reference only supports state_v_first=True") + + q_flat = _flatten_bsnd(q, layout).float() + k_flat = _flatten_bsnd(k, layout).float() + v_flat = _flatten_bsnd(v, layout).float() + g_flat = _flatten_bsnd(g, layout).float() + beta_flat = _flatten_bsnd(beta, layout).float() + total_tokens, h, dk = q_flat.shape + _, hv, dv = v_flat.shape + scale = dk**-0.5 if scale is None else scale + + if use_qk_l2norm_in_kernel: + q_flat = F.normalize(q_flat, p=2, dim=-1) + k_flat = F.normalize(k_flat, p=2, dim=-1) + q_flat = q_flat * scale + + if use_gate_in_kernel: + if A_log is None: + raise ValueError("A_log is required when use_gate_in_kernel=True") + gate = g_flat + if dt_bias is not None: + gate = gate + dt_bias.float().reshape(hv, dk).unsqueeze(0) + exp_a = torch.exp(A_log.float()).reshape(1, hv, 1) + gate = lower_bound * torch.sigmoid(exp_a * gate) if safe_gate else -exp_a * F.softplus(gate) + else: + gate = g_flat + gate_decay = torch.exp(gate.float()) + + beta_eff = beta_flat + if use_beta_sigmoid_in_kernel: + beta_eff = torch.sigmoid(beta_eff) + if allow_neg_eigval: + beta_eff = beta_eff * 2.0 + + ranges = _seq_ranges(total_tokens, cu_seqlens) + state_dtype = initial_state.dtype if initial_state is not None else torch.float32 + state = ( + torch.zeros((len(ranges), hv, dv, dk), dtype=torch.float32, device=q.device) + if initial_state is None + else initial_state.float().clone() + ) + out_flat = torch.zeros_like(v_flat, dtype=torch.float32) + + for seq_idx, (start, end) in enumerate(ranges): + if start == end: + continue + state_slot = seq_idx + if ssm_state_indices is not None: + token = start + if num_accepted_tokens is not None: + token = start + int(num_accepted_tokens[seq_idx].item()) - 1 + state_slot = _state_slot(ssm_state_indices, seq_idx, start, token) + for hv_idx in range(hv): + h_idx = hv_idx // (hv // h) + state_cur = state[state_slot, hv_idx].clone() + for token in range(start, end): + state_cur = state_cur * gate_decay[token, hv_idx].unsqueeze(0) + delta = v_flat[token, hv_idx] - torch.mv(state_cur, k_flat[token, h_idx]) + state_cur = state_cur + torch.outer(delta * beta_eff[token, hv_idx], k_flat[token, h_idx]) + out_flat[token, hv_idx] = torch.mv(state_cur, q_flat[token, h_idx]) + out_slot = ( + _state_slot(ssm_state_indices, seq_idx, start, token) if ssm_state_indices is not None else seq_idx + ) + state[out_slot, hv_idx] = state_cur + + return _restore_layout(out_flat.to(q.dtype), v, layout), state.to(state_dtype) + + +@torch.inference_mode() +def test_kimi_k3_tp16_recurrent_kda_non_contiguous_state_pool(): + """Preserve a strided cache view while updating only selected Kimi slots.""" + torch.manual_seed(20260806) + device = torch.device("npu") + batch, heads, dim = 4, 6, 128 + state_capacity = 17 + cu_seqlens_host = list(range(batch + 1)) + state_indices_cpu = torch.tensor([9, 2, 15, 4], dtype=torch.int64) + + q_cpu = torch.randn(1, batch, heads, dim, dtype=torch.bfloat16) + k_cpu = torch.randn_like(q_cpu) + v_cpu = torch.randn_like(q_cpu) + raw_gate_cpu = torch.randn(1, batch, heads, dim, dtype=torch.bfloat16) * 0.25 + beta_cpu = torch.rand(1, batch, heads, dtype=torch.float32).sigmoid() + state_cpu = torch.randn(state_capacity, heads, dim, dim, dtype=torch.float32) * 0.01 + a_log_cpu = torch.randn(heads, dtype=torch.float32) * 0.05 + dt_bias_cpu = torch.randn(heads, dim, dtype=torch.float32) * 0.05 + + ref_out, ref_state = recurrent_kda_reference( + q_cpu, + k_cpu, + v_cpu, + raw_gate_cpu, + beta_cpu, + state_cpu, + cu_seqlens=cu_seqlens_host, + ssm_state_indices=state_indices_cpu, + A_log=a_log_cpu, + dt_bias=dt_bias_cpu, + layout="BSND", + scale=dim**-0.5, + use_qk_l2norm_in_kernel=True, + use_gate_in_kernel=True, + safe_gate=True, + ) + + state_pool = torch.full( + (state_capacity + 1, 2, heads, dim, dim), + 7.0, + dtype=torch.float32, + device=device, + ) + state_view = state_pool[1:, 0] + state_view.copy_(state_cpu.to(device)) + guard_layer = state_pool[1:, 1].clone() + state_before = state_view.clone() + state_stride = state_view.stride() + state_storage = state_view.untyped_storage().data_ptr() + assert not state_view.is_contiguous() + assert state_view.storage_offset() > 0 + + out = torch.ops._C_ascend.recurrent_kda( + q_cpu.to(device), + k_cpu.to(device), + v_cpu.to(device), + raw_gate_cpu.to(device), + beta_cpu.to(device), + state_view, + torch.tensor(cu_seqlens_host, dtype=torch.int32, device=device), + state_indices_cpu.to(device), + a_log_cpu.to(device), + dt_bias_cpu.to(device), + scale=dim**-0.5, + use_qk_l2norm_in_kernel=True, + use_gate_in_kernel=True, + use_beta_sigmoid_in_kernel=False, + allow_neg_eigval=False, + safe_gate=True, + lower_bound=-5.0, + ) + torch.npu.synchronize() + + assert state_view.stride() == state_stride + assert state_view.untyped_storage().data_ptr() == state_storage + torch.testing.assert_close(out.cpu(), ref_out, rtol=0.02, atol=0.02) + torch.testing.assert_close(state_view.cpu(), ref_state, rtol=0.02, atol=0.02) + torch.testing.assert_close(state_pool[1:, 1], guard_layer, rtol=0, atol=0) + used_slots = set(state_indices_cpu.tolist()) + untouched_slots = [slot for slot in range(state_capacity) if slot not in used_slots] + torch.testing.assert_close(state_view[untouched_slots], state_before[untouched_slots], rtol=0, atol=0) From 0cc801af664da12ef7974c87e5668d3ff1aa3bc5 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Tue, 25 Aug 2026 03:53:36 -0500 Subject: [PATCH 02/50] refactor(ops): move KDA torch adapters into operator directories Signed-off-by: maoxx241 --- .../chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h | 142 ++++++++ .../kda_gate_cumsum_torch_adpt.h | 74 ++++ .../kda_layout_swap12_torch_adpt.h | 43 +++ .../kda_layout_swap12/op_host/CMakeLists.txt | 2 + .../op_host/kda_layout_swap12_def.cpp | 2 + csrc/attention/kda_torch_adpt_common.h | 121 +++++++ .../situ_mx_quant/op_kernel/inc/platform.h | 5 +- csrc/torch_binding.cpp | 317 +----------------- 8 files changed, 390 insertions(+), 316 deletions(-) create mode 100644 csrc/attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h create mode 100644 csrc/attention/kda_gate_cumsum/kda_gate_cumsum_torch_adpt.h create mode 100644 csrc/attention/kda_layout_swap12/kda_layout_swap12_torch_adpt.h create mode 100644 csrc/attention/kda_torch_adpt_common.h diff --git a/csrc/attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h b/csrc/attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h new file mode 100644 index 000000000000..ca968bafcea4 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h @@ -0,0 +1,142 @@ +/* + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ +#ifndef CHUNK_KDA_FWD_TORCH_ADPT_H +#define CHUNK_KDA_FWD_TORCH_ADPT_H + +#include +#include + +#include "attention/kda_torch_adpt_common.h" + +namespace vllm_ascend { + +std::tuple +chunk_kda_fwd( + const at::Tensor &q, + const at::Tensor &k, + const at::Tensor &v, + const at::Tensor &gk, + const at::Tensor &beta, + double scale, + int64_t chunk_size, + c10::string_view layout, + const c10::optional &initial_state, + c10::optional output_final_state, + c10::optional cu_seqlens, + c10::optional chunk_indices, + c10::optional return_intermediate, + c10::optional safe_gate, + c10::optional transpose_state_layout) +{ + std::string layout_str(layout.data(), layout.size()); + TORCH_CHECK(layout_str == "BSND" || layout_str == "BNSD" || layout_str == "TND" || layout_str == "NTD", + "chunk_kda_fwd: layout must be one of BSND, BNSD, TND, NTD and must be uppercase."); + TORCH_CHECK(!safe_gate.value_or(false), "chunk_kda_fwd: safe_gate=True is not supported."); + TORCH_CHECK(!transpose_state_layout.value_or(false), + "chunk_kda_fwd: transpose_state_layout=True is not supported."); + TORCH_CHECK(chunk_size == 32 || chunk_size == 64 || chunk_size == 128, + "chunk_kda_fwd: chunk_size must be 32, 64 or 128."); + + bool is_tnd = layout_str == "TND"; + bool is_ntd = layout_str == "NTD"; + bool is_bsnd = layout_str == "BSND"; + bool is_bnsd = layout_str == "BNSD"; + bool is_rank3 = is_tnd || is_ntd; + TORCH_CHECK((is_rank3 && q.dim() == 3 && k.dim() == 3 && v.dim() == 3 && gk.dim() == 3 && beta.dim() == 2) || + (!is_rank3 && q.dim() == 4 && k.dim() == 4 && v.dim() == 4 && gk.dim() == 4 && beta.dim() == 3), + "chunk_kda_fwd: layout/rank mismatch."); + TORCH_CHECK(q.sizes() == k.sizes(), "chunk_kda_fwd: q and k must have identical shape."); + + auto q_sizes = q.sizes(); + auto v_sizes = v.sizes(); + bool is_internal_layout = is_bnsd || is_ntd; + int64_t B = is_rank3 ? 1 : q_sizes[0]; + int64_t T = is_tnd ? q_sizes[0] : (is_ntd ? q_sizes[1] : (is_bnsd ? q_sizes[2] : q_sizes[1])); + int64_t H = is_tnd ? q_sizes[1] : (is_ntd ? q_sizes[0] : (is_bnsd ? q_sizes[1] : q_sizes[2])); + int64_t K = is_rank3 ? q_sizes[2] : q_sizes[3]; + int64_t HV = is_tnd ? v_sizes[1] : (is_ntd ? v_sizes[0] : (is_bnsd ? v_sizes[1] : v_sizes[2])); + int64_t V = is_rank3 ? v_sizes[2] : v_sizes[3]; + TORCH_CHECK(H > 0 && HV >= H, "chunk_kda_fwd: H and HV must be positive and H must be <= HV."); + TORCH_CHECK(H <= 128 && HV <= 128, "chunk_kda_fwd: H and HV must be <= 128."); + TORCH_CHECK(!is_tnd || H == 1, + "chunk_kda_fwd: TND layout with H > 1 is not supported; use NTD for multi-head rank3 input."); + check_kda_cu_seqlens(cu_seqlens, T, "chunk_kda_fwd"); + check_kda_chunk_indices(chunk_indices, cu_seqlens, chunk_size, "chunk_kda_fwd"); + TORCH_CHECK(!cu_seqlens.has_value() || is_rank3 || B == 1, + "chunk_kda_fwd: rank4 varlen input with cu_seqlens currently requires B=1."); + TORCH_CHECK(HV % H == 0, "chunk_kda_fwd: HV must be divisible by H."); + TORCH_CHECK(q.scalar_type() == at::kHalf || q.scalar_type() == at::kBFloat16, + "chunk_kda_fwd: q/k/v must use float16 or bfloat16."); + TORCH_CHECK(k.scalar_type() == q.scalar_type() && v.scalar_type() == q.scalar_type(), + "chunk_kda_fwd: q/k/v dtype must match."); + TORCH_CHECK(chunk_size == 64 && K >= 16 && V >= 16 && K % 16 == 0 && V % 16 == 0 && V <= 256 && + K * V >= 4 * 64 * 64 && K * V >= chunk_size * (K + V), + "chunk_kda_fwd: shape is outside the supported split cube/vector template."); + + int64_t seq_num = get_kda_seq_num(B, cu_seqlens); + at::Tensor initial_state_tensor = initial_state.value_or(at::Tensor()); + if (initial_state_tensor.defined()) { + TORCH_CHECK(initial_state_tensor.scalar_type() == at::kFloat, + "chunk_kda_fwd: initial_state must be float32 when provided."); + TORCH_CHECK(initial_state_tensor.dim() == 4 && initial_state_tensor.size(0) == seq_num && + initial_state_tensor.size(1) == HV && initial_state_tensor.size(2) == K && + initial_state_tensor.size(3) == V, + "chunk_kda_fwd: initial_state must be [seq_num,Hv,K,V]."); + } + + std::vector generated_chunk_indices; + c10::optional chunk_indices_for_call; + if (chunk_indices.has_value()) { + chunk_indices_for_call = chunk_indices.value(); + } else if (cu_seqlens.has_value()) { + generated_chunk_indices = build_kda_chunk_indices(cu_seqlens.value(), chunk_size); + chunk_indices_for_call = at::IntArrayRef(generated_chunk_indices); + } else { + chunk_indices_for_call = c10::nullopt; + } + + int64_t total_chunks = get_kda_total_chunks(B, T, chunk_size, cu_seqlens, chunk_indices_for_call); + at::Tensor o = at::empty_like(v); + at::Tensor final_state_work = at::empty({seq_num, HV, K, V}, q.options().dtype(at::kFloat)); + at::Tensor aqk = is_rank3 ? (is_internal_layout ? at::empty({HV, T, chunk_size}, q.options()) : + at::empty({T, HV, chunk_size}, q.options())) : (is_internal_layout ? + at::empty({B, HV, T, chunk_size}, q.options()) : at::empty({B, T, HV, chunk_size}, q.options())); + at::Tensor akk = at::empty_like(aqk); + at::Tensor w = is_rank3 ? (is_internal_layout ? at::empty({HV, T, K}, q.options()) : + at::empty({T, HV, K}, q.options())) : (is_internal_layout ? + at::empty({B, HV, T, K}, q.options()) : at::empty({B, T, HV, K}, q.options())); + at::Tensor u = at::empty_like(v); + at::Tensor qg = at::empty_like(w); + at::Tensor kg = at::empty_like(w); + at::Tensor v_new = at::empty_like(v); + at::Tensor h = is_rank3 ? (is_internal_layout ? at::empty({HV, total_chunks, K, V}, q.options()) : + at::empty({total_chunks, HV, K, V}, q.options())) : (is_internal_layout ? + at::empty({B, HV, total_chunks, K, V}, q.options()) : + at::empty({B, total_chunks, HV, K, V}, q.options())); + + bool recompute_output_final_state = true; + char *layout_cstr = const_cast(layout_str.c_str()); + EXEC_NPU_CMD( + aclnnChunkKdaFwd, + q, k, v, gk, beta, initial_state_tensor, + cu_seqlens, chunk_indices_for_call, + layout_cstr, scale, chunk_size, recompute_output_final_state, total_chunks, + o, final_state_work, aqk, akk, w, u, qg, kg, v_new, h + ); + + at::Tensor final_state = output_final_state.value_or(false) ? + final_state_work : at::empty({0}, q.options().dtype(at::kFloat)); + at::Tensor empty = at::empty({0}, q.options()); + at::Tensor g = gk.scalar_type() == at::kFloat ? gk : gk.to(at::kFloat); + at::Tensor initial_state_out = initial_state_tensor.defined() ? initial_state_tensor : empty; + (void)return_intermediate; + return std::make_tuple(o, final_state, g, aqk, akk, w, u, qg, kg, v_new, h, initial_state_out); +} + +} // namespace vllm_ascend + +#endif diff --git a/csrc/attention/kda_gate_cumsum/kda_gate_cumsum_torch_adpt.h b/csrc/attention/kda_gate_cumsum/kda_gate_cumsum_torch_adpt.h new file mode 100644 index 000000000000..ce9772eea260 --- /dev/null +++ b/csrc/attention/kda_gate_cumsum/kda_gate_cumsum_torch_adpt.h @@ -0,0 +1,74 @@ +/* + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ +#ifndef KDA_GATE_CUMSUM_TORCH_ADPT_H +#define KDA_GATE_CUMSUM_TORCH_ADPT_H + +#include + +#include "attention/kda_torch_adpt_common.h" + +namespace vllm_ascend { + +at::Tensor kda_gate_cumsum( + const at::Tensor &g, + int64_t chunk_size, + const c10::optional &A_log, + const c10::optional &dt_bias, + c10::optional cu_seqlens, + c10::optional use_gate_in_kernel, + c10::optional safe_gate, + c10::optional lower_bound, + c10::string_view layout) +{ + TORCH_CHECK(g.dim() == 3 || g.dim() == 4, + "kda_gate_cumsum: g must be BSND/BNSD rank4 or TND/NTD rank3."); + TORCH_CHECK(chunk_size == 32 || chunk_size == 64 || chunk_size == 128, + "kda_gate_cumsum: chunk_size must be 32, 64 or 128."); + auto gate_dtype = g.scalar_type(); + TORCH_CHECK(gate_dtype == at::kFloat || gate_dtype == at::kBFloat16 || gate_dtype == at::kHalf, + "kda_gate_cumsum: g must be float32, bfloat16 or float16."); + std::string layout_str(layout.data(), layout.size()); + TORCH_CHECK(layout_str == "BSND" || layout_str == "BNSD" || layout_str == "TND" || layout_str == "NTD", + "kda_gate_cumsum: layout must be uppercase and one of BSND, BNSD, TND or NTD."); + bool is_bnsd = layout_str == "BNSD"; + bool is_ntd = layout_str == "NTD"; + int64_t T = is_bnsd ? g.sizes()[2] : (is_ntd ? g.sizes()[1] : (g.dim() == 4 ? g.sizes()[1] : g.sizes()[0])); + int64_t K = g.dim() == 4 ? g.sizes()[3] : g.sizes()[2]; + int64_t HV = is_bnsd ? g.sizes()[1] : (is_ntd ? g.sizes()[0] : (g.dim() == 4 ? g.sizes()[2] : g.sizes()[1])); + TORCH_CHECK(K <= 256, "kda_gate_cumsum: K must be <= 256."); + check_kda_cu_seqlens(cu_seqlens, T, "kda_gate_cumsum"); + TORCH_CHECK(!cu_seqlens.has_value() || g.dim() == 3 || g.sizes()[0] == 1, + "kda_gate_cumsum: rank4 varlen input with cu_seqlens currently requires B=1."); + + bool use_gate = use_gate_in_kernel.value_or(false); + bool safe = safe_gate.value_or(false); + double lower = lower_bound.value_or(-5.0); + at::Tensor A_log_tensor = A_log.value_or(at::Tensor()); + at::Tensor dt_bias_tensor = dt_bias.value_or(at::Tensor()); + if (use_gate) { + TORCH_CHECK(A_log_tensor.defined(), "kda_gate_cumsum: A_log is required when use_gate_in_kernel=True."); + TORCH_CHECK(A_log_tensor.scalar_type() == at::kFloat && + A_log_tensor.dim() == 1 && A_log_tensor.sizes()[0] == HV, + "kda_gate_cumsum: A_log must be float32 with shape [HV]."); + TORCH_CHECK(safe, "kda_gate_cumsum: raw gate path currently requires safe_gate=True."); + TORCH_CHECK(lower >= -5.0 && lower < 0.0, "kda_gate_cumsum: lower_bound must be in [-5, 0)."); + } else { + TORCH_CHECK(!safe, "kda_gate_cumsum: safe_gate only applies when use_gate_in_kernel=True."); + } + + at::Tensor gk = at::empty(g.sizes(), g.options().dtype(at::kFloat)); + char *layout_cstr = const_cast(layout_str.c_str()); + EXEC_NPU_CMD( + aclnnKdaGateCumsum, + g, A_log_tensor, dt_bias_tensor, cu_seqlens, + chunk_size, use_gate, safe, lower, layout_cstr, gk + ); + return gk; +} + +} // namespace vllm_ascend + +#endif diff --git a/csrc/attention/kda_layout_swap12/kda_layout_swap12_torch_adpt.h b/csrc/attention/kda_layout_swap12/kda_layout_swap12_torch_adpt.h new file mode 100644 index 000000000000..6704b4b412a0 --- /dev/null +++ b/csrc/attention/kda_layout_swap12/kda_layout_swap12_torch_adpt.h @@ -0,0 +1,43 @@ +/* + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ +#ifndef KDA_LAYOUT_SWAP12_TORCH_ADPT_H +#define KDA_LAYOUT_SWAP12_TORCH_ADPT_H + +#include + +namespace vllm_ascend { + +at::Tensor kda_layout_swap12( + const at::Tensor &x, + const c10::optional &dependency) +{ + TORCH_CHECK(x.dim() >= 3, "kda_layout_swap12: x must have rank >= 3."); + auto dtype = x.scalar_type(); + TORCH_CHECK(dtype == at::kFloat || dtype == at::kHalf || dtype == at::kBFloat16, + "kda_layout_swap12: x must be float32, float16 or bfloat16."); + + std::vector y_sizes(x.sizes().begin(), x.sizes().end()); + // Swap logical axes 1 and 2 for rank-4 tensors. For the rank-3 TND/NTD + // layouts, the equivalent operation swaps axes 0 and 1. + if (x.dim() == 3) { + std::swap(y_sizes[0], y_sizes[1]); + } else { + std::swap(y_sizes[1], y_sizes[2]); + } + at::Tensor y = at::empty(y_sizes, x.options()); + at::Tensor dependency_tensor = dependency.value_or(at::Tensor()); + if (dependency_tensor.defined()) { + TORCH_CHECK(dependency_tensor.sizes() == y.sizes(), + "kda_layout_swap12: dependency must have the same shape as output."); + } + + EXEC_NPU_CMD(aclnnKdaLayoutSwap12, x, dependency_tensor, y); + return y; +} + +} // namespace vllm_ascend + +#endif diff --git a/csrc/attention/kda_layout_swap12/op_host/CMakeLists.txt b/csrc/attention/kda_layout_swap12/op_host/CMakeLists.txt index 0095d7a4e283..b3360fa4be5e 100644 --- a/csrc/attention/kda_layout_swap12/op_host/CMakeLists.txt +++ b/csrc/attention/kda_layout_swap12/op_host/CMakeLists.txt @@ -12,6 +12,8 @@ if (BUILD_OPEN_PROJECT) endif() add_modules_sources(OPTYPE kda_layout_swap12 ACLNNTYPE aclnn_exclude) +# The kernel explicitly orders MTE2/MTE3 with SetFlag/WaitFlag. Disable +# compiler-inserted synchronization so it does not duplicate those barriers. add_ops_compile_options( OP_NAME KdaLayoutSwap12 OPTIONS --cce-auto-sync=off diff --git a/csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_def.cpp b/csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_def.cpp index 1e6c7cd6236a..0f6b45f5c003 100644 --- a/csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_def.cpp +++ b/csrc/attention/kda_layout_swap12/op_host/kda_layout_swap12_def.cpp @@ -8,6 +8,8 @@ #include "register/op_def_registry.h" namespace ops { +// Swap logical axes 1 and 2 for rank-4 KDA tensors. The rank-3 TND/NTD +// equivalent swaps axes 0 and 1, so both layouts share the same operator. class KdaLayoutSwap12 : public OpDef { public: explicit KdaLayoutSwap12(const char *name) : OpDef(name) diff --git a/csrc/attention/kda_torch_adpt_common.h b/csrc/attention/kda_torch_adpt_common.h new file mode 100644 index 000000000000..74d592963291 --- /dev/null +++ b/csrc/attention/kda_torch_adpt_common.h @@ -0,0 +1,121 @@ +/* + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * SPDX-License-Identifier: Apache-2.0 + * SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project + */ +#ifndef KDA_TORCH_ADPT_COMMON_H +#define KDA_TORCH_ADPT_COMMON_H + +#include + +namespace vllm_ascend { +namespace { + +int64_t kda_ceil_div(int64_t x, int64_t y) +{ + return (x + y - 1) / y; +} + +int64_t get_kda_seq_num(int64_t batch, const c10::optional &cu_seqlens) +{ + if (!cu_seqlens.has_value()) { + return batch; + } + return static_cast(cu_seqlens.value().size()) - 1; +} + +void check_kda_cu_seqlens(const c10::optional &cu_seqlens, + int64_t total_tokens, + const char *op_name) +{ + if (!cu_seqlens.has_value()) { + return; + } + auto cu = cu_seqlens.value(); + TORCH_CHECK(cu.size() >= 2, op_name, ": cu_seqlens must contain at least [0, total_tokens]."); + TORCH_CHECK(cu[0] == 0, op_name, ": cu_seqlens[0] must be 0, but got ", cu[0], "."); + TORCH_CHECK(cu[cu.size() - 1] == total_tokens, + op_name, ": cu_seqlens[-1] must equal sequence length ", + total_tokens, ", but got ", cu[cu.size() - 1], "."); + for (size_t i = 0; i + 1 < cu.size(); ++i) { + TORCH_CHECK(cu[i] <= cu[i + 1], + op_name, ": cu_seqlens must be nondecreasing, but cu_seqlens[", + i, "]=", cu[i], " > cu_seqlens[", i + 1, "]=", cu[i + 1], "."); + } +} + +void check_kda_chunk_indices(const c10::optional &chunk_indices, + const c10::optional &cu_seqlens, + int64_t chunk_size, + const char *op_name) +{ + if (!chunk_indices.has_value()) { + return; + } + auto indices = chunk_indices.value(); + TORCH_CHECK(indices.size() % 2 == 0, + op_name, ": chunk_indices must contain (seq_id, chunk_id) pairs, but got ", + indices.size(), " elements."); + TORCH_CHECK(cu_seqlens.has_value(), op_name, ": chunk_indices requires cu_seqlens."); + auto cu = cu_seqlens.value(); + int64_t expected_chunks = 0; + for (size_t seq = 0; seq + 1 < cu.size(); ++seq) { + expected_chunks += kda_ceil_div(cu[seq + 1] - cu[seq], chunk_size); + } + TORCH_CHECK(static_cast(indices.size() / 2) == expected_chunks, + op_name, ": chunk_indices must contain exactly one pair per chunk."); + for (size_t idx = 0; idx < indices.size(); idx += 2) { + int64_t seq = indices[idx]; + int64_t chunk = indices[idx + 1]; + TORCH_CHECK(seq >= 0 && seq + 1 < static_cast(cu.size()), + op_name, ": chunk_indices seq_id is out of range."); + int64_t chunks = kda_ceil_div(cu[seq + 1] - cu[seq], chunk_size); + TORCH_CHECK(chunk >= 0 && chunk < chunks, + op_name, ": chunk_indices chunk_id is out of range."); + } +} + +int64_t get_kda_total_chunks(int64_t batch, + int64_t seqlen, + int64_t chunk_size, + const c10::optional &cu_seqlens, + const c10::optional &chunk_indices) +{ + if (chunk_indices.has_value()) { + return static_cast(chunk_indices.value().size()) / 2; + } + if (!cu_seqlens.has_value()) { + return kda_ceil_div(seqlen, chunk_size); + } + (void)batch; + int64_t total = 0; + auto cu = cu_seqlens.value(); + for (size_t i = 0; i + 1 < cu.size(); ++i) { + total += kda_ceil_div(cu[i + 1] - cu[i], chunk_size); + } + return total; +} + +std::vector build_kda_chunk_indices(at::IntArrayRef cu_seqlens, int64_t chunk_size) +{ + std::vector indices; + int64_t total_chunks = 0; + for (size_t i = 0; i + 1 < cu_seqlens.size(); ++i) { + total_chunks += kda_ceil_div(cu_seqlens[i + 1] - cu_seqlens[i], chunk_size); + } + indices.reserve(static_cast(total_chunks * 2)); + for (size_t seq = 0; seq + 1 < cu_seqlens.size(); ++seq) { + int64_t seq_len = cu_seqlens[seq + 1] - cu_seqlens[seq]; + int64_t chunks = kda_ceil_div(seq_len, chunk_size); + for (int64_t chunk = 0; chunk < chunks; ++chunk) { + indices.push_back(static_cast(seq)); + indices.push_back(chunk); + } + } + return indices; +} + +} // namespace +} // namespace vllm_ascend + +#endif diff --git a/csrc/moe/situ_mx_quant/op_kernel/inc/platform.h b/csrc/moe/situ_mx_quant/op_kernel/inc/platform.h index 19d59e8c04a0..0110abd178e0 100644 --- a/csrc/moe/situ_mx_quant/op_kernel/inc/platform.h +++ b/csrc/moe/situ_mx_quant/op_kernel/inc/platform.h @@ -41,7 +41,8 @@ __aicore__ inline constexpr bool IsDataCopyPadSupport() } /** - * Get the block size of unified buffer in bytes + * Get the UB data-movement alignment block size in bytes. The total UB + * capacity is queried by the host tiling code through GetCoreMemSize. */ __aicore__ inline constexpr uint32_t GetUbBlockSize() { return 32U; } @@ -78,4 +79,4 @@ __aicore__ inline constexpr bool IsDataCopyPadSupport() { return platform::IsDat } // namespace PlatformSocInfo -#endif // OPS_BUILT_IN_OP_ASCENDC_PLATFORM_INFO_H_ \ No newline at end of file +#endif // OPS_BUILT_IN_OP_ASCENDC_PLATFORM_INFO_H_ diff --git a/csrc/torch_binding.cpp b/csrc/torch_binding.cpp index 21d747ae0eb3..d651a11333d9 100644 --- a/csrc/torch_binding.cpp +++ b/csrc/torch_binding.cpp @@ -45,6 +45,9 @@ #include "moe/causal_conv1d_v310/causal_conv1d_310_torch_adpt.h" #include "attention/recurrent_gated_delta_rule/recurrent_gated_delta_rule_torch_adpt.h" #include "attention/recurrent_kda/recurrent_kda_torch_adpt.h" +#include "attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h" +#include "attention/kda_gate_cumsum/kda_gate_cumsum_torch_adpt.h" +#include "attention/kda_layout_swap12/kda_layout_swap12_torch_adpt.h" #include "attention/recurrent_gated_delta_rule_v310/recurrent_gated_delta_rule_310_torch_adpt.h" #include "attention/k2q_csr/k2q_csr_torch_adpt.h" #include "attention/msa_index_score/msa_index_score_torch_adpt.h" @@ -1923,320 +1926,6 @@ at::Tensor npu_sparse_attention_score_prefill( return output; } -namespace { - -int64_t kda_ceil_div(int64_t x, int64_t y) -{ - return (x + y - 1) / y; -} - -int64_t get_kda_seq_num(int64_t batch, const c10::optional &cu_seqlens) -{ - if (!cu_seqlens.has_value()) { - return batch; - } - return static_cast(cu_seqlens.value().size()) - 1; -} - -void check_kda_cu_seqlens(const c10::optional &cu_seqlens, - int64_t total_tokens, - const char *op_name) -{ - if (!cu_seqlens.has_value()) { - return; - } - auto cu = cu_seqlens.value(); - TORCH_CHECK(cu.size() >= 2, op_name, ": cu_seqlens must contain at least [0, total_tokens]."); - TORCH_CHECK(cu[0] == 0, op_name, ": cu_seqlens[0] must be 0, but got ", cu[0], "."); - TORCH_CHECK(cu[cu.size() - 1] == total_tokens, - op_name, ": cu_seqlens[-1] must equal sequence length ", - total_tokens, ", but got ", cu[cu.size() - 1], "."); - for (size_t i = 0; i + 1 < cu.size(); ++i) { - TORCH_CHECK(cu[i] <= cu[i + 1], - op_name, ": cu_seqlens must be nondecreasing, but cu_seqlens[", - i, "]=", cu[i], " > cu_seqlens[", i + 1, "]=", cu[i + 1], "."); - } -} - -void check_kda_chunk_indices(const c10::optional &chunk_indices, - const c10::optional &cu_seqlens, - int64_t chunk_size, - const char *op_name) -{ - if (!chunk_indices.has_value()) { - return; - } - auto indices = chunk_indices.value(); - TORCH_CHECK(indices.size() % 2 == 0, - op_name, ": chunk_indices must contain (seq_id, chunk_id) pairs, but got ", - indices.size(), " elements."); - TORCH_CHECK(cu_seqlens.has_value(), op_name, ": chunk_indices requires cu_seqlens."); - auto cu = cu_seqlens.value(); - int64_t expected_chunks = 0; - for (size_t seq = 0; seq + 1 < cu.size(); ++seq) { - expected_chunks += kda_ceil_div(cu[seq + 1] - cu[seq], chunk_size); - } - TORCH_CHECK(static_cast(indices.size() / 2) == expected_chunks, - op_name, ": chunk_indices must contain exactly one pair per chunk."); - for (size_t idx = 0; idx < indices.size(); idx += 2) { - int64_t seq = indices[idx]; - int64_t chunk = indices[idx + 1]; - TORCH_CHECK(seq >= 0 && seq + 1 < static_cast(cu.size()), - op_name, ": chunk_indices seq_id is out of range."); - int64_t chunks = kda_ceil_div(cu[seq + 1] - cu[seq], chunk_size); - TORCH_CHECK(chunk >= 0 && chunk < chunks, - op_name, ": chunk_indices chunk_id is out of range."); - } -} - -int64_t get_kda_total_chunks(int64_t batch, - int64_t seqlen, - int64_t chunk_size, - const c10::optional &cu_seqlens, - const c10::optional &chunk_indices) -{ - if (chunk_indices.has_value()) { - return static_cast(chunk_indices.value().size()) / 2; - } - if (!cu_seqlens.has_value()) { - return kda_ceil_div(seqlen, chunk_size); - } - (void)batch; - int64_t total = 0; - auto cu = cu_seqlens.value(); - for (size_t i = 0; i + 1 < cu.size(); ++i) { - total += kda_ceil_div(cu[i + 1] - cu[i], chunk_size); - } - return total; -} - -std::vector build_kda_chunk_indices(at::IntArrayRef cu_seqlens, int64_t chunk_size) -{ - std::vector indices; - int64_t total_chunks = 0; - for (size_t i = 0; i + 1 < cu_seqlens.size(); ++i) { - total_chunks += kda_ceil_div(cu_seqlens[i + 1] - cu_seqlens[i], chunk_size); - } - indices.reserve(static_cast(total_chunks * 2)); - for (size_t seq = 0; seq + 1 < cu_seqlens.size(); ++seq) { - int64_t seq_len = cu_seqlens[seq + 1] - cu_seqlens[seq]; - int64_t chunks = kda_ceil_div(seq_len, chunk_size); - for (int64_t chunk = 0; chunk < chunks; ++chunk) { - indices.push_back(static_cast(seq)); - indices.push_back(chunk); - } - } - return indices; -} - -} // namespace - -std::tuple -chunk_kda_fwd( - const at::Tensor &q, - const at::Tensor &k, - const at::Tensor &v, - const at::Tensor &gk, - const at::Tensor &beta, - double scale, - int64_t chunk_size, - c10::string_view layout, - const c10::optional &initial_state, - c10::optional output_final_state, - c10::optional cu_seqlens, - c10::optional chunk_indices, - c10::optional return_intermediate, - c10::optional safe_gate, - c10::optional transpose_state_layout) -{ - std::string layout_str(layout.data(), layout.size()); - TORCH_CHECK(layout_str == "BSND" || layout_str == "BNSD" || layout_str == "TND" || layout_str == "NTD", - "chunk_kda_fwd: layout must be one of BSND, BNSD, TND, NTD and must be uppercase."); - TORCH_CHECK(!safe_gate.value_or(false), "chunk_kda_fwd: safe_gate=True is not supported."); - TORCH_CHECK(!transpose_state_layout.value_or(false), - "chunk_kda_fwd: transpose_state_layout=True is not supported."); - TORCH_CHECK(chunk_size == 32 || chunk_size == 64 || chunk_size == 128, - "chunk_kda_fwd: chunk_size must be 32, 64 or 128."); - - bool is_tnd = layout_str == "TND"; - bool is_ntd = layout_str == "NTD"; - bool is_bsnd = layout_str == "BSND"; - bool is_bnsd = layout_str == "BNSD"; - bool is_rank3 = is_tnd || is_ntd; - TORCH_CHECK((is_rank3 && q.dim() == 3 && k.dim() == 3 && v.dim() == 3 && gk.dim() == 3 && beta.dim() == 2) || - (!is_rank3 && q.dim() == 4 && k.dim() == 4 && v.dim() == 4 && gk.dim() == 4 && beta.dim() == 3), - "chunk_kda_fwd: layout/rank mismatch."); - TORCH_CHECK(q.sizes() == k.sizes(), "chunk_kda_fwd: q and k must have identical shape."); - - auto q_sizes = q.sizes(); - auto v_sizes = v.sizes(); - bool is_internal_layout = is_bnsd || is_ntd; - int64_t B = is_rank3 ? 1 : q_sizes[0]; - int64_t T = is_tnd ? q_sizes[0] : (is_ntd ? q_sizes[1] : (is_bnsd ? q_sizes[2] : q_sizes[1])); - int64_t H = is_tnd ? q_sizes[1] : (is_ntd ? q_sizes[0] : (is_bnsd ? q_sizes[1] : q_sizes[2])); - int64_t K = is_rank3 ? q_sizes[2] : q_sizes[3]; - int64_t HV = is_tnd ? v_sizes[1] : (is_ntd ? v_sizes[0] : (is_bnsd ? v_sizes[1] : v_sizes[2])); - int64_t V = is_rank3 ? v_sizes[2] : v_sizes[3]; - TORCH_CHECK(H > 0 && HV >= H, "chunk_kda_fwd: H and HV must be positive and H must be <= HV."); - TORCH_CHECK(H <= 128 && HV <= 128, "chunk_kda_fwd: H and HV must be <= 128."); - TORCH_CHECK(!is_tnd || H == 1, - "chunk_kda_fwd: TND layout with H > 1 is not supported; use NTD for multi-head rank3 input."); - check_kda_cu_seqlens(cu_seqlens, T, "chunk_kda_fwd"); - check_kda_chunk_indices(chunk_indices, cu_seqlens, chunk_size, "chunk_kda_fwd"); - TORCH_CHECK(!cu_seqlens.has_value() || is_rank3 || B == 1, - "chunk_kda_fwd: rank4 varlen input with cu_seqlens currently requires B=1."); - TORCH_CHECK(HV % H == 0, "chunk_kda_fwd: HV must be divisible by H."); - TORCH_CHECK(q.scalar_type() == at::kHalf || q.scalar_type() == at::kBFloat16, - "chunk_kda_fwd: q/k/v must use float16 or bfloat16."); - TORCH_CHECK(k.scalar_type() == q.scalar_type() && v.scalar_type() == q.scalar_type(), - "chunk_kda_fwd: q/k/v dtype must match."); - TORCH_CHECK(chunk_size == 64 && K >= 16 && V >= 16 && K % 16 == 0 && V % 16 == 0 && V <= 256 && - K * V >= 4 * 64 * 64 && K * V >= chunk_size * (K + V), - "chunk_kda_fwd: shape is outside the supported split cube/vector template."); - - int64_t seq_num = get_kda_seq_num(B, cu_seqlens); - at::Tensor initial_state_tensor = initial_state.value_or(at::Tensor()); - if (initial_state_tensor.defined()) { - TORCH_CHECK(initial_state_tensor.scalar_type() == at::kFloat, - "chunk_kda_fwd: initial_state must be float32 when provided."); - TORCH_CHECK(initial_state_tensor.dim() == 4 && initial_state_tensor.size(0) == seq_num && - initial_state_tensor.size(1) == HV && initial_state_tensor.size(2) == K && - initial_state_tensor.size(3) == V, - "chunk_kda_fwd: initial_state must be [seq_num,Hv,K,V]."); - } - - std::vector generated_chunk_indices; - c10::optional chunk_indices_for_call; - if (chunk_indices.has_value()) { - chunk_indices_for_call = chunk_indices.value(); - } else if (cu_seqlens.has_value()) { - generated_chunk_indices = build_kda_chunk_indices(cu_seqlens.value(), chunk_size); - chunk_indices_for_call = at::IntArrayRef(generated_chunk_indices); - } else { - chunk_indices_for_call = c10::nullopt; - } - - int64_t total_chunks = get_kda_total_chunks(B, T, chunk_size, cu_seqlens, chunk_indices_for_call); - at::Tensor o = at::empty_like(v); - at::Tensor final_state_work = at::empty({seq_num, HV, K, V}, q.options().dtype(at::kFloat)); - at::Tensor aqk = is_rank3 ? (is_internal_layout ? at::empty({HV, T, chunk_size}, q.options()) : - at::empty({T, HV, chunk_size}, q.options())) : (is_internal_layout ? - at::empty({B, HV, T, chunk_size}, q.options()) : at::empty({B, T, HV, chunk_size}, q.options())); - at::Tensor akk = at::empty_like(aqk); - at::Tensor w = is_rank3 ? (is_internal_layout ? at::empty({HV, T, K}, q.options()) : - at::empty({T, HV, K}, q.options())) : (is_internal_layout ? - at::empty({B, HV, T, K}, q.options()) : at::empty({B, T, HV, K}, q.options())); - at::Tensor u = at::empty_like(v); - at::Tensor qg = at::empty_like(w); - at::Tensor kg = at::empty_like(w); - at::Tensor v_new = at::empty_like(v); - at::Tensor h = is_rank3 ? (is_internal_layout ? at::empty({HV, total_chunks, K, V}, q.options()) : - at::empty({total_chunks, HV, K, V}, q.options())) : (is_internal_layout ? - at::empty({B, HV, total_chunks, K, V}, q.options()) : - at::empty({B, total_chunks, HV, K, V}, q.options())); - - bool recompute_output_final_state = true; - char *layout_cstr = const_cast(layout_str.c_str()); - EXEC_NPU_CMD( - aclnnChunkKdaFwd, - q, k, v, gk, beta, initial_state_tensor, - cu_seqlens, chunk_indices_for_call, - layout_cstr, scale, chunk_size, recompute_output_final_state, total_chunks, - o, final_state_work, aqk, akk, w, u, qg, kg, v_new, h - ); - - at::Tensor final_state = output_final_state.value_or(false) ? - final_state_work : at::empty({0}, q.options().dtype(at::kFloat)); - at::Tensor empty = at::empty({0}, q.options()); - at::Tensor g = gk.scalar_type() == at::kFloat ? gk : gk.to(at::kFloat); - at::Tensor initial_state_out = initial_state_tensor.defined() ? initial_state_tensor : empty; - (void)return_intermediate; - return std::make_tuple(o, final_state, g, aqk, akk, w, u, qg, kg, v_new, h, initial_state_out); -} - -at::Tensor kda_gate_cumsum( - const at::Tensor &g, - int64_t chunk_size, - const c10::optional &A_log, - const c10::optional &dt_bias, - c10::optional cu_seqlens, - c10::optional use_gate_in_kernel, - c10::optional safe_gate, - c10::optional lower_bound, - c10::string_view layout) -{ - TORCH_CHECK(g.dim() == 3 || g.dim() == 4, - "kda_gate_cumsum: g must be BSND/BNSD rank4 or TND/NTD rank3."); - TORCH_CHECK(chunk_size == 32 || chunk_size == 64 || chunk_size == 128, - "kda_gate_cumsum: chunk_size must be 32, 64 or 128."); - auto gate_dtype = g.scalar_type(); - TORCH_CHECK(gate_dtype == at::kFloat || gate_dtype == at::kBFloat16 || gate_dtype == at::kHalf, - "kda_gate_cumsum: g must be float32, bfloat16 or float16."); - std::string layout_str(layout.data(), layout.size()); - TORCH_CHECK(layout_str == "BSND" || layout_str == "BNSD" || layout_str == "TND" || layout_str == "NTD", - "kda_gate_cumsum: layout must be uppercase and one of BSND, BNSD, TND or NTD."); - bool is_bnsd = layout_str == "BNSD"; - bool is_ntd = layout_str == "NTD"; - int64_t T = is_bnsd ? g.sizes()[2] : (is_ntd ? g.sizes()[1] : (g.dim() == 4 ? g.sizes()[1] : g.sizes()[0])); - int64_t K = g.dim() == 4 ? g.sizes()[3] : g.sizes()[2]; - int64_t HV = is_bnsd ? g.sizes()[1] : (is_ntd ? g.sizes()[0] : (g.dim() == 4 ? g.sizes()[2] : g.sizes()[1])); - TORCH_CHECK(K <= 256, "kda_gate_cumsum: K must be <= 256."); - check_kda_cu_seqlens(cu_seqlens, T, "kda_gate_cumsum"); - TORCH_CHECK(!cu_seqlens.has_value() || g.dim() == 3 || g.sizes()[0] == 1, - "kda_gate_cumsum: rank4 varlen input with cu_seqlens currently requires B=1."); - - bool use_gate = use_gate_in_kernel.value_or(false); - bool safe = safe_gate.value_or(false); - double lower = lower_bound.value_or(-5.0); - at::Tensor A_log_tensor = A_log.value_or(at::Tensor()); - at::Tensor dt_bias_tensor = dt_bias.value_or(at::Tensor()); - if (use_gate) { - TORCH_CHECK(A_log_tensor.defined(), "kda_gate_cumsum: A_log is required when use_gate_in_kernel=True."); - TORCH_CHECK(A_log_tensor.scalar_type() == at::kFloat && A_log_tensor.dim() == 1 && A_log_tensor.sizes()[0] == HV, - "kda_gate_cumsum: A_log must be float32 with shape [HV]."); - TORCH_CHECK(safe, "kda_gate_cumsum: raw gate path currently requires safe_gate=True."); - TORCH_CHECK(lower >= -5.0 && lower < 0.0, "kda_gate_cumsum: lower_bound must be in [-5, 0)."); - } else { - TORCH_CHECK(!safe, "kda_gate_cumsum: safe_gate only applies when use_gate_in_kernel=True."); - } - - at::Tensor gk = at::empty(g.sizes(), g.options().dtype(at::kFloat)); - char *layout_cstr = const_cast(layout_str.c_str()); - EXEC_NPU_CMD( - aclnnKdaGateCumsum, - g, A_log_tensor, dt_bias_tensor, cu_seqlens, - chunk_size, use_gate, safe, lower, layout_cstr, gk - ); - return gk; -} - -at::Tensor kda_layout_swap12( - const at::Tensor &x, - const c10::optional &dependency) -{ - TORCH_CHECK(x.dim() >= 3, "kda_layout_swap12: x must have rank >= 3."); - auto dtype = x.scalar_type(); - TORCH_CHECK(dtype == at::kFloat || dtype == at::kHalf || dtype == at::kBFloat16, - "kda_layout_swap12: x must be float32, float16 or bfloat16."); - - std::vector y_sizes(x.sizes().begin(), x.sizes().end()); - if (x.dim() == 3) { - std::swap(y_sizes[0], y_sizes[1]); - } else { - std::swap(y_sizes[1], y_sizes[2]); - } - at::Tensor y = at::empty(y_sizes, x.options()); - at::Tensor dependency_tensor = dependency.value_or(at::Tensor()); - if (dependency_tensor.defined()) { - TORCH_CHECK(dependency_tensor.sizes() == y.sizes(), - "kda_layout_swap12: dependency must have the same shape as output."); - } - - EXEC_NPU_CMD(aclnnKdaLayoutSwap12, x, dependency_tensor, y); - return y; -} - std::vector get_npu_storage_shape(const at::Tensor& tensor) { TORCH_CHECK( From 80282f85172ea34f2752f35ef70ea8699ab05df8 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Mon, 24 Aug 2026 04:20:48 -0500 Subject: [PATCH 03/50] fix(mamba): align hybrid state copy on Ascend Stage Mamba copy metadata during input preparation, execute the device copy after KV loading, and constrain the temporal launch grid to Ascend's coreDim limit. Recognize the upstream Kimi linear-attention configuration without adding model-specific runtime branches. Signed-off-by: maoxx241 --- .../ut/patch/worker/test_patch_mamba_utils.py | 174 ++++++++++++++++++ .../worker/test_patch_mamba_utils_source.py | 12 -- tests/ut/test_utils.py | 52 ++++++ vllm_ascend/patch/worker/patch_mamba_utils.py | 52 +++++- vllm_ascend/utils.py | 20 +- 5 files changed, 281 insertions(+), 29 deletions(-) create mode 100644 tests/ut/patch/worker/test_patch_mamba_utils.py diff --git a/tests/ut/patch/worker/test_patch_mamba_utils.py b/tests/ut/patch/worker/test_patch_mamba_utils.py new file mode 100644 index 000000000000..f43b5b9e0454 --- /dev/null +++ b/tests/ut/patch/worker/test_patch_mamba_utils.py @@ -0,0 +1,174 @@ +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace +from unittest.mock import MagicMock, call, patch + +import numpy as np + +from vllm_ascend.patch.worker.patch_mamba_utils import ( + _do_mamba_copy_block_npu, + _stage_mamba_copy_metadata, + preprocess_mamba, +) + + +def _copy_buffer(): + gpu = MagicMock() + gpu_view = MagicMock() + gpu.__getitem__.return_value = gpu_view + return SimpleNamespace(copy_to_gpu=MagicMock(), gpu=gpu), gpu_view + + +def test_mamba_copy_metadata_is_staged_asynchronously_during_preprocess(): + src_ptrs, _ = _copy_buffer() + dst_ptrs, _ = _copy_buffer() + sizes, _ = _copy_buffer() + copy_bufs = SimpleNamespace( + offset=2, + src_ptrs=src_ptrs, + dst_ptrs=dst_ptrs, + sizes=sizes, + ) + + _stage_mamba_copy_metadata(copy_bufs) + + for buffer in (src_ptrs, dst_ptrs, sizes): + buffer.copy_to_gpu.assert_called_once_with(2) + + +def test_mamba_state_copy_uses_previously_staged_metadata(): + src_ptrs, src_view = _copy_buffer() + dst_ptrs, dst_view = _copy_buffer() + sizes, sizes_view = _copy_buffer() + copy_bufs = SimpleNamespace( + offset=2, + src_ptrs=src_ptrs, + dst_ptrs=dst_ptrs, + sizes=sizes, + ) + + with patch("vllm_ascend.patch.worker.patch_mamba_utils._batch_memcpy_triton") as batch_memcpy: + _do_mamba_copy_block_npu(copy_bufs) + + for buffer in (src_ptrs, dst_ptrs, sizes): + buffer.copy_to_gpu.assert_not_called() + assert buffer.gpu.__getitem__.call_args_list == [call(slice(None, 2))] + batch_memcpy.assert_called_once_with(src_view, dst_view, sizes_view) + + +def test_preprocess_stages_metadata_but_defers_state_copy(): + copy_bufs = SimpleNamespace( + offset=0, + mamba_group_ids=[0], + mamba_spec=SimpleNamespace(num_speculative_blocks=1, block_size=7), + ) + scheduler_output = SimpleNamespace( + finished_req_ids=[], + preempted_req_ids=set(), + scheduled_cached_reqs=SimpleNamespace(resumed_req_ids=[]), + num_scheduled_tokens={"req": 7}, + ) + input_batch = SimpleNamespace( + req_ids=["req"], + num_accepted_tokens_cpu=np.array([2], dtype=np.int32), + ) + requests = {"req": SimpleNamespace(num_computed_tokens=7)} + mamba_state_idx = {"req": 0} + + def collect_metadata(copy_buffers, *_args): + copy_buffers.offset = 1 + + with ( + patch( + "vllm_ascend.patch.worker.patch_mamba_utils.mamba_utils.collect_mamba_copy_meta", + side_effect=collect_metadata, + ) as collect, + patch("vllm_ascend.patch.worker.patch_mamba_utils._can_launch_triton_batch_memcpy", return_value=True), + patch("vllm_ascend.patch.worker.patch_mamba_utils._stage_mamba_copy_metadata") as stage, + patch("vllm_ascend.patch.worker.patch_mamba_utils._do_mamba_copy_block_npu") as state_copy, + ): + preprocess_mamba( + scheduler_output, + SimpleNamespace(), + SimpleNamespace(), + mamba_state_idx, + input_batch, + requests, + {}, + (), + copy_bufs, + ) + + collect.assert_called_once() + stage.assert_called_once_with(copy_bufs) + state_copy.assert_not_called() + assert input_batch.num_accepted_tokens_cpu.tolist() == [1] + + +def test_load_only_step_does_not_hide_remote_state_copy_on_next_forward(): + copy_bufs = SimpleNamespace( + offset=0, + mamba_group_ids=[0], + mamba_spec=SimpleNamespace(num_speculative_blocks=7, block_size=128), + ) + scheduler_output = SimpleNamespace( + finished_req_ids=[], + preempted_req_ids=set(), + scheduled_cached_reqs=SimpleNamespace(resumed_req_ids=[]), + num_scheduled_tokens={"req": 0}, + ) + input_batch = SimpleNamespace( + req_ids=["req"], + num_accepted_tokens_cpu=np.array([1], dtype=np.int32), + ) + requests = {"req": SimpleNamespace(num_computed_tokens=0)} + mamba_state_idx: dict[str, int] = {} + + with ( + patch("vllm_ascend.patch.worker.patch_mamba_utils.mamba_utils.collect_mamba_copy_meta") as collect, + patch( + "vllm_ascend.patch.worker.patch_mamba_utils._can_launch_triton_batch_memcpy", + return_value=True, + ), + patch("vllm_ascend.patch.worker.patch_mamba_utils._stage_mamba_copy_metadata") as stage, + ): + preprocess_mamba( + scheduler_output, + SimpleNamespace(), + SimpleNamespace(), + mamba_state_idx, + input_batch, + requests, + {}, + (), + copy_bufs, + ) + + assert "req" not in mamba_state_idx + collect.assert_not_called() + stage.assert_called_once_with(copy_bufs) + + scheduler_output.num_scheduled_tokens["req"] = 8 + requests["req"].num_computed_tokens = 8191 + stage.reset_mock() + + def collect_metadata(copy_buffers, *_args): + copy_buffers.offset = 1 + + collect.side_effect = collect_metadata + preprocess_mamba( + scheduler_output, + SimpleNamespace(), + SimpleNamespace(), + mamba_state_idx, + input_batch, + requests, + {}, + (), + copy_bufs, + ) + + collect.assert_called_once() + assert collect.call_args.args[4:7] == (63, 64, 0) + stage.assert_called_once_with(copy_bufs) + assert mamba_state_idx["req"] == 64 diff --git a/tests/ut/patch/worker/test_patch_mamba_utils_source.py b/tests/ut/patch/worker/test_patch_mamba_utils_source.py index e0efa4d041bf..88dc71282be5 100644 --- a/tests/ut/patch/worker/test_patch_mamba_utils_source.py +++ b/tests/ut/patch/worker/test_patch_mamba_utils_source.py @@ -10,8 +10,6 @@ ROOT = Path(__file__).resolve().parents[4] POSTPROCESS = ROOT / "vllm_ascend" / "ops" / "triton" / "mamba" / "postprocess.py" -PATCH_MAMBA_UTILS = ROOT / "vllm_ascend" / "patch" / "worker" / "patch_mamba_utils.py" -PATCH_TRITON = ROOT / "vllm_ascend" / "patch" / "worker" / "patch_v2" / "patch_triton.py" def _top_level_functions(path: Path) -> dict[str, ast.FunctionDef]: @@ -61,13 +59,3 @@ def test_postprocess_keeps_only_existing_ascend_precision_kernel() -> None: assert "tile_idx" in postprocess_source assert "if tile_idx == 0:" in postprocess_source assert "and state_idx == 0 and tile_idx == 0" not in postprocess_source - - -def test_patch_only_installs_existing_ascend_postprocess_kernel() -> None: - patch_source = PATCH_MAMBA_UTILS.read_text() - patch_triton = PATCH_TRITON.read_text() - - assert "mamba_utils.postprocess_mamba_fused_kernel = postprocess_mamba_fused_kernel" in patch_source - assert "MambaBase.bind_kv_cache" not in patch_source - assert "mamba_utils._copy_mamba_state_block" not in patch_source - assert "mamba_utils.precopy_mamba_align_fused_kernel" in patch_triton diff --git a/tests/ut/test_utils.py b/tests/ut/test_utils.py index 3923b1645d44..e4e30e040a68 100644 --- a/tests/ut/test_utils.py +++ b/tests/ut/test_utils.py @@ -505,3 +505,55 @@ def test_is_pd_decode_recompute_scheduler_enabled_decode_consumer_disabled(): ascend_config.scheduler_config.recompute_scheduler_enable = False with mock.patch("vllm_ascend.utils.get_ascend_config", return_value=ascend_config): assert utils.is_pd_decode_recompute_scheduler_enabled(vllm_config) is False + + +def test_check_gdn_layer_supports_kimi_linear_config_property(): + from vllm.transformers_utils.configs.kimi_linear import KimiLinearConfig + + vllm_config = SimpleNamespace( + model_config=SimpleNamespace( + hf_config=SimpleNamespace( + text_config=KimiLinearConfig( + linear_attn_config={ + "kda_layers": [2], + "full_attn_layers": [1], + } + ) + ) + ) + ) + + assert utils.check_gdn_layer(vllm_config) is True + + +@pytest.mark.parametrize( + "hf_config", + [ + pytest.param( + SimpleNamespace(text_config=SimpleNamespace(layer_types=["linear_attention"])), + id="qwen3-5-nested-text-config", + ), + ], +) +def test_check_gdn_layer_supports_layer_types(hf_config): + vllm_config = SimpleNamespace(model_config=SimpleNamespace(hf_config=hf_config)) + + assert utils.check_gdn_layer(vllm_config) is True + + +def test_check_gdn_layer_supports_qwen3_next_config(): + from transformers import Qwen3NextConfig + + vllm_config = SimpleNamespace(model_config=SimpleNamespace(hf_config=Qwen3NextConfig())) + + assert utils.check_gdn_layer(vllm_config) is True + + +def test_check_gdn_layer_returns_false_without_linear_attention(): + from transformers import Qwen3Config + + # Dense Qwen3 configs, including Qwen3-8B, expose only full-attention + # layer types and must not be classified as hybrid GDN models. + vllm_config = SimpleNamespace(model_config=SimpleNamespace(hf_config=Qwen3Config())) + + assert utils.check_gdn_layer(vllm_config) is False diff --git a/vllm_ascend/patch/worker/patch_mamba_utils.py b/vllm_ascend/patch/worker/patch_mamba_utils.py index ae359e070ccf..3b241ec965f6 100644 --- a/vllm_ascend/patch/worker/patch_mamba_utils.py +++ b/vllm_ascend/patch/worker/patch_mamba_utils.py @@ -18,6 +18,14 @@ from vllm_ascend.ops.triton.mamba.postprocess import postprocess_mamba_fused_kernel from vllm_ascend.utils import is_310p +# Upstream uses 16 temporal-copy tiles to saturate H100/GB200. K3 already +# exposes 138 independent state programs per request, while Triton-Ascend +# flattens all grid dimensions into a coreDim that cannot exceed 65535. Keep +# the pre-tiling launch shape on Ascend: it has enough state-level parallelism +# and remains valid at the configured request limit (for example, +# 32 * 138 * 1 instead of 32 * 138 * 16). +mamba_utils._TEMPORAL_TILES = 1 + def _can_launch_triton_batch_memcpy() -> bool: return not is_310p() @@ -34,6 +42,28 @@ def _batch_memcpy_triton(src_ptrs, dst_ptrs, sizes): batch_memcpy_kernel[grid](src_ptrs, dst_ptrs, sizes, BLOCK_SIZE=BLOCK_SIZE) +def _stage_mamba_copy_metadata(copy_bufs: mamba_utils.MambaCopyBuffers) -> None: + """Stage pointer metadata while input-preparation buffers are protected.""" + n = copy_bufs.offset + if n == 0: + return + copy_bufs.src_ptrs.copy_to_gpu(n) + copy_bufs.dst_ptrs.copy_to_gpu(n) + copy_bufs.sizes.copy_to_gpu(n) + + +def _do_mamba_copy_block_npu(copy_bufs: mamba_utils.MambaCopyBuffers) -> None: + """Copy state after KV load using metadata staged during preprocessing.""" + n = copy_bufs.offset + if n == 0: + return + _batch_memcpy_triton( + copy_bufs.src_ptrs.gpu[:n], + copy_bufs.dst_ptrs.gpu[:n], + copy_bufs.sizes.gpu[:n], + ) + + def _tensor_view_from_data_ptr(state: torch.Tensor, start_addr: int, num_elements: int) -> torch.Tensor: byte_offset = start_addr - state.data_ptr() element_size = state.element_size() @@ -195,8 +225,7 @@ def _batch_memcpy_unavailable(src_ptrs, dst_ptrs, sizes): if _can_launch_triton_batch_memcpy(): mamba_utils.batch_memcpy_kernel = batch_memcpy_kernel mamba_utils.batch_memcpy = _batch_memcpy_triton - # Keep the existing Ascend postprocess precision fix. The shared copy - # helper and align pre-copy continue to use the upstream implementation. + mamba_utils.do_mamba_copy_block = _do_mamba_copy_block_npu mamba_utils.postprocess_mamba_fused_kernel = postprocess_mamba_fused_kernel else: mamba_utils.batch_memcpy = _batch_memcpy_unavailable @@ -253,13 +282,22 @@ def preprocess_mamba( copy_bufs.offset = 0 for i, req_id in enumerate(input_batch.req_ids): req_state = requests[req_id] + num_scheduled_tokens = scheduler_output.num_scheduled_tokens[req_id] + if num_scheduled_tokens == 0: + # Async KV connectors can surface a request in a load-only step + # before any model tokens are scheduled. Persisting the derived + # ``-1`` state index here makes the next real forward skip the + # copy from the remotely loaded h(N-1) state into its running + # block. Re-resolve the index from the updated computed-token + # count when the request is actually scheduled instead. + mamba_state_idx.pop(req_id, None) + continue prev_state_idx = mamba_state_idx.get(req_id) if prev_state_idx is None: # new / resumed request, no previous state # if num_computed_tokens is 0, prev_state_idx will be -1 prev_state_idx = (req_state.num_computed_tokens - 1) // block_size - num_scheduled_tokens = scheduler_output.num_scheduled_tokens[req_id] num_blocks: int = ( cdiv(req_state.num_computed_tokens + num_scheduled_tokens, block_size) + num_speculative_blocks ) @@ -289,8 +327,12 @@ def preprocess_mamba( forward_context, ) input_batch.num_accepted_tokens_cpu[i] = 1 - # do not copy here, since kv_transfer still not load - # do_mamba_copy_block(copy_bufs) + if _can_launch_triton_batch_memcpy(): + # Only stage the pointer table here. This runs inside the existing + # input-preparation event scope, so its pinned CPU buffers cannot be + # reused until the asynchronous H2D copies finish. The state copy must + # remain after KV transfer and is executed by do_mamba_copy_block(). + _stage_mamba_copy_metadata(copy_bufs) mamba_utils.preprocess_mamba = preprocess_mamba diff --git a/vllm_ascend/utils.py b/vllm_ascend/utils.py index 8841b2458f64..9555f27774dc 100644 --- a/vllm_ascend/utils.py +++ b/vllm_ascend/utils.py @@ -1339,8 +1339,7 @@ def enable_dsa_cp_with_o_proj_tp() -> bool: def check_gdn_layer(vllm_config) -> bool: """ - gdn layer is marked with `linear_attention`. - So, if `linear_attention` is detected, we think the model has gdn-attention. + Detect a model with GDN attention from either supported HF config shape. """ if not hasattr(vllm_config, "model_config"): return False @@ -1350,16 +1349,13 @@ def check_gdn_layer(vllm_config) -> bool: return False hf_config = model_config.hf_config - - # Use `or []` to prevent errors when layer_types is None - layer_types = getattr(hf_config, "layer_types", None) or [] - if "linear_attention" in layer_types: - return True - - text_config = getattr(hf_config, "text_config", None) - if text_config: - text_layer_types = getattr(text_config, "layer_types", None) or [] - if "linear_attention" in text_layer_types: + for config in (hf_config, getattr(hf_config, "text_config", None)): + if config is None: + continue + # Most hybrid models expose layer_types. Kimi Linear/K3 instead + # exposes the equivalent is_linear_attn property. + layer_types = getattr(config, "layer_types", None) or [] + if "linear_attention" in layer_types or bool(getattr(config, "is_linear_attn", False)): return True return False From 814c2fd0597b33d0326c533bc72bf374e5d2eafa Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Mon, 24 Aug 2026 04:21:54 -0500 Subject: [PATCH 04/50] feat(attention): add Kimi K3 KDA execution Implement Ascend KDA prefill and recurrent decode with the existing GDN metadata path, AscendC kernels, graph padding, mixed-precision projection loading, and focused boundary tests. Signed-off-by: maoxx241 --- tests/ut/ops/test_gdn_attn_builder.py | 298 +++++++++++++- tests/ut/ops/test_kimi_kda.py | 271 ++++++++++++ vllm_ascend/ops/gdn_attn_builder.py | 158 ++++++- vllm_ascend/ops/kimi_kda.py | 570 ++++++++++++++++++++++++++ 4 files changed, 1271 insertions(+), 26 deletions(-) create mode 100644 tests/ut/ops/test_kimi_kda.py create mode 100644 vllm_ascend/ops/kimi_kda.py diff --git a/tests/ut/ops/test_gdn_attn_builder.py b/tests/ut/ops/test_gdn_attn_builder.py index 909be456a385..0b0a28191b6c 100644 --- a/tests/ut/ops/test_gdn_attn_builder.py +++ b/tests/ut/ops/test_gdn_attn_builder.py @@ -9,6 +9,7 @@ from vllm.config.compilation import CUDAGraphMode from vllm.third_party.flash_linear_attention.ops import index as _fla_index from vllm.v1.attention.backend import CommonAttentionMetadata +from vllm.v1.attention.backends.utils import NULL_BLOCK_ID, PAD_SLOT_ID from vllm.v1.kv_cache_interface import MambaSpec from vllm_ascend.attention.utils import AscendCommonAttentionMetadata @@ -226,7 +227,10 @@ def _build_attn_metadata( def _assert_chunk_meta_matches_runtime(builder, chunk_meta, cu_seqlens: torch.Tensor) -> None: hf_text_config = getattr(builder.vllm_config.model_config, "hf_text_config", None) - if hf_text_config is not None and hasattr(hf_text_config, "linear_num_value_heads"): + linear_attn_config = getattr(hf_text_config, "linear_attn_config", None) + if isinstance(linear_attn_config, dict) and linear_attn_config.get("num_heads") is not None: + gdn_num_heads = linear_attn_config["num_heads"] // builder.vllm_config.parallel_config.tensor_parallel_size + elif hf_text_config is not None and hasattr(hf_text_config, "linear_num_value_heads"): gdn_num_heads = ( hf_text_config.linear_num_value_heads // builder.vllm_config.parallel_config.tensor_parallel_size ) @@ -287,6 +291,26 @@ def _patch_missing_runtime_cdiv(monkeypatch: pytest.MonkeyPatch) -> None: ) +def test_kimi_chunk_metadata_uses_linear_attention_head_count() -> None: + builder = _make_builder( + device=torch.device("cpu"), + num_heads=128, + num_speculative_tokens=0, + ) + builder.vllm_config.model_config.hf_text_config = SimpleNamespace( + linear_attn_config={"num_heads": 32}, + ) + cu_seqlens = torch.tensor([0, 130], dtype=torch.int32) + + chunk_meta = ascend_gdn_attn_builder._build_non_spec_chunked_prefill_metadata( + builder, + cu_seqlens, + torch.device("cpu"), + ) + + _assert_chunk_meta_matches_runtime(builder, chunk_meta, cu_seqlens) + + def test_ascend_gdn_attention_uses_ascend_backend(): assert AscendGatedDeltaNetAttention.get_attn_backend(object()) is AscendGDNAttentionBackend assert AscendGDNAttentionBackend.get_builder_cls() is AscendGDNAttentionMetadataBuilder @@ -431,7 +455,6 @@ def test_non_spec_prefill_metadata_uses_prefill_tail_for_chunk_metadata( assert torch.equal(conv1d_meta.query_start_loc, torch.tensor([0, 1, 9, 13], dtype=torch.int32)) assert torch.equal(_cache_index_first_column(conv1d_meta.cache_indices), torch.tensor([0, 1, 2], dtype=torch.int32)) assert torch.equal(conv1d_meta.initial_state_mode, torch.tensor([True, True, True])) - assert prefill_metadata.chunk.num_decodes == 0 _assert_chunk_meta_matches_runtime( builder, prefill_metadata.chunk, @@ -570,8 +593,8 @@ def test_full_graph_spec_actual_seq_lengths_use_padded_builder_buffer(): def test_full_graph_non_spec_actual_seq_lengths_use_padded_builder_buffer(): batch_spec = BatchSpec( - seq_lens=[1, 1], - query_lens=[1, 1], + seq_lens=[1, 1, 0, 0], + query_lens=[1, 1, 0, 0], name="full_graph_padded_non_spec_actual_seq_lengths", ) common_attn_metadata = create_common_attn_metadata( @@ -579,7 +602,6 @@ def test_full_graph_non_spec_actual_seq_lengths_use_padded_builder_buffer(): block_size=16, device=torch.device("cpu"), ) - common_attn_metadata.num_actual_tokens = 4 builder = _make_builder( device=torch.device("cpu"), num_heads=32, @@ -759,3 +781,269 @@ def test_builder_skips_prebuilt_meta_without_non_spec_prefill(batch_spec: BatchS spec_decode_metadata.actual_seq_lengths, torch.tensor([0, 4, 4], dtype=torch.int32), ) + + +def test_mixed_spec_prefill_chunk_metadata_preserves_single_token_count( + monkeypatch: pytest.MonkeyPatch, +): + _patch_missing_runtime_cdiv(monkeypatch) + batch_spec = BatchSpec( + seq_lens=[1, 4, 8], + query_lens=[1, 4, 8], + name="mixed_spec_prefill_with_single_token_non_spec", + ) + builder, _, attn_metadata = _build_attn_metadata( + batch_spec, + num_speculative_tokens=3, + num_decode_draft_tokens_cpu=torch.tensor([-1, 3, -1], dtype=torch.int32), + ) + + assert attn_metadata.num_decodes == 0 + assert attn_metadata.num_prefills == 2 + assert torch.equal( + attn_metadata.prefill_query_start_loc, + torch.tensor([0, 1, 9], dtype=torch.int32), + ) + chunk_metadata = attn_metadata.non_spec_prefill_metadata.chunk + _assert_chunk_meta_matches_runtime( + builder, + chunk_metadata, + attn_metadata.prefill_query_start_loc, + ) + + +@pytest.mark.parametrize( + ("seq_len", "expected_decodes", "expected_prefills"), + [ + pytest.param(1, 0, 1, id="first_token_stays_prefill"), + pytest.param(17, 1, 0, id="block-size-plus-one-becomes-decode"), + ], +) +def test_one_token_prefill_selection_respects_recurrent_state( + monkeypatch: pytest.MonkeyPatch, + seq_len: int, + expected_decodes: int, + expected_prefills: int, +): + _patch_missing_runtime_cdiv(monkeypatch) + common_attn_metadata = create_common_attn_metadata( + BatchSpec(seq_lens=[seq_len], query_lens=[1]), + block_size=16, + device=torch.device("cpu"), + ) + # Model a prompt chunk explicitly. The helper normally classifies a + # one-token row as decode when synthesizing test metadata. + common_attn_metadata.is_prefilling = torch.tensor([True]) + builder = _make_builder( + device=torch.device("cpu"), + num_heads=32, + num_speculative_tokens=0, + ) + + attn_metadata = builder.build(0, common_attn_metadata) + + assert common_attn_metadata.is_prefilling.tolist() == [True] + assert attn_metadata.num_decodes == expected_decodes + assert attn_metadata.num_prefills == expected_prefills + + +@pytest.mark.parametrize( + ("seq_len", "expected_spec_decodes", "expected_prefills"), + [ + pytest.param(4, 0, 1, id="first_chunk_stays_prefill"), + pytest.param(8, 1, 0, id="stateful_chunk_folds_into_spec"), + ], +) +def test_spec_sized_prefill_fold_requires_recurrent_state( + monkeypatch: pytest.MonkeyPatch, + seq_len: int, + expected_spec_decodes: int, + expected_prefills: int, +): + _patch_missing_runtime_cdiv(monkeypatch) + common_attn_metadata = create_common_attn_metadata( + BatchSpec(seq_lens=[seq_len], query_lens=[4]), + block_size=16, + device=torch.device("cpu"), + ) + builder = _make_builder( + device=torch.device("cpu"), + num_heads=32, + num_speculative_tokens=3, + ) + + attn_metadata = builder.build( + 0, + common_attn_metadata, + num_accepted_tokens=torch.ones(1, dtype=torch.int32), + num_decode_draft_tokens_cpu=torch.full((1,), -1, dtype=torch.int32), + ) + + assert attn_metadata.num_spec_decodes == expected_spec_decodes + assert attn_metadata.num_prefills == expected_prefills + if expected_spec_decodes: + assert attn_metadata.spec_sequence_masks.tolist() == [True] + assert attn_metadata.num_accepted_tokens.tolist() == [4] + else: + assert attn_metadata.spec_sequence_masks is None + assert attn_metadata.num_accepted_tokens is None + + +def test_full_graph_without_runtime_spec_resets_captured_spec_inputs(): + capture_common_metadata = create_common_attn_metadata( + batch_spec=BatchSpec( + seq_lens=[4, 4], + query_lens=[4, 4], + name="full_graph_spec_capture", + ), + block_size=16, + device=torch.device("cpu"), + ) + capture_common_metadata.num_reqs = 4 + capture_common_metadata.block_table_tensor = torch.tensor( + [[10, 11, 12, 13], [20, 21, 22, 23]], + dtype=torch.int32, + ) + builder = _make_builder( + device=torch.device("cpu"), + num_heads=32, + num_speculative_tokens=3, + cudagraph_mode=CUDAGraphMode.FULL_DECODE_ONLY, + ) + captured_metadata = builder.build( + 0, + capture_common_metadata, + num_accepted_tokens=torch.tensor([2, 4], dtype=torch.int32), + num_decode_draft_tokens_cpu=torch.tensor([3, 3], dtype=torch.int32), + ) + captured_spec_metadata = captured_metadata.spec_decode_metadata + captured_conv1d_metadata = captured_spec_metadata.spec_causal_conv1d + + assert torch.count_nonzero(captured_conv1d_metadata.query_start_loc) > 0 + assert torch.count_nonzero(captured_spec_metadata.actual_seq_lengths) > 0 + + replay_common_metadata = create_common_attn_metadata( + batch_spec=BatchSpec( + seq_lens=[1, 1, 0, 0], + query_lens=[1, 1, 0, 0], + name="full_graph_replay_without_spec", + ), + block_size=16, + device=torch.device("cpu"), + ) + replay_metadata = builder.build( + 0, + replay_common_metadata, + num_accepted_tokens=torch.ones(4, dtype=torch.int32), + num_decode_draft_tokens_cpu=torch.full((4,), -1, dtype=torch.int32), + ) + + assert replay_metadata.spec_sequence_masks is None + assert replay_metadata.spec_decode_metadata is None + assert torch.equal( + captured_conv1d_metadata.cache_indices, + torch.full((4, 4), PAD_SLOT_ID, dtype=torch.int32), + ) + assert torch.count_nonzero(captured_conv1d_metadata.query_start_loc) == 0 + assert torch.count_nonzero(captured_conv1d_metadata.num_accepted_tokens) == 0 + assert torch.count_nonzero(captured_spec_metadata.actual_seq_lengths) == 0 + + +def test_full_graph_idle_dummy_uses_zero_length_recurrent_metadata(): + common_attn_metadata = create_common_attn_metadata( + batch_spec=BatchSpec( + seq_lens=[8, 8, 8, 8], + query_lens=[0, 0, 0, 0], + name="full_graph_idle_dummy", + ), + block_size=16, + device=torch.device("cpu"), + ) + common_attn_metadata.block_table_tensor[:, 0] = torch.tensor( + [10, 11, 98, 99], + dtype=torch.int32, + ) + builder = _make_builder( + device=torch.device("cpu"), + num_heads=32, + num_speculative_tokens=3, + cudagraph_mode=CUDAGraphMode.FULL_DECODE_ONLY, + ) + builder.spec_state_indices_tensor.fill_(77) + builder.spec_query_start_loc.fill_(77) + builder.non_spec_state_indices_tensor.fill_(77) + builder.non_spec_query_start_loc.fill_(77) + + attn_metadata = builder.build(0, common_attn_metadata) + + assert attn_metadata.num_actual_tokens == 0 + assert attn_metadata.num_decode_tokens == 0 + assert torch.count_nonzero(attn_metadata.non_spec_query_start_loc) == 0 + assert torch.all(attn_metadata.non_spec_state_indices_tensor == NULL_BLOCK_ID) + assert torch.count_nonzero(builder.spec_query_start_loc[:5]) == 0 + assert torch.all(builder.spec_state_indices_tensor[:4] == PAD_SLOT_ID) + + +@pytest.mark.parametrize( + ("num_speculative_tokens", "num_decode_draft_tokens_cpu"), + [ + pytest.param(0, None, id="without_spec_decode"), + pytest.param( + 3, + torch.full((4,), -1, dtype=torch.int32), + id="spec_decode_without_runtime_spec_requests", + ), + ], +) +def test_full_graph_non_spec_metadata_nulls_padded_state_indices( + num_speculative_tokens: int, + num_decode_draft_tokens_cpu: torch.Tensor | None, +): + common_attn_metadata = create_common_attn_metadata( + batch_spec=BatchSpec( + seq_lens=[1, 1, 0, 0], + query_lens=[1, 1, 0, 0], + name="full_graph_padded_non_spec_actual_seq_lengths", + ), + block_size=16, + device=torch.device("cpu"), + ) + common_attn_metadata.block_table_tensor[:, 0] = torch.tensor([10, 11, 98, 99]) + builder = _make_builder( + device=torch.device("cpu"), + num_heads=32, + num_speculative_tokens=num_speculative_tokens, + cudagraph_mode=CUDAGraphMode.FULL_DECODE_ONLY, + ) + builder.non_spec_state_indices_tensor.fill_(77) + builder.non_spec_query_start_loc.fill_(77) + builder.non_spec_actual_seq_lengths.fill_(77) + + attn_metadata = builder.build( + 0, + common_attn_metadata, + num_decode_draft_tokens_cpu=num_decode_draft_tokens_cpu, + ) + + assert attn_metadata.num_decodes == 4 + assert attn_metadata.num_decode_tokens == 2 + assert torch.equal( + attn_metadata.non_spec_query_start_loc, + torch.tensor([0, 1, 2, 2, 2], dtype=torch.int32), + ) + assert torch.equal( + attn_metadata.non_spec_state_indices_tensor, + torch.tensor( + [10, 11, NULL_BLOCK_ID, NULL_BLOCK_ID], + dtype=torch.int32, + ), + ) + decode_metadata = attn_metadata.non_spec_decode_metadata + conv1d_metadata = decode_metadata.causal_conv1d + assert conv1d_metadata.query_start_loc.data_ptr() == attn_metadata.non_spec_query_start_loc.data_ptr() + assert conv1d_metadata.cache_indices.data_ptr() == attn_metadata.non_spec_state_indices_tensor.data_ptr() + assert decode_metadata.actual_seq_lengths.data_ptr() == builder.non_spec_actual_seq_lengths.data_ptr() + assert torch.equal( + decode_metadata.actual_seq_lengths, + torch.tensor([0, 1, 1, 0, 0], dtype=torch.int32), + ) diff --git a/tests/ut/ops/test_kimi_kda.py b/tests/ut/ops/test_kimi_kda.py new file mode 100644 index 000000000000..ed9d845f628d --- /dev/null +++ b/tests/ut/ops/test_kimi_kda.py @@ -0,0 +1,271 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import SimpleNamespace +from unittest.mock import patch + +import torch +from torch import nn + +from vllm_ascend.ops.kimi_kda import ( + _PACKED_CONV_WEIGHT_NAME, + AscendKimiK3DeltaAttention, + AscendKimiK3MergedGateProjection, + _prepare_beta, + _zero_padded_output, + _zero_padded_recurrent_output, +) + + +def test_zero_padded_recurrent_output_clears_uncovered_tail(): + output = torch.randn(1, 8, 2, 3) + expected = output[:, :5].clone() + output[:, 5:] = torch.nan + + actual = _zero_padded_recurrent_output( + output, + torch.tensor([0, 3, 5, 5], dtype=torch.int32), + ) + + torch.testing.assert_close(actual[:, :5], expected) + assert torch.equal(actual[:, 5:], torch.zeros_like(actual[:, 5:])) + assert torch.isfinite(actual).all() + + +def test_zero_padded_output_uses_combined_live_token_count(): + output = torch.full((1, 8, 1, 1), torch.nan) + output[:, :6] = torch.arange(6).view(1, 6, 1, 1) + + actual = _zero_padded_output(output, torch.tensor(6, dtype=torch.int32)) + + torch.testing.assert_close(actual[:, :6], output[:, :6]) + assert torch.equal(actual[:, 6:], torch.zeros_like(actual[:, 6:])) + + +def test_run_causal_conv1d_returns_declared_output_alias(): + mixed_qkv = torch.randn(3, 8) + conv_weights = torch.randn(4, 8) + conv_state = torch.randn(2, 8, 4) + query_start_loc = torch.tensor([0, 3], dtype=torch.int32) + cache_indices = torch.tensor([1], dtype=torch.int32) + returned_alias = torch.full_like(mixed_qkv, 7) + + with patch.object( + torch.ops._C_ascend, + "npu_causal_conv1d_custom", + return_value=returned_alias, + create=True, + ) as causal_conv: + actual = AscendKimiK3DeltaAttention._run_causal_conv1d( + mixed_qkv, + conv_weights, + conv_state, + query_start_loc, + cache_indices, + None, + run_mode=1, + num_accepted_tokens=torch.tensor([3], dtype=torch.int32), + ) + + assert actual is returned_alias + assert causal_conv.call_args.kwargs["query_start_loc_opt"] is query_start_loc + assert causal_conv.call_args.kwargs["cache_indices_opt"] is cache_indices + assert causal_conv.call_args.kwargs["initial_state_mode_opt"] is None + + +def test_kda_output_norm_uses_checkpoint_epsilon(): + def fake_upstream_init(attention, _config, _vllm_config, _prefix): + nn.Module.__init__(attention) + attention.o_norm = SimpleNamespace(eps=1e-5) + attention.conv_size = 4 + attention.local_projection_size = 2 + attention.model_config = SimpleNamespace(dtype=torch.bfloat16) + attention.conv1d = nn.Module() + attention.conv1d.weight = nn.Parameter(torch.empty(6, 1, 4)) + attention.conv1d.quant_method = SimpleNamespace(process_weights_after_loading=lambda: None) + + config = SimpleNamespace(rms_norm_eps=1e-6) + vllm_config = SimpleNamespace( + model_config=SimpleNamespace( + multimodal_config=None, + enable_prompt_embeds=False, + ) + ) + with ( + patch( + "vllm_ascend.ops.kimi_kda.KimiK3DeltaAttention.__init__", + new=fake_upstream_init, + ), + patch("vllm_ascend.ops.kimi_kda.is_vl_model", return_value=False), + ): + attention = AscendKimiK3DeltaAttention(config, vllm_config) + + assert attention.o_norm.eps == config.rms_norm_eps + + +def test_prepare_beta_slices_and_applies_sigmoid_in_fp32(): + raw_beta = torch.tensor( + [[[-20.0], [0.0], [20.0], [100.0]]], + dtype=torch.bfloat16, + ) + + beta = _prepare_beta(raw_beta, num_actual_tokens=3) + + assert beta.dtype == torch.float32 + assert beta.shape == (1, 3, 1) + torch.testing.assert_close(beta, raw_beta[:, :3].float().sigmoid()) + assert torch.all((beta >= 0.0) & (beta <= 1.0)) + + +def test_recurrent_gate_uses_unbounded_kda_transform(): + attention = AscendKimiK3DeltaAttention.__new__(AscendKimiK3DeltaAttention) + nn.Module.__init__(attention) + attention.head_dim = 3 + attention.A_log = nn.Parameter(torch.randn(2)) + attention.dt_bias = nn.Parameter(torch.randn(6)) + raw_gate = torch.randn(1, 4, 2, 3) + expected = torch.randn(4, 2, 3) + + with patch( + "vllm_ascend.ops.kimi_kda.fused_kda_gate", + return_value=expected, + ) as fused_gate: + actual = attention._recurrent_gate(raw_gate) + + torch.testing.assert_close(actual, expected.unsqueeze(0)) + fused_gate.assert_called_once() + torch.testing.assert_close(fused_gate.call_args.args[0], raw_gate.reshape(4, 6)) + assert fused_gate.call_args.args[1] is attention.A_log + assert fused_gate.call_args.args[2] == attention.head_dim + assert fused_gate.call_args.kwargs["g_bias"] is attention.dt_bias + + +def test_prefill_accepts_unbounded_gate(): + attention = AscendKimiK3DeltaAttention.__new__(AscendKimiK3DeltaAttention) + nn.Module.__init__(attention) + attention.head_dim = 2 + attention.gate_lower_bound = None + attention.A_log = nn.Parameter(torch.randn(1)) + attention.dt_bias = nn.Parameter(torch.randn(2)) + + q = torch.randn(1, 2, 1, 2) + k = torch.randn_like(q) + v = torch.randn_like(q) + raw_gate = torch.randn_like(q) + beta = torch.randn(1, 2, 1) + recurrent_state = torch.randn(1, 1, 2, 2) + state_indices = torch.tensor([0], dtype=torch.int32) + has_initial_state = torch.tensor([True]) + metadata = SimpleNamespace( + cu_seqlens_host=(0, 2), + cu_seqlens_kern=None, + keep_meta=None, + chunk_indices_chunk64_host=(0, 0), + ) + transformed_gate = torch.randn_like(raw_gate) + gate_cumsum = torch.randn_like(raw_gate, dtype=torch.float32) + output = torch.randn_like(v) + final_state = torch.randn(1, 1, 2, 2) + + with ( + patch("vllm_ascend.ops.kimi_kda.clear_ssm_states"), + patch("vllm_ascend.ops.kimi_kda.l2norm_fwd", side_effect=lambda x: x), + patch.object(attention, "_recurrent_gate", return_value=transformed_gate) as recurrent_gate, + patch.object( + torch.ops._C_ascend, + "kda_gate_cumsum", + return_value=gate_cumsum, + create=True, + ) as kda_gate_cumsum, + patch.object( + torch.ops._C_ascend, + "chunk_kda_fwd", + return_value=(output, final_state), + create=True, + ), + ): + actual = attention._run_prefill( + q, + k, + v, + raw_gate, + beta, + recurrent_state, + state_indices, + has_initial_state, + metadata, + ) + + assert actual is output + recurrent_gate.assert_called_once_with(raw_gate) + assert kda_gate_cumsum.call_args.args[0] is transformed_gate + assert kda_gate_cumsum.call_args.args[1] == 64 + assert "use_gate_in_kernel" not in kda_gate_cumsum.call_args.kwargs + + +def test_merged_gate_projection_uses_vllm_shard_loader(): + projection = AscendKimiK3MergedGateProjection.__new__( + AscendKimiK3MergedGateProjection, + ) + nn.Module.__init__(projection) + param = nn.Parameter(torch.empty(4, 3)) + loaded_weight = torch.empty(2, 3) + + with patch( + "vllm_ascend.ops.kimi_kda._KimiGDNMergedColumnParallelLinear.weight_loader", + autospec=True, + ) as weight_loader: + projection.load_shard_weight(param, loaded_weight, shard_id=2) + + weight_loader.assert_called_once_with( + projection, + param, + loaded_weight, + 2, + ) + + +def test_kda_empty_forward_context_clears_preallocated_output(): + attention = AscendKimiK3DeltaAttention.__new__(AscendKimiK3DeltaAttention) + core_attn_out = torch.full((1, 4, 2, 3), torch.nan) + + with patch( + "vllm_ascend.ops.kimi_kda.get_forward_context", + return_value=SimpleNamespace(attn_metadata=None), + ): + attention._forward( + mixed_qkv=torch.empty(4, 18), + g1=torch.empty(1, 4, 2, 3), + g2=torch.empty(4, 2, 3), + beta=torch.empty(1, 4, 2), + core_attn_out=core_attn_out, + ) + + assert torch.equal(core_attn_out, torch.zeros_like(core_attn_out)) + + +def test_kda_conv_weight_is_packed_once_in_kernel_layout(): + attention = AscendKimiK3DeltaAttention.__new__(AscendKimiK3DeltaAttention) + nn.Module.__init__(attention) + attention.conv_size = 4 + attention.local_projection_size = 6 + attention.conv1d = nn.Module() + source = torch.arange(18 * 4, dtype=torch.float32).reshape(18, 1, 4) + attention.conv1d.weight = nn.Parameter(source) + attention.register_parameter( + _PACKED_CONV_WEIGHT_NAME, + nn.Parameter(torch.empty(4, 18, dtype=torch.bfloat16), requires_grad=False), + ) + original = attention.get_parameter(_PACKED_CONV_WEIGHT_NAME) + original_ptr = original.data_ptr() + + attention._pack_conv_weights() + + packed = attention.get_parameter(_PACKED_CONV_WEIGHT_NAME) + assert packed.data_ptr() == original_ptr + assert packed.dtype == torch.bfloat16 + assert packed.is_contiguous() + torch.testing.assert_close( + packed, + source[:, 0, :].transpose(0, 1).to(torch.bfloat16), + ) diff --git a/vllm_ascend/ops/gdn_attn_builder.py b/vllm_ascend/ops/gdn_attn_builder.py index 0168d1a3e072..3ef3ceb82277 100644 --- a/vllm_ascend/ops/gdn_attn_builder.py +++ b/vllm_ascend/ops/gdn_attn_builder.py @@ -26,6 +26,7 @@ ) from vllm.v1.attention.backends.utils import ( NULL_BLOCK_ID, + PAD_SLOT_ID, compute_causal_conv1d_metadata, mamba_get_block_table_tensor, split_decodes_and_prefills, @@ -51,6 +52,32 @@ def _stable_argsort_for_npu(tensor: torch.Tensor) -> torch.Tensor: return torch.argsort(tensor, stable=True) +def _treat_single_token_prefills_with_state_as_decodes( + common_attn_metadata: CommonAttentionMetadata, +) -> CommonAttentionMetadata: + """Use decode metadata for one-token stateful prompt chunks. + + A final one-token prompt chunk can replay the same fixed graph as an + ordinary decode. Once recurrent state exists, both paths must construct + identical GDN metadata so the graph consumes the current state indices. + First-token prefills remain on the prefill path because they have no state + to update. + """ + is_prefilling = common_attn_metadata.is_prefilling + seq_lens_cpu = common_attn_metadata.seq_lens_cpu_upper_bound + if is_prefilling is None or seq_lens_cpu is None: + return common_attn_metadata + + query_lens_cpu = torch.diff(common_attn_metadata.query_start_loc_cpu) + prefill_to_decode = is_prefilling & (query_lens_cpu == 1) & (seq_lens_cpu > 1) + if not torch.any(prefill_to_decode).item(): + return common_attn_metadata + + is_prefilling = is_prefilling.clone() + is_prefilling[prefill_to_decode] = False + return common_attn_metadata.replace(is_prefilling=is_prefilling) + + @dataclass class GDNChunkedPrefillMetadata: cu_seqlens_host: tuple[int, ...] @@ -148,7 +175,10 @@ def _build_non_spec_chunked_prefill_metadata( device: torch.device, ) -> GDNChunkedPrefillMetadata: hf_text_config = getattr(builder.vllm_config.model_config, "hf_text_config", None) - if hf_text_config is not None and hasattr(hf_text_config, "linear_num_value_heads"): + linear_attn_config = getattr(hf_text_config, "linear_attn_config", None) + if isinstance(linear_attn_config, dict) and linear_attn_config.get("num_heads") is not None: + gdn_num_heads = linear_attn_config["num_heads"] // builder.vllm_config.parallel_config.tensor_parallel_size + elif hf_text_config is not None and hasattr(hf_text_config, "linear_num_value_heads"): gdn_num_heads = ( hf_text_config.linear_num_value_heads // builder.vllm_config.parallel_config.tensor_parallel_size ) @@ -297,6 +327,45 @@ def _copy_sequence_indices_to_device( return spec_indices, non_spec_indices + def _pad_non_spec_decode_graph_inputs( + self, + state_indices: torch.Tensor, + query_start_loc: torch.Tensor, + *, + num_decode_tokens: int, + graph_batch_size: int, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Refresh fixed buffers consumed by a non-spec decode graph.""" + assert num_decode_tokens <= graph_batch_size + + padded_state_indices = self.non_spec_state_indices_tensor[:graph_batch_size] + padded_state_indices[num_decode_tokens:].fill_(NULL_BLOCK_ID) + padded_state_indices[:num_decode_tokens].copy_( + state_indices[:num_decode_tokens], + non_blocking=True, + ) + + padded_query_start_loc = self.non_spec_query_start_loc[: graph_batch_size + 1] + padded_query_start_loc[: num_decode_tokens + 1].copy_( + query_start_loc[: num_decode_tokens + 1], + non_blocking=True, + ) + query_padding = padded_query_start_loc[num_decode_tokens + 1 :] + if query_padding.numel() > 0: + query_padding.copy_( + padded_query_start_loc[num_decode_tokens].expand_as(query_padding), + non_blocking=True, + ) + + return padded_state_indices, padded_query_start_loc + + def _reset_spec_decode_graph_inputs(self, graph_batch_size: int) -> None: + """Make a captured speculative branch a no-op for this replay.""" + self.spec_state_indices_tensor[:graph_batch_size].fill_(PAD_SLOT_ID) + self.spec_query_start_loc[: graph_batch_size + 1].zero_() + self.num_accepted_tokens[:graph_batch_size].zero_() + self.spec_actual_seq_lengths[: graph_batch_size + 1].zero_() + def _attach_non_spec_prefill_metadata( self, attn_metadata: GDNAttentionMetadata, @@ -405,6 +474,42 @@ def _attach_non_spec_decode_metadata( ) return attn_metadata + def _fold_spec_sized_prefill_chunks_into_spec( + self, + common_attn_metadata: CommonAttentionMetadata, + spec_sequence_masks_cpu: torch.Tensor, + num_accepted_tokens: torch.Tensor | None, + ) -> tuple[torch.Tensor, torch.Tensor | None]: + """Advance stateful spec-width prompt chunks through live spec inputs.""" + is_prefilling = common_attn_metadata.is_prefilling + seq_lens_cpu = common_attn_metadata.seq_lens_cpu_upper_bound + if is_prefilling is None or seq_lens_cpu is None or num_accepted_tokens is None: + return spec_sequence_masks_cpu, num_accepted_tokens + + num_reqs = min( + spec_sequence_masks_cpu.numel(), + is_prefilling.numel(), + seq_lens_cpu.numel(), + ) + is_prefilling = is_prefilling[:num_reqs] + seq_lens_cpu = seq_lens_cpu[:num_reqs] + query_lens_cpu = torch.diff(common_attn_metadata.query_start_loc_cpu)[:num_reqs] + fold = ( + is_prefilling + & ~spec_sequence_masks_cpu + & (query_lens_cpu == self.num_spec + 1) + & (seq_lens_cpu > query_lens_cpu) + ) + fold_indices = fold.nonzero(as_tuple=True)[0] + if fold_indices.numel() == 0: + return spec_sequence_masks_cpu, num_accepted_tokens + + spec_sequence_masks_cpu = spec_sequence_masks_cpu.clone() + spec_sequence_masks_cpu[fold_indices] = True + num_accepted_tokens = num_accepted_tokens.clone() + num_accepted_tokens[fold_indices.to(num_accepted_tokens.device)] = self.num_spec + 1 + return spec_sequence_masks_cpu, num_accepted_tokens + def build( # type: ignore[override] self, common_prefix_len: int, @@ -413,7 +518,7 @@ def build( # type: ignore[override] num_decode_draft_tokens_cpu: torch.Tensor | None = None, fast_build: bool = False, ) -> GDNAttentionMetadata: - m = common_attn_metadata + m = _treat_single_token_prefills_with_state_as_decodes(common_attn_metadata) query_start_loc = m.query_start_loc query_start_loc_cpu = m.query_start_loc_cpu @@ -436,10 +541,22 @@ def build( # type: ignore[override] else: num_reqs = num_decode_draft_tokens_cpu.numel() spec_sequence_masks_cpu = self.spec_sequence_masks_cpu[:num_reqs] - torch.ge( - num_decode_draft_tokens_cpu, - 0, - out=spec_sequence_masks_cpu, + runtime_draft_tokens = num_decode_draft_tokens_cpu[num_decode_draft_tokens_cpu >= 0] + if runtime_draft_tokens.sum().item() > 0: + torch.ge( + num_decode_draft_tokens_cpu, + 0, + out=spec_sequence_masks_cpu, + ) + else: + # Dynamic speculative decoding can be enabled while this batch + # carries no draft tokens. Treat it as ordinary decode unless a + # stateful spec-width prompt chunk must use the spec branch. + spec_sequence_masks_cpu.zero_() + spec_sequence_masks_cpu, num_accepted_tokens = self._fold_spec_sized_prefill_chunks_into_spec( + m, + spec_sequence_masks_cpu, + num_accepted_tokens, ) num_spec_decodes = spec_sequence_masks_cpu.sum().item() if num_spec_decodes == 0: @@ -584,6 +701,12 @@ def build( # type: ignore[override] spec_sequence_indices, ) + # A FULL graph retains captured speculative conv/recurrent tasks. Clear + # their stable inputs on every no-spec replay so an idle or prefill + # batch cannot mutate state belonging to the preceding request. + if self.use_full_cuda_graph and self.use_spec_decode and num_spec_decodes == 0: + self._reset_spec_decode_graph_inputs(m.num_reqs) + chunk_indices: torch.Tensor | None = None chunk_offsets: torch.Tensor | None = None prefill_query_start_loc: torch.Tensor | None = None @@ -642,8 +765,6 @@ def build( # type: ignore[override] f"num_decodes: {num_decodes}, num_spec_decodes: {num_spec_decodes}" ) - batch_size = m.num_actual_tokens - if ( self.use_full_cuda_graph and num_prefills == 0 @@ -707,22 +828,17 @@ def build( # type: ignore[override] and num_spec_decodes == 0 and num_decodes <= self.decode_cudagraph_max_bs ): - self.non_spec_state_indices_tensor[batch_size:].fill_(NULL_BLOCK_ID) - self.non_spec_state_indices_tensor[:num_decodes].copy_( + graph_batch_size = m.num_reqs + ( non_spec_state_indices_tensor, - non_blocking=True, - ) - non_spec_state_indices_tensor = self.non_spec_state_indices_tensor[:batch_size] - non_spec_state_indices_tensor[num_decodes:].fill_(NULL_BLOCK_ID) - non_spec_conv1d_cache_indices = non_spec_state_indices_tensor - - self.non_spec_query_start_loc[: num_decodes + 1].copy_( non_spec_query_start_loc, - non_blocking=True, + ) = self._pad_non_spec_decode_graph_inputs( + non_spec_state_indices_tensor, + non_spec_query_start_loc, + num_decode_tokens=num_decode_tokens, + graph_batch_size=graph_batch_size, ) - non_spec_num_query_tokens = non_spec_query_start_loc[-1] - non_spec_query_start_loc = self.non_spec_query_start_loc[: batch_size + 1] - non_spec_query_start_loc[num_decodes + 1 :].fill_(non_spec_num_query_tokens) + non_spec_conv1d_cache_indices = non_spec_state_indices_tensor attn_metadata = GDNAttentionMetadata( num_prefills=num_prefills, diff --git a/vllm_ascend/ops/kimi_kda.py b/vllm_ascend/ops/kimi_kda.py new file mode 100644 index 000000000000..5ef2694c8d4e --- /dev/null +++ b/vllm_ascend/ops/kimi_kda.py @@ -0,0 +1,570 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Ascend backend for the vLLM 0.27 Kimi K3 delta-attention layer. + +The projections, weight loading, and cache specification stay owned by +upstream vLLM. Only the CUDA-specific convolution and KDA execution is +replaced here with the Ascend metadata builder and AscendC operators. +""" + +from functools import wraps + +import torch +from einops import rearrange +from torch import nn +from vllm.compilation.breakable_cudagraph import eager_break_during_capture +from vllm.forward_context import get_forward_context +from vllm.model_executor.layers.linear import ( + MergedColumnParallelLinear, +) +from vllm.model_executor.utils import replace_parameter +from vllm.models.kimi_k3.nvidia.kda import ( + KimiK3DeltaAttention, + _KimiGDNMergedColumnParallelLinear, +) +from vllm.third_party.flash_linear_attention.ops.l2norm import l2norm_fwd +from vllm.v1.attention.backend import AttentionBackend +from vllm.v1.attention.backends.gdn_attn import GDNAttentionMetadata +from vllm.v1.attention.backends.utils import PAD_SLOT_ID + +from vllm_ascend.ops.gdn_attn_builder import AscendGDNAttentionBackend +from vllm_ascend.ops.triton.fla.utils import clear_ssm_states +from vllm_ascend.ops.triton.kda.kda import fused_kda_gate + +_KDA_CHUNK_SIZE = 64 +_PACKED_CONV_WEIGHT_NAME = "ascend_conv1d_weight" + + +def _zero_padded_output( + output: torch.Tensor, + num_live_tokens: torch.Tensor, +) -> torch.Tensor: + """Clear graph-padding rows using a device-side live-token count.""" + token_indices = torch.arange( + output.shape[1], + dtype=num_live_tokens.dtype, + device=output.device, + ) + valid_tokens = token_indices < num_live_tokens + return torch.where(valid_tokens.view(1, -1, 1, 1), output, 0.0) + + +def _zero_padded_recurrent_output( + output: torch.Tensor, + query_start_loc: torch.Tensor, +) -> torch.Tensor: + """Clear graph-padding rows skipped by recurrent KDA.""" + return _zero_padded_output(output, query_start_loc[-1]) + + +def _prepare_beta( + raw_beta: torch.Tensor, + num_actual_tokens: int, +) -> torch.Tensor: + """Convert vLLM 0.27's packed raw beta to the AscendC contract.""" + return raw_beta[:, :num_actual_tokens].float().sigmoid() + + +class AscendKimiK3DeltaAttention(KimiK3DeltaAttention): + """Kimi K3 KDA using AscendC prefill and recurrent kernels.""" + + def __init__(self, config, vllm_config, prefix: str = "") -> None: + quant_config = getattr(vllm_config, "quant_config", None) + uses_mixed_projection = bool( + quant_config is not None + and getattr( + quant_config, + "uses_kimi_k3_mixed_kda_projection", + lambda _prefix: False, + )(f"{prefix}.in_proj_qkvgfab") + ) + super().__init__(config, vllm_config, prefix) + self.uses_mixed_projection = uses_mixed_projection + if uses_mixed_projection: + # vLLM 0.27 packs all KDA input projections into one linear. A + # QuaRot checkpoint instead stores q/k/v as W8A8 and keeps the + # three gates in floating point, so form one fused GEMM per + # precision group instead of falling back to four projections. + self.in_proj_qkvgfab = MergedColumnParallelLinear( + self.hidden_size, + [self.projection_size] * 3, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.in_proj_qkv", + ) + gate_output_sizes = [ + self.projection_size, + self.head_dim, + self.num_heads, + ] + if self.in_proj_padding: + gate_output_sizes.append(self.in_proj_padding * self.tp_size) + self.in_proj_gfab = _KimiGDNMergedColumnParallelLinear( + self.hidden_size, + gate_output_sizes, + replicated_shard_id=1, + tp_size=self.tp_size, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.in_proj_gfab", + ) + if self.in_proj_padding: + self.in_proj_gfab.weight.data[-self.in_proj_padding :].zero_() + # Upstream's FusedRMSNormGated constructor defaults to 1e-5, while + # Kimi K3 checkpoints use the model-configured RMS epsilon (1e-6 for + # the production checkpoint). Preserve the checkpoint contract used + # by the validated v0.26 implementation. + self.o_norm.eps = config.rms_norm_eps + # vLLM keeps the checkpoint-compatible FP32 [3C, 1, W] weight, while + # npu_causal_conv1d_custom consumes an activation-dtype [W, 3C] + # tensor. Materialize that kernel layout once after weight loading. + self.register_parameter( + _PACKED_CONV_WEIGHT_NAME, + nn.Parameter( + torch.empty( + self.conv_size, + 3 * self.local_projection_size, + dtype=self.model_config.dtype, + device=self.conv1d.weight.device, + ), + requires_grad=False, + ), + ) + original_process_weights = self.conv1d.quant_method.process_weights_after_loading + + @wraps(original_process_weights) + def process_weights_and_pack(*args, **kwargs): + result = original_process_weights(*args, **kwargs) + self._pack_conv_weights() + return result + + self.conv1d.quant_method.process_weights_after_loading = process_weights_and_pack + + def get_attn_backend(self) -> type[AttentionBackend]: + return AscendGDNAttentionBackend + + def forward( + self, + hidden_states: torch.Tensor, + positions: torch.Tensor, + ) -> torch.Tensor: + if self.uses_mixed_projection: + num_tokens = hidden_states.size(0) + mixed_qkv = self.in_proj_qkvgfab(hidden_states)[0] + projected_gfab = self.in_proj_gfab(hidden_states)[0] + split_sizes = [ + self.local_projection_size, + self.head_dim, + self.local_num_heads, + ] + if self.in_proj_padding: + split_sizes.append(self.in_proj_padding) + g_proj_states, f_a, beta = projected_gfab.split(split_sizes, dim=-1)[:3] + beta = beta.unsqueeze(0) + + g1 = self.f_b_proj(f_a)[0] + g1 = rearrange(g1, "n (h d) -> 1 n h d", d=self.head_dim) + g2 = rearrange(g_proj_states, "... (h d) -> ... h d", d=self.head_dim) + core_attn_out = torch.empty( + (1, num_tokens, self.local_num_heads, self.head_dim), + dtype=hidden_states.dtype, + device=hidden_states.device, + ) + self._forward( + mixed_qkv=mixed_qkv, + g1=g1, + g2=g2, + beta=beta, + core_attn_out=core_attn_out, + ) + core_attn_out = rearrange(core_attn_out, "1 n h d -> n (h d)") + return self.o_proj(core_attn_out)[0] + return super().forward(hidden_states, positions) + + @staticmethod + def _run_causal_conv1d( + mixed_qkv: torch.Tensor, + conv_weights_t: torch.Tensor, + conv_state: torch.Tensor, + query_start_loc: torch.Tensor, + cache_indices: torch.Tensor, + initial_state_mode: torch.Tensor | None, + *, + run_mode: int, + num_accepted_tokens: torch.Tensor | None = None, + ) -> torch.Tensor: + output = torch.empty_like(mixed_qkv) + # Consume the operator's declared output alias. Returning ``output`` + # independently would let graph functionalization treat the custom-op + # result as dead and expose the uninitialized allocation instead. + return torch.ops._C_ascend.npu_causal_conv1d_custom( + output, + mixed_qkv, + conv_weights_t, + conv_state=conv_state, + bias_opt=None, + query_start_loc_opt=query_start_loc, + cache_indices_opt=cache_indices, + initial_state_mode_opt=initial_state_mode, + num_accepted_tokens_opt=num_accepted_tokens, + activation_mode=1, + pad_slot_id=PAD_SLOT_ID, + run_mode=run_mode, + ) + + @torch.no_grad() + def _pack_conv_weights(self) -> None: + if self.conv1d.weight.is_meta: + return + packed_param = self.get_parameter(_PACKED_CONV_WEIGHT_NAME) + packed_weight = ( + self.conv1d.weight.view(self.conv1d.weight.size(0), self.conv1d.weight.size(2)) + .transpose(0, 1) + .to(device=packed_param.device, dtype=packed_param.dtype) + .contiguous() + ) + replace_parameter( + self, + _PACKED_CONV_WEIGHT_NAME, + packed_weight, + prefer_copy=True, + ) + + def _recurrent_gate(self, raw_gate: torch.Tensor) -> torch.Tensor: + flat_gate = rearrange(raw_gate, "1 n h d -> n (h d)") + gate = fused_kda_gate( + flat_gate, + self.A_log, + self.head_dim, + g_bias=self.dt_bias, + ) + return gate.unsqueeze(0) + + def _run_recurrent( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + raw_gate: torch.Tensor, + beta: torch.Tensor, + recurrent_state: torch.Tensor, + cu_seqlens: torch.Tensor, + state_indices: torch.Tensor, + *, + num_accepted_tokens: torch.Tensor | None = None, + ) -> torch.Tensor: + return torch.ops._C_ascend.recurrent_kda( + q.contiguous(), + k.contiguous(), + v.contiguous(), + raw_gate.contiguous(), + beta.contiguous(), + recurrent_state, + cu_seqlens, + state_indices, + self.A_log.reshape(-1).contiguous(), + self.dt_bias.contiguous(), + num_accepted_tokens=num_accepted_tokens, + scale=self.head_dim**-0.5, + use_qk_l2norm_in_kernel=True, + use_gate_in_kernel=True, + use_beta_sigmoid_in_kernel=False, + allow_neg_eigval=False, + safe_gate=self.gate_lower_bound is not None, + lower_bound=(self.gate_lower_bound if self.gate_lower_bound is not None else -5.0), + ) + + def _run_prefill( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + raw_gate: torch.Tensor, + beta: torch.Tensor, + recurrent_state: torch.Tensor, + state_indices: torch.Tensor, + has_initial_state: torch.Tensor, + prebuilt_metadata, + ) -> torch.Tensor: + cu_seqlens = ( + prebuilt_metadata.cu_seqlens_host + if prebuilt_metadata.cu_seqlens_kern is None + else prebuilt_metadata.cu_seqlens_kern + ) + keep = prebuilt_metadata.keep_meta + if keep is not None: + state_indices = state_indices[keep] + has_initial_state = has_initial_state[keep] + + # The recurrent cache is [H, V, K], while chunk_kda_fwd consumes + # [H, K, V]. Keep the conversion at this operator boundary. + initial_state_vk = recurrent_state[state_indices].contiguous() + clear_ssm_states(initial_state_vk, has_initial_state) + initial_state_kv = initial_state_vk.transpose(-1, -2).contiguous() + + q = l2norm_fwd(q.contiguous()) + k = l2norm_fwd(k.contiguous()) + if self.gate_lower_bound is not None: + gate_cumsum = torch.ops._C_ascend.kda_gate_cumsum( + raw_gate.contiguous(), + _KDA_CHUNK_SIZE, + A_log=self.A_log.reshape(-1).contiguous(), + dt_bias=self.dt_bias.contiguous(), + cu_seqlens=cu_seqlens, + use_gate_in_kernel=True, + safe_gate=True, + lower_bound=self.gate_lower_bound, + layout="BSND", + ) + else: + gate = self._recurrent_gate(raw_gate) + gate_cumsum = torch.ops._C_ascend.kda_gate_cumsum( + gate.contiguous(), + _KDA_CHUNK_SIZE, + cu_seqlens=cu_seqlens, + layout="BSND", + ) + result = torch.ops._C_ascend.chunk_kda_fwd( + q, + k, + v.contiguous(), + gate_cumsum, + beta.contiguous(), + self.head_dim**-0.5, + _KDA_CHUNK_SIZE, + layout="BSND", + initial_state=initial_state_kv, + output_final_state=True, + cu_seqlens=cu_seqlens, + chunk_indices=prebuilt_metadata.chunk_indices_chunk64_host, + return_intermediate=False, + ) + recurrent_state[state_indices] = result[1].transpose(-1, -2).contiguous().to(recurrent_state.dtype) + return result[0] + + @eager_break_during_capture + def _forward( + self, + mixed_qkv: torch.Tensor, + g1: torch.Tensor, + g2: torch.Tensor, + beta: torch.Tensor, + core_attn_out: torch.Tensor, + ) -> None: + """Dispatch speculative, prefill, and decode tokens through KDA kernels.""" + forward_context = get_forward_context() + attn_metadata_raw = forward_context.attn_metadata + if attn_metadata_raw is None: + core_attn_out.zero_() + return + + assert isinstance(attn_metadata_raw, dict) + attn_metadata = attn_metadata_raw[self.prefix] + assert isinstance(attn_metadata, GDNAttentionMetadata) + + num_actual_tokens = attn_metadata.num_actual_tokens + mixed_qkv = mixed_qkv[:num_actual_tokens] + g1 = g1[:, :num_actual_tokens] + g2 = g2[:num_actual_tokens] + beta = _prepare_beta(beta, num_actual_tokens) + + conv_state, recurrent_state = self.kv_cache + conv_weights_t = self.get_parameter(_PACKED_CONV_WEIGHT_NAME) + spec_masks = attn_metadata.spec_sequence_masks + spec_token_indices = attn_metadata.spec_token_indx + non_spec_token_indices = attn_metadata.non_spec_token_indx + + if spec_masks is not None: + if attn_metadata.num_prefills == 0 and attn_metadata.num_decodes == 0: + mixed_spec = mixed_qkv + raw_gate_spec = g1 + beta_spec = beta + mixed_non_spec = raw_gate_non_spec = beta_non_spec = None + else: + assert spec_token_indices is not None + assert non_spec_token_indices is not None + mixed_spec = mixed_qkv.index_select(0, spec_token_indices) + raw_gate_spec = g1.index_select(1, spec_token_indices) + beta_spec = beta.index_select(1, spec_token_indices) + mixed_non_spec = mixed_qkv.index_select(0, non_spec_token_indices) + raw_gate_non_spec = g1.index_select(1, non_spec_token_indices) + beta_non_spec = beta.index_select(1, non_spec_token_indices) + else: + mixed_spec = raw_gate_spec = beta_spec = None + mixed_non_spec = mixed_qkv + raw_gate_non_spec = g1 + beta_non_spec = beta + + core_spec = None + if mixed_spec is not None: + spec_meta = attn_metadata.spec_decode_metadata + assert spec_meta is not None + spec_conv_meta = spec_meta.spec_causal_conv1d + mixed_spec = self._run_causal_conv1d( + mixed_spec, + conv_weights_t, + conv_state, + spec_conv_meta.query_start_loc, + spec_conv_meta.cache_indices, + None, + run_mode=1, + num_accepted_tokens=spec_conv_meta.num_accepted_tokens, + ) + q_spec, k_spec, v_spec = ( + rearrange(x, "n (h d) -> 1 n h d", d=self.head_dim) for x in mixed_spec.chunk(3, dim=-1) + ) + assert raw_gate_spec is not None and beta_spec is not None + assert attn_metadata.spec_query_start_loc is not None + assert attn_metadata.spec_state_indices_tensor is not None + core_spec = self._run_recurrent( + q_spec, + k_spec, + v_spec, + raw_gate_spec, + beta_spec, + recurrent_state, + attn_metadata.spec_query_start_loc, + attn_metadata.spec_state_indices_tensor, + num_accepted_tokens=spec_conv_meta.num_accepted_tokens, + ) + core_spec = _zero_padded_recurrent_output( + core_spec, + attn_metadata.spec_query_start_loc, + ) + + core_non_spec = None + if mixed_non_spec is not None and mixed_non_spec.shape[0] > 0: + if attn_metadata.num_prefills > 0: + prefill_meta = attn_metadata.non_spec_prefill_metadata + assert prefill_meta is not None + mixed_non_spec = self._run_causal_conv1d( + mixed_non_spec, + conv_weights_t, + conv_state, + prefill_meta.causal_conv1d.query_start_loc, + prefill_meta.causal_conv1d.cache_indices, + prefill_meta.causal_conv1d.initial_state_mode, + run_mode=0, + ) + elif attn_metadata.num_decodes > 0: + decode_meta = attn_metadata.non_spec_decode_metadata + assert decode_meta is not None + mixed_non_spec = self._run_causal_conv1d( + mixed_non_spec, + conv_weights_t, + conv_state, + decode_meta.causal_conv1d.query_start_loc, + decode_meta.causal_conv1d.cache_indices, + None, + run_mode=1, + ) + + q_non_spec, k_non_spec, v_non_spec = ( + rearrange(x, "n (h d) -> 1 n h d", d=self.head_dim) for x in mixed_non_spec.chunk(3, dim=-1) + ) + assert raw_gate_non_spec is not None + assert beta_non_spec is not None + + split_non_spec = spec_masks is None and attn_metadata.num_prefills > 0 and attn_metadata.num_decodes > 0 + num_decode_tokens = attn_metadata.num_decode_tokens + core_decode = None + if split_non_spec: + assert attn_metadata.non_spec_query_start_loc is not None + assert attn_metadata.non_spec_state_indices_tensor is not None + core_decode = self._run_recurrent( + q_non_spec[:, :num_decode_tokens], + k_non_spec[:, :num_decode_tokens], + v_non_spec[:, :num_decode_tokens], + raw_gate_non_spec[:, :num_decode_tokens], + beta_non_spec[:, :num_decode_tokens], + recurrent_state, + attn_metadata.non_spec_query_start_loc[: attn_metadata.num_decodes + 1], + attn_metadata.non_spec_state_indices_tensor[: attn_metadata.num_decodes], + ) + + if attn_metadata.num_prefills > 0: + if split_non_spec: + q_non_spec = q_non_spec[:, num_decode_tokens:] + k_non_spec = k_non_spec[:, num_decode_tokens:] + v_non_spec = v_non_spec[:, num_decode_tokens:] + raw_gate_non_spec = raw_gate_non_spec[:, num_decode_tokens:] + beta_non_spec = beta_non_spec[:, num_decode_tokens:] + + assert attn_metadata.prefill_state_indices is not None + assert attn_metadata.prefill_has_initial_state is not None + prefill_meta = attn_metadata.non_spec_prefill_metadata + assert prefill_meta is not None + core_prefill = self._run_prefill( + q_non_spec, + k_non_spec, + v_non_spec, + raw_gate_non_spec, + beta_non_spec, + recurrent_state, + attn_metadata.prefill_state_indices, + attn_metadata.prefill_has_initial_state, + prefill_meta.chunk, + ) + core_non_spec = ( + torch.cat((core_decode, core_prefill), dim=1) if core_decode is not None else core_prefill + ) + elif attn_metadata.num_decodes > 0: + assert attn_metadata.non_spec_query_start_loc is not None + assert attn_metadata.non_spec_state_indices_tensor is not None + core_non_spec = self._run_recurrent( + q_non_spec, + k_non_spec, + v_non_spec, + raw_gate_non_spec, + beta_non_spec, + recurrent_state, + attn_metadata.non_spec_query_start_loc[: attn_metadata.num_decodes + 1], + attn_metadata.non_spec_state_indices_tensor, + ) + + if core_non_spec is not None: + assert attn_metadata.non_spec_query_start_loc is not None + core_non_spec = _zero_padded_recurrent_output( + core_non_spec, + attn_metadata.non_spec_query_start_loc, + ) + + if core_spec is None and core_non_spec is None: + # Idle DP dummy runs carry graph-shaped metadata with no live work. + # Do not feed a previous replay's output through the norm gate. + core_attn_out.zero_() + return + + num_live_tokens = None + if core_spec is not None: + assert attn_metadata.spec_query_start_loc is not None + num_live_tokens = attn_metadata.spec_query_start_loc[-1] + if core_non_spec is not None: + assert attn_metadata.non_spec_query_start_loc is not None + num_non_spec_tokens = attn_metadata.non_spec_query_start_loc[-1] + num_live_tokens = num_non_spec_tokens if num_live_tokens is None else num_live_tokens + num_non_spec_tokens + assert num_live_tokens is not None + + # Reuse the caller-owned result buffer. FULL graphs can leave rows + # outside the live spec/non-spec index sets, so define them before the + # two index copies rather than allocating a temporary merged tensor. + core_attn_out[:, :num_actual_tokens].zero_() + if core_spec is not None and core_non_spec is not None: + assert spec_token_indices is not None + assert non_spec_token_indices is not None + assert spec_token_indices.numel() + non_spec_token_indices.numel() <= num_actual_tokens + core_attn_out[:, :num_actual_tokens].index_copy_(1, spec_token_indices, core_spec) + core_attn_out[:, :num_actual_tokens].index_copy_(1, non_spec_token_indices, core_non_spec) + elif core_spec is not None: + core_attn_out[:, :num_actual_tokens] = core_spec + elif core_non_spec is not None: + core_attn_out[:, :num_actual_tokens] = core_non_spec + + # Let vLLM's CustomOp dispatch select FusedRMSNormGated.forward_native + # on Ascend. Calling the CUDA/Triton helper directly bypasses platform + # dispatch. + normalized = self.o_norm(core_attn_out[:, :num_actual_tokens], g2) + # Mask again after the norm gate: zero * sigmoid(NaN) is still NaN in + # static padding rows whose captured gate values are not live. + core_attn_out[:, :num_actual_tokens].copy_(_zero_padded_output(normalized, num_live_tokens)) + core_attn_out[:, num_actual_tokens:].zero_() From 37cf61c8e26026866c977e29d79e03d82fa254ee Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Mon, 24 Aug 2026 04:23:19 -0500 Subject: [PATCH 05/50] feat(mla): support Kimi K3 attention on Ascend Preserve per-group RoPE semantics, non-causal multi-token decode metadata, no-RoPE execution, A5 MLA preprocessing, and identity rotary inputs required by Kimi K3 target and draft layers. Signed-off-by: maoxx241 --- tests/ut/attention/a2/test_mla_v1.py | 140 ++++++++++++++-- tests/ut/ops/test_rotary_embedding.py | 20 +++ vllm_ascend/attention/mla_v1.py | 228 ++++++++++++++++++++------ vllm_ascend/ops/mla.py | 8 + vllm_ascend/ops/rotary_embedding.py | 9 + 5 files changed, 345 insertions(+), 60 deletions(-) diff --git a/tests/ut/attention/a2/test_mla_v1.py b/tests/ut/attention/a2/test_mla_v1.py index 79b48257d787..0d11c879ba35 100644 --- a/tests/ut/attention/a2/test_mla_v1.py +++ b/tests/ut/attention/a2/test_mla_v1.py @@ -426,6 +426,7 @@ def test_ascend_mla_metadata_default(self): self.assertEqual(metadata.head_dim, head_dim) self.assertEqual(metadata.attn_mask, attn_mask) self.assertEqual(metadata.attn_state, attn_state) + self.assertTrue(metadata.causal) self.assertEqual(metadata.decode, decode) self.assertEqual(metadata.prefill, prefill) @@ -463,6 +464,7 @@ def test_ascend_mla_metadata_builder_default(self): mock_vllm_config.model_config.get_head_size.return_value = 64 mock_vllm_config.model_config.dtype = torch.float16 mock_vllm_config.model_config.hf_text_config.qk_rope_head_dim = 64 + mock_vllm_config.model_config.hf_text_config.mla_use_nope = False mock_vllm_config.cache_config.block_size = 16 mock_vllm_config.scheduler_config.max_num_seqs = 4 mock_vllm_config.scheduler_config.enable_chunked_prefill = False @@ -472,11 +474,111 @@ def test_ascend_mla_metadata_builder_default(self): ascend_config = MagicMock() with patch("vllm_ascend.attention.mla_v1.get_ascend_config", return_value=ascend_config): - builder = AscendMLAMetadataBuilder(None, None, mock_vllm_config, mock_device) + builder = AscendMLAMetadataBuilder(None, [], mock_vllm_config, mock_device) self.assertEqual(builder.block_size, mock_vllm_config.cache_config.block_size) self.assertEqual(builder.chunked_prefill_enabled, mock_vllm_config.scheduler_config.enable_chunked_prefill) + def test_metadata_builder_uses_draft_layer_rope_mode(self): + mock_vllm_config = MagicMock() + mock_vllm_config.model_config.max_model_len = 1024 + mock_vllm_config.model_config.get_head_size.return_value = 64 + mock_vllm_config.model_config.dtype = torch.float16 + mock_vllm_config.model_config.hf_text_config = SimpleNamespace( + qk_rope_head_dim=64, + mla_use_nope=True, + ) + mock_vllm_config.cache_config.block_size = 16 + mock_vllm_config.scheduler_config.max_num_seqs = 4 + mock_vllm_config.scheduler_config.enable_chunked_prefill = False + mock_vllm_config.speculative_config = None + mock_vllm_config.compilation_config.static_forward_context = { + "draft.self_attn": SimpleNamespace( + impl=SimpleNamespace(use_mla_rope=True), + ), + } + + with patch( + "vllm_ascend.attention.mla_v1.get_ascend_config", + return_value=MagicMock(), + ): + builder = AscendMLAMetadataBuilder( + None, + ["draft.self_attn"], + mock_vllm_config, + "cpu", + ) + + self.assertTrue(builder.use_mla_rope) + + def test_metadata_builder_uses_target_layer_nope_mode(self): + mock_vllm_config = MagicMock() + mock_vllm_config.model_config.max_model_len = 1024 + mock_vllm_config.model_config.get_head_size.return_value = 64 + mock_vllm_config.model_config.dtype = torch.float16 + mock_vllm_config.model_config.hf_text_config = SimpleNamespace( + qk_rope_head_dim=64, + mla_use_nope=False, + ) + mock_vllm_config.cache_config.block_size = 16 + mock_vllm_config.scheduler_config.max_num_seqs = 4 + mock_vllm_config.scheduler_config.enable_chunked_prefill = False + mock_vllm_config.speculative_config = None + mock_vllm_config.compilation_config.static_forward_context = { + "target.self_attn": SimpleNamespace( + impl=SimpleNamespace(use_mla_rope=False), + ), + } + + with patch( + "vllm_ascend.attention.mla_v1.get_ascend_config", + return_value=MagicMock(), + ): + builder = AscendMLAMetadataBuilder( + None, + ["target.self_attn"], + mock_vllm_config, + "cpu", + ) + + self.assertFalse(builder.use_mla_rope) + + def test_metadata_builder_rejects_mixed_rope_modes(self): + mock_vllm_config = MagicMock() + mock_vllm_config.model_config.max_model_len = 1024 + mock_vllm_config.model_config.get_head_size.return_value = 64 + mock_vllm_config.model_config.dtype = torch.float16 + mock_vllm_config.model_config.hf_text_config = SimpleNamespace( + qk_rope_head_dim=64, + mla_use_nope=False, + ) + mock_vllm_config.cache_config.block_size = 16 + mock_vllm_config.scheduler_config.max_num_seqs = 4 + mock_vllm_config.scheduler_config.enable_chunked_prefill = False + mock_vllm_config.speculative_config = None + mock_vllm_config.compilation_config.static_forward_context = { + "rope.self_attn": SimpleNamespace( + impl=SimpleNamespace(use_mla_rope=True), + ), + "nope.self_attn": SimpleNamespace( + impl=SimpleNamespace(use_mla_rope=False), + ), + } + + with ( + patch( + "vllm_ascend.attention.mla_v1.get_ascend_config", + return_value=MagicMock(), + ), + self.assertRaisesRegex(AssertionError, "separate KV cache groups"), + ): + AscendMLAMetadataBuilder( + None, + ["rope.self_attn", "nope.self_attn"], + mock_vllm_config, + "cpu", + ) + def test_ascend_mla_metadata_builder_spec_decode(self): mock_vllm_config = MagicMock() mock_vllm_config.model_config.max_model_len = 1024 @@ -494,7 +596,7 @@ def test_ascend_mla_metadata_builder_spec_decode(self): ascend_config = MagicMock() with patch("vllm_ascend.attention.mla_v1.get_ascend_config", return_value=ascend_config): - builder = AscendMLAMetadataBuilder(None, None, mock_vllm_config, mock_device) + builder = AscendMLAMetadataBuilder(None, [], mock_vllm_config, mock_device) self.assertEqual(builder.block_size, mock_vllm_config.cache_config.block_size) self.assertEqual(builder.chunked_prefill_enabled, mock_vllm_config.scheduler_config.enable_chunked_prefill) @@ -506,6 +608,7 @@ def test_ascend_mla_metadata_builder_build_full_graph(self, mock_get_cos_and_sin mock_vllm_config.model_config.get_head_size.return_value = 64 mock_vllm_config.model_config.dtype = torch.float16 mock_vllm_config.model_config.hf_text_config.qk_rope_head_dim = 64 + mock_vllm_config.model_config.hf_text_config.mla_use_nope = False mock_vllm_config.cache_config.block_size = 16 mock_vllm_config.scheduler_config.max_num_seqs = 4 mock_vllm_config.scheduler_config.chunked_prefill_enabled = False @@ -520,7 +623,7 @@ def test_ascend_mla_metadata_builder_build_full_graph(self, mock_get_cos_and_sin mock_spec_config.disable_padded_drafter_batch = True mock_vllm_config.speculative_config = mock_spec_config - builder = AscendMLAMetadataBuilder(None, None, mock_vllm_config, mock_device) + builder = AscendMLAMetadataBuilder(None, [], mock_vllm_config, mock_device) common_metadata = MagicMock() common_metadata.graph_pad_size = 8 common_metadata.num_reqs = 4 @@ -531,10 +634,12 @@ def test_ascend_mla_metadata_builder_build_full_graph(self, mock_get_cos_and_sin common_metadata.query_start_loc = torch.Tensor([0, 1, 2, 4, 5]).int() common_metadata.query_start_loc_cpu = torch.Tensor([0, 1, 2, 4, 5]).int() common_metadata.positions = torch.Tensor([1, 2, 3, 4, 5, 6]).int() + common_metadata.causal = False block_table = torch.Tensor([[1, 0], [2, 0], [3, 0], [4, 0]]).int() common_metadata.block_table_tensor = block_table mock_get_cos_and_sin_mla.return_value = (torch.tensor([6, 6]), torch.Tensor([6, 6])) metadata = builder.build(0, common_metadata) + self.assertFalse(metadata.causal) self.assertEqual(metadata.decode.actual_seq_lengths_q, [1, 2, 4, 5, 6, 6, 7, 8]) self.assertEqual(metadata.decode.block_table.shape[0], 8) @@ -555,7 +660,7 @@ def test_reorder_batch(self): mock_vllm_config.speculative_config = None with patch("vllm_ascend.attention.mla_v1.get_ascend_config", return_value=ascend_config): - builder = AscendMLAMetadataBuilder(None, None, mock_vllm_config, mock_device) + builder = AscendMLAMetadataBuilder(None, [], mock_vllm_config, mock_device) builder.decode_threshold = 1 input_batch = MagicMock() @@ -606,7 +711,7 @@ def test_set_num_actual_tokens(self): mock_vllm_config.speculative_config = None - builder = AscendMLAMetadataBuilder(None, None, mock_vllm_config, mock_device) + builder = AscendMLAMetadataBuilder(None, [], mock_vllm_config, mock_device) common_attn_metadata = MagicMock() common_attn_metadata.num_actual_tokens = 100 @@ -626,7 +731,7 @@ def test_pad_actual_seq_lens_q_mtp_disable_pad(self): mock_device = "cpu" mock_vllm_config.speculative_config = None - builder = AscendMLAMetadataBuilder(None, None, mock_vllm_config, mock_device) + builder = AscendMLAMetadataBuilder(None, [], mock_vllm_config, mock_device) input_seq_lens = [1, 2, 4, 5] expect_output = [1, 2, 4, 5, 6, 6, 7, 8] num_reqs = 4 @@ -650,7 +755,7 @@ def test_pad_actual_seq_lens_q_mtp_enable_pad(self): common_metadata = MagicMock() common_metadata.actual_seq_lengths_q = [2, 4, 6, 8] - builder = AscendMLAMetadataBuilder(None, None, mock_vllm_config, mock_device) + builder = AscendMLAMetadataBuilder(None, [], mock_vllm_config, mock_device) input_seq_lens = [2, 4, 6] expect_output = [2, 4, 6, 8] num_reqs = 3 @@ -676,7 +781,7 @@ def test_pad_actual_seq_lens_q_mtp_enable_pad_with_padding(self): common_metadata = MagicMock() common_metadata.actual_seq_lengths_q = [2, 4, 6, 100] - builder = AscendMLAMetadataBuilder(None, None, mock_vllm_config, mock_device) + builder = AscendMLAMetadataBuilder(None, [], mock_vllm_config, mock_device) input_seq_lens = [2, 4, 6] num_reqs = 3 num_reqs_pad_size = 1 @@ -1122,6 +1227,8 @@ def setUp(self, get_current_vllm_config, mock_tp): "fused_qkv_a_proj": MagicMock(), "kv_a_layernorm": kv_a_layernorm, "rotary_emb": MagicMock(), + "g_proj": None, + "use_mla_rope": True, } self.impl = AscendMLAImpl( @@ -1180,6 +1287,8 @@ def test_init_head_padding_for_non_power_of_two(self, mock_get_current_vllm_conf "fused_qkv_a_proj": MagicMock(), "kv_a_layernorm": MagicMock(), "rotary_emb": MagicMock(), + "g_proj": None, + "use_mla_rope": True, } impl = AscendMLAImpl( num_heads=20, @@ -1296,7 +1405,6 @@ def test_update_graph_params( mock_attn_metadata.decode.seq_lens_list = [10, 20, 30] mock_attn_metadata.decode.actual_seq_lengths_q = [10, 20, 30] mock_attn_metadata.decode.block_table = torch.randint(0, 100, (3, 4)) - mock_graph_params = MagicMock() mock_graph_params.attn_params = { @@ -1338,13 +1446,14 @@ def test_update_graph_params( mock_speculative_config = MagicMock() mock_speculative_config.disable_padded_drafter_batch = False - # Test non-draft model - mock_ctx.is_draft_model = False + mock_ctx.is_draft_model = True + mock_ctx.is_draft_model_prefill = False AscendMLAImpl.update_graph_params( mock_update_stream, mock_forward_context, 100, speculative_config=mock_speculative_config, + draft_attn_metadatas=[{"layer_0": mock_attn_metadata}], ) @patch("vllm_ascend.ascend_forward_context.get_forward_context") @@ -1541,6 +1650,7 @@ def test_process_weights_for_fused_mlapo_a5(self, mock_format_cast, mock_get_asc self.impl.q_proj.weight.data = torch.randn(128, 128) self.impl.q_proj.weight_scale.data = torch.randn(128, 128, 128) self.impl.q_lora_rank = 32 + self.impl._mlapo_quant_type = object from vllm_ascend.attention.mla_v1 import AscendDeviceType @@ -1626,6 +1736,8 @@ def test_forward_prefill_non_power_of_two_heads(self, mock_fia, mock_device_oper "fused_qkv_a_proj": MagicMock(), "kv_a_layernorm": MagicMock(), "rotary_emb": MagicMock(), + "g_proj": None, + "use_mla_rope": True, } impl = AscendMLAImpl( num_heads=num_heads, @@ -1972,6 +2084,8 @@ def test_compute_prefill_context_non_power_of_two_heads( "fused_qkv_a_proj": MagicMock(), "kv_a_layernorm": MagicMock(), "rotary_emb": MagicMock(), + "g_proj": None, + "use_mla_rope": True, } impl = AscendMLAImpl( num_heads=num_heads, @@ -2251,6 +2365,8 @@ def test_forward_decode_non_power_of_two_heads( "fused_qkv_a_proj": MagicMock(), "kv_a_layernorm": MagicMock(), "rotary_emb": MagicMock(), + "g_proj": None, + "use_mla_rope": True, } impl = AscendMLAImpl( num_heads=num_heads, @@ -2325,6 +2441,8 @@ def test_forward_decode_non_power_of_two_heads_normal( "fused_qkv_a_proj": MagicMock(), "kv_a_layernorm": MagicMock(), "rotary_emb": MagicMock(), + "g_proj": None, + "use_mla_rope": True, } impl = AscendMLAImpl( num_heads=num_heads, diff --git a/tests/ut/ops/test_rotary_embedding.py b/tests/ut/ops/test_rotary_embedding.py index 37db24ec55bc..ec6484701b06 100644 --- a/tests/ut/ops/test_rotary_embedding.py +++ b/tests/ut/ops/test_rotary_embedding.py @@ -21,6 +21,7 @@ import torch from vllm.model_executor.layers.rotary_embedding import RotaryEmbedding, YaRNScalingRotaryEmbedding +from vllm_ascend.ops import rotary_embedding as rotary_embedding_ops from vllm_ascend.ops.rotary_embedding import AscendRotaryEmbedding, AscendYaRNRotaryEmbedding HEAD_SIZE = 64 @@ -32,6 +33,25 @@ NUM_HEADS = 2 +def test_get_identity_cos_and_sin_mla_skips_rotary_cache(monkeypatch): + cos_buffer = torch.ones(4, 1, 1, ROTARY_DIM) + sin_buffer = torch.zeros(4, 1, 1, ROTARY_DIM) + monkeypatch.setattr(rotary_embedding_ops, "_cos_mla", cos_buffer) + monkeypatch.setattr(rotary_embedding_ops, "_sin_mla", sin_buffer) + monkeypatch.setattr(rotary_embedding_ops, "_cos_cache", None) + monkeypatch.setattr(rotary_embedding_ops, "_sin_cache", None) + + cos, sin = rotary_embedding_ops.get_identity_cos_and_sin_mla( + torch.tensor([1, 3]), + use_cache=True, + ) + + assert cos.data_ptr() == cos_buffer.data_ptr() + assert sin.data_ptr() == sin_buffer.data_ptr() + torch.testing.assert_close(cos, torch.ones_like(cos)) + torch.testing.assert_close(sin, torch.zeros_like(sin)) + + def _make_tensors(seq_len=SEQ_LEN, num_heads=NUM_HEADS, head_size=HEAD_SIZE): positions = torch.arange(seq_len, dtype=torch.long) query = torch.randn(seq_len, num_heads * head_size) diff --git a/vllm_ascend/attention/mla_v1.py b/vllm_ascend/attention/mla_v1.py index 0cc6563d7ac8..72f10ab275ba 100644 --- a/vllm_ascend/attention/mla_v1.py +++ b/vllm_ascend/attention/mla_v1.py @@ -46,7 +46,10 @@ from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec from vllm_ascend.device.device_op import DeviceOperator from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.attention_fence import record_attention_compute_start -from vllm_ascend.ops.rotary_embedding import get_cos_and_sin_mla +from vllm_ascend.ops.rotary_embedding import ( + get_cos_and_sin_mla, + get_identity_cos_and_sin_mla, +) from vllm_ascend.quantization.methods.w8a8_mxfp8 import AscendW8A8MXFP8DynamicLinearMethod from vllm_ascend.quantization.methods.w8a8_static import AscendW8A8LinearMethod from vllm_ascend.quantization.utils import enable_fa_quant @@ -217,6 +220,7 @@ class AscendMLAMetadata: attn_mask: torch.Tensor = None # chunked prefill by default if no attn_states passed attn_state: AscendAttentionState = AscendAttentionState.ChunkedPrefill + causal: bool = True decode: AscendMLADecodeMetadata | None = None prefill: AscendMLAPrefillMetadata | None = None @@ -289,6 +293,10 @@ def __init__( self.reorder_batch_threshold = self.decode_threshold self.rope_dim = self.model_config.hf_text_config.qk_rope_head_dim + static_forward_context = vllm_config.compilation_config.static_forward_context + layer_rope_modes = {static_forward_context[layer_name].impl.use_mla_rope for layer_name in layer_names} + assert len(layer_rope_modes) <= 1, "MLA layers with and without RoPE must use separate KV cache groups." + self.use_mla_rope = layer_rope_modes.pop() if layer_rope_modes else True self.cos_cache = None self.sin_cache = None @@ -473,8 +481,8 @@ def build( query_seq_lens_cpu = query_start_loc_cpu[1:] - query_start_loc_cpu[:-1] self.query_lens = query_seq_lens_cpu[:num_reqs] - # Prefer _seq_lens_cpu (always available, updated during draft - # iterations) over seq_lens_cpu (None in async spec decode mode). + # Prefer _seq_lens_cpu, which remains populated in async speculative + # decode, over seq_lens_cpu, which is intentionally None in that mode. if common_attn_metadata._seq_lens_cpu is not None: self.seq_lens = common_attn_metadata._seq_lens_cpu[:num_reqs] elif common_attn_metadata.seq_lens_cpu is not None: @@ -504,6 +512,7 @@ def build( num_prefills=self.num_prefills, attn_mask=self.attn_mask_builder.get_splitfuse_attn_mask(), attn_state=common_attn_metadata.attn_state, + causal=common_attn_metadata.causal, prefill=prefill_metadata, decode=decode_metadata, query_start_loc=query_start_loc, @@ -588,7 +597,8 @@ def build_prefill_metadata( prefill_query_start_loc = query_start_loc[reqs_start:] - query_start_loc[reqs_start] prefill_input_positions = input_positions[tokens_start:] - cos, sin = get_cos_and_sin_mla(prefill_input_positions) + cos_sin_getter = get_cos_and_sin_mla if self.use_mla_rope else get_identity_cos_and_sin_mla + cos, sin = cos_sin_getter(prefill_input_positions) prefill_query_lens = self.query_lens[reqs_start:].to(torch.int32) actual_seq_lengths_q = torch.cumsum(prefill_query_lens, dim=0).tolist() return AscendMLAPrefillMetadata( @@ -672,7 +682,8 @@ def build_decode_metadata( num_reqs_pad_size, num_reqs, actual_seq_lengths_q, common_attn_metadata ) - cos, sin = get_cos_and_sin_mla(input_positions, use_cache=True) + cos_sin_getter = get_cos_and_sin_mla if self.use_mla_rope else get_identity_cos_and_sin_mla + cos, sin = cos_sin_getter(input_positions, use_cache=True) decode_metadata = self.decode_metadata_cls( input_positions=input_positions, block_table=self.block_table, @@ -826,6 +837,9 @@ def __init__( self.q_proj = kwargs["q_proj"] if self.q_lora_rank is None else kwargs["q_b_proj"] self.kv_b_proj = kwargs["kv_b_proj"] self.o_proj = kwargs["o_proj"] + self.g_proj = kwargs["g_proj"] + self.use_output_gate = self.g_proj is not None + self.use_mla_rope = kwargs["use_mla_rope"] self.vllm_config = get_current_vllm_config() self.kv_a_proj_with_mqa = kwargs.get("kv_a_proj_with_mqa") self.kv_a_layernorm = kwargs.get("kv_a_layernorm") @@ -851,6 +865,8 @@ def __init__( # next power of 2. self.num_heads_padded = 1 << (self.num_heads - 1).bit_length() self.head_padding = self.num_heads_padded - self.num_heads + self.mlapo_num_heads = self.num_heads + self.mlapo_weight_quant_mode = 3 @staticmethod def update_graph_params( @@ -914,15 +930,16 @@ def update_graph_params( else: attn_metadata_current = attn_metadata - seq_lens_list = attn_metadata_current[key].decode.seq_lens_list + layer_metadata = attn_metadata_current[key] + seq_lens_list = layer_metadata.decode.seq_lens_list if speculative_config and speculative_config.use_eagle() and not _EXTRA_CTX.is_draft_model: - actual_seq_lengths = attn_metadata_current[key].decode.actual_seq_lengths_q + actual_seq_lengths = layer_metadata.decode.actual_seq_lengths_q spec_multiple = speculative_config.num_speculative_tokens + 1 seq_lens_list = seq_lens_list + [0] * (num_tokens // spec_multiple - len(seq_lens_list)) actual_seq_lengths = [spec_multiple * (i + 1) for i in range(num_tokens // spec_multiple)] elif _EXTRA_CTX.is_draft_model: - actual_seq_lengths = attn_metadata_current[key].decode.actual_seq_lengths_q - block_table = attn_metadata_current[key].decode.block_table + actual_seq_lengths = layer_metadata.decode.actual_seq_lengths_q + block_table = layer_metadata.decode.block_table # TODO: This is a hack and should be fixed in the future. if speculative_config.disable_padded_drafter_batch: block_table = block_table[: len(actual_seq_lengths)] @@ -1038,26 +1055,36 @@ def process_weights_after_loading(self, act_dtype: torch.dtype): else: self.W_UV.copy_(W_UV.transpose(0, 1).contiguous()) self.W_UK_T.copy_(W_UK.permute(1, 2, 0).contiguous()) + self.mlapo_W_UK_T = self.W_UK_T # TODO(zzzzwwjj): Currently, torch.ops._C_ascend.batch_matmul_transpose cannot support weight nz # self.W_UV = maybe_trans_nz(self.W_UV) if self.enable_mlapo: - # Currently mlapo only supports W8A8 and W8A8MXFP8 quantization in MLA scenario - # TODO(whx): modify this limitation when mlapo supports floating point - if self.fused_qkv_a_proj is None or ( - not isinstance( - getattr(self.fused_qkv_a_proj.quant_method, "quant_method", None), AscendW8A8LinearMethod - ) - and not isinstance( - getattr(self.fused_qkv_a_proj.quant_method, "quant_method", None), - AscendW8A8MXFP8DynamicLinearMethod, - ) - ): + device_type = get_ascend_device_type() + layer_quant_method = None if self.fused_qkv_a_proj is None else self.fused_qkv_a_proj.quant_method + if layer_quant_method is None or isinstance(layer_quant_method, UnquantizedLinearMethod): + quant_method = None + else: + # Quantized Ascend linears always expose their concrete scheme + # through AscendLinearMethod.quant_method. Let an unsupported + # wrapper fail here instead of silently disabling MLAPO. + quant_method = layer_quant_method.quant_method + self._mlapo_uses_native_weights = quant_method is None + supports_quantized_weights = isinstance( + quant_method, + (AscendW8A8LinearMethod, AscendW8A8MXFP8DynamicLinearMethod), + ) + supports_native_weights = device_type == AscendDeviceType.A5 and isinstance( + layer_quant_method, + UnquantizedLinearMethod, + ) + if self.fused_qkv_a_proj is None or not (supports_quantized_weights or supports_native_weights): self.enable_mlapo = False logger.warning_once( - "mlapo only supports W8A8 quantization in MLA. " - "Some layers not W8A8 quantized, mlapo disabled for these layers." + "MLAPO supports W8A8/W8A8-MXFP8 weights, plus native " + "floating-point weights on A5. Some layers use an " + "unsupported weight type, so MLAPO is disabled for these layers." ) if self.enable_mlapo or self.fa_quant_layer: self._process_weights_for_fused(act_dtype) @@ -1076,18 +1103,41 @@ def _process_weights_for_fused(self, act_dtype: torch.dtype): self._load_fa_quant_scales() assert self.q_proj is not None - assert hasattr(self.q_proj, "weight_scale") - assert hasattr(self.fused_qkv_a_proj, "weight_scale") assert self.fused_qkv_a_proj is not None - self.weight_dq = self.fused_qkv_a_proj.weight.data[..., : self.q_lora_rank].contiguous() # type: ignore[union-attr] - self.weight_dkv_kr = self.fused_qkv_a_proj.weight.data[..., self.q_lora_rank :].contiguous() # type: ignore[union-attr] - self.weight_uq_qr = self.q_proj.weight.data + is_native = self._mlapo_uses_native_weights + fused_weight = self.fused_qkv_a_proj.weight.data + weight_uq_qr = self.q_proj.weight.data + if is_native: + # Native Linear stores [out_features, in_features], while the + # prolog consumes [in_features, out_features]. + fused_weight = fused_weight.T + weight_uq_qr = weight_uq_qr.T.contiguous() + if self.head_padding > 0: + weight_uq_qr = weight_uq_qr.view( + self.q_lora_rank, + self.num_heads, + self.qk_head_dim, + ) + weight_uq_qr = F.pad(weight_uq_qr, (0, 0, 0, self.head_padding)) + weight_uq_qr = weight_uq_qr.view( + self.q_lora_rank, + self.num_heads_padded * self.qk_head_dim, + ) + + self.weight_dq = fused_weight[..., : self.q_lora_rank].contiguous() + self.weight_dkv_kr = fused_weight[..., self.q_lora_rank :].contiguous() + self.weight_uq_qr = weight_uq_qr.contiguous() self.weight_dq = torch_npu.npu_format_cast(self.weight_dq, ACL_FORMAT_FRACTAL_NZ) - self.weight_uq_qr = torch_npu.npu_format_cast(self.weight_uq_qr.contiguous(), ACL_FORMAT_FRACTAL_NZ) + self.weight_uq_qr = torch_npu.npu_format_cast(self.weight_uq_qr, ACL_FORMAT_FRACTAL_NZ) self.weight_dkv_kr = torch_npu.npu_format_cast(self.weight_dkv_kr, ACL_FORMAT_FRACTAL_NZ) - if get_ascend_device_type() == AscendDeviceType.A5: + self.mlapo_weight_quant_mode = 0 if is_native else 3 + if is_native: + self.mlapo_num_heads = self.num_heads_padded + if self.head_padding > 0: + self.mlapo_W_UK_T = F.pad(self.W_UK_T, (0, 0, 0, 0, 0, self.head_padding)) + elif get_ascend_device_type() == AscendDeviceType.A5: self.dequant_scale_w_uq_qr = self.q_proj.weight_scale.data.transpose(0, 1).flatten(1) weight_scale = self.fused_qkv_a_proj.weight_scale.transpose(0, 1).flatten(1) self.dequant_scale_w_dq = weight_scale[: self.q_lora_rank, ...] @@ -1323,6 +1373,37 @@ def _forward_prefill( return attn_output + def _exec_kv_no_rope( + self, + kv_no_split: torch.Tensor, + kv_cache: tuple, + slots: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Normalize and cache K3 MLA KV without rotating its raw slice.""" + assert self.kv_a_layernorm is not None + assert len(kv_cache) > 1 + num_tokens = kv_no_split.shape[0] + kv_no_split = kv_no_split.view( + num_tokens, + self.num_kv_heads, + self.kv_lora_rank + self.qk_rope_head_dim, + ) + kv_c, k_pe = kv_no_split.split( + [self.kv_lora_rank, self.qk_rope_head_dim], + dim=-1, + ) + kv_c_normed = self.kv_a_layernorm(kv_c.contiguous()) + kv_c_normed = kv_c_normed.view(num_tokens, self.num_kv_heads, self.kv_lora_rank) + k_pe = k_pe.view(num_tokens, self.num_kv_heads, self.qk_rope_head_dim) + DeviceOperator.reshape_and_cache( + key=kv_c_normed, + value=k_pe, + key_cache=kv_cache[0], + value_cache=kv_cache[1], + slot_mapping=slots, + ) + return k_pe, kv_c_normed + def exec_kv_decode( self, kv_no_split: torch.Tensor, @@ -1331,6 +1412,10 @@ def exec_kv_decode( kv_cache: tuple, slots: torch.Tensor, ): + if not self.use_mla_rope: + self._exec_kv_no_rope(kv_no_split, kv_cache, slots) + return kv_cache[1], kv_cache[0] + assert self.kv_a_layernorm is not None B = kv_no_split.shape[0] N = self.num_kv_heads @@ -1365,6 +1450,9 @@ def exec_kv_prefill( *, attn_metadata: AscendMLAMetadata | None = None, ): + if not self.use_mla_rope: + return self._exec_kv_no_rope(kv_no_split, kv_cache, slots) + assert self.kv_a_layernorm is not None B = kv_no_split.shape[0] N = self.num_kv_heads @@ -1396,6 +1484,8 @@ def rope_single( cos: torch.Tensor, sin: torch.Tensor, ) -> torch.Tensor: + if not self.use_mla_rope: + return x B, N, D = x.shape S = 1 x = x.view(B, N, S, D) @@ -1465,8 +1555,15 @@ def _forward_decode( q_nope = F.pad(q_nope, (0, 0, 0, self.head_padding), "constant", 0) # Output shape: [num_heads, num_tokens, dim] attn_output_shape = (self.num_heads_padded, num_tokens, self.kv_lora_rank) - sparse_mode = 3 - attn_mask = attn_metadata.decode.attn_mask # type:ignore + if not attn_metadata.causal: + # K3's DSpark draft block is bidirectional. With FIA this is + # sparse_mode=0 and no mask; a default mask here would hide the + # upper triangle and silently make the draft causal. + sparse_mode = 0 + attn_mask = None + else: + sparse_mode = 3 + attn_mask = decode_meta.attn_mask actual_seq_lengths = decode_meta.actual_seq_lengths_q if self.fa_quant_layer: dequant_scale_q_nope = dequant_scale_q_nope.view(num_tokens, self.num_heads) @@ -1609,31 +1706,53 @@ def _forward_decode( return self._v_up_proj(attn_output) def reorg_decode_q(self, decode_q_nope, decode_q_pe): + if self.mlapo_num_heads > self.num_heads: + decode_q_nope = decode_q_nope[:, : self.num_heads] + decode_q_pe = decode_q_pe[:, : self.num_heads] return decode_q_nope, decode_q_pe def mla_preprocess_only_decode(self, hidden_states, kv_cache, attn_metadata): bsz = attn_metadata.num_decode_tokens - cos_shape = attn_metadata.decode.cos.shape cache_index = attn_metadata.slot_mapping[:bsz].to(torch.int64) decode_k_nope, decode_k_pe = kv_cache[0], kv_cache[1] hidden_states = hidden_states[:bsz] if get_ascend_device_type() == AscendDeviceType.A5: - hidden_states = hidden_states.unsqueeze(1) - quantized_x, dynamic_scale = torch_npu.npu_dynamic_mx_quant(hidden_states, dst_type=torch.float8_e4m3fn) - dequant_scale_x = dynamic_scale.reshape(quantized_x.shape[0] * quantized_x.shape[1], -1).view( - torch.float8_e8m0fnu - ) - dequant_scale_w_dq = self.dequant_scale_w_dq.view(torch.float8_e8m0fnu) - dequant_scale_w_uq_qr = self.dequant_scale_w_uq_qr.view(torch.float8_e8m0fnu) - dequant_scale_w_dkv_kr = self.dequant_scale_w_dkv_kr.view(torch.float8_e8m0fnu) - cos = attn_metadata.decode.cos.view(cos_shape[0], 1, cos_shape[-1]) - sin = attn_metadata.decode.sin.view(cos_shape[0], 1, cos_shape[-1]) - cache_index = cache_index.view(bsz, -1) + if self.mlapo_weight_quant_mode == 0: + quantized_x = hidden_states + dequant_scale_x = None + dequant_scale_w_dq = None + dequant_scale_w_uq_qr = None + dequant_scale_w_dkv_kr = None + else: + hidden_states = hidden_states.unsqueeze(1) + quantized_x, dynamic_scale = torch_npu.npu_dynamic_mx_quant( + hidden_states, dst_type=torch.float8_e4m3fn + ) + dequant_scale_x = dynamic_scale.reshape(quantized_x.shape[0] * quantized_x.shape[1], -1).view( + torch.float8_e8m0fnu + ) + dequant_scale_w_dq = self.dequant_scale_w_dq.view(torch.float8_e8m0fnu) + dequant_scale_w_uq_qr = self.dequant_scale_w_uq_qr.view(torch.float8_e8m0fnu) + dequant_scale_w_dkv_kr = self.dequant_scale_w_dkv_kr.view(torch.float8_e8m0fnu) + if self.use_mla_rope: + cos_shape = attn_metadata.decode.cos.shape + rope_shape = ( + (cos_shape[0], 1, cos_shape[-1]) + if quantized_x.dim() == 3 + else (cos_shape[0], cos_shape[-1]) + ) + cos = attn_metadata.decode.cos.view(rope_shape) + sin = attn_metadata.decode.sin.view(rope_shape) + else: + cos = quantized_x.new_empty((0,), dtype=torch.bfloat16) + sin = quantized_x.new_empty((0,), dtype=torch.bfloat16) + cache_index = cache_index.view(bsz, -1) if quantized_x.dim() == 3 else cache_index.view(-1) cache_mode = "PA_BSND" - weight_quant_mode = 3 + weight_quant_mode = self.mlapo_weight_quant_mode quant_scale_ckv = self.fak_descale_reciprocal if self.fa_quant_layer else None else: + cos_shape = attn_metadata.decode.cos.shape quantized_x, dynamic_scale = torch_npu.npu_dynamic_quant(hidden_states) dequant_scale_x = dynamic_scale.view(-1, 1) dequant_scale_w_dq = self.dequant_scale_w_dq @@ -1653,7 +1772,7 @@ def mla_preprocess_only_decode(self, hidden_states, kv_cache, attn_metadata): token_x=quantized_x, weight_dq=self.weight_dq, weight_uq_qr=self.weight_uq_qr, - weight_uk=self.W_UK_T, + weight_uk=self.mlapo_W_UK_T, weight_dkv_kr=self.weight_dkv_kr, rmsnorm_gamma_cq=self.q_a_layernorm.weight.data, # type: ignore[union-attr] rmsnorm_gamma_ckv=self.kv_a_layernorm.weight.data, # type: ignore[union-attr] @@ -1671,8 +1790,8 @@ def mla_preprocess_only_decode(self, hidden_states, kv_cache, attn_metadata): quant_scale_ckv=quant_scale_ckv, ) - decode_q_nope = decode_q_nope.view(bsz, self.num_heads, self.kv_lora_rank) - decode_q_pe = decode_q_pe.view(bsz, self.num_heads, -1) + decode_q_nope = decode_q_nope.view(bsz, self.mlapo_num_heads, self.kv_lora_rank) + decode_q_pe = decode_q_pe.view(bsz, self.mlapo_num_heads, -1) decode_q_nope, decode_q_pe = self.reorg_decode_q(decode_q_nope, decode_q_pe) decode_preprocess_res = DecodeMLAPreprocessResult( @@ -1834,9 +1953,18 @@ def forward( o_proj_input_shape = (_EXTRA_CTX.num_tokens, self.num_heads * self.v_head_dim) o_proj_input = torch.zeros(o_proj_input_shape, dtype=hidden_states.dtype, device=hidden_states.device) + gate = None + if self.use_output_gate: + assert self.g_proj is not None + gate = self.g_proj(hidden_states.contiguous())[0] + # MLA Preprocess - if (self.fa_quant_layer or self.enable_mlapo) and ( - attn_metadata.num_decode_tokens <= MLAPO_MAX_SUPPORTED_TOKENS and attn_metadata.num_prefills == 0 + can_use_decode_prolog = self.use_mla_rope or get_ascend_device_type() == AscendDeviceType.A5 + if ( + (self.fa_quant_layer or self.enable_mlapo) + and can_use_decode_prolog + and attn_metadata.num_decode_tokens <= MLAPO_MAX_SUPPORTED_TOKENS + and attn_metadata.num_prefills == 0 ): decode_preprocess_res, prefill_preprocess_res = self.mla_preprocess_only_decode( hidden_states, kv_cache, attn_metadata @@ -1874,6 +2002,8 @@ def forward( ) o_proj_input[num_decode_tokens:num_actual_tokens] = output_prefill + if gate is not None: + o_proj_input.mul_(torch.sigmoid(gate)) # O proj output[...] = self.o_proj(o_proj_input, is_prefill=prefill_preprocess_res is not None)[0] diff --git a/vllm_ascend/ops/mla.py b/vllm_ascend/ops/mla.py index 800b09a6f6f0..d656f272d63b 100644 --- a/vllm_ascend/ops/mla.py +++ b/vllm_ascend/ops/mla.py @@ -77,6 +77,8 @@ def __init__( quant_config: QuantizationConfig | None = None, prefix: str = "", skip_topk: bool = False, + non_causal_multi_token_decode: bool = False, + allow_short_prefill_indexer_scoring_skip: bool = False, ) -> None: nn.Module.__init__(self) self.hidden_size = hidden_size @@ -88,6 +90,9 @@ def __init__( self.v_head_dim = v_head_dim self.prefix = prefix self.skip_topk = skip_topk + # This is an upstream CUDA indexer hint. Ascend accepts it to preserve + # constructor compatibility, but its indexer does not consume it. + del allow_short_prefill_indexer_scoring_skip hf_config = get_current_vllm_config().model_config.hf_text_config self.tp_size = get_tensor_model_parallel_world_size() self.layers = hf_config.num_hidden_layers @@ -111,6 +116,7 @@ def __init__( indexer=ascend_indexer, skip_topk=skip_topk, topk_indices_buffer=getattr(mla_modules, "topk_indices_buffer", None), + non_causal_multi_token_decode=non_causal_multi_token_decode, # extra args rotary_emb=mla_modules.rotary_emb, fused_qkv_a_proj=mla_modules.fused_qkv_a_proj, @@ -120,6 +126,8 @@ def __init__( kv_a_proj_with_mqa=mla_modules.kv_a_proj_with_mqa, kv_a_layernorm=mla_modules.kv_a_layernorm, o_proj=mla_modules.o_proj, + g_proj=mla_modules.g_proj, + use_mla_rope=mla_modules.rotary_emb is not None, layer_name=f"{prefix}.attn", ) diff --git a/vllm_ascend/ops/rotary_embedding.py b/vllm_ascend/ops/rotary_embedding.py index 92776d186b83..11243724ad20 100644 --- a/vllm_ascend/ops/rotary_embedding.py +++ b/vllm_ascend/ops/rotary_embedding.py @@ -103,6 +103,15 @@ def get_cos_and_sin_mla(positions, use_cache=False): return _cos_mla[:num_tokens, ...], _sin_mla[:num_tokens, ...] +def get_identity_cos_and_sin_mla(positions, use_cache=False): + """Return the existing stable-shape identity rotation for no-RoPE MLA.""" + del use_cache + global _cos_mla + global _sin_mla + num_tokens = positions.size(0) + return _cos_mla[:num_tokens, ...], _sin_mla[:num_tokens, ...] + + def _record_cos_sin_cache(cos_sin_cache): global _cos_sin_cache if _cos_sin_cache is not None: From 3db03ec3e62fede7a965b9c125fc553b394c80b4 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Mon, 24 Aug 2026 04:23:56 -0500 Subject: [PATCH 06/50] perf(ops): add Kimi K3 attention residual fusion Port the validated attention-residual Triton computation to vLLM 0.27's preallocated contiguous buffer contract and specialize the profile shapes used during graph capture. Signed-off-by: maoxx241 --- .../triton/test_kimi_k3_fusions.py | 97 +++++++++++++++ vllm_ascend/ops/triton/kimi_k3/__init__.py | 1 + .../ops/triton/kimi_k3/attention_residual.py | 111 ++++++++++++++++++ 3 files changed, 209 insertions(+) create mode 100644 tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_kimi_k3_fusions.py create mode 100644 vllm_ascend/ops/triton/kimi_k3/__init__.py create mode 100644 vllm_ascend/ops/triton/kimi_k3/attention_residual.py diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_kimi_k3_fusions.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_kimi_k3_fusions.py new file mode 100644 index 000000000000..67d37349aaa8 --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_kimi_k3_fusions.py @@ -0,0 +1,97 @@ +# Copyright (c) 2026 Huawei Technologies Co., Ltd. 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. +# 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. + +"""Numerical regression coverage for Kimi K3 attention residual fusion.""" + +from types import SimpleNamespace + +import pytest +import torch +import torch_npu # noqa: F401 +from vllm.triton_utils import HAS_TRITON + +if HAS_TRITON: + from vllm_ascend.ops.triton.kimi_k3.attention_residual import apply_attn_res + + +pytestmark = [ + pytest.mark.skipif(not HAS_TRITON, reason="Triton is not available"), + pytest.mark.skipif(not torch.npu.is_available(), reason="NPU required"), + pytest.mark.skip_global_cleanup, +] + + +@torch.inference_mode() +@pytest.mark.parametrize( + ("num_tokens", "num_blocks", "block_capacity"), + [ + pytest.param(7, 4, 7, id="partial-capacity"), + pytest.param(512, 8, 8, id="profile-shape"), + ], +) +def test_kimi_k3_attention_residual_triton_matches_reference( + num_tokens, + num_blocks, + block_capacity, +): + torch.manual_seed(1) + hidden_size = 7168 + eps = 1e-6 + prefix_sum = torch.randn( + (num_tokens, hidden_size), + dtype=torch.bfloat16, + device="npu", + ) + block_residual = torch.randn( + (num_tokens, block_capacity, hidden_size), + dtype=torch.bfloat16, + device="npu", + ) + projection = SimpleNamespace( + weight=torch.randn( + (1, hidden_size), + dtype=torch.bfloat16, + device="npu", + ) + ) + norm = SimpleNamespace( + weight=torch.randn( + (hidden_size,), + dtype=torch.bfloat16, + device="npu", + ), + variance_epsilon=eps, + ) + + actual = apply_attn_res( + prefix_sum, + block_residual, + projection, + norm, + num_blocks, + ) + + values = torch.cat((block_residual[:, :num_blocks, :], prefix_sum.unsqueeze(1)), dim=1).float() + normalized = values * torch.rsqrt(values.square().mean(dim=-1, keepdim=True) + eps) + score_weight = norm.weight.float() * projection.weight.squeeze(0).float() + scores = (normalized * score_weight).sum(dim=-1) + probabilities = scores.softmax(-1).unsqueeze(1) + expected = torch.matmul(probabilities, values).squeeze(1).to(prefix_sum.dtype) + + torch.testing.assert_close( + actual.cpu(), + expected.cpu(), + rtol=1e-2, + atol=1e-2, + ) diff --git a/vllm_ascend/ops/triton/kimi_k3/__init__.py b/vllm_ascend/ops/triton/kimi_k3/__init__.py new file mode 100644 index 000000000000..43e799bd0bd0 --- /dev/null +++ b/vllm_ascend/ops/triton/kimi_k3/__init__.py @@ -0,0 +1 @@ +"""Triton fusion kernels specific to Kimi K3.""" diff --git a/vllm_ascend/ops/triton/kimi_k3/attention_residual.py b/vllm_ascend/ops/triton/kimi_k3/attention_residual.py new file mode 100644 index 000000000000..dac78b9e8627 --- /dev/null +++ b/vllm_ascend/ops/triton/kimi_k3/attention_residual.py @@ -0,0 +1,111 @@ +"""Fused Kimi K3 attention-residual mixture. + +This ports the validated v0.26 computation to vLLM 0.27's preallocated +residual-buffer contract. +""" + +import torch +from vllm.triton_utils import tl, triton + +from vllm_ascend.ops.triton.triton_utils import ( + get_vectorcore_num, + init_device_properties_triton, +) + + +@triton.jit +def _apply_attn_res_kernel( + block_residual_ptr, + prefix_sum_ptr, + norm_w_ptr, + proj_w_ptr, + out_ptr, + N: tl.constexpr, + H: tl.constexpr, + B: tl.constexpr, + BLOCK_CAPACITY: tl.constexpr, + EPS: tl.constexpr, + NUM_CORES: tl.constexpr, + NB: tl.constexpr, +): + block_size = (N - 1) // NUM_CORES + 1 + pid = tl.program_id(0) + tok0 = pid * block_size + if tok0 >= N: + return + tok1 = tl.minimum(tok0 + block_size, N) + + cols = tl.arange(0, H) + s_idx = tl.arange(0, NB) + block_residual_stride = BLOCK_CAPACITY * H + + norm_w = tl.load(norm_w_ptr + cols).to(tl.float32) + proj_w = tl.load(proj_w_ptr + cols).to(tl.float32) + w = norm_w * proj_w + + for tok in range(tok0, tok1): + scores = tl.full([NB], -float("inf"), dtype=tl.float32) + for s in range(B + 1): + if s < B: + v = tl.load(block_residual_ptr + tok * block_residual_stride + s * H + cols).to(tl.float32) + else: + v = tl.load(prefix_sum_ptr + tok * H + cols).to(tl.float32) + ms = tl.sum(v * v) / H + rstd = tl.rsqrt(ms + EPS) + k = v * rstd + scores = tl.where(s_idx == s, tl.sum(k * w), scores) + + scores_max = tl.max(scores) + exp_scores = tl.exp(scores - scores_max) + weights = exp_scores / tl.sum(exp_scores) + + out = tl.zeros([H], dtype=tl.float32) + for s in range(B + 1): + if s < B: + v = tl.load(block_residual_ptr + tok * block_residual_stride + s * H + cols).to(tl.float32) + else: + v = tl.load(prefix_sum_ptr + tok * H + cols).to(tl.float32) + w_s = tl.sum(tl.where(s_idx == s, weights, 0.0)) + out += w_s * v + + tl.store(out_ptr + tok * H + cols, out.to(out_ptr.dtype.element_ty)) + + +def apply_attn_res( + prefix_sum: torch.Tensor, + block_residual: torch.Tensor, + proj: torch.nn.Module, + norm: torch.nn.Module, + num_valid_blocks: int, +) -> torch.Tensor: + """Return K3's learned softmax mixture of residual streams.""" + num_tokens, hidden_size = prefix_sum.shape + block_capacity = block_residual.shape[1] + proj_w = proj.weight.squeeze(0) + norm_w = norm.weight + eps = norm.variance_epsilon + + out = torch.empty( + (num_tokens, hidden_size), + dtype=prefix_sum.dtype, + device=prefix_sum.device, + ) + num_streams = triton.next_power_of_2(num_valid_blocks + 1) + init_device_properties_triton() + num_vectorcore = get_vectorcore_num() + _apply_attn_res_kernel[(num_vectorcore,)]( + block_residual, + prefix_sum, + norm_w, + proj_w, + out, + N=num_tokens, + H=hidden_size, + B=num_valid_blocks, + BLOCK_CAPACITY=block_capacity, + EPS=eps, + NUM_CORES=num_vectorcore, + NB=num_streams, + multibuffer=True, + ) + return out From 88252bc6ec936c935c104e8c628b221fc574655e Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Tue, 25 Aug 2026 03:53:36 -0500 Subject: [PATCH 07/50] docs(ops): document K3 attention residual Signed-off-by: maoxx241 --- .../ops/triton/kimi_k3/attention_residual.md | 37 +++++++++++++++++++ .../ops/triton/kimi_k3/attention_residual.py | 8 +++- 2 files changed, 43 insertions(+), 2 deletions(-) create mode 100644 vllm_ascend/ops/triton/kimi_k3/attention_residual.md diff --git a/vllm_ascend/ops/triton/kimi_k3/attention_residual.md b/vllm_ascend/ops/triton/kimi_k3/attention_residual.md new file mode 100644 index 000000000000..5bc4a4cc608d --- /dev/null +++ b/vllm_ascend/ops/triton/kimi_k3/attention_residual.md @@ -0,0 +1,37 @@ +# Kimi K3 Attention Residual 算子说明 + +## 功能 + +`apply_attn_res` 将 Kimi K3 每个 token 的有效 block residual 与 +`prefix_sum` residual 做可学习的 softmax 加权融合。实现位于 +`vllm_ascend/ops/triton/kimi_k3/attention_residual.py`。 + +对每条 residual stream `v_s`,算子先计算 RMSNorm,再通过 +`norm.weight * proj.weight` 得到标量分数: + +```text +score_s = sum(RMSNorm(v_s) * norm.weight * proj.weight) +weight_s = softmax(score)_s +output = sum(weight_s * v_s) +``` + +## 输入与输出 + +| 参数 | 形状 | 说明 | +| --- | --- | --- | +| `prefix_sum` | `[num_tokens, hidden_size]` | 每个 token 的 prefix-sum residual,也是最后一条参与融合的 stream。 | +| `block_residual` | `[num_tokens, block_capacity, hidden_size]` | vLLM 预分配的 residual buffer。只有前 `num_valid_blocks` 个 block 已初始化。 | +| `proj` | `[1, hidden_size]` | 将归一化 residual 投影为标量分数的线性层。 | +| `norm` | `[hidden_size]` | RMSNorm 权重及 epsilon。 | +| `num_valid_blocks` | `int` | `block_residual` 中有效 block 的数量。 | +| 返回值 | `[num_tokens, hidden_size]` | 所有有效 residual stream 的加权和。 | + +## 实现约束 + +- kernel 的 `B` 等于 `num_valid_blocks`,`BLOCK_CAPACITY` 来自预分配 + buffer 的第二维;kernel 只读取 `[0, B)`,不会读取未初始化容量。 +- `prefix_sum` 使用逻辑索引 `s == B`。启动端将 stream 数设置为 + `next_power_of_2(B + 1)`,因此 `NB` 始终覆盖 `B` 个 block residual + 加一条 prefix stream。 +- softmax 计算使用 FP32,输出再转换为 `prefix_sum` 的 dtype。 +- 启动 grid 使用设备 vector core 数,每个 program 处理一段连续 token。 diff --git a/vllm_ascend/ops/triton/kimi_k3/attention_residual.py b/vllm_ascend/ops/triton/kimi_k3/attention_residual.py index dac78b9e8627..ba148e833ab9 100644 --- a/vllm_ascend/ops/triton/kimi_k3/attention_residual.py +++ b/vllm_ascend/ops/triton/kimi_k3/attention_residual.py @@ -1,7 +1,10 @@ """Fused Kimi K3 attention-residual mixture. -This ports the validated v0.26 computation to vLLM 0.27's preallocated -residual-buffer contract. +For every token, the operator RMS-normalizes each valid block residual and the +prefix-sum residual, projects them to scalar scores, applies a softmax across +those streams, and returns their weighted sum. ``block_residual`` follows +vLLM's preallocated ``[num_tokens, block_capacity, hidden_size]`` contract; +``num_valid_blocks`` identifies the initialized prefix of that capacity. """ import torch @@ -90,6 +93,7 @@ def apply_attn_res( dtype=prefix_sum.dtype, device=prefix_sum.device, ) + # The extra stream is prefix_sum, so NB must cover num_valid_blocks + 1. num_streams = triton.next_power_of_2(num_valid_blocks + 1) init_device_properties_triton() num_vectorcore = get_vectorcore_num() From 7d21d38215b1a44719a9be8ffa85547e40a96f91 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Tue, 25 Aug 2026 03:54:44 -0500 Subject: [PATCH 08/50] fix(ops): enforce attention residual stream capacity Signed-off-by: maoxx241 --- vllm_ascend/ops/triton/kimi_k3/attention_residual.py | 1 + 1 file changed, 1 insertion(+) diff --git a/vllm_ascend/ops/triton/kimi_k3/attention_residual.py b/vllm_ascend/ops/triton/kimi_k3/attention_residual.py index ba148e833ab9..0a666fd766f4 100644 --- a/vllm_ascend/ops/triton/kimi_k3/attention_residual.py +++ b/vllm_ascend/ops/triton/kimi_k3/attention_residual.py @@ -31,6 +31,7 @@ def _apply_attn_res_kernel( NUM_CORES: tl.constexpr, NB: tl.constexpr, ): + tl.static_assert(NB >= B + 1, "NB must include all block residuals and prefix_sum") block_size = (N - 1) // NUM_CORES + 1 pid = tl.program_id(0) tok0 = pid * block_size From bdec553c006bf062331636502c8bb2a25a28787d Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Mon, 24 Aug 2026 04:24:58 -0500 Subject: [PATCH 09/50] feat(moe): support Kimi K3 SiTU execution Propagate SiTU parameters through routed and shared experts, add quantized A2/A3 and A5 execution paths, support ModelSlim mixed KDA projections, and order shared-expert sequence-parallel collectives against routed communication. Signed-off-by: maoxx241 --- tests/ut/ops/test_fused_moe.py | 334 +++++++++++++++++- tests/ut/ops/test_moe_mlp.py | 205 +++++++++++ tests/ut/ops/test_moe_runtime_args.py | 44 +++ .../ut/quantization/test_modelslim_config.py | 119 +++++++ .../ops/fused_moe/dataclass/moe_mlp.py | 6 + vllm_ascend/ops/fused_moe/fused_moe.py | 26 +- vllm_ascend/ops/fused_moe/moe_mlp.py | 200 ++++++++++- vllm_ascend/ops/fused_moe/shared_experts.py | 155 +++++--- vllm_ascend/quantization/modelslim_config.py | 52 ++- 9 files changed, 1079 insertions(+), 62 deletions(-) diff --git a/tests/ut/ops/test_fused_moe.py b/tests/ut/ops/test_fused_moe.py index d60bb486f3c1..a34a74b36d37 100644 --- a/tests/ut/ops/test_fused_moe.py +++ b/tests/ut/ops/test_fused_moe.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -from contextlib import nullcontext +from contextlib import contextmanager, nullcontext from types import SimpleNamespace from unittest.mock import MagicMock, patch @@ -7,6 +7,8 @@ import torch from torch import nn from torch.nn import functional as F +from vllm.config import VllmConfig, set_current_vllm_config +from vllm.model_executor.layers.activation import SituAndMul from vllm_ascend.ascend_forward_context import MoECommType from vllm_ascend.ops.fused_moe import fused_moe as fused_moe_module @@ -855,6 +857,150 @@ def test_shared_experts_part2_applies_optional_gate(with_gate): torch.testing.assert_close(output, expected) +def test_shared_expert_consistency_uses_projection_input_width(monkeypatch): + shared_experts = AscendSharedExperts.__new__(AscendSharedExperts) + shared_experts.hidden_size = 3584 + shared_experts.shared_expert_input_size = 7168 + shared_experts.in_dtype = torch.float16 + output = torch.ones(10, 7168) + shared_experts.layer = MagicMock(return_value=output) + shared_experts.part1 = MagicMock(return_value=output) + shared_experts.part2 = MagicMock(return_value=output) + random_input = torch.ones(10, 7168) + random = MagicMock(return_value=random_input) + monkeypatch.setattr(shared_experts_module.torch, "rand", random) + + shared_experts.validate_consistency() + + random.assert_called_once_with( + 10, + 7168, + device="npu", + dtype=torch.float16, + ) + shared_experts.layer.assert_called_once() + torch.testing.assert_close( + shared_experts.layer.call_args.args[0], + random_input, + ) + + +def _make_quantized_situ_shared_experts(quant_type, gate_up_proj, down_proj): + shared_experts = AscendSharedExperts.__new__(AscendSharedExperts) + shared_experts.layer = SimpleNamespace( + gate_up_proj=gate_up_proj, + down_proj=down_proj, + ) + shared_experts.multistream_overlap = False + shared_experts.quant_type = quant_type + with set_current_vllm_config(VllmConfig()): + shared_experts.situ_activation = SituAndMul(beta=4.0, linear_beta=25.0) + shared_experts.lora_context = None + shared_experts.parallel_mode = MagicMock( + return_value=SharedExpertParallelMode.TENSOR_PARALLEL, + ) + return shared_experts + + +def _make_shared_expert_events(): + event = MagicMock() + return FusedMoEEvents( + before_routed_experts=event, + after_routed_experts=event, + before_dispatch=event, + before_gmm2=event, + before_combine=event, + ) + + +def test_w8a8_shared_situ_uses_dequant_situ_quant(monkeypatch): + gate_up_proj = SimpleNamespace( + weight=torch.ones(4, 4, dtype=torch.int8), + weight_scale=torch.ones(4), + weight_scale_fp32=torch.ones(4), + ) + down_proj = SimpleNamespace( + weight=torch.ones(2, 4, dtype=torch.int8), + weight_scale=torch.ones(2), + ) + shared_experts = _make_quantized_situ_shared_experts(QuantType.W8A8, gate_up_proj, down_proj) + hidden_states = torch.randn(2, 4, dtype=torch.bfloat16) + quantized_input = torch.ones(2, 4, dtype=torch.int8) + input_scale = torch.ones(2) + gate_up_out = torch.ones(2, 4, dtype=torch.int32) + quantized_situ = torch.ones(2, 2, dtype=torch.int8) + situ_scale = torch.ones(2) + expected = torch.randn(2, 2, dtype=torch.bfloat16) + dequant_situ_quant = MagicMock(return_value=(quantized_situ, situ_scale)) + + monkeypatch.setattr(shared_experts_module, "has_lora", lambda _: False) + monkeypatch.setattr(shared_experts_module, "npu_stream_switch", lambda *_args, **_kwargs: nullcontext()) + monkeypatch.setattr(shared_experts_module, "shared_experts_calculation_stream", MagicMock()) + monkeypatch.setattr(shared_experts_module.torch.npu, "current_stream", MagicMock(return_value=MagicMock())) + monkeypatch.setattr( + shared_experts_module.torch_npu, + "npu_dynamic_quant", + MagicMock(return_value=(quantized_input, input_scale)), + ) + monkeypatch.setattr( + shared_experts_module.torch_npu, + "npu_quant_matmul", + MagicMock(side_effect=[gate_up_out, expected]), + ) + monkeypatch.setattr( + shared_experts_module.torch.ops, + "_C_ascend", + SimpleNamespace(dequant_situ_quant=dequant_situ_quant), + ) + + output = shared_experts.forward(hidden_states, _make_shared_expert_events()) + + assert output is expected + situ_kwargs = dequant_situ_quant.call_args.kwargs + assert situ_kwargs["x"] is gate_up_out + assert situ_kwargs["beta"] == 4.0 + assert situ_kwargs["linear_beta"] == 25.0 + + +def test_w4a8_mxfp_shared_situ_uses_situ_mx_quant(monkeypatch): + quantized_input = torch.ones(2, 4, dtype=torch.float8_e4m3fn) + input_scale = torch.ones(2, 1) + gate_up_out = torch.randn(2, 4, dtype=torch.bfloat16) + quantized_situ = torch.ones(2, 2, dtype=torch.float8_e4m3fn) + situ_scale = torch.ones(2, 1) + expected = torch.randn(2, 2, dtype=torch.bfloat16) + gate_up_proj = MagicMock(return_value=(gate_up_out, None)) + gate_up_proj.weight_scale = torch.ones(1) + down_proj = MagicMock(return_value=(expected, None)) + down_proj.weight_scale = torch.ones(1) + shared_experts = _make_quantized_situ_shared_experts(QuantType.W4A8MXFP, gate_up_proj, down_proj) + situ_mx_quant = MagicMock(return_value=(quantized_situ, situ_scale)) + + monkeypatch.setattr(shared_experts_module, "has_lora", lambda _: False) + monkeypatch.setattr(shared_experts_module, "npu_stream_switch", lambda *_args, **_kwargs: nullcontext()) + monkeypatch.setattr(shared_experts_module, "shared_experts_calculation_stream", MagicMock()) + monkeypatch.setattr(shared_experts_module.torch.npu, "current_stream", MagicMock(return_value=MagicMock())) + monkeypatch.setattr( + shared_experts_module.torch_npu, + "npu_dynamic_mx_quant", + MagicMock(return_value=(quantized_input, input_scale)), + ) + monkeypatch.setattr( + shared_experts_module.torch.ops, + "_C_ascend", + SimpleNamespace(situ_mx_quant=situ_mx_quant), + ) + + output = shared_experts.forward(torch.randn(2, 4, dtype=torch.bfloat16), _make_shared_expert_events()) + + assert output is expected + situ_kwargs = situ_mx_quant.call_args.kwargs + assert situ_kwargs["x"] is gate_up_out + assert situ_kwargs["beta"] == 4.0 + assert situ_kwargs["linear_beta"] == 25.0 + assert situ_kwargs["dst_type"] == shared_experts_module.SITU_MX_DST_TYPE_E4M3FN + + @pytest.mark.parametrize( ( "weights_replicated", @@ -959,6 +1105,125 @@ def test_sequence_parallel_only_gathers_unpads_pads_and_reduce_scatters(monkeypa ) +@pytest.mark.parametrize( + ("mode", "multistream_overlap", "starts_all_gather"), + [ + (SharedExpertParallelMode.SEQUENCE_PARALLEL_ONLY, True, True), + (SharedExpertParallelMode.SEQUENCE_PARALLEL_ONLY, False, False), + (SharedExpertParallelMode.SEQUENCE_PARALLEL_SEDP, True, False), + ], +) +def test_prepare_shared_expert_input_only_starts_sp_tp_all_gather( + monkeypatch, + mode, + multistream_overlap, + starts_all_gather, +): + shared_experts = AscendSharedExperts.__new__(AscendSharedExperts) + shared_experts.multistream_overlap = multistream_overlap + shared_experts.parallel_mode = MagicMock(return_value=mode) + hidden_states = torch.randn(2, 4) + gathered_states = torch.randn(4, 4) + shared_experts._gather_sp_input = MagicMock(return_value=gathered_states) + default_stream = MagicMock() + auxiliary_stream = MagicMock() + input_ready = MagicMock() + all_gather_done = MagicMock() + default_stream.record_event.return_value = input_ready + auxiliary_stream.record_event.return_value = all_gather_done + stream_state = {"current": default_stream} + + @contextmanager + def switch_stream(stream, enabled): + previous_stream = stream_state["current"] + if enabled: + stream_state["current"] = stream + try: + yield + finally: + stream_state["current"] = previous_stream + + monkeypatch.setattr(shared_experts_module, "npu_stream_switch", switch_stream) + monkeypatch.setattr(shared_experts_module, "shared_experts_calculation_stream", lambda: auxiliary_stream) + monkeypatch.setattr(shared_experts_module.torch.npu, "current_stream", lambda: stream_state["current"]) + + result, done_event = shared_experts.prepare_input_before_routed_experts(hidden_states) + + if starts_all_gather: + assert result is gathered_states + assert done_event is all_gather_done + default_stream.record_event.assert_called_once_with() + auxiliary_stream.wait_event.assert_called_once_with(input_ready) + shared_experts._gather_sp_input.assert_called_once_with(hidden_states) + auxiliary_stream.record_event.assert_called_once_with() + else: + assert result is hidden_states + assert done_event is None + default_stream.record_event.assert_not_called() + shared_experts._gather_sp_input.assert_not_called() + + +def test_sp_multistream_down_projection_and_reduce_scatter_wait_for_routed_finalize( + monkeypatch, +): + shared_experts = AscendSharedExperts.__new__(AscendSharedExperts) + shared_experts.layer = SimpleNamespace( + gate_up_proj=SimpleNamespace(), + down_proj=SimpleNamespace(), + ) + shared_experts.multistream_overlap = True + shared_experts.quant_type = QuantType.NONE + shared_experts.lora_context = None + shared_experts.parallel_mode = MagicMock(return_value=SharedExpertParallelMode.SEQUENCE_PARALLEL_ONLY) + hidden_states = torch.randn(4, 4) + part1_out = torch.randn(4, 8) + shared_out = torch.randn(4, 4) + reduced_out = torch.randn(2, 4) + shared_experts.part1 = MagicMock(return_value=part1_out) + shared_experts.part2 = MagicMock(return_value=shared_out) + shared_experts._gather_sp_input = MagicMock() + shared_experts._pad_and_reduce_scatter = MagicMock(return_value=reduced_out) + default_stream = MagicMock() + auxiliary_stream = MagicMock() + stream_state = {"current": default_stream} + events = FusedMoEEvents( + before_routed_experts=MagicMock(), + before_dispatch=MagicMock(), + before_combine=MagicMock(), + after_routed_finalize=MagicMock(), + ) + + @contextmanager + def switch_stream(stream, enabled): + previous_stream = stream_state["current"] + if enabled: + stream_state["current"] = stream + try: + yield + finally: + stream_state["current"] = previous_stream + + monkeypatch.setattr(shared_experts_module, "npu_stream_switch", switch_stream) + monkeypatch.setattr(shared_experts_module, "shared_experts_calculation_stream", lambda: auxiliary_stream) + monkeypatch.setattr(shared_experts_module.torch.npu, "current_stream", lambda: stream_state["current"]) + + result = shared_experts.forward( + hidden_states, + events, + input_is_gathered=True, + ) + + assert result is reduced_out + shared_experts._gather_sp_input.assert_not_called() + auxiliary_stream.wait_event.assert_any_call(events.after_routed_finalize) + assert not any( + event_wait.args == (events.before_combine,) for event_wait in auxiliary_stream.wait_event.call_args_list + ) + shared_experts.part2.assert_called_once_with(hidden_states, part1_out) + shared_experts._pad_and_reduce_scatter.assert_called_once_with(shared_out) + default_stream.wait_stream.assert_called_once_with(auxiliary_stream) + + def test_sequence_parallel_sedp_forward_skips_token_comms(monkeypatch): """SP+DP (replicated weights) computes directly on the SP shard without any token gather/scatter.""" @@ -1064,7 +1329,10 @@ def test_forward_impl_returns_current_runner_contract(monkeypatch, has_shared_ex input_ids = torch.tensor([11, 22]) routed_out = torch.randn(2, 4) shared_out = torch.randn(2, 4) - ascend_shared_experts = SimpleNamespace(forward=MagicMock(return_value=shared_out)) + ascend_shared_experts = SimpleNamespace( + prepare_input_before_routed_experts=MagicMock(return_value=(hidden_states, None)), + forward=MagicMock(return_value=shared_out), + ) routed_events = FusedMoEEvents( before_routed_experts=None, after_routed_experts=None, @@ -1090,6 +1358,7 @@ def test_forward_impl_returns_current_runner_contract(monkeypatch, has_shared_ex ) if has_shared_experts: + ascend_shared_experts.prepare_input_before_routed_experts.assert_called_once_with(hidden_states) runner.routed_experts.forward_impl.assert_called_once_with( hidden_states=hidden_states, router_logits=router_logits, @@ -1106,3 +1375,64 @@ def test_forward_impl_returns_current_runner_contract(monkeypatch, has_shared_ex ) assert result is routed_out ascend_shared_experts.forward.assert_not_called() + + +def test_forward_impl_keeps_full_width_input_for_shared_experts(monkeypatch): + runner = AscendMoERunner.__new__(AscendMoERunner) + nn.Module.__init__(runner) + routed_hidden_states = torch.randn(2, 4) + shared_hidden_states = torch.randn(2, 8) + router_logits = torch.randn(2, 3) + routed_out = torch.randn(2, 4) + shared_out = torch.randn(2, 8) + routed_events = FusedMoEEvents( + before_routed_experts=None, + after_routed_experts=None, + before_dispatch=None, + before_gmm2=None, + before_combine=None, + ) + runner.routed_experts = SimpleNamespace(forward_impl=MagicMock(return_value=(routed_out, routed_events))) + shared_input_all_gather_done = MagicMock() + prepared_shared_hidden_states = torch.randn(4, 8) + runner.ascend_shared_experts = SimpleNamespace( + prepare_input_before_routed_experts=MagicMock( + return_value=(prepared_shared_hidden_states, shared_input_all_gather_done), + ), + forward=MagicMock(return_value=shared_out), + ) + runner._sequence_parallel_context = MagicMock(return_value=nullcontext()) + current_stream = MagicMock() + + monkeypatch.setattr( + AscendMoERunner, + "is_internal_router", + property(lambda _: False), + ) + monkeypatch.setattr( + fused_moe_module.torch.npu, + "current_stream", + lambda: current_stream, + ) + + result = runner._forward_impl( + routed_hidden_states, + router_logits, + shared_experts_input=shared_hidden_states, + ) + + runner.routed_experts.forward_impl.assert_called_once_with( + hidden_states=routed_hidden_states, + router_logits=router_logits, + input_ids=None, + ) + runner.ascend_shared_experts.prepare_input_before_routed_experts.assert_called_once_with(shared_hidden_states) + current_stream.wait_event.assert_called_once_with(shared_input_all_gather_done) + assert routed_events.after_routed_finalize is current_stream.record_event.return_value + runner.ascend_shared_experts.forward.assert_called_once_with( + prepared_shared_hidden_states, + routed_events, + input_is_gathered=True, + ) + assert result[0] is shared_out + assert result[1] is routed_out diff --git a/tests/ut/ops/test_moe_mlp.py b/tests/ut/ops/test_moe_mlp.py index b314ae3e0fc3..f3e1ce2b0693 100644 --- a/tests/ut/ops/test_moe_mlp.py +++ b/tests/ut/ops/test_moe_mlp.py @@ -7,6 +7,8 @@ import torch import torch_npu # noqa: F401 -- registers torch.npu used by the module under test from torch.nn import functional as F +from vllm.config import VllmConfig, set_current_vllm_config +from vllm.model_executor.layers.activation import SituAndMul from vllm.model_executor.layers.fused_moe.activation import MoEActivation from vllm_ascend.ascend_forward_context import MoECommType @@ -117,6 +119,51 @@ def test_uses_small_op_activation_and_dynamic_mx_quant(self): class TestUnifiedApplyMlpRequest(unittest.TestCase): + def test_unquant_situ_uses_upstream_activation_contract(self): + hidden_states = torch.randn(2, 8, dtype=torch.bfloat16) + gate_up_out = torch.randn(2, 16, dtype=torch.bfloat16) + expected_output = torch.randn(2, 8, dtype=torch.bfloat16) + with set_current_vllm_config(VllmConfig()): + expected_activation = SituAndMul(beta=4.0, linear_beta=25.0)(gate_up_out) + + with patch( + f"{MOE_MLP}.torch_npu.npu_grouped_matmul", + side_effect=[[gate_up_out], [expected_output]], + create=True, + ) as grouped_matmul: + output, _ = unquant_apply_mlp( + hidden_states=hidden_states, + w1=torch.randn(1, 8, 16), + w2=torch.randn(1, 8, 8), + group_list=torch.tensor([1, 1]), + activation=MoEActivation.SITU, + activation_situ_beta=4.0, + activation_situ_linear_beta=25.0, + need_trans=False, + ) + + self.assertIs(output, expected_output) + torch.testing.assert_close(grouped_matmul.call_args_list[1].kwargs["x"][0], expected_activation) + + def test_quant_situ_dispatches_explicit_beta_parameters(self): + expected = torch.randn(2, 8) + with patch(f"{MOE_MLP}._w4a8_situ_apply_mlp", return_value=expected) as situ_mlp: + output = quant_apply_mlp( + hidden_states=torch.randn(2, 8), + w1=torch.randn(1, 8, 16), + w1_scale=torch.ones(1), + w2=torch.randn(1, 8, 8), + w2_scale=torch.ones(1), + group_list=torch.tensor([1, 1]), + activation=MoEActivation.SITU, + activation_situ_beta=4.0, + activation_situ_linear_beta=25.0, + ) + + self.assertIs(output, expected) + self.assertEqual(situ_mlp.call_args.kwargs["activation_situ_beta"], 4.0) + self.assertEqual(situ_mlp.call_args.kwargs["activation_situ_linear_beta"], 25.0) + def test_unquant_swigluoai_uninterleave_falls_back_on_a5(self): hidden_states = torch.randn(2, 8, dtype=torch.bfloat16) gate_up_out = torch.randn(2, 16, dtype=torch.bfloat16) @@ -407,6 +454,39 @@ def test_request_quant_path_passes_swiglustep_activation(self): self.assertEqual(quant_kwargs["swiglu_limit"], 5.0) mock_unquant.assert_not_called() + def test_request_quant_path_passes_situ_parameters(self): + expected = torch.randn(1, 2) + mlp_compute_input = MoEMlpComputeInput( + hidden_states=torch.ones((1, 2), dtype=torch.float32), + group_list=torch.tensor([1], dtype=torch.int64), + group_list_type=1, + dynamic_scale=None, + topk_scales=None, + weights=MoEWeights( + w1=[torch.ones((1, 2, 4), dtype=torch.float32)], + w2=[torch.ones((1, 2, 2), dtype=torch.float32)], + w1_scale=[torch.ones((1,), dtype=torch.float32)], + w2_scale=[torch.ones((1,), dtype=torch.float32)], + ), + quant=MoEQuantParams(quant_type=QuantType.W8A8), + fusion=False, + activation=MoEActivation.SITU, + activation_situ_beta=4.0, + activation_situ_linear_beta=25.0, + ) + + with ( + patch(f"{MOE_MLP}.quant_apply_mlp", return_value=expected) as mock_quant, + patch(f"{MOE_MLP}.unquant_apply_mlp") as mock_unquant, + ): + output = unified_apply_mlp(mlp_compute_input=mlp_compute_input) + + self.assertIs(output, expected) + self.assertEqual(mock_quant.call_args.kwargs["activation"], MoEActivation.SITU) + self.assertEqual(mock_quant.call_args.kwargs["activation_situ_beta"], 4.0) + self.assertEqual(mock_quant.call_args.kwargs["activation_situ_linear_beta"], 25.0) + mock_unquant.assert_not_called() + def test_request_quant_path_passes_gelu_activation(self): expected = torch.randn(1, 2) mlp_compute_input = MoEMlpComputeInput( @@ -519,6 +599,131 @@ def _common_w8a8_kwargs( ) +class TestQuantApplyMlpSituEplb(_GeluPathBase): + def test_dynamic_eplb_tensor_lists_reach_both_grouped_matmuls(self): + hidden_states = torch.ones(2, 4, dtype=torch.int8) + w1 = [torch.randn(8, 4), torch.randn(8, 4)] + w2 = [torch.randn(4, 4), torch.randn(4, 4)] + w1_scale = [torch.randn(8), torch.randn(8)] + w2_scale = [torch.randn(4), torch.randn(4)] + w1_scale_bias = [torch.randn(8), torch.randn(8)] + w2_scale_bias = [torch.randn(4), torch.randn(4)] + gate_up_out = torch.randn(2, 8, dtype=torch.bfloat16) + quantized_situ_out = torch.ones(2, 4, dtype=torch.int8) + situ_out_scale = torch.ones(2, 1) + expected = torch.randn(2, 4, dtype=torch.bfloat16) + stream_patch, evt = _patch_npu_stream() + + with ( + stream_patch, + patch("torch_npu.npu_grouped_matmul", return_value=[gate_up_out], create=True) as mock_gmm1, + patch( + "torch.ops._C_ascend.dequant_situ_quant", + return_value=(quantized_situ_out, situ_out_scale), + create=True, + ), + patch.object(DeviceOperator, "npu_grouped_matmul_gmm2", return_value=expected) as mock_gmm2, + patch(f"{MOE_MLP}.dispose_tensor"), + ): + output, before_gmm2_evt = quant_apply_mlp( + hidden_states=hidden_states, + w1=w1, + w1_scale=w1_scale, + w2=w2, + w2_scale=w2_scale, + group_list=torch.tensor([1, 1], dtype=torch.int64), + dynamic_scale=torch.ones(2, 1), + w1_scale_bias=w1_scale_bias, + w2_scale_bias=w2_scale_bias, + dynamic_eplb=True, + act_quant_type=torch.int8, + activation=MoEActivation.SITU, + activation_situ_beta=4.0, + activation_situ_linear_beta=25.0, + use_w4a8_per_channel_gmm_swiglu=True, + ) + + self.assertIs(output, expected) + self.assertIs(before_gmm2_evt, evt) + self.assertIs(mock_gmm1.call_args.kwargs["weight"], w1) + self.assertEqual(len(mock_gmm1.call_args.kwargs["scale"]), 2) + self.assertIs(mock_gmm1.call_args.kwargs["bias"], w1_scale_bias) + self.assertIs(mock_gmm2.call_args.kwargs["weight"], w2) + self.assertIs(mock_gmm2.call_args.kwargs["weight_scale"], w2_scale) + self.assertIs(mock_gmm2.call_args.kwargs["bias"], w2_scale_bias) + + def test_antiquant_weights_use_native_situ_between_grouped_matmuls(self): + gate_up_out = torch.tensor([[1.0, -1.0, 0.5, 2.0]]) + gate, up = gate_up_out.chunk(2, dim=-1) + expected_activation = 4.0 * torch.tanh(gate / 4.0) * torch.sigmoid(gate) + expected_activation *= 25.0 * torch.tanh(up / 25.0) + expected = torch.tensor([[3.0]]) + stream_patch, evt = _patch_npu_stream() + + with ( + set_current_vllm_config(VllmConfig()), + stream_patch, + patch(f"{MOE_MLP}._EXTRA_CTX", MagicMock(moe_comm_type=-1)), + patch( + "torch_npu.npu_grouped_matmul", + side_effect=[[gate_up_out], [expected]], + create=True, + ) as grouped_matmul, + patch("torch_npu.npu_dynamic_quant", create=True) as dynamic_quant, + patch.object(DeviceOperator, "npu_grouped_matmul_gmm2") as gmm2, + patch(f"{MOE_MLP}.dispose_tensor"), + ): + kwargs = self._common_w8a8_kwargs(activation=MoEActivation.SITU) + kwargs.update( + activation_situ_beta=4.0, + activation_situ_linear_beta=25.0, + w1_offset=torch.randn(1, 8, 4), + w2_offset=torch.randn(1, 4, 1), + ) + output, before_gmm2_evt = quant_apply_mlp(**kwargs) + + self.assertIs(output, expected) + self.assertIs(before_gmm2_evt, evt) + torch.testing.assert_close(grouped_matmul.call_args_list[1].kwargs["x"][0], expected_activation) + for call in grouped_matmul.call_args_list: + self.assertIn("antiquant_scale", call.kwargs) + self.assertIn("antiquant_offset", call.kwargs) + dynamic_quant.assert_not_called() + gmm2.assert_not_called() + + def test_w4a16_mxfp_uses_native_situ_without_activation_requant(self): + gate_up_out = torch.tensor([[1.0, -1.0, 0.5, 2.0]]) + gate, up = gate_up_out.chunk(2, dim=-1) + expected_activation = 4.0 * torch.tanh(gate / 4.0) * torch.sigmoid(gate) + expected_activation *= 25.0 * torch.tanh(up / 25.0) + expected = torch.tensor([[3.0]]) + stream_patch, evt = _patch_npu_stream() + + with ( + set_current_vllm_config(VllmConfig()), + stream_patch, + patch(f"{MOE_MLP}._EXTRA_CTX", MagicMock(moe_comm_type=-1)), + patch("torch_npu.npu_grouped_matmul", return_value=[gate_up_out], create=True), + patch("torch_npu.npu_dynamic_quant", create=True) as dynamic_quant, + patch.object(DeviceOperator, "npu_grouped_matmul_gmm2", return_value=expected) as gmm2, + patch(f"{MOE_MLP}.dispose_tensor"), + ): + kwargs = self._common_w8a8_kwargs(activation=MoEActivation.SITU) + kwargs.update( + activation_situ_beta=4.0, + activation_situ_linear_beta=25.0, + use_mxfp_quant=True, + mxfp_quant_dtype=QuantType.W4A16MXFP, + ) + output, before_gmm2_evt = quant_apply_mlp(**kwargs) + + self.assertIs(output, expected) + self.assertIs(before_gmm2_evt, evt) + torch.testing.assert_close(gmm2.call_args.kwargs["hidden_states"], expected_activation) + self.assertIsNone(gmm2.call_args.kwargs["per_token_scale"]) + dynamic_quant.assert_not_called() + + class TestQuantApplyMlpMxfpSwigluOAI(_GeluPathBase): def setUp(self): self._ctx_mock = MagicMock() diff --git a/tests/ut/ops/test_moe_runtime_args.py b/tests/ut/ops/test_moe_runtime_args.py index 487656246353..0e4c54665f04 100644 --- a/tests/ut/ops/test_moe_runtime_args.py +++ b/tests/ut/ops/test_moe_runtime_args.py @@ -15,8 +15,10 @@ # limitations under the License. # import unittest +from types import SimpleNamespace import torch +from vllm.model_executor.layers.fused_moe.activation import MoEActivation from vllm_ascend.ops.fused_moe.dataclass.fused_experts import MoEWeights, build_fused_experts_input from vllm_ascend.ops.fused_moe.dataclass.moe_mlp import build_mlp_compute_input @@ -41,6 +43,48 @@ def _get_test_mxfp_dtype(quant_type: QuantType) -> torch.dtype | None: class TestMoERuntimeArgs(unittest.TestCase): + def test_build_mlp_compute_input_preserves_situ_parameters(self): + fused_experts_input = build_fused_experts_input( + hidden_states=torch.randn(2, 4), + topk_weights=torch.ones(2, 1), + topk_ids=torch.zeros(2, 1, dtype=torch.int32), + w1=torch.randn(1, 4, 8), + w2=torch.randn(1, 8, 4), + quant_type=QuantType.NONE, + dynamic_eplb=False, + activation=MoEActivation.SITU, + ) + token_dispatch_output = MoETokenDispatchOutput( + hidden_states=fused_experts_input.hidden_states, + group_list=torch.tensor([2], dtype=torch.int64), + group_list_type=1, + dynamic_scale=None, + combine_metadata=MoEAllGatherCombineMetadata( + topk_weights=fused_experts_input.topk_weights, + expanded_row_idx=torch.arange(2, dtype=torch.int32), + restore_shape=torch.Size([2, 4]), + ), + ) + moe_config = SimpleNamespace( + activation=MoEActivation.SITU, + activation_situ_beta=4.0, + activation_situ_linear_beta=25.0, + swiglu_limit=None, + swiglu_alpha=None, + swiglu_beta=None, + ) + + mlp_compute_input = build_mlp_compute_input( + fused_experts_input=fused_experts_input, + token_dispatch_output=token_dispatch_output, + moe_config=moe_config, + use_fusion_ops=False, + ) + + self.assertEqual(mlp_compute_input.activation, MoEActivation.SITU) + self.assertEqual(mlp_compute_input.activation_situ_beta, 4.0) + self.assertEqual(mlp_compute_input.activation_situ_linear_beta, 25.0) + def test_build_fused_experts_input_preserves_runtime_semantics(self): for quant_type in ( QuantType.NONE, diff --git a/tests/ut/quantization/test_modelslim_config.py b/tests/ut/quantization/test_modelslim_config.py index c0c85ed05643..e56bb0c463db 100644 --- a/tests/ut/quantization/test_modelslim_config.py +++ b/tests/ut/quantization/test_modelslim_config.py @@ -195,6 +195,42 @@ def test_get_quant_method_for_moe_installs_modelslim_weight_loader(self): return_success=True, ) + def test_get_quant_method_for_kimi_linear_moe_uses_kimi_k3_mapping(self): + prefix = "language_model.model.layers.1.block_sparse_moe.experts" + quant_description = {f"{prefix}.0.{name}.weight": "W4A8_DYNAMIC" for name in ("w1", "w2", "w3")} + config = AscendModelSlimConfig(quant_description) + layer = RoutedExperts.__new__(RoutedExperts) + torch.nn.Module.__init__(layer) + layer.moe_config = MagicMock() + layer.weight_loader = MagicMock(return_value=True) + mock_vllm_config = MagicMock() + mock_vllm_config.model_config.hf_config.model_type = "kimi_linear" + mock_scheme = MagicMock() + + with ( + patch( + "vllm_ascend.quantization.modelslim_config.get_current_vllm_config", + return_value=mock_vllm_config, + ), + patch( + "vllm_ascend.quantization.modelslim_config.create_scheme_for_layer", + return_value=mock_scheme, + ) as create_scheme, + patch( + "vllm_ascend.quantization.method_adapters.AscendFusedMoEMethod", + return_value=MagicMock(), + ), + ): + config.get_quant_method(layer, prefix) + + self.assertEqual(config.packed_modules_mapping, get_packed_modules_mapping("kimi_k3")) + create_scheme.assert_called_once_with( + quant_description, + prefix, + "moe", + get_packed_modules_mapping("kimi_k3"), + ) + def test_get_quant_method_for_c8_kv_cache_attention(self): c8_config = AscendModelSlimConfig( { @@ -572,6 +608,89 @@ def test_gemma4_packed_modules_mapping_covers_attention_mlp_and_moe(self): with self.subTest(model_type=model_type): self.assertEqual(get_packed_modules_mapping(model_type), expected_mapping) + def test_kimi_k3_packed_modules_mapping_covers_kda_and_moe(self): + expected_mapping = { + "gate_up_proj": ["gate_proj", "up_proj"], + "experts": ["experts.0.w1", "experts.0.w2", "experts.0.w3"], + "in_proj_qkvgfab": [ + "q_proj", + "k_proj", + "v_proj", + "g_proj", + "f_a_proj", + "b_proj", + ], + "in_proj_qkv": ["q_proj", "k_proj", "v_proj"], + "in_proj_gfab": ["g_proj", "f_a_proj", "b_proj"], + "conv1d": ["q_conv1d", "k_conv1d", "v_conv1d"], + "fused_qkv_a_proj": ["q_a_proj", "kv_a_proj_with_mqa"], + } + + for model_type in ("kimi_k3", "kimi_linear"): + with self.subTest(model_type=model_type): + self.assertEqual(get_packed_modules_mapping(model_type), expected_mapping) + + def test_kimi_k3_modelslim_resolves_fused_kda_and_moe_types(self): + layer_prefix = "language_model.model.layers.1" + quant_description = { + **{ + f"{layer_prefix}.self_attn.{name}.weight": "FLOAT" + for name in ("q_proj", "k_proj", "v_proj", "g_proj", "f_a_proj", "b_proj") + }, + **{ + f"{layer_prefix}.block_sparse_moe.experts.0.{name}.weight": "W4A8_DYNAMIC" + for name in ("w1", "w2", "w3") + }, + } + packed_mapping = get_packed_modules_mapping("kimi_k3") + + self.assertEqual( + get_linear_quant_type( + quant_description, + f"{layer_prefix}.self_attn.in_proj_qkvgfab", + packed_mapping, + ), + "FLOAT", + ) + self.assertEqual( + get_linear_quant_type( + quant_description, + f"{layer_prefix}.block_sparse_moe.experts", + packed_mapping, + ), + "W4A8_DYNAMIC", + ) + + def test_kimi_k3_quarot_splits_mixed_kda_projection(self): + attention_prefix = "language_model.model.layers.0.self_attn" + quant_description = { + **{f"{attention_prefix}.{name}.weight": "W8A8_DYNAMIC" for name in ("q_proj", "k_proj", "v_proj")}, + **{f"{attention_prefix}.{name}.weight": "FLOAT" for name in ("g_proj", "f_a_proj", "b_proj")}, + } + config = AscendModelSlimConfig(quant_description) + fused_prefix = f"{attention_prefix}.in_proj_qkvgfab" + + self.assertTrue(config.uses_kimi_k3_mixed_kda_projection(fused_prefix)) + self.assertEqual( + get_linear_quant_type( + quant_description, + f"{attention_prefix}.in_proj_qkv", + get_packed_modules_mapping("kimi_k3"), + ), + "W8A8_DYNAMIC", + ) + + layer = MagicMock(spec=LinearBase) + mock_vllm_config = MagicMock() + mock_vllm_config.model_config.hf_config.model_type = "kimi_linear" + with patch( + "vllm_ascend.quantization.modelslim_config.get_current_vllm_config", + return_value=mock_vllm_config, + ): + method = config.get_quant_method(layer, fused_prefix) + + self.assertIsInstance(method, AscendUnquantizedLinearMethod) + def test_gemma4_moe_experts_float_shards_are_skipped_together(self): quant_description = { "language_model.model.layers.0.experts.0.gate_proj.weight": "FLOAT", diff --git a/vllm_ascend/ops/fused_moe/dataclass/moe_mlp.py b/vllm_ascend/ops/fused_moe/dataclass/moe_mlp.py index 413c1ac172c2..f29fa6534e20 100644 --- a/vllm_ascend/ops/fused_moe/dataclass/moe_mlp.py +++ b/vllm_ascend/ops/fused_moe/dataclass/moe_mlp.py @@ -48,6 +48,8 @@ class MoEMlpComputeInput: activation: MoEActivation = MoEActivation.SILU need_trans: bool = False dynamic_eplb: bool = False + activation_situ_beta: float | None = None + activation_situ_linear_beta: float | None = None swiglu_limit: float = 0.0 swiglu_alpha: float = 1.0 swiglu_beta: float = 0.0 @@ -73,6 +75,8 @@ def build_mlp_compute_input( if moe_config is None else getattr(moe_config, "activation", fused_experts_input.activation) ) + activation_situ_beta = None if moe_config is None else moe_config.activation_situ_beta + activation_situ_linear_beta = None if moe_config is None else moe_config.activation_situ_linear_beta swiglu_limit = 0.0 if moe_config is None else getattr(moe_config, "swiglu_limit", 0.0) or 0.0 swiglu_alpha = 1.0 if moe_config is None else getattr(moe_config, "swiglu_alpha", 1.0) or 1.0 swiglu_beta = 0.0 if moe_config is None else getattr(moe_config, "swiglu_beta", 0.0) or 0.0 @@ -99,6 +103,8 @@ def build_mlp_compute_input( activation=activation, need_trans=fused_experts_input.need_trans, dynamic_eplb=fused_experts_input.dynamic_eplb, + activation_situ_beta=activation_situ_beta, + activation_situ_linear_beta=activation_situ_linear_beta, swiglu_limit=swiglu_limit, swiglu_alpha=swiglu_alpha, swiglu_beta=swiglu_beta, diff --git a/vllm_ascend/ops/fused_moe/fused_moe.py b/vllm_ascend/ops/fused_moe/fused_moe.py index d6a71f674968..24f30b4d3a51 100644 --- a/vllm_ascend/ops/fused_moe/fused_moe.py +++ b/vllm_ascend/ops/fused_moe/fused_moe.py @@ -223,12 +223,16 @@ def _forward_impl( input_ids: torch.Tensor | None = None, ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: with self._sequence_parallel_context(): + shared_hidden_states = shared_experts_input if shared_experts_input is not None else hidden_states if self.ascend_shared_experts is None: return self.routed_experts.forward_impl( hidden_states=hidden_states, router_logits=router_logits, input_ids=input_ids, ) + shared_expert_input, shared_input_all_gather_done = ( + self.ascend_shared_experts.prepare_input_before_routed_experts(shared_hidden_states) + ) if self.is_internal_router: gate = self.gate assert gate is not None @@ -236,7 +240,7 @@ def _forward_impl( # increase with extra hidden states. We also assume that all gate # linear is unquantized so that we the weight is pre-casted in # process_weights_after_loading of AscendUnquantizedLinearMethod. - hidden_states_fp32 = hidden_states.float() + hidden_states_fp32 = shared_hidden_states.float() before_routed_experts = torch.npu.current_stream().record_event() # v0.27.1: weight_fp32 is guaranteed by is_internal_router. router_logits = F.linear(hidden_states_fp32, gate.weight_fp32) @@ -245,6 +249,8 @@ def _forward_impl( before_routed_experts = torch.npu.current_stream().record_event() after_routed_experts = None + if shared_input_all_gather_done is not None: + torch.npu.current_stream().wait_event(shared_input_all_gather_done) routed_out, fused_moe_events = self.routed_experts.forward_impl( hidden_states=hidden_states, router_logits=router_logits, @@ -252,10 +258,13 @@ def _forward_impl( ) fused_moe_events.before_routed_experts = before_routed_experts fused_moe_events.after_routed_experts = after_routed_experts + if shared_input_all_gather_done is not None: + fused_moe_events.after_routed_finalize = torch.npu.current_stream().record_event() shared_out = self.ascend_shared_experts.forward( - hidden_states, + shared_expert_input, fused_moe_events, + input_is_gathered=shared_input_all_gather_done is not None, ) return shared_out, routed_out @@ -269,6 +278,7 @@ def _forward_impl( input_ids: torch.Tensor | None = None, ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: with self._sequence_parallel_context(): + shared_hidden_states = shared_experts_input if shared_experts_input is not None else hidden_states if self.ascend_shared_experts is None: if self.is_internal_router: gate = self.gate @@ -283,6 +293,9 @@ def _forward_impl( router_logits=router_logits, input_ids=input_ids, ) + shared_expert_input, shared_input_all_gather_done = ( + self.ascend_shared_experts.prepare_input_before_routed_experts(shared_hidden_states) + ) if self.is_internal_router: gate = self.gate assert gate is not None @@ -290,7 +303,7 @@ def _forward_impl( # increase with extra hidden states. We also assume that all gate # linear is unquantized so that we the weight is pre-casted in # process_weights_after_loading of AscendUnquantizedLinearMethod. - hidden_states_fp32 = hidden_states.float() + hidden_states_fp32 = shared_hidden_states.float() before_routed_experts = torch.npu.current_stream().record_event() # main (cdc4824a21): is_internal_router only checks self.gate, # weight_fp32 may be absent, fall back to gate.weight. @@ -303,6 +316,8 @@ def _forward_impl( before_routed_experts = torch.npu.current_stream().record_event() after_routed_experts = None + if shared_input_all_gather_done is not None: + torch.npu.current_stream().wait_event(shared_input_all_gather_done) routed_out, fused_moe_events = self.routed_experts.forward_impl( hidden_states=hidden_states, router_logits=router_logits, @@ -310,9 +325,12 @@ def _forward_impl( ) fused_moe_events.before_routed_experts = before_routed_experts fused_moe_events.after_routed_experts = after_routed_experts + if shared_input_all_gather_done is not None: + fused_moe_events.after_routed_finalize = torch.npu.current_stream().record_event() shared_out = self.ascend_shared_experts.forward( - hidden_states, + shared_expert_input, fused_moe_events, + input_is_gathered=shared_input_all_gather_done is not None, ) return shared_out, routed_out diff --git a/vllm_ascend/ops/fused_moe/moe_mlp.py b/vllm_ascend/ops/fused_moe/moe_mlp.py index f85e440c113a..2fefd1f301b6 100644 --- a/vllm_ascend/ops/fused_moe/moe_mlp.py +++ b/vllm_ascend/ops/fused_moe/moe_mlp.py @@ -18,6 +18,7 @@ import torch import torch_npu from torch.nn.functional import pad +from vllm.model_executor.layers.activation import SituAndMul from vllm.model_executor.layers.fused_moe.activation import MoEActivation from vllm.triton_utils import HAS_TRITON @@ -34,20 +35,20 @@ ) ASCEND_DEVICE_TYPE = get_ascend_device_type() +# CANN uses 36 to select FP8 E4M3FN output for situ_mx_quant. +SITU_MX_DST_TYPE_E4M3FN = 36 def _custom_gmm_swiglu_enabled(fusion, dynamic_eplb, activation=None): - return ( - fusion - and dynamic_eplb - and getattr(activation, "value", activation) != "swigluoai_uninterleave" - and enable_custom_op() - ) + activation_name = getattr(activation, "value", activation) + return fusion and dynamic_eplb and activation_name not in ("situ", "swigluoai_uninterleave") and enable_custom_op() def _gmm_swiglu_quant_fusion_enabled(use_mxfp_quant, fusion, dynamic_eplb, activation=None): - return (use_mxfp_quant or (fusion and not dynamic_eplb)) and ( - getattr(activation, "value", activation) != "swigluoai_uninterleave" + activation_name = getattr(activation, "value", activation) + return (use_mxfp_quant or (fusion and not dynamic_eplb)) and activation_name not in ( + "situ", + "swigluoai_uninterleave", ) @@ -120,6 +121,131 @@ def _prepare_swigluoai_grouped_matmul_scales( return [scale.to(output_dtype) if scale.dtype != output_dtype else scale for scale in scales] +def _as_grouped_matmul_weights( + tensor_or_list: list[torch.Tensor] | torch.Tensor, +) -> list[torch.Tensor]: + return tensor_or_list if isinstance(tensor_or_list, list) else [tensor_or_list] + + +def _quantized_situ_apply_mlp( + *, + hidden_states: torch.Tensor, + w1: list[torch.Tensor] | torch.Tensor, + w1_scale: list[torch.Tensor] | torch.Tensor, + w2: list[torch.Tensor] | torch.Tensor, + w2_scale: list[torch.Tensor] | torch.Tensor, + group_list: torch.Tensor, + group_list_type: int, + dynamic_scale: torch.Tensor | None, + w1_scale_bias: torch.Tensor | None, + w2_scale_bias: torch.Tensor | None, + activation_situ_beta: float, + activation_situ_linear_beta: float | None, + act_quant_type: torch.dtype, + weight_quant_type: torch.dtype | None, + scale_type: torch.dtype | None, + per_token_scale_type: torch.dtype | None, + use_bf16: bool, + use_mxfp_quant: bool, + is_per_channel_weight: bool, + mxfp_quant_dtype: QuantType | None = None, +) -> tuple[torch.Tensor, torch.npu.Event]: + """Run GMM1 -> SiTU quant -> GMM2 without the SwiGLU fusions.""" + input_hidden_dtype = hidden_states.dtype + if dynamic_scale is None: + unquantized_hidden_states = hidden_states + hidden_states, pertoken_scale = DeviceOperator.npu_dynamic_quant( + hidden_states=hidden_states, + dynamic_scale=None, + act_quant_type=act_quant_type, + use_mxfp_quant=False, + ) + dispose_tensor(unquantized_hidden_states) + externally_quantized_hidden_states = None + else: + pertoken_scale = ( + DeviceOperator.maybe_normalize_mxfp_scale_layout(dynamic_scale) if use_mxfp_quant else dynamic_scale + ) + externally_quantized_hidden_states = hidden_states + + w1_scale_list = _as_grouped_matmul_weights(w1_scale) + w2_scale_list = _as_grouped_matmul_weights(w2_scale) + output_dtype = w2_scale_list[0].dtype + bias1, bias2 = None, None + if w1_scale_bias is not None: + if group_list_type == 0: + group_list = torch.cat([group_list[:1], torch.diff(group_list, dim=0)]) + group_list_type = 1 + bias1 = w1_scale_bias + bias2 = w2_scale_bias + output_dtype = torch.bfloat16 + + gmm1_scale = [scale.to(w2_scale_list[0].dtype) for scale in w1_scale_list] + if is_per_channel_weight: + gmm1_scale = [scale.unsqueeze(-2) for scale in gmm1_scale] + + gate_up_out = torch_npu.npu_grouped_matmul( + x=[hidden_states], + weight=_as_grouped_matmul_weights(w1), + antiquant_scale=gmm1_scale if use_mxfp_quant else None, + scale=gmm1_scale if not use_mxfp_quant else None, + bias=bias1, + per_token_scale=[pertoken_scale], + split_item=2, + group_list_type=group_list_type, + group_type=0, + group_list=group_list, + per_token_scale_dtype=torch_npu.float8_e8m0fnu if use_mxfp_quant else None, + weight_dtype=torch_npu.float4_e2m1fn_x2 if use_mxfp_quant else None, + output_dtype=torch.bfloat16 if use_mxfp_quant else output_dtype, + )[0] + if externally_quantized_hidden_states is not None: + dispose_tensor(externally_quantized_hidden_states) + + if use_mxfp_quant: + hidden_states, situ_out_scale = torch.ops._C_ascend.situ_mx_quant( + x=gate_up_out, + beta=activation_situ_beta, + linear_beta=activation_situ_linear_beta or 0.0, + activate_left=True, + dst_type=SITU_MX_DST_TYPE_E4M3FN, + ) + else: + hidden_states, situ_out_scale = torch.ops._C_ascend.dequant_situ_quant( + x=gate_up_out, + weight_scale=None, + activation_scale=None, + bias=None, + quant_scale=None, + quant_offset=None, + group_index=None, + beta=activation_situ_beta, + linear_beta=activation_situ_linear_beta or 0.0, + activate_left=True, + quant_mode="dynamic", + ) + before_gmm2_evt = torch.npu.current_stream().record_event() + hidden_states = DeviceOperator.npu_grouped_matmul_gmm2( + hidden_states=hidden_states, + weight=w2, + weight_scale=w2_scale, + per_token_scale=situ_out_scale, + group_list=group_list, + group_list_type=group_list_type, + input_dtype=input_hidden_dtype, + act_quant_type=act_quant_type, + weight_quant_type=weight_quant_type, + scale_type=scale_type, + per_token_scale_type=per_token_scale_type, + use_bf16=use_bf16, + use_mxfp_quant=use_mxfp_quant, + bias=bias2, + fallback_output_dtype=output_dtype, + mxfp_quant_dtype=mxfp_quant_dtype, + ) + return hidden_states, before_gmm2_evt + + def _apply_clipped_swiglu( hidden_states: torch.Tensor, *, @@ -193,13 +319,41 @@ def quant_apply_mlp( scale_type: torch.dtype | None = None, per_token_scale_type: torch.dtype | None = None, use_bf16: bool = True, - activation: str | None = None, + activation: str | MoEActivation | None = None, + activation_situ_beta: float | None = None, + activation_situ_linear_beta: float | None = None, swiglu_limit: float = 0.0, swiglu_alpha: float = 1.0, swiglu_beta: float = 0.0, use_w4a8_per_channel_gmm_swiglu: bool = False, ) -> torch.Tensor: input_hidden_dtype = hidden_states.dtype + situ_beta = 1.0 if activation_situ_beta is None else activation_situ_beta + if activation == MoEActivation.SITU: + use_antiquant_situ = mxfp_quant_dtype == QuantType.W4A16MXFP or w1_offset is not None + if not use_antiquant_situ: + return _quantized_situ_apply_mlp( + hidden_states=hidden_states, + w1=w1, + w1_scale=w1_scale, + w2=w2, + w2_scale=w2_scale, + group_list=group_list, + group_list_type=group_list_type, + dynamic_scale=dynamic_scale, + w1_scale_bias=w1_scale_bias, + w2_scale_bias=w2_scale_bias, + activation_situ_beta=situ_beta, + activation_situ_linear_beta=activation_situ_linear_beta, + act_quant_type=act_quant_type, + weight_quant_type=weight_quant_type, + scale_type=scale_type, + per_token_scale_type=per_token_scale_type, + use_bf16=use_bf16, + is_per_channel_weight=use_w4a8_per_channel_gmm_swiglu, + use_mxfp_quant=use_mxfp_quant, + mxfp_quant_dtype=mxfp_quant_dtype, + ) act_name = getattr(activation, "value", activation) use_gmm_swiglu_quant_fusion = _gmm_swiglu_quant_fusion_enabled( use_mxfp_quant, @@ -399,6 +553,11 @@ def quant_apply_mlp( gate, up = hidden_states.chunk(2, dim=-1) approximate = "tanh" if activation == MoEActivation.GELU_TANH else "none" hidden_states = torch.nn.functional.gelu(gate, approximate=approximate) * up + elif activation == MoEActivation.SITU: + hidden_states = SituAndMul( + beta=situ_beta, + linear_beta=activation_situ_linear_beta, + )(hidden_states) elif is_swigluoai_uninterleave: hidden_states = _apply_clipped_swiglu( hidden_states, @@ -519,6 +678,12 @@ def quant_apply_mlp( approximate = "tanh" if activation == MoEActivation.GELU_TANH else "none" hidden_states = torch.nn.functional.gelu(gate, approximate=approximate) * up hidden_states, swiglu_out_scale = torch_npu.npu_dynamic_quant(hidden_states) + elif activation == MoEActivation.SITU: + hidden_states = SituAndMul( + beta=situ_beta, + linear_beta=activation_situ_linear_beta, + )(hidden_states) + swiglu_out_scale = None elif is_swigluoai_uninterleave: if use_mxfp_quant: hidden_states, swiglu_out_scale = _swiglu_oai_dynamic_mx_quant( @@ -577,7 +742,9 @@ def unquant_apply_mlp( group_list: torch.Tensor, w1_bias: torch.Tensor = None, w2_bias: torch.Tensor = None, - activation: str | None = None, + activation: str | MoEActivation | None = None, + activation_situ_beta: float | None = None, + activation_situ_linear_beta: float | None = None, group_list_type: int = 1, topk_scales: torch.Tensor | None = None, need_trans: bool = True, @@ -641,7 +808,12 @@ def unquant_apply_mlp( ) act_name = getattr(activation, "value", activation) - if activation == MoEActivation.SWIGLUOAI: + if activation == MoEActivation.SITU: + gate_up_out = SituAndMul( + beta=activation_situ_beta, + linear_beta=activation_situ_linear_beta, + )(gate_up_out) + elif activation == MoEActivation.SWIGLUOAI: num_experts, _, hidden_size = w1.shape gate_up_out = AscendSwigluOAIAndMul.swiglu_oai_forward(gate_up_out.view(-1, hidden_size)) elif act_name == "swigluoai_uninterleave": @@ -714,6 +886,8 @@ def unified_apply_mlp(*, mlp_compute_input: MoEMlpComputeInput) -> torch.Tensor: swiglu_limit = mlp_compute_input.swiglu_limit swiglu_alpha = mlp_compute_input.swiglu_alpha swiglu_beta = mlp_compute_input.swiglu_beta + activation_situ_beta = mlp_compute_input.activation_situ_beta + activation_situ_linear_beta = mlp_compute_input.activation_situ_linear_beta if not mlp_compute_input.quant.is_quant: return unquant_apply_mlp( @@ -723,6 +897,8 @@ def unified_apply_mlp(*, mlp_compute_input: MoEMlpComputeInput) -> torch.Tensor: w1_bias=w1_bias, w2_bias=w2_bias, activation=activation, + activation_situ_beta=activation_situ_beta, + activation_situ_linear_beta=activation_situ_linear_beta, group_list=group_list, group_list_type=group_list_type, topk_scales=topk_scales, @@ -787,6 +963,8 @@ def unified_apply_mlp(*, mlp_compute_input: MoEMlpComputeInput) -> torch.Tensor: per_token_scale_type=per_token_scale_type, use_bf16=use_bf16, activation=activation, + activation_situ_beta=activation_situ_beta, + activation_situ_linear_beta=activation_situ_linear_beta, swiglu_limit=swiglu_limit, swiglu_alpha=swiglu_alpha, swiglu_beta=swiglu_beta, diff --git a/vllm_ascend/ops/fused_moe/shared_experts.py b/vllm_ascend/ops/fused_moe/shared_experts.py index 14f72b45ed98..3dfc78d8fd4a 100644 --- a/vllm_ascend/ops/fused_moe/shared_experts.py +++ b/vllm_ascend/ops/fused_moe/shared_experts.py @@ -23,6 +23,7 @@ import torch_npu from vllm.distributed import tensor_model_parallel_all_gather, tensor_model_parallel_reduce_scatter from vllm.logger import logger +from vllm.model_executor.layers.activation import SituAndMul from vllm.model_executor.layers.fused_moe import FusedMoEConfig, FusedMoEMethodBase from vllm_ascend.ascend_config import get_ascend_config @@ -36,6 +37,9 @@ shared_experts_calculation_stream, ) +# CANN uses 36 to select FP8 E4M3FN output for situ_mx_quant. +SITU_MX_DST_TYPE_E4M3FN = 36 + @dataclass class FusedMoEEvents: @@ -44,6 +48,7 @@ class FusedMoEEvents: before_dispatch: torch.npu.Event | None = field(default=None) before_gmm2: torch.npu.Event | None = field(default=None) before_combine: torch.npu.Event | None = field(default=None) + after_routed_finalize: torch.npu.Event | None = field(default=None) class SharedExpertParallelMode(Enum): @@ -73,11 +78,17 @@ def __init__( self.layer = layer self.moe_config = moe_config self.hidden_size = moe_config.hidden_dim + self.shared_expert_input_size = getattr( + layer.gate_up_proj, + "input_size", + self.hidden_size, + ) self.in_dtype = moe_config.in_dtype self.swiglu_limit = 0.0 if moe_config.swiglu_limit is None else moe_config.swiglu_limit self.swiglu_alpha = 1.0 if moe_config.swiglu_alpha is None else moe_config.swiglu_alpha self.swiglu_beta = 0.0 if moe_config.swiglu_beta is None else moe_config.swiglu_beta self.is_sequence_parallel = moe_config.is_sequence_parallel + self.situ_activation = layer.act_fn if isinstance(layer.act_fn, SituAndMul) else None self.quant_type = quant_type self.lora_context = None ascend_config = get_ascend_config() @@ -105,7 +116,14 @@ def set_lora_context(self, lora_context) -> None: def validate_consistency(self): """Validate that split shared expert computation matches integrated computation.""" test_input = ( - torch.rand(10, self.hidden_size, device="npu", dtype=self.in_dtype) * 2 - 1 + torch.rand( + 10, + self.shared_expert_input_size, + device="npu", + dtype=self.in_dtype, + ) + * 2 + - 1 ) # Random input for testing, scoped to [-1, 1] integrated_out = self.layer(test_input) @@ -124,7 +142,7 @@ def validate_consistency(self): integrated_out.norm().item(), split_out.sum().item(), split_out.norm().item(), - self.hidden_size, + self.shared_expert_input_size, self.in_dtype, ) raise ValueError("FusedMoE shared experts split computation does not match the integrated computation.") @@ -225,9 +243,34 @@ def _pad_and_reduce_scatter(self, shared_out: torch.Tensor) -> torch.Tensor: shared_out = F.pad(shared_out, (0, 0, 0, pad_size)) return tensor_model_parallel_reduce_scatter(shared_out, dim=0) - def forward(self, hidden_states: torch.Tensor, fused_moe_evts: FusedMoEEvents): + def prepare_input_before_routed_experts( + self, + hidden_states: torch.Tensor, + ) -> tuple[torch.Tensor, torch.npu.Event | None]: + """Start the SP-only input all-gather on the shared-expert stream.""" + if not (self.multistream_overlap and self.parallel_mode() is SharedExpertParallelMode.SEQUENCE_PARALLEL_ONLY): + return hidden_states, None + + input_ready = torch.npu.current_stream().record_event() + with npu_stream_switch(shared_experts_calculation_stream(), enabled=True): + torch.npu.current_stream().wait_event(input_ready) + hidden_states = self._gather_sp_input(hidden_states) + all_gather_done = torch.npu.current_stream().record_event() + return hidden_states, all_gather_done + + def forward( + self, + hidden_states: torch.Tensor, + fused_moe_evts: FusedMoEEvents, + input_is_gathered: bool = False, + ): mode = self.parallel_mode() local_dp_metadata = None + down_projection_ready = ( + fused_moe_evts.after_routed_finalize + if self.multistream_overlap and mode is SharedExpertParallelMode.SEQUENCE_PARALLEL_ONLY + else fused_moe_evts.before_combine + ) def maybe_wait_event(evt: torch.npu.Event | None): if evt is not None: @@ -243,8 +286,9 @@ def maybe_wait_event(evt: torch.npu.Event | None): # Sharded activations + TP-sharded weights: gather the SP # shard to full activations before the MLP; the output is # padded and reduce-scattered back below. - maybe_wait_event(fused_moe_evts.before_routed_experts) - hidden_states = self._gather_sp_input(hidden_states) + if not input_is_gathered: + maybe_wait_event(fused_moe_evts.before_routed_experts) + hidden_states = self._gather_sp_input(hidden_states) # Only used for int quantization has_quantized_shared_without_lora = ( not has_lora(self.lora_context) @@ -270,27 +314,40 @@ def maybe_wait_event(evt: torch.npu.Event | None): # Execute activation concurrently with gmm2. maybe_wait_event(fused_moe_evts.before_gmm2) - quantized_x, swiglu_out_scale = torch.ops._C_ascend.npu_dequant_swiglu_quant( - x=hidden_states, - weight_scale=self.layer.gate_up_proj.weight_scale_fp32, - activation_scale=pertoken_scale, - bias=None, - quant_scale=None, - quant_offset=None, - group_index=None, - activate_left=True, - quant_mode=1, - swiglu_mode=1, - clamp_limit=self.swiglu_limit, - **( - {} - if get_ascend_device_type() == AscendDeviceType.A5 - else {"glu_alpha": self.swiglu_alpha, "glu_bias": self.swiglu_beta} - ), - ) - # Execute the down projection concurrently with the combine - # communication. - maybe_wait_event(fused_moe_evts.before_combine) + if self.situ_activation is not None: + quantized_x, swiglu_out_scale = torch.ops._C_ascend.dequant_situ_quant( + x=hidden_states, + weight_scale=self.layer.gate_up_proj.weight_scale_fp32, + activation_scale=pertoken_scale, + bias=None, + quant_scale=None, + quant_offset=None, + group_index=None, + beta=self.situ_activation.beta, + linear_beta=self.situ_activation.linear_beta or 0.0, + activate_left=True, + quant_mode="dynamic", + ) + else: + quantized_x, swiglu_out_scale = torch.ops._C_ascend.npu_dequant_swiglu_quant( + x=hidden_states, + weight_scale=self.layer.gate_up_proj.weight_scale_fp32, + activation_scale=pertoken_scale, + bias=None, + quant_scale=None, + quant_offset=None, + group_index=None, + activate_left=True, + quant_mode=1, + swiglu_mode=1, + clamp_limit=self.swiglu_limit, + **( + {} + if get_ascend_device_type() == AscendDeviceType.A5 + else {"glu_alpha": self.swiglu_alpha, "glu_bias": self.swiglu_beta} + ), + ) + maybe_wait_event(down_projection_ready) shared_out = torch_npu.npu_quant_matmul( quantized_x, self.layer.down_proj.weight, @@ -312,17 +369,24 @@ def maybe_wait_event(evt: torch.npu.Event | None): hidden_states = self.layer.gate_up_proj((quantized_x, pertoken_scale))[0] # Execute activation concurrently with gmm2. maybe_wait_event(fused_moe_evts.before_gmm2) - quantized_x, swiglu_out_scale, _ = torch.ops._C_ascend.npu_swiglu_group_quant( - hidden_states, - topk_weight=None, - group_index=None, - dst_type=torch.float8_e4m3fn, - quant_mode=2, - clamp_value=self.swiglu_limit, - ) - # Execute the down projection concurrently with the combine - # communication. - maybe_wait_event(fused_moe_evts.before_combine) + if self.situ_activation is not None: + quantized_x, swiglu_out_scale = torch.ops._C_ascend.situ_mx_quant( + x=hidden_states, + beta=self.situ_activation.beta, + linear_beta=self.situ_activation.linear_beta or 0.0, + activate_left=True, + dst_type=SITU_MX_DST_TYPE_E4M3FN, + ) + else: + quantized_x, swiglu_out_scale, _ = torch.ops._C_ascend.npu_swiglu_group_quant( + hidden_states, + topk_weight=None, + group_index=None, + dst_type=torch.float8_e4m3fn, + quant_mode=2, + clamp_value=self.swiglu_limit, + ) + maybe_wait_event(down_projection_ready) shared_out = self.layer.down_proj((quantized_x, swiglu_out_scale))[0] else: # Ensure the shared experts wait for hidden_states to be ready. @@ -331,19 +395,24 @@ def maybe_wait_event(evt: torch.npu.Event | None): # dispatch communication. maybe_wait_event(fused_moe_evts.before_dispatch) part1_out = self.part1(hidden_states) - # Execute the down projection concurrently with the combine - # communication. - maybe_wait_event(fused_moe_evts.before_combine) + maybe_wait_event(down_projection_ready) shared_out = self.part2(hidden_states, part1_out) - # Make sure the default stream waits for the shared experts stream to - # finish. + if self.multistream_overlap and mode is SharedExpertParallelMode.SEQUENCE_PARALLEL_ONLY: + # Keep the shared-expert output collective on the auxiliary stream, + # but start it only after routed dispatch/combine/finalize + # communication has completed. + with npu_stream_switch(shared_experts_calculation_stream(), enabled=True): + maybe_wait_event(fused_moe_evts.after_routed_finalize) + shared_out = self._pad_and_reduce_scatter(shared_out) + + # Make sure the default stream waits for the shared experts stream to finish. if self.multistream_overlap: torch.npu.current_stream().wait_stream(shared_experts_calculation_stream()) if mode is SharedExpertParallelMode.SHARED_EXPERT_DATA_PARALLEL_ONLY: assert local_dp_metadata is not None shared_out = self._finalize_local_dp_output(shared_out, local_dp_metadata) - elif mode is SharedExpertParallelMode.SEQUENCE_PARALLEL_ONLY: + elif mode is SharedExpertParallelMode.SEQUENCE_PARALLEL_ONLY and not self.multistream_overlap: shared_out = self._pad_and_reduce_scatter(shared_out) return shared_out diff --git a/vllm_ascend/quantization/modelslim_config.py b/vllm_ascend/quantization/modelslim_config.py index fbd0b9188d9c..08912fff59a8 100644 --- a/vllm_ascend/quantization/modelslim_config.py +++ b/vllm_ascend/quantization/modelslim_config.py @@ -425,6 +425,27 @@ def get_packed_modules_mapping(model_type: str) -> dict[str, list[str]]: Dictionary mapping fused module names to their component module names. Returns empty dict if model_type is not found. """ + if model_type in ("kimi_k3", "kimi_linear"): + # Start from vLLM's model-owned mapping so changes to its packed KDA, + # MLA, convolution, or MLP modules remain the source of truth. Add only + # the ModelSlim MoE convention and Ascend's mixed-precision KDA split. + from vllm.models.kimi_k3.nvidia.model import KimiLinearModel + + mapping = { + packed_name: list(shard_names) + for packed_name, shard_names in KimiLinearModel.packed_modules_mapping.items() + } + qkv_shards = [name for name in mapping["in_proj_qkvgfab"] if name in ("q_proj", "k_proj", "v_proj")] + gate_shards = ["g_proj", "f_a_proj", "b_proj"] + mapping.update( + { + "experts": ["experts.0.w1", "experts.0.w2", "experts.0.w3"], + "in_proj_qkvgfab": qkv_shards + gate_shards, + "in_proj_qkv": qkv_shards, + "in_proj_gfab": gate_shards, + } + ) + return mapping return packed_modules_model_mapping.get(model_type, {}) @@ -670,6 +691,25 @@ def get_cache_scale_mapper(self) -> "WeightsMapper": cache_scale_mapper = WeightsMapper(orig_to_new_suffix=suffix_map, orig_to_new_regex=regex_map) return cache_scale_mapper | QuantizationConfig.get_cache_scale_mapper() + def uses_kimi_k3_mixed_kda_projection(self, prefix: str) -> bool: + """Return whether Kimi K3's packed KDA input must be split. + + vLLM 0.27 packs q/k/v and the full-rank KDA gates into one linear + layer. ModelSlim QuaRot checkpoints intentionally keep q/k/v W8A8 + while storing g/f_a/b in floating point, which cannot be represented + by one quantization method. Only accept that known layout here; other + mixed packed layouts retain the normal validation error. + """ + suffix = ".in_proj_qkvgfab" + attention_prefix = prefix[: -len(suffix)] if prefix.endswith(suffix) else prefix + quant_types = { + name: self.quant_description.get(f"{attention_prefix}.{name}.weight") + for name in ("q_proj", "k_proj", "v_proj", "g_proj", "f_a_proj", "b_proj") + } + qkv_types = {quant_types[name] for name in ("q_proj", "k_proj", "v_proj")} + gate_types = {quant_types[name] for name in ("g_proj", "f_a_proj", "b_proj")} + return len(qkv_types) == 1 and None not in qkv_types and qkv_types != {"FLOAT"} and gate_types == {"FLOAT"} + def _has_quant_weight(self, prefix: str, packed_modules_mapping: Mapping[str, list[str]]) -> bool: proj_name = prefix.split(".")[-1] if proj_name in packed_modules_mapping: @@ -740,11 +780,19 @@ def get_quant_method(self, layer: torch.nn.Module, prefix: str, tid2eid=None) -> # Adapt to bailing_hybrid architecture: update layer names to MoE convention prefix = prefix.replace("linear_attn", "attention") prefix = prefix.replace("self_attn", "attention") - if model_type in packed_modules_model_mapping: - self.packed_modules_mapping = packed_modules_model_mapping[model_type] + if model_type in packed_modules_model_mapping or model_type in ("kimi_k3", "kimi_linear"): + self.packed_modules_mapping = get_packed_modules_mapping(model_type) prefix = self.quant_prefix_mapper(model_type, prefix) if isinstance(layer, LinearBase): + if model_type in ("kimi_k3", "kimi_linear") and self.uses_kimi_k3_mixed_kda_projection(prefix): + # The Ascend K3 adapter replaces this temporary module with a + # W8A8 q/k/v projection plus a FLOAT gate projection directly + # after upstream construction. + from vllm_ascend.ops.linear import AscendUnquantizedLinearMethod + + logger.debug("Temporarily select unquantized Kimi K3 mixed KDA projection for %s", prefix) + return AscendUnquantizedLinearMethod() if self.is_layer_skipped_ascend(prefix, self.packed_modules_mapping): # Delayed import to avoid circular import from vllm_ascend.ops.linear import AscendUnquantizedLinearMethod From 8fa59cdd0c9dc12004f61023ea07888295e9d91c Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Tue, 25 Aug 2026 03:53:36 -0500 Subject: [PATCH 10/50] refactor(moe): integrate SiTU quantization into common flow Signed-off-by: maoxx241 --- tests/ut/ops/test_moe_mlp.py | 56 +- vllm_ascend/ops/fused_moe/moe_mlp.py | 1838 +++++++++--------- vllm_ascend/quantization/modelslim_config.py | 13 +- 3 files changed, 908 insertions(+), 999 deletions(-) diff --git a/tests/ut/ops/test_moe_mlp.py b/tests/ut/ops/test_moe_mlp.py index f3e1ce2b0693..643d4fb327c1 100644 --- a/tests/ut/ops/test_moe_mlp.py +++ b/tests/ut/ops/test_moe_mlp.py @@ -145,25 +145,6 @@ def test_unquant_situ_uses_upstream_activation_contract(self): self.assertIs(output, expected_output) torch.testing.assert_close(grouped_matmul.call_args_list[1].kwargs["x"][0], expected_activation) - def test_quant_situ_dispatches_explicit_beta_parameters(self): - expected = torch.randn(2, 8) - with patch(f"{MOE_MLP}._w4a8_situ_apply_mlp", return_value=expected) as situ_mlp: - output = quant_apply_mlp( - hidden_states=torch.randn(2, 8), - w1=torch.randn(1, 8, 16), - w1_scale=torch.ones(1), - w2=torch.randn(1, 8, 8), - w2_scale=torch.ones(1), - group_list=torch.tensor([1, 1]), - activation=MoEActivation.SITU, - activation_situ_beta=4.0, - activation_situ_linear_beta=25.0, - ) - - self.assertIs(output, expected) - self.assertEqual(situ_mlp.call_args.kwargs["activation_situ_beta"], 4.0) - self.assertEqual(situ_mlp.call_args.kwargs["activation_situ_linear_beta"], 25.0) - def test_unquant_swigluoai_uninterleave_falls_back_on_a5(self): hidden_states = torch.randn(2, 8, dtype=torch.bfloat16) gate_up_out = torch.randn(2, 16, dtype=torch.bfloat16) @@ -652,6 +633,43 @@ def test_dynamic_eplb_tensor_lists_reach_both_grouped_matmuls(self): self.assertIs(mock_gmm2.call_args.kwargs["weight_scale"], w2_scale) self.assertIs(mock_gmm2.call_args.kwargs["bias"], w2_scale_bias) + def test_w4a8_mxfp_situ_stays_in_common_grouped_matmul_flow(self): + gate_up_out = torch.randn(2, 8, dtype=torch.bfloat16) + quantized_situ_out = torch.ones(2, 4, dtype=torch.float8_e4m3fn) + situ_out_scale = torch.ones(2, 1) + expected = torch.randn(2, 4, dtype=torch.bfloat16) + stream_patch, evt = _patch_npu_stream() + + with ( + stream_patch, + patch(f"{MOE_MLP}._EXTRA_CTX", MagicMock(moe_comm_type=-1)), + patch("torch_npu.npu_grouped_matmul", return_value=[gate_up_out], create=True) as mock_gmm1, + patch( + "torch.ops._C_ascend.situ_mx_quant", + return_value=(quantized_situ_out, situ_out_scale), + create=True, + ) as situ_mx_quant, + patch.object(DeviceOperator, "maybe_normalize_mxfp_scale_layout", side_effect=lambda scale: scale), + patch.object(DeviceOperator, "npu_grouped_matmul_gmm2", return_value=expected) as mock_gmm2, + patch(f"{MOE_MLP}.dispose_tensor"), + ): + kwargs = self._common_w8a8_kwargs(activation=MoEActivation.SITU) + kwargs.update( + activation_situ_beta=4.0, + activation_situ_linear_beta=25.0, + use_mxfp_quant=True, + mxfp_quant_dtype=QuantType.W4A8MXFP, + ) + output, before_gmm2_evt = quant_apply_mlp(**kwargs) + + self.assertIs(output, expected) + self.assertIs(before_gmm2_evt, evt) + self.assertIsNone(mock_gmm1.call_args.kwargs["scale"]) + self.assertIsNotNone(mock_gmm1.call_args.kwargs["antiquant_scale"]) + self.assertEqual(situ_mx_quant.call_args.kwargs["beta"], 4.0) + self.assertEqual(situ_mx_quant.call_args.kwargs["linear_beta"], 25.0) + self.assertIs(mock_gmm2.call_args.kwargs["per_token_scale"], situ_out_scale) + def test_antiquant_weights_use_native_situ_between_grouped_matmuls(self): gate_up_out = torch.tensor([[1.0, -1.0, 0.5, 2.0]]) gate, up = gate_up_out.chunk(2, dim=-1) diff --git a/vllm_ascend/ops/fused_moe/moe_mlp.py b/vllm_ascend/ops/fused_moe/moe_mlp.py index 2fefd1f301b6..51160ed12dd6 100644 --- a/vllm_ascend/ops/fused_moe/moe_mlp.py +++ b/vllm_ascend/ops/fused_moe/moe_mlp.py @@ -1,972 +1,866 @@ -# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved. -# Copyright 2023 The vLLM team. -# -# 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. -# This file is a part of the vllm-ascend project. - - -import torch -import torch_npu -from torch.nn.functional import pad -from vllm.model_executor.layers.activation import SituAndMul -from vllm.model_executor.layers.fused_moe.activation import MoEActivation -from vllm.triton_utils import HAS_TRITON - -from vllm_ascend.ascend_forward_context import _EXTRA_CTX, MoECommType -from vllm_ascend.device.device_op import DeviceOperator -from vllm_ascend.ops.activation import AscendSwigluOAIAndMul, AscendSwigluStepAndMul -from vllm_ascend.ops.fused_moe.dataclass.moe_mlp import MoEMlpComputeInput -from vllm_ascend.quantization.quant_type import QuantType -from vllm_ascend.utils import ( - AscendDeviceType, - dispose_tensor, - enable_custom_op, - get_ascend_device_type, -) - -ASCEND_DEVICE_TYPE = get_ascend_device_type() -# CANN uses 36 to select FP8 E4M3FN output for situ_mx_quant. -SITU_MX_DST_TYPE_E4M3FN = 36 - - -def _custom_gmm_swiglu_enabled(fusion, dynamic_eplb, activation=None): - activation_name = getattr(activation, "value", activation) - return fusion and dynamic_eplb and activation_name not in ("situ", "swigluoai_uninterleave") and enable_custom_op() - - -def _gmm_swiglu_quant_fusion_enabled(use_mxfp_quant, fusion, dynamic_eplb, activation=None): - activation_name = getattr(activation, "value", activation) - return (use_mxfp_quant or (fusion and not dynamic_eplb)) and activation_name not in ( - "situ", - "swigluoai_uninterleave", - ) - - -def cumsum_group_list( - group_list: torch.Tensor, src_list_type: int, dst_list_type: int, active_num: int = 0, expert_num: int = 0 -) -> torch.Tensor: - if src_list_type not in [0, 1, 2]: - raise ValueError(f"group_list_type should be in [0, 1, 2], but received {src_list_type}") - - if src_list_type == dst_list_type: - return group_list - if src_list_type == 1 and dst_list_type == 0: - return group_list.cumsum(dim=0) - if src_list_type == 0 and dst_list_type == 1: - group_diff = torch.diff(group_list) - new_group = torch.cat([group_list[0].unsqueeze(0), group_diff], dim=0) - return new_group - if src_list_type == 2 and dst_list_type == 0: - experts = pad(group_list[:, 0], (1, 0)) - tokens = pad(group_list[:, 1].cumsum(dim=0), (1, 0)) - cumsum_group_list = torch.full( - size=(expert_num,), fill_value=active_num, dtype=group_list.dtype, device=group_list.device - ) - - for i, (start, end) in enumerate(zip(experts[:-1], experts[1:])): - if end > start: - cumsum_group_list[start:end] = tokens[i] - - return cumsum_group_list - raise NotImplementedError( - f"Conversion from src_list_type={src_list_type} to dst_list_type={dst_list_type} is not implemented yet. " - "This feature is under development." - ) - - -def _require_single_tensor_for_swiglu_quant( - tensor_or_list: list[torch.Tensor] | torch.Tensor, *, name: str -) -> torch.Tensor: - if isinstance(tensor_or_list, list): - if len(tensor_or_list) != 1: - raise ValueError(f"{name} must be a tensor or a single-element list, but got {len(tensor_or_list)}.") - return tensor_or_list[0] - return tensor_or_list - - -def _prepare_dequant_swiglu_weight_scale( - w1_scale: list[torch.Tensor] | torch.Tensor, - is_swigluoai_uninterleave: bool, -) -> torch.Tensor: - if not is_swigluoai_uninterleave: - weight_scale = w1_scale[0] - elif isinstance(w1_scale, list): - if len(w1_scale) == 1: - weight_scale = w1_scale[0] - else: - weight_scale = torch.stack([scale.reshape(-1) for scale in w1_scale], dim=0) - else: - weight_scale = w1_scale - if weight_scale.dtype != torch.float32: - weight_scale = weight_scale.to(torch.float32) - if is_swigluoai_uninterleave and weight_scale.dim() == 1: - weight_scale = weight_scale.reshape(1, -1) - return weight_scale - - -def _prepare_swigluoai_grouped_matmul_scales( - weight_scale: list[torch.Tensor] | torch.Tensor, output_dtype: torch.dtype -) -> list[torch.Tensor]: - scales = weight_scale if isinstance(weight_scale, list) else [weight_scale] - return [scale.to(output_dtype) if scale.dtype != output_dtype else scale for scale in scales] - - -def _as_grouped_matmul_weights( - tensor_or_list: list[torch.Tensor] | torch.Tensor, -) -> list[torch.Tensor]: - return tensor_or_list if isinstance(tensor_or_list, list) else [tensor_or_list] - - -def _quantized_situ_apply_mlp( - *, - hidden_states: torch.Tensor, - w1: list[torch.Tensor] | torch.Tensor, - w1_scale: list[torch.Tensor] | torch.Tensor, - w2: list[torch.Tensor] | torch.Tensor, - w2_scale: list[torch.Tensor] | torch.Tensor, - group_list: torch.Tensor, - group_list_type: int, - dynamic_scale: torch.Tensor | None, - w1_scale_bias: torch.Tensor | None, - w2_scale_bias: torch.Tensor | None, - activation_situ_beta: float, - activation_situ_linear_beta: float | None, - act_quant_type: torch.dtype, - weight_quant_type: torch.dtype | None, - scale_type: torch.dtype | None, - per_token_scale_type: torch.dtype | None, - use_bf16: bool, - use_mxfp_quant: bool, - is_per_channel_weight: bool, - mxfp_quant_dtype: QuantType | None = None, -) -> tuple[torch.Tensor, torch.npu.Event]: - """Run GMM1 -> SiTU quant -> GMM2 without the SwiGLU fusions.""" - input_hidden_dtype = hidden_states.dtype - if dynamic_scale is None: - unquantized_hidden_states = hidden_states - hidden_states, pertoken_scale = DeviceOperator.npu_dynamic_quant( - hidden_states=hidden_states, - dynamic_scale=None, - act_quant_type=act_quant_type, - use_mxfp_quant=False, - ) - dispose_tensor(unquantized_hidden_states) - externally_quantized_hidden_states = None - else: - pertoken_scale = ( - DeviceOperator.maybe_normalize_mxfp_scale_layout(dynamic_scale) if use_mxfp_quant else dynamic_scale - ) - externally_quantized_hidden_states = hidden_states - - w1_scale_list = _as_grouped_matmul_weights(w1_scale) - w2_scale_list = _as_grouped_matmul_weights(w2_scale) - output_dtype = w2_scale_list[0].dtype - bias1, bias2 = None, None - if w1_scale_bias is not None: - if group_list_type == 0: - group_list = torch.cat([group_list[:1], torch.diff(group_list, dim=0)]) - group_list_type = 1 - bias1 = w1_scale_bias - bias2 = w2_scale_bias - output_dtype = torch.bfloat16 - - gmm1_scale = [scale.to(w2_scale_list[0].dtype) for scale in w1_scale_list] - if is_per_channel_weight: - gmm1_scale = [scale.unsqueeze(-2) for scale in gmm1_scale] - - gate_up_out = torch_npu.npu_grouped_matmul( - x=[hidden_states], - weight=_as_grouped_matmul_weights(w1), - antiquant_scale=gmm1_scale if use_mxfp_quant else None, - scale=gmm1_scale if not use_mxfp_quant else None, - bias=bias1, - per_token_scale=[pertoken_scale], - split_item=2, - group_list_type=group_list_type, - group_type=0, - group_list=group_list, - per_token_scale_dtype=torch_npu.float8_e8m0fnu if use_mxfp_quant else None, - weight_dtype=torch_npu.float4_e2m1fn_x2 if use_mxfp_quant else None, - output_dtype=torch.bfloat16 if use_mxfp_quant else output_dtype, - )[0] - if externally_quantized_hidden_states is not None: - dispose_tensor(externally_quantized_hidden_states) - - if use_mxfp_quant: - hidden_states, situ_out_scale = torch.ops._C_ascend.situ_mx_quant( - x=gate_up_out, - beta=activation_situ_beta, - linear_beta=activation_situ_linear_beta or 0.0, - activate_left=True, - dst_type=SITU_MX_DST_TYPE_E4M3FN, - ) - else: - hidden_states, situ_out_scale = torch.ops._C_ascend.dequant_situ_quant( - x=gate_up_out, - weight_scale=None, - activation_scale=None, - bias=None, - quant_scale=None, - quant_offset=None, - group_index=None, - beta=activation_situ_beta, - linear_beta=activation_situ_linear_beta or 0.0, - activate_left=True, - quant_mode="dynamic", - ) - before_gmm2_evt = torch.npu.current_stream().record_event() - hidden_states = DeviceOperator.npu_grouped_matmul_gmm2( - hidden_states=hidden_states, - weight=w2, - weight_scale=w2_scale, - per_token_scale=situ_out_scale, - group_list=group_list, - group_list_type=group_list_type, - input_dtype=input_hidden_dtype, - act_quant_type=act_quant_type, - weight_quant_type=weight_quant_type, - scale_type=scale_type, - per_token_scale_type=per_token_scale_type, - use_bf16=use_bf16, - use_mxfp_quant=use_mxfp_quant, - bias=bias2, - fallback_output_dtype=output_dtype, - mxfp_quant_dtype=mxfp_quant_dtype, - ) - return hidden_states, before_gmm2_evt - - -def _apply_clipped_swiglu( - hidden_states: torch.Tensor, - *, - swiglu_limit: float, - swiglu_alpha: float, - swiglu_beta: float, -) -> torch.Tensor: - if ASCEND_DEVICE_TYPE == AscendDeviceType.A5: - hidden_size = hidden_states.shape[-1] // 2 - gate = hidden_states[..., :hidden_size].clamp(max=swiglu_limit) - up = hidden_states[..., hidden_size:].clamp( - min=-swiglu_limit, - max=swiglu_limit, - ) - return gate * torch.sigmoid(swiglu_alpha * gate) * (up + swiglu_beta) - - return torch_npu.npu_clipped_swiglu( - hidden_states, - interleaved=False, - alpha=swiglu_alpha, - limit=swiglu_limit, - bias=swiglu_beta, - ) - - -def _swiglu_oai_dynamic_mx_quant( - hidden_states: torch.Tensor, - *, - act_quant_type: torch.dtype, - swiglu_limit: float, - swiglu_alpha: float, - swiglu_beta: float, -) -> tuple[torch.Tensor, torch.Tensor]: - if ASCEND_DEVICE_TYPE != AscendDeviceType.A5: - raise RuntimeError("The MiniMax-M3 SwiGLU-OAI MX quant path is only expected on Ascend A5.") - - hidden_states = _apply_clipped_swiglu( - hidden_states, - swiglu_limit=swiglu_limit, - swiglu_alpha=swiglu_alpha, - swiglu_beta=swiglu_beta, - ) - hidden_states, swiglu_out_scale = DeviceOperator.npu_dynamic_quant( - hidden_states, - act_quant_type=act_quant_type, - use_mxfp_quant=True, - ) - assert swiglu_out_scale is not None - return hidden_states, swiglu_out_scale - - -def quant_apply_mlp( - hidden_states: torch.Tensor, - w1: list[torch.Tensor] | torch.Tensor, - w1_scale: list[torch.Tensor] | torch.Tensor, - w2: list[torch.Tensor] | torch.Tensor, - w2_scale: list[torch.Tensor] | torch.Tensor, - group_list: torch.Tensor, - group_list_type: int = 1, - dynamic_scale: torch.Tensor = None, - w1_scale_bias: torch.Tensor = None, - w2_scale_bias: torch.Tensor = None, - w1_offset: torch.Tensor | None = None, - w2_offset: torch.Tensor | None = None, - fusion: bool = False, - dynamic_eplb: bool = False, - use_mxfp_quant: bool = False, - mxfp_quant_dtype: QuantType | None = None, - act_quant_type: torch.dtype = torch.float8_e4m3fn, - weight_quant_type: torch.dtype | None = None, - scale_type: torch.dtype | None = None, - per_token_scale_type: torch.dtype | None = None, - use_bf16: bool = True, - activation: str | MoEActivation | None = None, - activation_situ_beta: float | None = None, - activation_situ_linear_beta: float | None = None, - swiglu_limit: float = 0.0, - swiglu_alpha: float = 1.0, - swiglu_beta: float = 0.0, - use_w4a8_per_channel_gmm_swiglu: bool = False, -) -> torch.Tensor: - input_hidden_dtype = hidden_states.dtype - situ_beta = 1.0 if activation_situ_beta is None else activation_situ_beta - if activation == MoEActivation.SITU: - use_antiquant_situ = mxfp_quant_dtype == QuantType.W4A16MXFP or w1_offset is not None - if not use_antiquant_situ: - return _quantized_situ_apply_mlp( - hidden_states=hidden_states, - w1=w1, - w1_scale=w1_scale, - w2=w2, - w2_scale=w2_scale, - group_list=group_list, - group_list_type=group_list_type, - dynamic_scale=dynamic_scale, - w1_scale_bias=w1_scale_bias, - w2_scale_bias=w2_scale_bias, - activation_situ_beta=situ_beta, - activation_situ_linear_beta=activation_situ_linear_beta, - act_quant_type=act_quant_type, - weight_quant_type=weight_quant_type, - scale_type=scale_type, - per_token_scale_type=per_token_scale_type, - use_bf16=use_bf16, - is_per_channel_weight=use_w4a8_per_channel_gmm_swiglu, - use_mxfp_quant=use_mxfp_quant, - mxfp_quant_dtype=mxfp_quant_dtype, - ) - act_name = getattr(activation, "value", activation) - use_gmm_swiglu_quant_fusion = _gmm_swiglu_quant_fusion_enabled( - use_mxfp_quant, - fusion, - dynamic_eplb, - activation, - ) - # GELU can't use the fused SwiGLU+quant ops below; fall back to the - # non-fused GMM -> GELU -> (re)quant -> GMM2 path for GELU activations. - is_gelu_activation = activation in (MoEActivation.GELU, MoEActivation.GELU_TANH) - is_swigluoai_uninterleave = act_name == "swigluoai_uninterleave" - - if use_mxfp_quant: - if w1_scale_bias is not None or w2_scale_bias is not None: - raise NotImplementedError("MXFP path does not support scale_bias yet.") - if w1_offset is not None or w2_offset is not None: - raise NotImplementedError("MXFP path does not support antiquant offset yet.") - - if w1_offset is not None: - unquantized_hidden_states = hidden_states - quantized_hidden_states = None - elif mxfp_quant_dtype == QuantType.W4A16MXFP: - quantized_hidden_states = None - pertoken_scale = None - elif dynamic_scale is None: - unquantized_hidden_states = hidden_states - hidden_states, pertoken_scale = DeviceOperator.npu_dynamic_quant( - hidden_states=hidden_states, - dynamic_scale=None, - act_quant_type=act_quant_type, - use_mxfp_quant=use_mxfp_quant, - ) - dispose_tensor(unquantized_hidden_states) - quantized_hidden_states = None - else: - unquantized_hidden_states = None - pertoken_scale = ( - DeviceOperator.maybe_normalize_mxfp_scale_layout(dynamic_scale) if use_mxfp_quant else dynamic_scale - ) - quantized_hidden_states = hidden_states - - bias1, bias2 = None, None - _output_dtype = w2_scale[0].dtype if isinstance(w2_scale, list) else w2_scale.dtype - - is_mc2 = _EXTRA_CTX.moe_comm_type == MoECommType.MC2 - if w1_scale_bias is None and w1_offset is None and is_mc2 and not is_gelu_activation: - if _custom_gmm_swiglu_enabled(fusion, dynamic_eplb, activation) and not use_mxfp_quant: - # gmm1: gate_up_proj & act_fn: swiglu - hidden_states, swiglu_out_scale, _ = torch.ops._C_ascend.grouped_matmul_swiglu_quant_weight_nz_tensor_list( - x=hidden_states, - weight=w1, - weight_scale=w1_scale, - x_scale=pertoken_scale, - group_list=cumsum_group_list(group_list, group_list_type, 0), - swiglu_limit=swiglu_limit, - ) - elif use_gmm_swiglu_quant_fusion and activation != MoEActivation.SWIGLUSTEP: - # gmm1: gate_up_proj & act_fn: swiglu - hidden_states, swiglu_out_scale, _ = DeviceOperator.npu_grouped_matmul_swiglu_quant( - x=hidden_states, - weight=_require_single_tensor_for_swiglu_quant(w1, name="w1"), - group_list=cumsum_group_list(group_list, group_list_type, 0), - weight_scale=_require_single_tensor_for_swiglu_quant(w1_scale, name="w1_scale"), - x_scale=pertoken_scale, - bias=None, - use_mxfp_quant=use_mxfp_quant, - act_quant_type=act_quant_type, - weight_quant_type=weight_quant_type, - swiglu_limit=swiglu_limit, - mxfp_quant_dtype=mxfp_quant_dtype, - ) - if quantized_hidden_states is not None: - dispose_tensor(quantized_hidden_states) - elif activation == MoEActivation.SWIGLUSTEP: - # Step3.5/3.7 needs to clamp in swiglu: out = silu(gate).clamp(max=limit) * up.clamp(-limit, limit) - gmm1_kwargs = { - "x": [hidden_states], - "weight": w1 if isinstance(w1, list) else [w1], - "scale": [w1_scale[0].to(w2_scale[0].dtype)] if isinstance(w1_scale, list) else [w1_scale], - "bias": None, - "per_token_scale": [pertoken_scale], - "split_item": 2, - "group_type": 0, - "group_list": group_list, - "output_dtype": torch.bfloat16, - } - if use_mxfp_quant: - gmm1_kwargs.update( - { - "scale_dtype": torch_npu.float8_e8m0fnu, - "per_token_scale_dtype": torch_npu.float8_e8m0fnu, - } - ) - hidden_states = torch_npu.npu_grouped_matmul(**gmm1_kwargs)[0] - hidden_states = AscendSwigluStepAndMul.swiglustep_forward(hidden_states, limit=7.0) - hidden_states, swiglu_out_scale = DeviceOperator.npu_dynamic_quant( - hidden_states, act_quant_type=act_quant_type, use_mxfp_quant=use_mxfp_quant - ) - else: - if use_mxfp_quant and is_swigluoai_uninterleave: - hidden_states = torch_npu.npu_grouped_matmul( - x=[hidden_states], - weight=w1 if isinstance(w1, list) else [w1], - scale=w1_scale if isinstance(w1_scale, list) else [w1_scale], - bias=None, - per_token_scale=[pertoken_scale], - split_item=2, - group_list_type=group_list_type, - group_type=0, - group_list=group_list, - output_dtype=torch.bfloat16, - scale_dtype=torch_npu.float8_e8m0fnu, - per_token_scale_dtype=torch_npu.float8_e8m0fnu, - )[0] - if quantized_hidden_states is not None: - dispose_tensor(quantized_hidden_states) - hidden_states, swiglu_out_scale = _swiglu_oai_dynamic_mx_quant( - hidden_states, - act_quant_type=act_quant_type, - swiglu_limit=swiglu_limit, - swiglu_alpha=swiglu_alpha, - swiglu_beta=swiglu_beta, - ) - else: - # gmm1: gate_up_proj - hidden_states = torch_npu.npu_grouped_matmul( - x=[hidden_states], - weight=w1, - split_item=3, - group_list_type=group_list_type, - group_type=0, - group_list=group_list, - output_dtype=torch.int32, - )[0] - if quantized_hidden_states is not None: - dispose_tensor(quantized_hidden_states) - # act_fn: swiglu - dequant_swiglu_kwargs = { - "x": hidden_states, - "weight_scale": _prepare_dequant_swiglu_weight_scale(w1_scale, is_swigluoai_uninterleave), - "activation_scale": pertoken_scale, - "bias": None, - "quant_scale": None, - "quant_offset": None, - "group_index": cumsum_group_list(group_list, group_list_type, 1), - "activate_left": True, - "quant_mode": 1, - } - if is_swigluoai_uninterleave: - dequant_swiglu_kwargs.update( - { - "swiglu_mode": 1, - "clamp_limit": swiglu_limit, - "glu_alpha": swiglu_alpha, - "glu_bias": swiglu_beta, - } - ) - hidden_states, swiglu_out_scale = torch.ops._C_ascend.npu_dequant_swiglu_quant(**dequant_swiglu_kwargs) - before_gmm2_evt = torch.npu.current_stream().record_event() - # gmm2: down_proj - hidden_states = DeviceOperator.npu_grouped_matmul_gmm2( - hidden_states=hidden_states, - weight=w2, - weight_scale=w2_scale, - per_token_scale=swiglu_out_scale, - group_list=group_list, - group_list_type=group_list_type, - input_dtype=input_hidden_dtype, - act_quant_type=act_quant_type, - weight_quant_type=weight_quant_type, - scale_type=scale_type, - per_token_scale_type=per_token_scale_type, - use_bf16=use_bf16, - use_mxfp_quant=use_mxfp_quant, - bias=None, - fallback_output_dtype=w2_scale[0].dtype if isinstance(w2_scale, list) else w2_scale.dtype, - mxfp_quant_dtype=mxfp_quant_dtype, - ) - elif w1_offset is not None: - # gmm1: gate_up_proj - hidden_states = torch_npu.npu_grouped_matmul( - x=[unquantized_hidden_states], - weight=[w1], - antiquant_scale=[w1_scale], - antiquant_offset=[w1_offset], - split_item=2, - group_list_type=group_list_type, - group_type=0, - group_list=group_list, - output_dtype=_output_dtype, - )[0] - dispose_tensor(unquantized_hidden_states) - # act_fn: swiglu - if activation == MoEActivation.SWIGLUSTEP: - hidden_states = AscendSwigluStepAndMul.swiglustep_forward(hidden_states, limit=swiglu_limit or 7.0) - elif is_gelu_activation: - gate, up = hidden_states.chunk(2, dim=-1) - approximate = "tanh" if activation == MoEActivation.GELU_TANH else "none" - hidden_states = torch.nn.functional.gelu(gate, approximate=approximate) * up - elif activation == MoEActivation.SITU: - hidden_states = SituAndMul( - beta=situ_beta, - linear_beta=activation_situ_linear_beta, - )(hidden_states) - elif is_swigluoai_uninterleave: - hidden_states = _apply_clipped_swiglu( - hidden_states, - swiglu_limit=swiglu_limit, - swiglu_alpha=swiglu_alpha, - swiglu_beta=swiglu_beta, - ) - else: - hidden_states = torch_npu.npu_swiglu(hidden_states) - before_gmm2_evt = torch.npu.current_stream().record_event() - # gmm2: down_proj - hidden_states = torch_npu.npu_grouped_matmul( - x=[hidden_states], - weight=[w2], - antiquant_scale=[w2_scale], - antiquant_offset=[w2_offset], - split_item=2, - group_list_type=group_list_type, - group_type=0, - group_list=group_list, - output_dtype=_output_dtype, - )[0] - else: - if w1_scale_bias is not None: - if group_list_type == 0: - group_list = torch.cat([group_list[:1], torch.diff(group_list, dim=0)]) - group_list_type = 1 - bias1 = w1_scale_bias - bias2 = w2_scale_bias - # TODO w4a8 scene: dynamic acquisition of dtype in the future - _output_dtype = torch.bfloat16 - - if ( - use_w4a8_per_channel_gmm_swiglu - and enable_custom_op() - and activation != MoEActivation.SWIGLUSTEP - and not is_gelu_activation - and not is_swigluoai_uninterleave - ): - hidden_states, swiglu_out_scale = torch.ops._C_ascend.grouped_matmul_swiglu_quant_v2( - x=hidden_states, - weight=w1, - weight_scale=w1_scale if isinstance(w1_scale, list) else [w1_scale], - x_scale=pertoken_scale, - group_list=group_list, - weight_assist_matrix=bias1, - dequant_mode=0, - group_list_type=group_list_type, - swiglu_limit=swiglu_limit, - ) - elif ( - _custom_gmm_swiglu_enabled(fusion, dynamic_eplb, activation) - and not use_mxfp_quant - and not is_gelu_activation - ): - # gmm1: gate_up_proj & act_fn: swiglu - hidden_states, swiglu_out_scale, _ = torch.ops._C_ascend.grouped_matmul_swiglu_quant_weight_nz_tensor_list( - x=hidden_states, - weight=w1, - weight_scale=w1_scale, - x_scale=pertoken_scale, - group_list=cumsum_group_list(group_list, group_list_type, 0), - bias=bias1, - swiglu_limit=swiglu_limit, - ) - elif use_gmm_swiglu_quant_fusion and activation != MoEActivation.SWIGLUSTEP and not is_gelu_activation: - hidden_states, swiglu_out_scale, _ = DeviceOperator.npu_grouped_matmul_swiglu_quant( - x=hidden_states, - weight=_require_single_tensor_for_swiglu_quant(w1, name="w1"), - group_list=cumsum_group_list(group_list, group_list_type, 0), - weight_scale=_require_single_tensor_for_swiglu_quant(w1_scale, name="w1_scale"), - x_scale=pertoken_scale, - bias=bias1, - use_mxfp_quant=use_mxfp_quant, - act_quant_type=act_quant_type, - weight_quant_type=weight_quant_type, - swiglu_limit=swiglu_limit, - mxfp_quant_dtype=mxfp_quant_dtype, - ) - if quantized_hidden_states is not None: - dispose_tensor(quantized_hidden_states) - else: - # gmm1: gate_up_proj - scale = [w1_scale[0].to(w2_scale[0].dtype)] if isinstance(w1_scale, list) else [w1_scale] - if is_swigluoai_uninterleave: - scale = _prepare_swigluoai_grouped_matmul_scales(w1_scale, _output_dtype) - gmm1_kwargs = { - "x": [hidden_states], - "weight": w1 if isinstance(w1, list) else [w1], - "scale": scale, - "bias": bias1, - "per_token_scale": [pertoken_scale], - "split_item": 2, - "group_type": 0, - "group_list": group_list, - "group_list_type": group_list_type, - "output_dtype": _output_dtype, - } - if use_mxfp_quant: - gmm1_kwargs.update( - { - "scale_dtype": torch_npu.float8_e8m0fnu, - "per_token_scale_dtype": torch_npu.float8_e8m0fnu, - "output_dtype": torch.bfloat16, - } - ) - hidden_states = torch_npu.npu_grouped_matmul(**gmm1_kwargs)[0] - if quantized_hidden_states is not None: - dispose_tensor(quantized_hidden_states) - # act_fn: swiglu - if activation == MoEActivation.SWIGLUSTEP: - hidden_states = AscendSwigluStepAndMul.swiglustep_forward(hidden_states, limit=swiglu_limit or 7.0) - hidden_states, swiglu_out_scale = DeviceOperator.npu_dynamic_quant( - hidden_states, act_quant_type=act_quant_type, use_mxfp_quant=use_mxfp_quant - ) - elif is_gelu_activation: - gate, up = hidden_states.chunk(2, dim=-1) - approximate = "tanh" if activation == MoEActivation.GELU_TANH else "none" - hidden_states = torch.nn.functional.gelu(gate, approximate=approximate) * up - hidden_states, swiglu_out_scale = torch_npu.npu_dynamic_quant(hidden_states) - elif activation == MoEActivation.SITU: - hidden_states = SituAndMul( - beta=situ_beta, - linear_beta=activation_situ_linear_beta, - )(hidden_states) - swiglu_out_scale = None - elif is_swigluoai_uninterleave: - if use_mxfp_quant: - hidden_states, swiglu_out_scale = _swiglu_oai_dynamic_mx_quant( - hidden_states, - act_quant_type=act_quant_type, - swiglu_limit=swiglu_limit, - swiglu_alpha=swiglu_alpha, - swiglu_beta=swiglu_beta, - ) - else: - hidden_states = _apply_clipped_swiglu( - hidden_states, - swiglu_limit=swiglu_limit, - swiglu_alpha=swiglu_alpha, - swiglu_beta=swiglu_beta, - ) - hidden_states, swiglu_out_scale = DeviceOperator.npu_dynamic_quant( - hidden_states, act_quant_type=act_quant_type, use_mxfp_quant=False - ) - elif HAS_TRITON: - from vllm_ascend.ops.triton.activation.swiglu_quant import swiglu_quant - - hidden_states, swiglu_out_scale = swiglu_quant( - hidden_states, group_list=group_list, group_list_type=group_list_type - ) - else: - hidden_states = torch_npu.npu_swiglu(hidden_states) - hidden_states, swiglu_out_scale = torch_npu.npu_dynamic_quant(hidden_states) - before_gmm2_evt = torch.npu.current_stream().record_event() - # gmm2: down_proj - hidden_states = DeviceOperator.npu_grouped_matmul_gmm2( - hidden_states=hidden_states, - weight=w2, - weight_scale=w2_scale, - per_token_scale=swiglu_out_scale, - group_list=group_list, - group_list_type=group_list_type, - input_dtype=input_hidden_dtype, - act_quant_type=act_quant_type, - weight_quant_type=weight_quant_type, - scale_type=scale_type, - per_token_scale_type=per_token_scale_type, - use_bf16=use_bf16, - use_mxfp_quant=use_mxfp_quant, - bias=bias2, - fallback_output_dtype=_output_dtype, - mxfp_quant_dtype=mxfp_quant_dtype, - ) - return hidden_states, before_gmm2_evt - - -def unquant_apply_mlp( - hidden_states: torch.Tensor, - w1: torch.Tensor, - w2: torch.Tensor, - group_list: torch.Tensor, - w1_bias: torch.Tensor = None, - w2_bias: torch.Tensor = None, - activation: str | MoEActivation | None = None, - activation_situ_beta: float | None = None, - activation_situ_linear_beta: float | None = None, - group_list_type: int = 1, - topk_scales: torch.Tensor | None = None, - need_trans: bool = True, - swiglu_limit: float = 0.0, - swiglu_alpha: float = 1.0, - swiglu_beta: float = 0.0, - lora_context=None, - expanded_row_idx: torch.Tensor | None = None, - topk_ids: torch.Tensor | None = None, -) -> torch.Tensor: - if need_trans: - w1 = w1.transpose(1, 2) - w2 = w2.transpose(1, 2) - - gate_up_out = torch_npu.npu_grouped_matmul( - x=[hidden_states], - weight=[w1], - bias=[w1_bias.to(dtype=torch.float32)] if w1_bias is not None else None, - split_item=2, - group_list_type=group_list_type, - group_type=0, - group_list=group_list, - )[0] - - # MoE LoRA: only attempt injection when an adapter wraps this layer and - # the comm method provided routing metadata in lora_context. - # Two paths are supported: - # - AllGather: expanded_row_idx + topk_ids from npu_moe_init_routing - # - AlltoAll: lora_context.exchanged_lora_indices + group_list after all_to_all - lora_routing = None - if lora_context is not None: # LoRA applied - from vllm_ascend.lora.fused_moe import ( - _recover_moe_lora_routing_all2all, - _recover_moe_lora_routing_allgather, - moe_lora_apply_w2, - moe_lora_apply_w13, - ) - - if expanded_row_idx is not None and topk_ids is not None: - # AllGather path: use npu_moe_init_routing's expanded_row_idx. - lora_routing = _recover_moe_lora_routing_allgather(lora_context, expanded_row_idx, topk_ids) - elif getattr(lora_context, "exchanged_lora_indices", None) is not None: - # AlltoAll path: tokens already sorted by expert after exchange. - # Build per-row (expert_id, lora_id) directly from group_list. - lora_routing = _recover_moe_lora_routing_all2all( - lora_context, - group_list=group_list, - ) - else: - raise AssertionError( - "MoE LoRA requires either expanded_row_idx+topk_ids " - "(AllGather) or lora_context.exchanged_lora_indices " - "(AlltoAll). Neither was provided." - ) - - moe_lora_apply_w13( - lora_context, - gate_up_out=gate_up_out, - hidden_states=hidden_states, - lora_routing=lora_routing, - ) - - act_name = getattr(activation, "value", activation) - if activation == MoEActivation.SITU: - gate_up_out = SituAndMul( - beta=activation_situ_beta, - linear_beta=activation_situ_linear_beta, - )(gate_up_out) - elif activation == MoEActivation.SWIGLUOAI: - num_experts, _, hidden_size = w1.shape - gate_up_out = AscendSwigluOAIAndMul.swiglu_oai_forward(gate_up_out.view(-1, hidden_size)) - elif act_name == "swigluoai_uninterleave": - gate_up_out = _apply_clipped_swiglu( - gate_up_out, - swiglu_limit=swiglu_limit, - swiglu_alpha=swiglu_alpha, - swiglu_beta=swiglu_beta, - ) - elif activation == MoEActivation.SWIGLUSTEP: - gate_up_out = AscendSwigluStepAndMul.swiglustep_forward(gate_up_out, limit=swiglu_limit or 7.0) - elif activation == MoEActivation.GELU: - gate, up = gate_up_out.chunk(2, dim=-1) - gate_up_out = torch.nn.functional.gelu(gate) * up - elif activation == MoEActivation.GELU_TANH: - gate, up = gate_up_out.chunk(2, dim=-1) - gate_up_out = torch.nn.functional.gelu(gate, approximate="tanh") * up - else: - gate_up_out = torch_npu.npu_swiglu(gate_up_out) - - if topk_scales is not None: - gate_up_out *= topk_scales - - hidden_states = torch_npu.npu_grouped_matmul( - x=[gate_up_out], - weight=[w2], - bias=[w2_bias.to(dtype=torch.float32)] if w2_bias is not None else None, - split_item=2, - group_list_type=group_list_type, - group_type=0, - group_list=group_list, - )[0] - - # LoRA w2 delta: applied to the down-proj output, with the activation output - # as the lora_a input. Reuses the per-row routing computed for w13. - if lora_routing is not None: - moe_lora_apply_w2( - lora_context, - down_out=hidden_states, - silu_out=gate_up_out, - lora_routing=lora_routing, - ) - return hidden_states, None - - -def unified_apply_mlp(*, mlp_compute_input: MoEMlpComputeInput) -> torch.Tensor: - """ - Unified MoE MLP entry. - Quant path is dispatched by DeviceOperator with explicit typed kernel flags. - """ - hidden_states = mlp_compute_input.hidden_states - group_list = mlp_compute_input.group_list - group_list_type = mlp_compute_input.group_list_type - dynamic_scale = mlp_compute_input.dynamic_scale - topk_scales = mlp_compute_input.topk_scales - w1 = mlp_compute_input.weights.w1 - w2 = mlp_compute_input.weights.w2 - w1_bias = mlp_compute_input.weights.w1_bias - w2_bias = mlp_compute_input.weights.w2_bias - w1_scale = mlp_compute_input.weights.w1_scale - w2_scale = mlp_compute_input.weights.w2_scale - w1_scale_bias = mlp_compute_input.weights.w1_scale_bias - w2_scale_bias = mlp_compute_input.weights.w2_scale_bias - w1_offset = mlp_compute_input.weights.w1_offset - w2_offset = mlp_compute_input.weights.w2_offset - activation = mlp_compute_input.activation - need_trans = mlp_compute_input.need_trans - dynamic_eplb = mlp_compute_input.dynamic_eplb - fusion = mlp_compute_input.fusion - swiglu_limit = mlp_compute_input.swiglu_limit - swiglu_alpha = mlp_compute_input.swiglu_alpha - swiglu_beta = mlp_compute_input.swiglu_beta - activation_situ_beta = mlp_compute_input.activation_situ_beta - activation_situ_linear_beta = mlp_compute_input.activation_situ_linear_beta - - if not mlp_compute_input.quant.is_quant: - return unquant_apply_mlp( - hidden_states=hidden_states, - w1=w1, - w2=w2, - w1_bias=w1_bias, - w2_bias=w2_bias, - activation=activation, - activation_situ_beta=activation_situ_beta, - activation_situ_linear_beta=activation_situ_linear_beta, - group_list=group_list, - group_list_type=group_list_type, - topk_scales=topk_scales, - need_trans=need_trans, - swiglu_limit=swiglu_limit, - swiglu_alpha=swiglu_alpha, - swiglu_beta=swiglu_beta, - lora_context=mlp_compute_input.lora_context, - expanded_row_idx=mlp_compute_input.expanded_row_idx, - topk_ids=mlp_compute_input.topk_ids, - ) - - from vllm_ascend.lora.fused_moe import has_lora - - if has_lora(mlp_compute_input.lora_context): - from vllm_ascend.lora.quant_moe import quant_apply_mlp_with_moe_lora - - return quant_apply_mlp_with_moe_lora(mlp_compute_input=mlp_compute_input) - - assert w1_scale is not None and w2_scale is not None - act_quant_type = torch.int8 if mlp_compute_input.quant.is_int_quant else torch.float8_e4m3fn - weight_quant_type = torch.float8_e4m3fn - scale_type = None - per_token_scale_type = None - use_bf16 = hidden_states.dtype == torch.bfloat16 - use_mxfp_quant = mlp_compute_input.quant.is_mxfp - mxfp_quant_dtype = mlp_compute_input.quant.quant_type - - if use_mxfp_quant: - mxfp = mlp_compute_input.quant.mxfp - assert mxfp is not None, "mlp_compute_input.quant.mxfp is required when quant_type is W8A8MXFP." - act_quant_type = mxfp.act_quant_type or act_quant_type - if mxfp_quant_dtype == QuantType.W4A16MXFP: - act_quant_type = mxfp.act_quant_type - weight_quant_type = mxfp.weight_quant_type or weight_quant_type - if mxfp_quant_dtype in [QuantType.W4A8MXFP, QuantType.W4A16MXFP]: - weight_quant_type = mxfp.weight_quant_type - scale_type = mxfp.scale_dtype - per_token_scale_type = mxfp.per_token_scale_dtype - use_bf16 = mxfp.use_bf16 - - return quant_apply_mlp( - hidden_states=hidden_states, - w1=w1, - w1_scale=w1_scale, - w2=w2, - w2_scale=w2_scale, - group_list=group_list, - dynamic_scale=dynamic_scale, - group_list_type=group_list_type, - w1_scale_bias=w1_scale_bias, - w2_scale_bias=w2_scale_bias, - w1_offset=w1_offset, - w2_offset=w2_offset, - fusion=fusion, - dynamic_eplb=dynamic_eplb, - use_mxfp_quant=use_mxfp_quant, - mxfp_quant_dtype=mxfp_quant_dtype, - act_quant_type=act_quant_type, - weight_quant_type=weight_quant_type, - scale_type=scale_type, - per_token_scale_type=per_token_scale_type, - use_bf16=use_bf16, - activation=activation, - activation_situ_beta=activation_situ_beta, - activation_situ_linear_beta=activation_situ_linear_beta, - swiglu_limit=swiglu_limit, - swiglu_alpha=swiglu_alpha, - swiglu_beta=swiglu_beta, - use_w4a8_per_channel_gmm_swiglu=mlp_compute_input.quant.use_w4a8_per_channel_gmm_swiglu, - ) +# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2023 The vLLM team. +# +# 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. +# This file is a part of the vllm-ascend project. + + +import torch +import torch_npu +from torch.nn.functional import pad +from vllm.model_executor.layers.activation import SituAndMul +from vllm.model_executor.layers.fused_moe.activation import MoEActivation +from vllm.triton_utils import HAS_TRITON + +from vllm_ascend.ascend_forward_context import _EXTRA_CTX, MoECommType +from vllm_ascend.device.device_op import DeviceOperator +from vllm_ascend.ops.activation import AscendSwigluOAIAndMul, AscendSwigluStepAndMul +from vllm_ascend.ops.fused_moe.dataclass.moe_mlp import MoEMlpComputeInput +from vllm_ascend.quantization.quant_type import QuantType +from vllm_ascend.utils import ( + AscendDeviceType, + dispose_tensor, + enable_custom_op, + get_ascend_device_type, +) + +ASCEND_DEVICE_TYPE = get_ascend_device_type() +# CANN uses 36 to select FP8 E4M3FN output for situ_mx_quant. +SITU_MX_DST_TYPE_E4M3FN = 36 + + +def _custom_gmm_swiglu_enabled(fusion, dynamic_eplb, activation=None): + activation_name = getattr(activation, "value", activation) + return fusion and dynamic_eplb and activation_name not in ("situ", "swigluoai_uninterleave") and enable_custom_op() + + +def _gmm_swiglu_quant_fusion_enabled(use_mxfp_quant, fusion, dynamic_eplb, activation=None): + activation_name = getattr(activation, "value", activation) + return (use_mxfp_quant or (fusion and not dynamic_eplb)) and activation_name not in ( + "situ", + "swigluoai_uninterleave", + ) + + +def cumsum_group_list( + group_list: torch.Tensor, src_list_type: int, dst_list_type: int, active_num: int = 0, expert_num: int = 0 +) -> torch.Tensor: + if src_list_type not in [0, 1, 2]: + raise ValueError(f"group_list_type should be in [0, 1, 2], but received {src_list_type}") + + if src_list_type == dst_list_type: + return group_list + if src_list_type == 1 and dst_list_type == 0: + return group_list.cumsum(dim=0) + if src_list_type == 0 and dst_list_type == 1: + group_diff = torch.diff(group_list) + new_group = torch.cat([group_list[0].unsqueeze(0), group_diff], dim=0) + return new_group + if src_list_type == 2 and dst_list_type == 0: + experts = pad(group_list[:, 0], (1, 0)) + tokens = pad(group_list[:, 1].cumsum(dim=0), (1, 0)) + cumsum_group_list = torch.full( + size=(expert_num,), fill_value=active_num, dtype=group_list.dtype, device=group_list.device + ) + + for i, (start, end) in enumerate(zip(experts[:-1], experts[1:])): + if end > start: + cumsum_group_list[start:end] = tokens[i] + + return cumsum_group_list + raise NotImplementedError( + f"Conversion from src_list_type={src_list_type} to dst_list_type={dst_list_type} is not implemented yet. " + "This feature is under development." + ) + + +def _require_single_tensor_for_swiglu_quant( + tensor_or_list: list[torch.Tensor] | torch.Tensor, *, name: str +) -> torch.Tensor: + if isinstance(tensor_or_list, list): + if len(tensor_or_list) != 1: + raise ValueError(f"{name} must be a tensor or a single-element list, but got {len(tensor_or_list)}.") + return tensor_or_list[0] + return tensor_or_list + + +def _prepare_dequant_swiglu_weight_scale( + w1_scale: list[torch.Tensor] | torch.Tensor, + is_swigluoai_uninterleave: bool, +) -> torch.Tensor: + if not is_swigluoai_uninterleave: + weight_scale = w1_scale[0] + elif isinstance(w1_scale, list): + if len(w1_scale) == 1: + weight_scale = w1_scale[0] + else: + weight_scale = torch.stack([scale.reshape(-1) for scale in w1_scale], dim=0) + else: + weight_scale = w1_scale + if weight_scale.dtype != torch.float32: + weight_scale = weight_scale.to(torch.float32) + if is_swigluoai_uninterleave and weight_scale.dim() == 1: + weight_scale = weight_scale.reshape(1, -1) + return weight_scale + + +def _prepare_swigluoai_grouped_matmul_scales( + weight_scale: list[torch.Tensor] | torch.Tensor, output_dtype: torch.dtype +) -> list[torch.Tensor]: + scales = weight_scale if isinstance(weight_scale, list) else [weight_scale] + return [scale.to(output_dtype) if scale.dtype != output_dtype else scale for scale in scales] + + +def _apply_clipped_swiglu( + hidden_states: torch.Tensor, + *, + swiglu_limit: float, + swiglu_alpha: float, + swiglu_beta: float, +) -> torch.Tensor: + if ASCEND_DEVICE_TYPE == AscendDeviceType.A5: + hidden_size = hidden_states.shape[-1] // 2 + gate = hidden_states[..., :hidden_size].clamp(max=swiglu_limit) + up = hidden_states[..., hidden_size:].clamp( + min=-swiglu_limit, + max=swiglu_limit, + ) + return gate * torch.sigmoid(swiglu_alpha * gate) * (up + swiglu_beta) + + return torch_npu.npu_clipped_swiglu( + hidden_states, + interleaved=False, + alpha=swiglu_alpha, + limit=swiglu_limit, + bias=swiglu_beta, + ) + + +def _swiglu_oai_dynamic_mx_quant( + hidden_states: torch.Tensor, + *, + act_quant_type: torch.dtype, + swiglu_limit: float, + swiglu_alpha: float, + swiglu_beta: float, +) -> tuple[torch.Tensor, torch.Tensor]: + if ASCEND_DEVICE_TYPE != AscendDeviceType.A5: + raise RuntimeError("The MiniMax-M3 SwiGLU-OAI MX quant path is only expected on Ascend A5.") + + hidden_states = _apply_clipped_swiglu( + hidden_states, + swiglu_limit=swiglu_limit, + swiglu_alpha=swiglu_alpha, + swiglu_beta=swiglu_beta, + ) + hidden_states, swiglu_out_scale = DeviceOperator.npu_dynamic_quant( + hidden_states, + act_quant_type=act_quant_type, + use_mxfp_quant=True, + ) + assert swiglu_out_scale is not None + return hidden_states, swiglu_out_scale + + +def quant_apply_mlp( + hidden_states: torch.Tensor, + w1: list[torch.Tensor] | torch.Tensor, + w1_scale: list[torch.Tensor] | torch.Tensor, + w2: list[torch.Tensor] | torch.Tensor, + w2_scale: list[torch.Tensor] | torch.Tensor, + group_list: torch.Tensor, + group_list_type: int = 1, + dynamic_scale: torch.Tensor = None, + w1_scale_bias: torch.Tensor = None, + w2_scale_bias: torch.Tensor = None, + w1_offset: torch.Tensor | None = None, + w2_offset: torch.Tensor | None = None, + fusion: bool = False, + dynamic_eplb: bool = False, + use_mxfp_quant: bool = False, + mxfp_quant_dtype: QuantType | None = None, + act_quant_type: torch.dtype = torch.float8_e4m3fn, + weight_quant_type: torch.dtype | None = None, + scale_type: torch.dtype | None = None, + per_token_scale_type: torch.dtype | None = None, + use_bf16: bool = True, + activation: str | MoEActivation | None = None, + activation_situ_beta: float | None = None, + activation_situ_linear_beta: float | None = None, + swiglu_limit: float = 0.0, + swiglu_alpha: float = 1.0, + swiglu_beta: float = 0.0, + use_w4a8_per_channel_gmm_swiglu: bool = False, +) -> torch.Tensor: + input_hidden_dtype = hidden_states.dtype + situ_beta = 1.0 if activation_situ_beta is None else activation_situ_beta + act_name = getattr(activation, "value", activation) + is_situ_activation = activation == MoEActivation.SITU + quantize_situ_output = is_situ_activation and mxfp_quant_dtype != QuantType.W4A16MXFP + use_gmm_swiglu_quant_fusion = _gmm_swiglu_quant_fusion_enabled( + use_mxfp_quant, + fusion, + dynamic_eplb, + activation, + ) + # GELU can't use the fused SwiGLU+quant ops below; fall back to the + # non-fused GMM -> GELU -> (re)quant -> GMM2 path for GELU activations. + is_gelu_activation = activation in (MoEActivation.GELU, MoEActivation.GELU_TANH) + is_swigluoai_uninterleave = act_name == "swigluoai_uninterleave" + + if use_mxfp_quant: + if w1_scale_bias is not None or w2_scale_bias is not None: + raise NotImplementedError("MXFP path does not support scale_bias yet.") + if w1_offset is not None or w2_offset is not None: + raise NotImplementedError("MXFP path does not support antiquant offset yet.") + + if w1_offset is not None: + unquantized_hidden_states = hidden_states + quantized_hidden_states = None + elif mxfp_quant_dtype == QuantType.W4A16MXFP: + quantized_hidden_states = None + pertoken_scale = None + elif dynamic_scale is None: + unquantized_hidden_states = hidden_states + hidden_states, pertoken_scale = DeviceOperator.npu_dynamic_quant( + hidden_states=hidden_states, + dynamic_scale=None, + act_quant_type=act_quant_type, + use_mxfp_quant=use_mxfp_quant, + ) + dispose_tensor(unquantized_hidden_states) + quantized_hidden_states = None + else: + unquantized_hidden_states = None + pertoken_scale = ( + DeviceOperator.maybe_normalize_mxfp_scale_layout(dynamic_scale) if use_mxfp_quant else dynamic_scale + ) + quantized_hidden_states = hidden_states + + bias1, bias2 = None, None + _output_dtype = w2_scale[0].dtype if isinstance(w2_scale, list) else w2_scale.dtype + + is_mc2 = _EXTRA_CTX.moe_comm_type == MoECommType.MC2 + if w1_scale_bias is None and w1_offset is None and is_mc2 and not is_gelu_activation and not is_situ_activation: + if _custom_gmm_swiglu_enabled(fusion, dynamic_eplb, activation) and not use_mxfp_quant: + # gmm1: gate_up_proj & act_fn: swiglu + hidden_states, swiglu_out_scale, _ = torch.ops._C_ascend.grouped_matmul_swiglu_quant_weight_nz_tensor_list( + x=hidden_states, + weight=w1, + weight_scale=w1_scale, + x_scale=pertoken_scale, + group_list=cumsum_group_list(group_list, group_list_type, 0), + swiglu_limit=swiglu_limit, + ) + elif use_gmm_swiglu_quant_fusion and activation != MoEActivation.SWIGLUSTEP: + # gmm1: gate_up_proj & act_fn: swiglu + hidden_states, swiglu_out_scale, _ = DeviceOperator.npu_grouped_matmul_swiglu_quant( + x=hidden_states, + weight=_require_single_tensor_for_swiglu_quant(w1, name="w1"), + group_list=cumsum_group_list(group_list, group_list_type, 0), + weight_scale=_require_single_tensor_for_swiglu_quant(w1_scale, name="w1_scale"), + x_scale=pertoken_scale, + bias=None, + use_mxfp_quant=use_mxfp_quant, + act_quant_type=act_quant_type, + weight_quant_type=weight_quant_type, + swiglu_limit=swiglu_limit, + mxfp_quant_dtype=mxfp_quant_dtype, + ) + if quantized_hidden_states is not None: + dispose_tensor(quantized_hidden_states) + elif activation == MoEActivation.SWIGLUSTEP: + # Step3.5/3.7 needs to clamp in swiglu: out = silu(gate).clamp(max=limit) * up.clamp(-limit, limit) + gmm1_kwargs = { + "x": [hidden_states], + "weight": w1 if isinstance(w1, list) else [w1], + "scale": [w1_scale[0].to(w2_scale[0].dtype)] if isinstance(w1_scale, list) else [w1_scale], + "bias": None, + "per_token_scale": [pertoken_scale], + "split_item": 2, + "group_type": 0, + "group_list": group_list, + "output_dtype": torch.bfloat16, + } + if use_mxfp_quant: + gmm1_kwargs.update( + { + "scale_dtype": torch_npu.float8_e8m0fnu, + "per_token_scale_dtype": torch_npu.float8_e8m0fnu, + } + ) + hidden_states = torch_npu.npu_grouped_matmul(**gmm1_kwargs)[0] + hidden_states = AscendSwigluStepAndMul.swiglustep_forward(hidden_states, limit=7.0) + hidden_states, swiglu_out_scale = DeviceOperator.npu_dynamic_quant( + hidden_states, act_quant_type=act_quant_type, use_mxfp_quant=use_mxfp_quant + ) + else: + if use_mxfp_quant and is_swigluoai_uninterleave: + hidden_states = torch_npu.npu_grouped_matmul( + x=[hidden_states], + weight=w1 if isinstance(w1, list) else [w1], + scale=w1_scale if isinstance(w1_scale, list) else [w1_scale], + bias=None, + per_token_scale=[pertoken_scale], + split_item=2, + group_list_type=group_list_type, + group_type=0, + group_list=group_list, + output_dtype=torch.bfloat16, + scale_dtype=torch_npu.float8_e8m0fnu, + per_token_scale_dtype=torch_npu.float8_e8m0fnu, + )[0] + if quantized_hidden_states is not None: + dispose_tensor(quantized_hidden_states) + hidden_states, swiglu_out_scale = _swiglu_oai_dynamic_mx_quant( + hidden_states, + act_quant_type=act_quant_type, + swiglu_limit=swiglu_limit, + swiglu_alpha=swiglu_alpha, + swiglu_beta=swiglu_beta, + ) + else: + # gmm1: gate_up_proj + hidden_states = torch_npu.npu_grouped_matmul( + x=[hidden_states], + weight=w1, + split_item=3, + group_list_type=group_list_type, + group_type=0, + group_list=group_list, + output_dtype=torch.int32, + )[0] + if quantized_hidden_states is not None: + dispose_tensor(quantized_hidden_states) + # act_fn: swiglu + dequant_swiglu_kwargs = { + "x": hidden_states, + "weight_scale": _prepare_dequant_swiglu_weight_scale(w1_scale, is_swigluoai_uninterleave), + "activation_scale": pertoken_scale, + "bias": None, + "quant_scale": None, + "quant_offset": None, + "group_index": cumsum_group_list(group_list, group_list_type, 1), + "activate_left": True, + "quant_mode": 1, + } + if is_swigluoai_uninterleave: + dequant_swiglu_kwargs.update( + { + "swiglu_mode": 1, + "clamp_limit": swiglu_limit, + "glu_alpha": swiglu_alpha, + "glu_bias": swiglu_beta, + } + ) + hidden_states, swiglu_out_scale = torch.ops._C_ascend.npu_dequant_swiglu_quant(**dequant_swiglu_kwargs) + before_gmm2_evt = torch.npu.current_stream().record_event() + # gmm2: down_proj + hidden_states = DeviceOperator.npu_grouped_matmul_gmm2( + hidden_states=hidden_states, + weight=w2, + weight_scale=w2_scale, + per_token_scale=swiglu_out_scale, + group_list=group_list, + group_list_type=group_list_type, + input_dtype=input_hidden_dtype, + act_quant_type=act_quant_type, + weight_quant_type=weight_quant_type, + scale_type=scale_type, + per_token_scale_type=per_token_scale_type, + use_bf16=use_bf16, + use_mxfp_quant=use_mxfp_quant, + bias=None, + fallback_output_dtype=w2_scale[0].dtype if isinstance(w2_scale, list) else w2_scale.dtype, + mxfp_quant_dtype=mxfp_quant_dtype, + ) + elif w1_offset is not None: + # gmm1: gate_up_proj + hidden_states = torch_npu.npu_grouped_matmul( + x=[unquantized_hidden_states], + weight=[w1], + antiquant_scale=[w1_scale], + antiquant_offset=[w1_offset], + split_item=2, + group_list_type=group_list_type, + group_type=0, + group_list=group_list, + output_dtype=_output_dtype, + )[0] + dispose_tensor(unquantized_hidden_states) + # act_fn: swiglu + if activation == MoEActivation.SWIGLUSTEP: + hidden_states = AscendSwigluStepAndMul.swiglustep_forward(hidden_states, limit=swiglu_limit or 7.0) + elif is_gelu_activation: + gate, up = hidden_states.chunk(2, dim=-1) + approximate = "tanh" if activation == MoEActivation.GELU_TANH else "none" + hidden_states = torch.nn.functional.gelu(gate, approximate=approximate) * up + elif activation == MoEActivation.SITU: + hidden_states = SituAndMul( + beta=situ_beta, + linear_beta=activation_situ_linear_beta, + )(hidden_states) + elif is_swigluoai_uninterleave: + hidden_states = _apply_clipped_swiglu( + hidden_states, + swiglu_limit=swiglu_limit, + swiglu_alpha=swiglu_alpha, + swiglu_beta=swiglu_beta, + ) + else: + hidden_states = torch_npu.npu_swiglu(hidden_states) + before_gmm2_evt = torch.npu.current_stream().record_event() + # gmm2: down_proj + hidden_states = torch_npu.npu_grouped_matmul( + x=[hidden_states], + weight=[w2], + antiquant_scale=[w2_scale], + antiquant_offset=[w2_offset], + split_item=2, + group_list_type=group_list_type, + group_type=0, + group_list=group_list, + output_dtype=_output_dtype, + )[0] + else: + if w1_scale_bias is not None: + if group_list_type == 0: + group_list = torch.cat([group_list[:1], torch.diff(group_list, dim=0)]) + group_list_type = 1 + bias1 = w1_scale_bias + bias2 = w2_scale_bias + # TODO w4a8 scene: dynamic acquisition of dtype in the future + _output_dtype = torch.bfloat16 + + if ( + use_w4a8_per_channel_gmm_swiglu + and enable_custom_op() + and activation != MoEActivation.SWIGLUSTEP + and not is_gelu_activation + and not is_situ_activation + and not is_swigluoai_uninterleave + ): + hidden_states, swiglu_out_scale = torch.ops._C_ascend.grouped_matmul_swiglu_quant_v2( + x=hidden_states, + weight=w1, + weight_scale=w1_scale if isinstance(w1_scale, list) else [w1_scale], + x_scale=pertoken_scale, + group_list=group_list, + weight_assist_matrix=bias1, + dequant_mode=0, + group_list_type=group_list_type, + swiglu_limit=swiglu_limit, + ) + elif ( + _custom_gmm_swiglu_enabled(fusion, dynamic_eplb, activation) + and not use_mxfp_quant + and not is_gelu_activation + ): + # gmm1: gate_up_proj & act_fn: swiglu + hidden_states, swiglu_out_scale, _ = torch.ops._C_ascend.grouped_matmul_swiglu_quant_weight_nz_tensor_list( + x=hidden_states, + weight=w1, + weight_scale=w1_scale, + x_scale=pertoken_scale, + group_list=cumsum_group_list(group_list, group_list_type, 0), + bias=bias1, + swiglu_limit=swiglu_limit, + ) + elif use_gmm_swiglu_quant_fusion and activation != MoEActivation.SWIGLUSTEP and not is_gelu_activation: + hidden_states, swiglu_out_scale, _ = DeviceOperator.npu_grouped_matmul_swiglu_quant( + x=hidden_states, + weight=_require_single_tensor_for_swiglu_quant(w1, name="w1"), + group_list=cumsum_group_list(group_list, group_list_type, 0), + weight_scale=_require_single_tensor_for_swiglu_quant(w1_scale, name="w1_scale"), + x_scale=pertoken_scale, + bias=bias1, + use_mxfp_quant=use_mxfp_quant, + act_quant_type=act_quant_type, + weight_quant_type=weight_quant_type, + swiglu_limit=swiglu_limit, + mxfp_quant_dtype=mxfp_quant_dtype, + ) + if quantized_hidden_states is not None: + dispose_tensor(quantized_hidden_states) + else: + # gmm1: gate_up_proj + if quantize_situ_output: + output_scale = w2_scale[0] if isinstance(w2_scale, list) else w2_scale + scale = w1_scale if isinstance(w1_scale, list) else [w1_scale] + scale = [item.to(output_scale.dtype) if item.dtype != output_scale.dtype else item for item in scale] + if use_w4a8_per_channel_gmm_swiglu: + scale = [item.unsqueeze(-2) for item in scale] + elif is_swigluoai_uninterleave: + scale = _prepare_swigluoai_grouped_matmul_scales(w1_scale, _output_dtype) + else: + scale = [w1_scale[0].to(w2_scale[0].dtype)] if isinstance(w1_scale, list) else [w1_scale] + gmm1_kwargs = { + "x": [hidden_states], + "weight": w1 if isinstance(w1, list) else [w1], + "scale": scale, + "bias": bias1, + "per_token_scale": [pertoken_scale], + "split_item": 2, + "group_type": 0, + "group_list": group_list, + "group_list_type": group_list_type, + "output_dtype": _output_dtype, + } + if use_mxfp_quant: + if quantize_situ_output: + gmm1_kwargs.update( + { + "scale": None, + "antiquant_scale": scale, + "per_token_scale_dtype": torch_npu.float8_e8m0fnu, + "weight_dtype": torch_npu.float4_e2m1fn_x2, + "output_dtype": torch.bfloat16, + } + ) + else: + gmm1_kwargs.update( + { + "scale_dtype": torch_npu.float8_e8m0fnu, + "per_token_scale_dtype": torch_npu.float8_e8m0fnu, + "output_dtype": torch.bfloat16, + } + ) + hidden_states = torch_npu.npu_grouped_matmul(**gmm1_kwargs)[0] + if quantized_hidden_states is not None: + dispose_tensor(quantized_hidden_states) + # act_fn: swiglu + if activation == MoEActivation.SWIGLUSTEP: + hidden_states = AscendSwigluStepAndMul.swiglustep_forward(hidden_states, limit=swiglu_limit or 7.0) + hidden_states, swiglu_out_scale = DeviceOperator.npu_dynamic_quant( + hidden_states, act_quant_type=act_quant_type, use_mxfp_quant=use_mxfp_quant + ) + elif is_gelu_activation: + gate, up = hidden_states.chunk(2, dim=-1) + approximate = "tanh" if activation == MoEActivation.GELU_TANH else "none" + hidden_states = torch.nn.functional.gelu(gate, approximate=approximate) * up + hidden_states, swiglu_out_scale = torch_npu.npu_dynamic_quant(hidden_states) + elif is_situ_activation: + if quantize_situ_output and use_mxfp_quant: + hidden_states, swiglu_out_scale = torch.ops._C_ascend.situ_mx_quant( + x=hidden_states, + beta=situ_beta, + linear_beta=activation_situ_linear_beta or 0.0, + activate_left=True, + dst_type=SITU_MX_DST_TYPE_E4M3FN, + ) + elif quantize_situ_output: + hidden_states, swiglu_out_scale = torch.ops._C_ascend.dequant_situ_quant( + x=hidden_states, + weight_scale=None, + activation_scale=None, + bias=None, + quant_scale=None, + quant_offset=None, + group_index=None, + beta=situ_beta, + linear_beta=activation_situ_linear_beta or 0.0, + activate_left=True, + quant_mode="dynamic", + ) + else: + hidden_states = SituAndMul( + beta=situ_beta, + linear_beta=activation_situ_linear_beta, + )(hidden_states) + swiglu_out_scale = None + elif is_swigluoai_uninterleave: + if use_mxfp_quant: + hidden_states, swiglu_out_scale = _swiglu_oai_dynamic_mx_quant( + hidden_states, + act_quant_type=act_quant_type, + swiglu_limit=swiglu_limit, + swiglu_alpha=swiglu_alpha, + swiglu_beta=swiglu_beta, + ) + else: + hidden_states = _apply_clipped_swiglu( + hidden_states, + swiglu_limit=swiglu_limit, + swiglu_alpha=swiglu_alpha, + swiglu_beta=swiglu_beta, + ) + hidden_states, swiglu_out_scale = DeviceOperator.npu_dynamic_quant( + hidden_states, act_quant_type=act_quant_type, use_mxfp_quant=False + ) + elif HAS_TRITON: + from vllm_ascend.ops.triton.activation.swiglu_quant import swiglu_quant + + hidden_states, swiglu_out_scale = swiglu_quant( + hidden_states, group_list=group_list, group_list_type=group_list_type + ) + else: + hidden_states = torch_npu.npu_swiglu(hidden_states) + hidden_states, swiglu_out_scale = torch_npu.npu_dynamic_quant(hidden_states) + before_gmm2_evt = torch.npu.current_stream().record_event() + # gmm2: down_proj + hidden_states = DeviceOperator.npu_grouped_matmul_gmm2( + hidden_states=hidden_states, + weight=w2, + weight_scale=w2_scale, + per_token_scale=swiglu_out_scale, + group_list=group_list, + group_list_type=group_list_type, + input_dtype=input_hidden_dtype, + act_quant_type=act_quant_type, + weight_quant_type=weight_quant_type, + scale_type=scale_type, + per_token_scale_type=per_token_scale_type, + use_bf16=use_bf16, + use_mxfp_quant=use_mxfp_quant, + bias=bias2, + fallback_output_dtype=_output_dtype, + mxfp_quant_dtype=mxfp_quant_dtype, + ) + return hidden_states, before_gmm2_evt + + +def unquant_apply_mlp( + hidden_states: torch.Tensor, + w1: torch.Tensor, + w2: torch.Tensor, + group_list: torch.Tensor, + w1_bias: torch.Tensor = None, + w2_bias: torch.Tensor = None, + activation: str | MoEActivation | None = None, + activation_situ_beta: float | None = None, + activation_situ_linear_beta: float | None = None, + group_list_type: int = 1, + topk_scales: torch.Tensor | None = None, + need_trans: bool = True, + swiglu_limit: float = 0.0, + swiglu_alpha: float = 1.0, + swiglu_beta: float = 0.0, + lora_context=None, + expanded_row_idx: torch.Tensor | None = None, + topk_ids: torch.Tensor | None = None, +) -> torch.Tensor: + if need_trans: + w1 = w1.transpose(1, 2) + w2 = w2.transpose(1, 2) + + gate_up_out = torch_npu.npu_grouped_matmul( + x=[hidden_states], + weight=[w1], + bias=[w1_bias.to(dtype=torch.float32)] if w1_bias is not None else None, + split_item=2, + group_list_type=group_list_type, + group_type=0, + group_list=group_list, + )[0] + + # MoE LoRA: only attempt injection when an adapter wraps this layer and + # the comm method provided routing metadata in lora_context. + # Two paths are supported: + # - AllGather: expanded_row_idx + topk_ids from npu_moe_init_routing + # - AlltoAll: lora_context.exchanged_lora_indices + group_list after all_to_all + lora_routing = None + if lora_context is not None: # LoRA applied + from vllm_ascend.lora.fused_moe import ( + _recover_moe_lora_routing_all2all, + _recover_moe_lora_routing_allgather, + moe_lora_apply_w2, + moe_lora_apply_w13, + ) + + if expanded_row_idx is not None and topk_ids is not None: + # AllGather path: use npu_moe_init_routing's expanded_row_idx. + lora_routing = _recover_moe_lora_routing_allgather(lora_context, expanded_row_idx, topk_ids) + elif getattr(lora_context, "exchanged_lora_indices", None) is not None: + # AlltoAll path: tokens already sorted by expert after exchange. + # Build per-row (expert_id, lora_id) directly from group_list. + lora_routing = _recover_moe_lora_routing_all2all( + lora_context, + group_list=group_list, + ) + else: + raise AssertionError( + "MoE LoRA requires either expanded_row_idx+topk_ids " + "(AllGather) or lora_context.exchanged_lora_indices " + "(AlltoAll). Neither was provided." + ) + + moe_lora_apply_w13( + lora_context, + gate_up_out=gate_up_out, + hidden_states=hidden_states, + lora_routing=lora_routing, + ) + + act_name = getattr(activation, "value", activation) + if activation == MoEActivation.SITU: + gate_up_out = SituAndMul( + beta=activation_situ_beta, + linear_beta=activation_situ_linear_beta, + )(gate_up_out) + elif activation == MoEActivation.SWIGLUOAI: + num_experts, _, hidden_size = w1.shape + gate_up_out = AscendSwigluOAIAndMul.swiglu_oai_forward(gate_up_out.view(-1, hidden_size)) + elif act_name == "swigluoai_uninterleave": + gate_up_out = _apply_clipped_swiglu( + gate_up_out, + swiglu_limit=swiglu_limit, + swiglu_alpha=swiglu_alpha, + swiglu_beta=swiglu_beta, + ) + elif activation == MoEActivation.SWIGLUSTEP: + gate_up_out = AscendSwigluStepAndMul.swiglustep_forward(gate_up_out, limit=swiglu_limit or 7.0) + elif activation == MoEActivation.GELU: + gate, up = gate_up_out.chunk(2, dim=-1) + gate_up_out = torch.nn.functional.gelu(gate) * up + elif activation == MoEActivation.GELU_TANH: + gate, up = gate_up_out.chunk(2, dim=-1) + gate_up_out = torch.nn.functional.gelu(gate, approximate="tanh") * up + else: + gate_up_out = torch_npu.npu_swiglu(gate_up_out) + + if topk_scales is not None: + gate_up_out *= topk_scales + + hidden_states = torch_npu.npu_grouped_matmul( + x=[gate_up_out], + weight=[w2], + bias=[w2_bias.to(dtype=torch.float32)] if w2_bias is not None else None, + split_item=2, + group_list_type=group_list_type, + group_type=0, + group_list=group_list, + )[0] + + # LoRA w2 delta: applied to the down-proj output, with the activation output + # as the lora_a input. Reuses the per-row routing computed for w13. + if lora_routing is not None: + moe_lora_apply_w2( + lora_context, + down_out=hidden_states, + silu_out=gate_up_out, + lora_routing=lora_routing, + ) + return hidden_states, None + + +def unified_apply_mlp(*, mlp_compute_input: MoEMlpComputeInput) -> torch.Tensor: + """ + Unified MoE MLP entry. + Quant path is dispatched by DeviceOperator with explicit typed kernel flags. + """ + hidden_states = mlp_compute_input.hidden_states + group_list = mlp_compute_input.group_list + group_list_type = mlp_compute_input.group_list_type + dynamic_scale = mlp_compute_input.dynamic_scale + topk_scales = mlp_compute_input.topk_scales + w1 = mlp_compute_input.weights.w1 + w2 = mlp_compute_input.weights.w2 + w1_bias = mlp_compute_input.weights.w1_bias + w2_bias = mlp_compute_input.weights.w2_bias + w1_scale = mlp_compute_input.weights.w1_scale + w2_scale = mlp_compute_input.weights.w2_scale + w1_scale_bias = mlp_compute_input.weights.w1_scale_bias + w2_scale_bias = mlp_compute_input.weights.w2_scale_bias + w1_offset = mlp_compute_input.weights.w1_offset + w2_offset = mlp_compute_input.weights.w2_offset + activation = mlp_compute_input.activation + need_trans = mlp_compute_input.need_trans + dynamic_eplb = mlp_compute_input.dynamic_eplb + fusion = mlp_compute_input.fusion + swiglu_limit = mlp_compute_input.swiglu_limit + swiglu_alpha = mlp_compute_input.swiglu_alpha + swiglu_beta = mlp_compute_input.swiglu_beta + activation_situ_beta = mlp_compute_input.activation_situ_beta + activation_situ_linear_beta = mlp_compute_input.activation_situ_linear_beta + + if not mlp_compute_input.quant.is_quant: + return unquant_apply_mlp( + hidden_states=hidden_states, + w1=w1, + w2=w2, + w1_bias=w1_bias, + w2_bias=w2_bias, + activation=activation, + activation_situ_beta=activation_situ_beta, + activation_situ_linear_beta=activation_situ_linear_beta, + group_list=group_list, + group_list_type=group_list_type, + topk_scales=topk_scales, + need_trans=need_trans, + swiglu_limit=swiglu_limit, + swiglu_alpha=swiglu_alpha, + swiglu_beta=swiglu_beta, + lora_context=mlp_compute_input.lora_context, + expanded_row_idx=mlp_compute_input.expanded_row_idx, + topk_ids=mlp_compute_input.topk_ids, + ) + + from vllm_ascend.lora.fused_moe import has_lora + + if has_lora(mlp_compute_input.lora_context): + from vllm_ascend.lora.quant_moe import quant_apply_mlp_with_moe_lora + + return quant_apply_mlp_with_moe_lora(mlp_compute_input=mlp_compute_input) + + assert w1_scale is not None and w2_scale is not None + act_quant_type = torch.int8 if mlp_compute_input.quant.is_int_quant else torch.float8_e4m3fn + weight_quant_type = torch.float8_e4m3fn + scale_type = None + per_token_scale_type = None + use_bf16 = hidden_states.dtype == torch.bfloat16 + use_mxfp_quant = mlp_compute_input.quant.is_mxfp + mxfp_quant_dtype = mlp_compute_input.quant.quant_type + + if use_mxfp_quant: + mxfp = mlp_compute_input.quant.mxfp + assert mxfp is not None, "mlp_compute_input.quant.mxfp is required when quant_type is W8A8MXFP." + act_quant_type = mxfp.act_quant_type or act_quant_type + if mxfp_quant_dtype == QuantType.W4A16MXFP: + act_quant_type = mxfp.act_quant_type + weight_quant_type = mxfp.weight_quant_type or weight_quant_type + if mxfp_quant_dtype in [QuantType.W4A8MXFP, QuantType.W4A16MXFP]: + weight_quant_type = mxfp.weight_quant_type + scale_type = mxfp.scale_dtype + per_token_scale_type = mxfp.per_token_scale_dtype + use_bf16 = mxfp.use_bf16 + + return quant_apply_mlp( + hidden_states=hidden_states, + w1=w1, + w1_scale=w1_scale, + w2=w2, + w2_scale=w2_scale, + group_list=group_list, + dynamic_scale=dynamic_scale, + group_list_type=group_list_type, + w1_scale_bias=w1_scale_bias, + w2_scale_bias=w2_scale_bias, + w1_offset=w1_offset, + w2_offset=w2_offset, + fusion=fusion, + dynamic_eplb=dynamic_eplb, + use_mxfp_quant=use_mxfp_quant, + mxfp_quant_dtype=mxfp_quant_dtype, + act_quant_type=act_quant_type, + weight_quant_type=weight_quant_type, + scale_type=scale_type, + per_token_scale_type=per_token_scale_type, + use_bf16=use_bf16, + activation=activation, + activation_situ_beta=activation_situ_beta, + activation_situ_linear_beta=activation_situ_linear_beta, + swiglu_limit=swiglu_limit, + swiglu_alpha=swiglu_alpha, + swiglu_beta=swiglu_beta, + use_w4a8_per_channel_gmm_swiglu=mlp_compute_input.quant.use_w4a8_per_channel_gmm_swiglu, + ) diff --git a/vllm_ascend/quantization/modelslim_config.py b/vllm_ascend/quantization/modelslim_config.py index 08912fff59a8..e8f1f83e16df 100644 --- a/vllm_ascend/quantization/modelslim_config.py +++ b/vllm_ascend/quantization/modelslim_config.py @@ -785,14 +785,6 @@ def get_quant_method(self, layer: torch.nn.Module, prefix: str, tid2eid=None) -> prefix = self.quant_prefix_mapper(model_type, prefix) if isinstance(layer, LinearBase): - if model_type in ("kimi_k3", "kimi_linear") and self.uses_kimi_k3_mixed_kda_projection(prefix): - # The Ascend K3 adapter replaces this temporary module with a - # W8A8 q/k/v projection plus a FLOAT gate projection directly - # after upstream construction. - from vllm_ascend.ops.linear import AscendUnquantizedLinearMethod - - logger.debug("Temporarily select unquantized Kimi K3 mixed KDA projection for %s", prefix) - return AscendUnquantizedLinearMethod() if self.is_layer_skipped_ascend(prefix, self.packed_modules_mapping): # Delayed import to avoid circular import from vllm_ascend.ops.linear import AscendUnquantizedLinearMethod @@ -839,6 +831,11 @@ def get_quant_method(self, layer: torch.nn.Module, prefix: str, tid2eid=None) -> def is_layer_skipped_ascend(self, prefix: str, fused_mapping: Mapping[str, list[str]] = MappingProxyType({})): # adapted from vllm.model_executor.layers.quantization.utils.quant_utils.is_layer_skipped + if self.model_type in ("kimi_k3", "kimi_linear") and self.uses_kimi_k3_mixed_kda_projection(prefix): + # The model adapter replaces this temporary packed projection with + # separate quantized q/k/v and floating-point gate projections. + return True + proj_name = prefix.split(".")[-1] if proj_name in fused_mapping: shard_prefixes = [ From dae7c134eaa8c42c742b9f472b21baa968a347e0 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Mon, 24 Aug 2026 04:25:58 -0500 Subject: [PATCH 11/50] feat(model): register Kimi K3 adapters Compose the upstream Kimi text, multimodal, MTP, and DSpark model structures with Ascend KDA, MLA, MoE, quantization, sequence-parallel, and auxiliary-state contracts. Keep ViT FIA key/value inputs contiguous at the operator boundary, matching the validated v0.26 path. Add a reduced nightly execution fixture and focused model tests. Signed-off-by: maoxx241 --- .github/workflows/configs/nightly_config.yaml | 4 + .github/workflows/scripts/test_config.yaml | 1 + .../kimi_k3_5layers_16experts/config.json | 63 ++ .../models/test_kimi_k3_execution_parity.py | 110 +++ tests/ut/model_executor/test_qwen3_dspark.py | 1 + tests/ut/models/test_kimi_k3_adapter.py | 799 ++++++++++++++++ tests/ut/worker/a2/test_model_runner_v1.py | 46 + vllm_ascend/models/__init__.py | 22 + vllm_ascend/models/kimi_k3.py | 860 ++++++++++++++++++ vllm_ascend/models/kimi_k3_dspark.py | 320 +++++++ vllm_ascend/models/kimi_k3_mtp.py | 93 ++ vllm_ascend/models/llama_eagle3.py | 32 + vllm_ascend/models/qwen3_dspark.py | 29 +- vllm_ascend/ops/mm_encoder_attention.py | 4 +- vllm_ascend/worker/model_runner_v1.py | 24 + 15 files changed, 2405 insertions(+), 3 deletions(-) create mode 100644 tests/e2e/nightly/single_node/models/fixtures/kimi_k3_5layers_16experts/config.json create mode 100644 tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py create mode 100644 tests/ut/models/test_kimi_k3_adapter.py create mode 100644 vllm_ascend/models/kimi_k3.py create mode 100644 vllm_ascend/models/kimi_k3_dspark.py create mode 100644 vllm_ascend/models/kimi_k3_mtp.py diff --git a/.github/workflows/configs/nightly_config.yaml b/.github/workflows/configs/nightly_config.yaml index 716adc58dab9..8f5adeaff85e 100644 --- a/.github/workflows/configs/nightly_config.yaml +++ b/.github/workflows/configs/nightly_config.yaml @@ -231,6 +231,10 @@ a3: multi_card: test_config: # pytest-driven tests + - name: kimi-k3-execution-parity + os: linux-aarch64-nightly-a3-16 + tests: tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py + testcase_timeout: 180 - name: qwen3-30b-acc os: linux-aarch64-nightly-a3-4 tests: tests/e2e/weekly/single_node/models/test_qwen3_30b_acc.py diff --git a/.github/workflows/scripts/test_config.yaml b/.github/workflows/scripts/test_config.yaml index daacd3bf75c1..ea08d20fd912 100644 --- a/.github/workflows/scripts/test_config.yaml +++ b/.github/workflows/scripts/test_config.yaml @@ -308,6 +308,7 @@ - tests/ut/models/test_deepseek_v4_compressor.py - tests/ut/models/test_deepseek_v4_indexer.py - tests/ut/models/test_deepseek_v4_moe.py + - tests/ut/models/test_kimi_k3_adapter.py - tests/e2e/pull_request/four_card/test_deepseek_v4.py - name: models_minimax_m3 diff --git a/tests/e2e/nightly/single_node/models/fixtures/kimi_k3_5layers_16experts/config.json b/tests/e2e/nightly/single_node/models/fixtures/kimi_k3_5layers_16experts/config.json new file mode 100644 index 000000000000..0dd1a5e42c0e --- /dev/null +++ b/tests/e2e/nightly/single_node/models/fixtures/kimi_k3_5layers_16experts/config.json @@ -0,0 +1,63 @@ +{ + "activation_situ_beta": 4, + "activation_situ_linear_beta": 25, + "architectures": [ + "KimiLinearForCausalLM" + ], + "attn_res_block_size": 12, + "bos_token_id": 163584, + "dtype": "bfloat16", + "eos_token_id": 163586, + "first_k_dense_replace": 1, + "hidden_act": "situ", + "hidden_size": 7168, + "intermediate_size": 33792, + "kv_lora_rank": 512, + "latent_moe_use_norm": true, + "linear_attn_config": { + "full_attn_layers": [ + 4 + ], + "gate_lower_bound": -5, + "head_dim": 128, + "kda_layers": [ + 1, + 2, + 3, + 5 + ], + "num_heads": 96, + "short_conv_kernel_size": 4, + "use_full_rank_gate": true + }, + "max_position_embeddings": 1048576, + "mla_use_nope": true, + "mla_use_output_gate": true, + "model_type": "kimi_linear", + "moe_intermediate_size": 3072, + "moe_layer_freq": 1, + "moe_renormalize": true, + "moe_router_activation_func": "sigmoid", + "num_attention_heads": 96, + "num_expert_group": 1, + "num_experts": 16, + "num_experts_per_token": 16, + "num_hidden_layers": 5, + "num_key_value_heads": 96, + "num_nextn_predict_layers": 0, + "num_shared_experts": 2, + "pad_token_id": 163839, + "q_lora_rank": 1536, + "qk_nope_head_dim": 128, + "qk_rope_head_dim": 64, + "rms_norm_eps": 1e-05, + "routed_expert_hidden_size": 3584, + "routed_scaling_factor": 1, + "tie_word_embeddings": false, + "topk_group": 1, + "topk_method": "noaux_tc", + "use_cache": true, + "use_grouped_topk": true, + "v_head_dim": 128, + "vocab_size": 163840 +} diff --git a/tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py b/tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py new file mode 100644 index 000000000000..a28c793923b5 --- /dev/null +++ b/tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py @@ -0,0 +1,110 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved. +"""Storage-light Kimi K3 execution parity guard. + +The committed fixture keeps Kimi K3's production dimensions and its mixed +KDA/MLA layout, but limits the model to five layers and sixteen experts. Dummy +weights deliberately make this an execution-parity test, not a semantic +accuracy test. Full-checkpoint GPQA remains a separate release gate. +""" + +from pathlib import Path + +import pytest +import torch +from vllm import SamplingParams +from vllm.inputs import TokensPrompt + +from tests.e2e.conftest import VllmRunner, wait_until_npu_memory_free + +MODEL_CONFIG = Path(__file__).parent / "fixtures" / "kimi_k3_5layers_16experts" +SCHEDULER_BLOCK_SIZE = 16 +PROMPT_TOKEN_IDS = [163584, *range(100, 100 + SCHEDULER_BLOCK_SIZE)] +MAX_TOKENS = 4 + + +def _assert_complete_output(request_output): + assert request_output is not None + assert request_output.finished + assert request_output.outputs is not None + assert len(request_output.outputs) == 1 + + completion = request_output.outputs[0] + assert completion is not None + assert completion.token_ids is not None + assert len(completion.token_ids) == MAX_TOKENS + assert completion.logprobs is not None + assert len(completion.logprobs) == MAX_TOKENS + + chosen_logprobs = [] + for token_id, step_logprobs in zip(completion.token_ids, completion.logprobs): + assert step_logprobs is not None + assert token_id in step_logprobs + logprob = step_logprobs[token_id].logprob + assert logprob is not None + assert torch.isfinite(torch.tensor(logprob)) + chosen_logprobs.append(logprob) + + return list(completion.token_ids), torch.tensor(chosen_logprobs, dtype=torch.float32) + + +@pytest.mark.e2e_model("sgl-npu/Kimi-K3-W4A8") +@pytest.mark.e2e_coverage( + arch="moe", + feature="aclgraph,prefix_caching,logprobs", + parallel="TP,EP", + deploy="pd_mix", + hardware="A3", + quantization="BF16", + graph_mode="full_decode_only", +) +@wait_until_npu_memory_free() +def test_kimi_k3_dummy_prefix_cache_one_token_prefill_parity(): + """Compare cold prefill with the cached block-size-plus-one path.""" + sampling_params = SamplingParams( + temperature=0, + max_tokens=MAX_TOKENS, + logprobs=1, + ignore_eos=True, + seed=0, + ) + prompt = TokensPrompt(prompt_token_ids=PROMPT_TOKEN_IDS) + + with VllmRunner( + str(MODEL_CONFIG), + skip_tokenizer_init=True, + load_format="dummy", + dtype="bfloat16", + seed=0, + block_size=SCHEDULER_BLOCK_SIZE, + max_model_len=64, + max_num_seqs=1, + max_num_batched_tokens=64, + tensor_parallel_size=16, + enable_expert_parallel=True, + enable_prefix_caching=True, + gpu_memory_utilization=0.75, + compilation_config={ + "cudagraph_mode": "FULL_DECODE_ONLY", + "cudagraph_capture_sizes": [1], + }, + ) as vllm_model: + cold = vllm_model.model.generate([prompt], sampling_params, use_tqdm=False)[0] + hit = vllm_model.model.generate([prompt], sampling_params, use_tqdm=False)[0] + + assert cold.num_cached_tokens in (None, 0) + assert hit.num_cached_tokens == SCHEDULER_BLOCK_SIZE + assert len(PROMPT_TOKEN_IDS) - hit.num_cached_tokens == 1 + + cold_tokens, cold_logprobs = _assert_complete_output(cold) + hit_tokens, hit_logprobs = _assert_complete_output(hit) + assert hit_tokens == cold_tokens + torch.testing.assert_close(hit_logprobs, cold_logprobs, rtol=5e-3, atol=5e-3) + + assert vllm_model.model.reset_prefix_cache() + reset = vllm_model.model.generate([prompt], sampling_params, use_tqdm=False)[0] + assert reset.num_cached_tokens in (None, 0) + reset_tokens, reset_logprobs = _assert_complete_output(reset) + assert reset_tokens == cold_tokens + torch.testing.assert_close(reset_logprobs, cold_logprobs, rtol=5e-3, atol=5e-3) diff --git a/tests/ut/model_executor/test_qwen3_dspark.py b/tests/ut/model_executor/test_qwen3_dspark.py index 7a799917570c..d1cc7893dd1a 100644 --- a/tests/ut/model_executor/test_qwen3_dspark.py +++ b/tests/ut/model_executor/test_qwen3_dspark.py @@ -59,6 +59,7 @@ def test_rotates_only_fc_weights(self) -> None: mock_get_rotation_matrix.assert_called_once_with(rotation_path) mock_parent_load_weights.assert_called_once() + assert model._shared_layer_rotation is rotation_matrix processed_weights = mock_parent_load_weights.call_args.args[0] torch.testing.assert_close(processed_weights[0][1], expected_fc_weight) diff --git a/tests/ut/models/test_kimi_k3_adapter.py b/tests/ut/models/test_kimi_k3_adapter.py new file mode 100644 index 000000000000..505db8091a5a --- /dev/null +++ b/tests/ut/models/test_kimi_k3_adapter.py @@ -0,0 +1,799 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import MethodType, SimpleNamespace +from unittest.mock import MagicMock, patch + +import torch +from torch import nn +from vllm.config import VllmConfig, set_current_vllm_config + +from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec +from vllm_ascend.models import kimi_k3 +from vllm_ascend.models.kimi_k3 import ( + AscendKimiK3ForConditionalGeneration, + AscendKimiK3MultiModalProjector, + AscendKimiLinearForCausalLM, + AscendKimiLinearModel, + AscendKimiMLAAttention, + AscendKimiMoE, +) +from vllm_ascend.models.kimi_k3_dspark import ( + AscendK3DSparkDecoderLayer, + AscendK3DSparkForCausalLM, +) + + +def test_ascend_attn_res_matches_canonical_k3_math(): + prefix_sum = torch.tensor([[1.0, 2.0], [3.0, 4.0]]) + block_residual = torch.tensor( + [ + [[0.5, 1.5], [2.5, 3.5], [1000.0, 1000.0]], + [[1.0, 0.0], [0.0, 1.0], [1000.0, 1000.0]], + ] + ) + norm = SimpleNamespace(weight=torch.tensor([1.0, 1.5]), variance_epsilon=1e-5) + proj = SimpleNamespace(weight=torch.tensor([[0.25, -0.5]])) + + output = kimi_k3._apply_ascend_attn_res( + prefix_sum, + block_residual, + proj, + norm, + num_valid_blocks=2, + ) + + values = torch.cat( + (block_residual[:, :2], prefix_sum.unsqueeze(1)), + dim=1, + ).float() + inverse_rms = torch.rsqrt(values.square().mean(-1, keepdim=True) + norm.variance_epsilon) + normalized_without_gamma = values * inverse_rms + score_weight = norm.weight.float() * proj.weight.squeeze(0).float() + probabilities = (normalized_without_gamma * score_weight).sum(-1).softmax(-1).unsqueeze(1) + expected = torch.matmul(probabilities, values).squeeze(1).to(prefix_sum.dtype) + torch.testing.assert_close(output, expected) + + +def _make_moe_config(**overrides): + values = { + "hidden_size": 16, + "moe_intermediate_size": 32, + "num_experts": 8, + "num_experts_per_token": 2, + "moe_renormalize": True, + "routed_expert_hidden_size": None, + "latent_moe_use_norm": False, + "routed_scaling_factor": 1.0, + "num_shared_experts": None, + "hidden_act": "silu", + "activation_situ_beta": None, + "activation_situ_linear_beta": None, + "use_grouped_topk": False, + "num_expert_group": None, + "topk_group": None, + "moe_router_activation_func": "softmax", + "rms_norm_eps": 1e-6, + } + values.update(overrides) + return SimpleNamespace(**values) + + +def test_ascend_kimi_moe_uses_standard_runner_dispatch(monkeypatch): + class FakeGate(nn.Module): + def __init__(self, **kwargs): + super().__init__() + self.kwargs = kwargs + + factory = MagicMock(return_value=nn.Identity()) + monkeypatch.setattr(kimi_k3, "GateLinear", FakeGate) + monkeypatch.setattr(kimi_k3, "FusedMoEFactory", factory) + + AscendKimiMoE( + config=_make_moe_config(), + prefix="model.layers.1.block_sparse_moe", + use_sequence_parallel=True, + ) + + assert factory.call_args.kwargs["intermediate_size"] == 32 + assert factory.call_args.kwargs["is_sequence_parallel"] is True + assert "runner_cls" not in factory.call_args.kwargs + + +def test_dspark_decoder_uses_upstream_mlp_activation_contract( + monkeypatch, +): + config = SimpleNamespace( + hidden_size=8, + num_attention_heads=2, + qk_nope_head_dim=2, + qk_rope_head_dim=2, + v_head_dim=2, + q_lora_rank=4, + kv_lora_rank=4, + intermediate_size=16, + hidden_act="silu", + rms_norm_eps=1e-6, + ) + vllm_config = SimpleNamespace(cache_config=None) + mlp_factory = MagicMock(return_value=nn.Identity()) + monkeypatch.setattr( + "vllm_ascend.models.kimi_k3_dspark.get_draft_quant_config", + lambda _: None, + ) + monkeypatch.setattr( + "vllm_ascend.models.kimi_k3_dspark.AscendKimiMLAAttention", + lambda **_: nn.Identity(), + ) + monkeypatch.setattr( + "vllm_ascend.models.kimi_k3_dspark.KimiMLP", + mlp_factory, + ) + + with set_current_vllm_config(VllmConfig()): + AscendK3DSparkDecoderLayer( + vllm_config=vllm_config, + config=config, + layer_idx=0, + start_layer_id=4, + prefix="model", + ) + + assert mlp_factory.call_args.kwargs["hidden_act"] == "silu" + assert "activation_situ_beta" not in mlp_factory.call_args.kwargs + assert "activation_situ_linear_beta" not in mlp_factory.call_args.kwargs + + +def test_ascend_kimi_moe_quantizes_modelslim_latent_projections(monkeypatch): + class FakeGate(nn.Module): + def __init__(self, **kwargs): + super().__init__() + + class FakeLinear(nn.Module): + def __init__(self, input_size, output_size, **kwargs): + super().__init__() + self.input_size = input_size + self.output_size = output_size + self.kwargs = kwargs + + config = _make_moe_config(routed_expert_hidden_size=8) + quant_config = MagicMock() + quant_config.get_name.return_value = "ascend" + factory = MagicMock(return_value=nn.Identity()) + monkeypatch.setattr(kimi_k3, "GateLinear", FakeGate) + monkeypatch.setattr(kimi_k3, "ReplicatedLinear", FakeLinear) + monkeypatch.setattr(kimi_k3, "FusedMoEFactory", factory) + + moe = AscendKimiMoE( + config=config, + quant_config=quant_config, + prefix="model.layers.1.block_sparse_moe", + ) + + assert moe.routed_expert_down_proj.input_size == 16 + assert moe.routed_expert_down_proj.output_size == 8 + assert moe.routed_expert_down_proj.kwargs == { + "bias": False, + "quant_config": quant_config, + "prefix": "model.layers.1.block_sparse_moe.routed_expert_down_proj", + } + assert moe.routed_expert_up_proj.input_size == 8 + assert moe.routed_expert_up_proj.output_size == 16 + assert moe.routed_expert_up_proj.kwargs == { + "bias": False, + "quant_config": quant_config, + "prefix": "model.layers.1.block_sparse_moe.routed_expert_up_proj", + } + assert factory.call_args.kwargs["routed_input_transform"] is moe.routed_expert_down_proj + assert factory.call_args.kwargs["routed_output_transform"] is moe.routed_output_transform + assert "runner_cls" not in factory.call_args.kwargs + + +def test_kimi_mixed_kda_gate_weights_use_upstream_packed_loader(monkeypatch): + model = AscendKimiLinearModel.__new__(AscendKimiLinearModel) + nn.Module.__init__(model) + layer = nn.Module() + layer.self_attn = nn.Module() + layer.self_attn.in_proj_gfab = nn.Module() + packed_weight = nn.Parameter(torch.empty(6, 4)) + layer.self_attn.in_proj_gfab.register_parameter("weight", packed_weight) + layer.router = nn.Linear(4, 1, bias=False) + model.layers = nn.ModuleList([layer]) + + remaining = [] + + def fake_upstream_load_weights(_self, weights): + remaining.extend(weights) + return {name for name, *_ in remaining} + + monkeypatch.setattr( + kimi_k3.UpstreamKimiLinearModel, + "load_weights", + fake_upstream_load_weights, + ) + source_weights = [ + ("layers.0.router.weight", torch.full((1, 4), 0.5)), + ("layers.0.self_attn.g_proj.weight", torch.full((1,), 1.0)), + ("layers.0.self_attn.f_a_proj.weight", torch.full((1,), 2.0)), + ("layers.0.self_attn.b_proj.weight", torch.full((1,), 3.0)), + ("layers.0.self_attn.o_proj.weight", torch.full((1,), 4.0)), + ] + + loaded = model.load_weights(iter(source_weights)) + + assert remaining[0] == source_weights[0] + assert remaining[-1] == source_weights[-1] + assert [name for name, _, _ in remaining[1:4]] == [ + "layers.0.self_attn.in_proj_gfab.weight", + ] * 3 + assert [loaded_weight.item() for _, loaded_weight, _ in remaining[1:4]] == [1.0, 2.0, 3.0] + assert [kwargs["loaded_shard_id"] for _, _, kwargs in remaining[1:4]] == [0, 1, 2] + assert loaded == { + "layers.0.self_attn.in_proj_gfab.weight", + "layers.0.router.weight", + "layers.0.self_attn.o_proj.weight", + } + + +def test_kimi_text_model_layer_factory_accepts_prefix_keyword(monkeypatch): + config = SimpleNamespace( + vocab_size=64, + hidden_size=16, + num_hidden_layers=1, + rms_norm_eps=1e-5, + attn_res_block_size=None, + num_attention_heads=1, + ) + vllm_config = MagicMock() + vllm_config.model_config.hf_text_config = config + vllm_config.parallel_config = SimpleNamespace( + pipeline_parallel_size=1, + enable_expert_parallel=True, + tensor_parallel_size=2, + ) + pp_group = SimpleNamespace(is_first_rank=False, is_last_rank=False) + decoder_layer = nn.Identity() + decoder_layer_factory = MagicMock(return_value=decoder_layer) + + def fake_make_layers(num_hidden_layers, layer_fn, *, prefix): + assert num_hidden_layers == 1 + assert layer_fn(prefix=f"{prefix}.0") is decoder_layer + return 0, 1, nn.ModuleList([decoder_layer]) + + monkeypatch.setattr(kimi_k3, "get_pp_group", lambda: pp_group) + monkeypatch.setattr(kimi_k3, "get_tensor_model_parallel_world_size", lambda: 1) + monkeypatch.setattr(kimi_k3, "AscendKimiDecoderLayer", decoder_layer_factory) + monkeypatch.setattr(kimi_k3, "make_layers", fake_make_layers) + + model = AscendKimiLinearModel(vllm_config=vllm_config, prefix="model") + + assert model.start_layer == 0 + assert model.end_layer == 1 + decoder_layer_factory.assert_called_once_with( + config, + vllm_config, + "model.layers.0", + use_sequence_parallel=True, + ) + + +def test_kimi_mla_cache_spec_preserves_hybrid_page_padding(): + real_page_size = 128 * 576 * torch.bfloat16.itemsize + padded_page_size = real_page_size + 128 + spec = AscendMLAAttentionSpec( + block_size=128, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + page_size_padded=padded_page_size, + ) + + assert spec.real_page_size_bytes == real_page_size + assert spec.page_size_bytes == padded_page_size + assert AscendMLAAttentionSpec.merge([spec, spec]).page_size_bytes == padded_page_size + + +def test_ascend_mla_exposes_layer_and_cache_contract(): + attention = AscendKimiMLAAttention.__new__(AscendKimiMLAAttention) + layer = MagicMock() + layer.layer_name = "model.layers.1.self_attn.attn" + layer.impl = object() + layer.kv_cache = (object(), object()) + layer.kv_cache_dtype = "auto" + layer._k_scale = 1.0 + attention.mla_attn = MagicMock() + attention.mla_attn.mla_attn = layer + + assert attention.layer_name == layer.layer_name + assert attention.impl is layer.impl + assert attention.kv_cache is layer.kv_cache + assert attention.kv_cache_dtype == layer.kv_cache_dtype + assert attention._k_scale == layer._k_scale + + +def test_kimi_attention_residual_stays_sequence_sharded(monkeypatch): + class IdentityAttention(nn.Module): + def forward(self, *, hidden_states, positions): + del positions + return hidden_states + + layer = kimi_k3.AscendKimiDecoderLayer.__new__(kimi_k3.AscendKimiDecoderLayer) + nn.Module.__init__(layer) + layer.use_sequence_parallel = True + layer.prev_valid_blocks = 0 + layer.is_block_write_layer = False + layer.input_layernorm = nn.Identity() + layer.post_attention_layernorm = nn.Identity() + layer.mlp = nn.Identity() + layer.self_attention_res_proj = object() + layer.self_attention_res_norm = object() + layer.mlp_res_proj = object() + layer.mlp_res_norm = object() + layer.self_attn = IdentityAttention() + + collective_shapes = [] + + def fake_all_gather(hidden_states): + collective_shapes.append(("gather", hidden_states.shape)) + return torch.cat((hidden_states, hidden_states), dim=0) + + def fake_reduce_scatter(hidden_states): + collective_shapes.append(("reduce_scatter", hidden_states.shape)) + return hidden_states.chunk(2, dim=0)[0] + + monkeypatch.setattr(kimi_k3, "sp_all_gather", fake_all_gather) + monkeypatch.setattr(kimi_k3, "sp_reduce_scatter", fake_reduce_scatter) + monkeypatch.setattr( + kimi_k3, + "_apply_ascend_attn_res", + lambda prefix_sum, *_args, **_kwargs: prefix_sum, + ) + + hidden_states = torch.arange(4, dtype=torch.float32).view(2, 2) + block_residual = torch.zeros(2, 1, 2) + output, returned_residual = layer.forward_attn_residual( + positions=torch.arange(3), + hidden_states=hidden_states, + block_residual=block_residual, + ) + + assert collective_shapes == [ + ("gather", torch.Size([2, 2])), + ("reduce_scatter", torch.Size([3, 2])), + ] + assert output.shape == torch.Size([2, 2]) + assert returned_residual.shape == torch.Size([2, 1, 2]) + + +def test_kimi_model_allocates_attention_residual_after_sp_shard(monkeypatch): + class RecordingLayer(nn.Module): + def __init__(self): + super().__init__() + self.residual_shape = None + + def forward(self, *, positions, hidden_states, residual): + self.residual_shape = residual.shape + return hidden_states, residual + + model = AscendKimiLinearModel.__new__(AscendKimiLinearModel) + nn.Module.__init__(model) + model.config = SimpleNamespace(attn_res_block_size=12) + model.start_layer = 0 + model.end_layer = 1 + layer = RecordingLayer() + model.layers = nn.ModuleList([layer]) + model.use_sequence_parallel = True + model.aux_hidden_state_layers = set() + model.output_attn_res_proj = object() + model.output_attn_res_norm = object() + model._maybe_add_hidden_state = MethodType( + lambda self, states, *_args: states, + model, + ) + + monkeypatch.setattr( + kimi_k3, + "get_pp_group", + lambda: SimpleNamespace(is_first_rank=True, is_last_rank=True), + ) + monkeypatch.setattr( + kimi_k3, + "sp_shard", + lambda hidden_states: torch.nn.functional.pad(hidden_states, (0, 0, 0, 1))[:2], + ) + monkeypatch.setattr( + kimi_k3, + "sp_all_gather", + lambda hidden_states: torch.cat((hidden_states, hidden_states), dim=0), + ) + monkeypatch.setattr( + kimi_k3, + "_apply_ascend_attn_res", + lambda hidden_states, *_args, **_kwargs: hidden_states, + ) + + output = model( + input_ids=None, + positions=torch.arange(3), + intermediate_tensors=None, + inputs_embeds=torch.arange(6, dtype=torch.float32).view(3, 2), + ) + + assert layer.residual_shape == torch.Size([2, 1, 2]) + assert output.shape == torch.Size([3, 2]) + + +def test_kimi_model_selects_materialized_or_raw_dspark_aux_stream(monkeypatch): + class Marker(nn.Module): + def __init__(self, value: int) -> None: + super().__init__() + self.value = value + + class RecordingLayer(nn.Module): + def __init__(self, layer_idx: int) -> None: + super().__init__() + self.layer_idx = layer_idx + self.prev_valid_blocks = layer_idx + self.self_attention_res_proj = Marker(layer_idx) + self.self_attention_res_norm = nn.Identity() + + def forward(self, *, positions, hidden_states, residual): + del positions + materialized = kimi_k3._apply_ascend_attn_res( + hidden_states, + residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + self.prev_valid_blocks, + ) + return materialized + 10, residual + + residual_calls: list[int] = [] + + def fake_attn_res(prefix_sum, _residual, projection, _norm, num_valid_blocks): + residual_calls.append(projection.value) + return prefix_sum + 100 * num_valid_blocks + + monkeypatch.setattr(kimi_k3, "_apply_ascend_attn_res", fake_attn_res) + monkeypatch.setattr( + kimi_k3, + "get_pp_group", + lambda: SimpleNamespace(is_first_rank=True, is_last_rank=True), + ) + + model = AscendKimiLinearModel.__new__(AscendKimiLinearModel) + nn.Module.__init__(model) + model.config = SimpleNamespace(attn_res_block_size=1) + model.start_layer = 0 + model.end_layer = 2 + model.layers = nn.ModuleList([RecordingLayer(0), RecordingLayer(1)]) + model.use_sequence_parallel = False + model.output_attn_res_proj = Marker(2) + model.output_attn_res_norm = nn.Identity() + model._set_aux_hidden_state_layers((1,)) + + model.dspark_aux_capture_materialized = True + _, materialized_aux = model( + input_ids=None, + positions=torch.tensor([0]), + intermediate_tensors=None, + inputs_embeds=torch.tensor([[1.0]]), + ) + torch.testing.assert_close(materialized_aux[0], torch.tensor([[111.0]])) + + residual_calls.clear() + model.dspark_aux_capture_materialized = False + _, raw_aux = model( + input_ids=None, + positions=torch.tensor([0]), + intermediate_tensors=None, + inputs_embeds=torch.tensor([[1.0]]), + ) + torch.testing.assert_close(raw_aux[0], torch.tensor([[11.0]])) + + +def test_kimi_dspark_aux_capture_mode_is_forwarded(): + causal_model = AscendKimiLinearForCausalLM.__new__(AscendKimiLinearForCausalLM) + nn.Module.__init__(causal_model) + causal_model.model = SimpleNamespace(dspark_aux_capture_materialized=False) + + causal_model.set_dspark_aux_capture_materialized(True) + + assert causal_model.model.dspark_aux_capture_materialized is True + + wrapper = AscendKimiK3ForConditionalGeneration.__new__(AscendKimiK3ForConditionalGeneration) + nn.Module.__init__(wrapper) + wrapper.language_model = MagicMock() + + wrapper.set_dspark_aux_capture_materialized(True) + + wrapper.language_model.set_dspark_aux_capture_materialized.assert_called_once_with(True) + + +def test_dspark_configures_upstream_mla_without_rebuilding(monkeypatch): + impl = SimpleNamespace( + scale=0.0, + rotary_emb=None, + use_mla_rope=False, + ) + layer = SimpleNamespace( + scale=0.0, + non_causal_multi_token_decode=False, + impl=impl, + ) + upstream_wrapper = SimpleNamespace(mla_attn=layer) + + def fake_upstream_init(self, **_kwargs): + nn.Module.__init__(self) + self.scaling = 0.125 + self.mla_attn = upstream_wrapper + + rotary_emb = object() + monkeypatch.setattr( + kimi_k3.UpstreamKimiMLAAttention, + "__init__", + fake_upstream_init, + ) + monkeypatch.setattr(kimi_k3, "get_rope", lambda *_args, **_kwargs: rotary_emb) + + attention = AscendKimiMLAAttention( + config=SimpleNamespace( + rope_parameters={"rope_type": "default"}, + max_position_embeddings=4096, + ), + hidden_size=16, + num_heads=2, + qk_nope_head_dim=4, + qk_rope_head_dim=4, + v_head_dim=4, + q_lora_rank=8, + kv_lora_rank=8, + use_output_gate=False, + use_rope=True, + prefix="model.layers.1.self_attn", + non_causal_multi_token_decode=True, + ) + + assert attention.mla_attn is upstream_wrapper + assert layer.scale == attention.scaling + assert layer.non_causal_multi_token_decode is True + assert impl.scale == attention.scaling + assert impl.rotary_emb is rotary_emb + assert impl.use_mla_rope is True + + +def test_projector_applies_optional_modelslim_rotation(): + class ScaleLinear(nn.Module): + def forward(self, hidden_states): + return hidden_states * 2, None + + projector = AscendKimiK3MultiModalProjector.__new__(AscendKimiK3MultiModalProjector) + nn.Module.__init__(projector) + image_features = torch.tensor([[1.0, 2.0]]) + + with patch.object( + kimi_k3.KimiK25MultiModalProjector, + "forward", + lambda self, hidden_states: hidden_states, + ): + projector.rot_proj = ScaleLinear() + torch.testing.assert_close( + projector(image_features), + image_features * 2, + ) + projector.rot_proj = None + torch.testing.assert_close(projector(image_features), image_features) + + +def test_projector_creates_rotation_only_when_enabled(monkeypatch): + def fake_upstream_init(self, *_args, **_kwargs): + nn.Module.__init__(self) + + monkeypatch.setattr( + kimi_k3.KimiK25MultiModalProjector, + "__init__", + fake_upstream_init, + ) + rotation = nn.Linear(1, 1, bias=False) + rotation_factory = MagicMock(return_value=rotation) + monkeypatch.setattr(kimi_k3, "ReplicatedLinear", rotation_factory) + config = SimpleNamespace(text_hidden_size=16) + + plain_projector = AscendKimiK3MultiModalProjector(config, prefix="mm_projector") + + assert plain_projector.rot_proj is None + rotation_factory.assert_not_called() + + rotated_projector = AscendKimiK3MultiModalProjector( + config, + prefix="mm_projector", + enable_rotation=True, + ) + + assert rotated_projector.rot_proj is rotation + rotation_factory.assert_called_once_with( + 16, + 16, + bias=False, + quant_config=None, + prefix="mm_projector.rot_proj", + ) + + +class _DraftTokenEmbedder(nn.Module): + def __init__(self) -> None: + super().__init__() + self.embedding = nn.Embedding.from_pretrained( + torch.tensor( + [ + [0.0, 0.0], + [1.0, 2.0], + [3.0, 4.0], + ] + ) + ) + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.embedding(input_ids) + + +def _make_k3_dspark_for_embedding_test() -> AscendK3DSparkForCausalLM: + model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) + nn.Module.__init__(model) + model.model = _DraftTokenEmbedder() + return model + + +def test_k3_dspark_load_weights_keeps_per_layer_context_kv(monkeypatch): + model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) + nn.Module.__init__(model) + model.rotation_path = None + source_weights = [ + ( + "layers.0.self_attn.kv_a_proj_with_mqa.weight", + torch.ones(1, 1), + ) + ] + seen_names: list[str] = [] + + class CapturingLoader: + def __init__(self, loaded_model): + assert loaded_model is model + + def load_weights(self, weights, *, mapper): + assert mapper is model.hf_to_vllm_mapper + seen_names.extend(name for name, _ in weights) + return {"model.layers.0.self_attn.fused_qkv_a_proj.weight"} + + monkeypatch.setattr( + "vllm_ascend.models.kimi_k3_dspark.AutoWeightsLoader", + CapturingLoader, + ) + + loaded = model.load_weights(iter(source_weights)) + + assert seen_names == [source_weights[0][0]] + assert loaded == {"model.layers.0.self_attn.fused_qkv_a_proj.weight"} + + +def test_k3_dspark_reuses_modelslim_rotation_loader(monkeypatch): + model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) + nn.Module.__init__(model) + model.rotation_path = "rotation.safetensors" + source_weights = [ + ("context_proj.weight", torch.ones(2, 4)), + ("context_norm.weight", torch.ones(2)), + ] + rotated_weight = torch.full((2, 4), 2.0) + seen_weights: list[tuple[str, torch.Tensor]] = [] + + class CapturingLoader: + def __init__(self, loaded_model): + assert loaded_model is model + + def load_weights(self, weights, *, mapper): + assert mapper is model.hf_to_vllm_mapper + seen_weights.extend(weights) + return {name for name, _ in seen_weights} + + monkeypatch.setattr( + "vllm_ascend.models.kimi_k3_dspark.AutoWeightsLoader", + CapturingLoader, + ) + rotation = torch.eye(4) + monkeypatch.setattr( + "vllm_ascend.models.kimi_k3_dspark.get_rotation_matrix", + lambda path: rotation if path == model.rotation_path else None, + ) + process_weight = MagicMock(return_value=rotated_weight) + monkeypatch.setattr( + "vllm_ascend.models.kimi_k3_dspark.process_weight", + process_weight, + ) + + model.load_weights(iter(source_weights)) + + process_weight.assert_called_once() + torch.testing.assert_close(process_weight.call_args.args[0], source_weights[0][1]) + torch.testing.assert_close(process_weight.call_args.args[1], rotation) + assert model._shared_layer_rotation is rotation + assert seen_weights[0][0] == "context_proj.weight" + assert seen_weights[0][1] is rotated_weight + assert seen_weights[1][0] == source_weights[1][0] + assert seen_weights[1][1] is source_weights[1][1] + + +def test_k3_dspark_prepares_unrotated_shared_layer(): + class NonCopyableCommGroup: + def __deepcopy__(self, memo): + del memo + raise TypeError("cannot pickle ProcessGroup") + + model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) + nn.Module.__init__(model) + rotation = torch.tensor([[0.0, 1.0], [-1.0, 0.0]]) + model._shared_layer_rotation = rotation + target = nn.Linear(2, 3, bias=False) + target.comm_group = NonCopyableCommGroup() + target_weight = torch.tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]) + target.weight.data.copy_(target_weight) + + prepared = model.prepare_shared_layer( + None, + target, + "draft embed_tokens.weight", + ) + + assert prepared is not None + assert prepared is not target + assert prepared.comm_group is target.comm_group + torch.testing.assert_close(prepared.weight, target_weight @ rotation.T) + torch.testing.assert_close(target.weight, target_weight) + model.finish_shared_layer_preparation() + assert model._shared_layer_rotation is None + + +def test_k3_dspark_embed_input_ids_keeps_text_only_path(): + model = _make_k3_dspark_for_embedding_test() + + output = model.embed_input_ids(torch.tensor([1, 2])) + + torch.testing.assert_close( + output, + torch.tensor([[1.0, 2.0], [3.0, 4.0]]), + ) + + +def test_k3_dspark_embed_input_ids_merges_multimodal_embeddings(): + model = _make_k3_dspark_for_embedding_test() + input_ids = torch.tensor([1, 999, 2]) + is_multimodal = torch.tensor([False, True, False]) + image_embedding = torch.tensor([[9.0, 10.0]]) + + output = model.embed_input_ids( + input_ids, + multimodal_embeddings=(image_embedding,), + is_multimodal=is_multimodal, + ) + + torch.testing.assert_close( + output, + torch.tensor( + [ + [1.0, 2.0], + [9.0, 10.0], + [3.0, 4.0], + ] + ), + ) + + +def test_k3_dspark_embed_input_ids_without_multimodal_mask_uses_text_path(): + model = _make_k3_dspark_for_embedding_test() + + output = model.embed_input_ids( + torch.tensor([1]), + multimodal_embeddings=(torch.tensor([[9.0, 10.0]]),), + ) + + torch.testing.assert_close(output, torch.tensor([[1.0, 2.0]])) diff --git a/tests/ut/worker/a2/test_model_runner_v1.py b/tests/ut/worker/a2/test_model_runner_v1.py index 107c65361d98..d99cfb308aa3 100644 --- a/tests/ut/worker/a2/test_model_runner_v1.py +++ b/tests/ut/worker/a2/test_model_runner_v1.py @@ -20,6 +20,52 @@ from vllm_ascend.worker.model_runner_v1 import NPUModelRunner +class TestDSparkAuxCaptureMode(unittest.TestCase): + def _build_runner( + self, + *, + model_type: str, + architecture: str, + use_dspark: bool = True, + ): + runner = NPUModelRunner.__new__(NPUModelRunner) + runner.speculative_config = SimpleNamespace( + use_dspark=MagicMock(return_value=use_dspark), + draft_model_config=SimpleNamespace( + hf_config=SimpleNamespace( + model_type=model_type, + architectures=[architecture], + ) + ), + ) + return runner + + def test_qwen3_gqa_dspark_uses_materialized_stream(self): + runner = self._build_runner( + model_type="qwen3", + architecture="Qwen3DSparkModel", + ) + + self.assertTrue(runner._draft_uses_qwen3_gqa_dspark()) + + def test_mla_dspark_keeps_raw_stream(self): + runner = self._build_runner( + model_type="kimi_k3_dspark", + architecture="KimiK3DSparkForCausalLM", + ) + + self.assertFalse(runner._draft_uses_qwen3_gqa_dspark()) + + def test_non_dspark_keeps_raw_stream(self): + runner = self._build_runner( + model_type="qwen3", + architecture="Qwen3DSparkModel", + use_dspark=False, + ) + + self.assertFalse(runner._draft_uses_qwen3_gqa_dspark()) + + class TestNPUModelRunnerKVCache(unittest.TestCase): def _build_runner(self): runner = NPUModelRunner.__new__(NPUModelRunner) diff --git a/vllm_ascend/models/__init__.py b/vllm_ascend/models/__init__.py index 2aa47828be4a..d0792863c00a 100644 --- a/vllm_ascend/models/__init__.py +++ b/vllm_ascend/models/__init__.py @@ -2,6 +2,28 @@ def register_model(): + ModelRegistry.register_model( + "KimiLinearForCausalLM", + "vllm_ascend.models.kimi_k3:AscendKimiLinearForCausalLM", + ) + # Keep the release-branch text architecture as a compatibility alias for + # checkpoints whose config predates vLLM's KimiLinear rename. + ModelRegistry.register_model( + "KimiK3ForCausalLM", + "vllm_ascend.models.kimi_k3:AscendKimiLinearForCausalLM", + ) + ModelRegistry.register_model( + "KimiK3ForConditionalGeneration", + "vllm_ascend.models.kimi_k3:AscendKimiK3ForConditionalGeneration", + ) + ModelRegistry.register_model( + "KimiK3MTPModel", + "vllm_ascend.models.kimi_k3_mtp:AscendKimiK3MTP", + ) + ModelRegistry.register_model( + "K3DSparkModel", + "vllm_ascend.models.kimi_k3_dspark:AscendK3DSparkForCausalLM", + ) ModelRegistry.register_model( "DeepseekV4ForCausalLM", "vllm_ascend.models.deepseek_v4.model:AscendDeepseekV4ForCausalLM" ) diff --git a/vllm_ascend/models/kimi_k3.py b/vllm_ascend/models/kimi_k3.py new file mode 100644 index 000000000000..b9569d095556 --- /dev/null +++ b/vllm_ascend/models/kimi_k3.py @@ -0,0 +1,860 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Kimi K3 model adapters for vLLM 0.27 on Ascend. + +vLLM owns Kimi's configuration, multimodal processor, weight mappings, and +model-level forward contract. This module composes those upstream pieces with +the generic MLA/MoE implementation and the Ascend KDA backend. +""" + +import math +from copy import copy + +import torch +import vllm.envs as envs +from torch import nn +from vllm.config import CacheConfig, VllmConfig +from vllm.distributed import ( + get_pp_group, + get_tensor_model_parallel_world_size, +) +from vllm.forward_context import get_forward_context, is_forward_context_available +from vllm.model_executor.layers.fused_moe import FusedMoEFactory +from vllm.model_executor.layers.fused_moe.router.gate_linear import GateLinear +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.linear import ( + ReplicatedLinear, +) +from vllm.model_executor.layers.logits_processor import LogitsProcessor +from vllm.model_executor.layers.quantization import QuantizationConfig +from vllm.model_executor.layers.rotary_embedding import get_rope +from vllm.model_executor.layers.vocab_parallel_embedding import ( + ParallelLMHead, + VocabParallelEmbedding, +) +from vllm.model_executor.models.kimi_k25_vit import ( + KimiK25MultiModalProjector, + MoonViT3dPretrainedModel, +) +from vllm.model_executor.models.utils import ( + PPMissingLayer, + init_vllm_registered_model, + make_layers, + maybe_prefix, +) +from vllm.model_executor.models.vision import is_vit_use_data_parallel +from vllm.models.common.ops.sequence_parallel import ( + sp_all_gather, + sp_padding_mask, + sp_reduce_scatter, + sp_shard, +) +from vllm.models.kimi_k3.amd.linear import ( + KimiDecoderLayer as UpstreamKimiDecoderLayer, +) +from vllm.models.kimi_k3.amd.linear import KimiLinearForCausalLM as UpstreamKimiLinearForCausalLM +from vllm.models.kimi_k3.amd.linear import KimiLinearModel as UpstreamKimiLinearModel +from vllm.models.kimi_k3.amd.linear import ( + KimiMLAAttention as UpstreamKimiMLAAttention, +) +from vllm.models.kimi_k3.amd.linear import ( + KimiMLP, + KimiRoutedOutputTransform, +) +from vllm.models.kimi_k3.amd.model import ( + KimiK3ForConditionalGeneration as UpstreamKimiK3ForConditionalGeneration, +) +from vllm.models.kimi_k3.common.mm_preprocess import ( + KimiK3DummyInputsBuilder, + KimiK3MultiModalProcessor, + KimiK3ProcessingInfo, +) +from vllm.models.kimi_k3.nvidia.model import ( + KimiLinearModel as UpstreamPackedKimiLinearModel, +) +from vllm.multimodal import MULTIMODAL_REGISTRY +from vllm.platforms import current_platform +from vllm.sequence import IntermediateTensors +from vllm.triton_utils import HAS_TRITON +from vllm.utils.math_utils import cdiv + +from vllm_ascend.models.llama_eagle3 import get_rotation_path +from vllm_ascend.ops.kimi_kda import AscendKimiK3DeltaAttention # type: ignore[import-untyped] + +if HAS_TRITON: + from vllm_ascend.ops.triton.kimi_k3.attention_residual import ( # type: ignore[import-untyped] + apply_attn_res, + ) +else: + apply_attn_res = None # type: ignore[assignment] + + +def _apply_ascend_attn_res( + prefix_sum: torch.Tensor, + block_residual: torch.Tensor, + proj: ReplicatedLinear, + norm: RMSNorm, + num_valid_blocks: int, +) -> torch.Tensor: + """Apply Kimi's canonical learned residual mixture with native ops.""" + if num_valid_blocks <= 0: + return prefix_sum + + if apply_attn_res is not None and prefix_sum.device.type == "npu" and prefix_sum.numel() > 0: + return apply_attn_res( + prefix_sum, + block_residual, + proj, + norm, + num_valid_blocks, + ) + + values = torch.cat( + ( + block_residual[:, :num_valid_blocks, :], + prefix_sum.unsqueeze(1), + ), + dim=1, + ) + values_fp32 = values.float() + inverse_rms = torch.rsqrt(values_fp32.square().mean(-1, keepdim=True) + norm.variance_epsilon) + normalized_without_gamma = values_fp32 * inverse_rms + score_weight = norm.weight.float() * proj.weight.squeeze(0).float() + scores = (normalized_without_gamma * score_weight).sum(-1) + probabilities = scores.softmax(-1).unsqueeze(1) + return torch.matmul(probabilities, values_fp32).squeeze(1).to(values.dtype) + + +class AscendKimiMoE(nn.Module): + """Kimi K3 MoE assembled from the standard vLLM MoE interfaces.""" + + def __init__( + self, + *, + config, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + use_sequence_parallel: bool = False, + ) -> None: + super().__init__() + hidden_size = config.hidden_size + moe_intermediate_size = config.moe_intermediate_size + num_experts = config.num_experts + num_experts_per_token = config.num_experts_per_token + assert moe_intermediate_size is not None + assert num_experts is not None + assert num_experts_per_token is not None + + routed_expert_hidden_size = config.routed_expert_hidden_size + self.use_latent_moe = routed_expert_hidden_size is not None + self.moe_hidden_size = routed_expert_hidden_size or hidden_size + self.latent_moe_use_norm = config.latent_moe_use_norm + self.routed_scaling_factor = config.routed_scaling_factor + self.num_shared_experts = config.num_shared_experts + activation_situ_beta = config.activation_situ_beta if config.hidden_act == "situ" else None + activation_situ_linear_beta = config.activation_situ_linear_beta if config.hidden_act == "situ" else None + + self.gate = GateLinear( + input_size=hidden_size, + output_size=num_experts, + bias=False, + out_dtype=torch.float32, + prefix=f"{prefix}.gate", + ) + self.gate.e_score_correction_bias = nn.Parameter(torch.empty(num_experts, dtype=torch.float32)) + + if self.num_shared_experts is not None: + self.shared_experts = KimiMLP( + hidden_size=hidden_size, + intermediate_size=moe_intermediate_size * self.num_shared_experts, + hidden_act=config.hidden_act, + quant_config=quant_config, + reduce_results=False, + prefix=f"{prefix}.shared_experts", + activation_situ_beta=activation_situ_beta, + activation_situ_linear_beta=activation_situ_linear_beta, + ) + else: + self.shared_experts = None + + latent_quant_config = quant_config if quant_config is not None and quant_config.get_name() == "ascend" else None + if self.use_latent_moe: + self.routed_expert_down_proj = ReplicatedLinear( + hidden_size, + self.moe_hidden_size, + bias=False, + quant_config=latent_quant_config, + prefix=f"{prefix}.routed_expert_down_proj", + ) + self.routed_expert_norm = ( + RMSNorm(self.moe_hidden_size, eps=config.rms_norm_eps) if self.latent_moe_use_norm else None + ) + self.routed_expert_up_proj = ReplicatedLinear( + self.moe_hidden_size, + hidden_size, + bias=False, + quant_config=latent_quant_config, + prefix=f"{prefix}.routed_expert_up_proj", + ) + self.routed_output_transform = KimiRoutedOutputTransform( + self.routed_expert_norm, + self.routed_expert_up_proj, + ) + else: + self.routed_expert_down_proj = None + self.routed_expert_norm = None + self.routed_expert_up_proj = None + self.routed_output_transform = None + + self.experts = FusedMoEFactory( + shared_experts=self.shared_experts, + num_experts=num_experts, + top_k=num_experts_per_token, + hidden_size=self.moe_hidden_size, + intermediate_size=moe_intermediate_size, + activation=config.hidden_act, + activation_situ_beta=activation_situ_beta, + activation_situ_linear_beta=activation_situ_linear_beta, + renormalize=config.moe_renormalize, + quant_config=quant_config, + use_grouped_topk=config.use_grouped_topk, + num_expert_group=config.num_expert_group, + topk_group=config.topk_group, + prefix=f"{prefix}.experts", + scoring_func=config.moe_router_activation_func, + e_score_correction_bias=self.gate.e_score_correction_bias, + routed_scaling_factor=self.routed_scaling_factor, + routed_input_transform=self.routed_expert_down_proj, + routed_output_transform=self.routed_output_transform, + is_sequence_parallel=use_sequence_parallel, + ) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + num_tokens, hidden_size = hidden_states.shape + hidden_states = hidden_states.view(-1, hidden_size) + router_logits, _ = self.gate(hidden_states) + final_hidden_states = self.experts( + hidden_states=hidden_states, + router_logits=router_logits, + ) + return final_hidden_states.view(num_tokens, hidden_size) + + +class AscendKimiMLAAttention(UpstreamKimiMLAAttention): + """Extend vLLM's generic Kimi MLA only for DSpark RoPE metadata.""" + + def __init__( + self, + config, + hidden_size: int, + num_heads: int, + qk_nope_head_dim: int, + qk_rope_head_dim: int, + v_head_dim: int, + q_lora_rank: int | None, + kv_lora_rank: int, + use_output_gate: bool, + use_rope: bool, + cache_config: CacheConfig | None = None, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + non_causal_multi_token_decode: bool = False, + ) -> None: + upstream_config = copy(config) + upstream_config.mla_use_output_gate = use_output_gate + super().__init__( + config=upstream_config, + hidden_size=hidden_size, + num_heads=num_heads, + qk_nope_head_dim=qk_nope_head_dim, + qk_rope_head_dim=qk_rope_head_dim, + v_head_dim=v_head_dim, + q_lora_rank=q_lora_rank, + kv_lora_rank=kv_lora_rank, + use_nope=True, + cache_config=cache_config, + quant_config=quant_config, + prefix=prefix, + ) + if not use_rope and not non_causal_multi_token_decode: + return + + rotary_emb = None + if use_rope: + rope_parameters = dict(config.rope_parameters) + if rope_parameters["rope_type"] != "default": + rope_parameters["rope_type"] = ( + "deepseek_yarn" if rope_parameters.get("apply_yarn_scaling", True) else "deepseek_llama_scaling" + ) + rotary_emb = get_rope( + qk_rope_head_dim, + max_position=config.max_position_embeddings, + rope_parameters=rope_parameters, + is_neox_style=False, + ) + if rope_parameters["rope_type"] == "deepseek_yarn": + scaling_factor = float(rope_parameters["factor"]) + mscale_all_dim = float(rope_parameters.get("mscale_all_dim", 0.0)) + if scaling_factor > 1 and mscale_all_dim: + mscale = 0.1 * mscale_all_dim * math.log(scaling_factor) + 1.0 + self.scaling *= mscale * mscale + + # The upstream Kimi module has already constructed the platform- + # registered MLA wrapper, including all projections and weight loaders. + # Configure that existing Ascend attention layer for DSpark instead of + # constructing and registering a second wrapper with the same prefix. + attention_layer = self._attention_layer + attention_layer.scale = self.scaling + attention_layer.non_causal_multi_token_decode = non_causal_multi_token_decode + attention_layer.impl.scale = float(self.scaling) + attention_layer.impl.rotary_emb = rotary_emb + attention_layer.impl.use_mla_rope = use_rope + + @property + def _attention_layer(self): + return self.mla_attn.mla_attn + + @property + def is_vl_first_layer(self) -> bool: + return self.mla_attn.is_vl_first_layer + + @property + def layer_name(self) -> str: + return self._attention_layer.layer_name + + @property + def impl(self): + return self._attention_layer.impl + + @property + def kv_cache(self): + return self._attention_layer.kv_cache + + @property + def kv_cache_dtype(self): + return self._attention_layer.kv_cache_dtype + + @property + def _k_scale(self): + return self._attention_layer._k_scale + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + ) -> torch.Tensor: + return self.mla_attn(positions, hidden_states) + + +class AscendKimiDecoderLayer(UpstreamKimiDecoderLayer): + """Upstream Kimi decoder structure with Ascend attention backends.""" + + def __init__( + self, + config, + vllm_config: VllmConfig, + prefix: str = "", + use_sequence_parallel: bool = False, + ) -> None: + """Select KDA or no-RoPE MLA and configure the layer residual path.""" + nn.Module.__init__(self) + self.hidden_size = config.hidden_size + self.layer_idx = int(prefix.rsplit(".", 1)[1]) + self.is_moe = config.is_moe + self.use_sequence_parallel = use_sequence_parallel + layer_idx = self.layer_idx + cache_config = vllm_config.cache_config + quant_config = vllm_config.quant_config + + if config.is_kda_layer(layer_idx): + self.self_attn = AscendKimiK3DeltaAttention( + config, + vllm_config, + prefix=f"{prefix}.self_attn", + ) + else: + qk_nope_head_dim = config.qk_nope_head_dim + qk_rope_head_dim = config.qk_rope_head_dim + v_head_dim = config.v_head_dim + kv_lora_rank = config.kv_lora_rank + assert qk_nope_head_dim is not None + assert qk_rope_head_dim is not None + assert v_head_dim is not None + assert kv_lora_rank is not None + assert config.mla_use_nope is True + self.self_attn = AscendKimiMLAAttention( + config=config, + hidden_size=self.hidden_size, + num_heads=config.num_attention_heads, + qk_nope_head_dim=qk_nope_head_dim, + qk_rope_head_dim=qk_rope_head_dim, + v_head_dim=v_head_dim, + q_lora_rank=config.q_lora_rank, + kv_lora_rank=kv_lora_rank, + use_output_gate=bool(config.mla_use_output_gate), + use_rope=False, + cache_config=cache_config, + quant_config=quant_config, + prefix=f"{prefix}.self_attn", + ) + + self.is_moe_layer = ( + self.is_moe + and config.num_experts is not None + and layer_idx >= config.first_k_dense_replace + and layer_idx % config.moe_layer_freq == 0 + ) + if self.is_moe_layer: + self.block_sparse_moe = AscendKimiMoE( + config=config, + quant_config=quant_config, + prefix=f"{prefix}.block_sparse_moe", + use_sequence_parallel=use_sequence_parallel, + ) + self.mlp = self.block_sparse_moe + else: + self.mlp = KimiMLP( + hidden_size=self.hidden_size, + intermediate_size=config.intermediate_size, + hidden_act=config.hidden_act, + quant_config=quant_config, + prefix=f"{prefix}.mlp", + activation_situ_beta=config.activation_situ_beta, + activation_situ_linear_beta=config.activation_situ_linear_beta, + ) + self.input_layernorm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + self.post_attention_layernorm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + + attn_res_block_size = config.attn_res_block_size + self.use_attn_residuals = attn_res_block_size is not None + if attn_res_block_size is not None: + self.attn_res_block_size = attn_res_block_size + self.is_block_write_layer = layer_idx % attn_res_block_size == 0 + self.block_write_idx = layer_idx // attn_res_block_size + self.prev_valid_blocks = cdiv(layer_idx, attn_res_block_size) + self.self_attention_res_norm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + self.mlp_res_norm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + self.self_attention_res_proj = ReplicatedLinear( + config.hidden_size, + 1, + bias=False, + quant_config=None, + prefix=f"{prefix}.self_attention_res_proj", + ) + self.mlp_res_proj = ReplicatedLinear( + config.hidden_size, + 1, + bias=False, + quant_config=None, + prefix=f"{prefix}.mlp_res_proj", + ) + + if self.use_sequence_parallel: + self.self_attn.o_proj.reduce_results = False + + def forward_attn_residual( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + block_residual: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Run Kimi attention residuals with Ascend attention and MoE.""" + prefix_sum: torch.Tensor | None = hidden_states + hidden_states = _apply_ascend_attn_res( + prefix_sum, + block_residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + self.prev_valid_blocks, + ) + if self.is_block_write_layer: + assert prefix_sum is not None + block_residual[:, self.block_write_idx, :].copy_(prefix_sum) + prefix_sum = None + + hidden_states = self.input_layernorm(hidden_states) + if self.use_sequence_parallel: + hidden_states = sp_all_gather(hidden_states) + hidden_states = hidden_states[: positions.shape[0]] + hidden_states = self.self_attn( + hidden_states=hidden_states, + positions=positions, + ) + if self.use_sequence_parallel: + hidden_states = sp_reduce_scatter(hidden_states) + + prefix_sum = hidden_states if prefix_sum is None else prefix_sum + hidden_states + mlp_valid_blocks = self.prev_valid_blocks + (1 if self.is_block_write_layer else 0) + hidden_states = _apply_ascend_attn_res( + prefix_sum, + block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + mlp_valid_blocks, + ) + hidden_states = self.post_attention_layernorm(hidden_states) + hidden_states = self.mlp(hidden_states) + hidden_states = prefix_sum + hidden_states + return hidden_states, block_residual + + +class AscendKimiLinearModel(UpstreamKimiLinearModel): + """Kimi text model assembled from the Ascend decoder layer.""" + + packed_modules_mapping = UpstreamPackedKimiLinearModel.packed_modules_mapping + # Legacy Qwen3 GQA DSpark checkpoints consume the materialized input + # to each selected Kimi layer. MLA DSpark checkpoints consume the raw + # prefix-sum stream used by upstream vLLM, so keep that as the default. + dspark_aux_capture_materialized = False + + def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: + nn.Module.__init__(self) + config = vllm_config.model_config.hf_text_config + self.config = config + self.vocab_size = config.vocab_size + parallel_config = vllm_config.parallel_config + # vLLM's generic MoE SP switch currently requires DP > 1. K3 also + # needs the same rank-local token layout for the TP/EP, DP=1 topology + # that FlashComm used before the standard SP operators were available. + self.use_sequence_parallel = ( + parallel_config.pipeline_parallel_size == 1 + and parallel_config.enable_expert_parallel + and parallel_config.tensor_parallel_size > 1 + ) + + if get_pp_group().is_first_rank: + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + prefix=f"{prefix}.embed_tokens", + ) + else: + self.embed_tokens = PPMissingLayer() + + def get_layer(prefix: str): + return AscendKimiDecoderLayer( + config, + vllm_config, + prefix, + use_sequence_parallel=self.use_sequence_parallel, + ) + + self.start_layer, self.end_layer, self.layers = make_layers( + config.num_hidden_layers, + get_layer, + prefix=f"{prefix}.layers", + ) + + if get_pp_group().is_last_rank: + self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + if config.attn_res_block_size is not None: + self.output_attn_res_norm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + self.output_attn_res_proj = ReplicatedLinear( + config.hidden_size, + 1, + bias=False, + quant_config=None, + prefix=f"{prefix}.output_attn_res_proj", + ) + else: + self.norm = PPMissingLayer() + if config.attn_res_block_size is not None: + self.output_attn_res_norm = PPMissingLayer() + self.output_attn_res_proj = PPMissingLayer() + + world_size = get_tensor_model_parallel_world_size() + assert config.num_attention_heads % world_size == 0, "num_attention_heads must be divisible by world_size" + + def load_weights(self, weights): + """Route mixed-precision KDA gates through vLLM's packed loader.""" + params_dict = dict(self.named_parameters()) + gate_mapping = ( + (".g_proj", ".in_proj_gfab", 0), + (".f_a_proj", ".in_proj_gfab", 1), + (".b_proj", ".in_proj_gfab", 2), + ) + + def remap_mixed_gate_weights(): + for args in weights: + name, loaded_weight = args[:2] + for source, target, shard_id in gate_mapping: + if source not in name: + continue + mapped_name = name.replace(source, target) + if mapped_name in params_dict: + kwargs = dict(args[2]) if len(args) > 2 else {} + kwargs["loaded_shard_id"] = shard_id + yield mapped_name, loaded_weight, kwargs + break + else: + yield args + + return super().load_weights(remap_mixed_gate_weights()) + + def forward( + self, + input_ids: torch.Tensor | None, + positions: torch.Tensor, + intermediate_tensors: IntermediateTensors | None, + inputs_embeds: torch.Tensor | None = None, + **kwargs, + ) -> torch.Tensor | IntermediateTensors | tuple[torch.Tensor, list[torch.Tensor]]: + if self.config.attn_res_block_size is None: + return super().forward( + input_ids=input_ids, + positions=positions, + intermediate_tensors=intermediate_tensors, + inputs_embeds=inputs_embeds, + **kwargs, + ) + + if get_pp_group().is_first_rank: + hidden_states = inputs_embeds if inputs_embeds is not None else self.embed_input_ids(input_ids) + residual = None + else: + assert intermediate_tensors is not None + hidden_states = intermediate_tensors["hidden_states"] + residual = intermediate_tensors["residual"] + + full_num_tokens = positions.shape[0] + if self.use_sequence_parallel: + if envs.VLLM_MOE_SKIP_PADDING and is_forward_context_available(): + forward_context = get_forward_context() + forward_context.is_padding = sp_padding_mask( + forward_context.is_padding, + hidden_states, + ) + hidden_states = sp_shard(hidden_states) + assert residual is None, "Sequence parallelism is not supported with pipeline parallelism" + + if self.dspark_aux_capture_materialized: + aux_hidden_states: list[torch.Tensor] = [] + else: + aux_hidden_states = self._maybe_add_hidden_state( + [], + self.start_layer, + hidden_states, + residual, + ) + attn_res_block_num = cdiv( + self.end_layer, + self.config.attn_res_block_size, + ) + block_residual = hidden_states.new_empty( + hidden_states.size(0), + attn_res_block_num, + hidden_states.size(1), + ) + if residual is not None: + block_residual[:, : residual.size(1), :].copy_(residual) + residual = block_residual + + for layer_idx, layer in enumerate( + self.layers[self.start_layer : self.end_layer], + start=self.start_layer, + ): + if self.dspark_aux_capture_materialized and layer_idx in self.aux_hidden_state_layers: + aux_hidden_states.append( + _apply_ascend_attn_res( + hidden_states, + residual, + layer.self_attention_res_proj, + layer.self_attention_res_norm, + layer.prev_valid_blocks, + ) + ) + hidden_states, residual = layer( + positions=positions, + hidden_states=hidden_states, + residual=residual, + ) + if not self.dspark_aux_capture_materialized and (layer_idx + 1) in self.aux_hidden_state_layers: + self._maybe_add_hidden_state( + aux_hidden_states, + layer_idx + 1, + hidden_states, + residual, + ) + + if not get_pp_group().is_last_rank: + assert not self.use_sequence_parallel, "Sequence parallelism is not supported with pipeline parallelism" + return IntermediateTensors( + { + "hidden_states": hidden_states, + "residual": residual, + } + ) + + hidden_states = _apply_ascend_attn_res( + hidden_states, + residual, + self.output_attn_res_proj, + self.output_attn_res_norm, + attn_res_block_num, + ) + if self.use_sequence_parallel: + if aux_hidden_states: + hidden_size = hidden_states.shape[-1] + packed_hidden_states = torch.cat( + [hidden_states, *aux_hidden_states], + dim=-1, + ) + packed_hidden_states = sp_all_gather(packed_hidden_states) + packed_hidden_states = packed_hidden_states[:full_num_tokens] + hidden_states, *aux_hidden_states = packed_hidden_states.split( + hidden_size, + dim=-1, + ) + else: + hidden_states = sp_all_gather(hidden_states) + hidden_states = hidden_states[:full_num_tokens] + if aux_hidden_states: + return hidden_states, aux_hidden_states + return hidden_states + + +class AscendKimiLinearForCausalLM(UpstreamKimiLinearForCausalLM): + """Causal-LM wrapper retaining vLLM 0.27 state/cache interfaces.""" + + packed_modules_mapping = AscendKimiLinearModel.packed_modules_mapping + + def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: + nn.Module.__init__(self) + self.model_config = vllm_config.model_config + self.vllm_config = vllm_config + self.config = self.model_config.hf_config + self.quant_config = vllm_config.quant_config + self.model = AscendKimiLinearModel( + vllm_config=vllm_config, + prefix=maybe_prefix(prefix, "model"), + ) + if get_pp_group().is_last_rank: + self.lm_head = ParallelLMHead( + self.config.vocab_size, + self.config.hidden_size, + quant_config=self.quant_config, + prefix=maybe_prefix(prefix, "lm_head"), + ) + else: + self.lm_head = PPMissingLayer() + self.logits_processor = LogitsProcessor( + self.config.vocab_size, + scale=getattr(self.config, "logit_scale", 1.0), + ) + + def set_dspark_aux_capture_materialized(self, enabled: bool) -> None: + self.model.dspark_aux_capture_materialized = enabled + + +class AscendKimiK3MultiModalProjector(KimiK25MultiModalProjector): + """Kimi projector with the optional ModelSlim output rotation.""" + + def __init__( + self, + config, + *args, + prefix: str = "", + enable_rotation: bool = False, + **kwargs, + ) -> None: + super().__init__(config, *args, prefix=prefix, **kwargs) + self.rot_proj: ReplicatedLinear | None = None + if enable_rotation: + output_size = config.text_hidden_size + self.rot_proj = ReplicatedLinear( + output_size, + output_size, + bias=False, + quant_config=None, + prefix=f"{prefix}.rot_proj", + ) + + def forward(self, image_features: torch.Tensor) -> torch.Tensor: + hidden_states = super().forward(image_features) + rot_proj = self.rot_proj + if rot_proj is not None: + hidden_states = rot_proj(hidden_states)[0] + return hidden_states + + +@MULTIMODAL_REGISTRY.register_processor( + KimiK3MultiModalProcessor, + info=KimiK3ProcessingInfo, + dummy_inputs=KimiK3DummyInputsBuilder, +) +class AscendKimiK3ForConditionalGeneration(UpstreamKimiK3ForConditionalGeneration): + """Upstream Kimi K3 multimodal wrapper with Ascend text/projector layers.""" + + def __init__(self, vllm_config: VllmConfig, prefix: str = "") -> None: + nn.Module.__init__(self) + model_config = vllm_config.model_config + self.config = model_config.hf_config + self.quant_config = vllm_config.quant_config + multimodal_config = model_config.multimodal_config + assert multimodal_config is not None + + self.use_data_parallel = is_vit_use_data_parallel( + self.config.vision_config.num_attention_heads, + ) + self.hidden_size = self.config.text_config.hidden_size + self.device = current_platform.current_device() + vision_quant_config = self._maybe_ignore_quant_config(self.quant_config) + + with self._mark_tower_model(vllm_config, "image"): + self.vision_tower = MoonViT3dPretrainedModel( + self.config.vision_config, + quant_config=vision_quant_config, + prefix=maybe_prefix(prefix, "vision_tower"), + ) + if vision_quant_config is not None: + self.vision_tower = self.vision_tower.to(device=self.device) + else: + self.vision_tower = self.vision_tower.to( + device=self.device, + dtype=model_config.dtype, + ) + + self.mm_projector = AscendKimiK3MultiModalProjector( + self.config.vision_config, + use_data_parallel=self.use_data_parallel, + quant_config=vision_quant_config, + prefix=maybe_prefix(prefix, "mm_projector"), + enable_rotation=get_rotation_path(vllm_config) is not None, + ) + if vision_quant_config is not None: + self.mm_projector = self.mm_projector.to(device=self.device) + else: + self.mm_projector = self.mm_projector.to( + device=self.device, + dtype=model_config.dtype, + ) + + with self._mark_language_model(vllm_config): + self.language_model = init_vllm_registered_model( + vllm_config=vllm_config, + hf_config=self.config.text_config, + prefix=maybe_prefix(prefix, "language_model"), + architectures=["KimiLinearForCausalLM"], + ) + self.make_empty_intermediate_tensors = ( # type: ignore[method-assign] + self.language_model.make_empty_intermediate_tensors + ) + self.media_placeholder = self.config.media_placeholder_token_id + + def set_dspark_aux_capture_materialized(self, enabled: bool) -> None: + self.language_model.set_dspark_aux_capture_materialized(enabled) diff --git a/vllm_ascend/models/kimi_k3_dspark.py b/vllm_ascend/models/kimi_k3_dspark.py new file mode 100644 index 000000000000..5be9bd920d4c --- /dev/null +++ b/vllm_ascend/models/kimi_k3_dspark.py @@ -0,0 +1,320 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Kimi K3 MLA DSpark draft model for Ascend.""" + +from collections.abc import Iterable + +import torch +from torch import nn +from vllm.config import VllmConfig +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.linear import ReplicatedLinear +from vllm.model_executor.layers.logits_processor import LogitsProcessor +from vllm.model_executor.models.interfaces import MultiModalEmbeddings +from vllm.model_executor.models.qwen3_dspark import DSparkMarkovHead +from vllm.model_executor.models.utils import ( + AutoWeightsLoader, + _merge_multimodal_embeddings, + get_draft_quant_config, + maybe_prefix, +) +from vllm.models.kimi_k3.amd.linear import KimiMLP +from vllm.models.kimi_k3.nvidia.dspark_mla import ( + K3DSparkDecoderLayer as UpstreamK3DSparkDecoderLayer, +) +from vllm.models.kimi_k3.nvidia.dspark_mla import ( + K3DSparkForCausalLM as UpstreamK3DSparkForCausalLM, +) +from vllm.models.kimi_k3.nvidia.dspark_mla import ( + K3DSparkModel as UpstreamK3DSparkModel, +) + +from vllm_ascend.models.kimi_k3 import ( + AscendKimiMLAAttention, +) +from vllm_ascend.models.llama_eagle3 import ( + get_rotation_matrix, + get_rotation_path, + prepare_quarot_shared_layer, +) +from vllm_ascend.models.qwen3_dspark import process_weight +from vllm_ascend.ops.rotary_embedding import get_cos_and_sin_mla + + +class AscendK3DSparkDecoderLayer(UpstreamK3DSparkDecoderLayer): + def __init__( + self, + *, + vllm_config: VllmConfig, + config, + layer_idx: int, + start_layer_id: int, + prefix: str, + ) -> None: + # The upstream constructor hard-codes NVIDIA attention and MLP + # components. Keep its class contract while constructing the Ascend + # equivalents below. + nn.Module.__init__(self) + quant_config = get_draft_quant_config(vllm_config) + layer_prefix = maybe_prefix( + prefix, + f"layers.{start_layer_id + layer_idx}", + ) + self.self_attn = AscendKimiMLAAttention( + config=config, + hidden_size=config.hidden_size, + num_heads=config.num_attention_heads, + qk_nope_head_dim=config.qk_nope_head_dim, + qk_rope_head_dim=config.qk_rope_head_dim, + v_head_dim=config.v_head_dim, + q_lora_rank=config.q_lora_rank, + kv_lora_rank=config.kv_lora_rank, + use_output_gate=False, + use_rope=True, + cache_config=vllm_config.cache_config, + quant_config=quant_config, + prefix=f"{layer_prefix}.self_attn", + non_causal_multi_token_decode=True, + ) + self.mlp = KimiMLP( + hidden_size=config.hidden_size, + intermediate_size=config.intermediate_size, + hidden_act=config.hidden_act, + quant_config=quant_config, + prefix=f"{layer_prefix}.mlp", + ) + self.input_layernorm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + self.post_attention_layernorm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + residual: torch.Tensor | None, + ) -> tuple[torch.Tensor, torch.Tensor]: + if residual is None: + residual = hidden_states + hidden_states = self.input_layernorm(hidden_states) + else: + hidden_states, residual = self.input_layernorm( + hidden_states, + residual, + ) + hidden_states = self.self_attn( + positions=positions, + hidden_states=hidden_states, + ) + hidden_states, residual = self.post_attention_layernorm( + hidden_states, + residual, + ) + hidden_states = self.mlp(hidden_states) + return hidden_states, residual + + +class AscendK3DSparkModel(UpstreamK3DSparkModel): + def __init__( + self, + *, + vllm_config: VllmConfig, + start_layer_id: int, + prefix: str, + ) -> None: + # The upstream constructor hard-codes CUDA/NVIDIA attention and decoder + # classes. Initialize only the module base, then keep the upstream + # model methods while constructing the Ascend-specific components. + nn.Module.__init__(self) + assert vllm_config.speculative_config is not None + draft_model_config = vllm_config.speculative_config.draft_model_config + assert draft_model_config is not None + self.config = draft_model_config.hf_config + self.quant_config = get_draft_quant_config(vllm_config) + self.embed_tokens: nn.Module | None = None + + self.context_proj = ReplicatedLinear( + self.config.target_hidden_size * self.config.num_target_layers, + self.config.hidden_size, + bias=False, + return_bias=False, + quant_config=self.quant_config, + prefix=maybe_prefix(prefix, "context_proj"), + ) + self.context_norm = RMSNorm( + self.config.hidden_size, + eps=self.config.rms_norm_eps, + ) + self.layers = nn.ModuleList( + [ + AscendK3DSparkDecoderLayer( + vllm_config=vllm_config, + config=self.config, + layer_idx=layer_idx, + start_layer_id=start_layer_id, + prefix=prefix, + ) + for layer_idx in range(self.config.num_hidden_layers) + ] + ) + self.final_norm = RMSNorm( + self.config.hidden_size, + eps=self.config.rms_norm_eps, + ) + self.markov_head = DSparkMarkovHead( + self.config.vocab_size, + self.config.draft_vocab_size, + self.config.markov_rank, + prefix=maybe_prefix(prefix, "markov_head"), + ) + + @torch.inference_mode() + def precompute_and_store_context_kv( + self, + context_states: torch.Tensor, + context_positions: torch.Tensor, + context_slot_mapping: ( + torch.Tensor | list[torch.Tensor | None] | tuple[torch.Tensor | None, ...] | None + ) = None, + ) -> None: + if context_slot_mapping is None or context_states.numel() == 0: + return + per_layer_slot_mapping = isinstance(context_slot_mapping, (list, tuple)) + cos, sin = get_cos_and_sin_mla(context_positions) + for layer_idx, layer in enumerate(self.layers): + attn = layer.self_attn + assert attn.fused_qkv_a_proj is not None + assert attn.q_lora_rank is not None + qkv_lora = attn.fused_qkv_a_proj(context_states)[0] + kv_no_split = qkv_lora[..., attn.q_lora_rank :].contiguous() + slots = context_slot_mapping[layer_idx] if per_layer_slot_mapping else context_slot_mapping + if slots is None: + continue + attn.impl.exec_kv_prefill( + kv_no_split, + cos, + sin, + attn.kv_cache, + slots, + ) + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + inputs_embeds: torch.Tensor | None = None, + ) -> torch.Tensor: + if inputs_embeds is None: + inputs_embeds = self.embed_input_ids(input_ids) + hidden_states = inputs_embeds + residual = None + for layer in self.layers: + hidden_states, residual = layer( + positions=positions, + hidden_states=hidden_states, + residual=residual, + ) + hidden_states, _ = self.final_norm(hidden_states, residual) + return hidden_states + + +class AscendK3DSparkForCausalLM(UpstreamK3DSparkForCausalLM): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: + nn.Module.__init__(self) + assert vllm_config.speculative_config is not None + self.draft_model_config = vllm_config.speculative_config.draft_model_config + assert self.draft_model_config is not None + self.config = self.draft_model_config.hf_config + target_layer_num = vllm_config.model_config.get_num_layers(vllm_config.parallel_config) + self.model = AscendK3DSparkModel( + vllm_config=vllm_config, + start_layer_id=target_layer_num, + prefix=maybe_prefix(prefix, "model"), + ) + self.lm_head: nn.Module | None = None + self.logits_processor = LogitsProcessor( + self.config.draft_vocab_size, + scale=getattr(self.config, "logit_scale", 1.0), + ) + self.rotation_path = get_rotation_path(vllm_config) + self._shared_layer_rotation: torch.Tensor | None = None + + def prepare_shared_layer( + self, + draft_layer: nn.Module | None, + target_layer: nn.Module, + label: str, + ) -> nn.Module | None: + rotation = self._shared_layer_rotation + if rotation is None: + return None + draft_layer, self._shared_layer_rotation = prepare_quarot_shared_layer( + draft_layer, + target_layer, + rotation, + label, + ) + return draft_layer + + def finish_shared_layer_preparation(self) -> None: + self._shared_layer_rotation = None + + def load_weights( + self, + weights: Iterable[tuple[str, torch.Tensor]], + ) -> set[str]: + """Load the per-layer KV projections used by the Ascend draft model. + + Upstream additionally duplicates these weights into a CUDA-specific + cross-layer ``context_kv_proj``. Ascend deliberately retains the + quantization-aware per-layer projections, so use vLLM's public loader + interface without creating that extra packed parameter. + """ + loader = AutoWeightsLoader(self) + self._shared_layer_rotation = None + if self.rotation_path is not None: + rotation_weight = get_rotation_matrix(self.rotation_path) + self._shared_layer_rotation = rotation_weight + weights = ( + ( + name, + process_weight(loaded_weight, rotation_weight) if "context_proj." in name else loaded_weight, + ) + for name, loaded_weight in weights + ) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) + + def embed_input_ids( + self, + input_ids: torch.Tensor, + multimodal_embeddings: MultiModalEmbeddings | None = None, + *, + is_multimodal: torch.Tensor | None = None, + ) -> torch.Tensor: + """Embed draft tokens and replace multimodal placeholder positions. + + vLLM 0.27 passes the target model's precomputed multimodal embeddings + through the speculative proposer. K3 DSpark shares the target token + embedding but upstream still exposes the older text-only method + signature, so adapt that interface without duplicating the vision + tower in the draft model. + """ + if multimodal_embeddings is None or len(multimodal_embeddings) == 0 or is_multimodal is None: + return self.model.embed_input_ids(input_ids) + + # Placeholder ids are overwritten below. Mask them before the shared + # vocabulary lookup so out-of-vocabulary multimodal ids are safe too. + text_input_ids = input_ids.masked_fill( + is_multimodal.to(device=input_ids.device, non_blocking=True), + 0, + ) + inputs_embeds = self.model.embed_input_ids(text_input_ids) + return _merge_multimodal_embeddings( + inputs_embeds=inputs_embeds, + multimodal_embeddings=multimodal_embeddings, + is_multimodal=is_multimodal, + ) diff --git a/vllm_ascend/models/kimi_k3_mtp.py b/vllm_ascend/models/kimi_k3_mtp.py new file mode 100644 index 000000000000..d56e8a124105 --- /dev/null +++ b/vllm_ascend/models/kimi_k3_mtp.py @@ -0,0 +1,93 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Kimi K3 MTP draft model for Ascend.""" + +import copy + +from torch import nn +from vllm.config import VllmConfig +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.logits_processor import LogitsProcessor +from vllm.model_executor.layers.vocab_parallel_embedding import VocabParallelEmbedding +from vllm.model_executor.models.utils import maybe_prefix +from vllm.models.kimi_k3.amd.mtp import ( + KimiK3MTP as UpstreamKimiK3MTP, +) +from vllm.models.kimi_k3.amd.mtp import ( + KimiK3MultiTokenPredictor as UpstreamKimiK3MultiTokenPredictor, +) +from vllm.models.kimi_k3.amd.mtp import ( + KimiK3MultiTokenPredictorLayer as UpstreamKimiK3MultiTokenPredictorLayer, +) +from vllm.models.kimi_k3.amd.mtp import SharedHead + +from vllm_ascend.models.kimi_k3 import AscendKimiDecoderLayer + + +class AscendKimiK3MultiTokenPredictorLayer( + UpstreamKimiK3MultiTokenPredictorLayer, +): + def __init__(self, config, vllm_config: VllmConfig, prefix: str) -> None: + # The upstream constructor hard-codes the AMD decoder layer. Build the + # same container with the Ascend decoder and inherit its forward path. + nn.Module.__init__(self) + self.config = config + self.enorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.hnorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.eh_proj = nn.Linear( + config.hidden_size * 2, + config.hidden_size, + bias=False, + ) + self.shared_head = SharedHead( + config=config, + prefix=prefix, + quant_config=vllm_config.quant_config, + ) + block_config = copy.copy(config) + block_config.attn_res_block_size = None + self.mtp_block = AscendKimiDecoderLayer( + block_config, + vllm_config, + prefix=prefix, + ) + + +class AscendKimiK3MultiTokenPredictor(UpstreamKimiK3MultiTokenPredictor): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: + # The upstream constructor hard-codes its predictor-layer class. + nn.Module.__init__(self) + config = vllm_config.model_config.hf_text_config + self.config = config + self.mtp_start_layer_idx = config.num_hidden_layers + self.num_mtp_layers = config.num_nextn_predict_layers + self.layers = nn.ModuleDict( + { + str(idx): AscendKimiK3MultiTokenPredictorLayer( + config, + vllm_config, + f"{prefix}.layers.{idx}", + ) + for idx in range( + self.mtp_start_layer_idx, + self.mtp_start_layer_idx + self.num_mtp_layers, + ) + } + ) + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + prefix=maybe_prefix(prefix, "embed_tokens"), + ) + self.logits_processor = LogitsProcessor(config.vocab_size) + + +class AscendKimiK3MTP(UpstreamKimiK3MTP): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: + nn.Module.__init__(self) + self.config = vllm_config.model_config.hf_text_config + self.quant_config = vllm_config.quant_config + self.model = AscendKimiK3MultiTokenPredictor( + vllm_config=vllm_config, + prefix=maybe_prefix(prefix, "model"), + ) diff --git a/vllm_ascend/models/llama_eagle3.py b/vllm_ascend/models/llama_eagle3.py index e938978a2c15..3f56e1e3965d 100644 --- a/vllm_ascend/models/llama_eagle3.py +++ b/vllm_ascend/models/llama_eagle3.py @@ -1,3 +1,4 @@ +import copy import logging import os from collections.abc import Iterable @@ -5,6 +6,7 @@ import torch from safetensors.torch import load_file +from torch import nn from vllm.config import VllmConfig from vllm.model_executor.models.llama_eagle3 import Eagle3LlamaForCausalLM @@ -52,6 +54,36 @@ def get_rotation_matrix(rotation_path: Path | None) -> torch.Tensor: raise e +@torch.inference_mode() +def prepare_quarot_shared_layer( + draft_layer: nn.Module | None, + target_layer: nn.Module, + rotation: torch.Tensor, + label: str, +) -> tuple[nn.Module, torch.Tensor]: + """Create a draft-owned target layer in the unrotated hidden basis.""" + if draft_layer is None: + comm_group = getattr(target_layer, "comm_group", None) + memo = {id(comm_group): comm_group} if comm_group is not None else None + draft_layer = copy.deepcopy(target_layer, memo) + + rotation = rotation.to( + device=target_layer.weight.device, + dtype=torch.float32, + ) + unrotated = torch.matmul( + target_layer.weight.data.to(torch.float32), + rotation.T, + ) + draft_layer.weight.data.copy_(unrotated.to(draft_layer.weight.dtype)) + logger.info( + "[spec_decode/quarot] Copied and aligned shared %s (weight=%s).", + label, + tuple(draft_layer.weight.shape), + ) + return draft_layer, rotation + + def compute_rotation_matrix3(Q: torch.Tensor) -> torch.Tensor: """Anti-rotate matrix for 3 layers of hidden_states.""" return torch.block_diag(Q, Q, Q) diff --git a/vllm_ascend/models/qwen3_dspark.py b/vllm_ascend/models/qwen3_dspark.py index 0d7587fc05a1..ef50bd196f89 100644 --- a/vllm_ascend/models/qwen3_dspark.py +++ b/vllm_ascend/models/qwen3_dspark.py @@ -7,7 +7,11 @@ from vllm.model_executor.models.qwen3_dspark import Qwen3DSparkForCausalLM from vllm.model_executor.models.utils import AutoWeightsLoader, maybe_prefix -from vllm_ascend.models.llama_eagle3 import get_rotation_matrix, get_rotation_path +from vllm_ascend.models.llama_eagle3 import ( + get_rotation_matrix, + get_rotation_path, + prepare_quarot_shared_layer, +) from vllm_ascend.utils import vllm_version_is @@ -67,6 +71,27 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: prefix=maybe_prefix(model_prefix, "confidence_head"), ) self.rotation_path = get_rotation_path(vllm_config) if vllm_config.quant_config is not None else None + self._shared_layer_rotation: torch.Tensor | None = None + + def prepare_shared_layer( + self, + draft_layer: nn.Module | None, + target_layer: nn.Module, + label: str, + ) -> nn.Module | None: + rotation = self._shared_layer_rotation + if rotation is None: + return None + draft_layer, self._shared_layer_rotation = prepare_quarot_shared_layer( + draft_layer, + target_layer, + rotation, + label, + ) + return draft_layer + + def finish_shared_layer_preparation(self) -> None: + self._shared_layer_rotation = None @staticmethod def _get_confidence_relative_name( @@ -89,9 +114,11 @@ def compute_confidence(self, head_hidden: torch.Tensor, markov_embed: torch.Tens def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): all_weights = list(weights) + self._shared_layer_rotation = None if self.rotation_path is not None: processed_weights: list[tuple[str, torch.Tensor]] = [] rotation_weight = get_rotation_matrix(self.rotation_path) + self._shared_layer_rotation = rotation_weight for name, loaded_weight in all_weights: if "fc." in name: loaded_weight = process_weight(loaded_weight, rotation_weight) diff --git a/vllm_ascend/ops/mm_encoder_attention.py b/vllm_ascend/ops/mm_encoder_attention.py index a5bbd0735ef5..498a08e57f70 100644 --- a/vllm_ascend/ops/mm_encoder_attention.py +++ b/vllm_ascend/ops/mm_encoder_attention.py @@ -166,8 +166,8 @@ def _run_vit_fia( ) -> torch.Tensor: fia_kwargs = dict( query=query, - key=key, - value=value, + key=key.contiguous(), + value=value.contiguous(), atten_mask=None, block_table=None, input_layout="TND", diff --git a/vllm_ascend/worker/model_runner_v1.py b/vllm_ascend/worker/model_runner_v1.py index a0e3d81e61c5..c7d29b770e27 100644 --- a/vllm_ascend/worker/model_runner_v1.py +++ b/vllm_ascend/worker/model_runner_v1.py @@ -665,6 +665,17 @@ def _get_eagle3_aux_layers_from_config(self) -> tuple[int, ...] | None: return tuple(i + 1 for i in dspark_layer_ids) return None + def _draft_uses_qwen3_gqa_dspark(self) -> bool: + """Return whether the draft expects Kimi's materialized residual.""" + if self.speculative_config is None or not self.speculative_config.use_dspark(): + return False + draft_model_config = self.speculative_config.draft_model_config + if draft_model_config is None: + return False + hf_config = draft_model_config.hf_config + architectures = getattr(hf_config, "architectures", ()) or () + return getattr(hf_config, "model_type", None) == "qwen3" and "Qwen3DSparkModel" in architectures + def _use_aclgraph(self) -> bool: return ( self.compilation_config.cudagraph_mode != CUDAGraphMode.NONE @@ -3550,6 +3561,19 @@ def mock_pass(param1, param2): if not aux_layers: aux_layers = self.model.get_eagle3_default_aux_hidden_state_layers() self.model.set_aux_hidden_state_layers(aux_layers) + if self.speculative_config.use_dspark(): + set_capture_mode = getattr( + self.model, + "set_dspark_aux_capture_materialized", + None, + ) + if set_capture_mode is not None: + materialized = self._draft_uses_qwen3_gqa_dspark() + set_capture_mode(materialized) + logger.info( + "Kimi K3 DSpark auxiliary capture uses %s stream.", + "materialized GQA" if materialized else "raw MLA", + ) if pp_group.world_size > 1: inner_model = self.model From c2c3af577f6bb8eeaf3dd5873bd33bd0edbe023a Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Mon, 24 Aug 2026 04:28:32 -0500 Subject: [PATCH 12/50] feat(spec-decode): add Kimi K3 DSpark execution Add the Kimi K3 DSpark proposer, model-specific speculative configuration, and proposer lifecycle handling for Ascend. Cover token preparation, proposal verification, and rotation-aware execution with focused unit tests. Signed-off-by: maoxx241 --- .../test_patch_speculative_config_dspark.py | 23 ++ tests/ut/spec_decode/test_dspark_proposer.py | 239 ++++++++++-------- .../ut/spec_decode/test_llm_base_proposer.py | 72 +++++- tests/ut/spec_decode/test_utils.py | 28 ++ .../platform/patch_speculative_config.py | 18 ++ vllm_ascend/spec_decode/dspark_proposer.py | 91 ++++--- vllm_ascend/spec_decode/llm_base_proposer.py | 102 ++++++-- 7 files changed, 396 insertions(+), 177 deletions(-) create mode 100644 tests/ut/patch/platform/test_patch_speculative_config_dspark.py diff --git a/tests/ut/patch/platform/test_patch_speculative_config_dspark.py b/tests/ut/patch/platform/test_patch_speculative_config_dspark.py new file mode 100644 index 000000000000..aabc1d694e3f --- /dev/null +++ b/tests/ut/patch/platform/test_patch_speculative_config_dspark.py @@ -0,0 +1,23 @@ +from transformers import Qwen3Config +from vllm.config.speculative import SpeculativeConfig + +import vllm_ascend.patch.platform.patch_speculative_config # noqa: F401 + + +def test_legacy_qwen3_dspark_config_uses_qwen3_loader(): + config = Qwen3Config( + architectures=["DSparkDraftModel"], + block_size=7, + dflash_config={ + "mask_token_id": 163824, + "target_layer_ids": [7, 23, 51, 67, 83], + }, + ) + + normalized = SpeculativeConfig.hf_config_override(config) + + assert normalized is config + assert normalized.architectures == ["Qwen3DSparkModel"] + assert normalized.mask_token_id == 163824 + assert normalized.target_layer_ids == [7, 23, 51, 67, 83] + assert normalized.block_size == 7 diff --git a/tests/ut/spec_decode/test_dspark_proposer.py b/tests/ut/spec_decode/test_dspark_proposer.py index ac486c9aebc9..73857de027f3 100644 --- a/tests/ut/spec_decode/test_dspark_proposer.py +++ b/tests/ut/spec_decode/test_dspark_proposer.py @@ -26,6 +26,7 @@ import numpy as np import pytest import torch +from vllm.v1.worker.utils import AttentionGroup from vllm_ascend.attention.attention_v1 import AscendAttentionState from vllm_ascend.spec_decode.dflash_proposer import AscendDflashProposer @@ -62,7 +63,10 @@ def _make_vllm_config(hf_config: SimpleNamespace) -> SimpleNamespace: """Build the minimal config consumed by the DSpark initializer.""" draft_model_config = SimpleNamespace(hf_config=hf_config, get_hidden_size=lambda: _HIDDEN_SIZE) return SimpleNamespace( - speculative_config=SimpleNamespace(draft_sample_method="greedy", draft_model_config=draft_model_config) + speculative_config=SimpleNamespace( + draft_sample_method="greedy", + draft_model_config=draft_model_config, + ) ) @classmethod @@ -100,7 +104,16 @@ def mock_parent_init( else SimpleNamespace() ) - with patch.object(AscendDSparkProposer.__base__, "__init__", mock_parent_init): + dynamic_spec_config = SimpleNamespace(method="", method_params={}) + with ( + patch.object(AscendDSparkProposer.__base__, "__init__", mock_parent_init), + patch( + "vllm_ascend.spec_decode.dspark_proposer.get_ascend_config", + return_value=SimpleNamespace( + dynamic_spec_config=dynamic_spec_config, + ), + ), + ): proposer = AscendDSparkProposer(vllm_config, device) num_query_total = num_reqs * proposer.num_query_per_req proposer.positions = torch.zeros(max_num_tokens, dtype=torch.int32, device=device) @@ -131,11 +144,11 @@ def mock_parent_init( proposer._per_group_block_table_buffers = {gid: block_table} slot = torch.zeros(max_num_tokens, dtype=torch.int32, device=device) proposer._per_group_slot_mappings = {gid: slot} + proposer._per_group_kernel_block_sizes = {gid: block_size} proposer._per_group_query_slot_mapping_buffers = {gid: slot.clone()} proposer._per_group_context_slot_mapping_buffers = {gid: slot.clone()} return proposer - # fmt: off @staticmethod def _invoke_set_inputs_first_pass( proposer, @@ -143,6 +156,8 @@ def _invoke_set_inputs_first_pass( num_reqs, block_size, seq_len=128, + host_seq_len=None, + async_metadata=False, context=None, num_rejected=None, with_optional_attrs=False, @@ -155,17 +170,20 @@ def _invoke_set_inputs_first_pass( next_token_ids, target_hidden_states)``. """ next_token_ids = torch.arange(1, num_reqs + 1, dtype=torch.int64) - target_hidden_states = torch.arange( - num_reqs * 8, dtype=torch.float32 - ).reshape(num_reqs, 8) + target_hidden_states = torch.arange(num_reqs * 8, dtype=torch.float32).reshape(num_reqs, 8) query_start_loc_cpu = torch.zeros(num_reqs + 1, dtype=torch.int32) if context is not None: query_start_loc_cpu[num_reqs] = context + if host_seq_len is None: + host_seq_len = seq_len + seq_lens_cpu = torch.full((num_reqs,), host_seq_len, dtype=torch.int32) cad = SimpleNamespace( num_reqs=num_reqs, query_start_loc=torch.arange(num_reqs + 1, dtype=torch.int32) * block_size, query_start_loc_cpu=query_start_loc_cpu, seq_lens=torch.full((num_reqs,), seq_len, dtype=torch.int32), + _seq_lens_cpu=seq_lens_cpu, + seq_lens_cpu=None if async_metadata else seq_lens_cpu, max_seq_len=seq_len, ) if with_optional_attrs: @@ -183,9 +201,6 @@ def _invoke_set_inputs_first_pass( return num_query_total, token_indices, cad, extra, next_token_ids, target_hidden_states -# fmt: on - - class TestDSparkPositionsFullUnderMultiDp(_DSparkProposerTestBase): """Guard: under multi-DP the dspark draft proposer must hand DSA attention a full-length positions buffer so ``positions[:num_input_tokens]`` never reads @@ -368,7 +383,6 @@ def test_configures_anchor_sampling( assert proposer.max_query_tokens == expected_max_query_tokens -# fmt: off class TestSetPerGroupAttnMetadata(_DSparkProposerTestBase): """``set_per_group_attn_metadata`` stores the runner-provided per-group block table / slot mapping into the read-only dicts the proposer consults @@ -376,9 +390,7 @@ class TestSetPerGroupAttnMetadata(_DSparkProposerTestBase): def test_stores_block_table_and_slot_mapping(self): num_reqs, block_size, max_num_tokens = 4, 5, 256 - proposer = self._make_proposer( - max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size - ) + proposer = self._make_proposer(max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size) # a gid not pre-populated by _make_proposer (which only seeds gid=0) gid = 7 block_table = torch.zeros((num_reqs, 16), dtype=torch.int32) @@ -391,9 +403,7 @@ def test_stores_block_table_and_slot_mapping(self): def test_overwrites_existing_gid(self): num_reqs, block_size, max_num_tokens = 2, 5, 256 - proposer = self._make_proposer( - max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size - ) + proposer = self._make_proposer(max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size) gid = 0 # already populated by _make_proposer old_block_table = proposer._per_group_block_tables[gid] new_block_table = torch.ones((num_reqs, 16), dtype=torch.int32) @@ -424,10 +434,7 @@ def _make_vllm_config( speculative_config = SimpleNamespace( num_speculative_tokens=num_speculative_tokens, draft_sample_method=draft_sample_method, - draft_model_config=SimpleNamespace( - hf_config=SimpleNamespace(), - get_hidden_size=lambda: hidden_size - ), + draft_model_config=SimpleNamespace(hf_config=SimpleNamespace(), get_hidden_size=lambda: hidden_size), ) return SimpleNamespace(speculative_config=speculative_config) @@ -451,7 +458,6 @@ def _stub(self, vllm_config, device, runner=None): self.dtype = dtype self.device = device self.draft_model_config = vllm_config.speculative_config.draft_model_config - # present so the ``del`` in DSpark.__init__ succeeds self.hidden_size = 0 self.hidden_states = None self._dflash_hidden_states = None @@ -498,7 +504,14 @@ def test_greedy_allocates_dspark_buffers(self, monkeypatch): draft_sample_method="greedy", hidden_size=hidden, ) - proposer = AscendDSparkProposer(vllm_config, device) + dynamic_spec_config = SimpleNamespace(method="", method_params={}) + with patch( + "vllm_ascend.spec_decode.dspark_proposer.get_ascend_config", + return_value=SimpleNamespace( + dynamic_spec_config=dynamic_spec_config, + ), + ): + proposer = AscendDSparkProposer(vllm_config, device) blk = 1 + num_spec max_query_tokens = max_batch * num_spec @@ -533,21 +546,16 @@ class TestSetInputsFirstPassOutputs(_DSparkProposerTestBase): @pytest.fixture(autouse=True) def _mock_kernel(self, monkeypatch): monkeypatch.setattr( - "vllm_ascend.spec_decode.dspark_proposer." - "copy_and_expand_dflash_and_dspark_inputs_kernel", + "vllm_ascend.spec_decode.dspark_proposer.copy_and_expand_dflash_and_dspark_inputs_kernel", MagicMock(), ) def test_return_value_and_token_indices(self): num_reqs, block_size, max_num_tokens = 4, 5, 256 - proposer = self._make_proposer( - max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size - ) - num_query_total, token_indices, _cad, extra = ( - self._invoke_set_inputs_first_pass( - proposer, num_reqs=num_reqs, block_size=block_size - )[:4] - ) + proposer = self._make_proposer(max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size) + num_query_total, token_indices, _cad, extra = self._invoke_set_inputs_first_pass( + proposer, num_reqs=num_reqs, block_size=block_size + )[:4] assert num_query_total == num_reqs * block_size assert token_indices.shape == (num_reqs * block_size,) assert token_indices.dtype == torch.int32 @@ -556,33 +564,49 @@ def test_return_value_and_token_indices(self): def test_seed_buffer_copied_from_next_tokens(self): num_reqs, block_size, max_num_tokens = 4, 5, 256 - proposer = self._make_proposer( - max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size - ) - self._invoke_set_inputs_first_pass( - proposer, num_reqs=num_reqs, block_size=block_size - ) + proposer = self._make_proposer(max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size) + self._invoke_set_inputs_first_pass(proposer, num_reqs=num_reqs, block_size=block_size) expected = torch.arange(1, num_reqs + 1, dtype=torch.int64) assert torch.equal(proposer._dspark_seed_buffer[:num_reqs], expected) assert torch.all(proposer._dspark_seed_buffer[num_reqs:] == 0) def test_context_hidden_states_copied(self): num_reqs, block_size, max_num_tokens = 4, 5, 256 + proposer = self._make_proposer(max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size) + self._invoke_set_inputs_first_pass(proposer, num_reqs=num_reqs, block_size=block_size, context=num_reqs) + assert proposer._dflash_num_context == num_reqs + expected = torch.arange(num_reqs * 8, dtype=torch.float32).reshape(num_reqs, 8) + assert torch.equal(proposer._dflash_hidden_states[:num_reqs], expected) + + def test_query_slot_kernel_uses_logical_block_size(self, monkeypatch): + kernel = MagicMock() + monkeypatch.setattr( + "vllm_ascend.spec_decode.dspark_proposer.copy_and_expand_dflash_and_dspark_inputs_kernel", + kernel, + ) + num_reqs, num_speculative_tokens, max_num_tokens = 1, 7, 32 proposer = self._make_proposer( - max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size + max_num_tokens=max_num_tokens, + num_reqs=num_reqs, + block_size=num_speculative_tokens, ) + proposer.draft_attn_groups[0].kv_cache_spec.block_size = 384 + proposer._per_group_kernel_block_sizes[0] = 128 + self._invoke_set_inputs_first_pass( - proposer, num_reqs=num_reqs, block_size=block_size, context=num_reqs + proposer, + num_reqs=num_reqs, + block_size=num_speculative_tokens, + seq_len=720, ) - assert proposer._dflash_num_context == num_reqs - expected = torch.arange(num_reqs * 8, dtype=torch.float32).reshape(num_reqs, 8) - assert torch.equal(proposer._dflash_hidden_states[:num_reqs], expected) + + kwargs = kernel[1,].call_args.kwargs + assert proposer.draft_attn_groups[0].kv_cache_spec.block_size == 384 + assert kwargs["block_size"] == 128 def test_cad_rewritten_to_cross_attention_shape(self): num_reqs, block_size, max_num_tokens = 4, 5, 256 - proposer = self._make_proposer( - max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size - ) + proposer = self._make_proposer(max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size) num_query_total, _, cad, _ = self._invoke_set_inputs_first_pass( proposer, num_reqs=num_reqs, block_size=block_size, with_optional_attrs=True )[:4] @@ -600,10 +624,7 @@ def test_cad_rewritten_to_cross_attention_shape(self): # slot mapping is a slice of the primary group's query buffer (shares # storage from offset 0); a fresh slice is not identity-equal, so check # the underlying storage and length instead. - assert ( - cad.slot_mapping.data_ptr() - == proposer._per_group_query_slot_mapping_buffers[0].data_ptr() - ) + assert cad.slot_mapping.data_ptr() == proposer._per_group_query_slot_mapping_buffers[0].data_ptr() assert cad.slot_mapping.shape[0] == num_query_total # optional attrs the proposer rewrites when present. assert cad.actual_seq_lengths_q == [block_size] * num_reqs @@ -617,25 +638,24 @@ def test_cad_uses_model_reported_causality(self): block_size=block_size, draft_attn_causal=True, ) - _, _, cad, _ = self._invoke_set_inputs_first_pass( - proposer, num_reqs=num_reqs, block_size=block_size - )[:4] + _, _, cad, _ = self._invoke_set_inputs_first_pass(proposer, num_reqs=num_reqs, block_size=block_size)[:4] assert cad.causal is True def test_cad_query_start_loc_and_seq_lens(self): num_reqs, block_size, max_num_tokens = 4, 5, 256 - proposer = self._make_proposer( - max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size - ) - _nqt, _ti, cad, _extra = self._invoke_set_inputs_first_pass( - proposer, num_reqs=num_reqs, block_size=block_size - )[:4] + proposer = self._make_proposer(max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size) + _nqt, _ti, cad, _extra = self._invoke_set_inputs_first_pass(proposer, num_reqs=num_reqs, block_size=block_size)[ + :4 + ] expected_qsl = torch.arange(num_reqs + 1, dtype=torch.int32) * block_size assert torch.equal(cad.query_start_loc, expected_qsl) assert torch.equal(cad.query_start_loc_cpu, expected_qsl) # seq_lens grow by block_size when no tokens were rejected. - assert torch.equal(cad.seq_lens, torch.full((num_reqs,), 128 + block_size, dtype=torch.int32)) + expected = torch.full((num_reqs,), 128 + block_size, dtype=torch.int32) + assert torch.equal(cad.seq_lens, expected) + assert torch.equal(cad._seq_lens_cpu, expected) + assert torch.equal(cad.seq_lens_cpu, expected) class TestSetInputsFirstPassRejectedTokens(_DSparkProposerTestBase): @@ -644,38 +664,36 @@ class TestSetInputsFirstPassRejectedTokens(_DSparkProposerTestBase): def test_seq_lens_subtracts_rejected(self, monkeypatch): monkeypatch.setattr( - "vllm_ascend.spec_decode.dspark_proposer." - "copy_and_expand_dflash_and_dspark_inputs_kernel", + "vllm_ascend.spec_decode.dspark_proposer.copy_and_expand_dflash_and_dspark_inputs_kernel", MagicMock(), ) num_reqs, block_size, max_num_tokens = 4, 5, 256 - proposer = self._make_proposer( - max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size - ) + proposer = self._make_proposer(max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size) rejected = torch.full((num_reqs,), 2, dtype=torch.int32) _nqt, _ti, cad, _extra = self._invoke_set_inputs_first_pass( - proposer, num_reqs=num_reqs, block_size=block_size, num_rejected=rejected + proposer, + num_reqs=num_reqs, + block_size=block_size, + host_seq_len=126, + async_metadata=True, + num_rejected=rejected, )[:4] # effective = seq_lens(128) - rejected(2) = 126; then + block_size(5) = 131. - assert torch.equal( - cad.seq_lens, torch.full((num_reqs,), 128 - 2 + block_size, dtype=torch.int32) - ) + assert torch.equal(cad.seq_lens, torch.full((num_reqs,), 128 - 2 + block_size, dtype=torch.int32)) + expected_host = torch.full((num_reqs,), 126 + block_size, dtype=torch.int32) + assert torch.equal(cad._seq_lens_cpu, expected_host) + assert cad.seq_lens_cpu is None def test_kernel_called_with_has_num_rejected(self, monkeypatch): kernel = MagicMock() monkeypatch.setattr( - "vllm_ascend.spec_decode.dspark_proposer." - "copy_and_expand_dflash_and_dspark_inputs_kernel", + "vllm_ascend.spec_decode.dspark_proposer.copy_and_expand_dflash_and_dspark_inputs_kernel", kernel, ) num_reqs, block_size, max_num_tokens = 4, 5, 256 - proposer = self._make_proposer( - max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size - ) + proposer = self._make_proposer(max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size) rejected = torch.full((num_reqs,), 2, dtype=torch.int32) - self._invoke_set_inputs_first_pass( - proposer, num_reqs=num_reqs, block_size=block_size, num_rejected=rejected - ) + self._invoke_set_inputs_first_pass(proposer, num_reqs=num_reqs, block_size=block_size, num_rejected=rejected) # The proposer calls the kernel as ``kernel[1,](...)`` (Triton-style # grid indexing), so the call lands on the indexed sub-mock. sub = kernel[1,] @@ -686,10 +704,8 @@ def test_kernel_called_with_has_num_rejected(self, monkeypatch): assert kwargs["SAMPLE_FROM_ANCHOR"] is True -class TestInitializeAttnBackendErrors(_DSparkProposerTestBase): - """``initialize_attn_backend`` raises clearly when the draft model does not - expose the DSpark layer-name API, or when no draft attention groups can be - built from the kv-cache groups.""" +class TestInitializeAttnBackend(_DSparkProposerTestBase): + """Initialization preserves each group's logical kernel block size.""" @staticmethod def _make_proposer_for_init(): @@ -698,32 +714,45 @@ def _make_proposer_for_init(): proposer.device = torch.device("cpu") return proposer - def test_model_without_draft_layer_names_raises(self, monkeypatch): - # get_layers_from_vllm_config is called first; stub it so the model - # check is what actually fails. + def test_initialization_tracks_logical_block_size_per_gid(self, monkeypatch): + manager_specs = [MagicMock(), MagicMock()] + for spec in manager_specs: + spec.block_size = 384 + + backend = MagicMock() + backend.full_cls_name.return_value = "fake.backend" + layers = {} + for gid in range(2): + layer = MagicMock() + layer.get_attn_backend.return_value = backend + layers[f"L{gid}"] = layer monkeypatch.setattr( "vllm_ascend.spec_decode.dspark_proposer.get_layers_from_vllm_config", - lambda *a, **k: {}, + lambda *a, **k: layers, ) - proposer = self._make_proposer_for_init() - # model lacks get_draft_kv_cache_layer_names entirely. - proposer.model = SimpleNamespace() - kv_cache_config = SimpleNamespace(kv_cache_groups=[]) - with pytest.raises(RuntimeError, match="get_draft_kv_cache_layer_names"): - proposer.initialize_attn_backend(kv_cache_config) - - def test_no_draft_attn_groups_raises(self, monkeypatch): - monkeypatch.setattr( - "vllm_ascend.spec_decode.dspark_proposer.get_layers_from_vllm_config", - lambda *a, **k: {}, - ) proposer = self._make_proposer_for_init() - # draft layer names exist, but no kv-cache group names overlap them. - proposer.model = SimpleNamespace(get_draft_kv_cache_layer_names=lambda: {"L0"}) - - non_overlapping_group = SimpleNamespace(layer_names=["OTHER_LAYER"]) - kv_cache_config = SimpleNamespace(kv_cache_groups=[non_overlapping_group]) - with pytest.raises(RuntimeError, match="registered draft attention groups"): - proposer.initialize_attn_backend(kv_cache_config) -# fmt: on + proposer.model = SimpleNamespace(get_draft_kv_cache_layer_names=lambda: {"L0", "L1"}) + proposer.max_query_tokens = 8 + proposer.max_num_tokens = 16 + kv_cache_config = SimpleNamespace( + kv_cache_groups=[ + SimpleNamespace( + layer_names=[f"L{gid}"], + kv_cache_spec=manager_specs[gid], + ) + for gid in range(2) + ], + ) + + with patch.object(AttentionGroup, "create_metadata_builders") as create_builders: + proposer.initialize_attn_backend( + kv_cache_config, + kernel_block_sizes=[128, 64], + ) + + assert [spec.block_size for spec in manager_specs] == [384, 384] + assert proposer._per_group_kernel_block_sizes == {0: 128, 1: 64} + assert [group.kv_cache_group_id for group in proposer.draft_attn_groups] == [0, 1] + assert proposer.kernel_block_size == 128 + assert [call.kwargs["kernel_block_size"] for call in create_builders.call_args_list] == [128, 64] diff --git a/tests/ut/spec_decode/test_llm_base_proposer.py b/tests/ut/spec_decode/test_llm_base_proposer.py index 0f78198a8ac2..91f458876e6b 100644 --- a/tests/ut/spec_decode/test_llm_base_proposer.py +++ b/tests/ut/spec_decode/test_llm_base_proposer.py @@ -23,6 +23,7 @@ from unittest.mock import MagicMock, patch import pytest +import torch.nn as nn from vllm.config import CUDAGraphMode from vllm_ascend.spec_decode.llm_base_proposer import AscendSpecDecodeBaseProposer @@ -76,16 +77,22 @@ def test_pixtral_uses_vision_config_image_token_id(self): assert image_token_index == 789 - def test_kimi_uses_media_placeholder_token_id(self): + @pytest.mark.parametrize( + "model_name", + [ + "KimiK25ForConditionalGeneration", + "KimiK3ForConditionalGeneration", + "AscendKimiK3ForConditionalGeneration", + ], + ) + def test_kimi_uses_media_placeholder_token_id(self, model_name: str): config = SimpleNamespace( image_token_id=123, image_token_index=456, media_placeholder_token_id=789, ) - image_token_index = AscendSpecDecodeBaseProposer._get_multimodal_image_token_index( - "KimiK25ForConditionalGeneration", config - ) + image_token_index = AscendSpecDecodeBaseProposer._get_multimodal_image_token_index(model_name, config) assert image_token_index == 789 @@ -103,7 +110,8 @@ def test_load_model_reads_validated_draft_window_size(): proposer = AscendSpecDecodeBaseProposer.__new__(AscendSpecDecodeBaseProposer) proposer.vllm_config = SimpleNamespace(additional_config={"draft_window_size": 64}) proposer.maybe_eager_context = nullcontext() - proposer._get_model = MagicMock(return_value=MagicMock()) + draft_model = MagicMock() + proposer._get_model = MagicMock(return_value=draft_model) proposer.method = "eagle3" proposer.num_speculative_tokens = 4 proposer.runner = SimpleNamespace(max_num_reqs=8) @@ -134,6 +142,60 @@ def test_load_model_reads_validated_draft_window_size(): assert proposer.draft_window_size == 4096 mock_adapter.assert_called_once_with(4096, 16, 8, 4, "cpu") + draft_model.finish_shared_layer_preparation.assert_called_once_with() + + +def test_dspark_embedding_sharing_uses_model_preparation_hook(): + proposer = AscendSpecDecodeBaseProposer.__new__(AscendSpecDecodeBaseProposer) + draft_embed = nn.Linear(2, 3, bias=False) + target_embed = nn.Linear(2, 3, bias=False) + prepared_embed = nn.Linear(2, 3, bias=False) + proposer.method = "dspark" + proposer.model = SimpleNamespace( + has_own_embed_tokens=False, + model=SimpleNamespace(embed_tokens=draft_embed), + prepare_shared_layer=MagicMock(return_value=prepared_embed), + ) + target_model = SimpleNamespace( + model=SimpleNamespace(embed_tokens=target_embed), + ) + + with patch("vllm_ascend.spec_decode.llm_base_proposer.get_pp_group") as mock_pp_group: + mock_pp_group.return_value.world_size = 1 + proposer._maybe_share_embeddings(target_model) + + proposer.model.prepare_shared_layer.assert_called_once_with( + draft_embed, + target_embed, + "draft embed_tokens.weight", + ) + assert proposer.model.model.embed_tokens is prepared_embed + + +def test_dspark_lm_head_sharing_uses_model_preparation_hook(): + proposer = AscendSpecDecodeBaseProposer.__new__(AscendSpecDecodeBaseProposer) + draft_lm_head = nn.Linear(2, 3, bias=False) + target_lm_head = nn.Linear(2, 3, bias=False) + prepared_lm_head = nn.Linear(2, 3, bias=False) + proposer.method = "dspark" + proposer.model = SimpleNamespace( + has_own_lm_head=False, + lm_head=draft_lm_head, + prepare_shared_layer=MagicMock(return_value=prepared_lm_head), + ) + proposer.vllm_config = SimpleNamespace( + compilation_config=SimpleNamespace(cudagraph_mode=CUDAGraphMode.NONE), + ) + proposer.use_cuda_graph = False + + proposer._maybe_share_lm_head(SimpleNamespace(lm_head=target_lm_head)) + + proposer.model.prepare_shared_layer.assert_called_once_with( + draft_lm_head, + target_lm_head, + "draft lm_head.weight", + ) + assert proposer.model.lm_head is prepared_lm_head class TestDisablePaddedDrafterBatchWithFullGraph: diff --git a/tests/ut/spec_decode/test_utils.py b/tests/ut/spec_decode/test_utils.py index 494082f4bd3f..575d53cb89c7 100644 --- a/tests/ut/spec_decode/test_utils.py +++ b/tests/ut/spec_decode/test_utils.py @@ -8,10 +8,13 @@ from __future__ import annotations +from types import SimpleNamespace + import numpy as np import torch from vllm_ascend.spec_decode.utils import ( + SlidingWindowAdapter, correct_optimistic_seq_lens_cpu, update_num_computed_tokens_for_batch_change, ) @@ -184,3 +187,28 @@ def test_cpu_and_gpu_corrections_agree(): gpu_seq_lens = num_computed_gpu.numpy() + num_scheduled_step_n np.testing.assert_array_equal(optimistic, gpu_seq_lens) + + +def test_sliding_window_updates_host_length_mirrors(): + adapter = SlidingWindowAdapter( + window_size=128, + block_size=64, + max_num_reqs=1, + future_offset=0, + device=torch.device("cpu"), + ) + metadata = SimpleNamespace( + block_table_tensor=torch.arange(10, dtype=torch.int32).view(1, 10), + seq_lens=torch.tensor([300], dtype=torch.int32), + seq_lens_cpu=torch.tensor([300], dtype=torch.int32), + _seq_lens_cpu=torch.tensor([300], dtype=torch.int32), + seq_lens_cpu_upper_bound=torch.tensor([300], dtype=torch.int32), + ) + + adapter.apply(metadata) + + expected = torch.tensor([172], dtype=torch.int32) + torch.testing.assert_close(metadata.seq_lens, expected) + torch.testing.assert_close(metadata.seq_lens_cpu, expected) + torch.testing.assert_close(metadata._seq_lens_cpu, expected) + torch.testing.assert_close(metadata.seq_lens_cpu_upper_bound, expected) diff --git a/vllm_ascend/patch/platform/patch_speculative_config.py b/vllm_ascend/patch/platform/patch_speculative_config.py index c4f4737c8043..269607a336de 100644 --- a/vllm_ascend/patch/platform/patch_speculative_config.py +++ b/vllm_ascend/patch/platform/patch_speculative_config.py @@ -1,6 +1,23 @@ +from transformers import PretrainedConfig from vllm.config.speculative import SpeculativeConfig _orig_post_init = SpeculativeConfig.__post_init__ +_orig_hf_config_override = SpeculativeConfig.hf_config_override + + +def _normalize_legacy_qwen3_dspark_config(hf_config: PretrainedConfig) -> PretrainedConfig: + hf_config = _orig_hf_config_override(hf_config) + architectures = hf_config.architectures or () + if hf_config.model_type == "qwen3" and "DSparkDraftModel" in architectures: + dflash_config = hf_config.dflash_config + hf_config.update( + { + "architectures": ["Qwen3DSparkModel"], + "mask_token_id": dflash_config["mask_token_id"], + "target_layer_ids": dflash_config["target_layer_ids"], + } + ) + return hf_config def _dspark_post_init(self): @@ -16,4 +33,5 @@ def _dspark_post_init(self): draft_hf_config.ptd_token_id = getattr(draft_hf_config, "mask_token_id", None) # type: ignore +SpeculativeConfig.hf_config_override = staticmethod(_normalize_legacy_qwen3_dspark_config) SpeculativeConfig.__post_init__ = _dspark_post_init diff --git a/vllm_ascend/spec_decode/dspark_proposer.py b/vllm_ascend/spec_decode/dspark_proposer.py index ac1485c8268f..643f2128aae1 100644 --- a/vllm_ascend/spec_decode/dspark_proposer.py +++ b/vllm_ascend/spec_decode/dspark_proposer.py @@ -43,8 +43,8 @@ def __init__( blk = 1 + self.num_speculative_tokens self._dspark_draft_buffer = torch.zeros((self.max_batch_size, blk), dtype=torch.int64, device=device) self._dspark_seed_buffer = torch.zeros(self.max_batch_size, dtype=torch.int64, device=device) - # DSpark is not supported in vllm v1, so related property needs to be reset here. - del self.hidden_size, self.hidden_states, self._dflash_hidden_states # type: ignore[has-type] + # Replace the target-sized DFlash buffers with the draft model's hidden + # size. Assignment releases the old tensors without an explicit del. self.hidden_size = vllm_config.speculative_config.draft_model_config.get_hidden_size() self.hidden_states = torch.zeros( (self.max_num_tokens, self.hidden_size), @@ -87,27 +87,20 @@ def __init__( device=device, ) - # TODO simplify these comments - # block_table / slot_mapping bookkeeping (10 dicts below). v1 self- - # manages per kv_cache_group_id / per layer because it lacks v2's - # BlockTables scaffold; v2 injects a single self.block_tables - # (BlockTables, with .slot_mappings) + build_slot_mappings_by_layer, - # so the speculator holds none of these. P2 refactor target (move to - # runner). - - # per-gid block_table from runner (just read) + # The v1 runner owns block tables and slot mappings. Keep per-group + # references here because K3 draft layers can span multiple cache + # groups with different logical block sizes. self._per_group_block_tables: dict[int, torch.Tensor] = {} - # per-gid slot_mapping from runner (just read) self._per_group_slot_mappings: dict[int, torch.Tensor] = {} + # Per-gid logical block size used to expand slot mappings. The KV + # manager's physical page can be larger when hybrid cache groups share + # one allocation, so kv_cache_spec.block_size is not interchangeable + # with the attention kernel's block size. + self._per_group_kernel_block_sizes: dict[int, int] = {} - # per-gid block_table (use in proposer) self._per_group_block_table_buffers: dict[int, torch.Tensor] = {} - # per-gid query slot_mapping buffer self._per_group_query_slot_mapping_buffers: dict[int, torch.Tensor] = {} - # per-gid context slot_mapping buffer self._per_group_context_slot_mapping_buffers: dict[int, torch.Tensor] = {} - - # per-layer context slot mappings as a flat list self._context_slot_mapping_buffers: list[torch.Tensor | None] | None = None def _compute_confidence( @@ -129,24 +122,22 @@ def _compute_confidence( confidence.copy_(conf_raw.reshape(num_reqs, self.num_speculative_tokens)) return confidence - def initialize_attn_backend(self, kv_cache_config, kernel_block_sizes=None) -> None: + def initialize_attn_backend( + self, + kv_cache_config, + kernel_block_sizes: list[int] | None = None, + ) -> None: # Find draft layers (attention layers added by draft model) all_attn_layers = get_layers_from_vllm_config( self.vllm_config, AttentionLayerBase, # type: ignore[type-abstract] ) - attention_groups_list: list[dict[tuple[str, str], AttentionGroup]] = [] - # the draft layers have multiple kv_cache_groups - if not hasattr(self.model, "get_draft_kv_cache_layer_names"): - raise RuntimeError( - "DSpark standard-cache path requires the draft model to expose get_draft_kv_cache_layer_names" - ) - self._draft_attn_layer_names = set(self.model.get_draft_kv_cache_layer_names()) self.attn_layer_names = list(sorted(self._draft_attn_layer_names)) + self._per_group_kernel_block_sizes = {} + self.draft_attn_groups: list[AttentionGroup] = [] - # there are many kv groups other than one for kv_cache_gid, kv_cache_group_spec in enumerate(kv_cache_config.kv_cache_groups): draft_layer_names_in_group = set(kv_cache_group_spec.layer_names) & self._draft_attn_layer_names if not draft_layer_names_in_group: @@ -162,33 +153,31 @@ def initialize_attn_backend(self, kv_cache_config, kernel_block_sizes=None) -> N key = (attn_backend.full_cls_name(), layer_kv_cache_spec) if key not in attention_groups: + kernel_block_size = int( + kernel_block_sizes[kv_cache_gid] + if kernel_block_sizes is not None and kv_cache_gid < len(kernel_block_sizes) + else layer_kv_cache_spec.block_size + ) attn_group = AttentionGroup( attn_backend, [layer_name], layer_kv_cache_spec, kv_cache_gid, ) - attn_group.create_metadata_builders(self.vllm_config, self.device) + attn_group.create_metadata_builders( + self.vllm_config, + self.device, + kernel_block_size=kernel_block_size, + ) + self._per_group_kernel_block_sizes[kv_cache_gid] = kernel_block_size attention_groups[key] = attn_group else: attention_groups[key].layer_names.append(layer_name) - attention_groups_list.append(attention_groups) - - self.draft_attn_groups = [ - attention_group - for attention_groups in attention_groups_list - for attention_group in attention_groups.values() - ] - self.kv_cache_gid = 0 - if not self.draft_attn_groups: - raise RuntimeError( - "DSpark standard-cache path requires registered draft attention " - f"groups. Missing layers: {self.attn_layer_names}" - ) + self.draft_attn_groups.extend(attention_groups.values()) self.kv_cache_gid = self.draft_attn_groups[0].kv_cache_group_id - self.kernel_block_size = int(self.draft_attn_groups[0].kv_cache_spec.block_size) + self.kernel_block_size = self._per_group_kernel_block_sizes[self.kv_cache_gid] name_to_gid = { ln: gid @@ -257,13 +246,10 @@ def set_inputs_first_pass( # Query block: reuse the DFlash inputs kernel logic (host-side ref) # per kv-cache-group to fill positions / input_ids / query slot_mapping # / token_indices. - draft_attn_groups = getattr(self, "draft_attn_groups", []) - for attn_group in draft_attn_groups: + for attn_group in self.draft_attn_groups: gid = attn_group.kv_cache_group_id - gid_block_table = self._per_group_block_table_buffers.get(gid) - if gid_block_table is None: - continue - kv_block_size = int(attn_group.kv_cache_spec.block_size) + gid_block_table = self._per_group_block_table_buffers[gid] + kernel_block_size = self._per_group_kernel_block_sizes[gid] copy_and_expand_dflash_and_dspark_inputs_kernel[ (_compute_num_programs(self._dflash_num_context, num_query_total),) ]( @@ -287,7 +273,7 @@ def set_inputs_first_pass( num_rejected_tokens_ptr=num_rejected_tokens_gpu, # Scalars parallel_drafting_token_id=self.parallel_drafting_token_id, - block_size=kv_block_size, + block_size=kernel_block_size, num_query_per_req=self.num_query_per_req, num_speculative_tokens=self.num_speculative_tokens, total_input_tokens=self._dflash_num_context, @@ -306,6 +292,15 @@ def set_inputs_first_pass( cad.query_start_loc = self.arange_dflash[: batch_size + 1] * self.num_query_per_req cad.seq_lens = effective_seq_lens + self.num_query_per_req + # The model runner has already corrected this canonical host mirror + # with the accepted-token count. Extend it on CPU alongside the device + # lengths, without another reject D2H copy or attention-side wait. + if cad._seq_lens_cpu is not None: + draft_seq_lens_cpu = cad._seq_lens_cpu.clone() + draft_seq_lens_cpu[:batch_size].add_(self.num_query_per_req) + cad._seq_lens_cpu = draft_seq_lens_cpu + if getattr(cad, "seq_lens_cpu", None) is not None: + cad.seq_lens_cpu = draft_seq_lens_cpu cad.query_start_loc_cpu = ( torch.from_numpy(self.token_arange_np[: batch_size + 1]).clone() * self.num_query_per_req ).to(torch.int32) diff --git a/vllm_ascend/spec_decode/llm_base_proposer.py b/vllm_ascend/spec_decode/llm_base_proposer.py index 8c5121e20a6d..f66431a5c68a 100644 --- a/vllm_ascend/spec_decode/llm_base_proposer.py +++ b/vllm_ascend/spec_decode/llm_base_proposer.py @@ -26,6 +26,7 @@ from vllm.model_executor.models.llama_eagle3 import Eagle3LlamaForCausalLM from vllm.model_executor.models.qwen3_dflash import DFlashQwen3ForCausalLM from vllm.model_executor.models.qwen3_dspark import Qwen3DSparkForCausalLM +from vllm.models.kimi_k3.nvidia.dspark_mla import K3DSparkForCausalLM from vllm.triton_utils import HAS_TRITON, triton from vllm.utils.platform_utils import is_pin_memory_available from vllm.v1.attention.backends.utils import CommonAttentionMetadata @@ -66,6 +67,16 @@ # Currently we will fix block size to a small one since `num_reqs` can't be too large _PREPARE_INPUTS_BLOCK_SIZE = 4 +_HIDDEN_STATE_DRAFTER_TYPES = ( + Eagle3LlamaForCausalLM, + DFlashQwen3ForCausalLM, + Qwen3DSparkForCausalLM, + K3DSparkForCausalLM, + Eagle3VwnLlamaForCausalLM, + Eagle3DeepseekV2ForCausalLM, + DSparkDeepseekV4ForCausalLM, +) + def greedy_sample(logits: torch.Tensor) -> torch.Tensor: tp_group = get_tp_group() @@ -115,7 +126,11 @@ def _get_multimodal_image_token_index(model_name: str, config: Any) -> int: return config.image_token_id if model_name == "PixtralForConditionalGeneration": return config.vision_config.image_token_id - if model_name == "KimiK25ForConditionalGeneration": + if model_name in { + "KimiK25ForConditionalGeneration", + "KimiK3ForConditionalGeneration", + "AscendKimiK3ForConditionalGeneration", + }: return config.media_placeholder_token_id return config.image_token_index @@ -363,6 +378,13 @@ def load_model(self, model: nn.Module) -> None: self._maybe_share_embeddings(target_language_model) self._maybe_share_topk_indices(target_language_model) self._maybe_share_lm_head(model) + finish_shared_layer_preparation = getattr( + self.model, + "finish_shared_layer_preparation", + None, + ) + if callable(finish_shared_layer_preparation): + finish_shared_layer_preparation() if ( self.parallel_drafting @@ -447,9 +469,31 @@ def _maybe_share_embeddings(self, target_language_model: nn.Module) -> None: ) if share_embeddings: - if hasattr(self.model.model, "embed_tokens"): - del self.model.model.embed_tokens - self.model.model.embed_tokens = target_embed_tokens + draft_embed_tokens = getattr( + self.model.model, + "embed_tokens", + None, + ) + prepare_shared_layer = getattr( + self.model, + "prepare_shared_layer", + None, + ) + prepared_embed_tokens = ( + prepare_shared_layer( + draft_embed_tokens, + target_embed_tokens, + "draft embed_tokens.weight", + ) + if callable(prepare_shared_layer) + else None + ) + if prepared_embed_tokens is not None: + self.model.model.embed_tokens = prepared_embed_tokens + else: + if hasattr(self.model.model, "embed_tokens"): + del self.model.model.embed_tokens + self.model.model.embed_tokens = target_embed_tokens else: logger.info( "[spec_decode/base] PP>1: draft model loaded its own vocab embedding" @@ -484,17 +528,34 @@ def _maybe_share_lm_head(self, model: nn.Module) -> None: ) else: logger.info("[spec_decode/base] Loading EAGLE/DFLASH LM head weights from the target model.") + target_lm_head = None if hasattr(model, "lm_head"): - self.model.lm_head = model.lm_head + target_lm_head = model.lm_head elif hasattr(model, "get_language_model") and hasattr(model.get_language_model(), "lm_head"): - self.model.lm_head = model.get_language_model().lm_head - else: + target_lm_head = model.get_language_model().lm_head + if target_lm_head is None: logger.warning( "[spec_decode/base] Target model has no accessible lm_head" " for sharing. Draft model will use its own lm_head." " This may cause incorrect logits if the draft lm_head" " is not trained." ) + else: + prepare_shared_layer = getattr( + self.model, + "prepare_shared_layer", + None, + ) + prepared_lm_head = ( + prepare_shared_layer( + getattr(self.model, "lm_head", None), + target_lm_head, + "draft lm_head.weight", + ) + if callable(prepare_shared_layer) + else None + ) + self.model.lm_head = target_lm_head if prepared_lm_head is None else prepared_lm_head if self.method == "mtp" and self.vllm_config.model_config.is_deepseek_mla: for _, layer_module in self.model.model.layers.items(): @@ -793,18 +854,8 @@ def _propose( model = self.model if isinstance(model, BreakableACLGraphWrapper): model = model.unwrap() - assert isinstance( - model, - ( - Eagle3LlamaForCausalLM, - DFlashQwen3ForCausalLM, - Qwen3DSparkForCausalLM, - Eagle3VwnLlamaForCausalLM, - Eagle3DeepseekV2ForCausalLM, - DSparkDeepseekV4ForCausalLM, - ), - ) - target_hidden_states = self.model.combine_hidden_states(target_hidden_states) + assert isinstance(model, _HIDDEN_STATE_DRAFTER_TYPES) + target_hidden_states = model.combine_hidden_states(target_hidden_states) assert target_hidden_states.shape[-1] == self.hidden_size num_tokens, token_indices_to_sample, common_attn_metadata, long_seq_args = self.set_inputs_first_pass( @@ -878,6 +929,19 @@ def _propose( ) if self.method == "dflash": common_attn_metadata.seq_lens = self._adjust_tensor(common_attn_metadata.seq_lens, num_reqs_padded) + elif self.method == "dspark": + # DSpark already rewrote both device and host sequence lengths + # in set_inputs_first_pass. Preserve those values while + # extending only the padded tail for full-graph replay. + common_attn_metadata.seq_lens = self._adjust_tensor(common_attn_metadata.seq_lens, num_reqs_padded) + if common_attn_metadata.seq_lens_cpu is not None: + common_attn_metadata.seq_lens_cpu = self._adjust_tensor( + common_attn_metadata.seq_lens_cpu, num_reqs_padded + ) + if common_attn_metadata._seq_lens_cpu is not None: + common_attn_metadata._seq_lens_cpu = self._adjust_tensor( + common_attn_metadata._seq_lens_cpu, num_reqs_padded + ) else: common_attn_metadata.seq_lens = self._adjust_tensor(self.runner.seq_lens, num_reqs_padded) common_attn_metadata.seq_lens_cpu = self._adjust_tensor( From b2aa0ec8a99b41d91a85daf941b8913fa1e3b7ce Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Mon, 24 Aug 2026 04:31:25 -0500 Subject: [PATCH 13/50] feat(kv-cache): group Kimi K3 hybrid caches Group target attention, draft attention, and recurrent caches by runtime topology while preserving exact cache-spec semantics. Carry logical block sizes through scheduler, runner, and DSpark metadata paths, and retain Mamba request capacity for prefix caching and speculative decoding. Signed-off-by: maoxx241 --- .../platform/test_prefix_cache_cp_patches.py | 232 ++++++++++++++++++ .../test_patch_mamba_utils_uniform_groups.py | 61 +++++ tests/ut/spec_decode/test_dspark_proposer.py | 78 ++++++ tests/ut/test_compressed_prefix_cache.py | 102 ++++++++ tests/ut/worker/a2/test_block_table.py | 43 ++++ tests/ut/worker/a2/test_model_runner_v1.py | 177 +++++++++++++ vllm_ascend/core/kv_cache_interface.py | 60 +++-- .../platform/patch_kv_cache_coordinator.py | 5 +- .../patch/platform/patch_kv_cache_utils.py | 149 +++++++++++ vllm_ascend/patch/worker/patch_mamba_utils.py | 36 ++- vllm_ascend/worker/block_table.py | 20 +- vllm_ascend/worker/model_runner_v1.py | 43 ++-- 12 files changed, 942 insertions(+), 64 deletions(-) create mode 100644 tests/ut/patch/worker/test_patch_mamba_utils_uniform_groups.py diff --git a/tests/ut/patch/platform/test_prefix_cache_cp_patches.py b/tests/ut/patch/platform/test_prefix_cache_cp_patches.py index c353ff85939a..2bce0fb1e0ba 100644 --- a/tests/ut/patch/platform/test_prefix_cache_cp_patches.py +++ b/tests/ut/patch/platform/test_prefix_cache_cp_patches.py @@ -1,12 +1,14 @@ # SPDX-License-Identifier: Apache-2.0 import math +from dataclasses import replace from types import SimpleNamespace from unittest.mock import MagicMock import pytest import torch from vllm.v1.core.block_pool import BlockPool +from vllm.v1.core.kv_cache_utils import generate_scheduler_kv_cache_config from vllm.v1.core.single_type_kv_cache_manager import ( FullAttentionManager, SlidingWindowManager, @@ -30,6 +32,8 @@ ) from vllm_ascend.patch.platform.patch_kv_cache_utils import ( _ascend_resolve_kv_cache_block_sizes, + _get_kimi_k3_dspark_mixed_kv_cache_groups, + _get_kv_cache_config_deepseek_v4, group_and_unify_kv_cache_specs, ) from vllm_ascend.patch.platform.patch_mamba_manager import AscendMambaManager @@ -64,6 +68,65 @@ def _make_hybrid_kv_cache_config( ) +def _make_kimi_k3_dspark_kv_cache_specs( + *, + block_size: int = 384, + page_size: int = 488448, + target_layer_count: int = 24, + draft_layer_count: int = 5, + mamba_layer_count: int = 69, + draft_uses_mla: bool = False, +) -> dict: + target_mla_spec = MLAAttentionSpec( + block_size=block_size, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + page_size_padded=page_size, + cache_dtype_str="auto", + ) + if draft_uses_mla: + draft_attention_spec = MLAAttentionSpec( + block_size=block_size, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + page_size_padded=page_size, + cache_dtype_str="auto", + non_causal_multi_token_decode=True, + ) + else: + draft_attention_spec = FullAttentionSpec( + block_size=block_size, + num_kv_heads=1, + head_size=64, + dtype=torch.bfloat16, + page_size_padded=page_size, + ) + mamba_spec = MambaSpec( + block_size=block_size, + shapes=((10, 2304), (6, 128, 128)), + dtypes=(torch.bfloat16, torch.float32), + page_size_padded=page_size, + mamba_cache_mode="align", + num_speculative_blocks=7, + ) + specs = { + f"language_model.model.layers.{layer_idx}.self_attn.attn": target_mla_spec + for layer_idx in range(target_layer_count) + } + specs.update( + { + f"model.layers.{layer_idx}.self_attn.attn": draft_attention_spec + for layer_idx in range(93, 93 + draft_layer_count) + } + ) + specs.update( + {f"language_model.model.layers.{layer_idx}.self_attn": mamba_spec for layer_idx in range(mamba_layer_count)} + ) + return specs + + def _make_deepseek_v4_kv_cache_config() -> KVCacheConfig: c4_spec = MLAAttentionSpec( block_size=128 * 4, @@ -108,6 +171,7 @@ def _make_vllm_config( cache_config=SimpleNamespace( block_size=block_size, enable_prefix_caching=enable_prefix_caching, + mamba_cache_mode="align", prefix_match_unit=None, ), parallel_config=SimpleNamespace( @@ -195,6 +259,134 @@ def test_deepseek_v4_groups_use_logical_sizes_and_full_attention_manager() -> No assert KVCacheSpecRegistry.get_manager_class(spec) is FullAttentionManager +@pytest.mark.parametrize( + ("block_size", "page_size"), + [ + pytest.param(384, 488448, id="tp16"), + pytest.param(768, 976896, id="tp8"), + ], +) +def test_kimi_k3_gqa_uses_four_mixed_kv_groups_and_five_builders( + block_size: int, + page_size: int, +) -> None: + groups = _get_kimi_k3_dspark_mixed_kv_cache_groups( + _make_kimi_k3_dspark_kv_cache_specs( + block_size=block_size, + page_size=page_size, + ) + ) + + assert groups is not None + assert [len(group.layer_names) for group in groups] == [29, 23, 23, 23] + assert all(isinstance(group.kv_cache_spec, UniformTypeKVCacheSpecs) for group in groups) + + mixed_specs = groups[0].kv_cache_spec.kv_cache_specs + assert sum(isinstance(spec, MLAAttentionSpec) for spec in mixed_specs.values()) == 24 + assert ( + sum( + isinstance(spec, FullAttentionSpec) and not isinstance(spec, MLAAttentionSpec) + for spec in mixed_specs.values() + ) + == 5 + ) + + # The model runner keys builders by backend and exact inner spec. The + # mixed attention group contributes MLA + GQA, and each Mamba group one. + metadata_builder_count = sum(len(set(group.kv_cache_spec.kv_cache_specs.values())) for group in groups) + assert metadata_builder_count == 5 + + +def test_kimi_k3_mla_dspark_uses_four_groups_and_five_builders() -> None: + groups = _get_kimi_k3_dspark_mixed_kv_cache_groups(_make_kimi_k3_dspark_kv_cache_specs(draft_uses_mla=True)) + + assert groups is not None + assert [len(group.layer_names) for group in groups] == [29, 23, 23, 23] + # Target MLA is causal while draft MLA enables non-causal multi-token + # decode, so the mixed group needs two exact-spec builders. The three + # recurrent groups each contribute one more builder. + metadata_builder_count = sum(len(set(group.kv_cache_spec.kv_cache_specs.values())) for group in groups) + assert metadata_builder_count == 5 + + +def test_kimi_k3_gqa_mixed_groups_preserve_scheduler_and_mamba_contracts() -> None: + groups = _get_kimi_k3_dspark_mixed_kv_cache_groups(_make_kimi_k3_dspark_kv_cache_specs()) + assert groups is not None + worker_config = KVCacheConfig( + num_blocks=100, + kv_cache_tensors=[], + kv_cache_groups=groups, + ) + + assert worker_config.has_mamba_layers + assert worker_config.needs_kv_cache_zeroing + assert ( + groups[1].kv_cache_spec.max_num_blocks_per_req( + _make_vllm_config(enable_prefix_caching=True, dcp=1, block_size=384), + 3840, + ) + == 17 + ) + + scheduler_config = generate_scheduler_kv_cache_config([worker_config]) + assert isinstance(scheduler_config.kv_cache_groups[0].kv_cache_spec, MLAAttentionSpec) + assert all(isinstance(group.kv_cache_spec, MambaSpec) for group in scheduler_config.kv_cache_groups[1:]) + assert scheduler_config.needs_kv_cache_zeroing + + +def test_kimi_k3_gqa_mixed_groups_use_expected_physical_layout(monkeypatch) -> None: + groups = _get_kimi_k3_dspark_mixed_kv_cache_groups(_make_kimi_k3_dspark_kv_cache_specs()) + assert groups is not None + page_size = 488448 + expected_num_blocks = 100 + available_memory = page_size * 29 * expected_num_blocks + monkeypatch.setattr( + "vllm_ascend.patch.platform.patch_kv_cache_utils.may_override_num_blocks", + lambda _config, num_blocks: num_blocks, + ) + + num_blocks, tensors = _get_kv_cache_config_deepseek_v4( + SimpleNamespace(), + groups, + available_memory, + ) + + assert num_blocks == expected_num_blocks + assert len(tensors) == 29 + assert [len(tensor.shared_by) for tensor in tensors] == [4] * 23 + [1] * 6 + assert all(tensor.size == page_size * expected_num_blocks for tensor in tensors) + assert sum(tensor.size for tensor in tensors) == available_memory + + +def test_kimi_k3_gqa_mixed_grouping_falls_back_on_partial_signature() -> None: + specs = _make_kimi_k3_dspark_kv_cache_specs() + draft_layer = "model.layers.93.self_attn.attn" + specs[draft_layer] = replace(specs[draft_layer], non_causal=True) + + assert _get_kimi_k3_dspark_mixed_kv_cache_groups(specs) is None + + +def test_kimi_k3_dspark_group_count_is_derived_from_layer_ratio() -> None: + groups = _get_kimi_k3_dspark_mixed_kv_cache_groups( + _make_kimi_k3_dspark_kv_cache_specs( + target_layer_count=20, + draft_layer_count=4, + mamba_layer_count=70, + ) + ) + + assert groups is not None + assert [len(group.layer_names) for group in groups] == [24, 24, 23, 23] + + +def test_kimi_k3_dspark_mixed_grouping_falls_back_on_unaligned_pages() -> None: + specs = _make_kimi_k3_dspark_kv_cache_specs() + draft_layer = "model.layers.93.self_attn.attn" + specs[draft_layer] = replace(specs[draft_layer], page_size_padded=976896) + + assert _get_kimi_k3_dspark_mixed_kv_cache_groups(specs) is None + + def test_deepseek_v4_scheduler_lcm_uses_logical_group_sizes() -> None: kv_cache_config = _make_deepseek_v4_kv_cache_config() vllm_config = _make_vllm_config( @@ -326,6 +518,46 @@ def _fake_orig(*args, **kwargs): assert coordinator is sentinel +@pytest.mark.parametrize( + ("num_prefill_lookahead", "expected"), + [(None, 0), (8, 8)], +) +def test_get_kv_cache_coordinator_normalizes_prefill_lookahead( + monkeypatch, + num_prefill_lookahead: int | None, + expected: int, +) -> None: + kv_cache_config = _make_hybrid_kv_cache_config( + full_block_size=16, + mamba_block_size=16, + ) + captured_kwargs = {} + + def _fake_ascend_coordinator(*args, **kwargs): + captured_kwargs.update(kwargs) + return object() + + monkeypatch.setattr( + "vllm_ascend.patch.platform.patch_kv_cache_coordinator.AscendHybridKVCacheCoordinator", + _fake_ascend_coordinator, + ) + + get_kv_cache_coordinator( + kv_cache_config, + max_model_len=1024, + max_num_batched_tokens=1024, + use_eagle=False, + enable_caching=True, + enable_kv_cache_events=False, + dcp_world_size=1, + pcp_world_size=1, + hash_block_size=16, + num_prefill_lookahead=num_prefill_lookahead, + ) + + assert captured_kwargs["num_prefill_lookahead"] == expected + + def test_get_kv_cache_coordinator_uses_ascend_for_deepseek_v4(monkeypatch) -> None: sentinel = object() kv_cache_config = _make_deepseek_v4_kv_cache_config() diff --git a/tests/ut/patch/worker/test_patch_mamba_utils_uniform_groups.py b/tests/ut/patch/worker/test_patch_mamba_utils_uniform_groups.py new file mode 100644 index 000000000000..e41547642a9f --- /dev/null +++ b/tests/ut/patch/worker/test_patch_mamba_utils_uniform_groups.py @@ -0,0 +1,61 @@ +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace + +import torch +from vllm.v1.kv_cache_interface import ( + KVCacheConfig, + KVCacheGroupSpec, + MambaSpec, + UniformTypeKVCacheSpecs, +) +from vllm.v1.worker.mamba_utils import MambaCopyBuffers + +from vllm_ascend.patch.worker.patch_mamba_utils import _get_mamba_groups + + +def test_uniform_mamba_groups_are_visible_to_all_mamba_buffers() -> None: + mamba_spec = MambaSpec( + block_size=384, + shapes=((10, 2304), (6, 128, 128)), + dtypes=(torch.bfloat16, torch.float32), + page_size_padded=488448, + mamba_cache_mode="align", + num_speculative_blocks=7, + ) + groups = [] + for group_id in range(3): + layer_specs = {f"mamba.{group_id}.{layer_id}": mamba_spec for layer_id in range(23)} + uniform_spec = UniformTypeKVCacheSpecs.from_specs(layer_specs) + assert uniform_spec is not None + groups.append( + KVCacheGroupSpec( + layer_names=list(layer_specs), + kv_cache_spec=uniform_spec, + ) + ) + kv_cache_config = KVCacheConfig( + num_blocks=100, + kv_cache_tensors=[], + kv_cache_groups=groups, + ) + + group_ids, resolved_spec = _get_mamba_groups(kv_cache_config) + assert group_ids == [0, 1, 2] + assert resolved_spec == mamba_spec + + def make_buffer(n: int, dtype: torch.dtype) -> SimpleNamespace: + return SimpleNamespace(n=n, dtype=dtype) + + copy_bufs = MambaCopyBuffers.create( + max_num_reqs=2, + kv_cache_config=kv_cache_config, + copy_funcs=(object(), object()), + make_buffer=make_buffer, + ) + assert copy_bufs.mamba_group_ids == [0, 1, 2] + assert copy_bufs.mamba_spec == mamba_spec + assert copy_bufs.src_ptrs.n == 2 * 69 * 2 + assert copy_bufs.src_ptrs.dtype == torch.int64 + assert copy_bufs.dst_ptrs.dtype == torch.int64 + assert copy_bufs.sizes.dtype == torch.int32 diff --git a/tests/ut/spec_decode/test_dspark_proposer.py b/tests/ut/spec_decode/test_dspark_proposer.py index 73857de027f3..bed7ab8d69f0 100644 --- a/tests/ut/spec_decode/test_dspark_proposer.py +++ b/tests/ut/spec_decode/test_dspark_proposer.py @@ -26,6 +26,11 @@ import numpy as np import pytest import torch +from vllm.v1.kv_cache_interface import ( + FullAttentionSpec, + MLAAttentionSpec, + UniformTypeKVCacheSpecs, +) from vllm.v1.worker.utils import AttentionGroup from vllm_ascend.attention.attention_v1 import AscendAttentionState @@ -756,3 +761,76 @@ def test_initialization_tracks_logical_block_size_per_gid(self, monkeypatch): assert [group.kv_cache_group_id for group in proposer.draft_attn_groups] == [0, 1] assert proposer.kernel_block_size == 128 assert [call.kwargs["kernel_block_size"] for call in create_builders.call_args_list] == [128, 64] + + @pytest.mark.parametrize("draft_uses_mla", [False, True], ids=["gqa", "mla"]) + def test_mixed_target_and_dspark_group_creates_one_draft_attention_group(self, monkeypatch, draft_uses_mla: bool): + page_size = 488448 + target_layer = "language_model.model.layers.3.self_attn.attn" + draft_layers = [f"model.layers.{layer_idx}.self_attn.attn" for layer_idx in range(93, 98)] + target_spec = MLAAttentionSpec( + block_size=384, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + page_size_padded=page_size, + ) + if draft_uses_mla: + draft_spec = MLAAttentionSpec( + block_size=384, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + page_size_padded=page_size, + non_causal_multi_token_decode=True, + ) + else: + draft_spec = FullAttentionSpec( + block_size=384, + num_kv_heads=1, + head_size=64, + dtype=torch.bfloat16, + page_size_padded=page_size, + ) + mixed_spec = UniformTypeKVCacheSpecs.from_specs( + { + target_layer: target_spec, + **{layer_name: draft_spec for layer_name in draft_layers}, + } + ) + assert mixed_spec is not None + + backend = MagicMock() + backend.full_cls_name.return_value = "fake.gqa.backend" + layers = {} + for layer_name in draft_layers: + layer = MagicMock() + layer.get_attn_backend.return_value = backend + layers[layer_name] = layer + monkeypatch.setattr( + "vllm_ascend.spec_decode.dspark_proposer.get_layers_from_vllm_config", + lambda *args, **kwargs: layers, + ) + + proposer = self._make_proposer_for_init() + proposer.model = SimpleNamespace(get_draft_kv_cache_layer_names=lambda: set(draft_layers)) + proposer.max_query_tokens = 16 + proposer.max_num_tokens = 32 + kv_cache_config = SimpleNamespace( + kv_cache_groups=[ + SimpleNamespace( + layer_names=[target_layer, *draft_layers], + kv_cache_spec=mixed_spec, + ) + ] + ) + + with patch.object(AttentionGroup, "create_metadata_builders"): + proposer.initialize_attn_backend( + kv_cache_config, + kernel_block_sizes=[128], + ) + + assert len(proposer.draft_attn_groups) == 1 + assert set(proposer.draft_attn_groups[0].layer_names) == set(draft_layers) + assert proposer.draft_attn_groups[0].kv_cache_group_id == 0 + assert proposer._layer_group_idx == [0] * 5 diff --git a/tests/ut/test_compressed_prefix_cache.py b/tests/ut/test_compressed_prefix_cache.py index 2fd791fb1b20..b50786fdc264 100644 --- a/tests/ut/test_compressed_prefix_cache.py +++ b/tests/ut/test_compressed_prefix_cache.py @@ -13,6 +13,7 @@ get_block_hash, get_request_block_hasher, init_none_hash, + is_kv_cache_spec_uniform, ) from vllm.v1.core.single_type_kv_cache_manager import ( FullAttentionManager, @@ -22,6 +23,7 @@ FullAttentionSpec, KVCacheConfig, KVCacheGroupSpec, + MambaSpec, MLAAttentionSpec, ) from vllm.v1.request import Request @@ -80,6 +82,22 @@ def _make_full_manager( return spec, block_pool, manager +def test_ascend_mla_spec_is_not_uniform_with_mamba() -> None: + mla_spec, _, _ = _make_full_manager() + mamba_spec = MambaSpec( + block_size=1, + shapes=((1,),), + dtypes=(torch.float32,), + ) + + assert not is_kv_cache_spec_uniform( + { + "mla": mla_spec, + "mamba": mamba_spec, + } + ) + + @pytest.mark.parametrize("physical_block_size", [32, 64, 128]) @pytest.mark.parametrize("compress_ratio", [4, 128]) def test_compressed_spec_separates_logical_and_storage_blocks( @@ -273,3 +291,87 @@ def test_hybrid_coordinator_rejects_partial_compressed_prefix_hit() -> None: assert hit_length == 0 assert hit_blocks == ([], []) + + +def test_hybrid_coordinator_truncates_every_full_attention_group() -> None: + hash_block_size = 2 + block_size = 2 * hash_block_size + coordinator = AscendHybridKVCacheCoordinator( + kv_cache_config=KVCacheConfig( + num_blocks=32, + kv_cache_tensors=[], + kv_cache_groups=[ + KVCacheGroupSpec( + ["full_a"], + FullAttentionSpec( + block_size=block_size, + num_kv_heads=1, + head_size=1, + dtype=torch.float32, + ), + ), + KVCacheGroupSpec( + ["full_b"], + FullAttentionSpec( + block_size=2 * block_size, + num_kv_heads=1, + head_size=1, + dtype=torch.float32, + ), + ), + KVCacheGroupSpec( + ["mamba"], + MambaSpec( + block_size=block_size, + shapes=((1,),), + dtypes=(torch.float32,), + mamba_cache_mode="align", + ), + ), + ], + ), + max_model_len=8192, + use_eagle=False, + enable_caching=True, + enable_kv_cache_events=False, + dcp_world_size=1, + pcp_world_size=1, + hash_block_size=hash_block_size, + scheduler_block_size=block_size, + max_num_batched_tokens=8192, + ) + request = _make_request( + "a", + [index // hash_block_size for index in range(24)], + hash_block_size, + ) + + for group_id in (0, 1): + group_block_size = block_size * (1 + group_id) + num_full_blocks = len(request.prompt_token_ids) // group_block_size + blocks = coordinator.block_pool.get_new_blocks(num_full_blocks) + coordinator.block_pool.cache_full_blocks( + request=request, + blocks=blocks, + num_cached_blocks=0, + num_full_blocks=num_full_blocks, + block_size=group_block_size, + kv_cache_group_id=group_id, + ) + + mamba_block = coordinator.block_pool.get_new_blocks(1)[0] + coordinator.block_pool.cache_partial_block( + request=request, + block=mamba_block, + num_tokens=6, + kv_cache_group_id=2, + block_size=block_size, + ) + + hit_blocks, hit_length, _ = coordinator.find_longest_cache_hit( + request.block_hashes, + max_cache_hit_length=len(request.prompt_token_ids), + ) + + assert hit_length == 6 + assert [len(blocks) for blocks in hit_blocks] == [2, 1, 2] diff --git a/tests/ut/worker/a2/test_block_table.py b/tests/ut/worker/a2/test_block_table.py index 73c10a5da8ea..a746c25b249e 100644 --- a/tests/ut/worker/a2/test_block_table.py +++ b/tests/ut/worker/a2/test_block_table.py @@ -14,6 +14,7 @@ # import unittest +from types import SimpleNamespace from unittest.mock import MagicMock, patch import numpy as np @@ -21,6 +22,11 @@ # import vllm.utils.cpu_triton_utils as cpu_tl from vllm.distributed.parallel_state import GroupCoordinator +from vllm.v1.kv_cache_interface import ( + KVCacheGroupSpec, + MambaSpec, + UniformTypeKVCacheSpecs, +) from tests.ut.base import TestBase @@ -98,6 +104,43 @@ def test_compute_slot_mapping_draft_reserves_mtp_slots(self): self.assertEqual(block_table.slot_mapping.cpu.numel(), 128) self.assertEqual(block_table.slot_mapping.cpu[: req_indices.size].numel(), 110) + def test_uniform_mamba_group_is_recognized_as_mamba(self): + mamba_spec = MambaSpec( + block_size=self.block_size, + shapes=((4, 8),), + dtypes=(torch.float32,), + page_size_padded=128, + mamba_cache_mode="align", + num_speculative_blocks=7, + ) + layer_specs = {f"mamba.{i}": mamba_spec for i in range(3)} + uniform_spec = UniformTypeKVCacheSpecs.from_specs(layer_specs) + self.assertIsNotNone(uniform_spec) + kv_cache_group = KVCacheGroupSpec( + layer_names=list(layer_specs), + kv_cache_spec=uniform_spec, + ) + + with patch("vllm_ascend.worker.block_table.get_dcp_group") as mock_get_dcp_group: + mock_get_dcp_group.return_value = SimpleNamespace( + world_size=1, + rank_in_group=0, + ) + from vllm_ascend.worker.block_table import BlockTable + + block_table = BlockTable( + block_size=self.block_size, + max_num_reqs=self.max_num_reqs, + max_num_blocks_per_req=self.max_num_blocks_per_req, + max_num_batched_tokens=self.max_num_batched_tokens, + pin_memory=self.pin_memory, + device=self.device, + kernel_sizes=[0], + kv_cache_group=kv_cache_group, + ) + + self.assertTrue(block_table.is_mamba_group) + def setup_block_table_data(self, block_table, num_reqs=2): """Helper method to populate block table with test data""" # Add block IDs for each request diff --git a/tests/ut/worker/a2/test_model_runner_v1.py b/tests/ut/worker/a2/test_model_runner_v1.py index d99cfb308aa3..170daf9ef525 100644 --- a/tests/ut/worker/a2/test_model_runner_v1.py +++ b/tests/ut/worker/a2/test_model_runner_v1.py @@ -11,11 +11,13 @@ KVCacheConfig, KVCacheGroupSpec, KVCacheTensor, + MLAAttentionSpec, UniformTypeKVCacheSpecs, ) from vllm_ascend.attention.utils import get_sfa_qsfa_packed_head_dim from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec, AscendSFAIndexerCacheSpec +from vllm_ascend.spec_decode.dspark_proposer import AscendDSparkProposer from vllm_ascend.utils import AscendDeviceType from vllm_ascend.worker.model_runner_v1 import NPUModelRunner @@ -97,6 +99,49 @@ def _build_runner(self): runner.attn_backend = backend return runner + @patch("vllm_ascend.worker.model_runner_v1.has_kv_transfer_group", return_value=False) + @patch("vllm_ascend.worker.model_runner_v1.apply_layerwise_kv_cache_plan") + def test_drafter_receives_logical_block_size_for_every_cache_group( + self, + _mock_apply_layerwise_plan, + _mock_has_kv_transfer_group, + ): + runner = self._build_runner() + runner.attn_groups = [] + runner.model_config.enable_return_routed_experts = False + runner.speculative_config = SimpleNamespace( + use_eagle=lambda: False, + uses_draft_model=lambda: True, + uses_extract_hidden_states=lambda: False, + ) + drafter = AscendDSparkProposer.__new__(AscendDSparkProposer) + drafter.initialize_attn_backend = MagicMock() + runner.drafter = drafter + runner.may_add_encoder_only_layers_to_kv_cache_config = MagicMock() + runner.maybe_add_kv_sharing_layers_to_kv_cache_groups = MagicMock() + + def initialize_attn_backend(_kv_cache_config): + runner.attn_groups = [ + [SimpleNamespace(kv_cache_spec=object())], + [SimpleNamespace(kv_cache_spec=object())], + ] + + runner.initialize_attn_backend = MagicMock(side_effect=initialize_attn_backend) + + def reinitialize_input_batch(_kv_cache_config): + runner.kernel_block_sizes = [[0], [128]] + + runner.may_reinitialize_input_batch = MagicMock(side_effect=reinitialize_input_batch) + runner.initialize_kv_cache_tensors = MagicMock(return_value={}) + + runner.initialize_kv_cache(SimpleNamespace(kv_cache_groups=[])) + + drafter.initialize_attn_backend.assert_called_once() + self.assertEqual( + drafter.initialize_attn_backend.call_args.args[1], + [0, 128], + ) + def test_allocate_kv_cache_uses_layer_spec_for_draft_gqa(self): runner = self._build_runner() runner.sparse_kv_offload_enabled = False @@ -119,6 +164,53 @@ def test_allocate_kv_cache_uses_layer_spec_for_draft_gqa(self): self.assertEqual(k_cache_raw.numel(), kv_cache_spec.page_size_bytes) self.assertEqual(v_cache_raw.numel(), kv_cache_spec.page_size_bytes) + @patch("vllm_ascend.worker.model_runner_v1.has_ec_transfer", return_value=False) + @patch("vllm_ascend.worker.model_runner_v1.get_layers_from_vllm_config") + def test_draft_mla_uses_separate_target_kv_cache_group( + self, + mock_get_layers, + _mock_has_ec_transfer, + ): + runner = self._build_runner() + runner.block_size = 16 + runner.shared_kv_cache_layers = {} + + draft_attn = MLAAttention.__new__(MLAAttention) + torch.nn.Module.__init__(draft_attn) + draft_attn.impl = SimpleNamespace(fa_quant_layer=False) + draft_attn.get_kv_cache_spec = MagicMock( + return_value=MLAAttentionSpec( + block_size=16, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + non_causal_multi_token_decode=True, + ) + ) + mock_get_layers.return_value = {"draft.self_attn": draft_attn} + + draft_spec = runner.get_kv_cache_spec()["draft.self_attn"] + target_spec = AscendMLAAttentionSpec( + block_size=16, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + ) + + self.assertTrue(draft_spec.non_causal_multi_token_decode) + uniform_spec = UniformTypeKVCacheSpecs.from_specs( + { + "target.self_attn": target_spec, + "draft.self_attn": draft_spec, + } + ) + self.assertIsNone(uniform_spec) + with self.assertRaisesRegex( + AssertionError, + "Causal target layers and non-causal multi-token draft layers", + ): + AscendMLAAttentionSpec.merge([target_spec, draft_spec]) + def test_sparse_c8_indexer_reuses_raw_cache_from_shared_descriptor(self): runner = self._build_runner() layer_names = [ @@ -185,6 +277,91 @@ def test_reshape_kv_cache_uses_layer_spec_for_draft_gqa(self): self.assertEqual(k_cache.shape, (2, 16, 8, 64)) self.assertEqual(v_cache.shape, (2, 16, 8, 64)) + @patch("vllm_ascend.worker.model_runner_v1.get_layers_from_vllm_config") + def test_hybrid_mla_cache_uses_logical_kernel_block_shape( + self, + mock_get_layers, + ): + """A 384-token scheduler page is exposed as three 128-token blocks.""" + runner = self._build_runner() + runner.use_hybrid_blocks = True + runner.hybrid_with_attn_and_mamba = False + runner.model_config.hf_text_config = SimpleNamespace( + kv_lora_rank=512, + qk_rope_head_dim=64, + ) + + layer_name = "draft_attn" + attn_module = MLAAttention.__new__(MLAAttention) + torch.nn.Module.__init__(attn_module) + attn_module.kv_lora_rank = 512 + attn_module.qk_rope_head_dim = 64 + mock_get_layers.return_value = {layer_name: attn_module} + + physical_block_size = 384 + kernel_block_size = 128 + num_physical_blocks = 2 + kv_cache_spec = AscendMLAAttentionSpec( + block_size=physical_block_size, + num_kv_heads=1, + head_size=512 + 64, + dtype=torch.bfloat16, + ) + kv_cache_config = KVCacheConfig( + num_blocks=num_physical_blocks, + kv_cache_tensors=[ + KVCacheTensor( + size=kv_cache_spec.page_size_bytes * num_physical_blocks, + shared_by=[layer_name], + ) + ], + kv_cache_groups=[ + KVCacheGroupSpec( + layer_names=[layer_name], + kv_cache_spec=kv_cache_spec, + ) + ], + ) + + # Raw cache tensors are byte buffers. Together they contain two + # physical pages: 512 latent dimensions plus 64 RoPE dimensions. + raw_k_cache = torch.empty( + num_physical_blocks * physical_block_size * 512 * 2, + dtype=torch.uint8, + ) + raw_v_cache = torch.empty( + num_physical_blocks * physical_block_size * 64 * 2, + dtype=torch.uint8, + ) + backend = MagicMock() + backend.get_supported_kernel_block_sizes.return_value = [kernel_block_size] + backend.get_kv_cache_shape.side_effect = lambda num_blocks, block_size, num_kv_heads, head_size: ( + num_blocks, + block_size, + num_kv_heads, + head_size, + ) + runner._kv_cache_spec_attn_group_iterator = lambda: [ + SimpleNamespace( + kv_cache_spec=kv_cache_spec, + backend=backend, + layer_names=[layer_name], + ) + ] + + k_cache, v_cache = runner._reshape_kv_cache_tensors( + kv_cache_config, + {layer_name: (raw_k_cache, raw_v_cache)}, + )[layer_name] + + num_kernel_blocks = num_physical_blocks * physical_block_size // kernel_block_size + self.assertEqual(k_cache.shape, (num_kernel_blocks, 128, 1, 512)) + self.assertEqual(v_cache.shape, (num_kernel_blocks, 128, 1, 64)) + self.assertEqual( + backend.get_kv_cache_shape.call_args.args[:2], + (num_kernel_blocks, kernel_block_size), + ) + @patch("vllm_ascend.worker.model_runner_v1.has_ec_transfer", return_value=False) @patch("vllm_ascend.worker.model_runner_v1.get_layers_from_vllm_config") def test_sparse_layer_without_indexer_allocates_only_mla_kv_cache( diff --git a/vllm_ascend/core/kv_cache_interface.py b/vllm_ascend/core/kv_cache_interface.py index b740a5bbf07b..a4d11379663c 100644 --- a/vllm_ascend/core/kv_cache_interface.py +++ b/vllm_ascend/core/kv_cache_interface.py @@ -1,7 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -from dataclasses import dataclass +from dataclasses import dataclass, replace import torch from typing_extensions import Self @@ -53,7 +53,7 @@ def storage_block_size(self) -> int: return self.block_size // self.compress_ratio @property - def page_size_bytes(self) -> int: + def real_page_size_bytes(self) -> int: return ( self.storage_block_size * self.num_kv_heads @@ -65,46 +65,42 @@ def merge(cls, specs: list[Self]) -> Self: assert all(isinstance(spec, MLAAttentionSpec) for spec in specs), ( "All attention layers in the same KV cache group must be MLAAttentionSpec." ) - layout_set = { + ascend_layouts = { ( - spec.block_size, - spec.num_kv_heads, - spec.head_size, spec.scale_dim, spec.scale_dtype, - spec.dtype, - spec.compress_ratio, + spec.cache_sparse_sfa_c8, + spec.store_on_host, + spec.alignment, ) for spec in specs } - assert len(layout_set) == 1, ( - "All attention layers in the same KV cache group must use the same KV cache layout." - ) - cache_dtype_str_set = set(spec.cache_dtype_str for spec in specs) - assert len(cache_dtype_str_set) == 1, ( - "All attention layers in the same KV cache group must use the same quantization method." - ) - cache_sparse_sfa_c8_set = set(spec.cache_sparse_sfa_c8 for spec in specs) - assert len(cache_sparse_sfa_c8_set) == 1, ( - "All attention layers in the same KV cache group must use the same sparse SFA C8 setting." + assert len(ascend_layouts) == 1, ( + "All attention layers in the same KV cache group must use the same Ascend KV cache layout." ) - store_on_host_set = set(spec.store_on_host for spec in specs) - assert len(store_on_host_set) == 1, ( - "All attention layers in the same KV cache group must use the same host storage setting." + non_causal_multi_token_decode_set = set(spec.non_causal_multi_token_decode for spec in specs) + assert len(non_causal_multi_token_decode_set) == 1, ( + "Causal target layers and non-causal multi-token draft layers must use separate KV cache groups." ) - return cls( - block_size=specs[0].block_size, - num_kv_heads=specs[0].num_kv_heads, - head_size=specs[0].head_size, - scale_dim=specs[0].scale_dim, - scale_dtype=specs[0].scale_dtype, - dtype=specs[0].dtype, - cache_dtype_str=cache_dtype_str_set.pop(), - compress_ratio=specs[0].compress_ratio, - cache_sparse_sfa_c8=specs[0].cache_sparse_sfa_c8, - store_on_host=store_on_host_set.pop(), + first_spec = specs[0] + merged = super().merge(specs) + return replace( + merged, + scale_dim=first_spec.scale_dim, + scale_dtype=first_spec.scale_dtype, + alignment=first_spec.alignment, + cache_sparse_sfa_c8=first_spec.cache_sparse_sfa_c8, + store_on_host=first_spec.store_on_host, ) + def is_uniform_with_collection(self, kv_cache_specs: dict[str, KVCacheSpec]) -> bool: + if any( + getattr(spec, "non_causal_multi_token_decode", False) != self.non_causal_multi_token_decode + for spec in kv_cache_specs.values() + ): + return False + return super().is_uniform_with_collection(kv_cache_specs) + def max_memory_usage_bytes(self, vllm_config: VllmConfig) -> int: max_model_len = vllm_config.model_config.max_model_len dcp_world_size = vllm_config.parallel_config.decode_context_parallel_size diff --git a/vllm_ascend/patch/platform/patch_kv_cache_coordinator.py b/vllm_ascend/patch/platform/patch_kv_cache_coordinator.py index 981e75459df7..467da6aab184 100644 --- a/vllm_ascend/patch/platform/patch_kv_cache_coordinator.py +++ b/vllm_ascend/patch/platform/patch_kv_cache_coordinator.py @@ -590,6 +590,7 @@ def get_kv_cache_coordinator( # type: ignore[misc] # compatibility; platform validation guarantees that it is one. del pcp_world_size token_budget = _select_kv_token_budget(max_model_len, max_in_flight_tokens, max_num_batched_tokens) + prefill_lookahead = 0 if num_prefill_lookahead is None else num_prefill_lookahead if _is_deepseek_v4_kv_cache_config(kv_cache_config): return AscendHybridKVCacheCoordinator( # type: ignore[call-arg] kv_cache_config, @@ -605,7 +606,7 @@ def get_kv_cache_coordinator( # type: ignore[misc] max_in_flight_tokens=token_budget, max_num_batched_tokens=token_budget, scheduler_block_size=scheduler_block_size, - num_prefill_lookahead=num_prefill_lookahead, + num_prefill_lookahead=prefill_lookahead, ) if len(kv_cache_config.kv_cache_groups) == 1 or not enable_caching: @@ -640,7 +641,7 @@ def get_kv_cache_coordinator( # type: ignore[misc] max_in_flight_tokens=token_budget, max_num_batched_tokens=token_budget, scheduler_block_size=scheduler_block_size, - num_prefill_lookahead=num_prefill_lookahead, + num_prefill_lookahead=prefill_lookahead, ) diff --git a/vllm_ascend/patch/platform/patch_kv_cache_utils.py b/vllm_ascend/patch/platform/patch_kv_cache_utils.py index 068bf777a0fa..734277d590f4 100644 --- a/vllm_ascend/patch/platform/patch_kv_cache_utils.py +++ b/vllm_ascend/patch/platform/patch_kv_cache_utils.py @@ -5,19 +5,37 @@ import vllm.v1.core.kv_cache_utils from vllm.config import VllmConfig +from vllm.logger import logger from vllm.utils.math_utils import cdiv, round_up from vllm.v1.core.kv_cache_utils import _approximate_gcd, may_override_num_blocks from vllm.v1.kv_cache_interface import ( + FullAttentionSpec, KVCacheConfig, KVCacheGroupSpec, KVCacheSpec, + KVCacheSpecKind, KVCacheTensor, + MambaSpec, MLAAttentionSpec, SlidingWindowMLASpec, UniformTypeKVCacheSpecs, + get_kv_cache_spec_kind, ) +_KIMI_K3_TARGET_LAYER_PREFIX = "language_model.model.layers." +_KIMI_K3_DRAFT_LAYER_PREFIX = "model.layers." +_ATTENTION_LAYER_SUFFIX = ".self_attn.attn" +_MAMBA_LAYER_SUFFIX = ".self_attn" + _orig_resolve_kv_cache_block_sizes = vllm.v1.core.kv_cache_utils.resolve_kv_cache_block_sizes +_orig_get_kv_cache_groups_uniform_page_size = vllm.v1.core.kv_cache_utils._get_kv_cache_groups_uniform_page_size + + +class _MambaUniformKVCacheSpecs(UniformTypeKVCacheSpecs): + """Preserve Mamba-specific worker block-table capacity after grouping.""" + + def max_num_blocks_per_req(self, vllm_config: VllmConfig, max_len: int) -> int: + return max(spec.max_num_blocks_per_req(vllm_config, max_len) for spec in self.kv_cache_specs.values()) def _ascend_resolve_kv_cache_block_sizes( @@ -56,6 +74,133 @@ def _ascend_resolve_kv_cache_block_sizes( return _orig_resolve_kv_cache_block_sizes(kv_cache_config, vllm_config) +def _is_kimi_k3_target_attention_layer(layer_name: str, spec: KVCacheSpec) -> bool: + return ( + layer_name.startswith(_KIMI_K3_TARGET_LAYER_PREFIX) + and layer_name.endswith(_ATTENTION_LAYER_SUFFIX) + and isinstance(spec, FullAttentionSpec) + and not spec.non_causal + and spec.sliding_window is None + and spec.attention_chunk_size is None + and not getattr(spec, "non_causal_multi_token_decode", False) + ) + + +def _is_kimi_k3_dspark_attention_layer(layer_name: str, spec: KVCacheSpec) -> bool: + return ( + layer_name.startswith(_KIMI_K3_DRAFT_LAYER_PREFIX) + and layer_name.endswith(_ATTENTION_LAYER_SUFFIX) + and isinstance(spec, FullAttentionSpec) + and not spec.non_causal + and spec.sliding_window is None + and spec.attention_chunk_size is None + ) + + +def _is_kimi_k3_mamba_layer(layer_name: str, spec: KVCacheSpec) -> bool: + return ( + layer_name.startswith(_KIMI_K3_TARGET_LAYER_PREFIX) + and layer_name.endswith(_MAMBA_LAYER_SUFFIX) + and not layer_name.endswith(_ATTENTION_LAYER_SUFFIX) + and isinstance(spec, MambaSpec) + and spec.mamba_cache_mode == "align" + and spec.num_speculative_blocks > 0 + ) + + +def _get_kimi_k3_dspark_mixed_kv_cache_groups( + kv_cache_spec: dict[str, KVCacheSpec], +) -> list[KVCacheGroupSpec] | None: + """Build topology-independent Kimi K3 DSpark scheduler groups. + + Target and causal draft attention layers require the same full-sequence + block ownership. Putting them in one UniformType group lets them share one + scheduler block table while preserving a separate physical page per layer. + Recurrent layers are split into the fewest balanced groups whose size does + not exceed the attention group. This minimizes scheduler groups while + keeping the recurrent groups balanced. + + Block and page sizes are resolved by the runtime and intentionally not + fixed here: TP8 and TP16 produce different sizes but the same ownership + relation. A partial or incompatible signature falls back to vLLM's generic + hybrid grouping. + """ + target_attention_specs = { + name: spec for name, spec in kv_cache_spec.items() if _is_kimi_k3_target_attention_layer(name, spec) + } + draft_attention_specs = { + name: spec for name, spec in kv_cache_spec.items() if _is_kimi_k3_dspark_attention_layer(name, spec) + } + mamba_specs = {name: spec for name, spec in kv_cache_spec.items() if _is_kimi_k3_mamba_layer(name, spec)} + + matched_layer_count = len(target_attention_specs) + len(draft_attention_specs) + len(mamba_specs) + if ( + not target_attention_specs + or not draft_attention_specs + or not mamba_specs + or matched_layer_count != len(kv_cache_spec) + ): + return None + + all_specs = [*target_attention_specs.values(), *draft_attention_specs.values(), *mamba_specs.values()] + if len({spec.block_size for spec in all_specs}) != 1 or len({spec.page_size_bytes for spec in all_specs}) != 1: + return None + + first_mamba_spec = next(iter(mamba_specs.values())) + if any(spec != first_mamba_spec for spec in mamba_specs.values()): + return None + + # Insert target attention first. generate_scheduler_kv_cache_config unwraps a + # UniformType group to its first spec, and this representative is registered + # with the FullAttentionManager needed by both target and draft attention. + mixed_attention_specs = {**target_attention_specs, **draft_attention_specs} + mixed_attention_spec = UniformTypeKVCacheSpecs.from_specs(mixed_attention_specs) + if mixed_attention_spec is None: + return None + + groups = [ + KVCacheGroupSpec( + layer_names=list(mixed_attention_specs), + kv_cache_spec=mixed_attention_spec, + ) + ] + mamba_layer_names = list(mamba_specs) + mamba_group_count = cdiv(len(mamba_layer_names), len(mixed_attention_specs)) + for group_idx in range(mamba_group_count): + layer_names = mamba_layer_names[group_idx::mamba_group_count] + group_specs = {name: mamba_specs[name] for name in layer_names} + uniform_mamba_spec = _MambaUniformKVCacheSpecs.from_specs(group_specs) + assert uniform_mamba_spec is not None + groups.append( + KVCacheGroupSpec( + layer_names=layer_names, + kv_cache_spec=uniform_mamba_spec, + ) + ) + + logger.info( + "Using Kimi K3 DSpark mixed KV grouping: %d target + %d draft attention layers, followed by Mamba groups %s", + len(target_attention_specs), + len(draft_attention_specs), + [len(group.layer_names) for group in groups[1:]], + ) + return groups + + +def _get_kv_cache_groups_uniform_page_size( + kv_cache_spec: dict[str, KVCacheSpec], +) -> list[KVCacheGroupSpec]: + kimi_k3_groups = _get_kimi_k3_dspark_mixed_kv_cache_groups(kv_cache_spec) + if kimi_k3_groups is not None: + return kimi_k3_groups + return _orig_get_kv_cache_groups_uniform_page_size(kv_cache_spec) + + +def _kv_cache_config_has_mamba_layers(self: KVCacheConfig) -> bool: + """Recognize Mamba layers nested in UniformType cache groups.""" + return any(get_kv_cache_spec_kind(group.kv_cache_spec) == KVCacheSpecKind.MAMBA for group in self.kv_cache_groups) + + def group_and_unify_kv_cache_specs( kv_cache_spec: dict[str, KVCacheSpec], ) -> list[UniformTypeKVCacheSpecs] | None: @@ -248,10 +393,14 @@ def _get_kv_cache_config_deepseek_v4( vllm.v1.core.kv_cache_utils.resolve_kv_cache_block_sizes = _ascend_resolve_kv_cache_block_sizes vllm.v1.core.kv_cache_utils.group_and_unify_kv_cache_specs = group_and_unify_kv_cache_specs vllm.v1.core.kv_cache_utils._get_kv_cache_groups_uniform_groups = _get_kv_cache_groups_uniform_groups +vllm.v1.core.kv_cache_utils._get_kv_cache_groups_uniform_page_size = _get_kv_cache_groups_uniform_page_size # vLLM v0.24.0 renamed _get_kv_cache_config_deepseek_v4 to _get_kv_cache_config_packed and # get_kv_cache_config_from_groups now calls _get_kv_cache_config_packed directly, bypassing # the alias patch above. Patch the canonical name so Ascend's non-packed layout is used. vllm.v1.core.kv_cache_utils._get_kv_cache_config_packed = _get_kv_cache_config_deepseek_v4 +KVCacheConfig.has_mamba_layers = property( # type: ignore[assignment] + _kv_cache_config_has_mamba_layers +) # Also patch the reference used by engine/core.py which imports the function directly. import vllm.v1.engine.core # noqa: E402 diff --git a/vllm_ascend/patch/worker/patch_mamba_utils.py b/vllm_ascend/patch/worker/patch_mamba_utils.py index 3b241ec965f6..a16fa0a22a4c 100644 --- a/vllm_ascend/patch/worker/patch_mamba_utils.py +++ b/vllm_ascend/patch/worker/patch_mamba_utils.py @@ -8,7 +8,11 @@ from vllm.model_executor.layers.mamba.mamba_utils import MambaStateCopyFunc from vllm.utils.math_utils import cdiv from vllm.v1.core.sched.output import SchedulerOutput -from vllm.v1.kv_cache_interface import KVCacheConfig +from vllm.v1.kv_cache_interface import ( + KVCacheConfig, + MambaSpec, + UniformTypeKVCacheSpecs, +) from vllm.v1.worker import mamba_utils from vllm.v1.worker.gpu_input_batch import CachedRequestState from vllm.v1.worker.lora_model_runner_mixin import GPUInputBatch @@ -31,6 +35,31 @@ def _can_launch_triton_batch_memcpy() -> bool: return not is_310p() +def _get_mamba_groups( + kv_cache_config: KVCacheConfig, +) -> tuple[list[int], MambaSpec]: + """Find Mamba groups, including uniform worker-side group wrappers.""" + mamba_group_ids: list[int] = [] + mamba_specs: list[MambaSpec] = [] + for group_id, group in enumerate(kv_cache_config.kv_cache_groups): + group_spec = group.kv_cache_spec + if isinstance(group_spec, MambaSpec): + mamba_group_ids.append(group_id) + mamba_specs.append(group_spec) + continue + if not isinstance(group_spec, UniformTypeKVCacheSpecs): + continue + + inner_specs = list(group_spec.kv_cache_specs.values()) + if inner_specs and all(isinstance(spec, MambaSpec) for spec in inner_specs): + mamba_group_ids.append(group_id) + mamba_specs.append(inner_specs[0]) + + assert mamba_group_ids, "no mamba layers in the model" + assert all(mamba_specs[0] == spec for spec in mamba_specs) + return mamba_group_ids, mamba_specs[0] + + def _batch_memcpy_triton(src_ptrs, dst_ptrs, sizes): batch = src_ptrs.shape[0] assert dst_ptrs.shape[0] == batch @@ -233,6 +262,11 @@ def _batch_memcpy_unavailable(src_ptrs, dst_ptrs, sizes): mamba_utils.do_mamba_copy_block = _do_mamba_copy_block_torch mamba_utils.postprocess_mamba_align_gpu = _postprocess_mamba_align_gpu_cpu_fallback +# Worker KV configs retain UniformTypeKVCacheSpecs so per-layer physical page +# layouts are available while the scheduler receives unwrapped representative +# specs. Teach all upstream Mamba buffer/context helpers to see those groups. +mamba_utils.get_mamba_groups = _get_mamba_groups + # Ascend NPU does not support DT_UINT64 in aclnnInplaceZero. # MambaCopyBuffers.create() uses torch.uint64 for src_ptrs/dst_ptrs, # which triggers a runtime error. Remap to int64 at the source. diff --git a/vllm_ascend/worker/block_table.py b/vllm_ascend/worker/block_table.py index 42e3a805ef82..b200c1bbeb70 100644 --- a/vllm_ascend/worker/block_table.py +++ b/vllm_ascend/worker/block_table.py @@ -3,7 +3,11 @@ from vllm.distributed import get_dcp_group from vllm.utils.math_utils import cdiv from vllm.v1.attention.backends.utils import PAD_SLOT_ID -from vllm.v1.kv_cache_interface import KVCacheGroupSpec, MambaSpec +from vllm.v1.kv_cache_interface import ( + KVCacheGroupSpec, + KVCacheSpecKind, + get_kv_cache_spec_kind, +) from vllm.v1.utils import CpuGpuBuffer from vllm_ascend.distributed.utils import get_decode_context_model_parallel_world_size @@ -30,23 +34,19 @@ def __init__( self.max_num_reqs = max_num_reqs self.dcp_world_size = get_dcp_group().world_size self.dcp_rank = get_dcp_group().rank_in_group - if ( + is_mamba_group = ( kv_cache_group is not None and hasattr(kv_cache_group, "kv_cache_spec") - and self.dcp_world_size > 1 - and isinstance(kv_cache_group.kv_cache_spec, MambaSpec) - ): + and get_kv_cache_spec_kind(kv_cache_group.kv_cache_spec) == KVCacheSpecKind.MAMBA + ) + if self.dcp_world_size > 1 and is_mamba_group: max_num_blocks_per_req = max_num_blocks_per_req * self.dcp_world_size self.max_num_blocks_per_req = max_num_blocks_per_req self.max_num_batched_tokens = max_num_batched_tokens self.pin_memory = pin_memory self.device = device self.physical_block_size = block_size - self.is_mamba_group = ( - kv_cache_group is not None - and hasattr(kv_cache_group, "kv_cache_spec") - and isinstance(kv_cache_group.kv_cache_spec, MambaSpec) - ) + self.is_mamba_group = is_mamba_group # If kernel_sizes is None or [0], use physical block size (no splitting) if kernel_sizes is None or kernel_sizes == [0]: diff --git a/vllm_ascend/worker/model_runner_v1.py b/vllm_ascend/worker/model_runner_v1.py index c7d29b770e27..f1adceb2a79b 100644 --- a/vllm_ascend/worker/model_runner_v1.py +++ b/vllm_ascend/worker/model_runner_v1.py @@ -143,7 +143,6 @@ reshape_kv_cache_tensors_for_sparse_kv_offload, update_sparse_kv_offload_metadata, ) -from vllm_ascend.distributed.utils import get_decode_context_model_parallel_world_size from vllm_ascend.eplb.adaptor.vllm_adaptor import VllmEplbAdaptor from vllm_ascend.eplb.core.eplb_device_transfer_loader import D2DExpertWeightLoader from vllm_ascend.eplb.core.eplb_worker import EplbProcess @@ -3690,10 +3689,7 @@ def initialize_kv_cache(self, kv_cache_config: KVCacheConfig) -> None: # NOTE(cmq): initialize_attn_backend must before using self.attn_groups self.initialize_attn_backend(kv_cache_config) self.use_hybrid_blocks = len(self.attn_groups) > 1 - # NOTE: Currently, we determine whether we need `num_accepted_tokens` through `MambaSpec`. - self.need_accepted_tokens = any( - [isinstance(attn_group[0].kv_cache_spec, MambaSpec) for attn_group in self.attn_groups] - ) + self.need_accepted_tokens = kv_cache_config.has_mamba_layers self.may_reinitialize_input_batch(kv_cache_config) if self.sparse_kv_offload_enabled: @@ -3716,10 +3712,26 @@ def initialize_kv_cache(self, kv_cache_config: KVCacheConfig) -> None: self.drafter, AscendEagleProposer | AscendDflashProposer | AscendDSparkProposer | AscendDraftModelProposer, ) - block_size = (self.kernel_block_sizes[0] if isinstance( - self.kernel_block_sizes, list) else self.kernel_block_sizes) - self.drafter.initialize_attn_backend(kv_cache_config, block_size) - + if isinstance(self.drafter, AscendDSparkProposer): + if isinstance(self.kernel_block_sizes, list): + draft_kernel_block_sizes = [ + int(sizes[0] if isinstance(sizes, (list, tuple)) else sizes) + for sizes in self.kernel_block_sizes + ] + else: + draft_kernel_block_sizes = [int(self.kernel_block_sizes)] + self.drafter.initialize_attn_backend( + kv_cache_config, + draft_kernel_block_sizes, + ) + else: + block_size = ( + self.kernel_block_sizes[0] + if isinstance(self.kernel_block_sizes, list) + else self.kernel_block_sizes + ) + self.drafter.initialize_attn_backend(kv_cache_config, block_size) + if ( self.speculative_config and self.speculative_config.uses_extract_hidden_states() @@ -4551,18 +4563,10 @@ def may_reinitialize_input_batch(self, kv_cache_config: KVCacheConfig) -> None: max_num_blocks = [] max_model_len = max(self.max_model_len, self.max_encoder_len) for kv_cache_group in non_encoder_groups: - max_num_blocks_per_req = cdiv( + max_num_blocks_per_req = kv_cache_group.kv_cache_spec.max_num_blocks_per_req( + self.vllm_config, max_model_len, - kv_cache_group.kv_cache_spec.block_size - * get_decode_context_model_parallel_world_size(), ) - if isinstance(kv_cache_group.kv_cache_spec, MambaSpec): - mamba_blocks_per_req = ( - max_num_blocks_per_req if self.cache_config.enable_prefix_caching else 1 - ) - - max_num_blocks_per_req = max(max_num_blocks_per_req, mamba_blocks_per_req) - max_num_blocks_per_req += kv_cache_group.kv_cache_spec.num_speculative_blocks max_num_blocks.append(max_num_blocks_per_req) if (block_sizes != [self.cache_config.block_size] @@ -4777,6 +4781,7 @@ def get_kv_cache_spec(self) -> dict[str, KVCacheSpec]: head_size=head_size, dtype=dtype, cache_dtype_str=cache_dtype_str, + non_causal_multi_token_decode=spec.non_causal_multi_token_decode, ) attn_layer_names.add(layer_name) From 8f2cb3108218cf557e2a0dfb2cae9320fde62e06 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Tue, 25 Aug 2026 03:53:36 -0500 Subject: [PATCH 14/50] refactor(kv-cache): address K3 grouping review Signed-off-by: maoxx241 --- .../platform/test_prefix_cache_cp_patches.py | 59 ++++++++++++++---- vllm_ascend/core/kv_cache_interface.py | 8 --- .../platform/patch_kv_cache_coordinator.py | 10 ++-- .../patch/platform/patch_kv_cache_utils.py | 60 ++++--------------- vllm_ascend/worker/model_runner_v1.py | 26 ++++---- 5 files changed, 74 insertions(+), 89 deletions(-) diff --git a/tests/ut/patch/platform/test_prefix_cache_cp_patches.py b/tests/ut/patch/platform/test_prefix_cache_cp_patches.py index 2bce0fb1e0ba..a992b9dc757b 100644 --- a/tests/ut/patch/platform/test_prefix_cache_cp_patches.py +++ b/tests/ut/patch/platform/test_prefix_cache_cp_patches.py @@ -25,6 +25,7 @@ ) from vllm.v1.kv_cache_spec_registry import KVCacheSpecRegistry +from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec from vllm_ascend.patch.platform.patch_kv_cache_coordinator import ( AscendHybridKVCacheCoordinator, _is_deepseek_v4_kv_cache_spec, @@ -77,7 +78,7 @@ def _make_kimi_k3_dspark_kv_cache_specs( mamba_layer_count: int = 69, draft_uses_mla: bool = False, ) -> dict: - target_mla_spec = MLAAttentionSpec( + target_mla_spec = AscendMLAAttentionSpec( block_size=block_size, num_kv_heads=1, head_size=576, @@ -86,7 +87,7 @@ def _make_kimi_k3_dspark_kv_cache_specs( cache_dtype_str="auto", ) if draft_uses_mla: - draft_attention_spec = MLAAttentionSpec( + draft_attention_spec = AscendMLAAttentionSpec( block_size=block_size, num_kv_heads=1, head_size=576, @@ -192,6 +193,46 @@ def _make_coordinator_for_effective_block_size( return coordinator +def test_ascend_mla_page_size_includes_scale_storage() -> None: + spec = AscendMLAAttentionSpec( + block_size=16, + num_kv_heads=1, + head_size=128, + dtype=torch.bfloat16, + scale_dim=1, + scale_dtype=torch.float16, + ) + + expected_page_size = 16 * (128 * 2 + 2) + assert spec.unpadded_page_size_bytes == expected_page_size + assert spec.real_page_size_bytes == expected_page_size + assert spec.page_size_bytes == expected_page_size + + +def test_ascend_mla_merge_preserves_upstream_layout_fields() -> None: + spec = AscendMLAAttentionSpec( + block_size=512, + num_kv_heads=1, + head_size=128, + dtype=torch.bfloat16, + cache_dtype_str="fp8_ds_mla", + compress_ratio=4, + model_version="deepseek_v4", + indexes_kv_by_block_stride=True, + scale_dim=1, + scale_dtype=torch.float16, + ) + + merged = AscendMLAAttentionSpec.merge([spec, replace(spec)]) + + assert merged.block_size == spec.block_size + assert merged.compress_ratio == spec.compress_ratio + assert merged.model_version == spec.model_version + assert merged.indexes_kv_by_block_stride == spec.indexes_kv_by_block_stride + assert merged.scale_dim == spec.scale_dim + assert merged.scale_dtype == spec.scale_dtype + + @pytest.mark.parametrize( ("enable_prefix_caching", "expected_hash_block_size"), [ @@ -361,7 +402,7 @@ def test_kimi_k3_gqa_mixed_groups_use_expected_physical_layout(monkeypatch) -> N def test_kimi_k3_gqa_mixed_grouping_falls_back_on_partial_signature() -> None: specs = _make_kimi_k3_dspark_kv_cache_specs() draft_layer = "model.layers.93.self_attn.attn" - specs[draft_layer] = replace(specs[draft_layer], non_causal=True) + specs.pop(draft_layer) assert _get_kimi_k3_dspark_mixed_kv_cache_groups(specs) is None @@ -518,14 +559,10 @@ def _fake_orig(*args, **kwargs): assert coordinator is sentinel -@pytest.mark.parametrize( - ("num_prefill_lookahead", "expected"), - [(None, 0), (8, 8)], -) -def test_get_kv_cache_coordinator_normalizes_prefill_lookahead( +@pytest.mark.parametrize("num_prefill_lookahead", [0, 8]) +def test_get_kv_cache_coordinator_forwards_prefill_lookahead( monkeypatch, - num_prefill_lookahead: int | None, - expected: int, + num_prefill_lookahead: int, ) -> None: kv_cache_config = _make_hybrid_kv_cache_config( full_block_size=16, @@ -555,7 +592,7 @@ def _fake_ascend_coordinator(*args, **kwargs): num_prefill_lookahead=num_prefill_lookahead, ) - assert captured_kwargs["num_prefill_lookahead"] == expected + assert captured_kwargs["num_prefill_lookahead"] == num_prefill_lookahead def test_get_kv_cache_coordinator_uses_ascend_for_deepseek_v4(monkeypatch) -> None: diff --git a/vllm_ascend/core/kv_cache_interface.py b/vllm_ascend/core/kv_cache_interface.py index a4d11379663c..cab6c85bae84 100644 --- a/vllm_ascend/core/kv_cache_interface.py +++ b/vllm_ascend/core/kv_cache_interface.py @@ -93,14 +93,6 @@ def merge(cls, specs: list[Self]) -> Self: store_on_host=first_spec.store_on_host, ) - def is_uniform_with_collection(self, kv_cache_specs: dict[str, KVCacheSpec]) -> bool: - if any( - getattr(spec, "non_causal_multi_token_decode", False) != self.non_causal_multi_token_decode - for spec in kv_cache_specs.values() - ): - return False - return super().is_uniform_with_collection(kv_cache_specs) - def max_memory_usage_bytes(self, vllm_config: VllmConfig) -> int: max_model_len = vllm_config.model_config.max_model_len dcp_world_size = vllm_config.parallel_config.decode_context_parallel_size diff --git a/vllm_ascend/patch/platform/patch_kv_cache_coordinator.py b/vllm_ascend/patch/platform/patch_kv_cache_coordinator.py index 467da6aab184..c74915923192 100644 --- a/vllm_ascend/patch/platform/patch_kv_cache_coordinator.py +++ b/vllm_ascend/patch/platform/patch_kv_cache_coordinator.py @@ -584,13 +584,12 @@ def get_kv_cache_coordinator( # type: ignore[misc] eagle_attn_layer_names: list[str] | None = None, metrics_collector: KVCacheMetricsCollector | None = None, max_num_batched_tokens: int | None = None, - num_prefill_lookahead: int | None = None, + num_prefill_lookahead: int = 0, ) -> KVCacheCoordinator: # Keep pcp_world_size in this patched function for upstream call # compatibility; platform validation guarantees that it is one. del pcp_world_size token_budget = _select_kv_token_budget(max_model_len, max_in_flight_tokens, max_num_batched_tokens) - prefill_lookahead = 0 if num_prefill_lookahead is None else num_prefill_lookahead if _is_deepseek_v4_kv_cache_config(kv_cache_config): return AscendHybridKVCacheCoordinator( # type: ignore[call-arg] kv_cache_config, @@ -606,7 +605,7 @@ def get_kv_cache_coordinator( # type: ignore[misc] max_in_flight_tokens=token_budget, max_num_batched_tokens=token_budget, scheduler_block_size=scheduler_block_size, - num_prefill_lookahead=prefill_lookahead, + num_prefill_lookahead=num_prefill_lookahead, ) if len(kv_cache_config.kv_cache_groups) == 1 or not enable_caching: @@ -623,8 +622,7 @@ def get_kv_cache_coordinator( # type: ignore[misc] ) orig_kwargs["max_in_flight_tokens"] = token_budget orig_kwargs["scheduler_block_size"] = scheduler_block_size - if num_prefill_lookahead is not None: - orig_kwargs["num_prefill_lookahead"] = num_prefill_lookahead + orig_kwargs["num_prefill_lookahead"] = num_prefill_lookahead return _orig_get_kv_cache_coordinator(**orig_kwargs) return AscendHybridKVCacheCoordinator( # type: ignore[call-arg] @@ -641,7 +639,7 @@ def get_kv_cache_coordinator( # type: ignore[misc] max_in_flight_tokens=token_budget, max_num_batched_tokens=token_budget, scheduler_block_size=scheduler_block_size, - num_prefill_lookahead=prefill_lookahead, + num_prefill_lookahead=num_prefill_lookahead, ) diff --git a/vllm_ascend/patch/platform/patch_kv_cache_utils.py b/vllm_ascend/patch/platform/patch_kv_cache_utils.py index 734277d590f4..6dc24f26e2d3 100644 --- a/vllm_ascend/patch/platform/patch_kv_cache_utils.py +++ b/vllm_ascend/patch/platform/patch_kv_cache_utils.py @@ -24,20 +24,10 @@ _KIMI_K3_TARGET_LAYER_PREFIX = "language_model.model.layers." _KIMI_K3_DRAFT_LAYER_PREFIX = "model.layers." -_ATTENTION_LAYER_SUFFIX = ".self_attn.attn" -_MAMBA_LAYER_SUFFIX = ".self_attn" - _orig_resolve_kv_cache_block_sizes = vllm.v1.core.kv_cache_utils.resolve_kv_cache_block_sizes _orig_get_kv_cache_groups_uniform_page_size = vllm.v1.core.kv_cache_utils._get_kv_cache_groups_uniform_page_size -class _MambaUniformKVCacheSpecs(UniformTypeKVCacheSpecs): - """Preserve Mamba-specific worker block-table capacity after grouping.""" - - def max_num_blocks_per_req(self, vllm_config: VllmConfig, max_len: int) -> int: - return max(spec.max_num_blocks_per_req(vllm_config, max_len) for spec in self.kv_cache_specs.values()) - - def _ascend_resolve_kv_cache_block_sizes( kv_cache_config: KVCacheConfig, vllm_config: VllmConfig, @@ -74,40 +64,6 @@ def _ascend_resolve_kv_cache_block_sizes( return _orig_resolve_kv_cache_block_sizes(kv_cache_config, vllm_config) -def _is_kimi_k3_target_attention_layer(layer_name: str, spec: KVCacheSpec) -> bool: - return ( - layer_name.startswith(_KIMI_K3_TARGET_LAYER_PREFIX) - and layer_name.endswith(_ATTENTION_LAYER_SUFFIX) - and isinstance(spec, FullAttentionSpec) - and not spec.non_causal - and spec.sliding_window is None - and spec.attention_chunk_size is None - and not getattr(spec, "non_causal_multi_token_decode", False) - ) - - -def _is_kimi_k3_dspark_attention_layer(layer_name: str, spec: KVCacheSpec) -> bool: - return ( - layer_name.startswith(_KIMI_K3_DRAFT_LAYER_PREFIX) - and layer_name.endswith(_ATTENTION_LAYER_SUFFIX) - and isinstance(spec, FullAttentionSpec) - and not spec.non_causal - and spec.sliding_window is None - and spec.attention_chunk_size is None - ) - - -def _is_kimi_k3_mamba_layer(layer_name: str, spec: KVCacheSpec) -> bool: - return ( - layer_name.startswith(_KIMI_K3_TARGET_LAYER_PREFIX) - and layer_name.endswith(_MAMBA_LAYER_SUFFIX) - and not layer_name.endswith(_ATTENTION_LAYER_SUFFIX) - and isinstance(spec, MambaSpec) - and spec.mamba_cache_mode == "align" - and spec.num_speculative_blocks > 0 - ) - - def _get_kimi_k3_dspark_mixed_kv_cache_groups( kv_cache_spec: dict[str, KVCacheSpec], ) -> list[KVCacheGroupSpec] | None: @@ -126,12 +82,20 @@ def _get_kimi_k3_dspark_mixed_kv_cache_groups( hybrid grouping. """ target_attention_specs = { - name: spec for name, spec in kv_cache_spec.items() if _is_kimi_k3_target_attention_layer(name, spec) + name: spec + for name, spec in kv_cache_spec.items() + if name.startswith(_KIMI_K3_TARGET_LAYER_PREFIX) and isinstance(spec, FullAttentionSpec) } draft_attention_specs = { - name: spec for name, spec in kv_cache_spec.items() if _is_kimi_k3_dspark_attention_layer(name, spec) + name: spec + for name, spec in kv_cache_spec.items() + if name.startswith(_KIMI_K3_DRAFT_LAYER_PREFIX) and isinstance(spec, FullAttentionSpec) + } + mamba_specs = { + name: spec + for name, spec in kv_cache_spec.items() + if name.startswith(_KIMI_K3_TARGET_LAYER_PREFIX) and isinstance(spec, MambaSpec) } - mamba_specs = {name: spec for name, spec in kv_cache_spec.items() if _is_kimi_k3_mamba_layer(name, spec)} matched_layer_count = len(target_attention_specs) + len(draft_attention_specs) + len(mamba_specs) if ( @@ -169,7 +133,7 @@ def _get_kimi_k3_dspark_mixed_kv_cache_groups( for group_idx in range(mamba_group_count): layer_names = mamba_layer_names[group_idx::mamba_group_count] group_specs = {name: mamba_specs[name] for name in layer_names} - uniform_mamba_spec = _MambaUniformKVCacheSpecs.from_specs(group_specs) + uniform_mamba_spec = UniformTypeKVCacheSpecs.from_specs(group_specs) assert uniform_mamba_spec is not None groups.append( KVCacheGroupSpec( diff --git a/vllm_ascend/worker/model_runner_v1.py b/vllm_ascend/worker/model_runner_v1.py index f1adceb2a79b..dcfe417988a6 100644 --- a/vllm_ascend/worker/model_runner_v1.py +++ b/vllm_ascend/worker/model_runner_v1.py @@ -3689,6 +3689,8 @@ def initialize_kv_cache(self, kv_cache_config: KVCacheConfig) -> None: # NOTE(cmq): initialize_attn_backend must before using self.attn_groups self.initialize_attn_backend(kv_cache_config) self.use_hybrid_blocks = len(self.attn_groups) > 1 + # K3's packed layout keeps Mamba specs inside UniformType groups, so the + # old first-spec scan over attn_groups cannot recognize them. self.need_accepted_tokens = kv_cache_config.has_mamba_layers self.may_reinitialize_input_batch(kv_cache_config) @@ -3712,25 +3714,17 @@ def initialize_kv_cache(self, kv_cache_config: KVCacheConfig) -> None: self.drafter, AscendEagleProposer | AscendDflashProposer | AscendDSparkProposer | AscendDraftModelProposer, ) + kernel_block_sizes = self.kernel_block_sizes if isinstance(self.drafter, AscendDSparkProposer): - if isinstance(self.kernel_block_sizes, list): - draft_kernel_block_sizes = [ - int(sizes[0] if isinstance(sizes, (list, tuple)) else sizes) - for sizes in self.kernel_block_sizes - ] - else: - draft_kernel_block_sizes = [int(self.kernel_block_sizes)] - self.drafter.initialize_attn_backend( - kv_cache_config, - draft_kernel_block_sizes, - ) + sizes = kernel_block_sizes if isinstance(kernel_block_sizes, list) else [kernel_block_sizes] + draft_kernel_block_sizes = [ + int(size[0] if isinstance(size, (list, tuple)) else size) for size in sizes + ] else: - block_size = ( - self.kernel_block_sizes[0] - if isinstance(self.kernel_block_sizes, list) - else self.kernel_block_sizes + draft_kernel_block_sizes = ( + kernel_block_sizes[0] if isinstance(kernel_block_sizes, list) else kernel_block_sizes ) - self.drafter.initialize_attn_backend(kv_cache_config, block_size) + self.drafter.initialize_attn_backend(kv_cache_config, draft_kernel_block_sizes) if ( self.speculative_config From 0e83f7121cdc5fcca2ee0ace558383a6b7901af1 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Mon, 24 Aug 2026 04:32:43 -0500 Subject: [PATCH 15/50] fix(pd): harden Kimi K3 disaggregated serving Handle Mooncake transfer completion and proxy routing for disaggregated prefill, and classify one-token prompt handoff batches as prefill until the prompt is complete. Cover connector state transitions, proxy behavior, and graph selection. Signed-off-by: maoxx241 --- .github/workflows/scripts/test_config.yaml | 2 + ..._balance_proxy_layerwise_server_example.py | 5 +- .../load_balance_proxy_server_example.py | 3 +- .../disaggregated_prefill_v1/proxy_utils.py | 11 + .../ut/kv_offload/test_mooncake_connector.py | 439 +++++++++++++++++- tests/ut/test_disaggregated_prefill_proxy.py | 34 ++ .../a2/test_model_runner_v1_with_device.py | 34 ++ .../kv_transfer/kv_p2p/mooncake_connector.py | 200 +++++--- vllm_ascend/worker/model_runner_v1.py | 9 +- 9 files changed, 671 insertions(+), 66 deletions(-) create mode 100644 examples/disaggregated_prefill_v1/proxy_utils.py create mode 100644 tests/ut/test_disaggregated_prefill_proxy.py diff --git a/.github/workflows/scripts/test_config.yaml b/.github/workflows/scripts/test_config.yaml index ea08d20fd912..eb5c5523af27 100644 --- a/.github/workflows/scripts/test_config.yaml +++ b/.github/workflows/scripts/test_config.yaml @@ -223,9 +223,11 @@ - name: distributed optional: false source_file_dependencies: + - examples/disaggregated_prefill_v1 - vllm_ascend/distributed - vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/attention_fence.py tests: + - tests/ut/test_disaggregated_prefill_proxy.py - tests/ut/distributed - tests/e2e/pull_request/two_card/test_data_parallel.py - tests/e2e/pull_request/four_card/test_data_parallel_tp2.py diff --git a/examples/disaggregated_prefill_v1/load_balance_proxy_layerwise_server_example.py b/examples/disaggregated_prefill_v1/load_balance_proxy_layerwise_server_example.py index 6a6d79977920..9d2521dc3cc7 100644 --- a/examples/disaggregated_prefill_v1/load_balance_proxy_layerwise_server_example.py +++ b/examples/disaggregated_prefill_v1/load_balance_proxy_layerwise_server_example.py @@ -99,6 +99,7 @@ import httpx from fastapi import FastAPI, HTTPException, Request from fastapi.responses import StreamingResponse +from proxy_utils import append_generated_text from vllm.logger import init_logger from vllm_ascend.distributed.kv_transfer.kv_p2p.sfa_pd_rd2h.protocol import ( @@ -519,8 +520,6 @@ async def _handle_completions(api: str, request: Request): elif chat_flag: messages = req_data["messages"] origin_prompt = messages[0].get("content", "") - if isinstance(origin_prompt, list): - origin_prompt = origin_prompt[0].get("text", "") else: origin_prompt = "" # refer to vLLM sampling_params: max_token default value @@ -584,7 +583,7 @@ async def generate_stream(): retry = True retry_count += 1 if chat_flag: - messages[0]["content"] = origin_prompt + generated_token + messages[0]["content"] = append_generated_text(origin_prompt, generated_token) else: req_data["prompt"] = origin_prompt + generated_token req_data["max_tokens"] = origin_max_tokens - completion_tokens + retry_count diff --git a/examples/disaggregated_prefill_v1/load_balance_proxy_server_example.py b/examples/disaggregated_prefill_v1/load_balance_proxy_server_example.py index 055df3d25049..14d3e6a8e375 100644 --- a/examples/disaggregated_prefill_v1/load_balance_proxy_server_example.py +++ b/examples/disaggregated_prefill_v1/load_balance_proxy_server_example.py @@ -136,6 +136,7 @@ import httpx from fastapi import FastAPI, Request from fastapi.responses import JSONResponse, Response, StreamingResponse +from proxy_utils import append_generated_text logger = logging.getLogger(__name__) @@ -1088,7 +1089,7 @@ async def release_prefill_kv_once() -> None: retry = True retry_count += 1 if chat_flag: - messages[0]["content"] = origin_prompt + generated_token + messages[0]["content"] = append_generated_text(origin_prompt, generated_token) else: req_data["prompt"] = origin_prompt + generated_token req_data["max_tokens"] = origin_max_tokens - completion_tokens + retry_count diff --git a/examples/disaggregated_prefill_v1/proxy_utils.py b/examples/disaggregated_prefill_v1/proxy_utils.py new file mode 100644 index 000000000000..f2c8321d4991 --- /dev/null +++ b/examples/disaggregated_prefill_v1/proxy_utils.py @@ -0,0 +1,11 @@ +# SPDX-License-Identifier: Apache-2.0 + +from typing import Any + + +def append_generated_text(origin_prompt: Any, generated_text: str) -> Any: + """Append partial output without dropping multimodal content parts.""" + if isinstance(origin_prompt, list): + text_part = [{"type": "text", "text": generated_text}] + return origin_prompt + (text_part if generated_text else []) + return (origin_prompt or "") + generated_text diff --git a/tests/ut/kv_offload/test_mooncake_connector.py b/tests/ut/kv_offload/test_mooncake_connector.py index 996fd2c3bdff..2fa3774f43ef 100644 --- a/tests/ut/kv_offload/test_mooncake_connector.py +++ b/tests/ut/kv_offload/test_mooncake_connector.py @@ -7,7 +7,7 @@ import types import unittest from collections import OrderedDict, defaultdict, deque -from typing import Any, cast +from typing import Any, TypedDict, cast from unittest.mock import MagicMock, patch import msgspec @@ -95,6 +95,19 @@ DONE_RECVING_MSG = b"done_recving_msg" +class KimiMambaTransferCase(TypedDict): + decode_tp: int + pulls: int + conv_shape: list[int] + ssm_shape: list[int] + local_conv_len: int + local_ssm_len: int + remote_conv_stride: int + remote_ssm_stride: int + remote_tp_offset: int + expected_segment: int + + def make_mock_kv_caches() -> dict[str, Any]: kv_cache = MagicMock(device=torch.device("npu:0")) return {"layer_0": (kv_cache, kv_cache)} @@ -616,6 +629,58 @@ def test_hybrid_rank_pulls_use_transfer_group_kv_heads(self): self.assertEqual(len(qga_pulls), 2) self.assertTrue(all(pull.num_group_pulls == 2 for pull in qga_pulls)) + def test_hybrid_mla_rank_stays_with_mamba_owner_for_unequal_tp(self): + request_id = "cmpl-ab51f3aa-8754-4ceb-93cd-e7bfee331f2d-0-9edb7c3a" + consumers_by_prefill_rank: defaultdict[int, set[int]] = defaultdict(set) + + for decode_tp_rank in range(8): + worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker) + worker.tp_rank = decode_tp_rank + worker.tp_size = 8 + worker._decode_tp_size = 8 + worker._prefill_tp_size = 16 + worker._prefill_pp_size = 1 + worker.num_key_value_heads = 1 + worker.use_sparse = False + worker.kv_group2layeridx = { + 0: ( + { + "kv_cache_spec_type": "AscendMLAAttentionSpec", + "kv_cache_group_id": 0, + "kv_cache_spec": {"total_num_kv_heads": 1}, + }, + [3], + ), + 1: ( + { + "kv_cache_spec_type": "MambaSpec", + "kv_cache_group_id": 1, + "kv_cache_spec": {}, + }, + [0, 1, 2, 4], + ), + } + + chosen_ranks, rank_group_pulls = worker._get_hybrid_remote_rank_group_pulls( + request_id, + prefill_tp_size=16, + ) + owned_mamba_ranks = {decode_tp_rank * 2, decode_tp_rank * 2 + 1} + mla_ranks = { + rank + for rank, group_pulls in rank_group_pulls.items() + if any(group_pull.group_id == 0 for group_pull in group_pulls) + } + + self.assertEqual(set(chosen_ranks), owned_mamba_ranks) + self.assertEqual(len(mla_ranks), 1) + self.assertTrue(mla_ranks.issubset(owned_mamba_ranks)) + for rank in chosen_ranks: + consumers_by_prefill_rank[rank].add(decode_tp_rank) + + self.assertEqual(set(consumers_by_prefill_rank), set(range(16))) + self.assertTrue(all(len(consumers) == 1 for consumers in consumers_by_prefill_rank.values())) + def test_hybrid_group_pulls_metadata_filters_groups_per_remote_card(self): worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker) worker.vllm_config = MockVllmConfig() @@ -1007,6 +1072,15 @@ def setUp(self): self.thread.remote_te_port = {"remote_engine": {6666: 7777}} self.thread.remote_block_stride_per_addr["remote_engine"][6666] = [[1024]] + def _configure_mock_mamba_transfer(self): + self.thread.kv_group2layeridx = {0: ({"kv_cache_spec_type": "MambaSpec"}, [0])} + self.thread.kv_caches_base_addr["local_engine"][5555] = [[0x1000, 0x2000]] + self.thread.kv_caches_base_addr["remote_engine"] = {6666: [[0x3000, 0x4000]]} + self.thread.block_len_per_addr = [[100, 200]] + self.thread.block_stride_per_addr = [[128, 256]] + self.thread.remote_block_size_scale["remote_engine"] = {6666: [[1, 1]]} + self.thread.remote_block_stride_per_addr["remote_engine"][6666] = [[160, 512]] + @patch.object(KVCacheRecvingThread, "_transfer_kv_cache_all_groups") @patch.object(KVCacheRecvingThread, "_send_done_recv_signal") def test_handle_request(self, mock_send, mock_transfer): @@ -1052,6 +1126,83 @@ def test_transfer_kv_cache(self, mock_get_meta): self.assertEqual(len(call_args[1]), len(call_args[3])) mock_get_meta.assert_not_called() + @patch.object(KVCacheRecvingThread, "_get_remote_metadata") + def test_transfer_mamba_uses_normalized_prompt_state_block(self, mock_get_meta): + # Pure metadata/address arithmetic: no NPU tensor or torch.npu call. + self._configure_mock_mamba_transfer() + self.thread.mamba_cache_mode = "align" + self.thread.num_speculative_tokens = 7 + req = dict(self.test_req) + # The aligned allocator returns the running-state destination first, + # followed by speculative blocks. The transfer must target block 2. + req["local_block_ids"] = [[2, 20, 21, 22, 23, 24, 25, 26]] + req["remote_block_ids"] = [[3]] + req["group_pulls"] = [ + GroupPull( + group_id=0, + remote_tp_offset=0, + num_group_pulls=1, + is_group_transfer_end=True, + ) + ] + + self.thread._transfer_kv_cache_all_groups(req) + + call_args, _ = self.engine.batch_transfer_sync_read.call_args + self.assertEqual(call_args[1], [0x1000 + 2 * 128, 0x2000 + 2 * 256]) + self.assertEqual(call_args[2], [0x3000 + 3 * 160, 0x4000 + 3 * 512]) + self.assertEqual(call_args[3], [100, 200]) + mock_get_meta.assert_not_called() + + @patch.object(KVCacheRecvingThread, "_get_remote_metadata") + def test_transfer_mamba_rejects_unnormalized_remote_blocks(self, mock_get_meta): + self._configure_mock_mamba_transfer() + self.thread.mamba_cache_mode = "align" + req = dict(self.test_req) + # The old token-count based selector crashed for one to three blocks + # and silently wrapped for four to seven blocks. Reject both forms. + req["local_block_ids"] = [[2]] + req["remote_block_ids"] = [[3, 4, 5]] + req["group_pulls"] = [ + GroupPull( + group_id=0, + remote_tp_offset=0, + num_group_pulls=1, + is_group_transfer_end=True, + ) + ] + + with self.assertRaisesRegex(RuntimeError, "exactly one normalized remote state block"): + self.thread._transfer_kv_cache_all_groups(req) + self.engine.batch_transfer_sync_read.assert_not_called() + mock_get_meta.assert_not_called() + + @patch.object(KVCacheRecvingThread, "_get_remote_metadata") + def test_transfer_mamba_non_aligned_selects_prompt_state_before_draft_blocks(self, mock_get_meta): + # Pure metadata/address arithmetic: no NPU tensor or torch.npu call. + self._configure_mock_mamba_transfer() + self.thread.mamba_cache_mode = "none" + self.thread.num_speculative_tokens = 7 + req = dict(self.test_req) + req["local_block_ids"] = [[2, 20, 21, 22, 23, 24, 25, 26]] + req["remote_block_ids"] = [[3, 30, 31, 32, 33, 34, 35, 36]] + req["group_pulls"] = [ + GroupPull( + group_id=0, + remote_tp_offset=0, + num_group_pulls=1, + is_group_transfer_end=True, + ) + ] + + self.thread._transfer_kv_cache_all_groups(req) + + call_args, _ = self.engine.batch_transfer_sync_read.call_args + self.assertEqual(call_args[1], [0x1000 + 2 * 128, 0x2000 + 2 * 256]) + self.assertEqual(call_args[2], [0x3000 + 3 * 160, 0x4000 + 3 * 512]) + self.assertEqual(call_args[3], [100, 200]) + mock_get_meta.assert_not_called() + @patch.object(KVCacheRecvingThread, "_get_remote_metadata") def test_transfer_groups_contiguous_kernel_blocks(self, mock_get_meta): # Kernel-level ids now arrive pre-expanded from _get_kv_split_metadata; the @@ -1202,6 +1353,205 @@ def test_append_mamba_transfer_meta_uses_block_stride_for_block_offsets(self): self.assertEqual(dst_list, [0x3000 + 3 * 160, 0x4000 + 3 * 512]) self.assertEqual(length_list, [100, 200]) + @patch( + "vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.is_conv_state_dim_first", + return_value=False, + ) + def test_append_mamba_transfer_meta_kimi_k3_sd_unequal_tp(self, _mock_layout): + """Cover K3 conv/SSM address slicing when prefill and decode TP differ.""" + cases: list[KimiMambaTransferCase] = [ + { + "decode_tp": 8, + "pulls": 2, + "conv_shape": [3, 4608], + "ssm_shape": [12, 128, 128], + "local_conv_len": 27648, + "local_ssm_len": 786432, + "remote_conv_stride": 15000, + "remote_ssm_stride": 400000, + "remote_tp_offset": 1, + "expected_segment": 768, + }, + { + "decode_tp": 8, + "pulls": 4, + "conv_shape": [3, 4608], + "ssm_shape": [12, 128, 128], + "local_conv_len": 27648, + "local_ssm_len": 786432, + "remote_conv_stride": 6912, + "remote_ssm_stride": 196608, + "remote_tp_offset": 3, + "expected_segment": 384, + }, + { + "decode_tp": 16, + "pulls": 2, + "conv_shape": [3, 2304], + "ssm_shape": [6, 128, 128], + "local_conv_len": 13824, + "local_ssm_len": 393216, + "remote_conv_stride": 6912, + "remote_ssm_stride": 196608, + "remote_tp_offset": 1, + "expected_segment": 384, + }, + ] + + for case in cases: + with self.subTest(case=case): + self.thread.tp_size = case["decode_tp"] + self.thread.vllm_config.model_config.hf_text_config = types.SimpleNamespace( + linear_attn_config={"num_heads": 96, "head_dim": 128} + ) + src_list: list[int] = [] + dst_list: list[int] = [] + length_list: list[int] = [] + local_bases = [0x100000, 0x300000] + remote_bases = [0x200000, 0x400000] + local_strides = [case["local_conv_len"] + 4096, case["local_ssm_len"] + 8192] + local_block_id = 2 + remote_block_id = 3 + + self.thread._append_mamba_transfer_meta( + src_list, + dst_list, + length_list, + group_spec={ + "kv_cache_spec_type": "MambaSpec", + "shapes": [case["conv_shape"], case["ssm_shape"]], + "dtype_sizes": [2, 4], + }, + src_layer_base_addr=local_bases, + dst_layer_base_addr=remote_bases, + block_len=[case["local_conv_len"], case["local_ssm_len"]], + block_stride=local_strides, + remote_block_stride=[case["remote_conv_stride"], case["remote_ssm_stride"]], + remote_block_id=remote_block_id, + local_block_id=local_block_id, + tp_num_need_pulls=case["pulls"], + remote_tp_offset=case["remote_tp_offset"], + ) + + remote_segment = case["expected_segment"] + local_segment = remote_segment * case["pulls"] + local_conv_base = local_bases[0] + local_block_id * local_strides[0] + remote_conv_base = remote_bases[0] + remote_block_id * case["remote_conv_stride"] + expected_src: list[int] = [] + expected_dst: list[int] = [] + for state_idx in range(3): + for segment_idx in range(3): + expected_src.append( + local_conv_base + + ( + state_idx * case["conv_shape"][1] + + segment_idx * local_segment + + case["remote_tp_offset"] * remote_segment + ) + * 2 + ) + expected_dst.append( + remote_conv_base + (state_idx * remote_segment * 3 + segment_idx * remote_segment) * 2 + ) + expected_src.append( + local_bases[1] + + local_block_id * local_strides[1] + + case["remote_tp_offset"] * case["local_ssm_len"] // case["pulls"] + ) + expected_dst.append(remote_bases[1] + remote_block_id * case["remote_ssm_stride"]) + + self.assertEqual(src_list, expected_src) + self.assertEqual(dst_list, expected_dst) + self.assertEqual( + length_list, + [remote_segment * 2] * 9 + [case["local_ssm_len"] // case["pulls"]], + ) + + @patch( + "vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.is_conv_state_dim_first", + return_value=False, + ) + def test_append_mamba_transfer_meta_legacy_qwen_layout(self, _mock_layout): + self.thread.tp_size = 4 + self.thread.vllm_config.model_config.hf_text_config = types.SimpleNamespace( + linear_num_key_heads=16, + linear_key_head_dim=128, + linear_num_value_heads=32, + linear_value_head_dim=128, + ) + src_list: list[int] = [] + dst_list: list[int] = [] + length_list: list[int] = [] + + self.thread._append_mamba_transfer_meta( + src_list, + dst_list, + length_list, + group_spec={ + "kv_cache_spec_type": "MambaSpec", + "shapes": [[3, 2048], [8, 128, 128]], + "dtype_sizes": [2, 4], + }, + src_layer_base_addr=[0x100000, 0x300000], + dst_layer_base_addr=[0x200000, 0x400000], + block_len=[12288, 524288], + block_stride=[12288, 524288], + remote_block_stride=[6144, 262144], + remote_block_id=0, + local_block_id=0, + tp_num_need_pulls=2, + remote_tp_offset=0, + ) + + self.assertEqual(len(src_list), 10) + self.assertEqual(len(dst_list), 10) + self.assertEqual(length_list, [512, 512, 1024] * 3 + [262144]) + + @patch( + "vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.is_conv_state_dim_first", + return_value=True, + ) + def test_append_mamba_transfer_meta_kimi_k3_ds_unequal_tp(self, _mock_layout): + self.thread.tp_size = 8 + self.thread.vllm_config.model_config.hf_text_config = types.SimpleNamespace( + linear_attn_config={"num_heads": 96, "head_dim": 128} + ) + src_list: list[int] = [] + dst_list: list[int] = [] + length_list: list[int] = [] + + self.thread._append_mamba_transfer_meta( + src_list, + dst_list, + length_list, + group_spec={ + "kv_cache_spec_type": "MambaSpec", + "shapes": [[4608, 3], [12, 128, 128]], + "dtype_sizes": [2, 4], + }, + src_layer_base_addr=[0x100000, 0x300000], + dst_layer_base_addr=[0x200000, 0x400000], + block_len=[27648, 786432], + block_stride=[27648, 786432], + remote_block_stride=[13824, 393216], + remote_block_id=0, + local_block_id=0, + tp_num_need_pulls=2, + remote_tp_offset=1, + ) + + self.assertEqual( + src_list, + [ + 0x100000 + 4608, + 0x100000 + 13824, + 0x100000 + 23040, + 0x300000 + 393216, + ], + ) + self.assertEqual(dst_list, [0x200000, 0x200000 + 4608, 0x200000 + 9216, 0x400000]) + self.assertEqual(length_list, [4608, 4608, 4608, 393216]) + def test_transfer_kv_cache_failure(self): self.engine.batch_transfer_sync_read.return_value = -1 self.thread.kv_caches_base_addr["remote_engine"] = {6666: [[0x3000]]} @@ -1827,7 +2177,22 @@ def test_get_transfer_block_ids_trims_attention_mtp_blocks(self): self.assertEqual(block_ids, ([10, 11, 12],)) - def test_get_transfer_block_ids_keeps_state_group(self): + def test_get_transfer_block_ids_selects_aligned_prompt_state_block(self): + self.scheduler.vllm_config.cache_config.mamba_cache_mode = "align" + self.scheduler.group_transfer_info = [ + types.SimpleNamespace( # type: ignore[list-item] + tokens_per_block=16, + blocks_per_window=0, + is_state_group=True, + ) + ] + + block_ids = self.scheduler._get_transfer_block_ids(([20, 21, 22, 23, 24],), prompt_len=33) + + self.assertEqual(block_ids, ([22],)) + + def test_get_transfer_block_ids_keeps_non_aligned_state_group(self): + self.scheduler.vllm_config.cache_config.mamba_cache_mode = "none" self.scheduler.group_transfer_info = [ types.SimpleNamespace( # type: ignore[list-item] tokens_per_block=16, @@ -1840,6 +2205,19 @@ def test_get_transfer_block_ids_keeps_state_group(self): self.assertEqual(block_ids, ([20, 21, 22, 23],)) + def test_get_transfer_block_ids_rejects_short_aligned_state_metadata(self): + self.scheduler.vllm_config.cache_config.mamba_cache_mode = "align" + self.scheduler.group_transfer_info = [ + types.SimpleNamespace( # type: ignore[list-item] + tokens_per_block=16, + blocks_per_window=0, + is_state_group=True, + ) + ] + + with self.assertRaisesRegex(RuntimeError, "Invalid aligned Mamba state block metadata"): + self.scheduler._get_transfer_block_ids(([20, 21],), prompt_len=48) + def test_get_transfer_block_ids_uses_compressed_prompt_len(self): self.scheduler.group_transfer_info = [ types.SimpleNamespace( # type: ignore[list-item] @@ -2001,6 +2379,7 @@ def test_request_finished_trims_mtp_before_swa_tail_clip(self): self.assertIn("req_mtp_swa", self.scheduler._reqs_need_send) def test_request_finished_handles_mtp_swa_and_state_groups_together(self): + self.scheduler.vllm_config.cache_config.mamba_cache_mode = "align" self.scheduler.group_transfer_info = [ types.SimpleNamespace( tokens_per_block=16, @@ -2037,7 +2416,7 @@ def test_request_finished_handles_mtp_swa_and_state_groups_together(self): ( [100, 101, 102, 103], [200, 201, 202], - [300, 301, 302, 303, 304], + [303], ), ) self.assertEqual(params["num_prompt_blocks"], 4) @@ -2289,6 +2668,60 @@ def test_register_kv_caches_mla_case(self): self.assertTrue(worker.use_mla) self.assertEqual(len(worker.block_len_per_addr[0]), 2) + def test_registered_hybrid_buffer_recovers_aligned_raw_base(self): + alignment = 2 * 1024 * 1024 + tensor_size = 4 * alignment + logical_view_offset = 0x17200 + raw_tensor = torch.empty(tensor_size + alignment, dtype=torch.uint8) + aligned_offset = (-raw_tensor.data_ptr()) % alignment + aligned_tensor = raw_tensor[aligned_offset : aligned_offset + tensor_size] + logical_tensor = aligned_tensor[logical_view_offset:] + layer_name = "model.layers.0.self_attn" + + worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker) + worker.num_blocks = 1579 + worker._layer_specs = {layer_name: MagicMock()} + worker.kv_cache_config = types.SimpleNamespace( + kv_cache_tensors=[types.SimpleNamespace(size=tensor_size, shared_by=[layer_name])] + ) + + self.assertEqual(aligned_tensor.data_ptr() % alignment, 0) + self.assertNotEqual(logical_tensor.data_ptr() % alignment, 0) + ptrs, lengths = worker._get_registered_kv_tensor_buffers({layer_name: logical_tensor}) + + self.assertEqual(ptrs, [aligned_tensor.data_ptr()]) + self.assertEqual(lengths, [tensor_size]) + + def test_registered_mtp_buffer_ignores_aligned_stale_group_padding(self): + alignment = 2 * 1024 * 1024 + tensor_size = 4 * alignment + raw_tensor = torch.empty(tensor_size + alignment, dtype=torch.uint8) + aligned_offset = (-raw_tensor.data_ptr()) % alignment + aligned_tensor = raw_tensor[aligned_offset : aligned_offset + tensor_size] + logical_tensor = aligned_tensor[2 * alignment :] + layer_name = "model.layers.0.mtp_attn" + + worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker) + worker.kv_cache_config = types.SimpleNamespace( + kv_cache_tensors=[ + types.SimpleNamespace( + size=tensor_size, + shared_by=[layer_name], + ) + ] + ) + + # Subtracting a stale one-group padding value from this view would + # produce another aligned address one alignment after the real base. + stale_padding_base = logical_tensor.data_ptr() - alignment + self.assertEqual(stale_padding_base % alignment, 0) + self.assertNotEqual(stale_padding_base, aligned_tensor.data_ptr()) + + ptrs, lengths = worker._get_registered_kv_tensor_buffers({layer_name: logical_tensor}) + + self.assertEqual(ptrs, [aligned_tensor.data_ptr()]) + self.assertEqual(lengths, [tensor_size]) + def test_device_id_selection_with_physical_devices(self): # Test with physical devices set worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig()) diff --git a/tests/ut/test_disaggregated_prefill_proxy.py b/tests/ut/test_disaggregated_prefill_proxy.py new file mode 100644 index 000000000000..71e21556066e --- /dev/null +++ b/tests/ut/test_disaggregated_prefill_proxy.py @@ -0,0 +1,34 @@ +# SPDX-License-Identifier: Apache-2.0 + +import importlib.util +from pathlib import Path +from typing import Any + + +def _load_proxy_utils(): + path = Path(__file__).parents[2] / "examples" / "disaggregated_prefill_v1" / "proxy_utils.py" + spec = importlib.util.spec_from_file_location("disaggregated_prefill_proxy_utils", path) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def test_append_generated_text_preserves_multimodal_parts(): + append_generated_text = _load_proxy_utils().append_generated_text + prompt: list[dict[str, Any]] = [ + {"type": "text", "text": "Describe the image."}, + {"type": "image_url", "image_url": {"url": "file:///tmp/image.png"}}, + ] + + result = append_generated_text(prompt, "Partial answer") + + assert result == prompt + [{"type": "text", "text": "Partial answer"}] + assert prompt[-1]["type"] == "image_url" + + +def test_append_generated_text_keeps_text_prompt_behavior(): + append_generated_text = _load_proxy_utils().append_generated_text + + assert append_generated_text("Prompt: ", "answer") == "Prompt: answer" + assert append_generated_text(None, "answer") == "answer" diff --git a/tests/ut/worker/a2/test_model_runner_v1_with_device.py b/tests/ut/worker/a2/test_model_runner_v1_with_device.py index 51a198d9af16..bd7875ba728e 100644 --- a/tests/ut/worker/a2/test_model_runner_v1_with_device.py +++ b/tests/ut/worker/a2/test_model_runner_v1_with_device.py @@ -389,3 +389,37 @@ def test_determine_batch_execution_and_padding( finally: runner.speculative_config = saved_spec_config runner.uniform_decode_query_len = saved_query_len + + +@pytest.mark.parametrize( + ("num_prompt_tokens", "expected_uniform_decode"), + [ + pytest.param(8, False, id="one_token_pd_prefill"), + pytest.param(7, True, id="decode_after_prompt"), + ], +) +def test_uniform_decode_requires_completed_prompt_without_spec_decode( + model_runner, + num_prompt_tokens: int, + expected_uniform_decode: bool, +): + runner = model_runner + runner.speculative_config = None + runner.uniform_decode_query_len = 1 + runner.input_batch.num_computed_tokens_cpu[0] = 7 + runner.input_batch.num_prompt_tokens[0] = num_prompt_tokens + + with patch.object( + runner.cudagraph_dispatcher, + "dispatch", + wraps=runner.cudagraph_dispatcher.dispatch, + ) as dispatch: + runner._determine_batch_execution_and_padding( + num_tokens=1, + num_reqs=1, + num_scheduled_tokens_np=np.array([1], dtype=np.int32), + max_num_scheduled_tokens=1, + use_cascade_attn=False, + ) + + assert bool(dispatch.call_args.kwargs["uniform_decode"]) is expected_uniform_decode diff --git a/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py b/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py index c61cd17ec94c..595242fb9c1c 100644 --- a/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py +++ b/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py @@ -42,6 +42,7 @@ ) from vllm.distributed.utils import get_pp_indices from vllm.logger import logger +from vllm.model_executor.layers.mamba.mamba_utils import is_conv_state_dim_first from vllm.utils.math_utils import cdiv from vllm.utils.network_utils import get_ip, make_zmq_path, make_zmq_socket from vllm.v1.core.sched.output import SchedulerOutput @@ -62,6 +63,7 @@ RegisterRegions, collect_storage_merged_register_regions, get_transfer_timeout_value, + tensor_storage_key, validate_register_region_count, ) from vllm_ascend.distributed.utils import ( @@ -86,6 +88,7 @@ # number of peers is larger than max_workers. Yield after a small FIFO batch so # other peers already waiting in the global executor queue can make progress. MAX_REQUESTS_PER_PEER_HANDLER = 5 +KV_CACHE_BUFFER_ALIGNMENT = 2 * 1024 * 1024 class RemotePortInfo(TypedDict): @@ -505,6 +508,7 @@ def __init__( assert vllm_config is not None self.vllm_config: VllmConfig = vllm_config self.model_config = self.vllm_config.model_config + self.mamba_cache_mode = getattr(self.vllm_config.cache_config, "mamba_cache_mode", None) self.num_speculative_tokens = ( self.vllm_config.speculative_config.num_speculative_tokens if self.vllm_config.speculative_config is not None @@ -899,10 +903,28 @@ def get_remote_layer_idx( ) ) else: - # When Prefix Caching is enabled on both P and D nodes, num_block should not be forced to match, - # as the D-node requires dynamic allocation based on its specific cache hit rate. - transfer_block_idx = len(remote_group_block_ids) - self.num_speculative_tokens - 1 - grouped_remote_block_ids = [[remote_group_block_ids[transfer_block_idx]]] + # Ascend Hybrid Mamba supports "align" (prefix caching) and + # "none" (no prefix caching), but not "all". + if self.mamba_cache_mode == "align": + if len(remote_group_block_ids) != 1: + raise RuntimeError( + "Mooncake Mamba transfer requires exactly one normalized remote state block; " + f"request_id={remote_request_id}, group_idx={group_idx}, " + f"remote_block_count={len(remote_group_block_ids)}, " + f"local_block_count={len(local_group_block_ids)}." + ) + remote_state_block_id = remote_group_block_ids[0] + else: + transfer_block_idx = len(remote_group_block_ids) - self.num_speculative_tokens - 1 + if transfer_block_idx < 0: + raise RuntimeError( + "Invalid non-aligned Mamba state block metadata: " + f"request_id={remote_request_id}, group_idx={group_idx}, " + f"remote_block_count={len(remote_group_block_ids)}, " + f"num_speculative_tokens={self.num_speculative_tokens}." + ) + remote_state_block_id = remote_group_block_ids[transfer_block_idx] + grouped_remote_block_ids = [[remote_state_block_id]] grouped_local_block_ids = [[local_group_block_ids[0]]] if is_mamba_group: @@ -1208,35 +1230,44 @@ def _append_mamba_transfer_meta( conv_shape = group_spec["shapes"][0] conv_dtype_size = group_spec["dtype_sizes"][0] - linear_key_head_dim = self.vllm_config.model_config.hf_text_config.linear_key_head_dim - linear_num_key_heads = self.vllm_config.model_config.hf_text_config.linear_num_key_heads - linear_value_head_dim = self.vllm_config.model_config.hf_text_config.linear_value_head_dim - linear_num_value_heads = self.vllm_config.model_config.hf_text_config.linear_num_value_heads - remote_num_key_heads = linear_num_key_heads // remote_tp_size - remote_num_value_heads = linear_num_value_heads // remote_tp_size - remote_conv_width = ( - remote_num_key_heads * 2 * linear_key_head_dim + remote_num_value_heads * linear_value_head_dim - ) - remote_conv_offsets = [ - 0, - remote_num_key_heads * linear_key_head_dim, - remote_num_key_heads * 2 * linear_key_head_dim, - ] - remote_conv_sizes = [ - remote_num_key_heads * linear_key_head_dim, - remote_num_key_heads * linear_key_head_dim, - remote_num_value_heads * linear_value_head_dim, - ] + hf_text_config = self.vllm_config.model_config.hf_text_config + linear_attn_config = getattr(hf_text_config, "linear_attn_config", None) + if linear_attn_config is not None: + projection_width = linear_attn_config["num_heads"] * linear_attn_config["head_dim"] + remote_conv_sizes = [projection_width // remote_tp_size] * 3 + else: + remote_num_key_heads = hf_text_config.linear_num_key_heads // remote_tp_size + remote_num_value_heads = hf_text_config.linear_num_value_heads // remote_tp_size + remote_conv_sizes = [ + remote_num_key_heads * hf_text_config.linear_key_head_dim, + remote_num_key_heads * hf_text_config.linear_key_head_dim, + remote_num_value_heads * hf_text_config.linear_value_head_dim, + ] - for i in range(conv_shape[0]): + remote_conv_width = sum(remote_conv_sizes) + remote_conv_offsets = [0, remote_conv_sizes[0], sum(remote_conv_sizes[:2])] + if is_conv_state_dim_first(): + state_len = conv_shape[1] for remote_conv_offset, remote_conv_size in zip(remote_conv_offsets, remote_conv_sizes): - remote_addr_offset = (i * remote_conv_width + remote_conv_offset) * conv_dtype_size local_addr_offset = ( - (i * remote_conv_width + remote_conv_offset) * tp_ratio + remote_tp_offset * remote_conv_size - ) * conv_dtype_size + (remote_conv_offset * tp_ratio + remote_tp_offset * remote_conv_size) * state_len * conv_dtype_size + ) + remote_addr_offset = remote_conv_offset * state_len * conv_dtype_size src_list.append(local_conv_addr + local_block_id * local_conv_stride + local_addr_offset) dst_list.append(remote_conv_addr + remote_block_id * remote_conv_stride + remote_addr_offset) - length_list.append(remote_conv_size * conv_dtype_size) + length_list.append(remote_conv_size * state_len * conv_dtype_size) + else: + state_len = conv_shape[0] + for state_idx in range(state_len): + for remote_conv_offset, remote_conv_size in zip(remote_conv_offsets, remote_conv_sizes): + remote_addr_offset = (state_idx * remote_conv_width + remote_conv_offset) * conv_dtype_size + local_addr_offset = ( + (state_idx * remote_conv_width + remote_conv_offset) * tp_ratio + + remote_tp_offset * remote_conv_size + ) * conv_dtype_size + src_list.append(local_conv_addr + local_block_id * local_conv_stride + local_addr_offset) + dst_list.append(remote_conv_addr + remote_block_id * remote_conv_stride + remote_addr_offset) + length_list.append(remote_conv_size * conv_dtype_size) src_list.append( local_ssm_addr + local_block_id * local_ssm_stride + remote_tp_offset * local_ssm_len // tp_num_need_pulls @@ -1768,9 +1799,11 @@ def _get_group_unique_specs(self, group: Any) -> list[Any]: def _get_transfer_block_ids(self, block_ids: BlockIds, prompt_len: int) -> BlockIds: """Return blocks that contain prompt KV, dropping MTP extra blocks. - State groups such as Mamba are not context-block aligned with attention - KV, so keep them unchanged and only clip attention-like groups here. - SWA tail clipping is handled as a separate step after this. + In aligned Mamba mode, normalize each state group to the single block + containing the final prompt state. This prevents the receiver from + inferring a block index from a speculative *token* count. Non-aligned + state groups keep their existing behavior. SWA tail clipping is handled + as a separate step after this. """ if len(block_ids) == 0: return block_ids @@ -1780,8 +1813,23 @@ def _get_transfer_block_ids(self, block_ids: BlockIds, prompt_len: int) -> Block transfer_block_ids = [] cp_size = max(1, self.pcp_size * self.dcp_size) for blocks, group_info in zip(block_ids, self.group_transfer_info): - if group_info.is_state_group: + is_aligned_state_group = group_info.is_state_group and ( + getattr(self.vllm_config.cache_config, "mamba_cache_mode", None) == "align" + ) + if group_info.is_state_group and not is_aligned_state_group: transfer_block_ids.append(blocks) + elif is_aligned_state_group: + # Mamba state is not CP-sharded like attention KV. Its aligned + # block index is derived from the actual (already truncated) + # prompt length, without multiplying by the CP size. + num_prompt_state_blocks = cdiv(prompt_len, group_info.tokens_per_block) + if num_prompt_state_blocks <= 0 or num_prompt_state_blocks > len(blocks): + raise RuntimeError( + "Invalid aligned Mamba state block metadata: " + f"prompt_len={prompt_len}, tokens_per_block={group_info.tokens_per_block}, " + f"required_block_count={num_prompt_state_blocks}, available_block_count={len(blocks)}." + ) + transfer_block_ids.append(blocks[num_prompt_state_blocks - 1 : num_prompt_state_blocks]) else: # In context parallelism, each scheduler-visible block id is a # CP-grouped/virtual block shared by all CP ranks. It therefore @@ -2377,34 +2425,52 @@ def _get_layer_spec(self, layer_name: str) -> Any: layer_spec = layer_spec.kv_cache_specs[layer_name] return layer_spec - def _get_mamba_conv_padding(self, layer_spec: Any) -> int: - if not isinstance(layer_spec, MambaSpec): - return 0 - conv_nbytes = torch.tensor([], dtype=layer_spec.dtypes[0]).element_size() # type: ignore[misc] - conv_shape = torch.Size(layer_spec.shapes[0]) - return self.num_blocks * conv_shape.numel() * conv_nbytes + @staticmethod + def _recover_aligned_kv_tensor_base( + shared_tensors: list[torch.Tensor], + tensor_size: int, + ) -> int: + """Recover the aligned raw buffer base behind hybrid cache views.""" + candidates: set[int] = set() + for tensor in shared_tensors: + storage = tensor.untyped_storage() + storage_base = tensor_storage_key(tensor) + aligned_base = ( + (storage_base + KV_CACHE_BUFFER_ALIGNMENT - 1) // KV_CACHE_BUFFER_ALIGNMENT * KV_CACHE_BUFFER_ALIGNMENT + ) + storage_end = storage_base + storage.nbytes() + if aligned_base <= tensor.data_ptr() and aligned_base + tensor_size <= storage_end: + candidates.add(aligned_base) + + if len(candidates) != 1: + raise RuntimeError( + "Unable to recover one aligned KV tensor base from hybrid cache views: " + f"candidates={sorted(candidates)}, tensor_size={tensor_size}." + ) + return candidates.pop() def _get_registered_kv_tensor_buffers(self, kv_caches: dict[str, torch.Tensor]) -> tuple[list[int], list[int]]: ptrs: list[int] = [] lengths: list[int] = [] - conv_padding = 0 for kv_cache_tensor in self.kv_cache_config.kv_cache_tensors: - shared_addrs: list[int] = [] - has_mtp = False + shared_tensors: list[torch.Tensor] = [] for layer_name in kv_cache_tensor.shared_by: - has_mtp = has_mtp or "mtp" in layer_name - layer_spec = self._get_layer_spec(layer_name) - conv_padding = max(conv_padding, self._get_mamba_conv_padding(layer_spec)) for single_kv_cache in self._as_kv_cache_tuple(kv_caches[layer_name]): - shared_addrs.append(single_kv_cache.data_ptr()) + shared_tensors.append(single_kv_cache) - if not shared_addrs: + if not shared_tensors: continue - base_addr = min(shared_addrs) - if has_mtp: - base_addr -= conv_padding - assert base_addr % (2 * 1024 * 1024) == 0, f"Tensor start addr {base_addr} is not align with 2M." + # Hybrid cache views can begin after Mamba padding, and target and + # draft MLA groups can use different padding. Arithmetic based on + # a previous group's padding can yield a wrong but still aligned + # address, so always recover the allocation base from storage. + base_addr = self._recover_aligned_kv_tensor_base( + shared_tensors, + kv_cache_tensor.size, + ) + if base_addr % KV_CACHE_BUFFER_ALIGNMENT != 0: + raise RuntimeError(f"Tensor start addr {base_addr} is not aligned to 2 MiB.") ptrs.append(base_addr) lengths.append(kv_cache_tensor.size) @@ -2425,7 +2491,8 @@ def _get_registered_kv_tensor_buffers_hybrid( if not shared_addrs: continue base_addr = min(shared_addrs) - assert base_addr % (2 * 1024 * 1024) == 0, f"Tensor start addr {base_addr} is not align with 2M." + if base_addr % KV_CACHE_BUFFER_ALIGNMENT != 0: + raise RuntimeError(f"Tensor start addr {base_addr} is not aligned to 2 MiB.") ptrs.append(base_addr) lengths.append(kv_cache_tensor.size) @@ -3338,6 +3405,18 @@ def _get_hybrid_remote_rank_group_pulls( prefill_tp_size: int, ) -> tuple[list[int], dict[int, list[GroupPull]]]: rank_group_pulls: OrderedDict[int, list[GroupPull]] = OrderedDict() + has_mamba_group = any( + layer_indices and group_spec["kv_cache_spec_type"] == "MambaSpec" + for group_spec, layer_indices in self.kv_group2layeridx.values() + ) + mamba_num_group_pulls = 0 + if has_mamba_group: + if prefill_tp_size % self.tp_size != 0: + raise ValueError( + f"Hybrid Mamba prefill tp size({prefill_tp_size}) must be divisible by " + f"decode tp size({self.tp_size})." + ) + mamba_num_group_pulls = prefill_tp_size // self.tp_size def add_group_pull(remote_rank: int, group_pull: GroupPull) -> None: rank_group_pulls.setdefault(remote_rank, []).append(group_pull) @@ -3347,11 +3426,7 @@ def add_group_pull(remote_rank: int, group_pull: GroupPull) -> None: continue if group_spec["kv_cache_spec_type"] == "MambaSpec": - assert prefill_tp_size % self.tp_size == 0, ( - f"Hybrid Mamba prefill tp size({prefill_tp_size}) must be divisible by " - f"decode tp size({self.tp_size})." - ) - num_group_pulls = prefill_tp_size // self.tp_size + num_group_pulls = mamba_num_group_pulls for pp_rank in range(self._prefill_pp_size): pp_rank_offset = pp_rank * prefill_tp_size local_tp_offset = self.tp_rank * num_group_pulls @@ -3370,7 +3445,18 @@ def add_group_pull(remote_rank: int, group_pull: GroupPull) -> None: continue num_group_pulls = self._get_attention_group_num_need_pulls(group_spec, prefill_tp_size) - chosen_rank_list = self._get_attention_group_remote_rank(req_id, group_spec, prefill_tp_size) + num_key_value_heads = self._get_attention_group_num_key_value_heads(group_spec) + if has_mamba_group and num_key_value_heads == 1 and num_group_pulls == 1: + # Keep a replicated single-head attention cache inside this + # D rank's Mamba owner group so only its owner can signal that + # all request state has finished transferring. + replica_offset = random.Random(string_to_int64_hash(req_id)).randrange(mamba_num_group_pulls) + chosen_rank_list = [ + pp_rank * prefill_tp_size + self.tp_rank * mamba_num_group_pulls + replica_offset + for pp_rank in range(self._prefill_pp_size) + ] + else: + chosen_rank_list = self._get_attention_group_remote_rank(req_id, group_spec, prefill_tp_size) assert len(chosen_rank_list) == num_group_pulls * self._prefill_pp_size, ( f"chosen_rank_list({chosen_rank_list}) does not match num_group_pulls({num_group_pulls}) " f"and prefill pp size({self._prefill_pp_size})." diff --git a/vllm_ascend/worker/model_runner_v1.py b/vllm_ascend/worker/model_runner_v1.py index dcfe417988a6..5938d1d6c3b8 100644 --- a/vllm_ascend/worker/model_runner_v1.py +++ b/vllm_ascend/worker/model_runner_v1.py @@ -2746,10 +2746,15 @@ def _determine_batch_execution_and_padding( num_encoder_reqs: int = 0, ) -> tuple[CUDAGraphMode, BatchDescriptor, bool, torch.Tensor | None, CUDAGraphStat | None]: num_tokens_padded = self._pad_for_sequence_parallelism(num_tokens) - is_all_decode = np.all(self.input_batch.num_computed_tokens_cpu[:num_reqs] > 0) + # A one-token chunk can still be prefill at a P/D handoff. Decode graph + # replay is valid only after every prompt has been fully computed. + is_all_decode = np.all( + self.input_batch.num_computed_tokens_cpu[:num_reqs] + >= self.input_batch.num_prompt_tokens[:num_reqs] + ) uniform_decode = ( ( - (is_all_decode if self.speculative_config else True) + is_all_decode and (max_num_scheduled_tokens == self.uniform_decode_query_len) and (num_tokens == max_num_scheduled_tokens * num_reqs) ) From bf40dadde158433f41355461650d7030b33e4b08 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Tue, 25 Aug 2026 03:53:36 -0500 Subject: [PATCH 16/50] refactor(proxy): inline multimodal retry handling Signed-off-by: maoxx241 --- .github/workflows/scripts/test_config.yaml | 2 -- ..._balance_proxy_layerwise_server_example.py | 8 +++-- .../load_balance_proxy_server_example.py | 8 +++-- .../disaggregated_prefill_v1/proxy_utils.py | 11 ------ tests/ut/test_disaggregated_prefill_proxy.py | 34 ------------------- 5 files changed, 12 insertions(+), 51 deletions(-) delete mode 100644 examples/disaggregated_prefill_v1/proxy_utils.py delete mode 100644 tests/ut/test_disaggregated_prefill_proxy.py diff --git a/.github/workflows/scripts/test_config.yaml b/.github/workflows/scripts/test_config.yaml index eb5c5523af27..ea08d20fd912 100644 --- a/.github/workflows/scripts/test_config.yaml +++ b/.github/workflows/scripts/test_config.yaml @@ -223,11 +223,9 @@ - name: distributed optional: false source_file_dependencies: - - examples/disaggregated_prefill_v1 - vllm_ascend/distributed - vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/attention_fence.py tests: - - tests/ut/test_disaggregated_prefill_proxy.py - tests/ut/distributed - tests/e2e/pull_request/two_card/test_data_parallel.py - tests/e2e/pull_request/four_card/test_data_parallel_tp2.py diff --git a/examples/disaggregated_prefill_v1/load_balance_proxy_layerwise_server_example.py b/examples/disaggregated_prefill_v1/load_balance_proxy_layerwise_server_example.py index 9d2521dc3cc7..efda19ab8c43 100644 --- a/examples/disaggregated_prefill_v1/load_balance_proxy_layerwise_server_example.py +++ b/examples/disaggregated_prefill_v1/load_balance_proxy_layerwise_server_example.py @@ -99,7 +99,6 @@ import httpx from fastapi import FastAPI, HTTPException, Request from fastapi.responses import StreamingResponse -from proxy_utils import append_generated_text from vllm.logger import init_logger from vllm_ascend.distributed.kv_transfer.kv_p2p.sfa_pd_rd2h.protocol import ( @@ -583,7 +582,12 @@ async def generate_stream(): retry = True retry_count += 1 if chat_flag: - messages[0]["content"] = append_generated_text(origin_prompt, generated_token) + messages[0]["content"] = ( + origin_prompt + + ([{"type": "text", "text": generated_token}] if generated_token else []) + if isinstance(origin_prompt, list) + else (origin_prompt or "") + generated_token + ) else: req_data["prompt"] = origin_prompt + generated_token req_data["max_tokens"] = origin_max_tokens - completion_tokens + retry_count diff --git a/examples/disaggregated_prefill_v1/load_balance_proxy_server_example.py b/examples/disaggregated_prefill_v1/load_balance_proxy_server_example.py index 14d3e6a8e375..83d213ce78e7 100644 --- a/examples/disaggregated_prefill_v1/load_balance_proxy_server_example.py +++ b/examples/disaggregated_prefill_v1/load_balance_proxy_server_example.py @@ -136,7 +136,6 @@ import httpx from fastapi import FastAPI, Request from fastapi.responses import JSONResponse, Response, StreamingResponse -from proxy_utils import append_generated_text logger = logging.getLogger(__name__) @@ -1089,7 +1088,12 @@ async def release_prefill_kv_once() -> None: retry = True retry_count += 1 if chat_flag: - messages[0]["content"] = append_generated_text(origin_prompt, generated_token) + messages[0]["content"] = ( + origin_prompt + + ([{"type": "text", "text": generated_token}] if generated_token else []) + if isinstance(origin_prompt, list) + else (origin_prompt or "") + generated_token + ) else: req_data["prompt"] = origin_prompt + generated_token req_data["max_tokens"] = origin_max_tokens - completion_tokens + retry_count diff --git a/examples/disaggregated_prefill_v1/proxy_utils.py b/examples/disaggregated_prefill_v1/proxy_utils.py deleted file mode 100644 index f2c8321d4991..000000000000 --- a/examples/disaggregated_prefill_v1/proxy_utils.py +++ /dev/null @@ -1,11 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 - -from typing import Any - - -def append_generated_text(origin_prompt: Any, generated_text: str) -> Any: - """Append partial output without dropping multimodal content parts.""" - if isinstance(origin_prompt, list): - text_part = [{"type": "text", "text": generated_text}] - return origin_prompt + (text_part if generated_text else []) - return (origin_prompt or "") + generated_text diff --git a/tests/ut/test_disaggregated_prefill_proxy.py b/tests/ut/test_disaggregated_prefill_proxy.py deleted file mode 100644 index 71e21556066e..000000000000 --- a/tests/ut/test_disaggregated_prefill_proxy.py +++ /dev/null @@ -1,34 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 - -import importlib.util -from pathlib import Path -from typing import Any - - -def _load_proxy_utils(): - path = Path(__file__).parents[2] / "examples" / "disaggregated_prefill_v1" / "proxy_utils.py" - spec = importlib.util.spec_from_file_location("disaggregated_prefill_proxy_utils", path) - assert spec is not None and spec.loader is not None - module = importlib.util.module_from_spec(spec) - spec.loader.exec_module(module) - return module - - -def test_append_generated_text_preserves_multimodal_parts(): - append_generated_text = _load_proxy_utils().append_generated_text - prompt: list[dict[str, Any]] = [ - {"type": "text", "text": "Describe the image."}, - {"type": "image_url", "image_url": {"url": "file:///tmp/image.png"}}, - ] - - result = append_generated_text(prompt, "Partial answer") - - assert result == prompt + [{"type": "text", "text": "Partial answer"}] - assert prompt[-1]["type"] == "image_url" - - -def test_append_generated_text_keeps_text_prompt_behavior(): - append_generated_text = _load_proxy_utils().append_generated_text - - assert append_generated_text("Prompt: ", "answer") == "Prompt: answer" - assert append_generated_text(None, "answer") == "answer" From 60c6cfc8ae68b4391016f1129d4629ff9133c354 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Fri, 21 Aug 2026 07:00:47 -0500 Subject: [PATCH 17/50] docs(kimi-k3): add deployment and validation guide Document Kimi K3 prerequisites, quantized and P/D deployment, DSpark usage, storage-light nightly coverage, and validation boundaries for main. Signed-off-by: maoxx241 --- docs/source/tutorials/models/Kimi-K3.md | 455 ++++++++++++++++++++++++ docs/source/tutorials/models/index.md | 2 + 2 files changed, 457 insertions(+) create mode 100644 docs/source/tutorials/models/Kimi-K3.md diff --git a/docs/source/tutorials/models/Kimi-K3.md b/docs/source/tutorials/models/Kimi-K3.md new file mode 100644 index 000000000000..56aca600db56 --- /dev/null +++ b/docs/source/tutorials/models/Kimi-K3.md @@ -0,0 +1,455 @@ +# Kimi K3 + +## 1 Introduction + +Kimi K3 is a multimodal mixture-of-experts model that combines Kimi Delta +Attention (KDA), gated Multi-head Latent Attention (MLA), attention residuals, +SiTU activations, and latent MoE layers. This guide describes the W4A8 +deployment validated by the Kimi K3 integration for the vLLM 0.27-based +vLLM-Ascend branch. + +The full W4A8 checkpoint is approximately 1.49 TB. Download it from +[ModelScope](https://www.modelscope.cn/models/sgl-npu/Kimi-K3-W4A8), or make it +available from shared storage at the same path on every serving node. + +## 2 Supported and Validated Features + +The integration contains the following model paths. The table distinguishes +runtime validation from implementation-only coverage. + +| Capability | Status in this integration | +| --- | --- | +| Text and multimodal serving | Validated on Atlas A3 | +| TP16 with expert parallelism | Validated on one and multiple nodes | +| `FULL_DECODE_ONLY` ACL Graph | Validated | +| Prefix Cache and hybrid KDA/MLA state | Validated | +| Prefill-Decode disaggregation | Validated on two nodes | +| GQA and MLA DSpark adapters | Supported with a matching draft checkpoint; see Section 6 | +| MTP adapter | Implemented; validate the target and draft checkpoint pair separately | +| Atlas A5 SiTU MX quantization | Cross-built; target-hardware execution is still required | + +Refer to the [supported features](../../user_guide/support_matrix/supported_features.md) +for the project-wide feature matrix. + +## 3 Choosing a Checkpoint + +Kimi K3 checkpoints serve different validation purposes. A reduced or dummy +checkpoint must not be used to report GPQA or other semantic accuracy. + +| Checkpoint | Typical storage | What it can validate | What it cannot validate | +| --- | ---: | --- | --- | +| Full 93-layer, 896-expert W4A8 | About 1.49 TB | Deployment and benchmark accuracy | Not applicable | +| Full 93-layer, 16-expert derivative | About 113 GB | Single-node integration and long-context execution | Full-model semantics and full expert routing | +| Five-layer, 16-expert W4A8 derivative | About 12.1 GiB | Real quantized loading, KDA/MLA/MoE execution, graph and cache parity | Full-depth behavior and benchmark accuracy | +| Five-layer, 16-expert dummy | Configuration only | Model construction, TP16/EP, graph replay, cache-state parity, finite outputs | Real weight loading, W4A8 numerics, or semantic accuracy | + +The nightly test uses the last option so it does not depend on a large model +cache. Its configuration preserves the production hidden size, head geometry, +mixed KDA/MLA layout, attention residuals, and top-16 expert routing. It reduces +only the layer and expert counts, then initializes BF16 dummy weights. + +For a storage-limited real-weight nightly artifact, derive the five-layer, +16-expert checkpoint from the W4A8 checkpoint as follows: + +1. Keep tokenizer and configuration metadata. +1. Keep all non-expert tensors needed by the first five layers. +1. Keep routed experts 0 through 15 in those layers and rewrite the Safetensors + index without renaming tensors. +1. Set `num_hidden_layers`, `num_experts`, and `num_experts_per_token` to 5, + 16, and 16 respectively. Keep KDA layers 1, 2, 3, and 5, and MLA layer 4. +1. Load the result with `--quantization ascend` and record deterministic token + and log-probability parity. Do not create a task-accuracy baseline from it. + +:::{note} +The reduced checkpoint is a CI fixture, not a model release. Keep the source +checkpoint revision and a manifest of retained tensors with the artifact so +that it can be reproduced when the quantized weights change. +::: + +## 4 Installation + +Use an Atlas A3 image containing the vLLM version pinned by this release and +the matching vLLM-Ascend build. Mount the checkpoint at the same path on all +nodes. + +For multi-node serving, first follow the +[multi-node communication check](../../installation.md#verify-multi-node-communication). +The commands below cover the A3 configurations validated for this integration. +The A2 deployment from the vLLM-Ascend 0.23 guide is not carried forward as a +validated main-branch configuration until it is rerun with the current runtime. + +```shell +export IMAGE= +export MODEL_ROOT= + +docker run --rm -it \ + --name vllm-kimi-k3 \ + --net=host \ + --shm-size=1g \ + --device /dev/davinci0 \ + --device /dev/davinci1 \ + --device /dev/davinci2 \ + --device /dev/davinci3 \ + --device /dev/davinci4 \ + --device /dev/davinci5 \ + --device /dev/davinci6 \ + --device /dev/davinci7 \ + --device /dev/davinci8 \ + --device /dev/davinci9 \ + --device /dev/davinci10 \ + --device /dev/davinci11 \ + --device /dev/davinci12 \ + --device /dev/davinci13 \ + --device /dev/davinci14 \ + --device /dev/davinci15 \ + --device /dev/davinci_manager \ + --device /dev/devmm_svm \ + --device /dev/hisi_hdc \ + -v /usr/local/dcmi:/usr/local/dcmi \ + -v /usr/local/Ascend/driver:/usr/local/Ascend/driver \ + -v "$MODEL_ROOT:$MODEL_ROOT" \ + "$IMAGE" bash +``` + +Verify the installed revisions before starting the service: + +```shell +python -c "import vllm, vllm_ascend; print(vllm.__version__, vllm_ascend.__version__)" +``` + +## 5 Single-Node Functional Deployment + +Use a full-depth 16-expert derivative for single-node functional testing. It +preserves every transformer layer but is not semantically equivalent to the +full 896-expert checkpoint. + +```shell +export MODEL_PATH= +export TOKENIZER_PATH= +export ASCEND_RT_VISIBLE_DEVICES=0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15 +export PYTORCH_NPU_ALLOC_CONF=expandable_segments:True +export HCCL_OP_EXPANSION_MODE=AIV +export HCCL_BUFFSIZE=1024 +export HCCL_BUFFSIZE_EP=2048 +export OMP_PROC_BIND=false +export OPENBLAS_NUM_THREADS=1 + +vllm serve "$MODEL_PATH" \ + --host 0.0.0.0 \ + --port 8000 \ + --served-model-name kimi-k3 \ + --tokenizer "$TOKENIZER_PATH" \ + --quantization ascend \ + --safetensors-load-strategy lazy \ + --tensor-parallel-size 16 \ + --enable-expert-parallel \ + --enable-prefix-caching \ + --max-model-len 133120 \ + --max-num-seqs 16 \ + --max-num-batched-tokens 8192 \ + --gpu-memory-utilization 0.85 \ + --reasoning-parser kimi_k3 \ + --tool-call-parser kimi_k3 \ + --compilation-config '{"cudagraph_mode":"FULL_DECODE_ONLY"}' +``` + +Lower `--max-model-len` and `--max-num-batched-tokens` for a short smoke test. +The larger values above require a capacity check with the exact checkpoint and +runner memory configuration. + +## 6 Four-Node Full-Checkpoint Deployment + +The full checkpoint uses four Atlas 800 A3 nodes in a DP4/TP16/EP64 mixed +deployment. Start Node 0 first. Each worker owns one global DP rank and joins +Node 0 through the DP RPC address. + +Set these variables on every node: + +```shell +export MODEL_PATH= +export TOKENIZER_PATH= +export LOCAL_IP= +export NODE0_IP= +export NIC_NAME= +export SERVICE_PORT=8000 +export RPC_PORT=13345 +export DP_SIZE=4 +export TP_SIZE=16 + +export HCCL_IF_IP=$LOCAL_IP +export GLOO_SOCKET_IFNAME=$NIC_NAME +export TP_SOCKET_IFNAME=$NIC_NAME +export HCCL_SOCKET_IFNAME=$NIC_NAME +export VLLM_ENGINE_READY_TIMEOUT_S=7200 +export VLLM_EXECUTE_MODEL_TIMEOUT_SECONDS=3000 +export PYTORCH_NPU_ALLOC_CONF=expandable_segments:True +export HCCL_BUFFSIZE=1024 +export HCCL_BUFFSIZE_EP=2048 +export HCCL_INTRA_PCIE_ENABLE=1 +export HCCL_INTRA_ROCE_ENABLE=0 +export HCCL_OP_EXPANSION_MODE=AIV +export OMP_PROC_BIND=false +export OPENBLAS_NUM_THREADS=1 +export ASCEND_RT_VISIBLE_DEVICES=0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15 +``` + +Run on Node 0: + +```shell +vllm serve "$MODEL_PATH" \ + --host 0.0.0.0 \ + --port $SERVICE_PORT \ + --served-model-name kimi-k3 \ + --tokenizer "$TOKENIZER_PATH" \ + --quantization ascend \ + --safetensors-load-strategy lazy \ + --tensor-parallel-size $TP_SIZE \ + --data-parallel-size $DP_SIZE \ + --data-parallel-size-local 1 \ + --data-parallel-address $LOCAL_IP \ + --data-parallel-rpc-port $RPC_PORT \ + --enable-expert-parallel \ + --enable-prefix-caching \ + --max-model-len 133120 \ + --max-num-seqs 16 \ + --max-num-batched-tokens 8192 \ + --gpu-memory-utilization 0.85 \ + --reasoning-parser kimi_k3 \ + --tool-call-parser kimi_k3 \ + --compilation-config '{"cudagraph_mode":"FULL_DECODE_ONLY"}' +``` + +Run on Nodes 1 through 3. Set `DP_START_RANK` to 1, 2, or 3 respectively: + +```shell +export DP_START_RANK=<1_OR_2_OR_3> + +vllm serve "$MODEL_PATH" \ + --headless \ + --host 0.0.0.0 \ + --port $SERVICE_PORT \ + --served-model-name kimi-k3 \ + --tokenizer "$TOKENIZER_PATH" \ + --quantization ascend \ + --safetensors-load-strategy lazy \ + --tensor-parallel-size $TP_SIZE \ + --data-parallel-size $DP_SIZE \ + --data-parallel-size-local 1 \ + --data-parallel-start-rank $DP_START_RANK \ + --data-parallel-address $NODE0_IP \ + --data-parallel-rpc-port $RPC_PORT \ + --enable-expert-parallel \ + --enable-prefix-caching \ + --max-model-len 133120 \ + --max-num-seqs 16 \ + --max-num-batched-tokens 8192 \ + --gpu-memory-utilization 0.85 \ + --reasoning-parser kimi_k3 \ + --tool-call-parser kimi_k3 \ + --compilation-config '{"cudagraph_mode":"FULL_DECODE_ONLY"}' +``` + +### 6.1 Enabling DSpark + +For the GQA path, a public draft checkpoint is +[`RadixArk/Kimi-K3-DSpark`](https://huggingface.co/RadixArk/Kimi-K3-DSpark). +For the MLA path, use a matching Kimi K3 MLA draft checkpoint. Add the same +speculative configuration to the Node 0 and headless worker commands: + +```shell +--speculative-config \ +'{ + "method": "dspark", + "model": "", + "num_speculative_tokens": 7, + "draft_tensor_parallel_size": 16, + "max_model_len": 4096, + "draft_sample_method": "greedy", + "enforce_eager": true +}' +``` + +The example uses seven draft tokens, so each proposal cycle contains seven +draft steps followed by one target-model verification step. Set +`draft_tensor_parallel_size` to the topology used to shard the draft model. + +K3 cache grouping is derived from the target and draft layer layouts rather +than a hard-coded TP16 layout. The grouping contracts cover TP8 and TP16, but +the full-checkpoint deployment documented above was validated with TP16. For a +different TP size, first verify that both checkpoints' head and hidden +dimensions are divisible by that TP size, then rerun the functional and +accuracy checks in Sections 8 and 9. + +## 7 Two-Node Prefill-Decode Deployment + +The functional Prefill-Decode (P/D) check uses one 16-NPU A3 Prefill node and +one 16-NPU A3 Decode node. Both engines use TP16/EP and the same checkpoint, +tokenizer, KDA/MLA cache layout, and model revision. Install Mooncake and check +the data-plane network as described in the +[multi-node Mooncake guide](../features/pd_disaggregation_mooncake_multi_node.md). + +Use these K3-specific settings in addition to the common model and environment +arguments from Section 5: + +| Setting | Prefill | Decode | +| --- | --- | --- | +| Parallelism | TP16/EP | TP16/EP | +| Execution mode | `--enforce-eager` | `FULL_DECODE_ONLY` | +| Hybrid state layout | `--mamba-cache-mode align` | `--mamba-cache-mode align` | +| KV role | `kv_producer` | `kv_consumer` | +| Prefix Cache | Enabled | Enabled | + +Add the following arguments to the Prefill service. Select a `kv_port` outside +Mooncake's reserved AscendDirectTransport range; for a 16-NPU node, use a port +of at least 36000. + +```shell +--enforce-eager \ +--mamba-cache-mode align \ +--kv-transfer-config \ +'{ + "kv_connector": "MooncakeConnectorV1", + "kv_role": "kv_producer", + "kv_port": "", + "kv_connector_extra_config": { + "prefill": {"dp_size": 1, "tp_size": 16}, + "decode": {"dp_size": 1, "tp_size": 16} + } +}' +``` + +Add the following arguments to the Decode service: + +```shell +--mamba-cache-mode align \ +--compilation-config '{"cudagraph_mode":"FULL_DECODE_ONLY"}' \ +--kv-transfer-config \ +'{ + "kv_connector": "MooncakeConnectorV1", + "kv_role": "kv_consumer", + "kv_port": "", + "kv_connector_extra_config": { + "prefill": {"dp_size": 1, "tp_size": 16}, + "decode": {"dp_size": 1, "tp_size": 16} + } +}' +``` + +Start the standard Mooncake proxy with the Prefill and Decode endpoints, then +send requests to the proxy rather than directly to an engine: + +```shell +python examples/disaggregated_prefill_v1/load_balance_proxy_server_example.py \ + --host 0.0.0.0 \ + --port 9000 \ + --prefiller-hosts \ + --prefiller-ports \ + --decoder-hosts \ + --decoder-ports +``` + +The proxy CLI may evolve with the shared P/D implementation. Treat the linked +Mooncake guide and `--help` output from the checked-out revision as +authoritative for proxy-only arguments. Do not change the K3 model, tokenizer, +TP size, or hybrid cache mode between the two engines. + +## 8 Functional Verification + +Run request generation inside the serving environment or its trusted service +network. Avoid sending a large benchmark load across a developer workstation +or VPN. + +```shell +curl http://:8000/v1/chat/completions \ + -H 'Content-Type: application/json' \ + -d '{ + "model": "kimi-k3", + "messages": [ + {"role": "user", "content": "Explain why prefix caching helps repeated long prompts."} + ], + "temperature": 0, + "max_tokens": 128, + "logprobs": true + }' +``` + +The response must be HTTP 200, contain one non-null choice, and contain the +requested number of finite token log probabilities unless generation reaches a +configured stop token. Repeat an identical prompt and confirm that Prefix Cache +metrics increase without changing deterministic output tokens. + +For a multimodal smoke test, replace the message content with an image and a +text instruction: + +```json +{ + "role": "user", + "content": [ + {"type": "image_url", "image_url": {"url": ""}}, + {"type": "text", "text": "Describe the image."} + ] +} +``` + +Run concurrency and benchmark requests from a host inside the trusted serving +network. A developer workstation or VPN should be used only for bounded smoke +requests. + +## 9 Accuracy and Nightly Validation + +Use the following validation ladder: + +1. The storage-free nightly guard loads the committed five-layer, 16-expert + configuration with dummy BF16 weights on one 16-NPU A3 node. In one model + instance it compares cold prefill, a Prefix Cache hit that leaves exactly + one token to prefill, and a post-reset cold run. It requires identical token + IDs, close and finite chosen-token log probabilities, complete outputs, and + `FULL_DECODE_ONLY` replay. +1. A hosted five-layer, 16-expert W4A8 fixture should run the same parity case + with real weights. This adds quantized loader and W4A8 numerical coverage + while staying small enough for a nightly worker. +1. Run GPQA with the full 93-layer, 896-expert checkpoint on the four-node + deployment. This is the semantic accuracy gate and cannot be replaced by + either reduced fixture. + +Keep the model revision, tokenizer, chat rendering, reasoning mode, sampling +parameters, dataset revision, and evaluator revision fixed when comparing +GPQA results. Record completed, failed, missing, and unparsed samples in +addition to the final score. Refer to [AISBench](../../developer_guide/evaluation/using_ais_bench.md) +or [lm_eval](../../developer_guide/evaluation/using_lm_eval.md) for evaluator +setup. + +## 10 Performance Evaluation + +Use [AISBench](../../developer_guide/evaluation/using_ais_bench.md) or the +[vLLM benchmark tools](https://docs.vllm.ai/en/latest/benchmarking/) from a +server-side load-generator environment. Record the checkpoint revision, +topology, graph mode, Prefix Cache setting, input/output lengths, concurrency, +completed requests, and error count together with throughput and latency. + +Reduced checkpoints are useful for execution and scaling comparisons, but +their throughput is not representative of the full 896-expert model. + +## 11 FAQ + +### The service returns an incomplete or null choice + +First check every rank log for a worker abort, a non-finite tensor, or a graph +replay failure. Then repeat the same deterministic request with log +probabilities and verify that every generated token has a finite chosen-token +log probability. An HTTP 200 response alone is not a sufficient pass condition. + +### A Prefix Cache hit fails only at block-size plus one token + +Use a prompt that leaves exactly one uncached token after a full cached block. +Compare its output tokens and chosen-token log probabilities against a cold +prefill and a post-reset cold run in the same model instance. This exercises +the one-token prefill classification without conflating it with decode. + +### P/D starts but requests hang + +Verify that both engines use `--mamba-cache-mode align`, the same model +revision and TP size, complementary Mooncake roles, reachable non-overlapping +KV ports, and topology values matching the actual Prefill and Decode groups. +Check both engine logs and the proxy response; engine health alone does not +prove hybrid KDA/MLA state transfer. diff --git a/docs/source/tutorials/models/index.md b/docs/source/tutorials/models/index.md index 57f8968bfca2..7e4071cf8e76 100644 --- a/docs/source/tutorials/models/index.md +++ b/docs/source/tutorials/models/index.md @@ -1,3 +1,5 @@ # Model Tutorials This section provides tutorials for different models of vLLM Ascend. + +- [Kimi K3](Kimi-K3.md) From 3f43cb8cc21aaa42c41d60102e41bcaee30b4d98 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Tue, 25 Aug 2026 08:42:23 -0500 Subject: [PATCH 18/50] fix(ops): use relative path for shared KDA adapter header Signed-off-by: maoxx241 --- csrc/attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h | 2 +- csrc/attention/kda_gate_cumsum/kda_gate_cumsum_torch_adpt.h | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/csrc/attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h b/csrc/attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h index ca968bafcea4..f76d474e34d4 100644 --- a/csrc/attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h +++ b/csrc/attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h @@ -9,7 +9,7 @@ #include #include -#include "attention/kda_torch_adpt_common.h" +#include "../kda_torch_adpt_common.h" namespace vllm_ascend { diff --git a/csrc/attention/kda_gate_cumsum/kda_gate_cumsum_torch_adpt.h b/csrc/attention/kda_gate_cumsum/kda_gate_cumsum_torch_adpt.h index ce9772eea260..19d15322189f 100644 --- a/csrc/attention/kda_gate_cumsum/kda_gate_cumsum_torch_adpt.h +++ b/csrc/attention/kda_gate_cumsum/kda_gate_cumsum_torch_adpt.h @@ -8,7 +8,7 @@ #include -#include "attention/kda_torch_adpt_common.h" +#include "../kda_torch_adpt_common.h" namespace vllm_ascend { From d533d5d3fc298e526629dc29fa5107915f7f09da Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Tue, 25 Aug 2026 14:59:40 -0500 Subject: [PATCH 19/50] fix(model): load QuaRot DSpark boundaries in modeling Load missing target embed and LM-head shards during DSpark weight loading, anti-rotate them in the draft basis, and mark the draft-owned boundaries explicitly. This keeps model-specific QuaRot handling out of the proposer while reading only the local TP vocab slice. Signed-off-by: maoxx241 --- tests/ut/model_executor/test_qwen3_dspark.py | 95 +++++++++++++- tests/ut/models/test_kimi_k3_adapter.py | 50 +++----- vllm_ascend/models/kimi_k3_dspark.py | 75 +++++++---- vllm_ascend/models/llama_eagle3.py | 74 ++++++++--- vllm_ascend/models/qwen3_dspark.py | 123 ++++++++++--------- 5 files changed, 283 insertions(+), 134 deletions(-) diff --git a/tests/ut/model_executor/test_qwen3_dspark.py b/tests/ut/model_executor/test_qwen3_dspark.py index d1cc7893dd1a..02c0f9bc31b8 100644 --- a/tests/ut/model_executor/test_qwen3_dspark.py +++ b/tests/ut/model_executor/test_qwen3_dspark.py @@ -19,11 +19,16 @@ from __future__ import annotations +import json +from types import SimpleNamespace from unittest.mock import patch import torch +from safetensors.torch import save_file +from torch import nn import vllm_ascend.models.qwen3_dspark as qwen3_dspark +from vllm_ascend.models.llama_eagle3 import load_quarot_target_layer class TestQwen3DSparkWeightLoading: @@ -45,7 +50,11 @@ def test_rotates_only_fc_weights(self) -> None: rotation_matrix = torch.tensor([[0.0, 1.0], [1.0, 0.0]]) fc_weight = torch.tensor([[1.0, 2.0], [3.0, 4.0]]) non_fc_weight = torch.tensor([[5.0, 6.0]]) - weights_to_load = [("model.fc.weight", fc_weight), ("model.embed_tokens.weight", non_fc_weight)] + weights_to_load = [ + ("model.fc.weight", fc_weight), + ("model.embed_tokens.weight", non_fc_weight), + ("lm_head.weight", non_fc_weight), + ] expected_fc_weight = torch.matmul(fc_weight, rotation_matrix) # Capture the final delegation without invoking the real model loader. @@ -59,8 +68,90 @@ def test_rotates_only_fc_weights(self) -> None: mock_get_rotation_matrix.assert_called_once_with(rotation_path) mock_parent_load_weights.assert_called_once() - assert model._shared_layer_rotation is rotation_matrix processed_weights = mock_parent_load_weights.call_args.args[0] torch.testing.assert_close(processed_weights[0][1], expected_fc_weight) torch.testing.assert_close(processed_weights[1][1], non_fc_weight) + torch.testing.assert_close(processed_weights[2][1], non_fc_weight) + + def test_quarot_loads_missing_boundaries_in_modeling(self) -> None: + model_cls = qwen3_dspark.AscendQwen3DSparkForCausalLM + model = model_cls.__new__(model_cls) + nn.Module.__init__(model) + model.rotation_path = "quarot.safetensors" + model.target_model_path = "/target" + model.enable_confidence_head = False + model.model = SimpleNamespace(embed_tokens=object()) + model.lm_head = object() + rotation = torch.eye(2) + + with ( + patch.object( + qwen3_dspark, + "get_rotation_matrix", + return_value=rotation, + ), + patch.object( + qwen3_dspark, + "load_quarot_target_layer", + ) as load_target_layer, + patch.object( + qwen3_dspark.Qwen3DSparkForCausalLM, + "load_weights", + ), + ): + model.load_weights([("model.fc.weight", torch.eye(2))]) + + assert load_target_layer.call_count == 2 + assert load_target_layer.call_args_list[0].args[:2] == ( + model.model.embed_tokens, + model.target_model_path, + ) + assert load_target_layer.call_args_list[1].args[:2] == ( + model.lm_head, + model.target_model_path, + ) + assert model.has_own_embed_tokens + assert model.has_own_lm_head + + +def test_load_quarot_target_layer_reads_local_vocab_shard(tmp_path) -> None: + weight_name = "language_model.model.embed_tokens.weight" + shard_name = "model-00001-of-00001.safetensors" + target_weight = torch.tensor( + [ + [1.0, 2.0], + [3.0, 4.0], + [5.0, 6.0], + [7.0, 8.0], + ] + ) + save_file({weight_name: target_weight}, tmp_path / shard_name) + (tmp_path / "model.safetensors.index.json").write_text( + json.dumps({"weight_map": {weight_name: shard_name}}), + encoding="utf-8", + ) + + layer = nn.Linear(2, 3, bias=False) + layer.weight.data.fill_(99) + layer.shard_indices = SimpleNamespace( + org_vocab_start_index=1, + org_vocab_end_index=3, + ) + rotation = torch.tensor([[0.0, 1.0], [1.0, 0.0]]) + + load_quarot_target_layer( + layer, + tmp_path, + (weight_name,), + rotation, + "test embedding", + ) + + expected = torch.cat( + ( + target_weight[1:3] @ rotation.T, + torch.zeros(1, 2), + ) + ) + torch.testing.assert_close(layer.weight, expected) diff --git a/tests/ut/models/test_kimi_k3_adapter.py b/tests/ut/models/test_kimi_k3_adapter.py index 505db8091a5a..8f466c54ccf8 100644 --- a/tests/ut/models/test_kimi_k3_adapter.py +++ b/tests/ut/models/test_kimi_k3_adapter.py @@ -680,6 +680,9 @@ def test_k3_dspark_reuses_modelslim_rotation_loader(monkeypatch): model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) nn.Module.__init__(model) model.rotation_path = "rotation.safetensors" + model.target_model_path = "/target" + model.model = SimpleNamespace(embed_tokens=object()) + model.lm_head = object() source_weights = [ ("context_proj.weight", torch.ones(2, 4)), ("context_norm.weight", torch.ones(2)), @@ -710,49 +713,34 @@ def load_weights(self, weights, *, mapper): "vllm_ascend.models.kimi_k3_dspark.process_weight", process_weight, ) + load_target_layer = MagicMock() + monkeypatch.setattr( + "vllm_ascend.models.kimi_k3_dspark.load_quarot_target_layer", + load_target_layer, + ) model.load_weights(iter(source_weights)) process_weight.assert_called_once() torch.testing.assert_close(process_weight.call_args.args[0], source_weights[0][1]) torch.testing.assert_close(process_weight.call_args.args[1], rotation) - assert model._shared_layer_rotation is rotation + assert load_target_layer.call_count == 2 + assert load_target_layer.call_args_list[0].args[:2] == ( + model.model.embed_tokens, + model.target_model_path, + ) + assert load_target_layer.call_args_list[1].args[:2] == ( + model.lm_head, + model.target_model_path, + ) + assert model.has_own_embed_tokens + assert model.has_own_lm_head assert seen_weights[0][0] == "context_proj.weight" assert seen_weights[0][1] is rotated_weight assert seen_weights[1][0] == source_weights[1][0] assert seen_weights[1][1] is source_weights[1][1] -def test_k3_dspark_prepares_unrotated_shared_layer(): - class NonCopyableCommGroup: - def __deepcopy__(self, memo): - del memo - raise TypeError("cannot pickle ProcessGroup") - - model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) - nn.Module.__init__(model) - rotation = torch.tensor([[0.0, 1.0], [-1.0, 0.0]]) - model._shared_layer_rotation = rotation - target = nn.Linear(2, 3, bias=False) - target.comm_group = NonCopyableCommGroup() - target_weight = torch.tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]) - target.weight.data.copy_(target_weight) - - prepared = model.prepare_shared_layer( - None, - target, - "draft embed_tokens.weight", - ) - - assert prepared is not None - assert prepared is not target - assert prepared.comm_group is target.comm_group - torch.testing.assert_close(prepared.weight, target_weight @ rotation.T) - torch.testing.assert_close(target.weight, target_weight) - model.finish_shared_layer_preparation() - assert model._shared_layer_rotation is None - - def test_k3_dspark_embed_input_ids_keeps_text_only_path(): model = _make_k3_dspark_for_embedding_test() diff --git a/vllm_ascend/models/kimi_k3_dspark.py b/vllm_ascend/models/kimi_k3_dspark.py index 5be9bd920d4c..5754bbac535d 100644 --- a/vllm_ascend/models/kimi_k3_dspark.py +++ b/vllm_ascend/models/kimi_k3_dspark.py @@ -10,6 +10,10 @@ from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ReplicatedLinear from vllm.model_executor.layers.logits_processor import LogitsProcessor +from vllm.model_executor.layers.vocab_parallel_embedding import ( + ParallelLMHead, + VocabParallelEmbedding, +) from vllm.model_executor.models.interfaces import MultiModalEmbeddings from vllm.model_executor.models.qwen3_dspark import DSparkMarkovHead from vllm.model_executor.models.utils import ( @@ -35,9 +39,13 @@ from vllm_ascend.models.llama_eagle3 import ( get_rotation_matrix, get_rotation_path, - prepare_quarot_shared_layer, + load_quarot_target_layer, +) +from vllm_ascend.models.qwen3_dspark import ( + TARGET_EMBED_WEIGHT_NAMES, + TARGET_LM_HEAD_WEIGHT_NAMES, + process_weight, ) -from vllm_ascend.models.qwen3_dspark import process_weight from vllm_ascend.ops.rotary_embedding import get_cos_and_sin_mla @@ -241,27 +249,20 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: scale=getattr(self.config, "logit_scale", 1.0), ) self.rotation_path = get_rotation_path(vllm_config) - self._shared_layer_rotation: torch.Tensor | None = None - - def prepare_shared_layer( - self, - draft_layer: nn.Module | None, - target_layer: nn.Module, - label: str, - ) -> nn.Module | None: - rotation = self._shared_layer_rotation - if rotation is None: - return None - draft_layer, self._shared_layer_rotation = prepare_quarot_shared_layer( - draft_layer, - target_layer, - rotation, - label, - ) - return draft_layer - - def finish_shared_layer_preparation(self) -> None: - self._shared_layer_rotation = None + self.target_model_path = vllm_config.model_config.model + if self.rotation_path is not None: + target_config = vllm_config.model_config.hf_text_config + model_prefix = maybe_prefix(prefix, "model") + self.model.embed_tokens = VocabParallelEmbedding( + target_config.vocab_size, + target_config.hidden_size, + prefix=maybe_prefix(model_prefix, "embed_tokens"), + ) + self.lm_head = ParallelLMHead( + target_config.vocab_size, + target_config.hidden_size, + prefix=maybe_prefix(prefix, "lm_head"), + ) def load_weights( self, @@ -275,10 +276,9 @@ def load_weights( interface without creating that extra packed parameter. """ loader = AutoWeightsLoader(self) - self._shared_layer_rotation = None + rotation_weight = None if self.rotation_path is not None: rotation_weight = get_rotation_matrix(self.rotation_path) - self._shared_layer_rotation = rotation_weight weights = ( ( name, @@ -286,7 +286,30 @@ def load_weights( ) for name, loaded_weight in weights ) - return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) + loaded_weights = loader.load_weights( + weights, + mapper=self.hf_to_vllm_mapper, + ) + if rotation_weight is not None: + assert self.model.embed_tokens is not None + assert self.lm_head is not None + load_quarot_target_layer( + self.model.embed_tokens, + self.target_model_path, + TARGET_EMBED_WEIGHT_NAMES, + rotation_weight, + "draft embed_tokens.weight", + ) + load_quarot_target_layer( + self.lm_head, + self.target_model_path, + TARGET_LM_HEAD_WEIGHT_NAMES, + rotation_weight, + "draft lm_head.weight", + ) + self.has_own_embed_tokens = True + self.has_own_lm_head = True + return loaded_weights def embed_input_ids( self, diff --git a/vllm_ascend/models/llama_eagle3.py b/vllm_ascend/models/llama_eagle3.py index 3f56e1e3965d..2bcfa0a01690 100644 --- a/vllm_ascend/models/llama_eagle3.py +++ b/vllm_ascend/models/llama_eagle3.py @@ -1,10 +1,11 @@ -import copy +import json import logging import os from collections.abc import Iterable from pathlib import Path import torch +from safetensors import safe_open from safetensors.torch import load_file from torch import nn from vllm.config import VllmConfig @@ -54,34 +55,71 @@ def get_rotation_matrix(rotation_path: Path | None) -> torch.Tensor: raise e +def _find_safetensors_weight( + model_path: Path, + weight_names: tuple[str, ...], +) -> tuple[Path, str]: + """Locate one target tensor without loading unrelated checkpoint shards.""" + for index_path in sorted(model_path.glob("*.safetensors.index.json")): + with index_path.open(encoding="utf-8") as index_file: + weight_map = json.load(index_file).get("weight_map", {}) + for weight_name in weight_names: + if shard_name := weight_map.get(weight_name): + return model_path / shard_name, weight_name + + for shard_path in sorted(model_path.glob("*.safetensors")): + with safe_open(shard_path, framework="pt", device="cpu") as shard: + shard_keys = set(shard.keys()) + for weight_name in weight_names: + if weight_name in shard_keys: + return shard_path, weight_name + + raise KeyError(f"None of {weight_names!r} was found in the target checkpoint at {model_path}.") + + @torch.inference_mode() -def prepare_quarot_shared_layer( - draft_layer: nn.Module | None, - target_layer: nn.Module, +def load_quarot_target_layer( + layer: nn.Module, + target_model_path: Path | str, + weight_names: tuple[str, ...], rotation: torch.Tensor, label: str, -) -> tuple[nn.Module, torch.Tensor]: - """Create a draft-owned target layer in the unrotated hidden basis.""" - if draft_layer is None: - comm_group = getattr(target_layer, "comm_group", None) - memo = {id(comm_group): comm_group} if comm_group is not None else None - draft_layer = copy.deepcopy(target_layer, memo) +) -> None: + """Load one target vocab shard into the draft's unrotated hidden basis.""" + target_model_path = Path(target_model_path) + shard_path, weight_name = _find_safetensors_weight( + target_model_path, + weight_names, + ) + shard_indices = getattr(layer, "shard_indices", None) + if shard_indices is None: + start_index = 0 + end_index = layer.weight.shape[0] + else: + start_index = shard_indices.org_vocab_start_index + end_index = shard_indices.org_vocab_end_index + + with safe_open(shard_path, framework="pt", device="cpu") as shard: + target_weight = shard.get_slice(weight_name)[start_index:end_index] rotation = rotation.to( - device=target_layer.weight.device, + device=layer.weight.device, dtype=torch.float32, ) - unrotated = torch.matmul( - target_layer.weight.data.to(torch.float32), - rotation.T, + target_weight = target_weight.to( + device=layer.weight.device, + dtype=torch.float32, ) - draft_layer.weight.data.copy_(unrotated.to(draft_layer.weight.dtype)) + aligned_weight = torch.matmul(target_weight, rotation.T) + loaded_rows = aligned_weight.shape[0] + layer.weight.data[:loaded_rows].copy_(aligned_weight.to(layer.weight.dtype)) + layer.weight.data[loaded_rows:].zero_() logger.info( - "[spec_decode/quarot] Copied and aligned shared %s (weight=%s).", + "[spec_decode/quarot] Loaded and aligned %s from %s (%s).", label, - tuple(draft_layer.weight.shape), + shard_path.name, + tuple(layer.weight.shape), ) - return draft_layer, rotation def compute_rotation_matrix3(Q: torch.Tensor) -> torch.Tensor: diff --git a/vllm_ascend/models/qwen3_dspark.py b/vllm_ascend/models/qwen3_dspark.py index ef50bd196f89..cd0866f7dc7b 100644 --- a/vllm_ascend/models/qwen3_dspark.py +++ b/vllm_ascend/models/qwen3_dspark.py @@ -1,4 +1,5 @@ from collections.abc import Iterable +from pathlib import Path import torch from torch import nn @@ -10,10 +11,19 @@ from vllm_ascend.models.llama_eagle3 import ( get_rotation_matrix, get_rotation_path, - prepare_quarot_shared_layer, + load_quarot_target_layer, ) from vllm_ascend.utils import vllm_version_is +TARGET_EMBED_WEIGHT_NAMES = ( + "language_model.model.embed_tokens.weight", + "model.embed_tokens.weight", +) +TARGET_LM_HEAD_WEIGHT_NAMES = ( + "language_model.lm_head.weight", + "lm_head.weight", +) + # Process the first linear weight with rotation matrix, if the target model uses rotary quantization def process_weight(linear_weight: torch.Tensor, rotation_weight: torch.Tensor): @@ -71,27 +81,7 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: prefix=maybe_prefix(model_prefix, "confidence_head"), ) self.rotation_path = get_rotation_path(vllm_config) if vllm_config.quant_config is not None else None - self._shared_layer_rotation: torch.Tensor | None = None - - def prepare_shared_layer( - self, - draft_layer: nn.Module | None, - target_layer: nn.Module, - label: str, - ) -> nn.Module | None: - rotation = self._shared_layer_rotation - if rotation is None: - return None - draft_layer, self._shared_layer_rotation = prepare_quarot_shared_layer( - draft_layer, - target_layer, - rotation, - label, - ) - return draft_layer - - def finish_shared_layer_preparation(self) -> None: - self._shared_layer_rotation = None + self.target_model_path = Path(vllm_config.model_config.model) @staticmethod def _get_confidence_relative_name( @@ -114,11 +104,12 @@ def compute_confidence(self, head_hidden: torch.Tensor, markov_embed: torch.Tens def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): all_weights = list(weights) - self._shared_layer_rotation = None + includes_embed_tokens = any("embed_tokens" in name for name, _ in all_weights) + includes_lm_head = any("lm_head" in name for name, _ in all_weights) + rotation_weight = None if self.rotation_path is not None: processed_weights: list[tuple[str, torch.Tensor]] = [] rotation_weight = get_rotation_matrix(self.rotation_path) - self._shared_layer_rotation = rotation_weight for name, loaded_weight in all_weights: if "fc." in name: loaded_weight = process_weight(loaded_weight, rotation_weight) @@ -128,37 +119,55 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): if not vllm_version_is("0.27.1"): # main (cdc4824a21): upstream load_weights already manages # confidence_head (vllm#47808). - super().load_weights(all_weights) - return - - base_weights: list[tuple[str, torch.Tensor]] = [] - confidence_weights: list[tuple[str, torch.Tensor]] = [] - - for name, loaded_weight in all_weights: - confidence_name = self._get_confidence_relative_name(name) - if confidence_name is None: - base_weights.append((name, loaded_weight)) - else: - confidence_weights.append((confidence_name, loaded_weight)) - - super().load_weights(base_weights) + result = super().load_weights(all_weights) + else: + base_weights: list[tuple[str, torch.Tensor]] = [] + confidence_weights: list[tuple[str, torch.Tensor]] = [] - if not self.enable_confidence_head: - return - - if not confidence_weights: - self.enable_confidence_head = False - return - - confidence_weights.sort(key=lambda item: item[0]) - loaded_parameters = AutoWeightsLoader(self.model.confidence_head).load_weights(confidence_weights) - expected_parameters = set(self.model.confidence_head.state_dict().keys()) - missing_parameters = expected_parameters - loaded_parameters - - if missing_parameters: - raise RuntimeError( - "Failed to load all confidence-head " - "parameters. Missing: " - f"{sorted(missing_parameters)}; loaded: " - f"{sorted(loaded_parameters)}" - ) + for name, loaded_weight in all_weights: + confidence_name = self._get_confidence_relative_name(name) + if confidence_name is None: + base_weights.append((name, loaded_weight)) + else: + confidence_weights.append((confidence_name, loaded_weight)) + + result = super().load_weights(base_weights) + + if self.enable_confidence_head: + if not confidence_weights: + self.enable_confidence_head = False + else: + confidence_weights.sort(key=lambda item: item[0]) + loaded_parameters = AutoWeightsLoader(self.model.confidence_head).load_weights(confidence_weights) + expected_parameters = set(self.model.confidence_head.state_dict().keys()) + missing_parameters = expected_parameters - loaded_parameters + + if missing_parameters: + raise RuntimeError( + "Failed to load all confidence-head " + "parameters. Missing: " + f"{sorted(missing_parameters)}; loaded: " + f"{sorted(loaded_parameters)}" + ) + + if rotation_weight is not None: + if not includes_embed_tokens: + load_quarot_target_layer( + self.model.embed_tokens, + self.target_model_path, + TARGET_EMBED_WEIGHT_NAMES, + rotation_weight, + "draft embed_tokens.weight", + ) + self.has_own_embed_tokens = True + if not includes_lm_head: + load_quarot_target_layer( + self.lm_head, + self.target_model_path, + TARGET_LM_HEAD_WEIGHT_NAMES, + rotation_weight, + "draft lm_head.weight", + ) + self.has_own_lm_head = True + + return result From 5b5512f04e121dd99a170b5439874753fb682d55 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Tue, 25 Aug 2026 14:59:47 -0500 Subject: [PATCH 20/50] refactor(spec_decode): remove QuaRot model hooks Keep the generic DSpark sharing path focused on sharing semantics. QuaRot boundary preparation is now owned by the model weight loaders, so remove the optional proposer callbacks and their hook-only tests. Signed-off-by: maoxx241 --- .../ut/spec_decode/test_llm_base_proposer.py | 55 ------------------- vllm_ascend/spec_decode/llm_base_proposer.py | 51 ++--------------- 2 files changed, 4 insertions(+), 102 deletions(-) diff --git a/tests/ut/spec_decode/test_llm_base_proposer.py b/tests/ut/spec_decode/test_llm_base_proposer.py index 91f458876e6b..2c8be1bfcd38 100644 --- a/tests/ut/spec_decode/test_llm_base_proposer.py +++ b/tests/ut/spec_decode/test_llm_base_proposer.py @@ -23,7 +23,6 @@ from unittest.mock import MagicMock, patch import pytest -import torch.nn as nn from vllm.config import CUDAGraphMode from vllm_ascend.spec_decode.llm_base_proposer import AscendSpecDecodeBaseProposer @@ -142,60 +141,6 @@ def test_load_model_reads_validated_draft_window_size(): assert proposer.draft_window_size == 4096 mock_adapter.assert_called_once_with(4096, 16, 8, 4, "cpu") - draft_model.finish_shared_layer_preparation.assert_called_once_with() - - -def test_dspark_embedding_sharing_uses_model_preparation_hook(): - proposer = AscendSpecDecodeBaseProposer.__new__(AscendSpecDecodeBaseProposer) - draft_embed = nn.Linear(2, 3, bias=False) - target_embed = nn.Linear(2, 3, bias=False) - prepared_embed = nn.Linear(2, 3, bias=False) - proposer.method = "dspark" - proposer.model = SimpleNamespace( - has_own_embed_tokens=False, - model=SimpleNamespace(embed_tokens=draft_embed), - prepare_shared_layer=MagicMock(return_value=prepared_embed), - ) - target_model = SimpleNamespace( - model=SimpleNamespace(embed_tokens=target_embed), - ) - - with patch("vllm_ascend.spec_decode.llm_base_proposer.get_pp_group") as mock_pp_group: - mock_pp_group.return_value.world_size = 1 - proposer._maybe_share_embeddings(target_model) - - proposer.model.prepare_shared_layer.assert_called_once_with( - draft_embed, - target_embed, - "draft embed_tokens.weight", - ) - assert proposer.model.model.embed_tokens is prepared_embed - - -def test_dspark_lm_head_sharing_uses_model_preparation_hook(): - proposer = AscendSpecDecodeBaseProposer.__new__(AscendSpecDecodeBaseProposer) - draft_lm_head = nn.Linear(2, 3, bias=False) - target_lm_head = nn.Linear(2, 3, bias=False) - prepared_lm_head = nn.Linear(2, 3, bias=False) - proposer.method = "dspark" - proposer.model = SimpleNamespace( - has_own_lm_head=False, - lm_head=draft_lm_head, - prepare_shared_layer=MagicMock(return_value=prepared_lm_head), - ) - proposer.vllm_config = SimpleNamespace( - compilation_config=SimpleNamespace(cudagraph_mode=CUDAGraphMode.NONE), - ) - proposer.use_cuda_graph = False - - proposer._maybe_share_lm_head(SimpleNamespace(lm_head=target_lm_head)) - - proposer.model.prepare_shared_layer.assert_called_once_with( - draft_lm_head, - target_lm_head, - "draft lm_head.weight", - ) - assert proposer.model.lm_head is prepared_lm_head class TestDisablePaddedDrafterBatchWithFullGraph: diff --git a/vllm_ascend/spec_decode/llm_base_proposer.py b/vllm_ascend/spec_decode/llm_base_proposer.py index f66431a5c68a..d6bfd42f7131 100644 --- a/vllm_ascend/spec_decode/llm_base_proposer.py +++ b/vllm_ascend/spec_decode/llm_base_proposer.py @@ -378,13 +378,6 @@ def load_model(self, model: nn.Module) -> None: self._maybe_share_embeddings(target_language_model) self._maybe_share_topk_indices(target_language_model) self._maybe_share_lm_head(model) - finish_shared_layer_preparation = getattr( - self.model, - "finish_shared_layer_preparation", - None, - ) - if callable(finish_shared_layer_preparation): - finish_shared_layer_preparation() if ( self.parallel_drafting @@ -469,31 +462,9 @@ def _maybe_share_embeddings(self, target_language_model: nn.Module) -> None: ) if share_embeddings: - draft_embed_tokens = getattr( - self.model.model, - "embed_tokens", - None, - ) - prepare_shared_layer = getattr( - self.model, - "prepare_shared_layer", - None, - ) - prepared_embed_tokens = ( - prepare_shared_layer( - draft_embed_tokens, - target_embed_tokens, - "draft embed_tokens.weight", - ) - if callable(prepare_shared_layer) - else None - ) - if prepared_embed_tokens is not None: - self.model.model.embed_tokens = prepared_embed_tokens - else: - if hasattr(self.model.model, "embed_tokens"): - del self.model.model.embed_tokens - self.model.model.embed_tokens = target_embed_tokens + if hasattr(self.model.model, "embed_tokens"): + del self.model.model.embed_tokens + self.model.model.embed_tokens = target_embed_tokens else: logger.info( "[spec_decode/base] PP>1: draft model loaded its own vocab embedding" @@ -541,21 +512,7 @@ def _maybe_share_lm_head(self, model: nn.Module) -> None: " is not trained." ) else: - prepare_shared_layer = getattr( - self.model, - "prepare_shared_layer", - None, - ) - prepared_lm_head = ( - prepare_shared_layer( - getattr(self.model, "lm_head", None), - target_lm_head, - "draft lm_head.weight", - ) - if callable(prepare_shared_layer) - else None - ) - self.model.lm_head = target_lm_head if prepared_lm_head is None else prepared_lm_head + self.model.lm_head = target_lm_head if self.method == "mtp" and self.vllm_config.model_config.is_deepseek_mla: for _, layer_module in self.model.model.layers.items(): From 6c039a1b6426fcd556d2654fdf1e64e932b1637f Mon Sep 17 00:00:00 2001 From: weinachuan Date: Fri, 21 Aug 2026 01:15:57 +0800 Subject: [PATCH 21/50] feat(ops): update fused chunk KDA forward Fuse gate processing into chunk_kda_fwd, add Ascend 950 support, and fix final-state accuracy in generic tail and dense arch35 paths. Add independent NPU regression coverage for tail metadata and determinism. Signed-off-by: weinachuan --- .pre-commit-config.yaml | 2 +- csrc/attention/chunk_kda_fwd/README.md | 114 + .../chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h | 184 +- csrc/attention/chunk_kda_fwd/docs/api.md | 150 + csrc/attention/chunk_kda_fwd/docs/design.md | 145 + .../chunk_kda_fwd/op_host/CMakeLists.txt | 21 +- .../arch35/chunk_kda_fwd_tiling_impl.h | 43 + .../op_host/chunk_kda_fwd_def.cpp | 89 +- .../op_host/chunk_kda_fwd_tiling.cpp | 466 +- .../op_host/chunk_kda_fwd_tiling.h | 68 +- .../op_host/op_api/aclnn_chunk_kda_fwd.cpp | 1299 ++--- .../op_host/op_api/aclnn_chunk_kda_fwd.h | 13 +- .../op_host/op_api/chunk_kda_fwd.cpp | 127 +- .../op_host/op_api/chunk_kda_fwd.h | 43 +- .../op_kernel/arch35/chunk_kda_fwd_finalize.h | 1446 +++++ .../op_kernel/arch35/chunk_kda_fwd_fwd_h.h | 1009 ++++ .../op_kernel/arch35/chunk_kda_fwd_impl.h | 51 + .../op_kernel/arch35/chunk_kda_fwd_post_wu.h | 2024 +++++++ .../op_kernel/arch35/chunk_kda_fwd_prepare.h | 5149 +++++++++++++++++ .../chunk_kda_fwd/op_kernel/chunk_kda_fwd.cpp | 2764 +-------- .../op_kernel/chunk_kda_fwd_common.h | 332 ++ .../op_kernel/chunk_kda_fwd_finalize.h | 847 +++ .../op_kernel/chunk_kda_fwd_post_wu.h | 982 ++++ .../op_kernel/chunk_kda_fwd_prepare.h | 2613 +++++++++ .../op_kernel/chunk_kda_fwd_varlen.h | 68 + .../op_kernel/kda_gate_cumsum_kernel.h | 746 +++ .../block/block_epilogue_gdn_fwdh_regbase.hpp | 216 + .../block/block_epilogue_gdn_fwdh_update.hpp | 261 +- .../block/block_epilogue_gdn_fwdh_vnew.hpp | 165 +- .../gemm/block/block_scheduler_gdn_fwd_h.hpp | 106 +- .../arch35/gemm/kernel/gdn_fwd_h_kernel.hpp | 832 ++- .../block/block_epilogue_gdn_fwdh_update.hpp | 571 ++ .../block/block_epilogue_gdn_fwdh_vnew.hpp | 480 ++ .../epilogue/gdn_fwd_h_epilogue_policies.hpp | 27 + .../gemm/block/block_scheduler_gdn_fwd_h.hpp | 433 ++ .../gemm/kernel/gdn_fwd_h_kernel.hpp | 814 +++ .../block/block_mmad_pingpong_tla.hpp | 1054 ++++ .../block/block_mmad_pingpong_tla_multi.hpp | 135 +- .../kernel_utils/tile/copy_l0c_to_ub.hpp | 414 ++ .../common/kernel_utils/vector/regbase.hpp | 147 + csrc/torch_binding.cpp | 4 +- csrc/torch_binding_meta.cpp | 109 +- .../singlecard_ops/test_chunk_kda_aclnn.py | 320 +- .../test_kimi_k3_chunk_kda_tail_npu.py | 280 + .../test_kimi_kda_ascendc_npu.py | 401 +- tests/ut/ops/test_kimi_kda.py | 79 +- vllm_ascend/ops/kimi_kda.py | 51 +- 47 files changed, 23449 insertions(+), 4245 deletions(-) create mode 100644 csrc/attention/chunk_kda_fwd/README.md create mode 100644 csrc/attention/chunk_kda_fwd/docs/api.md create mode 100644 csrc/attention/chunk_kda_fwd/docs/design.md create mode 100644 csrc/attention/chunk_kda_fwd/op_host/arch35/chunk_kda_fwd_tiling_impl.h create mode 100644 csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_finalize.h create mode 100644 csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_fwd_h.h create mode 100644 csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_impl.h create mode 100644 csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_post_wu.h create mode 100644 csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_prepare.h create mode 100644 csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_common.h create mode 100644 csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_finalize.h create mode 100644 csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_post_wu.h create mode 100644 csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_prepare.h create mode 100644 csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_varlen.h create mode 100644 csrc/attention/kda_gate_cumsum/op_kernel/kda_gate_cumsum_kernel.h create mode 100644 csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_regbase.hpp create mode 100644 csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/epilogue/block/block_epilogue_gdn_fwdh_update.hpp create mode 100644 csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp create mode 100644 csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/epilogue/gdn_fwd_h_epilogue_policies.hpp create mode 100644 csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/gemm/block/block_scheduler_gdn_fwd_h.hpp create mode 100644 csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/gemm/kernel/gdn_fwd_h_kernel.hpp create mode 100644 csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla.hpp create mode 100644 csrc/moe/common/kernel_utils/tile/copy_l0c_to_ub.hpp create mode 100644 csrc/moe/common/kernel_utils/vector/regbase.hpp create mode 100644 tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_k3_chunk_kda_tail_npu.py diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 979623fbd043..a290e95b0910 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -22,7 +22,7 @@ repos: args: [ --toml, pyproject.toml, '--skip', 'tests/prompts/**,./benchmarks/sonnet.txt,*tests/lora/data/**,build/**,./vllm_ascend.egg-info/**,typos.toml', - '-L', 'CANN,cann,NNAL,nnal,ASCEND,ascend,EnQue,CopyIn,ArchType,AND,ND,tbe,copyin,alog,outter,mata,PARD' + '-L', 'CANN,cann,NNAL,nnal,ASCEND,ascend,EnQue,CopyIn,ArchType,AND,ND,tbe,copyin,alog,outter,mata,PARD,uSeed,LoadIn' ] additional_dependencies: - tomli diff --git a/csrc/attention/chunk_kda_fwd/README.md b/csrc/attention/chunk_kda_fwd/README.md new file mode 100644 index 000000000000..eb622491c4bc --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/README.md @@ -0,0 +1,114 @@ +# ChunkKdaFwd + +## 功能 + +`ChunkKdaFwd` 对齐不涉及 CP 切分的 FLA `chunk_kda_fwd` 顶层语义。公共接口接收 raw gate 或已激活的 +自然对数 gate;Gate、Prepare、PostWu、FwdH 和 Finalize 均在一个物理 `ChunkKdaFwd` L0 内完成, +L2 不再拼接或依次发射多个阶段 L0。 + +Shape 符号与布局约定见 [KDA 模型符号表](../README.md#model-shape-symbols)。 + +## Gate 公式 + +令 `x = g + dt_bias`。逐 token、逐 K 维的自然对数衰减为: + +```text +use_gate_in_kernel = false: + gate = g + +use_gate_in_kernel = true, safe_gate = false: + gate = -exp(A_log) * softplus(x) + +use_gate_in_kernel = true, safe_gate = true: + gate = lower_bound * sigmoid(exp(A_log) * x) +``` + +随后在每个 chunk 内计算: + +```text +gk_i = cumsum(gate)_i / ln(2) +``` + +因此后续 `exp2(gk)` 与自然指数 gate 严格绑定,不暴露额外 gate scale。 + +## 输入 + +| 名称 | 必选性 | Shape/Dtype | 说明 | +| --- | --- | --- | --- | +| `q/k` | 必选 | 输入 layout 对应 Shape;FP16/BF16 | Query/Key | +| `v` | 必选 | 输入 layout 对应 Shape;与 q 同 dtype | Value | +| `g` | 必选 | 输入 layout 对应 K 维 Shape;FP32/BF16 | raw gate 或已激活自然对数 gate | +| `beta` | 必选 | 去掉 g 的 K 维;FP32/BF16 | Delta 系数 | +| `A_log` | 条件必选 | `[H_v]`,FP32 | `use_gate_in_kernel=true` 时必选 | +| `dt_bias` | 可选 | `[H_v*K]`,FP32 | gate bias | +| `initial_state` | 可选 | `[N,H_v,K,V]` 或 `[N,H_v,V,K]`,FP32 | 由 `state_v_first` 解释 | +| `cu_seqlens` | 可选 | `[N+1]`,INT64 | 变长序列 | +| `chunk_indices` | 可选 | `[2*N_c]`,INT64 | canonical chunk 顺序 | + +`layout` 只描述上述输入。BSND/TND 由 L2 使用 `l0op::Transpose` 转为内部 BNSD/NTD。 + +## 输出 + +Python 返回顺序为: + +```text +(attn_out, final_state, gk, Aqk, Akk, w, u, qg, kg, v_new, h, initial_state) +``` + +- `attn_out` 固定为 BSND/TND。 +- `final_state` 固定按序列排列,末两维服从 `state_v_first`。 +- `Aqk/Akk` 始终返回,固定为 head-major。 +- `gk/w/u/qg/kg/v_new` 是供反向使用的 head-major 中间量。 +- 公开 `h` 固定为 sequence-major;内部 `hCompute` 保持 head-major 供 Finalize 使用。 +- 第 12 个返回值是 Python 层对 `initial_state` 的原对象透传,不是 aclnn 输出。 + +输出保留策略对齐 fla-org +[`chunk_kda_fwd`](https://github.com/fla-org/flash-linear-attention/blob/0f0f0c97af39343855b43bbbaddcedfda5cb9d77/fla/ops/kda/chunk_fwd.py) +提交 `0f0f0c97af39343855b43bbbaddcedfda5cb9d77`: + +| 条件 | 返回 | +| --- | --- | +| `output_final_state=true` | 返回 `final_state`,否则为 `None` | +| `use_gate_in_kernel=false` 或 `disable_recompute=true` | 返回 `gk` | +| 始终 | 返回 `Aqk/Akk` | +| `disable_recompute=true` | 返回 `w/u/qg/kg/v_new` | +| `disable_recompute=true` 或 `return_intermediate_states=true` | 返回 `h` | + +这是 `fla_npu.ops.ascendc.chunk_kda_fwd` 的低层 12 返回值语义;不涉及 CP。aclnn L2 不接收 +`output_final_state/disable_recompute/return_intermediate_states`,每个可选输出是否写出仅由对应 +输出指针是否为空决定。`w/u/qg/kg/v_new/h` 的 L0 阶段固定写内部 compute 张量,L2 仅在 +对应指针非空时通过 `ViewCopy` 导出;`gkOut` 非空时直接复用为 `gkCompute`,避免目标场景 +额外复制整张 FP32 gate。内部 `hCompute` 是 FwdH 到 Finalize 的必需 head-major 阶段结果; +公开 `hOut` 非空时,L2 转为 sequence-major 后导出。`hOut` 为空时仍创建 `hCompute`,但不 +作为第 11 个 Python 返回值公开。 + +## 属性 + +| 名称 | 默认值 | 支持范围 | +| --- | --- | --- | +| `layout` | `BSND` | `BSND/BNSD/TND/NTD` | +| `scale` | 必传 | 通常为 `K**-0.5` | +| `chunk_size` | `64` | `64/128` | +| `output_final_state` | `false` | bool | +| `safe_gate` | `false` | bool | +| `lower_bound` | `-5.0` | safe raw gate 时 `[-5,0)` | +| `use_gate_in_kernel` | `false` | bool | +| `disable_recompute` | `false` | bool | +| `return_intermediate_states` | `false` | bool | +| `state_v_first` | `false` | bool | + +## 支持范围 + +- A2 (`ascend910b`)、A3 (`ascend910_93`)、A5 (`ascend950`)。 +- `K/V` 为 `[16,256]` 内 16 的倍数;交付重点覆盖 K=128、V=128/256。 +- `chunk_size` 为 64/128。 +- TND/NTD 均支持多 head。 +- 变长调用最多 1024 条逻辑序列,rank-4 变长输入要求 B=1。 + +## 验证 + +唯一用例规格是 `tests/op_cases/chunk_kda_fwd.json`。数值测试位于 +`tests/operators/chunk_kda_fwd/accuracy/`,性能使用 `tests/operators/chunk_kda_fwd/performance/profile.py` +和 `msopprof`。 + +完整 API 见 [API 文档](docs/api.md),阶段和内存设计见 [设计文档](docs/design.md)。 diff --git a/csrc/attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h b/csrc/attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h index f76d474e34d4..e7abfac6e20f 100644 --- a/csrc/attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h +++ b/csrc/attention/chunk_kda_fwd/chunk_kda_fwd_torch_adpt.h @@ -13,13 +13,15 @@ namespace vllm_ascend { -std::tuple +std::tuple, c10::optional, at::Tensor, at::Tensor, + c10::optional, c10::optional, c10::optional, + c10::optional, c10::optional, c10::optional, + c10::optional> chunk_kda_fwd( const at::Tensor &q, const at::Tensor &k, const at::Tensor &v, - const at::Tensor &gk, + const at::Tensor &g, const at::Tensor &beta, double scale, int64_t chunk_size, @@ -28,64 +30,99 @@ chunk_kda_fwd( c10::optional output_final_state, c10::optional cu_seqlens, c10::optional chunk_indices, - c10::optional return_intermediate, c10::optional safe_gate, - c10::optional transpose_state_layout) + c10::optional lower_bound, + c10::optional use_gate_in_kernel, + const c10::optional &A_log, + const c10::optional &dt_bias, + c10::optional disable_recompute, + c10::optional return_intermediate_states, + c10::optional state_v_first) { std::string layout_str(layout.data(), layout.size()); TORCH_CHECK(layout_str == "BSND" || layout_str == "BNSD" || layout_str == "TND" || layout_str == "NTD", "chunk_kda_fwd: layout must be one of BSND, BNSD, TND, NTD and must be uppercase."); - TORCH_CHECK(!safe_gate.value_or(false), "chunk_kda_fwd: safe_gate=True is not supported."); - TORCH_CHECK(!transpose_state_layout.value_or(false), - "chunk_kda_fwd: transpose_state_layout=True is not supported."); - TORCH_CHECK(chunk_size == 32 || chunk_size == 64 || chunk_size == 128, - "chunk_kda_fwd: chunk_size must be 32, 64 or 128."); + TORCH_CHECK(chunk_size == 64 || chunk_size == 128, "chunk_kda_fwd: chunk_size must be 64 or 128."); + bool output_final_state_ = output_final_state.value_or(false); + bool safe_gate_ = safe_gate.value_or(false); + double lower_bound_ = lower_bound.value_or(-5.0); + bool use_gate_in_kernel_ = use_gate_in_kernel.value_or(false); + bool disable_recompute_ = disable_recompute.value_or(false); + bool return_intermediate_states_ = return_intermediate_states.value_or(false); + bool state_v_first_ = state_v_first.value_or(false); bool is_tnd = layout_str == "TND"; bool is_ntd = layout_str == "NTD"; bool is_bsnd = layout_str == "BSND"; bool is_bnsd = layout_str == "BNSD"; bool is_rank3 = is_tnd || is_ntd; - TORCH_CHECK((is_rank3 && q.dim() == 3 && k.dim() == 3 && v.dim() == 3 && gk.dim() == 3 && beta.dim() == 2) || - (!is_rank3 && q.dim() == 4 && k.dim() == 4 && v.dim() == 4 && gk.dim() == 4 && beta.dim() == 3), - "chunk_kda_fwd: layout/rank mismatch."); + TORCH_CHECK( + (is_rank3 && q.dim() == 3 && k.dim() == 3 && v.dim() == 3 && g.dim() == 3 && beta.dim() == 2) || + (!is_rank3 && q.dim() == 4 && k.dim() == 4 && v.dim() == 4 && g.dim() == 4 && beta.dim() == 3), + "chunk_kda_fwd: input ranks do not match layout."); TORCH_CHECK(q.sizes() == k.sizes(), "chunk_kda_fwd: q and k must have identical shape."); auto q_sizes = q.sizes(); auto v_sizes = v.sizes(); - bool is_internal_layout = is_bnsd || is_ntd; int64_t B = is_rank3 ? 1 : q_sizes[0]; int64_t T = is_tnd ? q_sizes[0] : (is_ntd ? q_sizes[1] : (is_bnsd ? q_sizes[2] : q_sizes[1])); int64_t H = is_tnd ? q_sizes[1] : (is_ntd ? q_sizes[0] : (is_bnsd ? q_sizes[1] : q_sizes[2])); int64_t K = is_rank3 ? q_sizes[2] : q_sizes[3]; int64_t HV = is_tnd ? v_sizes[1] : (is_ntd ? v_sizes[0] : (is_bnsd ? v_sizes[1] : v_sizes[2])); int64_t V = is_rank3 ? v_sizes[2] : v_sizes[3]; - TORCH_CHECK(H > 0 && HV >= H, "chunk_kda_fwd: H and HV must be positive and H must be <= HV."); - TORCH_CHECK(H <= 128 && HV <= 128, "chunk_kda_fwd: H and HV must be <= 128."); - TORCH_CHECK(!is_tnd || H == 1, - "chunk_kda_fwd: TND layout with H > 1 is not supported; use NTD for multi-head rank3 input."); - check_kda_cu_seqlens(cu_seqlens, T, "chunk_kda_fwd"); - check_kda_chunk_indices(chunk_indices, cu_seqlens, chunk_size, "chunk_kda_fwd"); - TORCH_CHECK(!cu_seqlens.has_value() || is_rank3 || B == 1, - "chunk_kda_fwd: rank4 varlen input with cu_seqlens currently requires B=1."); - TORCH_CHECK(HV % H == 0, "chunk_kda_fwd: HV must be divisible by H."); + TORCH_CHECK(H > 0 && HV >= H && HV % H == 0 && H <= 128 && HV <= 128, + "chunk_kda_fwd: H/HV must satisfy 0 < H <= HV <= 128 and HV % H == 0."); + TORCH_CHECK(K >= 16 && K <= 256 && K % 16 == 0 && V >= 16 && V <= 256 && V % 16 == 0, + "chunk_kda_fwd: K/V must be multiples of 16 and no greater than 256."); TORCH_CHECK(q.scalar_type() == at::kHalf || q.scalar_type() == at::kBFloat16, "chunk_kda_fwd: q/k/v must use float16 or bfloat16."); TORCH_CHECK(k.scalar_type() == q.scalar_type() && v.scalar_type() == q.scalar_type(), "chunk_kda_fwd: q/k/v dtype must match."); - TORCH_CHECK(chunk_size == 64 && K >= 16 && V >= 16 && K % 16 == 0 && V % 16 == 0 && V <= 256 && - K * V >= 4 * 64 * 64 && K * V >= chunk_size * (K + V), - "chunk_kda_fwd: shape is outside the supported split cube/vector template."); + TORCH_CHECK((g.scalar_type() == at::kFloat || g.scalar_type() == at::kBFloat16) && + (beta.scalar_type() == at::kFloat || beta.scalar_type() == at::kBFloat16), + "chunk_kda_fwd: g and beta must be float32 or bfloat16."); + + check_kda_cu_seqlens(cu_seqlens, T, "chunk_kda_fwd"); + check_kda_chunk_indices(chunk_indices, cu_seqlens, chunk_size, "chunk_kda_fwd"); + TORCH_CHECK(!cu_seqlens.has_value() || is_rank3 || B == 1, + "chunk_kda_fwd: rank4 varlen input requires B=1."); + auto g_sizes = g.sizes(); + TORCH_CHECK((is_tnd && beta.sizes()[0] == T && beta.sizes()[1] == HV) || + (is_ntd && beta.sizes()[0] == HV && beta.sizes()[1] == T) || + (is_bsnd && beta.sizes()[0] == B && beta.sizes()[1] == T && beta.sizes()[2] == HV) || + (is_bnsd && beta.sizes()[0] == B && beta.sizes()[1] == HV && beta.sizes()[2] == T), + "chunk_kda_fwd: beta shape mismatch."); + TORCH_CHECK((is_tnd && v_sizes == at::IntArrayRef({T, HV, V}) && g_sizes == at::IntArrayRef({T, HV, K})) || + (is_ntd && v_sizes == at::IntArrayRef({HV, T, V}) && + g_sizes == at::IntArrayRef({HV, T, K})) || + (is_bsnd && v_sizes == at::IntArrayRef({B, T, HV, V}) && + g_sizes == at::IntArrayRef({B, T, HV, K})) || + (is_bnsd && v_sizes == at::IntArrayRef({B, HV, T, V}) && + g_sizes == at::IntArrayRef({B, HV, T, K})), + "chunk_kda_fwd: v/g shapes do not match layout."); int64_t seq_num = get_kda_seq_num(B, cu_seqlens); - at::Tensor initial_state_tensor = initial_state.value_or(at::Tensor()); - if (initial_state_tensor.defined()) { - TORCH_CHECK(initial_state_tensor.scalar_type() == at::kFloat, + TORCH_CHECK(seq_num <= 1024, "chunk_kda_fwd: at most 1024 sequences are supported."); + std::vector state_shape = state_v_first_ ? std::vector{seq_num, HV, V, K} + : std::vector{seq_num, HV, K, V}; + if (initial_state.has_value() && initial_state->defined()) { + TORCH_CHECK(initial_state->scalar_type() == at::kFloat, "chunk_kda_fwd: initial_state must be float32 when provided."); - TORCH_CHECK(initial_state_tensor.dim() == 4 && initial_state_tensor.size(0) == seq_num && - initial_state_tensor.size(1) == HV && initial_state_tensor.size(2) == K && - initial_state_tensor.size(3) == V, - "chunk_kda_fwd: initial_state must be [seq_num,Hv,K,V]."); + TORCH_CHECK(initial_state->sizes() == at::IntArrayRef(state_shape), + "chunk_kda_fwd: initial_state shape does not match state_v_first."); + } + if (use_gate_in_kernel_) { + TORCH_CHECK(A_log.has_value() && A_log->defined() && A_log->scalar_type() == at::kFloat && + A_log->sizes() == at::IntArrayRef({HV}), + "chunk_kda_fwd: A_log must be float32 [HV] when use_gate_in_kernel=True."); + if (dt_bias.has_value() && dt_bias->defined()) { + TORCH_CHECK(dt_bias->scalar_type() == at::kFloat && dt_bias->sizes() == at::IntArrayRef({HV * K}), + "chunk_kda_fwd: dt_bias must be float32 [HV*K]."); + } + if (safe_gate_) { + TORCH_CHECK(lower_bound_ >= -5.0 && lower_bound_ < 0.0, + "chunk_kda_fwd: lower_bound must be in [-5, 0)."); + } } std::vector generated_chunk_indices; @@ -100,43 +137,64 @@ chunk_kda_fwd( } int64_t total_chunks = get_kda_total_chunks(B, T, chunk_size, cu_seqlens, chunk_indices_for_call); - at::Tensor o = at::empty_like(v); - at::Tensor final_state_work = at::empty({seq_num, HV, K, V}, q.options().dtype(at::kFloat)); - at::Tensor aqk = is_rank3 ? (is_internal_layout ? at::empty({HV, T, chunk_size}, q.options()) : - at::empty({T, HV, chunk_size}, q.options())) : (is_internal_layout ? - at::empty({B, HV, T, chunk_size}, q.options()) : at::empty({B, T, HV, chunk_size}, q.options())); + std::vector attn_shape = is_rank3 ? std::vector{T, HV, V} + : std::vector{B, T, HV, V}; + std::vector matrix_shape = is_rank3 ? std::vector{HV, T, chunk_size} + : std::vector{B, HV, T, chunk_size}; + std::vector k_shape = is_rank3 ? std::vector{HV, T, K} + : std::vector{B, HV, T, K}; + std::vector v_shape = is_rank3 ? std::vector{HV, T, V} + : std::vector{B, HV, T, V}; + std::vector h_shape = + is_rank3 ? (state_v_first_ ? std::vector{total_chunks, HV, V, K} + : std::vector{total_chunks, HV, K, V}) + : (state_v_first_ ? std::vector{B, total_chunks, HV, V, K} + : std::vector{B, total_chunks, HV, K, V}); + + at::Tensor attn_out = at::empty(attn_shape, v.options()); + at::Tensor final_state = + output_final_state_ ? at::empty(state_shape, q.options().dtype(at::kFloat)) : at::Tensor(); + at::Tensor gk_out = (!use_gate_in_kernel_ || disable_recompute_) + ? at::empty(k_shape, q.options().dtype(at::kFloat)) + : at::Tensor(); + at::Tensor aqk = at::empty(matrix_shape, q.options()); at::Tensor akk = at::empty_like(aqk); - at::Tensor w = is_rank3 ? (is_internal_layout ? at::empty({HV, T, K}, q.options()) : - at::empty({T, HV, K}, q.options())) : (is_internal_layout ? - at::empty({B, HV, T, K}, q.options()) : at::empty({B, T, HV, K}, q.options())); - at::Tensor u = at::empty_like(v); - at::Tensor qg = at::empty_like(w); - at::Tensor kg = at::empty_like(w); - at::Tensor v_new = at::empty_like(v); - at::Tensor h = is_rank3 ? (is_internal_layout ? at::empty({HV, total_chunks, K, V}, q.options()) : - at::empty({total_chunks, HV, K, V}, q.options())) : (is_internal_layout ? - at::empty({B, HV, total_chunks, K, V}, q.options()) : - at::empty({B, total_chunks, HV, K, V}, q.options())); + at::Tensor w = disable_recompute_ ? at::empty(k_shape, q.options()) : at::Tensor(); + at::Tensor u = disable_recompute_ ? at::empty(v_shape, q.options()) : at::Tensor(); + at::Tensor qg = disable_recompute_ ? at::empty(k_shape, q.options()) : at::Tensor(); + at::Tensor kg = disable_recompute_ ? at::empty(k_shape, q.options()) : at::Tensor(); + at::Tensor v_new = disable_recompute_ ? at::empty(v_shape, q.options()) : at::Tensor(); + at::Tensor h = (disable_recompute_ || return_intermediate_states_) + ? at::empty(h_shape, q.options()) + : at::Tensor(); - bool recompute_output_final_state = true; - char *layout_cstr = const_cast(layout_str.c_str()); + const at::Tensor &initial_state_ = c10::value_or_else(initial_state, [] { return at::Tensor(); }); + const at::Tensor &A_log_ = c10::value_or_else(A_log, [] { return at::Tensor(); }); + const at::Tensor &dt_bias_ = c10::value_or_else(dt_bias, [] { return at::Tensor(); }); + const char *layout_cstr = layout_str.c_str(); EXEC_NPU_CMD( aclnnChunkKdaFwd, - q, k, v, gk, beta, initial_state_tensor, - cu_seqlens, chunk_indices_for_call, - layout_cstr, scale, chunk_size, recompute_output_final_state, total_chunks, - o, final_state_work, aqk, akk, w, u, qg, kg, v_new, h + q, k, v, g, beta, A_log_, dt_bias_, initial_state_, cu_seqlens, chunk_indices_for_call, + layout_cstr, scale, chunk_size, safe_gate_, lower_bound_, use_gate_in_kernel_, state_v_first_, + attn_out, final_state, gk_out, aqk, akk, w, u, qg, kg, v_new, h ); - at::Tensor final_state = output_final_state.value_or(false) ? - final_state_work : at::empty({0}, q.options().dtype(at::kFloat)); - at::Tensor empty = at::empty({0}, q.options()); - at::Tensor g = gk.scalar_type() == at::kFloat ? gk : gk.to(at::kFloat); - at::Tensor initial_state_out = initial_state_tensor.defined() ? initial_state_tensor : empty; - (void)return_intermediate; - return std::make_tuple(o, final_state, g, aqk, akk, w, u, qg, kg, v_new, h, initial_state_out); + c10::optional final_state_out = + final_state.defined() ? c10::optional(final_state) : c10::nullopt; + c10::optional gk_optional = + gk_out.defined() ? c10::optional(gk_out) : c10::nullopt; + c10::optional w_optional = w.defined() ? c10::optional(w) : c10::nullopt; + c10::optional u_optional = u.defined() ? c10::optional(u) : c10::nullopt; + c10::optional qg_optional = qg.defined() ? c10::optional(qg) : c10::nullopt; + c10::optional kg_optional = kg.defined() ? c10::optional(kg) : c10::nullopt; + c10::optional v_new_optional = + v_new.defined() ? c10::optional(v_new) : c10::nullopt; + c10::optional h_optional = h.defined() ? c10::optional(h) : c10::nullopt; + c10::optional initial_state_out = + initial_state.has_value() && initial_state->defined() ? initial_state : c10::nullopt; + return std::make_tuple(attn_out, final_state_out, gk_optional, aqk, akk, w_optional, u_optional, + qg_optional, kg_optional, v_new_optional, h_optional, initial_state_out); } - } // namespace vllm_ascend #endif diff --git a/csrc/attention/chunk_kda_fwd/docs/api.md b/csrc/attention/chunk_kda_fwd/docs/api.md new file mode 100644 index 000000000000..f07fdb8e77a3 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/docs/api.md @@ -0,0 +1,150 @@ +# ChunkKdaFwd API + +## Python 主入口 + +```python +from fla_npu.ops.ascendc import chunk_kda_fwd + +outputs = chunk_kda_fwd( + q, k, v, g, beta, scale, chunk_size, + layout="BSND", + initial_state=None, + output_final_state=False, + cu_seqlens=None, + chunk_indices=None, + safe_gate=False, + lower_bound=None, + use_gate_in_kernel=False, + A_log=None, + dt_bias=None, + disable_recompute=False, + return_intermediate_states=False, + state_v_first=False, +) +``` + +返回: + +```text +(attn_out, final_state, gk, Aqk, Akk, w, u, qg, kg, v_new, h, initial_state) +``` + +可选输出在 Python 层返回 `None`。`Aqk/Akk` 始终存在;其余保留策略见算子 README。 + +## aclnn + +```cpp +aclnnStatus aclnnChunkKdaFwdGetWorkspaceSize( + const aclTensor *q, + const aclTensor *k, + const aclTensor *v, + const aclTensor *g, + const aclTensor *beta, + const aclTensor *aLogOptional, + const aclTensor *dtBiasOptional, + const aclTensor *initialStateOptional, + const aclIntArray *cuSeqlensOptional, + const aclIntArray *chunkIndicesOptional, + const char *layout, + double scale, + int64_t chunkSize, + bool safeGate, + double lowerBound, + bool useGateInKernel, + bool stateVFirst, + const aclTensor *attnOut, + const aclTensor *finalStateOut, + const aclTensor *gkOut, + const aclTensor *aqkOut, + const aclTensor *akkOut, + const aclTensor *wOut, + const aclTensor *uOut, + const aclTensor *qgOut, + const aclTensor *kgOut, + const aclTensor *vNewOut, + const aclTensor *hOut, + uint64_t *workspaceSize, + aclOpExecutor **executor); + +aclnnStatus aclnnChunkKdaFwd( + void *workspace, + uint64_t workspaceSize, + aclOpExecutor *executor, + aclrtStream stream); +``` + +aclnn L2 只描述张量与算法契约,不接收或解释 autograd 重计算策略: + +- `attnOut/aqkOut/akkOut` 是必选输出。 +- `finalStateOut/gkOut/wOut/uOut/qgOut/kgOut/vNewOut/hOut` 均为相互独立的可选输出。 +- `w/u/qg/kg/vNew/h` 的 L0 阶段固定写内部 compute 张量;对应可选输出非空时,L2 通过 + `ViewCopy` 导出,为空时只保留前向内部生命周期。`gkOut` 非空时直接复用为 `gkCompute`, + 避免目标场景额外复制整张 FP32 gate。 +- `finalStateOut != nullptr` 同时表示本次需要计算并写出最终状态。 +- `hCompute` 是 FwdH 到 Finalize 的内部必需 head-major 张量;`hOut` 是独立的公开可选输出。 + `hOut == nullptr` 不会跳过内部 `hCompute`,只是不向调用方公开该中间状态;非空时由 + L2 转为固定 sequence-major 后导出。 + +`output_final_state/disable_recompute/return_intermediate_states` 只存在于 Python 和 legacy torch +包装层,由上层按 FLA 的保留策略决定向 L2 传入哪些输出指针。 + +## 输入与输出布局 + +`layout` 只解释 q/k/v/g/beta 输入。输出固定为: + +- `attnOut`: BSND 或 TND。 +- `finalStateOut`: `[N,H_v,K,V]` 或 `stateVFirst=true` 时 `[N,H_v,V,K]`。 +- `gkOut/AqkOut/AkkOut/wOut/uOut/qgOut/kgOut/vNewOut`: BNSD/NTD。 +- `hOut`: dense 为 `[B,N_c,H_v,K,V]`,varlen 为 `[N_c,H_v,K,V]`; + `stateVFirst=true` 时交换末两维。 + +完整 Shape 表见 [KDA 模型符号表](../../README.md#model-shape-symbols)。 + +## Gate 语义 + +```text +useGateInKernel=false: + gate = g +useGateInKernel=true, safeGate=false: + gate = -exp(A_log) * softplus(g + dt_bias) +useGateInKernel=true, safeGate=true: + gate = lowerBound * sigmoid(exp(A_log) * (g + dt_bias)) +gk = chunk_local_cumsum(gate) / ln(2) +``` + +`safeGate` 的 true/false 都支持;`useGateInKernel=false` 时仍支持 `safeGate=true` 的后续稳定计算路径。 + +## 示例 + +```python +import torch +from fla_npu.ops.ascendc import chunk_kda_fwd + +B, T, H, K, V = 1, 128, 4, 128, 128 +q = torch.randn(B, T, H, K, device="npu", dtype=torch.float16) +k = torch.randn_like(q) +v = torch.randn(B, T, H, V, device="npu", dtype=torch.float16) +g = -torch.rand(B, T, H, K, device="npu", dtype=torch.float32) * 0.01 +beta = torch.rand(B, T, H, device="npu", dtype=torch.float32) + +attn_out, final_state, *_ = chunk_kda_fwd( + q, k, v, g, beta, K ** -0.5, 64, + layout="BSND", + output_final_state=True, + safe_gate=True, +) +assert attn_out.shape == (B, T, H, V) +assert final_state.shape == (B, H, K, V) +``` + +## 调用途径 + +| 路径 | 入口 | +| --- | --- | +| 稳定 Python | `fla_npu.ops.ascendc.chunk_kda_fwd` | +| aclnn | `aclnnChunkKdaFwdGetWorkspaceSize/aclnnChunkKdaFwd` | +| legacy | 显式加载后的 `torch.ops.npu.npu_chunk_kda_fwd` | +| 受限直调样例 | `torch.ops.ascend_ops.chunk_kda_fwd_direct` | + +直调样例仅覆盖 dense BNSD、K=128、V=128/256,并保留“调用方传入已累计 gk”的低层测试接口; +公开顶层语义以稳定 Python/aclnn 接口为准。 diff --git a/csrc/attention/chunk_kda_fwd/docs/design.md b/csrc/attention/chunk_kda_fwd/docs/design.md new file mode 100644 index 000000000000..6bc8f8284fec --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/docs/design.md @@ -0,0 +1,145 @@ +# ChunkKdaFwd 设计 + +## 目标 + +1. 顶层接口对齐非 CP 的 FLA `chunk_kda_fwd`。 +2. 不新增公开算子原型;A5 快路径复用既有 `ChunkKdaFwd` 原型和外层 kernel 入口。 +3. A2/A3/A5 使用同一数学定义;A5 保留 regbase 双发射特化。 +4. 输入 layout 与输出 layout 解耦。 +5. FwdH 同时服务 KDA 与 GDN,并支持可选 scalar gate、key-wise gate 和 `state_v_first`。 + +## L2 调度 + +```text +raw g -> ChunkKdaFwd[ + gate cumsum -> Prepare/Post-WU -> FwdH -> Finalize +] -> attn_out +``` + +`aclnnChunkKdaFwd` 做公开 layout 的连续化和必要视图转换。A5 的 BF16、chunk=64、K=V=128 +dense 对齐快路径保持单次物理 `ChunkKdaFwd` L0;A5 其他多 chunk 场景将同一个私有 L0 按 +Gate/Prepare、Post-WU、FwdH、Finalize 四个阶段依次提交,使阶段间通过物理 launch 边界重置事件状态。 +A2/A3 和单 chunk 场景仍使用单次物理 L0。阶段选择仅使用私有 `stage` 属性,不增加公开属性、 +接口字段或独立算子原型。 + +## 阶段职责 + +### KdaGateCumsum + +将 raw/已激活 gate 转为 FP32 chunk-local log2 累计值: + +```text +gk = cumsum(gate) / ln(2) +``` + +该算子同时保留独立 L2 接口供 GDN2 调用,输入和输出固定为 BNSD/NTD。 + +### Prepare + +只读取 `q/k/v/gk/beta` 及变长元数据,产生: + +```text +Aqk, Akk, qg, qg_scaled, w_seed, u_seed +``` + +矩阵计算和三角求逆使用 FP32 累积;公开中间量在写回时转为 q dtype。 + +### Post-WU + +只读取 `k/gk/w_seed/Akk/u_seed`,产生: + +```text +w, u, kg, v_new_seed +``` + +`Akk` 的 head 循环按 `H_v` 执行,GQA 映射只在读取 q/k head 时换算,避免按 `H_k` 重复或漏算。 + +### FwdH state propagation + +读取 `kg/w/u/gk` 和可选 `initial_state`,计算 chunk 间递推: + +```text +v_new = u - w @ h_prev +h_next = exp2(gk_last) * h_prev + kg^T @ v_new +``` + +arch35 路径复用与 `ChunkGatedDeltaRuleFwdH` 相同的数学实现;其他场景在 `ChunkKdaFwd` 内嵌 +共享 FwdH 实现。独立 GDN L0 原型继续保留给其他调用方,key-wise `gk` 固定使用 `exp2`。 + +### Finalize + +只读取 `qg_scaled/Aqk/v_new/h`,计算: + +```text +attn_out = qg_scaled @ h + Aqk @ v_new +``` + +kernel 内直接按 BSND/TND 写出 `attn_out`。供反向使用的中间量保持 BNSD/NTD。 + +## 状态布局 + +内部递推统一使用 `[...,K,V]`。`state_v_first=true` 时,L2 在进入 FwdH 前转置 initial state。 +内部 `hCompute` 始终保持 head-major 供 Finalize 消费;公开 `hOut` 在 L2 导出边界转为 +sequence-major,并按 `state_v_first` 决定末两维顺序。`final_state` 按序列排列,与 FLA 顶层 +输出一致。 + +## 重计算策略 + +L2 不理解 autograd 重计算策略。`final_state/gk/w/u/qg/kg/v_new/h` 是相互独立的 +`OPTIONAL_OUTPUT`;非空指针表示导出,空指针表示不公开该结果。单 launch 路径为隐藏输出传递 +固定 ABI 占位,并由 tiling 在 kernel workspace 中承接实际中间结果。A5 四段 launch 路径将 +阶段间依赖的 `gk/w/u/qg/kg/v_new/h/final_state` 和私有 `qg_scaled/u_seed` 物化为 executor 内部张量, +使后续 launch 不依赖前一 launch 的 kernel workspace。公开输出存在时直接作为内部目标使用。 + +Python/legacy 包装层对齐 fla-org `chunk_kda_fwd` 提交 +`0f0f0c97af39343855b43bbbaddcedfda5cb9d77`: + +- `Aqk/Akk` 始终返回。 +- `disable_recompute=false` 时不保留 `w/u/qg/kg/v_new`。 +- `disable_recompute=true` 或 `return_intermediate_states=true` 时保留公开 `hOut`。 +- `use_gate_in_kernel=false` 或 `disable_recompute=true` 时保留 `gk`。 +- `final_state` 只在 `output_final_state=true` 时创建公开输出。 + +内部 `hCompute` 与公开 `hOut` 是两个生命周期:`hCompute` 是 FwdH 到 Finalize 的必需 +head-major 阶段结果;`hOut` 为空时,单 launch 路径由 kernel workspace 承接,四段路径由 executor +内部张量承接。`hOut` 非空时,L2 提供 head-major 临时输出并在导出边界转为 sequence-major。 +该规则对齐非 CP 的低层 12 返回值接口; +第 12 项 `initial_state` 由 Python 层原对象透传。 + +## 模板化方案与 tiling key + +`ChunkKdaFwd` 只有一个外层 `op_kernel/chunk_kda_fwd.cpp` 入口和一个私有 L0 类型。A5 实现位于 +`op_kernel/arch35/*.h`,host 侧 A5 模板选择位于 `op_host/arch35/*.h`。Prepare、Post-WU、 +Finalize 的内部实现头与统一 kernel 入口同属 `chunk_kda_fwd/op_kernel/`,不存在对应的独立 L0 +原型或 `.cpp` 入口。A5 四段路径只是用不同私有 `stage` 属性连续调用该入口。 + +- `tiling key=1`:非 chunk=64、K=V=128 场景的通用模板族。 +- `tiling key=2`:chunk=64、K=V=128 模板族,包括 dense、tail 和 varlen。 + +两个 key 是同一 L0 的编译期场景变体,不是平台编号、独立算子或独立接口。A2/A3/A5 +均生成两个 key;同一个 key 内再由编译架构选择根目录通用实现或 `arch35/` 实现。host 的 +`SetTilingKey` 只检查 chunk、K、V,不检查 SoC。 + +在 arch35 上,key2 的 dense 对齐场景使用单 launch 和 arch35 FwdH;融合 score 写回在跳过 +共享 PostWU 时会额外物化以块尾 gate 为参考的最终 `kg`,供 FwdH 和可选公开输出共同使用。 +A5 多 chunk 的 tail/varlen 以及 key1 泛化场景使用四段 launch。其他架构在同一 key 下使用其 +对应单 launch 实现。tiling key 和私有 `stage` 均不改变公开算子原型、输出契约或数学定义。 + +## 性能设计 + +- Prepare 的右矩阵在 L1 驻留,避免 K/K^T 重复搬运和重复转置。 +- AIC 使用 L1/L0 双缓冲组织 MTE2、MTE1、Cube、Fixpipe。 +- AIV 使用输入 staging ping-pong,使下一 tile MTE2 与当前 tile VEC 重叠。 +- A5 VEC 路径使用 regbase 双发射;数值主计算仍保持 FP32。 +- inter-sub-chunk 合并使用独立 workspace 区域,避免阻塞主 tile 流水。 + +性能结论只使用 `msopprof`。目标回归 case 定义在 `tests/op_cases/chunk_kda_fwd.json`。 + +## 验证矩阵 + +- 平台:A2/A3/A5。 +- dtype:FP16/BF16。 +- layout:BSND/BNSD/TND/NTD。 +- gate:raw/已激活、safe true/false。 +- Shape:K=128,V=128/256,chunk=64/128,dense/varlen/tail/GQA。 +- 属性:final state、重计算策略、`state_v_first`。 diff --git a/csrc/attention/chunk_kda_fwd/op_host/CMakeLists.txt b/csrc/attention/chunk_kda_fwd/op_host/CMakeLists.txt index e6a6a541264d..2ebc27e7fdf9 100644 --- a/csrc/attention/chunk_kda_fwd/op_host/CMakeLists.txt +++ b/csrc/attention/chunk_kda_fwd/op_host/CMakeLists.txt @@ -5,11 +5,15 @@ # THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, # INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. # ----------------------------------------------------------------------------------------------------------- -set(CURRENT_CMAKE_DIR ${CMAKE_CURRENT_SOURCE_DIR}) set(CATLASS_INCLUDE_DIR "${CMAKE_SOURCE_DIR}/third_party/catlass/include") get_filename_component(CATLASS_INCLUDE_DIR_ABS ${CATLASS_INCLUDE_DIR} ABSOLUTE) +set(COMMON_KERNEL_UTILS_DIR "${CMAKE_SOURCE_DIR}/moe/common") +get_filename_component(COMMON_KERNEL_UTILS_DIR_ABS ${COMMON_KERNEL_UTILS_DIR} ABSOLUTE) add_op_to_compiled_list() +set(chunk_kda_fwd_depends + "attention/kda_gate_cumsum;moe/chunk_gated_delta_rule_fwd_h" + CACHE STRING "Kernel source dependencies for chunk_kda_fwd" FORCE) if (BUILD_OPEN_PROJECT) target_sources(op_host_aclnnExc PRIVATE chunk_kda_fwd_def.cpp @@ -18,11 +22,18 @@ endif() add_ops_compile_options( OP_NAME ChunkKdaFwd - OPTIONS - --cce-auto-sync=off - -Wno-deprecated-declarations - -I${CATLASS_INCLUDE_DIR_ABS} + OPTIONS --cce-auto-sync=off + -Wno-deprecated-declarations + -I${CATLASS_INCLUDE_DIR_ABS} + -I${COMMON_KERNEL_UTILS_DIR_ABS} ) +if("${ASCEND_COMPUTE_UNIT}" STREQUAL "ascend950") + add_ops_compile_options( + OP_NAME ChunkKdaFwd + COMPUTE_UNIT Ascend950PR_9599 + OPTIONS -DENABLE_CV_COMM_VIA_SSBUF=true + ) +endif() if (NOT BUILD_OPS_RTY_KERNEL) add_modules_sources(OPTYPE chunk_kda_fwd ACLNNTYPE aclnn_exclude) diff --git a/csrc/attention/chunk_kda_fwd/op_host/arch35/chunk_kda_fwd_tiling_impl.h b/csrc/attention/chunk_kda_fwd/op_host/arch35/chunk_kda_fwd_tiling_impl.h new file mode 100644 index 000000000000..a8b8cc2a4a8c --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_host/arch35/chunk_kda_fwd_tiling_impl.h @@ -0,0 +1,43 @@ +#pragma once + +namespace optiling::arch35 { + +struct ChunkKdaFwdArch35Options { + bool computeGateInPrepare = false; + bool fusePostWu = false; + bool fusePostWuIntoFwdH = false; + bool useDenseFwdH = false; +}; + +inline ChunkKdaFwdArch35Options ConfigureChunkKdaFwdArch35( + bool isAscend950, bool qIsBf16, bool rawGIsFp32, bool hasALog, + bool useGateInKernel, bool safeGate, bool isVarLen, int64_t seqlen, + int64_t vHeads, int64_t chunkSize, int64_t kDim, int64_t vDim, + bool storeQG, bool storeVNew, bool storeH) +{ + ChunkKdaFwdArch35Options options; + const bool shapeSupported = + isAscend950 && chunkSize == 64 && kDim == 128 && vDim == 128; + if (!shapeSupported) { + return options; + } + + // Tiling keys describe shape families independently of the SoC. These + // options only enable arch35 sub-pipelines within the selected family. + options.computeGateInPrepare = + qIsBf16 && rawGIsFp32 && hasALog && + useGateInKernel && safeGate; + const bool denseAligned = !isVarLen && seqlen % chunkSize == 0; + options.useDenseFwdH = denseAligned && qIsBf16; + const bool canFusePreparePostWu = + denseAligned && qIsBf16 && safeGate && vHeads % 2 == 0; + options.fusePostWuIntoFwdH = + options.useDenseFwdH && canFusePreparePostWu && + options.computeGateInPrepare && + !storeQG && !storeVNew && !storeH; + options.fusePostWu = + canFusePreparePostWu && !options.fusePostWuIntoFwdH; + return options; +} + +} // namespace optiling::arch35 diff --git a/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_def.cpp b/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_def.cpp index 90f6080f2c39..a18b7ce4f853 100644 --- a/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_def.cpp +++ b/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_def.cpp @@ -15,68 +15,72 @@ class ChunkKdaFwd : public OpDef { explicit ChunkKdaFwd(const char *name) : OpDef(name) { const std::initializer_list dataTypes = { - ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, - ge::DT_FLOAT16, ge::DT_BF16 + ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, + ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16 }; - const std::initializer_list stateTypes = { - ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, - ge::DT_FLOAT, ge::DT_FLOAT + const std::initializer_list gateTypes = { + ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_BF16, + ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_BF16 }; - const std::initializer_list akkTypes = { - ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, - ge::DT_FLOAT16, ge::DT_BF16 + const std::initializer_list betaTypes = { + ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT, ge::DT_BF16, + ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT, ge::DT_BF16 }; - const std::initializer_list outputDataTypes = { - ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT, - ge::DT_FLOAT16, ge::DT_BF16 + const std::initializer_list stateTypes = { + ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, + ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT }; const std::initializer_list formats = { - ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, - ge::FORMAT_ND, ge::FORMAT_ND + ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, + ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND }; this->Input("q").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); this->Input("k").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); this->Input("v").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); - this->Input("gk").ParamType(REQUIRED) - .DataType(stateTypes) + this->Input("g").ParamType(REQUIRED) + .DataType(gateTypes) .Format(formats).UnknownShapeFormat(formats); this->Input("beta").ParamType(REQUIRED) - .DataType(stateTypes) + .DataType(betaTypes) .Format(formats).UnknownShapeFormat(formats); + this->Input("a_log").ParamType(OPTIONAL).DataType(stateTypes).Format(formats).UnknownShapeFormat(formats); + this->Input("dt_bias").ParamType(OPTIONAL).DataType(stateTypes).Format(formats).UnknownShapeFormat(formats); this->Input("initial_state").ParamType(OPTIONAL).DataType(stateTypes).Format(formats).UnknownShapeFormat(formats); this->Input("cu_seqlens").ParamType(OPTIONAL).ValueDepend(OPTIONAL) - .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, - ge::DT_INT64, ge::DT_INT64}) + .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, + ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64}) .Format(formats).UnknownShapeFormat(formats); this->Input("chunk_indices").ParamType(OPTIONAL).ValueDepend(OPTIONAL) - .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, - ge::DT_INT64, ge::DT_INT64}) + .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, + ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64}) .Format(formats).UnknownShapeFormat(formats); - this->Input("stage_qg").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); - this->Input("stage_aqk").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); - this->Input("stage_v_new").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); - this->Input("stage_h").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); - this->Output("o").ParamType(REQUIRED).DataType(outputDataTypes).Format(formats).UnknownShapeFormat(formats); - this->Output("final_state").ParamType(REQUIRED).DataType(stateTypes).Format(formats).UnknownShapeFormat(formats); - this->Output("Aqk").ParamType(REQUIRED).DataType(stateTypes).Format(formats).UnknownShapeFormat(formats); - this->Output("Akk").ParamType(REQUIRED).DataType(akkTypes).Format(formats).UnknownShapeFormat(formats); - this->Output("w").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); - this->Output("u").ParamType(REQUIRED).DataType(outputDataTypes).Format(formats).UnknownShapeFormat(formats); - this->Output("qg").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); - this->Output("kg").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); - this->Output("v_new").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); - this->Output("h").ParamType(REQUIRED).DataType(stateTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("attn_out").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("final_state").ParamType(OPTIONAL).DataType(stateTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("gk").ParamType(OPTIONAL).DataType(stateTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("Aqk").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("Akk").ParamType(REQUIRED).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("w").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("u").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("qg").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("kg").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("v_new").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("h").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("qg_scaled").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Output("u_seed").ParamType(OPTIONAL).DataType(dataTypes).Format(formats).UnknownShapeFormat(formats); + this->Attr("layout").AttrType(OPTIONAL).String("BSND"); this->Attr("scale").AttrType(REQUIRED).Float(1.0); this->Attr("chunk_size").AttrType(REQUIRED).Int(64); - this->Attr("output_final_state").AttrType(REQUIRED).Bool(false); - this->Attr("total_chunks").AttrType(REQUIRED).Int(1); - this->Attr("stage").AttrType(OPTIONAL).Int(0); + this->Attr("safe_gate").AttrType(REQUIRED).Bool(false); + this->Attr("lower_bound").AttrType(OPTIONAL).Float(-5.0); + this->Attr("use_gate_in_kernel").AttrType(REQUIRED).Bool(false); + this->Attr("state_v_first").AttrType(OPTIONAL).Bool(false); + this->Attr("stage").AttrType(OPTIONAL).Int(-1); - OpAICoreConfig aicoreConfig; - aicoreConfig.DynamicCompileStaticFlag(true) + OpAICoreConfig config; + config.DynamicCompileStaticFlag(true) .DynamicFormatFlag(true) .DynamicRankSupportFlag(true) .DynamicShapeSupportFlag(true) @@ -85,10 +89,9 @@ class ChunkKdaFwd : public OpDef { .ExtendCfgInfo("prebuildPattern.value", "Opaque") .ExtendCfgInfo("coreType.value", "AiCore") .ExtendCfgInfo("aclnnSupport.value", "support_aclnn"); - - this->AICore().AddConfig("ascend910b", aicoreConfig); - this->AICore().AddConfig("ascend910_93", aicoreConfig); - this->AICore().AddConfig("ascend950", aicoreConfig); + this->AICore().AddConfig("ascend910b", config); + this->AICore().AddConfig("ascend910_93", config); + this->AICore().AddConfig("ascend950", config); } }; diff --git a/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.cpp b/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.cpp index 43d1b8e9c633..ff138b0c1b0f 100644 --- a/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.cpp +++ b/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.cpp @@ -1,181 +1,368 @@ -/** - * Copyright (c) 2026 Tianjin University, Ltd. - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of - * the BSD 3-Clause License (the "License"). - * Please refer to the License for details. You may not use this file except in compliance with the License. - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. - */ - #include "chunk_kda_fwd_tiling.h" + #include -#include +#include +#include #include +#include "arch35/chunk_kda_fwd_tiling_impl.h" #include "tiling/platform/platform_ascendc.h" namespace optiling { namespace { constexpr size_t INPUT_Q_IDX = 0; constexpr size_t INPUT_V_IDX = 2; -constexpr size_t INPUT_GK_IDX = 3; -constexpr size_t INPUT_INITIAL_IDX = 5; -constexpr size_t INPUT_CU_SEQLENS_IDX = 6; -constexpr size_t INPUT_CHUNK_INDICES_IDX = 7; -constexpr size_t ATTR_SCALE_IDX = 0; -constexpr size_t ATTR_CHUNK_SIZE_IDX = 1; -constexpr size_t ATTR_OUTPUT_FINAL_STATE_IDX = 2; -constexpr size_t ATTR_TOTAL_CHUNKS_IDX = 3; -constexpr size_t ATTR_STAGE_IDX = 4; +constexpr size_t INPUT_G_IDX = 3; +constexpr size_t INPUT_A_LOG_IDX = 5; +constexpr size_t INPUT_DT_BIAS_IDX = 6; +constexpr size_t INPUT_INITIAL_STATE_IDX = 7; +constexpr size_t INPUT_CU_SEQLENS_IDX = 8; +constexpr size_t INPUT_CHUNK_INDICES_IDX = 9; + +constexpr size_t OUTPUT_FINAL_STATE_IDX = 1; +constexpr size_t OUTPUT_GK_IDX = 2; +constexpr size_t OUTPUT_W_IDX = 5; +constexpr size_t OUTPUT_U_IDX = 6; +constexpr size_t OUTPUT_QG_IDX = 7; +constexpr size_t OUTPUT_KG_IDX = 8; +constexpr size_t OUTPUT_V_NEW_IDX = 9; +constexpr size_t OUTPUT_H_IDX = 10; + +constexpr size_t ATTR_LAYOUT_IDX = 0; +constexpr size_t ATTR_SCALE_IDX = 1; +constexpr size_t ATTR_CHUNK_SIZE_IDX = 2; +constexpr size_t ATTR_SAFE_GATE_IDX = 3; +constexpr size_t ATTR_LOWER_BOUND_IDX = 4; +constexpr size_t ATTR_USE_GATE_IDX = 5; +constexpr size_t ATTR_STAGE_IDX = 7; +constexpr int64_t KDA_STAGE_FULL = -1; +constexpr int64_t KDA_STAGE_FINALIZE = 3; + +constexpr uint64_t KDA_ALIGN = 512; constexpr uint64_t KDA_SOLVE_SCRATCH_SLOTS = 5; -constexpr uint64_t KDA_SCORE_QUEUE_SLOTS = 2; +constexpr uint64_t KDA_SOLVE_PIPELINE_DEPTH = 4; +constexpr uint64_t KDA_SCORE_QUEUE_SLOTS = 4; constexpr uint64_t KDA_SCORE_SCRATCH_PLANES = 3; -constexpr uint64_t KDA_FP32_BYTES = sizeof(float); -constexpr uint64_t KDA_WORKSPACE_ALIGN = 512; +constexpr uint64_t KDA_GDN_PIPELINE_DEPTH = 2; +constexpr uint32_t KDA_BATCH_MODE = 1; + +uint64_t AlignWorkspace(uint64_t bytes) +{ + return (bytes + KDA_ALIGN - 1) / KDA_ALIGN * KDA_ALIGN; +} + +uint64_t AllocateWorkspace(uint64_t &cursor, uint64_t bytes) +{ + const uint64_t offset = AlignWorkspace(cursor); + cursor = offset + bytes; + return offset; +} + +bool HasOutput(gert::TilingContext *context, size_t index) +{ + const auto instanceInfo = context->GetIrOutputInstanceInfo(index); + if (instanceInfo == nullptr || instanceInfo->GetInstanceNum() == 0) { + return false; + } + const auto outputShape = context->GetOutputShape(instanceInfo->GetInstanceStart()); + return outputShape != nullptr && + outputShape->GetStorageShape().GetShapeSize() != 1; +} -constexpr size_t DIM_B = 0; -constexpr size_t DIM_H = 1; -constexpr size_t DIM_T = 2; -constexpr size_t DIM_D = 3; +struct ShapeInfo { + int64_t rank = 0; + int64_t batch = 0; + int64_t seqlen = 0; + int64_t qHeads = 0; + int64_t vHeads = 0; + int64_t kDim = 0; + int64_t vDim = 0; + bool sequenceMajor = false; +}; -int64_t DTypeCode(ge::DataType dtype) +bool ResolveShape(gert::TilingContext *context, const char *layout, ShapeInfo &info) { - if (dtype == ge::DT_BF16) { - return 1; + const auto qShapePtr = context->GetInputShape(INPUT_Q_IDX); + const auto vShapePtr = context->GetInputShape(INPUT_V_IDX); + if (qShapePtr == nullptr || vShapePtr == nullptr || layout == nullptr) { + return false; + } + const auto &qShape = qShapePtr->GetStorageShape(); + const auto &vShape = vShapePtr->GetStorageShape(); + info.rank = qShape.GetDimNum(); + if (info.rank != vShape.GetDimNum() || (info.rank != 3 && info.rank != 4)) { + return false; + } + + info.sequenceMajor = std::strcmp(layout, "BSND") == 0 || std::strcmp(layout, "TND") == 0; + if (info.rank == 4) { + info.batch = qShape.GetDim(0); + if (info.sequenceMajor) { + info.seqlen = qShape.GetDim(1); + info.qHeads = qShape.GetDim(2); + info.vHeads = vShape.GetDim(2); + } else { + info.qHeads = qShape.GetDim(1); + info.vHeads = vShape.GetDim(1); + info.seqlen = qShape.GetDim(2); + } + info.kDim = qShape.GetDim(3); + info.vDim = vShape.GetDim(3); + } else { + info.batch = 1; + if (info.sequenceMajor) { + info.seqlen = qShape.GetDim(0); + info.qHeads = qShape.GetDim(1); + info.vHeads = vShape.GetDim(1); + } else { + info.qHeads = qShape.GetDim(0); + info.vHeads = vShape.GetDim(0); + info.seqlen = qShape.GetDim(1); + } + info.kDim = qShape.GetDim(2); + info.vDim = vShape.GetDim(2); + } + return info.batch > 0 && info.seqlen > 0 && info.qHeads > 0 && info.vHeads > 0 && + info.kDim > 0 && info.vDim > 0 && info.vHeads % info.qHeads == 0; +} + +bool ResolveSequenceInfo(gert::TilingContext *context, int64_t seqlen, int64_t chunkSize, + int64_t batch, bool &isVarLen, int64_t &seqNum, + int64_t &totalChunks) +{ + const auto cuTensor = context->GetOptionalInputTensor(INPUT_CU_SEQLENS_IDX); + isVarLen = cuTensor != nullptr; + seqNum = batch; + totalChunks = (seqlen + chunkSize - 1) / chunkSize; + if (!isVarLen) { + return totalChunks > 0; + } + + seqNum = cuTensor->GetStorageShape().GetDim(0) - 1; + const int64_t *cu = cuTensor->GetData(); + if (seqNum <= 0 || cu == nullptr || cu[0] != 0 || cu[seqNum] > seqlen) { + return false; + } + totalChunks = 0; + for (int64_t seq = 0; seq < seqNum; ++seq) { + if (cu[seq] < 0 || cu[seq + 1] < cu[seq]) { + return false; + } + totalChunks += (cu[seq + 1] - cu[seq] + chunkSize - 1) / chunkSize; } - if (dtype == ge::DT_FLOAT) { - return 2; + + const auto chunkShape = context->GetOptionalInputShape(INPUT_CHUNK_INDICES_IDX); + if (chunkShape != nullptr && + chunkShape->GetStorageShape().GetShapeSize() != totalChunks * 2) { + return false; } - return 0; + return totalChunks > 0; } } // namespace ge::graphStatus Tiling4ChunkKdaFwd(gert::TilingContext *context) { - ChunkKdaFwdTilingData tiling; + const auto qDesc = context->GetInputDesc(INPUT_Q_IDX); + const auto gDesc = context->GetInputDesc(INPUT_G_IDX); + const auto attrs = context->GetAttrs(); + if (qDesc == nullptr || gDesc == nullptr || attrs == nullptr) { + return ge::GRAPH_FAILED; + } - auto qShape = context->GetOptionalInputShape(INPUT_Q_IDX)->GetStorageShape(); - auto vShape = context->GetOptionalInputShape(INPUT_V_IDX)->GetStorageShape(); - auto qDesc = context->GetInputDesc(INPUT_Q_IDX); - auto gDesc = context->GetInputDesc(INPUT_GK_IDX); - if (qDesc == nullptr || gDesc == nullptr) { + const char *layout = attrs->GetStr(ATTR_LAYOUT_IDX); + const float scale = static_cast(*attrs->GetAttrPointer(ATTR_SCALE_IDX)); + const int64_t chunkSize = *attrs->GetAttrPointer(ATTR_CHUNK_SIZE_IDX); + const bool safeGate = *attrs->GetAttrPointer(ATTR_SAFE_GATE_IDX); + const float lowerBound = *attrs->GetAttrPointer(ATTR_LOWER_BOUND_IDX); + const bool useGateInKernel = *attrs->GetAttrPointer(ATTR_USE_GATE_IDX); + const int64_t stage = *attrs->GetAttrPointer(ATTR_STAGE_IDX); + if (chunkSize <= 0 || stage < KDA_STAGE_FULL || + stage > KDA_STAGE_FINALIZE) { return ge::GRAPH_FAILED; } - auto attrPtr = context->GetAttrs(); - if (attrPtr == nullptr) { + ShapeInfo shape; + if (!ResolveShape(context, layout, shape)) { return ge::GRAPH_FAILED; } - float scale = static_cast(*(attrPtr->GetAttrPointer(ATTR_SCALE_IDX))); - int64_t chunkSize = *(attrPtr->GetAttrPointer(ATTR_CHUNK_SIZE_IDX)); - bool outputFinalState = *(attrPtr->GetAttrPointer(ATTR_OUTPUT_FINAL_STATE_IDX)); - int64_t totalChunks = *(attrPtr->GetAttrPointer(ATTR_TOTAL_CHUNKS_IDX)); - int64_t stage = 0; - const int64_t *stagePtr = attrPtr->GetAttrPointer(ATTR_STAGE_IDX); - if (stagePtr != nullptr) { - stage = *stagePtr; + bool isVarLen = false; + int64_t seqNum = 0; + int64_t totalChunks = 0; + if (!ResolveSequenceInfo(context, shape.seqlen, chunkSize, shape.batch, + isVarLen, seqNum, totalChunks)) { + return ge::GRAPH_FAILED; } - bool isVarLen = context->GetOptionalInputTensor(INPUT_CU_SEQLENS_IDX) != nullptr; - int64_t batch = qShape.GetDim(DIM_B); - int64_t seqNum = batch; - std::array seqStart{}; - std::array seqEnd{}; - std::array seqChunkOffset{}; - if (isVarLen) { - auto cuTensor = context->GetOptionalInputTensor(INPUT_CU_SEQLENS_IDX); - seqNum = cuTensor->GetStorageShape().GetDim(0) - 1; - auto chunkMetadata = context->GetOptionalInputTensor(INPUT_CHUNK_INDICES_IDX); - if (seqNum <= 0 || seqNum > KDA_MAX_TILING_SEQUENCES || chunkMetadata == nullptr || - chunkMetadata->GetStorageShape().GetShapeSize() != totalChunks * 4) { - return ge::GRAPH_FAILED; - } - const int64_t *cu = cuTensor->GetData(); - if (cu == nullptr) { - return ge::GRAPH_FAILED; - } - int64_t chunkOffset = 0; - for (int64_t seq = 0; seq < seqNum; ++seq) { - if (cu[seq] < 0 || cu[seq + 1] < cu[seq]) { - return ge::GRAPH_FAILED; - } - seqStart[seq] = cu[seq]; - seqEnd[seq] = cu[seq + 1]; - seqChunkOffset[seq] = chunkOffset; - const int64_t seqLength = cu[seq + 1] - cu[seq]; - chunkOffset += (seqLength + chunkSize - 1) / chunkSize; - } - seqChunkOffset[seqNum] = chunkOffset; - if (chunkOffset != totalChunks) { - return ge::GRAPH_FAILED; - } - } - bool hasInitialState = context->GetOptionalInputTensor(INPUT_INITIAL_IDX) != nullptr; + const bool hasALog = context->GetOptionalInputDesc(INPUT_A_LOG_IDX) != nullptr; + const bool hasDtBias = context->GetOptionalInputDesc(INPUT_DT_BIAS_IDX) != nullptr; + const bool hasInitialState = context->GetOptionalInputDesc(INPUT_INITIAL_STATE_IDX) != nullptr; + const bool storeFinalState = HasOutput(context, OUTPUT_FINAL_STATE_IDX); + const bool storeGk = HasOutput(context, OUTPUT_GK_IDX); + const bool storeW = HasOutput(context, OUTPUT_W_IDX); + const bool storeU = HasOutput(context, OUTPUT_U_IDX); + const bool storeQG = HasOutput(context, OUTPUT_QG_IDX); + const bool storeKg = HasOutput(context, OUTPUT_KG_IDX); + const bool storeVNew = HasOutput(context, OUTPUT_V_NEW_IDX); + const bool storeH = HasOutput(context, OUTPUT_H_IDX); - const auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo()); - uint32_t coreNum = ascendcPlatform.GetCoreNumAic(); - int64_t taskNum = seqNum * vShape.GetDim(DIM_H); - if (stage == 1 || stage == 2 || stage == 3) { - taskNum = (isVarLen ? totalChunks : batch * totalChunks) * vShape.GetDim(DIM_H); - } - uint32_t blockDim = static_cast(std::min(taskNum, coreNum)); - if (stage == 1 || stage == 2 || stage == 3 || - (qDesc->GetDataType() != ge::DT_FLOAT && qShape.GetDim(DIM_D) >= 16)) { - blockDim = coreNum; - } - context->SetBlockDim(blockDim == 0 ? 1 : blockDim); - size_t *workspace = context->GetWorkspaceSizes(1); - uint64_t kernelScratch = 0; - if (stage == 1) { - const uint64_t usedCoreNum = static_cast(blockDim == 0 ? 1 : blockDim); - const uint64_t solveScratch = usedCoreNum * KDA_SOLVE_SCRATCH_SLOTS * - static_cast(chunkSize) * static_cast(chunkSize) * - KDA_FP32_BYTES; - const uint64_t alignedSolveScratch = - (solveScratch + KDA_WORKSPACE_ALIGN - 1) / KDA_WORKSPACE_ALIGN * KDA_WORKSPACE_ALIGN; - const uint64_t scoreElementBytes = qDesc->GetDataType() == ge::DT_FLOAT ? sizeof(float) : sizeof(uint16_t); - const uint64_t scoreScratch = usedCoreNum * KDA_SCORE_QUEUE_SLOTS * KDA_SCORE_SCRATCH_PLANES * - static_cast(chunkSize) * - static_cast(qShape.GetDim(DIM_D)) * scoreElementBytes; - kernelScratch = alignedSolveScratch + scoreScratch; - } else if (stage == 2) { - const uint64_t outputElements = static_cast(batch) * - static_cast(vShape.GetDim(DIM_H)) * - static_cast(qShape.GetDim(DIM_T)) * - static_cast(vShape.GetDim(DIM_D)); - kernelScratch = 2 * outputElements * KDA_FP32_BYTES; + const auto platform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo()); + const uint32_t blockDim = std::max(platform.GetCoreNumAic(), 1); + const bool isAscend950 = + platform.GetSocVersion() == platform_ascendc::SocVersion::ASCEND950; + const bool useChunk64K128V128Template = + chunkSize == 64 && shape.kDim == 128 && shape.vDim == 128; + const auto arch35Options = arch35::ConfigureChunkKdaFwdArch35( + isAscend950, qDesc->GetDataType() == ge::DT_BF16, + gDesc->GetDataType() == ge::DT_FLOAT, hasALog, useGateInKernel, + safeGate, isVarLen, shape.seqlen, shape.vHeads, chunkSize, + shape.kDim, shape.vDim, storeQG, storeVNew, storeH); + + const uint64_t dataBytes = + qDesc->GetDataType() == ge::DT_FLOAT ? sizeof(float) : sizeof(uint16_t); + const uint64_t tokenHeads = static_cast(shape.batch) * + shape.vHeads * shape.seqlen; + const uint64_t kTensorBytes = tokenHeads * shape.kDim * dataBytes; + const uint64_t vTensorBytes = tokenHeads * shape.vDim * dataBytes; + const uint64_t gkBytes = tokenHeads * shape.kDim * sizeof(float); + const uint64_t stateElements = static_cast(seqNum) * + shape.vHeads * shape.kDim * shape.vDim; + const uint64_t hChunkCount = isVarLen + ? static_cast(totalChunks) + : static_cast(shape.batch) * totalChunks; + const uint64_t hBytes = hChunkCount * shape.vHeads * shape.kDim * + shape.vDim * dataBytes; + + uint64_t cursor = 0; + const uint64_t gkStorageOffset = storeGk ? 0 : AllocateWorkspace(cursor, gkBytes); + const uint64_t finalStateStorageOffset = storeFinalState ? 0 : + AllocateWorkspace(cursor, stateElements * sizeof(float)); + const uint64_t wStorageOffset = storeW ? 0 : AllocateWorkspace(cursor, kTensorBytes); + const uint64_t uStorageOffset = storeU ? 0 : AllocateWorkspace(cursor, vTensorBytes); + const uint64_t qgStorageOffset = storeQG ? 0 : AllocateWorkspace(cursor, kTensorBytes); + const uint64_t kgStorageOffset = storeKg ? 0 : AllocateWorkspace(cursor, kTensorBytes); + const uint64_t vNewStorageBytes = + arch35Options.useDenseFwdH && !storeVNew + ? static_cast(shape.batch) * shape.vHeads * chunkSize * + shape.vDim * dataBytes + : vTensorBytes; + const uint64_t vNewStorageOffset = storeVNew ? 0 : + AllocateWorkspace(cursor, vNewStorageBytes); + const uint64_t hStorageBytes = + arch35Options.useDenseFwdH && !storeH + ? static_cast(shape.batch) * shape.vHeads * shape.kDim * + shape.vDim * dataBytes + : hBytes; + const uint64_t hStorageOffset = storeH ? 0 : AllocateWorkspace(cursor, hStorageBytes); + const uint64_t qgScaledOffset = AllocateWorkspace(cursor, kTensorBytes); + + const uint64_t matrixBytes = tokenHeads * chunkSize * sizeof(float); + const uint64_t prepareAqkFp32Offset = AllocateWorkspace(cursor, matrixBytes); + const uint64_t prepareAkkFp32Offset = AllocateWorkspace(cursor, matrixBytes); + const uint64_t prepareScratchOffset = AlignWorkspace(cursor); + const uint64_t solveDepth = safeGate ? KDA_SOLVE_PIPELINE_DEPTH : 1; + const uint64_t solveBytes = static_cast(blockDim) * solveDepth * + KDA_SOLVE_SCRATCH_SLOTS * chunkSize * chunkSize * sizeof(float); + const uint64_t scoreBytes = static_cast(blockDim) * + KDA_SCORE_QUEUE_SLOTS * KDA_SCORE_SCRATCH_PLANES * chunkSize * + shape.kDim * dataBytes; + cursor = prepareScratchOffset + AlignWorkspace(solveBytes) + scoreBytes; + + const uint64_t postWuScratchOffset = AlignWorkspace(cursor); + if (!arch35Options.fusePostWu && !arch35Options.fusePostWuIntoFwdH) { + cursor = postWuScratchOffset + tokenHeads * shape.kDim * sizeof(float); } - kernelScratch = (kernelScratch + KDA_WORKSPACE_ALIGN - 1) / KDA_WORKSPACE_ALIGN * KDA_WORKSPACE_ALIGN; - workspace[0] = ascendcPlatform.GetLibApiWorkSpaceSize() + kernelScratch; - tiling.set_batch(batch); + const uint64_t fwdHWorkspaceBaseOffset = AlignWorkspace(cursor); + uint64_t fwdHCursor = 0; + const uint64_t vWorkspaceOffset = AllocateWorkspace( + fwdHCursor, static_cast(blockDim) * chunkSize * shape.vDim * + sizeof(float) * KDA_GDN_PIPELINE_DEPTH); + const uint64_t vUpdateWorkspaceOffset = AllocateWorkspace( + fwdHCursor, static_cast(blockDim) * chunkSize * shape.vDim * + sizeof(float) * KDA_GDN_PIPELINE_DEPTH); + const uint64_t kDecayWorkspaceOffset = AllocateWorkspace( + fwdHCursor, static_cast(blockDim) * chunkSize * shape.kDim * + sizeof(float) * KDA_GDN_PIPELINE_DEPTH); + const uint64_t hWorkspaceOffset = AllocateWorkspace( + fwdHCursor, static_cast(blockDim) * shape.kDim * shape.vDim * + sizeof(float) * KDA_GDN_PIPELINE_DEPTH); + const uint64_t tokenBatch = isVarLen ? static_cast(seqNum) : 1; + const uint64_t numSeqWorkspaceOffset = AllocateWorkspace( + fwdHCursor, (tokenBatch + 1) * sizeof(int64_t)); + const uint64_t numChunksWorkspaceOffset = AllocateWorkspace( + fwdHCursor, (tokenBatch + 1) * sizeof(int64_t)); + cursor = fwdHWorkspaceBaseOffset + AlignWorkspace(fwdHCursor); + + const uint64_t outputScratchOffset = AllocateWorkspace( + cursor, 2 * tokenHeads * shape.vDim * sizeof(float)); + const uint64_t totalWorkspace = AlignWorkspace(cursor); + + context->SetBlockDim(blockDim); + context->SetTilingKey(useChunk64K128V128Template ? 2 : 1); + context->SetScheduleMode(KDA_BATCH_MODE); + context->GetWorkspaceSizes(1)[0] = platform.GetLibApiWorkSpaceSize() + totalWorkspace; + + ChunkKdaFwdTilingData tiling; + tiling.set_batch(shape.batch); tiling.set_seqNum(seqNum); - tiling.set_qHeadNum(qShape.GetDim(DIM_H)); - tiling.set_vHeadNum(vShape.GetDim(DIM_H)); - tiling.set_seqlen(qShape.GetDim(DIM_T)); - tiling.set_kHeadDim(qShape.GetDim(DIM_D)); - tiling.set_vHeadDim(vShape.GetDim(DIM_D)); + tiling.set_qHeadNum(shape.qHeads); + tiling.set_vHeadNum(shape.vHeads); + tiling.set_seqlen(shape.seqlen); + tiling.set_kHeadDim(shape.kDim); + tiling.set_vHeadDim(shape.vDim); tiling.set_chunkSize(chunkSize); tiling.set_totalChunks(totalChunks); + tiling.set_inputRank(shape.rank); tiling.set_scale(scale); + tiling.set_lowerBound(lowerBound); tiling.set_hasInitialState(hasInitialState); - tiling.set_outputFinalState(outputFinalState); tiling.set_isVarLen(isVarLen); - tiling.set_dataType(DTypeCode(qDesc->GetDataType())); - tiling.set_gateDataType(DTypeCode(gDesc->GetDataType())); - tiling.set_usedCoreNum(blockDim == 0 ? 1 : blockDim); + tiling.set_safeGate(safeGate); + tiling.set_inputSequenceMajor(shape.sequenceMajor); + tiling.set_useGateInKernel(useGateInKernel); + tiling.set_hasALog(hasALog); + tiling.set_hasDtBias(hasDtBias); + tiling.set_computeGateInPrepare(arch35Options.computeGateInPrepare); + tiling.set_fusePostWu(arch35Options.fusePostWu); + tiling.set_fusePostWuIntoFwdH(arch35Options.fusePostWuIntoFwdH); + tiling.set_useDenseFwdH(arch35Options.useDenseFwdH); + tiling.set_storeFinalState(storeFinalState); + tiling.set_storeGk(storeGk); + tiling.set_storeW(storeW); + tiling.set_storeU(storeU); + tiling.set_storeQG(storeQG); + tiling.set_storeKg(storeKg); + tiling.set_storeVNew(storeVNew); + tiling.set_storeH(storeH); tiling.set_stage(stage); - tiling.set_seqStart(seqStart.data()); - tiling.set_seqEnd(seqEnd.data()); - tiling.set_seqChunkOffset(seqChunkOffset.data()); - - if (qDesc->GetDataType() == ge::DT_FLOAT) { - context->SetTilingKey(0); - } else if (qShape.GetDim(DIM_D) < 16) { - context->SetTilingKey(2); - } else { - context->SetTilingKey(1); - } - tiling.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity()); + tiling.set_gateDataType(gDesc->GetDataType() == ge::DT_FLOAT ? 2 : + (gDesc->GetDataType() == ge::DT_BF16 ? 1 : 0)); + tiling.set_gateUsedCoreNum(static_cast(blockDim) * 2); + tiling.set_prepareUsedCoreNum(blockDim); + tiling.set_postWuUsedCoreNum(blockDim); + tiling.set_outputUsedCoreNum(blockDim); + tiling.set_gkStorageOffset(gkStorageOffset); + tiling.set_finalStateStorageOffset(finalStateStorageOffset); + tiling.set_wStorageOffset(wStorageOffset); + tiling.set_uStorageOffset(uStorageOffset); + tiling.set_qgStorageOffset(qgStorageOffset); + tiling.set_kgStorageOffset(kgStorageOffset); + tiling.set_vNewStorageOffset(vNewStorageOffset); + tiling.set_hStorageOffset(hStorageOffset); + tiling.set_qgScaledOffset(qgScaledOffset); + tiling.set_prepareAqkFp32Offset(prepareAqkFp32Offset); + tiling.set_prepareAkkFp32Offset(prepareAkkFp32Offset); + tiling.set_prepareScratchOffset(prepareScratchOffset); + tiling.set_postWuScratchOffset(postWuScratchOffset); + tiling.set_outputScratchOffset(outputScratchOffset); + tiling.set_fwdHWorkspaceBaseOffset(fwdHWorkspaceBaseOffset); + tiling.set_vWorkspaceOffset(vWorkspaceOffset); + tiling.set_vUpdateWorkspaceOffset(vUpdateWorkspaceOffset); + tiling.set_kDecayWorkspaceOffset(kDecayWorkspaceOffset); + tiling.set_hWorkspaceOffset(hWorkspaceOffset); + tiling.set_numSeqWorkspaceOffset(numSeqWorkspaceOffset); + tiling.set_numChunksWorkspaceOffset(numChunksWorkspaceOffset); + tiling.SaveToBuffer(context->GetRawTilingData()->GetData(), + context->GetRawTilingData()->GetCapacity()); context->GetRawTilingData()->SetDataSize(tiling.GetDataSize()); return ge::GRAPH_SUCCESS; } @@ -189,5 +376,4 @@ ge::graphStatus TilingPrepare4ChunkKdaFwd(gert::TilingParseContext *context) IMPL_OP_OPTILING(ChunkKdaFwd) .Tiling(Tiling4ChunkKdaFwd) .TilingParse(TilingPrepare4ChunkKdaFwd); - } // namespace optiling diff --git a/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.h b/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.h index b22d2de71464..fd435d09acb7 100644 --- a/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.h +++ b/csrc/attention/chunk_kda_fwd/op_host/chunk_kda_fwd_tiling.h @@ -1,23 +1,10 @@ -/** - * Copyright (c) 2026 Tianjin University, Ltd. - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of - * the BSD 3-Clause License (the "License"). - * Please refer to the License for details. You may not use this file except in compliance with the License. - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. - */ - #pragma once #include #include -#include namespace optiling { -constexpr int64_t KDA_MAX_TILING_SEQUENCES = 1024; -constexpr int64_t KDA_MAX_TILING_SEQUENCE_OFFSETS = KDA_MAX_TILING_SEQUENCES + 1; - BEGIN_TILING_DATA_DEF(ChunkKdaFwdTilingData) TILING_DATA_FIELD_DEF(int64_t, batch); TILING_DATA_FIELD_DEF(int64_t, seqNum); @@ -28,17 +15,58 @@ TILING_DATA_FIELD_DEF(int64_t, kHeadDim); TILING_DATA_FIELD_DEF(int64_t, vHeadDim); TILING_DATA_FIELD_DEF(int64_t, chunkSize); TILING_DATA_FIELD_DEF(int64_t, totalChunks); +TILING_DATA_FIELD_DEF(int64_t, inputRank); TILING_DATA_FIELD_DEF(float, scale); +TILING_DATA_FIELD_DEF(float, lowerBound); TILING_DATA_FIELD_DEF(bool, hasInitialState); -TILING_DATA_FIELD_DEF(bool, outputFinalState); TILING_DATA_FIELD_DEF(bool, isVarLen); -TILING_DATA_FIELD_DEF(int64_t, dataType); -TILING_DATA_FIELD_DEF(int64_t, gateDataType); -TILING_DATA_FIELD_DEF(int64_t, usedCoreNum); +TILING_DATA_FIELD_DEF(bool, safeGate); +TILING_DATA_FIELD_DEF(bool, inputSequenceMajor); +TILING_DATA_FIELD_DEF(bool, useGateInKernel); +TILING_DATA_FIELD_DEF(bool, hasALog); +TILING_DATA_FIELD_DEF(bool, hasDtBias); +TILING_DATA_FIELD_DEF(bool, computeGateInPrepare); +TILING_DATA_FIELD_DEF(bool, fusePostWu); +TILING_DATA_FIELD_DEF(bool, fusePostWuIntoFwdH); +TILING_DATA_FIELD_DEF(bool, useDenseFwdH); +TILING_DATA_FIELD_DEF(bool, storeFinalState); +TILING_DATA_FIELD_DEF(bool, storeGk); +TILING_DATA_FIELD_DEF(bool, storeW); +TILING_DATA_FIELD_DEF(bool, storeU); +TILING_DATA_FIELD_DEF(bool, storeQG); +TILING_DATA_FIELD_DEF(bool, storeKg); +TILING_DATA_FIELD_DEF(bool, storeVNew); +TILING_DATA_FIELD_DEF(bool, storeH); TILING_DATA_FIELD_DEF(int64_t, stage); -TILING_DATA_FIELD_DEF_ARR(int64_t, KDA_MAX_TILING_SEQUENCES, seqStart); -TILING_DATA_FIELD_DEF_ARR(int64_t, KDA_MAX_TILING_SEQUENCES, seqEnd); -TILING_DATA_FIELD_DEF_ARR(int64_t, KDA_MAX_TILING_SEQUENCE_OFFSETS, seqChunkOffset); + +TILING_DATA_FIELD_DEF(int64_t, gateDataType); +TILING_DATA_FIELD_DEF(int64_t, gateUsedCoreNum); +TILING_DATA_FIELD_DEF(int64_t, prepareUsedCoreNum); +TILING_DATA_FIELD_DEF(int64_t, postWuUsedCoreNum); +TILING_DATA_FIELD_DEF(int64_t, outputUsedCoreNum); + +TILING_DATA_FIELD_DEF(int64_t, gkStorageOffset); +TILING_DATA_FIELD_DEF(int64_t, finalStateStorageOffset); +TILING_DATA_FIELD_DEF(int64_t, wStorageOffset); +TILING_DATA_FIELD_DEF(int64_t, uStorageOffset); +TILING_DATA_FIELD_DEF(int64_t, qgStorageOffset); +TILING_DATA_FIELD_DEF(int64_t, kgStorageOffset); +TILING_DATA_FIELD_DEF(int64_t, vNewStorageOffset); +TILING_DATA_FIELD_DEF(int64_t, hStorageOffset); +TILING_DATA_FIELD_DEF(int64_t, qgScaledOffset); +TILING_DATA_FIELD_DEF(int64_t, prepareAqkFp32Offset); +TILING_DATA_FIELD_DEF(int64_t, prepareAkkFp32Offset); +TILING_DATA_FIELD_DEF(int64_t, prepareScratchOffset); +TILING_DATA_FIELD_DEF(int64_t, postWuScratchOffset); +TILING_DATA_FIELD_DEF(int64_t, outputScratchOffset); + +TILING_DATA_FIELD_DEF(int64_t, fwdHWorkspaceBaseOffset); +TILING_DATA_FIELD_DEF(int64_t, vWorkspaceOffset); +TILING_DATA_FIELD_DEF(int64_t, vUpdateWorkspaceOffset); +TILING_DATA_FIELD_DEF(int64_t, kDecayWorkspaceOffset); +TILING_DATA_FIELD_DEF(int64_t, hWorkspaceOffset); +TILING_DATA_FIELD_DEF(int64_t, numSeqWorkspaceOffset); +TILING_DATA_FIELD_DEF(int64_t, numChunksWorkspaceOffset); END_TILING_DATA_DEF; REGISTER_TILING_DATA_CLASS(ChunkKdaFwd, ChunkKdaFwdTilingData) diff --git a/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.cpp b/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.cpp index 7b387cb0a4eb..c6896f1e63a3 100644 --- a/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.cpp +++ b/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.cpp @@ -2,17 +2,15 @@ * Copyright (c) 2026 Tianjin University, Ltd. * This program is free software, you can redistribute it and/or modify it under the terms and conditions of * the BSD 3-Clause License (the "License"). - * Please refer to the License for details. You may not use this file except in compliance with the License. - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. */ #include "aclnn_chunk_kda_fwd.h" #include "chunk_kda_fwd.h" #include "../../../kda_layout_swap12/op_host/op_api/kda_layout_swap12.h" -#include "moe/chunk_gated_delta_rule_fwd_h/op_host/op_api/chunk_gated_delta_rule_fwd_h.h" +#include #include +#include #include "acl/acl.h" #include "aclnn/aclnn_base.h" @@ -20,6 +18,7 @@ #include "aclnn_kernels/common/op_error_check.h" #include "aclnn_kernels/contiguous.h" #include "aclnn_kernels/reshape.h" +#include "aclnn_kernels/transpose.h" #include "opdev/make_op_executor.h" #include "opdev/op_dfx.h" #include "opdev/op_executor.h" @@ -28,11 +27,6 @@ using namespace op; -namespace l0op { -const aclTensor *Muls(const aclTensor *self, float alpha, aclOpExecutor *executor); -const aclTensor *ZerosLike(const aclTensor *self, aclOpExecutor *executor); -} - #ifdef __cplusplus extern "C" { #endif @@ -40,24 +34,40 @@ extern "C" { namespace { constexpr int64_t MAX_KDA_K_DIM = 256; constexpr int64_t MAX_KDA_HEAD_NUM = 128; +constexpr int64_t KDA_STAGE_FULL = -1; +constexpr int64_t KDA_STAGE_GATE_PREPARE = 0; +constexpr int64_t KDA_STAGE_COUNT = 4; + constexpr int64_t MAX_KDA_VARLEN_SEQUENCES = 1024; +enum class KdaFwdLayout { + BSND, + BNSD, + TND, + NTD, +}; + struct ChunkKdaFwdParams { const aclTensor *q = nullptr; const aclTensor *k = nullptr; const aclTensor *v = nullptr; - const aclTensor *gk = nullptr; + const aclTensor *g = nullptr; const aclTensor *beta = nullptr; + const aclTensor *aLogOptional = nullptr; + const aclTensor *dtBiasOptional = nullptr; const aclTensor *initialStateOptional = nullptr; const aclIntArray *cuSeqlensOptional = nullptr; const aclIntArray *chunkIndicesOptional = nullptr; const char *layout = "BSND"; double scale = 1.0; int64_t chunkSize = 64; - bool outputFinalState = false; - int64_t totalChunks = 1; - const aclTensor *oOut = nullptr; + bool safeGate = false; + double lowerBound = -5.0; + bool useGateInKernel = false; + bool stateVFirst = false; + const aclTensor *attnOut = nullptr; const aclTensor *finalStateOut = nullptr; + const aclTensor *gkOut = nullptr; const aclTensor *aqkOut = nullptr; const aclTensor *akkOut = nullptr; const aclTensor *wOut = nullptr; @@ -68,17 +78,19 @@ struct ChunkKdaFwdParams { const aclTensor *hOut = nullptr; }; -aclnnStatus KdaFwdDataContiguous(const aclTensor *&tensor, aclOpExecutor *executor) -{ - if (tensor == nullptr) { - return ACLNN_SUCCESS; - } - tensor = l0op::Contiguous(tensor, executor); - CHECK_RET(tensor != nullptr, ACLNN_ERR_INNER_NULLPTR); - return ACLNN_SUCCESS; -} +struct KdaShapeInfo { + bool isRank3 = false; + int64_t batch = 0; + int64_t seqlen = 0; + int64_t hNum = 0; + int64_t hvNum = 0; + int64_t kDim = 0; + int64_t vDim = 0; + int64_t seqNum = 0; + int64_t totalChunks = 0; +}; -op::Shape KdaFwdMakeShape(std::initializer_list dims) +op::Shape MakeShape(std::initializer_list dims) { op::Shape shape; for (int64_t dim : dims) { @@ -87,24 +99,104 @@ op::Shape KdaFwdMakeShape(std::initializer_list dims) return shape; } -int64_t KdaFwdDim(const aclTensor *tensor, size_t idx) +const aclTensor *Transpose(const aclTensor *input, const std::vector &perm, aclOpExecutor *executor) +{ + const aclIntArray *permArray = executor->AllocIntArray(perm.data(), perm.size()); + if (permArray == nullptr) { + return nullptr; + } + const aclTensor *transposed = l0op::Transpose(input, permArray, executor); + if (transposed == nullptr) { + return nullptr; + } + const aclTensor *materialized = l0op::Contiguous(transposed, executor); + if (materialized == nullptr) { + return nullptr; + } + const aclTensor *reshaped = + l0op::Reshape(materialized, transposed->GetViewShape(), executor); + if (reshaped == nullptr) { + return nullptr; + } + reshaped->SetStorageShape(reshaped->GetViewShape()); + reshaped->SetOriginalShape(reshaped->GetViewShape()); + return reshaped; +} + +const aclTensor *TransposeLastTwo(const aclTensor *input, aclOpExecutor *executor) +{ + const size_t rank = input->GetViewShape().GetDimNum(); + std::vector perm(rank); + for (size_t idx = 0; idx < rank; ++idx) { + perm[idx] = static_cast(idx); + } + std::swap(perm[rank - 2], perm[rank - 1]); + return Transpose(input, perm, executor); +} + +static int64_t Dim(const aclTensor *tensor, size_t idx) { return tensor->GetViewShape().GetDim(idx); } -const aclTensor *KdaFwdMaybeCast(const aclTensor *tensor, DataType dataType, aclOpExecutor *executor) +static size_t Rank(const aclTensor *tensor) { - if (tensor == nullptr || tensor->GetDataType() == dataType) { - return tensor; + return tensor->GetViewShape().GetDimNum(); +} + +static bool SameShape(const aclTensor *lhs, const aclTensor *rhs) +{ + if (lhs == nullptr || rhs == nullptr || Rank(lhs) != Rank(rhs)) { + return false; } - return l0op::Cast(tensor, dataType, executor); + for (size_t idx = 0; idx < Rank(lhs); ++idx) { + if (Dim(lhs, idx) != Dim(rhs, idx)) { + return false; + } + } + return true; } -aclnnStatus KdaFwdViewCopyMaybeCast(const aclTensor *src, const aclTensor *dst, aclOpExecutor *executor) +bool HasShape(const aclTensor *tensor, std::initializer_list expected) { - const aclTensor *castSrc = KdaFwdMaybeCast(src, dst->GetDataType(), executor); - CHECK_RET(castSrc != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::ViewCopy(castSrc, dst, executor) != nullptr, ACLNN_ERR_INNER_NULLPTR); + if (tensor == nullptr || Rank(tensor) != expected.size()) { + return false; + } + size_t idx = 0; + for (int64_t dim : expected) { + if (Dim(tensor, idx++) != dim) { + return false; + } + } + return true; +} + +aclnnStatus MakeContiguous(const aclTensor *&tensor, aclOpExecutor *executor) +{ + if (tensor == nullptr) { + return ACLNN_SUCCESS; + } + tensor = l0op::Contiguous(tensor, executor); + CHECK_RET(tensor != nullptr, ACLNN_ERR_INNER_NULLPTR); + return ACLNN_SUCCESS; +} + +static aclnnStatus ParseLayout(const char *layout, KdaFwdLayout &parsed) +{ + CHECK_COND(layout != nullptr, ACLNN_ERR_PARAM_INVALID, + "layout must be uppercase and one of BSND, BNSD, TND or NTD."); + if (std::strcmp(layout, "BSND") == 0) { + parsed = KdaFwdLayout::BSND; + } else if (std::strcmp(layout, "BNSD") == 0) { + parsed = KdaFwdLayout::BNSD; + } else if (std::strcmp(layout, "TND") == 0) { + parsed = KdaFwdLayout::TND; + } else if (std::strcmp(layout, "NTD") == 0) { + parsed = KdaFwdLayout::NTD; + } else { + CHECK_COND(false, ACLNN_ERR_PARAM_INVALID, + "layout must be uppercase and one of BSND, BNSD, TND or NTD."); + } return ACLNN_SUCCESS; } @@ -118,6 +210,15 @@ int64_t KdaFwdNumel(const aclTensor *tensor) return numel; } +const aclTensor *KdaFwdMaybeCast(const aclTensor *tensor, DataType dataType, + aclOpExecutor *executor) +{ + if (tensor == nullptr || tensor->GetDataType() == dataType) { + return tensor; + } + return l0op::Cast(tensor, dataType, executor); +} + aclnnStatus KdaFwdCopyMaybeCastAfter(const aclTensor *src, const aclTensor *dependency, const aclTensor *dst, aclOpExecutor *executor) { @@ -128,275 +229,337 @@ aclnnStatus KdaFwdCopyMaybeCastAfter(const aclTensor *src, const aclTensor *depe // dim-1/dim-2 swap an identity by flattening both swapped dimensions to 1. // This deliberately calls the l0op directly; the public aclnn swap shape // contract applies to layout conversion, not to this internal copy barrier. - const aclTensor *linearSrc = l0op::Reshape(castSrc, KdaFwdMakeShape({1, 1, 1, KdaFwdNumel(castSrc)}), executor); + const aclTensor *linearSrc = + l0op::Reshape(castSrc, MakeShape({1, 1, 1, KdaFwdNumel(castSrc)}), executor); CHECK_RET(linearSrc != nullptr, ACLNN_ERR_INNER_NULLPTR); CHECK_RET(l0op::KdaLayoutSwap12(linearSrc, dependency, dst, executor)[0] != nullptr, ACLNN_ERR_INNER_NULLPTR); return ACLNN_SUCCESS; } -size_t KdaFwdRank(const aclTensor *tensor) +aclnnStatus CheckCuSeqlens(const aclIntArray *cuSeqlens, int64_t seqlen) { - return tensor->GetViewShape().GetDimNum(); -} - -aclnnStatus KdaFwdCheckCuSeqlens(const aclIntArray *cuSeqlensOptional, int64_t seqlen) -{ - if (cuSeqlensOptional == nullptr) { + if (cuSeqlens == nullptr) { return ACLNN_SUCCESS; } - const aclIntArray &cu = *cuSeqlensOptional; - CHECK_COND(cu.Size() >= 2, ACLNN_ERR_PARAM_INVALID, + CHECK_COND(cuSeqlens->Size() >= 2, ACLNN_ERR_PARAM_INVALID, "cuSeqlensOptional must contain at least [0, total_tokens]."); - CHECK_COND(cu[0] == 0, ACLNN_ERR_PARAM_INVALID, "cuSeqlensOptional[0] must be 0."); - CHECK_COND(cu[cu.Size() - 1] == seqlen, ACLNN_ERR_PARAM_INVALID, + CHECK_COND((*cuSeqlens)[0] == 0, ACLNN_ERR_PARAM_INVALID, + "cuSeqlensOptional[0] must be 0."); + CHECK_COND((*cuSeqlens)[cuSeqlens->Size() - 1] == seqlen, ACLNN_ERR_PARAM_INVALID, "cuSeqlensOptional last element must equal the sequence length."); - for (size_t idx = 0; idx + 1 < cu.Size(); ++idx) { - CHECK_COND(cu[idx] <= cu[idx + 1], ACLNN_ERR_PARAM_INVALID, + for (size_t idx = 0; idx + 1 < cuSeqlens->Size(); ++idx) { + CHECK_COND((*cuSeqlens)[idx] <= (*cuSeqlens)[idx + 1], ACLNN_ERR_PARAM_INVALID, "cuSeqlensOptional must be nondecreasing."); } return ACLNN_SUCCESS; } -int64_t KdaFwdExpectedChunks(const aclIntArray *cuSeqlensOptional, int64_t seqlen, int64_t chunkSize) +int64_t CountChunks(const aclIntArray *cuSeqlens, int64_t seqlen, int64_t chunkSize) { - if (cuSeqlensOptional == nullptr) { + if (cuSeqlens == nullptr) { return (seqlen + chunkSize - 1) / chunkSize; } - int64_t total = 0; - const aclIntArray &cu = *cuSeqlensOptional; - for (size_t idx = 0; idx + 1 < cu.Size(); ++idx) { - int64_t length = cu[idx + 1] - cu[idx]; - total += (length + chunkSize - 1) / chunkSize; + int64_t chunks = 0; + for (size_t idx = 0; idx + 1 < cuSeqlens->Size(); ++idx) { + chunks += ((*cuSeqlens)[idx + 1] - (*cuSeqlens)[idx] + chunkSize - 1) / chunkSize; } - return total; + return chunks; } -aclnnStatus KdaFwdCheckChunkIndices(const aclIntArray *chunkIndicesOptional, - const aclIntArray *cuSeqlensOptional, - int64_t totalChunks, - int64_t expectedChunks, - int64_t chunkSize) +aclnnStatus CheckChunkIndices(const aclIntArray *chunkIndices, const aclIntArray *cuSeqlens, + int64_t totalChunks, int64_t chunkSize) { - CHECK_COND(totalChunks == expectedChunks, ACLNN_ERR_PARAM_INVALID, - "totalChunks must equal the number of chunks derived from sequence lengths and chunkSize."); - if (chunkIndicesOptional == nullptr) { + if (chunkIndices == nullptr) { return ACLNN_SUCCESS; } - CHECK_COND(cuSeqlensOptional != nullptr, ACLNN_ERR_PARAM_INVALID, - "chunkIndicesOptional is only valid when cuSeqlensOptional is provided."); - CHECK_COND(chunkIndicesOptional->Size() % 2 == 0, ACLNN_ERR_PARAM_INVALID, - "chunkIndicesOptional must contain (seq_id, chunk_id) pairs."); - CHECK_COND(static_cast(chunkIndicesOptional->Size() / 2) == expectedChunks, + CHECK_COND(cuSeqlens != nullptr, ACLNN_ERR_PARAM_INVALID, + "chunkIndicesOptional requires cuSeqlensOptional."); + CHECK_COND(chunkIndices->Size() == static_cast(totalChunks) * 2, ACLNN_ERR_PARAM_INVALID, - "chunkIndicesOptional must contain exactly totalChunks (seq_id, chunk_id) pairs."); - const aclIntArray &indices = *chunkIndicesOptional; - const aclIntArray &cu = *cuSeqlensOptional; - int64_t seqNum = static_cast(cu.Size()) - 1; - for (size_t idx = 0; idx < indices.Size(); idx += 2) { - int64_t seq = indices[idx]; - int64_t localChunk = indices[idx + 1]; - CHECK_COND(seq >= 0 && seq < seqNum, ACLNN_ERR_PARAM_INVALID, - "chunkIndicesOptional seq_id must be in [0, seq_num)."); - int64_t seqLength = cu[seq + 1] - cu[seq]; - int64_t seqChunks = (seqLength + chunkSize - 1) / chunkSize; - CHECK_COND(localChunk >= 0 && localChunk < seqChunks, ACLNN_ERR_PARAM_INVALID, - "chunkIndicesOptional chunk_id is outside the selected sequence."); - } - size_t expectedIdx = 0; - for (int64_t seq = 0; seq < seqNum; ++seq) { - int64_t seqLength = cu[seq + 1] - cu[seq]; - int64_t seqChunks = (seqLength + chunkSize - 1) / chunkSize; - for (int64_t localChunk = 0; localChunk < seqChunks; ++localChunk) { - CHECK_COND(indices[expectedIdx] == seq && indices[expectedIdx + 1] == localChunk, + "chunkIndicesOptional must contain exactly one (seq_id, chunk_id) pair per chunk."); + size_t offset = 0; + for (size_t seq = 0; seq + 1 < cuSeqlens->Size(); ++seq) { + const int64_t length = (*cuSeqlens)[seq + 1] - (*cuSeqlens)[seq]; + const int64_t chunks = (length + chunkSize - 1) / chunkSize; + for (int64_t chunk = 0; chunk < chunks; ++chunk) { + CHECK_COND((*chunkIndices)[offset] == static_cast(seq) && + (*chunkIndices)[offset + 1] == chunk, ACLNN_ERR_PARAM_INVALID, "chunkIndicesOptional must use canonical sequence-major chunk order."); - expectedIdx += 2; + offset += 2; } } return ACLNN_SUCCESS; } -int64_t KdaFwdSeqNum(int64_t batch, const aclIntArray *cuSeqlensOptional) +aclnnStatus ResolveShapeInfo(const ChunkKdaFwdParams ¶ms, KdaFwdLayout layout, KdaShapeInfo &info) { - if (cuSeqlensOptional == nullptr) { - return batch; + info.isRank3 = layout == KdaFwdLayout::TND || layout == KdaFwdLayout::NTD; + const size_t tensorRank = info.isRank3 ? 3 : 4; + const size_t betaRank = info.isRank3 ? 2 : 3; + CHECK_COND(Rank(params.q) == tensorRank && Rank(params.k) == tensorRank && + Rank(params.v) == tensorRank && Rank(params.g) == tensorRank && + Rank(params.beta) == betaRank, + ACLNN_ERR_PARAM_INVALID, + "q/k/v/g and beta ranks must match layout: rank3/rank2 for TND/NTD, rank4/rank3 for BSND/BNSD."); + CHECK_COND(SameShape(params.q, params.k), ACLNN_ERR_PARAM_INVALID, + "q and k must have identical shape."); + + if (layout == KdaFwdLayout::TND) { + info.batch = 1; + info.seqlen = Dim(params.q, 0); + info.hNum = Dim(params.q, 1); + info.kDim = Dim(params.q, 2); + info.hvNum = Dim(params.v, 1); + info.vDim = Dim(params.v, 2); + CHECK_COND(HasShape(params.v, {info.seqlen, info.hvNum, info.vDim}) && + HasShape(params.g, {info.seqlen, info.hvNum, info.kDim}) && + HasShape(params.beta, {info.seqlen, info.hvNum}), + ACLNN_ERR_PARAM_INVALID, "TND expects v/g/beta as [T,HV,V], [T,HV,K], [T,HV]."); + } else if (layout == KdaFwdLayout::NTD) { + info.batch = 1; + info.hNum = Dim(params.q, 0); + info.seqlen = Dim(params.q, 1); + info.kDim = Dim(params.q, 2); + info.hvNum = Dim(params.v, 0); + info.vDim = Dim(params.v, 2); + CHECK_COND(HasShape(params.v, {info.hvNum, info.seqlen, info.vDim}) && + HasShape(params.g, {info.hvNum, info.seqlen, info.kDim}) && + HasShape(params.beta, {info.hvNum, info.seqlen}), + ACLNN_ERR_PARAM_INVALID, "NTD expects v/g/beta as [HV,T,V], [HV,T,K], [HV,T]."); + } else if (layout == KdaFwdLayout::BSND) { + info.batch = Dim(params.q, 0); + info.seqlen = Dim(params.q, 1); + info.hNum = Dim(params.q, 2); + info.kDim = Dim(params.q, 3); + info.hvNum = Dim(params.v, 2); + info.vDim = Dim(params.v, 3); + CHECK_COND(HasShape(params.v, {info.batch, info.seqlen, info.hvNum, info.vDim}) && + HasShape(params.g, {info.batch, info.seqlen, info.hvNum, info.kDim}) && + HasShape(params.beta, {info.batch, info.seqlen, info.hvNum}), + ACLNN_ERR_PARAM_INVALID, + "BSND expects v/g/beta as [B,T,HV,V], [B,T,HV,K], [B,T,HV]."); + } else { + info.batch = Dim(params.q, 0); + info.hNum = Dim(params.q, 1); + info.seqlen = Dim(params.q, 2); + info.kDim = Dim(params.q, 3); + info.hvNum = Dim(params.v, 1); + info.vDim = Dim(params.v, 3); + CHECK_COND(HasShape(params.v, {info.batch, info.hvNum, info.seqlen, info.vDim}) && + HasShape(params.g, {info.batch, info.hvNum, info.seqlen, info.kDim}) && + HasShape(params.beta, {info.batch, info.hvNum, info.seqlen}), + ACLNN_ERR_PARAM_INVALID, + "BNSD expects v/g/beta as [B,HV,T,V], [B,HV,T,K], [B,HV,T]."); } - return static_cast(cuSeqlensOptional->Size()) - 1; + info.seqNum = params.cuSeqlensOptional == nullptr + ? info.batch + : static_cast(params.cuSeqlensOptional->Size()) - 1; + info.totalChunks = CountChunks(params.cuSeqlensOptional, info.seqlen, params.chunkSize); + return ACLNN_SUCCESS; } -aclnnStatus KdaFwdCheckStateShape(const aclTensor *state, const char *name, int64_t seqNum, int64_t hvNum, - int64_t kDim, int64_t vDim) +aclnnStatus CheckDtypes(const ChunkKdaFwdParams ¶ms) { - if (state == nullptr) { - return ACLNN_SUCCESS; + const DataType dataType = params.q->GetDataType(); + CHECK_COND((dataType == DataType::DT_FLOAT16 || dataType == DataType::DT_BF16) && + params.k->GetDataType() == dataType && params.v->GetDataType() == dataType, + ACLNN_ERR_PARAM_INVALID, "q, k and v must use the same float16 or bfloat16 dtype."); + const DataType gateType = params.g->GetDataType(); + CHECK_COND(gateType == DataType::DT_FLOAT || gateType == DataType::DT_BF16, + ACLNN_ERR_PARAM_INVALID, "g must be float32 or bfloat16."); + const DataType betaType = params.beta->GetDataType(); + CHECK_COND(betaType == DataType::DT_FLOAT || betaType == DataType::DT_BF16, + ACLNN_ERR_PARAM_INVALID, "beta must be float32 or bfloat16."); + if (params.aLogOptional != nullptr) { + CHECK_COND(params.aLogOptional->GetDataType() == DataType::DT_FLOAT, ACLNN_ERR_PARAM_INVALID, + "aLogOptional must be float32."); + } + if (params.dtBiasOptional != nullptr) { + CHECK_COND(params.dtBiasOptional->GetDataType() == DataType::DT_FLOAT, ACLNN_ERR_PARAM_INVALID, + "dtBiasOptional must be float32."); + } + if (params.initialStateOptional != nullptr) { + CHECK_COND(params.initialStateOptional->GetDataType() == DataType::DT_FLOAT, + ACLNN_ERR_PARAM_INVALID, "initialStateOptional must be float32."); } - const auto shape = state->GetViewShape(); - CHECK_COND(shape.GetDimNum() == 4 && shape.GetDim(0) == seqNum && shape.GetDim(1) == hvNum && - shape.GetDim(2) == kDim && shape.GetDim(3) == vDim, - ACLNN_ERR_PARAM_INVALID, - "%s must be [seq_num, HV, K, V], where seq_num is batch for dense input or " - "len(cuSeqlensOptional)-1 for varlen input.", - name); return ACLNN_SUCCESS; } -enum class KdaFwdLayout { - BSND, - BNSD, - TND, - NTD, -}; - -bool KdaFwdSameShape(const aclTensor *lhs, const aclTensor *rhs) +aclnnStatus CheckStateShape(const aclTensor *state, const char *name, const KdaShapeInfo &info, bool stateVFirst) { - if (KdaFwdRank(lhs) != KdaFwdRank(rhs)) { - return false; - } - for (size_t idx = 0; idx < KdaFwdRank(lhs); ++idx) { - if (KdaFwdDim(lhs, idx) != KdaFwdDim(rhs, idx)) { - return false; - } + if (state == nullptr) { + return ACLNN_SUCCESS; } - return true; + const bool valid = stateVFirst + ? HasShape(state, {info.seqNum, info.hvNum, info.vDim, info.kDim}) + : HasShape(state, {info.seqNum, info.hvNum, info.kDim, info.vDim}); + CHECK_COND(valid, ACLNN_ERR_PARAM_INVALID, + "%s must be [N,HV,K,V] when stateVFirst=false and [N,HV,V,K] otherwise.", name); + return ACLNN_SUCCESS; } -aclnnStatus KdaFwdParseLayout(const char *layout, KdaFwdLayout &parsed) +aclnnStatus CheckOutputShapes(const ChunkKdaFwdParams ¶ms, const KdaShapeInfo &info) { - CHECK_COND(layout != nullptr, ACLNN_ERR_PARAM_INVALID, - "layout must not be nullptr and must be one of BSND, BNSD, TND, NTD."); - if (std::strcmp(layout, "BSND") == 0) { - parsed = KdaFwdLayout::BSND; - return ACLNN_SUCCESS; + const DataType dataType = params.q->GetDataType(); + const bool attnShapeValid = info.isRank3 + ? HasShape(params.attnOut, {info.seqlen, info.hvNum, info.vDim}) + : HasShape(params.attnOut, + {info.batch, info.seqlen, info.hvNum, info.vDim}); + CHECK_COND(attnShapeValid && params.attnOut->GetDataType() == dataType, + ACLNN_ERR_PARAM_INVALID, + "attnOut must match q dtype and use fixed sequence-major TND/BSND layout."); + if (params.gkOut != nullptr) { + const bool valid = info.isRank3 + ? HasShape(params.gkOut, {info.hvNum, info.seqlen, info.kDim}) + : HasShape(params.gkOut, + {info.batch, info.hvNum, info.seqlen, info.kDim}); + CHECK_COND(valid && params.gkOut->GetDataType() == DataType::DT_FLOAT, + ACLNN_ERR_PARAM_INVALID, "gkOut must be float32 in fixed head-major NTD/BNSD layout."); } - if (std::strcmp(layout, "BNSD") == 0) { - parsed = KdaFwdLayout::BNSD; - return ACLNN_SUCCESS; + const aclTensor *matrixOutputs[] = {params.aqkOut, params.akkOut}; + for (const aclTensor *output : matrixOutputs) { + const bool valid = info.isRank3 + ? HasShape(output, {info.hvNum, info.seqlen, params.chunkSize}) + : HasShape(output, + {info.batch, info.hvNum, info.seqlen, params.chunkSize}); + CHECK_COND(valid && output->GetDataType() == dataType, ACLNN_ERR_PARAM_INVALID, + "Aqk/Akk must match q dtype and use fixed head-major NTD/BNSD layout."); } - if (std::strcmp(layout, "TND") == 0) { - parsed = KdaFwdLayout::TND; - return ACLNN_SUCCESS; + const aclTensor *kOutputs[] = {params.wOut, params.qgOut, params.kgOut}; + for (const aclTensor *output : kOutputs) { + if (output == nullptr) { + continue; + } + const bool valid = info.isRank3 + ? HasShape(output, {info.hvNum, info.seqlen, info.kDim}) + : HasShape(output, + {info.batch, info.hvNum, info.seqlen, info.kDim}); + CHECK_COND(valid && output->GetDataType() == dataType, ACLNN_ERR_PARAM_INVALID, + "w/qg/kg must match q dtype and use fixed head-major NTD/BNSD layout."); } - if (std::strcmp(layout, "NTD") == 0) { - parsed = KdaFwdLayout::NTD; - return ACLNN_SUCCESS; + const aclTensor *vOutputs[] = {params.uOut, params.vNewOut}; + for (const aclTensor *output : vOutputs) { + if (output == nullptr) { + continue; + } + const bool valid = info.isRank3 + ? HasShape(output, {info.hvNum, info.seqlen, info.vDim}) + : HasShape(output, + {info.batch, info.hvNum, info.seqlen, info.vDim}); + CHECK_COND(valid && output->GetDataType() == dataType, ACLNN_ERR_PARAM_INVALID, + "u/vNew must match q dtype and use fixed head-major NTD/BNSD layout."); + } + if (params.hOut != nullptr) { + const bool valid = params.stateVFirst + ? (info.isRank3 + ? HasShape(params.hOut, + {info.totalChunks, info.hvNum, info.vDim, info.kDim}) + : HasShape(params.hOut, + {info.batch, info.totalChunks, info.hvNum, + info.vDim, info.kDim})) + : (info.isRank3 + ? HasShape(params.hOut, + {info.totalChunks, info.hvNum, info.kDim, info.vDim}) + : HasShape(params.hOut, + {info.batch, info.totalChunks, info.hvNum, + info.kDim, info.vDim})); + CHECK_COND(valid && params.hOut->GetDataType() == dataType, ACLNN_ERR_PARAM_INVALID, + "hOut must match q dtype, use fixed sequence-major layout, and follow stateVFirst."); } - CHECK_COND(false, ACLNN_ERR_PARAM_INVALID, - "layout must be one of BSND, BNSD, TND, NTD and must be uppercase."); - return ACLNN_ERR_PARAM_INVALID; + CHECK_RET(CheckStateShape(params.finalStateOut, "finalStateOut", info, params.stateVFirst) == + ACLNN_SUCCESS, + ACLNN_ERR_PARAM_INVALID); + if (params.finalStateOut != nullptr) { + CHECK_COND(params.finalStateOut->GetDataType() == DataType::DT_FLOAT, + ACLNN_ERR_PARAM_INVALID, "finalStateOut must be float32."); + } + return ACLNN_SUCCESS; } -aclnnStatus KdaFwdCheckLayoutShape(const ChunkKdaFwdParams ¶ms, KdaFwdLayout layout) +aclnnStatus CheckParams(const ChunkKdaFwdParams ¶ms, KdaFwdLayout &layout, KdaShapeInfo &info) { - CHECK_COND(KdaFwdSameShape(params.q, params.k), ACLNN_ERR_PARAM_INVALID, - "q and k must have identical shape."); - if (layout == KdaFwdLayout::TND) { - CHECK_COND(KdaFwdRank(params.q) == 3 && KdaFwdRank(params.v) == 3 && - KdaFwdRank(params.gk) == 3 && KdaFwdRank(params.beta) == 2, - ACLNN_ERR_PARAM_INVALID, - "layout TND expects q/k [T,H,K], v [T,HV,V], gk [T,HV,K], beta [T,HV]."); - CHECK_COND(KdaFwdDim(params.v, 0) == KdaFwdDim(params.q, 0) && - KdaFwdDim(params.gk, 0) == KdaFwdDim(params.q, 0) && - KdaFwdDim(params.beta, 0) == KdaFwdDim(params.q, 0) && - KdaFwdDim(params.gk, 1) == KdaFwdDim(params.v, 1) && - KdaFwdDim(params.beta, 1) == KdaFwdDim(params.v, 1) && - KdaFwdDim(params.gk, 2) == KdaFwdDim(params.q, 2), - ACLNN_ERR_PARAM_INVALID, - "layout TND shape mismatch."); - } else if (layout == KdaFwdLayout::NTD) { - CHECK_COND(KdaFwdRank(params.q) == 3 && KdaFwdRank(params.v) == 3 && - KdaFwdRank(params.gk) == 3 && KdaFwdRank(params.beta) == 2, - ACLNN_ERR_PARAM_INVALID, - "layout NTD expects q/k [H,T,K], v [HV,T,V], gk [HV,T,K], beta [HV,T]."); - CHECK_COND(KdaFwdDim(params.v, 1) == KdaFwdDim(params.q, 1) && - KdaFwdDim(params.gk, 0) == KdaFwdDim(params.v, 0) && - KdaFwdDim(params.beta, 0) == KdaFwdDim(params.v, 0) && - KdaFwdDim(params.gk, 1) == KdaFwdDim(params.q, 1) && - KdaFwdDim(params.beta, 1) == KdaFwdDim(params.q, 1) && - KdaFwdDim(params.gk, 2) == KdaFwdDim(params.q, 2), - ACLNN_ERR_PARAM_INVALID, - "layout NTD shape mismatch."); - } else if (layout == KdaFwdLayout::BSND) { - CHECK_COND(KdaFwdRank(params.q) == 4 && KdaFwdRank(params.v) == 4 && - KdaFwdRank(params.gk) == 4 && KdaFwdRank(params.beta) == 3, - ACLNN_ERR_PARAM_INVALID, - "layout BSND expects q/k [B,T,H,K], v [B,T,HV,V], gk [B,T,HV,K], beta [B,T,HV]."); - CHECK_COND(KdaFwdDim(params.v, 0) == KdaFwdDim(params.q, 0) && - KdaFwdDim(params.v, 1) == KdaFwdDim(params.q, 1) && - KdaFwdDim(params.gk, 0) == KdaFwdDim(params.q, 0) && - KdaFwdDim(params.gk, 1) == KdaFwdDim(params.q, 1) && - KdaFwdDim(params.beta, 0) == KdaFwdDim(params.q, 0) && - KdaFwdDim(params.beta, 1) == KdaFwdDim(params.q, 1) && - KdaFwdDim(params.gk, 2) == KdaFwdDim(params.v, 2) && - KdaFwdDim(params.beta, 2) == KdaFwdDim(params.v, 2) && - KdaFwdDim(params.gk, 3) == KdaFwdDim(params.q, 3), - ACLNN_ERR_PARAM_INVALID, - "layout BSND shape mismatch."); - } else { - CHECK_COND(KdaFwdRank(params.q) == 4 && KdaFwdRank(params.v) == 4 && - KdaFwdRank(params.gk) == 4 && KdaFwdRank(params.beta) == 3, - ACLNN_ERR_PARAM_INVALID, - "layout BNSD expects q/k [B,H,T,K], v [B,HV,T,V], gk [B,HV,T,K], beta [B,HV,T]."); - CHECK_COND(KdaFwdDim(params.v, 0) == KdaFwdDim(params.q, 0) && - KdaFwdDim(params.v, 2) == KdaFwdDim(params.q, 2) && - KdaFwdDim(params.gk, 0) == KdaFwdDim(params.q, 0) && - KdaFwdDim(params.gk, 1) == KdaFwdDim(params.v, 1) && - KdaFwdDim(params.beta, 0) == KdaFwdDim(params.q, 0) && - KdaFwdDim(params.beta, 1) == KdaFwdDim(params.v, 1) && - KdaFwdDim(params.gk, 2) == KdaFwdDim(params.q, 2) && - KdaFwdDim(params.beta, 2) == KdaFwdDim(params.q, 2) && - KdaFwdDim(params.gk, 3) == KdaFwdDim(params.q, 3), - ACLNN_ERR_PARAM_INVALID, - "layout BNSD shape mismatch."); + CHECK_COND(params.q != nullptr && params.k != nullptr && params.v != nullptr && + params.g != nullptr && params.beta != nullptr, + ACLNN_ERR_PARAM_NULLPTR, "q, k, v, g and beta must not be nullptr."); + CHECK_COND(params.attnOut != nullptr, ACLNN_ERR_PARAM_NULLPTR, "attnOut must not be nullptr."); + CHECK_COND(params.aqkOut != nullptr && params.akkOut != nullptr, + ACLNN_ERR_PARAM_NULLPTR, "aqkOut and akkOut must not be nullptr."); + CHECK_COND(params.chunkSize == 64 || params.chunkSize == 128, ACLNN_ERR_PARAM_INVALID, + "chunkSize must be 64 or 128."); + CHECK_RET(ParseLayout(params.layout, layout) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(ResolveShapeInfo(params, layout, info) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_COND(info.hNum > 0 && info.hvNum >= info.hNum && info.hvNum % info.hNum == 0, + ACLNN_ERR_PARAM_INVALID, + "H and HV must be positive, HV must be greater than or equal to H, and HV must be divisible by H."); + CHECK_COND(info.hNum <= MAX_KDA_HEAD_NUM && info.hvNum <= MAX_KDA_HEAD_NUM, + ACLNN_ERR_PARAM_INVALID, "H and HV must be less than or equal to 128."); + CHECK_COND(info.kDim >= 16 && info.kDim <= MAX_KDA_K_DIM && info.kDim % 16 == 0 && + info.vDim >= 16 && info.vDim <= 256 && info.vDim % 16 == 0, + ACLNN_ERR_PARAM_INVALID, + "K/V must be multiples of 16, K must be <=256, and V must be <=256."); + CHECK_RET(CheckDtypes(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(CheckCuSeqlens(params.cuSeqlensOptional, info.seqlen) == ACLNN_SUCCESS, + ACLNN_ERR_PARAM_INVALID); + CHECK_COND(params.cuSeqlensOptional == nullptr || info.isRank3 || info.batch == 1, + ACLNN_ERR_PARAM_INVALID, + "rank4 varlen input with cuSeqlensOptional requires B=1."); + CHECK_COND(params.cuSeqlensOptional == nullptr || info.seqNum <= MAX_KDA_VARLEN_SEQUENCES, + ACLNN_ERR_PARAM_INVALID, "varlen input supports at most 1024 sequences."); + CHECK_RET(CheckChunkIndices(params.chunkIndicesOptional, params.cuSeqlensOptional, + info.totalChunks, params.chunkSize) == ACLNN_SUCCESS, + ACLNN_ERR_PARAM_INVALID); + CHECK_RET(CheckStateShape(params.initialStateOptional, "initialStateOptional", info, + params.stateVFirst) == ACLNN_SUCCESS, + ACLNN_ERR_PARAM_INVALID); + if (params.useGateInKernel) { + CHECK_COND(params.aLogOptional != nullptr, ACLNN_ERR_PARAM_NULLPTR, + "aLogOptional is required when useGateInKernel is true."); + CHECK_COND(HasShape(params.aLogOptional, {info.hvNum}), ACLNN_ERR_PARAM_INVALID, + "aLogOptional must have shape [HV]."); + if (params.dtBiasOptional != nullptr) { + CHECK_COND(HasShape(params.dtBiasOptional, {info.hvNum * info.kDim}), + ACLNN_ERR_PARAM_INVALID, "dtBiasOptional must have shape [HV*K]."); + } + if (params.safeGate) { + CHECK_COND(params.lowerBound >= -5.0 && params.lowerBound < 0.0, + ACLNN_ERR_PARAM_INVALID, + "lowerBound must be in [-5, 0) when safeGate is true."); + } } + CHECK_RET(CheckOutputShapes(params, info) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); return ACLNN_SUCCESS; } -aclnnStatus KdaFwdCheckParams(const ChunkKdaFwdParams ¶ms) +aclnnStatus ContiguousInputs(ChunkKdaFwdParams ¶ms, aclOpExecutor *executor) { - CHECK_COND(params.q != nullptr, ACLNN_ERR_PARAM_NULLPTR, "q must not be nullptr."); - CHECK_COND(params.k != nullptr, ACLNN_ERR_PARAM_NULLPTR, "k must not be nullptr."); - CHECK_COND(params.v != nullptr, ACLNN_ERR_PARAM_NULLPTR, "v must not be nullptr."); - CHECK_COND(params.gk != nullptr, ACLNN_ERR_PARAM_NULLPTR, "gk must not be nullptr."); - CHECK_COND(params.beta != nullptr, ACLNN_ERR_PARAM_NULLPTR, "beta must not be nullptr."); - CHECK_COND(params.oOut != nullptr && params.finalStateOut != nullptr && params.aqkOut != nullptr && - params.akkOut != nullptr && params.wOut != nullptr && params.uOut != nullptr && - params.qgOut != nullptr && params.kgOut != nullptr && params.vNewOut != nullptr && - params.hOut != nullptr, - ACLNN_ERR_PARAM_NULLPTR, "ChunkKdaFwd outputs must not be nullptr."); - CHECK_COND(params.chunkSize > 0, ACLNN_ERR_PARAM_INVALID, "chunkSize must be positive."); - CHECK_COND(params.totalChunks > 0, ACLNN_ERR_PARAM_INVALID, "totalChunks must be positive."); - size_t qRank = KdaFwdRank(params.q); - size_t betaRank = KdaFwdRank(params.beta); - CHECK_COND((qRank == 4 && betaRank == 3) || (qRank == 3 && betaRank == 2), ACLNN_ERR_PARAM_INVALID, - "q/k/v/gk must be BSND/BNSD rank4 with beta rank3, or TND/NTD rank3 with beta rank2."); - size_t kDimIdx = (qRank == 4) ? 3 : 2; - CHECK_COND(params.q->GetViewShape().GetDim(kDimIdx) <= MAX_KDA_K_DIM, ACLNN_ERR_PARAM_INVALID, - "k head dimension must be less than or equal to 256."); + CHECK_RET(MakeContiguous(params.q, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(MakeContiguous(params.k, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(MakeContiguous(params.v, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(MakeContiguous(params.g, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(MakeContiguous(params.beta, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(MakeContiguous(params.aLogOptional, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(MakeContiguous(params.dtBiasOptional, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(MakeContiguous(params.initialStateOptional, executor) == ACLNN_SUCCESS, + ACLNN_ERR_PARAM_INVALID); return ACLNN_SUCCESS; } -bool KdaFwdSplitCubePathSupported(const ChunkKdaFwdParams ¶ms, int64_t kDim, int64_t vDim) +const aclTensor *AllocTensor(aclOpExecutor *executor, const op::Shape &shape, DataType dtype) { - auto qDtype = params.q->GetDataType(); - auto kDtype = params.k->GetDataType(); - auto vDtype = params.v->GetDataType(); - bool dataDtypeSupported = (qDtype == DataType::DT_FLOAT16 || qDtype == DataType::DT_BF16) && - kDtype == qDtype && vDtype == qDtype; - return dataDtypeSupported && - (params.chunkSize == 64 || params.chunkSize == 128) && kDim >= 16 && vDim >= 16 && - kDim % 16 == 0 && vDim % 16 == 0 && vDim <= 256; + return executor->AllocTensor(shape, dtype, Format::FORMAT_ND); } -aclnnStatus KdaFwdParamsDataContiguous(ChunkKdaFwdParams ¶ms, aclOpExecutor *executor) +const aclTensor *AsRank4(const aclTensor *tensor, const op::Shape &shape, aclOpExecutor *executor) { - CHECK_RET(KdaFwdDataContiguous(params.q, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); - CHECK_RET(KdaFwdDataContiguous(params.k, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); - CHECK_RET(KdaFwdDataContiguous(params.v, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); - CHECK_RET(KdaFwdDataContiguous(params.gk, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); - CHECK_RET(KdaFwdDataContiguous(params.beta, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); - CHECK_RET(KdaFwdDataContiguous(params.initialStateOptional, executor) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); - return ACLNN_SUCCESS; + return l0op::Reshape(tensor, shape, executor); +} + +bool IsAscend950() +{ + const char *socName = aclrtGetSocName(); + return socName != nullptr && std::strstr(socName, "Ascend950") != nullptr; } } // namespace @@ -404,18 +567,23 @@ aclnnStatus aclnnChunkKdaFwdGetWorkspaceSize( const aclTensor *q, const aclTensor *k, const aclTensor *v, - const aclTensor *gk, + const aclTensor *g, const aclTensor *beta, + const aclTensor *aLogOptional, + const aclTensor *dtBiasOptional, const aclTensor *initialStateOptional, const aclIntArray *cuSeqlensOptional, const aclIntArray *chunkIndicesOptional, const char *layout, double scale, int64_t chunkSize, - bool outputFinalState, - int64_t totalChunks, - const aclTensor *oOut, + bool safeGate, + double lowerBound, + bool useGateInKernel, + bool stateVFirst, + const aclTensor *attnOut, const aclTensor *finalStateOut, + const aclTensor *gkOut, const aclTensor *aqkOut, const aclTensor *akkOut, const aclTensor *wOut, @@ -427,482 +595,237 @@ aclnnStatus aclnnChunkKdaFwdGetWorkspaceSize( uint64_t *workspaceSize, aclOpExecutor **executor) { - ChunkKdaFwdParams params{q, k, v, gk, beta, initialStateOptional, cuSeqlensOptional, chunkIndicesOptional, layout, - scale, chunkSize, outputFinalState, totalChunks, oOut, finalStateOut, aqkOut, - akkOut, wOut, uOut, qgOut, kgOut, vNewOut, hOut}; - L2_DFX_PHASE_1(aclnnChunkKdaFwd, - DFX_IN(q, k, v, gk, beta, initialStateOptional, cuSeqlensOptional, chunkIndicesOptional), - DFX_OUT(oOut, finalStateOut, aqkOut, akkOut, wOut, uOut, qgOut, kgOut, vNewOut, hOut)); + ChunkKdaFwdParams params{ + q, k, v, g, beta, aLogOptional, dtBiasOptional, initialStateOptional, + cuSeqlensOptional, chunkIndicesOptional, layout, scale, chunkSize, + safeGate, lowerBound, useGateInKernel, stateVFirst, attnOut, finalStateOut, gkOut, + aqkOut, akkOut, wOut, uOut, qgOut, kgOut, vNewOut, hOut}; + L2_DFX_PHASE_1( + aclnnChunkKdaFwd, + DFX_IN(q, k, v, g, beta, aLogOptional, dtBiasOptional, initialStateOptional, + cuSeqlensOptional, chunkIndicesOptional, layout, scale, chunkSize, + safeGate, lowerBound, useGateInKernel, stateVFirst), + DFX_OUT(attnOut, finalStateOut, gkOut, aqkOut, akkOut, wOut, uOut, + qgOut, kgOut, vNewOut, hOut)); + auto uniqueExecutor = CREATE_EXECUTOR(); CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); auto executorPtr = uniqueExecutor.get(); - CHECK_RET(KdaFwdCheckParams(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); - CHECK_RET(KdaFwdParamsDataContiguous(params, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); - KdaFwdLayout parsedLayout = KdaFwdLayout::BSND; - CHECK_RET(KdaFwdParseLayout(params.layout, parsedLayout) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); - CHECK_RET(KdaFwdCheckLayoutShape(params, parsedLayout) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); - bool isTnd = parsedLayout == KdaFwdLayout::TND || parsedLayout == KdaFwdLayout::NTD; - bool isInternalLayout = parsedLayout == KdaFwdLayout::BNSD || parsedLayout == KdaFwdLayout::NTD; - int64_t batch = isTnd ? 1 : KdaFwdDim(params.q, 0); - int64_t seqlen = parsedLayout == KdaFwdLayout::TND ? KdaFwdDim(params.q, 0) : - (parsedLayout == KdaFwdLayout::NTD ? KdaFwdDim(params.q, 1) : - (parsedLayout == KdaFwdLayout::BNSD ? KdaFwdDim(params.q, 2) : KdaFwdDim(params.q, 1))); - int64_t hNum = parsedLayout == KdaFwdLayout::TND ? KdaFwdDim(params.q, 1) : - (parsedLayout == KdaFwdLayout::NTD ? KdaFwdDim(params.q, 0) : - (parsedLayout == KdaFwdLayout::BNSD ? KdaFwdDim(params.q, 1) : KdaFwdDim(params.q, 2))); - int64_t kDim = isTnd ? KdaFwdDim(params.q, 2) : KdaFwdDim(params.q, 3); - int64_t hvNum = parsedLayout == KdaFwdLayout::TND ? KdaFwdDim(params.v, 1) : - (parsedLayout == KdaFwdLayout::NTD ? KdaFwdDim(params.v, 0) : - (parsedLayout == KdaFwdLayout::BNSD ? KdaFwdDim(params.v, 1) : KdaFwdDim(params.v, 2))); - int64_t vDim = isTnd ? KdaFwdDim(params.v, 2) : KdaFwdDim(params.v, 3); - int64_t seqNum = KdaFwdSeqNum(batch, params.cuSeqlensOptional); - CHECK_COND(hNum <= MAX_KDA_HEAD_NUM && hvNum <= MAX_KDA_HEAD_NUM, ACLNN_ERR_PARAM_INVALID, - "H and HV must be less than or equal to 128."); - CHECK_COND(hNum > 0 && hvNum >= hNum && hvNum % hNum == 0, ACLNN_ERR_PARAM_INVALID, - "H and HV must be positive, HV must be greater than or equal to H, and HV must be divisible by H."); - CHECK_COND(parsedLayout != KdaFwdLayout::TND || hNum == 1, ACLNN_ERR_PARAM_INVALID, - "TND layout with H > 1 is not supported by npu_chunk_kda_fwd; use NTD [H,T,D] layout " - "for multi-head rank3 input."); - CHECK_RET(KdaFwdCheckCuSeqlens(params.cuSeqlensOptional, seqlen) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); - int64_t expectedChunks = KdaFwdExpectedChunks(params.cuSeqlensOptional, seqlen, params.chunkSize); - CHECK_COND(params.cuSeqlensOptional == nullptr || seqNum <= MAX_KDA_VARLEN_SEQUENCES, - ACLNN_ERR_PARAM_INVALID, - "varlen input supports at most 1024 sequences in one call; split a larger request at sequence " - "boundaries."); - CHECK_RET(KdaFwdCheckChunkIndices(params.chunkIndicesOptional, params.cuSeqlensOptional, params.totalChunks, - expectedChunks, params.chunkSize) == ACLNN_SUCCESS, - ACLNN_ERR_PARAM_INVALID); - CHECK_COND(params.cuSeqlensOptional == nullptr || isTnd || batch == 1, ACLNN_ERR_PARAM_INVALID, - "rank4 varlen input with cuSeqlensOptional currently requires B=1."); - CHECK_RET(KdaFwdCheckStateShape(params.initialStateOptional, "initialStateOptional", seqNum, hvNum, kDim, vDim) == - ACLNN_SUCCESS, - ACLNN_ERR_PARAM_INVALID); - CHECK_RET(KdaFwdCheckStateShape(params.finalStateOut, "finalStateOut", seqNum, hvNum, kDim, vDim) == - ACLNN_SUCCESS, - ACLNN_ERR_PARAM_INVALID); - CHECK_COND(KdaFwdSplitCubePathSupported(params, kDim, vDim), ACLNN_ERR_PARAM_INVALID, - "npu_chunk_kda_fwd only supports the AscendC split cube/vector path: q/k/v dtype must be the same " - "fp16/bf16 type, chunkSize must be 64 or 128, K/V must be multiples of 16, and V must be <= 256."); - - const aclTensor *qBsnd = params.q; - const aclTensor *kBsnd = params.k; - const aclTensor *vBsnd = params.v; - const aclTensor *gkBsnd = params.gk; - const aclTensor *betaBsn = params.beta; - if (parsedLayout == KdaFwdLayout::TND) { - qBsnd = l0op::Reshape(params.q, KdaFwdMakeShape({1, seqlen, hNum, kDim}), executorPtr); - kBsnd = l0op::Reshape(params.k, KdaFwdMakeShape({1, seqlen, hNum, kDim}), executorPtr); - vBsnd = l0op::Reshape(params.v, KdaFwdMakeShape({1, seqlen, hvNum, vDim}), executorPtr); - gkBsnd = l0op::Reshape(params.gk, KdaFwdMakeShape({1, seqlen, hvNum, kDim}), executorPtr); - betaBsn = l0op::Reshape(params.beta, KdaFwdMakeShape({1, seqlen, hvNum}), executorPtr); - CHECK_RET(qBsnd != nullptr && kBsnd != nullptr && vBsnd != nullptr && gkBsnd != nullptr && betaBsn != nullptr, - ACLNN_ERR_INNER_NULLPTR); - } else if (parsedLayout == KdaFwdLayout::NTD) { - qBsnd = l0op::Reshape(params.q, KdaFwdMakeShape({1, hNum, seqlen, kDim}), executorPtr); - kBsnd = l0op::Reshape(params.k, KdaFwdMakeShape({1, hNum, seqlen, kDim}), executorPtr); - vBsnd = l0op::Reshape(params.v, KdaFwdMakeShape({1, hvNum, seqlen, vDim}), executorPtr); - gkBsnd = l0op::Reshape(params.gk, KdaFwdMakeShape({1, hvNum, seqlen, kDim}), executorPtr); - betaBsn = l0op::Reshape(params.beta, KdaFwdMakeShape({1, hvNum, seqlen}), executorPtr); - CHECK_RET(qBsnd != nullptr && kBsnd != nullptr && vBsnd != nullptr && gkBsnd != nullptr && betaBsn != nullptr, + KdaShapeInfo info; + CHECK_RET(CheckParams(params, parsedLayout, info) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + CHECK_RET(ContiguousInputs(params, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID); + + const aclTensor *qHead = params.q; + const aclTensor *kHead = params.k; + const aclTensor *vHead = params.v; + const aclTensor *gHead = params.g; + const aclTensor *betaHead = params.beta; + if (parsedLayout == KdaFwdLayout::BSND) { + // beta is small and its scalar-per-token rows are not DMA friendly in + // sequence-major form. Keep only this lightweight transpose. + betaHead = Transpose(params.beta, {0, 2, 1}, executorPtr); + } else if (parsedLayout == KdaFwdLayout::TND) { + qHead = Transpose(params.q, {1, 0, 2}, executorPtr); + kHead = Transpose(params.k, {1, 0, 2}, executorPtr); + vHead = Transpose(params.v, {1, 0, 2}, executorPtr); + gHead = Transpose(params.g, {1, 0, 2}, executorPtr); + betaHead = Transpose(params.beta, {1, 0}, executorPtr); + } + CHECK_RET(qHead != nullptr && kHead != nullptr && vHead != nullptr && + gHead != nullptr && betaHead != nullptr, + ACLNN_ERR_INNER_NULLPTR); + + if (info.isRank3) { + qHead = AsRank4(qHead, MakeShape({1, info.hNum, info.seqlen, info.kDim}), executorPtr); + kHead = AsRank4(kHead, MakeShape({1, info.hNum, info.seqlen, info.kDim}), executorPtr); + vHead = AsRank4(vHead, MakeShape({1, info.hvNum, info.seqlen, info.vDim}), executorPtr); + gHead = AsRank4(gHead, MakeShape({1, info.hvNum, info.seqlen, info.kDim}), executorPtr); + betaHead = AsRank4(betaHead, MakeShape({1, info.hvNum, info.seqlen}), executorPtr); + CHECK_RET(qHead != nullptr && kHead != nullptr && vHead != nullptr && + gHead != nullptr && betaHead != nullptr, ACLNN_ERR_INNER_NULLPTR); } - const aclTensor *qBnsd = isInternalLayout ? qBsnd : - executorPtr->AllocTensor(KdaFwdMakeShape({batch, hNum, seqlen, kDim}), - params.q->GetDataType(), Format::FORMAT_ND); - const aclTensor *kBnsd = isInternalLayout ? kBsnd : - executorPtr->AllocTensor(KdaFwdMakeShape({batch, hNum, seqlen, kDim}), - params.k->GetDataType(), Format::FORMAT_ND); - const aclTensor *vBnsd = isInternalLayout ? vBsnd : - executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, vDim}), - params.v->GetDataType(), Format::FORMAT_ND); - const aclTensor *gkBnsdRaw = isInternalLayout ? gkBsnd : - executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), - params.gk->GetDataType(), Format::FORMAT_ND); - const aclTensor *betaBnsRaw = isInternalLayout ? betaBsn : - executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen}), - params.beta->GetDataType(), Format::FORMAT_ND); - bool returnIntermediates = KdaFwdNumel(params.aqkOut) != 0; - const aclTensor *oBnsd = nullptr; - const aclTensor *aqkBnst = nullptr; - const aclTensor *akkBnst = nullptr; - const aclTensor *wBnsd = nullptr; - const aclTensor *uBnsd = nullptr; - const aclTensor *qgBnsd = nullptr; - const aclTensor *kgBnsd = nullptr; - const aclTensor *vNewBnsd = nullptr; - const aclTensor *hBnst = nullptr; - if (isInternalLayout) { - if (parsedLayout == KdaFwdLayout::NTD) { - oBnsd = l0op::Reshape(params.oOut, KdaFwdMakeShape({1, hvNum, seqlen, vDim}), executorPtr); - if (returnIntermediates) { - aqkBnst = l0op::Reshape(params.aqkOut, KdaFwdMakeShape({1, hvNum, seqlen, params.chunkSize}), - executorPtr); - akkBnst = l0op::Reshape(params.akkOut, KdaFwdMakeShape({1, hvNum, seqlen, params.chunkSize}), - executorPtr); - wBnsd = l0op::Reshape(params.wOut, KdaFwdMakeShape({1, hvNum, seqlen, kDim}), executorPtr); - uBnsd = l0op::Reshape(params.uOut, KdaFwdMakeShape({1, hvNum, seqlen, vDim}), executorPtr); - qgBnsd = l0op::Reshape(params.qgOut, KdaFwdMakeShape({1, hvNum, seqlen, kDim}), executorPtr); - kgBnsd = l0op::Reshape(params.kgOut, KdaFwdMakeShape({1, hvNum, seqlen, kDim}), executorPtr); - vNewBnsd = l0op::Reshape(params.vNewOut, KdaFwdMakeShape({1, hvNum, seqlen, vDim}), executorPtr); - hBnst = l0op::Reshape(params.hOut, KdaFwdMakeShape({1, hvNum, params.totalChunks, kDim, vDim}), - executorPtr); - } - } else { - oBnsd = params.oOut; - if (returnIntermediates) { - aqkBnst = params.aqkOut; - akkBnst = params.akkOut; - wBnsd = params.wOut; - uBnsd = params.uOut; - qgBnsd = params.qgOut; - kgBnsd = params.kgOut; - vNewBnsd = params.vNewOut; - hBnst = params.hOut; - } + const op::Shape gkShape4 = MakeShape({info.batch, info.hvNum, info.seqlen, info.kDim}); + const op::Shape matrixShape4 = + MakeShape({info.batch, info.hvNum, info.seqlen, params.chunkSize}); + const op::Shape kShape4 = MakeShape({info.batch, info.hvNum, info.seqlen, info.kDim}); + const op::Shape vShape4 = MakeShape({info.batch, info.hvNum, info.seqlen, info.vDim}); + const op::Shape hShape5 = + MakeShape({info.batch, info.hvNum, info.totalChunks, info.kDim, info.vDim}); + const op::Shape hExportShape5 = + params.stateVFirst + ? MakeShape({info.batch, info.totalChunks, info.hvNum, info.vDim, info.kDim}) + : MakeShape({info.batch, info.totalChunks, info.hvNum, info.kDim, info.vDim}); + const op::Shape stateShape4 = + MakeShape({info.seqNum, info.hvNum, info.kDim, info.vDim}); + const op::Shape placeholderShape = MakeShape({1}); + const bool useDenseA5FastPath = + params.cuSeqlensOptional == nullptr && params.q->GetDataType() == DataType::DT_BF16 && + params.chunkSize == 64 && info.kDim == 128 && info.vDim == 128 && + info.seqlen % params.chunkSize == 0; + const bool splitStages = + IsAscend950() && info.totalChunks > 1 && !useDenseA5FastPath; + + const aclTensor *gkCompute = params.gkOut; + if (gkCompute != nullptr && info.isRank3) { + gkCompute = AsRank4(gkCompute, gkShape4, executorPtr); + } + if (gkCompute == nullptr) { + gkCompute = AllocTensor( + executorPtr, splitStages ? gkShape4 : placeholderShape, + DataType::DT_FLOAT); + } + CHECK_RET(gkCompute != nullptr, ACLNN_ERR_INNER_NULLPTR); + + const aclTensor *aqkCompute = params.aqkOut; + const aclTensor *akkCompute = params.akkOut; + const aclTensor *wExport = params.wOut; + const aclTensor *uExport = params.uOut; + const aclTensor *qgExport = params.qgOut; + const aclTensor *kgExport = params.kgOut; + const aclTensor *vNewExport = params.vNewOut; + const aclTensor *hExport = params.hOut; + if (info.isRank3) { + aqkCompute = aqkCompute == nullptr ? nullptr : AsRank4(aqkCompute, matrixShape4, executorPtr); + akkCompute = akkCompute == nullptr ? nullptr : AsRank4(akkCompute, matrixShape4, executorPtr); + wExport = wExport == nullptr ? nullptr : AsRank4(wExport, kShape4, executorPtr); + uExport = uExport == nullptr ? nullptr : AsRank4(uExport, vShape4, executorPtr); + qgExport = qgExport == nullptr ? nullptr : AsRank4(qgExport, kShape4, executorPtr); + kgExport = kgExport == nullptr ? nullptr : AsRank4(kgExport, kShape4, executorPtr); + vNewExport = vNewExport == nullptr ? nullptr : AsRank4(vNewExport, vShape4, executorPtr); + if (hExport != nullptr) { + hExport = AsRank4(hExport, hExportShape5, executorPtr); } - } else { - oBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, vDim}), - params.oOut->GetDataType(), Format::FORMAT_ND); - } - if (!isInternalLayout) { - wBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), - params.wOut->GetDataType(), Format::FORMAT_ND); - uBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, vDim}), - params.uOut->GetDataType(), Format::FORMAT_ND); - qgBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), - params.qgOut->GetDataType(), Format::FORMAT_ND); - kgBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), - params.kgOut->GetDataType(), Format::FORMAT_ND); - vNewBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, vDim}), - params.vNewOut->GetDataType(), Format::FORMAT_ND); - hBnst = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, params.totalChunks, kDim, vDim}), - params.hOut->GetDataType(), Format::FORMAT_ND); - } - const bool internalIntermediateOutputsReady = !isInternalLayout || !returnIntermediates || - (aqkBnst != nullptr && akkBnst != nullptr && wBnsd != nullptr && uBnsd != nullptr && - qgBnsd != nullptr && kgBnsd != nullptr && vNewBnsd != nullptr && hBnst != nullptr); - const bool externalComputeBuffersReady = isInternalLayout || - (wBnsd != nullptr && uBnsd != nullptr && qgBnsd != nullptr && kgBnsd != nullptr && - vNewBnsd != nullptr && hBnst != nullptr); - CHECK_RET(qBnsd != nullptr && kBnsd != nullptr && vBnsd != nullptr && gkBnsdRaw != nullptr && - betaBnsRaw != nullptr && oBnsd != nullptr && internalIntermediateOutputsReady && - externalComputeBuffersReady, + } + CHECK_RET((params.wOut == nullptr || wExport != nullptr) && + (params.uOut == nullptr || uExport != nullptr) && + (params.qgOut == nullptr || qgExport != nullptr) && + (params.kgOut == nullptr || kgExport != nullptr) && + (params.vNewOut == nullptr || vNewExport != nullptr) && + (params.hOut == nullptr || hExport != nullptr), + ACLNN_ERR_INNER_NULLPTR); + if (aqkCompute == nullptr) { + aqkCompute = AllocTensor(executorPtr, matrixShape4, params.q->GetDataType()); + } + if (akkCompute == nullptr) { + akkCompute = AllocTensor(executorPtr, matrixShape4, params.q->GetDataType()); + } + const aclTensor *wCompute = wExport == nullptr + ? AllocTensor(executorPtr, splitStages ? kShape4 : placeholderShape, + params.q->GetDataType()) + : wExport; + const aclTensor *uCompute = uExport == nullptr + ? AllocTensor(executorPtr, splitStages ? vShape4 : placeholderShape, + params.q->GetDataType()) + : uExport; + const aclTensor *qgCompute = qgExport == nullptr + ? AllocTensor(executorPtr, splitStages ? kShape4 : placeholderShape, + params.q->GetDataType()) + : qgExport; + const aclTensor *kgCompute = kgExport == nullptr + ? AllocTensor(executorPtr, splitStages ? kShape4 : placeholderShape, + params.q->GetDataType()) + : kgExport; + const aclTensor *vNewCompute = vNewExport == nullptr + ? AllocTensor(executorPtr, splitStages ? vShape4 : placeholderShape, + params.q->GetDataType()) + : vNewExport; + const aclTensor *hCompute = AllocTensor( + executorPtr, hExport == nullptr && !splitStages ? placeholderShape : hShape5, + params.q->GetDataType()); + CHECK_RET(aqkCompute != nullptr && akkCompute != nullptr && wCompute != nullptr && + uCompute != nullptr && qgCompute != nullptr && kgCompute != nullptr && + vNewCompute != nullptr && hCompute != nullptr, ACLNN_ERR_INNER_NULLPTR); - if (!isInternalLayout) { - CHECK_RET(l0op::KdaLayoutSwap12(qBsnd, qBnsd, executorPtr)[0] != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(kBsnd, kBnsd, executorPtr)[0] != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(vBsnd, vBnsd, executorPtr)[0] != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(gkBsnd, gkBnsdRaw, executorPtr)[0] != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(betaBsn, betaBnsRaw, executorPtr)[0] != nullptr, ACLNN_ERR_INNER_NULLPTR); - } - - const aclTensor *gkBnsd = gkBnsdRaw; - const aclTensor *betaBns = betaBnsRaw; - if (gkBnsd->GetDataType() != DataType::DT_FLOAT) { - gkBnsd = l0op::Cast(gkBnsd, DataType::DT_FLOAT, executorPtr); - CHECK_RET(gkBnsd != nullptr, ACLNN_ERR_INNER_NULLPTR); - } - if (betaBns->GetDataType() != DataType::DT_FLOAT) { - betaBns = l0op::Cast(betaBns, DataType::DT_FLOAT, executorPtr); - CHECK_RET(betaBns != nullptr, ACLNN_ERR_INNER_NULLPTR); - } - - std::array result; - bool useSplitForward = true; - const aclTensor *aqkComputeBnst = aqkBnst; - const aclTensor *akkComputeBnst = akkBnst; - const aclTensor *wComputeBnsd = wBnsd; - const aclTensor *uComputeBnsd = uBnsd; - const aclTensor *qgComputeBnsd = qgBnsd; - const aclTensor *kgComputeBnsd = kgBnsd; - const aclTensor *vNewComputeBnsd = vNewBnsd; - const aclTensor *hComputeBnst = hBnst; - const aclTensor *oOutComputeBnsd = nullptr; - const aclTensor *wPreComputeBnsd = nullptr; - const aclTensor *kgScratchComputeBnsd = nullptr; - if (useSplitForward && isInternalLayout) { - aqkComputeBnst = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, params.chunkSize}), - DataType::DT_FLOAT, Format::FORMAT_ND); - akkComputeBnst = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, params.chunkSize}), - DataType::DT_FLOAT, Format::FORMAT_ND); - wComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), - params.wOut->GetDataType(), Format::FORMAT_ND); - uComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, vDim}), - params.uOut->GetDataType(), Format::FORMAT_ND); - qgComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), - params.qgOut->GetDataType(), Format::FORMAT_ND); - kgComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), - params.kgOut->GetDataType(), Format::FORMAT_ND); - vNewComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, vDim}), - params.vNewOut->GetDataType(), Format::FORMAT_ND); - hComputeBnst = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, params.totalChunks, kDim, vDim}), - params.hOut->GetDataType(), Format::FORMAT_ND); - CHECK_RET(aqkComputeBnst != nullptr && akkComputeBnst != nullptr && wComputeBnsd != nullptr && - uComputeBnsd != nullptr && qgComputeBnsd != nullptr && - kgComputeBnsd != nullptr && vNewComputeBnsd != nullptr && hComputeBnst != nullptr, - ACLNN_ERR_INNER_NULLPTR); + const aclTensor *initialStateCompute = params.initialStateOptional; + if (params.stateVFirst && initialStateCompute != nullptr) { + initialStateCompute = TransposeLastTwo(initialStateCompute, executorPtr); + CHECK_RET(initialStateCompute != nullptr, ACLNN_ERR_INNER_NULLPTR); + } + const bool outputFinalState = params.finalStateOut != nullptr; + const aclTensor *finalStateCompute = AllocTensor( + executorPtr, outputFinalState || splitStages ? stateShape4 : placeholderShape, + DataType::DT_FLOAT); + CHECK_RET(finalStateCompute != nullptr, ACLNN_ERR_INNER_NULLPTR); + + const aclTensor *attnCompute = params.attnOut; + if (info.isRank3) { + attnCompute = AsRank4( + params.attnOut, MakeShape({1, info.seqlen, info.hvNum, info.vDim}), executorPtr); + CHECK_RET(attnCompute != nullptr, ACLNN_ERR_INNER_NULLPTR); } - if (useSplitForward) { - if (!isInternalLayout) { - aqkComputeBnst = executorPtr->AllocTensor( - KdaFwdMakeShape({batch, hvNum, seqlen, params.chunkSize}), - DataType::DT_FLOAT, Format::FORMAT_ND); - akkComputeBnst = executorPtr->AllocTensor( - KdaFwdMakeShape({batch, hvNum, seqlen, params.chunkSize}), - DataType::DT_FLOAT, Format::FORMAT_ND); - CHECK_RET(aqkComputeBnst != nullptr && akkComputeBnst != nullptr, ACLNN_ERR_INNER_NULLPTR); - } - oOutComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, Format::FORMAT_ND); - CHECK_RET(oOutComputeBnsd != nullptr, ACLNN_ERR_INNER_NULLPTR); - wPreComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), - params.wOut->GetDataType(), Format::FORMAT_ND); - kgScratchComputeBnsd = executorPtr->AllocTensor(KdaFwdMakeShape({batch, hvNum, seqlen, kDim}), - params.kgOut->GetDataType(), Format::FORMAT_ND); - auto stage1ODummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.oOut->GetDataType(), - Format::FORMAT_ND); - auto stage1FinalStateDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, - Format::FORMAT_ND); - auto stage1UDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.uOut->GetDataType(), - Format::FORMAT_ND); - auto stage1HDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, Format::FORMAT_ND); - CHECK_RET(wPreComputeBnsd != nullptr && kgScratchComputeBnsd != nullptr && stage1ODummy != nullptr && - stage1FinalStateDummy != nullptr && stage1UDummy != nullptr && stage1HDummy != nullptr, - ACLNN_ERR_INNER_NULLPTR); - auto prepResult = l0op::ChunkKdaFwd(qBnsd, kBnsd, vBnsd, gkBnsd, betaBns, params.initialStateOptional, - params.cuSeqlensOptional, params.chunkIndicesOptional, nullptr, nullptr, - nullptr, nullptr, params.scale, params.chunkSize, params.outputFinalState, - params.totalChunks, 1, stage1ODummy, stage1FinalStateDummy, - aqkComputeBnst, akkComputeBnst, wPreComputeBnsd, stage1UDummy, - qgComputeBnsd, kgScratchComputeBnsd, - vNewComputeBnsd, stage1HDummy, executorPtr); - for (auto tensor : prepResult) { - CHECK_RET(tensor != nullptr, ACLNN_ERR_PARAM_NULLPTR); - } - const aclTensor *aqkScaledBnst = - l0op::Muls(aqkComputeBnst, static_cast(params.scale), executorPtr); - CHECK_RET(aqkScaledBnst != nullptr, ACLNN_ERR_INNER_NULLPTR); - const aclTensor *aqkForOutBnst = KdaFwdMaybeCast(aqkScaledBnst, qBnsd->GetDataType(), executorPtr); - CHECK_RET(aqkForOutBnst != nullptr, ACLNN_ERR_INNER_NULLPTR); - const aclTensor *qgScaledBnsd = - l0op::Muls(qgComputeBnsd, static_cast(params.scale), executorPtr); - CHECK_RET(qgScaledBnsd != nullptr, ACLNN_ERR_INNER_NULLPTR); - - auto wScratchBntd = executorPtr->AllocTensor( - KdaFwdMakeShape({batch, hvNum, params.totalChunks, params.chunkSize, kDim}), - DataType::DT_FLOAT, Format::FORMAT_ND); - const aclTensor *akkPostBnst = KdaFwdMaybeCast(akkComputeBnst, qBnsd->GetDataType(), executorPtr); - CHECK_RET(wScratchBntd != nullptr && akkPostBnst != nullptr, ACLNN_ERR_INNER_NULLPTR); - - auto stage3ODummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.oOut->GetDataType(), - Format::FORMAT_ND); - auto stage3FinalStateDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, - Format::FORMAT_ND); - auto stage3AqkDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, - Format::FORMAT_ND); - auto stage3AkkDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), qBnsd->GetDataType(), - Format::FORMAT_ND); - auto stage3QGDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.qgOut->GetDataType(), - Format::FORMAT_ND); - auto stage3VNewDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.vNewOut->GetDataType(), - Format::FORMAT_ND); - CHECK_RET(stage3ODummy != nullptr && stage3FinalStateDummy != nullptr && stage3AqkDummy != nullptr && - stage3AkkDummy != nullptr && stage3QGDummy != nullptr && stage3VNewDummy != nullptr, - ACLNN_ERR_INNER_NULLPTR); - auto postResult = l0op::ChunkKdaFwd( - qBnsd, kBnsd, vBnsd, gkBnsd, betaBns, params.initialStateOptional, - params.cuSeqlensOptional, params.chunkIndicesOptional, wPreComputeBnsd, akkPostBnst, - vNewComputeBnsd, nullptr, params.scale, params.chunkSize, params.outputFinalState, - params.totalChunks, 3, stage3ODummy, stage3FinalStateDummy, stage3AqkDummy, stage3AkkDummy, - wComputeBnsd, uComputeBnsd, stage3QGDummy, kgComputeBnsd, stage3VNewDummy, wScratchBntd, - executorPtr); - for (auto tensor : postResult) { - CHECK_RET(tensor != nullptr, ACLNN_ERR_PARAM_NULLPTR); - } + const aclTensor *qgScaledCompute = AllocTensor( + executorPtr, splitStages ? kShape4 : placeholderShape, + params.q->GetDataType()); + const aclTensor *uSeedCompute = AllocTensor( + executorPtr, splitStages ? vShape4 : placeholderShape, + params.q->GetDataType()); + CHECK_RET(qgScaledCompute != nullptr && uSeedCompute != nullptr, + ACLNN_ERR_INNER_NULLPTR); - const aclTensor *neutralGForH = l0op::ZerosLike(betaBns, executorPtr); - CHECK_RET(neutralGForH != nullptr, ACLNN_ERR_INNER_NULLPTR); - auto hResult = l0op::ChunkGatedDeltaRuleFwdH( - kgComputeBnsd, wComputeBnsd, uComputeBnsd, neutralGForH, gkBnsd, - params.initialStateOptional, params.cuSeqlensOptional, params.chunkIndicesOptional, - params.outputFinalState, params.chunkSize, hComputeBnst, vNewComputeBnsd, - params.finalStateOut, executorPtr); - for (auto tensor : hResult) { - CHECK_RET(tensor != nullptr, ACLNN_ERR_PARAM_NULLPTR); + auto launchStage = [&](int64_t stage) { + return l0op::KdaChunkForward( + qHead, kHead, vHead, gHead, betaHead, params.aLogOptional, + params.dtBiasOptional, initialStateCompute, params.cuSeqlensOptional, + params.chunkIndicesOptional, params.scale, params.chunkSize, + params.safeGate, parsedLayout == KdaFwdLayout::BSND, + params.useGateInKernel, params.lowerBound, attnCompute, + finalStateCompute, gkCompute, aqkCompute, akkCompute, wCompute, + uCompute, qgCompute, kgCompute, vNewCompute, hCompute, + qgScaledCompute, uSeedCompute, stage, executorPtr); + }; + l0op::KdaCoreOutputs result{}; + if (splitStages) { + // Physical launch boundaries reset the A5 event state between the + // prepare, post-WU, recurrent, and output pipelines. + for (int64_t stage = KDA_STAGE_GATE_PREPARE; stage < KDA_STAGE_COUNT; + ++stage) { + result = launchStage(stage); + for (const aclTensor *tensor : result) { + CHECK_RET(tensor != nullptr, ACLNN_ERR_INNER_NULLPTR); + } } - - auto oLocalDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, - Format::FORMAT_ND); - auto stage2FinalStateDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, - Format::FORMAT_ND); - auto stage2AqkDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, - Format::FORMAT_ND); - auto stage2AkkDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, - Format::FORMAT_ND); - auto stage2WDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.wOut->GetDataType(), - Format::FORMAT_ND); - auto stage2QGDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.qgOut->GetDataType(), - Format::FORMAT_ND); - auto stage2KGDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), params.kgOut->GetDataType(), - Format::FORMAT_ND); - auto stage2HDummy = executorPtr->AllocTensor(KdaFwdMakeShape({1}), DataType::DT_FLOAT, - Format::FORMAT_ND); - CHECK_RET(oLocalDummy != nullptr && stage2FinalStateDummy != nullptr && - stage2AqkDummy != nullptr && stage2AkkDummy != nullptr && stage2WDummy != nullptr && - stage2QGDummy != nullptr && stage2KGDummy != nullptr && stage2HDummy != nullptr, - ACLNN_ERR_INNER_NULLPTR); - auto outResult = l0op::ChunkKdaFwd( - qBnsd, kBnsd, vBnsd, gkBnsd, betaBns, params.initialStateOptional, - params.cuSeqlensOptional, params.chunkIndicesOptional, qgScaledBnsd, aqkForOutBnst, - vNewComputeBnsd, hComputeBnst, params.scale, params.chunkSize, false, params.totalChunks, 2, - oOutComputeBnsd, stage2FinalStateDummy, stage2AqkDummy, stage2AkkDummy, stage2WDummy, - oLocalDummy, stage2QGDummy, stage2KGDummy, oBnsd, stage2HDummy, executorPtr); - for (auto tensor : outResult) { - CHECK_RET(tensor != nullptr, ACLNN_ERR_PARAM_NULLPTR); + } else { + result = launchStage(KDA_STAGE_FULL); + for (const aclTensor *tensor : result) { + CHECK_RET(tensor != nullptr, ACLNN_ERR_INNER_NULLPTR); } - result = {oBnsd, params.finalStateOut, aqkScaledBnst, akkComputeBnst, wComputeBnsd, - uComputeBnsd, qgComputeBnsd, kgComputeBnsd, vNewComputeBnsd, hComputeBnst}; } - for (auto tensor : result) { - CHECK_RET(tensor != nullptr, ACLNN_ERR_PARAM_NULLPTR); - } - if (isInternalLayout) { - if (useSplitForward) { - if (result[0] != oBnsd) { - CHECK_RET(KdaFwdViewCopyMaybeCast(result[0], oBnsd, executorPtr) == ACLNN_SUCCESS, - ACLNN_ERR_INNER_NULLPTR); - } - if (returnIntermediates) { - CHECK_RET(KdaFwdCopyMaybeCastAfter(result[2], oBnsd, aqkBnst, executorPtr) == ACLNN_SUCCESS, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(KdaFwdCopyMaybeCastAfter(result[3], aqkBnst, akkBnst, executorPtr) == ACLNN_SUCCESS, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(KdaFwdCopyMaybeCastAfter(result[4], akkBnst, wBnsd, executorPtr) == ACLNN_SUCCESS, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(KdaFwdCopyMaybeCastAfter(result[5], wBnsd, uBnsd, executorPtr) == ACLNN_SUCCESS, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(KdaFwdCopyMaybeCastAfter(result[6], uBnsd, qgBnsd, executorPtr) == ACLNN_SUCCESS, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(KdaFwdCopyMaybeCastAfter(result[7], qgBnsd, kgBnsd, executorPtr) == ACLNN_SUCCESS, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(KdaFwdCopyMaybeCastAfter(result[8], kgBnsd, vNewBnsd, executorPtr) == ACLNN_SUCCESS, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(KdaFwdCopyMaybeCastAfter(result[9], vNewBnsd, hBnst, executorPtr) == ACLNN_SUCCESS, - ACLNN_ERR_INNER_NULLPTR); - } + + if (outputFinalState) { + const aclTensor *finalStateResult = result[1]; + if (params.stateVFirst) { + finalStateResult = TransposeLastTwo(finalStateResult, executorPtr); + CHECK_RET(finalStateResult != nullptr, ACLNN_ERR_INNER_NULLPTR); } - } else if (isTnd) { - auto oBsnd = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, vDim}), - params.oOut->GetDataType(), Format::FORMAT_ND); - CHECK_RET(oBsnd != nullptr, ACLNN_ERR_INNER_NULLPTR); - const aclTensor *oForLayout = KdaFwdMaybeCast(result[0], params.oOut->GetDataType(), executorPtr); - CHECK_RET(oForLayout != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(oForLayout, nullptr, oBsnd, executorPtr)[0] != nullptr, + CHECK_RET(l0op::ViewCopy(finalStateResult, params.finalStateOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::ViewCopy(l0op::Reshape(oBsnd, KdaFwdMakeShape({seqlen, hvNum, vDim}), executorPtr), - params.oOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); - if (returnIntermediates) { - const aclTensor *aqkForLayout = KdaFwdMaybeCast(result[2], params.aqkOut->GetDataType(), executorPtr); - CHECK_RET(aqkForLayout != nullptr, ACLNN_ERR_INNER_NULLPTR); - const aclTensor *akkForLayout = KdaFwdMaybeCast(result[3], params.akkOut->GetDataType(), executorPtr); - CHECK_RET(akkForLayout != nullptr, ACLNN_ERR_INNER_NULLPTR); - auto aqkBsnt = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, params.chunkSize}), - params.aqkOut->GetDataType(), Format::FORMAT_ND); - auto akkBsnt = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, params.chunkSize}), - params.akkOut->GetDataType(), Format::FORMAT_ND); - auto wBsnd = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, kDim}), - params.wOut->GetDataType(), Format::FORMAT_ND); - auto uBsnd = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, vDim}), - params.uOut->GetDataType(), Format::FORMAT_ND); - auto qgBsnd = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, kDim}), - params.qgOut->GetDataType(), Format::FORMAT_ND); - auto kgBsnd = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, kDim}), - params.kgOut->GetDataType(), Format::FORMAT_ND); - auto vNewBsnd = executorPtr->AllocTensor(KdaFwdMakeShape({1, seqlen, hvNum, vDim}), - params.vNewOut->GetDataType(), Format::FORMAT_ND); - auto hBsnt = executorPtr->AllocTensor(KdaFwdMakeShape({1, params.totalChunks, hvNum, kDim, vDim}), - params.hOut->GetDataType(), Format::FORMAT_ND); - CHECK_RET(aqkBsnt != nullptr && akkBsnt != nullptr && wBsnd != nullptr && uBsnd != nullptr && - qgBsnd != nullptr && kgBsnd != nullptr && vNewBsnd != nullptr && hBsnt != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(aqkForLayout, oBsnd, aqkBsnt, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(akkForLayout, aqkBsnt, akkBsnt, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(result[4], akkBsnt, wBsnd, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(result[5], wBsnd, uBsnd, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(result[6], uBsnd, qgBsnd, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(result[7], qgBsnd, kgBsnd, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(result[8], kgBsnd, vNewBsnd, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(result[9], vNewBsnd, hBsnt, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::ViewCopy(l0op::Reshape(aqkBsnt, KdaFwdMakeShape({seqlen, hvNum, params.chunkSize}), - executorPtr), - params.aqkOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::ViewCopy(l0op::Reshape(akkBsnt, KdaFwdMakeShape({seqlen, hvNum, params.chunkSize}), - executorPtr), - params.akkOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::ViewCopy(l0op::Reshape(wBsnd, KdaFwdMakeShape({seqlen, hvNum, kDim}), executorPtr), - params.wOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::ViewCopy(l0op::Reshape(uBsnd, KdaFwdMakeShape({seqlen, hvNum, vDim}), executorPtr), - params.uOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::ViewCopy(l0op::Reshape(qgBsnd, KdaFwdMakeShape({seqlen, hvNum, kDim}), executorPtr), - params.qgOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::ViewCopy(l0op::Reshape(kgBsnd, KdaFwdMakeShape({seqlen, hvNum, kDim}), executorPtr), - params.kgOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::ViewCopy(l0op::Reshape(vNewBsnd, KdaFwdMakeShape({seqlen, hvNum, vDim}), executorPtr), - params.vNewOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::ViewCopy(l0op::Reshape(hBsnt, KdaFwdMakeShape({params.totalChunks, hvNum, kDim, vDim}), - executorPtr), - params.hOut, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); - } - } else { - const aclTensor *oBsnd = params.oOut; - const aclTensor *oForLayout = KdaFwdMaybeCast(result[0], params.oOut->GetDataType(), executorPtr); - CHECK_RET(oForLayout != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(oForLayout, nullptr, oBsnd, executorPtr)[0] != nullptr, + } + if (hExport != nullptr) { + const std::vector hPerm = + params.stateVFirst ? std::vector{0, 2, 1, 4, 3} + : std::vector{0, 2, 1, 3, 4}; + const aclTensor *hResult = Transpose(result[10], hPerm, executorPtr); + CHECK_RET(hResult != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(l0op::ViewCopy(hResult, hExport, executorPtr) != nullptr, ACLNN_ERR_INNER_NULLPTR); - if (returnIntermediates) { - const aclTensor *aqkForLayout = KdaFwdMaybeCast(result[2], params.aqkOut->GetDataType(), executorPtr); - CHECK_RET(aqkForLayout != nullptr, ACLNN_ERR_INNER_NULLPTR); - const aclTensor *akkForLayout = KdaFwdMaybeCast(result[3], params.akkOut->GetDataType(), executorPtr); - CHECK_RET(akkForLayout != nullptr, ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(aqkForLayout, oBsnd, params.aqkOut, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(akkForLayout, params.aqkOut, params.akkOut, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(result[4], params.akkOut, params.wOut, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(result[5], params.wOut, params.uOut, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(result[6], params.uOut, params.qgOut, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(result[7], params.qgOut, params.kgOut, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(result[8], params.kgOut, params.vNewOut, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - CHECK_RET(l0op::KdaLayoutSwap12(result[9], params.vNewOut, params.hOut, executorPtr)[0] != nullptr, - ACLNN_ERR_INNER_NULLPTR); - } } + *workspaceSize = uniqueExecutor->GetWorkspaceSize(); uniqueExecutor.ReleaseTo(executor); return ACLNN_SUCCESS; } -aclnnStatus aclnnChunkKdaFwd(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream) +aclnnStatus aclnnChunkKdaFwd(void *workspace, uint64_t workspaceSize, + aclOpExecutor *executor, aclrtStream stream) { L2_DFX_PHASE_2(aclnnChunkKdaFwd); - CHECK_COND(CommonOpExecutorRun(workspace, workspaceSize, executor, stream) == ACLNN_SUCCESS, ACLNN_ERR_INNER, - "ChunkKdaFwd launch failed."); + CHECK_COND(CommonOpExecutorRun(workspace, workspaceSize, executor, stream) == ACLNN_SUCCESS, + ACLNN_ERR_INNER, "ChunkKdaFwd launch failed."); return ACLNN_SUCCESS; } diff --git a/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.h b/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.h index 63301d16d227..735a898d0a0f 100644 --- a/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.h +++ b/csrc/attention/chunk_kda_fwd/op_host/op_api/aclnn_chunk_kda_fwd.h @@ -20,18 +20,23 @@ aclnnStatus aclnnChunkKdaFwdGetWorkspaceSize( const aclTensor *q, const aclTensor *k, const aclTensor *v, - const aclTensor *gk, + const aclTensor *g, const aclTensor *beta, + const aclTensor *aLogOptional, + const aclTensor *dtBiasOptional, const aclTensor *initialStateOptional, const aclIntArray *cuSeqlensOptional, const aclIntArray *chunkIndicesOptional, const char *layout, double scale, int64_t chunkSize, - bool outputFinalState, - int64_t totalChunks, - const aclTensor *oOut, + bool safeGate, + double lowerBound, + bool useGateInKernel, + bool stateVFirst, + const aclTensor *attnOut, const aclTensor *finalStateOut, + const aclTensor *gkOut, const aclTensor *aqkOut, const aclTensor *akkOut, const aclTensor *wOut, diff --git a/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.cpp b/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.cpp index 999c8e2b367b..8a76c9a48814 100644 --- a/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.cpp +++ b/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.cpp @@ -13,92 +13,31 @@ #include "opdev/op_dfx.h" #include "opdev/op_log.h" -#include - using namespace op; namespace l0op { OP_TYPE_REGISTER(ChunkKdaFwd); -namespace { -const aclIntArray *BuildPackedChunkMetadata(const aclIntArray *cuSeqlens, - const aclIntArray *chunkIndices, - int64_t chunkSize, - int64_t totalChunks, - aclOpExecutor *executor) -{ - if (cuSeqlens == nullptr || cuSeqlens->Size() < 2 || chunkSize <= 0 || totalChunks <= 0) { - return nullptr; - } - - const aclIntArray &cu = *cuSeqlens; - std::vector packed; - packed.reserve(static_cast(totalChunks) * 4); - auto appendChunk = [&](int64_t seq, int64_t localChunk) -> bool { - if (seq < 0 || static_cast(seq + 1) >= cu.Size() || localChunk < 0) { - return false; - } - int64_t seqStart = cu[static_cast(seq)]; - int64_t seqEnd = cu[static_cast(seq + 1)]; - int64_t start = seqStart + localChunk * chunkSize; - if (start < seqStart || start >= seqEnd) { - return false; - } - int64_t end = start + chunkSize; - if (end > seqEnd) { - end = seqEnd; - } - packed.insert(packed.end(), {seq, start, end, 0}); - return true; - }; - - if (chunkIndices != nullptr) { - if (chunkIndices->Size() != static_cast(totalChunks) * 2) { - return nullptr; - } - for (size_t idx = 0; idx < chunkIndices->Size(); idx += 2) { - if (!appendChunk((*chunkIndices)[idx], (*chunkIndices)[idx + 1])) { - return nullptr; - } - } - } else { - for (size_t seq = 0; seq + 1 < cu.Size(); ++seq) { - int64_t seqLength = cu[seq + 1] - cu[seq]; - int64_t chunkCount = (seqLength + chunkSize - 1) / chunkSize; - for (int64_t localChunk = 0; localChunk < chunkCount; ++localChunk) { - if (!appendChunk(static_cast(seq), localChunk)) { - return nullptr; - } - } - } - } - if (packed.size() != static_cast(totalChunks) * 4) { - return nullptr; - } - return executor->AllocIntArray(packed.data(), packed.size()); -} -} // namespace - -const std::array ChunkKdaFwd( +KdaCoreOutputs KdaChunkForward( const aclTensor *q, const aclTensor *k, const aclTensor *v, - const aclTensor *gk, + const aclTensor *g, const aclTensor *beta, + const aclTensor *aLogOptional, + const aclTensor *dtBiasOptional, const aclTensor *initialStateOptional, const aclIntArray *cuSeqlensOptional, const aclIntArray *chunkIndicesOptional, - const aclTensor *stageQGInputOptional, - const aclTensor *stageAqkInputOptional, - const aclTensor *stageVNewInputOptional, - const aclTensor *stageHInputOptional, double scale, int64_t chunkSize, - bool outputFinalState, - int64_t totalChunks, - int64_t stage, - const aclTensor *oOut, + bool safeGate, + bool inputSequenceMajor, + bool useGateInKernel, + double lowerBound, + const aclTensor *attnOut, const aclTensor *finalStateOut, + const aclTensor *gkOut, const aclTensor *aqkOut, const aclTensor *akkOut, const aclTensor *wOut, @@ -107,12 +46,17 @@ const std::array ChunkKdaFwd( const aclTensor *kgOut, const aclTensor *vNewOut, const aclTensor *hOut, + const aclTensor *qgScaledOut, + const aclTensor *uSeedOut, + int64_t stage, aclOpExecutor *executor) { - L0_DFX(ChunkKdaFwd, q, k, v, gk, beta, initialStateOptional, cuSeqlensOptional, chunkIndicesOptional, - stageQGInputOptional, stageAqkInputOptional, stageVNewInputOptional, stageHInputOptional, scale, chunkSize, - outputFinalState, totalChunks, stage, oOut, finalStateOut, aqkOut, akkOut, wOut, uOut, qgOut, kgOut, - vNewOut, hOut); + L0_DFX(KdaChunkForward, q, k, v, g, beta, aLogOptional, dtBiasOptional, + initialStateOptional, cuSeqlensOptional, chunkIndicesOptional, + scale, chunkSize, safeGate, inputSequenceMajor, useGateInKernel, + lowerBound, attnOut, finalStateOut, gkOut, aqkOut, akkOut, + wOut, uOut, qgOut, kgOut, vNewOut, hOut, qgScaledOut, uSeedOut, + stage); const aclTensor *actualCuSeqlens = nullptr; if (cuSeqlensOptional != nullptr) { @@ -123,17 +67,11 @@ const std::array ChunkKdaFwd( } const aclTensor *actualChunkIndices = nullptr; - if (cuSeqlensOptional != nullptr) { - const aclIntArray *packedChunkMetadata = BuildPackedChunkMetadata( - cuSeqlensOptional, chunkIndicesOptional, chunkSize, totalChunks, executor); - if (packedChunkMetadata == nullptr) { - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "failed to build packed chunk metadata."); - return {nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr}; - } - actualChunkIndices = executor->ConvertToTensor(packedChunkMetadata, DataType::DT_INT64); + if (chunkIndicesOptional != nullptr) { + actualChunkIndices = executor->ConvertToTensor(chunkIndicesOptional, DataType::DT_INT64); if (actualChunkIndices == nullptr) { - OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "failed to convert packed chunk metadata to tensor."); - return {nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr}; + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "failed to convert chunk metadata to tensor."); + return {}; } const_cast(actualChunkIndices)->SetStorageFormat(Format::FORMAT_ND); const_cast(actualChunkIndices)->SetViewFormat(Format::FORMAT_ND); @@ -142,15 +80,20 @@ const std::array ChunkKdaFwd( auto ret = ADD_TO_LAUNCHER_LIST_AICORE( ChunkKdaFwd, - OP_INPUT(q, k, v, gk, beta, initialStateOptional, actualCuSeqlens, actualChunkIndices, - stageQGInputOptional, stageAqkInputOptional, stageVNewInputOptional, stageHInputOptional), - OP_OUTPUT(oOut, finalStateOut, aqkOut, akkOut, wOut, uOut, qgOut, kgOut, vNewOut, hOut), - OP_ATTR(scale, chunkSize, outputFinalState, totalChunks, stage)); + OP_INPUT(q, k, v, g, beta, aLogOptional, dtBiasOptional, + initialStateOptional, actualCuSeqlens, actualChunkIndices), + OP_OUTPUT(attnOut, finalStateOut, gkOut, aqkOut, akkOut, wOut, uOut, + qgOut, kgOut, vNewOut, hOut, qgScaledOut, uSeedOut), + OP_ATTR(inputSequenceMajor ? "BSND" : "BNSD", scale, chunkSize, + safeGate, static_cast(lowerBound), useGateInKernel, + false, stage)); if (ret != ACLNN_SUCCESS) { - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "ADD_TO_LAUNCHER_LIST_AICORE ChunkKdaFwd failed."); - return {nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr}; + OP_LOGE(ACLNN_ERR_PARAM_INVALID, + "ADD_TO_LAUNCHER_LIST_AICORE ChunkKdaFwd failed."); + return {}; } - return {oOut, finalStateOut, aqkOut, akkOut, wOut, uOut, qgOut, kgOut, vNewOut, hOut}; + return {attnOut, finalStateOut, gkOut, aqkOut, akkOut, wOut, uOut, + qgOut, kgOut, vNewOut, hOut, qgScaledOut, uSeedOut}; } } // namespace l0op diff --git a/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.h b/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.h index b606715c96a8..78aa4ad5e755 100644 --- a/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.h +++ b/csrc/attention/chunk_kda_fwd/op_host/op_api/chunk_kda_fwd.h @@ -13,35 +13,20 @@ #include "opdev/op_executor.h" namespace l0op { -const std::array ChunkKdaFwd( - const aclTensor *q, - const aclTensor *k, - const aclTensor *v, - const aclTensor *gk, - const aclTensor *beta, - const aclTensor *initialStateOptional, - const aclIntArray *cuSeqlensOptional, - const aclIntArray *chunkIndicesOptional, - const aclTensor *stageQGInputOptional, - const aclTensor *stageAqkInputOptional, - const aclTensor *stageVNewInputOptional, - const aclTensor *stageHInputOptional, - double scale, - int64_t chunkSize, - bool outputFinalState, - int64_t totalChunks, - int64_t stage, - const aclTensor *oOut, - const aclTensor *finalStateOut, - const aclTensor *aqkOut, - const aclTensor *akkOut, - const aclTensor *wOut, - const aclTensor *uOut, - const aclTensor *qgOut, - const aclTensor *kgOut, - const aclTensor *vNewOut, - const aclTensor *hOut, +using KdaCoreOutputs = std::array; + +KdaCoreOutputs KdaChunkForward( + const aclTensor *q, const aclTensor *k, const aclTensor *v, const aclTensor *g, const aclTensor *beta, + const aclTensor *aLogOptional, const aclTensor *dtBiasOptional, + const aclTensor *initialStateOptional, const aclIntArray *cuSeqlensOptional, + const aclIntArray *chunkIndicesOptional, double scale, int64_t chunkSize, + bool safeGate, bool inputSequenceMajor, bool useGateInKernel, double lowerBound, + const aclTensor *attnOut, + const aclTensor *finalStateOut, const aclTensor *gkOut, const aclTensor *aqkOut, + const aclTensor *akkOut, const aclTensor *wOut, const aclTensor *uOut, const aclTensor *qgOut, + const aclTensor *kgOut, const aclTensor *vNewOut, const aclTensor *hOut, + const aclTensor *qgScaledOut, const aclTensor *uSeedOut, int64_t stage, aclOpExecutor *executor); -} +} // namespace l0op #endif diff --git a/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_finalize.h b/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_finalize.h new file mode 100644 index 000000000000..5045855c0486 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_finalize.h @@ -0,0 +1,1446 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#pragma once + +#ifndef CATLASS_ARCH +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +#define CATLASS_ARCH 3510 +#else +#define CATLASS_ARCH 2201 +#endif +#endif + +#include "catlass/arch/arch.hpp" +#include "catlass/arch/cross_core_sync.hpp" +#include "catlass/arch/resource.hpp" +#include "catlass/catlass.hpp" +#include "catlass/gemm/block/block_mmad.hpp" +#include "catlass/gemm/dispatch_policy.hpp" +#include "catlass/gemm/tile/tile_copy.hpp" +#include "catlass/gemm_coord.hpp" +#include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp" +#include "catlass/layout/layout.hpp" +#include "kernel_operator.h" +#include "../chunk_kda_fwd_varlen.h" +#include "tla/layout.hpp" +#include "tla/tensor.hpp" + +using namespace AscendC; + +namespace KdaFinalize { +namespace { +using KdaInt64 = tla::Int<64>; +using KdaInt128 = tla::Int<128>; +constexpr float LN2 = 0.69314718055994530942f; +constexpr float KDA_EXP2_CLAMP = 80.0f; +constexpr float KDA_EXP_INPUT_MAX = KDA_EXP2_CLAMP * LN2; +constexpr float KDA_EXP_INPUT_MIN = -KDA_EXP2_CLAMP * LN2; +constexpr float KDA_FP16_MAX = 65504.0f; +constexpr uint32_t EXP2_UB_ELEMENTS = 256; +constexpr uint32_t EXP2_UB_BYTES = EXP2_UB_ELEMENTS * (sizeof(float) + sizeof(uint16_t)); +constexpr uint32_t EXP2_EVENT_ID = 0; +constexpr uint32_t KDA_SOLVE_BT = 64; +constexpr uint32_t KDA_SOLVE_MATRIX_ELEMENTS = KDA_SOLVE_BT * KDA_SOLVE_BT; +constexpr uint32_t KDA_SOLVE_SCRATCH_X = 0; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y0 = 1; +constexpr uint32_t KDA_SOLVE_SCRATCH_TMP = 2; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y1 = 3; +constexpr uint32_t KDA_SOLVE_SCRATCH_IDENTITY = 4; +constexpr uint32_t KDA_SOLVE_SCRATCH_SLOTS = 5; +constexpr uint32_t KDA_SOLVE_DIAG_BT = 16; +constexpr uint32_t KDA_SOLVE_DIAG_BLOCKS = KDA_SOLVE_BT / KDA_SOLVE_DIAG_BT; +constexpr uint32_t KDA_SOLVE_DIAG_MCH_ITERS = 3; +constexpr uint32_t KDA_SCORE_REF_BC = 16; +constexpr uint32_t KDA_VEC_ARENA_ELEMENTS = 32768; +constexpr uint32_t KDA_BITS_PER_MASK_BYTE = 8; +constexpr uint32_t KDA_SELECT_COL_BLOCKS = 2; +constexpr uint32_t KDA_SELECT_COL_MASK_BYTES = KDA_SOLVE_MATRIX_ELEMENTS / KDA_BITS_PER_MASK_BYTE; +constexpr uint32_t KDA_SELECT_MASK_BYTES = KDA_SELECT_COL_BLOCKS * KDA_SELECT_COL_MASK_BYTES; +constexpr uint32_t KDA_SELECT_AQK_MASK_BYTE_OFFSET = 120 * 1024; +constexpr uint32_t KDA_SELECT_AKK_MASK_BYTE_OFFSET = KDA_SELECT_AQK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_BYTE_OFFSET = KDA_SELECT_AKK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_FLOAT_OFFSET = KDA_SELECT_ZERO_BYTE_OFFSET / sizeof(float); +constexpr uint8_t KDA_SCORE_DONE_FLAG0 = 2; +constexpr uint8_t KDA_SCORE_DONE_FLAG1 = 3; +constexpr uint8_t KDA_SCORE_READY_FLAG0 = 4; +constexpr uint8_t KDA_SCORE_READY_FLAG1 = 5; +constexpr uint32_t KDA_SCORE_QUEUE_DEPTH = 2; +constexpr uint32_t KDA_SYNC_REVERSE_DEPTH = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_PLANES = 3; +constexpr uint32_t KDA_SCORE_SCRATCH_QG = 0; +constexpr uint32_t KDA_SCORE_SCRATCH_W = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_KG = 2; +constexpr uint64_t KDA_WORKSPACE_ALIGN = 512; +constexpr uint32_t KDA_GATE_TILE_ROWS = 32; +constexpr uint32_t KDA_CUBE_MIN_REDUCTION = 16; + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +using KdaArchTag = Catlass::Arch::Ascend950; +#else +using KdaArchTag = Catlass::Arch::AtlasA2; +#endif +using KdaDispatchPolicy = Catlass::Gemm::MmadPingpong; +using KdaScoreDispatchPolicy = + Catlass::Gemm::MmadPingpongTlaMulti; +static_assert(KdaScoreDispatchPolicy::ENABLE_L1_RESIDENT, + "KDA Aqk/Akk score MMAD must keep the shared right matrix resident in L1"); +static_assert(KdaScoreDispatchPolicy::L1B_STAGES == 1, + "KDA Aqk/Akk score MMAD needs one L1 B slot so the second MMAD reuses it"); +using KdaSolveDispatchPolicy = Catlass::Gemm::MmadPingpong; +static_assert(!KdaSolveDispatchPolicy::USE_HF32_MODE, "KDA triangular solve must use IEEE FP32 Cube mode"); +using KdaL1TileShape = tla::Shape; +using KdaL0TileShape = KdaL1TileShape; +using KdaSolveL1TileShape = tla::Shape; +using KdaSolveL0TileShape = KdaSolveL1TileShape; + +__aicore__ inline uint32_t FloatToBits(float value) +{ + union Bits { + __aicore__ Bits() {} + float f; + uint32_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline float BitsToFloat(uint32_t value) +{ + union Bits { + __aicore__ Bits() {} + uint32_t u; + float f; + } bits; + bits.u = value; + return bits.f; +} + +__aicore__ inline uint16_t Bf16ToBits(bfloat16_t value) +{ + union Bits { + __aicore__ Bits() {} + bfloat16_t f; + uint16_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline bfloat16_t BitsToBf16(uint16_t value) +{ + union Bits { + __aicore__ Bits() {} + uint16_t u; + bfloat16_t f; + } bits; + bits.u = value; + return bits.f; +} + +template +__aicore__ inline T FloatToType(float value) +{ + if constexpr (IsSameType::value) { + uint32_t bits = FloatToBits(value); + uint32_t bias = 0x7FFFu + ((bits >> 16) & 1u); + return BitsToBf16(static_cast((bits + bias) >> 16)); + } + return static_cast(value); +} + +template +class ChunkKdaFwdFinalizeKernel { +public: + using OUT_T = float; + using AKK_T = float; + template + __aicore__ inline void Init(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR preparedQG, GM_ADDR preparedAqk, + GM_ADDR propagatedVNew, GM_ADDR propagatedH, GM_ADDR o, GM_ADDR finalState, GM_ADDR aqk, + GM_ADDR akk, GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR workspace, const TilingData &tiling, TPipe *pipe, + bool initVecBuffers = true) + { + pipe_ = pipe; + q_.SetGlobalBuffer((__gm__ T *)q); + k_.SetGlobalBuffer((__gm__ T *)k); + v_.SetGlobalBuffer((__gm__ T *)v); + gk_.SetGlobalBuffer((__gm__ GK_T *)gk); + beta_.SetGlobalBuffer((__gm__ BETA_T *)beta); + if (initialState != nullptr) { + initialState_.SetGlobalBuffer((__gm__ float *)initialState); + } + cuSeqlensAddr_ = reinterpret_cast<__gm__ int64_t *>(cuSeqlens); + if (preparedQG != nullptr) { + preparedQG_.SetGlobalBuffer((__gm__ T *)preparedQG); + } + if (preparedAqk != nullptr) { + preparedAqk_.SetGlobalBuffer((__gm__ T *)preparedAqk); + } + if (propagatedVNew != nullptr) { + propagatedVNew_.SetGlobalBuffer((__gm__ T *)propagatedVNew); + } + if (propagatedH != nullptr) { + propagatedH_.SetGlobalBuffer((__gm__ T *)propagatedH); + } + chunkIndicesAddr_ = reinterpret_cast<__gm__ int64_t *>(chunkIndices); + o_.SetGlobalBuffer((__gm__ OUT_T *)o); + finalState_.SetGlobalBuffer((__gm__ float *)finalState); + aqk_.SetGlobalBuffer((__gm__ float *)aqk); + akk_.SetGlobalBuffer((__gm__ AKK_T *)akk); + w_.SetGlobalBuffer((__gm__ T *)w); + u_.SetGlobalBuffer((__gm__ OUT_T *)u); + qg_.SetGlobalBuffer((__gm__ T *)qg); + kg_.SetGlobalBuffer((__gm__ T *)kg); + vNew_.SetGlobalBuffer((__gm__ T *)vNew); + h_.SetGlobalBuffer((__gm__ float *)h); + solveWorkspace_.SetGlobalBuffer((__gm__ float *)workspace); + + B_ = tiling.batch; + N_ = tiling.seqNum; + H_ = tiling.qHeadNum; + HV_ = tiling.vHeadNum; + T_ = tiling.seqlen; + K_ = tiling.kHeadDim; + V_ = tiling.vHeadDim; + BT_ = tiling.chunkSize; + NT_ = tiling.totalChunks; + scale_ = tiling.scale; + hasInitial_ = tiling.hasInitialState; + isVarLen_ = tiling.isVarLen; + usedCoreNum_ = tiling.outputUsedCoreNum; + const uint64_t outputElements = B_ * HV_ * T_ * V_; + o_.SetGlobalBuffer((__gm__ OUT_T *)workspace); + u_.SetGlobalBuffer((__gm__ OUT_T *)workspace + outputElements); + if ASCEND_IS_AIV { + uint64_t subBlockNum = static_cast(GetSubBlockNum()); + solveCoreIdx_ = subBlockNum == 0 ? 0 : static_cast(GetBlockIdx()) / subBlockNum; + } else { + solveCoreIdx_ = static_cast(GetBlockIdx()); + } + if (pipe_ != nullptr && initVecBuffers) { + pipe_->InitBuffer(exp2Buf_, EXP2_UB_BYTES); + pipe_->InitBuffer(vecBuf_, KDA_VEC_ARENA_ELEMENTS * sizeof(float)); + const uint64_t gateWritebackRows = + ScoreVectorMaxRows(5 * sizeof(float) + 2 * sizeof(T) + sizeof(GK_T)); + pipe_->InitBuffer(gateWritebackBuf_, + static_cast(gateWritebackRows * K_ * + (3 * sizeof(T) + sizeof(GK_T)))); + AllocVectorEvents(); + } + } + __aicore__ inline void ProcessAiv() + { + ProcessOutAiv(); + ReleaseVectorEvents(); + } + + __aicore__ inline void ProcessAic() + { + ProcessOutAic(); + } + +private: + __aicore__ inline void AllocVectorEvents() + { + mte2ToVEvent_ = pipe_->AllocEventID(); + vToMte2Event_ = pipe_->AllocEventID(); + vToMte3Event_ = pipe_->AllocEventID(); + mte3ToVEvent_ = pipe_->AllocEventID(); + mte2ToMte3Event_ = pipe_->AllocEventID(); + mte3ToMte2Event_ = pipe_->AllocEventID(); + sToVEvent_ = pipe_->AllocEventID(); + sToMte2Event_ = pipe_->AllocEventID(); + vectorEventsAllocated_ = true; + } + + __aicore__ inline void ReleaseVectorEvents() + { + if (!vectorEventsAllocated_) { + return; + } + pipe_->ReleaseEventID(mte2ToVEvent_); + pipe_->ReleaseEventID(vToMte2Event_); + pipe_->ReleaseEventID(vToMte3Event_); + pipe_->ReleaseEventID(mte3ToVEvent_); + pipe_->ReleaseEventID(mte2ToMte3Event_); + pipe_->ReleaseEventID(mte3ToMte2Event_); + pipe_->ReleaseEventID(sToVEvent_); + pipe_->ReleaseEventID(sToMte2Event_); + vectorEventsAllocated_ = false; + } + + __aicore__ inline uint64_t QOffset(uint64_t b, uint64_t h, uint64_t t, uint64_t d) const + { + return ((b * H_ + h) * T_ + t) * K_ + d; + } + + __aicore__ inline uint64_t KVOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t d, uint64_t dim) const + { + return ((b * HV_ + hv) * T_ + t) * dim + d; + } + + __aicore__ inline uint64_t OutputOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t d) const + { + return ((b * T_ + t) * HV_ + hv) * V_ + d; + } + + __aicore__ inline uint64_t BetaOffset(uint64_t b, uint64_t hv, uint64_t t) const + { + return (b * HV_ + hv) * T_ + t; + } + + __aicore__ inline uint64_t AOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t j) const + { + return ((b * HV_ + hv) * T_ + t) * BT_ + j; + } + + __aicore__ inline uint64_t HOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t d, uint64_t r) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * K_ + d) * V_ + r; + } + + __aicore__ inline uint64_t WScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t t, uint64_t d) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * BT_ + t) * K_ + d; + } + + __aicore__ inline uint64_t SolveScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t slot) const + { + (void)b; + (void)hv; + (void)chunkIdx; + uint64_t matrixElements = BT_ * BT_; + return solveCoreIdx_ * KDA_SOLVE_SCRATCH_SLOTS * matrixElements + slot * matrixElements; + } + + __aicore__ inline uint64_t ScoreScratchOffset(uint64_t slot, uint64_t plane, uint64_t t = 0, + uint64_t d = 0) const + { + return (((solveCoreIdx_ * KDA_SCORE_QUEUE_DEPTH + slot) * KDA_SCORE_SCRATCH_PLANES + plane) * BT_ + t) * + K_ + + d; + } + + + + __aicore__ inline uint64_t ScoreRefBlockSize() const + { + return KDA_SCORE_REF_BC; + } + + __aicore__ inline uint64_t ScoreRowBlockCount(uint64_t curT, uint64_t rowBegin) const + { + uint64_t blockSize = ScoreRefBlockSize(); + uint64_t rowCount = curT - rowBegin; + if (rowCount > blockSize) { + rowCount = blockSize; + } + return rowCount; + } + + __aicore__ inline uint64_t ScoreRefToken(uint64_t start, uint64_t curT, uint64_t rowBegin, + uint64_t rowCount) const + { + uint64_t ref = rowBegin + rowCount / 2; + if (ref >= curT) { + ref = curT - 1; + } + return start + ref; + } + + __aicore__ inline void RunExp2(LocalTensor &tensor, uint32_t count) + { + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + ClampExpInput(tensor, count); + Exp(tensor, tensor, count); + PipeBarrier(); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + } + + __aicore__ inline void ClampExpInput(LocalTensor &tensor, uint32_t count) + { + Mins(tensor, tensor, KDA_EXP_INPUT_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, KDA_EXP_INPUT_MIN, count); + PipeBarrier(); + } + + __aicore__ inline void ClampFp32ToOutputType(LocalTensor &tensor, uint32_t count) + { + if constexpr (IsSameType::value) { + Mins(tensor, tensor, KDA_FP16_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, -KDA_FP16_MAX, count); + PipeBarrier(); + } + } + + template + __aicore__ inline void CopyVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst, src[offset], static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + + template + __aicore__ inline void CopyVectorOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst[offset], src, static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPad(dst[offset], src, params); + } + + template + __aicore__ inline void CopyRowsOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src, + uint64_t rows, uint64_t cols, uint64_t dstStride) + { + if (cols == dstStride) { + CopyVectorOut(dst, offset, src, rows * cols); + return; + } + constexpr uint64_t blockBytes = 32; + const uint64_t rowBytes = cols * sizeof(CopyT); + const uint64_t gapBytes = (dstStride - cols) * sizeof(CopyT); + DataCopyParams params{ +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + 1, +#else + static_cast(rows), +#endif + static_cast(rowBytes / blockBytes), + 0, +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + 0 +#else + static_cast(gapBytes / blockBytes) +#endif + }; +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + const uint64_t dstRowBytes = dstStride * sizeof(CopyT); + LoopModeParams loopParams{ + static_cast(rows), 1, rowBytes, dstRowBytes, 0, 0}; + // Loop-mode registers are core-local state and must not leak across DMA calls. + ResetLoopModePara(DataCopyMVType::UB_TO_OUT); + SetLoopModePara(loopParams, DataCopyMVType::UB_TO_OUT); + DataCopy(dst[offset], src, params); + ResetLoopModePara(DataCopyMVType::UB_TO_OUT); +#else + DataCopy(dst[offset], src, params); +#endif + } + + template + __aicore__ inline void CopyRowIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset) + { + CopyVectorIn(dst, src, offset, K_); + } + + template + __aicore__ inline void CopyRowOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src) + { + CopyVectorOut(dst, offset, src, K_); + } + + __aicore__ inline LocalTensor VecScratch(uint64_t slot) + { + return vecBuf_.Get()[slot * EXP2_UB_ELEMENTS]; + } + + template + __aicore__ inline void LoadAsFloatRow(GlobalTensor &src, uint64_t srcOffset, LocalTensor &dst, + uint64_t count) + { + if constexpr (IsSameType::value) { + CopyVectorIn(dst, src, srcOffset, count); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Adds(dst, dst, 0.0f, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } else { + constexpr uint32_t typedOffset = EXP2_UB_ELEMENTS * sizeof(float) / sizeof(CopyT); + LocalTensor rowLocal = exp2Buf_.Get()[typedOffset]; + CopyVectorIn(rowLocal, src, srcOffset, count); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast(dst, rowLocal, RoundMode::CAST_NONE, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } + PipeBarrier(); + } + + __aicore__ inline void ComputeTailLocalRows(LocalTensor &dst, uint64_t b, uint64_t hv, + uint64_t start, uint64_t curT, uint64_t rowBegin, + uint64_t rows) + { + LocalTensor vRow = exp2Buf_.Get(); + LocalTensor coefficientTyped = gateWritebackBuf_.Get(); + LocalTensor coefficients = gateWritebackBuf_.Get()[BT_]; + for (uint64_t localRow = 0; localRow < rows; ++localRow) { + LocalTensor dstRow = dst[localRow * V_]; + CopyVectorIn( + coefficientTyped, preparedAqk_, + AOffset(b, hv, start + rowBegin + localRow, 0), curT); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast( + coefficients, coefficientTyped, RoundMode::CAST_NONE, + static_cast(curT)); + // Earlier megakernel stages reuse V_S event IDs. Drain the cast + // before the first scalar coefficient read in this tail row. + PipeBarrier(); + Duplicate(dstRow, 0.0f, static_cast(V_)); + PipeBarrier(); + for (uint64_t j = 0; j < curT; ++j) { + LoadAsFloatRow( + propagatedVNew_, KVOffset(b, hv, start + j, 0, V_), vRow, V_); + float weight = coefficients.GetValue(j); + SetFlag(sToVEvent_); + WaitFlag(sToVEvent_); + Muls(vRow, vRow, weight, static_cast(V_)); + PipeBarrier(); + Add(dstRow, dstRow, vRow, static_cast(V_)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } + SetFlag(sToMte2Event_); + WaitFlag(sToMte2Event_); + } + } + + __aicore__ inline void ComputeTailStateRows(LocalTensor &dst, uint64_t b, uint64_t hv, + uint64_t chunkIdx, uint64_t start, uint64_t rowBegin, + uint64_t rows) + { + LocalTensor hRow = exp2Buf_.Get(); + LocalTensor coefficientTyped = gateWritebackBuf_.Get(); + LocalTensor coefficients = gateWritebackBuf_.Get()[BT_]; + for (uint64_t localRow = 0; localRow < rows; ++localRow) { + LocalTensor dstRow = dst[localRow * V_]; + CopyVectorIn( + coefficientTyped, preparedQG_, + KVOffset(b, hv, start + rowBegin + localRow, 0, K_), K_); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast( + coefficients, coefficientTyped, RoundMode::CAST_NONE, + static_cast(K_)); + // Earlier megakernel stages reuse V_S event IDs. Drain the cast + // before the first scalar coefficient read in this tail row. + PipeBarrier(); + Duplicate(dstRow, 0.0f, static_cast(V_)); + PipeBarrier(); + for (uint64_t d = 0; d < K_; ++d) { + LoadAsFloatRow( + propagatedH_, HOffset(b, hv, chunkIdx, d, 0), hRow, V_); + float weight = coefficients.GetValue(d); + SetFlag(sToVEvent_); + WaitFlag(sToVEvent_); + Muls(hRow, hRow, weight, static_cast(V_)); + PipeBarrier(); + Add(dstRow, dstRow, hRow, static_cast(V_)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } + SetFlag(sToMte2Event_); + WaitFlag(sToMte2Event_); + } + } + + template + __aicore__ inline void LoadAsFloatVector(GlobalTensor &src, uint64_t srcOffset, + LocalTensor &dst, LocalTensor &typedScratch, + uint64_t count) + { + if constexpr (IsSameType::value) { + CopyVectorIn(dst, src, srcOffset, count); + } else { + CopyVectorIn(typedScratch, src, srcOffset, count); + } + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + if constexpr (!IsSameType::value) { + Cast(dst, typedScratch, RoundMode::CAST_NONE, static_cast(count)); + PipeBarrier(); + } + } + + template + __aicore__ inline void StoreFloatRow(GlobalTensor &dst, uint64_t dstOffset, LocalTensor &src, + uint64_t count) + { + if constexpr (IsSameType::value) { + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(dst, dstOffset, src, count); + } else { + constexpr uint32_t typedOffset = EXP2_UB_ELEMENTS * sizeof(float) / sizeof(CopyT); + LocalTensor rowLocal = exp2Buf_.Get()[typedOffset]; + Cast(rowLocal, src, RoundMode::CAST_RINT, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(dst, dstOffset, rowLocal, count); + } + SetFlag(mte3ToMte2Event_); + WaitFlag(mte3ToMte2Event_); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + + + + + __aicore__ inline LocalTensor Exp2NegG(uint64_t b, uint64_t hv, uint64_t t) + { + LocalTensor exp2Local = exp2Buf_.Get(); + LoadAsFloatRow(gk_, KVOffset(b, hv, t, 0, K_), exp2Local, K_); + Muls(exp2Local, exp2Local, -LN2, static_cast(K_)); + PipeBarrier(); + RunExp2(exp2Local, static_cast(K_)); + return exp2Local; + } + + + __aicore__ inline uint64_t ScoreVectorMaxRows(uint64_t bytesPerElem) const + { + constexpr uint64_t arenaBytes = static_cast(KDA_VEC_ARENA_ELEMENTS) * sizeof(float); + uint64_t maxRows = (arenaBytes / bytesPerElem) / K_; + if (K_ >= 128 && maxRows > 32) { + maxRows = 32; + } + return maxRows; + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + __aicore__ inline void ComputeOutputCubeStagedArch35(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT) + { + SetMMLayoutTransform(true); + using ElementA = T; + using ElementB = T; + using ElementC = OUT_T; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + using LayoutTagL0A = typename TileCopy::LayoutTagL0A; + using LayoutTagL0B = typename TileCopy::LayoutTagL0B; + using CopyL1ToL0A = typename TileCopy::CopyL1ToL0A; + using CopyL1ToL0B = typename TileCopy::CopyL1ToL0B; + using TileMmad = Catlass::Gemm::Tile::TileMmadTla; + + constexpr uint16_t kMte2Event = 0; + constexpr uint16_t kMte1Event = 0; + constexpr uint16_t kMmadEvent = 0; + constexpr uint32_t kL1A0Offset = 0; + constexpr uint32_t kL1B0Offset = 64 * 128 * sizeof(ElementA); + constexpr uint32_t kL1A1Offset = kL1B0Offset + 128 * 128 * sizeof(ElementB); + constexpr uint32_t kL1B1Offset = kL1A1Offset + 64 * 64 * sizeof(ElementA); + + Catlass::Arch::Resource resource; + LocalTensor l1A0 = resource.l1Buf.template GetBufferByByte(kL1A0Offset); + LocalTensor l1B0 = resource.l1Buf.template GetBufferByByte(kL1B0Offset); + LocalTensor l1A1 = resource.l1Buf.template GetBufferByByte(kL1A1Offset); + LocalTensor l1B1 = resource.l1Buf.template GetBufferByByte(kL1B1Offset); + LocalTensor l0A = resource.l0ABuf.template GetBufferByByte(0); + LocalTensor l0B = resource.l0BBuf.template GetBufferByByte(0); + LocalTensor l0C = resource.l0CBuf.template GetBufferByByte(0); + + CopyL1ToL0A copyL1ToL0A; + CopyL1ToL0B copyL1ToL0B; + TileMmad tileMmad; + + auto layoutQ = tla::MakeLayout(BT_, K_); + auto layoutH = tla::MakeLayout(K_, V_); + auto layoutO = tla::MakeLayout(BT_, V_); + auto layoutAqk = tla::MakeLayout(BT_, BT_); + auto layoutV = tla::MakeLayout(BT_, V_); + + for (uint64_t nOffset = 0; nOffset < V_; nOffset += 128) { + const uint32_t curN = static_cast((V_ - nOffset) > 128 ? 128 : (V_ - nOffset)); + auto tensorH = tla::MakeTensor(propagatedH_[HOffset(b, hv, chunkIdx, 0, nOffset)], layoutH, + Catlass::Arch::PositionGM{}); + auto tensorVNew = tla::MakeTensor(propagatedVNew_[KVOffset(b, hv, start, nOffset, V_)], layoutV, + Catlass::Arch::PositionGM{}); + + for (uint64_t mOffset = 0; mOffset < curT; mOffset += 64) { + const uint32_t curM = static_cast((curT - mOffset) > 64 ? 64 : (curT - mOffset)); + auto tensorQ = tla::MakeTensor(preparedQG_[KVOffset(b, hv, start + mOffset, 0, K_)], layoutQ, + Catlass::Arch::PositionGM{}); + auto tensorAqk = tla::MakeTensor(preparedAqk_[AOffset(b, hv, start + mOffset, 0)], layoutAqk, + Catlass::Arch::PositionGM{}); + auto tensorO = tla::MakeTensor(o_[KVOffset(b, hv, start + mOffset, nOffset, V_)], layoutO, + Catlass::Arch::PositionGM{}); + auto tensorLocal = tla::MakeTensor(u_[KVOffset(b, hv, start + mOffset, nOffset, V_)], layoutO, + Catlass::Arch::PositionGM{}); + + auto blockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(curM, K_)); + auto blockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(K_, curN)); + auto blockAqk = GetTile(tensorAqk, tla::MakeCoord(0, 0), tla::MakeShape(curM, curT)); + auto blockVNew = GetTile(tensorVNew, tla::MakeCoord(0, 0), tla::MakeShape(curT, curN)); + auto blockO = GetTile(tensorO, tla::MakeCoord(0, 0), tla::MakeShape(curM, curN)); + auto blockLocal = + GetTile(tensorLocal, tla::MakeCoord(0, 0), tla::MakeShape(curM, curN)); + + using CopyGmToL1A0 = typename TileCopy::template CopyGmToL1A; + using CopyGmToL1B0 = typename TileCopy::template CopyGmToL1B; + using CopyGmToL1A1 = typename TileCopy::template CopyGmToL1A; + using CopyGmToL1B1 = typename TileCopy::template CopyGmToL1B; + using CopyL0CToDst = typename TileCopy::template CopyL0CToDst; + CopyGmToL1A0 copyGmToL1A0; + CopyGmToL1B0 copyGmToL1B0; + CopyGmToL1A1 copyGmToL1A1; + CopyGmToL1B1 copyGmToL1B1; + CopyL0CToDst copyL0CToDst; + + auto layoutL1A0 = tla::MakeLayout(curM, K_); + auto layoutL1B0 = tla::MakeLayout(K_, curN); + auto layoutL1A1 = tla::MakeLayout(curM, curT); + auto layoutL1B1 = tla::MakeLayout(curT, curN); + auto layoutL0A0 = tla::MakeLayout(curM, K_); + auto layoutL0B0 = tla::MakeLayout(K_, curN); + auto layoutL0A1 = tla::MakeLayout(curM, curT); + auto layoutL0B1 = tla::MakeLayout(curT, curN); + auto layoutL0C = tla::MakeLayoutL0C(curM, curN); + + auto tensorL1A0 = tla::MakeTensor(l1A0, layoutL1A0, Catlass::Arch::PositionL1{}); + auto tensorL1B0 = tla::MakeTensor(l1B0, layoutL1B0, Catlass::Arch::PositionL1{}); + auto tensorL1A1 = tla::MakeTensor(l1A1, layoutL1A1, Catlass::Arch::PositionL1{}); + auto tensorL1B1 = tla::MakeTensor(l1B1, layoutL1B1, Catlass::Arch::PositionL1{}); + auto tensorL0A0 = tla::MakeTensor(l0A, layoutL0A0, Catlass::Arch::PositionL0A{}); + auto tensorL0B0 = tla::MakeTensor(l0B, layoutL0B0, Catlass::Arch::PositionL0B{}); + auto tensorL0A1 = tla::MakeTensor(l0A, layoutL0A1, Catlass::Arch::PositionL0A{}); + auto tensorL0B1 = tla::MakeTensor(l0B, layoutL0B1, Catlass::Arch::PositionL0B{}); + auto tensorL0C = tla::MakeTensor(l0C, layoutL0C, Catlass::Arch::PositionL0C{}); + uint32_t localRow = 0; + uint32_t localColumn = 0; + auto tileL1A0 = GetTile(tensorL1A0, tla::MakeCoord(localRow, localColumn), + tla::MakeShape(curM, K_)); + auto tileL1B0 = GetTile(tensorL1B0, tla::MakeCoord(localRow, localColumn), + tla::MakeShape(K_, curN)); + auto tileL1A1 = GetTile(tensorL1A1, tla::MakeCoord(localRow, localColumn), + tla::MakeShape(curM, curT)); + auto tileL1B1 = GetTile(tensorL1B1, tla::MakeCoord(localRow, localColumn), + tla::MakeShape(curT, curN)); + auto tileL0A0 = GetTile(tensorL0A0, tla::MakeCoord(localRow, localColumn), + tla::MakeShape(curM, K_)); + auto tileL0B0 = GetTile(tensorL0B0, tla::MakeCoord(localRow, localColumn), + tla::MakeShape(K_, curN)); + auto tileL0A1 = GetTile(tensorL0A1, tla::MakeCoord(localRow, localColumn), + tla::MakeShape(curM, curT)); + auto tileL0B1 = GetTile(tensorL0B1, tla::MakeCoord(localRow, localColumn), + tla::MakeShape(curT, curN)); + auto tileL0C = GetTile(tensorL0C, tla::MakeCoord(localRow, localColumn), + tla::MakeShape(curM, curN)); + + copyGmToL1A0(tensorL1A0, blockQ); + copyGmToL1B0(tensorL1B0, blockH); + copyGmToL1A1(tensorL1A1, blockAqk); + copyGmToL1B1(tensorL1B1, blockVNew); + SetFlag(kMte2Event); + WaitFlag(kMte2Event); + + copyL1ToL0A(tileL0A0, tileL1A0); + copyL1ToL0B(tileL0B0, tileL1B0); + SetFlag(kMte1Event); + WaitFlag(kMte1Event); + tileMmad(tileL0C, tileL0A0, tileL0B0, curM, curN, static_cast(K_), true, 0b11); + SetFlag(kMmadEvent); + WaitFlag(kMmadEvent); + copyL0CToDst(blockO, tileL0C, 0b11); + PipeBarrier(); + + copyL1ToL0A(tileL0A1, tileL1A1); + copyL1ToL0B(tileL0B1, tileL1B1); + SetFlag(kMte1Event); + SetFlag(kMte2Event); + WaitFlag(kMte1Event); + WaitFlag(kMte2Event); + tileMmad(tileL0C, tileL0A1, tileL0B1, curM, curN, static_cast(curT), true, 0b11); + SetFlag(kMmadEvent); + WaitFlag(kMmadEvent); + copyL0CToDst(blockLocal, tileL0C, 0b11); + PipeBarrier(); + } + } + SetMMLayoutTransform(false); + } + + __aicore__ inline void PrefetchOutputTileArch35(Catlass::Arch::Resource &resource, uint32_t slot, + uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT, uint64_t nOffset, bool reuseSlot) + { + using ElementA = T; + using ElementB = T; + using ElementC = OUT_T; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + + constexpr uint32_t kL1SlotBytes = 96 * 1024; + constexpr uint32_t kL1A0Offset = 0; + constexpr uint32_t kL1B0Offset = 64 * 128 * sizeof(ElementA); + constexpr uint32_t kL1A1Offset = kL1B0Offset + 128 * 128 * sizeof(ElementB); + constexpr uint32_t kL1B1Offset = kL1A1Offset + 64 * 64 * sizeof(ElementA); + const uint32_t slotBase = slot * kL1SlotBytes; + const uint32_t curM = static_cast(curT); + const uint32_t curN = static_cast((V_ - nOffset) > 128 ? 128 : (V_ - nOffset)); + + auto layoutQ = tla::MakeLayout(BT_, K_); + auto layoutH = tla::MakeLayout(K_, V_); + auto layoutAqk = tla::MakeLayout(BT_, BT_); + auto layoutV = tla::MakeLayout(BT_, V_); + auto tensorQ = tla::MakeTensor(preparedQG_[KVOffset(b, hv, start, 0, K_)], layoutQ, + Catlass::Arch::PositionGM{}); + auto tensorH = tla::MakeTensor(propagatedH_[HOffset(b, hv, chunkIdx, 0, nOffset)], layoutH, + Catlass::Arch::PositionGM{}); + auto tensorAqk = tla::MakeTensor(preparedAqk_[AOffset(b, hv, start, 0)], layoutAqk, + Catlass::Arch::PositionGM{}); + auto tensorVNew = tla::MakeTensor(propagatedVNew_[KVOffset(b, hv, start, nOffset, V_)], layoutV, + Catlass::Arch::PositionGM{}); + auto blockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(curM, K_)); + auto blockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(K_, curN)); + auto blockAqk = GetTile(tensorAqk, tla::MakeCoord(0, 0), tla::MakeShape(curM, curT)); + auto blockVNew = GetTile(tensorVNew, tla::MakeCoord(0, 0), tla::MakeShape(curT, curN)); + + using CopyGmToL1A0 = typename TileCopy::template CopyGmToL1A; + using CopyGmToL1B0 = typename TileCopy::template CopyGmToL1B; + using CopyGmToL1A1 = typename TileCopy::template CopyGmToL1A; + using CopyGmToL1B1 = typename TileCopy::template CopyGmToL1B; + CopyGmToL1A0 copyGmToL1A0; + CopyGmToL1B0 copyGmToL1B0; + CopyGmToL1A1 copyGmToL1A1; + CopyGmToL1B1 copyGmToL1B1; + + LocalTensor l1A0 = + resource.l1Buf.template GetBufferByByte(slotBase + kL1A0Offset); + LocalTensor l1B0 = + resource.l1Buf.template GetBufferByByte(slotBase + kL1B0Offset); + LocalTensor l1A1 = + resource.l1Buf.template GetBufferByByte(slotBase + kL1A1Offset); + LocalTensor l1B1 = + resource.l1Buf.template GetBufferByByte(slotBase + kL1B1Offset); + auto tensorL1A0 = tla::MakeTensor( + l1A0, tla::MakeLayout(curM, K_), Catlass::Arch::PositionL1{}); + auto tensorL1B0 = tla::MakeTensor( + l1B0, tla::MakeLayout(K_, curN), Catlass::Arch::PositionL1{}); + auto tensorL1A1 = tla::MakeTensor( + l1A1, tla::MakeLayout(curM, curT), Catlass::Arch::PositionL1{}); + auto tensorL1B1 = tla::MakeTensor( + l1B1, tla::MakeLayout(curT, curN), Catlass::Arch::PositionL1{}); + + if (reuseSlot) { + WaitFlag(slot); + } + copyGmToL1A0(tensorL1A0, blockQ); + copyGmToL1B0(tensorL1B0, blockH); + copyGmToL1A1(tensorL1A1, blockAqk); + copyGmToL1B1(tensorL1B1, blockVNew); + SetFlag(slot); + } + + __aicore__ inline void ComputePrefetchedOutputTileArch35(Catlass::Arch::Resource &resource, + uint32_t slot, uint64_t b, uint64_t hv, uint64_t start, + uint64_t curT, uint64_t nOffset) + { + using ElementA = T; + using ElementB = T; + using ElementC = OUT_T; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + using LayoutTagL0A = typename TileCopy::LayoutTagL0A; + using LayoutTagL0B = typename TileCopy::LayoutTagL0B; + using CopyL1ToL0A = typename TileCopy::CopyL1ToL0A; + using CopyL1ToL0B = typename TileCopy::CopyL1ToL0B; + using TileMmad = Catlass::Gemm::Tile::TileMmadTla; + + constexpr uint16_t kMte1Event = 0; + constexpr uint16_t kMmadEvent = 0; + constexpr uint16_t kFixEvent = 0; + constexpr uint32_t kL1SlotBytes = 96 * 1024; + constexpr uint32_t kL1A0Offset = 0; + constexpr uint32_t kL1B0Offset = 64 * 128 * sizeof(ElementA); + constexpr uint32_t kL1A1Offset = kL1B0Offset + 128 * 128 * sizeof(ElementB); + constexpr uint32_t kL1B1Offset = kL1A1Offset + 64 * 64 * sizeof(ElementA); + const uint32_t slotBase = slot * kL1SlotBytes; + const uint32_t curM = static_cast(curT); + const uint32_t curN = static_cast((V_ - nOffset) > 128 ? 128 : (V_ - nOffset)); + + LocalTensor l1A0 = + resource.l1Buf.template GetBufferByByte(slotBase + kL1A0Offset); + LocalTensor l1B0 = + resource.l1Buf.template GetBufferByByte(slotBase + kL1B0Offset); + LocalTensor l1A1 = + resource.l1Buf.template GetBufferByByte(slotBase + kL1A1Offset); + LocalTensor l1B1 = + resource.l1Buf.template GetBufferByByte(slotBase + kL1B1Offset); + LocalTensor l0A = resource.l0ABuf.template GetBufferByByte(0); + LocalTensor l0B = resource.l0BBuf.template GetBufferByByte(0); + LocalTensor l0C = resource.l0CBuf.template GetBufferByByte(0); + + auto tensorL1A0 = tla::MakeTensor( + l1A0, tla::MakeLayout(curM, K_), Catlass::Arch::PositionL1{}); + auto tensorL1B0 = tla::MakeTensor( + l1B0, tla::MakeLayout(K_, curN), Catlass::Arch::PositionL1{}); + auto tensorL1A1 = tla::MakeTensor( + l1A1, tla::MakeLayout(curM, curT), Catlass::Arch::PositionL1{}); + auto tensorL1B1 = tla::MakeTensor( + l1B1, tla::MakeLayout(curT, curN), Catlass::Arch::PositionL1{}); + auto tensorL0A0 = tla::MakeTensor( + l0A, tla::MakeLayout(curM, K_), Catlass::Arch::PositionL0A{}); + auto tensorL0B0 = tla::MakeTensor( + l0B, tla::MakeLayout(K_, curN), Catlass::Arch::PositionL0B{}); + auto tensorL0A1 = tla::MakeTensor( + l0A, tla::MakeLayout(curM, curT), Catlass::Arch::PositionL0A{}); + auto tensorL0B1 = tla::MakeTensor( + l0B, tla::MakeLayout(curT, curN), Catlass::Arch::PositionL0B{}); + auto tensorL0C = tla::MakeTensor(l0C, tla::MakeLayoutL0C(curM, curN), Catlass::Arch::PositionL0C{}); + + uint32_t localRow = 0; + uint32_t localColumn = 0; + auto tileL1A0 = GetTile(tensorL1A0, tla::MakeCoord(localRow, localColumn), tla::MakeShape(curM, K_)); + auto tileL1B0 = GetTile(tensorL1B0, tla::MakeCoord(localRow, localColumn), tla::MakeShape(K_, curN)); + auto tileL1A1 = GetTile(tensorL1A1, tla::MakeCoord(localRow, localColumn), tla::MakeShape(curM, curT)); + auto tileL1B1 = GetTile(tensorL1B1, tla::MakeCoord(localRow, localColumn), tla::MakeShape(curT, curN)); + auto tileL0A0 = GetTile(tensorL0A0, tla::MakeCoord(localRow, localColumn), tla::MakeShape(curM, K_)); + auto tileL0B0 = GetTile(tensorL0B0, tla::MakeCoord(localRow, localColumn), tla::MakeShape(K_, curN)); + auto tileL0A1 = GetTile(tensorL0A1, tla::MakeCoord(localRow, localColumn), tla::MakeShape(curM, curT)); + auto tileL0B1 = GetTile(tensorL0B1, tla::MakeCoord(localRow, localColumn), tla::MakeShape(curT, curN)); + auto tileL0C = GetTile(tensorL0C, tla::MakeCoord(localRow, localColumn), tla::MakeShape(curM, curN)); + + auto layoutO = tla::MakeLayout(BT_, V_); + auto tensorO = tla::MakeTensor(o_[KVOffset(b, hv, start, nOffset, V_)], layoutO, + Catlass::Arch::PositionGM{}); + auto tensorLocal = tla::MakeTensor(u_[KVOffset(b, hv, start, nOffset, V_)], layoutO, + Catlass::Arch::PositionGM{}); + auto blockO = GetTile(tensorO, tla::MakeCoord(0, 0), tla::MakeShape(curM, curN)); + auto blockLocal = GetTile(tensorLocal, tla::MakeCoord(0, 0), tla::MakeShape(curM, curN)); + using CopyL0CToDst = typename TileCopy::template CopyL0CToDst; + CopyL1ToL0A copyL1ToL0A; + CopyL1ToL0B copyL1ToL0B; + CopyL0CToDst copyL0CToDst; + TileMmad tileMmad; + + WaitFlag(slot); + if constexpr (IsSameType::value) { + WaitFlag(kFixEvent); + } + copyL1ToL0A(tileL0A0, tileL1A0); + copyL1ToL0B(tileL0B0, tileL1B0); + SetFlag(kMte1Event); + WaitFlag(kMte1Event); + tileMmad(tileL0C, tileL0A0, tileL0B0, curM, curN, static_cast(K_), true, 0b11); + SetFlag(kMmadEvent); + WaitFlag(kMmadEvent); + if constexpr (!IsSameType::value) { + copyL0CToDst(blockO, tileL0C, 0b11); + PipeBarrier(); + } + + copyL1ToL0A(tileL0A1, tileL1A1); + copyL1ToL0B(tileL0B1, tileL1B1); + SetFlag(kMte1Event); + SetFlag(slot); + WaitFlag(kMte1Event); + tileMmad(tileL0C, tileL0A1, tileL0B1, curM, curN, static_cast(curT), + !IsSameType::value, 0b11); + if constexpr (IsSameType::value) { + SetFlag(kFixEvent); + WaitFlag(kFixEvent); + auto fixParams = FixpipeParamsV220( + curN, curM, curN, static_cast(HV_ * V_), false); + fixParams.quantPre = QuantMode_t::F322BF16; + Fixpipe( + vNew_[OutputOffset(b, hv, start, nOffset)], l0C, fixParams); + SetFlag(kFixEvent); + } else { + SetFlag(kMmadEvent); + WaitFlag(kMmadEvent); + copyL0CToDst(blockLocal, tileL0C, 0b11); + PipeBarrier(); + } + } + + __aicore__ inline void ProcessOutAicPipelinedArch35() + { + SetLoadDataPaddingValue(static_cast(0)); + SetMMLayoutTransform(true); + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + uint64_t currentTask = static_cast(GetBlockIdx()); + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + while (currentTask < taskNum && + !ResolveFlatChunk(currentTask, seq, b, h, hv, chunkIdx, start, end)) { + currentTask += coreNum; + } + if (currentTask >= taskNum) { + SetMMLayoutTransform(false); + return; + } + + Catlass::Arch::Resource resource; + uint64_t nOffset = 0; + uint32_t slot = 0; + if constexpr (IsSameType::value) { + SetFlag(0); + } + PrefetchOutputTileArch35(resource, slot, b, hv, chunkIdx, start, end - start, nOffset, false); + uint64_t outputTileIdx = 0; + + while (true) { + uint64_t nextTask = currentTask; + uint64_t nextSeq = seq; + uint64_t nextB = b; + uint64_t nextH = h; + uint64_t nextHv = hv; + uint64_t nextChunkIdx = chunkIdx; + uint64_t nextStart = start; + uint64_t nextEnd = end; + uint64_t nextNOffset = nOffset + 128; + bool hasNext = nextNOffset < V_; + if (!hasNext) { + nextTask += coreNum; + nextNOffset = 0; + while (nextTask < taskNum && + !ResolveFlatChunk(nextTask, nextSeq, nextB, nextH, nextHv, nextChunkIdx, nextStart, + nextEnd)) { + nextTask += coreNum; + } + hasNext = nextTask < taskNum; + } + + const uint32_t nextSlot = slot ^ 1U; + if (hasNext) { + PrefetchOutputTileArch35(resource, nextSlot, nextB, nextHv, nextChunkIdx, nextStart, + nextEnd - nextStart, nextNOffset, + outputTileIdx + 1 >= 2); + } + ComputePrefetchedOutputTileArch35(resource, slot, b, hv, start, end - start, nOffset); + if constexpr (!IsSameType::value) { + if (nOffset + 128 >= V_) { + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); + } + } + if (!hasNext) { + break; + } + ++outputTileIdx; + + currentTask = nextTask; + seq = nextSeq; + b = nextB; + h = nextH; + hv = nextHv; + chunkIdx = nextChunkIdx; + start = nextStart; + end = nextEnd; + nOffset = nextNOffset; + slot = nextSlot; + } + if constexpr (IsSameType::value) { + WaitFlag(0); + } + SetMMLayoutTransform(false); + } +#endif + + __aicore__ inline void ComputeOutputCube(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT) + { +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if (curT < KDA_CUBE_MIN_REDUCTION) { + return; + } + SetLoadDataPaddingValue(static_cast(0)); + if (BT_ == 64 && curT == BT_) { + ComputeOutputCubeStagedArch35(b, hv, chunkIdx, start, curT); + return; + } +#endif + using ElementA = T; + using ElementB = T; + using ElementC = OUT_T; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using BlockMmad = Catlass::Gemm::Block::BlockMmadTla; + + Catlass::Arch::Resource resource; + BlockMmad blockMmad(resource); + + auto layoutQ = tla::MakeLayout(BT_, K_); + auto layoutH = tla::MakeLayout(K_, V_); + auto layoutO = tla::MakeLayout(BT_, V_); + for (uint64_t nOffset = 0; nOffset < V_; nOffset += 128) { + uint32_t curN = static_cast((V_ - nOffset) > 128 ? 128 : (V_ - nOffset)); + auto tensorH = tla::MakeTensor(propagatedH_[HOffset(b, hv, chunkIdx, 0, nOffset)], layoutH, + Catlass::Arch::PositionGM{}); + for (uint64_t mOffset = 0; mOffset < curT; mOffset += 64) { + uint32_t curM = static_cast((curT - mOffset) > 64 ? 64 : (curT - mOffset)); + Catlass::GemmCoord shapeQH{curM, curN, static_cast(K_)}; + auto tensorQ = tla::MakeTensor(preparedQG_[KVOffset(b, hv, start + mOffset, 0, K_)], layoutQ, + Catlass::Arch::PositionGM{}); + auto tensorO = tla::MakeTensor(o_[KVOffset(b, hv, start + mOffset, nOffset, V_)], layoutO, + Catlass::Arch::PositionGM{}); + auto blockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(shapeQH.m(), shapeQH.k())); + auto blockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(shapeQH.k(), shapeQH.n())); + auto blockO = GetTile(tensorO, tla::MakeCoord(0, 0), tla::MakeShape(shapeQH.m(), shapeQH.n())); + blockMmad(blockQ, blockH, blockO, shapeQH); + PipeBarrier(); + } + } + + if (curT < KDA_CUBE_MIN_REDUCTION) { + return; + } + + auto layoutAqk = tla::MakeLayout(BT_, BT_); + auto layoutV = tla::MakeLayout(BT_, V_); + for (uint64_t nOffset = 0; nOffset < V_; nOffset += 128) { + uint32_t curN = static_cast((V_ - nOffset) > 128 ? 128 : (V_ - nOffset)); + auto tensorVNew = tla::MakeTensor(propagatedVNew_[KVOffset(b, hv, start, nOffset, V_)], layoutV, + Catlass::Arch::PositionGM{}); + for (uint64_t mOffset = 0; mOffset < curT; mOffset += 64) { + uint32_t curM = static_cast((curT - mOffset) > 64 ? 64 : (curT - mOffset)); + Catlass::GemmCoord shapeAV{curM, curN, static_cast(curT)}; + auto tensorAqk = tla::MakeTensor(preparedAqk_[AOffset(b, hv, start + mOffset, 0)], layoutAqk, + Catlass::Arch::PositionGM{}); + auto tensorLocal = tla::MakeTensor(u_[KVOffset(b, hv, start + mOffset, nOffset, V_)], layoutO, + Catlass::Arch::PositionGM{}); + auto blockAqk = GetTile(tensorAqk, tla::MakeCoord(0, 0), tla::MakeShape(shapeAV.m(), shapeAV.k())); + auto blockVNew = GetTile(tensorVNew, tla::MakeCoord(0, 0), tla::MakeShape(shapeAV.k(), shapeAV.n())); + auto blockLocal = GetTile(tensorLocal, tla::MakeCoord(0, 0), tla::MakeShape(shapeAV.m(), shapeAV.n())); + blockMmad(blockAqk, blockVNew, blockLocal, shapeAV); + PipeBarrier(); + } + } + } + + __aicore__ inline void FinalizeOutputRows(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT, uint64_t subBlockIdx, uint64_t subBlockNum) + { + if (subBlockNum == 0 || subBlockIdx >= subBlockNum || V_ == 0) { + return; + } + const uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; + const uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + const uint64_t gateWritebackRows = + ScoreVectorMaxRows(5 * sizeof(float) + 2 * sizeof(T) + sizeof(GK_T)); + const uint64_t gateWritebackBytes = + gateWritebackRows * K_ * (3 * sizeof(T) + sizeof(GK_T)); + uint64_t maxRows = KDA_VEC_ARENA_ELEMENTS / (3 * V_); + const uint64_t typedMaxRows = gateWritebackBytes / (V_ * sizeof(T)); + if (maxRows > typedMaxRows) { + maxRows = typedMaxRows; + } + if (maxRows == 0) { + return; + } + + for (uint64_t tileRow = rowBegin; tileRow < rowEnd; tileRow += maxRows) { + uint64_t tileRows = rowEnd - tileRow; + if (tileRows > maxRows) { + tileRows = maxRows; + } + const uint64_t elems = tileRows * V_; + const uint64_t ti = start + tileRow; + LocalTensor arena = vecBuf_.Get(); + LocalTensor stateLocal = arena; + LocalTensor localLocal = arena[elems]; + LocalTensor outLocal = arena[2 * elems]; + LocalTensor outTyped = gateWritebackBuf_.Get(); + + if (curT < KDA_CUBE_MIN_REDUCTION) { + ComputeTailStateRows( + stateLocal, b, hv, chunkIdx, start, tileRow, tileRows); + ComputeTailLocalRows(localLocal, b, hv, start, curT, tileRow, tileRows); + } else { + CopyVectorIn(stateLocal, o_, KVOffset(b, hv, ti, 0, V_), elems); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + CopyVectorIn(localLocal, u_, KVOffset(b, hv, ti, 0, V_), elems); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + } + Add(outLocal, stateLocal, localLocal, static_cast(elems)); + PipeBarrier(); + ClampFp32ToOutputType(outLocal, static_cast(elems)); + Cast(outTyped, outLocal, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyRowsOut(vNew_, OutputOffset(b, hv, ti, 0), outTyped, tileRows, V_, HV_ * V_); + SetFlag(mte3ToMte2Event_); + WaitFlag(mte3ToMte2Event_); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + } + __aicore__ inline bool ResolveFlatChunk(uint64_t task, uint64_t &seq, uint64_t &b, uint64_t &h, uint64_t &hv, + uint64_t &chunkIdx, uint64_t &start, uint64_t &end) + { + hv = task % HV_; + uint64_t flatChunk = task / HV_; + if (!isVarLen_) { + seq = flatChunk / NT_; + b = seq; + chunkIdx = flatChunk % NT_; + start = chunkIdx * BT_; + end = start + BT_; + if (end > T_) { + end = T_; + } + } else { + if (!KdaVarlen::ResolveChunkRange( + cuSeqlensAddr_, chunkIndicesAddr_, N_, T_, BT_, flatChunk, + seq, start, end)) { + return false; + } + b = 0; + chunkIdx = flatChunk; + } + h = hv / (HV_ / H_); + return start < end; + } + + __aicore__ inline void ProcessChunkOutAiv(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end, uint64_t subBlockIdx, uint64_t subBlockNum) + { + uint64_t curT = end - start; + if (curT == 0) { + return; + } + if constexpr (IsSameType::value) { + return; + } + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(syncDoneFlag_); + FinalizeOutputRows(b, hv, chunkIdx, start, curT, subBlockIdx, subBlockNum); + } + + __aicore__ inline void ProcessChunkOutAic(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + uint64_t curT = end - start; + if (curT == 0) { + return; + } + ComputeOutputCube(b, hv, chunkIdx, start, curT); + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); + } + + __aicore__ inline void ProcessOutAiv() + { + if constexpr (IsSameType::value) { + return; + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (IsSameType::value) { + if (!isVarLen_ && T_ % BT_ == 0 && BT_ == 64 && K_ == 128 && V_ == 128) { + return; + } + } +#endif + uint64_t subBlockNum = static_cast(GetSubBlockNum()); + if (subBlockNum == 0) { + return; + } + uint64_t subBlockIdx = static_cast(GetSubBlockIdx()); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + uint64_t coreIdx = static_cast(GetBlockIdx()) / subBlockNum; + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + for (uint64_t task = coreIdx; task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + (void)h; + (void)chunkIdx; + ProcessChunkOutAiv(b, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); + } + } + } + + __aicore__ inline void ProcessOutAic() + { + if constexpr (IsSameType::value) { + return; + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if (!isVarLen_ && T_ % BT_ == 0 && BT_ == 64 && K_ == 128 && V_ == 128) { + ProcessOutAicPipelinedArch35(); + return; + } +#endif + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + (void)h; + ProcessChunkOutAic(b, hv, chunkIdx, start, end); + } + } + } + +private: + GlobalTensor q_; + GlobalTensor k_; + GlobalTensor v_; + GlobalTensor gk_; + GlobalTensor beta_; + GlobalTensor initialState_; + GlobalTensor o_; + GlobalTensor finalState_; + GlobalTensor aqk_; + GlobalTensor akk_; + GlobalTensor w_; + GlobalTensor u_; + GlobalTensor qg_; + GlobalTensor kg_; + GlobalTensor vNew_; + GlobalTensor h_; + GlobalTensor preparedQG_; + GlobalTensor preparedAqk_; + GlobalTensor propagatedVNew_; + GlobalTensor propagatedH_; + GlobalTensor solveWorkspace_; + GlobalTensor scoreWorkspace_; + TPipe *pipe_ = nullptr; + TBuf exp2Buf_; + TBuf vecBuf_; + TBuf gateWritebackBuf_; + TEventID mte2ToVEvent_ = 0; + TEventID vToMte2Event_ = 0; + TEventID vToMte3Event_ = 0; + TEventID mte3ToVEvent_ = 0; + TEventID mte2ToMte3Event_ = 0; + TEventID mte3ToMte2Event_ = 0; + TEventID sToVEvent_ = 0; + TEventID sToMte2Event_ = 0; + bool vectorEventsAllocated_ = false; + Catlass::Arch::CrossCoreFlagWithReverse scoreReadyFlag_{KDA_SCORE_READY_FLAG0, + KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse scoreDoneFlag_{KDA_SCORE_DONE_FLAG0, + KDA_SCORE_DONE_FLAG1}; + // Score production is fully drained before solve starts, so the solve handshake can safely reuse + // the existing score flags without consuming additional hardware flag IDs. + Catlass::Arch::CrossCoreFlagWithReverse syncReadyFlag_{KDA_SCORE_READY_FLAG0, + KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse syncDoneFlag_{KDA_SCORE_DONE_FLAG0, + KDA_SCORE_DONE_FLAG1}; + uint64_t B_ = 0; + uint64_t N_ = 0; + uint64_t H_ = 0; + uint64_t HV_ = 0; + uint64_t T_ = 0; + uint64_t K_ = 0; + uint64_t V_ = 0; + uint64_t BT_ = 0; + uint64_t NT_ = 0; + float scale_ = 1.0f; + bool hasInitial_ = false; + bool isVarLen_ = false; + bool isAivOnly_ = false; + uint64_t usedCoreNum_ = 1; + uint64_t solveCoreIdx_ = 0; + __gm__ int64_t *chunkIndicesAddr_ = nullptr; + __gm__ int64_t *cuSeqlensAddr_ = nullptr; +}; +} // namespace + +template +__aicore__ inline void RunChunkKdaOutput( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR qgScaled, GM_ADDR aqk, + GM_ADDR propagatedVNew, GM_ADDR propagatedH, GM_ADDR o, GM_ADDR userWorkspace, + const TilingData &tiling, TPipe &pipe) +{ + GM_ADDR outputScratch = userWorkspace + tiling.outputScratchOffset; + uint64_t outputElements = static_cast(tiling.batch) * + static_cast(tiling.vHeadNum) * + static_cast(tiling.seqlen) * + static_cast(tiling.vHeadDim); + GM_ADDR stateScratch = outputScratch; + GM_ADDR localScratch = outputScratch + outputElements * sizeof(float); + if ASCEND_IS_AIC { + ChunkKdaFwdFinalizeKernel op; + op.Init(q, k, v, gk, beta, initialState, cuSeqlens, chunkIndices, + qgScaled, aqk, propagatedVNew, propagatedH, stateScratch, userWorkspace, aqk, userWorkspace, + userWorkspace, localScratch, userWorkspace, userWorkspace, o, propagatedH, + outputScratch, tiling, &pipe, false); + op.ProcessAic(); + } + if ASCEND_IS_AIV { + ChunkKdaFwdFinalizeKernel op; + op.Init(q, k, v, gk, beta, initialState, cuSeqlens, chunkIndices, + qgScaled, aqk, propagatedVNew, propagatedH, stateScratch, userWorkspace, aqk, userWorkspace, + userWorkspace, localScratch, userWorkspace, userWorkspace, o, propagatedH, + outputScratch, tiling, &pipe); + op.ProcessAiv(); + } +} + +} // namespace KdaFinalize diff --git a/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_fwd_h.h b/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_fwd_h.h new file mode 100644 index 000000000000..249e54d96fe1 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_fwd_h.h @@ -0,0 +1,1009 @@ +#ifndef CHUNK_KDA_FWD_ARCH35_FWD_H_H +#define CHUNK_KDA_FWD_ARCH35_FWD_H_H + +#include "kernel_operator.h" +#include "catlass/arch/resource.hpp" +#include "catlass/gemm/tile/tile_copy.hpp" +#include "catlass/gemm/tile/tile_mmad.hpp" +#include "kernel_utils/tile/copy_l0c_to_ub.hpp" + +namespace KdaForward::arch35 { + +using namespace AscendC; + +// The direct UB is shared by one AIC and its two AIV subblocks. Mode 4 +// addresses each subblock explicitly: the second AIV uses flag + 16. +constexpr uint64_t KDA_FWD_H_DIRECT_FREE_FLAG = 6; +constexpr uint64_t KDA_FWD_H_DIRECT_READY_FLAG = 7; +constexpr uint64_t KDA_FWD_H_SUBBLOCK_FLAG_OFFSET = 16; +// Prepare is fully drained before FwdH starts. Keep FwdH mode-2 L1 +// traffic outside Matmul's possible 0..7 range and SyncAll's 11..14 range. +constexpr uint64_t KDA_FWD_H_L1_FREE_FLAG = 9; +constexpr uint64_t KDA_FWD_H_L1_READY_FLAG = 9; +constexpr uint64_t KDA_FWD_H_STATE_FREE_FLAG = KDA_FWD_H_L1_FREE_FLAG; +constexpr uint64_t KDA_FWD_H_STATE_READY_FLAG = KDA_FWD_H_L1_READY_FLAG; +constexpr uint64_t KDA_FWD_H_VNEW_FREE_FLAG = 10; +constexpr uint64_t KDA_FWD_H_VNEW_READY_FLAG = 10; + +constexpr TEventID KDA_FWD_H_MTE_W_EVENT = 0; +constexpr TEventID KDA_FWD_H_MTE_Q_EVENT = 1; +constexpr TEventID KDA_FWD_H_MTE_B_EVENT = 2; +constexpr TEventID KDA_FWD_H_MTE_A_EVENT = 3; +constexpr TEventID KDA_FWD_H_M_EVENT = 4; +constexpr TEventID KDA_FWD_H_FIX_EVENT = 5; +constexpr TEventID KDA_FWD_H_IO_REUSE_EVENT = 6; + +constexpr uint32_t KDA_FWD_H_CHUNK = 64; +constexpr uint32_t KDA_FWD_H_DIM = 128; +constexpr uint32_t KDA_FWD_H_SUB_CHUNK = KDA_FWD_H_CHUNK / 2; +constexpr uint32_t KDA_FWD_H_SUB_DIM = KDA_FWD_H_DIM / 2; +constexpr uint32_t KDA_FWD_H_STATE_SUB_ELEMS = KDA_FWD_H_SUB_DIM * KDA_FWD_H_DIM; +constexpr uint32_t KDA_FWD_H_TOKEN_SUB_ELEMS = KDA_FWD_H_SUB_CHUNK * KDA_FWD_H_DIM; + +constexpr uint32_t KDA_FWD_H_L1_W_OFFSET = 0; +constexpr uint32_t KDA_FWD_H_L1_Q_OFFSET = 16 * 1024; +constexpr uint32_t KDA_FWD_H_L1_H_OFFSET = 32 * 1024; +constexpr uint32_t KDA_FWD_H_L1_KG_OFFSET = 64 * 1024; +constexpr uint32_t KDA_FWD_H_L1_AQK_OFFSET = 80 * 1024; +constexpr uint32_t KDA_FWD_H_L1_V_OFFSET = 96 * 1024; +constexpr uint32_t KDA_FWD_H_L1_W1_OFFSET = 112 * 1024; +constexpr uint32_t KDA_FWD_H_L1_Q1_OFFSET = 128 * 1024; +constexpr uint32_t KDA_FWD_H_L1_KG1_OFFSET = 144 * 1024; +constexpr uint32_t KDA_FWD_H_L1_AQK1_OFFSET = 160 * 1024; +constexpr uint32_t KDA_FWD_H_L1_W2_OFFSET = 176 * 1024; +constexpr uint32_t KDA_FWD_H_L1_Q2_OFFSET = 192 * 1024; +constexpr uint32_t KDA_FWD_H_L1_KG2_OFFSET = 208 * 1024; +constexpr uint32_t KDA_FWD_H_L1_AQK2_OFFSET = 224 * 1024; +constexpr uint32_t KDA_FWD_H_L1_W3_OFFSET = 240 * 1024; +constexpr uint32_t KDA_FWD_H_L1_Q3_OFFSET = 256 * 1024; +constexpr uint32_t KDA_FWD_H_L1_KG3_OFFSET = 272 * 1024; +constexpr uint32_t KDA_FWD_H_L1_AQK3_OFFSET = 288 * 1024; +constexpr uint32_t KDA_FWD_H_L1_AKK_OFFSET = 304 * 1024; +constexpr uint32_t KDA_FWD_H_L1_U_OFFSET = 336 * 1024; +constexpr uint32_t KDA_FWD_H_L1_STAGING_DEPTH = 4; +constexpr uint32_t KDA_FWD_H_L1_AKK_SLOT_BYTES = 8 * 1024; +constexpr uint32_t KDA_FWD_H_L1_U_SLOT_BYTES = 16 * 1024; + +constexpr uint32_t KDA_FWD_H_L0A_STATE_OFFSET = 0; +constexpr uint32_t KDA_FWD_H_L0A_VNEW_OFFSET = 16 * 1024; +constexpr uint32_t KDA_FWD_H_L0A_POST_OFFSET = 32 * 1024; +constexpr uint32_t KDA_FWD_H_L0B_STATE_OFFSET = 0; +constexpr uint32_t KDA_FWD_H_L0B_VNEW_OFFSET = 32 * 1024; +constexpr uint32_t KDA_FWD_H_L0B_POST_OFFSET = 32 * 1024; + +constexpr uint32_t KDA_FWD_H_UB_STATE_OFFSET = 0; +constexpr uint32_t KDA_FWD_H_UB_STATE_TYPED_OFFSET = 32 * 1024; +constexpr uint32_t KDA_FWD_H_UB_DIRECT_OFFSET = 48 * 1024; +constexpr uint32_t KDA_FWD_H_UB_OUT1_OFFSET = 80 * 1024; +constexpr uint32_t KDA_FWD_H_UB_VNEW_OFFSET = 96 * 1024; +constexpr uint32_t KDA_FWD_H_UB_IO_OFFSET = 112 * 1024; +constexpr uint32_t KDA_FWD_H_UB_GATE_OFFSET = 128 * 1024; + +template +class ChunkKdaFwdFwdH { +public: + using ArchTag = Catlass::Arch::Ascend950; + using LayoutRM = Catlass::layout::RowMajor; + using LayoutCM = Catlass::layout::ColumnMajor; + using TileCopyRM = Catlass::Gemm::Tile::PackedTileCopyTla< + ArchTag, T, LayoutRM, T, LayoutRM, float, LayoutRM>; + using DirectTileCopyRM = Common::Tile::PackedTileCopyTlaToUB< + ArchTag, T, LayoutRM, T, LayoutRM, float, LayoutRM, void, + Catlass::Gemm::Tile::CopyL0CToUBMode::SPLIT_M>; + using TileCopyCM = Catlass::Gemm::Tile::PackedTileCopyTla< + ArchTag, T, LayoutCM, T, LayoutRM, float, LayoutRM>; + using DirectTileCopyCM = Common::Tile::PackedTileCopyTlaToUB< + ArchTag, T, LayoutCM, T, LayoutRM, float, LayoutRM, void, + Catlass::Gemm::Tile::CopyL0CToUBMode::SPLIT_M>; + + __aicore__ inline void Init( + GM_ADDR gk, GM_ADDR initialState, GM_ADDR attnOut, GM_ADDR finalState, + GM_ADDR aqk, GM_ADDR akk, GM_ADDR w, GM_ADDR u, GM_ADDR qgScaled, GM_ADDR kg, + GM_ADDR vNew, GM_ADDR h, const TilingData &tiling) + { + gk_.SetGlobalBuffer(reinterpret_cast<__gm__ GK_T *>(gk)); + if (initialState != nullptr) { + initialState_.SetGlobalBuffer(reinterpret_cast<__gm__ float *>(initialState)); + } + attnOut_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(attnOut)); + finalState_.SetGlobalBuffer(reinterpret_cast<__gm__ float *>(finalState)); + aqk_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(aqk)); + akk_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(akk)); + w_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(w)); + u_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(u)); + qgScaled_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(qgScaled)); + kg_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(kg)); + vNew_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(vNew)); + h_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(h)); + batch_ = tiling.batch; + heads_ = tiling.vHeadNum; + seqlen_ = tiling.seqlen; + totalChunks_ = tiling.totalChunks; + hasInitialState_ = tiling.hasInitialState; + storeFinalState_ = tiling.storeFinalState; + storeVNew_ = tiling.storeVNew; + storeH_ = tiling.storeH; + fusePostWuIntoFwdH_ = tiling.fusePostWuIntoFwdH; + coreNum_ = tiling.prepareUsedCoreNum; + statePublishCount_[0] = 0; + statePublishCount_[1] = 0; + vnewPublishCount_[0] = 0; + vnewPublishCount_[1] = 0; + } + + __aicore__ inline void Process() + { + if ASCEND_IS_AIC { + ProcessAic(); + } + if ASCEND_IS_AIV { + ProcessAiv(); + } + } + +private: + __aicore__ inline void WaitL1SlotFreeMte3( + uint64_t freeFlag, uint32_t publishCount) + { + if (publishCount != 0) { + CrossCoreWaitFlag(freeFlag); + } + } + + template + __aicore__ inline void SetL1SlotFlagAicToAiv(uint64_t flag) + { + CrossCoreSetFlag<0x2, PIPE>(flag); + } + + template + __aicore__ inline void SetL1SlotFlagAivToAic(uint64_t flag) + { + CrossCoreSetFlag<0x2, PIPE>(flag); + } + + __aicore__ inline void WaitL1SlotReadyMte1(uint64_t readyFlag) + { + CrossCoreWaitFlag(readyFlag); + } + + __aicore__ inline void WaitDirectFreeAic() + { + CrossCoreWaitFlag<0x4, PIPE_FIX>(KDA_FWD_H_DIRECT_FREE_FLAG); + CrossCoreWaitFlag<0x4, PIPE_FIX>( + KDA_FWD_H_DIRECT_FREE_FLAG + KDA_FWD_H_SUBBLOCK_FLAG_OFFSET); + } + + __aicore__ inline void SetDirectReadyAic() + { + CrossCoreSetFlag<0x4, PIPE_FIX>(KDA_FWD_H_DIRECT_READY_FLAG); + CrossCoreSetFlag<0x4, PIPE_FIX>( + KDA_FWD_H_DIRECT_READY_FLAG + KDA_FWD_H_SUBBLOCK_FLAG_OFFSET); + } + + __aicore__ inline void WaitDirectReadyAiv() + { + CrossCoreWaitFlag<0x4, PIPE_V>(KDA_FWD_H_DIRECT_READY_FLAG); + } + + __aicore__ inline void SetDirectFreeAiv() + { + CrossCoreSetFlag<0x4, PIPE_V>(KDA_FWD_H_DIRECT_FREE_FLAG); + } + + __aicore__ inline uint64_t MatrixOffset(uint64_t b, uint64_t hv, uint64_t t) const + { + return ((b * heads_ + hv) * seqlen_ + t) * KDA_FWD_H_DIM; + } + + __aicore__ inline uint64_t ChunkMatrixOffset( + uint64_t b, uint64_t hv, uint64_t chunk, uint64_t row = 0) const + { + return ((b * heads_ + hv) * seqlen_ + + chunk * KDA_FWD_H_CHUNK + row) * KDA_FWD_H_DIM; + } + + __aicore__ inline uint64_t ScoreOffset(uint64_t b, uint64_t hv, uint64_t t) const + { + return ((b * heads_ + hv) * seqlen_ + t) * KDA_FWD_H_CHUNK; + } + + __aicore__ inline uint64_t ChunkScoreOffset( + uint64_t b, uint64_t hv, uint64_t chunk) const + { + return ((b * heads_ + hv) * seqlen_ + + chunk * KDA_FWD_H_CHUNK) * KDA_FWD_H_CHUNK; + } + + __aicore__ inline uint64_t StateOffset(uint64_t b, uint64_t hv) const + { + return (b * heads_ + hv) * KDA_FWD_H_DIM * KDA_FWD_H_DIM; + } + + __aicore__ inline uint64_t HOffset(uint64_t b, uint64_t hv, uint64_t chunk) const + { + if (!storeH_) { + return StateOffset(b, hv); + } + return ((b * heads_ + hv) * totalChunks_ + chunk) * + KDA_FWD_H_DIM * KDA_FWD_H_DIM; + } + + __aicore__ inline uint64_t VNewOffset( + uint64_t b, uint64_t hv, uint64_t chunk, uint64_t row = 0) const + { + if (!storeVNew_) { + return ((b * heads_ + hv) * KDA_FWD_H_CHUNK + row) * KDA_FWD_H_DIM; + } + return ChunkMatrixOffset(b, hv, chunk, row); + } + + __aicore__ inline uint64_t OutputOffset(uint64_t b, uint64_t hv, uint64_t t) const + { + return ((b * seqlen_ + t) * heads_ + hv) * KDA_FWD_H_DIM; + } + + __aicore__ inline uint32_t L1WOffset(uint32_t slot) const + { + if (slot == 0) { + return KDA_FWD_H_L1_W_OFFSET; + } + if (slot == 1) { + return KDA_FWD_H_L1_W1_OFFSET; + } + return slot == 2 ? KDA_FWD_H_L1_W2_OFFSET : KDA_FWD_H_L1_W3_OFFSET; + } + + __aicore__ inline uint32_t L1QOffset(uint32_t slot) const + { + if (slot == 0) { + return KDA_FWD_H_L1_Q_OFFSET; + } + if (slot == 1) { + return KDA_FWD_H_L1_Q1_OFFSET; + } + return slot == 2 ? KDA_FWD_H_L1_Q2_OFFSET : KDA_FWD_H_L1_Q3_OFFSET; + } + + __aicore__ inline uint32_t L1KgOffset(uint32_t slot) const + { + if (slot == 0) { + return KDA_FWD_H_L1_KG_OFFSET; + } + if (slot == 1) { + return KDA_FWD_H_L1_KG1_OFFSET; + } + return slot == 2 ? KDA_FWD_H_L1_KG2_OFFSET : KDA_FWD_H_L1_KG3_OFFSET; + } + + __aicore__ inline uint32_t L1AqkOffset(uint32_t slot) const + { + if (slot == 0) { + return KDA_FWD_H_L1_AQK_OFFSET; + } + if (slot == 1) { + return KDA_FWD_H_L1_AQK1_OFFSET; + } + return slot == 2 ? KDA_FWD_H_L1_AQK2_OFFSET : KDA_FWD_H_L1_AQK3_OFFSET; + } + + __aicore__ inline uint32_t L1AkkOffset(uint32_t slot) const + { + return KDA_FWD_H_L1_AKK_OFFSET + slot * KDA_FWD_H_L1_AKK_SLOT_BYTES; + } + + __aicore__ inline uint32_t L1UOffset(uint32_t slot) const + { + return KDA_FWD_H_L1_U_OFFSET + slot * KDA_FWD_H_L1_U_SLOT_BYTES; + } + + template + __aicore__ inline void PublishDirectTile( + TensorL0C tensorL0C, uint32_t m, uint32_t n, + TEventID mToFixEvent, TEventID fixToMEvent) + { + auto layoutUb = tla::MakeLayout(m, n); + auto tensorUb = tla::MakeTensor( + resource_.ubBuf.template GetBufferByByte(KDA_FWD_H_UB_DIRECT_OFFSET), + layoutUb, Catlass::Arch::PositionUB{}); + using CopyL0CToDst = + typename DirectTileCopy::template CopyL0CToDst; + CopyL0CToDst copyL0CToDst; + + WaitDirectFreeAic(); + SetFlag(mToFixEvent); + WaitFlag(mToFixEvent); + copyL0CToDst(tensorUb, tensorL0C); + SetDirectReadyAic(); + SetFlag(fixToMEvent); + } + + template + __aicore__ inline void PublishDirect( + LocalTensor l0C, uint32_t m, uint32_t n, + TEventID mToFixEvent, TEventID fixToMEvent) + { + auto layoutL0C = tla::MakeLayoutL0C(m, n); + auto tensorL0C = tla::MakeTensor(l0C, layoutL0C, Catlass::Arch::PositionL0C{}); + PublishDirectTile( + tensorL0C, m, n, mToFixEvent, fixToMEvent); + } + + __aicore__ inline void PrefetchIndependentProductsAic( + uint64_t b, uint64_t hv, uint64_t chunk) + { + const uint32_t slot = static_cast(chunk & 3); + WaitFlag(aicL1ReuseEvents_[slot]); + using LayoutTagL1ARm = typename TileCopyRM::LayoutTagL1A; + using LayoutTagL1BRm = typename TileCopyRM::LayoutTagL1B; + using LayoutTagL1ACm = typename TileCopyCM::LayoutTagL1A; + + auto layoutToken = tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM); + auto layoutKg = tla::MakeLayout(KDA_FWD_H_DIM, KDA_FWD_H_CHUNK); + auto layoutAqk = tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK); + auto layoutAkk = tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK); + auto tensorW = tla::MakeTensor( + w_[ChunkMatrixOffset(b, hv, chunk)], layoutToken, Catlass::Arch::PositionGM{}); + auto tensorQ = tla::MakeTensor( + qgScaled_[ChunkMatrixOffset(b, hv, chunk)], layoutToken, + Catlass::Arch::PositionGM{}); + auto tensorKg = tla::MakeTensor( + kg_[ChunkMatrixOffset(b, hv, chunk)], layoutKg, Catlass::Arch::PositionGM{}); + auto tensorAqk = tla::MakeTensor( + aqk_[ChunkScoreOffset(b, hv, chunk)], layoutAqk, Catlass::Arch::PositionGM{}); + auto tensorAkk = tla::MakeTensor( + akk_[ChunkScoreOffset(b, hv, chunk)], layoutAkk, Catlass::Arch::PositionGM{}); + auto tensorU = tla::MakeTensor( + u_[ChunkMatrixOffset(b, hv, chunk)], layoutToken, Catlass::Arch::PositionGM{}); + auto blockW = GetTile(tensorW, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + auto blockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + auto blockKg = GetTile(tensorKg, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_DIM, KDA_FWD_H_CHUNK)); + auto blockAqk = GetTile(tensorAqk, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK)); + auto blockAkk = GetTile(tensorAkk, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK)); + auto blockU = GetTile(tensorU, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + using CopyGmToL1ARmW = typename TileCopyRM::template CopyGmToL1A; + using CopyGmToL1ARmQ = typename TileCopyRM::template CopyGmToL1A; + using CopyGmToL1ACm = typename TileCopyCM::template CopyGmToL1A; + using CopyGmToL1ARmAqk = typename TileCopyRM::template CopyGmToL1A; + using CopyGmToL1ARmAkk = typename TileCopyRM::template CopyGmToL1A; + using CopyGmToL1BRmU = typename TileCopyRM::template CopyGmToL1B; + + LocalTensor l1W = resource_.l1Buf.template GetBufferByByte( + L1WOffset(slot)); + LocalTensor l1Q = resource_.l1Buf.template GetBufferByByte( + L1QOffset(slot)); + LocalTensor l1Kg = resource_.l1Buf.template GetBufferByByte( + L1KgOffset(slot)); + LocalTensor l1Aqk = resource_.l1Buf.template GetBufferByByte( + L1AqkOffset(slot)); + LocalTensor l1Akk = resource_.l1Buf.template GetBufferByByte( + L1AkkOffset(slot)); + LocalTensor l1U = resource_.l1Buf.template GetBufferByByte( + L1UOffset(slot)); + auto tensorL1W = tla::MakeTensor( + l1W, tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM), + Catlass::Arch::PositionL1{}); + auto tensorL1Q = tla::MakeTensor( + l1Q, tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM), + Catlass::Arch::PositionL1{}); + auto tensorL1Kg = tla::MakeTensor( + l1Kg, tla::MakeLayout(KDA_FWD_H_DIM, KDA_FWD_H_CHUNK), + Catlass::Arch::PositionL1{}); + auto tensorL1Aqk = tla::MakeTensor( + l1Aqk, tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK), + Catlass::Arch::PositionL1{}); + auto tensorL1Akk = tla::MakeTensor( + l1Akk, tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK), + Catlass::Arch::PositionL1{}); + auto tensorL1U = tla::MakeTensor( + l1U, tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM), + Catlass::Arch::PositionL1{}); + + CopyGmToL1ARmW{}(tensorL1W, blockW); + CopyGmToL1ARmQ{}(tensorL1Q, blockQ); + CopyGmToL1ACm{}(tensorL1Kg, blockKg); + CopyGmToL1ARmAqk{}(tensorL1Aqk, blockAqk); + if (fusePostWuIntoFwdH_) { + CopyGmToL1ARmAkk{}(tensorL1Akk, blockAkk); + CopyGmToL1BRmU{}(tensorL1U, blockU); + } + SetFlag(aicMte2ToMte1Event_); + } + + __aicore__ inline void ComputePostWuAic(uint64_t chunk) + { + const uint32_t slot = static_cast(chunk & 3); + using LayoutTagL1A = typename TileCopyRM::LayoutTagL1A; + using LayoutTagL1B = typename TileCopyRM::LayoutTagL1B; + using LayoutTagL0A = typename TileCopyRM::LayoutTagL0A; + using LayoutTagL0B = typename TileCopyRM::LayoutTagL0B; + using CopyL1ToL0A = typename TileCopyRM::CopyL1ToL0A; + using CopyL1ToL0B = typename TileCopyRM::CopyL1ToL0B; + using TileMmad = Catlass::Gemm::Tile::TileMmadTla; + + LocalTensor l1Akk = resource_.l1Buf.template GetBufferByByte( + L1AkkOffset(slot)); + LocalTensor l1W = resource_.l1Buf.template GetBufferByByte( + L1WOffset(slot)); + LocalTensor l1U = resource_.l1Buf.template GetBufferByByte( + L1UOffset(slot)); + LocalTensor l0A = resource_.l0ABuf.template GetBufferByByte( + KDA_FWD_H_L0A_POST_OFFSET); + LocalTensor l0B = resource_.l0BBuf.template GetBufferByByte( + KDA_FWD_H_L0B_POST_OFFSET); + LocalTensor l0C = resource_.l0CBuf.template GetBufferByByte(0); + + auto tensorL1Akk = tla::MakeTensor( + l1Akk, tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK), + Catlass::Arch::PositionL1{}); + auto tensorL1W = tla::MakeTensor( + l1W, tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM), + Catlass::Arch::PositionL1{}); + auto tensorL1U = tla::MakeTensor( + l1U, tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM), + Catlass::Arch::PositionL1{}); + auto tensorL0A = tla::MakeTensor( + l0A, tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK), + Catlass::Arch::PositionL0A{}); + auto tensorL0B = tla::MakeTensor( + l0B, tla::MakeLayout(KDA_FWD_H_CHUNK, 2 * KDA_FWD_H_DIM), + Catlass::Arch::PositionL0B{}); + auto tensorL0C = tla::MakeTensor( + l0C, tla::MakeLayoutL0C(KDA_FWD_H_CHUNK, 2 * KDA_FWD_H_DIM), + Catlass::Arch::PositionL0C{}); + auto tileL1Akk = GetTile(tensorL1Akk, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK)); + auto tileL1W = GetTile(tensorL1W, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + auto tileL1U = GetTile(tensorL1U, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + auto tileL0A = GetTile(tensorL0A, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK)); + auto tileL0BW = GetTile(tensorL0B, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + auto tileL0BU = GetTile(tensorL0B, tla::MakeCoord(0, KDA_FWD_H_DIM), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + auto tileL0B = GetTile(tensorL0B, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, 2 * KDA_FWD_H_DIM)); + auto tileL0C = GetTile(tensorL0C, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, 2 * KDA_FWD_H_DIM)); + auto tileL0CW = GetTile(tensorL0C, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + auto tileL0CU = GetTile(tensorL0C, tla::MakeCoord(0, KDA_FWD_H_DIM), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + + WaitFlag(aicMte2ToMte1Event_); + WaitFlag(stateL0FreeEvent_); + CopyL1ToL0A{}(tileL0A, tileL1Akk); + CopyL1ToL0B{}(tileL0BW, tileL1W); + CopyL1ToL0B{}(tileL0BU, tileL1U); + SetFlag(aicMte1ToMEvent_); + WaitFlag(aicMte1ToMEvent_); + TileMmad{}(tileL0C, tileL0A, tileL0B, KDA_FWD_H_CHUNK, + 2 * KDA_FWD_H_DIM, KDA_FWD_H_CHUNK, true, 0); + SetFlag(stateL0FreeEvent_); + PublishDirectTile( + tileL0CW, KDA_FWD_H_CHUNK, KDA_FWD_H_DIM, + aicMToFixEvent_, aicFixToMEvent_); + WaitFlag(aicFixToMEvent_); + PublishDirectTile( + tileL0CU, KDA_FWD_H_CHUNK, KDA_FWD_H_DIM, + aicMToFixEvent_, aicFixToMEvent_); + WaitFlag(aicFixToMEvent_); + } + + __aicore__ inline void ComputeStateProductsAic( + uint64_t b, uint64_t hv, uint64_t chunk, bool inputsReady = false) + { + const uint32_t slot = static_cast(chunk & 3); + using LayoutTagL1A = typename TileCopyRM::LayoutTagL1A; + using LayoutTagL1B = typename TileCopyRM::LayoutTagL1B; + using LayoutTagL0A = typename TileCopyRM::LayoutTagL0A; + using LayoutTagL0B = typename TileCopyRM::LayoutTagL0B; + using CopyL1ToL0A = typename TileCopyRM::CopyL1ToL0A; + using CopyL1ToL0B = typename TileCopyRM::CopyL1ToL0B; + using TileMmad = Catlass::Gemm::Tile::TileMmadTla; + + LocalTensor l1A0 = resource_.l1Buf.template GetBufferByByte( + L1WOffset(slot)); + LocalTensor l1A1 = resource_.l1Buf.template GetBufferByByte( + L1QOffset(slot)); + LocalTensor l1B = + resource_.l1Buf.template GetBufferByByte(KDA_FWD_H_L1_H_OFFSET); + LocalTensor l0A = resource_.l0ABuf.template GetBufferByByte( + KDA_FWD_H_L0A_STATE_OFFSET); + LocalTensor l0B = resource_.l0BBuf.template GetBufferByByte( + KDA_FWD_H_L0B_STATE_OFFSET); + LocalTensor l0C = resource_.l0CBuf.template GetBufferByByte(0); + + auto layoutL1A = tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM); + auto layoutL1B = tla::MakeLayout(KDA_FWD_H_DIM, KDA_FWD_H_DIM); + auto layoutL0A = tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM); + auto layoutL0B = tla::MakeLayout(KDA_FWD_H_DIM, KDA_FWD_H_DIM); + auto layoutL0C = tla::MakeLayoutL0C(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM); + auto tensorL1A0 = tla::MakeTensor(l1A0, layoutL1A, Catlass::Arch::PositionL1{}); + auto tensorL1A1 = tla::MakeTensor(l1A1, layoutL1A, Catlass::Arch::PositionL1{}); + auto tensorL1B = tla::MakeTensor(l1B, layoutL1B, Catlass::Arch::PositionL1{}); + auto tensorL0A = tla::MakeTensor(l0A, layoutL0A, Catlass::Arch::PositionL0A{}); + auto tensorL0B = tla::MakeTensor(l0B, layoutL0B, Catlass::Arch::PositionL0B{}); + auto tensorL0C = tla::MakeTensor(l0C, layoutL0C, Catlass::Arch::PositionL0C{}); + auto tileL1B = GetTile(tensorL1B, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_DIM, KDA_FWD_H_DIM)); + auto tileL0A = GetTile(tensorL0A, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + auto tileL1W = GetTile(tensorL1A0, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + auto tileL1Q = GetTile(tensorL1A1, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + auto tileL0B = GetTile(tensorL0B, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_DIM, KDA_FWD_H_DIM)); + auto tileL0C = GetTile(tensorL0C, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + + CopyL1ToL0A copyL1ToL0A; + CopyL1ToL0B copyL1ToL0B; + TileMmad tileMmad; + + if (!inputsReady) { + WaitFlag(aicMte2ToMte1Event_); + } + WaitFlag(stateL0FreeEvent_); + copyL1ToL0A(tileL0A, tileL1W); + copyL1ToL0B(tileL0B, tileL1B); + SetFlag(aicMte1ToMEvent_); + WaitFlag(aicMte1ToMEvent_); + tileMmad(tileL0C, tileL0A, tileL0B, KDA_FWD_H_CHUNK, + KDA_FWD_H_DIM, KDA_FWD_H_DIM, true, 0); + SetFlag(stateL0FreeEvent_); + PublishDirect( + l0C, KDA_FWD_H_CHUNK, KDA_FWD_H_DIM, + aicMToFixEvent_, aicFixToMEvent_); + SetL1SlotFlagAicToAiv(KDA_FWD_H_STATE_FREE_FLAG); + + WaitFlag(stateL0FreeEvent_); + WaitFlag(aicFixToMEvent_); + copyL1ToL0A(tileL0A, tileL1Q); + SetFlag(aicMte1ToMEvent_); + WaitFlag(aicMte1ToMEvent_); + tileMmad(tileL0C, tileL0A, tileL0B, KDA_FWD_H_CHUNK, + KDA_FWD_H_DIM, KDA_FWD_H_DIM, true, 0); + SetFlag(stateL0FreeEvent_); + PublishDirect( + l0C, KDA_FWD_H_CHUNK, KDA_FWD_H_DIM, + aicMToFixEvent_, aicFixToMEvent_); + WaitFlag(aicFixToMEvent_); + } + + __aicore__ inline void ComputeVnewProductsAic( + uint64_t b, uint64_t hv, uint64_t chunk, bool prefetchNext) + { + const uint32_t slot = static_cast(chunk & 3); + using LayoutTagL1AK = typename TileCopyCM::LayoutTagL1A; + using LayoutTagL1AA = typename TileCopyRM::LayoutTagL1A; + using LayoutTagL1B = typename TileCopyRM::LayoutTagL1B; + using LayoutTagL0AK = typename TileCopyCM::LayoutTagL0A; + using LayoutTagL0AA = typename TileCopyRM::LayoutTagL0A; + using LayoutTagL0B = typename TileCopyRM::LayoutTagL0B; + using CopyL1ToL0AK = typename TileCopyCM::CopyL1ToL0A; + using CopyL1ToL0AA = typename TileCopyRM::CopyL1ToL0A; + using CopyL1ToL0B = typename TileCopyRM::CopyL1ToL0B; + using TileMmadK = Catlass::Gemm::Tile::TileMmadTla; + using TileMmadA = Catlass::Gemm::Tile::TileMmadTla; + + auto layoutKg = tla::MakeLayout(KDA_FWD_H_DIM, KDA_FWD_H_CHUNK); + auto layoutAqk = tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK); + LocalTensor l1Kg = resource_.l1Buf.template GetBufferByByte( + L1KgOffset(slot)); + LocalTensor l1Aqk = resource_.l1Buf.template GetBufferByByte( + L1AqkOffset(slot)); + LocalTensor l1V = + resource_.l1Buf.template GetBufferByByte(KDA_FWD_H_L1_V_OFFSET); + LocalTensor l0A = resource_.l0ABuf.template GetBufferByByte( + KDA_FWD_H_L0A_VNEW_OFFSET); + LocalTensor l0B = resource_.l0BBuf.template GetBufferByByte( + KDA_FWD_H_L0B_VNEW_OFFSET); + LocalTensor l0C = resource_.l0CBuf.template GetBufferByByte(0); + + auto layoutL1Kg = tla::MakeLayout(KDA_FWD_H_DIM, KDA_FWD_H_CHUNK); + auto layoutL1Aqk = tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK); + auto layoutL1V = tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM); + auto layoutL0Kg = tla::MakeLayout(KDA_FWD_H_DIM, KDA_FWD_H_CHUNK); + auto layoutL0Aqk = tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK); + auto layoutL0V = tla::MakeLayout(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM); + auto baseL1Kg = tla::MakeTensor(l1Kg, layoutL1Kg, Catlass::Arch::PositionL1{}); + auto baseL1Aqk = tla::MakeTensor(l1Aqk, layoutL1Aqk, Catlass::Arch::PositionL1{}); + auto baseL1V = tla::MakeTensor(l1V, layoutL1V, Catlass::Arch::PositionL1{}); + auto baseL0Kg = tla::MakeTensor(l0A, layoutL0Kg, Catlass::Arch::PositionL0A{}); + auto baseL0Aqk = tla::MakeTensor(l0A, layoutL0Aqk, Catlass::Arch::PositionL0A{}); + auto baseL0V = tla::MakeTensor(l0B, layoutL0V, Catlass::Arch::PositionL0B{}); + auto baseL0Update = tla::MakeTensor( + l0C, tla::MakeLayoutL0C(KDA_FWD_H_DIM, KDA_FWD_H_DIM), + Catlass::Arch::PositionL0C{}); + auto baseL0Out = tla::MakeTensor( + l0C, tla::MakeLayoutL0C(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM), + Catlass::Arch::PositionL0C{}); + auto tensorL1Kg = GetTile(baseL1Kg, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_DIM, KDA_FWD_H_CHUNK)); + auto tensorL1Aqk = GetTile(baseL1Aqk, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK)); + auto tensorL1V = GetTile(baseL1V, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + auto tensorL0Kg = GetTile(baseL0Kg, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_DIM, KDA_FWD_H_CHUNK)); + auto tensorL0Aqk = GetTile(baseL0Aqk, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_CHUNK)); + auto tensorL0V = GetTile(baseL0V, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + auto tensorL0Update = GetTile(baseL0Update, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_DIM, KDA_FWD_H_DIM)); + auto tensorL0Out = GetTile(baseL0Out, tla::MakeCoord(0, 0), + tla::MakeShape(KDA_FWD_H_CHUNK, KDA_FWD_H_DIM)); + + CopyL1ToL0AK copyL1ToL0AK; + CopyL1ToL0AA copyL1ToL0AA; + CopyL1ToL0B copyL1ToL0B; + TileMmadK tileMmadK; + TileMmadA tileMmadA; + + WaitFlag(vnewL0FreeEvent_); + copyL1ToL0AK(tensorL0Kg, tensorL1Kg); + copyL1ToL0B(tensorL0V, tensorL1V); + SetFlag(aicMte1ToMEvent_); + WaitFlag(aicMte1ToMEvent_); + tileMmadK(tensorL0Update, tensorL0Kg, tensorL0V, + KDA_FWD_H_DIM, KDA_FWD_H_DIM, KDA_FWD_H_CHUNK, true, 0); + SetFlag(vnewL0FreeEvent_); + PublishDirect( + l0C, KDA_FWD_H_DIM, KDA_FWD_H_DIM, + aicMToFixEvent_, aicFixToMEvent_); + SetL1SlotFlagAicToAiv(KDA_FWD_H_VNEW_FREE_FLAG); + + WaitFlag(vnewL0FreeEvent_); + WaitFlag(aicFixToMEvent_); + copyL1ToL0AA(tensorL0Aqk, tensorL1Aqk); + SetFlag(aicL1ReuseEvents_[slot]); + if (prefetchNext) { + PrefetchIndependentProductsAic(b, hv, chunk + 1); + } + SetFlag(aicMte1ToMEvent_); + WaitFlag(aicMte1ToMEvent_); + tileMmadA(tensorL0Out, tensorL0Aqk, tensorL0V, + KDA_FWD_H_CHUNK, KDA_FWD_H_DIM, KDA_FWD_H_CHUNK, true, 0); + SetFlag(vnewL0FreeEvent_); + PublishDirect( + l0C, KDA_FWD_H_CHUNK, KDA_FWD_H_DIM, + aicMToFixEvent_, aicFixToMEvent_); + WaitFlag(aicFixToMEvent_); + } + + __aicore__ inline void ProcessAic() + { + SetLoadDataPaddingValue(static_cast(0)); + SetFlag(stateL0FreeEvent_); + for (uint32_t slot = 0; slot < KDA_FWD_H_L1_STAGING_DEPTH; ++slot) { + SetFlag(aicL1ReuseEvents_[slot]); + } + const uint64_t coreIdx = static_cast(GetBlockIdx()); + const uint64_t coreNum = coreNum_ == 0 ? 1 : coreNum_; + for (uint64_t task = coreIdx; task < batch_ * heads_; task += coreNum) { + const uint64_t b = task / heads_; + const uint64_t hv = task % heads_; + PrefetchIndependentProductsAic(b, hv, 0); + for (uint64_t chunk = 0; chunk < totalChunks_; ++chunk) { + if (fusePostWuIntoFwdH_) { + ComputePostWuAic(chunk); + } + WaitL1SlotReadyMte1(KDA_FWD_H_STATE_READY_FLAG); + ComputeStateProductsAic(b, hv, chunk, fusePostWuIntoFwdH_); + WaitL1SlotReadyMte1(KDA_FWD_H_VNEW_READY_FLAG); + ComputeVnewProductsAic( + b, hv, chunk, chunk + 1 < totalChunks_); + } + WaitDirectFreeAic(); + } + WaitFlag(stateL0FreeEvent_); + for (uint32_t slot = 0; slot < KDA_FWD_H_L1_STAGING_DEPTH; ++slot) { + WaitFlag(aicL1ReuseEvents_[slot]); + } + } + + __aicore__ inline void CopyOutputRows( + uint64_t b, uint64_t hv, uint64_t tokenStart, + LocalTensor src, uint32_t rows) + { + DataCopyExtParams params{ + static_cast(rows), + static_cast(KDA_FWD_H_DIM * sizeof(T)), + 0, + static_cast((heads_ * KDA_FWD_H_DIM - KDA_FWD_H_DIM) * sizeof(T)), + 0}; + DataCopyPad(attnOut_[OutputOffset(b, hv, tokenStart)], src, params); + } + + __aicore__ inline void InitializeStateAiv( + uint64_t b, uint64_t hv, uint32_t rowBegin, + LocalTensor state) + { + if (hasInitialState_) { + DataCopy(state, initialState_[StateOffset(b, hv) + rowBegin * KDA_FWD_H_DIM], + KDA_FWD_H_STATE_SUB_ELEMS); + SetFlag(aivMte2ToVEvent_); + WaitFlag(aivMte2ToVEvent_); + } else { + Duplicate(state, 0.0f, KDA_FWD_H_STATE_SUB_ELEMS); + PipeBarrier(); + } + } + + __aicore__ inline void StoreCurrentStateAiv( + uint64_t b, uint64_t hv, uint64_t chunk, uint32_t rowBegin, + LocalTensor state, LocalTensor stateTyped) + { + if (storeH_) { + Cast(stateTyped, state, RoundMode::CAST_RINT, KDA_FWD_H_STATE_SUB_ELEMS); + PipeBarrier(); + SetFlag(aivVToMte3Event_); + WaitFlag(aivVToMte3Event_); + DataCopy(h_[HOffset(b, hv, chunk) + rowBegin * KDA_FWD_H_DIM], + stateTyped, KDA_FWD_H_STATE_SUB_ELEMS); + SetFlag(aivMte3ToVEvent_); + WaitFlag(aivMte3ToVEvent_); + } + + constexpr uint32_t columnGroup = 64; + constexpr uint32_t columnGroups = KDA_FWD_H_DIM / columnGroup; + for (uint32_t group = 0; group < columnGroups; ++group) { + Cast(stateTyped[group * KDA_FWD_H_SUB_DIM * columnGroup], + state[group * columnGroup], RoundMode::CAST_RINT, + columnGroup, KDA_FWD_H_SUB_DIM, + {static_cast(KDA_FWD_H_SUB_DIM), 1, 1, + static_cast(columnGroups * 8)}); + } + PipeBarrier(); + SetFlag(aivVToMte3Event_); + WaitFlag(aivVToMte3Event_); + LocalTensor l1State = + resource_.l1Buf.template GetBufferByByte(KDA_FWD_H_L1_H_OFFSET); + const uint32_t subBlockIdx = rowBegin / KDA_FWD_H_SUB_DIM; + WaitL1SlotFreeMte3( + KDA_FWD_H_STATE_FREE_FLAG, statePublishCount_[subBlockIdx]); + DataCopyParams stateCopyParams; + stateCopyParams.blockCount = KDA_FWD_H_DIM / 16; + stateCopyParams.blockLen = KDA_FWD_H_SUB_DIM; + stateCopyParams.srcGap = 0; + stateCopyParams.dstGap = KDA_FWD_H_DIM - KDA_FWD_H_SUB_DIM; + DataCopy(l1State[rowBegin * 16], stateTyped, stateCopyParams); + SetFlag(aivMte3ToVEvent_); + WaitFlag(aivMte3ToVEvent_); + if (!fusePostWuIntoFwdH_) { + SetL1SlotFlagAivToAic(KDA_FWD_H_STATE_READY_FLAG); + } + ++statePublishCount_[subBlockIdx]; + } + + __aicore__ inline void ProcessChunkAiv( + uint64_t b, uint64_t hv, uint32_t chunk, + uint32_t subBlockIdx, LocalTensor state, + LocalTensor stateTyped, LocalTensor direct, + LocalTensor out1, LocalTensor vnew, + LocalTensor ioTyped, LocalTensor gate) + { + const uint32_t tokenBegin = subBlockIdx * KDA_FWD_H_SUB_CHUNK; + const uint32_t stateRowBegin = subBlockIdx * KDA_FWD_H_SUB_DIM; + StoreCurrentStateAiv(b, hv, chunk, stateRowBegin, state, stateTyped); + + WaitDirectReadyAiv(); + WaitFlag(aivMte3ToMte2Event_); + if (fusePostWuIntoFwdH_) { + Cast(ioTyped, direct, RoundMode::CAST_RINT, KDA_FWD_H_TOKEN_SUB_ELEMS); + PipeBarrier(); + SetDirectFreeAiv(); + + SetFlag(aivVToMte3Event_); + WaitFlag(aivVToMte3Event_); + LocalTensor l1W = resource_.l1Buf.template GetBufferByByte( + L1WOffset(chunk & 3)); + DataCopyParams wCopyParams; + wCopyParams.blockCount = KDA_FWD_H_SUB_CHUNK; + wCopyParams.blockLen = 1; + wCopyParams.srcGap = KDA_FWD_H_DIM / 16 - 1; + wCopyParams.dstGap = 0; + for (uint32_t colBlock = 0; colBlock < KDA_FWD_H_DIM / 16; ++colBlock) { + const uint32_t dstOffset = + colBlock * KDA_FWD_H_CHUNK * 16 + tokenBegin * 16; + DataCopy(l1W[dstOffset], ioTyped[colBlock * 16], wCopyParams); + } + SetFlag(aivMte3ToVEvent_); + WaitFlag(aivMte3ToVEvent_); + SetL1SlotFlagAivToAic(KDA_FWD_H_STATE_READY_FLAG); + + WaitDirectReadyAiv(); + Cast(ioTyped, direct, RoundMode::CAST_RINT, KDA_FWD_H_TOKEN_SUB_ELEMS); + PipeBarrier(); + Cast(vnew, ioTyped, RoundMode::CAST_NONE, KDA_FWD_H_TOKEN_SUB_ELEMS); + PipeBarrier(); + SetDirectFreeAiv(); + WaitDirectReadyAiv(); + } else { + DataCopy(ioTyped, u_[ChunkMatrixOffset(b, hv, chunk, tokenBegin)], + KDA_FWD_H_TOKEN_SUB_ELEMS); + SetFlag(aivMte2ToVEvent_); + WaitFlag(aivMte2ToVEvent_); + Cast(vnew, ioTyped, RoundMode::CAST_NONE, KDA_FWD_H_TOKEN_SUB_ELEMS); + PipeBarrier(); + } + Sub(vnew, vnew, direct, KDA_FWD_H_TOKEN_SUB_ELEMS); + PipeBarrier(); + SetDirectFreeAiv(); + if (storeVNew_) { + Cast(ioTyped, vnew, RoundMode::CAST_RINT, KDA_FWD_H_TOKEN_SUB_ELEMS); + PipeBarrier(); + SetFlag(aivVToMte3Event_); + WaitFlag(aivVToMte3Event_); + DataCopy(vNew_[VNewOffset(b, hv, chunk, tokenBegin)], + ioTyped, KDA_FWD_H_TOKEN_SUB_ELEMS); + SetFlag(aivMte3ToVEvent_); + WaitFlag(aivMte3ToVEvent_); + } + + constexpr uint32_t columnGroup = 64; + constexpr uint32_t columnGroups = KDA_FWD_H_DIM / columnGroup; + for (uint32_t group = 0; group < columnGroups; ++group) { + Cast(ioTyped[group * KDA_FWD_H_SUB_CHUNK * columnGroup], + vnew[group * columnGroup], RoundMode::CAST_RINT, + columnGroup, KDA_FWD_H_SUB_CHUNK, + {static_cast(KDA_FWD_H_SUB_CHUNK), 1, 1, + static_cast(columnGroups * 8)}); + } + PipeBarrier(); + SetFlag(aivVToMte3Event_); + WaitFlag(aivVToMte3Event_); + LocalTensor l1Vnew = + resource_.l1Buf.template GetBufferByByte(KDA_FWD_H_L1_V_OFFSET); + WaitL1SlotFreeMte3( + KDA_FWD_H_VNEW_FREE_FLAG, vnewPublishCount_[subBlockIdx]); + DataCopyParams vnewL1CopyParams; + vnewL1CopyParams.blockCount = KDA_FWD_H_DIM / 16; + vnewL1CopyParams.blockLen = KDA_FWD_H_SUB_CHUNK; + vnewL1CopyParams.srcGap = 0; + vnewL1CopyParams.dstGap = KDA_FWD_H_CHUNK - KDA_FWD_H_SUB_CHUNK; + DataCopy(l1Vnew[tokenBegin * 16], ioTyped, vnewL1CopyParams); + SetFlag(aivMte3ToVEvent_); + WaitFlag(aivMte3ToVEvent_); + SetL1SlotFlagAivToAic(KDA_FWD_H_VNEW_READY_FLAG); + ++vnewPublishCount_[subBlockIdx]; + + WaitDirectReadyAiv(); + Adds(out1, direct, 0.0f, KDA_FWD_H_TOKEN_SUB_ELEMS); + PipeBarrier(); + SetDirectFreeAiv(); + + WaitDirectReadyAiv(); + DataCopy(gate, gk_[ChunkMatrixOffset(b, hv, chunk, KDA_FWD_H_CHUNK - 1) + + stateRowBegin], KDA_FWD_H_SUB_DIM); + SetFlag(aivMte2ToVEvent_); + WaitFlag(aivMte2ToVEvent_); + Muls(gate, gate, 0.6931471805599453f, KDA_FWD_H_SUB_DIM); + PipeBarrier(); + Exp(gate, gate, KDA_FWD_H_SUB_DIM); + PipeBarrier(); + AscendC::VF_CALL>( + reinterpret_cast<__ubuf__ float *>(direct.GetPhyAddr()), + reinterpret_cast<__ubuf__ float *>(state.GetPhyAddr()), + reinterpret_cast<__ubuf__ float *>(gate.GetPhyAddr()), + static_cast(KDA_FWD_H_SUB_DIM), + static_cast(KDA_FWD_H_DIM)); + PipeBarrier(); + Adds(state, direct, 0.0f, KDA_FWD_H_STATE_SUB_ELEMS); + PipeBarrier(); + + SetDirectFreeAiv(); + WaitDirectReadyAiv(); + Add(direct, direct, out1, KDA_FWD_H_TOKEN_SUB_ELEMS); + PipeBarrier(); + Cast(ioTyped, direct, RoundMode::CAST_RINT, KDA_FWD_H_TOKEN_SUB_ELEMS); + PipeBarrier(); + SetFlag(aivVToMte3Event_); + WaitFlag(aivVToMte3Event_); + CopyOutputRows( + b, hv, chunk * KDA_FWD_H_CHUNK + tokenBegin, + ioTyped, KDA_FWD_H_SUB_CHUNK); + SetDirectFreeAiv(); + SetFlag(aivMte3ToMte2Event_); + } + + __aicore__ inline void ProcessAiv() + { + const uint32_t subBlockIdx = static_cast(GetSubBlockIdx()); + const uint32_t subBlockNum = static_cast(GetSubBlockNum()); + const uint64_t coreIdx = static_cast(GetBlockIdx()) / subBlockNum; + const uint64_t coreNum = coreNum_ == 0 ? 1 : coreNum_; + LocalTensor state = + resource_.ubBuf.template GetBufferByByte(KDA_FWD_H_UB_STATE_OFFSET); + LocalTensor stateTyped = + resource_.ubBuf.template GetBufferByByte(KDA_FWD_H_UB_STATE_TYPED_OFFSET); + LocalTensor direct = + resource_.ubBuf.template GetBufferByByte(KDA_FWD_H_UB_DIRECT_OFFSET); + LocalTensor out1 = + resource_.ubBuf.template GetBufferByByte(KDA_FWD_H_UB_OUT1_OFFSET); + LocalTensor vnew = + resource_.ubBuf.template GetBufferByByte(KDA_FWD_H_UB_VNEW_OFFSET); + LocalTensor ioTyped = + resource_.ubBuf.template GetBufferByByte(KDA_FWD_H_UB_IO_OFFSET); + LocalTensor gate = + resource_.ubBuf.template GetBufferByByte(KDA_FWD_H_UB_GATE_OFFSET); + + SetFlag(aivMte3ToMte2Event_); + for (uint64_t task = coreIdx; task < batch_ * heads_; task += coreNum) { + const uint64_t b = task / heads_; + const uint64_t hv = task % heads_; + const uint32_t stateRowBegin = subBlockIdx * KDA_FWD_H_SUB_DIM; + InitializeStateAiv(b, hv, stateRowBegin, state); + SetDirectFreeAiv(); + for (uint32_t chunk = 0; + chunk < static_cast(totalChunks_); ++chunk) { + ProcessChunkAiv( + b, hv, chunk, subBlockIdx, + state, stateTyped, direct, out1, vnew, ioTyped, gate); + } + if (storeFinalState_) { + SetFlag(aivVToMte3Event_); + WaitFlag(aivVToMte3Event_); + DataCopy(finalState_[StateOffset(b, hv) + + stateRowBegin * KDA_FWD_H_DIM], + state, KDA_FWD_H_STATE_SUB_ELEMS); + SetFlag(aivMte3ToVEvent_); + WaitFlag(aivMte3ToVEvent_); + } + } + WaitFlag(aivMte3ToMte2Event_); + } + +private: + GlobalTensor gk_; + GlobalTensor initialState_; + GlobalTensor attnOut_; + GlobalTensor finalState_; + GlobalTensor aqk_; + GlobalTensor akk_; + GlobalTensor w_; + GlobalTensor u_; + GlobalTensor qgScaled_; + GlobalTensor kg_; + GlobalTensor vNew_; + GlobalTensor h_; + uint32_t statePublishCount_[2] = {0, 0}; + uint32_t vnewPublishCount_[2] = {0, 0}; + TEventID aicMte2ToMte1Event_ = KDA_FWD_H_MTE_W_EVENT; + TEventID aicL1ReuseEvents_[KDA_FWD_H_L1_STAGING_DEPTH] = {0, 1, 2, 3}; + TEventID stateL0FreeEvent_ = KDA_FWD_H_M_EVENT; + TEventID vnewL0FreeEvent_ = KDA_FWD_H_M_EVENT; + TEventID aicMte1ToMEvent_ = KDA_FWD_H_M_EVENT; + TEventID aicMToFixEvent_ = KDA_FWD_H_FIX_EVENT; + TEventID aicFixToMEvent_ = KDA_FWD_H_FIX_EVENT; + TEventID aivMte2ToVEvent_ = KDA_FWD_H_MTE_W_EVENT; + TEventID aivVToMte3Event_ = KDA_FWD_H_MTE_Q_EVENT; + TEventID aivMte3ToVEvent_ = KDA_FWD_H_MTE_B_EVENT; + TEventID aivMte3ToMte2Event_ = KDA_FWD_H_IO_REUSE_EVENT; + Catlass::Arch::Resource resource_; + uint64_t batch_ = 0; + uint64_t heads_ = 0; + uint64_t seqlen_ = 0; + uint64_t totalChunks_ = 0; + uint64_t coreNum_ = 1; + bool hasInitialState_ = false; + bool storeFinalState_ = false; + bool storeVNew_ = false; + bool storeH_ = false; + bool fusePostWuIntoFwdH_ = false; +}; + +} // namespace KdaForward::arch35 + +#endif diff --git a/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_impl.h b/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_impl.h new file mode 100644 index 000000000000..f1c698d648f0 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_impl.h @@ -0,0 +1,51 @@ +#pragma once + +#include "../chunk_kda_fwd_common.h" +#include "chunk_kda_fwd_fwd_h.h" + +namespace KdaForward::arch35 { + +template +__aicore__ inline void Run( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR g, GM_ADDR beta, + GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR attnOut, + GM_ADDR finalState, GM_ADDR gk, GM_ADDR aqk, GM_ADDR akk, + GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR userWorkspace, const TilingData &tiling, AscendC::TPipe &pipe) +{ + const auto addresses = ResolveAddresses( + finalState, gk, w, u, qg, kg, vNew, h, userWorkspace, tiling); + RunFrontEnd( + q, k, v, g, beta, aLog, dtBias, initialState, cuSeqlens, + chunkIndices, aqk, akk, addresses, userWorkspace, tiling, pipe); + if (tiling.useDenseFwdH) { + ChunkKdaFwdFwdH fwdH; + fwdH.Init( + addresses.gk, initialState, attnOut, addresses.finalState, + aqk, akk, addresses.w, addresses.u, addresses.qgScaled, + addresses.kg, addresses.vNew, addresses.h, tiling); + fwdH.Process(); + return; + } + + const int64_t fwdHTaskCount = + (tiling.isVarLen ? tiling.seqNum : tiling.batch) * tiling.vHeadNum; + const bool isolateGenericBackEnd = + (!tiling.isVarLen && tiling.seqlen % tiling.chunkSize == 0) || + fwdHTaskCount > tiling.prepareUsedCoreNum; + if (isolateGenericBackEnd) { + pipe.Destroy(); + RunGenericBackEnd( + q, k, v, beta, initialState, cuSeqlens, chunkIndices, aqk, + attnOut, addresses, userWorkspace, tiling); + } else { + RunGenericBackEnd( + q, k, v, beta, initialState, cuSeqlens, chunkIndices, aqk, + attnOut, addresses, userWorkspace, tiling, pipe); + } +} + +} // namespace KdaForward::arch35 diff --git a/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_post_wu.h b/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_post_wu.h new file mode 100644 index 000000000000..56adc06fdb22 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_post_wu.h @@ -0,0 +1,2024 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#pragma once + +#ifndef CATLASS_ARCH +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +#define CATLASS_ARCH 3510 +#else +#define CATLASS_ARCH 2201 +#endif +#endif + +#include "catlass/arch/arch.hpp" +#include "catlass/arch/cross_core_sync.hpp" +#include "catlass/arch/resource.hpp" +#include "catlass/catlass.hpp" +#include "catlass/gemm/block/block_mmad.hpp" +#include "catlass/gemm/dispatch_policy.hpp" +#include "catlass/gemm/tile/tile_copy.hpp" +#include "catlass/gemm/tile/tile_mmad.hpp" +#include "catlass/gemm_coord.hpp" +#include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp" +#include "catlass/layout/layout.hpp" +#include "kernel_operator.h" +#include "../chunk_kda_fwd_varlen.h" +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +#ifndef FLA_NPU_REGBASE_HPP_INCLUDED +#define FLA_NPU_REGBASE_HPP_INCLUDED +#include "kernel_utils/vector/regbase.hpp" +#endif +#endif +#include "tla/layout.hpp" +#include "tla/tensor.hpp" + +using namespace AscendC; + +namespace KdaPostWu { +namespace { +using KdaInt64 = tla::Int<64>; +using KdaInt128 = tla::Int<128>; +constexpr float LN2 = 0.69314718055994530942f; +constexpr float KDA_EXP2_CLAMP = 80.0f; +constexpr float KDA_EXP_INPUT_MAX = KDA_EXP2_CLAMP * LN2; +constexpr float KDA_EXP_INPUT_MIN = -KDA_EXP2_CLAMP * LN2; +constexpr float KDA_FP16_MAX = 65504.0f; +constexpr uint32_t EXP2_UB_ELEMENTS = 256; +constexpr uint32_t EXP2_UB_BYTES = EXP2_UB_ELEMENTS * (sizeof(float) + sizeof(uint16_t)); +constexpr uint32_t EXP2_EVENT_ID = 0; +constexpr uint32_t KDA_SOLVE_BT = 64; +constexpr uint32_t KDA_SOLVE_MATRIX_ELEMENTS = KDA_SOLVE_BT * KDA_SOLVE_BT; +constexpr uint32_t KDA_SOLVE_SCRATCH_X = 0; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y0 = 1; +constexpr uint32_t KDA_SOLVE_SCRATCH_TMP = 2; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y1 = 3; +constexpr uint32_t KDA_SOLVE_SCRATCH_IDENTITY = 4; +constexpr uint32_t KDA_SOLVE_SCRATCH_SLOTS = 5; +constexpr uint32_t KDA_SOLVE_DIAG_BT = 16; +constexpr uint32_t KDA_SOLVE_DIAG_BLOCKS = KDA_SOLVE_BT / KDA_SOLVE_DIAG_BT; +constexpr uint32_t KDA_SOLVE_DIAG_MCH_ITERS = 3; +constexpr uint32_t KDA_SCORE_REF_BC = 16; +constexpr uint32_t KDA_VEC_ARENA_ELEMENTS = 32768; +constexpr uint32_t KDA_BITS_PER_MASK_BYTE = 8; +constexpr uint32_t KDA_SELECT_COL_BLOCKS = 2; +constexpr uint32_t KDA_SELECT_COL_MASK_BYTES = KDA_SOLVE_MATRIX_ELEMENTS / KDA_BITS_PER_MASK_BYTE; +constexpr uint32_t KDA_SELECT_MASK_BYTES = KDA_SELECT_COL_BLOCKS * KDA_SELECT_COL_MASK_BYTES; +constexpr uint32_t KDA_SELECT_AQK_MASK_BYTE_OFFSET = 120 * 1024; +constexpr uint32_t KDA_SELECT_AKK_MASK_BYTE_OFFSET = KDA_SELECT_AQK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_BYTE_OFFSET = KDA_SELECT_AKK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_FLOAT_OFFSET = KDA_SELECT_ZERO_BYTE_OFFSET / sizeof(float); +constexpr uint8_t KDA_SCORE_DONE_FLAG0 = 2; +constexpr uint8_t KDA_SCORE_DONE_FLAG1 = 3; +constexpr uint8_t KDA_SCORE_READY_FLAG0 = 4; +constexpr uint8_t KDA_SCORE_READY_FLAG1 = 5; +constexpr uint32_t KDA_SCORE_QUEUE_DEPTH = 2; +constexpr uint32_t KDA_SYNC_REVERSE_DEPTH = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_PLANES = 3; +constexpr uint32_t KDA_SCORE_SCRATCH_QG = 0; +constexpr uint32_t KDA_SCORE_SCRATCH_W = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_KG = 2; +constexpr uint64_t KDA_WORKSPACE_ALIGN = 512; +constexpr uint32_t KDA_GATE_TILE_ROWS = 32; +constexpr uint32_t KDA_TYPICAL_GATE_TILE_ROWS = 16; +constexpr uint32_t KDA_TYPICAL_GATE_PIPELINE_ROWS = 32; +constexpr uint16_t KDA_TYPICAL_GATE_PIPELINE_STAGES = 3; +constexpr uint32_t KDA_POST_EVENT = 3; +constexpr uint32_t KDA_POST_EVENT_NEXT = 4; +constexpr uint32_t KDA_POST_EVENT_FIX = 5; +constexpr uint32_t KDA_POST_PIPELINE_L1_SLOT_BYTES = 24 * 1024; +constexpr uint32_t KDA_POST_PIPELINE_L1_A_BYTES = 64 * 64 * sizeof(uint16_t); +constexpr uint32_t KDA_POST_PIPELINE_L1_B_BYTES = 64 * 128 * sizeof(uint16_t); +constexpr uint32_t KDA_POST_PIPELINE_L1_U_SLOT_BYTES = 64 * 128 * sizeof(uint16_t); +constexpr uint32_t KDA_POST_PIPELINE_L0_A_SLOT_BYTES = 64 * 64 * sizeof(uint16_t); +constexpr uint32_t KDA_POST_PIPELINE_L0_B_SLOT_BYTES = 64 * 256 * sizeof(uint16_t); +constexpr uint32_t KDA_POST_PIPELINE_L0_C_SLOT_BYTES = 64 * 256 * sizeof(float); +constexpr uint16_t KDA_POST_PIPELINE_STAGE_COUNT = 2; +constexpr uint16_t KDA_POST_FUSED_BATCH_TASKS = 4; +constexpr uint16_t KDA_POST_HEAD_PAIR_LANES = 2; +constexpr uint16_t KDA_POST_PIPELINE_U_EVENT = KDA_POST_EVENT_FIX; +constexpr bool KDA_ENABLE_POST_AIC_PIPELINE = true; + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +template +__simd_callee__ inline void LoadPostKdaGateRegbasePair( + AscendC::MicroAPI::RegTensor &zeroReg, + AscendC::MicroAPI::RegTensor &oneReg, + __ubuf__ InputT *src, + AscendC::MicroAPI::MaskReg &inputMask) +{ + using namespace AscendC::MicroAPI; + if constexpr (std::is_same()) { + LoadAlign(zeroReg, oneReg, src); + } else { + RegTensor inputReg; + LoadIn(inputReg, src); + CastHalf2Float(zeroReg, oneReg, inputReg, inputMask); + } +} + +template +__simd_callee__ inline void StorePostKdaGateRegbasePair( + __ubuf__ OutputT *dst, + AscendC::MicroAPI::RegTensor &zeroReg, + AscendC::MicroAPI::RegTensor &oneReg, + AscendC::MicroAPI::MaskReg &inputMask, + AscendC::MicroAPI::MaskReg &floatMask) +{ + using namespace AscendC::MicroAPI; + if constexpr (std::is_same()) { + Mins(zeroReg, zeroReg, KDA_FP16_MAX, floatMask); + Mins(oneReg, oneReg, KDA_FP16_MAX, floatMask); + Maxs(zeroReg, zeroReg, -KDA_FP16_MAX, floatMask); + Maxs(oneReg, oneReg, -KDA_FP16_MAX, floatMask); + } + RegTensor outputReg; + CastFloat2Half(outputReg, zeroReg, oneReg, floatMask); + StoreAlign(dst, outputReg, inputMask); +} + +template +static __simd_vf__ inline void ComputePostKdaKgRegbase( + __ubuf__ T *kAndKg, __ubuf__ GK_T *gate, __ubuf__ float *ref, + uint16_t rows, uint16_t cols) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t ELEMENTS_PER_REG = AscendC::VECTOR_REG_WIDTH / sizeof(T); + MaskReg floatMask = CreateMask(); + for (uint16_t row = 0; row < rows; ++row) { + uint32_t rowOffset = static_cast(row) * cols; + for (uint16_t col = 0; col < cols; col += ELEMENTS_PER_REG) { + uint32_t activeCount = static_cast(cols - col); + MaskReg inputMask = UpdateMask(activeCount); + uint32_t offset = rowOffset + col; + + RegTensor gateZeroReg; + RegTensor gateOneReg; + RegTensor refZeroReg; + RegTensor refOneReg; + RegTensor expZeroReg; + RegTensor expOneReg; + RegTensor inputZeroReg; + RegTensor inputOneReg; + RegTensor outputZeroReg; + RegTensor outputOneReg; + + LoadPostKdaGateRegbasePair( + gateZeroReg, gateOneReg, gate + offset, inputMask); + LoadAlign( + refZeroReg, refOneReg, ref + col); + SubFloatTwoReg(expZeroReg, expOneReg, refZeroReg, refOneReg, + gateZeroReg, gateOneReg, floatMask); + Muls(expZeroReg, expZeroReg, LN2, floatMask); + Muls(expOneReg, expOneReg, LN2, floatMask); + MinsFloatTwoReg(expZeroReg, expOneReg, expZeroReg, expOneReg, + KDA_EXP_INPUT_MAX, floatMask); + Maxs(expZeroReg, expZeroReg, KDA_EXP_INPUT_MIN, floatMask); + Maxs(expOneReg, expOneReg, KDA_EXP_INPUT_MIN, floatMask); + ExpFloatTwoReg(expZeroReg, expOneReg, expZeroReg, expOneReg, floatMask); + + LoadPostKdaGateRegbasePair( + inputZeroReg, inputOneReg, kAndKg + offset, inputMask); + MulFloatTwoReg(outputZeroReg, outputOneReg, inputZeroReg, inputOneReg, + expZeroReg, expOneReg, floatMask); + StorePostKdaGateRegbasePair( + kAndKg + offset, outputZeroReg, outputOneReg, inputMask, floatMask); + } + } +} +#endif + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +using KdaArchTag = Catlass::Arch::Ascend950; +#else +using KdaArchTag = Catlass::Arch::AtlasA2; +#endif +using KdaDispatchPolicy = Catlass::Gemm::MmadPingpong; +using KdaScoreDispatchPolicy = + Catlass::Gemm::MmadPingpongTlaMulti; +static_assert(KdaScoreDispatchPolicy::ENABLE_L1_RESIDENT, + "KDA Aqk/Akk score MMAD must keep the shared right matrix resident in L1"); +static_assert(KdaScoreDispatchPolicy::L1B_STAGES == 1, + "KDA Aqk/Akk score MMAD needs one L1 B slot so the second MMAD reuses it"); +using KdaSolveDispatchPolicy = Catlass::Gemm::MmadPingpong; +static_assert(!KdaSolveDispatchPolicy::USE_HF32_MODE, "KDA triangular solve must use IEEE FP32 Cube mode"); +using KdaL1TileShape = tla::Shape; +using KdaL0TileShape = KdaL1TileShape; +using KdaSolveL1TileShape = tla::Shape; +using KdaSolveL0TileShape = KdaSolveL1TileShape; + +__aicore__ inline uint32_t FloatToBits(float value) +{ + union Bits { + __aicore__ Bits() {} + float f; + uint32_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline float BitsToFloat(uint32_t value) +{ + union Bits { + __aicore__ Bits() {} + uint32_t u; + float f; + } bits; + bits.u = value; + return bits.f; +} + +__aicore__ inline uint16_t Bf16ToBits(bfloat16_t value) +{ + union Bits { + __aicore__ Bits() {} + bfloat16_t f; + uint16_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline bfloat16_t BitsToBf16(uint16_t value) +{ + union Bits { + __aicore__ Bits() {} + uint16_t u; + bfloat16_t f; + } bits; + bits.u = value; + return bits.f; +} + +template +__aicore__ inline T FloatToType(float value) +{ + if constexpr (IsSameType::value) { + uint32_t bits = FloatToBits(value); + uint32_t bias = 0x7FFFu + ((bits >> 16) & 1u); + return BitsToBf16(static_cast((bits + bias) >> 16)); + } + return static_cast(value); +} + +template +class ChunkKdaFwdPostWuKernel { +public: + using OUT_T = T; + using AKK_T = T; + template + __aicore__ inline void Init(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR preparedQG, GM_ADDR preparedAqk, + GM_ADDR propagatedVNew, GM_ADDR propagatedH, GM_ADDR o, GM_ADDR finalState, GM_ADDR aqk, + GM_ADDR akk, GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR workspace, const TilingData &tiling, TPipe *pipe, + bool initVecBuffers = true) + { + pipe_ = pipe; + q_.SetGlobalBuffer((__gm__ T *)q); + k_.SetGlobalBuffer((__gm__ T *)k); + v_.SetGlobalBuffer((__gm__ T *)v); + gk_.SetGlobalBuffer((__gm__ GK_T *)gk); + beta_.SetGlobalBuffer((__gm__ BETA_T *)beta); + if (initialState != nullptr) { + initialState_.SetGlobalBuffer((__gm__ float *)initialState); + } + cuSeqlensAddr_ = reinterpret_cast<__gm__ int64_t *>(cuSeqlens); + if (preparedQG != nullptr) { + preparedQG_.SetGlobalBuffer((__gm__ T *)preparedQG); + } + if (preparedAqk != nullptr) { + preparedAqk_.SetGlobalBuffer((__gm__ T *)preparedAqk); + } + if (propagatedVNew != nullptr) { + propagatedVNew_.SetGlobalBuffer((__gm__ T *)propagatedVNew); + } + if (propagatedH != nullptr) { + propagatedH_.SetGlobalBuffer((__gm__ T *)propagatedH); + } + chunkIndicesAddr_ = reinterpret_cast<__gm__ int64_t *>(chunkIndices); + o_.SetGlobalBuffer((__gm__ OUT_T *)o); + finalState_.SetGlobalBuffer((__gm__ float *)finalState); + aqk_.SetGlobalBuffer((__gm__ float *)aqk); + akk_.SetGlobalBuffer((__gm__ AKK_T *)akk); + w_.SetGlobalBuffer((__gm__ T *)w); + u_.SetGlobalBuffer((__gm__ OUT_T *)u); + qg_.SetGlobalBuffer((__gm__ T *)qg); + kg_.SetGlobalBuffer((__gm__ T *)kg); + vNew_.SetGlobalBuffer((__gm__ T *)vNew); + h_.SetGlobalBuffer((__gm__ float *)h); + solveWorkspace_.SetGlobalBuffer((__gm__ float *)workspace); + + B_ = tiling.batch; + N_ = tiling.seqNum; + H_ = tiling.qHeadNum; + HV_ = tiling.vHeadNum; + T_ = tiling.seqlen; + K_ = tiling.kHeadDim; + V_ = tiling.vHeadDim; + BT_ = tiling.chunkSize; + NT_ = tiling.totalChunks; + scale_ = tiling.scale; + hasInitial_ = tiling.hasInitialState; + isVarLen_ = tiling.isVarLen; + inputSequenceMajor_ = tiling.inputSequenceMajor; + usedCoreNum_ = tiling.postWuUsedCoreNum; + if ASCEND_IS_AIV { + uint64_t subBlockNum = static_cast(GetSubBlockNum()); + solveCoreIdx_ = subBlockNum == 0 ? 0 : static_cast(GetBlockIdx()) / subBlockNum; + } else { + solveCoreIdx_ = static_cast(GetBlockIdx()); + } + if (pipe_ != nullptr && initVecBuffers) { + pipe_->InitBuffer(exp2Buf_, EXP2_UB_BYTES); + pipe_->InitBuffer(vecBuf_, KDA_VEC_ARENA_ELEMENTS * sizeof(float)); + const uint64_t gateWritebackRows = + ScoreVectorMaxRows(5 * sizeof(float) + 2 * sizeof(T) + sizeof(GK_T)); + pipe_->InitBuffer(gateWritebackBuf_, + static_cast(gateWritebackRows * K_ * + (3 * sizeof(T) + sizeof(GK_T)))); + AllocVectorEvents(); + } + } + __aicore__ inline void ProcessAiv() + { + ProcessPostAiv(); + ReleaseVectorEvents(); + } + + __aicore__ inline void ProcessAic() + { + ProcessPostAic(); + } + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + __aicore__ inline void ProcessPreparedFullHeadPairBatchArch35( + const uint64_t *batchB, const uint64_t *batchHvBase, + const uint64_t *batchStart, uint16_t taskCount) + { + if (taskCount == 0) { + return; + } + SetLoadDataPaddingValue(static_cast(0)); + Catlass::Arch::Resource resource; + const uint16_t itemCount = taskCount * KDA_POST_HEAD_PAIR_LANES; + uint16_t slot = 0; + uint16_t usedSlotCount = 1; + uint16_t taskIdx = 0; + uint16_t lane = 0; + uint64_t b = batchB[taskIdx]; + uint64_t hv = batchHvBase[taskIdx] + lane; + uint64_t start = batchStart[taskIdx]; + InitializePostWuPipelineEvents(); + PrefetchPostWuPipelineArch35(resource, slot, b, hv, start, BT_, false); + PrefetchPostWuPipelineU(resource, slot, b, hv, start, BT_, false); + + for (uint16_t item = 0; item < itemCount; ++item) { + const uint16_t nextItem = item + 1; + if (nextItem < itemCount) { + const uint16_t nextTaskIdx = nextItem / KDA_POST_HEAD_PAIR_LANES; + const uint16_t nextLane = nextItem % KDA_POST_HEAD_PAIR_LANES; + const uint16_t nextSlot = slot ^ 1; + const bool reuseSlot = nextItem >= KDA_POST_PIPELINE_STAGE_COUNT; + PrefetchPostWuPipelineArch35( + resource, nextSlot, batchB[nextTaskIdx], + batchHvBase[nextTaskIdx] + nextLane, batchStart[nextTaskIdx], BT_, reuseSlot); + PrefetchPostWuPipelineU( + resource, nextSlot, batchB[nextTaskIdx], + batchHvBase[nextTaskIdx] + nextLane, batchStart[nextTaskIdx], BT_, reuseSlot); + if (!reuseSlot) { + ++usedSlotCount; + } + } + + ComputePrefetchedPostWuPipelineArch35(resource, slot, b, hv, start, BT_); + if (nextItem < itemCount) { + taskIdx = nextItem / KDA_POST_HEAD_PAIR_LANES; + lane = nextItem % KDA_POST_HEAD_PAIR_LANES; + b = batchB[taskIdx]; + hv = batchHvBase[taskIdx] + lane; + start = batchStart[taskIdx]; + slot ^= 1; + } + } + FinalizePostWuPipelineEvents(usedSlotCount); + } + + __aicore__ inline void ProcessPreparedTailHeadPairArch35( + uint64_t b, uint64_t hvBase, uint64_t start, uint64_t curT) + { + SetLoadDataPaddingValue(static_cast(0)); + Catlass::Arch::Resource resource; + InitializePostWuPipelineEvents(); + for (uint16_t lane = 0; lane < KDA_POST_HEAD_PAIR_LANES; ++lane) { + PrefetchPostWuPipelineArch35( + resource, lane, b, hvBase + lane, start, curT, false); + PrefetchPostWuPipelineU( + resource, lane, b, hvBase + lane, start, curT, false); + } + for (uint16_t lane = 0; lane < KDA_POST_HEAD_PAIR_LANES; ++lane) { + ComputePrefetchedPostWuPipelineArch35( + resource, lane, b, hvBase + lane, start, curT); + } + FinalizePostWuPipelineEvents(KDA_POST_HEAD_PAIR_LANES); + } + + __aicore__ inline void ProcessPreparedTailSingleArch35( + uint64_t b, uint64_t hv, uint64_t start, uint64_t curT) + { + SetLoadDataPaddingValue(static_cast(0)); + Catlass::Arch::Resource resource; + InitializePostWuPipelineSlot(0); + PrefetchPostWuPipelineArch35(resource, 0, b, hv, start, curT, false); + PrefetchPostWuPipelineU(resource, 0, b, hv, start, curT, false); + ComputePrefetchedPostWuPipelineArch35(resource, 0, b, hv, start, curT); + FinalizePostWuPipelineEvents(1); + } + + __aicore__ inline void ProcessPreparedHeadPairBatchArch35( + const uint64_t *batchB, const uint64_t *batchHvBase, + const uint64_t *batchStart, const uint64_t *batchEnd, uint16_t taskCount) + { + uint16_t fullRunBegin = 0; + for (uint16_t task = 0; task < taskCount; ++task) { + if (batchEnd[task] - batchStart[task] == BT_) { + continue; + } + if (task > fullRunBegin) { + ProcessPreparedFullHeadPairBatchArch35( + batchB + fullRunBegin, batchHvBase + fullRunBegin, + batchStart + fullRunBegin, task - fullRunBegin); + } + ProcessPreparedTailHeadPairArch35( + batchB[task], batchHvBase[task], batchStart[task], + batchEnd[task] - batchStart[task]); + fullRunBegin = task + 1; + } + if (fullRunBegin < taskCount) { + ProcessPreparedFullHeadPairBatchArch35( + batchB + fullRunBegin, batchHvBase + fullRunBegin, + batchStart + fullRunBegin, taskCount - fullRunBegin); + } + } +#endif + +private: + __aicore__ inline void AllocVectorEvents() + { + mte2ToVEvent_ = pipe_->AllocEventID(); + vToMte2Event_ = pipe_->AllocEventID(); + vToMte3Event_ = pipe_->AllocEventID(); + mte3ToVEvent_ = pipe_->AllocEventID(); + mte2ToMte3Event_ = pipe_->AllocEventID(); + mte3ToMte2Event_ = pipe_->AllocEventID(); + vectorEventsAllocated_ = true; + } + + __aicore__ inline void ReleaseVectorEvents() + { + if (!vectorEventsAllocated_) { + return; + } + pipe_->ReleaseEventID(mte2ToVEvent_); + pipe_->ReleaseEventID(vToMte2Event_); + pipe_->ReleaseEventID(vToMte3Event_); + pipe_->ReleaseEventID(mte3ToVEvent_); + pipe_->ReleaseEventID(mte2ToMte3Event_); + pipe_->ReleaseEventID(mte3ToMte2Event_); + vectorEventsAllocated_ = false; + } + + __aicore__ inline uint64_t QOffset(uint64_t b, uint64_t h, uint64_t t, uint64_t d) const + { + if (inputSequenceMajor_) { + return ((b * T_ + t) * H_ + h) * K_ + d; + } + return ((b * H_ + h) * T_ + t) * K_ + d; + } + + __aicore__ inline uint64_t KVOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t d, uint64_t dim) const + { + return ((b * HV_ + hv) * T_ + t) * dim + d; + } + + __aicore__ inline uint64_t BetaOffset(uint64_t b, uint64_t hv, uint64_t t) const + { + return (b * HV_ + hv) * T_ + t; + } + + __aicore__ inline uint64_t AOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t j) const + { + return ((b * HV_ + hv) * T_ + t) * BT_ + j; + } + + __aicore__ inline uint64_t HOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t d, uint64_t r) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * K_ + d) * V_ + r; + } + + __aicore__ inline uint64_t WScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t t, uint64_t d) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * BT_ + t) * K_ + d; + } + + __aicore__ inline uint64_t SolveScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t slot) const + { + (void)b; + (void)hv; + (void)chunkIdx; + uint64_t matrixElements = BT_ * BT_; + return solveCoreIdx_ * KDA_SOLVE_SCRATCH_SLOTS * matrixElements + slot * matrixElements; + } + + __aicore__ inline uint64_t ScoreScratchOffset(uint64_t slot, uint64_t plane, uint64_t t = 0, + uint64_t d = 0) const + { + return (((solveCoreIdx_ * KDA_SCORE_QUEUE_DEPTH + slot) * KDA_SCORE_SCRATCH_PLANES + plane) * BT_ + t) * + K_ + + d; + } + + + + __aicore__ inline uint64_t ScoreRefBlockSize() const + { + return KDA_SCORE_REF_BC; + } + + __aicore__ inline uint64_t ScoreRowBlockCount(uint64_t curT, uint64_t rowBegin) const + { + uint64_t blockSize = ScoreRefBlockSize(); + uint64_t rowCount = curT - rowBegin; + if (rowCount > blockSize) { + rowCount = blockSize; + } + return rowCount; + } + + __aicore__ inline uint64_t ScoreRefToken(uint64_t start, uint64_t curT, uint64_t rowBegin, + uint64_t rowCount) const + { + uint64_t ref = rowBegin + rowCount / 2; + if (ref >= curT) { + ref = curT - 1; + } + return start + ref; + } + + __aicore__ inline void RunExp2(LocalTensor &tensor, uint32_t count) + { + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + ClampExpInput(tensor, count); + Exp(tensor, tensor, count); + PipeBarrier(); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + } + + __aicore__ inline void ClampExpInput(LocalTensor &tensor, uint32_t count) + { + Mins(tensor, tensor, KDA_EXP_INPUT_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, KDA_EXP_INPUT_MIN, count); + PipeBarrier(); + } + + __aicore__ inline void ClampFp32ToOutputType(LocalTensor &tensor, uint32_t count) + { + if constexpr (IsSameType::value) { + Mins(tensor, tensor, KDA_FP16_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, -KDA_FP16_MAX, count); + PipeBarrier(); + } + } + + template + __aicore__ inline void CopyVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst, src[offset], static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + + template + __aicore__ inline void CopyRowsIn(LocalTensor &dst, GlobalTensor &src, + uint64_t offset, uint64_t rows, uint64_t cols, + uint64_t rowStride) + { + if (rows == 0 || cols == 0) { + return; + } + if (rowStride == cols) { + CopyVectorIn(dst, src, offset, rows * cols); + return; + } + DataCopyExtParams params{ + static_cast(rows), + static_cast(cols * sizeof(CopyT)), + static_cast((rowStride - cols) * sizeof(CopyT)), + 0, + 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + + template + __aicore__ inline void CopyVectorOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst[offset], src, static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPad(dst[offset], src, params); + } + + template + __aicore__ inline void CopyRowIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset) + { + CopyVectorIn(dst, src, offset, K_); + } + + template + __aicore__ inline void CopyRowOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src) + { + CopyVectorOut(dst, offset, src, K_); + } + + __aicore__ inline LocalTensor VecScratch(uint64_t slot) + { + return vecBuf_.Get()[slot * EXP2_UB_ELEMENTS]; + } + + template + __aicore__ inline void LoadAsFloatRow(GlobalTensor &src, uint64_t srcOffset, LocalTensor &dst, + uint64_t count) + { + if constexpr (IsSameType::value) { + CopyVectorIn(dst, src, srcOffset, count); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Adds(dst, dst, 0.0f, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } else { + constexpr uint32_t typedOffset = EXP2_UB_ELEMENTS * sizeof(float) / sizeof(CopyT); + LocalTensor rowLocal = exp2Buf_.Get()[typedOffset]; + CopyVectorIn(rowLocal, src, srcOffset, count); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast(dst, rowLocal, RoundMode::CAST_NONE, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } + PipeBarrier(); + } + + template + __aicore__ inline void LoadAsFloatVector(GlobalTensor &src, uint64_t srcOffset, + LocalTensor &dst, LocalTensor &typedScratch, + uint64_t count) + { + if constexpr (IsSameType::value) { + CopyVectorIn(dst, src, srcOffset, count); + } else { + CopyVectorIn(typedScratch, src, srcOffset, count); + } + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + if constexpr (!IsSameType::value) { + Cast(dst, typedScratch, RoundMode::CAST_NONE, static_cast(count)); + PipeBarrier(); + } + } + + template + __aicore__ inline void StoreFloatRow(GlobalTensor &dst, uint64_t dstOffset, LocalTensor &src, + uint64_t count) + { + if constexpr (IsSameType::value) { + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(dst, dstOffset, src, count); + } else { + constexpr uint32_t typedOffset = EXP2_UB_ELEMENTS * sizeof(float) / sizeof(CopyT); + LocalTensor rowLocal = exp2Buf_.Get()[typedOffset]; + Cast(rowLocal, src, RoundMode::CAST_RINT, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(dst, dstOffset, rowLocal, count); + } + SetFlag(mte3ToMte2Event_); + WaitFlag(mte3ToMte2Event_); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + + + + + __aicore__ inline LocalTensor Exp2NegG(uint64_t b, uint64_t hv, uint64_t t) + { + LocalTensor exp2Local = exp2Buf_.Get(); + LoadAsFloatRow(gk_, KVOffset(b, hv, t, 0, K_), exp2Local, K_); + Muls(exp2Local, exp2Local, -LN2, static_cast(K_)); + PipeBarrier(); + RunExp2(exp2Local, static_cast(K_)); + return exp2Local; + } + + + __aicore__ inline uint64_t ScoreVectorMaxRows(uint64_t bytesPerElem) const + { + constexpr uint64_t arenaBytes = static_cast(KDA_VEC_ARENA_ELEMENTS) * sizeof(float); + uint64_t maxRows = (arenaBytes / bytesPerElem) / K_; + if (K_ >= 128 && maxRows > 32) { + maxRows = 32; + } + return maxRows; + } + __aicore__ inline bool UsePostWuCube(uint64_t curT) const + { + return curT > 0 && curT <= BT_ && (BT_ == 64 || BT_ == 128) && K_ >= 16 && V_ >= 16 && + V_ <= 256 && K_ % 16 == 0 && V_ % 16 == 0; + } + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + __aicore__ inline bool UsePostWuCubeArch35(uint64_t curT) const + { + return curT == 64 && BT_ == 64 && K_ == 128 && V_ == 128; + } + + __aicore__ inline bool UseFullPostWuPipelineArch35(uint64_t curT) const + { + return curT == 64 && BT_ == 64 && K_ == 128 && V_ == 128; + } + + __aicore__ inline void ComputePostWuCubeFusedArch35( + uint64_t b, uint64_t hv, uint64_t start, uint64_t curT) + { + using ElementA = T; + using ElementB = T; + using ElementC = T; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla< + KdaArchTag, ElementA, LayoutTagA, ElementB, LayoutTagB, ElementC, LayoutTagC>; + + constexpr uint32_t capacityM = 64; + constexpr uint32_t n = 128; + constexpr uint32_t capacityK = 64; + const uint32_t m = static_cast(curT); + const uint32_t k = static_cast(curT); + SetLoadDataPaddingValue(static_cast(0)); + + auto layoutA = tla::MakeLayout(capacityM, capacityK); + auto layoutB = tla::MakeLayout(capacityK, n); + auto layoutC = tla::MakeLayout(capacityM, n); + auto tensorA = tla::MakeTensor( + preparedAqk_[AOffset(b, hv, start, 0)], layoutA, Catlass::Arch::PositionGM{}); + auto tensorW = tla::MakeTensor( + preparedQG_[KVOffset(b, hv, start, 0, K_)], layoutB, Catlass::Arch::PositionGM{}); + auto tensorU = tla::MakeTensor( + propagatedVNew_[KVOffset(b, hv, start, 0, V_)], layoutB, Catlass::Arch::PositionGM{}); + auto tensorWOut = tla::MakeTensor( + w_[KVOffset(b, hv, start, 0, K_)], layoutC, Catlass::Arch::PositionGM{}); + auto tensorUOut = tla::MakeTensor( + u_[KVOffset(b, hv, start, 0, V_)], layoutC, Catlass::Arch::PositionGM{}); + auto blockA = GetTile(tensorA, tla::MakeCoord(0, 0), tla::MakeShape(m, k)); + auto blockW = GetTile(tensorW, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto blockU = GetTile(tensorU, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto blockWOut = GetTile(tensorWOut, tla::MakeCoord(0, 0), tla::MakeShape(m, n)); + auto blockUOut = GetTile(tensorUOut, tla::MakeCoord(0, 0), tla::MakeShape(m, n)); + + Catlass::Arch::Resource resource; + constexpr uint32_t aBytes = capacityM * capacityK * sizeof(ElementA); + constexpr uint32_t bBytes = capacityK * n * sizeof(ElementB); + LocalTensor l1A = resource.l1Buf.template GetBufferByByte(0); + LocalTensor l1B0 = resource.l1Buf.template GetBufferByByte(aBytes); + LocalTensor l1B1 = resource.l1Buf.template GetBufferByByte(aBytes + bBytes); + LocalTensor l0A = resource.l0ABuf.template GetBufferByByte(0); + LocalTensor l0B = resource.l0BBuf.template GetBufferByByte(0); + LocalTensor l0C = resource.l0CBuf.template GetBufferByByte(0); + + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + using LayoutTagL0A = typename TileCopy::LayoutTagL0A; + using LayoutTagL0B = typename TileCopy::LayoutTagL0B; + using CopyGmToL1A = typename TileCopy::template CopyGmToL1A; + using CopyGmToL1B = typename TileCopy::template CopyGmToL1B; + using CopyL1ToL0A = typename TileCopy::CopyL1ToL0A; + using CopyL1ToL0B = typename TileCopy::CopyL1ToL0B; + using CopyL0CToDst = typename TileCopy::template CopyL0CToDst; + using TileMmad = + Catlass::Gemm::Tile::TileMmadTla; + + auto layoutL1A = tla::MakeLayout(capacityM, capacityK); + auto layoutL1B = tla::MakeLayout(capacityK, n); + auto layoutL0A = tla::MakeLayout(capacityM, capacityK); + auto layoutL0B = tla::MakeLayout(capacityK, n); + auto layoutL0C = tla::MakeLayoutL0C(capacityM, n); + auto tensorL1A = tla::MakeTensor(l1A, layoutL1A, Catlass::Arch::PositionL1{}); + auto tensorL1B0 = tla::MakeTensor(l1B0, layoutL1B, Catlass::Arch::PositionL1{}); + auto tensorL1B1 = tla::MakeTensor(l1B1, layoutL1B, Catlass::Arch::PositionL1{}); + auto tensorL0A = tla::MakeTensor(l0A, layoutL0A, Catlass::Arch::PositionL0A{}); + auto tensorL0B = tla::MakeTensor(l0B, layoutL0B, Catlass::Arch::PositionL0B{}); + auto tensorL0C = tla::MakeTensor(l0C, layoutL0C, Catlass::Arch::PositionL0C{}); + auto tileL1A = GetTile(tensorL1A, tla::MakeCoord(0, 0), tla::MakeShape(m, k)); + auto tileL1B0 = GetTile(tensorL1B0, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto tileL1B1 = GetTile(tensorL1B1, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto tileL0A = GetTile(tensorL0A, tla::MakeCoord(0, 0), tla::MakeShape(m, k)); + auto tileL0B = GetTile(tensorL0B, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto tileL0C = GetTile(tensorL0C, tla::MakeCoord(0, 0), tla::MakeShape(m, n)); + + CopyGmToL1A copyGmToL1A; + CopyGmToL1B copyGmToL1B; + CopyL1ToL0A copyL1ToL0A; + CopyL1ToL0B copyL1ToL0B; + CopyL0CToDst copyL0CToDst; + TileMmad tileMmad; + + copyGmToL1A(tensorL1A, blockA); + copyGmToL1B(tensorL1B0, blockW); + SetFlag(KDA_POST_EVENT); + WaitFlag(KDA_POST_EVENT); + copyL1ToL0A(tileL0A, tileL1A); + copyL1ToL0B(tileL0B, tileL1B0); + SetFlag(KDA_POST_EVENT); + copyGmToL1B(tensorL1B1, blockU); + SetFlag(KDA_POST_EVENT_NEXT); + WaitFlag(KDA_POST_EVENT); + tileMmad(tileL0C, tileL0A, tileL0B, m, n, k, true, 0); + SetFlag(KDA_POST_EVENT); + SetFlag(KDA_POST_EVENT); + WaitFlag(KDA_POST_EVENT_NEXT); + WaitFlag(KDA_POST_EVENT); + WaitFlag(KDA_POST_EVENT); + copyL0CToDst(blockWOut, tileL0C); + SetFlag(KDA_POST_EVENT_FIX); + + copyL1ToL0B(tileL0B, tileL1B1); + SetFlag(KDA_POST_EVENT_NEXT); + WaitFlag(KDA_POST_EVENT_FIX); + WaitFlag(KDA_POST_EVENT_NEXT); + tileMmad(tileL0C, tileL0A, tileL0B, m, n, k, true, 0); + SetFlag(KDA_POST_EVENT); + WaitFlag(KDA_POST_EVENT); + copyL0CToDst(blockUOut, tileL0C); + SetFlag(KDA_POST_EVENT_FIX); + WaitFlag(KDA_POST_EVENT_FIX); + } + + __aicore__ inline void FinalizePostWuPipelineEvents(uint16_t usedSlotCount) + { + for (uint16_t slot = 0; slot < usedSlotCount; ++slot) { + WaitFlag(KDA_POST_EVENT + slot); + WaitFlag(KDA_POST_PIPELINE_U_EVENT + slot); + WaitFlag(KDA_POST_EVENT + slot); + WaitFlag(KDA_POST_EVENT + slot); + } + } + + __aicore__ inline void InitializePostWuPipelineEvents() + { + for (uint16_t slot = 0; slot < KDA_POST_PIPELINE_STAGE_COUNT; ++slot) { + InitializePostWuPipelineSlot(slot); + } + } + + __aicore__ inline void InitializePostWuPipelineSlot(uint16_t slot) + { + SetFlag(KDA_POST_EVENT + slot); + SetFlag(KDA_POST_EVENT + slot); + } + + __aicore__ inline void PrefetchPostWuPipelineArch35( + Catlass::Arch::Resource &resource, uint16_t slot, + uint64_t b, uint64_t hv, uint64_t start, uint64_t curT, bool reuseSlot) + { + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla< + KdaArchTag, T, LayoutTagA, T, LayoutTagB, T, LayoutTagC>; + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + + constexpr uint32_t capacityM = 64; + constexpr uint32_t n = 128; + constexpr uint32_t capacityK = 64; + const uint32_t m = static_cast(curT); + const uint32_t k = static_cast(curT); + auto layoutA = tla::MakeLayout(capacityM, capacityK); + auto layoutB = tla::MakeLayout(capacityK, n); + auto tensorA = tla::MakeTensor( + preparedAqk_[AOffset(b, hv, start, 0)], layoutA, Catlass::Arch::PositionGM{}); + auto tensorW = tla::MakeTensor( + preparedQG_[KVOffset(b, hv, start, 0, K_)], layoutB, Catlass::Arch::PositionGM{}); + auto blockA = GetTile(tensorA, tla::MakeCoord(0, 0), tla::MakeShape(m, k)); + auto blockW = GetTile(tensorW, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + using CopyGmToL1A = typename TileCopy::template CopyGmToL1A; + using CopyGmToL1B = typename TileCopy::template CopyGmToL1B; + CopyGmToL1A copyGmToL1A; + CopyGmToL1B copyGmToL1B; + + uint32_t slotBase = static_cast(slot) * KDA_POST_PIPELINE_L1_SLOT_BYTES; + LocalTensor l1A = resource.l1Buf.template GetBufferByByte(slotBase); + LocalTensor l1W = resource.l1Buf.template GetBufferByByte( + slotBase + KDA_POST_PIPELINE_L1_A_BYTES); + auto layoutL1A = tla::MakeLayout(capacityM, capacityK); + auto layoutL1B = tla::MakeLayout(capacityK, n); + auto tensorL1A = tla::MakeTensor(l1A, layoutL1A, Catlass::Arch::PositionL1{}); + auto tensorL1W = tla::MakeTensor(l1W, layoutL1B, Catlass::Arch::PositionL1{}); + + uint16_t pipelineEvent = KDA_POST_EVENT + slot; + if (reuseSlot) { + WaitFlag(pipelineEvent); + } + copyGmToL1A(tensorL1A, blockA); + copyGmToL1B(tensorL1W, blockW); + SetFlag(pipelineEvent); + } + + __aicore__ inline void PrefetchPostWuPipelineU( + Catlass::Arch::Resource &resource, uint16_t slot, + uint64_t b, uint64_t hv, uint64_t start, uint64_t curT, bool reuseStage) + { + using LayoutTagB = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla< + KdaArchTag, T, Catlass::layout::RowMajor, T, LayoutTagB, + T, Catlass::layout::RowMajor>; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + + constexpr uint32_t capacityK = 64; + constexpr uint32_t n = 128; + const uint32_t k = static_cast(curT); + auto layoutB = tla::MakeLayout(capacityK, n); + auto tensorU = tla::MakeTensor( + propagatedVNew_[KVOffset(b, hv, start, 0, V_)], + layoutB, Catlass::Arch::PositionGM{}); + auto blockU = GetTile(tensorU, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + using CopyGmToL1B = typename TileCopy::template CopyGmToL1B; + CopyGmToL1B copyGmToL1B; + + uint32_t uOffset = KDA_POST_PIPELINE_STAGE_COUNT * KDA_POST_PIPELINE_L1_SLOT_BYTES + + static_cast(slot) * KDA_POST_PIPELINE_L1_U_SLOT_BYTES; + LocalTensor l1U = resource.l1Buf.template GetBufferByByte(uOffset); + auto layoutL1B = tla::MakeLayout(capacityK, n); + auto tensorL1U = tla::MakeTensor(l1U, layoutL1B, Catlass::Arch::PositionL1{}); + + uint16_t pipelineEvent = KDA_POST_PIPELINE_U_EVENT + slot; + if (reuseStage) { + WaitFlag(pipelineEvent); + } + copyGmToL1B(tensorL1U, blockU); + SetFlag(pipelineEvent); + } + + __aicore__ inline void ComputePrefetchedPostWuPipelineArch35( + Catlass::Arch::Resource &resource, uint16_t slot, + uint64_t b, uint64_t hv, uint64_t start, uint64_t curT) + { + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla< + KdaArchTag, T, LayoutTagA, T, LayoutTagB, T, LayoutTagC>; + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + using LayoutTagL0A = typename TileCopy::LayoutTagL0A; + using LayoutTagL0B = typename TileCopy::LayoutTagL0B; + using CopyL1ToL0A = typename TileCopy::CopyL1ToL0A; + using CopyL1ToL0B = typename TileCopy::CopyL1ToL0B; + using TileMmad = Catlass::Gemm::Tile::TileMmadTla; + + constexpr uint32_t capacityM = 64; + constexpr uint32_t n = 128; + constexpr uint32_t packedN = 256; + constexpr uint32_t capacityK = 64; + const uint32_t m = static_cast(curT); + const uint32_t k = static_cast(curT); + auto layoutC = tla::MakeLayout(capacityM, n); + auto tensorWOut = tla::MakeTensor( + w_[KVOffset(b, hv, start, 0, K_)], layoutC, Catlass::Arch::PositionGM{}); + auto tensorUOut = tla::MakeTensor( + u_[KVOffset(b, hv, start, 0, V_)], layoutC, Catlass::Arch::PositionGM{}); + auto blockWOut = GetTile(tensorWOut, tla::MakeCoord(0, 0), tla::MakeShape(m, n)); + auto blockUOut = GetTile(tensorUOut, tla::MakeCoord(0, 0), tla::MakeShape(m, n)); + using CopyL0CToDst = typename TileCopy::template CopyL0CToDst; + + uint32_t l1Base = static_cast(slot) * KDA_POST_PIPELINE_L1_SLOT_BYTES; + LocalTensor l1A = resource.l1Buf.template GetBufferByByte(l1Base); + LocalTensor l1W = resource.l1Buf.template GetBufferByByte( + l1Base + KDA_POST_PIPELINE_L1_A_BYTES); + uint32_t uOffset = KDA_POST_PIPELINE_STAGE_COUNT * KDA_POST_PIPELINE_L1_SLOT_BYTES + + static_cast(slot) * KDA_POST_PIPELINE_L1_U_SLOT_BYTES; + LocalTensor l1U = resource.l1Buf.template GetBufferByByte(uOffset); + LocalTensor l0A = resource.l0ABuf.template GetBufferByByte( + static_cast(slot) * KDA_POST_PIPELINE_L0_A_SLOT_BYTES); + LocalTensor l0B = resource.l0BBuf.template GetBufferByByte( + static_cast(slot) * KDA_POST_PIPELINE_L0_B_SLOT_BYTES); + LocalTensor l0C = resource.l0CBuf.template GetBufferByByte( + static_cast(slot) * KDA_POST_PIPELINE_L0_C_SLOT_BYTES); + + auto layoutL1A = tla::MakeLayout(capacityM, capacityK); + auto layoutL1B = tla::MakeLayout(capacityK, n); + auto layoutL0A = tla::MakeLayout(capacityM, capacityK); + auto layoutL0B = tla::MakeLayout(capacityK, packedN); + auto layoutL0C = tla::MakeLayoutL0C(capacityM, packedN); + auto tensorL1A = tla::MakeTensor(l1A, layoutL1A, Catlass::Arch::PositionL1{}); + auto tensorL1W = tla::MakeTensor(l1W, layoutL1B, Catlass::Arch::PositionL1{}); + auto tensorL1U = tla::MakeTensor(l1U, layoutL1B, Catlass::Arch::PositionL1{}); + auto tensorL0A = tla::MakeTensor(l0A, layoutL0A, Catlass::Arch::PositionL0A{}); + auto tensorL0B = tla::MakeTensor(l0B, layoutL0B, Catlass::Arch::PositionL0B{}); + auto tensorL0C = tla::MakeTensor(l0C, layoutL0C, Catlass::Arch::PositionL0C{}); + auto tileL1A = GetTile(tensorL1A, tla::MakeCoord(0, 0), tla::MakeShape(m, k)); + auto tileL1W = GetTile(tensorL1W, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto tileL1U = GetTile(tensorL1U, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto tileL0A = GetTile(tensorL0A, tla::MakeCoord(0, 0), tla::MakeShape(m, k)); + auto tileL0BW = GetTile(tensorL0B, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto tileL0BU = GetTile(tensorL0B, tla::MakeCoord(0, n), tla::MakeShape(k, n)); + auto tileL0B = GetTile(tensorL0B, tla::MakeCoord(0, 0), tla::MakeShape(k, packedN)); + auto tileL0C = GetTile(tensorL0C, tla::MakeCoord(0, 0), tla::MakeShape(m, packedN)); + auto tileL0CW = GetTile(tensorL0C, tla::MakeCoord(0, 0), tla::MakeShape(m, n)); + auto tileL0CU = GetTile(tensorL0C, tla::MakeCoord(0, n), tla::MakeShape(m, n)); + + CopyL1ToL0A copyL1ToL0A; + CopyL1ToL0B copyL1ToL0B; + CopyL0CToDst copyL0CToDst; + TileMmad tileMmad; + + uint16_t pipelineEvent = KDA_POST_EVENT + slot; + uint16_t uPipelineEvent = KDA_POST_PIPELINE_U_EVENT + slot; + WaitFlag(pipelineEvent); + WaitFlag(uPipelineEvent); + WaitFlag(pipelineEvent); + copyL1ToL0A(tileL0A, tileL1A); + copyL1ToL0B(tileL0BW, tileL1W); + copyL1ToL0B(tileL0BU, tileL1U); + SetFlag(pipelineEvent); + SetFlag(pipelineEvent); + SetFlag(uPipelineEvent); + WaitFlag(pipelineEvent); + WaitFlag(pipelineEvent); + tileMmad(tileL0C, tileL0A, tileL0B, m, packedN, k, true, 0); + SetFlag(pipelineEvent); + SetFlag(pipelineEvent); + WaitFlag(pipelineEvent); + copyL0CToDst(blockWOut, tileL0CW); + copyL0CToDst(blockUOut, tileL0CU); + SetFlag(pipelineEvent); + } +#endif + + __aicore__ inline void ComputePostWuCube(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT) + { +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if (UsePostWuCubeArch35(curT)) { + ComputePostWuCubeFusedArch35(b, hv, start, curT); + return; + } +#endif + using ElementA = AKK_T; + using ElementB = T; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + SetLoadDataPaddingValue(static_cast(0)); + using WTileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; +#else + using WTileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; +#endif + using UTileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using PostL1TileShape128 = tla::Shape; + using PostL0TileShape128 = tla::Shape; + using PostL1TileShape256 = tla::Shape; + using PostL0TileShape256 = tla::Shape; +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + using WBlockMmad = Catlass::Gemm::Block::BlockMmadTla; +#else + using WBlockMmad = Catlass::Gemm::Block::BlockMmadTla; +#endif + using UBlockMmad128 = Catlass::Gemm::Block::BlockMmadTla; + using UBlockMmad256 = Catlass::Gemm::Block::BlockMmadTla; + LayoutTagA tagA = LayoutTagA::template MakeLayout(BT_, BT_); + auto layoutA = tla::MakeLayoutFromTag(tagA); + auto tensorA = tla::MakeTensor(preparedAqk_[AOffset(b, hv, start, 0)], layoutA, + Catlass::Arch::PositionGM{}); + + { + LayoutTagB tagB = LayoutTagB::template MakeLayout(BT_, K_); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + LayoutTagC tagC = LayoutTagC::template MakeLayout(BT_, K_); +#else + LayoutTagC tagC = LayoutTagC::template MakeLayout(BT_, K_); +#endif + auto layoutB = tla::MakeLayoutFromTag(tagB); + auto layoutC = tla::MakeLayoutFromTag(tagC); + Catlass::GemmCoord shape{static_cast(curT), static_cast(K_), + static_cast(curT)}; + auto tensorB = tla::MakeTensor(preparedQG_[KVOffset(b, hv, start, 0, K_)], layoutB, + Catlass::Arch::PositionGM{}); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + auto tensorC = tla::MakeTensor(w_[KVOffset(b, hv, start, 0, K_)], layoutC, + Catlass::Arch::PositionGM{}); +#else + auto tensorC = tla::MakeTensor(h_[WScratchOffset(b, hv, chunkIdx, 0, 0)], layoutC, + Catlass::Arch::PositionGM{}); +#endif + auto blockA = GetTile(tensorA, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.k())); + auto blockB = GetTile(tensorB, tla::MakeCoord(0, 0), tla::MakeShape(shape.k(), shape.n())); + auto blockC = GetTile(tensorC, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.n())); + Catlass::Arch::Resource wResource; + WBlockMmad wBlockMmad(wResource); + wBlockMmad(blockA, blockB, blockC, shape); + PipeBarrier(); + } + + { + LayoutTagB tagB = LayoutTagB::template MakeLayout(BT_, V_); + LayoutTagC tagC = LayoutTagC::template MakeLayout(BT_, V_); + auto layoutB = tla::MakeLayoutFromTag(tagB); + auto layoutC = tla::MakeLayoutFromTag(tagC); + Catlass::GemmCoord shape{static_cast(curT), static_cast(V_), + static_cast(curT)}; + auto tensorB = tla::MakeTensor(propagatedVNew_[KVOffset(b, hv, start, 0, V_)], layoutB, + Catlass::Arch::PositionGM{}); + auto tensorC = tla::MakeTensor(u_[KVOffset(b, hv, start, 0, V_)], layoutC, + Catlass::Arch::PositionGM{}); + auto blockA = GetTile(tensorA, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.k())); + auto blockB = GetTile(tensorB, tla::MakeCoord(0, 0), tla::MakeShape(shape.k(), shape.n())); + auto blockC = GetTile(tensorC, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.n())); + Catlass::Arch::Resource uResource; + if (V_ <= 128) { + UBlockMmad128 uBlockMmad(uResource); + uBlockMmad(blockA, blockB, blockC, shape); + } else { + UBlockMmad256 uBlockMmad(uResource); + uBlockMmad(blockA, blockB, blockC, shape); + } + PipeBarrier(); + } + + } + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + __aicore__ inline bool UseTypicalPostWuGate(uint64_t curT) const + { + return curT == 64 && BT_ == 64 && K_ == 128 && V_ == 128; + } + + __aicore__ inline uint64_t TypicalGateStageElems() const + { + return static_cast(KDA_TYPICAL_GATE_TILE_ROWS) * 128; + } + + __aicore__ inline uint64_t TypicalGateStageBytes() const + { + return TypicalGateStageElems() * (sizeof(T) + sizeof(GK_T)); + } + + __aicore__ inline LocalTensor TypicalGateK(uint64_t slot) + { + return gateWritebackBuf_.Get()[slot * TypicalGateStageBytes() / sizeof(T)]; + } + + __aicore__ inline LocalTensor TypicalGateG(uint64_t slot) + { + uint64_t byteOffset = slot * TypicalGateStageBytes() + TypicalGateStageElems() * sizeof(T); + return gateWritebackBuf_.Get()[byteOffset / sizeof(GK_T)]; + } + + __aicore__ inline void PrefetchTypicalKg(uint64_t slot, uint64_t b, uint64_t h, uint64_t hv, + uint64_t token, uint64_t rows) + { + uint64_t elems = rows * K_; + LocalTensor kStage = TypicalGateK(slot); + LocalTensor gateStage = TypicalGateG(slot); + CopyRowsIn(kStage, k_, QOffset(b, h, token, 0), rows, K_, + inputSequenceMajor_ ? H_ * K_ : K_); + DataCopy(gateStage, gk_[KVOffset(b, hv, token, 0, K_)], static_cast(elems)); + SetFlag(mte2ToVEvent_); + } + + __aicore__ inline void ComputeTypicalKg(uint64_t b, uint64_t h, uint64_t hv, uint64_t start, + uint64_t curT, uint64_t subBlockIdx, uint64_t subBlockNum) + { + uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; + uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + if (rowBegin >= rowEnd) { + return; + } + + LocalTensor gateLast = exp2Buf_.Get(); + LoadAsFloatRow(gk_, KVOffset(b, hv, start + curT - 1, 0, K_), gateLast, K_); + + uint64_t slot = 0; + uint64_t firstRows = rowEnd - rowBegin; + if (firstRows > KDA_TYPICAL_GATE_TILE_ROWS) { + firstRows = KDA_TYPICAL_GATE_TILE_ROWS; + } + PrefetchTypicalKg(slot, b, h, hv, start + rowBegin, firstRows); + + bool outputPending = false; + for (uint64_t tileRow = rowBegin; tileRow < rowEnd; tileRow += KDA_TYPICAL_GATE_TILE_ROWS) { + uint64_t tileRows = rowEnd - tileRow; + if (tileRows > KDA_TYPICAL_GATE_TILE_ROWS) { + tileRows = KDA_TYPICAL_GATE_TILE_ROWS; + } + uint64_t elems = tileRows * K_; + WaitFlag(mte2ToVEvent_); + + if (outputPending) { + WaitFlag(mte3ToMte2Event_); + } + uint64_t nextTileRow = tileRow + KDA_TYPICAL_GATE_TILE_ROWS; + if (nextTileRow < rowEnd) { + uint64_t nextRows = rowEnd - nextTileRow; + if (nextRows > KDA_TYPICAL_GATE_TILE_ROWS) { + nextRows = KDA_TYPICAL_GATE_TILE_ROWS; + } + PrefetchTypicalKg(slot ^ 1, b, h, hv, start + nextTileRow, nextRows); + } + + LocalTensor kAndKg = TypicalGateK(slot); + LocalTensor gateStage = TypicalGateG(slot); + ComputePostKdaKgRegbase( + (__ubuf__ T *)reinterpret_cast(kAndKg.GetPhyAddr()), + (__ubuf__ GK_T *)reinterpret_cast(gateStage.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(gateLast.GetPhyAddr()), + static_cast(tileRows), static_cast(K_)); + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(kg_[KVOffset(b, hv, start + tileRow, 0, K_)], kAndKg, + static_cast(elems)); + SetFlag(mte3ToMte2Event_); + outputPending = true; + slot ^= 1; + } + if (outputPending) { + WaitFlag(mte3ToMte2Event_); + } + } + + __aicore__ inline uint64_t TypicalGatePipelineStageElems() const + { + return static_cast(KDA_TYPICAL_GATE_PIPELINE_ROWS) * 128; + } + + __aicore__ inline uint64_t TypicalGatePipelineStageBytes() const + { + return TypicalGatePipelineStageElems() * (sizeof(T) + sizeof(float)) + + 128 * sizeof(float); + } + + __aicore__ inline LocalTensor TypicalGatePipelineK(uint64_t slot) + { + uint64_t byteOffset = slot * TypicalGatePipelineStageBytes(); + return vecBuf_.Get()[byteOffset / sizeof(T)]; + } + + __aicore__ inline LocalTensor TypicalGatePipelineG(uint64_t slot) + { + uint64_t byteOffset = slot * TypicalGatePipelineStageBytes() + + TypicalGatePipelineStageElems() * sizeof(T); + return vecBuf_.Get()[byteOffset / sizeof(float)]; + } + + __aicore__ inline LocalTensor TypicalGatePipelineRef(uint64_t slot) + { + uint64_t byteOffset = slot * TypicalGatePipelineStageBytes() + + TypicalGatePipelineStageElems() * (sizeof(T) + sizeof(float)); + return vecBuf_.Get()[byteOffset / sizeof(float)]; + } + + __aicore__ inline bool CanPipelineTypicalKg( + uint64_t curT, uint64_t subBlockIdx, uint64_t subBlockNum) const + { + if constexpr (!IsSameType::value) { + return false; + } + if (!UseTypicalPostWuGate(curT) || subBlockNum == 0) { + return false; + } + uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; + uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + return rowBegin < rowEnd && rowEnd - rowBegin <= KDA_TYPICAL_GATE_PIPELINE_ROWS; + } + + __aicore__ inline void PrefetchTypicalKgPipeline( + uint64_t slot, uint64_t b, uint64_t h, uint64_t hv, uint64_t start, + uint64_t curT, uint64_t rowBegin, uint64_t rowEnd) + { + if constexpr (IsSameType::value) { + uint64_t elems = (rowEnd - rowBegin) * K_; + LocalTensor kStage = TypicalGatePipelineK(slot); + LocalTensor gateStage = TypicalGatePipelineG(slot); + LocalTensor refStage = TypicalGatePipelineRef(slot); + CopyRowsIn(kStage, k_, QOffset(b, h, start + rowBegin, 0), rowEnd - rowBegin, K_, + inputSequenceMajor_ ? H_ * K_ : K_); + DataCopy(gateStage, gk_[KVOffset(b, hv, start + rowBegin, 0, K_)], + static_cast(elems)); + DataCopy(refStage, gk_[KVOffset(b, hv, start + curT - 1, 0, K_)], + static_cast(K_)); + SetFlag(mte2ToVEvent_); + } + } + + __aicore__ inline void ComputeTypicalKgPipelineRegs( + uint64_t slot, uint64_t rowBegin, uint64_t rowEnd) + { + uint64_t rows = rowEnd - rowBegin; + LocalTensor kAndKg = TypicalGatePipelineK(slot); + LocalTensor gateStage = TypicalGatePipelineG(slot); + LocalTensor refStage = TypicalGatePipelineRef(slot); + ComputePostKdaKgRegbase( + (__ubuf__ T *)reinterpret_cast(kAndKg.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(gateStage.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(refStage.GetPhyAddr()), + static_cast(rows), static_cast(K_)); + } + + __aicore__ inline void StoreTypicalKgPipeline( + uint64_t slot, uint64_t b, uint64_t hv, uint64_t start, + uint64_t rowBegin, uint64_t rowEnd) + { + uint64_t elems = (rowEnd - rowBegin) * K_; + LocalTensor kAndKg = TypicalGatePipelineK(slot); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(kg_[KVOffset(b, hv, start + rowBegin, 0, K_)], kAndKg, + static_cast(elems)); + SetFlag(mte3ToMte2Event_); + } +#endif + + __aicore__ inline void CopyScratchWAndFinalizeKg(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT, uint64_t subBlockIdx, + uint64_t subBlockNum) + { + constexpr uint64_t typedOffsetFloats = 20480; + constexpr uint64_t typedOffset = typedOffsetFloats * sizeof(float) / sizeof(T); + constexpr uint64_t kgFp32Planes = 4; + uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; + uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + if (rowBegin >= rowEnd) { + return; + } + uint64_t maxRows = (typedOffsetFloats / kgFp32Planes) / K_; + if (maxRows > 32) { + maxRows = 32; + } + if (maxRows == 0) { + return; + } + + uint64_t last = start + curT - 1; + LocalTensor arena = vecBuf_.Get(); + LocalTensor gateLast = exp2Buf_.Get(); + LocalTensor typedLocal = vecBuf_.Get()[typedOffset]; + LoadAsFloatRow(gk_, KVOffset(b, hv, last, 0, K_), gateLast, K_); + + for (uint64_t tileRow = rowBegin; tileRow < rowEnd; tileRow += maxRows) { + uint64_t tileRows = rowEnd - tileRow; + if (tileRows > maxRows) { + tileRows = maxRows; + } + uint64_t elemCount = tileRows * K_; +#if !defined(__CCE_AICORE__) || __CCE_AICORE__ != 310 + uint64_t scratchBase = WScratchOffset(b, hv, chunkIdx, tileRow, 0); +#else + (void)chunkIdx; +#endif + uint64_t token = start + tileRow; + +#if !defined(__CCE_AICORE__) || __CCE_AICORE__ != 310 + DataCopy(arena, h_[scratchBase], static_cast(elemCount)); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast(typedLocal, arena, RoundMode::CAST_RINT, static_cast(elemCount)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(w_[KVOffset(b, hv, token, 0, K_)], typedLocal, static_cast(elemCount)); + SetFlag(mte3ToMte2Event_); + WaitFlag(mte3ToMte2Event_); +#endif + + LocalTensor kLocal = arena; + LocalTensor gLocal = arena[elemCount]; + LocalTensor expLocal = arena[2 * elemCount]; + LocalTensor outLocal = arena[3 * elemCount]; + const uint64_t gateOffsetBytes = (typedOffset + elemCount) * sizeof(T); + LocalTensor gateTyped = vecBuf_.Get()[ + (gateOffsetBytes + sizeof(GK_T) - 1) / sizeof(GK_T)]; + CopyRowsIn(typedLocal, k_, QOffset(b, h, token, 0), tileRows, K_, + inputSequenceMajor_ ? H_ * K_ : K_); + LoadAsFloatVector(gk_, KVOffset(b, hv, token, 0, K_), gLocal, gateTyped, elemCount); + Cast(kLocal, typedLocal, RoundMode::CAST_NONE, static_cast(elemCount)); + PipeBarrier(); + + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expLocal[row * K_], gateLast, gLocal[row * K_], static_cast(K_)); + } + PipeBarrier(); + Muls(expLocal, expLocal, LN2, static_cast(elemCount)); + PipeBarrier(); + ClampExpInput(expLocal, static_cast(elemCount)); + Exp(expLocal, expLocal, static_cast(elemCount)); + PipeBarrier(); + Mul(outLocal, kLocal, expLocal, static_cast(elemCount)); + PipeBarrier(); + ClampFp32ToOutputType(outLocal, static_cast(elemCount)); + Cast(typedLocal, outLocal, RoundMode::CAST_RINT, static_cast(elemCount)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(kg_, KVOffset(b, hv, token, 0, K_), typedLocal, elemCount); + SetFlag(mte3ToMte2Event_); + WaitFlag(mte3ToMte2Event_); + } + if (rowEnd == curT) { + CopyVectorIn(typedLocal, k_, QOffset(b, h, last, 0), K_); + SetFlag(mte2ToMte3Event_); + WaitFlag(mte2ToMte3Event_); + CopyVectorOut(kg_, KVOffset(b, hv, last, 0, K_), typedLocal, K_); + SetFlag(mte3ToMte2Event_); + WaitFlag(mte3ToMte2Event_); + } + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + template + __aicore__ inline void ComputeTailWuRow(GlobalTensor &src, GlobalTensor &dst, + uint64_t akkBase, uint64_t srcBase, uint64_t dstBase, uint64_t curT, + uint64_t dim, uint64_t rowStride) + { + LocalTensor acc = vecBuf_.Get(); + LocalTensor value = vecBuf_.Get()[512]; + LocalTensor typed = vecBuf_.Get()[4096]; + LocalTensor coefficientTyped = exp2Buf_.Get(); + LocalTensor coefficients = exp2Buf_.Get()[128]; + CopyVectorIn(coefficientTyped, preparedAqk_, akkBase, curT); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast(coefficients, coefficientTyped, RoundMode::CAST_NONE, static_cast(curT)); + PipeBarrier(); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + for (uint64_t j = 0; j < curT; ++j) { + LoadAsFloatVector(src, srcBase + j * rowStride, value, typed, dim); + float coefficient = coefficients.GetValue(j); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + Muls(value, value, coefficient, static_cast(dim)); + PipeBarrier(); + if (j == 0) { + Adds(acc, value, 0.0f, static_cast(dim)); + } else { + Add(acc, acc, value, static_cast(dim)); + } + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + ClampFp32ToOutputType(acc, static_cast(dim)); + StoreFloatRow(dst, dstBase, acc, dim); + } + + __aicore__ inline void ComputeTailWuVector(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT, uint64_t subBlockIdx, uint64_t subBlockNum) + { + // preparedQG_ and w_ alias. Each subblock owns disjoint columns and + // writes rows from last to first, so every lower-triangular source row + // remains live through its final use without desynchronizing the AIVs. + uint64_t colBegin = (K_ * subBlockIdx) / subBlockNum; + uint64_t colEnd = (K_ * (subBlockIdx + 1)) / subBlockNum; + for (uint64_t row = curT; row > 0; --row) { + uint64_t rowIdx = row - 1; + ComputeTailWuRow( + preparedQG_, w_, AOffset(b, hv, start + rowIdx, 0), KVOffset(b, hv, start, colBegin, K_), + KVOffset(b, hv, start + rowIdx, colBegin, K_), curT, colEnd - colBegin, K_); + } + + uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; + uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + for (uint64_t row = rowBegin; row < rowEnd; ++row) { + ComputeTailWuRow( + propagatedVNew_, u_, AOffset(b, hv, start + row, 0), KVOffset(b, hv, start, 0, V_), + KVOffset(b, hv, start + row, 0, V_), curT, V_, V_); + } + } +#endif + __aicore__ inline bool ResolveFlatChunk(uint64_t task, uint64_t &seq, uint64_t &b, uint64_t &h, uint64_t &hv, + uint64_t &chunkIdx, uint64_t &start, uint64_t &end) + { + hv = task % HV_; + uint64_t flatChunk = task / HV_; + if (!isVarLen_) { + seq = flatChunk / NT_; + b = seq; + chunkIdx = flatChunk % NT_; + start = chunkIdx * BT_; + end = start + BT_; + if (end > T_) { + end = T_; + } + } else { + if (!KdaVarlen::ResolveChunkRange( + cuSeqlensAddr_, chunkIndicesAddr_, N_, T_, BT_, flatChunk, + seq, start, end)) { + return false; + } + b = 0; + chunkIdx = flatChunk; + } + h = hv / (HV_ / H_); + return start < end; + } + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + __aicore__ inline bool ResolveHeadMajorChunk( + uint64_t task, uint64_t &seq, uint64_t &b, uint64_t &h, uint64_t &hv, + uint64_t &chunkIdx, uint64_t &start, uint64_t &end) + { + uint64_t flatTask = 0; + if (isVarLen_) { + hv = task / NT_; + uint64_t flatChunk = task % NT_; + flatTask = flatChunk * HV_ + hv; + } else { + uint64_t entity = task / NT_; + uint64_t localChunk = task % NT_; + b = entity / HV_; + hv = entity % HV_; + uint64_t flatChunk = b * NT_ + localChunk; + flatTask = flatChunk * HV_ + hv; + } + return ResolveFlatChunk(flatTask, seq, b, h, hv, chunkIdx, start, end); + } + + __aicore__ inline void GetHeadMajorTaskRange( + uint64_t coreIdx, uint64_t coreNum, uint64_t taskNum, + uint64_t &taskBegin, uint64_t &taskEnd) const + { + uint64_t tasksPerCore = (taskNum + coreNum - 1) / coreNum; + taskBegin = coreIdx * tasksPerCore; + taskEnd = taskBegin + tasksPerCore; + if (taskBegin > taskNum) { + taskBegin = taskNum; + } + if (taskEnd > taskNum) { + taskEnd = taskNum; + } + } + + __aicore__ inline void ProcessPostAivPipelineArch35( + uint64_t taskBegin, uint64_t taskEnd, uint64_t subBlockIdx, uint64_t subBlockNum) + { + uint64_t task = taskBegin; + while (task < taskEnd) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + bool resolved = ResolveHeadMajorChunk(task, seq, b, h, hv, chunkIdx, start, end); + uint64_t curT = resolved ? end - start : 0; + if (!resolved || !CanPipelineTypicalKg(curT, subBlockIdx, subBlockNum)) { + if (resolved) { + (void)seq; + ProcessChunkPostAiv( + b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); + } + ++task; + continue; + } + + uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; + uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + uint16_t slot = 0; + bool outputPending = false; + PrefetchTypicalKgPipeline( + slot, b, h, hv, start, curT, rowBegin, rowEnd); + + while (true) { + WaitFlag(mte2ToVEvent_); + + uint64_t nextTask = task + 1; + uint64_t nextSeq = 0; + uint64_t nextB = 0; + uint64_t nextH = 0; + uint64_t nextHv = 0; + uint64_t nextChunkIdx = 0; + uint64_t nextStart = 0; + uint64_t nextEnd = 0; + bool nextResolved = nextTask < taskEnd && ResolveHeadMajorChunk( + nextTask, nextSeq, nextB, nextH, nextHv, nextChunkIdx, nextStart, nextEnd); + uint64_t nextCurT = nextResolved ? nextEnd - nextStart : 0; + bool nextIsTypical = nextResolved && + CanPipelineTypicalKg(nextCurT, subBlockIdx, subBlockNum); + uint64_t nextRowBegin = 0; + uint64_t nextRowEnd = 0; + if (nextIsTypical) { + nextRowBegin = (nextCurT * subBlockIdx) / subBlockNum; + nextRowEnd = (nextCurT * (subBlockIdx + 1)) / subBlockNum; + uint16_t nextSlot = (slot + 1) % KDA_TYPICAL_GATE_PIPELINE_STAGES; + PrefetchTypicalKgPipeline( + nextSlot, nextB, nextH, nextHv, nextStart, nextCurT, + nextRowBegin, nextRowEnd); + } + + ComputeTypicalKgPipelineRegs(slot, rowBegin, rowEnd); + if (outputPending) { + WaitFlag(mte3ToMte2Event_); + outputPending = false; + } + StoreTypicalKgPipeline(slot, b, hv, start, rowBegin, rowEnd); + outputPending = true; + + if (!nextIsTypical) { + WaitFlag(mte3ToMte2Event_); + task = nextTask; + break; + } + + task = nextTask; + seq = nextSeq; + b = nextB; + h = nextH; + hv = nextHv; + chunkIdx = nextChunkIdx; + start = nextStart; + end = nextEnd; + curT = nextCurT; + rowBegin = nextRowBegin; + rowEnd = nextRowEnd; + slot = (slot + 1) % KDA_TYPICAL_GATE_PIPELINE_STAGES; + (void)seq; + (void)chunkIdx; + (void)end; + (void)curT; + } + } + } +#endif + + __aicore__ inline void ProcessChunkPostAiv(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t end, uint64_t subBlockIdx, + uint64_t subBlockNum) + { + uint64_t curT = end - start; + if (curT == 0 || !UsePostWuCube(curT)) { + return; + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if (curT < BT_) { + ComputeTailWuVector(b, hv, chunkIdx, start, curT, subBlockIdx, subBlockNum); + CopyScratchWAndFinalizeKg( + b, h, hv, chunkIdx, start, curT, subBlockIdx, subBlockNum); + return; + } + if (UseTypicalPostWuGate(curT)) { + ComputeTypicalKg(b, h, hv, start, curT, subBlockIdx, subBlockNum); + return; + } else { + CopyScratchWAndFinalizeKg(b, h, hv, chunkIdx, start, curT, subBlockIdx, subBlockNum); + } +#else + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(syncDoneFlag_); + CopyScratchWAndFinalizeKg(b, h, hv, chunkIdx, start, curT, subBlockIdx, subBlockNum); +#endif + } + + __aicore__ inline void ProcessChunkPostAic(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + if constexpr (IsSameType::value) { + ProcessChunkPostAicTyped(b, hv, chunkIdx, start, end); + } + } + + __aicore__ inline void ProcessChunkPostAicTyped(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + uint64_t curT = end - start; + if (curT == 0 || !UsePostWuCube(curT)) { + return; + } + ComputePostWuCube(b, hv, chunkIdx, start, curT); +#if !defined(__CCE_AICORE__) || __CCE_AICORE__ != 310 + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); +#endif + } + + __aicore__ inline void ProcessPostAiv() + { + if constexpr (IsSameType::value) { + return; + } + uint64_t subBlockNum = static_cast(GetSubBlockNum()); + if (subBlockNum == 0) { + return; + } + uint64_t subBlockIdx = static_cast(GetSubBlockIdx()); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + uint64_t coreIdx = static_cast(GetBlockIdx()) / subBlockNum; + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if (BT_ == 64 && K_ == 128 && V_ == 128) { + uint64_t taskBegin = 0; + uint64_t taskEnd = 0; + GetHeadMajorTaskRange(coreIdx, coreNum, taskNum, taskBegin, taskEnd); + ProcessPostAivPipelineArch35(taskBegin, taskEnd, subBlockIdx, subBlockNum); + return; + } +#endif + for (uint64_t task = coreIdx; task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + ProcessChunkPostAiv(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); + } + } + } + + __aicore__ inline void ProcessPostAic() + { + if constexpr (IsSameType::value) { + return; + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if (KDA_ENABLE_POST_AIC_PIPELINE && BT_ == 64 && K_ == 128 && V_ == 128) { + ProcessPostAicPipelineArch35(); + return; + } +#endif + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + (void)h; + ProcessChunkPostAic(b, hv, chunkIdx, start, end); + } + } + } + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + __aicore__ inline void ProcessPostAicPipelineArch35() + { + static_assert(sizeof(T) == sizeof(uint16_t), + "arch35 PostWU pipeline is specialized for fp16/bf16 inputs"); + SetLoadDataPaddingValue(static_cast(0)); + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + uint64_t coreIdx = static_cast(GetBlockIdx()); + uint64_t taskBegin = 0; + uint64_t taskEnd = 0; + GetHeadMajorTaskRange(coreIdx, coreNum, taskNum, taskBegin, taskEnd); + if (taskEnd - taskBegin < KDA_POST_PIPELINE_STAGE_COUNT) { + for (uint64_t task = taskBegin; task < taskEnd; ++task) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveHeadMajorChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + (void)h; + if (end - start == BT_) { + ProcessPreparedTailSingleArch35(b, hv, start, BT_); + } + } + } + return; + } + uint64_t task = taskBegin; + Catlass::Arch::Resource resource; + + while (task < taskEnd) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + bool resolved = ResolveHeadMajorChunk(task, seq, b, h, hv, chunkIdx, start, end); + uint64_t curT = resolved ? end - start : 0; + if (!resolved || !UseFullPostWuPipelineArch35(curT)) { + (void)seq; + (void)h; + ++task; + continue; + } + + uint16_t slot = 0; + uint16_t usedSlotCount = 1; + InitializePostWuPipelineSlot(slot); + PrefetchPostWuPipelineArch35(resource, slot, b, hv, start, curT, false); + PrefetchPostWuPipelineU(resource, slot, b, hv, start, curT, false); + while (true) { + uint64_t nextTask = task + 1; + uint64_t nextSeq = 0; + uint64_t nextB = 0; + uint64_t nextH = 0; + uint64_t nextHv = 0; + uint64_t nextChunkIdx = 0; + uint64_t nextStart = 0; + uint64_t nextEnd = 0; + bool nextResolved = nextTask < taskEnd && ResolveHeadMajorChunk( + nextTask, nextSeq, nextB, nextH, nextHv, nextChunkIdx, nextStart, nextEnd); + uint64_t nextCurT = nextResolved ? nextEnd - nextStart : 0; + bool nextIsTypical = nextResolved && UseFullPostWuPipelineArch35(nextCurT); + if (nextIsTypical) { + uint16_t nextSlot = slot ^ 1; + bool reuseSlot = usedSlotCount == KDA_POST_PIPELINE_STAGE_COUNT; + if (!reuseSlot) { + InitializePostWuPipelineSlot(nextSlot); + } + PrefetchPostWuPipelineArch35( + resource, nextSlot, nextB, nextHv, nextStart, nextCurT, reuseSlot); + PrefetchPostWuPipelineU( + resource, nextSlot, nextB, nextHv, nextStart, nextCurT, reuseSlot); + if (!reuseSlot) { + ++usedSlotCount; + } + } + + ComputePrefetchedPostWuPipelineArch35(resource, slot, b, hv, start, curT); + if (!nextIsTypical) { + FinalizePostWuPipelineEvents(usedSlotCount); + task = nextTask; + break; + } + task = nextTask; + seq = nextSeq; + b = nextB; + h = nextH; + hv = nextHv; + chunkIdx = nextChunkIdx; + start = nextStart; + end = nextEnd; + curT = nextCurT; + slot ^= 1; + (void)seq; + (void)h; + (void)chunkIdx; + (void)end; + (void)curT; + } + } + } +#endif + + +private: + GlobalTensor q_; + GlobalTensor k_; + GlobalTensor v_; + GlobalTensor gk_; + GlobalTensor beta_; + GlobalTensor initialState_; + GlobalTensor o_; + GlobalTensor finalState_; + GlobalTensor aqk_; + GlobalTensor akk_; + GlobalTensor w_; + GlobalTensor u_; + GlobalTensor qg_; + GlobalTensor kg_; + GlobalTensor vNew_; + GlobalTensor h_; + GlobalTensor preparedQG_; + GlobalTensor preparedAqk_; + GlobalTensor propagatedVNew_; + GlobalTensor propagatedH_; + GlobalTensor solveWorkspace_; + GlobalTensor scoreWorkspace_; + TPipe *pipe_ = nullptr; + TBuf exp2Buf_; + TBuf vecBuf_; + TBuf gateWritebackBuf_; + TEventID mte2ToVEvent_ = 0; + TEventID vToMte2Event_ = 0; + TEventID vToMte3Event_ = 0; + TEventID mte3ToVEvent_ = 0; + TEventID mte2ToMte3Event_ = 0; + TEventID mte3ToMte2Event_ = 0; + bool vectorEventsAllocated_ = false; + Catlass::Arch::CrossCoreFlagWithReverse scoreReadyFlag_{KDA_SCORE_READY_FLAG0, + KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse scoreDoneFlag_{KDA_SCORE_DONE_FLAG0, + KDA_SCORE_DONE_FLAG1}; + // Score production is fully drained before solve starts, so the solve handshake can safely reuse + // the existing score flags without consuming additional hardware flag IDs. + Catlass::Arch::CrossCoreFlagWithReverse syncReadyFlag_{KDA_SCORE_READY_FLAG0, + KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse syncDoneFlag_{KDA_SCORE_DONE_FLAG0, + KDA_SCORE_DONE_FLAG1}; + uint64_t B_ = 0; + uint64_t N_ = 0; + uint64_t H_ = 0; + uint64_t HV_ = 0; + uint64_t T_ = 0; + uint64_t K_ = 0; + uint64_t V_ = 0; + uint64_t BT_ = 0; + uint64_t NT_ = 0; + float scale_ = 1.0f; + bool hasInitial_ = false; + bool isVarLen_ = false; + bool isAivOnly_ = false; + bool inputSequenceMajor_ = false; + uint64_t usedCoreNum_ = 1; + uint64_t solveCoreIdx_ = 0; + __gm__ int64_t *chunkIndicesAddr_ = nullptr; + __gm__ int64_t *cuSeqlensAddr_ = nullptr; +}; +} // namespace + +template +__aicore__ inline void RunChunkKdaPostWu( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR wSeed, GM_ADDR akk, GM_ADDR uSeed, + GM_ADDR w, GM_ADDR u, GM_ADDR kg, GM_ADDR vNew, GM_ADDR userWorkspace, + const TilingData &tiling, TPipe &pipe) +{ + GM_ADDR postScratch = userWorkspace + tiling.postWuScratchOffset; + if ASCEND_IS_AIC { + ChunkKdaFwdPostWuKernel op; + op.Init(q, k, v, gk, beta, initialState, cuSeqlens, chunkIndices, + wSeed, akk, uSeed, nullptr, userWorkspace, userWorkspace, userWorkspace, akk, w, u, + userWorkspace, kg, vNew, postScratch, postScratch, tiling, &pipe, false); + op.ProcessAic(); + } + if ASCEND_IS_AIV { + ChunkKdaFwdPostWuKernel op; + op.Init(q, k, v, gk, beta, initialState, cuSeqlens, chunkIndices, + wSeed, akk, uSeed, nullptr, userWorkspace, userWorkspace, userWorkspace, akk, w, u, + userWorkspace, kg, vNew, postScratch, postScratch, tiling, &pipe); + op.ProcessAiv(); + } +} + +} // namespace KdaPostWu diff --git a/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_prepare.h b/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_prepare.h new file mode 100644 index 000000000000..5f8a58189a5b --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_kernel/arch35/chunk_kda_fwd_prepare.h @@ -0,0 +1,5149 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#pragma once + +#ifndef CATLASS_ARCH +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +#define CATLASS_ARCH 3510 +#else +#define CATLASS_ARCH 2201 +#endif +#endif + +#include "catlass/arch/arch.hpp" +#include "catlass/arch/cross_core_sync.hpp" +#include "catlass/arch/resource.hpp" +#include "catlass/catlass.hpp" +#include "catlass/gemm/block/block_mmad.hpp" +#include "catlass/gemm/dispatch_policy.hpp" +#include "catlass/gemm/tile/tile_copy.hpp" +#include "catlass/gemm_coord.hpp" +#include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp" +#include "kernel_utils/tile/copy_l0c_to_ub.hpp" +#include "catlass/layout/layout.hpp" +#include "kernel_operator.h" +#include "../chunk_kda_fwd_varlen.h" +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +#ifndef FLA_NPU_REGBASE_HPP_INCLUDED +#define FLA_NPU_REGBASE_HPP_INCLUDED +#include "kernel_utils/vector/regbase.hpp" +#endif +#endif +#include "tla/layout.hpp" +#include "tla/tensor.hpp" +#include "chunk_kda_fwd_post_wu.h" + +using namespace AscendC; + +namespace KdaPrepare { +namespace { +using KdaInt64 = tla::Int<64>; +using KdaInt128 = tla::Int<128>; +constexpr float LN2 = 0.69314718055994530942f; +constexpr float RCP_LN2 = 1.44269504088896340736f; +constexpr float KDA_EXP2_CLAMP = 80.0f; +constexpr float KDA_EXP_INPUT_MAX = KDA_EXP2_CLAMP * LN2; +constexpr float KDA_EXP_INPUT_MIN = -KDA_EXP2_CLAMP * LN2; +constexpr float KDA_SCORE_EXP2_CLAMP = 120.0f; +constexpr float KDA_SCORE_EXP2_MIN_CLAMP = 126.0f; +constexpr float KDA_SCORE_EXP_INPUT_MAX = KDA_SCORE_EXP2_CLAMP * LN2; +constexpr float KDA_SCORE_EXP_INPUT_MIN = -KDA_SCORE_EXP2_MIN_CLAMP * LN2; +constexpr float KDA_FP16_MAX = 65504.0f; +constexpr uint32_t EXP2_UB_ELEMENTS = 256; +constexpr uint32_t EXP2_UB_BYTES = EXP2_UB_ELEMENTS * (sizeof(float) + sizeof(uint16_t)); +constexpr uint32_t EXP2_EVENT_ID = 0; +constexpr uint32_t KDA_SOLVE_BT = 64; +constexpr uint32_t KDA_SOLVE_MATRIX_ELEMENTS = KDA_SOLVE_BT * KDA_SOLVE_BT; +constexpr uint32_t KDA_SOLVE_SCRATCH_X = 0; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y0 = 1; +constexpr uint32_t KDA_SOLVE_SCRATCH_TMP = 2; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y1 = 3; +constexpr uint32_t KDA_SOLVE_SCRATCH_IDENTITY = 4; +constexpr uint32_t KDA_SOLVE_SCRATCH_RAW_AKK = KDA_SOLVE_SCRATCH_Y1; +constexpr uint32_t KDA_SOLVE_SCRATCH_RAW_AQK = KDA_SOLVE_SCRATCH_IDENTITY; +constexpr uint32_t KDA_SOLVE_SCRATCH_SLOTS = 5; +constexpr uint32_t KDA_SOLVE_PIPELINE_DEPTH = 4; +constexpr uint32_t KDA_SOLVE_DIAG_BT = 16; +constexpr uint32_t KDA_SOLVE_DIAG_BLOCKS = KDA_SOLVE_BT / KDA_SOLVE_DIAG_BT; +constexpr uint32_t KDA_SOLVE_DIAG_MCH_ITERS = 3; +// Keep the local safe-gate exponent span within the BF16 score range while +// reducing repeated gate-factor work and AIV/AIC handshakes. +constexpr uint32_t KDA_SCORE_REF_BC = 32; +constexpr uint32_t KDA_SAFE_SCORE_REF_BC = 32; +constexpr uint32_t KDA_VEC_ARENA_ELEMENTS = 32768; +constexpr uint32_t KDA_BITS_PER_MASK_BYTE = 8; +constexpr uint32_t KDA_SELECT_COL_BLOCKS = 2; +constexpr uint32_t KDA_SELECT_COL_MASK_BYTES = KDA_SOLVE_MATRIX_ELEMENTS / KDA_BITS_PER_MASK_BYTE; +constexpr uint32_t KDA_SELECT_MASK_BYTES = KDA_SELECT_COL_BLOCKS * KDA_SELECT_COL_MASK_BYTES; +constexpr uint32_t KDA_SELECT_AQK_MASK_BYTE_OFFSET = 120 * 1024; +constexpr uint32_t KDA_SELECT_AKK_MASK_BYTE_OFFSET = KDA_SELECT_AQK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_BYTE_OFFSET = KDA_SELECT_AKK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_FLOAT_OFFSET = KDA_SELECT_ZERO_BYTE_OFFSET / sizeof(float); +constexpr uint8_t KDA_SCORE_DONE_FLAG0 = 2; +constexpr uint8_t KDA_SCORE_DONE_FLAG1 = 3; +constexpr uint8_t KDA_SCORE_READY_FLAG0 = 4; +constexpr uint8_t KDA_SCORE_READY_FLAG1 = 5; +constexpr uint8_t KDA_SOLVE_DONE_FLAG = 6; +constexpr uint8_t KDA_SOLVE_READY_FLAG = 7; +constexpr uint8_t KDA_POST_READY_FLAG = 0; +constexpr uint8_t KDA_POST_FREE_FLAG = 1; +constexpr uint32_t KDA_SCORE_QUEUE_DEPTH = 2; +constexpr uint32_t KDA_SCORE_LANES = 2; +constexpr uint32_t KDA_POST_QUEUE_DEPTH = 4; +constexpr uint32_t KDA_POST_QUEUE_STORAGE = KDA_POST_QUEUE_DEPTH + 1; +constexpr uint32_t KDA_DIRECT_SCORE_QUEUE_DEPTH = 2; +constexpr uint32_t KDA_DIRECT_SCORE_ROWS = 32; +constexpr uint32_t KDA_DIRECT_SCORE_MATRIX_ELEMENTS = KDA_DIRECT_SCORE_ROWS * 64; +constexpr uint32_t KDA_DIRECT_SCORE_SLOT_ELEMENTS = 3 * KDA_DIRECT_SCORE_MATRIX_ELEMENTS; +constexpr uint32_t KDA_DIRECT_SCORE_FLOAT_OFFSET = 20 * 1024; +constexpr uint32_t KDA_DIRECT_SCORE_UB_BYTE_OFFSET = + EXP2_UB_BYTES + KDA_DIRECT_SCORE_FLOAT_OFFSET * sizeof(float); +constexpr uint32_t KDA_DIRECT_SCORE_L1_SLOT_ELEMENTS = 2 * KDA_DIRECT_SCORE_ROWS * 128; +constexpr uint32_t KDA_DIRECT_SCORE_L1_SLOT_BYTES = + KDA_DIRECT_SCORE_L1_SLOT_ELEMENTS * sizeof(uint16_t); +constexpr uint32_t KDA_DIRECT_SCORE_L1_B_OFFSET = + KDA_SCORE_QUEUE_DEPTH * KDA_SCORE_LANES * KDA_DIRECT_SCORE_L1_SLOT_BYTES; +constexpr uint64_t KDA_DIRECT_SCORE_FREE_FLAG = 8; +constexpr uint64_t KDA_DIRECT_SCORE_READY_FLAG = 10; +constexpr uint64_t KDA_DIRECT_SCORE_SUBBLOCK_FLAG_STRIDE = 16; +constexpr uint32_t KDA_SCORE_SCRATCH_SLOTS = KDA_SCORE_QUEUE_DEPTH * KDA_SCORE_LANES; +constexpr uint32_t KDA_SYNC_REVERSE_DEPTH = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_PLANES = 3; +constexpr uint32_t KDA_SCORE_SCRATCH_QG = 0; +constexpr uint32_t KDA_SCORE_SCRATCH_W = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_KG = 2; +constexpr uint64_t KDA_WORKSPACE_ALIGN = 512; +constexpr uint32_t KDA_GATE_TILE_ROWS = 16; +constexpr uint32_t KDA_GATE_PIPELINE_DEPTH = 3; +constexpr uint32_t KDA_AIV_UB_BUDGET_BYTES = 192 * 1024; +constexpr uint32_t KDA_LOCAL_GK_FLOAT_OFFSET = 10 * 1024; +constexpr uint32_t KDA_SCALED_QG_FLOAT_OFFSET = 18 * 1024; +constexpr bool KDA_ARCH35_ENABLE_HEAD_PAIR = true; +constexpr bool KDA_ARCH35_ENABLE_MANUAL_SCORE_PIPELINE = false; +constexpr bool KDA_ARCH35_ENABLE_DIRECT_SCORE_UB = true; +constexpr bool KDA_ARCH35_ENABLE_DIRECT_SCORE_L1 = false; +constexpr uint16_t KDA_ARCH35_SCORE_EVENT = 3; +constexpr uint16_t KDA_ARCH35_SCORE_W_EVENT = 4; +constexpr uint16_t KDA_ARCH35_SOLVE_FIX_EVENT = 7; + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +template +static __simd_vf__ inline void AccumulateRawSafeGateChunk128Regbase( + __ubuf__ float *input, __ubuf__ float *bias, __ubuf__ float *acc, + uint16_t rows, float expA, float lowerBound) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t FLOAT_ELEMENTS_PER_REG = AscendC::VECTOR_REG_WIDTH / sizeof(float); + constexpr uint16_t ROW_ELEMENTS = 2 * FLOAT_ELEMENTS_PER_REG; + + MaskReg floatMask = CreateMask(); + RegTensor accZeroReg; + RegTensor accOneReg; + RegTensor oneZeroReg; + RegTensor oneOneReg; + RegTensor biasZeroReg; + RegTensor biasOneReg; + LoadAlign(accZeroReg, acc); + LoadAlign(accOneReg, acc + FLOAT_ELEMENTS_PER_REG); + Duplicate(oneZeroReg, 1.0f, floatMask); + Duplicate(oneOneReg, 1.0f, floatMask); + if constexpr (HAS_BIAS) { + LoadAlign(biasZeroReg, bias); + LoadAlign(biasOneReg, bias + FLOAT_ELEMENTS_PER_REG); + } + + const float gateScale = lowerBound * RCP_LN2; + for (uint16_t row = 0; row < rows; ++row) { + const uint32_t rowOffset = static_cast(row) * ROW_ELEMENTS; + RegTensor gateZeroReg; + RegTensor gateOneReg; + RegTensor sigmoidZeroReg; + RegTensor sigmoidOneReg; + LoadAlign(gateZeroReg, input + rowOffset); + LoadAlign(gateOneReg, input + rowOffset + FLOAT_ELEMENTS_PER_REG); + if constexpr (HAS_BIAS) { + Add(gateZeroReg, gateZeroReg, biasZeroReg, floatMask); + Add(gateOneReg, gateOneReg, biasOneReg, floatMask); + } + Muls(gateZeroReg, gateZeroReg, -expA, floatMask); + Muls(gateOneReg, gateOneReg, -expA, floatMask); + Exp(gateZeroReg, gateZeroReg, floatMask); + Exp(gateOneReg, gateOneReg, floatMask); + Adds(gateZeroReg, gateZeroReg, 1.0f, floatMask); + Adds(gateOneReg, gateOneReg, 1.0f, floatMask); + Div(sigmoidZeroReg, oneZeroReg, gateZeroReg, floatMask); + Div(sigmoidOneReg, oneOneReg, gateOneReg, floatMask); + Muls(sigmoidZeroReg, sigmoidZeroReg, gateScale, floatMask); + Muls(sigmoidOneReg, sigmoidOneReg, gateScale, floatMask); + Add(accZeroReg, accZeroReg, sigmoidZeroReg, floatMask); + Add(accOneReg, accOneReg, sigmoidOneReg, floatMask); + StoreAlign(input + rowOffset, accZeroReg, floatMask); + StoreAlign(input + rowOffset + FLOAT_ELEMENTS_PER_REG, accOneReg, floatMask); + } + StoreAlign(acc, accZeroReg, floatMask); + StoreAlign(acc + FLOAT_ELEMENTS_PER_REG, accOneReg, floatMask); +} + +template +__simd_callee__ inline void LoadKdaGateRegbasePair( + AscendC::MicroAPI::RegTensor &zeroReg, + AscendC::MicroAPI::RegTensor &oneReg, + __ubuf__ InputT *src, + AscendC::MicroAPI::MaskReg &inputMask) +{ + using namespace AscendC::MicroAPI; + if constexpr (std::is_same()) { + LoadAlign(zeroReg, oneReg, src); + } else { + RegTensor inputReg; + LoadIn(inputReg, src); + CastHalf2Float(zeroReg, oneReg, inputReg, inputMask); + } +} + +template +__simd_callee__ inline void ClampKdaGateRegbaseOutput( + AscendC::MicroAPI::RegTensor &zeroReg, + AscendC::MicroAPI::RegTensor &oneReg, + AscendC::MicroAPI::MaskReg &floatMask) +{ + using namespace AscendC::MicroAPI; + if constexpr (std::is_same()) { + Mins(zeroReg, zeroReg, KDA_FP16_MAX, floatMask); + Mins(oneReg, oneReg, KDA_FP16_MAX, floatMask); + Maxs(zeroReg, zeroReg, -KDA_FP16_MAX, floatMask); + Maxs(oneReg, oneReg, -KDA_FP16_MAX, floatMask); + } +} + +template +__simd_callee__ inline void BuildKdaGateRegbaseExp( + AscendC::MicroAPI::RegTensor &expZeroReg, + AscendC::MicroAPI::RegTensor &expOneReg, + AscendC::MicroAPI::RegTensor &gateZeroReg, + AscendC::MicroAPI::RegTensor &gateOneReg, + __ubuf__ float *ref, + AscendC::MicroAPI::MaskReg &floatMask) +{ + using namespace AscendC::MicroAPI; + constexpr float expInputMax = + std::is_same() ? KDA_SCORE_EXP_INPUT_MAX : KDA_EXP_INPUT_MAX; + constexpr float expInputMin = + std::is_same() ? KDA_SCORE_EXP_INPUT_MIN : KDA_EXP_INPUT_MIN; + if constexpr (USE_REF) { + RegTensor refZeroReg; + RegTensor refOneReg; + LoadAlign(refZeroReg, refOneReg, ref); + if constexpr (NEGATIVE) { + SubFloatTwoReg(expZeroReg, expOneReg, refZeroReg, refOneReg, + gateZeroReg, gateOneReg, floatMask); + } else { + SubFloatTwoReg(expZeroReg, expOneReg, gateZeroReg, gateOneReg, + refZeroReg, refOneReg, floatMask); + } + } else if constexpr (NEGATIVE) { + Muls(expZeroReg, gateZeroReg, -1.0f, floatMask); + Muls(expOneReg, gateOneReg, -1.0f, floatMask); + } else { + Adds(expZeroReg, gateZeroReg, 0.0f, floatMask); + Adds(expOneReg, gateOneReg, 0.0f, floatMask); + } + Muls(expZeroReg, expZeroReg, LN2, floatMask); + Muls(expOneReg, expOneReg, LN2, floatMask); + MinsFloatTwoReg(expZeroReg, expOneReg, expZeroReg, expOneReg, + expInputMax, floatMask); + Maxs(expZeroReg, expZeroReg, expInputMin, floatMask); + Maxs(expOneReg, expOneReg, expInputMin, floatMask); + ExpFloatTwoReg(expZeroReg, expOneReg, expZeroReg, expOneReg, floatMask); +} + +template +__simd_callee__ inline void StoreKdaGateRegbasePair( + __ubuf__ OutputT *dst, + AscendC::MicroAPI::RegTensor &zeroReg, + AscendC::MicroAPI::RegTensor &oneReg, + AscendC::MicroAPI::MaskReg &inputMask, + AscendC::MicroAPI::MaskReg &floatMask) +{ + using namespace AscendC::MicroAPI; + RegTensor outputReg; + ClampKdaGateRegbaseOutput(zeroReg, oneReg, floatMask); + CastFloat2Half(outputReg, zeroReg, oneReg, floatMask); + StoreAlign(dst, outputReg, inputMask); +} + +template +static __simd_vf__ inline void PrepareKdaGateQwRegbase( + __ubuf__ InputT *q, __ubuf__ InputT *k, __ubuf__ OutputT *qOut, + __ubuf__ OutputT *kOut, __ubuf__ InputT *qDirect, __ubuf__ InputT *kDirect, + __ubuf__ GK_T *gate, __ubuf__ float *ref, uint16_t rows, uint16_t cols) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t ELEMENTS_PER_REG = AscendC::VECTOR_REG_WIDTH / sizeof(InputT); + + MaskReg floatMask = CreateMask(); + for (uint16_t row = 0; row < rows; ++row) { + uint32_t rowOffset = static_cast(row) * cols; + for (uint16_t col = 0; col < cols; col += ELEMENTS_PER_REG) { + uint32_t activeCount = static_cast(cols - col); + MaskReg inputMask = UpdateMask(activeCount); + uint32_t offset = rowOffset + col; + + RegTensor gateZeroReg; + RegTensor gateOneReg; + RegTensor expZeroReg; + RegTensor expOneReg; + RegTensor directZeroReg; + RegTensor directOneReg; + RegTensor inputZeroReg; + RegTensor inputOneReg; + RegTensor outputZeroReg; + RegTensor outputOneReg; + + LoadKdaGateRegbasePair(gateZeroReg, gateOneReg, gate + offset, inputMask); + BuildKdaGateRegbaseExp( + expZeroReg, expOneReg, gateZeroReg, gateOneReg, ref + col, floatMask); + BuildKdaGateRegbaseExp( + directZeroReg, directOneReg, gateZeroReg, gateOneReg, ref + col, floatMask); + + LoadKdaGateRegbasePair(inputZeroReg, inputOneReg, q + offset, inputMask); + MulFloatTwoReg(outputZeroReg, outputOneReg, inputZeroReg, inputOneReg, + directZeroReg, directOneReg, floatMask); + StoreKdaGateRegbasePair( + qDirect + offset, outputZeroReg, outputOneReg, inputMask, floatMask); + MulFloatTwoReg(outputZeroReg, outputOneReg, inputZeroReg, inputOneReg, + expZeroReg, expOneReg, floatMask); + StoreKdaGateRegbasePair( + qOut + offset, outputZeroReg, outputOneReg, inputMask, floatMask); + + LoadKdaGateRegbasePair(inputZeroReg, inputOneReg, k + offset, inputMask); + MulFloatTwoReg(outputZeroReg, outputOneReg, inputZeroReg, inputOneReg, + directZeroReg, directOneReg, floatMask); + StoreKdaGateRegbasePair( + kDirect + offset, outputZeroReg, outputOneReg, inputMask, floatMask); + MulFloatTwoReg(outputZeroReg, outputOneReg, inputZeroReg, inputOneReg, + expZeroReg, expOneReg, floatMask); + StoreKdaGateRegbasePair( + kOut + offset, outputZeroReg, outputOneReg, inputMask, floatMask); + } + } +} + +template +static __simd_vf__ inline void PrepareKdaGateKgRegbase( + __ubuf__ OutputT *kg, __ubuf__ InputT *k, __ubuf__ GK_T *gate, + __ubuf__ float *ref, uint16_t rows, uint16_t cols, uint16_t validRows) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t ELEMENTS_PER_REG = AscendC::VECTOR_REG_WIDTH / sizeof(InputT); + + MaskReg floatMask = CreateMask(); + for (uint16_t row = 0; row < rows; ++row) { + uint32_t rowOffset = static_cast(row) * cols; + for (uint16_t col = 0; col < cols; col += ELEMENTS_PER_REG) { + uint32_t activeCount = static_cast(cols - col); + MaskReg inputMask = UpdateMask(activeCount); + uint32_t offset = rowOffset + col; + + RegTensor gateZeroReg; + RegTensor gateOneReg; + RegTensor expZeroReg; + RegTensor expOneReg; + RegTensor inputZeroReg; + RegTensor inputOneReg; + RegTensor outputZeroReg; + RegTensor outputOneReg; + + LoadKdaGateRegbasePair(gateZeroReg, gateOneReg, gate + offset, inputMask); + BuildKdaGateRegbaseExp( + expZeroReg, expOneReg, gateZeroReg, gateOneReg, ref + col, floatMask); + LoadKdaGateRegbasePair(inputZeroReg, inputOneReg, k + offset, inputMask); + MulFloatTwoReg(outputZeroReg, outputOneReg, inputZeroReg, inputOneReg, + expZeroReg, expOneReg, floatMask); + if constexpr (USE_REF) { + if (row >= validRows) { + Duplicate(outputZeroReg, 0.0f, floatMask); + Duplicate(outputOneReg, 0.0f, floatMask); + } + } + StoreKdaGateRegbasePair( + kg + offset, outputZeroReg, outputOneReg, inputMask, floatMask); + } + } +} + +template +static __simd_vf__ inline void PrepareKdaGateQwKgRegbase( + __ubuf__ InputT *q, __ubuf__ InputT *k, __ubuf__ OutputT *qOut, + __ubuf__ OutputT *wOut, __ubuf__ OutputT *kgOut, __ubuf__ InputT *qDirect, + __ubuf__ InputT *wDirect, __ubuf__ InputT *v, __ubuf__ InputT *vDirect, + __ubuf__ InputT *finalKgOut, __ubuf__ float *beta, __ubuf__ GK_T *gate, + __ubuf__ float *ref, __ubuf__ float *finalRef, + uint16_t rows, uint16_t cols, uint16_t validRows) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t ELEMENTS_PER_REG = AscendC::VECTOR_REG_WIDTH / sizeof(InputT); + constexpr float scoreExpInputMax = + std::is_same() ? KDA_SCORE_EXP_INPUT_MAX : KDA_EXP_INPUT_MAX; + constexpr float scoreExpInputMin = + std::is_same() ? KDA_SCORE_EXP_INPUT_MIN : KDA_EXP_INPUT_MIN; + + MaskReg floatMask = CreateMask(); + for (uint16_t row = 0; row < rows; ++row) { + uint32_t rowOffset = static_cast(row) * cols; + for (uint16_t col = 0; col < cols; col += ELEMENTS_PER_REG) { + uint32_t activeCount = static_cast(cols - col); + MaskReg inputMask = UpdateMask(activeCount); + uint32_t offset = rowOffset + col; + + RegTensor gateZeroReg; + RegTensor gateOneReg; + RegTensor posZeroReg; + RegTensor posOneReg; + RegTensor negZeroReg; + RegTensor negOneReg; + RegTensor directZeroReg; + RegTensor directOneReg; + RegTensor finalNegZeroReg; + RegTensor finalNegOneReg; + RegTensor finalKZeroReg; + RegTensor finalKOneReg; + RegTensor inputZeroReg; + RegTensor inputOneReg; + RegTensor outputZeroReg; + RegTensor outputOneReg; + RegTensor betaReg; + + LoadKdaGateRegbasePair(gateZeroReg, gateOneReg, gate + offset, inputMask); + if constexpr (USE_REF) { + RegTensor refZeroReg; + RegTensor refOneReg; + LoadAlign(refZeroReg, refOneReg, ref + col); + SubFloatTwoReg(posZeroReg, posOneReg, gateZeroReg, gateOneReg, + refZeroReg, refOneReg, floatMask); + SubFloatTwoReg(negZeroReg, negOneReg, refZeroReg, refOneReg, + gateZeroReg, gateOneReg, floatMask); + } else { + Adds(posZeroReg, gateZeroReg, 0.0f, floatMask); + Adds(posOneReg, gateOneReg, 0.0f, floatMask); + Muls(negZeroReg, gateZeroReg, -1.0f, floatMask); + Muls(negOneReg, gateOneReg, -1.0f, floatMask); + } + if constexpr (STORE_DIRECT) { + Adds(directZeroReg, gateZeroReg, 0.0f, floatMask); + Adds(directOneReg, gateOneReg, 0.0f, floatMask); + Muls(directZeroReg, directZeroReg, LN2, floatMask); + Muls(directOneReg, directOneReg, LN2, floatMask); + MinsFloatTwoReg(directZeroReg, directOneReg, directZeroReg, directOneReg, + KDA_EXP_INPUT_MAX, floatMask); + Maxs(directZeroReg, directZeroReg, KDA_EXP_INPUT_MIN, floatMask); + Maxs(directOneReg, directOneReg, KDA_EXP_INPUT_MIN, floatMask); + ExpFloatTwoReg(directZeroReg, directOneReg, directZeroReg, directOneReg, floatMask); + } + Muls(posZeroReg, posZeroReg, LN2, floatMask); + Muls(posOneReg, posOneReg, LN2, floatMask); + Muls(negZeroReg, negZeroReg, LN2, floatMask); + Muls(negOneReg, negOneReg, LN2, floatMask); + MinsFloatTwoReg(posZeroReg, posOneReg, posZeroReg, posOneReg, + scoreExpInputMax, floatMask); + MinsFloatTwoReg(negZeroReg, negOneReg, negZeroReg, negOneReg, + scoreExpInputMax, floatMask); + Maxs(posZeroReg, posZeroReg, scoreExpInputMin, floatMask); + Maxs(posOneReg, posOneReg, scoreExpInputMin, floatMask); + Maxs(negZeroReg, negZeroReg, scoreExpInputMin, floatMask); + Maxs(negOneReg, negOneReg, scoreExpInputMin, floatMask); + ExpFloatTwoReg(posZeroReg, posOneReg, posZeroReg, posOneReg, floatMask); + ExpFloatTwoReg(negZeroReg, negOneReg, negZeroReg, negOneReg, floatMask); + + LoadKdaGateRegbasePair(inputZeroReg, inputOneReg, q + offset, inputMask); + if constexpr (STORE_DIRECT) { + MulFloatTwoReg(outputZeroReg, outputOneReg, inputZeroReg, inputOneReg, + directZeroReg, directOneReg, floatMask); + StoreKdaGateRegbasePair( + qDirect + offset, outputZeroReg, outputOneReg, inputMask, floatMask); + } + MulFloatTwoReg(outputZeroReg, outputOneReg, inputZeroReg, inputOneReg, + posZeroReg, posOneReg, floatMask); + StoreKdaGateRegbasePair( + qOut + offset, outputZeroReg, outputOneReg, inputMask, floatMask); + + LoadKdaGateRegbasePair(inputZeroReg, inputOneReg, k + offset, inputMask); + if constexpr (EXPORT_FINAL_KG) { + // wOut may alias k, and STORE_DIRECT later reuses inputZeroReg/inputOneReg for V. + Adds(finalKZeroReg, inputZeroReg, 0.0f, floatMask); + Adds(finalKOneReg, inputOneReg, 0.0f, floatMask); + } + if constexpr (STORE_DIRECT || SCALE_SCORE_W) { + LoadAlign(betaReg, beta + row); + } + if constexpr (STORE_DIRECT) { + MulFloatTwoReg(outputZeroReg, outputOneReg, inputZeroReg, inputOneReg, + directZeroReg, directOneReg, floatMask); + RegTensor roundedReg; + ClampKdaGateRegbaseOutput(outputZeroReg, outputOneReg, floatMask); + CastFloat2Half(roundedReg, outputZeroReg, outputOneReg, floatMask); + CastHalf2Float(outputZeroReg, outputOneReg, roundedReg, inputMask); + Mul(outputZeroReg, outputZeroReg, betaReg, floatMask); + Mul(outputOneReg, outputOneReg, betaReg, floatMask); + StoreKdaGateRegbasePair( + wDirect + offset, outputZeroReg, outputOneReg, inputMask, floatMask); + } + MulFloatTwoReg(outputZeroReg, outputOneReg, inputZeroReg, inputOneReg, + posZeroReg, posOneReg, floatMask); + if constexpr (SCALE_SCORE_W) { + Mul(outputZeroReg, outputZeroReg, betaReg, floatMask); + Mul(outputOneReg, outputOneReg, betaReg, floatMask); + } + StoreKdaGateRegbasePair( + wOut + offset, outputZeroReg, outputOneReg, inputMask, floatMask); + MulFloatTwoReg(outputZeroReg, outputOneReg, inputZeroReg, inputOneReg, + negZeroReg, negOneReg, floatMask); + if constexpr (USE_REF) { + if (row >= validRows) { + Duplicate(outputZeroReg, 0.0f, floatMask); + Duplicate(outputOneReg, 0.0f, floatMask); + } + } + StoreKdaGateRegbasePair( + kgOut + offset, outputZeroReg, outputOneReg, inputMask, floatMask); + + if constexpr (STORE_DIRECT) { + LoadKdaGateRegbasePair(inputZeroReg, inputOneReg, v + offset, inputMask); + Mul(outputZeroReg, inputZeroReg, betaReg, floatMask); + Mul(outputOneReg, inputOneReg, betaReg, floatMask); + StoreKdaGateRegbasePair( + vDirect + offset, outputZeroReg, outputOneReg, inputMask, floatMask); + } + if constexpr (EXPORT_FINAL_KG) { + RegTensor finalRefZeroReg; + RegTensor finalRefOneReg; + LoadAlign( + finalRefZeroReg, finalRefOneReg, finalRef + col); + SubFloatTwoReg(finalNegZeroReg, finalNegOneReg, + finalRefZeroReg, finalRefOneReg, + gateZeroReg, gateOneReg, floatMask); + Muls(finalNegZeroReg, finalNegZeroReg, LN2, floatMask); + Muls(finalNegOneReg, finalNegOneReg, LN2, floatMask); + MinsFloatTwoReg(finalNegZeroReg, finalNegOneReg, + finalNegZeroReg, finalNegOneReg, + KDA_EXP_INPUT_MAX, floatMask); + Maxs(finalNegZeroReg, finalNegZeroReg, KDA_EXP_INPUT_MIN, floatMask); + Maxs(finalNegOneReg, finalNegOneReg, KDA_EXP_INPUT_MIN, floatMask); + ExpFloatTwoReg(finalNegZeroReg, finalNegOneReg, + finalNegZeroReg, finalNegOneReg, floatMask); + MulFloatTwoReg(outputZeroReg, outputOneReg, finalKZeroReg, finalKOneReg, + finalNegZeroReg, finalNegOneReg, floatMask); + StoreKdaGateRegbasePair( + finalKgOut + offset, outputZeroReg, outputOneReg, inputMask, floatMask); + } + } + } +} + +static __simd_vf__ inline void ForwardSubDiag16Regbase(__ubuf__ float *diag, uint16_t valid) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t DIAG_SIZE = KDA_SOLVE_DIAG_BT; + uint32_t activeCount = DIAG_SIZE; + MaskReg rowMask = UpdateMask(activeCount); + + for (uint16_t row = 2; row < valid; ++row) { + RegTensor currentReg; + RegTensor scaleReg; + RegTensor matrixReg; + RegTensor productReg; + RegTensor sumReg; + LoadAlign(currentReg, diag + static_cast(row) * DIAG_SIZE); + Duplicate(sumReg, 0.0f, rowMask); + + for (uint16_t sourceRow = 0; sourceRow < row; ++sourceRow) { + LoadAlign( + scaleReg, diag + static_cast(row) * DIAG_SIZE + sourceRow); + LoadAlign(matrixReg, diag + static_cast(sourceRow) * DIAG_SIZE); + Mul(productReg, matrixReg, scaleReg, rowMask); + Add(sumReg, sumReg, productReg, rowMask); + } + Add(currentReg, currentReg, sumReg, rowMask); + StoreAlign(diag + static_cast(row) * DIAG_SIZE, currentReg, rowMask); + LocalMemBar(); + } +} + +static __simd_vf__ inline void SelectCausalRows64Regbase( + __ubuf__ float *aqk, __ubuf__ float *akk, uint16_t rowBegin, uint16_t rowCount) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t ROW_ELEMENTS = 64; + + MaskReg fullMask = CreateMask(); + for (uint16_t localRow = 0; localRow < rowCount; ++localRow) { + const uint16_t row = rowBegin + localRow; + const uint32_t rowOffset = static_cast(localRow) * ROW_ELEMENTS; + RegTensor zeroReg; + RegTensor aqkInputReg; + RegTensor akkInputReg; + RegTensor aqkReg; + RegTensor akkReg; + uint32_t aqkCount = static_cast(row) + 1; + uint32_t akkCount = static_cast(row); + MaskReg aqkMask = UpdateMask(aqkCount); + MaskReg akkMask = UpdateMask(akkCount); + Duplicate(zeroReg, 0.0f, fullMask); + LoadAlign(aqkInputReg, aqk + rowOffset); + LoadAlign(akkInputReg, akk + rowOffset); + Select(aqkReg, aqkInputReg, zeroReg, aqkMask); + Select(akkReg, akkInputReg, zeroReg, akkMask); + StoreAlign(aqk + rowOffset, aqkReg, fullMask); + StoreAlign(akk + rowOffset, akkReg, fullMask); + } +} + +static __simd_vf__ inline void ForwardSubDiag16StridedRegbase( + __ubuf__ float *matrix, uint16_t rowStride, uint16_t rowBegin, uint16_t colBegin, + uint16_t valid) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t DIAG_SIZE = KDA_SOLVE_DIAG_BT; + uint32_t activeCount = DIAG_SIZE; + MaskReg rowMask = UpdateMask(activeCount); + + for (uint16_t row = 2; row < valid; ++row) { + uint32_t currentOffset = + static_cast(rowBegin + row) * rowStride + colBegin; + RegTensor currentReg; + RegTensor scaleReg; + RegTensor matrixReg; + RegTensor productReg; + RegTensor sumReg; + LoadAlign(currentReg, matrix + currentOffset); + Duplicate(sumReg, 0.0f, rowMask); + + for (uint16_t sourceRow = 0; sourceRow < row; ++sourceRow) { + LoadAlign( + scaleReg, matrix + currentOffset + sourceRow); + uint32_t sourceOffset = + static_cast(rowBegin + sourceRow) * rowStride + colBegin; + LoadAlign(matrixReg, matrix + sourceOffset); + Mul(productReg, matrixReg, scaleReg, rowMask); + Add(sumReg, sumReg, productReg, rowMask); + } + Add(currentReg, currentReg, sumReg, rowMask); + StoreAlign(matrix + currentOffset, currentReg, rowMask); + LocalMemBar(); + } +} + +static __simd_vf__ inline void ForwardSubDiag16PairStridedRegbase( + __ubuf__ float *matrix, uint16_t rowStride, uint16_t colBegin, + uint16_t firstValid, uint16_t secondValid) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t DIAG_SIZE = KDA_SOLVE_DIAG_BT; + constexpr uint16_t SECOND_LOCAL_ROW = DIAG_SIZE; + uint32_t activeCount = DIAG_SIZE; + MaskReg rowMask = UpdateMask(activeCount); + + for (uint16_t row = 2; row < DIAG_SIZE; ++row) { + uint32_t firstCurrentOffset = + static_cast(row) * rowStride + colBegin; + uint32_t secondCurrentOffset = + static_cast(SECOND_LOCAL_ROW + row) * rowStride + + colBegin + DIAG_SIZE; + RegTensor firstCurrentReg; + RegTensor secondCurrentReg; + RegTensor firstSumReg; + RegTensor secondSumReg; + LoadAlign(firstCurrentReg, matrix + firstCurrentOffset); + LoadAlign(secondCurrentReg, matrix + secondCurrentOffset); + Duplicate(firstSumReg, 0.0f, rowMask); + Duplicate(secondSumReg, 0.0f, rowMask); + + if (row < firstValid || row < secondValid) { + for (uint16_t sourceRow = 0; sourceRow < row; ++sourceRow) { + RegTensor firstScaleReg; + RegTensor secondScaleReg; + RegTensor firstMatrixReg; + RegTensor secondMatrixReg; + RegTensor firstProductReg; + RegTensor secondProductReg; + LoadAlign( + firstScaleReg, matrix + firstCurrentOffset + sourceRow); + LoadAlign( + secondScaleReg, matrix + secondCurrentOffset + sourceRow); + uint32_t firstSourceOffset = + static_cast(sourceRow) * rowStride + colBegin; + uint32_t secondSourceOffset = + static_cast(SECOND_LOCAL_ROW + sourceRow) * rowStride + + colBegin + DIAG_SIZE; + LoadAlign(firstMatrixReg, matrix + firstSourceOffset); + LoadAlign(secondMatrixReg, matrix + secondSourceOffset); + Mul(firstProductReg, firstMatrixReg, firstScaleReg, rowMask); + Mul(secondProductReg, secondMatrixReg, secondScaleReg, rowMask); + Add(firstSumReg, firstSumReg, firstProductReg, rowMask); + Add(secondSumReg, secondSumReg, secondProductReg, rowMask); + } + } + if (row < firstValid) { + Add(firstCurrentReg, firstCurrentReg, firstSumReg, rowMask); + StoreAlign(matrix + firstCurrentOffset, firstCurrentReg, rowMask); + } + if (row < secondValid) { + Add(secondCurrentReg, secondCurrentReg, secondSumReg, rowMask); + StoreAlign(matrix + secondCurrentOffset, secondCurrentReg, rowMask); + } + LocalMemBar(); + } + + RegTensor indexReg; + Arange(indexReg, 0); + for (uint16_t row = 0; row < DIAG_SIZE; ++row) { + MaskReg diagMask; + CompareScalar( + diagMask, indexReg, static_cast(row), rowMask); + uint32_t firstOffset = static_cast(row) * rowStride + colBegin; + uint32_t secondOffset = + static_cast(SECOND_LOCAL_ROW + row) * rowStride + + colBegin + DIAG_SIZE; + RegTensor firstReg; + RegTensor secondReg; + LoadAlign(firstReg, matrix + firstOffset); + LoadAlign(secondReg, matrix + secondOffset); + Adds(firstReg, firstReg, 1.0f, diagMask); + Adds(secondReg, secondReg, 1.0f, diagMask); + StoreAlign(matrix + firstOffset, firstReg, rowMask); + StoreAlign(matrix + secondOffset, secondReg, rowMask); + } +} + +static __simd_vf__ inline void ApplyKdaRowScaleRegbase( + __ubuf__ float *matrix, __ubuf__ float *rowScale, uint16_t rows, uint16_t cols) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t FP32_PER_REG = AscendC::VECTOR_REG_WIDTH / sizeof(float); + RegTensor matrixReg0; + RegTensor matrixReg1; + RegTensor scaleReg0; + RegTensor scaleReg1; + + uint16_t row = 0; + for (; row + 1 < rows; row += 2) { + LoadAlign(scaleReg0, rowScale + row); + LoadAlign(scaleReg1, rowScale + row + 1); + for (uint16_t col = 0; col < cols; col += FP32_PER_REG) { + uint32_t activeCount0 = static_cast(cols - col); + uint32_t activeCount1 = activeCount0; + MaskReg mask0 = UpdateMask(activeCount0); + MaskReg mask1 = UpdateMask(activeCount1); + uint32_t offset0 = static_cast(row) * cols + col; + uint32_t offset1 = static_cast(row + 1) * cols + col; + LoadAlign(matrixReg0, matrix + offset0); + LoadAlign(matrixReg1, matrix + offset1); + Mul(matrixReg0, matrixReg0, scaleReg0, mask0); + Mul(matrixReg1, matrixReg1, scaleReg1, mask1); + StoreAlign(matrix + offset0, matrixReg0, mask0); + StoreAlign(matrix + offset1, matrixReg1, mask1); + } + } + if (row < rows) { + LoadAlign(scaleReg0, rowScale + row); + for (uint16_t col = 0; col < cols; col += FP32_PER_REG) { + uint32_t activeCount = static_cast(cols - col); + MaskReg mask = UpdateMask(activeCount); + uint32_t offset = static_cast(row) * cols + col; + LoadAlign(matrixReg0, matrix + offset); + Mul(matrixReg0, matrixReg0, scaleReg0, mask); + StoreAlign(matrix + offset, matrixReg0, mask); + } + } +} +#endif + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +using KdaArchTag = Catlass::Arch::Ascend950; +#else +using KdaArchTag = Catlass::Arch::AtlasA2; +#endif +using KdaDispatchPolicy = Catlass::Gemm::MmadPingpong; +using KdaScoreDispatchPolicy = + Catlass::Gemm::MmadPingpongTlaMulti; +static_assert(KdaScoreDispatchPolicy::ENABLE_L1_RESIDENT, + "KDA Aqk/Akk score MMAD must keep the shared right matrix resident in L1"); +static_assert(KdaScoreDispatchPolicy::L1B_STAGES == 1, + "KDA Aqk/Akk score MMAD needs one L1 B slot so the second MMAD reuses it"); +using KdaSolveDispatchPolicy = Catlass::Gemm::MmadPingpong; +static_assert(!KdaSolveDispatchPolicy::USE_HF32_MODE, "KDA triangular solve must use IEEE FP32 Cube mode"); +using KdaL1TileShape = tla::Shape; +using KdaL0TileShape = KdaL1TileShape; +using KdaSolveL1TileShape = tla::Shape; +using KdaSolveL0TileShape = KdaSolveL1TileShape; + +__aicore__ inline uint32_t FloatToBits(float value) +{ + union Bits { + __aicore__ Bits() {} + float f; + uint32_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline float BitsToFloat(uint32_t value) +{ + union Bits { + __aicore__ Bits() {} + uint32_t u; + float f; + } bits; + bits.u = value; + return bits.f; +} + +__aicore__ inline uint16_t Bf16ToBits(bfloat16_t value) +{ + union Bits { + __aicore__ Bits() {} + bfloat16_t f; + uint16_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline bfloat16_t BitsToBf16(uint16_t value) +{ + union Bits { + __aicore__ Bits() {} + uint16_t u; + bfloat16_t f; + } bits; + bits.u = value; + return bits.f; +} + +template +__aicore__ inline T FloatToType(float value) +{ + if constexpr (IsSameType::value) { + uint32_t bits = FloatToBits(value); + uint32_t bias = 0x7FFFu + ((bits >> 16) & 1u); + return BitsToBf16(static_cast((bits + bias) >> 16)); + } + return static_cast(value); +} + +template +class ChunkKdaFwdPrepareKernel { +public: + using OUT_T = T; + using AKK_T = float; + using SCORE_T = + std::conditional_t::value, bfloat16_t, T>; + template + __aicore__ inline void Init(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR rawG, + GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR preparedQG, GM_ADDR preparedAqk, + GM_ADDR propagatedVNew, GM_ADDR propagatedH, GM_ADDR o, GM_ADDR finalState, GM_ADDR aqk, + GM_ADDR akk, GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR finalKg, GM_ADDR workspace, const TilingData &tiling, TPipe *pipe, + bool initVecBuffers = true, bool storeQG = true) + { + pipe_ = pipe; + q_.SetGlobalBuffer((__gm__ T *)q); + k_.SetGlobalBuffer((__gm__ T *)k); + v_.SetGlobalBuffer((__gm__ T *)v); + gk_.SetGlobalBuffer((__gm__ GK_T *)gk); + rawG_.SetGlobalBuffer((__gm__ float *)rawG); + aLog_.SetGlobalBuffer((__gm__ float *)aLog); + dtBias_.SetGlobalBuffer((__gm__ float *)dtBias); + beta_.SetGlobalBuffer((__gm__ BETA_T *)beta); + if (initialState != nullptr) { + initialState_.SetGlobalBuffer((__gm__ float *)initialState); + } + cuSeqlensAddr_ = reinterpret_cast<__gm__ int64_t *>(cuSeqlens); + if (preparedQG != nullptr) { + preparedQG_.SetGlobalBuffer((__gm__ T *)preparedQG); + } + if (preparedAqk != nullptr) { + preparedAqk_.SetGlobalBuffer((__gm__ T *)preparedAqk); + } + if (propagatedVNew != nullptr) { + propagatedVNew_.SetGlobalBuffer((__gm__ T *)propagatedVNew); + } + if (propagatedH != nullptr) { + propagatedH_.SetGlobalBuffer((__gm__ T *)propagatedH); + } + chunkIndicesAddr_ = reinterpret_cast<__gm__ int64_t *>(chunkIndices); + o_.SetGlobalBuffer((__gm__ OUT_T *)o); + finalState_.SetGlobalBuffer((__gm__ float *)finalState); + aqk_.SetGlobalBuffer((__gm__ float *)aqk); + akk_.SetGlobalBuffer((__gm__ AKK_T *)akk); + w_.SetGlobalBuffer((__gm__ T *)w); + u_.SetGlobalBuffer((__gm__ OUT_T *)u); + qg_.SetGlobalBuffer((__gm__ T *)qg); + kg_.SetGlobalBuffer((__gm__ T *)kg); + vNew_.SetGlobalBuffer((__gm__ T *)vNew); + h_.SetGlobalBuffer((__gm__ float *)h); + finalKg_.SetGlobalBuffer((__gm__ T *)finalKg); + solveWorkspace_.SetGlobalBuffer((__gm__ float *)workspace); + + B_ = tiling.batch; + N_ = tiling.seqNum; + H_ = tiling.qHeadNum; + HV_ = tiling.vHeadNum; + T_ = tiling.seqlen; + K_ = COMPILE_K == 0 ? tiling.kHeadDim : COMPILE_K; + V_ = COMPILE_V == 0 ? tiling.vHeadDim : COMPILE_V; + BT_ = COMPILE_BT == 0 ? tiling.chunkSize : COMPILE_BT; + NT_ = tiling.totalChunks; + scale_ = tiling.scale; + hasInitial_ = tiling.hasInitialState; + isVarLen_ = tiling.isVarLen; + inputSequenceMajor_ = tiling.inputSequenceMajor; + fusePostWu_ = tiling.fusePostWu; + materializeFinalKg_ = tiling.fusePostWu || tiling.fusePostWuIntoFwdH; + computeGateInPrepare_ = tiling.computeGateInPrepare; + hasALog_ = tiling.hasALog; + hasDtBias_ = tiling.hasDtBias; + lowerBound_ = tiling.lowerBound; + storeQG_ = storeQG; + usedCoreNum_ = tiling.prepareUsedCoreNum; +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE && !IsSameType::value) { + headPairMode_ = KDA_ARCH35_ENABLE_HEAD_PAIR && + HV_ % KDA_SCORE_LANES == 0; + } +#endif + constexpr uint64_t solvePipelineDepth = SAFE_GATE ? KDA_SOLVE_PIPELINE_DEPTH : 1; + const uint64_t solveBytes = + usedCoreNum_ * solvePipelineDepth * KDA_SOLVE_SCRATCH_SLOTS * BT_ * BT_ * sizeof(float); + const uint64_t alignedSolveBytes = + (solveBytes + KDA_WORKSPACE_ALIGN - 1) / KDA_WORKSPACE_ALIGN * KDA_WORKSPACE_ALIGN; + scoreWorkspace_.SetGlobalBuffer((__gm__ SCORE_T *)(workspace + alignedSolveBytes)); + if ASCEND_IS_AIV { + uint64_t subBlockNum = static_cast(GetSubBlockNum()); + solveCoreIdx_ = subBlockNum == 0 ? 0 : static_cast(GetBlockIdx()) / subBlockNum; + } else { + solveCoreIdx_ = static_cast(GetBlockIdx()); + } + if (pipe_ != nullptr && initVecBuffers) { + pipe_->InitBuffer(exp2Buf_, EXP2_UB_BYTES); + pipe_->InitBuffer(vecBuf_, KDA_VEC_ARENA_ELEMENTS * sizeof(float)); + const uint64_t gateStageElems = GatePipelineRows() * K_; + const uint64_t gateInputSlotBytes = GateInputSlotBytes(); + const uint64_t gatePipelineBytes = + GateBufferDepth() * (gateInputSlotBytes + gateStageElems * sizeof(T)); + pipe_->InitBuffer(gateWritebackBuf_, static_cast(gatePipelineBytes)); + AllocVectorEvents(); + } + } + __aicore__ inline void ProcessAivOnly() + { + isAivOnly_ = true; + ProcessPreAiv(); + ReleaseVectorEvents(); + } + + __aicore__ inline void ProcessAiv() + { + ProcessPreAiv(); + ReleaseVectorEvents(); + } + + __aicore__ inline void ProcessAic() + { + ProcessPreAic(); + } + + template + __aicore__ inline void ProcessAicFused(PostWuOp &postWu) + { + ProcessPreAicHeadPairFused(postWu); + } + +private: + __aicore__ inline void AllocVectorEvents() + { + mte2ToVEvent_ = pipe_->AllocEventID(); + vToMte2Event_ = pipe_->AllocEventID(); + vToMte3Event_ = pipe_->AllocEventID(); + mte3ToVEvent_ = pipe_->AllocEventID(); + mte2ToMte3Event_ = pipe_->AllocEventID(); + vToSEvent_ = pipe_->AllocEventID(); + for (uint32_t slot = 0; slot < KDA_GATE_PIPELINE_DEPTH; ++slot) { + mte3ToMte2Events_[slot] = pipe_->AllocEventID(); + } + vectorEventsAllocated_ = true; + } + + __aicore__ inline void ReleaseVectorEvents() + { + if (!vectorEventsAllocated_) { + return; + } + pipe_->ReleaseEventID(mte2ToVEvent_); + pipe_->ReleaseEventID(vToMte2Event_); + pipe_->ReleaseEventID(vToMte3Event_); + pipe_->ReleaseEventID(mte3ToVEvent_); + pipe_->ReleaseEventID(mte2ToMte3Event_); + pipe_->ReleaseEventID(vToSEvent_); + for (uint32_t slot = 0; slot < KDA_GATE_PIPELINE_DEPTH; ++slot) { + pipe_->ReleaseEventID(mte3ToMte2Events_[slot]); + } + vectorEventsAllocated_ = false; + } + + __aicore__ inline uint64_t QOffset(uint64_t b, uint64_t h, uint64_t t, uint64_t d) const + { + if (inputSequenceMajor_) { + return ((b * T_ + t) * H_ + h) * K_ + d; + } + return ((b * H_ + h) * T_ + t) * K_ + d; + } + + __aicore__ inline uint64_t VInputOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t d) const + { + if (inputSequenceMajor_) { + return ((b * T_ + t) * HV_ + hv) * V_ + d; + } + return ((b * HV_ + hv) * T_ + t) * V_ + d; + } + + __aicore__ inline uint64_t RawGateOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t d) const + { + if (inputSequenceMajor_) { + return ((b * T_ + t) * HV_ + hv) * K_ + d; + } + return ((b * HV_ + hv) * T_ + t) * K_ + d; + } + + __aicore__ inline uint64_t KVOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t d, uint64_t dim) const + { + return ((b * HV_ + hv) * T_ + t) * dim + d; + } + + __aicore__ inline uint64_t BetaOffset(uint64_t b, uint64_t hv, uint64_t t) const + { + return (b * HV_ + hv) * T_ + t; + } + + __aicore__ inline uint64_t AOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t j) const + { + return ((b * HV_ + hv) * T_ + t) * BT_ + j; + } + + __aicore__ inline uint64_t HOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t d, uint64_t r) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * K_ + d) * V_ + r; + } + + __aicore__ inline uint64_t WScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t t, uint64_t d) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * BT_ + t) * K_ + d; + } + + __aicore__ inline uint64_t SolveScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t slot) const + { + (void)b; + (void)hv; + (void)chunkIdx; + constexpr uint64_t solvePipelineDepth = SAFE_GATE ? KDA_SOLVE_PIPELINE_DEPTH : 1; + uint64_t matrixElements = BT_ * BT_; + return ((solveCoreIdx_ * solvePipelineDepth + activeSolveSlot_) * KDA_SOLVE_SCRATCH_SLOTS + slot) * + matrixElements; + } + + __aicore__ inline uint64_t ScoreScratchOffset(uint64_t slot, uint64_t plane, uint64_t t = 0, + uint64_t d = 0) const + { + return (((solveCoreIdx_ * KDA_SCORE_SCRATCH_SLOTS + slot) * KDA_SCORE_SCRATCH_PLANES + plane) * BT_ + t) * + K_ + + d; + } + + __aicore__ inline uint64_t ScoreScratchSlot(uint64_t queueSlot, uint64_t lane, bool pairHeads) const + { + return pairHeads ? queueSlot * KDA_SCORE_LANES + lane : queueSlot; + } + + + + __aicore__ inline uint64_t ScoreRefBlockSize() const + { + if constexpr (SAFE_GATE) { + return KDA_SAFE_SCORE_REF_BC; + } + return KDA_SCORE_REF_BC; + } + + __aicore__ inline uint64_t ScoreRowBlockCount(uint64_t curT, uint64_t rowBegin) const + { + uint64_t blockSize = ScoreRefBlockSize(); + uint64_t rowCount = curT - rowBegin; + if (rowCount > blockSize) { + rowCount = blockSize; + } + return rowCount; + } + + __aicore__ inline uint64_t ScoreRefToken(uint64_t start, uint64_t curT, uint64_t rowBegin, + uint64_t rowCount) const + { + uint64_t ref = rowBegin + rowCount / 2; + if (ref >= curT) { + ref = curT - 1; + } + return start + ref; + } + + __aicore__ inline void RunExp2(LocalTensor &tensor, uint32_t count) + { + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + ClampExpInput(tensor, count); + Exp(tensor, tensor, count); + PipeBarrier(); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + } + + __aicore__ inline void ClampExpInput(LocalTensor &tensor, uint32_t count) + { + Mins(tensor, tensor, KDA_EXP_INPUT_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, KDA_EXP_INPUT_MIN, count); + PipeBarrier(); + } + + __aicore__ inline void ClampScoreExpInput(LocalTensor &tensor, uint32_t count) + { + constexpr float expInputMax = + IsSameType::value ? KDA_SCORE_EXP_INPUT_MAX : KDA_EXP_INPUT_MAX; + constexpr float expInputMin = + IsSameType::value ? KDA_SCORE_EXP_INPUT_MIN : KDA_EXP_INPUT_MIN; + Mins(tensor, tensor, expInputMax, count); + PipeBarrier(); + Maxs(tensor, tensor, expInputMin, count); + PipeBarrier(); + } + + template + __aicore__ inline void ClampFp32ForCast(LocalTensor &tensor, uint32_t count) + { + if constexpr (IsSameType::value) { + Mins(tensor, tensor, KDA_FP16_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, -KDA_FP16_MAX, count); + PipeBarrier(); + } + } + + __aicore__ inline void ClampFp32ToOutputType(LocalTensor &tensor, uint32_t count) + { + ClampFp32ForCast(tensor, count); + } + + template + __aicore__ inline void CopyVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst, src[offset], static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + + template + __aicore__ inline void CopyRowsIn(LocalTensor &dst, GlobalTensor &src, + uint64_t offset, uint64_t rows, uint64_t cols, + uint64_t rowStride) + { + if (rows == 0 || cols == 0) { + return; + } + if (rowStride == cols) { + CopyVectorIn(dst, src, offset, rows * cols); + return; + } + DataCopyExtParams params{ + static_cast(rows), + static_cast(cols * sizeof(CopyT)), + static_cast((rowStride - cols) * sizeof(CopyT)), + 0, + 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + + template + __aicore__ inline void CopyVectorOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst[offset], src, static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPad(dst[offset], src, params); + } + + template + __aicore__ inline void CopyRowIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset) + { + CopyVectorIn(dst, src, offset, K_); + } + + template + __aicore__ inline void CopyRowOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src) + { + CopyVectorOut(dst, offset, src, K_); + } + + __aicore__ inline LocalTensor VecScratch(uint64_t slot) + { + return vecBuf_.Get()[slot * EXP2_UB_ELEMENTS]; + } + + __aicore__ inline uint64_t GateStageElems() const + { + return GatePipelineRows() * K_; + } + + __aicore__ inline uint64_t GatePipelineRows() const + { + constexpr uint64_t fixedBytes = + static_cast(KDA_VEC_ARENA_ELEMENTS) * sizeof(float) + EXP2_UB_BYTES; + constexpr uint64_t availableBytes = KDA_AIV_UB_BUDGET_BYTES - fixedBytes; + uint64_t bytesPerRow = K_ * KDA_GATE_PIPELINE_DEPTH * (3 * sizeof(T) + sizeof(GK_T)); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128) { + bytesPerRow = GateBufferDepth() * + (K_ * (4 * sizeof(T) + sizeof(GK_T)) + sizeof(BETA_T)); + } +#endif + uint64_t rows = bytesPerRow == 0 ? 0 : availableBytes / bytesPerRow; + return rows < KDA_GATE_TILE_ROWS ? rows : KDA_GATE_TILE_ROWS; + } + + __aicore__ inline constexpr uint64_t GateBufferDepth() const + { +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128) { + return 2; + } +#endif + return KDA_GATE_PIPELINE_DEPTH; + } + + __aicore__ inline uint64_t GateInputSlotBytes() const + { +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128) { + return GateStageElems() * (3 * sizeof(T) + sizeof(GK_T)) + + GatePipelineRows() * sizeof(BETA_T); + } +#endif + return GateStageElems() * (2 * sizeof(T) + sizeof(GK_T)); + } + + __aicore__ inline LocalTensor GateQTyped(uint64_t slot) + { + uint64_t byteOffset = slot * GateInputSlotBytes(); + return gateWritebackBuf_.Get()[byteOffset / sizeof(T)]; + } + + __aicore__ inline LocalTensor GateKTyped(uint64_t slot) + { + uint64_t byteOffset = slot * GateInputSlotBytes() + GateStageElems() * sizeof(T); + return gateWritebackBuf_.Get()[byteOffset / sizeof(T)]; + } + + __aicore__ inline LocalTensor GateGTyped(uint64_t slot) + { + uint64_t byteOffset = slot * GateInputSlotBytes() + 2 * GateStageElems() * sizeof(T); + return gateWritebackBuf_.Get()[byteOffset / sizeof(GK_T)]; + } + + __aicore__ inline LocalTensor GateVTyped(uint64_t slot) + { + uint64_t byteOffset = slot * GateInputSlotBytes() + + GateStageElems() * (2 * sizeof(T) + sizeof(GK_T)); + return gateWritebackBuf_.Get()[byteOffset / sizeof(T)]; + } + + __aicore__ inline LocalTensor GateBetaTyped(uint64_t slot) + { + uint64_t byteOffset = slot * GateInputSlotBytes() + + GateStageElems() * (3 * sizeof(T) + sizeof(GK_T)); + return gateWritebackBuf_.Get()[byteOffset / sizeof(BETA_T)]; + } + + __aicore__ inline LocalTensor GateKgTyped(uint64_t slot) + { + uint64_t byteOffset = GateBufferDepth() * GateInputSlotBytes() + + slot * GateStageElems() * sizeof(T); + return gateWritebackBuf_.Get()[byteOffset / sizeof(T)]; + } + + __aicore__ inline LocalTensor LocalGateChunk() + { + constexpr uint64_t byteOffset = + static_cast(KDA_LOCAL_GK_FLOAT_OFFSET) * sizeof(float); + return vecBuf_.Get()[byteOffset / sizeof(GK_T)]; + } + + __aicore__ inline LocalTensor GateScoreTyped(uint64_t slot, uint64_t tileRow) + { + if constexpr (IsSameType::value) { + if (computeGateInPrepare_) { + return LocalGateChunk()[tileRow * K_]; + } + } + return GateGTyped(slot); + } + + __aicore__ inline void LoadGateScoreRef( + LocalTensor dst, uint64_t b, uint64_t hv, uint64_t token) + { + if constexpr (IsSameType::value) { + if (computeGateInPrepare_) { + const uint64_t tileRow = token - activeGateChunkStart_; + Adds(dst, LocalGateChunk()[tileRow * K_], 0.0f, static_cast(K_)); + PipeBarrier(); + return; + } + } + LoadAsFloatRow(gk_, KVOffset(b, hv, token, 0, K_), dst, K_); + } + + __aicore__ inline void PrefetchQKGate(uint64_t slot, uint64_t b, uint64_t h, uint64_t hv, + uint64_t token, uint64_t elems) + { + const uint64_t rows = elems / K_; + LocalTensor qTyped = GateQTyped(slot); + LocalTensor kTyped = GateKTyped(slot); + LocalTensor gateTyped = GateGTyped(slot); + CopyRowsIn(qTyped, q_, QOffset(b, h, token, 0), rows, K_, inputSequenceMajor_ ? H_ * K_ : K_); + CopyRowsIn(kTyped, k_, QOffset(b, h, token, 0), rows, K_, inputSequenceMajor_ ? H_ * K_ : K_); + if (!computeGateInPrepare_) { + CopyVectorIn(gateTyped, gk_, KVOffset(b, hv, token, 0, K_), elems); + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128) { + LocalTensor vTyped = GateVTyped(slot); + LocalTensor betaTyped = GateBetaTyped(slot); + CopyRowsIn(vTyped, v_, VInputOffset(b, hv, token, 0), rows, V_, + inputSequenceMajor_ ? HV_ * V_ : V_); + CopyVectorIn(betaTyped, beta_, BetaOffset(b, hv, token), rows); + } +#endif + SetFlag(mte2ToVEvent_); + } + + __aicore__ inline float LoadGateExpA(uint64_t hv) + { + if (!hasALog_) { + return 1.0f; + } + LocalTensor scalar = exp2Buf_.Get(); + CopyVectorIn(scalar, aLog_, hv, 1); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Exp(scalar, scalar, 1); + PipeBarrier(); + SetFlag(vToSEvent_); + WaitFlag(vToSEvent_); + __ubuf__ float *ptr = (__ubuf__ float *)scalar.GetPhyAddr(); + return ptr[0]; + } + + __aicore__ inline void PrefetchRawGateTile( + uint64_t slot, uint64_t b, uint64_t hv, uint64_t token, uint64_t rows) + { + (void)slot; + const uint64_t tileRow = token - activeGateChunkStart_; + LocalTensor gate = + LocalGateChunk().template ReinterpretCast()[tileRow * K_]; + CopyRowsIn(gate, rawG_, RawGateOffset(b, hv, token, 0), rows, K_, + inputSequenceMajor_ ? HV_ * K_ : K_); + SetFlag(mte2ToVEvent_); + } + + __aicore__ inline void MaterializeRawGateChunkArch35( + uint64_t b, uint64_t hv, uint64_t start, uint64_t rows) + { +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if (!computeGateInPrepare_) { + return; + } + if constexpr (!IsSameType::value) { + return; + } else { + activeGateChunkStart_ = start; + const float expA = LoadGateExpA(hv); + LocalTensor acc = exp2Buf_.Get(); + LocalTensor bias = exp2Buf_.Get()[K_]; + if (hasDtBias_) { + CopyVectorIn(bias, dtBias_, hv * K_, K_); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + } + Duplicate(acc, 0.0f, static_cast(K_)); + PipeBarrier(); + + const uint64_t tileRows = GatePipelineRows(); + const uint64_t tileCount = (rows + tileRows - 1) / tileRows; + uint64_t currentRows = rows < tileRows ? rows : tileRows; + PrefetchRawGateTile(0, b, hv, start, currentRows); + for (uint64_t tile = 0; tile < tileCount; ++tile) { + const uint64_t slot = tile & 1; + const uint64_t tileRow = tile * tileRows; + currentRows = rows - tileRow; + if (currentRows > tileRows) { + currentRows = tileRows; + } + WaitFlag(mte2ToVEvent_); + + const uint64_t nextTile = tile + 1; + if (nextTile < tileCount) { + const uint64_t nextSlot = nextTile & 1; + uint64_t nextRows = rows - nextTile * tileRows; + if (nextRows > tileRows) { + nextRows = tileRows; + } + PrefetchRawGateTile(nextSlot, b, hv, start + nextTile * tileRows, nextRows); + } + + LocalTensor gate = + LocalGateChunk().template ReinterpretCast()[tileRow * K_]; + if (hasDtBias_) { + AccumulateRawSafeGateChunk128Regbase( + (__ubuf__ float *)gate.GetPhyAddr(), (__ubuf__ float *)bias.GetPhyAddr(), + (__ubuf__ float *)acc.GetPhyAddr(), static_cast(currentRows), + expA, lowerBound_); + } else { + AccumulateRawSafeGateChunk128Regbase( + (__ubuf__ float *)gate.GetPhyAddr(), (__ubuf__ float *)bias.GetPhyAddr(), + (__ubuf__ float *)acc.GetPhyAddr(), static_cast(currentRows), + expA, lowerBound_); + } + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(gk_, KVOffset(b, hv, start + tileRow, 0, K_), gate, + currentRows * K_); + } + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } +#else + (void)b; + (void)hv; + (void)start; + (void)rows; +#endif + } + + __aicore__ inline LocalTensor GateDirectQ(uint64_t slot) + { + return vecBuf_.Get()[slot * 3 * GateStageElems()]; + } + + __aicore__ inline LocalTensor GateDirectW(uint64_t slot) + { + return GateDirectQ(slot)[GateStageElems()]; + } + + __aicore__ inline LocalTensor GateDirectV(uint64_t slot) + { + return GateDirectQ(slot)[2 * GateStageElems()]; + } + + __aicore__ inline LocalTensor GateBetaFloat(uint64_t slot) + { + constexpr uint64_t directBytes = + KDA_GATE_PIPELINE_DEPTH * 3 * KDA_GATE_TILE_ROWS * COMPILE_K * sizeof(T); + return vecBuf_.Get()[directBytes / sizeof(float) + slot * KDA_GATE_TILE_ROWS]; + } + + __aicore__ inline void StorePreparedQG(uint64_t b, uint64_t hv, uint64_t token, + LocalTensor directQ, uint64_t elems) + { + static_assert(KDA_SCALED_QG_FLOAT_OFFSET + KDA_GATE_TILE_ROWS * 128 <= + KDA_DIRECT_SCORE_FLOAT_OFFSET); + const uint64_t offset = KVOffset(b, hv, token, 0, K_); + if (storeQG_) { + CopyVectorOut(qg_, offset, directQ, elems); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + LocalTensor scaledQG = vecBuf_.Get()[KDA_SCALED_QG_FLOAT_OFFSET]; + Cast(scaledQG, directQ, RoundMode::CAST_NONE, static_cast(elems)); + PipeBarrier(); + Muls(scaledQG, scaledQG, scale_, static_cast(elems)); + PipeBarrier(); + ClampFp32ToOutputType(scaledQG, static_cast(elems)); + Cast(directQ, scaledQG, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(kg_, offset, directQ, elems); + } + + __aicore__ inline void PrefetchKGate(uint64_t slot, uint64_t b, uint64_t h, uint64_t hv, + uint64_t token, uint64_t elems) + { + const uint64_t rows = elems / K_; + LocalTensor kTyped = GateQTyped(slot); + LocalTensor gateTyped = GateGTyped(slot); + CopyRowsIn(kTyped, k_, QOffset(b, h, token, 0), rows, K_, inputSequenceMajor_ ? H_ * K_ : K_); + if (!computeGateInPrepare_) { + CopyVectorIn(gateTyped, gk_, KVOffset(b, hv, token, 0, K_), elems); + } + SetFlag(mte2ToVEvent_); + } + + __aicore__ inline void WaitGateInputReady() + { + WaitFlag(mte2ToVEvent_); + } + + __aicore__ inline void WaitGateOutputForMte2(uint64_t slot = 0) + { + WaitFlag(mte3ToMte2Events_[slot]); + } + + __aicore__ inline void WaitGateOutputForVector() + { + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void SignalGateOutputDone() + { + SetFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + } + + __aicore__ inline void SignalGateOutputDoneForMte2(uint64_t slot) + { + SetFlag(mte3ToMte2Events_[slot]); + } + + template + __aicore__ inline void LoadAsFloatRow(GlobalTensor &src, uint64_t srcOffset, LocalTensor &dst, + uint64_t count) + { + if constexpr (IsSameType::value) { + CopyVectorIn(dst, src, srcOffset, count); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Adds(dst, dst, 0.0f, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } else { + constexpr uint32_t typedOffset = EXP2_UB_ELEMENTS * sizeof(float) / sizeof(CopyT); + LocalTensor rowLocal = exp2Buf_.Get()[typedOffset]; + CopyVectorIn(rowLocal, src, srcOffset, count); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast(dst, rowLocal, RoundMode::CAST_NONE, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } + PipeBarrier(); + } + + template + __aicore__ inline void LoadAsFloatVector(GlobalTensor &src, uint64_t srcOffset, + LocalTensor &dst, LocalTensor &typedScratch, + uint64_t count) + { + if constexpr (IsSameType::value) { + CopyVectorIn(dst, src, srcOffset, count); + } else { + CopyVectorIn(typedScratch, src, srcOffset, count); + } + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + if constexpr (!IsSameType::value) { + Cast(dst, typedScratch, RoundMode::CAST_NONE, static_cast(count)); + PipeBarrier(); + } + } + + template + __aicore__ inline void StoreFloatRow(GlobalTensor &dst, uint64_t dstOffset, LocalTensor &src, + uint64_t count) + { + if constexpr (IsSameType::value) { + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(dst, dstOffset, src, count); + } else { + constexpr uint32_t typedOffset = EXP2_UB_ELEMENTS * sizeof(float) / sizeof(CopyT); + LocalTensor rowLocal = exp2Buf_.Get()[typedOffset]; + Cast(rowLocal, src, RoundMode::CAST_RINT, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(dst, dstOffset, rowLocal, count); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + + + + + __aicore__ inline LocalTensor Exp2NegG(uint64_t b, uint64_t hv, uint64_t t) + { + LocalTensor exp2Local = exp2Buf_.Get(); + LoadAsFloatRow(gk_, KVOffset(b, hv, t, 0, K_), exp2Local, K_); + Muls(exp2Local, exp2Local, -LN2, static_cast(K_)); + PipeBarrier(); + RunExp2(exp2Local, static_cast(K_)); + return exp2Local; + } + + __aicore__ inline void PrepareScoreFactorsBulk(uint64_t b, uint64_t h, uint64_t hv, uint64_t start, + uint64_t subBlockIdx, uint64_t subBlockNum, + uint64_t refToken, uint64_t scoreRowBegin, + uint64_t scoreRowCount, uint64_t validColEnd, + uint64_t finalRefToken, uint64_t scoreSlot) + { + const bool useDirectScoreL1 = +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + KDA_ARCH35_ENABLE_DIRECT_SCORE_L1 && KDA_ARCH35_ENABLE_DIRECT_SCORE_UB && + SAFE_GATE && BT_ == 64 && K_ == 128 && V_ == 128 && + subBlockNum == 1 && scoreRowCount == KDA_DIRECT_SCORE_ROWS && + finalRefToken == start + BT_ - 1; +#else + false; +#endif + LocalTensor refFp32 = exp2Buf_.Get(); + LoadGateScoreRef(refFp32, b, hv, refToken); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + constexpr bool exportFinalKg = + SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128; + LocalTensor finalRefFp32 = exp2Buf_.Get()[K_]; + if constexpr (exportFinalKg) { + LoadGateScoreRef(finalRefFp32, b, hv, finalRefToken); + } +#endif + + uint64_t qwBegin = scoreRowBegin + (scoreRowCount * subBlockIdx) / subBlockNum; + uint64_t qwEnd = scoreRowBegin + (scoreRowCount * (subBlockIdx + 1)) / subBlockNum; + uint64_t qwMaxRows = GatePipelineRows(); + bool qwOutputPending = false; + uint64_t qwSlot = 0; + if (qwBegin < qwEnd && qwMaxRows > 0) { + uint64_t firstRows = qwEnd - qwBegin; + if (firstRows > qwMaxRows) { + firstRows = qwMaxRows; + } + PrefetchQKGate(qwSlot, b, h, hv, start + qwBegin, firstRows * K_); + } + for (uint64_t tileRow = qwBegin; tileRow < qwEnd && qwMaxRows > 0; tileRow += qwMaxRows) { + uint64_t tileRows = qwEnd - tileRow; + if (tileRows > qwMaxRows) { + tileRows = qwMaxRows; + } + uint64_t elems = tileRows * K_; + LocalTensor qTyped = GateQTyped(qwSlot); + LocalTensor kTyped = GateKTyped(qwSlot); + LocalTensor qScore = qTyped.template ReinterpretCast(); + LocalTensor kScore = kTyped.template ReinterpretCast(); + LocalTensor gateTyped = GateScoreTyped(qwSlot, tileRow); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + LocalTensor kgScore = + GateKgTyped(qwSlot).template ReinterpretCast(); + LocalTensor vTyped = GateVTyped(qwSlot); + LocalTensor betaTyped = GateBetaTyped(qwSlot); + LocalTensor directQ = GateDirectQ(qwSlot); + LocalTensor directW = GateDirectW(qwSlot); + LocalTensor directV = GateDirectV(qwSlot); + LocalTensor betaFp32 = GateBetaFloat(qwSlot); +#endif +#if !defined(__CCE_AICORE__) || __CCE_AICORE__ != 310 + LocalTensor arena = vecBuf_.Get(); + LocalTensor qFp32 = arena; + LocalTensor kFp32 = arena[elems]; + LocalTensor gFp32 = arena[2 * elems]; + LocalTensor expFp32 = arena[3 * elems]; + LocalTensor outFp32 = arena[4 * elems]; +#endif + + WaitGateInputReady(); +#if !defined(__CCE_AICORE__) || __CCE_AICORE__ != 310 + Cast(qFp32, qTyped, RoundMode::CAST_NONE, static_cast(elems)); + Cast(kFp32, kTyped, RoundMode::CAST_NONE, static_cast(elems)); + if constexpr (IsSameType::value) { + gFp32 = gateTyped; + } else { + Cast(gFp32, gateTyped, RoundMode::CAST_NONE, static_cast(elems)); + } +#endif +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (IsSameType::value) { + Adds(betaFp32, betaTyped, 0.0f, static_cast(tileRows)); + } else { + Cast(betaFp32, betaTyped, RoundMode::CAST_NONE, static_cast(tileRows)); + } + PipeBarrier(); +#endif + if (qwOutputPending) { + WaitGateOutputForMte2(); + } + uint64_t nextTileRow = tileRow + qwMaxRows; + if (nextTileRow < qwEnd) { + uint64_t nextRows = qwEnd - nextTileRow; + if (nextRows > qwMaxRows) { + nextRows = qwMaxRows; + } + PrefetchQKGate(qwSlot ^ 1, b, h, hv, start + nextTileRow, nextRows * K_); + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + bool fuseQwKg = SAFE_GATE && BT_ == 64 && K_ == 128 && V_ == 128 && subBlockNum == 1; + if (fuseQwKg) { + PrepareKdaGateQwKgRegbase( + (__ubuf__ T *)reinterpret_cast(qTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(kTyped.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(qScore.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(kScore.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(kgScore.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(directQ.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(directW.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(vTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(directV.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(vTyped.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(betaFp32.GetPhyAddr()), + (__ubuf__ GK_T *)reinterpret_cast(gateTyped.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(refFp32.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(finalRefFp32.GetPhyAddr()), + static_cast(tileRows), static_cast(K_), + static_cast(tileRows)); + } else { + PrepareKdaGateQwRegbase( + (__ubuf__ T *)reinterpret_cast(qTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(kTyped.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(qScore.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(kScore.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(directQ.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(directW.GetPhyAddr()), + (__ubuf__ GK_T *)reinterpret_cast(gateTyped.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(refFp32.GetPhyAddr()), + static_cast(tileRows), static_cast(K_)); + } +#else + PipeBarrier(); + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expFp32[row * K_], gFp32[row * K_], refFp32, static_cast(K_)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(elems)); + PipeBarrier(); + ClampScoreExpInput(expFp32, static_cast(elems)); + Exp(expFp32, expFp32, static_cast(elems)); + PipeBarrier(); + + Mul(outFp32, qFp32, expFp32, static_cast(elems)); + PipeBarrier(); + ClampFp32ForCast(outFp32, static_cast(elems)); + Cast(qScore, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); + Mul(outFp32, kFp32, expFp32, static_cast(elems)); + PipeBarrier(); + ClampFp32ForCast(outFp32, static_cast(elems)); + Cast(kScore, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); +#endif + + if (qwOutputPending) { + WaitGateOutputForVector(); + } + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + if (useDirectScoreL1) { + Catlass::Arch::Resource resource; + LocalTensor scoreL1 = + resource.l1Buf.template GetBufferByByte( + scoreSlot * KDA_DIRECT_SCORE_L1_SLOT_BYTES); + const uint64_t localRow = tileRow - scoreRowBegin; + constexpr uint32_t c0Elements = 16; + constexpr uint32_t c0Blocks = 128 / c0Elements; + constexpr uint32_t rowBlockElements = 16 * c0Elements; + LocalTensor nzScratch = + resource.ubBuf.template GetBufferByByte( + KDA_DIRECT_SCORE_UB_BYTE_OFFSET + + (scoreSlot / KDA_SCORE_LANES) * + KDA_DIRECT_SCORE_SLOT_ELEMENTS * sizeof(float) + + 2 * KDA_DIRECT_SCORE_MATRIX_ELEMENTS * sizeof(float)); + LocalTensor qNz = nzScratch; + LocalTensor kNz = qNz[tileRows * K_]; + constexpr uint8_t srcRepeatStride = 128 * sizeof(SCORE_T) / 32; + for (uint32_t colBlock = 0; colBlock < c0Blocks; ++colBlock) { + const uint64_t srcOffset = colBlock * c0Elements; + const uint64_t dstOffset = colBlock * tileRows * c0Elements; + Copy(qNz[dstOffset], qScore[srcOffset], c0Elements, + static_cast(tileRows), {1, 1, 1, srcRepeatStride}); + Copy(kNz[dstOffset], kScore[srcOffset], c0Elements, + static_cast(tileRows), {1, 1, 1, srcRepeatStride}); + } + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + const uint64_t qL1Offset = (localRow / 16) * rowBlockElements; + const uint64_t kL1Offset = + ((KDA_DIRECT_SCORE_ROWS + localRow) / 16) * rowBlockElements; + DataCopyParams nzCopyParams{ + c0Blocks, static_cast(tileRows), 0, + static_cast(64 - tileRows)}; + DataCopy(scoreL1[qL1Offset], qNz, nzCopyParams); + DataCopy(scoreL1[kL1Offset], kNz, nzCopyParams); + } else { + CopyVectorOut(scoreWorkspace_, + ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG, tileRow), + qScore, elems); + CopyVectorOut(scoreWorkspace_, + ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W, tileRow), + kScore, elems); + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if (fuseQwKg) { + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG, tileRow), + kgScore, elems); + } + StorePreparedQG(b, hv, start + tileRow, directQ, elems); + CopyVectorOut(w_, KVOffset(b, hv, start + tileRow, 0, K_), directW, elems); + CopyVectorOut(vNew_, KVOffset(b, hv, start + tileRow, 0, V_), directV, + tileRows * V_); + if constexpr (exportFinalKg) { + CopyVectorOut(finalKg_, KVOffset(b, hv, start + tileRow, 0, K_), vTyped, elems); + } +#endif + SignalGateOutputDone(); + qwOutputPending = true; + qwSlot ^= 1; + } + if (qwOutputPending) { + WaitGateOutputForMte2(); + WaitGateOutputForVector(); + } + + bool fuseQwKg = false; +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + fuseQwKg = SAFE_GATE && BT_ == 64 && K_ == 128 && V_ == 128 && subBlockNum == 1; +#endif + uint64_t kgRows = fuseQwKg ? scoreRowBegin : validColEnd; + uint64_t kgBegin = (kgRows * subBlockIdx) / subBlockNum; + uint64_t kgEnd = (kgRows * (subBlockIdx + 1)) / subBlockNum; + uint64_t kgMaxRows = GatePipelineRows(); + bool kgOutputPending = false; + uint64_t kgSlot = 0; + if (kgBegin < kgEnd && kgMaxRows > 0) { + uint64_t firstRows = kgEnd - kgBegin; + if (firstRows > kgMaxRows) { + firstRows = kgMaxRows; + } + PrefetchKGate(kgSlot, b, h, hv, start + kgBegin, firstRows * K_); + } + for (uint64_t tileRow = kgBegin; tileRow < kgEnd && kgMaxRows > 0; tileRow += kgMaxRows) { + uint64_t tileRows = kgEnd - tileRow; + if (tileRows > kgMaxRows) { + tileRows = kgMaxRows; + } + uint64_t elems = tileRows * K_; + LocalTensor kTyped = GateQTyped(kgSlot); + LocalTensor kgScore = kTyped.template ReinterpretCast(); + LocalTensor gateTyped = GateScoreTyped(kgSlot, tileRow); +#if !defined(__CCE_AICORE__) || __CCE_AICORE__ != 310 + LocalTensor arena = vecBuf_.Get(); + LocalTensor kFp32 = arena; + LocalTensor gFp32 = arena[elems]; + LocalTensor expFp32 = arena[2 * elems]; + LocalTensor outFp32 = arena[3 * elems]; +#endif + + WaitGateInputReady(); +#if !defined(__CCE_AICORE__) || __CCE_AICORE__ != 310 + Cast(kFp32, kTyped, RoundMode::CAST_NONE, static_cast(elems)); + if constexpr (IsSameType::value) { + gFp32 = gateTyped; + } else { + Cast(gFp32, gateTyped, RoundMode::CAST_NONE, static_cast(elems)); + } +#endif + if (kgOutputPending) { + WaitGateOutputForMte2(); + } + uint64_t nextTileRow = tileRow + kgMaxRows; + if (nextTileRow < kgEnd) { + uint64_t nextRows = kgEnd - nextTileRow; + if (nextRows > kgMaxRows) { + nextRows = kgMaxRows; + } + PrefetchKGate(kgSlot ^ 1, b, h, hv, start + nextTileRow, nextRows * K_); + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + PrepareKdaGateKgRegbase( + (__ubuf__ SCORE_T *)reinterpret_cast(kgScore.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(kTyped.GetPhyAddr()), + (__ubuf__ GK_T *)reinterpret_cast(gateTyped.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(refFp32.GetPhyAddr()), + static_cast(tileRows), static_cast(K_), + static_cast(tileRows)); +#else + PipeBarrier(); + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expFp32[row * K_], refFp32, gFp32[row * K_], static_cast(K_)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(elems)); + PipeBarrier(); + ClampScoreExpInput(expFp32, static_cast(elems)); + Exp(expFp32, expFp32, static_cast(elems)); + PipeBarrier(); + Mul(outFp32, kFp32, expFp32, static_cast(elems)); + PipeBarrier(); + ClampFp32ForCast(outFp32, static_cast(elems)); + Cast(kgScore, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); +#endif + + if (kgOutputPending) { + WaitGateOutputForVector(); + } + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG, tileRow), + kgScore, elems); + SignalGateOutputDone(); + kgOutputPending = true; + kgSlot ^= 1; + } + if (kgOutputPending) { + WaitGateOutputForMte2(); + WaitGateOutputForVector(); + } + } + + __aicore__ inline void PrepareGateProductsBulk(uint64_t b, uint64_t h, uint64_t hv, uint64_t start, + uint64_t curT, uint64_t subBlockIdx, uint64_t subBlockNum, + bool useRef, uint64_t refToken, uint64_t validColEnd, + bool writeScoreScratch, uint64_t scoreSlot) + { + if constexpr (IsSameType::value) { + return; + } + if (subBlockNum == 0 || subBlockIdx >= subBlockNum || K_ == 0) { + return; + } + uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; + uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + if (rowBegin >= rowEnd) { + return; + } + + uint64_t maxRows = GatePipelineRows(); + if (maxRows == 0) { + return; + } + LocalTensor refFp32 = exp2Buf_.Get(); + if (useRef) { + LoadGateScoreRef(refFp32, b, hv, refToken); + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + LocalTensor finalRefFp32 = exp2Buf_.Get()[K_]; + if (materializeFinalKg_) { + LoadGateScoreRef(finalRefFp32, b, hv, start + curT - 1); + } + const bool fuseScoreWriteback = + writeScoreScratch && useRef && K_ * 2 <= EXP2_UB_ELEMENTS; +#endif + + bool outputPending = false; + uint64_t gateSlot = 0; + uint64_t firstRows = rowEnd - rowBegin; + if (firstRows > maxRows) { + firstRows = maxRows; + } + PrefetchQKGate(gateSlot, b, h, hv, start + rowBegin, firstRows * K_); + for (uint64_t tileRow = rowBegin; tileRow < rowEnd; tileRow += maxRows) { + uint64_t tileRows = rowEnd - tileRow; + if (tileRows > maxRows) { + tileRows = maxRows; + } + uint64_t elems = tileRows * K_; + LocalTensor qTyped = GateQTyped(gateSlot); + LocalTensor kTyped = GateKTyped(gateSlot); + LocalTensor kgTyped = GateKgTyped(gateSlot); + LocalTensor qScore = qTyped.template ReinterpretCast(); + LocalTensor wScore = kTyped.template ReinterpretCast(); + LocalTensor kgScore = kgTyped.template ReinterpretCast(); + LocalTensor gateTyped = GateScoreTyped(gateSlot, tileRow); +#if !defined(__CCE_AICORE__) || __CCE_AICORE__ != 310 + LocalTensor arena = vecBuf_.Get(); + LocalTensor qFp32 = arena; + LocalTensor kFp32 = arena[elems]; + LocalTensor gFp32 = arena[2 * elems]; + LocalTensor expFp32 = arena[3 * elems]; + LocalTensor outFp32 = arena[4 * elems]; +#else + LocalTensor vTyped = GateVTyped(gateSlot); + LocalTensor betaTyped = GateBetaTyped(gateSlot); + LocalTensor directQ = GateDirectQ(gateSlot); + LocalTensor directW = GateDirectW(gateSlot); + LocalTensor directV = GateDirectV(gateSlot); + LocalTensor betaFp32 = GateBetaFloat(gateSlot); +#endif + + uint64_t token = start + tileRow; + WaitGateInputReady(); +#if !defined(__CCE_AICORE__) || __CCE_AICORE__ != 310 + Cast(qFp32, qTyped, RoundMode::CAST_NONE, static_cast(elems)); + Cast(kFp32, kTyped, RoundMode::CAST_NONE, static_cast(elems)); + if constexpr (IsSameType::value) { + gFp32 = gateTyped; + } else { + Cast(gFp32, gateTyped, RoundMode::CAST_NONE, static_cast(elems)); + } +#endif +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (IsSameType::value) { + Adds(betaFp32, betaTyped, 0.0f, static_cast(tileRows)); + } else { + Cast(betaFp32, betaTyped, RoundMode::CAST_NONE, static_cast(tileRows)); + } + PipeBarrier(); +#endif + uint64_t nextTileRow = tileRow + maxRows; +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + uint64_t nextGateSlot = (gateSlot + 1) % KDA_GATE_PIPELINE_DEPTH; + if (nextTileRow < rowEnd) { + uint64_t tileIndex = (tileRow - rowBegin) / maxRows; + if (tileIndex + 1 >= KDA_GATE_PIPELINE_DEPTH) { + WaitGateOutputForMte2(nextGateSlot); + } + uint64_t nextRows = rowEnd - nextTileRow; + if (nextRows > maxRows) { + nextRows = maxRows; + } + PrefetchQKGate(nextGateSlot, b, h, hv, start + nextTileRow, nextRows * K_); + } +#else + if (outputPending) { + WaitGateOutputForMte2(); + } + if (nextTileRow < rowEnd) { + uint64_t nextRows = rowEnd - nextTileRow; + if (nextRows > maxRows) { + nextRows = maxRows; + } + PrefetchQKGate(gateSlot ^ 1, b, h, hv, start + nextTileRow, nextRows * K_); + } +#endif +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + uint16_t validRows = static_cast(tileRows); + if (useRef && tileRow >= validColEnd) { + validRows = 0; + } else if (useRef && tileRow + tileRows > validColEnd) { + validRows = static_cast(validColEnd - tileRow); + } + if (writeScoreScratch) { + if (fuseScoreWriteback) { + if (materializeFinalKg_) { + PrepareKdaGateQwKgRegbase( + (__ubuf__ T *)reinterpret_cast(qTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(kTyped.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(qScore.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(wScore.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(kgScore.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(directQ.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(directW.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(vTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(directV.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(vTyped.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(betaFp32.GetPhyAddr()), + (__ubuf__ GK_T *)reinterpret_cast(gateTyped.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(refFp32.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(finalRefFp32.GetPhyAddr()), + static_cast(tileRows), static_cast(K_), + validRows); + } else { + PrepareKdaGateQwKgRegbase( + (__ubuf__ T *)reinterpret_cast(qTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(kTyped.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(qScore.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(wScore.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(kgScore.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(directQ.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(directW.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(vTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(directV.GetPhyAddr()), + nullptr, + (__ubuf__ float *)reinterpret_cast(betaFp32.GetPhyAddr()), + (__ubuf__ GK_T *)reinterpret_cast(gateTyped.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(refFp32.GetPhyAddr()), + nullptr, + static_cast(tileRows), static_cast(K_), validRows); + } + } else { + PrepareKdaGateQwKgRegbase( + (__ubuf__ T *)reinterpret_cast(qTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(kTyped.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(qScore.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(wScore.GetPhyAddr()), + (__ubuf__ SCORE_T *)reinterpret_cast(kgScore.GetPhyAddr()), + nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, + (__ubuf__ GK_T *)reinterpret_cast(gateTyped.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(refFp32.GetPhyAddr()), + nullptr, + static_cast(tileRows), static_cast(K_), validRows); + } + } else if (useRef) { + PrepareKdaGateQwKgRegbase( + (__ubuf__ T *)reinterpret_cast(qTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(kTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(qTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(kTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(kgTyped.GetPhyAddr()), + nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, + (__ubuf__ GK_T *)reinterpret_cast(gateTyped.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(refFp32.GetPhyAddr()), + nullptr, + static_cast(tileRows), static_cast(K_), validRows); + } else { + PrepareKdaGateQwKgRegbase( + (__ubuf__ T *)reinterpret_cast(qTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(kTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(qTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(kTyped.GetPhyAddr()), + (__ubuf__ T *)reinterpret_cast(kgTyped.GetPhyAddr()), + nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, + (__ubuf__ GK_T *)reinterpret_cast(gateTyped.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(refFp32.GetPhyAddr()), + nullptr, + static_cast(tileRows), static_cast(K_), validRows); + } +#else + PipeBarrier(); + + if (useRef) { + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expFp32[row * K_], gFp32[row * K_], refFp32, static_cast(K_)); + } + } else { + Adds(expFp32, gFp32, 0.0f, static_cast(elems)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(elems)); + PipeBarrier(); + if (writeScoreScratch) { + ClampScoreExpInput(expFp32, static_cast(elems)); + } else { + ClampExpInput(expFp32, static_cast(elems)); + } + Exp(expFp32, expFp32, static_cast(elems)); + PipeBarrier(); + + Mul(outFp32, qFp32, expFp32, static_cast(elems)); + PipeBarrier(); + if (writeScoreScratch) { + ClampFp32ForCast(outFp32, static_cast(elems)); + Cast(qScore, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + } else { + ClampFp32ToOutputType(outFp32, static_cast(elems)); + Cast(qTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + } + PipeBarrier(); + + Mul(outFp32, kFp32, expFp32, static_cast(elems)); + PipeBarrier(); + if (writeScoreScratch) { + ClampFp32ForCast(outFp32, static_cast(elems)); + Cast(wScore, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + } else { + ClampFp32ToOutputType(outFp32, static_cast(elems)); + Cast(kTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + } + PipeBarrier(); + + if (useRef) { + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expFp32[row * K_], refFp32, gFp32[row * K_], static_cast(K_)); + } + } else { + Muls(expFp32, gFp32, -1.0f, static_cast(elems)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(elems)); + PipeBarrier(); + if (writeScoreScratch) { + ClampScoreExpInput(expFp32, static_cast(elems)); + } else { + ClampExpInput(expFp32, static_cast(elems)); + } + Exp(expFp32, expFp32, static_cast(elems)); + PipeBarrier(); + Mul(outFp32, kFp32, expFp32, static_cast(elems)); + PipeBarrier(); + if (useRef && tileRow + tileRows > validColEnd) { + for (uint64_t row = 0; row < tileRows; ++row) { + if (tileRow + row >= validColEnd) { + Duplicate(outFp32[row * K_], 0.0f, static_cast(K_)); + } + } + PipeBarrier(); + } + if (writeScoreScratch) { + ClampFp32ForCast(outFp32, static_cast(elems)); + } else { + ClampFp32ToOutputType(outFp32, static_cast(elems)); + } + if (outputPending) { + WaitGateOutputForVector(); + } + if (writeScoreScratch) { + Cast(kgScore, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + } else { + Cast(kgTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + } + PipeBarrier(); +#endif + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + if (writeScoreScratch) { + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG, tileRow), + qScore, elems); + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W, tileRow), + wScore, elems); + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG, tileRow), + kgScore, elems); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if (fuseScoreWriteback) { + StorePreparedQG(b, hv, token, directQ, elems); + CopyVectorOut(w_, KVOffset(b, hv, token, 0, K_), directW, elems); + CopyVectorOut(vNew_, KVOffset(b, hv, token, 0, V_), directV, tileRows * V_); + if (materializeFinalKg_) { + CopyVectorOut(finalKg_, KVOffset(b, hv, token, 0, K_), vTyped, elems); + } + } +#endif + } else { + CopyVectorOut(qg_, KVOffset(b, hv, token, 0, K_), qTyped, elems); + CopyVectorOut(w_, KVOffset(b, hv, token, 0, K_), kTyped, elems); + CopyVectorOut(kg_, KVOffset(b, hv, token, 0, K_), kgTyped, elems); + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + SignalGateOutputDoneForMte2(gateSlot); +#else + SignalGateOutputDone(); +#endif + outputPending = true; +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + gateSlot = (gateSlot + 1) % KDA_GATE_PIPELINE_DEPTH; +#else + gateSlot ^= 1; +#endif + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + uint64_t tileCount = (rowEnd - rowBegin + maxRows - 1) / maxRows; + uint64_t firstPending = + tileCount > KDA_GATE_PIPELINE_DEPTH ? tileCount - KDA_GATE_PIPELINE_DEPTH : 0; + for (uint64_t tile = firstPending; tile < tileCount; ++tile) { + WaitGateOutputForMte2(tile % KDA_GATE_PIPELINE_DEPTH); + } +#else + if (outputPending) { + WaitGateOutputForMte2(); + WaitGateOutputForVector(); + } +#endif + return; + } + + __aicore__ inline void ZeroScoreScratchRange(uint64_t scoreSlot, uint64_t planeBegin, + uint64_t planeEnd, uint64_t firstRow, + uint64_t rowEnd) + { + if (firstRow >= rowEnd) { + return; + } + const uint64_t maxRows = GatePipelineRows(); + LocalTensor zeroLocal = GateQTyped(0).template ReinterpretCast(); + for (uint64_t row = firstRow; row < rowEnd; row += maxRows) { + uint64_t rows = rowEnd - row; + if (rows > maxRows) { + rows = maxRows; + } + const uint64_t elems = rows * K_; + Duplicate(zeroLocal, static_cast(0), static_cast(elems)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + for (uint64_t plane = planeBegin; plane < planeEnd; ++plane) { + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, plane, row), + zeroLocal, elems); + } + SignalGateOutputDone(); + WaitGateOutputForMte2(); + WaitGateOutputForVector(); + } + } + + __aicore__ inline void ZeroScoreScratchPadding(uint64_t scoreSlot, + uint64_t scoreRowBegin, + uint64_t scoreRowCount, + uint64_t validColEnd, + uint64_t subBlockIdx, + uint64_t subBlockNum) + { + if (subBlockNum == 0 || subBlockIdx + 1 != subBlockNum) { + return; + } + const uint64_t validRowEnd = scoreRowBegin + scoreRowCount; + const uint64_t paddedRowEnd = (validRowEnd + 15) / 16 * 16; + const uint64_t paddedColEnd = BT_; + ZeroScoreScratchRange(scoreSlot, KDA_SCORE_SCRATCH_QG, + KDA_SCORE_SCRATCH_KG, validRowEnd, paddedRowEnd); + ZeroScoreScratchRange(scoreSlot, KDA_SCORE_SCRATCH_KG, + KDA_SCORE_SCRATCH_PLANES, validColEnd, paddedColEnd); + } + + __aicore__ inline void PrepareGateProducts(uint64_t b, uint64_t h, uint64_t hv, uint64_t start, uint64_t curT, + uint64_t subBlockIdx, uint64_t subBlockNum, bool useRef = false, + uint64_t refToken = 0, uint64_t validColEnd = 0, + bool writeScoreScratch = false, uint64_t scoreSlot = 0, + uint64_t scoreRowBegin = 0, uint64_t scoreRowCount = 0) + { + if (subBlockNum == 0 || subBlockIdx >= subBlockNum) { + return; + } + if (validColEnd == 0 || validColEnd > curT) { + validColEnd = curT; + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if (writeScoreScratch && curT == BT_ && scoreRowBegin == 0 && + scoreRowCount == curT && validColEnd == curT) { + PrepareGateProductsBulk(b, h, hv, start, curT, subBlockIdx, subBlockNum, useRef, refToken, + validColEnd, writeScoreScratch, scoreSlot); + ZeroScoreScratchPadding(scoreSlot, scoreRowBegin, scoreRowCount, validColEnd, + subBlockIdx, subBlockNum); + return; + } +#endif + if (writeScoreScratch) { + PrepareScoreFactorsBulk(b, h, hv, start, subBlockIdx, subBlockNum, refToken, scoreRowBegin, + scoreRowCount, validColEnd, start + curT - 1, scoreSlot); + ZeroScoreScratchPadding(scoreSlot, scoreRowBegin, scoreRowCount, validColEnd, + subBlockIdx, subBlockNum); + return; + } + PrepareGateProductsBulk(b, h, hv, start, curT, subBlockIdx, subBlockNum, useRef, refToken, + validColEnd, writeScoreScratch, scoreSlot); + } + + __aicore__ inline void ComputeRawAqkAkkCube(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT) + { + ComputeRawAqkAkkCubeBlock(b, hv, chunkIdx, start, curT, 0, curT); + } + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + template + __aicore__ inline void ComputeRawAqkAkkCubeStableBlockDirectUbArch35( + uint64_t b, uint64_t hv, uint64_t start, uint64_t rowBegin, + uint64_t scoreSlot, uint8_t subBlockIdx, uint32_t directSlot) + { + using ElementA = SCORE_T; + using ElementB = SCORE_T; + using ElementC = float; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::ColumnMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla< + KdaArchTag, ElementA, LayoutTagA, ElementB, LayoutTagB, ElementC, LayoutTagC>; + using DirectTileCopy = Common::Tile::PackedTileCopyTlaToUB< + KdaArchTag, ElementA, LayoutTagA, ElementB, LayoutTagB, ElementC, LayoutTagC>; + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + using LayoutTagL0A = typename TileCopy::LayoutTagL0A; + using LayoutTagL0B = typename TileCopy::LayoutTagL0B; + using CopyL1ToL0A = typename TileCopy::CopyL1ToL0A; + using CopyL1ToL0B = typename TileCopy::CopyL1ToL0B; + using TileMmad = Catlass::Gemm::Tile::TileMmadTla; + + constexpr uint32_t scoreRows = KDA_SAFE_SCORE_REF_BC; + constexpr uint32_t packedRows = scoreRows * 2; + constexpr uint32_t k = 128; + static_assert(scoreRows == 32 && packedRows == 64); + static_assert(N == 32 || N == 64); + + Catlass::Arch::Resource resource; + auto layoutA = tla::MakeLayout(BT_, K_); + auto layoutB = tla::MakeLayout(K_, BT_); + auto tensorQPos = tla::MakeTensor( + scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG)], + layoutA, Catlass::Arch::PositionGM{}); + auto tensorKPos = tla::MakeTensor( + scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W)], + layoutA, Catlass::Arch::PositionGM{}); + auto tensorKNeg = tla::MakeTensor( + scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG)], + layoutB, Catlass::Arch::PositionGM{}); + + auto blockQPos = GetTile( + tensorQPos, tla::MakeCoord(rowBegin, 0), tla::MakeShape(scoreRows, k)); + auto blockKPos = GetTile( + tensorKPos, tla::MakeCoord(rowBegin, 0), tla::MakeShape(scoreRows, k)); + auto blockKNeg = GetTile( + tensorKNeg, tla::MakeCoord(0, 0), tla::MakeShape(k, N)); + + using CopyGmToL1A = typename TileCopy::template CopyGmToL1A; + using CopyGmToL1B = typename TileCopy::template CopyGmToL1B; + + static_assert(sizeof(ElementA) == sizeof(uint16_t)); + LocalTensor l1A = resource.l1Buf.template GetBufferByByte( + scoreSlot * KDA_DIRECT_SCORE_L1_SLOT_BYTES); + LocalTensor l1B = resource.l1Buf.template GetBufferByByte( + KDA_DIRECT_SCORE_L1_B_OFFSET); + LocalTensor l0A = resource.l0ABuf.template GetBufferByByte(0); + LocalTensor l0B = resource.l0BBuf.template GetBufferByByte(0); + LocalTensor l0C = resource.l0CBuf.template GetBufferByByte(0); + auto layoutL1A = tla::MakeLayout(packedRows, k); + auto layoutL1B = tla::MakeLayout(k, N); + auto layoutL0A = tla::MakeLayout(packedRows, k); + auto layoutL0B = tla::MakeLayout(k, N); + auto layoutL0C = tla::MakeLayoutL0C(packedRows, N); + auto tensorL1A = tla::MakeTensor(l1A, layoutL1A, Catlass::Arch::PositionL1{}); + auto tensorL1B = tla::MakeTensor(l1B, layoutL1B, Catlass::Arch::PositionL1{}); + auto tensorL0A = tla::MakeTensor(l0A, layoutL0A, Catlass::Arch::PositionL0A{}); + auto tensorL0B = tla::MakeTensor(l0B, layoutL0B, Catlass::Arch::PositionL0B{}); + auto tensorL0C = tla::MakeTensor(l0C, layoutL0C, Catlass::Arch::PositionL0C{}); + auto tileL1A = GetTile( + tensorL1A, tla::MakeCoord(0, 0), tla::MakeShape(packedRows, k)); + auto tileL1AQ = GetTile( + tensorL1A, tla::MakeCoord(0, 0), tla::MakeShape(scoreRows, k)); + auto tileL1AK = GetTile( + tensorL1A, tla::MakeCoord(scoreRows, 0), tla::MakeShape(scoreRows, k)); + auto tileL1B = GetTile( + tensorL1B, tla::MakeCoord(0, 0), tla::MakeShape(k, N)); + auto tileL0A = GetTile( + tensorL0A, tla::MakeCoord(0, 0), tla::MakeShape(packedRows, k)); + auto tileL0B = GetTile( + tensorL0B, tla::MakeCoord(0, 0), tla::MakeShape(k, N)); + auto tileL0C = GetTile( + tensorL0C, tla::MakeCoord(0, 0), tla::MakeShape(packedRows, N)); + auto tileL0CTop = GetTile( + tensorL0C, tla::MakeCoord(0, 0), tla::MakeShape(scoreRows, N)); + auto tileL0CBottom = GetTile( + tensorL0C, tla::MakeCoord(scoreRows, 0), tla::MakeShape(scoreRows, N)); + LocalTensor directBase = resource.ubBuf.template GetBufferByByte( + KDA_DIRECT_SCORE_UB_BYTE_OFFSET + + directSlot * KDA_DIRECT_SCORE_SLOT_ELEMENTS * sizeof(ElementC)); + LocalTensor directAqk = directBase; + LocalTensor directAkk = directBase[KDA_DIRECT_SCORE_MATRIX_ELEMENTS]; + auto layoutDirect = tla::MakeLayout(scoreRows, BT_); + auto tensorDirectAqk = tla::MakeTensor( + directAqk, layoutDirect, Catlass::Arch::PositionUB{}); + auto tensorDirectAkk = tla::MakeTensor( + directAkk, layoutDirect, Catlass::Arch::PositionUB{}); + auto blockDirectAqk = GetTile( + tensorDirectAqk, tla::MakeCoord(0, 0), tla::MakeShape(scoreRows, N)); + auto blockDirectAkk = GetTile( + tensorDirectAkk, tla::MakeCoord(0, 0), tla::MakeShape(scoreRows, N)); + using CopyL0CToDirectUb = + typename DirectTileCopy::template CopyL0CToDst; + + CopyGmToL1A copyGmToL1A; + CopyGmToL1B copyGmToL1B; + CopyL1ToL0A copyL1ToL0A; + CopyL1ToL0B copyL1ToL0B; + CopyL0CToDirectUb copyL0CToDirectUb; + TileMmad tileMmad; + + copyGmToL1A(tileL1AQ, blockQPos); + copyGmToL1A(tileL1AK, blockKPos); + copyGmToL1B(tensorL1B, blockKNeg); + SetFlag(KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + copyL1ToL0A(tileL0A, tileL1A); + copyL1ToL0B(tileL0B, tileL1B); + SetFlag(KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + tileMmad(tileL0C, tileL0A, tileL0B, packedRows, N, k, true, 0); + SetFlag(KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + const uint64_t flagOffset = + directSlot + subBlockIdx * KDA_DIRECT_SCORE_SUBBLOCK_FLAG_STRIDE; + CrossCoreWaitFlag<0x4, PIPE_FIX>(KDA_DIRECT_SCORE_FREE_FLAG + flagOffset); + copyL0CToDirectUb( + blockDirectAqk, tileL0CTop, KDA_DIRECT_SCORE_ROWS, subBlockIdx, 1, 0); + copyL0CToDirectUb( + blockDirectAkk, tileL0CBottom, KDA_DIRECT_SCORE_ROWS, subBlockIdx, 1, 0); + CrossCoreSetFlag<0x4, PIPE_FIX>(KDA_DIRECT_SCORE_READY_FLAG + flagOffset); + SetFlag(KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + } + + template + __aicore__ inline void PrefetchRawAqkAkkHeadPairArch35( + uint64_t rowBegin, uint64_t scoreSlotBase, uint32_t l1BaseOffset, + TEventID readyEvent) + { + using ElementA = SCORE_T; + using ElementB = SCORE_T; + using ElementC = float; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::ColumnMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla< + KdaArchTag, ElementA, LayoutTagA, ElementB, LayoutTagB, + ElementC, LayoutTagC>; + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + + constexpr uint32_t scoreRows = KDA_SAFE_SCORE_REF_BC; + constexpr uint32_t packedRows = scoreRows * 2; + constexpr uint32_t n = N; + constexpr uint32_t k = 128; + constexpr uint32_t l1ABytes = packedRows * k * sizeof(ElementA); + constexpr uint32_t l1BBytes = k * n * sizeof(ElementB); + constexpr uint32_t l1LaneBytes = l1ABytes + l1BBytes; + static_assert(scoreRows == 32 && packedRows == 64); + static_assert(N == 32 || N == 64); + + Catlass::Arch::Resource resource; + auto layoutA = tla::MakeLayout(BT_, K_); + auto layoutB = tla::MakeLayout(K_, BT_); + auto layoutL1A = tla::MakeLayout(packedRows, k); + auto layoutL1B = tla::MakeLayout(k, n); + + for (uint64_t lane = 0; lane < KDA_SCORE_LANES; ++lane) { + const uint64_t scoreSlot = scoreSlotBase + lane; + LocalTensor l1A = resource.l1Buf.template GetBufferByByte( + l1BaseOffset + static_cast(lane) * l1LaneBytes); + LocalTensor l1B = resource.l1Buf.template GetBufferByByte( + l1BaseOffset + static_cast(lane) * l1LaneBytes + l1ABytes); + auto tensorL1A = tla::MakeTensor( + l1A, layoutL1A, Catlass::Arch::PositionL1{}); + auto tensorL1B = tla::MakeTensor( + l1B, layoutL1B, Catlass::Arch::PositionL1{}); + auto tileL1AQ = GetTile( + tensorL1A, tla::MakeCoord(0, 0), tla::MakeShape(scoreRows, k)); + auto tileL1AK = GetTile( + tensorL1A, tla::MakeCoord(scoreRows, 0), tla::MakeShape(scoreRows, k)); + + auto tensorQPos = tla::MakeTensor( + scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG)], + layoutA, Catlass::Arch::PositionGM{}); + auto tensorKPos = tla::MakeTensor( + scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W)], + layoutA, Catlass::Arch::PositionGM{}); + auto tensorKNeg = tla::MakeTensor( + scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG)], + layoutB, Catlass::Arch::PositionGM{}); + auto blockQPos = GetTile( + tensorQPos, tla::MakeCoord(rowBegin, 0), tla::MakeShape(scoreRows, k)); + auto blockKPos = GetTile( + tensorKPos, tla::MakeCoord(rowBegin, 0), tla::MakeShape(scoreRows, k)); + auto blockKNeg = GetTile( + tensorKNeg, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + using CopyGmToL1A = typename TileCopy::template CopyGmToL1A; + using CopyGmToL1B = typename TileCopy::template CopyGmToL1B; + CopyGmToL1A copyGmToL1A; + CopyGmToL1B copyGmToL1B; + copyGmToL1A(tileL1AQ, blockQPos); + copyGmToL1A(tileL1AK, blockKPos); + copyGmToL1B(tensorL1B, blockKNeg); + } + SetFlag(readyEvent); + } + + template + __aicore__ inline void ComputeRawAqkAkkCubeStableHeadPairDirectUbArch35( + uint64_t rowBegin, uint64_t scoreSlotBase, uint32_t directSlot) + { + using ElementA = SCORE_T; + using ElementB = SCORE_T; + using ElementC = float; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::ColumnMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla< + KdaArchTag, ElementA, LayoutTagA, ElementB, LayoutTagB, + ElementC, LayoutTagC>; + using DirectTileCopy = Common::Tile::PackedTileCopyTlaToUB< + KdaArchTag, ElementA, LayoutTagA, ElementB, LayoutTagB, + ElementC, LayoutTagC>; + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + using LayoutTagL0A = typename TileCopy::LayoutTagL0A; + using LayoutTagL0B = typename TileCopy::LayoutTagL0B; + using CopyL1ToL0A = typename TileCopy::CopyL1ToL0A; + using CopyL1ToL0B = typename TileCopy::CopyL1ToL0B; + using TileMmad = Catlass::Gemm::Tile::TileMmadTla< + KdaArchTag, ElementA, LayoutTagL1A>; + + constexpr uint32_t scoreRows = KDA_SAFE_SCORE_REF_BC; + constexpr uint32_t packedRows = scoreRows * 2; + constexpr uint32_t n = N; + constexpr uint32_t k = 128; + constexpr uint32_t l1ABytes = packedRows * k * sizeof(ElementA); + constexpr uint32_t l1BBytes = k * n * sizeof(ElementB); + constexpr uint32_t l1LaneBytes = l1ABytes + l1BBytes; + static_assert(scoreRows == 32 && packedRows == 64); + static_assert(N == 32 || N == 64); + + PrefetchRawAqkAkkHeadPairArch35( + rowBegin, scoreSlotBase, 0, KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + + Catlass::Arch::Resource resource; + auto layoutL1A = tla::MakeLayout(packedRows, k); + auto layoutL1B = tla::MakeLayout(k, n); + + constexpr uint32_t l0ABytes = packedRows * k * sizeof(ElementA); + constexpr uint32_t l0BBytes = k * n * sizeof(ElementB); + constexpr uint32_t l0CBytes = packedRows * n * sizeof(ElementC); + LocalTensor l0A0 = + resource.l0ABuf.template GetBufferByByte(0); + LocalTensor l0A1 = + resource.l0ABuf.template GetBufferByByte(l0ABytes); + LocalTensor l0B0 = + resource.l0BBuf.template GetBufferByByte(0); + LocalTensor l0B1 = + resource.l0BBuf.template GetBufferByByte(l0BBytes); + LocalTensor l0C0 = + resource.l0CBuf.template GetBufferByByte(0); + LocalTensor l0C1 = + resource.l0CBuf.template GetBufferByByte(l0CBytes); + auto layoutL0A = tla::MakeLayout(packedRows, k); + auto layoutL0B = tla::MakeLayout(k, n); + auto layoutL0C = tla::MakeLayoutL0C(packedRows, n); + auto tensorL0A0 = tla::MakeTensor(l0A0, layoutL0A, Catlass::Arch::PositionL0A{}); + auto tensorL0A1 = tla::MakeTensor(l0A1, layoutL0A, Catlass::Arch::PositionL0A{}); + auto tensorL0B0 = tla::MakeTensor(l0B0, layoutL0B, Catlass::Arch::PositionL0B{}); + auto tensorL0B1 = tla::MakeTensor(l0B1, layoutL0B, Catlass::Arch::PositionL0B{}); + auto tensorL0C0 = tla::MakeTensor(l0C0, layoutL0C, Catlass::Arch::PositionL0C{}); + auto tensorL0C1 = tla::MakeTensor(l0C1, layoutL0C, Catlass::Arch::PositionL0C{}); + auto tileL0A0 = GetTile( + tensorL0A0, tla::MakeCoord(0, 0), tla::MakeShape(packedRows, k)); + auto tileL0A1 = GetTile( + tensorL0A1, tla::MakeCoord(0, 0), tla::MakeShape(packedRows, k)); + auto tileL0B0 = GetTile( + tensorL0B0, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto tileL0B1 = GetTile( + tensorL0B1, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto tileL0C0 = GetTile( + tensorL0C0, tla::MakeCoord(0, 0), tla::MakeShape(packedRows, n)); + auto tileL0C1 = GetTile( + tensorL0C1, tla::MakeCoord(0, 0), tla::MakeShape(packedRows, n)); + auto tileL0C0Top = GetTile( + tensorL0C0, tla::MakeCoord(0, 0), tla::MakeShape(scoreRows, n)); + auto tileL0C0Bottom = GetTile( + tensorL0C0, tla::MakeCoord(scoreRows, 0), tla::MakeShape(scoreRows, n)); + auto tileL0C1Top = GetTile( + tensorL0C1, tla::MakeCoord(0, 0), tla::MakeShape(scoreRows, n)); + auto tileL0C1Bottom = GetTile( + tensorL0C1, tla::MakeCoord(scoreRows, 0), tla::MakeShape(scoreRows, n)); + CopyL1ToL0A copyL1ToL0A; + CopyL1ToL0B copyL1ToL0B; + TileMmad tileMmad; + + LocalTensor l1A0 = + resource.l1Buf.template GetBufferByByte(0); + LocalTensor l1A1 = + resource.l1Buf.template GetBufferByByte(l1LaneBytes); + LocalTensor l1B0 = + resource.l1Buf.template GetBufferByByte(l1ABytes); + LocalTensor l1B1 = + resource.l1Buf.template GetBufferByByte(l1LaneBytes + l1ABytes); + auto tensorL1A0 = tla::MakeTensor(l1A0, layoutL1A, Catlass::Arch::PositionL1{}); + auto tensorL1A1 = tla::MakeTensor(l1A1, layoutL1A, Catlass::Arch::PositionL1{}); + auto tensorL1B0 = tla::MakeTensor(l1B0, layoutL1B, Catlass::Arch::PositionL1{}); + auto tensorL1B1 = tla::MakeTensor(l1B1, layoutL1B, Catlass::Arch::PositionL1{}); + auto tileL1A0 = GetTile( + tensorL1A0, tla::MakeCoord(0, 0), tla::MakeShape(packedRows, k)); + auto tileL1A1 = GetTile( + tensorL1A1, tla::MakeCoord(0, 0), tla::MakeShape(packedRows, k)); + auto tileL1B0 = GetTile( + tensorL1B0, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto tileL1B1 = GetTile( + tensorL1B1, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + + LocalTensor directBase = resource.ubBuf.template GetBufferByByte( + KDA_DIRECT_SCORE_UB_BYTE_OFFSET + + directSlot * KDA_DIRECT_SCORE_SLOT_ELEMENTS * sizeof(ElementC)); + LocalTensor directAqk = directBase; + LocalTensor directAkk = directBase[KDA_DIRECT_SCORE_MATRIX_ELEMENTS]; + auto layoutDirect = tla::MakeLayout(scoreRows, BT_); + auto tensorDirectAqk = tla::MakeTensor( + directAqk, layoutDirect, Catlass::Arch::PositionUB{}); + auto tensorDirectAkk = tla::MakeTensor( + directAkk, layoutDirect, Catlass::Arch::PositionUB{}); + auto blockDirectAqk = GetTile( + tensorDirectAqk, tla::MakeCoord(0, 0), tla::MakeShape(scoreRows, n)); + auto blockDirectAkk = GetTile( + tensorDirectAkk, tla::MakeCoord(0, 0), tla::MakeShape(scoreRows, n)); + using CopyL0CToDirectUb = + typename DirectTileCopy::template CopyL0CToDst; + CopyL0CToDirectUb copyL0CToDirectUb; + + copyL1ToL0A(tileL0A0, tileL1A0); + copyL1ToL0B(tileL0B0, tileL1B0); + SetFlag(KDA_ARCH35_SCORE_EVENT); + copyL1ToL0A(tileL0A1, tileL1A1); + copyL1ToL0B(tileL0B1, tileL1B1); + SetFlag(KDA_ARCH35_SCORE_W_EVENT); + + WaitFlag(KDA_ARCH35_SCORE_EVENT); + tileMmad(tileL0C0, tileL0A0, tileL0B0, packedRows, n, k, true, 0); + SetFlag(KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + const uint64_t flagOffset0 = directSlot; + CrossCoreWaitFlag<0x4, PIPE_FIX>(KDA_DIRECT_SCORE_FREE_FLAG + flagOffset0); + copyL0CToDirectUb( + blockDirectAqk, tileL0C0Top, KDA_DIRECT_SCORE_ROWS, 0, 1, 0); + copyL0CToDirectUb( + blockDirectAkk, tileL0C0Bottom, KDA_DIRECT_SCORE_ROWS, 0, 1, 0); + CrossCoreSetFlag<0x4, PIPE_FIX>(KDA_DIRECT_SCORE_READY_FLAG + flagOffset0); + SetFlag(KDA_ARCH35_SCORE_EVENT); + + WaitFlag(KDA_ARCH35_SCORE_W_EVENT); + tileMmad(tileL0C1, tileL0A1, tileL0B1, packedRows, n, k, true, 0); + SetFlag(KDA_ARCH35_SCORE_W_EVENT); + WaitFlag(KDA_ARCH35_SCORE_W_EVENT); + const uint64_t flagOffset1 = + directSlot + KDA_DIRECT_SCORE_SUBBLOCK_FLAG_STRIDE; + CrossCoreWaitFlag<0x4, PIPE_FIX>(KDA_DIRECT_SCORE_FREE_FLAG + flagOffset1); + copyL0CToDirectUb( + blockDirectAqk, tileL0C1Top, KDA_DIRECT_SCORE_ROWS, 1, 1, 0); + copyL0CToDirectUb( + blockDirectAkk, tileL0C1Bottom, KDA_DIRECT_SCORE_ROWS, 1, 1, 0); + CrossCoreSetFlag<0x4, PIPE_FIX>(KDA_DIRECT_SCORE_READY_FLAG + flagOffset1); + SetFlag(KDA_ARCH35_SCORE_W_EVENT); + + WaitFlag(KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_W_EVENT); + } + + __aicore__ inline void ComputeRawAqkAkkCubeFullArch35( + uint64_t b, uint64_t hv, uint64_t start, uint64_t scoreSlot) + { + using ElementA = SCORE_T; + using ElementB = SCORE_T; + using ElementC = float; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::ColumnMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla< + KdaArchTag, ElementA, LayoutTagA, ElementB, LayoutTagB, ElementC, LayoutTagC>; + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + using LayoutTagL0A = typename TileCopy::LayoutTagL0A; + using LayoutTagL0B = typename TileCopy::LayoutTagL0B; + using CopyL1ToL0A = typename TileCopy::CopyL1ToL0A; + using CopyL1ToL0B = typename TileCopy::CopyL1ToL0B; + using TileMmad = Catlass::Gemm::Tile::TileMmadTla; + + constexpr uint32_t m = 64; + constexpr uint32_t n = 64; + constexpr uint32_t k = 128; + Catlass::Arch::Resource resource; + auto layoutA = tla::MakeLayout(m, k); + auto layoutB = tla::MakeLayout(k, n); + auto layoutC = tla::MakeLayout(m, n); + auto tensorQPos = tla::MakeTensor( + scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG)], + layoutA, Catlass::Arch::PositionGM{}); + auto tensorKPos = tla::MakeTensor( + scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W)], + layoutA, Catlass::Arch::PositionGM{}); + auto tensorKNeg = tla::MakeTensor( + scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG)], + layoutB, Catlass::Arch::PositionGM{}); + auto tensorAqk = tla::MakeTensor( + aqk_[AOffset(b, hv, start, 0)], layoutC, Catlass::Arch::PositionGM{}); + auto tensorAkk = tla::MakeTensor( + akk_[AOffset(b, hv, start, 0)], layoutC, Catlass::Arch::PositionGM{}); + auto blockQPos = GetTile(tensorQPos, tla::MakeCoord(0, 0), tla::MakeShape(m, k)); + auto blockKPos = GetTile(tensorKPos, tla::MakeCoord(0, 0), tla::MakeShape(m, k)); + auto blockKNeg = GetTile(tensorKNeg, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto blockAqk = GetTile(tensorAqk, tla::MakeCoord(0, 0), tla::MakeShape(m, n)); + auto blockAkk = GetTile(tensorAkk, tla::MakeCoord(0, 0), tla::MakeShape(m, n)); + + using CopyGmToL1A = typename TileCopy::template CopyGmToL1A; + using CopyGmToL1B = typename TileCopy::template CopyGmToL1B; + using CopyL0CToDst = typename TileCopy::template CopyL0CToDst; + + constexpr uint32_t l1ABytes = m * k * sizeof(ElementA); + LocalTensor l1A0 = resource.l1Buf.template GetBufferByByte(0); + LocalTensor l1A1 = resource.l1Buf.template GetBufferByByte(l1ABytes); + LocalTensor l1B = resource.l1Buf.template GetBufferByByte(2 * l1ABytes); + LocalTensor l0A = resource.l0ABuf.template GetBufferByByte(0); + LocalTensor l0B = resource.l0BBuf.template GetBufferByByte(0); + LocalTensor l0C = resource.l0CBuf.template GetBufferByByte(0); + auto layoutL1A = tla::MakeLayout(m, k); + auto layoutL1B = tla::MakeLayout(k, n); + auto layoutL0A = tla::MakeLayout(m, k); + auto layoutL0B = tla::MakeLayout(k, n); + auto layoutL0C = tla::MakeLayoutL0C(m, n); + auto tensorL1A0 = tla::MakeTensor(l1A0, layoutL1A, Catlass::Arch::PositionL1{}); + auto tensorL1A1 = tla::MakeTensor(l1A1, layoutL1A, Catlass::Arch::PositionL1{}); + auto tensorL1B = tla::MakeTensor(l1B, layoutL1B, Catlass::Arch::PositionL1{}); + auto tensorL0A = tla::MakeTensor(l0A, layoutL0A, Catlass::Arch::PositionL0A{}); + auto tensorL0B = tla::MakeTensor(l0B, layoutL0B, Catlass::Arch::PositionL0B{}); + auto tensorL0C = tla::MakeTensor(l0C, layoutL0C, Catlass::Arch::PositionL0C{}); + auto tileL1A0 = GetTile(tensorL1A0, tla::MakeCoord(0, 0), tla::MakeShape(m, k)); + auto tileL1A1 = GetTile(tensorL1A1, tla::MakeCoord(0, 0), tla::MakeShape(m, k)); + auto tileL1B = GetTile(tensorL1B, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto tileL0A = GetTile(tensorL0A, tla::MakeCoord(0, 0), tla::MakeShape(m, k)); + auto tileL0B = GetTile(tensorL0B, tla::MakeCoord(0, 0), tla::MakeShape(k, n)); + auto tileL0C = GetTile(tensorL0C, tla::MakeCoord(0, 0), tla::MakeShape(m, n)); + + CopyGmToL1A copyGmToL1A; + CopyGmToL1B copyGmToL1B; + CopyL1ToL0A copyL1ToL0A; + CopyL1ToL0B copyL1ToL0B; + CopyL0CToDst copyL0CToDst; + TileMmad tileMmad; + + copyGmToL1B(tensorL1B, blockKNeg); + copyGmToL1A(tensorL1A0, blockQPos); + SetFlag(KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + copyL1ToL0B(tileL0B, tileL1B); + copyL1ToL0A(tileL0A, tileL1A0); + SetFlag(KDA_ARCH35_SCORE_EVENT); + copyGmToL1A(tensorL1A1, blockKPos); + SetFlag(KDA_ARCH35_SCORE_W_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + tileMmad(tileL0C, tileL0A, tileL0B, m, n, k, true, 0); + SetFlag(KDA_ARCH35_SCORE_EVENT); + SetFlag(KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_W_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + copyL0CToDst(blockAqk, tensorL0C); + SetFlag(KDA_ARCH35_SCORE_EVENT); + copyL1ToL0B(tileL0B, tileL1B); + copyL1ToL0A(tileL0A, tileL1A1); + SetFlag(KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + tileMmad(tileL0C, tileL0A, tileL0B, m, n, k, true, 0); + SetFlag(KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + copyL0CToDst(blockAkk, tensorL0C); + SetFlag(KDA_ARCH35_SCORE_EVENT); + WaitFlag(KDA_ARCH35_SCORE_EVENT); + } +#endif + + __aicore__ inline void ComputeRawAqkAkkCubeBlock(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT, + uint64_t rowBegin, uint64_t rowCount, + bool readScoreScratch = false, uint64_t scoreSlot = 0, + uint64_t colCount = 0) + { + if (colCount == 0 || colCount > curT) { + colCount = curT; + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128) { + if (KDA_ARCH35_ENABLE_MANUAL_SCORE_PIPELINE && fusePostWu_ && + rowBegin == 0 && rowCount == 64 && colCount == 64) { + ComputeRawAqkAkkCubeFullArch35(b, hv, start, scoreSlot); + return; + } + } +#endif + using ElementA = SCORE_T; + using ElementB = SCORE_T; + using ElementC = float; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::ColumnMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using BlockMmad = Catlass::Gemm::Block::BlockMmadTla; + + Catlass::Arch::Resource resource; + BlockMmad blockMmad(resource); + auto layoutA = tla::MakeLayout(BT_, K_); + auto layoutB = tla::MakeLayout(K_, BT_); + auto layoutC = tla::MakeLayout(BT_, BT_); + const bool paddedTail = curT < BT_; + const uint64_t mmRowCount = paddedTail ? (rowCount + 15) / 16 * 16 : rowCount; + const uint64_t mmColCount = paddedTail ? BT_ : colCount; + Catlass::GemmCoord shape{static_cast(mmRowCount), static_cast(mmColCount), + static_cast(K_)}; + + (void)readScoreScratch; + auto tensorQPos = + tla::MakeTensor(scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG)], + layoutA, Catlass::Arch::PositionGM{}); + auto tensorKPos = + tla::MakeTensor(scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W)], + layoutA, Catlass::Arch::PositionGM{}); + auto tensorKNeg = + tla::MakeTensor(scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG)], + layoutB, Catlass::Arch::PositionGM{}); + auto aqkBase = paddedTail + ? solveWorkspace_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_RAW_AQK)] + : aqk_[AOffset(b, hv, start, 0)]; + auto akkBase = paddedTail + ? solveWorkspace_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_RAW_AKK)] + : akk_[AOffset(b, hv, start, 0)]; + auto tensorAqk = tla::MakeTensor(aqkBase, layoutC, Catlass::Arch::PositionGM{}); + auto tensorAkk = tla::MakeTensor(akkBase, layoutC, Catlass::Arch::PositionGM{}); + + auto blockQPos = GetTile(tensorQPos, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.k())); + auto blockKPos = GetTile(tensorKPos, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.k())); + auto blockKNeg = GetTile(tensorKNeg, tla::MakeCoord(0, 0), tla::MakeShape(shape.k(), shape.n())); + auto blockAqk = GetTile(tensorAqk, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.n())); + auto blockAkk = GetTile(tensorAkk, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.n())); + + blockMmad.preSetFlags(); + blockMmad(blockQPos, blockKNeg, blockAqk, shape); + blockMmad(blockKPos, blockKNeg, blockAkk, shape); + blockMmad.finalWaitFlags(); + } + + __aicore__ inline bool UseAkkCubeSolve(uint64_t curT) const + { + return curT > 0 && curT <= BT_ && (BT_ == 64 || BT_ == 128) && K_ >= 16 && V_ >= 16 && + V_ <= 256 && K_ % 16 == 0 && V_ % 16 == 0; + } + + __aicore__ inline bool UsePostWuCube(uint64_t curT) const + { + return curT > 0 && curT <= BT_ && (BT_ == 64 || BT_ == 128) && K_ >= 16 && V_ >= 16 && + V_ <= 256 && K_ % 16 == 0 && V_ % 16 == 0; + } + + __aicore__ inline void CopyLocalFloat(LocalTensor dst, LocalTensor src, uint64_t count) + { + if (count == 0) { + return; + } + Adds(dst, src, 0.0f, static_cast(count)); + PipeBarrier(); + } + + __aicore__ inline void FillLocalFloat(LocalTensor dst, float value, uint64_t count) + { + if (count == 0) { + return; + } + Duplicate(dst, value, static_cast(count)); + PipeBarrier(); + } + + __aicore__ inline void ForwardSubDiag16(LocalTensor diag, LocalTensor row, + LocalTensor prod, LocalTensor rowBrcb, + LocalTensor reduced, uint64_t valid) + { + constexpr uint32_t brcbStride = 8; + constexpr uint32_t diagSize = KDA_SOLVE_DIAG_BT; + constexpr uint8_t rowBlk = diagSize * sizeof(float) / 32; + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + ForwardSubDiag16Regbase( + (__ubuf__ float *)reinterpret_cast(diag.GetPhyAddr()), + static_cast(valid)); +#else + for (uint64_t i = 2; i < valid; ++i) { + uint32_t rowOffset = static_cast(i * diagSize); + DataCopy(row, diag[rowOffset], diagSize); + PipeBarrier(); + + Brcb(rowBrcb, row, diagSize / brcbStride, {1, 8}); + PipeBarrier(); + for (uint32_t col = 0; col < diagSize; col += brcbStride) { + Mul(prod[col], diag[col], rowBrcb, brcbStride, static_cast(diagSize), + {1, 1, 0, rowBlk, rowBlk, 1}); + } + PipeBarrier(); + + uint32_t remain = diagSize; + while (remain > 1) { + uint32_t calcCount = (remain / 2) * diagSize; + remain = (remain + 1) / 2; + Add(prod, prod, prod[remain * diagSize], calcCount); + PipeBarrier(); + } + DataCopy(reduced, prod, diagSize); + PipeBarrier(); + Add(row, row, reduced, diagSize); + PipeBarrier(); + DataCopy(diag[rowOffset], row, diagSize); + PipeBarrier(); + } +#endif + + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + for (uint32_t i = 0; i < diagSize; ++i) { + uint32_t diagOffset = i * diagSize + i; + if (i < valid) { + diag.SetValue(diagOffset, diag.GetValue(diagOffset) + 1.0f); + } else { + diag.SetValue(diagOffset, 1.0f); + } + } + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + } + + __aicore__ inline void SolveDiagonalBlocksInRows(LocalTensor akkMat, LocalTensor xMat, + LocalTensor arena, uint64_t scratchBase, + uint64_t curT, uint64_t rowBegin, uint64_t rowCount) + { + constexpr uint32_t diagSize = KDA_SOLVE_DIAG_BT; + constexpr uint32_t diagElements = diagSize * diagSize; + constexpr uint32_t brcbElements = diagSize * 8; + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + (void)akkMat; + (void)arena; + (void)scratchBase; + uint64_t rowEnd = rowBegin + rowCount; + for (uint64_t blockBegin = 0; blockBegin < BT_; blockBegin += diagSize) { + if (blockBegin < rowBegin || blockBegin + diagSize > rowEnd) { + continue; + } + uint64_t localBlockRow = blockBegin - rowBegin; + uint64_t valid = blockBegin < curT ? curT - blockBegin : 0; + if (valid > diagSize) { + valid = diagSize; + } + ForwardSubDiag16StridedRegbase( + (__ubuf__ float *)reinterpret_cast(xMat.GetPhyAddr()), + static_cast(BT_), static_cast(localBlockRow), + static_cast(blockBegin), static_cast(valid)); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + for (uint32_t rowIdx = 0; rowIdx < diagSize; ++rowIdx) { + uint32_t diagOffset = + static_cast((localBlockRow + rowIdx) * BT_ + blockBegin + rowIdx); + if (rowIdx < valid) { + xMat.SetValue(diagOffset, xMat.GetValue(diagOffset) + 1.0f); + } else { + xMat.SetValue(diagOffset, 1.0f); + } + } + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + } +#else + LocalTensor diag = arena[scratchBase]; + LocalTensor row = diag[diagElements]; + LocalTensor prod = row[diagSize]; + LocalTensor rowBrcb = prod[diagElements]; + LocalTensor reduced = rowBrcb[brcbElements]; + + uint64_t rowEnd = rowBegin + rowCount; + for (uint64_t blockBegin = 0; blockBegin < BT_; blockBegin += diagSize) { + if (blockBegin < rowBegin || blockBegin + diagSize > rowEnd) { + continue; + } + Duplicate(diag, 0.0f, diagElements); + PipeBarrier(); + + uint64_t localBlockRow = blockBegin - rowBegin; + uint64_t valid = blockBegin < curT ? curT - blockBegin : 0; + if (valid > diagSize) { + valid = diagSize; + } + for (uint32_t rowIdx = 0; rowIdx < diagSize; ++rowIdx) { + uint64_t srcOffset = (localBlockRow + rowIdx) * BT_ + blockBegin; + Muls(diag[rowIdx * diagSize], akkMat[srcOffset], -1.0f, diagSize); + } + PipeBarrier(); + + ForwardSubDiag16(diag, row, prod, rowBrcb, reduced, valid); + for (uint32_t rowIdx = 0; rowIdx < diagSize; ++rowIdx) { + uint64_t dstOffset = (localBlockRow + rowIdx) * BT_ + blockBegin; + Adds(xMat[dstOffset], diag[rowIdx * diagSize], 0.0f, diagSize); + } + PipeBarrier(); + } +#endif + } + + __aicore__ inline void BuildPrefixMask(LocalTensor dst, uint64_t prefix, uint64_t count) + { + if (prefix > count) { + prefix = count; + } + Duplicate(dst, 0.0f, static_cast(count)); + if (prefix > 0) { + Duplicate(dst, 1.0f, static_cast(prefix)); + } + PipeBarrier(); + } + + __aicore__ inline uint64_t BuildCausalMask(uint64_t threshold, uint64_t colBegin) const + { + if (threshold <= colBegin) { + return ~0ULL; + } + if (threshold >= colBegin + KDA_SOLVE_BT) { + return 0ULL; + } + return ~0ULL << (threshold - colBegin); + } + + __aicore__ inline void BuildCausalSelectMasks(LocalTensor aqkMask, LocalTensor akkMask, + uint64_t rowBegin, uint64_t rowCount, uint64_t colBegin) + { + __ubuf__ uint64_t *aqkMaskPtr = reinterpret_cast<__ubuf__ uint64_t *>(aqkMask.GetPhyAddr()); + __ubuf__ uint64_t *akkMaskPtr = reinterpret_cast<__ubuf__ uint64_t *>(akkMask.GetPhyAddr()); + for (uint32_t localRow = 0; localRow < rowCount; ++localRow) { + uint32_t row = static_cast(rowBegin + localRow); + aqkMaskPtr[localRow] = BuildCausalMask(static_cast(row) + 1, colBegin); + akkMaskPtr[localRow] = BuildCausalMask(static_cast(row), colBegin); + } + } + + __aicore__ inline void SelectCausalRows(LocalTensor aqkMat, LocalTensor akkMat, + uint64_t rowBegin, uint64_t rowCount) + { +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128) { + SelectCausalRows64Regbase( + (__ubuf__ float *)reinterpret_cast(aqkMat.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(akkMat.GetPhyAddr()), + static_cast(rowBegin), static_cast(rowCount)); + PipeBarrier(); + return; + } +#endif + LocalTensor aqkMask = vecBuf_.Get()[KDA_SELECT_AQK_MASK_BYTE_OFFSET]; + LocalTensor akkMask = vecBuf_.Get()[KDA_SELECT_AKK_MASK_BYTE_OFFSET]; + LocalTensor zeroLocal = vecBuf_.Get()[KDA_SELECT_ZERO_FLOAT_OFFSET]; + Duplicate(zeroLocal, 0.0f, 8); + PipeBarrier(); + + uint64_t colBlockCount = (BT_ + KDA_SOLVE_BT - 1) / KDA_SOLVE_BT; + for (uint64_t colBlock = 0; colBlock < colBlockCount; ++colBlock) { + uint64_t maskOffset = colBlock * KDA_SELECT_COL_MASK_BYTES; + uint64_t colBegin = colBlock * KDA_SOLVE_BT; + BuildCausalSelectMasks(aqkMask[maskOffset], akkMask[maskOffset], rowBegin, rowCount, colBegin); + } + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + + uint8_t rowStride = static_cast(BT_ * sizeof(float) / 32); + BinaryRepeatParams repeatParams = {1, 0, 1, rowStride, 0, rowStride}; + for (uint64_t colBlock = 0; colBlock < colBlockCount; ++colBlock) { + uint64_t maskOffset = colBlock * KDA_SELECT_COL_MASK_BYTES; + uint64_t colBegin = colBlock * KDA_SOLVE_BT; + Select(aqkMat[colBegin], aqkMask[maskOffset], zeroLocal, aqkMat[colBegin], + SELMODE::VSEL_TENSOR_TENSOR_MODE, KDA_SOLVE_BT, static_cast(rowCount), repeatParams); + Select(akkMat[colBegin], akkMask[maskOffset], zeroLocal, akkMat[colBegin], + SELMODE::VSEL_TENSOR_TENSOR_MODE, KDA_SOLVE_BT, static_cast(rowCount), repeatParams); + } + PipeBarrier(); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + } + + __aicore__ inline void PrepareAqkAkkSolveInput64(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start) + { + LocalTensor arena = vecBuf_.Get(); + LocalTensor aqkMat = arena; + LocalTensor akkMat = arena[KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor xMat = arena[2 * KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor betaLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor betaBrcb = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT]; + LocalTensor maskLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512]; + LocalTensor oneHotLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512 + KDA_SOLVE_BT]; + + LoadAsFloatRow(beta_, BetaOffset(b, hv, start), betaLocal, KDA_SOLVE_BT); + Brcb(betaBrcb, betaLocal, 8, {1, 8}); + PipeBarrier(); + + DataCopy(aqkMat, aqk_[AOffset(b, hv, start, 0)], KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(akkMat, akk_[AOffset(b, hv, start, 0)], KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + + for (uint64_t col = 0; col < KDA_SOLVE_BT; col += 8) { + Mul(akkMat[col], akkMat[col], betaBrcb, 8, KDA_SOLVE_BT, {1, 1, 1, 8, 8, 1}); + PipeBarrier(); + } + SelectCausalRows(aqkMat, akkMat, 0, KDA_SOLVE_BT); + + Muls(xMat, akkMat, -1.0f, KDA_SOLVE_MATRIX_ELEMENTS); + PipeBarrier(); + for (uint64_t row = 0; row < KDA_SOLVE_BT; ++row) { + BuildPrefixMask(maskLocal, row + 1, KDA_SOLVE_BT); + BuildPrefixMask(oneHotLocal, row, KDA_SOLVE_BT); + Sub(maskLocal, maskLocal, oneHotLocal, KDA_SOLVE_BT); + PipeBarrier(); + Add(xMat[row * KDA_SOLVE_BT], xMat[row * KDA_SOLVE_BT], maskLocal, KDA_SOLVE_BT); + PipeBarrier(); + } + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(aqk_[AOffset(b, hv, start, 0)], aqkMat, KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(akk_[AOffset(b, hv, start, 0)], akkMat, KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(h_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X)], xMat, + KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void PrepareAqkAkkSolveInputTail(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT) + { + uint64_t elemCount = curT * KDA_SOLVE_BT; + LocalTensor arena = vecBuf_.Get(); + LocalTensor aqkMat = arena; + LocalTensor akkMat = arena[KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor xMat = arena[2 * KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor betaLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor betaBrcb = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT]; + LocalTensor maskLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512]; + LocalTensor oneHotLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512 + KDA_SOLVE_BT]; + + FillLocalFloat(betaLocal, 0.0f, KDA_SOLVE_BT); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + LoadAsFloatRow(beta_, BetaOffset(b, hv, start), betaLocal, curT); + Brcb(betaBrcb, betaLocal, 8, {1, 8}); + PipeBarrier(); + + DataCopy(aqkMat, aqk_[AOffset(b, hv, start, 0)], static_cast(elemCount)); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + if (elemCount < KDA_SOLVE_MATRIX_ELEMENTS) { + FillLocalFloat(aqkMat[elemCount], 0.0f, KDA_SOLVE_MATRIX_ELEMENTS - elemCount); + } + DataCopy(akkMat, akk_[AOffset(b, hv, start, 0)], static_cast(elemCount)); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + if (elemCount < KDA_SOLVE_MATRIX_ELEMENTS) { + FillLocalFloat(akkMat[elemCount], 0.0f, KDA_SOLVE_MATRIX_ELEMENTS - elemCount); + } + + for (uint64_t col = 0; col < KDA_SOLVE_BT; col += 8) { + Mul(akkMat[col], akkMat[col], betaBrcb, 8, KDA_SOLVE_BT, {1, 1, 1, 8, 8, 1}); + PipeBarrier(); + } + SelectCausalRows(aqkMat, akkMat, 0, KDA_SOLVE_BT); + + Muls(xMat, akkMat, -1.0f, KDA_SOLVE_MATRIX_ELEMENTS); + PipeBarrier(); + for (uint64_t row = 0; row < KDA_SOLVE_BT; ++row) { + BuildPrefixMask(maskLocal, row + 1, KDA_SOLVE_BT); + BuildPrefixMask(oneHotLocal, row, KDA_SOLVE_BT); + Sub(maskLocal, maskLocal, oneHotLocal, KDA_SOLVE_BT); + PipeBarrier(); + Add(xMat[row * KDA_SOLVE_BT], xMat[row * KDA_SOLVE_BT], maskLocal, KDA_SOLVE_BT); + PipeBarrier(); + } + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(aqk_[AOffset(b, hv, start, 0)], aqkMat, static_cast(elemCount)); + DataCopy(h_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X)], xMat, + KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(h_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0)], akkMat, + KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void GetSolveRowRange(uint64_t curT, uint64_t subBlockIdx, uint64_t subBlockNum, + uint64_t &rowBegin, uint64_t &rowEnd) const + { + if (subBlockNum == 0 || subBlockIdx >= subBlockNum) { + rowBegin = 0; + rowEnd = 0; + return; + } + rowBegin = (curT * subBlockIdx) / subBlockNum; + rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + } + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + __aicore__ inline void InitializeDirectScoreUbArch35() + { + for (uint32_t slot = 0; slot < KDA_DIRECT_SCORE_QUEUE_DEPTH; ++slot) { + CrossCoreSetFlag<0x4, PIPE_V>(KDA_DIRECT_SCORE_FREE_FLAG + slot); + } + } + + __aicore__ inline void ProcessDirectScoreSolveRowsArch35( + uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t rowBegin, uint32_t directSlot) + { + CrossCoreWaitFlag<0x4, PIPE_V>(KDA_DIRECT_SCORE_READY_FLAG + directSlot); + + Catlass::Arch::Resource resource; + LocalTensor arena = vecBuf_.Get(); + LocalTensor directBase = resource.ubBuf.template GetBufferByByte( + KDA_DIRECT_SCORE_UB_BYTE_OFFSET + + directSlot * KDA_DIRECT_SCORE_SLOT_ELEMENTS * sizeof(float)); + LocalTensor aqkMat = directBase; + LocalTensor akkMat = directBase[KDA_DIRECT_SCORE_MATRIX_ELEMENTS]; + LocalTensor xMat = directBase[2 * KDA_DIRECT_SCORE_MATRIX_ELEMENTS]; + + SelectCausalRows64Regbase( + (__ubuf__ float *)reinterpret_cast(aqkMat.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(akkMat.GetPhyAddr()), + static_cast(rowBegin), KDA_DIRECT_SCORE_ROWS); + PipeBarrier(); + Muls(xMat, akkMat, -1.0f, KDA_DIRECT_SCORE_MATRIX_ELEMENTS); + PipeBarrier(); + SolveDiagonalBlocksInRows( + akkMat, xMat, arena, 0, BT_, rowBegin, KDA_DIRECT_SCORE_ROWS); + + LocalTensor aqkTyped = GateQTyped(0); + Muls(aqkMat, aqkMat, scale_, KDA_DIRECT_SCORE_MATRIX_ELEMENTS); + PipeBarrier(); + ClampFp32ToOutputType(aqkMat, KDA_DIRECT_SCORE_MATRIX_ELEMENTS); + Cast(aqkTyped, aqkMat, RoundMode::CAST_RINT, KDA_DIRECT_SCORE_MATRIX_ELEMENTS); + PipeBarrier(); + + const uint64_t token = start + rowBegin; + const uint64_t xBase = + SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(o_, AOffset(b, hv, token, 0), aqkTyped, + KDA_DIRECT_SCORE_MATRIX_ELEMENTS); + DataCopy(solveWorkspace_[xBase], xMat, KDA_DIRECT_SCORE_MATRIX_ELEMENTS); + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + CrossCoreSetFlag<0x4, PIPE_V>(KDA_DIRECT_SCORE_FREE_FLAG + directSlot); + } +#endif + + __aicore__ inline void PrepareAqkAkkSolveInputRows(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT, uint64_t rowBegin, + uint64_t rowEnd, bool storeLToAkk, bool storeLToScratch) + { + uint64_t rowCount = rowEnd - rowBegin; + if (rowCount == 0) { + return; + } + uint64_t validRowCount = rowBegin < curT ? curT - rowBegin : 0; + if (validRowCount > rowCount) { + validRowCount = rowCount; + } + uint64_t elemCount = rowCount * BT_; + uint64_t validElemCount = validRowCount * BT_; + LocalTensor arena = vecBuf_.Get(); + LocalTensor aqkMat = arena; + LocalTensor akkMat = arena[elemCount]; + LocalTensor xMat = arena[2 * elemCount]; + LocalTensor betaLocal = arena[3 * elemCount]; + LocalTensor betaBrcb = arena[3 * elemCount + BT_]; + LocalTensor maskLocal = arena[3 * elemCount + BT_ + 512]; + LocalTensor oneHotLocal = arena[3 * elemCount + BT_ + 512 + BT_]; + + uint64_t token = start + rowBegin; + + constexpr bool scoreCanIncludeBeta = +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128; +#else + false; +#endif + const bool scoreIncludesBeta = scoreCanIncludeBeta && HV_ % KDA_SCORE_LANES == 0; + if (validRowCount < rowCount) { + FillLocalFloat(aqkMat, 0.0f, elemCount); + FillLocalFloat(akkMat, 0.0f, elemCount); + if (!scoreIncludesBeta) { + FillLocalFloat(betaLocal, 0.0f, rowCount); + } + } + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + if (validRowCount > 0) { + if (!scoreIncludesBeta) { + LoadAsFloatRow(beta_, BetaOffset(b, hv, token), betaLocal, validRowCount); + } + if (curT < BT_) { + DataCopy(aqkMat, + solveWorkspace_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_RAW_AQK) + + rowBegin * BT_], + static_cast(validElemCount)); + DataCopy(akkMat, + solveWorkspace_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_RAW_AKK) + + rowBegin * BT_], + static_cast(validElemCount)); + } else { + DataCopy(aqkMat, aqk_[AOffset(b, hv, token, 0)], static_cast(validElemCount)); + DataCopy(akkMat, akk_[AOffset(b, hv, token, 0)], static_cast(validElemCount)); + } + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if (!scoreIncludesBeta) { + ApplyKdaRowScaleRegbase( + (__ubuf__ float *)reinterpret_cast(akkMat.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(betaLocal.GetPhyAddr()), + static_cast(rowCount), static_cast(BT_)); + PipeBarrier(); + } +#else + Brcb(betaBrcb, betaLocal, static_cast((rowCount + 7) / 8), {1, 8}); + PipeBarrier(); + uint8_t rowStride = static_cast(BT_ * sizeof(float) / 32); + for (uint64_t col = 0; col < BT_; col += 8) { + Mul(akkMat[col], akkMat[col], betaBrcb, 8, static_cast(rowCount), + {1, 1, 0, rowStride, rowStride, 1}); + } + PipeBarrier(); +#endif + if (validRowCount > 0) { + SelectCausalRows(aqkMat, akkMat, rowBegin, validRowCount); + } + + Muls(xMat, akkMat, -1.0f, static_cast(elemCount)); + PipeBarrier(); + if constexpr (SAFE_GATE) { + uint64_t scratchBase = 3 * elemCount + BT_ + 512 + 2 * BT_; + SolveDiagonalBlocksInRows(akkMat, xMat, arena, scratchBase, curT, rowBegin, rowCount); + } else if (curT < BT_) { + uint64_t scratchBase = 3 * elemCount + BT_ + 512 + 2 * BT_; + SolveDiagonalBlocksInRows(akkMat, xMat, arena, scratchBase, curT, rowBegin, rowCount); + } else { + for (uint64_t localRow = 0; localRow < rowCount; ++localRow) { + uint64_t row = rowBegin + localRow; + BuildPrefixMask(maskLocal, row + 1, BT_); + BuildPrefixMask(oneHotLocal, row, BT_); + Sub(maskLocal, maskLocal, oneHotLocal, static_cast(BT_)); + PipeBarrier(); + Add(xMat[localRow * BT_], xMat[localRow * BT_], maskLocal, static_cast(BT_)); + PipeBarrier(); + } + } + + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; + uint64_t lBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0) + rowBegin * BT_; +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128) { + LocalTensor aqkTyped = GateQTyped(0); + if (validElemCount > 0) { + Muls(aqkMat, aqkMat, scale_, static_cast(validElemCount)); + PipeBarrier(); + ClampFp32ToOutputType(aqkMat, static_cast(validElemCount)); + Cast(aqkTyped, aqkMat, RoundMode::CAST_RINT, static_cast(validElemCount)); + PipeBarrier(); + } + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + if (validElemCount > 0) { + CopyVectorOut(o_, AOffset(b, hv, token, 0), aqkTyped, validElemCount); + } + DataCopy(solveWorkspace_[xBase], xMat, static_cast(elemCount)); + if (storeLToScratch) { + DataCopy(solveWorkspace_[lBase], akkMat, static_cast(elemCount)); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + return; + } +#endif + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + if (validRowCount > 0) { + DataCopy(aqk_[AOffset(b, hv, token, 0)], aqkMat, static_cast(validElemCount)); + if (storeLToAkk) { + DataCopy(akk_[AOffset(b, hv, token, 0)], akkMat, static_cast(validElemCount)); + } + } + DataCopy(solveWorkspace_[xBase], xMat, static_cast(elemCount)); + if (storeLToScratch) { + DataCopy(solveWorkspace_[lBase], akkMat, static_cast(elemCount)); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void CubeGemmSolveSub(GlobalTensor &tensorA, uint64_t baseA, uint64_t rowA, uint64_t colA, + GlobalTensor &tensorB, uint64_t baseB, uint64_t rowB, uint64_t colB, + GlobalTensor &tensorC, uint64_t baseC, uint64_t rowC, uint64_t colC, + uint32_t m, uint32_t n, uint32_t k) + { + using ElementA = float; + using ElementB = float; + using ElementC = float; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using BlockMmad = Catlass::Gemm::Block::BlockMmadTla; + Catlass::Arch::Resource resource; + auto layoutA = tla::MakeLayout(BT_, BT_); + auto layoutB = tla::MakeLayout(BT_, BT_); + auto layoutC = tla::MakeLayout(BT_, BT_); + auto tensorLayoutA = tla::MakeTensor(tensorA[baseA], layoutA, Catlass::Arch::PositionGM{}); + auto tensorLayoutB = tla::MakeTensor(tensorB[baseB], layoutB, Catlass::Arch::PositionGM{}); + auto tensorLayoutC = tla::MakeTensor(tensorC[baseC], layoutC, Catlass::Arch::PositionGM{}); + Catlass::GemmCoord shape{m, n, k}; + auto blockA = GetTile(tensorLayoutA, tla::MakeCoord(rowA, colA), tla::MakeShape(shape.m(), shape.k())); + auto blockB = GetTile(tensorLayoutB, tla::MakeCoord(rowB, colB), tla::MakeShape(shape.k(), shape.n())); + auto blockC = GetTile(tensorLayoutC, tla::MakeCoord(rowC, colC), tla::MakeShape(shape.m(), shape.n())); + BlockMmad blockMmad(resource); + blockMmad(blockA, blockB, blockC, shape); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + SetFlag(KDA_ARCH35_SOLVE_FIX_EVENT); + WaitFlag(KDA_ARCH35_SOLVE_FIX_EVENT); +#else + PipeBarrier(); +#endif + } + + __aicore__ inline void AddSolveTmpToX(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + bool storeAkk) + { + LocalTensor arena = vecBuf_.Get(); + LocalTensor xLocal = arena; + LocalTensor tmpLocal = arena[KDA_SOLVE_MATRIX_ELEMENTS]; + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + DataCopy(xLocal, h_[xBase], KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(tmpLocal, h_[tmpBase], KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + + Add(xLocal, xLocal, tmpLocal, KDA_SOLVE_MATRIX_ELEMENTS); + PipeBarrier(); + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(h_[xBase], xLocal, KDA_SOLVE_MATRIX_ELEMENTS); + if (storeAkk) { + DataCopy(akk_[AOffset(b, hv, start, 0)], xLocal, KDA_SOLVE_MATRIX_ELEMENTS); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void AddSolveTmpToXTail(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT, bool storeAkk) + { + uint64_t elemCount = curT * KDA_SOLVE_BT; + LocalTensor arena = vecBuf_.Get(); + LocalTensor xLocal = arena; + LocalTensor tmpLocal = arena[KDA_SOLVE_MATRIX_ELEMENTS]; + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + DataCopy(xLocal, h_[xBase], KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(tmpLocal, h_[tmpBase], KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + + Add(xLocal, xLocal, tmpLocal, KDA_SOLVE_MATRIX_ELEMENTS); + PipeBarrier(); + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(h_[xBase], xLocal, KDA_SOLVE_MATRIX_ELEMENTS); + if (storeAkk) { + DataCopy(akk_[AOffset(b, hv, start, 0)], xLocal, static_cast(elemCount)); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void AddSolveTmpToXRows(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT, uint64_t rowBegin, uint64_t rowEnd, bool storeAkk) + { + uint64_t rowCount = rowEnd - rowBegin; + if (rowCount == 0) { + return; + } + uint64_t validRowCount = rowBegin < curT ? curT - rowBegin : 0; + if (validRowCount > rowCount) { + validRowCount = rowCount; + } + uint64_t elemCount = rowCount * BT_; + uint64_t validElemCount = validRowCount * BT_; + LocalTensor arena = vecBuf_.Get(); + LocalTensor xLocal = arena; + LocalTensor tmpLocal = arena[elemCount]; + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP) + rowBegin * BT_; + uint64_t token = start + rowBegin; + + DataCopy(xLocal, solveWorkspace_[xBase], static_cast(elemCount)); + DataCopy(tmpLocal, solveWorkspace_[tmpBase], static_cast(elemCount)); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + + Add(xLocal, xLocal, tmpLocal, static_cast(elemCount)); + PipeBarrier(); + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(solveWorkspace_[xBase], xLocal, static_cast(elemCount)); + if (storeAkk && validRowCount > 0) { + DataCopy(akk_[AOffset(b, hv, token, 0)], xLocal, static_cast(validElemCount)); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void AddSolveTmpToXDiagRows(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t rowBegin, uint64_t rowEnd, bool storeAkk) + { + uint64_t rowCount = rowEnd - rowBegin; + if (rowCount == 0) { + return; + } + uint64_t elemCount = rowCount * BT_; + LocalTensor arena = vecBuf_.Get(); + LocalTensor xLocal = arena; + LocalTensor tmpLocal = arena[elemCount]; + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP) + rowBegin * BT_; + uint64_t token = start + rowBegin; + + DataCopy(xLocal, solveWorkspace_[xBase], static_cast(elemCount)); + DataCopy(tmpLocal, solveWorkspace_[tmpBase], static_cast(elemCount)); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + + for (uint64_t localRow = 0; localRow < rowCount; ++localRow) { + uint64_t row = rowBegin + localRow; + uint64_t col = (row / KDA_SOLVE_DIAG_BT) * KDA_SOLVE_DIAG_BT; + uint64_t offset = localRow * BT_ + col; + Add(xLocal[offset], xLocal[offset], tmpLocal[offset], KDA_SOLVE_DIAG_BT); + PipeBarrier(); + } + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(solveWorkspace_[xBase], xLocal, static_cast(elemCount)); + if (storeAkk) { + DataCopy(akk_[AOffset(b, hv, token, 0)], xLocal, static_cast(elemCount)); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void StoreSolveXRowsToAkk(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT, uint64_t rowBegin, uint64_t rowEnd) + { + uint64_t validRowCount = rowBegin < curT ? curT - rowBegin : 0; + uint64_t rowCount = rowEnd - rowBegin; + if (validRowCount > rowCount) { + validRowCount = rowCount; + } + if (validRowCount == 0) { + return; + } + uint64_t elemCount = validRowCount * BT_; + LocalTensor xLocal = vecBuf_.Get(); + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; + + DataCopy(xLocal, solveWorkspace_[xBase], static_cast(elemCount)); + SetFlag(mte2ToMte3Event_); + WaitFlag(mte2ToMte3Event_); + DataCopy(akk_[AOffset(b, hv, start + rowBegin, 0)], xLocal, static_cast(elemCount)); + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + } + + __aicore__ inline void ComputeAkkMergeCube(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start) + { + uint64_t aiBase = AOffset(b, hv, start, 0); + uint64_t negABase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + for (uint32_t mergeSize = 2 * KDA_SOLVE_DIAG_BT; mergeSize <= BT_; mergeSize *= 2) { + uint32_t half = mergeSize / 2; + for (uint32_t block = 0; block < BT_; block += mergeSize) { + uint32_t lower = block + half; + CubeGemmSolveSub(akk_, aiBase, lower, lower, solveWorkspace_, negABase, lower, block, + solveWorkspace_, tmpBase, 0, 0, half, half, half); + CubeGemmSolveSub(solveWorkspace_, tmpBase, 0, 0, akk_, aiBase, block, block, + akk_, aiBase, lower, block, half, half, half); + } + } + } + + __aicore__ inline void ComputeAkkMergeCubeWorkspace(uint64_t b, uint64_t hv, uint64_t chunkIdx) + { + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + for (uint32_t mergeSize = 2 * KDA_SOLVE_DIAG_BT; mergeSize <= BT_; mergeSize *= 2) { + uint32_t half = mergeSize / 2; + for (uint32_t block = 0; block < BT_; block += mergeSize) { + uint32_t lower = block + half; + CubeGemmSolveSub(solveWorkspace_, xBase, lower, lower, solveWorkspace_, xBase, lower, block, + solveWorkspace_, tmpBase, 0, 0, half, half, half); + CubeGemmSolveSub(solveWorkspace_, tmpBase, 0, 0, solveWorkspace_, xBase, block, block, + solveWorkspace_, xBase, lower, block, half, half, half); + } + } + } + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + template + __aicore__ inline void LoadSolveTile(LocalTensor dst, LocalTensor src) + { + static_assert(TILE_SIZE == 16 || TILE_SIZE == 32, "arch35 solve tile must be 16 or 32 rows"); + constexpr uint32_t rowFractals = TILE_SIZE / 16; + constexpr uint32_t columnFractals = TILE_SIZE / 8; + LoadData2DParamsV2 loadParams; + loadParams.mStartPosition = 0; + loadParams.kStartPosition = 0; + loadParams.mStep = rowFractals; + loadParams.kStep = columnFractals; + loadParams.srcStride = rowFractals; + loadParams.dstStride = rowFractals; + loadParams.ifTranspose = TRANSPOSE; + LoadData(dst, src, loadParams); + } + + __aicore__ inline void ComputeAkkMergeCubeWorkspaceArch35(uint64_t b, uint64_t hv, uint64_t chunkIdx) + { + (void)b; + (void)hv; + SetMMLayoutTransform(true); + + using Element = float; + using LayoutTag = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + using LayoutTagL0A = typename TileCopy::LayoutTagL0A; + using LayoutTagL0B = typename TileCopy::LayoutTagL0B; + using TileMmad = Catlass::Gemm::Tile::TileMmadTla; + + constexpr uint32_t maxTile = 32; + constexpr uint32_t tileSlotBytes = maxTile * maxTile * sizeof(Element); + constexpr uint32_t a0Slot = 0; + constexpr uint32_t b0Slot = 1; + constexpr uint32_t a1Slot = 2; + constexpr uint32_t b1Slot = 3; + constexpr uint32_t diag0Slot = 4; + constexpr uint32_t diag1Slot = 5; + constexpr uint32_t tmp0Slot = 6; + constexpr uint32_t tmp1Slot = 7; + + Catlass::Arch::Resource resource; + LocalTensor l1A0 = resource.l1Buf.template GetBufferByByte(a0Slot * tileSlotBytes); + LocalTensor l1B0 = resource.l1Buf.template GetBufferByByte(b0Slot * tileSlotBytes); + LocalTensor l1A1 = resource.l1Buf.template GetBufferByByte(a1Slot * tileSlotBytes); + LocalTensor l1B1 = resource.l1Buf.template GetBufferByByte(b1Slot * tileSlotBytes); + LocalTensor l1Diag0 = + resource.l1Buf.template GetBufferByByte(diag0Slot * tileSlotBytes); + LocalTensor l1Diag1 = + resource.l1Buf.template GetBufferByByte(diag1Slot * tileSlotBytes); + LocalTensor l1Tmp0 = + resource.l1Buf.template GetBufferByByte(tmp0Slot * tileSlotBytes); + LocalTensor l1Tmp1 = + resource.l1Buf.template GetBufferByByte(tmp1Slot * tileSlotBytes); + + LocalTensor l0A0 = resource.l0ABuf.template GetBufferByByte(a0Slot * tileSlotBytes); + LocalTensor l0B0 = resource.l0BBuf.template GetBufferByByte(b0Slot * tileSlotBytes); + LocalTensor l0A1 = resource.l0ABuf.template GetBufferByByte(a1Slot * tileSlotBytes); + LocalTensor l0B1 = resource.l0BBuf.template GetBufferByByte(b1Slot * tileSlotBytes); + LocalTensor l0A2 = + resource.l0ABuf.template GetBufferByByte(diag0Slot * tileSlotBytes); + LocalTensor l0B2 = + resource.l0BBuf.template GetBufferByByte(diag0Slot * tileSlotBytes); + LocalTensor l0A3 = + resource.l0ABuf.template GetBufferByByte(diag1Slot * tileSlotBytes); + LocalTensor l0B3 = + resource.l0BBuf.template GetBufferByByte(diag1Slot * tileSlotBytes); + LocalTensor l0C0 = resource.l0CBuf.template GetBufferByByte(a0Slot * tileSlotBytes); + LocalTensor l0C1 = resource.l0CBuf.template GetBufferByByte(a1Slot * tileSlotBytes); + LocalTensor l0C2 = + resource.l0CBuf.template GetBufferByByte(diag0Slot * tileSlotBytes); + LocalTensor l0C3 = + resource.l0CBuf.template GetBufferByByte(diag1Slot * tileSlotBytes); + + TileMmad tileMmad; + const uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + const uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + auto gmLayout = tla::MakeLayout(BT_, BT_); + auto tensorX = tla::MakeTensor(solveWorkspace_[xBase], gmLayout, Catlass::Arch::PositionGM{}); + auto tensorTmp = tla::MakeTensor(solveWorkspace_[tmpBase], gmLayout, Catlass::Arch::PositionGM{}); + + constexpr uint32_t tile16 = 16; + auto shape16 = tla::MakeShape(tile16, tile16); + auto blockA0 = GetTile(tensorX, tla::MakeCoord(16, 16), shape16); + auto blockB0 = GetTile(tensorX, tla::MakeCoord(16, 0), shape16); + auto blockA1 = GetTile(tensorX, tla::MakeCoord(48, 48), shape16); + auto blockB1 = GetTile(tensorX, tla::MakeCoord(48, 32), shape16); + auto blockDiag0 = GetTile(tensorX, tla::MakeCoord(0, 0), shape16); + auto blockDiag1 = GetTile(tensorX, tla::MakeCoord(32, 32), shape16); + auto blockTmp0 = GetTile(tensorTmp, tla::MakeCoord(0, 0), shape16); + auto blockTmp1 = GetTile(tensorTmp, tla::MakeCoord(16, 0), shape16); + auto blockOut0 = GetTile(tensorX, tla::MakeCoord(16, 0), shape16); + auto blockOut1 = GetTile(tensorX, tla::MakeCoord(48, 32), shape16); + + using CopyGmToL1A16 = typename TileCopy::template CopyGmToL1A; + using CopyGmToL1B16 = typename TileCopy::template CopyGmToL1B; + using CopyL0CToDst16 = typename TileCopy::template CopyL0CToDst; + CopyGmToL1A16 copyGmToL1A16; + CopyGmToL1B16 copyGmToL1B16; + CopyL0CToDst16 copyL0CToDst16; + + auto tensorL1A0 = tla::MakeTensor( + l1A0, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL1{}); + auto tensorL1B0 = tla::MakeTensor( + l1B0, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL1{}); + auto tensorL1A1 = tla::MakeTensor( + l1A1, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL1{}); + auto tensorL1B1 = tla::MakeTensor( + l1B1, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL1{}); + auto tensorL1Diag0 = tla::MakeTensor( + l1Diag0, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL1{}); + auto tensorL1Diag1 = tla::MakeTensor( + l1Diag1, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL1{}); + auto tensorL1Tmp0 = tla::MakeTensor( + l1Tmp0, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL1{}); + auto tensorL1Tmp1 = tla::MakeTensor( + l1Tmp1, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL1{}); + + auto tensorL0A0 = tla::MakeTensor( + l0A0, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL0A{}); + auto tensorL0B0 = tla::MakeTensor( + l0B0, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL0B{}); + auto tensorL0A1 = tla::MakeTensor( + l0A1, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL0A{}); + auto tensorL0B1 = tla::MakeTensor( + l0B1, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL0B{}); + auto tensorL0A2 = tla::MakeTensor( + l0A2, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL0A{}); + auto tensorL0B2 = tla::MakeTensor( + l0B2, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL0B{}); + auto tensorL0A3 = tla::MakeTensor( + l0A3, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL0A{}); + auto tensorL0B3 = tla::MakeTensor( + l0B3, tla::MakeLayout(tile16, tile16), Catlass::Arch::PositionL0B{}); + auto tensorL0C0 = tla::MakeTensor(l0C0, tla::MakeLayoutL0C(tile16, tile16), Catlass::Arch::PositionL0C{}); + auto tensorL0C1 = tla::MakeTensor(l0C1, tla::MakeLayoutL0C(tile16, tile16), Catlass::Arch::PositionL0C{}); + auto tensorL0C2 = tla::MakeTensor(l0C2, tla::MakeLayoutL0C(tile16, tile16), Catlass::Arch::PositionL0C{}); + auto tensorL0C3 = tla::MakeTensor(l0C3, tla::MakeLayoutL0C(tile16, tile16), Catlass::Arch::PositionL0C{}); + uint32_t localRow16 = 0; + uint32_t localColumn16 = 0; + auto tileL0A0 = GetTile(tensorL0A0, tla::MakeCoord(localRow16, localColumn16), shape16); + auto tileL0B0 = GetTile(tensorL0B0, tla::MakeCoord(localRow16, localColumn16), shape16); + auto tileL0A1 = GetTile(tensorL0A1, tla::MakeCoord(localRow16, localColumn16), shape16); + auto tileL0B1 = GetTile(tensorL0B1, tla::MakeCoord(localRow16, localColumn16), shape16); + auto tileL0A2 = GetTile(tensorL0A2, tla::MakeCoord(localRow16, localColumn16), shape16); + auto tileL0B2 = GetTile(tensorL0B2, tla::MakeCoord(localRow16, localColumn16), shape16); + auto tileL0A3 = GetTile(tensorL0A3, tla::MakeCoord(localRow16, localColumn16), shape16); + auto tileL0B3 = GetTile(tensorL0B3, tla::MakeCoord(localRow16, localColumn16), shape16); + auto tileL0C0 = GetTile(tensorL0C0, tla::MakeCoord(localRow16, localColumn16), shape16); + auto tileL0C1 = GetTile(tensorL0C1, tla::MakeCoord(localRow16, localColumn16), shape16); + auto tileL0C2 = GetTile(tensorL0C2, tla::MakeCoord(localRow16, localColumn16), shape16); + auto tileL0C3 = GetTile(tensorL0C3, tla::MakeCoord(localRow16, localColumn16), shape16); + + copyGmToL1A16(tensorL1A0, blockA0); + copyGmToL1B16(tensorL1B0, blockB0); + copyGmToL1A16(tensorL1A1, blockA1); + copyGmToL1B16(tensorL1B1, blockB1); + copyGmToL1B16(tensorL1Diag0, blockDiag0); + copyGmToL1B16(tensorL1Diag1, blockDiag1); + SetFlag(0); + WaitFlag(0); + + LoadSolveTile(l0A0, l1A0); + LoadSolveTile(l0B0, l1B0); + LoadSolveTile(l0A1, l1A1); + LoadSolveTile(l0B1, l1B1); + SetFlag(0); + WaitFlag(0); + tileMmad(tileL0C0, tileL0A0, tileL0B0, tile16, tile16, tile16, true, 0b11); + SetFlag(0); + tileMmad(tileL0C1, tileL0A1, tileL0B1, tile16, tile16, tile16, true, 0b11); + SetFlag(1); + WaitFlag(0); + copyL0CToDst16(blockTmp0, tileL0C0, 0b11); + SetFlag(0); + WaitFlag(1); + copyL0CToDst16(blockTmp1, tileL0C1, 0b11); + SetFlag(1); + WaitFlag(0); + copyGmToL1A16(tensorL1Tmp0, blockTmp0); + WaitFlag(1); + copyGmToL1A16(tensorL1Tmp1, blockTmp1); + SetFlag(0); + WaitFlag(0); + + LoadSolveTile(l0A2, l1Tmp0); + LoadSolveTile(l0B2, l1Diag0); + LoadSolveTile(l0A3, l1Tmp1); + LoadSolveTile(l0B3, l1Diag1); + SetFlag(0); + WaitFlag(0); + tileMmad(tileL0C2, tileL0A2, tileL0B2, tile16, tile16, tile16, true, 0b11); + SetFlag(2); + tileMmad(tileL0C3, tileL0A3, tileL0B3, tile16, tile16, tile16, true, 0b11); + SetFlag(3); + WaitFlag(2); + copyL0CToDst16(blockOut0, tileL0C2, 0b11); + SetFlag(2); + WaitFlag(3); + copyL0CToDst16(blockOut1, tileL0C3, 0b11); + SetFlag(3); + + constexpr uint32_t tile32 = 32; + auto shape32 = tla::MakeShape(tile32, tile32); + auto blockA32 = GetTile(tensorX, tla::MakeCoord(32, 32), shape32); + auto blockB32 = GetTile(tensorX, tla::MakeCoord(32, 0), shape32); + auto blockDiag32 = GetTile(tensorX, tla::MakeCoord(0, 0), shape32); + auto blockTmp32 = GetTile(tensorTmp, tla::MakeCoord(0, 0), shape32); + auto blockOut32 = GetTile(tensorX, tla::MakeCoord(32, 0), shape32); + using CopyGmToL1A32 = typename TileCopy::template CopyGmToL1A; + using CopyGmToL1B32 = typename TileCopy::template CopyGmToL1B; + using CopyL0CToDst32 = typename TileCopy::template CopyL0CToDst; + CopyGmToL1A32 copyGmToL1A32; + CopyGmToL1B32 copyGmToL1B32; + CopyL0CToDst32 copyL0CToDst32; + + auto tensorL1A32 = tla::MakeTensor( + l1A0, tla::MakeLayout(tile32, tile32), Catlass::Arch::PositionL1{}); + auto tensorL1B32 = tla::MakeTensor( + l1B0, tla::MakeLayout(tile32, tile32), Catlass::Arch::PositionL1{}); + auto tensorL1Diag32 = tla::MakeTensor( + l1Diag0, tla::MakeLayout(tile32, tile32), Catlass::Arch::PositionL1{}); + auto tensorL1Tmp32 = tla::MakeTensor( + l1Tmp0, tla::MakeLayout(tile32, tile32), Catlass::Arch::PositionL1{}); + auto tensorL0A32 = tla::MakeTensor( + l0A0, tla::MakeLayout(tile32, tile32), Catlass::Arch::PositionL0A{}); + auto tensorL0B32 = tla::MakeTensor( + l0B0, tla::MakeLayout(tile32, tile32), Catlass::Arch::PositionL0B{}); + auto tensorL0C32 = tla::MakeTensor(l0C0, tla::MakeLayoutL0C(tile32, tile32), Catlass::Arch::PositionL0C{}); + uint32_t localRow32 = 0; + uint32_t localColumn32 = 0; + auto tileL0A32 = GetTile(tensorL0A32, tla::MakeCoord(localRow32, localColumn32), shape32); + auto tileL0B32 = GetTile(tensorL0B32, tla::MakeCoord(localRow32, localColumn32), shape32); + auto tileL0C32 = GetTile(tensorL0C32, tla::MakeCoord(localRow32, localColumn32), shape32); + + WaitFlag(2); + WaitFlag(3); + copyGmToL1A32(tensorL1A32, blockA32); + copyGmToL1B32(tensorL1B32, blockB32); + copyGmToL1B32(tensorL1Diag32, blockDiag32); + SetFlag(0); + WaitFlag(0); + LoadSolveTile(l0A0, l1A0); + LoadSolveTile(l0B0, l1B0); + SetFlag(0); + WaitFlag(0); + tileMmad(tileL0C32, tileL0A32, tileL0B32, tile32, tile32, tile32, true, 0b11); + SetFlag(0); + WaitFlag(0); + copyL0CToDst32(blockTmp32, tileL0C32, 0b11); + SetFlag(0); + WaitFlag(0); + copyGmToL1A32(tensorL1Tmp32, blockTmp32); + SetFlag(0); + WaitFlag(0); + LoadSolveTile(l0A0, l1Tmp0); + LoadSolveTile(l0B0, l1Diag0); + SetFlag(0); + WaitFlag(0); + tileMmad(tileL0C32, tileL0A32, tileL0B32, tile32, tile32, tile32, true, 0b11); + SetFlag(0); + WaitFlag(0); + copyL0CToDst32(blockOut32, tileL0C32, 0b11); + SetFlag(0); + WaitFlag(0); + SetMMLayoutTransform(false); + } +#endif + + __aicore__ inline void ComputeAkkMergeCubeWorkspaceDispatch(uint64_t b, uint64_t hv, uint64_t chunkIdx) + { +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128) { + ComputeAkkMergeCubeWorkspace(b, hv, chunkIdx); + return; + } +#endif + ComputeAkkMergeCubeWorkspace(b, hv, chunkIdx); + } + + __aicore__ inline void ComputeAkkInverseMchFull(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start) + { + uint64_t aBase = AOffset(b, hv, start, 0); + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t yBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0); + uint64_t yNextBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y1); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + uint32_t diagBlocks = static_cast(BT_ / KDA_SOLVE_DIAG_BT); + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(akk_, aBase, off, off, akk_, aBase, off, off, solveWorkspace_, yBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + for (uint32_t iter = 0; iter < KDA_SOLVE_DIAG_MCH_ITERS; ++iter) { + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(solveWorkspace_, xBase, off, off, solveWorkspace_, yBase, off, off, + solveWorkspace_, tmpBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(mchSyncDoneFlag_); + if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(solveWorkspace_, yBase, off, off, solveWorkspace_, yBase, off, off, + solveWorkspace_, yNextBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + } + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(mchSyncReadyFlag_); + if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { + uint64_t oldYBase = yBase; + yBase = yNextBase; + yNextBase = oldYBase; + } + } + ComputeAkkMergeCube(b, hv, chunkIdx, start); + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(mchSyncDoneFlag_); + } + + __aicore__ inline void ScaleRowsByBeta(GlobalTensor &src, GlobalTensor &dst, uint64_t b, uint64_t hv, + uint64_t start, uint64_t rowBegin, uint64_t rowCount, uint64_t dim, + LocalTensor &betaLocal, LocalTensor &betaBrcb, + LocalTensor &matrixLocal, bool sourceSequenceMajor = false) + { + constexpr uint64_t vecElemsPerRepeat = 64; + constexpr uint64_t typedOffsetFloats = 20480; + constexpr uint64_t typedOffset = typedOffsetFloats * sizeof(float) / sizeof(T); + uint64_t elemCount = rowCount * dim; + uint64_t baseOffset = KVOffset(b, hv, start + rowBegin, 0, dim); + uint64_t sourceOffset = sourceSequenceMajor + ? VInputOffset(b, hv, start + rowBegin, 0) + : baseOffset; + uint64_t sourceStride = sourceSequenceMajor ? HV_ * dim : dim; + + if constexpr (IsSameType::value) { + CopyRowsIn(matrixLocal, src, sourceOffset, rowCount, dim, sourceStride); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + } else { + LocalTensor matrixTyped = vecBuf_.Get()[typedOffset]; + CopyRowsIn(matrixTyped, src, sourceOffset, rowCount, dim, sourceStride); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast(matrixLocal, matrixTyped, RoundMode::CAST_NONE, static_cast(elemCount)); + PipeBarrier(); + } + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + ApplyKdaRowScaleRegbase( + (__ubuf__ float *)reinterpret_cast(matrixLocal.GetPhyAddr()), + (__ubuf__ float *)reinterpret_cast(betaLocal.GetPhyAddr()), + static_cast(rowCount), static_cast(dim)); +#else + uint8_t repeatStride = static_cast(dim * sizeof(float) / 32); + for (uint64_t col = 0; col < dim; col += vecElemsPerRepeat) { + uint64_t mask = dim - col; + if (mask > vecElemsPerRepeat) { + mask = vecElemsPerRepeat; + } + Mul(matrixLocal[col], matrixLocal[col], betaBrcb, mask, static_cast(rowCount), + {1, 1, 0, repeatStride, repeatStride, 1}); + } + PipeBarrier(); +#endif + + if constexpr (IsSameType::value) { + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(dst[baseOffset], matrixLocal, static_cast(elemCount)); + } else { + LocalTensor matrixTyped = vecBuf_.Get()[typedOffset]; + Cast(matrixTyped, matrixLocal, RoundMode::CAST_RINT, static_cast(elemCount)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(dst[baseOffset], matrixTyped, static_cast(elemCount)); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void PrepareWuCubeInputs(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT, + uint64_t subBlockIdx, uint64_t subBlockNum) + { + uint64_t rowsPerSubBlock = (curT + subBlockNum - 1) / subBlockNum; + uint64_t rowBegin = subBlockIdx * rowsPerSubBlock; + if (rowBegin >= curT) { + return; + } + uint64_t rowCount = curT - rowBegin; + if (rowCount > rowsPerSubBlock) { + rowCount = rowsPerSubBlock; + } + LocalTensor arena = vecBuf_.Get(); + LocalTensor betaLocal = arena; + LocalTensor betaBrcb = arena[KDA_SOLVE_BT]; + LocalTensor matrixLocal = arena[KDA_SOLVE_BT + 512]; +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + // Keep each arch35 Cast/regbase panel within an 8K-element UB instruction span. + constexpr uint64_t maxScaleElements = 8192; + if (rowCount * K_ > maxScaleElements || rowCount * V_ > maxScaleElements) { + constexpr uint64_t tileRows = 16; + for (uint64_t tileRow = 0; tileRow < rowCount; tileRow += tileRows) { + uint64_t tileCount = rowCount - tileRow; + if (tileCount > tileRows) { + tileCount = tileRows; + } + LoadAsFloatRow(beta_, BetaOffset(b, hv, start + rowBegin + tileRow), betaLocal, tileCount); + ScaleRowsByBeta(w_, w_, b, hv, start, rowBegin + tileRow, tileCount, K_, + betaLocal, betaBrcb, matrixLocal); + ScaleRowsByBeta(v_, vNew_, b, hv, start, rowBegin + tileRow, tileCount, V_, + betaLocal, betaBrcb, matrixLocal, inputSequenceMajor_); + } + return; + } +#endif + LoadAsFloatRow(beta_, BetaOffset(b, hv, start + rowBegin), betaLocal, rowCount); +#if !defined(__CCE_AICORE__) || __CCE_AICORE__ != 310 + Brcb(betaBrcb, betaLocal, static_cast((rowCount + 7) / 8), {1, 8}); + PipeBarrier(); +#endif + ScaleRowsByBeta(w_, w_, b, hv, start, rowBegin, rowCount, K_, betaLocal, betaBrcb, matrixLocal); + ScaleRowsByBeta(v_, vNew_, b, hv, start, rowBegin, rowCount, V_, betaLocal, betaBrcb, + matrixLocal, inputSequenceMajor_); + } + + __aicore__ inline void FinalizePrepareIntermediates(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT, + uint64_t subBlockIdx, uint64_t subBlockNum) + { +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + constexpr bool qgScaledAlreadyStored = + SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128; +#else + constexpr bool qgScaledAlreadyStored = false; +#endif + constexpr uint64_t tileRows = 32; + // Keep tail rows on the same AIV that owns their padded solve rows. Splitting by curT would + // move short-tail export to AIV1 while AIV0 is still writing the solved matrix. + const uint64_t rowBegin = (BT_ * subBlockIdx) / subBlockNum; + uint64_t rowEnd = (BT_ * (subBlockIdx + 1)) / subBlockNum; + if (rowEnd > curT) { + rowEnd = curT; + } + if (rowBegin >= rowEnd) { + return; + } + for (uint64_t tileRow = rowBegin; tileRow < rowEnd; tileRow += tileRows) { + const uint64_t rows = (rowEnd - tileRow) > tileRows ? tileRows : (rowEnd - tileRow); + const uint64_t matrixElems = rows * BT_; + const uint64_t qgElems = rows * K_; + LocalTensor arena = vecBuf_.Get(); + LocalTensor aqkLocal = arena; + LocalTensor akkLocal = arena[matrixElems]; + LocalTensor qgLocal = arena[2 * matrixElems]; + const uint64_t typedOffset = + (2 * matrixElems + qgElems) * sizeof(float) / sizeof(T); + LocalTensor typedBase = vecBuf_.Get()[typedOffset]; + LocalTensor aqkTyped = typedBase; + LocalTensor akkTyped = typedBase[matrixElems]; + LocalTensor qgTyped = typedBase[2 * matrixElems]; + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128) { + const uint64_t xBase = + SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + tileRow * BT_; + CopyVectorIn(akkLocal, solveWorkspace_, xBase, matrixElems); + } else { + CopyVectorIn(aqkLocal, aqk_, AOffset(b, hv, start + tileRow, 0), matrixElems); + CopyVectorIn(akkLocal, akk_, AOffset(b, hv, start + tileRow, 0), matrixElems); + } +#else + CopyVectorIn(aqkLocal, aqk_, AOffset(b, hv, start + tileRow, 0), matrixElems); + CopyVectorIn(akkLocal, akk_, AOffset(b, hv, start + tileRow, 0), matrixElems); +#endif + if constexpr (!qgScaledAlreadyStored) { + CopyVectorIn(qgTyped, qg_, KVOffset(b, hv, start + tileRow, 0, K_), qgElems); + } + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (!(SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128)) { + Muls(aqkLocal, aqkLocal, scale_, static_cast(matrixElems)); + } +#else + Muls(aqkLocal, aqkLocal, scale_, static_cast(matrixElems)); +#endif + if constexpr (!qgScaledAlreadyStored) { + Cast(qgLocal, qgTyped, RoundMode::CAST_NONE, static_cast(qgElems)); + PipeBarrier(); + Muls(qgLocal, qgLocal, scale_, static_cast(qgElems)); + PipeBarrier(); + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (!(SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128)) { + ClampFp32ToOutputType(aqkLocal, static_cast(matrixElems)); + } +#else + ClampFp32ToOutputType(aqkLocal, static_cast(matrixElems)); +#endif + ClampFp32ToOutputType(akkLocal, static_cast(matrixElems)); + if constexpr (!qgScaledAlreadyStored) { + ClampFp32ToOutputType(qgLocal, static_cast(qgElems)); + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (!(SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128)) { + Cast(aqkTyped, aqkLocal, RoundMode::CAST_RINT, static_cast(matrixElems)); + } +#else + Cast(aqkTyped, aqkLocal, RoundMode::CAST_RINT, static_cast(matrixElems)); +#endif + Cast(akkTyped, akkLocal, RoundMode::CAST_RINT, static_cast(matrixElems)); + if constexpr (!qgScaledAlreadyStored) { + Cast(qgTyped, qgLocal, RoundMode::CAST_RINT, static_cast(qgElems)); + } + PipeBarrier(); + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (!(SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128)) { + CopyVectorOut(o_, AOffset(b, hv, start + tileRow, 0), aqkTyped, matrixElems); + } +#else + CopyVectorOut(o_, AOffset(b, hv, start + tileRow, 0), aqkTyped, matrixElems); +#endif + CopyVectorOut(u_, AOffset(b, hv, start + tileRow, 0), akkTyped, matrixElems); + if constexpr (!qgScaledAlreadyStored) { + CopyVectorOut(kg_, KVOffset(b, hv, start + tileRow, 0, K_), qgTyped, qgElems); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + } + + __aicore__ inline bool ResolveFlatChunk(uint64_t task, uint64_t &seq, uint64_t &b, uint64_t &h, uint64_t &hv, + uint64_t &chunkIdx, uint64_t &start, uint64_t &end) + { + hv = task % HV_; + uint64_t flatChunk = task / HV_; + if (!isVarLen_) { + seq = flatChunk / NT_; + b = seq; + chunkIdx = flatChunk % NT_; + start = chunkIdx * BT_; + end = start + BT_; + if (end > T_) { + end = T_; + } + } else { + if (!KdaVarlen::ResolveChunkRange( + cuSeqlensAddr_, chunkIndicesAddr_, N_, T_, BT_, flatChunk, + seq, start, end)) { + return false; + } + b = 0; + chunkIdx = flatChunk; + } + h = hv / (HV_ / H_); + return start < end; + } + + __aicore__ inline void ProcessChunkPreAiv(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t end, uint64_t subBlockIdx, + uint64_t subBlockNum) + { + if constexpr (IsSameType::value) { + ProcessChunkPreAivFp32(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); + } + } + + template + __aicore__ inline void JoinAivMte3() + { + if constexpr (CORE_TYPE == AscendC::AIV) { +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if (!isAivOnly_) { + if (!headPairMode_) { + Catlass::Arch::CrossCoreBarrier<0x1, PIPE_MTE3>(); + } + PipeBarrier(); + } +#endif + } + } + + template + __aicore__ inline void RunAicAfterBothAivReady(uint64_t subBlockIdx, uint64_t subBlockNum) + { + if constexpr (CORE_TYPE == AscendC::AIV) { + (void)subBlockIdx; + (void)subBlockNum; + JoinAivMte3(); + if constexpr (SAFE_GATE) { + Catlass::Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(syncReadyFlag_); + Catlass::Arch::CrossCoreWaitFlag(syncDoneFlag_); + } else { + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_MTE3>(mchSyncReadyFlag_); + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(mchSyncDoneFlag_); + } + } + } + + template + __aicore__ inline void SignalAicSolveReady() + { + if constexpr (CORE_TYPE == AscendC::AIV) { + JoinAivMte3(); + Catlass::Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(syncReadyFlag_); + } + } + + template + __aicore__ inline void WaitAicSolveDone() + { + if constexpr (CORE_TYPE == AscendC::AIV) { + Catlass::Arch::CrossCoreWaitFlag(syncDoneFlag_); + } + } + + template + __aicore__ inline void SignalPostWuReady() + { + if constexpr (CORE_TYPE == AscendC::AIV) { + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_MTE3>(postWuReadyFlag_); + } + } + + __aicore__ inline void ProcessChunkPreAivFp32(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t end, uint64_t subBlockIdx, + uint64_t subBlockNum, bool deferSafeSolve = false, + bool waitPendingSafeSolve = false, + uint64_t scoreLane = 0, bool pairHeads = false) + { + uint64_t curT = end - start; + if (curT == 0) { + return; + } + if constexpr (IsSameType::value) { + return; + } + + if (K_ < 16) { + return; + } + MaterializeRawGateChunkArch35(b, hv, start, curT); + bool usePostWuCube = UsePostWuCube(curT); + bool useAkkCubeSolve = UseAkkCubeSolve(curT); + uint64_t solveRowBegin = 0; + uint64_t solveRowEnd = 0; + GetSolveRowRange(BT_, subBlockIdx, subBlockNum, solveRowBegin, solveRowEnd); + // Safe-gate score factors need the bounded 16-row reference span. + // A single 64-row reference loses BF16 dynamic range for valid gates. + const bool useFullChunkScore = false; + uint64_t scoreBlockSize = useFullChunkScore ? curT : ScoreRefBlockSize(); + uint64_t scoreBlockCount = (curT + scoreBlockSize - 1) / scoreBlockSize; + uint64_t pipelineBlockCount = useFullChunkScore + ? scoreBlockCount + : (scoreBlockCount + KDA_SCORE_QUEUE_DEPTH - 1) / KDA_SCORE_QUEUE_DEPTH * + KDA_SCORE_QUEUE_DEPTH; + const bool useDirectScoreUb = +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + KDA_ARCH35_ENABLE_DIRECT_SCORE_UB && pairHeads && curT == 64 && scoreBlockCount == 2; +#else + false; +#endif + bool firstSolveRowsPrepared = false; + for (uint64_t block = 0; block < pipelineBlockCount; ++block) { + if (block < scoreBlockCount) { + uint64_t rowBegin = block * scoreBlockSize; + uint64_t rowCount = useFullChunkScore + ? curT - rowBegin + : ScoreRowBlockCount(curT, rowBegin); + uint64_t refToken = ScoreRefToken(start, curT, rowBegin, rowCount); + uint64_t queueSlot = useFullChunkScore + ? activeSolveSlot_ / KDA_SCORE_LANES + : block % KDA_SCORE_QUEUE_DEPTH; + uint64_t scoreSlot = ScoreScratchSlot(queueSlot, scoreLane, pairHeads); + PrepareGateProducts(b, h, hv, start, curT, subBlockIdx, subBlockNum, true, refToken, + rowBegin + rowCount, true, scoreSlot, + rowBegin, rowCount); + } + JoinAivMte3(); + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_MTE3>(scoreReadyFlag_); + if (block > 0) { + if constexpr (SAFE_GATE) { + if (waitPendingSafeSolve) { + WaitAicSolveDone(); + waitPendingSafeSolve = false; + } + } + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(scoreDoneFlag_); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128) { + if (useDirectScoreUb && useAkkCubeSolve && block == 1) { + ProcessDirectScoreSolveRowsArch35( + b, hv, chunkIdx, start, 0, 0); + firstSolveRowsPrepared = true; + } else if (pairHeads && useAkkCubeSolve && block == 1) { + uint64_t firstRowBegin = 0; + uint64_t firstRowEnd = 0; + GetSolveRowRange( + BT_, 0, KDA_SCORE_LANES, firstRowBegin, firstRowEnd); + PrepareAqkAkkSolveInputRows( + b, hv, chunkIdx, start, curT, + firstRowBegin, firstRowEnd, false, false); + firstSolveRowsPrepared = true; + } + } +#endif + } + } + bool fusedScoreWriteback = false; +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + fusedScoreWriteback = SAFE_GATE && BT_ == 64 && K_ == 128 && V_ == 128; +#endif + if (!fusedScoreWriteback) { + // The final score MMAD only consumes scoreWorkspace_. Run the + // independent gate writeback while AIC drains its MMAD/Fixpipe path. + PrepareGateProducts(b, h, hv, start, curT, subBlockIdx, subBlockNum); + } + if (pipelineBlockCount > 0) { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(scoreDoneFlag_); + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if (useDirectScoreUb && useAkkCubeSolve) { + ProcessDirectScoreSolveRowsArch35( + b, hv, chunkIdx, start, KDA_DIRECT_SCORE_ROWS, 1); + } +#endif + if constexpr (SAFE_GATE) { + if (waitPendingSafeSolve) { + WaitAicSolveDone(); + waitPendingSafeSolve = false; + } + } + if (useAkkCubeSolve) { + bool fullChunk = curT == BT_; + if constexpr (SAFE_GATE) { + if (pairHeads) { + if (!useDirectScoreUb) { + uint64_t firstRowPart = firstSolveRowsPrepared ? 1 : 0; + for (uint64_t rowPart = firstRowPart; + rowPart < KDA_SCORE_LANES; ++rowPart) { + uint64_t pairRowBegin = 0; + uint64_t pairRowEnd = 0; + GetSolveRowRange( + BT_, rowPart, KDA_SCORE_LANES, pairRowBegin, pairRowEnd); + PrepareAqkAkkSolveInputRows( + b, hv, chunkIdx, start, curT, + pairRowBegin, pairRowEnd, false, false); + } + } + } else { + PrepareAqkAkkSolveInputRows( + b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd, false, false); + } + if (deferSafeSolve) { + SignalAicSolveReady(); + return; + } + RunAicAfterBothAivReady(subBlockIdx, subBlockNum); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (!(SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128)) { + StoreSolveXRowsToAkk(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd); + } +#else + StoreSolveXRowsToAkk(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd); +#endif + } else { + PrepareAqkAkkSolveInputRows(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd, + fullChunk, false); + if (!fullChunk) { + RunAicAfterBothAivReady(subBlockIdx, subBlockNum); + StoreSolveXRowsToAkk(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd); + } else { + uint32_t solveIters = KDA_SOLVE_DIAG_MCH_ITERS; + RunAicAfterBothAivReady(subBlockIdx, subBlockNum); + for (uint32_t iter = 0; iter < solveIters; ++iter) { + AddSolveTmpToXDiagRows(b, hv, chunkIdx, start, solveRowBegin, solveRowEnd, + iter + 1 == solveIters); + RunAicAfterBothAivReady(subBlockIdx, subBlockNum); + } + } + } + } + // Host validation guarantees every accepted shape has enough workspace for this cube path. +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128) { + if (HV_ % KDA_SCORE_LANES != 0) { + PrepareWuCubeInputs(b, hv, start, curT, subBlockIdx, subBlockNum); + } + } else { + PrepareWuCubeInputs(b, hv, start, curT, subBlockIdx, subBlockNum); + } +#else + PrepareWuCubeInputs(b, hv, start, curT, subBlockIdx, subBlockNum); +#endif + FinalizePrepareIntermediates(b, hv, chunkIdx, start, curT, subBlockIdx, subBlockNum); + } + + __aicore__ inline void FinishDeferredSafeChunk(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t end, uint64_t subBlockIdx, + uint64_t subBlockNum) + { + uint64_t curT = end - start; + uint64_t solveRowBegin = 0; + uint64_t solveRowEnd = 0; + GetSolveRowRange(BT_, subBlockIdx, subBlockNum, solveRowBegin, solveRowEnd); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (!(SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128)) { + StoreSolveXRowsToAkk(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd); + } +#else + StoreSolveXRowsToAkk(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd); +#endif +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128) { + if (HV_ % KDA_SCORE_LANES != 0) { + PrepareWuCubeInputs(b, hv, start, curT, subBlockIdx, subBlockNum); + } + } else { + PrepareWuCubeInputs(b, hv, start, curT, subBlockIdx, subBlockNum); + } +#else + PrepareWuCubeInputs(b, hv, start, curT, subBlockIdx, subBlockNum); +#endif + FinalizePrepareIntermediates(b, hv, chunkIdx, start, curT, subBlockIdx, subBlockNum); + } + + __aicore__ inline void FinishDeferredSafeChunkPair( + uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, uint64_t end) + { + FinishDeferredSafeChunk(b, hv, chunkIdx, start, end, 0, 1); + } + + __aicore__ inline void ProcessChunkPreAic(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + if constexpr (IsSameType::value) { + ProcessChunkPreAicFp32(b, hv, chunkIdx, start, end); + } + } + + __aicore__ inline void ProcessChunkPreAicFp32(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + uint64_t curT = end - start; + if (curT == 0 || K_ < 16) { + return; + } + uint64_t scoreBlockSize = ScoreRefBlockSize(); + uint64_t scoreBlockCount = (curT + scoreBlockSize - 1) / scoreBlockSize; + uint64_t pipelineBlockCount = + (scoreBlockCount + KDA_SCORE_QUEUE_DEPTH - 1) / KDA_SCORE_QUEUE_DEPTH * KDA_SCORE_QUEUE_DEPTH; + for (uint64_t block = 0; block < pipelineBlockCount; ++block) { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(scoreReadyFlag_); + if (block < scoreBlockCount) { + uint64_t rowBegin = block * scoreBlockSize; + uint64_t rowCount = ScoreRowBlockCount(curT, rowBegin); + ComputeRawAqkAkkCubeBlock(b, hv, chunkIdx, start, curT, rowBegin, rowCount, true, + block % KDA_SCORE_QUEUE_DEPTH, rowBegin + rowCount); + } + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(scoreDoneFlag_); + } + bool usePostWuCube = UsePostWuCube(curT); + bool useAkkCubeSolve = UseAkkCubeSolve(curT); + if (useAkkCubeSolve) { + if constexpr (SAFE_GATE) { + Catlass::Arch::CrossCoreWaitFlag(syncReadyFlag_); + ComputeAkkMergeCubeWorkspaceDispatch(b, hv, chunkIdx); + Catlass::Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(syncDoneFlag_); + } else { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(mchSyncReadyFlag_); + if (curT == BT_) { + ComputeAkkInverseMchFull(b, hv, chunkIdx, start); + } else { + ComputeAkkMergeCubeWorkspace(b, hv, chunkIdx); + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(mchSyncDoneFlag_); + } + } + } + (void)usePostWuCube; + (void)chunkIdx; + } + + __aicore__ inline void ProcessChunkPreAicHeadPairFp32( + uint64_t b, uint64_t hvBase, uint64_t chunkIdx, uint64_t start, uint64_t end, + uint64_t localTaskIdx) + { + uint64_t curT = end - start; + if (curT == 0 || K_ < 16) { + return; + } + const bool useFullChunkScore = false; + uint64_t scoreBlockSize = useFullChunkScore ? curT : ScoreRefBlockSize(); + uint64_t scoreBlockCount = (curT + scoreBlockSize - 1) / scoreBlockSize; + uint64_t pipelineBlockCount = useFullChunkScore + ? scoreBlockCount + : (scoreBlockCount + KDA_SCORE_QUEUE_DEPTH - 1) / KDA_SCORE_QUEUE_DEPTH * + KDA_SCORE_QUEUE_DEPTH; + for (uint64_t block = 0; block < pipelineBlockCount; ++block) { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(scoreReadyFlag_); + if (block < scoreBlockCount) { + uint64_t rowBegin = block * scoreBlockSize; + uint64_t rowCount = useFullChunkScore + ? curT - rowBegin + : ScoreRowBlockCount(curT, rowBegin); + uint64_t queueSlot = useFullChunkScore + ? localTaskIdx % KDA_SCORE_QUEUE_DEPTH + : block % KDA_SCORE_QUEUE_DEPTH; +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + bool directScoreDispatched = false; + if constexpr (SAFE_GATE) { + if (KDA_ARCH35_ENABLE_DIRECT_SCORE_UB && curT == 64 && rowCount == 32) { + uint64_t scoreSlotBase = ScoreScratchSlot(queueSlot, 0, true); + if (rowBegin == 0) { + ComputeRawAqkAkkCubeStableHeadPairDirectUbArch35<32>( + rowBegin, scoreSlotBase, static_cast(block)); + } else { + ComputeRawAqkAkkCubeStableHeadPairDirectUbArch35<64>( + rowBegin, scoreSlotBase, static_cast(block)); + } + directScoreDispatched = true; + } + } + if (!directScoreDispatched) +#endif + { + for (uint64_t lane = 0; lane < KDA_SCORE_LANES; ++lane) { + uint64_t hv = hvBase + lane; + uint64_t scoreSlot = ScoreScratchSlot(queueSlot, lane, true); + activeSolveSlot_ = + (localTaskIdx % (KDA_SOLVE_PIPELINE_DEPTH / KDA_SCORE_LANES)) * + KDA_SCORE_LANES + lane; + ComputeRawAqkAkkCubeBlock( + b, hv, chunkIdx, start, curT, rowBegin, rowCount, true, + scoreSlot, rowBegin + rowCount); + } + } + } + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(scoreDoneFlag_); + } + + if (UseAkkCubeSolve(curT)) { + Catlass::Arch::CrossCoreWaitFlag(syncReadyFlag_); + for (uint64_t lane = 0; lane < KDA_SCORE_LANES; ++lane) { + activeSolveSlot_ = + (localTaskIdx % (KDA_SOLVE_PIPELINE_DEPTH / KDA_SCORE_LANES)) * KDA_SCORE_LANES + lane; + ComputeAkkMergeCubeWorkspaceDispatch(b, hvBase + lane, chunkIdx); + } + Catlass::Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(syncDoneFlag_); + } + } + + __aicore__ inline bool ResolveFlatChunkForHv( + uint64_t flatChunk, uint64_t hv, uint64_t &seq, uint64_t &b, uint64_t &h, + uint64_t &chunkIdx, uint64_t &start, uint64_t &end) + { + if (!isVarLen_) { + seq = flatChunk / NT_; + b = seq; + chunkIdx = flatChunk % NT_; + start = chunkIdx * BT_; + end = start + BT_; + if (end > T_) { + end = T_; + } + } else { + if (!KdaVarlen::ResolveChunkRange( + cuSeqlensAddr_, chunkIndicesAddr_, N_, T_, BT_, flatChunk, + seq, start, end)) { + return false; + } + b = 0; + chunkIdx = flatChunk; + } + h = hv / (HV_ / H_); + return start < end; + } + + __aicore__ inline void ProcessPreAivHeadPair() + { + const uint64_t subBlockIdx = static_cast(GetSubBlockIdx()); + const uint64_t subBlockNum = static_cast(GetSubBlockNum()); + const uint64_t coreNum = usedCoreNum_; + const uint64_t coreIdx = static_cast(GetBlockIdx()) / subBlockNum; + const uint64_t chunkCount = isVarLen_ ? NT_ : B_ * NT_; + const uint64_t headWindows = HV_ / KDA_SCORE_LANES; + const uint64_t taskNum = chunkCount * headWindows; + bool pendingValid = false; + uint64_t pendingB = 0; + uint64_t pendingHv = 0; + uint64_t pendingChunkIdx = 0; + uint64_t pendingStart = 0; + uint64_t pendingEnd = 0; + uint64_t pendingSlot = 0; + uint64_t localTaskIdx = 0; + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE && COMPILE_BT == 64 && COMPILE_K == 128 && COMPILE_V == 128) { + if (KDA_ARCH35_ENABLE_DIRECT_SCORE_UB) { + InitializeDirectScoreUbArch35(); + } + } +#endif + + for (uint64_t task = coreIdx; task < taskNum; task += coreNum, ++localTaskIdx) { + uint64_t flatChunk = task / headWindows; + uint64_t hv = (task % headWindows) * KDA_SCORE_LANES + subBlockIdx; + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (!ResolveFlatChunkForHv(flatChunk, hv, seq, b, h, chunkIdx, start, end)) { + continue; + } + (void)seq; + uint64_t currentSlot = + (localTaskIdx % (KDA_SOLVE_PIPELINE_DEPTH / KDA_SCORE_LANES)) * + KDA_SCORE_LANES + + subBlockIdx; + activeSolveSlot_ = currentSlot; + bool deferSolve = UseAkkCubeSolve(end - start); + ProcessChunkPreAivFp32( + b, h, hv, chunkIdx, start, end, 0, 1, deferSolve, pendingValid, subBlockIdx, true); + if (pendingValid) { + activeSolveSlot_ = pendingSlot; + FinishDeferredSafeChunkPair( + pendingB, pendingHv, pendingChunkIdx, pendingStart, pendingEnd); + if (fusePostWu_) { + SignalPostWuReady(); + } + } + pendingValid = deferSolve; + if (pendingValid) { + pendingB = b; + pendingHv = hv; + pendingChunkIdx = chunkIdx; + pendingStart = start; + pendingEnd = end; + pendingSlot = currentSlot; + } + } + if (pendingValid) { + WaitAicSolveDone(); + activeSolveSlot_ = pendingSlot; + FinishDeferredSafeChunkPair( + pendingB, pendingHv, pendingChunkIdx, pendingStart, pendingEnd); + if (fusePostWu_) { + SignalPostWuReady(); + } + } + } + + __aicore__ inline void ProcessPreAicHeadPair() + { + const uint64_t chunkCount = isVarLen_ ? NT_ : B_ * NT_; + const uint64_t headWindows = HV_ / KDA_SCORE_LANES; + const uint64_t taskNum = chunkCount * headWindows; + const uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + uint64_t localTaskIdx = 0; + for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum, ++localTaskIdx) { + uint64_t flatChunk = task / headWindows; + uint64_t hvBase = (task % headWindows) * KDA_SCORE_LANES; + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunkForHv(flatChunk, hvBase, seq, b, h, chunkIdx, start, end)) { + (void)seq; + (void)h; + ProcessChunkPreAicHeadPairFp32(b, hvBase, chunkIdx, start, end, localTaskIdx); + } + } + } + + template + __aicore__ inline void ProcessPreAicHeadPairFused(PostWuOp &postWu) + { + const uint64_t chunkCount = isVarLen_ ? NT_ : B_ * NT_; + const uint64_t headWindows = HV_ / KDA_SCORE_LANES; + const uint64_t taskNum = chunkCount * headWindows; + const uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + uint64_t batchB[KDA_POST_QUEUE_STORAGE]; + uint64_t batchHvBase[KDA_POST_QUEUE_STORAGE]; + uint64_t batchStart[KDA_POST_QUEUE_STORAGE]; + uint64_t batchEnd[KDA_POST_QUEUE_STORAGE]; + uint16_t batchCount = 0; + uint64_t localTaskIdx = 0; + + for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum, ++localTaskIdx) { + uint64_t flatChunk = task / headWindows; + uint64_t hvBase = (task % headWindows) * KDA_SCORE_LANES; + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (!ResolveFlatChunkForHv(flatChunk, hvBase, seq, b, h, chunkIdx, start, end)) { + continue; + } + (void)seq; + (void)h; + ProcessChunkPreAicHeadPairFp32(b, hvBase, chunkIdx, start, end, localTaskIdx); + batchB[batchCount] = b; + batchHvBase[batchCount] = hvBase; + batchStart[batchCount] = start; + batchEnd[batchCount] = end; + ++batchCount; + + if (batchCount == KDA_POST_QUEUE_STORAGE) { + for (uint16_t i = 0; i < KDA_POST_QUEUE_DEPTH; ++i) { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(postWuReadyFlag_); + } + postWu.ProcessPreparedHeadPairBatchArch35( + batchB, batchHvBase, batchStart, batchEnd, + KDA_POST_QUEUE_DEPTH); + batchB[0] = batchB[KDA_POST_QUEUE_DEPTH]; + batchHvBase[0] = batchHvBase[KDA_POST_QUEUE_DEPTH]; + batchStart[0] = batchStart[KDA_POST_QUEUE_DEPTH]; + batchEnd[0] = batchEnd[KDA_POST_QUEUE_DEPTH]; + batchCount = 1; + } + } + + if (batchCount > 0) { + for (uint16_t i = 0; i < batchCount; ++i) { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(postWuReadyFlag_); + } + postWu.ProcessPreparedHeadPairBatchArch35( + batchB, batchHvBase, batchStart, batchEnd, + batchCount); + } + } + + __aicore__ inline void ProcessPreAiv() + { + if constexpr (IsSameType::value) { + isAivOnly_ = true; + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE && !IsSameType::value) { + if (headPairMode_ && !isAivOnly_) { + ProcessPreAivHeadPair(); + return; + } + } +#endif + uint64_t subBlockNum = isAivOnly_ ? 1 : static_cast(GetSubBlockNum()); + if (subBlockNum == 0) { + return; + } + uint64_t subBlockIdx = isAivOnly_ ? 0 : static_cast(GetSubBlockIdx()); + uint64_t coreNum = isAivOnly_ ? static_cast(GetBlockNum()) : usedCoreNum_; + uint64_t coreIdx = isAivOnly_ ? static_cast(GetBlockIdx()) : + static_cast(GetBlockIdx()) / subBlockNum; + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + if constexpr (SAFE_GATE && !IsSameType::value) { + bool pendingValid = false; + uint64_t pendingB = 0; + uint64_t pendingHv = 0; + uint64_t pendingChunkIdx = 0; + uint64_t pendingStart = 0; + uint64_t pendingEnd = 0; + uint64_t pendingSlot = 0; + uint64_t localTaskIdx = 0; + for (uint64_t task = coreIdx; task < taskNum; task += coreNum, ++localTaskIdx) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (!ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + continue; + } + (void)seq; + uint64_t currentSlot = localTaskIdx % KDA_SOLVE_PIPELINE_DEPTH; + activeSolveSlot_ = currentSlot; + bool deferSolve = UseAkkCubeSolve(end - start); + if (!deferSolve && pendingValid) { + WaitAicSolveDone(); + activeSolveSlot_ = pendingSlot; + FinishDeferredSafeChunk(pendingB, pendingHv, pendingChunkIdx, pendingStart, pendingEnd, + subBlockIdx, subBlockNum); + pendingValid = false; + activeSolveSlot_ = currentSlot; + } + ProcessChunkPreAivFp32(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum, + deferSolve, pendingValid); + if (pendingValid) { + activeSolveSlot_ = pendingSlot; + FinishDeferredSafeChunk(pendingB, pendingHv, pendingChunkIdx, pendingStart, pendingEnd, + subBlockIdx, subBlockNum); + } + pendingValid = deferSolve; + if (pendingValid) { + pendingB = b; + pendingHv = hv; + pendingChunkIdx = chunkIdx; + pendingStart = start; + pendingEnd = end; + pendingSlot = currentSlot; + } + } + if (pendingValid) { + WaitAicSolveDone(); + activeSolveSlot_ = pendingSlot; + FinishDeferredSafeChunk(pendingB, pendingHv, pendingChunkIdx, pendingStart, pendingEnd, + subBlockIdx, subBlockNum); + } + return; + } + for (uint64_t task = coreIdx; task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + ProcessChunkPreAiv(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); + } + } + } + + __aicore__ inline void ProcessPreAic() + { + if constexpr (IsSameType::value) { + return; + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (SAFE_GATE) { + if (headPairMode_) { + ProcessPreAicHeadPair(); + return; + } + } +#endif + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + uint64_t localTaskIdx = 0; + for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum, ++localTaskIdx) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + if constexpr (SAFE_GATE) { + activeSolveSlot_ = localTaskIdx % KDA_SOLVE_PIPELINE_DEPTH; + } + (void)seq; + (void)h; + ProcessChunkPreAic(b, hv, chunkIdx, start, end); + } + } + } + + +private: + GlobalTensor q_; + GlobalTensor k_; + GlobalTensor v_; + GlobalTensor gk_; + GlobalTensor rawG_; + GlobalTensor aLog_; + GlobalTensor dtBias_; + GlobalTensor beta_; + GlobalTensor initialState_; + GlobalTensor o_; + GlobalTensor finalState_; + GlobalTensor aqk_; + GlobalTensor akk_; + GlobalTensor w_; + GlobalTensor u_; + GlobalTensor qg_; + GlobalTensor kg_; + GlobalTensor vNew_; + GlobalTensor h_; + GlobalTensor finalKg_; + GlobalTensor preparedQG_; + GlobalTensor preparedAqk_; + GlobalTensor propagatedVNew_; + GlobalTensor propagatedH_; + GlobalTensor solveWorkspace_; + GlobalTensor scoreWorkspace_; + TPipe *pipe_ = nullptr; + TBuf exp2Buf_; + TBuf vecBuf_; + TBuf gateWritebackBuf_; + TEventID mte2ToVEvent_ = 0; + TEventID vToMte2Event_ = 0; + TEventID vToMte3Event_ = 0; + TEventID mte3ToVEvent_ = 0; + TEventID mte2ToMte3Event_ = 0; + TEventID vToSEvent_ = 0; + TEventID mte3ToMte2Events_[KDA_GATE_PIPELINE_DEPTH] = {0, 0, 0}; + bool vectorEventsAllocated_ = false; + Catlass::Arch::CrossCoreFlagWithReverse scoreReadyFlag_{KDA_SCORE_READY_FLAG0, + KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse scoreDoneFlag_{KDA_SCORE_DONE_FLAG0, + KDA_SCORE_DONE_FLAG1}; + // Solve has one outstanding task per core. Reuse the primary score IDs as an ordered token stream; + // score credits remain on the reverse IDs, so no additional hardware flag IDs are consumed. + Catlass::Arch::CrossCoreFlag syncReadyFlag_{KDA_SOLVE_READY_FLAG}; + Catlass::Arch::CrossCoreFlag syncDoneFlag_{KDA_SOLVE_DONE_FLAG}; + Catlass::Arch::CrossCoreFlagWithReverse postWuReadyFlag_{ + KDA_POST_READY_FLAG, KDA_POST_FREE_FLAG}; + Catlass::Arch::CrossCoreFlagWithReverse mchSyncReadyFlag_{ + KDA_SCORE_READY_FLAG0, KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse mchSyncDoneFlag_{ + KDA_SCORE_DONE_FLAG0, KDA_SCORE_DONE_FLAG1}; + uint64_t B_ = 0; + uint64_t N_ = 0; + uint64_t H_ = 0; + uint64_t HV_ = 0; + uint64_t T_ = 0; + uint64_t K_ = 0; + uint64_t V_ = 0; + uint64_t BT_ = 0; + uint64_t NT_ = 0; + float scale_ = 1.0f; + bool hasInitial_ = false; + bool isVarLen_ = false; + bool isAivOnly_ = false; + bool headPairMode_ = false; + bool inputSequenceMajor_ = false; + bool fusePostWu_ = false; + bool materializeFinalKg_ = false; + bool computeGateInPrepare_ = false; + bool hasALog_ = false; + bool hasDtBias_ = false; + bool storeQG_ = true; + float lowerBound_ = -5.0f; + uint64_t usedCoreNum_ = 1; + uint64_t solveCoreIdx_ = 0; + uint64_t activeSolveSlot_ = 0; + uint64_t activeGateChunkStart_ = 0; + __gm__ int64_t *chunkIndicesAddr_ = nullptr; + __gm__ int64_t *cuSeqlensAddr_ = nullptr; +}; +} // namespace + +template +__aicore__ inline void RunChunkKdaPrepare( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR rawG, GM_ADDR aLog, + GM_ADDR dtBias, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR aqk, GM_ADDR akk, GM_ADDR qg, + GM_ADDR qgScaled, GM_ADDR wSeed, GM_ADDR uSeed, GM_ADDR finalKg, GM_ADDR userWorkspace, + const TilingData &tiling, TPipe &pipe, bool storeQG = true) +{ + GM_ADDR aqkFp32 = userWorkspace + tiling.prepareAqkFp32Offset; + GM_ADDR akkFp32 = userWorkspace + tiling.prepareAkkFp32Offset; + GM_ADDR prepareScratch = userWorkspace + tiling.prepareScratchOffset; + + if ASCEND_IS_AIC { + ChunkKdaFwdPrepareKernel op; + op.Init(q, k, v, gk, rawG, aLog, dtBias, beta, initialState, cuSeqlens, chunkIndices, + nullptr, nullptr, nullptr, nullptr, aqk, userWorkspace, aqkFp32, akkFp32, + wSeed, akk, qg, qgScaled, uSeed, userWorkspace, finalKg, + prepareScratch, tiling, &pipe, false, storeQG); + if (tiling.fusePostWu) { +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + KdaPostWu::ChunkKdaFwdPostWuKernel postWu; + postWu.Init(nullptr, k, nullptr, gk, beta, initialState, cuSeqlens, chunkIndices, + wSeed, akk, uSeed, nullptr, userWorkspace, userWorkspace, userWorkspace, + akk, wSeed, uSeed, userWorkspace, finalKg, userWorkspace, + prepareScratch, prepareScratch, tiling, &pipe, false); + op.ProcessAicFused(postWu); +#else + op.ProcessAic(); +#endif + } else { + op.ProcessAic(); + } + } + if ASCEND_IS_AIV { + ChunkKdaFwdPrepareKernel op; + op.Init(q, k, v, gk, rawG, aLog, dtBias, beta, initialState, cuSeqlens, chunkIndices, + nullptr, nullptr, nullptr, nullptr, aqk, userWorkspace, aqkFp32, akkFp32, + wSeed, akk, qg, qgScaled, uSeed, userWorkspace, finalKg, + prepareScratch, tiling, &pipe, true, storeQG); + op.ProcessAiv(); + } +} + +} // namespace KdaPrepare diff --git a/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd.cpp b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd.cpp index 05ff5f4e607f..33987a5fb24e 100644 --- a/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd.cpp +++ b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd.cpp @@ -1,2584 +1,236 @@ -/** - * Copyright (c) 2026 Tianjin University, Ltd. - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of - * the BSD 3-Clause License (the "License"). - * Please refer to the License for details. You may not use this file except in compliance with the License. - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. - */ - #include "kernel_operator.h" - -#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 -#define CATLASS_ARCH 3510 -#include "catlass/arch/arch.hpp" -#include "catlass/catlass.hpp" -#include "catlass/gemm/block/block_mmad.hpp" -#include "catlass/gemm/dispatch_policy.hpp" -#include "catlass/gemm/gemm_type.hpp" -#include "catlass/gemm_coord.hpp" -#include "catlass/layout/layout.hpp" -#include "catlass/arch/cross_core_sync.hpp" -#include "tla/layout.hpp" -#include "tla/tensor.hpp" -using _128 = tla::Int<128>; -#else -#define CATLASS_ARCH 2201 -#include "catlass/arch/arch.hpp" -#include "catlass/catlass.hpp" -#include "catlass/gemm/block/block_mmad.hpp" -#include "catlass/gemm/dispatch_policy.hpp" -#include "catlass/gemm/gemm_type.hpp" -#include "catlass/gemm_coord.hpp" -#include "catlass/layout/layout.hpp" -#include "catlass/arch/cross_core_sync.hpp" -#include "tla/layout.hpp" -#include "tla/tensor.hpp" -#endif - -#ifndef TORCH_MODE #include "lib/matmul_intf.h" -#endif - -using namespace AscendC; -using _64 = tla::Int<64>; -namespace { -constexpr float LN2 = 0.69314718055994530942f; -constexpr float KDA_EXP2_CLAMP = 80.0f; -constexpr float KDA_EXP_INPUT_MAX = KDA_EXP2_CLAMP * LN2; -constexpr float KDA_EXP_INPUT_MIN = -KDA_EXP2_CLAMP * LN2; -constexpr float KDA_FP16_MAX = 65504.0f; -constexpr uint32_t EXP2_UB_ELEMENTS = 256; -constexpr uint32_t EXP2_EVENT_ID = 0; -constexpr uint32_t KDA_MTE2_V_EVENT_ID = 1; -constexpr uint32_t KDA_SCALAR_V_MTE3_EVENT_ID = 4; -constexpr uint32_t KDA_SCALAR_MTE3_V_EVENT_ID = 5; -constexpr uint32_t KDA_MTE2_MTE3_EVENT_ID = 6; -constexpr uint32_t KDA_MTE3_MTE2_EVENT_ID = 7; -constexpr uint32_t KDA_VEC_BUFFER_NUM = 2; -constexpr uint32_t KDA_SOLVE_BT = 64; -constexpr uint32_t KDA_SOLVE_MATRIX_ELEMENTS = KDA_SOLVE_BT * KDA_SOLVE_BT; -constexpr uint32_t KDA_SOLVE_SCRATCH_X = 0; -constexpr uint32_t KDA_SOLVE_SCRATCH_Y0 = 1; -constexpr uint32_t KDA_SOLVE_SCRATCH_TMP = 2; -constexpr uint32_t KDA_SOLVE_SCRATCH_Y1 = 3; -constexpr uint32_t KDA_SOLVE_SCRATCH_IDENTITY = 4; -constexpr uint32_t KDA_SOLVE_SCRATCH_SLOTS = 5; -constexpr uint32_t KDA_SOLVE_DIAG_BT = 16; -constexpr uint32_t KDA_SOLVE_DIAG_BLOCKS = KDA_SOLVE_BT / KDA_SOLVE_DIAG_BT; -constexpr uint32_t KDA_SOLVE_DIAG_MCH_ITERS = 3; -constexpr uint32_t KDA_SCORE_REF_BC = 16; -constexpr uint32_t KDA_VEC_ARENA_ELEMENTS = 32768; -constexpr uint32_t KDA_BITS_PER_MASK_BYTE = 8; -constexpr uint32_t KDA_SELECT_COL_BLOCKS = 2; -constexpr uint32_t KDA_SELECT_COL_MASK_BYTES = KDA_SOLVE_MATRIX_ELEMENTS / KDA_BITS_PER_MASK_BYTE; -constexpr uint32_t KDA_SELECT_MASK_BYTES = KDA_SELECT_COL_BLOCKS * KDA_SELECT_COL_MASK_BYTES; -constexpr uint32_t KDA_SELECT_AQK_MASK_BYTE_OFFSET = 120 * 1024; -constexpr uint32_t KDA_SELECT_AKK_MASK_BYTE_OFFSET = KDA_SELECT_AQK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; -constexpr uint32_t KDA_SELECT_ZERO_BYTE_OFFSET = KDA_SELECT_AKK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; -constexpr uint32_t KDA_SELECT_ZERO_FLOAT_OFFSET = KDA_SELECT_ZERO_BYTE_OFFSET / sizeof(float); -constexpr uint8_t KDA_SCORE_DONE_FLAG0 = 2; -constexpr uint8_t KDA_SCORE_DONE_FLAG1 = 3; -constexpr uint8_t KDA_SCORE_READY_FLAG0 = 4; -constexpr uint8_t KDA_SCORE_READY_FLAG1 = 5; -constexpr uint32_t KDA_SCORE_QUEUE_DEPTH = 2; -constexpr uint32_t KDA_SYNC_REVERSE_DEPTH = 1; -constexpr uint32_t KDA_SCORE_SCRATCH_PLANES = 3; -constexpr uint32_t KDA_SCORE_SCRATCH_QG = 0; -constexpr uint32_t KDA_SCORE_SCRATCH_W = 1; -constexpr uint32_t KDA_SCORE_SCRATCH_KG = 2; -constexpr uint64_t KDA_WORKSPACE_ALIGN = 512; - -#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 -using KdaArchTag = Catlass::Arch::Ascend950; +#include "chunk_kda_fwd_common.h" +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 && \ + (!defined(TILING_KEY_VAR) || TILING_KEY_VAR == 2UL) +#define KDA_COMPILE_ARCH35_FAST_PATH 1 +#include "arch35/chunk_kda_fwd_impl.h" #else -using KdaArchTag = Catlass::Arch::AtlasA2; +#define KDA_COMPILE_ARCH35_FAST_PATH 0 #endif -using KdaDispatchPolicy = Catlass::Gemm::MmadPingpong; -using KdaSolveDispatchPolicy = Catlass::Gemm::MmadPingpong; -static_assert(!KdaSolveDispatchPolicy::USE_HF32_MODE, "KDA triangular solve must use IEEE FP32 Cube mode"); -using KdaL1TileShape = tla::Shape<_64, _128, _128>; -using KdaL0TileShape = KdaL1TileShape; -using KdaSolveL1TileShape = tla::Shape<_64, _64, _64>; -using KdaSolveL0TileShape = KdaSolveL1TileShape; -__aicore__ inline uint32_t FloatToBits(float value) +namespace KdaForward { + +constexpr int64_t KDA_STAGE_FULL = -1; +constexpr int64_t KDA_STAGE_GATE_PREPARE = 0; +constexpr int64_t KDA_STAGE_POST_WU = 1; +constexpr int64_t KDA_STAGE_FWD_H = 2; +constexpr int64_t KDA_STAGE_FINALIZE = 3; + +template +__aicore__ inline void DispatchStage( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR g, GM_ADDR beta, + GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR attnOut, + GM_ADDR finalState, GM_ADDR gk, GM_ADDR aqk, GM_ADDR akk, + GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR qgScaled, GM_ADDR uSeed, GM_ADDR userWorkspace, + const TilingData &tiling) { - union Bits { - __aicore__ Bits() {} - float f; - uint32_t u; - } bits; - bits.f = value; - return bits.u; + auto addresses = ResolveAddresses( + finalState, gk, w, u, qg, kg, vNew, h, userWorkspace, tiling); + addresses.qgScaled = qgScaled; + if (tiling.stage == KDA_STAGE_GATE_PREPARE) { + RunGateCumsum(g, aLog, dtBias, cuSeqlens, addresses.gk, tiling); + if (!tiling.computeGateInPrepare) { + SyncAll(); + } + TPipe pipe; + KdaPrepare::RunChunkKdaPrepare( + q, k, v, addresses.gk, g, aLog, dtBias, beta, initialState, + cuSeqlens, chunkIndices, aqk, akk, addresses.qg, + addresses.qgScaled, addresses.w, uSeed, addresses.kg, + userWorkspace, tiling, pipe, tiling.storeQG); + } else if (tiling.stage == KDA_STAGE_POST_WU) { + TPipe pipe; + KdaPostWu::RunChunkKdaPostWu( + q, k, v, addresses.gk, beta, initialState, cuSeqlens, + chunkIndices, addresses.w, akk, uSeed, addresses.w, + addresses.u, addresses.kg, addresses.vNew, userWorkspace, + tiling, pipe); + } else if (tiling.stage == KDA_STAGE_FWD_H) { + if (tiling.vHeadDim > 128) { + RunFwdH( + initialState, cuSeqlens, chunkIndices, addresses, + userWorkspace, tiling); + } else { + RunFwdH( + initialState, cuSeqlens, chunkIndices, addresses, + userWorkspace, tiling); + } + } else if (tiling.stage == KDA_STAGE_FINALIZE) { + TPipe pipe; + KdaFinalize::RunChunkKdaOutput( + q, k, v, addresses.gk, beta, initialState, cuSeqlens, + chunkIndices, addresses.qgScaled, aqk, addresses.vNew, + addresses.h, attnOut, userWorkspace, tiling, pipe); + } } -__aicore__ inline float BitsToFloat(uint32_t value) +template +__aicore__ inline void DispatchStageSafeGate( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR g, GM_ADDR beta, + GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR attnOut, + GM_ADDR finalState, GM_ADDR gk, GM_ADDR aqk, GM_ADDR akk, + GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR qgScaled, GM_ADDR uSeed, GM_ADDR userWorkspace, + const TilingData &tiling) { - union Bits { - __aicore__ Bits() {} - uint32_t u; - float f; - } bits; - bits.u = value; - return bits.f; + if (tiling.safeGate) { + DispatchStage( + q, k, v, g, beta, aLog, dtBias, initialState, cuSeqlens, + chunkIndices, attnOut, finalState, gk, aqk, akk, w, u, qg, + kg, vNew, h, qgScaled, uSeed, userWorkspace, tiling); + } else { + DispatchStage( + q, k, v, g, beta, aLog, dtBias, initialState, cuSeqlens, + chunkIndices, attnOut, finalState, gk, aqk, akk, w, u, qg, + kg, vNew, h, qgScaled, uSeed, userWorkspace, tiling); + } } -__aicore__ inline uint16_t Bf16ToBits(bfloat16_t value) +template +__aicore__ inline void DispatchGeneric( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR g, GM_ADDR beta, + GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR attnOut, + GM_ADDR finalState, GM_ADDR gk, GM_ADDR aqk, GM_ADDR akk, + GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR userWorkspace, const TilingData &tiling) { - union Bits { - __aicore__ Bits() {} - bfloat16_t f; - uint16_t u; - } bits; - bits.f = value; - return bits.u; + RunGeneric( + q, k, v, g, beta, aLog, dtBias, initialState, cuSeqlens, + chunkIndices, attnOut, finalState, gk, aqk, akk, w, u, qg, kg, + vNew, h, userWorkspace, tiling); } -__aicore__ inline bfloat16_t BitsToBf16(uint16_t value) +#if KDA_COMPILE_ARCH35_FAST_PATH +template +__aicore__ inline void DispatchArch35SafeGate( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR g, GM_ADDR beta, + GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR attnOut, + GM_ADDR finalState, GM_ADDR gk, GM_ADDR aqk, GM_ADDR akk, + GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR userWorkspace, const TilingData &tiling) { - union Bits { - __aicore__ Bits() {} - uint16_t u; - bfloat16_t f; - } bits; - bits.u = value; - return bits.f; + AscendC::TPipe pipe; + if (tiling.safeGate) { + arch35::Run( + q, k, v, g, beta, aLog, dtBias, initialState, cuSeqlens, + chunkIndices, attnOut, finalState, gk, aqk, akk, w, u, qg, + kg, vNew, h, userWorkspace, tiling, pipe); + } else { + arch35::Run( + q, k, v, g, beta, aLog, dtBias, initialState, cuSeqlens, + chunkIndices, attnOut, finalState, gk, aqk, akk, w, u, qg, + kg, vNew, h, userWorkspace, tiling, pipe); + } } - -template -__aicore__ inline T FloatToType(float value) +#elif defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +template +__aicore__ inline void DispatchArch35SafeGate( + GM_ADDR, GM_ADDR, GM_ADDR, GM_ADDR, GM_ADDR, + GM_ADDR, GM_ADDR, GM_ADDR, + GM_ADDR, GM_ADDR, + GM_ADDR, GM_ADDR, GM_ADDR, GM_ADDR, GM_ADDR, + GM_ADDR, GM_ADDR, GM_ADDR, GM_ADDR, GM_ADDR, GM_ADDR, + GM_ADDR, + const TilingData &) { - if constexpr (IsSameType::value) { - uint32_t bits = FloatToBits(value); - uint32_t bias = 0x7FFFu + ((bits >> 16) & 1u); - return BitsToBf16(static_cast((bits + bias) >> 16)); - } - return static_cast(value); } +#endif -template -class ChunkKdaFwdKernel { -public: - __aicore__ inline void Init(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, GM_ADDR initialState, - GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR stageQG, GM_ADDR stageAqk, - GM_ADDR stageVNew, GM_ADDR stageH, GM_ADDR o, GM_ADDR finalState, GM_ADDR aqk, - GM_ADDR akk, GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, - GM_ADDR workspace, const ChunkKdaFwdTilingData &tiling, TPipe *pipe, - bool initVecBuffers = true) - { - pipe_ = pipe; - q_.SetGlobalBuffer((__gm__ T *)q); - k_.SetGlobalBuffer((__gm__ T *)k); - v_.SetGlobalBuffer((__gm__ T *)v); - gk_.SetGlobalBuffer((__gm__ float *)gk); - beta_.SetGlobalBuffer((__gm__ float *)beta); - if (initialState != nullptr) { - initialState_.SetGlobalBuffer((__gm__ float *)initialState); - } - if (cuSeqlens != nullptr) { - cuSeqlens_.SetGlobalBuffer((__gm__ int64_t *)cuSeqlens); - } - if (stageQG != nullptr) { - stageQG_.SetGlobalBuffer((__gm__ T *)stageQG); - } - if (stageAqk != nullptr) { - stageAqk_.SetGlobalBuffer((__gm__ T *)stageAqk); - } - if (stageVNew != nullptr) { - stageVNew_.SetGlobalBuffer((__gm__ T *)stageVNew); - } - if (stageH != nullptr) { - stageH_.SetGlobalBuffer((__gm__ T *)stageH); - } - hasChunkIndices_ = chunkIndices != nullptr; - o_.SetGlobalBuffer((__gm__ OUT_T *)o); - finalState_.SetGlobalBuffer((__gm__ float *)finalState); - aqk_.SetGlobalBuffer((__gm__ float *)aqk); - akk_.SetGlobalBuffer((__gm__ AKK_T *)akk); - w_.SetGlobalBuffer((__gm__ T *)w); - u_.SetGlobalBuffer((__gm__ OUT_T *)u); - qg_.SetGlobalBuffer((__gm__ T *)qg); - kg_.SetGlobalBuffer((__gm__ T *)kg); - vNew_.SetGlobalBuffer((__gm__ T *)vNew); - h_.SetGlobalBuffer((__gm__ float *)h); - solveWorkspace_.SetGlobalBuffer((__gm__ float *)workspace); - - B_ = tiling.batch; - N_ = tiling.seqNum; - H_ = tiling.qHeadNum; - HV_ = tiling.vHeadNum; - T_ = tiling.seqlen; - K_ = tiling.kHeadDim; - V_ = tiling.vHeadDim; - BT_ = tiling.chunkSize; - NT_ = tiling.totalChunks; - scale_ = tiling.scale; - hasInitial_ = tiling.hasInitialState; - isVarLen_ = tiling.isVarLen; - usedCoreNum_ = tiling.usedCoreNum; - stage_ = tiling.stage; - if (stage_ == 1) { - const uint64_t solveBytes = usedCoreNum_ * KDA_SOLVE_SCRATCH_SLOTS * BT_ * BT_ * sizeof(float); - const uint64_t alignedSolveBytes = - (solveBytes + KDA_WORKSPACE_ALIGN - 1) / KDA_WORKSPACE_ALIGN * KDA_WORKSPACE_ALIGN; - scoreWorkspace_.SetGlobalBuffer((__gm__ T *)(workspace + alignedSolveBytes)); - } - if (stage_ == 2) { - const uint64_t outputElements = B_ * HV_ * T_ * V_; - o_.SetGlobalBuffer((__gm__ OUT_T *)workspace); - u_.SetGlobalBuffer((__gm__ OUT_T *)workspace + outputElements); - } - if ASCEND_IS_AIV { - uint64_t subBlockNum = static_cast(GetSubBlockNum()); - solveCoreIdx_ = subBlockNum == 0 ? 0 : static_cast(GetBlockIdx()) / subBlockNum; - } else { - solveCoreIdx_ = static_cast(GetBlockIdx()); - } - seqStart_ = tiling.seqStart; - seqEnd_ = tiling.seqEnd; - seqChunkOffset_ = tiling.seqChunkOffset; - - if (pipe_ != nullptr && initVecBuffers) { - pipe_->InitBuffer(exp2Buf_, EXP2_UB_ELEMENTS * sizeof(float)); - pipe_->InitBuffer(vecBuf_, KDA_VEC_ARENA_ELEMENTS * sizeof(float)); - pipe_->InitBuffer(qInQue_, KDA_VEC_BUFFER_NUM, EXP2_UB_ELEMENTS * sizeof(T)); - pipe_->InitBuffer(kInQue_, KDA_VEC_BUFFER_NUM, EXP2_UB_ELEMENTS * sizeof(T)); - pipe_->InitBuffer(gInQue_, KDA_VEC_BUFFER_NUM, EXP2_UB_ELEMENTS * sizeof(float)); - pipe_->InitBuffer(qgOutQue_, KDA_VEC_BUFFER_NUM, EXP2_UB_ELEMENTS * sizeof(T)); - pipe_->InitBuffer(wOutQue_, KDA_VEC_BUFFER_NUM, EXP2_UB_ELEMENTS * sizeof(T)); - pipe_->InitBuffer(kgOutQue_, KDA_VEC_BUFFER_NUM, EXP2_UB_ELEMENTS * sizeof(T)); - } - } - - __aicore__ inline void ProcessAivOnly() - { - if (stage_ == 1) { - isAivOnly_ = true; - ProcessPreAiv(); - return; - } - if (stage_ == 2) { - isAivOnly_ = true; - ProcessOutAiv(); - return; - } - if (stage_ == 3) { - isAivOnly_ = true; - ProcessPostAiv(); - return; - } - return; - } - - __aicore__ inline void ProcessAiv() - { - if (stage_ == 1) { - ProcessPreAiv(); - return; - } - if (stage_ == 2) { - ProcessOutAiv(); - return; - } - if (stage_ == 3) { - ProcessPostAiv(); - return; - } - return; - } - - __aicore__ inline void ProcessAic() - { - if (stage_ == 1) { - ProcessPreAic(); - return; - } - if (stage_ == 2) { - ProcessOutAic(); - return; - } - if (stage_ == 3) { - ProcessPostAic(); - return; - } - return; - } - -private: - __aicore__ inline uint64_t QOffset(uint64_t b, uint64_t h, uint64_t t, uint64_t d) const - { - return ((b * H_ + h) * T_ + t) * K_ + d; - } - - __aicore__ inline uint64_t KVOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t d, uint64_t dim) const - { - return ((b * HV_ + hv) * T_ + t) * dim + d; - } - - __aicore__ inline uint64_t BetaOffset(uint64_t b, uint64_t hv, uint64_t t) const - { - return (b * HV_ + hv) * T_ + t; - } - - __aicore__ inline uint64_t AOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t j) const - { - return ((b * HV_ + hv) * T_ + t) * BT_ + j; - } - - __aicore__ inline uint64_t HOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t d, uint64_t r) const - { - return (((b * HV_ + hv) * NT_ + chunkIdx) * K_ + d) * V_ + r; - } - - __aicore__ inline uint64_t WScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t t, uint64_t d) const - { - return (((b * HV_ + hv) * NT_ + chunkIdx) * BT_ + t) * K_ + d; - } - - __aicore__ inline uint64_t SolveScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, - uint64_t slot) const - { - (void)b; - (void)hv; - (void)chunkIdx; - uint64_t matrixElements = BT_ * BT_; - return solveCoreIdx_ * KDA_SOLVE_SCRATCH_SLOTS * matrixElements + slot * matrixElements; - } - - __aicore__ inline uint64_t ScoreScratchOffset(uint64_t slot, uint64_t plane, uint64_t t = 0, - uint64_t d = 0) const - { - return (((solveCoreIdx_ * KDA_SCORE_QUEUE_DEPTH + slot) * KDA_SCORE_SCRATCH_PLANES + plane) * BT_ + t) * - K_ + - d; - } - - - - __aicore__ inline uint64_t ScoreRefBlockSize() const - { - if constexpr (IsSameType::value) { - return 2; - } - return KDA_SCORE_REF_BC; - } - - __aicore__ inline uint64_t ScoreRowBlockCount(uint64_t curT, uint64_t rowBegin) const - { - uint64_t blockSize = ScoreRefBlockSize(); - uint64_t rowCount = curT - rowBegin; - if (rowCount > blockSize) { - rowCount = blockSize; - } - return rowCount; - } - - __aicore__ inline uint64_t ScoreRefToken(uint64_t start, uint64_t curT, uint64_t rowBegin, - uint64_t rowCount) const - { - uint64_t ref = rowBegin + rowCount / 2; - if (ref >= curT) { - ref = curT - 1; - } - return start + ref; - } - - __aicore__ inline void RunExp2(LocalTensor &tensor, uint32_t count) - { - SetFlag(EXP2_EVENT_ID); - WaitFlag(EXP2_EVENT_ID); - ClampExpInput(tensor, count); - Exp(tensor, tensor, count); - PipeBarrier(); - SetFlag(EXP2_EVENT_ID); - WaitFlag(EXP2_EVENT_ID); - } - - __aicore__ inline void ClampExpInput(LocalTensor &tensor, uint32_t count) - { - Mins(tensor, tensor, KDA_EXP_INPUT_MAX, count); - PipeBarrier(); - Maxs(tensor, tensor, KDA_EXP_INPUT_MIN, count); - PipeBarrier(); - } - - __aicore__ inline void ClampFp32ToOutputType(LocalTensor &tensor, uint32_t count) - { - if constexpr (IsSameType::value) { - Mins(tensor, tensor, KDA_FP16_MAX, count); - PipeBarrier(); - Maxs(tensor, tensor, -KDA_FP16_MAX, count); - PipeBarrier(); - } - } - - template - __aicore__ inline void CopyVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, - uint64_t count) - { - uint64_t rowBytes = count * static_cast(sizeof(CopyT)); - if (rowBytes >= 32 && rowBytes % 32 == 0) { - DataCopy(dst, src[offset], static_cast(count)); - return; - } - DataCopyParams params{1, static_cast(rowBytes), 0, 0}; - DataCopyPadParams padParams{false, 0, 0, 0}; - DataCopyPad(dst, src[offset], params, padParams); - } - - template - __aicore__ inline void CopyVectorOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src, - uint64_t count) - { - uint64_t rowBytes = count * static_cast(sizeof(CopyT)); - if (rowBytes >= 32 && rowBytes % 32 == 0) { - DataCopy(dst[offset], src, static_cast(count)); - return; - } - DataCopyParams params{1, static_cast(rowBytes), 0, 0}; - DataCopyPad(dst[offset], src, params); - } - - template - __aicore__ inline void CopyRowIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset) - { - CopyVectorIn(dst, src, offset, K_); - } - - template - __aicore__ inline void CopyRowOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src) - { - CopyVectorOut(dst, offset, src, K_); - } - - __aicore__ inline LocalTensor VecScratch(uint64_t slot) - { - return vecBuf_.Get()[slot * EXP2_UB_ELEMENTS]; - } - - template - __aicore__ inline void LoadAsFloatRow(GlobalTensor &src, uint64_t srcOffset, LocalTensor &dst, - uint64_t count) - { - if constexpr (IsSameType::value) { - CopyVectorIn(dst, src, srcOffset, count); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - Adds(dst, dst, 0.0f, static_cast(count)); - PipeBarrier(); - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - } else { - LocalTensor rowLocal = exp2Buf_.Get(); - CopyVectorIn(rowLocal, src, srcOffset, count); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - Cast(dst, rowLocal, RoundMode::CAST_NONE, static_cast(count)); - PipeBarrier(); - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - } - PipeBarrier(); - } - - template - __aicore__ inline void StoreFloatRow(GlobalTensor &dst, uint64_t dstOffset, LocalTensor &src, - uint64_t count) - { - if constexpr (IsSameType::value) { - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - CopyVectorOut(dst, dstOffset, src, count); - } else { - LocalTensor rowLocal = exp2Buf_.Get(); - Cast(rowLocal, src, RoundMode::CAST_RINT, static_cast(count)); - PipeBarrier(); - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - CopyVectorOut(dst, dstOffset, rowLocal, count); - } - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - } - - - - - - __aicore__ inline LocalTensor Exp2NegG(uint64_t b, uint64_t hv, uint64_t t) - { - LocalTensor exp2Local = exp2Buf_.Get(); - CopyRowIn(exp2Local, gk_, KVOffset(b, hv, t, 0, K_)); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - Muls(exp2Local, exp2Local, -LN2, static_cast(K_)); - PipeBarrier(); - RunExp2(exp2Local, static_cast(K_)); - return exp2Local; - } - - - __aicore__ inline uint64_t GateProductToken(uint64_t start, uint64_t logicalIdx, uint64_t subBlockIdx, - uint64_t subBlockNum) const - { - return start + subBlockIdx + logicalIdx * subBlockNum; - } - - __aicore__ inline void LoadGateProductRow(uint64_t b, uint64_t h, uint64_t hv, uint64_t ti) - { - LocalTensor qLocal = qInQue_.AllocTensor(); - LocalTensor kLocal = kInQue_.AllocTensor(); - LocalTensor gLocal = gInQue_.AllocTensor(); - CopyRowIn(qLocal, q_, QOffset(b, h, ti, 0)); - CopyRowIn(kLocal, k_, QOffset(b, h, ti, 0)); - CopyRowIn(gLocal, gk_, KVOffset(b, hv, ti, 0, K_)); - qInQue_.EnQue(qLocal); - kInQue_.EnQue(kLocal); - gInQue_.EnQue(gLocal); - } - - __aicore__ inline void StoreGateProductRow(uint64_t b, uint64_t hv, uint64_t ti) - { - LocalTensor qPosLocal = qgOutQue_.DeQue(); - LocalTensor kPosLocal = wOutQue_.DeQue(); - LocalTensor kNegLocal = kgOutQue_.DeQue(); - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - CopyRowOut(qg_, KVOffset(b, hv, ti, 0, K_), qPosLocal); - CopyRowOut(w_, KVOffset(b, hv, ti, 0, K_), kPosLocal); - CopyRowOut(kg_, KVOffset(b, hv, ti, 0, K_), kNegLocal); - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - qgOutQue_.FreeTensor(qPosLocal); - wOutQue_.FreeTensor(kPosLocal); - kgOutQue_.FreeTensor(kNegLocal); - } - - __aicore__ inline void ComputeGateProductRow(LocalTensor &qFp32, LocalTensor &kFp32, - LocalTensor &gFp32, LocalTensor &refFp32, - LocalTensor &expFp32, LocalTensor &outFp32, - bool useRef, bool zeroKg) - { - LocalTensor qPosLocal = qgOutQue_.AllocTensor(); - LocalTensor kPosLocal = wOutQue_.AllocTensor(); - LocalTensor kNegLocal = kgOutQue_.AllocTensor(); - - if (useRef) { - Sub(expFp32, gFp32, refFp32, static_cast(K_)); - } else { - Adds(expFp32, gFp32, 0.0f, static_cast(K_)); - } - PipeBarrier(); - Muls(expFp32, expFp32, LN2, static_cast(K_)); - PipeBarrier(); - ClampExpInput(expFp32, static_cast(K_)); - Exp(expFp32, expFp32, static_cast(K_)); - PipeBarrier(); - - Mul(outFp32, qFp32, expFp32, static_cast(K_)); - PipeBarrier(); - ClampFp32ToOutputType(outFp32, static_cast(K_)); - if constexpr (IsSameType::value) { - DataCopy(qPosLocal, outFp32, static_cast(K_)); - } else { - Cast(qPosLocal, outFp32, RoundMode::CAST_RINT, static_cast(K_)); - } - PipeBarrier(); - - Mul(outFp32, kFp32, expFp32, static_cast(K_)); - PipeBarrier(); - ClampFp32ToOutputType(outFp32, static_cast(K_)); - if constexpr (IsSameType::value) { - DataCopy(kPosLocal, outFp32, static_cast(K_)); - } else { - Cast(kPosLocal, outFp32, RoundMode::CAST_RINT, static_cast(K_)); - } - PipeBarrier(); - - if (zeroKg) { - Duplicate(outFp32, 0.0f, static_cast(K_)); - PipeBarrier(); - } else { - if (useRef) { - Sub(expFp32, refFp32, gFp32, static_cast(K_)); - } else { - Muls(expFp32, gFp32, -1.0f, static_cast(K_)); - } - PipeBarrier(); - Muls(expFp32, expFp32, LN2, static_cast(K_)); - PipeBarrier(); - ClampExpInput(expFp32, static_cast(K_)); - Exp(expFp32, expFp32, static_cast(K_)); - PipeBarrier(); - Mul(outFp32, kFp32, expFp32, static_cast(K_)); - PipeBarrier(); - } - ClampFp32ToOutputType(outFp32, static_cast(K_)); - if constexpr (IsSameType::value) { - DataCopy(kNegLocal, outFp32, static_cast(K_)); - } else { - Cast(kNegLocal, outFp32, RoundMode::CAST_RINT, static_cast(K_)); - } - - qgOutQue_.EnQue(qPosLocal); - wOutQue_.EnQue(kPosLocal); - kgOutQue_.EnQue(kNegLocal); - } - - __aicore__ inline uint64_t ScoreVectorMaxRows(uint64_t bytesPerElem) const - { - constexpr uint64_t arenaBytes = static_cast(KDA_VEC_ARENA_ELEMENTS) * sizeof(float); - uint64_t maxRows = (arenaBytes / bytesPerElem) / K_; - if (K_ >= 128 && maxRows > 32) { - maxRows = 32; - } - return maxRows; - } - - __aicore__ inline void PrepareScoreFactorsBulk(uint64_t b, uint64_t h, uint64_t hv, uint64_t start, - uint64_t subBlockIdx, uint64_t subBlockNum, - uint64_t refToken, uint64_t scoreRowBegin, - uint64_t scoreRowCount, uint64_t validColEnd, - uint64_t scoreSlot) - { - LocalTensor refFp32 = exp2Buf_.Get(); - LoadAsFloatRow(gk_, KVOffset(b, hv, refToken, 0, K_), refFp32, K_); - - uint64_t qwBegin = scoreRowBegin + (scoreRowCount * subBlockIdx) / subBlockNum; - uint64_t qwEnd = scoreRowBegin + (scoreRowCount * (subBlockIdx + 1)) / subBlockNum; - uint64_t qwMaxRows = ScoreVectorMaxRows(5 * sizeof(float) + 2 * sizeof(T)); - for (uint64_t tileRow = qwBegin; tileRow < qwEnd; tileRow += qwMaxRows) { - uint64_t tileRows = qwEnd - tileRow; - if (tileRows > qwMaxRows) { - tileRows = qwMaxRows; - } - uint64_t elems = tileRows * K_; - LocalTensor arena = vecBuf_.Get(); - LocalTensor qFp32 = arena; - LocalTensor kFp32 = arena[elems]; - LocalTensor gFp32 = arena[2 * elems]; - LocalTensor expFp32 = arena[3 * elems]; - LocalTensor outFp32 = arena[4 * elems]; - uint64_t typedOffset = (5 * elems * sizeof(float) + sizeof(T) - 1) / sizeof(T); - LocalTensor typedBase = vecBuf_.Get()[typedOffset]; - LocalTensor qTyped = typedBase; - LocalTensor kTyped = typedBase[elems]; - - uint64_t token = start + tileRow; - CopyVectorIn(qTyped, q_, QOffset(b, h, token, 0), elems); - CopyVectorIn(kTyped, k_, QOffset(b, h, token, 0), elems); - CopyVectorIn(gFp32, gk_, KVOffset(b, hv, token, 0, K_), elems); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - Cast(qFp32, qTyped, RoundMode::CAST_NONE, static_cast(elems)); - Cast(kFp32, kTyped, RoundMode::CAST_NONE, static_cast(elems)); - PipeBarrier(); - for (uint64_t row = 0; row < tileRows; ++row) { - Sub(expFp32[row * K_], gFp32[row * K_], refFp32, static_cast(K_)); - } - PipeBarrier(); - Muls(expFp32, expFp32, LN2, static_cast(elems)); - PipeBarrier(); - ClampExpInput(expFp32, static_cast(elems)); - Exp(expFp32, expFp32, static_cast(elems)); - PipeBarrier(); - - Mul(outFp32, qFp32, expFp32, static_cast(elems)); - PipeBarrier(); - ClampFp32ToOutputType(outFp32, static_cast(elems)); - Cast(qTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); - PipeBarrier(); - Mul(outFp32, kFp32, expFp32, static_cast(elems)); - PipeBarrier(); - ClampFp32ToOutputType(outFp32, static_cast(elems)); - Cast(kTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); - PipeBarrier(); - - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG, tileRow), - qTyped, elems); - CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W, tileRow), - kTyped, elems); - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - } - - uint64_t kgBegin = (validColEnd * subBlockIdx) / subBlockNum; - uint64_t kgEnd = (validColEnd * (subBlockIdx + 1)) / subBlockNum; - uint64_t kgMaxRows = ScoreVectorMaxRows(4 * sizeof(float) + sizeof(T)); - for (uint64_t tileRow = kgBegin; tileRow < kgEnd; tileRow += kgMaxRows) { - uint64_t tileRows = kgEnd - tileRow; - if (tileRows > kgMaxRows) { - tileRows = kgMaxRows; - } - uint64_t elems = tileRows * K_; - LocalTensor arena = vecBuf_.Get(); - LocalTensor kFp32 = arena; - LocalTensor gFp32 = arena[elems]; - LocalTensor expFp32 = arena[2 * elems]; - LocalTensor outFp32 = arena[3 * elems]; - uint64_t typedOffset = (4 * elems * sizeof(float) + sizeof(T) - 1) / sizeof(T); - LocalTensor kTyped = vecBuf_.Get()[typedOffset]; - - uint64_t token = start + tileRow; - CopyVectorIn(kTyped, k_, QOffset(b, h, token, 0), elems); - CopyVectorIn(gFp32, gk_, KVOffset(b, hv, token, 0, K_), elems); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - Cast(kFp32, kTyped, RoundMode::CAST_NONE, static_cast(elems)); - PipeBarrier(); - for (uint64_t row = 0; row < tileRows; ++row) { - Sub(expFp32[row * K_], refFp32, gFp32[row * K_], static_cast(K_)); - } - PipeBarrier(); - Muls(expFp32, expFp32, LN2, static_cast(elems)); - PipeBarrier(); - ClampExpInput(expFp32, static_cast(elems)); - Exp(expFp32, expFp32, static_cast(elems)); - PipeBarrier(); - Mul(outFp32, kFp32, expFp32, static_cast(elems)); - PipeBarrier(); - ClampFp32ToOutputType(outFp32, static_cast(elems)); - Cast(kTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); - PipeBarrier(); - - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG, tileRow), - kTyped, elems); - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - } - } - - __aicore__ inline bool PrepareGateProductsBulk(uint64_t b, uint64_t h, uint64_t hv, uint64_t start, - uint64_t curT, uint64_t subBlockIdx, uint64_t subBlockNum, - bool useRef, uint64_t refToken, uint64_t validColEnd, - bool writeScoreScratch, uint64_t scoreSlot) - { - if constexpr (IsSameType::value) { - return false; - } - if (subBlockNum == 0 || subBlockIdx >= subBlockNum || K_ == 0) { - return false; - } - uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; - uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; - if (rowBegin >= rowEnd) { - return true; - } - - constexpr uint64_t arenaBytes = static_cast(KDA_VEC_ARENA_ELEMENTS) * sizeof(float); - constexpr uint64_t bytesPerElem = 5 * sizeof(float) + 3 * sizeof(T); - uint64_t maxElems = arenaBytes / bytesPerElem; - uint64_t maxRows = maxElems / K_; - // Keep the multi-row SIMD tile below the 192 KiB per-core UB budget. - // K=128 uses five FP32 work planes plus three typed planes; 32 rows - // leaves headroom for alignment and the surrounding pipeline buffers. - if (K_ >= 128 && maxRows > 32) { - maxRows = 32; - } - if (maxRows == 0) { - return false; - } - LocalTensor refFp32 = exp2Buf_.Get(); - if (useRef) { - LoadAsFloatRow(gk_, KVOffset(b, hv, refToken, 0, K_), refFp32, K_); - } - - for (uint64_t tileRow = rowBegin; tileRow < rowEnd; tileRow += maxRows) { - uint64_t tileRows = rowEnd - tileRow; - if (tileRows > maxRows) { - tileRows = maxRows; - } - uint64_t elems = tileRows * K_; - LocalTensor arena = vecBuf_.Get(); - LocalTensor qFp32 = arena; - LocalTensor kFp32 = arena[elems]; - LocalTensor gFp32 = arena[2 * elems]; - LocalTensor expFp32 = arena[3 * elems]; - LocalTensor outFp32 = arena[4 * elems]; - - uint64_t typedOffset = (5 * elems * sizeof(float) + sizeof(T) - 1) / sizeof(T); - uint64_t typedCapacity = arenaBytes / sizeof(T); - if (typedOffset + 3 * elems > typedCapacity) { - return false; - } - LocalTensor typedBase = vecBuf_.Get()[typedOffset]; - LocalTensor qTyped = typedBase; - LocalTensor kTyped = typedBase[elems]; - LocalTensor kgTyped = typedBase[2 * elems]; - - uint64_t token = start + tileRow; - CopyVectorIn(qTyped, q_, QOffset(b, h, token, 0), elems); - CopyVectorIn(kTyped, k_, QOffset(b, h, token, 0), elems); - CopyVectorIn(gFp32, gk_, KVOffset(b, hv, token, 0, K_), elems); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - - Cast(qFp32, qTyped, RoundMode::CAST_NONE, static_cast(elems)); - Cast(kFp32, kTyped, RoundMode::CAST_NONE, static_cast(elems)); - PipeBarrier(); - - if (useRef) { - for (uint64_t row = 0; row < tileRows; ++row) { - Sub(expFp32[row * K_], gFp32[row * K_], refFp32, static_cast(K_)); - } - } else { - Adds(expFp32, gFp32, 0.0f, static_cast(elems)); - } - PipeBarrier(); - Muls(expFp32, expFp32, LN2, static_cast(elems)); - PipeBarrier(); - ClampExpInput(expFp32, static_cast(elems)); - Exp(expFp32, expFp32, static_cast(elems)); - PipeBarrier(); - - Mul(outFp32, qFp32, expFp32, static_cast(elems)); - PipeBarrier(); - ClampFp32ToOutputType(outFp32, static_cast(elems)); - Cast(qTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); - PipeBarrier(); - - Mul(outFp32, kFp32, expFp32, static_cast(elems)); - PipeBarrier(); - ClampFp32ToOutputType(outFp32, static_cast(elems)); - Cast(kTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); - PipeBarrier(); - - if (useRef) { - for (uint64_t row = 0; row < tileRows; ++row) { - Sub(expFp32[row * K_], refFp32, gFp32[row * K_], static_cast(K_)); - } - } else { - Muls(expFp32, gFp32, -1.0f, static_cast(elems)); - } - PipeBarrier(); - Muls(expFp32, expFp32, LN2, static_cast(elems)); - PipeBarrier(); - ClampExpInput(expFp32, static_cast(elems)); - Exp(expFp32, expFp32, static_cast(elems)); - PipeBarrier(); - Mul(outFp32, kFp32, expFp32, static_cast(elems)); - PipeBarrier(); - if (useRef && tileRow + tileRows > validColEnd) { - for (uint64_t row = 0; row < tileRows; ++row) { - if (tileRow + row >= validColEnd) { - Duplicate(outFp32[row * K_], 0.0f, static_cast(K_)); - } - } - PipeBarrier(); - } - ClampFp32ToOutputType(outFp32, static_cast(elems)); - Cast(kgTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); - PipeBarrier(); - - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - if (writeScoreScratch) { - CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG, tileRow), - qTyped, elems); - CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W, tileRow), - kTyped, elems); - CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG, tileRow), - kgTyped, elems); - } else { - CopyVectorOut(qg_, KVOffset(b, hv, token, 0, K_), qTyped, elems); - CopyVectorOut(w_, KVOffset(b, hv, token, 0, K_), kTyped, elems); - CopyVectorOut(kg_, KVOffset(b, hv, token, 0, K_), kgTyped, elems); - } - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - } - return true; - } - - __aicore__ inline void PrepareGateProducts(uint64_t b, uint64_t h, uint64_t hv, uint64_t start, uint64_t curT, - uint64_t subBlockIdx, uint64_t subBlockNum, bool useRef = false, - uint64_t refToken = 0, uint64_t validColEnd = 0, - bool writeScoreScratch = false, uint64_t scoreSlot = 0, - uint64_t scoreRowBegin = 0, uint64_t scoreRowCount = 0) - { - if (subBlockNum == 0 || subBlockIdx >= subBlockNum) { - return; - } - if (validColEnd == 0 || validColEnd > curT) { - validColEnd = curT; - } - if (writeScoreScratch) { - PrepareScoreFactorsBulk(b, h, hv, start, subBlockIdx, subBlockNum, refToken, scoreRowBegin, - scoreRowCount, validColEnd, scoreSlot); - return; - } - if (PrepareGateProductsBulk(b, h, hv, start, curT, subBlockIdx, subBlockNum, useRef, refToken, - validColEnd, writeScoreScratch, scoreSlot)) { - return; - } - - if (subBlockIdx >= curT) { - return; - } - - LocalTensor vecLocal = vecBuf_.Get(); - LocalTensor qFp32 = vecLocal; - LocalTensor kFp32 = vecLocal[EXP2_UB_ELEMENTS]; - LocalTensor gFp32 = vecLocal[2 * EXP2_UB_ELEMENTS]; - LocalTensor expFp32 = vecLocal[3 * EXP2_UB_ELEMENTS]; - LocalTensor outFp32 = vecLocal[4 * EXP2_UB_ELEMENTS]; - LocalTensor refFp32 = vecLocal[5 * EXP2_UB_ELEMENTS]; - if (useRef) { - LoadAsFloatRow(gk_, KVOffset(b, hv, refToken, 0, K_), refFp32, K_); - } - - uint64_t rowCount = 0; - for (uint64_t i = subBlockIdx; i < curT; i += subBlockNum) { - ++rowCount; - } - - LoadGateProductRow(b, h, hv, GateProductToken(start, 0, subBlockIdx, subBlockNum)); - for (uint64_t logicalIdx = 0; logicalIdx < rowCount; ++logicalIdx) { - uint64_t ti = GateProductToken(start, logicalIdx, subBlockIdx, subBlockNum); - LocalTensor qLocal = qInQue_.DeQue(); - LocalTensor kLocal = kInQue_.DeQue(); - LocalTensor gLocal = gInQue_.DeQue(); - - if (logicalIdx + 1 < rowCount) { - LoadGateProductRow(b, h, hv, GateProductToken(start, logicalIdx + 1, subBlockIdx, subBlockNum)); - } - - if constexpr (IsSameType::value) { - DataCopy(qFp32, qLocal, static_cast(K_)); - DataCopy(kFp32, kLocal, static_cast(K_)); - } else { - Cast(qFp32, qLocal, RoundMode::CAST_NONE, static_cast(K_)); - Cast(kFp32, kLocal, RoundMode::CAST_NONE, static_cast(K_)); - } - DataCopy(gFp32, gLocal, static_cast(K_)); - qInQue_.FreeTensor(qLocal); - kInQue_.FreeTensor(kLocal); - gInQue_.FreeTensor(gLocal); - PipeBarrier(); - - if (logicalIdx > 0) { - StoreGateProductRow(b, hv, GateProductToken(start, logicalIdx - 1, subBlockIdx, subBlockNum)); - } - bool zeroKg = useRef && (ti - start >= validColEnd); - ComputeGateProductRow(qFp32, kFp32, gFp32, refFp32, expFp32, outFp32, useRef, zeroKg); - } - StoreGateProductRow(b, hv, GateProductToken(start, rowCount - 1, subBlockIdx, subBlockNum)); - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - } - - __aicore__ inline void ComputeRawAqkAkkCube(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT) - { - ComputeRawAqkAkkCubeBlock(b, hv, start, curT, 0, curT); - } - - __aicore__ inline void ComputeRawAqkAkkCubeBlock(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT, - uint64_t rowBegin, uint64_t rowCount, - bool readScoreScratch = false, uint64_t scoreSlot = 0, - uint64_t colCount = 0) - { - using ElementA = T; - using ElementB = T; - using ElementC = float; - using LayoutTagA = Catlass::layout::RowMajor; - using LayoutTagB = Catlass::layout::ColumnMajor; - using LayoutTagC = Catlass::layout::RowMajor; - using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; - using BlockMmad = Catlass::Gemm::Block::BlockMmadTla; - - Catlass::Arch::Resource resource; - BlockMmad blockMmad(resource); - auto layoutA = tla::MakeLayout(BT_, K_); - auto layoutB = tla::MakeLayout(K_, BT_); - auto layoutC = tla::MakeLayout(BT_, BT_); - if (colCount == 0 || colCount > curT) { - colCount = curT; - } - Catlass::GemmCoord shape{static_cast(rowCount), static_cast(colCount), - static_cast(K_)}; - - auto tensorQPos = readScoreScratch ? - tla::MakeTensor(scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG)], - layoutA, Catlass::Arch::PositionGM{}) : - tla::MakeTensor(qg_[KVOffset(b, hv, start, 0, K_)], layoutA, - Catlass::Arch::PositionGM{}); - auto tensorKPos = readScoreScratch ? - tla::MakeTensor(scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W)], - layoutA, Catlass::Arch::PositionGM{}) : - tla::MakeTensor(w_[KVOffset(b, hv, start, 0, K_)], layoutA, - Catlass::Arch::PositionGM{}); - auto tensorKNeg = readScoreScratch ? - tla::MakeTensor(scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG)], - layoutB, Catlass::Arch::PositionGM{}) : - tla::MakeTensor(kg_[KVOffset(b, hv, start, 0, K_)], layoutB, - Catlass::Arch::PositionGM{}); - auto tensorAqk = tla::MakeTensor(aqk_[AOffset(b, hv, start, 0)], layoutC, - Catlass::Arch::PositionGM{}); - auto tensorAkk = tla::MakeTensor(akk_[AOffset(b, hv, start, 0)], layoutC, - Catlass::Arch::PositionGM{}); - - auto blockQPos = GetTile(tensorQPos, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.k())); - auto blockKPos = GetTile(tensorKPos, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.k())); - auto blockKNeg = GetTile(tensorKNeg, tla::MakeCoord(0, 0), tla::MakeShape(shape.k(), shape.n())); - auto blockAqk = GetTile(tensorAqk, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.n())); - auto blockAkk = GetTile(tensorAkk, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.n())); - - blockMmad(blockQPos, blockKNeg, blockAqk, shape); - PipeBarrier(); - blockMmad(blockKPos, blockKNeg, blockAkk, shape); - PipeBarrier(); - } - - __aicore__ inline bool UseAkkCubeSolve(uint64_t curT) const - { - return curT > 0 && curT <= BT_ && (BT_ == 64 || BT_ == 128) && K_ >= 16 && V_ >= 16 && - V_ <= 256 && K_ % 16 == 0 && V_ % 16 == 0; - } - - __aicore__ inline bool UsePostWuCube(uint64_t curT) const - { - return curT > 0 && curT <= BT_ && (BT_ == 64 || BT_ == 128) && K_ >= 16 && V_ >= 16 && - V_ <= 256 && K_ % 16 == 0 && V_ % 16 == 0; - } - - __aicore__ inline void CopyLocalFloat(LocalTensor dst, LocalTensor src, uint64_t count) - { - if (count == 0) { - return; - } - Adds(dst, src, 0.0f, static_cast(count)); - PipeBarrier(); - } - - __aicore__ inline void FillLocalFloat(LocalTensor dst, float value, uint64_t count) - { - if (count == 0) { - return; - } - Duplicate(dst, value, static_cast(count)); - PipeBarrier(); - } - - __aicore__ inline void BuildPrefixMask(LocalTensor dst, uint64_t prefix, uint64_t count) - { - if (prefix > count) { - prefix = count; - } - Duplicate(dst, 0.0f, static_cast(count)); - if (prefix > 0) { - Duplicate(dst, 1.0f, static_cast(prefix)); - } - PipeBarrier(); - } - - __aicore__ inline uint64_t BuildCausalMask(uint64_t threshold, uint64_t colBegin) const - { - if (threshold <= colBegin) { - return ~0ULL; - } - if (threshold >= colBegin + KDA_SOLVE_BT) { - return 0ULL; - } - return ~0ULL << (threshold - colBegin); - } - - __aicore__ inline void BuildCausalSelectMasks(LocalTensor aqkMask, LocalTensor akkMask, - uint64_t rowBegin, uint64_t rowCount, uint64_t colBegin) - { - __ubuf__ uint64_t *aqkMaskPtr = reinterpret_cast<__ubuf__ uint64_t *>(aqkMask.GetPhyAddr()); - __ubuf__ uint64_t *akkMaskPtr = reinterpret_cast<__ubuf__ uint64_t *>(akkMask.GetPhyAddr()); - for (uint32_t localRow = 0; localRow < rowCount; ++localRow) { - uint32_t row = static_cast(rowBegin + localRow); - aqkMaskPtr[localRow] = BuildCausalMask(static_cast(row) + 1, colBegin); - akkMaskPtr[localRow] = BuildCausalMask(static_cast(row), colBegin); - } - } - - __aicore__ inline void SelectCausalRows(LocalTensor aqkMat, LocalTensor akkMat, - uint64_t rowBegin, uint64_t rowCount) - { - LocalTensor aqkMask = vecBuf_.Get()[KDA_SELECT_AQK_MASK_BYTE_OFFSET]; - LocalTensor akkMask = vecBuf_.Get()[KDA_SELECT_AKK_MASK_BYTE_OFFSET]; - LocalTensor zeroLocal = vecBuf_.Get()[KDA_SELECT_ZERO_FLOAT_OFFSET]; - Duplicate(zeroLocal, 0.0f, 8); - PipeBarrier(); - - uint64_t colBlockCount = (BT_ + KDA_SOLVE_BT - 1) / KDA_SOLVE_BT; - for (uint64_t colBlock = 0; colBlock < colBlockCount; ++colBlock) { - uint64_t maskOffset = colBlock * KDA_SELECT_COL_MASK_BYTES; - uint64_t colBegin = colBlock * KDA_SOLVE_BT; - BuildCausalSelectMasks(aqkMask[maskOffset], akkMask[maskOffset], rowBegin, rowCount, colBegin); - } - SetFlag(EXP2_EVENT_ID); - WaitFlag(EXP2_EVENT_ID); - - uint8_t rowStride = static_cast(BT_ * sizeof(float) / 32); - BinaryRepeatParams repeatParams = {1, 0, 1, rowStride, 0, rowStride}; - for (uint64_t colBlock = 0; colBlock < colBlockCount; ++colBlock) { - uint64_t maskOffset = colBlock * KDA_SELECT_COL_MASK_BYTES; - uint64_t colBegin = colBlock * KDA_SOLVE_BT; - Select(aqkMat[colBegin], aqkMask[maskOffset], zeroLocal, aqkMat[colBegin], - SELMODE::VSEL_TENSOR_TENSOR_MODE, KDA_SOLVE_BT, static_cast(rowCount), repeatParams); - Select(akkMat[colBegin], akkMask[maskOffset], zeroLocal, akkMat[colBegin], - SELMODE::VSEL_TENSOR_TENSOR_MODE, KDA_SOLVE_BT, static_cast(rowCount), repeatParams); - } - PipeBarrier(); - SetFlag(EXP2_EVENT_ID); - WaitFlag(EXP2_EVENT_ID); - } - - __aicore__ inline void PrepareAqkAkkSolveInput64(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start) - { - LocalTensor arena = vecBuf_.Get(); - LocalTensor aqkMat = arena; - LocalTensor akkMat = arena[KDA_SOLVE_MATRIX_ELEMENTS]; - LocalTensor xMat = arena[2 * KDA_SOLVE_MATRIX_ELEMENTS]; - LocalTensor betaLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS]; - LocalTensor betaBrcb = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT]; - LocalTensor maskLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512]; - LocalTensor oneHotLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512 + KDA_SOLVE_BT]; - - LoadAsFloatRow(beta_, BetaOffset(b, hv, start), betaLocal, KDA_SOLVE_BT); - Brcb(betaBrcb, betaLocal, 8, {1, 8}); - PipeBarrier(); - - DataCopy(aqkMat, aqk_[AOffset(b, hv, start, 0)], KDA_SOLVE_MATRIX_ELEMENTS); - DataCopy(akkMat, akk_[AOffset(b, hv, start, 0)], KDA_SOLVE_MATRIX_ELEMENTS); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - - for (uint64_t col = 0; col < KDA_SOLVE_BT; col += 8) { - Mul(akkMat[col], akkMat[col], betaBrcb, 8, KDA_SOLVE_BT, {1, 1, 1, 8, 8, 1}); - PipeBarrier(); - } - SelectCausalRows(aqkMat, akkMat, 0, KDA_SOLVE_BT); - - Muls(xMat, akkMat, -1.0f, KDA_SOLVE_MATRIX_ELEMENTS); - PipeBarrier(); - for (uint64_t row = 0; row < KDA_SOLVE_BT; ++row) { - BuildPrefixMask(maskLocal, row + 1, KDA_SOLVE_BT); - BuildPrefixMask(oneHotLocal, row, KDA_SOLVE_BT); - Sub(maskLocal, maskLocal, oneHotLocal, KDA_SOLVE_BT); - PipeBarrier(); - Add(xMat[row * KDA_SOLVE_BT], xMat[row * KDA_SOLVE_BT], maskLocal, KDA_SOLVE_BT); - PipeBarrier(); - } - - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - DataCopy(aqk_[AOffset(b, hv, start, 0)], aqkMat, KDA_SOLVE_MATRIX_ELEMENTS); - DataCopy(akk_[AOffset(b, hv, start, 0)], akkMat, KDA_SOLVE_MATRIX_ELEMENTS); - DataCopy(h_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X)], xMat, - KDA_SOLVE_MATRIX_ELEMENTS); - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - } - - __aicore__ inline void PrepareAqkAkkSolveInputTail(uint64_t b, uint64_t hv, uint64_t chunkIdx, - uint64_t start, uint64_t curT) - { - uint64_t elemCount = curT * KDA_SOLVE_BT; - DataCopyParams aqkValidParams{1, static_cast(elemCount * sizeof(float)), 0, 0}; - DataCopyParams akkValidParams{1, static_cast(elemCount * sizeof(float)), 0, 0}; - DataCopyPadParams padParams{false, 0, 0, 0}; - LocalTensor arena = vecBuf_.Get(); - LocalTensor aqkMat = arena; - LocalTensor akkMat = arena[KDA_SOLVE_MATRIX_ELEMENTS]; - LocalTensor xMat = arena[2 * KDA_SOLVE_MATRIX_ELEMENTS]; - LocalTensor betaLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS]; - LocalTensor betaBrcb = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT]; - LocalTensor maskLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512]; - LocalTensor oneHotLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512 + KDA_SOLVE_BT]; - - FillLocalFloat(betaLocal, 0.0f, KDA_SOLVE_BT); - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - LoadAsFloatRow(beta_, BetaOffset(b, hv, start), betaLocal, curT); - Brcb(betaBrcb, betaLocal, 8, {1, 8}); - PipeBarrier(); - - DataCopyPad(aqkMat, aqk_[AOffset(b, hv, start, 0)], aqkValidParams, padParams); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - if (elemCount < KDA_SOLVE_MATRIX_ELEMENTS) { - FillLocalFloat(aqkMat[elemCount], 0.0f, KDA_SOLVE_MATRIX_ELEMENTS - elemCount); - } - DataCopyPad(akkMat, akk_[AOffset(b, hv, start, 0)], akkValidParams, padParams); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - if (elemCount < KDA_SOLVE_MATRIX_ELEMENTS) { - FillLocalFloat(akkMat[elemCount], 0.0f, KDA_SOLVE_MATRIX_ELEMENTS - elemCount); - } - - for (uint64_t col = 0; col < KDA_SOLVE_BT; col += 8) { - Mul(akkMat[col], akkMat[col], betaBrcb, 8, KDA_SOLVE_BT, {1, 1, 1, 8, 8, 1}); - PipeBarrier(); - } - SelectCausalRows(aqkMat, akkMat, 0, KDA_SOLVE_BT); - - Muls(xMat, akkMat, -1.0f, KDA_SOLVE_MATRIX_ELEMENTS); - PipeBarrier(); - for (uint64_t row = 0; row < KDA_SOLVE_BT; ++row) { - BuildPrefixMask(maskLocal, row + 1, KDA_SOLVE_BT); - BuildPrefixMask(oneHotLocal, row, KDA_SOLVE_BT); - Sub(maskLocal, maskLocal, oneHotLocal, KDA_SOLVE_BT); - PipeBarrier(); - Add(xMat[row * KDA_SOLVE_BT], xMat[row * KDA_SOLVE_BT], maskLocal, KDA_SOLVE_BT); - PipeBarrier(); - } - - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - DataCopyPad(aqk_[AOffset(b, hv, start, 0)], aqkMat, aqkValidParams); - DataCopy(h_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X)], xMat, - KDA_SOLVE_MATRIX_ELEMENTS); - DataCopy(h_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0)], akkMat, - KDA_SOLVE_MATRIX_ELEMENTS); - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - } - - __aicore__ inline void GetSolveRowRange(uint64_t curT, uint64_t subBlockIdx, uint64_t subBlockNum, - uint64_t &rowBegin, uint64_t &rowEnd) const - { - if (subBlockNum == 0 || subBlockIdx >= subBlockNum) { - rowBegin = 0; - rowEnd = 0; - return; - } - rowBegin = (curT * subBlockIdx) / subBlockNum; - rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; - } - - __aicore__ inline void PrepareAqkAkkSolveInputRows(uint64_t b, uint64_t hv, uint64_t chunkIdx, - uint64_t start, uint64_t curT, uint64_t rowBegin, - uint64_t rowEnd, bool storeLToAkk, bool storeLToScratch) - { - uint64_t rowCount = rowEnd - rowBegin; - if (rowCount == 0) { - return; - } - uint64_t validRowCount = rowBegin < curT ? curT - rowBegin : 0; - if (validRowCount > rowCount) { - validRowCount = rowCount; - } - uint64_t elemCount = rowCount * BT_; - uint64_t validElemCount = validRowCount * BT_; - DataCopyParams aqkValidParams{1, static_cast(validElemCount * sizeof(float)), 0, 0}; - DataCopyParams akkValidParams{1, static_cast(validElemCount * sizeof(float)), 0, 0}; - DataCopyPadParams padParams{false, 0, 0, 0}; - LocalTensor arena = vecBuf_.Get(); - LocalTensor aqkMat = arena; - LocalTensor akkMat = arena[elemCount]; - LocalTensor xMat = arena[2 * elemCount]; - LocalTensor betaLocal = arena[3 * elemCount]; - LocalTensor betaBrcb = arena[3 * elemCount + BT_]; - LocalTensor maskLocal = arena[3 * elemCount + BT_ + 512]; - LocalTensor oneHotLocal = arena[3 * elemCount + BT_ + 512 + BT_]; - - uint64_t token = start + rowBegin; - - FillLocalFloat(aqkMat, 0.0f, elemCount); - FillLocalFloat(akkMat, 0.0f, elemCount); - FillLocalFloat(betaLocal, 0.0f, rowCount); - SetFlag(KDA_MTE2_MTE3_EVENT_ID); - WaitFlag(KDA_MTE2_MTE3_EVENT_ID); - if (validRowCount > 0) { - LoadAsFloatRow(beta_, BetaOffset(b, hv, token), betaLocal, validRowCount); - DataCopyPad(aqkMat, aqk_[AOffset(b, hv, token, 0)], aqkValidParams, padParams); - DataCopyPad(akkMat, akk_[AOffset(b, hv, token, 0)], akkValidParams, padParams); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - } - Brcb(betaBrcb, betaLocal, static_cast((rowCount + 7) / 8), {1, 8}); - PipeBarrier(); - - uint8_t rowStride = static_cast(BT_ * sizeof(float) / 32); - for (uint64_t col = 0; col < BT_; col += 8) { - Mul(akkMat[col], akkMat[col], betaBrcb, 8, static_cast(rowCount), - {1, 1, 0, rowStride, rowStride, 1}); - PipeBarrier(); - } - if (validRowCount > 0) { - SelectCausalRows(aqkMat, akkMat, rowBegin, validRowCount); - } - - Muls(xMat, akkMat, -1.0f, static_cast(elemCount)); - PipeBarrier(); - for (uint64_t localRow = 0; localRow < rowCount; ++localRow) { - uint64_t row = rowBegin + localRow; - BuildPrefixMask(maskLocal, row + 1, BT_); - BuildPrefixMask(oneHotLocal, row, BT_); - Sub(maskLocal, maskLocal, oneHotLocal, static_cast(BT_)); - PipeBarrier(); - Add(xMat[localRow * BT_], xMat[localRow * BT_], maskLocal, static_cast(BT_)); - PipeBarrier(); - } - - uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; - uint64_t lBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0) + rowBegin * BT_; - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - if (validRowCount > 0) { - DataCopyPad(aqk_[AOffset(b, hv, token, 0)], aqkMat, aqkValidParams); - if (storeLToAkk) { - DataCopyPad(akk_[AOffset(b, hv, token, 0)], akkMat, akkValidParams); - } - } - DataCopy(solveWorkspace_[xBase], xMat, static_cast(elemCount)); - if (storeLToScratch) { - DataCopy(solveWorkspace_[lBase], akkMat, static_cast(elemCount)); - } - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - } - - __aicore__ inline void CubeGemmSolveSub(GlobalTensor &tensorA, uint64_t baseA, uint64_t rowA, uint64_t colA, - GlobalTensor &tensorB, uint64_t baseB, uint64_t rowB, uint64_t colB, - GlobalTensor &tensorC, uint64_t baseC, uint64_t rowC, uint64_t colC, - uint32_t m, uint32_t n, uint32_t k) - { - using ElementA = float; - using ElementB = float; - using ElementC = float; - using LayoutTagA = Catlass::layout::RowMajor; - using LayoutTagB = Catlass::layout::RowMajor; - using LayoutTagC = Catlass::layout::RowMajor; - using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; - using BlockMmad = Catlass::Gemm::Block::BlockMmadTla; - Catlass::Arch::Resource resource; - auto layoutA = tla::MakeLayout(BT_, BT_); - auto layoutB = tla::MakeLayout(BT_, BT_); - auto layoutC = tla::MakeLayout(BT_, BT_); - auto tensorLayoutA = tla::MakeTensor(tensorA[baseA], layoutA, Catlass::Arch::PositionGM{}); - auto tensorLayoutB = tla::MakeTensor(tensorB[baseB], layoutB, Catlass::Arch::PositionGM{}); - auto tensorLayoutC = tla::MakeTensor(tensorC[baseC], layoutC, Catlass::Arch::PositionGM{}); - Catlass::GemmCoord shape{m, n, k}; - auto blockA = GetTile(tensorLayoutA, tla::MakeCoord(rowA, colA), tla::MakeShape(shape.m(), shape.k())); - auto blockB = GetTile(tensorLayoutB, tla::MakeCoord(rowB, colB), tla::MakeShape(shape.k(), shape.n())); - auto blockC = GetTile(tensorLayoutC, tla::MakeCoord(rowC, colC), tla::MakeShape(shape.m(), shape.n())); - BlockMmad blockMmad(resource); - blockMmad(blockA, blockB, blockC, shape); - PipeBarrier(); - } - - __aicore__ inline void AddSolveTmpToX(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, - bool storeAkk) - { - LocalTensor arena = vecBuf_.Get(); - LocalTensor xLocal = arena; - LocalTensor tmpLocal = arena[KDA_SOLVE_MATRIX_ELEMENTS]; - uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); - uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); - - DataCopy(xLocal, h_[xBase], KDA_SOLVE_MATRIX_ELEMENTS); - DataCopy(tmpLocal, h_[tmpBase], KDA_SOLVE_MATRIX_ELEMENTS); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - - Add(xLocal, xLocal, tmpLocal, KDA_SOLVE_MATRIX_ELEMENTS); - PipeBarrier(); - - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - DataCopy(h_[xBase], xLocal, KDA_SOLVE_MATRIX_ELEMENTS); - if (storeAkk) { - DataCopy(akk_[AOffset(b, hv, start, 0)], xLocal, KDA_SOLVE_MATRIX_ELEMENTS); - } - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - } - - __aicore__ inline void AddSolveTmpToXTail(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, - uint64_t curT, bool storeAkk) - { - uint64_t elemCount = curT * KDA_SOLVE_BT; - DataCopyParams validParams{1, static_cast(elemCount * sizeof(float)), 0, 0}; - LocalTensor arena = vecBuf_.Get(); - LocalTensor xLocal = arena; - LocalTensor tmpLocal = arena[KDA_SOLVE_MATRIX_ELEMENTS]; - uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); - uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); - - DataCopy(xLocal, h_[xBase], KDA_SOLVE_MATRIX_ELEMENTS); - DataCopy(tmpLocal, h_[tmpBase], KDA_SOLVE_MATRIX_ELEMENTS); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - - Add(xLocal, xLocal, tmpLocal, KDA_SOLVE_MATRIX_ELEMENTS); - PipeBarrier(); - - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - DataCopy(h_[xBase], xLocal, KDA_SOLVE_MATRIX_ELEMENTS); - if (storeAkk) { - DataCopyPad(akk_[AOffset(b, hv, start, 0)], xLocal, validParams); - } - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - } - - __aicore__ inline void AddSolveTmpToXRows(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, - uint64_t curT, uint64_t rowBegin, uint64_t rowEnd, bool storeAkk) - { - uint64_t rowCount = rowEnd - rowBegin; - if (rowCount == 0) { - return; - } - uint64_t validRowCount = rowBegin < curT ? curT - rowBegin : 0; - if (validRowCount > rowCount) { - validRowCount = rowCount; - } - uint64_t elemCount = rowCount * BT_; - uint64_t validElemCount = validRowCount * BT_; - DataCopyParams validParams{1, static_cast(validElemCount * sizeof(float)), 0, 0}; - LocalTensor arena = vecBuf_.Get(); - LocalTensor xLocal = arena; - LocalTensor tmpLocal = arena[elemCount]; - uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; - uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP) + rowBegin * BT_; - uint64_t token = start + rowBegin; - - DataCopy(xLocal, solveWorkspace_[xBase], static_cast(elemCount)); - DataCopy(tmpLocal, solveWorkspace_[tmpBase], static_cast(elemCount)); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - - Add(xLocal, xLocal, tmpLocal, static_cast(elemCount)); - PipeBarrier(); - - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - DataCopy(solveWorkspace_[xBase], xLocal, static_cast(elemCount)); - if (storeAkk && validRowCount > 0) { - DataCopyPad(akk_[AOffset(b, hv, token, 0)], xLocal, validParams); - } - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - } - - __aicore__ inline void AddSolveTmpToXDiagRows(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, - uint64_t rowBegin, uint64_t rowEnd, bool storeAkk) - { - uint64_t rowCount = rowEnd - rowBegin; - if (rowCount == 0) { - return; - } - uint64_t elemCount = rowCount * BT_; - DataCopyParams validParams{1, static_cast(elemCount * sizeof(float)), 0, 0}; - LocalTensor arena = vecBuf_.Get(); - LocalTensor xLocal = arena; - LocalTensor tmpLocal = arena[elemCount]; - uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; - uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP) + rowBegin * BT_; - uint64_t token = start + rowBegin; - - DataCopy(xLocal, solveWorkspace_[xBase], static_cast(elemCount)); - DataCopy(tmpLocal, solveWorkspace_[tmpBase], static_cast(elemCount)); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - - for (uint64_t localRow = 0; localRow < rowCount; ++localRow) { - uint64_t row = rowBegin + localRow; - uint64_t col = (row / KDA_SOLVE_DIAG_BT) * KDA_SOLVE_DIAG_BT; - uint64_t offset = localRow * BT_ + col; - Add(xLocal[offset], xLocal[offset], tmpLocal[offset], KDA_SOLVE_DIAG_BT); - PipeBarrier(); - } - - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - DataCopy(solveWorkspace_[xBase], xLocal, static_cast(elemCount)); - if (storeAkk) { - DataCopyPad(akk_[AOffset(b, hv, token, 0)], xLocal, validParams); - } - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - } - - __aicore__ inline void StoreSolveXRowsToAkk(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, - uint64_t curT, uint64_t rowBegin, uint64_t rowEnd) - { - uint64_t validRowCount = rowBegin < curT ? curT - rowBegin : 0; - uint64_t rowCount = rowEnd - rowBegin; - if (validRowCount > rowCount) { - validRowCount = rowCount; - } - if (validRowCount == 0) { - return; - } - uint64_t elemCount = validRowCount * BT_; - DataCopyParams validParams{1, static_cast(elemCount * sizeof(float)), 0, 0}; - LocalTensor xLocal = vecBuf_.Get(); - uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; - - DataCopy(xLocal, solveWorkspace_[xBase], static_cast(elemCount)); - SetFlag(KDA_MTE2_MTE3_EVENT_ID); - WaitFlag(KDA_MTE2_MTE3_EVENT_ID); - DataCopyPad(akk_[AOffset(b, hv, start + rowBegin, 0)], xLocal, validParams); - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - } - - __aicore__ inline void ComputeAkkMergeCube(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start) - { - uint64_t aiBase = AOffset(b, hv, start, 0); - uint64_t negABase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); - uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); - - for (uint32_t mergeSize = 2 * KDA_SOLVE_DIAG_BT; mergeSize <= BT_; mergeSize *= 2) { - uint32_t half = mergeSize / 2; - for (uint32_t block = 0; block < BT_; block += mergeSize) { - uint32_t lower = block + half; - CubeGemmSolveSub(akk_, aiBase, lower, lower, solveWorkspace_, negABase, lower, block, - solveWorkspace_, tmpBase, 0, 0, half, half, half); - CubeGemmSolveSub(solveWorkspace_, tmpBase, 0, 0, akk_, aiBase, block, block, - akk_, aiBase, lower, block, half, half, half); - } - } - } - - __aicore__ inline void ComputeAkkMergeCubeWorkspace(uint64_t b, uint64_t hv, uint64_t chunkIdx) - { - uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); - uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); - - for (uint32_t mergeSize = 2 * KDA_SOLVE_DIAG_BT; mergeSize <= BT_; mergeSize *= 2) { - uint32_t half = mergeSize / 2; - for (uint32_t block = 0; block < BT_; block += mergeSize) { - uint32_t lower = block + half; - CubeGemmSolveSub(solveWorkspace_, xBase, lower, lower, solveWorkspace_, xBase, lower, block, - solveWorkspace_, tmpBase, 0, 0, half, half, half); - CubeGemmSolveSub(solveWorkspace_, tmpBase, 0, 0, solveWorkspace_, xBase, block, block, - solveWorkspace_, xBase, lower, block, half, half, half); - } - } - } - - __aicore__ inline void ComputeAkkInverseMchFull(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start) - { - uint64_t aBase = AOffset(b, hv, start, 0); - uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); - uint64_t yBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0); - uint64_t yNextBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y1); - uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); - - uint32_t diagBlocks = static_cast(BT_ / KDA_SOLVE_DIAG_BT); - for (uint32_t block = 0; block < diagBlocks; ++block) { - uint32_t off = block * KDA_SOLVE_DIAG_BT; - CubeGemmSolveSub(akk_, aBase, off, off, akk_, aBase, off, off, solveWorkspace_, yBase, off, off, - KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); - } - for (uint32_t iter = 0; iter < KDA_SOLVE_DIAG_MCH_ITERS; ++iter) { - for (uint32_t block = 0; block < diagBlocks; ++block) { - uint32_t off = block * KDA_SOLVE_DIAG_BT; - CubeGemmSolveSub(solveWorkspace_, xBase, off, off, solveWorkspace_, yBase, off, off, - solveWorkspace_, tmpBase, off, off, - KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); - } - Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); - if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { - for (uint32_t block = 0; block < diagBlocks; ++block) { - uint32_t off = block * KDA_SOLVE_DIAG_BT; - CubeGemmSolveSub(solveWorkspace_, yBase, off, off, solveWorkspace_, yBase, off, off, - solveWorkspace_, yNextBase, off, off, - KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); - } - } - Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(syncReadyFlag_); - if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { - uint64_t oldYBase = yBase; - yBase = yNextBase; - yNextBase = oldYBase; - } - } - ComputeAkkMergeCube(b, hv, chunkIdx, start); - Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); - } - - __aicore__ inline void ComputeAkkInverseMchTail(uint64_t b, uint64_t hv, uint64_t chunkIdx, - uint64_t start, uint64_t curT) - { - uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); - uint64_t lBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0); - uint64_t yBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y1); - uint64_t yNextBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0); - uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); - (void)start; - (void)curT; - - uint32_t diagBlocks = static_cast(BT_ / KDA_SOLVE_DIAG_BT); - for (uint32_t block = 0; block < diagBlocks; ++block) { - uint32_t off = block * KDA_SOLVE_DIAG_BT; - CubeGemmSolveSub(solveWorkspace_, lBase, off, off, solveWorkspace_, lBase, off, off, - solveWorkspace_, yBase, off, off, - KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); - } - for (uint32_t iter = 0; iter < KDA_SOLVE_DIAG_MCH_ITERS; ++iter) { - for (uint32_t block = 0; block < diagBlocks; ++block) { - uint32_t off = block * KDA_SOLVE_DIAG_BT; - CubeGemmSolveSub(solveWorkspace_, xBase, off, off, solveWorkspace_, yBase, off, off, - solveWorkspace_, tmpBase, off, off, - KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); - } - Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); - if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { - for (uint32_t block = 0; block < diagBlocks; ++block) { - uint32_t off = block * KDA_SOLVE_DIAG_BT; - CubeGemmSolveSub(solveWorkspace_, yBase, off, off, solveWorkspace_, yBase, off, off, - solveWorkspace_, yNextBase, off, off, - KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); - } - } - Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(syncReadyFlag_); - if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { - uint64_t oldYBase = yBase; - yBase = yNextBase; - yNextBase = oldYBase; - } - } - ComputeAkkMergeCubeWorkspace(b, hv, chunkIdx); - Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); - } - - - - __aicore__ inline void ScaleRowsByBeta(GlobalTensor &src, GlobalTensor &dst, uint64_t b, uint64_t hv, - uint64_t start, uint64_t rowBegin, uint64_t rowCount, uint64_t dim, - LocalTensor &betaBrcb, LocalTensor &matrixLocal) - { - constexpr uint64_t vecElemsPerRepeat = 64; - constexpr uint64_t typedOffsetFloats = 20480; - constexpr uint64_t typedOffset = typedOffsetFloats * sizeof(float) / sizeof(T); - uint64_t elemCount = rowCount * dim; - uint64_t baseOffset = KVOffset(b, hv, start + rowBegin, 0, dim); - - if constexpr (IsSameType::value) { - DataCopy(matrixLocal, src[baseOffset], static_cast(elemCount)); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - } else { - LocalTensor matrixTyped = vecBuf_.Get()[typedOffset]; - DataCopy(matrixTyped, src[baseOffset], static_cast(elemCount)); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - Cast(matrixLocal, matrixTyped, RoundMode::CAST_NONE, static_cast(elemCount)); - PipeBarrier(); - } - - uint8_t repeatStride = static_cast(dim * sizeof(float) / 32); - for (uint64_t col = 0; col < dim; col += vecElemsPerRepeat) { - uint64_t mask = dim - col; - if (mask > vecElemsPerRepeat) { - mask = vecElemsPerRepeat; - } - Mul(matrixLocal[col], matrixLocal[col], betaBrcb, mask, static_cast(rowCount), - {1, 1, 0, repeatStride, repeatStride, 1}); - PipeBarrier(); - } - - if constexpr (IsSameType::value) { - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - DataCopy(dst[baseOffset], matrixLocal, static_cast(elemCount)); - } else { - LocalTensor matrixTyped = vecBuf_.Get()[typedOffset]; - Cast(matrixTyped, matrixLocal, RoundMode::CAST_RINT, static_cast(elemCount)); - PipeBarrier(); - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - DataCopy(dst[baseOffset], matrixTyped, static_cast(elemCount)); - } - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - } - - __aicore__ inline void PrepareWuCubeInputs(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT, - uint64_t subBlockIdx, uint64_t subBlockNum) - { - uint64_t rowsPerSubBlock = (curT + subBlockNum - 1) / subBlockNum; - uint64_t rowBegin = subBlockIdx * rowsPerSubBlock; - if (rowBegin >= curT) { - return; - } - uint64_t rowCount = curT - rowBegin; - if (rowCount > rowsPerSubBlock) { - rowCount = rowsPerSubBlock; - } - LocalTensor arena = vecBuf_.Get(); - LocalTensor betaLocal = arena; - LocalTensor betaBrcb = arena[KDA_SOLVE_BT]; - LocalTensor matrixLocal = arena[KDA_SOLVE_BT + 512]; - LoadAsFloatRow(beta_, BetaOffset(b, hv, start + rowBegin), betaLocal, rowCount); - Brcb(betaBrcb, betaLocal, static_cast((rowCount + 7) / 8), {1, 8}); - PipeBarrier(); - ScaleRowsByBeta(w_, w_, b, hv, start, rowBegin, rowCount, K_, betaBrcb, matrixLocal); - ScaleRowsByBeta(v_, vNew_, b, hv, start, rowBegin, rowCount, V_, betaBrcb, matrixLocal); - } - - __aicore__ inline void ComputePostWuCube(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, - uint64_t curT) - { - using ElementA = AKK_T; - using ElementB = T; - using LayoutTagA = Catlass::layout::RowMajor; - using LayoutTagB = Catlass::layout::RowMajor; - using LayoutTagC = Catlass::layout::RowMajor; - using WTileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; - using UTileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; - using PostL1TileShape128 = tla::Shape<_128, _128, tla::_256>; - using PostL0TileShape128 = tla::Shape<_128, _128, _128>; - using PostL1TileShape256 = tla::Shape<_128, tla::_256, tla::_256>; - using PostL0TileShape256 = tla::Shape<_128, tla::_256, _64>; - using WBlockMmad = Catlass::Gemm::Block::BlockMmadTla; - using UBlockMmad128 = Catlass::Gemm::Block::BlockMmadTla; - using UBlockMmad256 = Catlass::Gemm::Block::BlockMmadTla; - - LayoutTagA tagA = LayoutTagA::template MakeLayout(BT_, BT_); - auto layoutA = tla::MakeLayoutFromTag(tagA); - auto tensorA = tla::MakeTensor(stageAqk_[AOffset(b, hv, start, 0)], layoutA, - Catlass::Arch::PositionGM{}); - - { - LayoutTagB tagB = LayoutTagB::template MakeLayout(BT_, K_); - LayoutTagC tagC = LayoutTagC::template MakeLayout(BT_, K_); - auto layoutB = tla::MakeLayoutFromTag(tagB); - auto layoutC = tla::MakeLayoutFromTag(tagC); - Catlass::GemmCoord shape{static_cast(curT), static_cast(K_), - static_cast(curT)}; - auto tensorB = tla::MakeTensor(stageQG_[KVOffset(b, hv, start, 0, K_)], layoutB, - Catlass::Arch::PositionGM{}); - auto tensorC = tla::MakeTensor(h_[WScratchOffset(b, hv, chunkIdx, 0, 0)], layoutC, - Catlass::Arch::PositionGM{}); - auto blockA = GetTile(tensorA, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.k())); - auto blockB = GetTile(tensorB, tla::MakeCoord(0, 0), tla::MakeShape(shape.k(), shape.n())); - auto blockC = GetTile(tensorC, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.n())); - Catlass::Arch::Resource wResource; - WBlockMmad wBlockMmad(wResource); - wBlockMmad(blockA, blockB, blockC, shape); - PipeBarrier(); - } - - { - LayoutTagB tagB = LayoutTagB::template MakeLayout(BT_, V_); - LayoutTagC tagC = LayoutTagC::template MakeLayout(BT_, V_); - auto layoutB = tla::MakeLayoutFromTag(tagB); - auto layoutC = tla::MakeLayoutFromTag(tagC); - Catlass::GemmCoord shape{static_cast(curT), static_cast(V_), - static_cast(curT)}; - auto tensorB = tla::MakeTensor(stageVNew_[KVOffset(b, hv, start, 0, V_)], layoutB, - Catlass::Arch::PositionGM{}); - auto tensorC = tla::MakeTensor(u_[KVOffset(b, hv, start, 0, V_)], layoutC, - Catlass::Arch::PositionGM{}); - auto blockA = GetTile(tensorA, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.k())); - auto blockB = GetTile(tensorB, tla::MakeCoord(0, 0), tla::MakeShape(shape.k(), shape.n())); - auto blockC = GetTile(tensorC, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.n())); - Catlass::Arch::Resource uResource; - if (V_ <= 128) { - UBlockMmad128 uBlockMmad(uResource); - uBlockMmad(blockA, blockB, blockC, shape); - } else { - UBlockMmad256 uBlockMmad(uResource); - uBlockMmad(blockA, blockB, blockC, shape); - } - PipeBarrier(); - } - - } - - __aicore__ inline void CopyScratchWAndFinalizeKg(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, - uint64_t start, uint64_t curT, uint64_t subBlockIdx, - uint64_t subBlockNum) - { - constexpr uint64_t typedOffsetFloats = 20480; - constexpr uint64_t typedOffset = typedOffsetFloats * sizeof(float) / sizeof(T); - constexpr uint64_t kgFp32Planes = 4; - uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; - uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; - if (rowBegin >= rowEnd) { - return; - } - uint64_t maxRows = (typedOffsetFloats / kgFp32Planes) / K_; - if (maxRows > 32) { - maxRows = 32; - } - if (maxRows == 0) { - return; - } - - uint64_t last = start + curT - 1; - LocalTensor arena = vecBuf_.Get(); - LocalTensor gateLast = exp2Buf_.Get(); - LocalTensor typedLocal = vecBuf_.Get()[typedOffset]; - LoadAsFloatRow(gk_, KVOffset(b, hv, last, 0, K_), gateLast, K_); - - for (uint64_t tileRow = rowBegin; tileRow < rowEnd; tileRow += maxRows) { - uint64_t tileRows = rowEnd - tileRow; - if (tileRows > maxRows) { - tileRows = maxRows; - } - uint64_t elemCount = tileRows * K_; - uint64_t scratchBase = WScratchOffset(b, hv, chunkIdx, tileRow, 0); - uint64_t token = start + tileRow; - - DataCopy(arena, h_[scratchBase], static_cast(elemCount)); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - Cast(typedLocal, arena, RoundMode::CAST_RINT, static_cast(elemCount)); - PipeBarrier(); - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - DataCopy(w_[KVOffset(b, hv, token, 0, K_)], typedLocal, static_cast(elemCount)); - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - - LocalTensor kLocal = arena; - LocalTensor gLocal = arena[elemCount]; - LocalTensor expLocal = arena[2 * elemCount]; - LocalTensor outLocal = arena[3 * elemCount]; - CopyVectorIn(typedLocal, k_, QOffset(b, h, token, 0), elemCount); - CopyVectorIn(gLocal, gk_, KVOffset(b, hv, token, 0, K_), elemCount); - SetFlag(KDA_MTE2_V_EVENT_ID); - WaitFlag(KDA_MTE2_V_EVENT_ID); - Cast(kLocal, typedLocal, RoundMode::CAST_NONE, static_cast(elemCount)); - PipeBarrier(); - - for (uint64_t row = 0; row < tileRows; ++row) { - Sub(expLocal[row * K_], gateLast, gLocal[row * K_], static_cast(K_)); - } - PipeBarrier(); - Muls(expLocal, expLocal, LN2, static_cast(elemCount)); - PipeBarrier(); - ClampExpInput(expLocal, static_cast(elemCount)); - Exp(expLocal, expLocal, static_cast(elemCount)); - PipeBarrier(); - Mul(outLocal, kLocal, expLocal, static_cast(elemCount)); - PipeBarrier(); - ClampFp32ToOutputType(outLocal, static_cast(elemCount)); - Cast(typedLocal, outLocal, RoundMode::CAST_RINT, static_cast(elemCount)); - PipeBarrier(); - SetFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - WaitFlag(KDA_SCALAR_V_MTE3_EVENT_ID); - CopyVectorOut(kg_, KVOffset(b, hv, token, 0, K_), typedLocal, elemCount); - SetFlag(KDA_MTE3_MTE2_EVENT_ID); - WaitFlag(KDA_MTE3_MTE2_EVENT_ID); - } - SetFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - WaitFlag(KDA_SCALAR_MTE3_V_EVENT_ID); - } - - - - - - __aicore__ inline void ComputeOutputCube(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, - uint64_t curT) - { - using ElementA = T; - using ElementB = T; - using ElementC = OUT_T; - using LayoutTagA = Catlass::layout::RowMajor; - using LayoutTagB = Catlass::layout::RowMajor; - using LayoutTagC = Catlass::layout::RowMajor; - using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; - using BlockMmad = Catlass::Gemm::Block::BlockMmadTla; - - Catlass::Arch::Resource resource; - BlockMmad blockMmad(resource); - - auto layoutQ = tla::MakeLayout(BT_, K_); - auto layoutH = tla::MakeLayout(K_, V_); - auto layoutO = tla::MakeLayout(BT_, V_); - for (uint64_t nOffset = 0; nOffset < V_; nOffset += 128) { - uint32_t curN = static_cast((V_ - nOffset) > 128 ? 128 : (V_ - nOffset)); - auto tensorH = tla::MakeTensor(stageH_[HOffset(b, hv, chunkIdx, 0, nOffset)], layoutH, - Catlass::Arch::PositionGM{}); - for (uint64_t mOffset = 0; mOffset < curT; mOffset += 64) { - uint32_t curM = static_cast((curT - mOffset) > 64 ? 64 : (curT - mOffset)); - Catlass::GemmCoord shapeQH{curM, curN, static_cast(K_)}; - auto tensorQ = tla::MakeTensor(stageQG_[KVOffset(b, hv, start + mOffset, 0, K_)], layoutQ, - Catlass::Arch::PositionGM{}); - auto tensorO = tla::MakeTensor(o_[KVOffset(b, hv, start + mOffset, nOffset, V_)], layoutO, - Catlass::Arch::PositionGM{}); - auto blockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(shapeQH.m(), shapeQH.k())); - auto blockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(shapeQH.k(), shapeQH.n())); - auto blockO = GetTile(tensorO, tla::MakeCoord(0, 0), tla::MakeShape(shapeQH.m(), shapeQH.n())); - blockMmad(blockQ, blockH, blockO, shapeQH); - PipeBarrier(); - } - } - - auto layoutAqk = tla::MakeLayout(BT_, BT_); - auto layoutV = tla::MakeLayout(BT_, V_); - for (uint64_t nOffset = 0; nOffset < V_; nOffset += 128) { - uint32_t curN = static_cast((V_ - nOffset) > 128 ? 128 : (V_ - nOffset)); - auto tensorVNew = tla::MakeTensor(stageVNew_[KVOffset(b, hv, start, nOffset, V_)], layoutV, - Catlass::Arch::PositionGM{}); - for (uint64_t mOffset = 0; mOffset < curT; mOffset += 64) { - uint32_t curM = static_cast((curT - mOffset) > 64 ? 64 : (curT - mOffset)); - Catlass::GemmCoord shapeAV{curM, curN, static_cast(curT)}; - auto tensorAqk = tla::MakeTensor(stageAqk_[AOffset(b, hv, start + mOffset, 0)], layoutAqk, - Catlass::Arch::PositionGM{}); - auto tensorLocal = tla::MakeTensor(u_[KVOffset(b, hv, start + mOffset, nOffset, V_)], layoutO, - Catlass::Arch::PositionGM{}); - auto blockAqk = GetTile(tensorAqk, tla::MakeCoord(0, 0), tla::MakeShape(shapeAV.m(), shapeAV.k())); - auto blockVNew = GetTile(tensorVNew, tla::MakeCoord(0, 0), tla::MakeShape(shapeAV.k(), shapeAV.n())); - auto blockLocal = GetTile(tensorLocal, tla::MakeCoord(0, 0), tla::MakeShape(shapeAV.m(), shapeAV.n())); - blockMmad(blockAqk, blockVNew, blockLocal, shapeAV); - PipeBarrier(); - } - } - } - - __aicore__ inline void FinalizeOutputRows(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT, - uint64_t subBlockIdx, uint64_t subBlockNum) - { - LocalTensor stateLocal = VecScratch(0); - LocalTensor localLocal = VecScratch(1); - LocalTensor outLocal = VecScratch(2); - for (uint64_t i = subBlockIdx; i < curT; i += subBlockNum) { - uint64_t ti = start + i; - LoadAsFloatRow(o_, KVOffset(b, hv, ti, 0, V_), stateLocal, V_); - LoadAsFloatRow(u_, KVOffset(b, hv, ti, 0, V_), localLocal, V_); - Add(outLocal, stateLocal, localLocal, static_cast(V_)); - PipeBarrier(); - ClampFp32ToOutputType(outLocal, static_cast(V_)); - StoreFloatRow(vNew_, KVOffset(b, hv, ti, 0, V_), outLocal, V_); - } - } - - - __aicore__ inline bool ResolveFlatChunk(uint64_t task, uint64_t &seq, uint64_t &b, uint64_t &h, uint64_t &hv, - uint64_t &chunkIdx, uint64_t &start, uint64_t &end) - { - hv = task % HV_; - uint64_t flatChunk = task / HV_; - if (!isVarLen_) { - seq = flatChunk / NT_; - b = seq; - chunkIdx = flatChunk % NT_; - start = chunkIdx * BT_; - end = start + BT_; - if (end > T_) { - end = T_; - } - } else { - if (hasChunkIndices_) { - uint64_t low = 0; - uint64_t high = N_; - while (low + 1 < high) { - uint64_t mid = (low + high) >> 1; - if (flatChunk < static_cast(seqChunkOffset_[mid])) { - high = mid; - } else { - low = mid; - } - } - seq = low; - uint64_t localChunk = flatChunk - static_cast(seqChunkOffset_[seq]); - start = static_cast(seqStart_[seq]) + localChunk * BT_; - end = start + BT_; - uint64_t seqEnd = static_cast(seqEnd_[seq]); - if (end > seqEnd) { - end = seqEnd; - } - b = 0; - chunkIdx = flatChunk; - h = hv / (HV_ / H_); - return start < end; - } - return false; - } - h = hv / (HV_ / H_); - return start < end; - } - - __aicore__ inline void ProcessChunkPreAiv(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, - uint64_t start, uint64_t end, uint64_t subBlockIdx, - uint64_t subBlockNum) - { - if constexpr (IsSameType::value) { - ProcessChunkPreAivFp32(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); - } - } - - template - __aicore__ inline void RunAicAfterBothAivReady(uint64_t subBlockIdx, uint64_t subBlockNum) - { - if constexpr (CORE_TYPE == AscendC::AIV) { - (void)subBlockIdx; - (void)subBlockNum; - Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_MTE3>(syncReadyFlag_); - Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(syncDoneFlag_); - } - } - - __aicore__ inline void ProcessChunkPreAivFp32(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, - uint64_t start, uint64_t end, uint64_t subBlockIdx, - uint64_t subBlockNum) - { - uint64_t curT = end - start; - if (curT == 0) { - return; - } - if constexpr (IsSameType::value) { - return; - } - - if (K_ < 16) { - return; - } - bool usePostWuCube = UsePostWuCube(curT); - bool useAkkCubeSolve = UseAkkCubeSolve(curT); - uint64_t solveRowBegin = 0; - uint64_t solveRowEnd = 0; - GetSolveRowRange(BT_, subBlockIdx, subBlockNum, solveRowBegin, solveRowEnd); - uint64_t scoreBlockSize = ScoreRefBlockSize(); - uint64_t scoreBlockCount = (curT + scoreBlockSize - 1) / scoreBlockSize; - uint64_t pipelineBlockCount = - (scoreBlockCount + KDA_SCORE_QUEUE_DEPTH - 1) / KDA_SCORE_QUEUE_DEPTH * KDA_SCORE_QUEUE_DEPTH; - for (uint64_t block = 0; block < pipelineBlockCount; ++block) { - if (block < scoreBlockCount) { - uint64_t rowBegin = block * scoreBlockSize; - uint64_t rowCount = ScoreRowBlockCount(curT, rowBegin); - uint64_t refToken = ScoreRefToken(start, curT, rowBegin, rowCount); - PrepareGateProducts(b, h, hv, start, curT, subBlockIdx, subBlockNum, true, refToken, - rowBegin + rowCount, true, block % KDA_SCORE_QUEUE_DEPTH, - rowBegin, rowCount); - } - Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_MTE3>(scoreReadyFlag_); - if (block > 0) { - Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(scoreDoneFlag_); - } - } - if (pipelineBlockCount > 0) { - Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(scoreDoneFlag_); - } - PrepareGateProducts(b, h, hv, start, curT, subBlockIdx, subBlockNum); - if (useAkkCubeSolve) { - bool fullChunk = curT == BT_; - PrepareAqkAkkSolveInputRows(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd, - fullChunk, !fullChunk); - } - if (useAkkCubeSolve) { - bool fullChunk = curT == BT_; - uint32_t solveIters = KDA_SOLVE_DIAG_MCH_ITERS; - RunAicAfterBothAivReady(subBlockIdx, subBlockNum); - for (uint32_t iter = 0; iter < solveIters; ++iter) { - AddSolveTmpToXDiagRows(b, hv, chunkIdx, start, solveRowBegin, solveRowEnd, - fullChunk && iter + 1 == solveIters); - RunAicAfterBothAivReady(subBlockIdx, subBlockNum); - } - if (!fullChunk) { - StoreSolveXRowsToAkk(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd); - } - } - // Host validation guarantees every accepted shape has enough workspace for this cube path. - PrepareWuCubeInputs(b, hv, start, curT, subBlockIdx, subBlockNum); - } - - __aicore__ inline void ProcessChunkPreAic(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, - uint64_t end) - { - if constexpr (IsSameType::value) { - ProcessChunkPreAicFp32(b, hv, chunkIdx, start, end); - } - } - - __aicore__ inline void ProcessChunkPreAicFp32(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, - uint64_t end) - { - uint64_t curT = end - start; - if (curT == 0 || K_ < 16) { - return; - } - uint64_t scoreBlockSize = ScoreRefBlockSize(); - uint64_t scoreBlockCount = (curT + scoreBlockSize - 1) / scoreBlockSize; - uint64_t pipelineBlockCount = - (scoreBlockCount + KDA_SCORE_QUEUE_DEPTH - 1) / KDA_SCORE_QUEUE_DEPTH * KDA_SCORE_QUEUE_DEPTH; - for (uint64_t block = 0; block < pipelineBlockCount; ++block) { - Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(scoreReadyFlag_); - if (block < scoreBlockCount) { - uint64_t rowBegin = block * scoreBlockSize; - uint64_t rowCount = ScoreRowBlockCount(curT, rowBegin); - ComputeRawAqkAkkCubeBlock(b, hv, start, curT, rowBegin, rowCount, true, - block % KDA_SCORE_QUEUE_DEPTH, rowBegin + rowCount); - } - Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(scoreDoneFlag_); - } - bool usePostWuCube = UsePostWuCube(curT); - bool useAkkCubeSolve = UseAkkCubeSolve(curT); - if (useAkkCubeSolve) { - Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(syncReadyFlag_); - if (curT == BT_) { - ComputeAkkInverseMchFull(b, hv, chunkIdx, start); - } else { - ComputeAkkInverseMchTail(b, hv, chunkIdx, start, curT); - } - } - (void)usePostWuCube; - (void)chunkIdx; - } - - __aicore__ inline void ProcessChunkPostAiv(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, - uint64_t start, uint64_t end, uint64_t subBlockIdx, - uint64_t subBlockNum) - { - uint64_t curT = end - start; - if (curT == 0 || !UsePostWuCube(curT)) { - return; - } - Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(syncDoneFlag_); - CopyScratchWAndFinalizeKg(b, h, hv, chunkIdx, start, curT, subBlockIdx, subBlockNum); - } - - __aicore__ inline void ProcessChunkPostAic(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, - uint64_t end) - { - if constexpr (IsSameType::value) { - ProcessChunkPostAicTyped(b, hv, chunkIdx, start, end); - } - } - - __aicore__ inline void ProcessChunkPostAicTyped(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, - uint64_t end) - { - uint64_t curT = end - start; - if (curT == 0 || !UsePostWuCube(curT)) { - return; - } - ComputePostWuCube(b, hv, chunkIdx, start, curT); - Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); - } - - __aicore__ inline void ProcessChunkOutAiv(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, - uint64_t end, uint64_t subBlockIdx, uint64_t subBlockNum) - { - uint64_t curT = end - start; - if (curT == 0) { - return; - } - if constexpr (IsSameType::value) { - return; - } - Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(syncDoneFlag_); - FinalizeOutputRows(b, hv, start, curT, subBlockIdx, subBlockNum); - } - - __aicore__ inline void ProcessChunkOutAic(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, - uint64_t end) - { - uint64_t curT = end - start; - if (curT == 0) { - return; - } - ComputeOutputCube(b, hv, chunkIdx, start, curT); - Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); - } - - __aicore__ inline void ProcessPreAiv() - { - if constexpr (IsSameType::value) { - isAivOnly_ = true; - } - uint64_t subBlockNum = isAivOnly_ ? 1 : static_cast(GetSubBlockNum()); - if (subBlockNum == 0) { - return; - } - uint64_t subBlockIdx = isAivOnly_ ? 0 : static_cast(GetSubBlockIdx()); - uint64_t coreNum = isAivOnly_ ? static_cast(GetBlockNum()) : usedCoreNum_; - uint64_t coreIdx = isAivOnly_ ? static_cast(GetBlockIdx()) : - static_cast(GetBlockIdx()) / subBlockNum; - uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); - for (uint64_t task = coreIdx; task < taskNum; task += coreNum) { - uint64_t seq = 0; - uint64_t b = 0; - uint64_t h = 0; - uint64_t hv = 0; - uint64_t chunkIdx = 0; - uint64_t start = 0; - uint64_t end = 0; - if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { - (void)seq; - ProcessChunkPreAiv(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); - } - } - } - - __aicore__ inline void ProcessPreAic() - { - if constexpr (IsSameType::value) { - return; - } - uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); - uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; - for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum) { - uint64_t seq = 0; - uint64_t b = 0; - uint64_t h = 0; - uint64_t hv = 0; - uint64_t chunkIdx = 0; - uint64_t start = 0; - uint64_t end = 0; - if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { - (void)seq; - (void)h; - ProcessChunkPreAic(b, hv, chunkIdx, start, end); - } - } - } - - __aicore__ inline void ProcessPostAiv() - { - if constexpr (IsSameType::value) { - return; - } - uint64_t subBlockNum = static_cast(GetSubBlockNum()); - if (subBlockNum == 0) { - return; - } - uint64_t subBlockIdx = static_cast(GetSubBlockIdx()); - uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; - uint64_t coreIdx = static_cast(GetBlockIdx()) / subBlockNum; - uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); - for (uint64_t task = coreIdx; task < taskNum; task += coreNum) { - uint64_t seq = 0; - uint64_t b = 0; - uint64_t h = 0; - uint64_t hv = 0; - uint64_t chunkIdx = 0; - uint64_t start = 0; - uint64_t end = 0; - if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { - (void)seq; - ProcessChunkPostAiv(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); - } - } - } - - __aicore__ inline void ProcessPostAic() - { - if constexpr (IsSameType::value) { - return; - } - uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); - uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; - for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum) { - uint64_t seq = 0; - uint64_t b = 0; - uint64_t h = 0; - uint64_t hv = 0; - uint64_t chunkIdx = 0; - uint64_t start = 0; - uint64_t end = 0; - if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { - (void)seq; - (void)h; - ProcessChunkPostAic(b, hv, chunkIdx, start, end); - } - } - } - - __aicore__ inline void ProcessOutAiv() - { - if constexpr (IsSameType::value) { - return; - } - uint64_t subBlockNum = static_cast(GetSubBlockNum()); - if (subBlockNum == 0) { - return; - } - uint64_t subBlockIdx = static_cast(GetSubBlockIdx()); - uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; - uint64_t coreIdx = static_cast(GetBlockIdx()) / subBlockNum; - uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); - for (uint64_t task = coreIdx; task < taskNum; task += coreNum) { - uint64_t seq = 0; - uint64_t b = 0; - uint64_t h = 0; - uint64_t hv = 0; - uint64_t chunkIdx = 0; - uint64_t start = 0; - uint64_t end = 0; - if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { - (void)seq; - (void)h; - (void)chunkIdx; - ProcessChunkOutAiv(b, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); - } - } - } - - __aicore__ inline void ProcessOutAic() - { - if constexpr (IsSameType::value) { - return; - } - uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); - uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; - for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum) { - uint64_t seq = 0; - uint64_t b = 0; - uint64_t h = 0; - uint64_t hv = 0; - uint64_t chunkIdx = 0; - uint64_t start = 0; - uint64_t end = 0; - if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { - (void)seq; - (void)h; - ProcessChunkOutAic(b, hv, chunkIdx, start, end); - } - } +template +__aicore__ inline void DispatchGenericSafeGate( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR g, GM_ADDR beta, + GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR attnOut, + GM_ADDR finalState, GM_ADDR gk, GM_ADDR aqk, GM_ADDR akk, + GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR userWorkspace, const TilingData &tiling) +{ + if (tiling.safeGate) { + DispatchGeneric( + q, k, v, g, beta, aLog, dtBias, initialState, cuSeqlens, + chunkIndices, attnOut, finalState, gk, aqk, akk, w, u, qg, + kg, vNew, h, userWorkspace, tiling); + } else { + DispatchGeneric( + q, k, v, g, beta, aLog, dtBias, initialState, cuSeqlens, + chunkIndices, attnOut, finalState, gk, aqk, akk, w, u, qg, + kg, vNew, h, userWorkspace, tiling); } +} +} // namespace KdaForward - - -private: - GlobalTensor q_; - GlobalTensor k_; - GlobalTensor v_; - GlobalTensor gk_; - GlobalTensor beta_; - GlobalTensor initialState_; - GlobalTensor cuSeqlens_; - GlobalTensor o_; - GlobalTensor finalState_; - GlobalTensor aqk_; - GlobalTensor akk_; - GlobalTensor w_; - GlobalTensor u_; - GlobalTensor qg_; - GlobalTensor kg_; - GlobalTensor vNew_; - GlobalTensor h_; - GlobalTensor stageQG_; - GlobalTensor stageAqk_; - GlobalTensor stageVNew_; - GlobalTensor stageH_; - GlobalTensor solveWorkspace_; - GlobalTensor scoreWorkspace_; - TPipe *pipe_ = nullptr; - TBuf exp2Buf_; - TBuf vecBuf_; - TQue qInQue_; - TQue kInQue_; - TQue gInQue_; - TQue qgOutQue_; - TQue wOutQue_; - TQue kgOutQue_; - Catlass::Arch::CrossCoreFlagWithReverse scoreReadyFlag_{KDA_SCORE_READY_FLAG0, - KDA_SCORE_READY_FLAG1}; - Catlass::Arch::CrossCoreFlagWithReverse scoreDoneFlag_{KDA_SCORE_DONE_FLAG0, - KDA_SCORE_DONE_FLAG1}; - Catlass::Arch::CrossCoreFlagWithReverse syncReadyFlag_{KDA_SCORE_READY_FLAG0, - KDA_SCORE_READY_FLAG1}; - Catlass::Arch::CrossCoreFlagWithReverse syncDoneFlag_{KDA_SCORE_DONE_FLAG0, - KDA_SCORE_DONE_FLAG1}; - uint64_t B_ = 0; - uint64_t N_ = 0; - uint64_t H_ = 0; - uint64_t HV_ = 0; - uint64_t T_ = 0; - uint64_t K_ = 0; - uint64_t V_ = 0; - uint64_t BT_ = 0; - uint64_t NT_ = 0; - float scale_ = 1.0f; - bool hasInitial_ = false; - bool isVarLen_ = false; - bool hasChunkIndices_ = false; - bool isAivOnly_ = false; - uint64_t usedCoreNum_ = 1; - uint64_t solveCoreIdx_ = 0; - int64_t stage_ = 0; - const int64_t *seqStart_ = nullptr; - const int64_t *seqEnd_ = nullptr; - const int64_t *seqChunkOffset_ = nullptr; -}; -} // namespace - -extern "C" __global__ __aicore__ void chunk_kda_fwd(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, - GM_ADDR initial_state, GM_ADDR cu_seqlens, - GM_ADDR chunk_indices, GM_ADDR stage_qg, GM_ADDR stage_aqk, - GM_ADDR stage_v_new, GM_ADDR stage_h, GM_ADDR o, - GM_ADDR final_state, GM_ADDR aqk, GM_ADDR akk, GM_ADDR w, - GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR v_new, GM_ADDR h, - GM_ADDR workspace, GM_ADDR tiling) +extern "C" __global__ __aicore__ void chunk_kda_fwd( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR g, GM_ADDR beta, + GM_ADDR a_log, GM_ADDR dt_bias, GM_ADDR initial_state, + GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR attn_out, + GM_ADDR final_state, GM_ADDR gk, GM_ADDR aqk, GM_ADDR akk, + GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR v_new, GM_ADDR h, + GM_ADDR qg_scaled, GM_ADDR u_seed, GM_ADDR workspace, GM_ADDR tiling) { - GM_ADDR userWS = AscendC::GetUserWorkspace(workspace); - (void)userWS; - GET_TILING_DATA(tilingData, tiling); - TPipe pipe; - if (TILING_KEY_IS(0)) { - KERNEL_TASK_TYPE(0, KERNEL_TYPE_AIV_ONLY); - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, stage_v_new, - stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, &pipe); - op.ProcessAivOnly(); - } else if (TILING_KEY_IS(1)) { - KERNEL_TASK_TYPE(1, KERNEL_TYPE_MIX_AIC_1_2); - if (tilingData.dataType == 1) { - if ASCEND_IS_AIC { - if (tilingData.stage == 2) { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, - stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, - &pipe, false); - op.ProcessAic(); - } else if (tilingData.stage == 3) { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, - stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, - &pipe, false); - op.ProcessAic(); - } else { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, - stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, - &pipe, false); - op.ProcessAic(); - } - } - if ASCEND_IS_AIV { - if (tilingData.stage == 2) { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, - stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, - &pipe); - op.ProcessAiv(); - } else if (tilingData.stage == 3) { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, - stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, - &pipe); - op.ProcessAiv(); - } else { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, - stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, - &pipe); - op.ProcessAiv(); - } - } - } else { - if ASCEND_IS_AIC { - if (tilingData.stage == 2) { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, - stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, - &pipe, false); - op.ProcessAic(); - } else if (tilingData.stage == 3) { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, - stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, - &pipe, false); - op.ProcessAic(); - } else { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, - stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, - &pipe, false); - op.ProcessAic(); - } - } - if ASCEND_IS_AIV { - if (tilingData.stage == 2) { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, - stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, - &pipe); - op.ProcessAiv(); - } else if (tilingData.stage == 3) { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, - stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, - &pipe); - op.ProcessAiv(); - } else { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, - stage_v_new, stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, - &pipe); - op.ProcessAiv(); - } - } - } + KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2); + KERNEL_TASK_TYPE(1, KERNEL_TYPE_MIX_AIC_1_2); + KERNEL_TASK_TYPE(2, KERNEL_TYPE_MIX_AIC_1_2); + GM_ADDR userWorkspace = AscendC::GetUserWorkspace(workspace); + GET_TILING_DATA_WITH_STRUCT(ChunkKdaFwdTilingData, tilingData, tiling); + if (TILING_KEY_IS(1)) { + if (tilingData.stage != KdaForward::KDA_STAGE_FULL) { + KdaForward::DispatchStageSafeGate( + q, k, v, g, beta, a_log, dt_bias, initial_state, cu_seqlens, + chunk_indices, attn_out, final_state, gk, aqk, akk, w, u, + qg, kg, v_new, h, qg_scaled, u_seed, userWorkspace, + tilingData); + return; + } + KdaForward::DispatchGenericSafeGate( + q, k, v, g, beta, a_log, dt_bias, initial_state, cu_seqlens, + chunk_indices, attn_out, final_state, gk, aqk, akk, w, u, qg, + kg, v_new, h, userWorkspace, tilingData); } else if (TILING_KEY_IS(2)) { - KERNEL_TASK_TYPE(2, KERNEL_TYPE_AIV_ONLY); - if (tilingData.dataType == 1) { - if (tilingData.stage == 3) { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, stage_v_new, - stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, &pipe); - op.ProcessAivOnly(); - } else { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, stage_v_new, - stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, &pipe); - op.ProcessAivOnly(); - } - } else { - if (tilingData.stage == 3) { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, stage_v_new, - stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, &pipe); - op.ProcessAivOnly(); - } else { - ChunkKdaFwdKernel op; - op.Init(q, k, v, gk, beta, initial_state, cu_seqlens, chunk_indices, stage_qg, stage_aqk, stage_v_new, - stage_h, o, final_state, aqk, akk, w, u, qg, kg, v_new, h, userWS, tilingData, &pipe); - op.ProcessAivOnly(); - } + if (tilingData.stage != KdaForward::KDA_STAGE_FULL) { + KdaForward::DispatchStageSafeGate( + q, k, v, g, beta, a_log, dt_bias, initial_state, cu_seqlens, + chunk_indices, attn_out, final_state, gk, aqk, akk, w, u, + qg, kg, v_new, h, qg_scaled, u_seed, userWorkspace, + tilingData); + return; } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + KdaForward::DispatchArch35SafeGate( +#else + KdaForward::DispatchGenericSafeGate( +#endif + q, k, v, g, beta, a_log, dt_bias, initial_state, cu_seqlens, + chunk_indices, attn_out, final_state, gk, aqk, akk, w, u, qg, + kg, v_new, h, userWorkspace, tilingData); } } + +#undef KDA_COMPILE_ARCH35_FAST_PATH diff --git a/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_common.h b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_common.h new file mode 100644 index 000000000000..d33feb0076e7 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_common.h @@ -0,0 +1,332 @@ +#pragma once + +#include "kernel_operator.h" +#include "../../kda_gate_cumsum/op_kernel/kda_gate_cumsum_kernel.h" +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +#include "arch35/chunk_kda_fwd_prepare.h" +#include "arch35/chunk_kda_fwd_post_wu.h" +#include "arch35/chunk_kda_fwd_finalize.h" +#else +#include "chunk_kda_fwd_prepare.h" +#include "chunk_kda_fwd_post_wu.h" +#include "chunk_kda_fwd_finalize.h" +#endif + +#if __has_include("../../../gdn/chunk_gdn_fwd/chunk_gated_delta_rule_fwd_h/op_kernel/chunk_gated_delta_rule_fwd_h_struct.h") +#include "../../../gdn/chunk_gdn_fwd/chunk_gated_delta_rule_fwd_h/op_kernel/chunk_gated_delta_rule_fwd_h_struct.h" +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +#include "../../../gdn/chunk_gdn_fwd/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/kernel/gdn_fwd_h_kernel.hpp" +#else +#include "../../../gdn/chunk_gdn_fwd/chunk_gated_delta_rule_fwd_h/op_kernel/gemm/kernel/gdn_fwd_h_kernel.hpp" +#endif +#else +#include "../../chunk_gated_delta_rule_fwd_h/op_kernel/chunk_gated_delta_rule_fwd_h_struct.h" +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +#include "../../chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/kernel/gdn_fwd_h_kernel.hpp" +#else +#include "../../chunk_gated_delta_rule_fwd_h/op_kernel/gemm/kernel/gdn_fwd_h_kernel.hpp" +#endif +#endif + +namespace KdaForward { + +using namespace AscendC; + +struct GateRuntimeTiling { + int64_t batch; + int64_t t; + int64_t hv; + int64_t k; + int64_t rank; + int64_t chunkSize; + int64_t seqNum; + int64_t hasCuSeqlens; + int64_t hasALog; + int64_t hasDtBias; + int64_t dataType; + int64_t useGateInKernel; + int64_t safeGate; + int64_t inputSequenceMajor; + float lowerBound; + int64_t usedCoreNum; +}; + +struct ChunkKdaFwdAddresses { + GM_ADDR gk; + GM_ADDR finalState; + GM_ADDR w; + GM_ADDR u; + GM_ADDR qg; + GM_ADDR kg; + GM_ADDR vNew; + GM_ADDR h; + GM_ADDR qgScaled; + GM_ADDR uSeed; +}; + +struct FwdHTilingView { + int64_t batch; + int64_t seqlen; + int64_t kNumHead; + int64_t vNumHead; + int64_t kHeadDim; + int64_t vHeadDim; + int64_t chunkSize; + bool useInitialState; + bool storeFinalState; + int64_t isVariedLen; + int64_t shapeBatch; + int64_t tokenBatch; + int64_t vWorkspaceOffset; + int64_t vUpdateWorkspaceOffset; + int64_t kDecayWorkspaceOffset; + int64_t hWorkspaceOffset; + int64_t numSeqWorkspaceOffset; + int64_t numChunksWorkspaceOffset; +}; + +template +__aicore__ inline FwdHTilingView MakeFwdHTiling(const TilingData &tiling) +{ + return { + tiling.isVarLen ? tiling.seqNum : tiling.batch, + tiling.seqlen, + tiling.vHeadNum, + tiling.vHeadNum, + tiling.kHeadDim, + tiling.vHeadDim, + tiling.chunkSize, + tiling.hasInitialState, + tiling.storeFinalState, + tiling.isVarLen ? 1 : 0, + tiling.isVarLen ? 1 : tiling.batch, + tiling.isVarLen ? tiling.seqNum : 1, + tiling.vWorkspaceOffset, + tiling.vUpdateWorkspaceOffset, + tiling.kDecayWorkspaceOffset, + tiling.hWorkspaceOffset, + tiling.numSeqWorkspaceOffset, + tiling.numChunksWorkspaceOffset, + }; +} + +__aicore__ inline GM_ADDR ResolveStorage( + GM_ADDR output, GM_ADDR userWorkspace, int64_t offset, bool storeOutput) +{ + return storeOutput ? output : userWorkspace + offset; +} + +template +__aicore__ inline ChunkKdaFwdAddresses ResolveAddresses( + GM_ADDR finalState, GM_ADDR gk, GM_ADDR w, GM_ADDR u, GM_ADDR qg, + GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, GM_ADDR userWorkspace, + const TilingData &tiling) +{ + return { + ResolveStorage(gk, userWorkspace, tiling.gkStorageOffset, tiling.storeGk), + ResolveStorage(finalState, userWorkspace, tiling.finalStateStorageOffset, + tiling.storeFinalState), + ResolveStorage(w, userWorkspace, tiling.wStorageOffset, tiling.storeW), + ResolveStorage(u, userWorkspace, tiling.uStorageOffset, tiling.storeU), + ResolveStorage(qg, userWorkspace, tiling.qgStorageOffset, tiling.storeQG), + ResolveStorage(kg, userWorkspace, tiling.kgStorageOffset, tiling.storeKg), + ResolveStorage(vNew, userWorkspace, tiling.vNewStorageOffset, tiling.storeVNew), + ResolveStorage(h, userWorkspace, tiling.hStorageOffset, tiling.storeH), + userWorkspace + tiling.qgScaledOffset, + userWorkspace + tiling.outputScratchOffset, + }; +} + +template +__aicore__ inline GateRuntimeTiling MakeGateTiling(const TilingData &tiling) +{ + return { + tiling.batch, + tiling.seqlen, + tiling.vHeadNum, + tiling.kHeadDim, + tiling.inputRank, + tiling.chunkSize, + tiling.seqNum, + tiling.isVarLen ? 1 : 0, + tiling.hasALog ? 1 : 0, + tiling.hasDtBias ? 1 : 0, + tiling.gateDataType, + tiling.useGateInKernel ? 1 : 0, + tiling.safeGate ? 1 : 0, + tiling.inputSequenceMajor ? 1 : 0, + tiling.lowerBound, + tiling.gateUsedCoreNum, + }; +} + +template +__aicore__ inline void RunGateCumsum( + GM_ADDR g, GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR cuSeqlens, + GM_ADDR gk, const TilingData &tiling) +{ + if (tiling.computeGateInPrepare) { + return; + } + if ASCEND_IS_AIV { + GateRuntimeTiling gateTiling = MakeGateTiling(tiling); + TPipe gatePipe; + if (gateTiling.dataType == 2) { + KdaGateCumsum::DispatchKdaGateCumsum( + g, aLog, dtBias, cuSeqlens, gk, gateTiling, &gatePipe); + } else if (gateTiling.dataType == 1) { + KdaGateCumsum::DispatchKdaGateCumsum( + g, aLog, dtBias, cuSeqlens, gk, gateTiling, &gatePipe); + } else { + KdaGateCumsum::DispatchKdaGateCumsum( + g, aLog, dtBias, cuSeqlens, gk, gateTiling, &gatePipe); + } + } +} + +template +__aicore__ inline void RunFrontEnd( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR g, GM_ADDR beta, + GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR aqk, GM_ADDR akk, + const ChunkKdaFwdAddresses &addresses, GM_ADDR userWorkspace, + const TilingData &tiling, TPipe &pipe) +{ + RunGateCumsum(g, aLog, dtBias, cuSeqlens, addresses.gk, tiling); + if (!tiling.computeGateInPrepare) { + SyncAll(); + } + GM_ADDR uSeed = (tiling.fusePostWu || tiling.fusePostWuIntoFwdH) + ? addresses.u + : addresses.uSeed; + + KdaPrepare::RunChunkKdaPrepare( + q, k, v, addresses.gk, g, aLog, dtBias, beta, initialState, + cuSeqlens, chunkIndices, aqk, akk, addresses.qg, + addresses.qgScaled, addresses.w, uSeed, addresses.kg, + userWorkspace, tiling, pipe, tiling.storeQG); + SyncAll(); + pipe.Reset(); + + if (!tiling.fusePostWu && !tiling.fusePostWuIntoFwdH) { + KdaPostWu::RunChunkKdaPostWu( + q, k, v, addresses.gk, beta, initialState, cuSeqlens, + chunkIndices, addresses.w, akk, uSeed, + addresses.w, addresses.u, addresses.kg, addresses.vNew, + userWorkspace, tiling, pipe); + SyncAll(); + pipe.Reset(); + } +} + +template +__aicore__ inline void RunFwdH( + GM_ADDR initialState, GM_ADDR cuSeqlens, GM_ADDR chunkIndices, + const ChunkKdaFwdAddresses &addresses, GM_ADDR userWorkspace, + const TilingData &tiling) +{ + using FwdHKernel = Catlass::Gemm::Kernel::GDNFwdHKernel< + T, float, float, float, TileShapes, true, false, true>; + const auto fwdHTiling = MakeFwdHTiling(tiling); + FwdHKernel stateOp; + stateOp.InitFromData( + addresses.kg, addresses.w, addresses.u, addresses.gk, addresses.gk, + initialState, cuSeqlens, chunkIndices, addresses.h, addresses.vNew, + addresses.finalState, fwdHTiling, + userWorkspace + tiling.fwdHWorkspaceBaseOffset); + stateOp.Process(); +} + +template +__aicore__ inline void RunGenericBackEnd( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR aqk, GM_ADDR attnOut, + const ChunkKdaFwdAddresses &addresses, GM_ADDR userWorkspace, + const TilingData &tiling) +{ + if (tiling.vHeadDim > 128) { + RunFwdH( + initialState, cuSeqlens, chunkIndices, addresses, + userWorkspace, tiling); + } else { + RunFwdH( + initialState, cuSeqlens, chunkIndices, addresses, + userWorkspace, tiling); + } + SyncAll(); + TPipe pipe; + KdaFinalize::RunChunkKdaOutput( + q, k, v, addresses.gk, beta, initialState, cuSeqlens, + chunkIndices, addresses.qgScaled, aqk, + addresses.vNew, addresses.h, attnOut, userWorkspace, tiling, pipe); +} + +template +__aicore__ inline void RunGenericBackEnd( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR aqk, GM_ADDR attnOut, + const ChunkKdaFwdAddresses &addresses, GM_ADDR userWorkspace, + const TilingData &tiling, TPipe &pipe) +{ + if (tiling.vHeadDim > 128) { + RunFwdH( + initialState, cuSeqlens, chunkIndices, addresses, + userWorkspace, tiling); + } else { + RunFwdH( + initialState, cuSeqlens, chunkIndices, addresses, + userWorkspace, tiling); + } + SyncAll(); + KdaFinalize::RunChunkKdaOutput( + q, k, v, addresses.gk, beta, initialState, cuSeqlens, + chunkIndices, addresses.qgScaled, aqk, + addresses.vNew, addresses.h, attnOut, userWorkspace, tiling, pipe); +} + +template +__aicore__ inline void RunGeneric( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR g, GM_ADDR beta, + GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR attnOut, + GM_ADDR finalState, GM_ADDR gk, GM_ADDR aqk, GM_ADDR akk, + GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR userWorkspace, const TilingData &tiling) +{ + const auto addresses = ResolveAddresses( + finalState, gk, w, u, qg, kg, vNew, h, userWorkspace, tiling); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + TPipe pipe; + RunFrontEnd( + q, k, v, g, beta, aLog, dtBias, initialState, cuSeqlens, + chunkIndices, aqk, akk, addresses, userWorkspace, tiling, pipe); + if (!tiling.isVarLen && tiling.seqlen % tiling.chunkSize == 0) { + pipe.Destroy(); + RunGenericBackEnd( + q, k, v, beta, initialState, cuSeqlens, chunkIndices, aqk, + attnOut, addresses, userWorkspace, tiling); + } else { + RunGenericBackEnd( + q, k, v, beta, initialState, cuSeqlens, chunkIndices, aqk, + attnOut, addresses, userWorkspace, tiling, pipe); + } +#else + { + TPipe pipe; + RunFrontEnd( + q, k, v, g, beta, aLog, dtBias, initialState, cuSeqlens, + chunkIndices, aqk, akk, addresses, userWorkspace, tiling, pipe); + } + RunGenericBackEnd( + q, k, v, beta, initialState, cuSeqlens, chunkIndices, aqk, + attnOut, addresses, userWorkspace, tiling); +#endif +} + +} // namespace KdaForward diff --git a/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_finalize.h b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_finalize.h new file mode 100644 index 000000000000..831be9868b84 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_finalize.h @@ -0,0 +1,847 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#pragma once + +#ifndef CATLASS_ARCH +#define CATLASS_ARCH 2201 +#endif + +#include "catlass/arch/arch.hpp" +#include "catlass/arch/cross_core_sync.hpp" +#include "catlass/arch/resource.hpp" +#include "catlass/catlass.hpp" +#include "catlass/gemm/block/block_mmad.hpp" +#include "catlass/gemm/dispatch_policy.hpp" +#include "catlass/gemm/tile/tile_copy.hpp" +#include "catlass/gemm_coord.hpp" +#include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp" +#include "catlass/layout/layout.hpp" +#include "kernel_operator.h" +#include "chunk_kda_fwd_varlen.h" +#include "tla/layout.hpp" +#include "tla/tensor.hpp" + +using namespace AscendC; + +namespace KdaFinalize { +namespace { +using KdaInt64 = tla::Int<64>; +using KdaInt128 = tla::Int<128>; +constexpr float LN2 = 0.69314718055994530942f; +constexpr float KDA_EXP2_CLAMP = 80.0f; +constexpr float KDA_EXP_INPUT_MAX = KDA_EXP2_CLAMP * LN2; +constexpr float KDA_EXP_INPUT_MIN = -KDA_EXP2_CLAMP * LN2; +constexpr float KDA_FP16_MAX = 65504.0f; +constexpr uint32_t EXP2_UB_ELEMENTS = 256; +constexpr uint32_t EXP2_UB_BYTES = EXP2_UB_ELEMENTS * (sizeof(float) + sizeof(uint16_t)); +constexpr uint32_t EXP2_EVENT_ID = 0; +constexpr uint32_t KDA_SOLVE_BT = 64; +constexpr uint32_t KDA_SOLVE_MATRIX_ELEMENTS = KDA_SOLVE_BT * KDA_SOLVE_BT; +constexpr uint32_t KDA_SOLVE_SCRATCH_X = 0; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y0 = 1; +constexpr uint32_t KDA_SOLVE_SCRATCH_TMP = 2; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y1 = 3; +constexpr uint32_t KDA_SOLVE_SCRATCH_IDENTITY = 4; +constexpr uint32_t KDA_SOLVE_SCRATCH_SLOTS = 5; +constexpr uint32_t KDA_SOLVE_DIAG_BT = 16; +constexpr uint32_t KDA_SOLVE_DIAG_BLOCKS = KDA_SOLVE_BT / KDA_SOLVE_DIAG_BT; +constexpr uint32_t KDA_SOLVE_DIAG_MCH_ITERS = 3; +constexpr uint32_t KDA_SCORE_REF_BC = 16; +constexpr uint32_t KDA_VEC_ARENA_ELEMENTS = 32768; +constexpr uint32_t KDA_BITS_PER_MASK_BYTE = 8; +constexpr uint32_t KDA_SELECT_COL_BLOCKS = 2; +constexpr uint32_t KDA_SELECT_COL_MASK_BYTES = KDA_SOLVE_MATRIX_ELEMENTS / KDA_BITS_PER_MASK_BYTE; +constexpr uint32_t KDA_SELECT_MASK_BYTES = KDA_SELECT_COL_BLOCKS * KDA_SELECT_COL_MASK_BYTES; +constexpr uint32_t KDA_SELECT_AQK_MASK_BYTE_OFFSET = 120 * 1024; +constexpr uint32_t KDA_SELECT_AKK_MASK_BYTE_OFFSET = KDA_SELECT_AQK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_BYTE_OFFSET = KDA_SELECT_AKK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_FLOAT_OFFSET = KDA_SELECT_ZERO_BYTE_OFFSET / sizeof(float); +constexpr uint8_t KDA_SCORE_DONE_FLAG0 = 2; +constexpr uint8_t KDA_SCORE_DONE_FLAG1 = 3; +constexpr uint8_t KDA_SCORE_READY_FLAG0 = 4; +constexpr uint8_t KDA_SCORE_READY_FLAG1 = 5; +constexpr uint32_t KDA_SCORE_QUEUE_DEPTH = 2; +constexpr uint32_t KDA_SYNC_REVERSE_DEPTH = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_PLANES = 3; +constexpr uint32_t KDA_SCORE_SCRATCH_QG = 0; +constexpr uint32_t KDA_SCORE_SCRATCH_W = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_KG = 2; +constexpr uint64_t KDA_WORKSPACE_ALIGN = 512; +constexpr uint32_t KDA_GATE_TILE_ROWS = 32; + +using KdaArchTag = Catlass::Arch::AtlasA2; +using KdaDispatchPolicy = Catlass::Gemm::MmadPingpong; +using KdaScoreDispatchPolicy = + Catlass::Gemm::MmadPingpongTlaMulti; +static_assert(KdaScoreDispatchPolicy::ENABLE_L1_RESIDENT, + "KDA Aqk/Akk score MMAD must keep the shared right matrix resident in L1"); +static_assert(KdaScoreDispatchPolicy::L1B_STAGES == 1, + "KDA Aqk/Akk score MMAD needs one L1 B slot so the second MMAD reuses it"); +using KdaSolveDispatchPolicy = Catlass::Gemm::MmadPingpong; +static_assert(!KdaSolveDispatchPolicy::USE_HF32_MODE, "KDA triangular solve must use IEEE FP32 Cube mode"); +using KdaL1TileShape = tla::Shape; +using KdaL0TileShape = KdaL1TileShape; +using KdaSolveL1TileShape = tla::Shape; +using KdaSolveL0TileShape = KdaSolveL1TileShape; + +__aicore__ inline uint32_t FloatToBits(float value) +{ + union Bits { + __aicore__ Bits() {} + float f; + uint32_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline float BitsToFloat(uint32_t value) +{ + union Bits { + __aicore__ Bits() {} + uint32_t u; + float f; + } bits; + bits.u = value; + return bits.f; +} + +__aicore__ inline uint16_t Bf16ToBits(bfloat16_t value) +{ + union Bits { + __aicore__ Bits() {} + bfloat16_t f; + uint16_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline bfloat16_t BitsToBf16(uint16_t value) +{ + union Bits { + __aicore__ Bits() {} + uint16_t u; + bfloat16_t f; + } bits; + bits.u = value; + return bits.f; +} + +template +__aicore__ inline T FloatToType(float value) +{ + if constexpr (IsSameType::value) { + uint32_t bits = FloatToBits(value); + uint32_t bias = 0x7FFFu + ((bits >> 16) & 1u); + return BitsToBf16(static_cast((bits + bias) >> 16)); + } + return static_cast(value); +} + +template +class ChunkKdaFwdFinalizeKernel { +public: + using OUT_T = float; + using AKK_T = float; + template + __aicore__ inline void Init(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR preparedQG, GM_ADDR preparedAqk, + GM_ADDR propagatedVNew, GM_ADDR propagatedH, GM_ADDR o, GM_ADDR finalState, GM_ADDR aqk, + GM_ADDR akk, GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR workspace, const TilingData &tiling, TPipe *pipe, + bool initVecBuffers = true) + { + pipe_ = pipe; + q_.SetGlobalBuffer((__gm__ T *)q); + k_.SetGlobalBuffer((__gm__ T *)k); + v_.SetGlobalBuffer((__gm__ T *)v); + gk_.SetGlobalBuffer((__gm__ GK_T *)gk); + beta_.SetGlobalBuffer((__gm__ BETA_T *)beta); + if (initialState != nullptr) { + initialState_.SetGlobalBuffer((__gm__ float *)initialState); + } + cuSeqlensAddr_ = reinterpret_cast<__gm__ int64_t *>(cuSeqlens); + if (preparedQG != nullptr) { + preparedQG_.SetGlobalBuffer((__gm__ T *)preparedQG); + } + if (preparedAqk != nullptr) { + preparedAqk_.SetGlobalBuffer((__gm__ T *)preparedAqk); + } + if (propagatedVNew != nullptr) { + propagatedVNew_.SetGlobalBuffer((__gm__ T *)propagatedVNew); + } + if (propagatedH != nullptr) { + propagatedH_.SetGlobalBuffer((__gm__ T *)propagatedH); + } + chunkIndicesAddr_ = reinterpret_cast<__gm__ int64_t *>(chunkIndices); + o_.SetGlobalBuffer((__gm__ OUT_T *)o); + finalState_.SetGlobalBuffer((__gm__ float *)finalState); + aqk_.SetGlobalBuffer((__gm__ float *)aqk); + akk_.SetGlobalBuffer((__gm__ AKK_T *)akk); + w_.SetGlobalBuffer((__gm__ T *)w); + u_.SetGlobalBuffer((__gm__ OUT_T *)u); + qg_.SetGlobalBuffer((__gm__ T *)qg); + kg_.SetGlobalBuffer((__gm__ T *)kg); + vNew_.SetGlobalBuffer((__gm__ T *)vNew); + h_.SetGlobalBuffer((__gm__ float *)h); + solveWorkspace_.SetGlobalBuffer((__gm__ float *)workspace); + + B_ = tiling.batch; + N_ = tiling.seqNum; + H_ = tiling.qHeadNum; + HV_ = tiling.vHeadNum; + T_ = tiling.seqlen; + K_ = tiling.kHeadDim; + V_ = tiling.vHeadDim; + BT_ = tiling.chunkSize; + NT_ = tiling.totalChunks; + scale_ = tiling.scale; + hasInitial_ = tiling.hasInitialState; + isVarLen_ = tiling.isVarLen; + usedCoreNum_ = tiling.outputUsedCoreNum; + const uint64_t outputElements = B_ * HV_ * T_ * V_; + o_.SetGlobalBuffer((__gm__ OUT_T *)workspace); + u_.SetGlobalBuffer((__gm__ OUT_T *)workspace + outputElements); + if ASCEND_IS_AIV { + uint64_t subBlockNum = static_cast(GetSubBlockNum()); + solveCoreIdx_ = subBlockNum == 0 ? 0 : static_cast(GetBlockIdx()) / subBlockNum; + } else { + solveCoreIdx_ = static_cast(GetBlockIdx()); + } + if (pipe_ != nullptr && initVecBuffers) { + pipe_->InitBuffer(exp2Buf_, EXP2_UB_BYTES); + pipe_->InitBuffer(vecBuf_, KDA_VEC_ARENA_ELEMENTS * sizeof(float)); + const uint64_t gateWritebackRows = + ScoreVectorMaxRows(5 * sizeof(float) + 2 * sizeof(T) + sizeof(GK_T)); + pipe_->InitBuffer(gateWritebackBuf_, + static_cast(gateWritebackRows * K_ * + (3 * sizeof(T) + sizeof(GK_T)))); + AllocVectorEvents(); + } + } + __aicore__ inline void ProcessAiv() + { + ProcessOutAiv(); + ReleaseVectorEvents(); + } + + __aicore__ inline void ProcessAic() + { + ProcessOutAic(); + } + +private: + __aicore__ inline void AllocVectorEvents() + { + mte2ToVEvent_ = pipe_->AllocEventID(); + vToMte2Event_ = pipe_->AllocEventID(); + vToMte3Event_ = pipe_->AllocEventID(); + mte3ToVEvent_ = pipe_->AllocEventID(); + mte2ToMte3Event_ = pipe_->AllocEventID(); + mte3ToMte2Event_ = pipe_->AllocEventID(); + vectorEventsAllocated_ = true; + } + + __aicore__ inline void ReleaseVectorEvents() + { + if (!vectorEventsAllocated_) { + return; + } + pipe_->ReleaseEventID(mte2ToVEvent_); + pipe_->ReleaseEventID(vToMte2Event_); + pipe_->ReleaseEventID(vToMte3Event_); + pipe_->ReleaseEventID(mte3ToVEvent_); + pipe_->ReleaseEventID(mte2ToMte3Event_); + pipe_->ReleaseEventID(mte3ToMte2Event_); + vectorEventsAllocated_ = false; + } + + __aicore__ inline uint64_t QOffset(uint64_t b, uint64_t h, uint64_t t, uint64_t d) const + { + return ((b * H_ + h) * T_ + t) * K_ + d; + } + + __aicore__ inline uint64_t KVOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t d, uint64_t dim) const + { + return ((b * HV_ + hv) * T_ + t) * dim + d; + } + + __aicore__ inline uint64_t OutputOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t d) const + { + return ((b * T_ + t) * HV_ + hv) * V_ + d; + } + + __aicore__ inline uint64_t BetaOffset(uint64_t b, uint64_t hv, uint64_t t) const + { + return (b * HV_ + hv) * T_ + t; + } + + __aicore__ inline uint64_t AOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t j) const + { + return ((b * HV_ + hv) * T_ + t) * BT_ + j; + } + + __aicore__ inline uint64_t HOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t d, uint64_t r) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * K_ + d) * V_ + r; + } + + __aicore__ inline uint64_t WScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t t, uint64_t d) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * BT_ + t) * K_ + d; + } + + __aicore__ inline uint64_t SolveScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t slot) const + { + (void)b; + (void)hv; + (void)chunkIdx; + uint64_t matrixElements = BT_ * BT_; + return solveCoreIdx_ * KDA_SOLVE_SCRATCH_SLOTS * matrixElements + slot * matrixElements; + } + + __aicore__ inline uint64_t ScoreScratchOffset(uint64_t slot, uint64_t plane, uint64_t t = 0, + uint64_t d = 0) const + { + return (((solveCoreIdx_ * KDA_SCORE_QUEUE_DEPTH + slot) * KDA_SCORE_SCRATCH_PLANES + plane) * BT_ + t) * + K_ + + d; + } + + + + __aicore__ inline uint64_t ScoreRefBlockSize() const + { + return KDA_SCORE_REF_BC; + } + + __aicore__ inline uint64_t ScoreRowBlockCount(uint64_t curT, uint64_t rowBegin) const + { + uint64_t blockSize = ScoreRefBlockSize(); + uint64_t rowCount = curT - rowBegin; + if (rowCount > blockSize) { + rowCount = blockSize; + } + return rowCount; + } + + __aicore__ inline uint64_t ScoreRefToken(uint64_t start, uint64_t curT, uint64_t rowBegin, + uint64_t rowCount) const + { + uint64_t ref = rowBegin + rowCount / 2; + if (ref >= curT) { + ref = curT - 1; + } + return start + ref; + } + + __aicore__ inline void RunExp2(LocalTensor &tensor, uint32_t count) + { + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + ClampExpInput(tensor, count); + Exp(tensor, tensor, count); + PipeBarrier(); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + } + + __aicore__ inline void ClampExpInput(LocalTensor &tensor, uint32_t count) + { + Mins(tensor, tensor, KDA_EXP_INPUT_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, KDA_EXP_INPUT_MIN, count); + PipeBarrier(); + } + + __aicore__ inline void ClampFp32ToOutputType(LocalTensor &tensor, uint32_t count) + { + if constexpr (IsSameType::value) { + Mins(tensor, tensor, KDA_FP16_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, -KDA_FP16_MAX, count); + PipeBarrier(); + } + } + + template + __aicore__ inline void CopyVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst, src[offset], static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + + template + __aicore__ inline void CopyVectorOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst[offset], src, static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPad(dst[offset], src, params); + } + + template + __aicore__ inline void CopyRowsOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src, + uint64_t rows, uint64_t cols, uint64_t dstStride) + { + if (cols == dstStride) { + CopyVectorOut(dst, offset, src, rows * cols); + return; + } + constexpr uint64_t blockBytes = 32; + const uint64_t rowBytes = cols * sizeof(CopyT); + const uint64_t gapBytes = (dstStride - cols) * sizeof(CopyT); + DataCopyParams params{ + static_cast(rows), + static_cast(rowBytes / blockBytes), + 0, + static_cast(gapBytes / blockBytes) + }; + DataCopy(dst[offset], src, params); + } + + template + __aicore__ inline void CopyRowIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset) + { + CopyVectorIn(dst, src, offset, K_); + } + + template + __aicore__ inline void CopyRowOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src) + { + CopyVectorOut(dst, offset, src, K_); + } + + __aicore__ inline LocalTensor VecScratch(uint64_t slot) + { + return vecBuf_.Get()[slot * EXP2_UB_ELEMENTS]; + } + + template + __aicore__ inline void LoadAsFloatRow(GlobalTensor &src, uint64_t srcOffset, LocalTensor &dst, + uint64_t count) + { + if constexpr (IsSameType::value) { + CopyVectorIn(dst, src, srcOffset, count); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Adds(dst, dst, 0.0f, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } else { + constexpr uint32_t typedOffset = EXP2_UB_ELEMENTS * sizeof(float) / sizeof(CopyT); + LocalTensor rowLocal = exp2Buf_.Get()[typedOffset]; + CopyVectorIn(rowLocal, src, srcOffset, count); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast(dst, rowLocal, RoundMode::CAST_NONE, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } + PipeBarrier(); + } + + template + __aicore__ inline void LoadAsFloatVector(GlobalTensor &src, uint64_t srcOffset, + LocalTensor &dst, LocalTensor &typedScratch, + uint64_t count) + { + if constexpr (IsSameType::value) { + CopyVectorIn(dst, src, srcOffset, count); + } else { + CopyVectorIn(typedScratch, src, srcOffset, count); + } + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + if constexpr (!IsSameType::value) { + Cast(dst, typedScratch, RoundMode::CAST_NONE, static_cast(count)); + PipeBarrier(); + } + } + + template + __aicore__ inline void StoreFloatRow(GlobalTensor &dst, uint64_t dstOffset, LocalTensor &src, + uint64_t count) + { + if constexpr (IsSameType::value) { + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(dst, dstOffset, src, count); + } else { + constexpr uint32_t typedOffset = EXP2_UB_ELEMENTS * sizeof(float) / sizeof(CopyT); + LocalTensor rowLocal = exp2Buf_.Get()[typedOffset]; + Cast(rowLocal, src, RoundMode::CAST_RINT, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(dst, dstOffset, rowLocal, count); + } + SetFlag(mte3ToMte2Event_); + WaitFlag(mte3ToMte2Event_); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + + + + + __aicore__ inline LocalTensor Exp2NegG(uint64_t b, uint64_t hv, uint64_t t) + { + LocalTensor exp2Local = exp2Buf_.Get(); + LoadAsFloatRow(gk_, KVOffset(b, hv, t, 0, K_), exp2Local, K_); + Muls(exp2Local, exp2Local, -LN2, static_cast(K_)); + PipeBarrier(); + RunExp2(exp2Local, static_cast(K_)); + return exp2Local; + } + + + __aicore__ inline uint64_t ScoreVectorMaxRows(uint64_t bytesPerElem) const + { + constexpr uint64_t arenaBytes = static_cast(KDA_VEC_ARENA_ELEMENTS) * sizeof(float); + uint64_t maxRows = (arenaBytes / bytesPerElem) / K_; + if (K_ >= 128 && maxRows > 32) { + maxRows = 32; + } + return maxRows; + } + + __aicore__ inline void ComputeOutputCube(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT) + { + using ElementA = T; + using ElementB = T; + using ElementC = OUT_T; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using BlockMmad = Catlass::Gemm::Block::BlockMmadTla; + + Catlass::Arch::Resource resource; + BlockMmad blockMmad(resource); + + auto layoutQ = tla::MakeLayout(BT_, K_); + auto layoutH = tla::MakeLayout(K_, V_); + auto layoutO = tla::MakeLayout(BT_, V_); + for (uint64_t nOffset = 0; nOffset < V_; nOffset += 128) { + uint32_t curN = static_cast((V_ - nOffset) > 128 ? 128 : (V_ - nOffset)); + auto tensorH = tla::MakeTensor(propagatedH_[HOffset(b, hv, chunkIdx, 0, nOffset)], layoutH, + Catlass::Arch::PositionGM{}); + for (uint64_t mOffset = 0; mOffset < curT; mOffset += 64) { + uint32_t curM = static_cast((curT - mOffset) > 64 ? 64 : (curT - mOffset)); + Catlass::GemmCoord shapeQH{curM, curN, static_cast(K_)}; + auto tensorQ = tla::MakeTensor(preparedQG_[KVOffset(b, hv, start + mOffset, 0, K_)], layoutQ, + Catlass::Arch::PositionGM{}); + auto tensorO = tla::MakeTensor(o_[KVOffset(b, hv, start + mOffset, nOffset, V_)], layoutO, + Catlass::Arch::PositionGM{}); + auto blockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(shapeQH.m(), shapeQH.k())); + auto blockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(shapeQH.k(), shapeQH.n())); + auto blockO = GetTile(tensorO, tla::MakeCoord(0, 0), tla::MakeShape(shapeQH.m(), shapeQH.n())); + blockMmad(blockQ, blockH, blockO, shapeQH); + PipeBarrier(); + } + } + + auto layoutAqk = tla::MakeLayout(BT_, BT_); + auto layoutV = tla::MakeLayout(BT_, V_); + for (uint64_t nOffset = 0; nOffset < V_; nOffset += 128) { + uint32_t curN = static_cast((V_ - nOffset) > 128 ? 128 : (V_ - nOffset)); + auto tensorVNew = tla::MakeTensor(propagatedVNew_[KVOffset(b, hv, start, nOffset, V_)], layoutV, + Catlass::Arch::PositionGM{}); + for (uint64_t mOffset = 0; mOffset < curT; mOffset += 64) { + uint32_t curM = static_cast((curT - mOffset) > 64 ? 64 : (curT - mOffset)); + Catlass::GemmCoord shapeAV{curM, curN, static_cast(curT)}; + auto tensorAqk = tla::MakeTensor(preparedAqk_[AOffset(b, hv, start + mOffset, 0)], layoutAqk, + Catlass::Arch::PositionGM{}); + auto tensorLocal = tla::MakeTensor(u_[KVOffset(b, hv, start + mOffset, nOffset, V_)], layoutO, + Catlass::Arch::PositionGM{}); + auto blockAqk = GetTile(tensorAqk, tla::MakeCoord(0, 0), tla::MakeShape(shapeAV.m(), shapeAV.k())); + auto blockVNew = GetTile(tensorVNew, tla::MakeCoord(0, 0), tla::MakeShape(shapeAV.k(), shapeAV.n())); + auto blockLocal = GetTile(tensorLocal, tla::MakeCoord(0, 0), tla::MakeShape(shapeAV.m(), shapeAV.n())); + blockMmad(blockAqk, blockVNew, blockLocal, shapeAV); + PipeBarrier(); + } + } + } + + __aicore__ inline void FinalizeOutputRows(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT, + uint64_t subBlockIdx, uint64_t subBlockNum) + { + if (subBlockNum == 0 || subBlockIdx >= subBlockNum || V_ == 0) { + return; + } + const uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; + const uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + const uint64_t gateWritebackRows = + ScoreVectorMaxRows(5 * sizeof(float) + 2 * sizeof(T) + sizeof(GK_T)); + const uint64_t gateWritebackBytes = + gateWritebackRows * K_ * (3 * sizeof(T) + sizeof(GK_T)); + uint64_t maxRows = KDA_VEC_ARENA_ELEMENTS / (3 * V_); + const uint64_t typedMaxRows = gateWritebackBytes / (V_ * sizeof(T)); + if (maxRows > typedMaxRows) { + maxRows = typedMaxRows; + } + if (maxRows == 0) { + return; + } + + for (uint64_t tileRow = rowBegin; tileRow < rowEnd; tileRow += maxRows) { + uint64_t tileRows = rowEnd - tileRow; + if (tileRows > maxRows) { + tileRows = maxRows; + } + const uint64_t elems = tileRows * V_; + const uint64_t ti = start + tileRow; + LocalTensor arena = vecBuf_.Get(); + LocalTensor stateLocal = arena; + LocalTensor localLocal = arena[elems]; + LocalTensor outLocal = arena[2 * elems]; + LocalTensor outTyped = gateWritebackBuf_.Get(); + + CopyVectorIn(stateLocal, o_, KVOffset(b, hv, ti, 0, V_), elems); + CopyVectorIn(localLocal, u_, KVOffset(b, hv, ti, 0, V_), elems); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Add(outLocal, stateLocal, localLocal, static_cast(elems)); + PipeBarrier(); + ClampFp32ToOutputType(outLocal, static_cast(elems)); + Cast(outTyped, outLocal, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyRowsOut(vNew_, OutputOffset(b, hv, ti, 0), outTyped, tileRows, V_, HV_ * V_); + SetFlag(mte3ToMte2Event_); + WaitFlag(mte3ToMte2Event_); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + } + __aicore__ inline bool ResolveFlatChunk(uint64_t task, uint64_t &seq, uint64_t &b, uint64_t &h, uint64_t &hv, + uint64_t &chunkIdx, uint64_t &start, uint64_t &end) + { + hv = task % HV_; + uint64_t flatChunk = task / HV_; + if (!isVarLen_) { + seq = flatChunk / NT_; + b = seq; + chunkIdx = flatChunk % NT_; + start = chunkIdx * BT_; + end = start + BT_; + if (end > T_) { + end = T_; + } + } else { + if (!KdaVarlen::ResolveChunkRange( + cuSeqlensAddr_, chunkIndicesAddr_, N_, T_, BT_, flatChunk, + seq, start, end)) { + return false; + } + b = 0; + chunkIdx = flatChunk; + } + h = hv / (HV_ / H_); + return start < end; + } + + __aicore__ inline void ProcessChunkOutAiv(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end, uint64_t subBlockIdx, uint64_t subBlockNum) + { + uint64_t curT = end - start; + if (curT == 0) { + return; + } + if constexpr (IsSameType::value) { + return; + } + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(syncDoneFlag_); + FinalizeOutputRows(b, hv, start, curT, subBlockIdx, subBlockNum); + } + + __aicore__ inline void ProcessChunkOutAic(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + uint64_t curT = end - start; + if (curT == 0) { + return; + } + ComputeOutputCube(b, hv, chunkIdx, start, curT); + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); + } + + __aicore__ inline void ProcessOutAiv() + { + if constexpr (IsSameType::value) { + return; + } + uint64_t subBlockNum = static_cast(GetSubBlockNum()); + if (subBlockNum == 0) { + return; + } + uint64_t subBlockIdx = static_cast(GetSubBlockIdx()); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + uint64_t coreIdx = static_cast(GetBlockIdx()) / subBlockNum; + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + for (uint64_t task = coreIdx; task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + (void)h; + (void)chunkIdx; + ProcessChunkOutAiv(b, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); + } + } + } + + __aicore__ inline void ProcessOutAic() + { + if constexpr (IsSameType::value) { + return; + } + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + (void)h; + ProcessChunkOutAic(b, hv, chunkIdx, start, end); + } + } + } + +private: + GlobalTensor q_; + GlobalTensor k_; + GlobalTensor v_; + GlobalTensor gk_; + GlobalTensor beta_; + GlobalTensor initialState_; + GlobalTensor o_; + GlobalTensor finalState_; + GlobalTensor aqk_; + GlobalTensor akk_; + GlobalTensor w_; + GlobalTensor u_; + GlobalTensor qg_; + GlobalTensor kg_; + GlobalTensor vNew_; + GlobalTensor h_; + GlobalTensor preparedQG_; + GlobalTensor preparedAqk_; + GlobalTensor propagatedVNew_; + GlobalTensor propagatedH_; + GlobalTensor solveWorkspace_; + GlobalTensor scoreWorkspace_; + TPipe *pipe_ = nullptr; + TBuf exp2Buf_; + TBuf vecBuf_; + TBuf gateWritebackBuf_; + TEventID mte2ToVEvent_ = 0; + TEventID vToMte2Event_ = 0; + TEventID vToMte3Event_ = 0; + TEventID mte3ToVEvent_ = 0; + TEventID mte2ToMte3Event_ = 0; + TEventID mte3ToMte2Event_ = 0; + bool vectorEventsAllocated_ = false; + Catlass::Arch::CrossCoreFlagWithReverse scoreReadyFlag_{KDA_SCORE_READY_FLAG0, + KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse scoreDoneFlag_{KDA_SCORE_DONE_FLAG0, + KDA_SCORE_DONE_FLAG1}; + // Score production is fully drained before solve starts, so the solve handshake can safely reuse + // the existing score flags without consuming additional hardware flag IDs. + Catlass::Arch::CrossCoreFlagWithReverse syncReadyFlag_{KDA_SCORE_READY_FLAG0, + KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse syncDoneFlag_{KDA_SCORE_DONE_FLAG0, + KDA_SCORE_DONE_FLAG1}; + uint64_t B_ = 0; + uint64_t N_ = 0; + uint64_t H_ = 0; + uint64_t HV_ = 0; + uint64_t T_ = 0; + uint64_t K_ = 0; + uint64_t V_ = 0; + uint64_t BT_ = 0; + uint64_t NT_ = 0; + float scale_ = 1.0f; + bool hasInitial_ = false; + bool isVarLen_ = false; + bool isAivOnly_ = false; + uint64_t usedCoreNum_ = 1; + uint64_t solveCoreIdx_ = 0; + __gm__ int64_t *chunkIndicesAddr_ = nullptr; + __gm__ int64_t *cuSeqlensAddr_ = nullptr; +}; +} // namespace + +template +__aicore__ inline void RunChunkKdaOutput( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR qgScaled, GM_ADDR aqk, + GM_ADDR propagatedVNew, GM_ADDR propagatedH, GM_ADDR o, GM_ADDR userWorkspace, + const TilingData &tiling, TPipe &pipe) +{ + GM_ADDR outputScratch = userWorkspace + tiling.outputScratchOffset; + uint64_t outputElements = static_cast(tiling.batch) * + static_cast(tiling.vHeadNum) * + static_cast(tiling.seqlen) * + static_cast(tiling.vHeadDim); + GM_ADDR stateScratch = outputScratch; + GM_ADDR localScratch = outputScratch + outputElements * sizeof(float); + if ASCEND_IS_AIC { + ChunkKdaFwdFinalizeKernel op; + op.Init(q, k, v, gk, beta, initialState, cuSeqlens, chunkIndices, + qgScaled, aqk, propagatedVNew, propagatedH, stateScratch, userWorkspace, aqk, userWorkspace, + userWorkspace, localScratch, userWorkspace, userWorkspace, o, propagatedH, + outputScratch, tiling, &pipe, false); + op.ProcessAic(); + } + if ASCEND_IS_AIV { + ChunkKdaFwdFinalizeKernel op; + op.Init(q, k, v, gk, beta, initialState, cuSeqlens, chunkIndices, + qgScaled, aqk, propagatedVNew, propagatedH, stateScratch, userWorkspace, aqk, userWorkspace, + userWorkspace, localScratch, userWorkspace, userWorkspace, o, propagatedH, + outputScratch, tiling, &pipe); + op.ProcessAiv(); + } +} + +} // namespace KdaFinalize diff --git a/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_post_wu.h b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_post_wu.h new file mode 100644 index 000000000000..f69405cce431 --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_post_wu.h @@ -0,0 +1,982 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#pragma once + +#ifndef CATLASS_ARCH +#define CATLASS_ARCH 2201 +#endif + +#include "catlass/arch/arch.hpp" +#include "catlass/arch/cross_core_sync.hpp" +#include "catlass/arch/resource.hpp" +#include "catlass/catlass.hpp" +#include "catlass/gemm/block/block_mmad.hpp" +#include "catlass/gemm/dispatch_policy.hpp" +#include "catlass/gemm/tile/tile_copy.hpp" +#include "catlass/gemm/tile/tile_mmad.hpp" +#include "catlass/gemm_coord.hpp" +#include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp" +#include "catlass/layout/layout.hpp" +#include "kernel_operator.h" +#include "chunk_kda_fwd_varlen.h" +#include "tla/layout.hpp" +#include "tla/tensor.hpp" + +using namespace AscendC; + +namespace KdaPostWu { +namespace { +using KdaInt64 = tla::Int<64>; +using KdaInt128 = tla::Int<128>; +constexpr float LN2 = 0.69314718055994530942f; +constexpr float KDA_EXP2_CLAMP = 80.0f; +constexpr float KDA_EXP_INPUT_MAX = KDA_EXP2_CLAMP * LN2; +constexpr float KDA_EXP_INPUT_MIN = -KDA_EXP2_CLAMP * LN2; +constexpr float KDA_FP16_MAX = 65504.0f; +constexpr uint32_t EXP2_UB_ELEMENTS = 256; +constexpr uint32_t EXP2_UB_BYTES = EXP2_UB_ELEMENTS * (sizeof(float) + sizeof(uint16_t)); +constexpr uint32_t EXP2_EVENT_ID = 0; +constexpr uint32_t KDA_SOLVE_BT = 64; +constexpr uint32_t KDA_SOLVE_MATRIX_ELEMENTS = KDA_SOLVE_BT * KDA_SOLVE_BT; +constexpr uint32_t KDA_SOLVE_SCRATCH_X = 0; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y0 = 1; +constexpr uint32_t KDA_SOLVE_SCRATCH_TMP = 2; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y1 = 3; +constexpr uint32_t KDA_SOLVE_SCRATCH_IDENTITY = 4; +constexpr uint32_t KDA_SOLVE_SCRATCH_SLOTS = 5; +constexpr uint32_t KDA_SOLVE_DIAG_BT = 16; +constexpr uint32_t KDA_SOLVE_DIAG_BLOCKS = KDA_SOLVE_BT / KDA_SOLVE_DIAG_BT; +constexpr uint32_t KDA_SOLVE_DIAG_MCH_ITERS = 3; +constexpr uint32_t KDA_SCORE_REF_BC = 16; +constexpr uint32_t KDA_VEC_ARENA_ELEMENTS = 32768; +constexpr uint32_t KDA_BITS_PER_MASK_BYTE = 8; +constexpr uint32_t KDA_SELECT_COL_BLOCKS = 2; +constexpr uint32_t KDA_SELECT_COL_MASK_BYTES = KDA_SOLVE_MATRIX_ELEMENTS / KDA_BITS_PER_MASK_BYTE; +constexpr uint32_t KDA_SELECT_MASK_BYTES = KDA_SELECT_COL_BLOCKS * KDA_SELECT_COL_MASK_BYTES; +constexpr uint32_t KDA_SELECT_AQK_MASK_BYTE_OFFSET = 120 * 1024; +constexpr uint32_t KDA_SELECT_AKK_MASK_BYTE_OFFSET = KDA_SELECT_AQK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_BYTE_OFFSET = KDA_SELECT_AKK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_FLOAT_OFFSET = KDA_SELECT_ZERO_BYTE_OFFSET / sizeof(float); +constexpr uint8_t KDA_SCORE_DONE_FLAG0 = 2; +constexpr uint8_t KDA_SCORE_DONE_FLAG1 = 3; +constexpr uint8_t KDA_SCORE_READY_FLAG0 = 4; +constexpr uint8_t KDA_SCORE_READY_FLAG1 = 5; +constexpr uint32_t KDA_SCORE_QUEUE_DEPTH = 2; +constexpr uint32_t KDA_SYNC_REVERSE_DEPTH = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_PLANES = 3; +constexpr uint32_t KDA_SCORE_SCRATCH_QG = 0; +constexpr uint32_t KDA_SCORE_SCRATCH_W = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_KG = 2; +constexpr uint64_t KDA_WORKSPACE_ALIGN = 512; +constexpr uint32_t KDA_GATE_TILE_ROWS = 32; +constexpr uint32_t KDA_POST_EVENT = 3; +constexpr uint32_t KDA_POST_EVENT_NEXT = 4; +constexpr uint32_t KDA_POST_EVENT_FIX = 5; + +using KdaArchTag = Catlass::Arch::AtlasA2; +using KdaDispatchPolicy = Catlass::Gemm::MmadPingpong; +using KdaScoreDispatchPolicy = + Catlass::Gemm::MmadPingpongTlaMulti; +static_assert(KdaScoreDispatchPolicy::ENABLE_L1_RESIDENT, + "KDA Aqk/Akk score MMAD must keep the shared right matrix resident in L1"); +static_assert(KdaScoreDispatchPolicy::L1B_STAGES == 1, + "KDA Aqk/Akk score MMAD needs one L1 B slot so the second MMAD reuses it"); +using KdaSolveDispatchPolicy = Catlass::Gemm::MmadPingpong; +static_assert(!KdaSolveDispatchPolicy::USE_HF32_MODE, "KDA triangular solve must use IEEE FP32 Cube mode"); +using KdaL1TileShape = tla::Shape; +using KdaL0TileShape = KdaL1TileShape; +using KdaSolveL1TileShape = tla::Shape; +using KdaSolveL0TileShape = KdaSolveL1TileShape; + +__aicore__ inline uint32_t FloatToBits(float value) +{ + union Bits { + __aicore__ Bits() {} + float f; + uint32_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline float BitsToFloat(uint32_t value) +{ + union Bits { + __aicore__ Bits() {} + uint32_t u; + float f; + } bits; + bits.u = value; + return bits.f; +} + +__aicore__ inline uint16_t Bf16ToBits(bfloat16_t value) +{ + union Bits { + __aicore__ Bits() {} + bfloat16_t f; + uint16_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline bfloat16_t BitsToBf16(uint16_t value) +{ + union Bits { + __aicore__ Bits() {} + uint16_t u; + bfloat16_t f; + } bits; + bits.u = value; + return bits.f; +} + +template +__aicore__ inline T FloatToType(float value) +{ + if constexpr (IsSameType::value) { + uint32_t bits = FloatToBits(value); + uint32_t bias = 0x7FFFu + ((bits >> 16) & 1u); + return BitsToBf16(static_cast((bits + bias) >> 16)); + } + return static_cast(value); +} + +template +class ChunkKdaFwdPostWuKernel { +public: + using OUT_T = T; + using AKK_T = T; + template + __aicore__ inline void Init(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR preparedQG, GM_ADDR preparedAqk, + GM_ADDR propagatedVNew, GM_ADDR propagatedH, GM_ADDR o, GM_ADDR finalState, GM_ADDR aqk, + GM_ADDR akk, GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR workspace, const TilingData &tiling, TPipe *pipe, + bool initVecBuffers = true) + { + pipe_ = pipe; + q_.SetGlobalBuffer((__gm__ T *)q); + k_.SetGlobalBuffer((__gm__ T *)k); + v_.SetGlobalBuffer((__gm__ T *)v); + gk_.SetGlobalBuffer((__gm__ GK_T *)gk); + beta_.SetGlobalBuffer((__gm__ BETA_T *)beta); + if (initialState != nullptr) { + initialState_.SetGlobalBuffer((__gm__ float *)initialState); + } + cuSeqlensAddr_ = reinterpret_cast<__gm__ int64_t *>(cuSeqlens); + if (preparedQG != nullptr) { + preparedQG_.SetGlobalBuffer((__gm__ T *)preparedQG); + } + if (preparedAqk != nullptr) { + preparedAqk_.SetGlobalBuffer((__gm__ T *)preparedAqk); + } + if (propagatedVNew != nullptr) { + propagatedVNew_.SetGlobalBuffer((__gm__ T *)propagatedVNew); + } + if (propagatedH != nullptr) { + propagatedH_.SetGlobalBuffer((__gm__ T *)propagatedH); + } + chunkIndicesAddr_ = reinterpret_cast<__gm__ int64_t *>(chunkIndices); + o_.SetGlobalBuffer((__gm__ OUT_T *)o); + finalState_.SetGlobalBuffer((__gm__ float *)finalState); + aqk_.SetGlobalBuffer((__gm__ float *)aqk); + akk_.SetGlobalBuffer((__gm__ AKK_T *)akk); + w_.SetGlobalBuffer((__gm__ T *)w); + u_.SetGlobalBuffer((__gm__ OUT_T *)u); + qg_.SetGlobalBuffer((__gm__ T *)qg); + kg_.SetGlobalBuffer((__gm__ T *)kg); + vNew_.SetGlobalBuffer((__gm__ T *)vNew); + h_.SetGlobalBuffer((__gm__ float *)h); + solveWorkspace_.SetGlobalBuffer((__gm__ float *)workspace); + + B_ = tiling.batch; + N_ = tiling.seqNum; + H_ = tiling.qHeadNum; + HV_ = tiling.vHeadNum; + T_ = tiling.seqlen; + K_ = tiling.kHeadDim; + V_ = tiling.vHeadDim; + BT_ = tiling.chunkSize; + NT_ = tiling.totalChunks; + scale_ = tiling.scale; + hasInitial_ = tiling.hasInitialState; + isVarLen_ = tiling.isVarLen; + inputSequenceMajor_ = tiling.inputSequenceMajor; + usedCoreNum_ = tiling.postWuUsedCoreNum; + if ASCEND_IS_AIV { + uint64_t subBlockNum = static_cast(GetSubBlockNum()); + solveCoreIdx_ = subBlockNum == 0 ? 0 : static_cast(GetBlockIdx()) / subBlockNum; + } else { + solveCoreIdx_ = static_cast(GetBlockIdx()); + } + if (pipe_ != nullptr && initVecBuffers) { + pipe_->InitBuffer(exp2Buf_, EXP2_UB_BYTES); + pipe_->InitBuffer(vecBuf_, KDA_VEC_ARENA_ELEMENTS * sizeof(float)); + const uint64_t gateWritebackRows = + ScoreVectorMaxRows(5 * sizeof(float) + 2 * sizeof(T) + sizeof(GK_T)); + pipe_->InitBuffer(gateWritebackBuf_, + static_cast(gateWritebackRows * K_ * + (3 * sizeof(T) + sizeof(GK_T)))); + AllocVectorEvents(); + } + } + __aicore__ inline void ProcessAiv() + { + ProcessPostAiv(); + ReleaseVectorEvents(); + } + + __aicore__ inline void ProcessAic() + { + ProcessPostAic(); + } + +private: + __aicore__ inline void AllocVectorEvents() + { + mte2ToVEvent_ = pipe_->AllocEventID(); + vToMte2Event_ = pipe_->AllocEventID(); + vToMte3Event_ = pipe_->AllocEventID(); + mte3ToVEvent_ = pipe_->AllocEventID(); + mte2ToMte3Event_ = pipe_->AllocEventID(); + mte3ToMte2Event_ = pipe_->AllocEventID(); + vectorEventsAllocated_ = true; + } + + __aicore__ inline void ReleaseVectorEvents() + { + if (!vectorEventsAllocated_) { + return; + } + pipe_->ReleaseEventID(mte2ToVEvent_); + pipe_->ReleaseEventID(vToMte2Event_); + pipe_->ReleaseEventID(vToMte3Event_); + pipe_->ReleaseEventID(mte3ToVEvent_); + pipe_->ReleaseEventID(mte2ToMte3Event_); + pipe_->ReleaseEventID(mte3ToMte2Event_); + vectorEventsAllocated_ = false; + } + + __aicore__ inline uint64_t QOffset(uint64_t b, uint64_t h, uint64_t t, uint64_t d) const + { + if (inputSequenceMajor_) { + return ((b * T_ + t) * H_ + h) * K_ + d; + } + return ((b * H_ + h) * T_ + t) * K_ + d; + } + + __aicore__ inline uint64_t KVOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t d, uint64_t dim) const + { + return ((b * HV_ + hv) * T_ + t) * dim + d; + } + + __aicore__ inline uint64_t BetaOffset(uint64_t b, uint64_t hv, uint64_t t) const + { + return (b * HV_ + hv) * T_ + t; + } + + __aicore__ inline uint64_t AOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t j) const + { + return ((b * HV_ + hv) * T_ + t) * BT_ + j; + } + + __aicore__ inline uint64_t HOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t d, uint64_t r) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * K_ + d) * V_ + r; + } + + __aicore__ inline uint64_t WScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t t, uint64_t d) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * BT_ + t) * K_ + d; + } + + __aicore__ inline uint64_t SolveScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t slot) const + { + (void)b; + (void)hv; + (void)chunkIdx; + uint64_t matrixElements = BT_ * BT_; + return solveCoreIdx_ * KDA_SOLVE_SCRATCH_SLOTS * matrixElements + slot * matrixElements; + } + + __aicore__ inline uint64_t ScoreScratchOffset(uint64_t slot, uint64_t plane, uint64_t t = 0, + uint64_t d = 0) const + { + return (((solveCoreIdx_ * KDA_SCORE_QUEUE_DEPTH + slot) * KDA_SCORE_SCRATCH_PLANES + plane) * BT_ + t) * + K_ + + d; + } + + + + __aicore__ inline uint64_t ScoreRefBlockSize() const + { + return KDA_SCORE_REF_BC; + } + + __aicore__ inline uint64_t ScoreRowBlockCount(uint64_t curT, uint64_t rowBegin) const + { + uint64_t blockSize = ScoreRefBlockSize(); + uint64_t rowCount = curT - rowBegin; + if (rowCount > blockSize) { + rowCount = blockSize; + } + return rowCount; + } + + __aicore__ inline uint64_t ScoreRefToken(uint64_t start, uint64_t curT, uint64_t rowBegin, + uint64_t rowCount) const + { + uint64_t ref = rowBegin + rowCount / 2; + if (ref >= curT) { + ref = curT - 1; + } + return start + ref; + } + + __aicore__ inline void RunExp2(LocalTensor &tensor, uint32_t count) + { + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + ClampExpInput(tensor, count); + Exp(tensor, tensor, count); + PipeBarrier(); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + } + + __aicore__ inline void ClampExpInput(LocalTensor &tensor, uint32_t count) + { + Mins(tensor, tensor, KDA_EXP_INPUT_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, KDA_EXP_INPUT_MIN, count); + PipeBarrier(); + } + + __aicore__ inline void ClampFp32ToOutputType(LocalTensor &tensor, uint32_t count) + { + if constexpr (IsSameType::value) { + Mins(tensor, tensor, KDA_FP16_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, -KDA_FP16_MAX, count); + PipeBarrier(); + } + } + + template + __aicore__ inline void CopyVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst, src[offset], static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + + template + __aicore__ inline void CopyVectorOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst[offset], src, static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPad(dst[offset], src, params); + } + + template + __aicore__ inline void CopyRowsIn(LocalTensor &dst, GlobalTensor &src, + uint64_t offset, uint64_t rows, uint64_t cols, + uint64_t rowStride) + { + if (rows == 0 || cols == 0) { + return; + } + if (rowStride == cols) { + CopyVectorIn(dst, src, offset, rows * cols); + return; + } + DataCopyExtParams params{ + static_cast(rows), + static_cast(cols * sizeof(CopyT)), + static_cast((rowStride - cols) * sizeof(CopyT)), + 0, + 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + + template + __aicore__ inline void CopyRowIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset) + { + CopyVectorIn(dst, src, offset, K_); + } + + template + __aicore__ inline void CopyRowOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src) + { + CopyVectorOut(dst, offset, src, K_); + } + + __aicore__ inline LocalTensor VecScratch(uint64_t slot) + { + return vecBuf_.Get()[slot * EXP2_UB_ELEMENTS]; + } + + template + __aicore__ inline void LoadAsFloatRow(GlobalTensor &src, uint64_t srcOffset, LocalTensor &dst, + uint64_t count) + { + if constexpr (IsSameType::value) { + CopyVectorIn(dst, src, srcOffset, count); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Adds(dst, dst, 0.0f, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } else { + constexpr uint32_t typedOffset = EXP2_UB_ELEMENTS * sizeof(float) / sizeof(CopyT); + LocalTensor rowLocal = exp2Buf_.Get()[typedOffset]; + CopyVectorIn(rowLocal, src, srcOffset, count); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast(dst, rowLocal, RoundMode::CAST_NONE, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } + PipeBarrier(); + } + + template + __aicore__ inline void LoadAsFloatVector(GlobalTensor &src, uint64_t srcOffset, + LocalTensor &dst, LocalTensor &typedScratch, + uint64_t count) + { + if constexpr (IsSameType::value) { + CopyVectorIn(dst, src, srcOffset, count); + } else { + CopyVectorIn(typedScratch, src, srcOffset, count); + } + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + if constexpr (!IsSameType::value) { + Cast(dst, typedScratch, RoundMode::CAST_NONE, static_cast(count)); + PipeBarrier(); + } + } + + template + __aicore__ inline void StoreFloatRow(GlobalTensor &dst, uint64_t dstOffset, LocalTensor &src, + uint64_t count) + { + if constexpr (IsSameType::value) { + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(dst, dstOffset, src, count); + } else { + constexpr uint32_t typedOffset = EXP2_UB_ELEMENTS * sizeof(float) / sizeof(CopyT); + LocalTensor rowLocal = exp2Buf_.Get()[typedOffset]; + Cast(rowLocal, src, RoundMode::CAST_RINT, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(dst, dstOffset, rowLocal, count); + } + SetFlag(mte3ToMte2Event_); + WaitFlag(mte3ToMte2Event_); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + + + + + __aicore__ inline LocalTensor Exp2NegG(uint64_t b, uint64_t hv, uint64_t t) + { + LocalTensor exp2Local = exp2Buf_.Get(); + LoadAsFloatRow(gk_, KVOffset(b, hv, t, 0, K_), exp2Local, K_); + Muls(exp2Local, exp2Local, -LN2, static_cast(K_)); + PipeBarrier(); + RunExp2(exp2Local, static_cast(K_)); + return exp2Local; + } + + + __aicore__ inline uint64_t ScoreVectorMaxRows(uint64_t bytesPerElem) const + { + constexpr uint64_t arenaBytes = static_cast(KDA_VEC_ARENA_ELEMENTS) * sizeof(float); + uint64_t maxRows = (arenaBytes / bytesPerElem) / K_; + if (K_ >= 128 && maxRows > 32) { + maxRows = 32; + } + return maxRows; + } + __aicore__ inline bool UsePostWuCube(uint64_t curT) const + { + return curT > 0 && curT <= BT_ && (BT_ == 64 || BT_ == 128) && K_ >= 16 && V_ >= 16 && + V_ <= 256 && K_ % 16 == 0 && V_ % 16 == 0; + } + + + __aicore__ inline void ComputePostWuCube(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT) + { + using ElementA = AKK_T; + using ElementB = T; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using WTileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using UTileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using PostL1TileShape128 = tla::Shape; + using PostL0TileShape128 = tla::Shape; + using PostL1TileShape256 = tla::Shape; + using PostL0TileShape256 = tla::Shape; + using WBlockMmad = Catlass::Gemm::Block::BlockMmadTla; + using UBlockMmad128 = Catlass::Gemm::Block::BlockMmadTla; + using UBlockMmad256 = Catlass::Gemm::Block::BlockMmadTla; + LayoutTagA tagA = LayoutTagA::template MakeLayout(BT_, BT_); + auto layoutA = tla::MakeLayoutFromTag(tagA); + auto tensorA = tla::MakeTensor(preparedAqk_[AOffset(b, hv, start, 0)], layoutA, + Catlass::Arch::PositionGM{}); + + { + LayoutTagB tagB = LayoutTagB::template MakeLayout(BT_, K_); + LayoutTagC tagC = LayoutTagC::template MakeLayout(BT_, K_); + auto layoutB = tla::MakeLayoutFromTag(tagB); + auto layoutC = tla::MakeLayoutFromTag(tagC); + Catlass::GemmCoord shape{static_cast(curT), static_cast(K_), + static_cast(curT)}; + auto tensorB = tla::MakeTensor(preparedQG_[KVOffset(b, hv, start, 0, K_)], layoutB, + Catlass::Arch::PositionGM{}); + auto tensorC = tla::MakeTensor(h_[WScratchOffset(b, hv, chunkIdx, 0, 0)], layoutC, + Catlass::Arch::PositionGM{}); + auto blockA = GetTile(tensorA, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.k())); + auto blockB = GetTile(tensorB, tla::MakeCoord(0, 0), tla::MakeShape(shape.k(), shape.n())); + auto blockC = GetTile(tensorC, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.n())); + Catlass::Arch::Resource wResource; + WBlockMmad wBlockMmad(wResource); + wBlockMmad(blockA, blockB, blockC, shape); + PipeBarrier(); + } + + { + LayoutTagB tagB = LayoutTagB::template MakeLayout(BT_, V_); + LayoutTagC tagC = LayoutTagC::template MakeLayout(BT_, V_); + auto layoutB = tla::MakeLayoutFromTag(tagB); + auto layoutC = tla::MakeLayoutFromTag(tagC); + Catlass::GemmCoord shape{static_cast(curT), static_cast(V_), + static_cast(curT)}; + auto tensorB = tla::MakeTensor(propagatedVNew_[KVOffset(b, hv, start, 0, V_)], layoutB, + Catlass::Arch::PositionGM{}); + auto tensorC = tla::MakeTensor(u_[KVOffset(b, hv, start, 0, V_)], layoutC, + Catlass::Arch::PositionGM{}); + auto blockA = GetTile(tensorA, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.k())); + auto blockB = GetTile(tensorB, tla::MakeCoord(0, 0), tla::MakeShape(shape.k(), shape.n())); + auto blockC = GetTile(tensorC, tla::MakeCoord(0, 0), tla::MakeShape(shape.m(), shape.n())); + Catlass::Arch::Resource uResource; + if (V_ <= 128) { + UBlockMmad128 uBlockMmad(uResource); + uBlockMmad(blockA, blockB, blockC, shape); + } else { + UBlockMmad256 uBlockMmad(uResource); + uBlockMmad(blockA, blockB, blockC, shape); + } + PipeBarrier(); + } + + } + + __aicore__ inline void CopyScratchWAndFinalizeKg(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT, uint64_t subBlockIdx, + uint64_t subBlockNum, bool copyScratchW) + { + constexpr uint64_t typedOffsetFloats = 20480; + constexpr uint64_t typedOffset = typedOffsetFloats * sizeof(float) / sizeof(T); + constexpr uint64_t kgFp32Planes = 4; + uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; + uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + if (rowBegin >= rowEnd) { + return; + } + uint64_t maxRows = (typedOffsetFloats / kgFp32Planes) / K_; + if (maxRows > 32) { + maxRows = 32; + } + if (maxRows == 0) { + return; + } + + uint64_t last = start + curT - 1; + LocalTensor arena = vecBuf_.Get(); + LocalTensor gateLast = exp2Buf_.Get(); + LocalTensor typedLocal = vecBuf_.Get()[typedOffset]; + LoadAsFloatRow(gk_, KVOffset(b, hv, last, 0, K_), gateLast, K_); + + for (uint64_t tileRow = rowBegin; tileRow < rowEnd; tileRow += maxRows) { + uint64_t tileRows = rowEnd - tileRow; + if (tileRows > maxRows) { + tileRows = maxRows; + } + uint64_t elemCount = tileRows * K_; + uint64_t token = start + tileRow; + + if (copyScratchW) { + uint64_t scratchBase = WScratchOffset(b, hv, chunkIdx, tileRow, 0); + DataCopy(arena, h_[scratchBase], static_cast(elemCount)); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast(typedLocal, arena, RoundMode::CAST_RINT, static_cast(elemCount)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(w_[KVOffset(b, hv, token, 0, K_)], typedLocal, static_cast(elemCount)); + SetFlag(mte3ToMte2Event_); + WaitFlag(mte3ToMte2Event_); + } + + LocalTensor kLocal = arena; + LocalTensor gLocal = arena[elemCount]; + LocalTensor expLocal = arena[2 * elemCount]; + LocalTensor outLocal = arena[3 * elemCount]; + const uint64_t gateOffsetBytes = (typedOffset + elemCount) * sizeof(T); + LocalTensor gateTyped = vecBuf_.Get()[ + (gateOffsetBytes + sizeof(GK_T) - 1) / sizeof(GK_T)]; + CopyRowsIn(typedLocal, k_, QOffset(b, h, token, 0), tileRows, K_, + inputSequenceMajor_ ? H_ * K_ : K_); + LoadAsFloatVector(gk_, KVOffset(b, hv, token, 0, K_), gLocal, gateTyped, elemCount); + Cast(kLocal, typedLocal, RoundMode::CAST_NONE, static_cast(elemCount)); + PipeBarrier(); + + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expLocal[row * K_], gateLast, gLocal[row * K_], static_cast(K_)); + } + PipeBarrier(); + Muls(expLocal, expLocal, LN2, static_cast(elemCount)); + PipeBarrier(); + ClampExpInput(expLocal, static_cast(elemCount)); + Exp(expLocal, expLocal, static_cast(elemCount)); + PipeBarrier(); + Mul(outLocal, kLocal, expLocal, static_cast(elemCount)); + PipeBarrier(); + ClampFp32ToOutputType(outLocal, static_cast(elemCount)); + Cast(typedLocal, outLocal, RoundMode::CAST_RINT, static_cast(elemCount)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(kg_, KVOffset(b, hv, token, 0, K_), typedLocal, elemCount); + SetFlag(mte3ToMte2Event_); + WaitFlag(mte3ToMte2Event_); + } + if (rowEnd == curT) { + CopyVectorIn(typedLocal, k_, QOffset(b, h, last, 0), K_); + SetFlag(mte2ToMte3Event_); + WaitFlag(mte2ToMte3Event_); + CopyVectorOut(kg_, KVOffset(b, hv, last, 0, K_), typedLocal, K_); + SetFlag(mte3ToMte2Event_); + WaitFlag(mte3ToMte2Event_); + } + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + template + __aicore__ inline void ComputeTailWuRow(GlobalTensor &src, GlobalTensor &dst, + uint64_t akkBase, uint64_t srcBase, uint64_t dstBase, uint64_t curT, + uint64_t dim, uint64_t rowStride) + { + LocalTensor acc = vecBuf_.Get(); + LocalTensor value = vecBuf_.Get()[512]; + LocalTensor typed = vecBuf_.Get()[4096]; + LocalTensor coefficientTyped = exp2Buf_.Get(); + LocalTensor coefficients = exp2Buf_.Get()[128]; + CopyVectorIn(coefficientTyped, preparedAqk_, akkBase, curT); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast(coefficients, coefficientTyped, RoundMode::CAST_NONE, static_cast(curT)); + PipeBarrier(); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + for (uint64_t j = 0; j < curT; ++j) { + LoadAsFloatVector(src, srcBase + j * rowStride, value, typed, dim); + float coefficient = coefficients.GetValue(j); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + Muls(value, value, coefficient, static_cast(dim)); + PipeBarrier(); + if (j == 0) { + Adds(acc, value, 0.0f, static_cast(dim)); + } else { + Add(acc, acc, value, static_cast(dim)); + } + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + ClampFp32ToOutputType(acc, static_cast(dim)); + StoreFloatRow(dst, dstBase, acc, dim); + } + + __aicore__ inline void ComputeTailWuVector(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT, + uint64_t subBlockIdx, uint64_t subBlockNum) + { + // preparedQG_ and w_ alias. Each subblock owns disjoint columns and + // writes rows from last to first, so every lower-triangular source row + // remains live through its final use without desynchronizing the AIVs. + uint64_t colBegin = (K_ * subBlockIdx) / subBlockNum; + uint64_t colEnd = (K_ * (subBlockIdx + 1)) / subBlockNum; + for (uint64_t row = curT; row > 0; --row) { + uint64_t rowIdx = row - 1; + ComputeTailWuRow( + preparedQG_, w_, AOffset(b, hv, start + rowIdx, 0), KVOffset(b, hv, start, colBegin, K_), + KVOffset(b, hv, start + rowIdx, colBegin, K_), curT, colEnd - colBegin, K_); + } + + uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; + uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + for (uint64_t row = rowBegin; row < rowEnd; ++row) { + ComputeTailWuRow( + propagatedVNew_, u_, AOffset(b, hv, start + row, 0), KVOffset(b, hv, start, 0, V_), + KVOffset(b, hv, start + row, 0, V_), curT, V_, V_); + } + } + + __aicore__ inline bool ResolveFlatChunk(uint64_t task, uint64_t &seq, uint64_t &b, uint64_t &h, uint64_t &hv, + uint64_t &chunkIdx, uint64_t &start, uint64_t &end) + { + hv = task % HV_; + uint64_t flatChunk = task / HV_; + if (!isVarLen_) { + seq = flatChunk / NT_; + b = seq; + chunkIdx = flatChunk % NT_; + start = chunkIdx * BT_; + end = start + BT_; + if (end > T_) { + end = T_; + } + } else { + if (!KdaVarlen::ResolveChunkRange( + cuSeqlensAddr_, chunkIndicesAddr_, N_, T_, BT_, flatChunk, + seq, start, end)) { + return false; + } + b = 0; + chunkIdx = flatChunk; + } + h = hv / (HV_ / H_); + return start < end; + } + + __aicore__ inline void ProcessChunkPostAiv(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t end, uint64_t subBlockIdx, + uint64_t subBlockNum) + { + uint64_t curT = end - start; + if (curT == 0 || !UsePostWuCube(curT)) { + return; + } + if (curT < BT_) { + ComputeTailWuVector(b, hv, start, curT, subBlockIdx, subBlockNum); + CopyScratchWAndFinalizeKg( + b, h, hv, chunkIdx, start, curT, subBlockIdx, subBlockNum, false); + return; + } + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(syncDoneFlag_); + CopyScratchWAndFinalizeKg( + b, h, hv, chunkIdx, start, curT, subBlockIdx, subBlockNum, true); + } + + __aicore__ inline void ProcessChunkPostAic(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + if constexpr (IsSameType::value) { + ProcessChunkPostAicTyped(b, hv, chunkIdx, start, end); + } + } + + __aicore__ inline void ProcessChunkPostAicTyped(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + uint64_t curT = end - start; + if (curT == 0 || !UsePostWuCube(curT)) { + return; + } + if (curT < BT_) { + return; + } + ComputePostWuCube(b, hv, chunkIdx, start, curT); + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(syncDoneFlag_); + } + + __aicore__ inline void ProcessPostAiv() + { + if constexpr (IsSameType::value) { + return; + } + uint64_t subBlockNum = static_cast(GetSubBlockNum()); + if (subBlockNum == 0) { + return; + } + uint64_t subBlockIdx = static_cast(GetSubBlockIdx()); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + uint64_t coreIdx = static_cast(GetBlockIdx()) / subBlockNum; + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + for (uint64_t task = coreIdx; task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + ProcessChunkPostAiv(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); + } + } + } + + __aicore__ inline void ProcessPostAic() + { + if constexpr (IsSameType::value) { + return; + } + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + (void)h; + ProcessChunkPostAic(b, hv, chunkIdx, start, end); + } + } + } + + +private: + GlobalTensor q_; + GlobalTensor k_; + GlobalTensor v_; + GlobalTensor gk_; + GlobalTensor beta_; + GlobalTensor initialState_; + GlobalTensor o_; + GlobalTensor finalState_; + GlobalTensor aqk_; + GlobalTensor akk_; + GlobalTensor w_; + GlobalTensor u_; + GlobalTensor qg_; + GlobalTensor kg_; + GlobalTensor vNew_; + GlobalTensor h_; + GlobalTensor preparedQG_; + GlobalTensor preparedAqk_; + GlobalTensor propagatedVNew_; + GlobalTensor propagatedH_; + GlobalTensor solveWorkspace_; + GlobalTensor scoreWorkspace_; + TPipe *pipe_ = nullptr; + TBuf exp2Buf_; + TBuf vecBuf_; + TBuf gateWritebackBuf_; + TEventID mte2ToVEvent_ = 0; + TEventID vToMte2Event_ = 0; + TEventID vToMte3Event_ = 0; + TEventID mte3ToVEvent_ = 0; + TEventID mte2ToMte3Event_ = 0; + TEventID mte3ToMte2Event_ = 0; + bool vectorEventsAllocated_ = false; + Catlass::Arch::CrossCoreFlagWithReverse scoreReadyFlag_{KDA_SCORE_READY_FLAG0, + KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse scoreDoneFlag_{KDA_SCORE_DONE_FLAG0, + KDA_SCORE_DONE_FLAG1}; + // Score production is fully drained before solve starts, so the solve handshake can safely reuse + // the existing score flags without consuming additional hardware flag IDs. + Catlass::Arch::CrossCoreFlagWithReverse syncReadyFlag_{KDA_SCORE_READY_FLAG0, + KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse syncDoneFlag_{KDA_SCORE_DONE_FLAG0, + KDA_SCORE_DONE_FLAG1}; + uint64_t B_ = 0; + uint64_t N_ = 0; + uint64_t H_ = 0; + uint64_t HV_ = 0; + uint64_t T_ = 0; + uint64_t K_ = 0; + uint64_t V_ = 0; + uint64_t BT_ = 0; + uint64_t NT_ = 0; + float scale_ = 1.0f; + bool hasInitial_ = false; + bool isVarLen_ = false; + bool isAivOnly_ = false; + bool inputSequenceMajor_ = false; + uint64_t usedCoreNum_ = 1; + uint64_t solveCoreIdx_ = 0; + __gm__ int64_t *chunkIndicesAddr_ = nullptr; + __gm__ int64_t *cuSeqlensAddr_ = nullptr; +}; +} // namespace + +template +__aicore__ inline void RunChunkKdaPostWu( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR wSeed, GM_ADDR akk, GM_ADDR uSeed, + GM_ADDR w, GM_ADDR u, GM_ADDR kg, GM_ADDR vNew, GM_ADDR userWorkspace, + const TilingData &tiling, TPipe &pipe) +{ + GM_ADDR postScratch = userWorkspace + tiling.postWuScratchOffset; + if ASCEND_IS_AIC { + ChunkKdaFwdPostWuKernel op; + op.Init(q, k, v, gk, beta, initialState, cuSeqlens, chunkIndices, + wSeed, akk, uSeed, nullptr, userWorkspace, userWorkspace, userWorkspace, akk, w, u, + userWorkspace, kg, vNew, postScratch, postScratch, tiling, &pipe, false); + op.ProcessAic(); + } + if ASCEND_IS_AIV { + ChunkKdaFwdPostWuKernel op; + op.Init(q, k, v, gk, beta, initialState, cuSeqlens, chunkIndices, + wSeed, akk, uSeed, nullptr, userWorkspace, userWorkspace, userWorkspace, akk, w, u, + userWorkspace, kg, vNew, postScratch, postScratch, tiling, &pipe); + op.ProcessAiv(); + } +} + +} // namespace KdaPostWu diff --git a/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_prepare.h b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_prepare.h new file mode 100644 index 000000000000..21b5c420fb8d --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_prepare.h @@ -0,0 +1,2613 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#pragma once + +#ifndef CATLASS_ARCH +#define CATLASS_ARCH 2201 +#endif + +#include "catlass/arch/arch.hpp" +#include "catlass/arch/cross_core_sync.hpp" +#include "catlass/arch/resource.hpp" +#include "catlass/catlass.hpp" +#include "catlass/gemm/block/block_mmad.hpp" +#include "catlass/gemm/dispatch_policy.hpp" +#include "catlass/gemm/tile/tile_copy.hpp" +#include "catlass/gemm_coord.hpp" +#include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp" +#include "catlass/layout/layout.hpp" +#include "kernel_operator.h" +#include "chunk_kda_fwd_varlen.h" +#include "tla/layout.hpp" +#include "tla/tensor.hpp" + +using namespace AscendC; + +namespace KdaPrepare { +namespace { +using KdaInt64 = tla::Int<64>; +using KdaInt128 = tla::Int<128>; +constexpr float LN2 = 0.69314718055994530942f; +constexpr float KDA_EXP2_CLAMP = 80.0f; +constexpr float KDA_EXP_INPUT_MAX = KDA_EXP2_CLAMP * LN2; +constexpr float KDA_EXP_INPUT_MIN = -KDA_EXP2_CLAMP * LN2; +constexpr float KDA_SCORE_EXP2_CLAMP = 120.0f; +constexpr float KDA_SCORE_EXP2_MIN_CLAMP = 126.0f; +constexpr float KDA_SCORE_EXP_INPUT_MAX = KDA_SCORE_EXP2_CLAMP * LN2; +constexpr float KDA_SCORE_EXP_INPUT_MIN = -KDA_SCORE_EXP2_MIN_CLAMP * LN2; +constexpr float KDA_FP16_MAX = 65504.0f; +constexpr uint32_t EXP2_UB_ELEMENTS = 256; +constexpr uint32_t EXP2_UB_BYTES = EXP2_UB_ELEMENTS * (sizeof(float) + sizeof(uint16_t)); +constexpr uint32_t EXP2_EVENT_ID = 0; +constexpr uint32_t KDA_SOLVE_BT = 64; +constexpr uint32_t KDA_SOLVE_MATRIX_ELEMENTS = KDA_SOLVE_BT * KDA_SOLVE_BT; +constexpr uint32_t KDA_SOLVE_SCRATCH_X = 0; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y0 = 1; +constexpr uint32_t KDA_SOLVE_SCRATCH_TMP = 2; +constexpr uint32_t KDA_SOLVE_SCRATCH_Y1 = 3; +constexpr uint32_t KDA_SOLVE_SCRATCH_IDENTITY = 4; +constexpr uint32_t KDA_SOLVE_SCRATCH_SLOTS = 5; +constexpr uint32_t KDA_SOLVE_PIPELINE_DEPTH = 4; +constexpr uint32_t KDA_SOLVE_DIAG_BT = 16; +constexpr uint32_t KDA_SOLVE_DIAG_BLOCKS = KDA_SOLVE_BT / KDA_SOLVE_DIAG_BT; +constexpr uint32_t KDA_SOLVE_DIAG_MCH_ITERS = 3; +// Keep the local safe-gate exponent span within the BF16 score range while +// reducing repeated gate-factor work and AIV/AIC handshakes. +constexpr uint32_t KDA_SCORE_REF_BC = 32; +constexpr uint32_t KDA_SAFE_SCORE_REF_BC = 32; +constexpr uint32_t KDA_VEC_ARENA_ELEMENTS = 32768; +constexpr uint32_t KDA_BITS_PER_MASK_BYTE = 8; +constexpr uint32_t KDA_SELECT_COL_BLOCKS = 2; +constexpr uint32_t KDA_SELECT_COL_MASK_BYTES = KDA_SOLVE_MATRIX_ELEMENTS / KDA_BITS_PER_MASK_BYTE; +constexpr uint32_t KDA_SELECT_MASK_BYTES = KDA_SELECT_COL_BLOCKS * KDA_SELECT_COL_MASK_BYTES; +constexpr uint32_t KDA_SELECT_AQK_MASK_BYTE_OFFSET = 120 * 1024; +constexpr uint32_t KDA_SELECT_AKK_MASK_BYTE_OFFSET = KDA_SELECT_AQK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_BYTE_OFFSET = KDA_SELECT_AKK_MASK_BYTE_OFFSET + KDA_SELECT_MASK_BYTES; +constexpr uint32_t KDA_SELECT_ZERO_FLOAT_OFFSET = KDA_SELECT_ZERO_BYTE_OFFSET / sizeof(float); +constexpr uint8_t KDA_SCORE_DONE_FLAG0 = 2; +constexpr uint8_t KDA_SCORE_DONE_FLAG1 = 3; +constexpr uint8_t KDA_SCORE_READY_FLAG0 = 4; +constexpr uint8_t KDA_SCORE_READY_FLAG1 = 5; +constexpr uint8_t KDA_SOLVE_DONE_FLAG = 6; +constexpr uint8_t KDA_SOLVE_READY_FLAG = 7; +constexpr uint32_t KDA_SCORE_QUEUE_DEPTH = 2; +constexpr uint32_t KDA_SCORE_LANES = 2; +constexpr uint32_t KDA_SCORE_SCRATCH_SLOTS = KDA_SCORE_QUEUE_DEPTH * KDA_SCORE_LANES; +constexpr uint32_t KDA_SYNC_REVERSE_DEPTH = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_PLANES = 3; +constexpr uint32_t KDA_SCORE_SCRATCH_QG = 0; +constexpr uint32_t KDA_SCORE_SCRATCH_W = 1; +constexpr uint32_t KDA_SCORE_SCRATCH_KG = 2; +constexpr uint64_t KDA_WORKSPACE_ALIGN = 512; +constexpr uint32_t KDA_GATE_TILE_ROWS = 16; +constexpr uint32_t KDA_GATE_PIPELINE_DEPTH = 3; +constexpr uint32_t KDA_AIV_UB_BUDGET_BYTES = 192 * 1024; +using KdaArchTag = Catlass::Arch::AtlasA2; +using KdaDispatchPolicy = Catlass::Gemm::MmadPingpong; +using KdaScoreDispatchPolicy = + Catlass::Gemm::MmadPingpongTlaMulti; +static_assert(KdaScoreDispatchPolicy::ENABLE_L1_RESIDENT, + "KDA Aqk/Akk score MMAD must keep the shared right matrix resident in L1"); +static_assert(KdaScoreDispatchPolicy::L1B_STAGES == 1, + "KDA Aqk/Akk score MMAD needs one L1 B slot so the second MMAD reuses it"); +using KdaSolveDispatchPolicy = Catlass::Gemm::MmadPingpong; +static_assert(!KdaSolveDispatchPolicy::USE_HF32_MODE, "KDA triangular solve must use IEEE FP32 Cube mode"); +using KdaL1TileShape = tla::Shape; +using KdaL0TileShape = KdaL1TileShape; +using KdaSolveL1TileShape = tla::Shape; +using KdaSolveL0TileShape = KdaSolveL1TileShape; + +__aicore__ inline uint32_t FloatToBits(float value) +{ + union Bits { + __aicore__ Bits() {} + float f; + uint32_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline float BitsToFloat(uint32_t value) +{ + union Bits { + __aicore__ Bits() {} + uint32_t u; + float f; + } bits; + bits.u = value; + return bits.f; +} + +__aicore__ inline uint16_t Bf16ToBits(bfloat16_t value) +{ + union Bits { + __aicore__ Bits() {} + bfloat16_t f; + uint16_t u; + } bits; + bits.f = value; + return bits.u; +} + +__aicore__ inline bfloat16_t BitsToBf16(uint16_t value) +{ + union Bits { + __aicore__ Bits() {} + uint16_t u; + bfloat16_t f; + } bits; + bits.u = value; + return bits.f; +} + +template +__aicore__ inline T FloatToType(float value) +{ + if constexpr (IsSameType::value) { + uint32_t bits = FloatToBits(value); + uint32_t bias = 0x7FFFu + ((bits >> 16) & 1u); + return BitsToBf16(static_cast((bits + bias) >> 16)); + } + return static_cast(value); +} + +template +class ChunkKdaFwdPrepareKernel { +public: + using OUT_T = T; + using AKK_T = float; + using SCORE_T = + std::conditional_t::value, bfloat16_t, T>; + template + __aicore__ inline void Init(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR preparedQG, GM_ADDR preparedAqk, + GM_ADDR propagatedVNew, GM_ADDR propagatedH, GM_ADDR o, GM_ADDR finalState, GM_ADDR aqk, + GM_ADDR akk, GM_ADDR w, GM_ADDR u, GM_ADDR qg, GM_ADDR kg, GM_ADDR vNew, GM_ADDR h, + GM_ADDR workspace, const TilingData &tiling, TPipe *pipe, + bool initVecBuffers = true) + { + pipe_ = pipe; + q_.SetGlobalBuffer((__gm__ T *)q); + k_.SetGlobalBuffer((__gm__ T *)k); + v_.SetGlobalBuffer((__gm__ T *)v); + gk_.SetGlobalBuffer((__gm__ GK_T *)gk); + beta_.SetGlobalBuffer((__gm__ BETA_T *)beta); + if (initialState != nullptr) { + initialState_.SetGlobalBuffer((__gm__ float *)initialState); + } + cuSeqlensAddr_ = reinterpret_cast<__gm__ int64_t *>(cuSeqlens); + if (preparedQG != nullptr) { + preparedQG_.SetGlobalBuffer((__gm__ T *)preparedQG); + } + if (preparedAqk != nullptr) { + preparedAqk_.SetGlobalBuffer((__gm__ T *)preparedAqk); + } + if (propagatedVNew != nullptr) { + propagatedVNew_.SetGlobalBuffer((__gm__ T *)propagatedVNew); + } + if (propagatedH != nullptr) { + propagatedH_.SetGlobalBuffer((__gm__ T *)propagatedH); + } + chunkIndicesAddr_ = reinterpret_cast<__gm__ int64_t *>(chunkIndices); + o_.SetGlobalBuffer((__gm__ OUT_T *)o); + finalState_.SetGlobalBuffer((__gm__ float *)finalState); + aqk_.SetGlobalBuffer((__gm__ float *)aqk); + akk_.SetGlobalBuffer((__gm__ AKK_T *)akk); + w_.SetGlobalBuffer((__gm__ T *)w); + u_.SetGlobalBuffer((__gm__ OUT_T *)u); + qg_.SetGlobalBuffer((__gm__ T *)qg); + kg_.SetGlobalBuffer((__gm__ T *)kg); + vNew_.SetGlobalBuffer((__gm__ T *)vNew); + h_.SetGlobalBuffer((__gm__ float *)h); + solveWorkspace_.SetGlobalBuffer((__gm__ float *)workspace); + + B_ = tiling.batch; + N_ = tiling.seqNum; + H_ = tiling.qHeadNum; + HV_ = tiling.vHeadNum; + T_ = tiling.seqlen; + K_ = tiling.kHeadDim; + V_ = tiling.vHeadDim; + BT_ = tiling.chunkSize; + NT_ = tiling.totalChunks; + scale_ = tiling.scale; + hasInitial_ = tiling.hasInitialState; + isVarLen_ = tiling.isVarLen; + inputSequenceMajor_ = tiling.inputSequenceMajor; + usedCoreNum_ = tiling.prepareUsedCoreNum; + constexpr uint64_t solvePipelineDepth = SAFE_GATE ? KDA_SOLVE_PIPELINE_DEPTH : 1; + const uint64_t solveBytes = + usedCoreNum_ * solvePipelineDepth * KDA_SOLVE_SCRATCH_SLOTS * BT_ * BT_ * sizeof(float); + const uint64_t alignedSolveBytes = + (solveBytes + KDA_WORKSPACE_ALIGN - 1) / KDA_WORKSPACE_ALIGN * KDA_WORKSPACE_ALIGN; + scoreWorkspace_.SetGlobalBuffer((__gm__ SCORE_T *)(workspace + alignedSolveBytes)); + if ASCEND_IS_AIV { + uint64_t subBlockNum = static_cast(GetSubBlockNum()); + solveCoreIdx_ = subBlockNum == 0 ? 0 : static_cast(GetBlockIdx()) / subBlockNum; + } else { + solveCoreIdx_ = static_cast(GetBlockIdx()); + } + if (pipe_ != nullptr && initVecBuffers) { + pipe_->InitBuffer(exp2Buf_, EXP2_UB_BYTES); + pipe_->InitBuffer(vecBuf_, KDA_VEC_ARENA_ELEMENTS * sizeof(float)); + const uint64_t gateStageElems = GatePipelineRows() * K_; + const uint64_t gateInputSlotBytes = gateStageElems * (2 * sizeof(T) + sizeof(GK_T)); + const uint64_t gatePipelineBytes = + KDA_GATE_PIPELINE_DEPTH * (gateInputSlotBytes + gateStageElems * sizeof(T)); + pipe_->InitBuffer(gateWritebackBuf_, static_cast(gatePipelineBytes)); + AllocVectorEvents(); + } + } + __aicore__ inline void ProcessAivOnly() + { + isAivOnly_ = true; + ProcessPreAiv(); + ReleaseVectorEvents(); + } + + __aicore__ inline void ProcessAiv() + { + ProcessPreAiv(); + ReleaseVectorEvents(); + } + + __aicore__ inline void ProcessAic() + { + ProcessPreAic(); + } + +private: + __aicore__ inline void AllocVectorEvents() + { + mte2ToVEvent_ = pipe_->AllocEventID(); + vToMte2Event_ = pipe_->AllocEventID(); + vToMte3Event_ = pipe_->AllocEventID(); + mte3ToVEvent_ = pipe_->AllocEventID(); + mte2ToMte3Event_ = pipe_->AllocEventID(); + for (uint32_t slot = 0; slot < KDA_GATE_PIPELINE_DEPTH; ++slot) { + mte3ToMte2Events_[slot] = pipe_->AllocEventID(); + } + vectorEventsAllocated_ = true; + } + + __aicore__ inline void ReleaseVectorEvents() + { + if (!vectorEventsAllocated_) { + return; + } + pipe_->ReleaseEventID(mte2ToVEvent_); + pipe_->ReleaseEventID(vToMte2Event_); + pipe_->ReleaseEventID(vToMte3Event_); + pipe_->ReleaseEventID(mte3ToVEvent_); + pipe_->ReleaseEventID(mte2ToMte3Event_); + for (uint32_t slot = 0; slot < KDA_GATE_PIPELINE_DEPTH; ++slot) { + pipe_->ReleaseEventID(mte3ToMte2Events_[slot]); + } + vectorEventsAllocated_ = false; + } + + __aicore__ inline uint64_t QOffset(uint64_t b, uint64_t h, uint64_t t, uint64_t d) const + { + if (inputSequenceMajor_) { + return ((b * T_ + t) * H_ + h) * K_ + d; + } + return ((b * H_ + h) * T_ + t) * K_ + d; + } + + __aicore__ inline uint64_t VInputOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t d) const + { + if (inputSequenceMajor_) { + return ((b * T_ + t) * HV_ + hv) * V_ + d; + } + return ((b * HV_ + hv) * T_ + t) * V_ + d; + } + + __aicore__ inline uint64_t KVOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t d, uint64_t dim) const + { + return ((b * HV_ + hv) * T_ + t) * dim + d; + } + + __aicore__ inline uint64_t BetaOffset(uint64_t b, uint64_t hv, uint64_t t) const + { + return (b * HV_ + hv) * T_ + t; + } + + __aicore__ inline uint64_t AOffset(uint64_t b, uint64_t hv, uint64_t t, uint64_t j) const + { + return ((b * HV_ + hv) * T_ + t) * BT_ + j; + } + + __aicore__ inline uint64_t HOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t d, uint64_t r) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * K_ + d) * V_ + r; + } + + __aicore__ inline uint64_t WScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t t, uint64_t d) const + { + return (((b * HV_ + hv) * NT_ + chunkIdx) * BT_ + t) * K_ + d; + } + + __aicore__ inline uint64_t SolveScratchOffset(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t slot) const + { + (void)b; + (void)hv; + (void)chunkIdx; + constexpr uint64_t solvePipelineDepth = SAFE_GATE ? KDA_SOLVE_PIPELINE_DEPTH : 1; + uint64_t matrixElements = BT_ * BT_; + return ((solveCoreIdx_ * solvePipelineDepth + activeSolveSlot_) * KDA_SOLVE_SCRATCH_SLOTS + slot) * + matrixElements; + } + + __aicore__ inline uint64_t ScoreScratchOffset(uint64_t slot, uint64_t plane, uint64_t t = 0, + uint64_t d = 0) const + { + return (((solveCoreIdx_ * KDA_SCORE_SCRATCH_SLOTS + slot) * KDA_SCORE_SCRATCH_PLANES + plane) * BT_ + t) * + K_ + + d; + } + + __aicore__ inline uint64_t ScoreScratchSlot(uint64_t queueSlot, uint64_t lane, bool pairHeads) const + { + return pairHeads ? queueSlot * KDA_SCORE_LANES + lane : queueSlot; + } + + + + __aicore__ inline uint64_t ScoreRefBlockSize() const + { + if constexpr (SAFE_GATE) { + return KDA_SAFE_SCORE_REF_BC; + } + return KDA_SCORE_REF_BC; + } + + __aicore__ inline uint64_t ScoreRowBlockCount(uint64_t curT, uint64_t rowBegin) const + { + uint64_t blockSize = ScoreRefBlockSize(); + uint64_t rowCount = curT - rowBegin; + if (rowCount > blockSize) { + rowCount = blockSize; + } + return rowCount; + } + + __aicore__ inline uint64_t ScoreRefToken(uint64_t start, uint64_t curT, uint64_t rowBegin, + uint64_t rowCount) const + { + uint64_t ref = rowBegin + rowCount / 2; + if (ref >= curT) { + ref = curT - 1; + } + return start + ref; + } + + __aicore__ inline void RunExp2(LocalTensor &tensor, uint32_t count) + { + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + ClampExpInput(tensor, count); + Exp(tensor, tensor, count); + PipeBarrier(); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + } + + __aicore__ inline void ClampExpInput(LocalTensor &tensor, uint32_t count) + { + Mins(tensor, tensor, KDA_EXP_INPUT_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, KDA_EXP_INPUT_MIN, count); + PipeBarrier(); + } + + __aicore__ inline void ClampScoreExpInput(LocalTensor &tensor, uint32_t count) + { + constexpr float expInputMax = + IsSameType::value ? KDA_SCORE_EXP_INPUT_MAX : KDA_EXP_INPUT_MAX; + constexpr float expInputMin = + IsSameType::value ? KDA_SCORE_EXP_INPUT_MIN : KDA_EXP_INPUT_MIN; + Mins(tensor, tensor, expInputMax, count); + PipeBarrier(); + Maxs(tensor, tensor, expInputMin, count); + PipeBarrier(); + } + + template + __aicore__ inline void ClampFp32ForCast(LocalTensor &tensor, uint32_t count) + { + if constexpr (IsSameType::value) { + Mins(tensor, tensor, KDA_FP16_MAX, count); + PipeBarrier(); + Maxs(tensor, tensor, -KDA_FP16_MAX, count); + PipeBarrier(); + } + } + + __aicore__ inline void ClampFp32ToOutputType(LocalTensor &tensor, uint32_t count) + { + ClampFp32ForCast(tensor, count); + } + + template + __aicore__ inline void CopyVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst, src[offset], static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + + template + __aicore__ inline void CopyRowsIn(LocalTensor &dst, GlobalTensor &src, + uint64_t offset, uint64_t rows, uint64_t cols, + uint64_t rowStride) + { + if (rows == 0 || cols == 0) { + return; + } + if (rowStride == cols) { + CopyVectorIn(dst, src, offset, rows * cols); + return; + } + DataCopyExtParams params{ + static_cast(rows), + static_cast(cols * sizeof(CopyT)), + static_cast((rowStride - cols) * sizeof(CopyT)), + 0, + 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + + template + __aicore__ inline void CopyVectorOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src, + uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(CopyT)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst[offset], src, static_cast(count)); + return; + } + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPad(dst[offset], src, params); + } + + template + __aicore__ inline void CopyRowIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset) + { + CopyVectorIn(dst, src, offset, K_); + } + + template + __aicore__ inline void CopyRowOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src) + { + CopyVectorOut(dst, offset, src, K_); + } + + __aicore__ inline LocalTensor VecScratch(uint64_t slot) + { + return vecBuf_.Get()[slot * EXP2_UB_ELEMENTS]; + } + + __aicore__ inline uint64_t GateStageElems() const + { + return GatePipelineRows() * K_; + } + + __aicore__ inline uint64_t GatePipelineRows() const + { + constexpr uint64_t fixedBytes = + static_cast(KDA_VEC_ARENA_ELEMENTS) * sizeof(float) + EXP2_UB_BYTES; + constexpr uint64_t availableBytes = KDA_AIV_UB_BUDGET_BYTES - fixedBytes; + uint64_t bytesPerRow = + K_ * KDA_GATE_PIPELINE_DEPTH * (3 * sizeof(T) + sizeof(GK_T)); + uint64_t rows = bytesPerRow == 0 ? 0 : availableBytes / bytesPerRow; + return rows < KDA_GATE_TILE_ROWS ? rows : KDA_GATE_TILE_ROWS; + } + + __aicore__ inline uint64_t GateInputSlotBytes() const + { + return GateStageElems() * (2 * sizeof(T) + sizeof(GK_T)); + } + + __aicore__ inline LocalTensor GateQTyped(uint64_t slot) + { + uint64_t byteOffset = slot * GateInputSlotBytes(); + return gateWritebackBuf_.Get()[byteOffset / sizeof(T)]; + } + + __aicore__ inline LocalTensor GateKTyped(uint64_t slot) + { + uint64_t byteOffset = slot * GateInputSlotBytes() + GateStageElems() * sizeof(T); + return gateWritebackBuf_.Get()[byteOffset / sizeof(T)]; + } + + __aicore__ inline LocalTensor GateGTyped(uint64_t slot) + { + uint64_t byteOffset = slot * GateInputSlotBytes() + 2 * GateStageElems() * sizeof(T); + return gateWritebackBuf_.Get()[byteOffset / sizeof(GK_T)]; + } + + __aicore__ inline LocalTensor GateKgTyped(uint64_t slot) + { + uint64_t byteOffset = KDA_GATE_PIPELINE_DEPTH * GateInputSlotBytes() + + slot * GateStageElems() * sizeof(T); + return gateWritebackBuf_.Get()[byteOffset / sizeof(T)]; + } + + __aicore__ inline void PrefetchQKGate(uint64_t slot, uint64_t b, uint64_t h, uint64_t hv, + uint64_t token, uint64_t elems) + { + const uint64_t rows = elems / K_; + LocalTensor qTyped = GateQTyped(slot); + LocalTensor kTyped = GateKTyped(slot); + LocalTensor gateTyped = GateGTyped(slot); + CopyRowsIn(qTyped, q_, QOffset(b, h, token, 0), rows, K_, + inputSequenceMajor_ ? H_ * K_ : K_); + CopyRowsIn(kTyped, k_, QOffset(b, h, token, 0), rows, K_, + inputSequenceMajor_ ? H_ * K_ : K_); + CopyVectorIn(gateTyped, gk_, KVOffset(b, hv, token, 0, K_), elems); + SetFlag(mte2ToVEvent_); + } + + __aicore__ inline void PrefetchKGate(uint64_t slot, uint64_t b, uint64_t h, uint64_t hv, + uint64_t token, uint64_t elems) + { + const uint64_t rows = elems / K_; + LocalTensor kTyped = GateQTyped(slot); + LocalTensor gateTyped = GateGTyped(slot); + CopyRowsIn(kTyped, k_, QOffset(b, h, token, 0), rows, K_, + inputSequenceMajor_ ? H_ * K_ : K_); + CopyVectorIn(gateTyped, gk_, KVOffset(b, hv, token, 0, K_), elems); + SetFlag(mte2ToVEvent_); + } + + __aicore__ inline void WaitGateInputReady() + { + WaitFlag(mte2ToVEvent_); + } + + __aicore__ inline void WaitGateOutputForMte2(uint64_t slot = 0) + { + WaitFlag(mte3ToMte2Events_[slot]); + } + + __aicore__ inline void WaitGateOutputForVector() + { + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void SignalGateOutputDone() + { + SetFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + } + + __aicore__ inline void SignalGateOutputDoneForMte2(uint64_t slot) + { + SetFlag(mte3ToMte2Events_[slot]); + } + + template + __aicore__ inline void LoadAsFloatRow(GlobalTensor &src, uint64_t srcOffset, LocalTensor &dst, + uint64_t count) + { + if constexpr (IsSameType::value) { + CopyVectorIn(dst, src, srcOffset, count); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Adds(dst, dst, 0.0f, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } else { + constexpr uint32_t typedOffset = EXP2_UB_ELEMENTS * sizeof(float) / sizeof(CopyT); + LocalTensor rowLocal = exp2Buf_.Get()[typedOffset]; + CopyVectorIn(rowLocal, src, srcOffset, count); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast(dst, rowLocal, RoundMode::CAST_NONE, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + } + PipeBarrier(); + } + + template + __aicore__ inline void LoadAsFloatVector(GlobalTensor &src, uint64_t srcOffset, + LocalTensor &dst, LocalTensor &typedScratch, + uint64_t count) + { + if constexpr (IsSameType::value) { + CopyVectorIn(dst, src, srcOffset, count); + } else { + CopyVectorIn(typedScratch, src, srcOffset, count); + } + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + if constexpr (!IsSameType::value) { + Cast(dst, typedScratch, RoundMode::CAST_NONE, static_cast(count)); + PipeBarrier(); + } + } + + template + __aicore__ inline void StoreFloatRow(GlobalTensor &dst, uint64_t dstOffset, LocalTensor &src, + uint64_t count) + { + if constexpr (IsSameType::value) { + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(dst, dstOffset, src, count); + } else { + constexpr uint32_t typedOffset = EXP2_UB_ELEMENTS * sizeof(float) / sizeof(CopyT); + LocalTensor rowLocal = exp2Buf_.Get()[typedOffset]; + Cast(rowLocal, src, RoundMode::CAST_RINT, static_cast(count)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(dst, dstOffset, rowLocal, count); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + + + + + __aicore__ inline LocalTensor Exp2NegG(uint64_t b, uint64_t hv, uint64_t t) + { + LocalTensor exp2Local = exp2Buf_.Get(); + LoadAsFloatRow(gk_, KVOffset(b, hv, t, 0, K_), exp2Local, K_); + Muls(exp2Local, exp2Local, -LN2, static_cast(K_)); + PipeBarrier(); + RunExp2(exp2Local, static_cast(K_)); + return exp2Local; + } + + + __aicore__ inline void PrepareScoreFactorsBulk(uint64_t b, uint64_t h, uint64_t hv, uint64_t start, + uint64_t subBlockIdx, uint64_t subBlockNum, + uint64_t refToken, uint64_t scoreRowBegin, + uint64_t scoreRowCount, uint64_t validColEnd, + uint64_t scoreSlot) + { + LocalTensor refFp32 = exp2Buf_.Get(); + LoadAsFloatRow(gk_, KVOffset(b, hv, refToken, 0, K_), refFp32, K_); + + uint64_t qwBegin = scoreRowBegin + (scoreRowCount * subBlockIdx) / subBlockNum; + uint64_t qwEnd = scoreRowBegin + (scoreRowCount * (subBlockIdx + 1)) / subBlockNum; + uint64_t qwMaxRows = GatePipelineRows(); + bool qwOutputPending = false; + uint64_t qwSlot = 0; + if (qwBegin < qwEnd && qwMaxRows > 0) { + uint64_t firstRows = qwEnd - qwBegin; + if (firstRows > qwMaxRows) { + firstRows = qwMaxRows; + } + PrefetchQKGate(qwSlot, b, h, hv, start + qwBegin, firstRows * K_); + } + for (uint64_t tileRow = qwBegin; tileRow < qwEnd && qwMaxRows > 0; tileRow += qwMaxRows) { + uint64_t tileRows = qwEnd - tileRow; + if (tileRows > qwMaxRows) { + tileRows = qwMaxRows; + } + uint64_t elems = tileRows * K_; + LocalTensor qTyped = GateQTyped(qwSlot); + LocalTensor kTyped = GateKTyped(qwSlot); + LocalTensor qScore = qTyped.template ReinterpretCast(); + LocalTensor kScore = kTyped.template ReinterpretCast(); + LocalTensor gateTyped = GateGTyped(qwSlot); + LocalTensor arena = vecBuf_.Get(); + LocalTensor qFp32 = arena; + LocalTensor kFp32 = arena[elems]; + LocalTensor gFp32 = arena[2 * elems]; + LocalTensor expFp32 = arena[3 * elems]; + LocalTensor outFp32 = arena[4 * elems]; + + WaitGateInputReady(); + Cast(qFp32, qTyped, RoundMode::CAST_NONE, static_cast(elems)); + Cast(kFp32, kTyped, RoundMode::CAST_NONE, static_cast(elems)); + if constexpr (IsSameType::value) { + gFp32 = gateTyped; + } else { + Cast(gFp32, gateTyped, RoundMode::CAST_NONE, static_cast(elems)); + } + if (qwOutputPending) { + WaitGateOutputForMte2(); + } + uint64_t nextTileRow = tileRow + qwMaxRows; + if (nextTileRow < qwEnd) { + uint64_t nextRows = qwEnd - nextTileRow; + if (nextRows > qwMaxRows) { + nextRows = qwMaxRows; + } + PrefetchQKGate(qwSlot ^ 1, b, h, hv, start + nextTileRow, nextRows * K_); + } + PipeBarrier(); + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expFp32[row * K_], gFp32[row * K_], refFp32, static_cast(K_)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(elems)); + PipeBarrier(); + ClampScoreExpInput(expFp32, static_cast(elems)); + Exp(expFp32, expFp32, static_cast(elems)); + PipeBarrier(); + + Mul(outFp32, qFp32, expFp32, static_cast(elems)); + PipeBarrier(); + ClampFp32ForCast(outFp32, static_cast(elems)); + Cast(qScore, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); + Mul(outFp32, kFp32, expFp32, static_cast(elems)); + PipeBarrier(); + ClampFp32ForCast(outFp32, static_cast(elems)); + Cast(kScore, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); + + if (qwOutputPending) { + WaitGateOutputForVector(); + } + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG, tileRow), + qScore, elems); + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W, tileRow), + kScore, elems); + SignalGateOutputDone(); + qwOutputPending = true; + qwSlot ^= 1; + } + if (qwOutputPending) { + WaitGateOutputForMte2(); + WaitGateOutputForVector(); + } + + uint64_t kgBegin = (validColEnd * subBlockIdx) / subBlockNum; + uint64_t kgEnd = (validColEnd * (subBlockIdx + 1)) / subBlockNum; + uint64_t kgMaxRows = GatePipelineRows(); + bool kgOutputPending = false; + uint64_t kgSlot = 0; + if (kgBegin < kgEnd && kgMaxRows > 0) { + uint64_t firstRows = kgEnd - kgBegin; + if (firstRows > kgMaxRows) { + firstRows = kgMaxRows; + } + PrefetchKGate(kgSlot, b, h, hv, start + kgBegin, firstRows * K_); + } + for (uint64_t tileRow = kgBegin; tileRow < kgEnd && kgMaxRows > 0; tileRow += kgMaxRows) { + uint64_t tileRows = kgEnd - tileRow; + if (tileRows > kgMaxRows) { + tileRows = kgMaxRows; + } + uint64_t elems = tileRows * K_; + LocalTensor kTyped = GateQTyped(kgSlot); + LocalTensor kgScore = kTyped.template ReinterpretCast(); + LocalTensor gateTyped = GateGTyped(kgSlot); + LocalTensor arena = vecBuf_.Get(); + LocalTensor kFp32 = arena; + LocalTensor gFp32 = arena[elems]; + LocalTensor expFp32 = arena[2 * elems]; + LocalTensor outFp32 = arena[3 * elems]; + + WaitGateInputReady(); + Cast(kFp32, kTyped, RoundMode::CAST_NONE, static_cast(elems)); + if constexpr (IsSameType::value) { + gFp32 = gateTyped; + } else { + Cast(gFp32, gateTyped, RoundMode::CAST_NONE, static_cast(elems)); + } + if (kgOutputPending) { + WaitGateOutputForMte2(); + } + uint64_t nextTileRow = tileRow + kgMaxRows; + if (nextTileRow < kgEnd) { + uint64_t nextRows = kgEnd - nextTileRow; + if (nextRows > kgMaxRows) { + nextRows = kgMaxRows; + } + PrefetchKGate(kgSlot ^ 1, b, h, hv, start + nextTileRow, nextRows * K_); + } + PipeBarrier(); + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expFp32[row * K_], refFp32, gFp32[row * K_], static_cast(K_)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(elems)); + PipeBarrier(); + ClampScoreExpInput(expFp32, static_cast(elems)); + Exp(expFp32, expFp32, static_cast(elems)); + PipeBarrier(); + Mul(outFp32, kFp32, expFp32, static_cast(elems)); + PipeBarrier(); + ClampFp32ForCast(outFp32, static_cast(elems)); + Cast(kgScore, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + PipeBarrier(); + + if (kgOutputPending) { + WaitGateOutputForVector(); + } + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG, tileRow), + kgScore, elems); + SignalGateOutputDone(); + kgOutputPending = true; + kgSlot ^= 1; + } + if (kgOutputPending) { + WaitGateOutputForMte2(); + WaitGateOutputForVector(); + } + } + + __aicore__ inline void PrepareGateProductsBulk(uint64_t b, uint64_t h, uint64_t hv, uint64_t start, + uint64_t curT, uint64_t subBlockIdx, uint64_t subBlockNum, + bool useRef, uint64_t refToken, uint64_t validColEnd, + bool writeScoreScratch, uint64_t scoreSlot) + { + if constexpr (IsSameType::value) { + return; + } + if (subBlockNum == 0 || subBlockIdx >= subBlockNum || K_ == 0) { + return; + } + uint64_t rowBegin = (curT * subBlockIdx) / subBlockNum; + uint64_t rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + if (rowBegin >= rowEnd) { + return; + } + + uint64_t maxRows = GatePipelineRows(); + if (maxRows == 0) { + return; + } + LocalTensor refFp32 = exp2Buf_.Get(); + if (useRef) { + LoadAsFloatRow(gk_, KVOffset(b, hv, refToken, 0, K_), refFp32, K_); + } + + bool outputPending = false; + uint64_t gateSlot = 0; + uint64_t firstRows = rowEnd - rowBegin; + if (firstRows > maxRows) { + firstRows = maxRows; + } + PrefetchQKGate(gateSlot, b, h, hv, start + rowBegin, firstRows * K_); + for (uint64_t tileRow = rowBegin; tileRow < rowEnd; tileRow += maxRows) { + uint64_t tileRows = rowEnd - tileRow; + if (tileRows > maxRows) { + tileRows = maxRows; + } + uint64_t elems = tileRows * K_; + LocalTensor qTyped = GateQTyped(gateSlot); + LocalTensor kTyped = GateKTyped(gateSlot); + LocalTensor kgTyped = GateKgTyped(gateSlot); + LocalTensor qScore = qTyped.template ReinterpretCast(); + LocalTensor wScore = kTyped.template ReinterpretCast(); + LocalTensor kgScore = kgTyped.template ReinterpretCast(); + LocalTensor gateTyped = GateGTyped(gateSlot); + LocalTensor arena = vecBuf_.Get(); + LocalTensor qFp32 = arena; + LocalTensor kFp32 = arena[elems]; + LocalTensor gFp32 = arena[2 * elems]; + LocalTensor expFp32 = arena[3 * elems]; + LocalTensor outFp32 = arena[4 * elems]; + + uint64_t token = start + tileRow; + WaitGateInputReady(); + Cast(qFp32, qTyped, RoundMode::CAST_NONE, static_cast(elems)); + Cast(kFp32, kTyped, RoundMode::CAST_NONE, static_cast(elems)); + if constexpr (IsSameType::value) { + gFp32 = gateTyped; + } else { + Cast(gFp32, gateTyped, RoundMode::CAST_NONE, static_cast(elems)); + } + uint64_t nextTileRow = tileRow + maxRows; + if (outputPending) { + WaitGateOutputForMte2(); + } + if (nextTileRow < rowEnd) { + uint64_t nextRows = rowEnd - nextTileRow; + if (nextRows > maxRows) { + nextRows = maxRows; + } + PrefetchQKGate(gateSlot ^ 1, b, h, hv, start + nextTileRow, nextRows * K_); + } + PipeBarrier(); + + if (useRef) { + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expFp32[row * K_], gFp32[row * K_], refFp32, static_cast(K_)); + } + } else { + Adds(expFp32, gFp32, 0.0f, static_cast(elems)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(elems)); + PipeBarrier(); + if (writeScoreScratch) { + ClampScoreExpInput(expFp32, static_cast(elems)); + } else { + ClampExpInput(expFp32, static_cast(elems)); + } + Exp(expFp32, expFp32, static_cast(elems)); + PipeBarrier(); + + Mul(outFp32, qFp32, expFp32, static_cast(elems)); + PipeBarrier(); + if (writeScoreScratch) { + ClampFp32ForCast(outFp32, static_cast(elems)); + Cast(qScore, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + } else { + ClampFp32ToOutputType(outFp32, static_cast(elems)); + Cast(qTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + } + PipeBarrier(); + + Mul(outFp32, kFp32, expFp32, static_cast(elems)); + PipeBarrier(); + if (writeScoreScratch) { + ClampFp32ForCast(outFp32, static_cast(elems)); + Cast(wScore, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + } else { + ClampFp32ToOutputType(outFp32, static_cast(elems)); + Cast(kTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + } + PipeBarrier(); + + if (useRef) { + for (uint64_t row = 0; row < tileRows; ++row) { + Sub(expFp32[row * K_], refFp32, gFp32[row * K_], static_cast(K_)); + } + } else { + Muls(expFp32, gFp32, -1.0f, static_cast(elems)); + } + PipeBarrier(); + Muls(expFp32, expFp32, LN2, static_cast(elems)); + PipeBarrier(); + if (writeScoreScratch) { + ClampScoreExpInput(expFp32, static_cast(elems)); + } else { + ClampExpInput(expFp32, static_cast(elems)); + } + Exp(expFp32, expFp32, static_cast(elems)); + PipeBarrier(); + Mul(outFp32, kFp32, expFp32, static_cast(elems)); + PipeBarrier(); + if (useRef && tileRow + tileRows > validColEnd) { + for (uint64_t row = 0; row < tileRows; ++row) { + if (tileRow + row >= validColEnd) { + Duplicate(outFp32[row * K_], 0.0f, static_cast(K_)); + } + } + PipeBarrier(); + } + if (writeScoreScratch) { + ClampFp32ForCast(outFp32, static_cast(elems)); + } else { + ClampFp32ToOutputType(outFp32, static_cast(elems)); + } + if (outputPending) { + WaitGateOutputForVector(); + } + if (writeScoreScratch) { + Cast(kgScore, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + } else { + Cast(kgTyped, outFp32, RoundMode::CAST_RINT, static_cast(elems)); + } + PipeBarrier(); + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + if (writeScoreScratch) { + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG, tileRow), + qScore, elems); + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W, tileRow), + wScore, elems); + CopyVectorOut(scoreWorkspace_, ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG, tileRow), + kgScore, elems); + } else { + CopyVectorOut(qg_, KVOffset(b, hv, token, 0, K_), qTyped, elems); + CopyVectorOut(w_, KVOffset(b, hv, token, 0, K_), kTyped, elems); + CopyVectorOut(kg_, KVOffset(b, hv, token, 0, K_), kgTyped, elems); + } + SignalGateOutputDone(); + outputPending = true; + gateSlot ^= 1; + } + if (outputPending) { + WaitGateOutputForMte2(); + WaitGateOutputForVector(); + } + return; + } + + __aicore__ inline void PrepareGateProducts(uint64_t b, uint64_t h, uint64_t hv, uint64_t start, uint64_t curT, + uint64_t subBlockIdx, uint64_t subBlockNum, bool useRef = false, + uint64_t refToken = 0, uint64_t validColEnd = 0, + bool writeScoreScratch = false, uint64_t scoreSlot = 0, + uint64_t scoreRowBegin = 0, uint64_t scoreRowCount = 0) + { + if (subBlockNum == 0 || subBlockIdx >= subBlockNum) { + return; + } + if (validColEnd == 0 || validColEnd > curT) { + validColEnd = curT; + } + if (writeScoreScratch) { + PrepareScoreFactorsBulk(b, h, hv, start, subBlockIdx, subBlockNum, refToken, scoreRowBegin, + scoreRowCount, validColEnd, scoreSlot); + return; + } + PrepareGateProductsBulk(b, h, hv, start, curT, subBlockIdx, subBlockNum, useRef, refToken, + validColEnd, writeScoreScratch, scoreSlot); + } + + __aicore__ inline void ComputeRawAqkAkkCube(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT) + { + ComputeRawAqkAkkCubeBlock(b, hv, start, curT, 0, curT); + } + + __aicore__ inline void ComputeRawAqkAkkCubeBlock(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT, + uint64_t rowBegin, uint64_t rowCount, + bool readScoreScratch = false, uint64_t scoreSlot = 0, + uint64_t colCount = 0) + { + using ElementA = SCORE_T; + using ElementB = SCORE_T; + using ElementC = float; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::ColumnMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using BlockMmad = Catlass::Gemm::Block::BlockMmadTla; + + Catlass::Arch::Resource resource; + BlockMmad blockMmad(resource); + auto layoutA = tla::MakeLayout(BT_, K_); + auto layoutB = tla::MakeLayout(K_, BT_); + auto layoutC = tla::MakeLayout(BT_, BT_); + if (colCount == 0 || colCount > curT) { + colCount = curT; + } + Catlass::GemmCoord shape{static_cast(rowCount), static_cast(colCount), + static_cast(K_)}; + + (void)readScoreScratch; + auto tensorQPos = + tla::MakeTensor(scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_QG)], + layoutA, Catlass::Arch::PositionGM{}); + auto tensorKPos = + tla::MakeTensor(scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_W)], + layoutA, Catlass::Arch::PositionGM{}); + auto tensorKNeg = + tla::MakeTensor(scoreWorkspace_[ScoreScratchOffset(scoreSlot, KDA_SCORE_SCRATCH_KG)], + layoutB, Catlass::Arch::PositionGM{}); + auto tensorAqk = tla::MakeTensor(aqk_[AOffset(b, hv, start, 0)], layoutC, + Catlass::Arch::PositionGM{}); + auto tensorAkk = tla::MakeTensor(akk_[AOffset(b, hv, start, 0)], layoutC, + Catlass::Arch::PositionGM{}); + + auto blockQPos = GetTile(tensorQPos, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.k())); + auto blockKPos = GetTile(tensorKPos, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.k())); + auto blockKNeg = GetTile(tensorKNeg, tla::MakeCoord(0, 0), tla::MakeShape(shape.k(), shape.n())); + auto blockAqk = GetTile(tensorAqk, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.n())); + auto blockAkk = GetTile(tensorAkk, tla::MakeCoord(rowBegin, 0), tla::MakeShape(shape.m(), shape.n())); + + blockMmad.preSetFlags(); + blockMmad(blockQPos, blockKNeg, blockAqk, shape); + blockMmad(blockKPos, blockKNeg, blockAkk, shape); + blockMmad.finalWaitFlags(); + } + + __aicore__ inline bool UseAkkCubeSolve(uint64_t curT) const + { + return curT > 0 && curT <= BT_ && (BT_ == 64 || BT_ == 128) && K_ >= 16 && V_ >= 16 && + V_ <= 256 && K_ % 16 == 0 && V_ % 16 == 0; + } + + __aicore__ inline bool UsePostWuCube(uint64_t curT) const + { + return curT > 0 && curT <= BT_ && (BT_ == 64 || BT_ == 128) && K_ >= 16 && V_ >= 16 && + V_ <= 256 && K_ % 16 == 0 && V_ % 16 == 0; + } + + __aicore__ inline void CopyLocalFloat(LocalTensor dst, LocalTensor src, uint64_t count) + { + if (count == 0) { + return; + } + Adds(dst, src, 0.0f, static_cast(count)); + PipeBarrier(); + } + + __aicore__ inline void FillLocalFloat(LocalTensor dst, float value, uint64_t count) + { + if (count == 0) { + return; + } + Duplicate(dst, value, static_cast(count)); + PipeBarrier(); + } + + __aicore__ inline void ForwardSubDiag16(LocalTensor diag, LocalTensor row, + LocalTensor prod, LocalTensor rowBrcb, + LocalTensor reduced, uint64_t valid) + { + constexpr uint32_t brcbStride = 8; + constexpr uint32_t diagSize = KDA_SOLVE_DIAG_BT; + constexpr uint8_t rowBlk = diagSize * sizeof(float) / 32; + + for (uint64_t i = 2; i < valid; ++i) { + uint32_t rowOffset = static_cast(i * diagSize); + DataCopy(row, diag[rowOffset], diagSize); + PipeBarrier(); + + Brcb(rowBrcb, row, diagSize / brcbStride, {1, 8}); + PipeBarrier(); + for (uint32_t col = 0; col < diagSize; col += brcbStride) { + Mul(prod[col], diag[col], rowBrcb, brcbStride, static_cast(diagSize), + {1, 1, 0, rowBlk, rowBlk, 1}); + } + PipeBarrier(); + + uint32_t remain = diagSize; + while (remain > 1) { + uint32_t calcCount = (remain / 2) * diagSize; + remain = (remain + 1) / 2; + Add(prod, prod, prod[remain * diagSize], calcCount); + PipeBarrier(); + } + DataCopy(reduced, prod, diagSize); + PipeBarrier(); + Add(row, row, reduced, diagSize); + PipeBarrier(); + DataCopy(diag[rowOffset], row, diagSize); + PipeBarrier(); + } + + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + for (uint32_t i = 0; i < diagSize; ++i) { + uint32_t diagOffset = i * diagSize + i; + if (i < valid) { + diag.SetValue(diagOffset, diag.GetValue(diagOffset) + 1.0f); + } else { + diag.SetValue(diagOffset, 1.0f); + } + } + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + } + + __aicore__ inline void SolveDiagonalBlocksInRows(LocalTensor akkMat, LocalTensor xMat, + LocalTensor arena, uint64_t scratchBase, + uint64_t curT, uint64_t rowBegin, uint64_t rowCount) + { + constexpr uint32_t diagSize = KDA_SOLVE_DIAG_BT; + constexpr uint32_t diagElements = diagSize * diagSize; + constexpr uint32_t brcbElements = diagSize * 8; + + LocalTensor diag = arena[scratchBase]; + LocalTensor row = diag[diagElements]; + LocalTensor prod = row[diagSize]; + LocalTensor rowBrcb = prod[diagElements]; + LocalTensor reduced = rowBrcb[brcbElements]; + + uint64_t rowEnd = rowBegin + rowCount; + for (uint64_t blockBegin = 0; blockBegin < BT_; blockBegin += diagSize) { + if (blockBegin < rowBegin || blockBegin + diagSize > rowEnd) { + continue; + } + Duplicate(diag, 0.0f, diagElements); + PipeBarrier(); + + uint64_t localBlockRow = blockBegin - rowBegin; + uint64_t valid = blockBegin < curT ? curT - blockBegin : 0; + if (valid > diagSize) { + valid = diagSize; + } + for (uint32_t rowIdx = 0; rowIdx < diagSize; ++rowIdx) { + uint64_t srcOffset = (localBlockRow + rowIdx) * BT_ + blockBegin; + Muls(diag[rowIdx * diagSize], akkMat[srcOffset], -1.0f, diagSize); + } + PipeBarrier(); + + ForwardSubDiag16(diag, row, prod, rowBrcb, reduced, valid); + for (uint32_t rowIdx = 0; rowIdx < diagSize; ++rowIdx) { + uint64_t dstOffset = (localBlockRow + rowIdx) * BT_ + blockBegin; + Adds(xMat[dstOffset], diag[rowIdx * diagSize], 0.0f, diagSize); + } + PipeBarrier(); + } + } + + __aicore__ inline void BuildPrefixMask(LocalTensor dst, uint64_t prefix, uint64_t count) + { + if (prefix > count) { + prefix = count; + } + Duplicate(dst, 0.0f, static_cast(count)); + if (prefix > 0) { + Duplicate(dst, 1.0f, static_cast(prefix)); + } + PipeBarrier(); + } + + __aicore__ inline uint64_t BuildCausalMask(uint64_t threshold, uint64_t colBegin) const + { + if (threshold <= colBegin) { + return ~0ULL; + } + if (threshold >= colBegin + KDA_SOLVE_BT) { + return 0ULL; + } + return ~0ULL << (threshold - colBegin); + } + + __aicore__ inline void BuildCausalSelectMasks(LocalTensor aqkMask, LocalTensor akkMask, + uint64_t rowBegin, uint64_t rowCount, uint64_t colBegin) + { + __ubuf__ uint64_t *aqkMaskPtr = reinterpret_cast<__ubuf__ uint64_t *>(aqkMask.GetPhyAddr()); + __ubuf__ uint64_t *akkMaskPtr = reinterpret_cast<__ubuf__ uint64_t *>(akkMask.GetPhyAddr()); + for (uint32_t localRow = 0; localRow < rowCount; ++localRow) { + uint32_t row = static_cast(rowBegin + localRow); + aqkMaskPtr[localRow] = BuildCausalMask(static_cast(row) + 1, colBegin); + akkMaskPtr[localRow] = BuildCausalMask(static_cast(row), colBegin); + } + } + + __aicore__ inline void SelectCausalRows(LocalTensor aqkMat, LocalTensor akkMat, + uint64_t rowBegin, uint64_t rowCount) + { + LocalTensor aqkMask = vecBuf_.Get()[KDA_SELECT_AQK_MASK_BYTE_OFFSET]; + LocalTensor akkMask = vecBuf_.Get()[KDA_SELECT_AKK_MASK_BYTE_OFFSET]; + LocalTensor zeroLocal = vecBuf_.Get()[KDA_SELECT_ZERO_FLOAT_OFFSET]; + Duplicate(zeroLocal, 0.0f, 8); + PipeBarrier(); + + uint64_t colBlockCount = (BT_ + KDA_SOLVE_BT - 1) / KDA_SOLVE_BT; + for (uint64_t colBlock = 0; colBlock < colBlockCount; ++colBlock) { + uint64_t maskOffset = colBlock * KDA_SELECT_COL_MASK_BYTES; + uint64_t colBegin = colBlock * KDA_SOLVE_BT; + BuildCausalSelectMasks(aqkMask[maskOffset], akkMask[maskOffset], rowBegin, rowCount, colBegin); + } + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + + uint8_t rowStride = static_cast(BT_ * sizeof(float) / 32); + BinaryRepeatParams repeatParams = {1, 0, 1, rowStride, 0, rowStride}; + for (uint64_t colBlock = 0; colBlock < colBlockCount; ++colBlock) { + uint64_t maskOffset = colBlock * KDA_SELECT_COL_MASK_BYTES; + uint64_t colBegin = colBlock * KDA_SOLVE_BT; + Select(aqkMat[colBegin], aqkMask[maskOffset], zeroLocal, aqkMat[colBegin], + SELMODE::VSEL_TENSOR_TENSOR_MODE, KDA_SOLVE_BT, static_cast(rowCount), repeatParams); + Select(akkMat[colBegin], akkMask[maskOffset], zeroLocal, akkMat[colBegin], + SELMODE::VSEL_TENSOR_TENSOR_MODE, KDA_SOLVE_BT, static_cast(rowCount), repeatParams); + } + PipeBarrier(); + SetFlag(EXP2_EVENT_ID); + WaitFlag(EXP2_EVENT_ID); + } + + __aicore__ inline void PrepareAqkAkkSolveInput64(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start) + { + LocalTensor arena = vecBuf_.Get(); + LocalTensor aqkMat = arena; + LocalTensor akkMat = arena[KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor xMat = arena[2 * KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor betaLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor betaBrcb = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT]; + LocalTensor maskLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512]; + LocalTensor oneHotLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512 + KDA_SOLVE_BT]; + + LoadAsFloatRow(beta_, BetaOffset(b, hv, start), betaLocal, KDA_SOLVE_BT); + Brcb(betaBrcb, betaLocal, 8, {1, 8}); + PipeBarrier(); + + DataCopy(aqkMat, aqk_[AOffset(b, hv, start, 0)], KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(akkMat, akk_[AOffset(b, hv, start, 0)], KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + + for (uint64_t col = 0; col < KDA_SOLVE_BT; col += 8) { + Mul(akkMat[col], akkMat[col], betaBrcb, 8, KDA_SOLVE_BT, {1, 1, 1, 8, 8, 1}); + PipeBarrier(); + } + SelectCausalRows(aqkMat, akkMat, 0, KDA_SOLVE_BT); + + Muls(xMat, akkMat, -1.0f, KDA_SOLVE_MATRIX_ELEMENTS); + PipeBarrier(); + for (uint64_t row = 0; row < KDA_SOLVE_BT; ++row) { + BuildPrefixMask(maskLocal, row + 1, KDA_SOLVE_BT); + BuildPrefixMask(oneHotLocal, row, KDA_SOLVE_BT); + Sub(maskLocal, maskLocal, oneHotLocal, KDA_SOLVE_BT); + PipeBarrier(); + Add(xMat[row * KDA_SOLVE_BT], xMat[row * KDA_SOLVE_BT], maskLocal, KDA_SOLVE_BT); + PipeBarrier(); + } + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(aqk_[AOffset(b, hv, start, 0)], aqkMat, KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(akk_[AOffset(b, hv, start, 0)], akkMat, KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(h_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X)], xMat, + KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void PrepareAqkAkkSolveInputTail(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT) + { + uint64_t elemCount = curT * KDA_SOLVE_BT; + LocalTensor arena = vecBuf_.Get(); + LocalTensor aqkMat = arena; + LocalTensor akkMat = arena[KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor xMat = arena[2 * KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor betaLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS]; + LocalTensor betaBrcb = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT]; + LocalTensor maskLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512]; + LocalTensor oneHotLocal = arena[3 * KDA_SOLVE_MATRIX_ELEMENTS + KDA_SOLVE_BT + 512 + KDA_SOLVE_BT]; + + FillLocalFloat(betaLocal, 0.0f, KDA_SOLVE_BT); + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + LoadAsFloatRow(beta_, BetaOffset(b, hv, start), betaLocal, curT); + Brcb(betaBrcb, betaLocal, 8, {1, 8}); + PipeBarrier(); + + DataCopy(aqkMat, aqk_[AOffset(b, hv, start, 0)], static_cast(elemCount)); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + if (elemCount < KDA_SOLVE_MATRIX_ELEMENTS) { + FillLocalFloat(aqkMat[elemCount], 0.0f, KDA_SOLVE_MATRIX_ELEMENTS - elemCount); + } + DataCopy(akkMat, akk_[AOffset(b, hv, start, 0)], static_cast(elemCount)); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + if (elemCount < KDA_SOLVE_MATRIX_ELEMENTS) { + FillLocalFloat(akkMat[elemCount], 0.0f, KDA_SOLVE_MATRIX_ELEMENTS - elemCount); + } + + for (uint64_t col = 0; col < KDA_SOLVE_BT; col += 8) { + Mul(akkMat[col], akkMat[col], betaBrcb, 8, KDA_SOLVE_BT, {1, 1, 1, 8, 8, 1}); + PipeBarrier(); + } + SelectCausalRows(aqkMat, akkMat, 0, KDA_SOLVE_BT); + + Muls(xMat, akkMat, -1.0f, KDA_SOLVE_MATRIX_ELEMENTS); + PipeBarrier(); + for (uint64_t row = 0; row < KDA_SOLVE_BT; ++row) { + BuildPrefixMask(maskLocal, row + 1, KDA_SOLVE_BT); + BuildPrefixMask(oneHotLocal, row, KDA_SOLVE_BT); + Sub(maskLocal, maskLocal, oneHotLocal, KDA_SOLVE_BT); + PipeBarrier(); + Add(xMat[row * KDA_SOLVE_BT], xMat[row * KDA_SOLVE_BT], maskLocal, KDA_SOLVE_BT); + PipeBarrier(); + } + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(aqk_[AOffset(b, hv, start, 0)], aqkMat, static_cast(elemCount)); + DataCopy(h_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X)], xMat, + KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(h_[SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0)], akkMat, + KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void GetSolveRowRange(uint64_t curT, uint64_t subBlockIdx, uint64_t subBlockNum, + uint64_t &rowBegin, uint64_t &rowEnd) const + { + if (subBlockNum == 0 || subBlockIdx >= subBlockNum) { + rowBegin = 0; + rowEnd = 0; + return; + } + rowBegin = (curT * subBlockIdx) / subBlockNum; + rowEnd = (curT * (subBlockIdx + 1)) / subBlockNum; + } + + __aicore__ inline void PrepareAqkAkkSolveInputRows(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT, uint64_t rowBegin, + uint64_t rowEnd, bool storeLToAkk, bool storeLToScratch) + { + uint64_t rowCount = rowEnd - rowBegin; + if (rowCount == 0) { + return; + } + uint64_t validRowCount = rowBegin < curT ? curT - rowBegin : 0; + if (validRowCount > rowCount) { + validRowCount = rowCount; + } + uint64_t elemCount = rowCount * BT_; + uint64_t validElemCount = validRowCount * BT_; + LocalTensor arena = vecBuf_.Get(); + LocalTensor aqkMat = arena; + LocalTensor akkMat = arena[elemCount]; + LocalTensor xMat = arena[2 * elemCount]; + LocalTensor betaLocal = arena[3 * elemCount]; + LocalTensor betaBrcb = arena[3 * elemCount + BT_]; + LocalTensor maskLocal = arena[3 * elemCount + BT_ + 512]; + LocalTensor oneHotLocal = arena[3 * elemCount + BT_ + 512 + BT_]; + + uint64_t token = start + rowBegin; + + if (validRowCount < rowCount) { + FillLocalFloat(aqkMat, 0.0f, elemCount); + FillLocalFloat(akkMat, 0.0f, elemCount); + FillLocalFloat(betaLocal, 0.0f, rowCount); + } + SetFlag(vToMte2Event_); + WaitFlag(vToMte2Event_); + if (validRowCount > 0) { + LoadAsFloatRow(beta_, BetaOffset(b, hv, token), betaLocal, validRowCount); + DataCopy(aqkMat, aqk_[AOffset(b, hv, token, 0)], static_cast(validElemCount)); + DataCopy(akkMat, akk_[AOffset(b, hv, token, 0)], static_cast(validElemCount)); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + } + Brcb(betaBrcb, betaLocal, static_cast((rowCount + 7) / 8), {1, 8}); + PipeBarrier(); + uint8_t rowStride = static_cast(BT_ * sizeof(float) / 32); + for (uint64_t col = 0; col < BT_; col += 8) { + Mul(akkMat[col], akkMat[col], betaBrcb, 8, static_cast(rowCount), + {1, 1, 0, rowStride, rowStride, 1}); + } + PipeBarrier(); + if (validRowCount > 0) { + SelectCausalRows(aqkMat, akkMat, rowBegin, validRowCount); + } + + Muls(xMat, akkMat, -1.0f, static_cast(elemCount)); + PipeBarrier(); + if constexpr (SAFE_GATE) { + uint64_t scratchBase = 3 * elemCount + BT_ + 512 + 2 * BT_; + SolveDiagonalBlocksInRows(akkMat, xMat, arena, scratchBase, curT, rowBegin, rowCount); + } else { + for (uint64_t localRow = 0; localRow < rowCount; ++localRow) { + uint64_t row = rowBegin + localRow; + BuildPrefixMask(maskLocal, row + 1, BT_); + BuildPrefixMask(oneHotLocal, row, BT_); + Sub(maskLocal, maskLocal, oneHotLocal, static_cast(BT_)); + PipeBarrier(); + Add(xMat[localRow * BT_], xMat[localRow * BT_], maskLocal, static_cast(BT_)); + PipeBarrier(); + } + } + + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; + uint64_t lBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0) + rowBegin * BT_; + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + if (validRowCount > 0) { + DataCopy(aqk_[AOffset(b, hv, token, 0)], aqkMat, static_cast(validElemCount)); + if (storeLToAkk) { + DataCopy(akk_[AOffset(b, hv, token, 0)], akkMat, static_cast(validElemCount)); + } + } + DataCopy(solveWorkspace_[xBase], xMat, static_cast(elemCount)); + if (storeLToScratch) { + DataCopy(solveWorkspace_[lBase], akkMat, static_cast(elemCount)); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void CubeGemmSolveSub(GlobalTensor &tensorA, uint64_t baseA, uint64_t rowA, uint64_t colA, + GlobalTensor &tensorB, uint64_t baseB, uint64_t rowB, uint64_t colB, + GlobalTensor &tensorC, uint64_t baseC, uint64_t rowC, uint64_t colC, + uint32_t m, uint32_t n, uint32_t k) + { + using ElementA = float; + using ElementB = float; + using ElementC = float; + using LayoutTagA = Catlass::layout::RowMajor; + using LayoutTagB = Catlass::layout::RowMajor; + using LayoutTagC = Catlass::layout::RowMajor; + using TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla; + using BlockMmad = Catlass::Gemm::Block::BlockMmadTla; + Catlass::Arch::Resource resource; + auto layoutA = tla::MakeLayout(BT_, BT_); + auto layoutB = tla::MakeLayout(BT_, BT_); + auto layoutC = tla::MakeLayout(BT_, BT_); + auto tensorLayoutA = tla::MakeTensor(tensorA[baseA], layoutA, Catlass::Arch::PositionGM{}); + auto tensorLayoutB = tla::MakeTensor(tensorB[baseB], layoutB, Catlass::Arch::PositionGM{}); + auto tensorLayoutC = tla::MakeTensor(tensorC[baseC], layoutC, Catlass::Arch::PositionGM{}); + Catlass::GemmCoord shape{m, n, k}; + auto blockA = GetTile(tensorLayoutA, tla::MakeCoord(rowA, colA), tla::MakeShape(shape.m(), shape.k())); + auto blockB = GetTile(tensorLayoutB, tla::MakeCoord(rowB, colB), tla::MakeShape(shape.k(), shape.n())); + auto blockC = GetTile(tensorLayoutC, tla::MakeCoord(rowC, colC), tla::MakeShape(shape.m(), shape.n())); + BlockMmad blockMmad(resource); + blockMmad(blockA, blockB, blockC, shape); + PipeBarrier(); + } + + __aicore__ inline void AddSolveTmpToX(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + bool storeAkk) + { + LocalTensor arena = vecBuf_.Get(); + LocalTensor xLocal = arena; + LocalTensor tmpLocal = arena[KDA_SOLVE_MATRIX_ELEMENTS]; + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + DataCopy(xLocal, h_[xBase], KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(tmpLocal, h_[tmpBase], KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + + Add(xLocal, xLocal, tmpLocal, KDA_SOLVE_MATRIX_ELEMENTS); + PipeBarrier(); + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(h_[xBase], xLocal, KDA_SOLVE_MATRIX_ELEMENTS); + if (storeAkk) { + DataCopy(akk_[AOffset(b, hv, start, 0)], xLocal, KDA_SOLVE_MATRIX_ELEMENTS); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void AddSolveTmpToXTail(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT, bool storeAkk) + { + uint64_t elemCount = curT * KDA_SOLVE_BT; + LocalTensor arena = vecBuf_.Get(); + LocalTensor xLocal = arena; + LocalTensor tmpLocal = arena[KDA_SOLVE_MATRIX_ELEMENTS]; + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + DataCopy(xLocal, h_[xBase], KDA_SOLVE_MATRIX_ELEMENTS); + DataCopy(tmpLocal, h_[tmpBase], KDA_SOLVE_MATRIX_ELEMENTS); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + + Add(xLocal, xLocal, tmpLocal, KDA_SOLVE_MATRIX_ELEMENTS); + PipeBarrier(); + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(h_[xBase], xLocal, KDA_SOLVE_MATRIX_ELEMENTS); + if (storeAkk) { + DataCopy(akk_[AOffset(b, hv, start, 0)], xLocal, static_cast(elemCount)); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void AddSolveTmpToXRows(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT, uint64_t rowBegin, uint64_t rowEnd, bool storeAkk) + { + uint64_t rowCount = rowEnd - rowBegin; + if (rowCount == 0) { + return; + } + uint64_t validRowCount = rowBegin < curT ? curT - rowBegin : 0; + if (validRowCount > rowCount) { + validRowCount = rowCount; + } + uint64_t elemCount = rowCount * BT_; + uint64_t validElemCount = validRowCount * BT_; + LocalTensor arena = vecBuf_.Get(); + LocalTensor xLocal = arena; + LocalTensor tmpLocal = arena[elemCount]; + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP) + rowBegin * BT_; + uint64_t token = start + rowBegin; + + DataCopy(xLocal, solveWorkspace_[xBase], static_cast(elemCount)); + DataCopy(tmpLocal, solveWorkspace_[tmpBase], static_cast(elemCount)); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + + Add(xLocal, xLocal, tmpLocal, static_cast(elemCount)); + PipeBarrier(); + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(solveWorkspace_[xBase], xLocal, static_cast(elemCount)); + if (storeAkk && validRowCount > 0) { + DataCopy(akk_[AOffset(b, hv, token, 0)], xLocal, static_cast(validElemCount)); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void AddSolveTmpToXDiagRows(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t rowBegin, uint64_t rowEnd, bool storeAkk) + { + uint64_t rowCount = rowEnd - rowBegin; + if (rowCount == 0) { + return; + } + uint64_t elemCount = rowCount * BT_; + LocalTensor arena = vecBuf_.Get(); + LocalTensor xLocal = arena; + LocalTensor tmpLocal = arena[elemCount]; + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP) + rowBegin * BT_; + uint64_t token = start + rowBegin; + + DataCopy(xLocal, solveWorkspace_[xBase], static_cast(elemCount)); + DataCopy(tmpLocal, solveWorkspace_[tmpBase], static_cast(elemCount)); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + + for (uint64_t localRow = 0; localRow < rowCount; ++localRow) { + uint64_t row = rowBegin + localRow; + uint64_t col = (row / KDA_SOLVE_DIAG_BT) * KDA_SOLVE_DIAG_BT; + uint64_t offset = localRow * BT_ + col; + Add(xLocal[offset], xLocal[offset], tmpLocal[offset], KDA_SOLVE_DIAG_BT); + PipeBarrier(); + } + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(solveWorkspace_[xBase], xLocal, static_cast(elemCount)); + if (storeAkk) { + DataCopy(akk_[AOffset(b, hv, token, 0)], xLocal, static_cast(elemCount)); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void StoreSolveXRowsToAkk(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t curT, uint64_t rowBegin, uint64_t rowEnd) + { + uint64_t validRowCount = rowBegin < curT ? curT - rowBegin : 0; + uint64_t rowCount = rowEnd - rowBegin; + if (validRowCount > rowCount) { + validRowCount = rowCount; + } + if (validRowCount == 0) { + return; + } + uint64_t elemCount = validRowCount * BT_; + LocalTensor xLocal = vecBuf_.Get(); + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X) + rowBegin * BT_; + + DataCopy(xLocal, solveWorkspace_[xBase], static_cast(elemCount)); + SetFlag(mte2ToMte3Event_); + WaitFlag(mte2ToMte3Event_); + DataCopy(akk_[AOffset(b, hv, start + rowBegin, 0)], xLocal, static_cast(elemCount)); + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + } + + __aicore__ inline void ComputeAkkMergeCube(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start) + { + uint64_t aiBase = AOffset(b, hv, start, 0); + uint64_t negABase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + for (uint32_t mergeSize = 2 * KDA_SOLVE_DIAG_BT; mergeSize <= BT_; mergeSize *= 2) { + uint32_t half = mergeSize / 2; + for (uint32_t block = 0; block < BT_; block += mergeSize) { + uint32_t lower = block + half; + CubeGemmSolveSub(akk_, aiBase, lower, lower, solveWorkspace_, negABase, lower, block, + solveWorkspace_, tmpBase, 0, 0, half, half, half); + CubeGemmSolveSub(solveWorkspace_, tmpBase, 0, 0, akk_, aiBase, block, block, + akk_, aiBase, lower, block, half, half, half); + } + } + } + + __aicore__ inline void ComputeAkkMergeCubeWorkspace(uint64_t b, uint64_t hv, uint64_t chunkIdx) + { + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + for (uint32_t mergeSize = 2 * KDA_SOLVE_DIAG_BT; mergeSize <= BT_; mergeSize *= 2) { + uint32_t half = mergeSize / 2; + for (uint32_t block = 0; block < BT_; block += mergeSize) { + uint32_t lower = block + half; + CubeGemmSolveSub(solveWorkspace_, xBase, lower, lower, solveWorkspace_, xBase, lower, block, + solveWorkspace_, tmpBase, 0, 0, half, half, half); + CubeGemmSolveSub(solveWorkspace_, tmpBase, 0, 0, solveWorkspace_, xBase, block, block, + solveWorkspace_, xBase, lower, block, half, half, half); + } + } + } + + __aicore__ inline void ComputeAkkInverseMchFull(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start) + { + uint64_t aBase = AOffset(b, hv, start, 0); + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t yBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0); + uint64_t yNextBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y1); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + + uint32_t diagBlocks = static_cast(BT_ / KDA_SOLVE_DIAG_BT); + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(akk_, aBase, off, off, akk_, aBase, off, off, solveWorkspace_, yBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + for (uint32_t iter = 0; iter < KDA_SOLVE_DIAG_MCH_ITERS; ++iter) { + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(solveWorkspace_, xBase, off, off, solveWorkspace_, yBase, off, off, + solveWorkspace_, tmpBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(mchSyncDoneFlag_); + if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(solveWorkspace_, yBase, off, off, solveWorkspace_, yBase, off, off, + solveWorkspace_, yNextBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + } + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(mchSyncReadyFlag_); + if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { + uint64_t oldYBase = yBase; + yBase = yNextBase; + yNextBase = oldYBase; + } + } + ComputeAkkMergeCube(b, hv, chunkIdx, start); + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(mchSyncDoneFlag_); + } + + __aicore__ inline void ComputeAkkInverseMchTail(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t curT) + { + uint64_t xBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_X); + uint64_t lBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0); + uint64_t yBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y1); + uint64_t yNextBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_Y0); + uint64_t tmpBase = SolveScratchOffset(b, hv, chunkIdx, KDA_SOLVE_SCRATCH_TMP); + (void)start; + (void)curT; + + uint32_t diagBlocks = static_cast(BT_ / KDA_SOLVE_DIAG_BT); + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(solveWorkspace_, lBase, off, off, solveWorkspace_, lBase, off, off, + solveWorkspace_, yBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + for (uint32_t iter = 0; iter < KDA_SOLVE_DIAG_MCH_ITERS; ++iter) { + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(solveWorkspace_, xBase, off, off, solveWorkspace_, yBase, off, off, + solveWorkspace_, tmpBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(mchSyncDoneFlag_); + if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { + for (uint32_t block = 0; block < diagBlocks; ++block) { + uint32_t off = block * KDA_SOLVE_DIAG_BT; + CubeGemmSolveSub(solveWorkspace_, yBase, off, off, solveWorkspace_, yBase, off, off, + solveWorkspace_, yNextBase, off, off, + KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT, KDA_SOLVE_DIAG_BT); + } + } + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(mchSyncReadyFlag_); + if (iter + 1 < KDA_SOLVE_DIAG_MCH_ITERS) { + uint64_t oldYBase = yBase; + yBase = yNextBase; + yNextBase = oldYBase; + } + } + ComputeAkkMergeCubeWorkspace(b, hv, chunkIdx); + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(mchSyncDoneFlag_); + } + + + + __aicore__ inline void ScaleRowsByBeta(GlobalTensor &src, GlobalTensor &dst, uint64_t b, uint64_t hv, + uint64_t start, uint64_t rowBegin, uint64_t rowCount, uint64_t dim, + LocalTensor &betaLocal, LocalTensor &betaBrcb, + LocalTensor &matrixLocal, bool sourceSequenceMajor = false) + { + constexpr uint64_t vecElemsPerRepeat = 64; + constexpr uint64_t typedOffsetFloats = 20480; + constexpr uint64_t typedOffset = typedOffsetFloats * sizeof(float) / sizeof(T); + uint64_t elemCount = rowCount * dim; + uint64_t baseOffset = KVOffset(b, hv, start + rowBegin, 0, dim); + uint64_t sourceOffset = sourceSequenceMajor + ? VInputOffset(b, hv, start + rowBegin, 0) + : baseOffset; + uint64_t sourceStride = sourceSequenceMajor ? HV_ * dim : dim; + + if constexpr (IsSameType::value) { + CopyRowsIn(matrixLocal, src, sourceOffset, rowCount, dim, sourceStride); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + } else { + LocalTensor matrixTyped = vecBuf_.Get()[typedOffset]; + CopyRowsIn(matrixTyped, src, sourceOffset, rowCount, dim, sourceStride); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + Cast(matrixLocal, matrixTyped, RoundMode::CAST_NONE, static_cast(elemCount)); + PipeBarrier(); + } + + uint8_t repeatStride = static_cast(dim * sizeof(float) / 32); + for (uint64_t col = 0; col < dim; col += vecElemsPerRepeat) { + uint64_t mask = dim - col; + if (mask > vecElemsPerRepeat) { + mask = vecElemsPerRepeat; + } + Mul(matrixLocal[col], matrixLocal[col], betaBrcb, mask, static_cast(rowCount), + {1, 1, 0, repeatStride, repeatStride, 1}); + } + PipeBarrier(); + + if constexpr (IsSameType::value) { + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(dst[baseOffset], matrixLocal, static_cast(elemCount)); + } else { + LocalTensor matrixTyped = vecBuf_.Get()[typedOffset]; + Cast(matrixTyped, matrixLocal, RoundMode::CAST_RINT, static_cast(elemCount)); + PipeBarrier(); + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + DataCopy(dst[baseOffset], matrixTyped, static_cast(elemCount)); + } + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + + __aicore__ inline void PrepareWuCubeInputs(uint64_t b, uint64_t hv, uint64_t start, uint64_t curT, + uint64_t subBlockIdx, uint64_t subBlockNum) + { + uint64_t rowsPerSubBlock = (curT + subBlockNum - 1) / subBlockNum; + uint64_t rowBegin = subBlockIdx * rowsPerSubBlock; + if (rowBegin >= curT) { + return; + } + uint64_t rowCount = curT - rowBegin; + if (rowCount > rowsPerSubBlock) { + rowCount = rowsPerSubBlock; + } + LocalTensor arena = vecBuf_.Get(); + LocalTensor betaLocal = arena; + LocalTensor betaBrcb = arena[KDA_SOLVE_BT]; + LocalTensor matrixLocal = arena[KDA_SOLVE_BT + 512]; + LoadAsFloatRow(beta_, BetaOffset(b, hv, start + rowBegin), betaLocal, rowCount); + Brcb(betaBrcb, betaLocal, static_cast((rowCount + 7) / 8), {1, 8}); + PipeBarrier(); + ScaleRowsByBeta(w_, w_, b, hv, start, rowBegin, rowCount, K_, betaLocal, betaBrcb, matrixLocal); + ScaleRowsByBeta(v_, vNew_, b, hv, start, rowBegin, rowCount, V_, betaLocal, betaBrcb, + matrixLocal, inputSequenceMajor_); + } + + __aicore__ inline void FinalizePrepareIntermediates(uint64_t b, uint64_t hv, uint64_t start, + uint64_t curT, uint64_t subBlockIdx, + uint64_t subBlockNum) + { + constexpr uint64_t tileRows = 32; + // Keep tail rows on the same AIV that owns their padded solve rows. Splitting by curT would + // move short-tail export to AIV1 while AIV0 is still writing the solved matrix. + const uint64_t rowBegin = (BT_ * subBlockIdx) / subBlockNum; + uint64_t rowEnd = (BT_ * (subBlockIdx + 1)) / subBlockNum; + if (rowEnd > curT) { + rowEnd = curT; + } + if (rowBegin >= rowEnd) { + return; + } + for (uint64_t tileRow = rowBegin; tileRow < rowEnd; tileRow += tileRows) { + const uint64_t rows = (rowEnd - tileRow) > tileRows ? tileRows : (rowEnd - tileRow); + const uint64_t matrixElems = rows * BT_; + const uint64_t qgElems = rows * K_; + LocalTensor arena = vecBuf_.Get(); + LocalTensor aqkLocal = arena; + LocalTensor akkLocal = arena[matrixElems]; + LocalTensor qgLocal = arena[2 * matrixElems]; + const uint64_t typedOffset = + (2 * matrixElems + qgElems) * sizeof(float) / sizeof(T); + LocalTensor typedBase = vecBuf_.Get()[typedOffset]; + LocalTensor aqkTyped = typedBase; + LocalTensor akkTyped = typedBase[matrixElems]; + LocalTensor qgTyped = typedBase[2 * matrixElems]; + + CopyVectorIn(aqkLocal, aqk_, AOffset(b, hv, start + tileRow, 0), matrixElems); + CopyVectorIn(akkLocal, akk_, AOffset(b, hv, start + tileRow, 0), matrixElems); + CopyVectorIn(qgTyped, qg_, KVOffset(b, hv, start + tileRow, 0, K_), qgElems); + SetFlag(mte2ToVEvent_); + WaitFlag(mte2ToVEvent_); + + Muls(aqkLocal, aqkLocal, scale_, static_cast(matrixElems)); + Cast(qgLocal, qgTyped, RoundMode::CAST_NONE, static_cast(qgElems)); + PipeBarrier(); + Muls(qgLocal, qgLocal, scale_, static_cast(qgElems)); + PipeBarrier(); + ClampFp32ToOutputType(aqkLocal, static_cast(matrixElems)); + ClampFp32ToOutputType(akkLocal, static_cast(matrixElems)); + ClampFp32ToOutputType(qgLocal, static_cast(qgElems)); + Cast(aqkTyped, aqkLocal, RoundMode::CAST_RINT, static_cast(matrixElems)); + Cast(akkTyped, akkLocal, RoundMode::CAST_RINT, static_cast(matrixElems)); + Cast(qgTyped, qgLocal, RoundMode::CAST_RINT, static_cast(qgElems)); + PipeBarrier(); + + SetFlag(vToMte3Event_); + WaitFlag(vToMte3Event_); + CopyVectorOut(o_, AOffset(b, hv, start + tileRow, 0), aqkTyped, matrixElems); + CopyVectorOut(u_, AOffset(b, hv, start + tileRow, 0), akkTyped, matrixElems); + CopyVectorOut(kg_, KVOffset(b, hv, start + tileRow, 0, K_), qgTyped, qgElems); + SetFlag(mte3ToMte2Events_[0]); + WaitFlag(mte3ToMte2Events_[0]); + SetFlag(mte3ToVEvent_); + WaitFlag(mte3ToVEvent_); + } + } + + __aicore__ inline bool ResolveFlatChunk(uint64_t task, uint64_t &seq, uint64_t &b, uint64_t &h, uint64_t &hv, + uint64_t &chunkIdx, uint64_t &start, uint64_t &end) + { + hv = task % HV_; + uint64_t flatChunk = task / HV_; + if (!isVarLen_) { + seq = flatChunk / NT_; + b = seq; + chunkIdx = flatChunk % NT_; + start = chunkIdx * BT_; + end = start + BT_; + if (end > T_) { + end = T_; + } + } else { + if (!KdaVarlen::ResolveChunkRange( + cuSeqlensAddr_, chunkIndicesAddr_, N_, T_, BT_, flatChunk, + seq, start, end)) { + return false; + } + b = 0; + chunkIdx = flatChunk; + } + h = hv / (HV_ / H_); + return start < end; + } + + __aicore__ inline void ProcessChunkPreAiv(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t end, uint64_t subBlockIdx, + uint64_t subBlockNum) + { + if constexpr (IsSameType::value) { + ProcessChunkPreAivFp32(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); + } + } + + template + __aicore__ inline void JoinAivMte3() + { + if constexpr (CORE_TYPE == AscendC::AIV) { + } + } + + template + __aicore__ inline void RunAicAfterBothAivReady(uint64_t subBlockIdx, uint64_t subBlockNum) + { + if constexpr (CORE_TYPE == AscendC::AIV) { + (void)subBlockIdx; + (void)subBlockNum; + JoinAivMte3(); + if constexpr (SAFE_GATE) { + Catlass::Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(syncReadyFlag_); + Catlass::Arch::CrossCoreWaitFlag(syncDoneFlag_); + } else { + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_MTE3>(mchSyncReadyFlag_); + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(mchSyncDoneFlag_); + } + } + } + + template + __aicore__ inline void SignalAicSolveReady() + { + if constexpr (CORE_TYPE == AscendC::AIV) { + JoinAivMte3(); + Catlass::Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(syncReadyFlag_); + } + } + + template + __aicore__ inline void WaitAicSolveDone() + { + if constexpr (CORE_TYPE == AscendC::AIV) { + Catlass::Arch::CrossCoreWaitFlag(syncDoneFlag_); + } + } + + __aicore__ inline void ProcessChunkPreAivFp32(uint64_t b, uint64_t h, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t end, uint64_t subBlockIdx, + uint64_t subBlockNum, bool deferSafeSolve = false, + bool waitPendingSafeSolve = false, + uint64_t scoreLane = 0, bool pairHeads = false) + { + uint64_t curT = end - start; + if (curT == 0) { + return; + } + if constexpr (IsSameType::value) { + return; + } + + if (K_ < 16) { + return; + } + bool usePostWuCube = UsePostWuCube(curT); + bool useAkkCubeSolve = UseAkkCubeSolve(curT); + uint64_t solveRowBegin = 0; + uint64_t solveRowEnd = 0; + GetSolveRowRange(BT_, subBlockIdx, subBlockNum, solveRowBegin, solveRowEnd); + uint64_t scoreBlockSize = ScoreRefBlockSize(); + uint64_t scoreBlockCount = (curT + scoreBlockSize - 1) / scoreBlockSize; + uint64_t pipelineBlockCount = + (scoreBlockCount + KDA_SCORE_QUEUE_DEPTH - 1) / KDA_SCORE_QUEUE_DEPTH * KDA_SCORE_QUEUE_DEPTH; + for (uint64_t block = 0; block < pipelineBlockCount; ++block) { + if (block < scoreBlockCount) { + uint64_t rowBegin = block * scoreBlockSize; + uint64_t rowCount = ScoreRowBlockCount(curT, rowBegin); + uint64_t refToken = ScoreRefToken(start, curT, rowBegin, rowCount); + uint64_t scoreSlot = + ScoreScratchSlot(block % KDA_SCORE_QUEUE_DEPTH, scoreLane, pairHeads); + PrepareGateProducts(b, h, hv, start, curT, subBlockIdx, subBlockNum, true, refToken, + rowBegin + rowCount, true, scoreSlot, + rowBegin, rowCount); + } + JoinAivMte3(); + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_MTE3>(scoreReadyFlag_); + if (block > 0) { + if constexpr (SAFE_GATE) { + if (waitPendingSafeSolve) { + WaitAicSolveDone(); + waitPendingSafeSolve = false; + } + } + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(scoreDoneFlag_); + } + } + bool fusedScoreWriteback = false; + if (!fusedScoreWriteback) { + // The final score MMAD only consumes scoreWorkspace_. Run the + // independent gate writeback while AIC drains its MMAD/Fixpipe path. + PrepareGateProducts(b, h, hv, start, curT, subBlockIdx, subBlockNum); + } + if (pipelineBlockCount > 0) { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_MTE2>(scoreDoneFlag_); + } + if constexpr (SAFE_GATE) { + if (waitPendingSafeSolve) { + WaitAicSolveDone(); + waitPendingSafeSolve = false; + } + } + if (useAkkCubeSolve) { + bool fullChunk = curT == BT_; + if constexpr (SAFE_GATE) { + if (pairHeads) { + for (uint64_t rowPart = 0; rowPart < KDA_SCORE_LANES; ++rowPart) { + uint64_t pairRowBegin = 0; + uint64_t pairRowEnd = 0; + GetSolveRowRange( + BT_, rowPart, KDA_SCORE_LANES, pairRowBegin, pairRowEnd); + PrepareAqkAkkSolveInputRows( + b, hv, chunkIdx, start, curT, pairRowBegin, pairRowEnd, false, false); + } + } else { + PrepareAqkAkkSolveInputRows( + b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd, false, false); + } + if (deferSafeSolve) { + SignalAicSolveReady(); + return; + } + RunAicAfterBothAivReady(subBlockIdx, subBlockNum); + StoreSolveXRowsToAkk(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd); + } else { + PrepareAqkAkkSolveInputRows(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd, + fullChunk, !fullChunk); + uint32_t solveIters = KDA_SOLVE_DIAG_MCH_ITERS; + RunAicAfterBothAivReady(subBlockIdx, subBlockNum); + for (uint32_t iter = 0; iter < solveIters; ++iter) { + AddSolveTmpToXDiagRows(b, hv, chunkIdx, start, solveRowBegin, solveRowEnd, + fullChunk && iter + 1 == solveIters); + RunAicAfterBothAivReady(subBlockIdx, subBlockNum); + } + if (!fullChunk) { + StoreSolveXRowsToAkk(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd); + } + } + } + // Host validation guarantees every accepted shape has enough workspace for this cube path. + PrepareWuCubeInputs(b, hv, start, curT, subBlockIdx, subBlockNum); + FinalizePrepareIntermediates(b, hv, start, curT, subBlockIdx, subBlockNum); + } + + __aicore__ inline void FinishDeferredSafeChunk(uint64_t b, uint64_t hv, uint64_t chunkIdx, + uint64_t start, uint64_t end, uint64_t subBlockIdx, + uint64_t subBlockNum) + { + uint64_t curT = end - start; + uint64_t solveRowBegin = 0; + uint64_t solveRowEnd = 0; + GetSolveRowRange(BT_, subBlockIdx, subBlockNum, solveRowBegin, solveRowEnd); + StoreSolveXRowsToAkk(b, hv, chunkIdx, start, curT, solveRowBegin, solveRowEnd); + PrepareWuCubeInputs(b, hv, start, curT, subBlockIdx, subBlockNum); + FinalizePrepareIntermediates(b, hv, start, curT, subBlockIdx, subBlockNum); + } + + __aicore__ inline void FinishDeferredSafeChunkPair( + uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, uint64_t end) + { + FinishDeferredSafeChunk(b, hv, chunkIdx, start, end, 0, 1); + } + + __aicore__ inline void ProcessChunkPreAic(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + if constexpr (IsSameType::value) { + ProcessChunkPreAicFp32(b, hv, chunkIdx, start, end); + } + } + + __aicore__ inline void ProcessChunkPreAicFp32(uint64_t b, uint64_t hv, uint64_t chunkIdx, uint64_t start, + uint64_t end) + { + uint64_t curT = end - start; + if (curT == 0 || K_ < 16) { + return; + } + uint64_t scoreBlockSize = ScoreRefBlockSize(); + uint64_t scoreBlockCount = (curT + scoreBlockSize - 1) / scoreBlockSize; + uint64_t pipelineBlockCount = + (scoreBlockCount + KDA_SCORE_QUEUE_DEPTH - 1) / KDA_SCORE_QUEUE_DEPTH * KDA_SCORE_QUEUE_DEPTH; + for (uint64_t block = 0; block < pipelineBlockCount; ++block) { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(scoreReadyFlag_); + if (block < scoreBlockCount) { + uint64_t rowBegin = block * scoreBlockSize; + uint64_t rowCount = ScoreRowBlockCount(curT, rowBegin); + ComputeRawAqkAkkCubeBlock(b, hv, start, curT, rowBegin, rowCount, true, + block % KDA_SCORE_QUEUE_DEPTH, rowBegin + rowCount); + } + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(scoreDoneFlag_); + } + bool usePostWuCube = UsePostWuCube(curT); + bool useAkkCubeSolve = UseAkkCubeSolve(curT); + if (useAkkCubeSolve) { + if constexpr (SAFE_GATE) { + Catlass::Arch::CrossCoreWaitFlag(syncReadyFlag_); + ComputeAkkMergeCubeWorkspace(b, hv, chunkIdx); + Catlass::Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(syncDoneFlag_); + } else { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(mchSyncReadyFlag_); + if (curT == BT_) { + ComputeAkkInverseMchFull(b, hv, chunkIdx, start); + } else { + ComputeAkkInverseMchTail(b, hv, chunkIdx, start, curT); + } + } + } + (void)usePostWuCube; + (void)chunkIdx; + } + + __aicore__ inline void ProcessChunkPreAicHeadPairFp32( + uint64_t b, uint64_t hvBase, uint64_t chunkIdx, uint64_t start, uint64_t end, + uint64_t localTaskIdx) + { + uint64_t curT = end - start; + if (curT == 0 || K_ < 16) { + return; + } + uint64_t scoreBlockSize = ScoreRefBlockSize(); + uint64_t scoreBlockCount = (curT + scoreBlockSize - 1) / scoreBlockSize; + uint64_t pipelineBlockCount = + (scoreBlockCount + KDA_SCORE_QUEUE_DEPTH - 1) / KDA_SCORE_QUEUE_DEPTH * KDA_SCORE_QUEUE_DEPTH; + for (uint64_t block = 0; block < pipelineBlockCount; ++block) { + Catlass::Arch::CrossCoreWaitFlagWithReverse<0x2, PIPE_FIX>(scoreReadyFlag_); + if (block < scoreBlockCount) { + uint64_t rowBegin = block * scoreBlockSize; + uint64_t rowCount = ScoreRowBlockCount(curT, rowBegin); + for (uint64_t lane = 0; lane < KDA_SCORE_LANES; ++lane) { + uint64_t hv = hvBase + lane; + uint64_t scoreSlot = + ScoreScratchSlot(block % KDA_SCORE_QUEUE_DEPTH, lane, true); + ComputeRawAqkAkkCubeBlock(b, hv, start, curT, rowBegin, rowCount, true, + scoreSlot, rowBegin + rowCount); + } + } + Catlass::Arch::CrossCoreSetFlagWithReverse<0x2, PIPE_FIX>(scoreDoneFlag_); + } + + if (UseAkkCubeSolve(curT)) { + Catlass::Arch::CrossCoreWaitFlag(syncReadyFlag_); + for (uint64_t lane = 0; lane < KDA_SCORE_LANES; ++lane) { + activeSolveSlot_ = + (localTaskIdx % (KDA_SOLVE_PIPELINE_DEPTH / KDA_SCORE_LANES)) * KDA_SCORE_LANES + lane; + ComputeAkkMergeCubeWorkspace(b, hvBase + lane, chunkIdx); + } + Catlass::Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(syncDoneFlag_); + } + } + + __aicore__ inline bool ResolveFlatChunkForHv( + uint64_t flatChunk, uint64_t hv, uint64_t &seq, uint64_t &b, uint64_t &h, + uint64_t &chunkIdx, uint64_t &start, uint64_t &end) + { + if (!isVarLen_) { + seq = flatChunk / NT_; + b = seq; + chunkIdx = flatChunk % NT_; + start = chunkIdx * BT_; + end = start + BT_; + if (end > T_) { + end = T_; + } + } else { + if (!KdaVarlen::ResolveChunkRange( + cuSeqlensAddr_, chunkIndicesAddr_, N_, T_, BT_, flatChunk, + seq, start, end)) { + return false; + } + b = 0; + chunkIdx = flatChunk; + } + h = hv / (HV_ / H_); + return start < end; + } + + __aicore__ inline void ProcessPreAivHeadPair() + { + const uint64_t subBlockIdx = static_cast(GetSubBlockIdx()); + const uint64_t subBlockNum = static_cast(GetSubBlockNum()); + const uint64_t coreNum = usedCoreNum_; + const uint64_t coreIdx = static_cast(GetBlockIdx()) / subBlockNum; + const uint64_t chunkCount = isVarLen_ ? NT_ : B_ * NT_; + const uint64_t headWindows = HV_ / KDA_SCORE_LANES; + const uint64_t taskNum = chunkCount * headWindows; + bool pendingValid = false; + uint64_t pendingB = 0; + uint64_t pendingHv = 0; + uint64_t pendingChunkIdx = 0; + uint64_t pendingStart = 0; + uint64_t pendingEnd = 0; + uint64_t pendingSlot = 0; + uint64_t localTaskIdx = 0; + + for (uint64_t task = coreIdx; task < taskNum; task += coreNum, ++localTaskIdx) { + uint64_t flatChunk = task / headWindows; + uint64_t hv = (task % headWindows) * KDA_SCORE_LANES + subBlockIdx; + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (!ResolveFlatChunkForHv(flatChunk, hv, seq, b, h, chunkIdx, start, end)) { + continue; + } + (void)seq; + uint64_t currentSlot = + (localTaskIdx % (KDA_SOLVE_PIPELINE_DEPTH / KDA_SCORE_LANES)) * + KDA_SCORE_LANES + + subBlockIdx; + activeSolveSlot_ = currentSlot; + bool deferSolve = UseAkkCubeSolve(end - start); + ProcessChunkPreAivFp32( + b, h, hv, chunkIdx, start, end, 0, 1, deferSolve, pendingValid, subBlockIdx, true); + if (pendingValid) { + activeSolveSlot_ = pendingSlot; + FinishDeferredSafeChunkPair( + pendingB, pendingHv, pendingChunkIdx, pendingStart, pendingEnd); + } + pendingValid = deferSolve; + if (pendingValid) { + pendingB = b; + pendingHv = hv; + pendingChunkIdx = chunkIdx; + pendingStart = start; + pendingEnd = end; + pendingSlot = currentSlot; + } + } + if (pendingValid) { + WaitAicSolveDone(); + activeSolveSlot_ = pendingSlot; + FinishDeferredSafeChunkPair( + pendingB, pendingHv, pendingChunkIdx, pendingStart, pendingEnd); + } + } + + __aicore__ inline void ProcessPreAicHeadPair() + { + const uint64_t chunkCount = isVarLen_ ? NT_ : B_ * NT_; + const uint64_t headWindows = HV_ / KDA_SCORE_LANES; + const uint64_t taskNum = chunkCount * headWindows; + const uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + uint64_t localTaskIdx = 0; + for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum, ++localTaskIdx) { + uint64_t flatChunk = task / headWindows; + uint64_t hvBase = (task % headWindows) * KDA_SCORE_LANES; + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunkForHv(flatChunk, hvBase, seq, b, h, chunkIdx, start, end)) { + (void)seq; + (void)h; + ProcessChunkPreAicHeadPairFp32(b, hvBase, chunkIdx, start, end, localTaskIdx); + } + } + } + + __aicore__ inline void ProcessPreAiv() + { + if constexpr (IsSameType::value) { + isAivOnly_ = true; + } + uint64_t subBlockNum = isAivOnly_ ? 1 : static_cast(GetSubBlockNum()); + if (subBlockNum == 0) { + return; + } + uint64_t subBlockIdx = isAivOnly_ ? 0 : static_cast(GetSubBlockIdx()); + uint64_t coreNum = isAivOnly_ ? static_cast(GetBlockNum()) : usedCoreNum_; + uint64_t coreIdx = isAivOnly_ ? static_cast(GetBlockIdx()) : + static_cast(GetBlockIdx()) / subBlockNum; + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + if constexpr (SAFE_GATE && !IsSameType::value) { + bool pendingValid = false; + uint64_t pendingB = 0; + uint64_t pendingHv = 0; + uint64_t pendingChunkIdx = 0; + uint64_t pendingStart = 0; + uint64_t pendingEnd = 0; + uint64_t pendingSlot = 0; + uint64_t localTaskIdx = 0; + for (uint64_t task = coreIdx; task < taskNum; task += coreNum, ++localTaskIdx) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (!ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + continue; + } + (void)seq; + uint64_t currentSlot = localTaskIdx % KDA_SOLVE_PIPELINE_DEPTH; + activeSolveSlot_ = currentSlot; + bool deferSolve = UseAkkCubeSolve(end - start); + if (!deferSolve && pendingValid) { + WaitAicSolveDone(); + activeSolveSlot_ = pendingSlot; + FinishDeferredSafeChunk(pendingB, pendingHv, pendingChunkIdx, pendingStart, pendingEnd, + subBlockIdx, subBlockNum); + pendingValid = false; + activeSolveSlot_ = currentSlot; + } + ProcessChunkPreAivFp32(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum, + deferSolve, pendingValid); + if (pendingValid) { + activeSolveSlot_ = pendingSlot; + FinishDeferredSafeChunk(pendingB, pendingHv, pendingChunkIdx, pendingStart, pendingEnd, + subBlockIdx, subBlockNum); + } + pendingValid = deferSolve; + if (pendingValid) { + pendingB = b; + pendingHv = hv; + pendingChunkIdx = chunkIdx; + pendingStart = start; + pendingEnd = end; + pendingSlot = currentSlot; + } + } + if (pendingValid) { + WaitAicSolveDone(); + activeSolveSlot_ = pendingSlot; + FinishDeferredSafeChunk(pendingB, pendingHv, pendingChunkIdx, pendingStart, pendingEnd, + subBlockIdx, subBlockNum); + } + return; + } + for (uint64_t task = coreIdx; task < taskNum; task += coreNum) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + (void)seq; + ProcessChunkPreAiv(b, h, hv, chunkIdx, start, end, subBlockIdx, subBlockNum); + } + } + } + + __aicore__ inline void ProcessPreAic() + { + if constexpr (IsSameType::value) { + return; + } + uint64_t taskNum = static_cast((isVarLen_ ? NT_ : B_ * NT_) * HV_); + uint64_t coreNum = usedCoreNum_ == 0 ? 1 : usedCoreNum_; + uint64_t localTaskIdx = 0; + for (uint64_t task = GetBlockIdx(); task < taskNum; task += coreNum, ++localTaskIdx) { + uint64_t seq = 0; + uint64_t b = 0; + uint64_t h = 0; + uint64_t hv = 0; + uint64_t chunkIdx = 0; + uint64_t start = 0; + uint64_t end = 0; + if (ResolveFlatChunk(task, seq, b, h, hv, chunkIdx, start, end)) { + if constexpr (SAFE_GATE) { + activeSolveSlot_ = localTaskIdx % KDA_SOLVE_PIPELINE_DEPTH; + } + (void)seq; + (void)h; + ProcessChunkPreAic(b, hv, chunkIdx, start, end); + } + } + } + + +private: + GlobalTensor q_; + GlobalTensor k_; + GlobalTensor v_; + GlobalTensor gk_; + GlobalTensor beta_; + GlobalTensor initialState_; + GlobalTensor o_; + GlobalTensor finalState_; + GlobalTensor aqk_; + GlobalTensor akk_; + GlobalTensor w_; + GlobalTensor u_; + GlobalTensor qg_; + GlobalTensor kg_; + GlobalTensor vNew_; + GlobalTensor h_; + GlobalTensor preparedQG_; + GlobalTensor preparedAqk_; + GlobalTensor propagatedVNew_; + GlobalTensor propagatedH_; + GlobalTensor solveWorkspace_; + GlobalTensor scoreWorkspace_; + TPipe *pipe_ = nullptr; + TBuf exp2Buf_; + TBuf vecBuf_; + TBuf gateWritebackBuf_; + TEventID mte2ToVEvent_ = 0; + TEventID vToMte2Event_ = 0; + TEventID vToMte3Event_ = 0; + TEventID mte3ToVEvent_ = 0; + TEventID mte2ToMte3Event_ = 0; + TEventID mte3ToMte2Events_[KDA_GATE_PIPELINE_DEPTH] = {0, 0, 0}; + bool vectorEventsAllocated_ = false; + Catlass::Arch::CrossCoreFlagWithReverse scoreReadyFlag_{KDA_SCORE_READY_FLAG0, + KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse scoreDoneFlag_{KDA_SCORE_DONE_FLAG0, + KDA_SCORE_DONE_FLAG1}; + // Solve has one outstanding task per core. Reuse the primary score IDs as an ordered token stream; + // score credits remain on the reverse IDs, so no additional hardware flag IDs are consumed. + Catlass::Arch::CrossCoreFlag syncReadyFlag_{KDA_SOLVE_READY_FLAG}; + Catlass::Arch::CrossCoreFlag syncDoneFlag_{KDA_SOLVE_DONE_FLAG}; + Catlass::Arch::CrossCoreFlagWithReverse mchSyncReadyFlag_{ + KDA_SCORE_READY_FLAG0, KDA_SCORE_READY_FLAG1}; + Catlass::Arch::CrossCoreFlagWithReverse mchSyncDoneFlag_{ + KDA_SCORE_DONE_FLAG0, KDA_SCORE_DONE_FLAG1}; + uint64_t B_ = 0; + uint64_t N_ = 0; + uint64_t H_ = 0; + uint64_t HV_ = 0; + uint64_t T_ = 0; + uint64_t K_ = 0; + uint64_t V_ = 0; + uint64_t BT_ = 0; + uint64_t NT_ = 0; + float scale_ = 1.0f; + bool hasInitial_ = false; + bool isVarLen_ = false; + bool inputSequenceMajor_ = false; + bool isAivOnly_ = false; + uint64_t usedCoreNum_ = 1; + uint64_t solveCoreIdx_ = 0; + uint64_t activeSolveSlot_ = 0; + __gm__ int64_t *chunkIndicesAddr_ = nullptr; + __gm__ int64_t *cuSeqlensAddr_ = nullptr; +}; +} // namespace + +template +__aicore__ inline void RunChunkKdaPrepare( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR beta, GM_ADDR initialState, + GM_ADDR cuSeqlens, GM_ADDR chunkIndices, GM_ADDR aqk, GM_ADDR akk, GM_ADDR qg, + GM_ADDR qgScaled, GM_ADDR wSeed, GM_ADDR uSeed, GM_ADDR userWorkspace, + const TilingData &tiling, TPipe &pipe) +{ + GM_ADDR aqkFp32 = userWorkspace + tiling.prepareAqkFp32Offset; + GM_ADDR akkFp32 = userWorkspace + tiling.prepareAkkFp32Offset; + GM_ADDR prepareScratch = userWorkspace + tiling.prepareScratchOffset; + + if ASCEND_IS_AIC { + ChunkKdaFwdPrepareKernel op; + op.Init(q, k, v, gk, beta, initialState, cuSeqlens, chunkIndices, + nullptr, nullptr, nullptr, nullptr, aqk, userWorkspace, aqkFp32, akkFp32, + wSeed, akk, qg, qgScaled, uSeed, userWorkspace, prepareScratch, tiling, &pipe, false); + op.ProcessAic(); + } + if ASCEND_IS_AIV { + ChunkKdaFwdPrepareKernel op; + op.Init(q, k, v, gk, beta, initialState, cuSeqlens, chunkIndices, + nullptr, nullptr, nullptr, nullptr, aqk, userWorkspace, aqkFp32, akkFp32, + wSeed, akk, qg, qgScaled, uSeed, userWorkspace, prepareScratch, tiling, &pipe); + op.ProcessAiv(); + } +} + +template +__aicore__ inline void RunChunkKdaPrepare( + GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR gk, GM_ADDR, GM_ADDR, + GM_ADDR, GM_ADDR beta, GM_ADDR initialState, GM_ADDR cuSeqlens, + GM_ADDR chunkIndices, GM_ADDR aqk, GM_ADDR akk, GM_ADDR qg, + GM_ADDR qgScaled, GM_ADDR wSeed, GM_ADDR uSeed, GM_ADDR, + GM_ADDR userWorkspace, const TilingData &tiling, TPipe &pipe, + bool = true) +{ + RunChunkKdaPrepare( + q, k, v, gk, beta, initialState, cuSeqlens, chunkIndices, aqk, + akk, qg, qgScaled, wSeed, uSeed, userWorkspace, tiling, pipe); +} + +} // namespace KdaPrepare diff --git a/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_varlen.h b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_varlen.h new file mode 100644 index 000000000000..65b80073060f --- /dev/null +++ b/csrc/attention/chunk_kda_fwd/op_kernel/chunk_kda_fwd_varlen.h @@ -0,0 +1,68 @@ +#pragma once + +#include "kernel_operator.h" + +namespace KdaVarlen { + +__aicore__ inline bool ResolveChunkRange( + const __gm__ int64_t *cuSeqlens, const __gm__ int64_t *chunkIndices, + uint64_t seqNum, uint64_t totalTokens, uint64_t chunkSize, + uint64_t flatChunk, uint64_t &seq, uint64_t &start, uint64_t &end) +{ + if (cuSeqlens == nullptr || chunkSize == 0) { + return false; + } + + int64_t seqValue = -1; + int64_t localChunkValue = -1; + if (chunkIndices != nullptr) { + const uint64_t metadataOffset = flatChunk * 2; + seqValue = chunkIndices[metadataOffset]; + localChunkValue = chunkIndices[metadataOffset + 1]; + } else { + uint64_t chunkPrefix = 0; + for (uint64_t seqIdx = 0; seqIdx < seqNum; ++seqIdx) { + const int64_t seqStart = cuSeqlens[seqIdx]; + const int64_t seqEnd = cuSeqlens[seqIdx + 1]; + if (seqStart < 0 || seqEnd < seqStart) { + return false; + } + const uint64_t seqLength = static_cast(seqEnd - seqStart); + const uint64_t seqChunks = (seqLength + chunkSize - 1) / chunkSize; + if (flatChunk < chunkPrefix + seqChunks) { + seqValue = static_cast(seqIdx); + localChunkValue = static_cast(flatChunk - chunkPrefix); + break; + } + chunkPrefix += seqChunks; + } + } + + if (seqValue < 0 || localChunkValue < 0 || + static_cast(seqValue) >= seqNum) { + return false; + } + const int64_t seqStartValue = cuSeqlens[seqValue]; + const int64_t seqEndValue = cuSeqlens[seqValue + 1]; + if (seqStartValue < 0 || seqEndValue < seqStartValue) { + return false; + } + const uint64_t seqStart = static_cast(seqStartValue); + const uint64_t seqEnd = static_cast(seqEndValue); + const uint64_t localChunk = static_cast(localChunkValue); + start = seqStart + localChunk * chunkSize; + if (start >= seqEnd || start >= totalTokens) { + return false; + } + end = start + chunkSize; + if (end > seqEnd) { + end = seqEnd; + } + if (end > totalTokens) { + return false; + } + seq = static_cast(seqValue); + return start < end; +} + +} // namespace KdaVarlen diff --git a/csrc/attention/kda_gate_cumsum/op_kernel/kda_gate_cumsum_kernel.h b/csrc/attention/kda_gate_cumsum/op_kernel/kda_gate_cumsum_kernel.h new file mode 100644 index 000000000000..e8894bbb59bd --- /dev/null +++ b/csrc/attention/kda_gate_cumsum/op_kernel/kda_gate_cumsum_kernel.h @@ -0,0 +1,746 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). Please refer to the License for details. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND. + */ + +#pragma once + +#include "kernel_operator.h" +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +#ifndef FLA_NPU_REGBASE_HPP_INCLUDED +#define FLA_NPU_REGBASE_HPP_INCLUDED +#include "kernel_utils/vector/regbase.hpp" +#endif +#endif + +namespace KdaGateCumsum { + +using namespace AscendC; + +constexpr float RCP_LN2 = 1.4426950408889634f; +constexpr uint32_t GATE_ROW_ELEMENTS = 256; +constexpr uint32_t GATE_PIPELINE_DEPTH = 2; +constexpr uint32_t GATE_BULK_ROWS = 64; +constexpr uint32_t GATE_BULK_COLS = 128; +constexpr uint32_t GATE_BULK_ELEMENTS = GATE_BULK_ROWS * GATE_BULK_COLS; + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 +static __simd_vf__ inline void AccumulateGateRowRegbase(__ubuf__ float *input, __ubuf__ float *acc, + __ubuf__ float *output, uint16_t count) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t FLOAT_ELEMENTS_PER_REG = AscendC::VECTOR_REG_WIDTH / sizeof(float); + constexpr uint16_t ELEMENTS_PER_PAIR = 2 * FLOAT_ELEMENTS_PER_REG; + + for (uint16_t offset = 0; offset < count; offset += ELEMENTS_PER_PAIR) { + RegTensor inputZeroReg; + RegTensor inputOneReg; + RegTensor accZeroReg; + RegTensor accOneReg; + MaskReg floatMask = CreateMask(); + + LoadAlign(inputZeroReg, input + offset); + LoadAlign(inputOneReg, input + offset + FLOAT_ELEMENTS_PER_REG); + LoadAlign(accZeroReg, acc + offset); + LoadAlign(accOneReg, acc + offset + FLOAT_ELEMENTS_PER_REG); + + Muls(inputZeroReg, inputZeroReg, RCP_LN2, floatMask); + Muls(inputOneReg, inputOneReg, RCP_LN2, floatMask); + Add(accZeroReg, accZeroReg, inputZeroReg, floatMask); + Add(accOneReg, accOneReg, inputOneReg, floatMask); + + StoreAlign(acc + offset, accZeroReg, floatMask); + StoreAlign(acc + offset + FLOAT_ELEMENTS_PER_REG, accOneReg, floatMask); + StoreAlign(output + offset, accZeroReg, floatMask); + StoreAlign(output + offset + FLOAT_ELEMENTS_PER_REG, accOneReg, floatMask); + } +} + +static __simd_vf__ inline void AccumulateGateChunk128Regbase(__ubuf__ float *input, + __ubuf__ float *output, + uint16_t rows) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t FLOAT_ELEMENTS_PER_REG = AscendC::VECTOR_REG_WIDTH / sizeof(float); + constexpr uint16_t ROW_ELEMENTS = 2 * FLOAT_ELEMENTS_PER_REG; + + MaskReg floatMask = CreateMask(); + RegTensor accZeroReg; + RegTensor accOneReg; + Duplicate(accZeroReg, 0.0f, floatMask); + Duplicate(accOneReg, 0.0f, floatMask); + for (uint16_t row = 0; row < rows; ++row) { + uint32_t rowOffset = static_cast(row) * ROW_ELEMENTS; + RegTensor inputZeroReg; + RegTensor inputOneReg; + LoadAlign(inputZeroReg, input + rowOffset); + LoadAlign( + inputOneReg, input + rowOffset + FLOAT_ELEMENTS_PER_REG); + Muls(inputZeroReg, inputZeroReg, RCP_LN2, floatMask); + Muls(inputOneReg, inputOneReg, RCP_LN2, floatMask); + Add(accZeroReg, accZeroReg, inputZeroReg, floatMask); + Add(accOneReg, accOneReg, inputOneReg, floatMask); + StoreAlign(output + rowOffset, accZeroReg, floatMask); + StoreAlign(output + rowOffset + FLOAT_ELEMENTS_PER_REG, accOneReg, floatMask); + } +} + +template +static __simd_vf__ inline void AccumulateSafeGateChunk128Regbase( + __ubuf__ float *input, __ubuf__ float *bias, __ubuf__ float *output, + uint16_t rows, float expA, float lowerBound) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t FLOAT_ELEMENTS_PER_REG = AscendC::VECTOR_REG_WIDTH / sizeof(float); + constexpr uint16_t ROW_ELEMENTS = 2 * FLOAT_ELEMENTS_PER_REG; + + MaskReg floatMask = CreateMask(); + RegTensor accZeroReg; + RegTensor accOneReg; + RegTensor oneZeroReg; + RegTensor oneOneReg; + RegTensor biasZeroReg; + RegTensor biasOneReg; + Duplicate(accZeroReg, 0.0f, floatMask); + Duplicate(accOneReg, 0.0f, floatMask); + Duplicate(oneZeroReg, 1.0f, floatMask); + Duplicate(oneOneReg, 1.0f, floatMask); + if constexpr (HAS_BIAS) { + LoadAlign(biasZeroReg, bias); + LoadAlign(biasOneReg, bias + FLOAT_ELEMENTS_PER_REG); + } + + const float gateScale = lowerBound * RCP_LN2; + for (uint16_t row = 0; row < rows; ++row) { + uint32_t rowOffset = static_cast(row) * ROW_ELEMENTS; + RegTensor gateZeroReg; + RegTensor gateOneReg; + RegTensor sigmoidZeroReg; + RegTensor sigmoidOneReg; + LoadAlign(gateZeroReg, input + rowOffset); + LoadAlign( + gateOneReg, input + rowOffset + FLOAT_ELEMENTS_PER_REG); + if constexpr (HAS_BIAS) { + Add(gateZeroReg, gateZeroReg, biasZeroReg, floatMask); + Add(gateOneReg, gateOneReg, biasOneReg, floatMask); + } + Muls(gateZeroReg, gateZeroReg, -expA, floatMask); + Muls(gateOneReg, gateOneReg, -expA, floatMask); + Exp(gateZeroReg, gateZeroReg, floatMask); + Exp(gateOneReg, gateOneReg, floatMask); + Adds(gateZeroReg, gateZeroReg, 1.0f, floatMask); + Adds(gateOneReg, gateOneReg, 1.0f, floatMask); + Div(sigmoidZeroReg, oneZeroReg, gateZeroReg, floatMask); + Div(sigmoidOneReg, oneOneReg, gateOneReg, floatMask); + Muls(sigmoidZeroReg, sigmoidZeroReg, gateScale, floatMask); + Muls(sigmoidOneReg, sigmoidOneReg, gateScale, floatMask); + Add(accZeroReg, accZeroReg, sigmoidZeroReg, floatMask); + Add(accOneReg, accOneReg, sigmoidOneReg, floatMask); + StoreAlign(output + rowOffset, accZeroReg, floatMask); + StoreAlign(output + rowOffset + FLOAT_ELEMENTS_PER_REG, accOneReg, floatMask); + } +} + +#endif + +template +class KdaGateCumsumKernel { +public: + template + __aicore__ inline void Init(GM_ADDR g, GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR cuSeqlens, GM_ADDR gk, + const TilingData &tiling, TPipe *pipe) + { + g_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(g)); + aLog_.SetGlobalBuffer(reinterpret_cast<__gm__ float *>(aLog)); + dtBias_.SetGlobalBuffer(reinterpret_cast<__gm__ float *>(dtBias)); + cuSeqlens_.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(cuSeqlens)); + gk_.SetGlobalBuffer(reinterpret_cast<__gm__ float *>(gk)); + pipe_ = pipe; + batch_ = static_cast(tiling.batch); + t_ = static_cast(tiling.t); + hv_ = static_cast(tiling.hv); + k_ = static_cast(tiling.k); + rank_ = static_cast(tiling.rank); + chunkSize_ = static_cast(tiling.chunkSize); + seqNum_ = static_cast(tiling.seqNum); + hasCuSeqlens_ = tiling.hasCuSeqlens != 0; + hasALog_ = tiling.hasALog != 0; + hasDtBias_ = tiling.hasDtBias != 0; + inputSequenceMajor_ = tiling.inputSequenceMajor != 0; + lowerBound_ = tiling.lowerBound; + usedCoreNum_ = static_cast(tiling.usedCoreNum); + maxChunks_ = (t_ + chunkSize_ - 1) / chunkSize_; + pipe_->InitBuffer(rowBuf_, GATE_ROW_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(accBuf_, GATE_ROW_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(outBuf_, GATE_PIPELINE_DEPTH * GATE_ROW_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(tmpBuf_, GATE_ROW_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(oneBuf_, GATE_ROW_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(biasBuf_, GATE_ROW_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(inBuf_, GATE_PIPELINE_DEPTH * GATE_ROW_ELEMENTS * sizeof(T)); + pipe_->InitBuffer(chunkBuf_, 2 * GATE_BULK_ELEMENTS * sizeof(float)); + pipe_->InitBuffer(scalarBuf_, 32); + pipe_->InitBuffer(scalarI64Buf_, 32); + AllocEvents(); + } + + __aicore__ inline void Process() + { + uint64_t taskCount = hasCuSeqlens_ ? seqNum_ * hv_ : batch_ * hv_ * maxChunks_; + uint64_t coreIdx = static_cast(GetBlockIdx()); + for (uint64_t task = coreIdx; task < taskCount; task += usedCoreNum_) { + ProcessTask(task); + } + ReleaseEvents(); + } + +private: + __aicore__ inline void AllocEvents() + { + for (uint32_t slot = 0; slot < GATE_PIPELINE_DEPTH; ++slot) { + inputMte2ToVEvent_[slot] = pipe_->AllocEventID(); + inputVToMte2Event_[slot] = pipe_->AllocEventID(); + outputVToMte3Event_[slot] = pipe_->AllocEventID(); + outputMte3ToVEvent_[slot] = pipe_->AllocEventID(); + } + auxMte2ToVEvent_ = pipe_->AllocEventID(); + scalarVToSEvent_ = pipe_->AllocEventID(); + bulkMte3ToMte2Event_ = pipe_->AllocEventID(); + } + + __aicore__ inline void ReleaseEvents() + { + for (uint32_t slot = 0; slot < GATE_PIPELINE_DEPTH; ++slot) { + pipe_->ReleaseEventID(inputMte2ToVEvent_[slot]); + pipe_->ReleaseEventID(inputVToMte2Event_[slot]); + pipe_->ReleaseEventID(outputVToMte3Event_[slot]); + pipe_->ReleaseEventID(outputMte3ToVEvent_[slot]); + } + pipe_->ReleaseEventID(auxMte2ToVEvent_); + pipe_->ReleaseEventID(scalarVToSEvent_); + pipe_->ReleaseEventID(bulkMte3ToMte2Event_); + } + + __aicore__ inline uint64_t InputOffset(uint64_t b, uint64_t t, uint64_t hv, uint64_t k) const + { + if (inputSequenceMajor_) { + if (rank_ == 4) { + return ((b * t_ + t) * hv_ + hv) * k_ + k; + } + return (t * hv_ + hv) * k_ + k; + } + if (rank_ == 4) { + return ((b * hv_ + hv) * t_ + t) * k_ + k; + } + return (hv * t_ + t) * k_ + k; + } + + __aicore__ inline uint64_t OutputOffset(uint64_t b, uint64_t t, uint64_t hv, uint64_t k) const + { + if (rank_ == 4) { + return ((b * hv_ + hv) * t_ + t) * k_ + k; + } + return (hv * t_ + t) * k_ + k; + } + + __aicore__ inline void CopyVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, uint64_t count) + { + uint64_t rowBytes = count * static_cast(sizeof(T)); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst, src[offset], static_cast(count)); + } else { + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + } + + __aicore__ inline void CopyFloatVectorIn(LocalTensor &dst, GlobalTensor &src, uint64_t offset, + uint64_t count) + { + uint64_t rowBytes = count * sizeof(float); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst, src[offset], static_cast(count)); + } else { + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, src[offset], params, padParams); + } + } + + __aicore__ inline void CopyFloatVectorOut(GlobalTensor &dst, uint64_t offset, LocalTensor &src, + uint64_t count) + { + uint64_t rowBytes = count * sizeof(float); + if (rowBytes >= 32 && rowBytes % 32 == 0) { + DataCopy(dst[offset], src, static_cast(count)); + } else { + DataCopyParams params{1, static_cast(rowBytes), 0, 0}; + DataCopyPad(dst[offset], src, params); + } + } + + __aicore__ inline void CopyGateRowsIn(LocalTensor &dst, uint64_t b, uint64_t start, + uint64_t hv, uint64_t rows) + { + if (!inputSequenceMajor_) { + CopyVectorIn(dst, g_, InputOffset(b, start, hv, 0), rows * k_); + return; + } + DataCopyExtParams params{ + static_cast(rows), + static_cast(k_ * sizeof(T)), + static_cast((hv_ - 1) * k_ * sizeof(T)), + 0, + 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + DataCopyPad(dst, g_[InputOffset(b, start, hv, 0)], params, padParams); + } + + __aicore__ inline void PrefetchGateRow(uint64_t offset, uint32_t slot) + { + LocalTensor input = inBuf_.Get()[slot * GATE_ROW_ELEMENTS]; + CopyVectorIn(input, g_, offset, k_); + SetFlag(inputMte2ToVEvent_[slot]); + } + + __aicore__ inline void MaterializeGateRow(uint32_t slot, LocalTensor &row) + { + LocalTensor input = inBuf_.Get()[slot * GATE_ROW_ELEMENTS]; + if constexpr (IsSameType::value) { + Adds(row, input, 0.0f, static_cast(k_)); + } else { + Cast(row, input, RoundMode::CAST_NONE, static_cast(k_)); + } + PipeBarrier(); + } + + __aicore__ inline float ReadFloat(GlobalTensor &tensor, uint64_t offset) + { + LocalTensor scalar = scalarBuf_.Get(); + DataCopyParams params{1, static_cast(sizeof(float)), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(scalar, tensor[offset], params, padParams); + SetFlag(auxMte2ToVEvent_); + WaitFlag(auxMte2ToVEvent_); + Adds(scalar, scalar, 0.0f, 1); + PipeBarrier(); + SetFlag(scalarVToSEvent_); + WaitFlag(scalarVToSEvent_); + __ubuf__ float *ptr = (__ubuf__ float *)scalar.GetPhyAddr(); + return ptr[0]; + } + + __aicore__ inline int64_t ReadInt64(GlobalTensor &tensor, uint64_t offset) + { + LocalTensor scalar = scalarI64Buf_.Get(); + DataCopyParams params{1, static_cast(sizeof(int64_t)), 0, 0}; + DataCopyPadParams padParams{false, 0, 0, 0}; + DataCopyPad(scalar, tensor[offset], params, padParams); + SetFlag(auxMte2ToVEvent_); + WaitFlag(auxMte2ToVEvent_); + SetFlag(scalarVToSEvent_); + WaitFlag(scalarVToSEvent_); + __ubuf__ int64_t *ptr = (__ubuf__ int64_t *)scalar.GetPhyAddr(); + return ptr[0]; + } + + __aicore__ inline float ExpScalar(float x) + { + LocalTensor scalar = scalarBuf_.Get(); + Duplicate(scalar, x, 1); + PipeBarrier(); + Exp(scalar, scalar, 1); + PipeBarrier(); + SetFlag(scalarVToSEvent_); + WaitFlag(scalarVToSEvent_); + __ubuf__ float *ptr = (__ubuf__ float *)scalar.GetPhyAddr(); + return ptr[0]; + } + + __aicore__ inline void PrepareGate(uint64_t hv) + { + if constexpr (USE_GATE_IN_KERNEL) { + expA_ = ExpScalar(ReadFloat(aLog_, hv)); + if (hasDtBias_) { + LocalTensor bias = biasBuf_.Get(); + CopyFloatVectorIn(bias, dtBias_, hv * k_, k_); + SetFlag(auxMte2ToVEvent_); + WaitFlag(auxMte2ToVEvent_); + } + } + } + + __aicore__ inline void ApplyGate(LocalTensor &row) + { + if constexpr (USE_GATE_IN_KERNEL) { + if (hasDtBias_) { + LocalTensor bias = biasBuf_.Get(); + Add(row, row, bias, static_cast(k_)); + PipeBarrier(); + } + if constexpr (SAFE_GATE) { + Muls(row, row, expA_, static_cast(k_)); + PipeBarrier(); + + LocalTensor tmp = tmpBuf_.Get(); + Muls(tmp, row, -1.0f, static_cast(k_)); + PipeBarrier(); + Exp(tmp, tmp, static_cast(k_)); + PipeBarrier(); + Adds(tmp, tmp, 1.0f, static_cast(k_)); + PipeBarrier(); + + LocalTensor one = oneBuf_.Get(); + Duplicate(one, 1.0f, static_cast(k_)); + PipeBarrier(); + Div(row, one, tmp, static_cast(k_)); + PipeBarrier(); + Muls(row, row, lowerBound_, static_cast(k_)); + PipeBarrier(); + } else { + LocalTensor positive = oneBuf_.Get(); + LocalTensor tmp = tmpBuf_.Get(); + Maxs(positive, row, 0.0f, static_cast(k_)); + Abs(tmp, row, static_cast(k_)); + PipeBarrier(); + Muls(tmp, tmp, -1.0f, static_cast(k_)); + PipeBarrier(); + Exp(tmp, tmp, static_cast(k_)); + PipeBarrier(); + Adds(tmp, tmp, 1.0f, static_cast(k_)); + PipeBarrier(); + Ln(tmp, tmp, static_cast(k_)); + PipeBarrier(); + Add(row, positive, tmp, static_cast(k_)); + PipeBarrier(); + Muls(row, row, -expA_, static_cast(k_)); + PipeBarrier(); + } + } + } + + __aicore__ inline void ProcessTask(uint64_t task) + { + uint64_t hv = hasCuSeqlens_ ? task % hv_ : (task / maxChunks_) % hv_; + PrepareGate(hv); + if (!hasCuSeqlens_) { + uint64_t chunk = task % maxChunks_; + uint64_t b = task / (maxChunks_ * hv_); + uint64_t start = chunk * chunkSize_; + uint64_t end = start + chunkSize_; + if (end > t_) { + end = t_; + } + ProcessChunk(b, hv, start, end); + return; + } + uint64_t seq = task / hv_; + uint64_t seqStart = static_cast(ReadInt64(cuSeqlens_, seq)); + uint64_t seqEnd = static_cast(ReadInt64(cuSeqlens_, seq + 1)); + for (uint64_t start = seqStart; start < seqEnd; start += chunkSize_) { + uint64_t end = start + chunkSize_; + if (end > seqEnd) { + end = seqEnd; + } + ProcessChunk(0, hv, start, end); + } + } + + __aicore__ inline void ProcessChunk(uint64_t b, uint64_t hv, uint64_t start, uint64_t end) + { +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 + if constexpr (USE_GATE_IN_KERNEL && SAFE_GATE && IsSameType::value) { + if (k_ == GATE_BULK_COLS && chunkSize_ == GATE_BULK_ROWS && end - start == GATE_BULK_ROWS) { + ProcessChunkBulkFp32SafeVec(b, hv, start); + return; + } + } +#endif +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (USE_GATE_IN_KERNEL && SAFE_GATE && IsSameType::value) { + if (k_ == GATE_BULK_COLS && chunkSize_ == GATE_BULK_ROWS) { + ProcessChunkBulkFp32Safe(b, hv, start, end); + return; + } + } + if constexpr (!USE_GATE_IN_KERNEL && IsSameType::value) { + if (k_ == GATE_BULK_COLS && chunkSize_ == GATE_BULK_ROWS) { + ProcessChunkBulkFp32(b, hv, start, end); + return; + } + } +#endif + LocalTensor acc = accBuf_.Get(); + LocalTensor row = rowBuf_.Get(); + Duplicate(acc, 0.0f, static_cast(k_)); + PipeBarrier(); + uint64_t rows = end - start; + if (rows == 0) { + return; + } + + PrefetchGateRow(InputOffset(b, start, hv, 0), 0); + + for (uint64_t rowIdx = 0; rowIdx < rows; ++rowIdx) { + uint64_t token = start + rowIdx; + uint32_t slot = static_cast(rowIdx & 1); + WaitFlag(inputMte2ToVEvent_[slot]); + + if (rowIdx + 1 < rows) { + uint32_t nextSlot = slot ^ 1; + if (rowIdx >= 1) { + WaitFlag(inputVToMte2Event_[nextSlot]); + } + PrefetchGateRow(InputOffset(b, token + 1, hv, 0), nextSlot); + } + + uint32_t outputSlot = slot; + if (rowIdx >= GATE_PIPELINE_DEPTH) { + WaitFlag(outputMte3ToVEvent_[outputSlot]); + } + LocalTensor output = + outBuf_.Get()[outputSlot * GATE_ROW_ELEMENTS]; + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + if constexpr (!USE_GATE_IN_KERNEL && IsSameType::value) { + if ((k_ % 128) == 0) { + LocalTensor input = inBuf_.Get()[slot * GATE_ROW_ELEMENTS]; + AccumulateGateRowRegbase( + (__ubuf__ float *)input.GetPhyAddr(), (__ubuf__ float *)acc.GetPhyAddr(), + (__ubuf__ float *)output.GetPhyAddr(), static_cast(k_)); + PipeBarrier(); + } else { + MaterializeGateRow(slot, row); + Muls(row, row, RCP_LN2, static_cast(k_)); + PipeBarrier(); + Add(acc, acc, row, static_cast(k_)); + PipeBarrier(); + Adds(output, acc, 0.0f, static_cast(k_)); + } + } else { + MaterializeGateRow(slot, row); + ApplyGate(row); + Muls(row, row, RCP_LN2, static_cast(k_)); + PipeBarrier(); + Add(acc, acc, row, static_cast(k_)); + PipeBarrier(); + Adds(output, acc, 0.0f, static_cast(k_)); + } +#else + MaterializeGateRow(slot, row); + ApplyGate(row); + Muls(row, row, RCP_LN2, static_cast(k_)); + PipeBarrier(); + Add(acc, acc, row, static_cast(k_)); + PipeBarrier(); + Adds(output, acc, 0.0f, static_cast(k_)); +#endif + + PipeBarrier(); + SetFlag(inputVToMte2Event_[slot]); + SetFlag(outputVToMte3Event_[outputSlot]); + WaitFlag(outputVToMte3Event_[outputSlot]); + CopyFloatVectorOut(gk_, OutputOffset(b, token, hv, 0), output, k_); + SetFlag(outputMte3ToVEvent_[outputSlot]); + } + + uint64_t drainStart = rows > GATE_PIPELINE_DEPTH ? rows - GATE_PIPELINE_DEPTH : 0; + for (uint64_t rowIdx = drainStart; rowIdx < rows; ++rowIdx) { + uint32_t outputSlot = static_cast(rowIdx & 1); + WaitFlag(inputVToMte2Event_[outputSlot]); + WaitFlag(outputMte3ToVEvent_[outputSlot]); + } + } + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 + __aicore__ inline void ScanGateChunk64x128(LocalTensor input, LocalTensor output) + { + constexpr uint32_t rowElements = GATE_BULK_COLS; + Adds(input, output, 0.0f, rowElements); + Add(input[rowElements], output[rowElements], output, (GATE_BULK_ROWS - 1) * rowElements); + PipeBarrier(); + + Adds(output, input, 0.0f, 2 * rowElements); + Add(output[2 * rowElements], input[2 * rowElements], input, (GATE_BULK_ROWS - 2) * rowElements); + PipeBarrier(); + + Adds(input, output, 0.0f, 4 * rowElements); + Add(input[4 * rowElements], output[4 * rowElements], output, (GATE_BULK_ROWS - 4) * rowElements); + PipeBarrier(); + + Adds(output, input, 0.0f, 8 * rowElements); + Add(output[8 * rowElements], input[8 * rowElements], input, (GATE_BULK_ROWS - 8) * rowElements); + PipeBarrier(); + + Adds(input, output, 0.0f, 16 * rowElements); + Add(input[16 * rowElements], output[16 * rowElements], output, (GATE_BULK_ROWS - 16) * rowElements); + PipeBarrier(); + + Adds(output, input, 0.0f, 32 * rowElements); + Add(output[32 * rowElements], input[32 * rowElements], input, (GATE_BULK_ROWS - 32) * rowElements); + PipeBarrier(); + } + + __aicore__ inline void ProcessChunkBulkFp32SafeVec(uint64_t b, uint64_t hv, uint64_t start) + { + constexpr uint32_t elems = GATE_BULK_ELEMENTS; + LocalTensor input = chunkBuf_.Get(); + LocalTensor output = chunkBuf_.Get()[GATE_BULK_ELEMENTS]; + CopyGateRowsIn(input, b, start, hv, GATE_BULK_ROWS); + SetFlag(inputMte2ToVEvent_[0]); + WaitFlag(inputMte2ToVEvent_[0]); + + Duplicate(output, 1.0f, elems); + if (hasDtBias_) { + LocalTensor bias = biasBuf_.Get(); + for (uint32_t row = 0; row < GATE_BULK_ROWS; ++row) { + Add(input[row * GATE_BULK_COLS], input[row * GATE_BULK_COLS], bias, GATE_BULK_COLS); + } + PipeBarrier(); + } + Muls(input, input, -expA_, elems); + PipeBarrier(); + Exp(input, input, elems); + PipeBarrier(); + Adds(input, input, 1.0f, elems); + PipeBarrier(); + Div(output, output, input, elems); + PipeBarrier(); + Muls(output, output, lowerBound_ * RCP_LN2, elems); + PipeBarrier(); + ScanGateChunk64x128(input, output); + + SetFlag(outputVToMte3Event_[0]); + WaitFlag(outputVToMte3Event_[0]); + CopyFloatVectorOut(gk_, OutputOffset(b, start, hv, 0), output, elems); + SetFlag(bulkMte3ToMte2Event_); + WaitFlag(bulkMte3ToMte2Event_); + SetFlag(outputMte3ToVEvent_[0]); + WaitFlag(outputMte3ToVEvent_[0]); + } +#endif + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310 + __aicore__ inline void ProcessChunkBulkFp32Safe(uint64_t b, uint64_t hv, uint64_t start, uint64_t end) + { + uint64_t rows = end - start; + if (rows == 0) { + return; + } + uint32_t elems = static_cast(rows * k_); + LocalTensor input = chunkBuf_.Get(); + LocalTensor output = chunkBuf_.Get()[GATE_BULK_ELEMENTS]; + CopyGateRowsIn(input, b, start, hv, rows); + SetFlag(inputMte2ToVEvent_[0]); + WaitFlag(inputMte2ToVEvent_[0]); + LocalTensor bias = biasBuf_.Get(); + if (hasDtBias_) { + AccumulateSafeGateChunk128Regbase( + (__ubuf__ float *)input.GetPhyAddr(), (__ubuf__ float *)bias.GetPhyAddr(), + (__ubuf__ float *)output.GetPhyAddr(), static_cast(rows), expA_, lowerBound_); + } else { + AccumulateSafeGateChunk128Regbase( + (__ubuf__ float *)input.GetPhyAddr(), (__ubuf__ float *)bias.GetPhyAddr(), + (__ubuf__ float *)output.GetPhyAddr(), static_cast(rows), expA_, lowerBound_); + } + PipeBarrier(); + SetFlag(outputVToMte3Event_[0]); + WaitFlag(outputVToMte3Event_[0]); + CopyFloatVectorOut(gk_, OutputOffset(b, start, hv, 0), output, elems); + SetFlag(bulkMte3ToMte2Event_); + WaitFlag(bulkMte3ToMte2Event_); + SetFlag(outputMte3ToVEvent_[0]); + WaitFlag(outputMte3ToVEvent_[0]); + } + + __aicore__ inline void ProcessChunkBulkFp32(uint64_t b, uint64_t hv, uint64_t start, uint64_t end) + { + uint64_t rows = end - start; + if (rows == 0) { + return; + } + uint32_t elems = static_cast(rows * k_); + LocalTensor buffer0 = chunkBuf_.Get(); + CopyGateRowsIn(buffer0, b, start, hv, rows); + SetFlag(inputMte2ToVEvent_[0]); + WaitFlag(inputMte2ToVEvent_[0]); + AccumulateGateChunk128Regbase( + (__ubuf__ float *)buffer0.GetPhyAddr(), (__ubuf__ float *)buffer0.GetPhyAddr(), + static_cast(rows)); + PipeBarrier(); + SetFlag(outputVToMte3Event_[0]); + WaitFlag(outputVToMte3Event_[0]); + CopyFloatVectorOut(gk_, OutputOffset(b, start, hv, 0), buffer0, elems); + SetFlag(bulkMte3ToMte2Event_); + WaitFlag(bulkMte3ToMte2Event_); + SetFlag(outputMte3ToVEvent_[0]); + WaitFlag(outputMte3ToVEvent_[0]); + } +#endif + + GlobalTensor g_; + GlobalTensor aLog_; + GlobalTensor dtBias_; + GlobalTensor cuSeqlens_; + GlobalTensor gk_; + TPipe *pipe_ = nullptr; + TBuf rowBuf_; + TBuf accBuf_; + TBuf outBuf_; + TBuf tmpBuf_; + TBuf oneBuf_; + TBuf biasBuf_; + TBuf inBuf_; + TBuf chunkBuf_; + TBuf scalarBuf_; + TBuf scalarI64Buf_; + TEventID inputMte2ToVEvent_[GATE_PIPELINE_DEPTH]; + TEventID inputVToMte2Event_[GATE_PIPELINE_DEPTH]; + TEventID outputVToMte3Event_[GATE_PIPELINE_DEPTH]; + TEventID outputMte3ToVEvent_[GATE_PIPELINE_DEPTH]; + TEventID auxMte2ToVEvent_; + TEventID scalarVToSEvent_; + TEventID bulkMte3ToMte2Event_; + uint64_t batch_ = 0; + uint64_t t_ = 0; + uint64_t hv_ = 0; + uint64_t k_ = 0; + uint64_t rank_ = 0; + uint64_t chunkSize_ = 0; + uint64_t seqNum_ = 0; + uint64_t maxChunks_ = 0; + bool hasCuSeqlens_ = false; + bool hasALog_ = false; + bool hasDtBias_ = false; + bool inputSequenceMajor_ = false; + float expA_ = 1.0f; + float lowerBound_ = -5.0f; + uint64_t usedCoreNum_ = 1; +}; + +template +__aicore__ inline void RunKdaGateCumsum(GM_ADDR g, GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR cuSeqlens, GM_ADDR gk, + const TilingData &tilingData, TPipe *pipe) +{ + KdaGateCumsumKernel op; + op.Init(g, aLog, dtBias, cuSeqlens, gk, tilingData, pipe); + op.Process(); +} + +template +__aicore__ inline void DispatchKdaGateCumsum(GM_ADDR g, GM_ADDR aLog, GM_ADDR dtBias, GM_ADDR cuSeqlens, + GM_ADDR gk, const TilingData &tilingData, TPipe *pipe) +{ + if (tilingData.useGateInKernel != 0) { + if (tilingData.safeGate != 0) { + RunKdaGateCumsum(g, aLog, dtBias, cuSeqlens, gk, tilingData, pipe); + } else { + RunKdaGateCumsum(g, aLog, dtBias, cuSeqlens, gk, tilingData, pipe); + } + } else { + RunKdaGateCumsum(g, aLog, dtBias, cuSeqlens, gk, tilingData, pipe); + } +} +} // namespace KdaGateCumsum diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_regbase.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_regbase.hpp new file mode 100644 index 000000000000..12c4049c18a3 --- /dev/null +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_regbase.hpp @@ -0,0 +1,216 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + */ + +#ifndef BLOCK_EPILOGUE_GDN_FWDH_REGBASE_HPP +#define BLOCK_EPILOGUE_GDN_FWDH_REGBASE_HPP + +#include "kernel_operator.h" +#ifndef FLA_NPU_REGBASE_HPP_INCLUDED +#define FLA_NPU_REGBASE_HPP_INCLUDED +#include "kernel_utils/vector/regbase.hpp" +#endif + +namespace Catlass::Epilogue::Block::detail { + +using namespace AscendC::MicroAPI; + +constexpr CastTrait KDA_B16_TO_F32_ZERO = { + RegLayout::ZERO, + SatMode::UNKNOWN, + MaskMergeMode::ZEROING, + AscendC::RoundMode::UNKNOWN, +}; + +template +__simd_callee__ inline void LoadKdaAsFloat( + RegTensor &dst, __ubuf__ T *src, MaskReg &mask) +{ + if constexpr (std::is_same()) { + DataCopy(dst, src); + } else if constexpr (std::is_same() || std::is_same()) { + RegTensor raw; + DataCopy(raw, src); + Cast(dst, raw, mask); + } else { + static_assert(!std::is_same::value, "KDA regbase only supports float/half/bfloat16_t"); + } +} + +__simd_callee__ inline void LoadKdaFloat(RegTensor &dst, __ubuf__ float *src) +{ + DataCopy(dst, src); +} + +__simd_callee__ inline void StoreKdaFloat( + __ubuf__ float *dst, RegTensor &src, MaskReg &mask) +{ + DataCopy(dst, src, mask); +} + +template +static __simd_vf__ inline void PrepareKGateRegbase( + __ubuf__ float *gateOutput, __ubuf__ T *gateInput, uint16_t count) +{ + constexpr uint16_t FP32_PER_REG = AscendC::VECTOR_REG_WIDTH / sizeof(float); + constexpr float LN2 = 0.6931471805599453f; + RegTensor gateReg; + MaskReg mask; + uint32_t remaining = count; + uint32_t offset = 0; + while (remaining > 0) { + mask = UpdateMask(remaining); + LoadKdaAsFloat(gateReg, gateInput + offset, mask); + if constexpr (USE_EXP2) { + Muls(gateReg, gateReg, LN2, mask); + } + Exp(gateReg, gateReg, mask); + StoreKdaFloat(gateOutput + offset, gateReg, mask); + offset += FP32_PER_REG; + } +} + +template +static __simd_vf__ inline void ComputeVNewRegbaseDualIssue( + __ubuf__ float *workspace, __ubuf__ T *uInput, uint32_t count) +{ + constexpr uint32_t FP32_PER_REG = AscendC::VECTOR_REG_WIDTH / sizeof(float); + RegTensor uReg0; + RegTensor uReg1; + RegTensor wsReg0; + RegTensor wsReg1; + MaskReg mask0; + MaskReg mask1; + + uint32_t remaining = count; + uint32_t offset = 0; + while (remaining > FP32_PER_REG) { + mask0 = UpdateMask(remaining); + mask1 = UpdateMask(remaining); + LoadKdaAsFloat(uReg0, uInput + offset, mask0); + LoadKdaAsFloat(uReg1, uInput + offset + FP32_PER_REG, mask1); + LoadKdaFloat(wsReg0, workspace + offset); + LoadKdaFloat(wsReg1, workspace + offset + FP32_PER_REG); + Sub(wsReg0, uReg0, wsReg0, mask0); + Sub(wsReg1, uReg1, wsReg1, mask1); + StoreKdaFloat(workspace + offset, wsReg0, mask0); + StoreKdaFloat(workspace + offset + FP32_PER_REG, wsReg1, mask1); + offset += 2 * FP32_PER_REG; + } + if (remaining > 0) { + mask0 = UpdateMask(remaining); + LoadKdaAsFloat(uReg0, uInput + offset, mask0); + LoadKdaFloat(wsReg0, workspace + offset); + Sub(wsReg0, uReg0, wsReg0, mask0); + StoreKdaFloat(workspace + offset, wsReg0, mask0); + } +} + +template +static __simd_vf__ inline void ApplyKGateUpdateRegbaseDualIssue( + __ubuf__ float *update, __ubuf__ T *state, __ubuf__ float *rowScale, + uint16_t rows, uint16_t cols) +{ + constexpr uint16_t FP32_PER_REG = AscendC::VECTOR_REG_WIDTH / sizeof(float); + RegTensor stateReg0; + RegTensor stateReg1; + RegTensor updateReg0; + RegTensor updateReg1; + RegTensor scaleReg0; + RegTensor scaleReg1; + MaskReg mask0; + MaskReg mask1; + + uint16_t row = 0; + for (; row + 1 < rows; row += 2) { + LoadAlign(scaleReg0, rowScale + row); + LoadAlign(scaleReg1, rowScale + row + 1); + uint32_t remaining0 = cols; + uint32_t remaining1 = cols; + uint16_t colLoops = static_cast((cols + FP32_PER_REG - 1) / FP32_PER_REG); + for (uint16_t colLoop = 0; colLoop < colLoops; ++colLoop) { + uint32_t col = static_cast(colLoop) * FP32_PER_REG; + mask0 = UpdateMask(remaining0); + mask1 = UpdateMask(remaining1); + LoadKdaAsFloat(stateReg0, state + static_cast(row) * cols + col, mask0); + LoadKdaAsFloat(stateReg1, state + static_cast(row + 1) * cols + col, mask1); + LoadKdaFloat(updateReg0, update + static_cast(row) * cols + col); + LoadKdaFloat(updateReg1, update + static_cast(row + 1) * cols + col); + Mul(stateReg0, stateReg0, scaleReg0, mask0); + Mul(stateReg1, stateReg1, scaleReg1, mask1); + Add(updateReg0, stateReg0, updateReg0, mask0); + Add(updateReg1, stateReg1, updateReg1, mask1); + StoreKdaFloat(update + static_cast(row) * cols + col, updateReg0, mask0); + StoreKdaFloat(update + static_cast(row + 1) * cols + col, updateReg1, mask1); + } + } + + if (row < rows) { + LoadAlign(scaleReg0, rowScale + row); + uint32_t remaining = cols; + uint16_t colLoops = static_cast((cols + FP32_PER_REG - 1) / FP32_PER_REG); + for (uint16_t colLoop = 0; colLoop < colLoops; ++colLoop) { + uint32_t col = static_cast(colLoop) * FP32_PER_REG; + mask0 = UpdateMask(remaining); + LoadKdaAsFloat(stateReg0, state + static_cast(row) * cols + col, mask0); + LoadKdaFloat(updateReg0, update + static_cast(row) * cols + col); + Mul(stateReg0, stateReg0, scaleReg0, mask0); + Add(updateReg0, stateReg0, updateReg0, mask0); + StoreKdaFloat(update + static_cast(row) * cols + col, updateReg0, mask0); + } + } +} + +static __simd_vf__ inline void ApplyRowScaleDualIssue( + __ubuf__ float *matrix, __ubuf__ float *rowScale, uint32_t rowScaleOffset, + uint16_t rows, uint16_t cols) +{ + using namespace AscendC::MicroAPI; + constexpr uint16_t FP32_PER_REG = AscendC::VECTOR_REG_WIDTH / sizeof(float); + + RegTensor matrixReg0; + RegTensor matrixReg1; + RegTensor scaleReg0; + RegTensor scaleReg1; + MaskReg mask0; + MaskReg mask1; + + uint16_t row = 0; + for (; row + 1 < rows; row += 2) { + LoadAlign(scaleReg0, rowScale + rowScaleOffset + row); + LoadAlign(scaleReg1, rowScale + rowScaleOffset + row + 1); + uint32_t remaining0 = cols; + uint32_t remaining1 = cols; + uint16_t colLoops = static_cast((cols + FP32_PER_REG - 1) / FP32_PER_REG); + for (uint16_t colLoop = 0; colLoop < colLoops; ++colLoop) { + uint32_t col = static_cast(colLoop) * FP32_PER_REG; + mask0 = UpdateMask(remaining0); + mask1 = UpdateMask(remaining1); + LoadAlign(matrixReg0, matrix + static_cast(row) * cols + col); + LoadAlign(matrixReg1, matrix + static_cast(row + 1) * cols + col); + Mul(matrixReg0, matrixReg0, scaleReg0, mask0); + Mul(matrixReg1, matrixReg1, scaleReg1, mask1); + StoreAlign(matrix + static_cast(row) * cols + col, matrixReg0, mask0); + StoreAlign(matrix + static_cast(row + 1) * cols + col, matrixReg1, mask1); + } + } + + if (row < rows) { + LoadAlign(scaleReg0, rowScale + rowScaleOffset + row); + uint32_t remaining = cols; + uint16_t colLoops = static_cast((cols + FP32_PER_REG - 1) / FP32_PER_REG); + for (uint16_t colLoop = 0; colLoop < colLoops; ++colLoop) { + uint32_t col = static_cast(colLoop) * FP32_PER_REG; + mask0 = UpdateMask(remaining); + LoadAlign(matrixReg0, matrix + static_cast(row) * cols + col); + Mul(matrixReg0, matrixReg0, scaleReg0, mask0); + StoreAlign(matrix + static_cast(row) * cols + col, matrixReg0, mask0); + } + } +} + +} // namespace Catlass::Epilogue::Block::detail + +#endif diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_update.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_update.hpp index 9d141633659e..a159fa86420a 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_update.hpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_update.hpp @@ -15,6 +15,7 @@ #include "catlass/gemm_coord.hpp" #include "catlass/matrix_coord.hpp" #include "catlass/epilogue/tile/tile_copy.hpp" +#include "block_epilogue_gdn_fwdh_regbase.hpp" namespace Catlass::Epilogue::Block { @@ -36,6 +37,9 @@ class BlockEpilogue < KGatedTag > { static constexpr bool kGated = KGatedTag::value; + static constexpr bool scalarGated = KGatedTag::scalarGated; + static constexpr bool useExp2 = KGatedTag::useExp2; + static constexpr float LN2 = 0.6931471805599453f; public: // Type aliases using DispatchPolicy = EpilogueAtlasGDNFwdHUpdate; @@ -145,26 +149,49 @@ class BlockEpilogue < void ApplyRowScale( AscendC::LocalTensor matrix, AscendC::LocalTensor rowScale, - AscendC::LocalTensor rowScaleBrcb, uint32_t rows, uint32_t cols) { - constexpr uint32_t FP32_PER_BLOCK = 8; - constexpr uint32_t FP32_PER_REPEAT = 64; - uint8_t rowStride = static_cast(cols / FP32_PER_BLOCK); - AscendC::BinaryRepeatParams params(1, 1, 0, rowStride, rowStride, 1); - for (uint32_t row = 0; row < rows; row += FP32_PER_BLOCK) { - uint32_t rowsThisBlock = Min(FP32_PER_BLOCK, rows - row); - AscendC::Brcb(rowScaleBrcb, rowScale[row], 1, {1, FP32_PER_BLOCK}); - AscendC::PipeBarrier(); - for (uint32_t col = 0; col < cols; col += FP32_PER_REPEAT) { - uint32_t count = Min(FP32_PER_REPEAT, cols - col); - uint32_t offset = row * cols + col; - AscendC::Mul(matrix[offset], matrix[offset], rowScaleBrcb, - count, rowsThisBlock, params); - } - AscendC::PipeBarrier(); - } + __ubuf__ float *matrixAddr = reinterpret_cast<__ubuf__ float *>(matrix.GetPhyAddr()); + __ubuf__ float *rowScaleAddr = reinterpret_cast<__ubuf__ float *>(rowScale.GetPhyAddr()); + AscendC::VF_CALL( + matrixAddr, rowScaleAddr, 0, + static_cast(rows), static_cast(cols)); + AscendC::PipeBarrier(); + } + + CATLASS_DEVICE + void PrepareKGate( + AscendC::LocalTensor gateOutput, + AscendC::LocalTensor gateInput, + uint32_t count) + { + __ubuf__ float *gateOutputAddr = + reinterpret_cast<__ubuf__ float *>(gateOutput.GetPhyAddr()); + __ubuf__ GElementInput *gateInputAddr = + reinterpret_cast<__ubuf__ GElementInput *>(gateInput.GetPhyAddr()); + AscendC::VF_CALL>( + gateOutputAddr, gateInputAddr, static_cast(count)); + AscendC::PipeBarrier(); + } + + template + CATLASS_DEVICE + void ApplyKGateUpdate( + AscendC::LocalTensor update, + AscendC::LocalTensor state, + AscendC::LocalTensor rowScale, + uint32_t rows, + uint32_t cols) + { + __ubuf__ float *updateAddr = reinterpret_cast<__ubuf__ float *>(update.GetPhyAddr()); + __ubuf__ StateElement *stateAddr = + reinterpret_cast<__ubuf__ StateElement *>(state.GetPhyAddr()); + __ubuf__ float *rowScaleAddr = reinterpret_cast<__ubuf__ float *>(rowScale.GetPhyAddr()); + AscendC::VF_CALL>( + updateAddr, stateAddr, rowScaleAddr, + static_cast(rows), static_cast(cols)); + AscendC::PipeBarrier(); } CATLASS_DEVICE @@ -175,6 +202,7 @@ class BlockEpilogue < AscendC::GlobalTensor hInput, AscendC::GlobalTensor hUpdateInput, AscendC::GlobalTensor gkInput, + AscendC::GlobalTensor initialState, uint32_t chunkSize, uint32_t kHeadDim, uint32_t vBlockDim, @@ -183,7 +211,12 @@ class BlockEpilogue < bool isInitialState, bool isFinalState, bool storeFinalState, - bool isPing + bool useInitialState, + bool isPing, + bool cube2AlreadyWaited, + bool useDirectFp32Ub, + uint64_t directUbFreeFlagBegin, + uint64_t directUbReadyFlagBegin ) { static constexpr uint32_t ROW_TILE = 16; @@ -199,7 +232,15 @@ class BlockEpilogue < rowEnd = mActual; } if (rowBegin >= mActual) { - Arch::CrossCoreWaitFlag(cube2Done); + if (useDirectFp32Ub) { + uint32_t directUbSlot = isPing ? 0 : 1; + AscendC::CrossCoreWaitFlag<0x4, PIPE_V>( + directUbReadyFlagBegin + directUbSlot); + AscendC::CrossCoreSetFlag<0x4, PIPE_V>( + directUbFreeFlagBegin + directUbSlot); + } else if (!cube2AlreadyWaited) { + Arch::CrossCoreWaitFlag(cube2Done); + } return; } @@ -212,34 +253,53 @@ class BlockEpilogue < AscendC::LocalTensor hUbTensor = isPing ? hUbTensor_ping : hUbTensor_pong; AscendC::LocalTensor finalOutputUbTensor = isPing ? finalOutputUbTensor_ping : finalOutputUbTensor_pong; AscendC::LocalTensor glastUbTensor = isPing ? glastUbTensor_ping : glastUbTensor_pong; + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + bool useFp32StateUpdate = storeFinalState && std::is_same::value && + (!isInitialState || useInitialState); + float muls = 1.0f; + if constexpr (scalarGated) { + GElementInput gLastVal = gInputThisSubBlock.GetValue(chunkSize-1); + float gLastFloat = 0.0f; + if constexpr(std::is_same::value) { + gLastFloat = gLastVal; + } else if constexpr(std::is_same::value) { + gLastFloat = (float)gLastVal; + } else if constexpr(std::is_same::value) { + gLastFloat = AscendC::ToFloat(gLastVal); + } + glastUbTensor.SetValue(0, gLastFloat); - GElementInput gLastVal = gInputThisSubBlock.GetValue(chunkSize-1); - float gLastFloat = 0.0f; - if constexpr(std::is_same::value) { - gLastFloat = gLastVal; - } else if constexpr(std::is_same::value) { - gLastFloat = (float)gLastVal; - } else if constexpr(std::is_same::value) { - gLastFloat = AscendC::ToFloat(gLastVal); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + if constexpr (useExp2) { + AscendC::Muls(glastUbTensor, glastUbTensor, LN2, 1); + AscendC::PipeBarrier(); + } + AscendC::Exp(glastUbTensor, glastUbTensor, 1); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + muls = glastUbTensor.GetValue(0); } - glastUbTensor.SetValue(0, gLastFloat); - - AscendC::SetFlag(EVENT_ID3 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); - AscendC::Exp(glastUbTensor, glastUbTensor, 1); - AscendC::SetFlag(EVENT_ID3 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); - float muls = glastUbTensor.GetValue(0); if constexpr (kGated) { AscendC::SetFlag(EVENT_ID1 + pingpongFlag); } - AscendC::SetFlag(EVENT_ID3 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + if constexpr (scalarGated) { + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + } - Arch::CrossCoreWaitFlag(cube2Done); + if (useDirectFp32Ub) { + uint32_t directUbSlot = isPing ? 0 : 1; + AscendC::CrossCoreWaitFlag<0x4, PIPE_V>( + directUbReadyFlagBegin + directUbSlot); + } else if (!cube2AlreadyWaited) { + Arch::CrossCoreWaitFlag(cube2Done); + } // fix: need to adapt kGated. issue: A5 do not have vdim128 branch. bool waitHFromV = storeFinalState && isInitialState && std::is_same::value; bool waitUpdateFromMte3 = false; + uint32_t updateReadyEvent = EVENT_ID3 + pingpongFlag; for (uint32_t rowStart = rowBegin; rowStart < rowEnd; rowStart += ROW_TILE) { uint32_t rowsThisTile = rowEnd - rowStart; if (rowsThisTile > ROW_TILE) { @@ -250,27 +310,42 @@ class BlockEpilogue < AscendC::GlobalTensor hInputThisTile = hInput[rowStart * outputStride]; AscendC::GlobalTensor hUpdateInputThisTile = hUpdateInput[rowStart * nActual]; AscendC::GlobalTensor finalStateThisTile = finalState[rowStart * outputStride]; + AscendC::LocalTensor hUpdateUbTensorThisTile = useDirectFp32Ub + ? hUpdateUbTensor[(rowStart - rowBegin) * nActual] + : hUpdateUbTensor; if (waitHFromV) { AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); } else { AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); } - CopyGmToUb(hUbTensor, hInputThisTile, rowsThisTile, nActual, outputStride); + if constexpr (std::is_same::value) { + if (useFp32StateUpdate) { + if (isInitialState) { + CopyGmToUb(calcUbTensor, initialState[rowStart * outputStride], + rowsThisTile, nActual, outputStride); + } else { + CopyGmToUb(calcUbTensor, finalStateThisTile, + rowsThisTile, nActual, outputStride); + } + } else { + CopyGmToUb(hUbTensor, hInputThisTile, rowsThisTile, nActual, outputStride); + } + } else { + CopyGmToUb(hUbTensor, hInputThisTile, rowsThisTile, nActual, outputStride); + } AscendC::SetFlag(EVENT_ID2 + pingpongFlag); AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); - AscendC::Cast(calcUbTensor, hUbTensor, AscendC::RoundMode::CAST_NONE, rowsThisTile * nActual); - AscendC::PipeBarrier(); - if (storeFinalState && isFinalState && std::is_same::value) { - AscendC::SetFlag(EVENT_ID2 + pingpongFlag); - waitHFromV = true; - } else { - waitHFromV = false; + if (!useFp32StateUpdate) { + AscendC::Cast(calcUbTensor, hUbTensor, AscendC::RoundMode::CAST_NONE, + rowsThisTile * nActual); + AscendC::PipeBarrier(); + } + if constexpr (scalarGated) { + AscendC::Muls(calcUbTensor, calcUbTensor, muls, rowsThisTile * nActual); + AscendC::PipeBarrier(); } - - AscendC::Muls(calcUbTensor, calcUbTensor, muls, rowsThisTile * nActual); - AscendC::PipeBarrier(); if constexpr (kGated) { AscendC::GlobalTensor gkLastInput = @@ -279,9 +354,6 @@ class BlockEpilogue < isPing ? gkLastUbTensor_ping : gkLastUbTensor_pong; AscendC::LocalTensor gkInputUbTensor = isPing ? gkInputUbTensor_ping : gkInputUbTensor_pong; - AscendC::LocalTensor gkBrcbUbTensor = - isPing ? gkBrcbUbTensor_ping : gkBrcbUbTensor_pong; - if (rowStart == rowBegin) { AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); } else { @@ -294,60 +366,82 @@ class BlockEpilogue < } AscendC::SetFlag(EVENT_ID2 + pingpongFlag); AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); - if constexpr(!std::is_same::value) { - AscendC::Cast(gkLastUbTensor, gkInputUbTensor, - AscendC::RoundMode::CAST_NONE, rowsThisTile); + if constexpr(std::is_same::value) { + PrepareKGate(gkLastUbTensor, gkLastUbTensor, rowsThisTile); + } else { + PrepareKGate(gkLastUbTensor, gkInputUbTensor, rowsThisTile); } - AscendC::PipeBarrier(); - AscendC::Muls(gkLastUbTensor, gkLastUbTensor, 0.6931471805599453f, - rowsThisTile); - AscendC::PipeBarrier(); - AscendC::Exp(gkLastUbTensor, gkLastUbTensor, rowsThisTile); - AscendC::PipeBarrier(); - - ApplyRowScale(calcUbTensor, gkLastUbTensor, gkBrcbUbTensor, - rowsThisTile, nActual); + ApplyRowScale(calcUbTensor, gkLastUbTensor, rowsThisTile, nActual); AscendC::SetFlag(EVENT_ID1 + pingpongFlag); } if (waitUpdateFromMte3) { - AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(updateReadyEvent); } else { AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); } - CopyGmToUb(hUpdateUbTensor, hUpdateInputThisTile, rowsThisTile, nActual, nActual); - AscendC::SetFlag(EVENT_ID0 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); - AscendC::Add(hUpdateUbTensor, calcUbTensor, hUpdateUbTensor, rowsThisTile * nActual); + if (!useDirectFp32Ub) { + CopyGmToUb(hUpdateUbTensorThisTile, hUpdateInputThisTile, rowsThisTile, nActual, nActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } + AscendC::Add( + hUpdateUbTensorThisTile, calcUbTensor, hUpdateUbTensorThisTile, + rowsThisTile * nActual); AscendC::PipeBarrier(); + if (storeFinalState && isFinalState && std::is_same::value) { + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + waitHFromV = true; + } else { + waitHFromV = false; + } if constexpr(std::is_same::value) { - if (storeFinalState && isFinalState) { + if (storeFinalState) { + if (!isFinalState) { + AscendC::Cast(hUbTensor, hUpdateUbTensorThisTile, + AscendC::RoundMode::CAST_RINT, + rowsThisTile * nActual); + AscendC::PipeBarrier(); + } AscendC::SetFlag(EVENT_ID0 + pingpongFlag); AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); - CopyUbToGm(finalStateThisTile, hUpdateUbTensor, rowsThisTile, nActual, outputStride); - AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + CopyUbToGm(finalStateThisTile, hUpdateUbTensorThisTile, + rowsThisTile, nActual, outputStride); + AscendC::PipeBarrier(); + AscendC::SetFlag(updateReadyEvent); + AscendC::SetFlag(updateReadyEvent); + AscendC::WaitFlag(updateReadyEvent); waitUpdateFromMte3 = true; + if (!isFinalState) { + CopyUbToGm(hOutputThisTile, hUbTensor, + rowsThisTile, nActual, outputStride); + AscendC::SetFlag( + EVENT_ID2 + pingpongFlag); + } } else { - AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); + AscendC::Cast(hUbTensor, hUpdateUbTensorThisTile, + AscendC::RoundMode::CAST_RINT, + rowsThisTile * nActual); AscendC::PipeBarrier(); AscendC::SetFlag(EVENT_ID0 + pingpongFlag); AscendC::SetFlag(EVENT_ID2 + pingpongFlag); AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); - CopyUbToGm(hOutputThisTile, hUbTensor, rowsThisTile, nActual, outputStride); + CopyUbToGm(hOutputThisTile, hUbTensor, + rowsThisTile, nActual, outputStride); AscendC::SetFlag(EVENT_ID2 + pingpongFlag); waitUpdateFromMte3 = false; } } else { if (storeFinalState && isFinalState) { - AscendC::Cast(finalOutputUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); + AscendC::Cast(finalOutputUbTensor, hUpdateUbTensorThisTile, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); AscendC::PipeBarrier(); AscendC::SetFlag(EVENT_ID0 + pingpongFlag); AscendC::SetFlag(EVENT_ID2 + pingpongFlag); AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); CopyUbToGm(finalStateThisTile, finalOutputUbTensor, rowsThisTile, nActual, outputStride); } else { - AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); + AscendC::Cast(hUbTensor, hUpdateUbTensorThisTile, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); AscendC::PipeBarrier(); AscendC::SetFlag(EVENT_ID0 + pingpongFlag); AscendC::SetFlag(EVENT_ID2 + pingpongFlag); @@ -359,9 +453,13 @@ class BlockEpilogue < } } - if (storeFinalState && isFinalState && std::is_same::value) { - AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + if (storeFinalState && std::is_same::value) { + AscendC::WaitFlag(updateReadyEvent); AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + if (!isFinalState) { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } } else { AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); AscendC::SetFlag(EVENT_ID2 + pingpongFlag); @@ -369,6 +467,13 @@ class BlockEpilogue < if constexpr (kGated) { AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); } + if (useDirectFp32Ub) { + uint32_t directUbSlot = isPing ? 0 : 1; + AscendC::CrossCoreSetFlag<0x4, PIPE_MTE3>( + directUbFreeFlagBegin + directUbSlot); + } + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); } diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp index 696af333ae59..26cb1f9f0b30 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp @@ -15,6 +15,7 @@ #include "catlass/gemm_coord.hpp" #include "catlass/matrix_coord.hpp" #include "catlass/epilogue/tile/tile_copy.hpp" +#include "block_epilogue_gdn_fwdh_regbase.hpp" @@ -40,6 +41,9 @@ class BlockEpilogue < KGatedTag > { static constexpr bool kGated = KGatedTag::value; + static constexpr bool scalarGated = KGatedTag::scalarGated; + static constexpr bool useExp2 = KGatedTag::useExp2; + static constexpr float LN2 = 0.6931471805599453f; public: using DispatchPolicy = EpilogueAtlasGDNFwdHVnew; using ArchTag = typename DispatchPolicy::ArchTag; @@ -157,6 +161,11 @@ class BlockEpilogue < uint32_t pingpongFlag) { AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + if constexpr (!scalarGated) { + AscendC::Duplicate(gUbTensor, 1.0f, mActual); + AscendC::PipeBarrier(); + return; + } if (mActual == 1) { AscendC::Duplicate(gUbTensor, 1.0f, 1); AscendC::PipeBarrier(); @@ -190,6 +199,10 @@ class BlockEpilogue < AscendC::Sub(gUbTensor, gLastUbTensor, gUbTensor, mActual); AscendC::PipeBarrier(); + if constexpr (useExp2) { + AscendC::Muls(gUbTensor, gUbTensor, LN2, mActual); + AscendC::PipeBarrier(); + } AscendC::Exp(gUbTensor, gUbTensor, mActual); AscendC::PipeBarrier(); } @@ -202,28 +215,26 @@ class BlockEpilogue < uint32_t rows, uint32_t cols) { - constexpr uint32_t FP32_PER_BLOCK = 8; - constexpr uint32_t FP32_PER_REPEAT = 64; - uint8_t rowStride = static_cast(cols / FP32_PER_BLOCK); - AscendC::BinaryRepeatParams params(1, 1, 0, rowStride, rowStride, 1); - uint32_t localRow = 0; - while (localRow < rows) { - uint32_t scaleRow = rowScaleOffset + localRow; - uint32_t alignedScaleRow = scaleRow & ~(FP32_PER_BLOCK - 1); - uint32_t firstScaleLane = scaleRow - alignedScaleRow; - uint32_t rowsThisBlock = Min(FP32_PER_BLOCK - firstScaleLane, rows - localRow); - AscendC::Brcb(gBrcbUbTensor_, rowScale[alignedScaleRow], 1, {1, FP32_PER_BLOCK}); - AscendC::PipeBarrier(); - for (uint32_t col = 0; col < cols; col += FP32_PER_REPEAT) { - uint32_t count = Min(FP32_PER_REPEAT, cols - col); - uint32_t matrixOffset = localRow * cols + col; - AscendC::Mul(matrix[matrixOffset], matrix[matrixOffset], - gBrcbUbTensor_[firstScaleLane * FP32_PER_BLOCK], count, - rowsThisBlock, params); - } - AscendC::PipeBarrier(); - localRow += rowsThisBlock; - } + __ubuf__ float *matrixAddr = reinterpret_cast<__ubuf__ float *>(matrix.GetPhyAddr()); + __ubuf__ float *rowScaleAddr = reinterpret_cast<__ubuf__ float *>(rowScale.GetPhyAddr()); + AscendC::VF_CALL( + matrixAddr, rowScaleAddr, rowScaleOffset, + static_cast(rows), static_cast(cols)); + AscendC::PipeBarrier(); + } + + CATLASS_DEVICE + void ComputeVNew( + AscendC::LocalTensor workspace, + AscendC::LocalTensor uInput, + uint32_t count) + { + __ubuf__ float *workspaceAddr = reinterpret_cast<__ubuf__ float *>(workspace.GetPhyAddr()); + __ubuf__ UElementInput *uInputAddr = + reinterpret_cast<__ubuf__ UElementInput *>(uInput.GetPhyAddr()); + AscendC::VF_CALL>( + workspaceAddr, uInputAddr, count); + AscendC::PipeBarrier(); } CATLASS_DEVICE @@ -247,7 +258,11 @@ class BlockEpilogue < bool isFinalState, bool storeFinalState, bool waitWsFromMte3, - bool isPing + bool isPing, + bool cube1AlreadyWaited, + bool useDirectFp32Ub, + uint64_t directUbFreeFlagBegin, + uint64_t directUbReadyFlagBegin ) { static constexpr uint32_t ROW_TILE = 16; @@ -264,8 +279,24 @@ class BlockEpilogue < if (rowEnd > mActual) { rowEnd = mActual; } + uint32_t pingpongFlag = isPing ? 0 : pongBaseEvent; if (rowBegin >= mActual) { - Arch::CrossCoreWaitFlag(cube1Done); + if (useDirectFp32Ub) { + uint32_t directUbSlot = isPing ? 0 : 1; + AscendC::CrossCoreWaitFlag<0x4, PIPE_V>( + directUbReadyFlagBegin + directUbSlot); + AscendC::CrossCoreSetFlag<0x4, PIPE_V>( + directUbFreeFlagBegin + directUbSlot); + } else if (!cube1AlreadyWaited) { + Arch::CrossCoreWaitFlag(cube1Done); + } + // A zero-row AIV lane still owns the EVENT0 hand-off consumed by V2. + if (waitWsFromMte3) { + AscendC::WaitFlag( + EVENT_ID0 + pingpongFlag); + AscendC::SetFlag( + EVENT_ID0 + pingpongFlag); + } Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); return; } @@ -273,7 +304,6 @@ class BlockEpilogue < AscendC::GlobalTensor gInputThisSubBlock = gInput; - uint32_t pingpongFlag = isPing ? 0 : pongBaseEvent; AscendC::LocalTensor uUbTensor = isPing ? uUbTensor_ping : uUbTensor_pong; AscendC::LocalTensor wsUbTensor = isPing ? wsUbTensor_ping : wsUbTensor_pong; AscendC::LocalTensor gUbTensor = isPing ? gUbTensor_ping : gUbTensor_pong; @@ -297,33 +327,51 @@ class BlockEpilogue < AscendC::DataCopy(uUbTensor, uInputThisSubBlock, mActualThisSubBlock * nvActual); AscendC::SetFlag(EVENT_ID1 + pingpongFlag); AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); - AscendC::Cast(calcUbTensor, uUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nvActual); - - PrepareG(gUbTensor, gLastUbTensor, gInputUbTensor, gInputThisSubBlock, mActual, pingpongFlag); - Arch::CrossCoreWaitFlag(cube1Done); + if constexpr (scalarGated) { + AscendC::Cast(calcUbTensor, uUbTensor, AscendC::RoundMode::CAST_NONE, + mActualThisSubBlock * nvActual); + } + if constexpr (scalarGated) { + PrepareG(gUbTensor, gLastUbTensor, gInputUbTensor, gInputThisSubBlock, mActual, pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + } if (waitWsFromMte3) { AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); } else { AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); } - AscendC::DataCopy(wsUbTensor, wsInputThisSubBlock, mActualThisSubBlock * nvActual); - AscendC::SetFlag(EVENT_ID0 + pingpongFlag); - AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); - - AscendC::Sub(wsUbTensor, calcUbTensor, wsUbTensor, mActualThisSubBlock * nvActual); - AscendC::PipeBarrier(); + if (useDirectFp32Ub) { + uint32_t directUbSlot = isPing ? 0 : 1; + AscendC::CrossCoreWaitFlag<0x4, PIPE_V>( + directUbReadyFlagBegin + directUbSlot); + } else { + if (!cube1AlreadyWaited) { + Arch::CrossCoreWaitFlag(cube1Done); + } + AscendC::DataCopy(wsUbTensor, wsInputThisSubBlock, mActualThisSubBlock * nvActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } - AscendC::Copy(calcUbTensor, wsUbTensor, mActualThisSubBlock * nvActual); - AscendC::PipeBarrier(); - ApplyRowScale(calcUbTensor, gUbTensor, rowBegin, mActualThisSubBlock, nvActual); + if constexpr (scalarGated) { + AscendC::Sub(wsUbTensor, calcUbTensor, wsUbTensor, mActualThisSubBlock * nvActual); + AscendC::PipeBarrier(); + AscendC::Copy(calcUbTensor, wsUbTensor, mActualThisSubBlock * nvActual); + AscendC::PipeBarrier(); + ApplyRowScale(calcUbTensor, gUbTensor, rowBegin, mActualThisSubBlock, nvActual); + } else { + ComputeVNew(wsUbTensor, uUbTensor, mActualThisSubBlock * nvActual); + } AscendC::SetFlag(EVENT_ID3 + pingpongFlag); uint32_t nvLoops = nvActual / FLOAT_NUM_PER_REPEAT; for (uint32_t nLoop = 0; nLoop < nvLoops; nLoop++) { uint32_t castSrcOffset = nLoop * FLOAT_NUM_PER_REPEAT; uint32_t castDstOffset = nLoop * mActualThisSubBlock * FLOAT_NUM_PER_REPEAT; - AscendC::Cast(vNewDecayUbTensor[castDstOffset], calcUbTensor[castSrcOffset], AscendC::RoundMode::CAST_RINT, FLOAT_NUM_PER_REPEAT, mActualThisSubBlock, {(uint16_t)mActualThisSubBlock, 1, 1, (uint8_t)(nvLoops * 8)}); + AscendC::LocalTensor decayInput = scalarGated ? calcUbTensor : wsUbTensor; + AscendC::Cast(vNewDecayUbTensor[castDstOffset], decayInput[castSrcOffset], AscendC::RoundMode::CAST_RINT, FLOAT_NUM_PER_REPEAT, mActualThisSubBlock, {(uint16_t)mActualThisSubBlock, 1, 1, (uint8_t)(nvLoops * 8)}); } AscendC::PipeBarrier(); @@ -368,11 +416,22 @@ class BlockEpilogue < Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); } + if (useDirectFp32Ub) { + uint32_t directUbSlot = isPing ? 0 : 1; + AscendC::CrossCoreSetFlag<0x4, PIPE_V>( + directUbFreeFlagBegin + directUbSlot); + } return; } - PrepareG(gUbTensor, gLastUbTensor, gInputUbTensor, gInputThisSubBlock, mActual, pingpongFlag); - Arch::CrossCoreWaitFlag(cube1Done); + if constexpr (scalarGated) { + PrepareG(gUbTensor, gLastUbTensor, gInputUbTensor, gInputThisSubBlock, mActual, pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + } + if (!cube1AlreadyWaited) { + Arch::CrossCoreWaitFlag(cube1Done); + } uint32_t mActualPadded = (mActual + NZ_BLOCK_SIZE - 1) / NZ_BLOCK_SIZE * NZ_BLOCK_SIZE; bool waitWsThisTileFromMte3 = waitWsFromMte3; @@ -394,8 +453,10 @@ class BlockEpilogue < CopyGmToUb(uUbTensor, uInputThisTile, rowsThisTile, nvActual, inputStride); AscendC::SetFlag(EVENT_ID1 + pingpongFlag); AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); - AscendC::Cast(calcUbTensor, uUbTensor, AscendC::RoundMode::CAST_NONE, rowsThisTile * nvActual); - AscendC::PipeBarrier(); + if constexpr (scalarGated) { + AscendC::Cast(calcUbTensor, uUbTensor, AscendC::RoundMode::CAST_NONE, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + } if (waitWsThisTileFromMte3) { AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); @@ -406,18 +467,22 @@ class BlockEpilogue < CopyGmToUb(wsUbTensorThisTile, wsInputThisTile, rowsThisTile, nvActual, nvActual); AscendC::SetFlag(EVENT_ID0 + pingpongFlag); AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); - AscendC::Sub(wsUbTensorThisTile, calcUbTensor, wsUbTensorThisTile, rowsThisTile * nvActual); - AscendC::PipeBarrier(); - - AscendC::Copy(calcUbTensor, wsUbTensorThisTile, rowsThisTile * nvActual); - AscendC::PipeBarrier(); - ApplyRowScale(calcUbTensor, gUbTensor, rowStart, rowsThisTile, nvActual); + if constexpr (scalarGated) { + AscendC::Sub(wsUbTensorThisTile, calcUbTensor, wsUbTensorThisTile, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + AscendC::Copy(calcUbTensor, wsUbTensorThisTile, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + ApplyRowScale(calcUbTensor, gUbTensor, rowStart, rowsThisTile, nvActual); + } else { + ComputeVNew(wsUbTensorThisTile, uUbTensor, rowsThisTile * nvActual); + } uint32_t nvLoops = nvActual / FLOAT_NUM_PER_REPEAT; for (uint32_t nLoop = 0; nLoop < nvLoops; nLoop++) { uint32_t castSrcOffset = nLoop * FLOAT_NUM_PER_REPEAT; uint32_t castDstOffset = nLoop * rowsThisTile * FLOAT_NUM_PER_REPEAT; - AscendC::Cast(vNewDecayUbTensor[castDstOffset], calcUbTensor[castSrcOffset], AscendC::RoundMode::CAST_RINT, FLOAT_NUM_PER_REPEAT, rowsThisTile, {(uint16_t)rowsThisTile, 1, 1, (uint8_t)(nvLoops * 8)}); + AscendC::LocalTensor decayInput = scalarGated ? calcUbTensor : wsUbTensorThisTile; + AscendC::Cast(vNewDecayUbTensor[castDstOffset], decayInput[castSrcOffset], AscendC::RoundMode::CAST_RINT, FLOAT_NUM_PER_REPEAT, rowsThisTile, {(uint16_t)rowsThisTile, 1, 1, (uint8_t)(nvLoops * 8)}); } AscendC::PipeBarrier(); diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/block/block_scheduler_gdn_fwd_h.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/block/block_scheduler_gdn_fwd_h.hpp index e6e21fd1b355..5702037406a1 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/block/block_scheduler_gdn_fwd_h.hpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/block/block_scheduler_gdn_fwd_h.hpp @@ -99,6 +99,7 @@ struct BlockSchedulerGdnFwdH { uint32_t isVariedLen; uint32_t shapeBatch; uint32_t tokenBatch; + uint32_t inputTokenBatch; bool useInitialState; bool storeFinalState; uint32_t numSeqWorkspaceOffset; @@ -148,30 +149,58 @@ struct BlockSchedulerGdnFwdH { numSeqWorkspaceOffset = gdnFwdHTilingData->numSeqWorkspaceOffset; numChunksWorkspaceOffset = gdnFwdHTilingData->numChunksWorkspaceOffset; + InitRuntime(cu_seqlens, chunk_indices, user, coreIdx, coreNum); + } + + template + CATLASS_DEVICE + void InitFromData(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, const TilingData& tilingData, + GM_ADDR user, uint32_t coreIdx, uint32_t coreNum) { + batch = tilingData.batch; + seqlen = tilingData.seqlen; + kNumHead = tilingData.kNumHead; + vNumHead = tilingData.vNumHead; + kHeadDim = tilingData.kHeadDim; + vHeadDim = tilingData.vHeadDim; + chunkSize = tilingData.chunkSize; + isVariedLen = tilingData.isVariedLen; + shapeBatch = tilingData.shapeBatch; + tokenBatch = tilingData.tokenBatch; + useInitialState = tilingData.useInitialState; + storeFinalState = tilingData.storeFinalState; + numSeqWorkspaceOffset = tilingData.numSeqWorkspaceOffset; + numChunksWorkspaceOffset = tilingData.numChunksWorkspaceOffset; + + InitRuntime(cu_seqlens, chunk_indices, user, coreIdx, coreNum); + } + + CATLASS_DEVICE + void InitRuntime(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR user, + uint32_t coreIdx, uint32_t coreNum) { + gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens); gmNumSeq.SetGlobalBuffer((__gm__ int64_t *)(user + numSeqWorkspaceOffset)); gmNumChunks.SetGlobalBuffer((__gm__ int64_t *)(user + numChunksWorkspaceOffset)); if (isVariedLen) { - gmNumChunks.SetValue(0, 0); - gmNumSeq.SetValue(0, 0); + inputTokenBatch = tokenBatch; uint32_t actualBatch = 0; + int64_t chunkPrefix = 0; int64_t prevSeq = 0, currSeq; - for (uint32_t b = 1; b <= tokenBatch; b++) { + for (uint32_t b = 1; b <= inputTokenBatch; b++) { currSeq = gmSeqlen.GetValue(b); int64_t batchSeqLen = currSeq - prevSeq; if (batchSeqLen > 0) { actualBatch++; - gmNumSeq.SetValue(actualBatch, currSeq); int64_t batchChunk = (batchSeqLen + chunkSize - 1) / chunkSize; - gmNumChunks.SetValue(actualBatch, gmNumChunks.GetValue(actualBatch - 1) + batchChunk); + chunkPrefix += batchChunk; } prevSeq = currSeq; } tokenBatch = actualBatch; batch = actualBatch; - totalChunks = gmNumChunks.GetValue(tokenBatch); - totalTokens = gmNumSeq.GetValue(tokenBatch); + totalChunks = chunkPrefix; + totalTokens = prevSeq; } else { totalChunks = (seqlen + chunkSize - 1) / chunkSize; totalTokens = seqlen; @@ -196,6 +225,41 @@ struct BlockSchedulerGdnFwdH { } + CATLASS_DEVICE + void ResolveVarlenSequence(uint32_t compactBatchIdx, GDNFwdHStream& stream) { + uint32_t actualBatch = 0; + int64_t chunkPrefix = 0; + int64_t prevSeq = 0; + for (uint32_t b = 1; b <= inputTokenBatch; ++b) { + int64_t currSeq = gmSeqlen.GetValue(b); + int64_t batchTokens = currSeq - prevSeq; + if (batchTokens > 0) { + int64_t batchChunks = (batchTokens + chunkSize - 1) / chunkSize; + if (actualBatch == compactBatchIdx) { + stream.chunkOffset = static_cast(chunkPrefix); + stream.batchChunks = static_cast(batchChunks); + stream.tokenOffset = static_cast(prevSeq); + stream.batchTokens = static_cast(batchTokens); + return; + } + ++actualBatch; + chunkPrefix += batchChunks; + } + prevSeq = currSeq; + } + stream.chunkOffset = 0; + stream.batchChunks = 0; + stream.tokenOffset = 0; + stream.batchTokens = 0; + } + + CATLASS_DEVICE + uint32_t GetVarlenChunkOffset(uint32_t compactBatchIdx) { + GDNFwdHStream stream; + ResolveVarlenSequence(compactBatchIdx, stream); + return stream.chunkOffset; + } + CATLASS_DEVICE void InitNewStream(GDNFwdHStream& newStream) { newStream.batchIdx = taskIdx / vNumHead; @@ -203,10 +267,14 @@ struct BlockSchedulerGdnFwdH { newStream.kHeadIdx = newStream.vHeadIdx / headGroups; newStream.shapeBatchIdx = isVariedLen ? 0 : newStream.batchIdx; newStream.tokenBatchIdx = isVariedLen ? newStream.batchIdx : 0; - newStream.chunkOffset = isVariedLen ? gmNumChunks.GetValue(newStream.tokenBatchIdx) : 0; - newStream.batchChunks = isVariedLen ? (gmNumChunks.GetValue(newStream.tokenBatchIdx + 1) - newStream.chunkOffset) : totalChunks; - newStream.tokenOffset = isVariedLen ? gmNumSeq.GetValue(newStream.tokenBatchIdx) : 0; - newStream.batchTokens = isVariedLen ? (gmNumSeq.GetValue(newStream.tokenBatchIdx + 1) - newStream.tokenOffset) : totalTokens; + if (isVariedLen) { + ResolveVarlenSequence(newStream.tokenBatchIdx, newStream); + } else { + newStream.chunkOffset = 0; + newStream.batchChunks = totalChunks; + newStream.tokenOffset = 0; + newStream.batchTokens = totalTokens; + } newStream.chunkIdx = 0; } @@ -319,6 +387,13 @@ struct BlockSchedulerGdnFwdHCube : public BlockSchedulerGdnFwdH { BlockSchedulerGdnFwdH::Init(cu_seqlens, chunk_indices, tiling, user, AscendC::GetBlockIdx(), AscendC::GetBlockNum()); } + template + CATLASS_DEVICE + void InitFromData(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, const TilingData& tilingData, GM_ADDR user) { + BlockSchedulerGdnFwdH::InitFromData( + cu_seqlens, chunk_indices, tilingData, user, AscendC::GetBlockIdx(), AscendC::GetBlockNum()); + } + }; struct BlockSchedulerGdnFwdHVec : public BlockSchedulerGdnFwdH { @@ -333,6 +408,15 @@ struct BlockSchedulerGdnFwdHVec : public BlockSchedulerGdnFwdH { AscendC::GetBlockNum()); } + template + CATLASS_DEVICE + void InitFromData(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, const TilingData& tilingData, GM_ADDR user) { + BlockSchedulerGdnFwdH::InitFromData( + cu_seqlens, chunk_indices, tilingData, user, + AscendC::GetBlockIdx() / AscendC::GetSubBlockNum(), + AscendC::GetBlockNum()); + } + }; } // namespace Catlass::Gemm::Block diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/kernel/gdn_fwd_h_kernel.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/kernel/gdn_fwd_h_kernel.hpp index 03919be1d948..b04500c6a2e9 100644 --- a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/kernel/gdn_fwd_h_kernel.hpp +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/arch35/gemm/kernel/gdn_fwd_h_kernel.hpp @@ -19,6 +19,7 @@ #include "../../epilogue/block/block_epilogue_gdn_fwdh_update.hpp" #include "../../epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp" #include "catlass/gemm/block/block_mmad.hpp" +#include "kernel_utils/block/block_mmad_pingpong_tla.hpp" #include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp" #include "kernel_utils/block/block_mmad_pingpong_tla_preloadA_l1B.hpp" #include "catlass/gemm/block/block_swizzle.hpp" @@ -57,13 +58,20 @@ using namespace tla; namespace Catlass::Gemm::Kernel { struct GDNFwdHTileShapes128 { - using L1TileShape = Shape<_128, _128, _128>; + using L1TileShape = tla::Shape<_128, _128, _128>; using L0TileShape = L1TileShape; }; struct GDNFwdHTileShapes256 { - using L1TileShape = Shape<_128, _256, _128>; - using L0TileShape = Shape<_128, _256, _64>; + using L1TileShape = tla::Shape<_128, _256, _128>; + using L0TileShape = tla::Shape<_128, _256, _64>; +}; + +template +struct GDNFwdHGateTag { + static constexpr bool value = KGated; + static constexpr bool scalarGated = ScalarGated; + static constexpr bool useExp2 = UseExp2; }; template< @@ -72,7 +80,9 @@ template< typename STATE_TYPE, typename WORKSPACE_TYPE, typename TileShapes = GDNFwdHTileShapes128, - bool kGated = false + bool kGated = false, + bool scalarGated = true, + bool useExp2 = false > class GDNFwdHKernel { public: @@ -81,7 +91,9 @@ class GDNFwdHKernel { using CubeScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdHCube; using VecScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdHVec; - using DispatchPolicyTlaMulti = Gemm::MmadPingpongTlaMulti; + using DispatchPolicyTlaMulti = Gemm::MmadPingpongTlaMulti; + using DispatchPolicyTlaTail = Gemm::MmadPingpongTlaMulti; + using DispatchPolicyDirectUb = Common::MmadPingpong; using DispatchPolicyTlaPreloadAL1B = Gemm::MmadPingpongTlaPreloadAL1B; using L1TileShapeVTla = typename TileShapes::L1TileShape; using L0TileShapeVTla = typename TileShapes::L0TileShape; @@ -99,20 +111,35 @@ class GDNFwdHKernel { // cube 1 using TileCopyWH = Catlass::Gemm::Tile::PackedTileCopyTla; + using TileCopyWHDirectUb = Common::Tile::PackedTileCopyTlaToUB< + ArchTag, INPUT_TYPE, layout::RowMajor, INPUT_TYPE, layout::RowMajor, + WORKSPACE_TYPE, layout::RowMajor, void, Gemm::Tile::CopyL0CToUBMode::NO_SPLIT>; using BlockMmadWH = Gemm::Block::BlockMmadTla; + using BlockMmadWHTail = Gemm::Block::BlockMmadTla; + using BlockMmadWHDirectUb = Common::BlockMmadTla< + DispatchPolicyDirectUb, L1TileShapeVTla, L0TileShapeVTla, + INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyWHDirectUb>; // cube 2 using TileCopyKV = Catlass::Gemm::Tile::PackedTileCopyTla; + using TileCopyKVDirectUb = Common::Tile::PackedTileCopyTlaToUB< + ArchTag, INPUT_TYPE, layout::ColumnMajor, INPUT_TYPE, layout::zN, + WORKSPACE_TYPE, layout::RowMajor, void, Gemm::Tile::CopyL0CToUBMode::NO_SPLIT>; using TileMmadKV = Gemm::Tile::TileMmadTla; using BlockMmadKV = Gemm::Block::BlockMmadTla; + using BlockMmadKVTail = Gemm::Block::BlockMmadTla; + using BlockMmadKVDirectUb = Common::BlockMmadTla< + DispatchPolicyDirectUb, L1TileShapeVTla, L0TileShapeVTla, + INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyKVDirectUb>; // vec 1 using DispatchPolicyGDNFwdHVnew = Epilogue::EpilogueAtlasGDNFwdHVnew; - using EpilogueGDNFwdHVnew = Epilogue::Block::BlockEpilogue>; + using GateTag = GDNFwdHGateTag; + using EpilogueGDNFwdHVnew = Epilogue::Block::BlockEpilogue; // vec 2 using DispatchPolicyGDNFwdHUpdate = Epilogue::EpilogueAtlasGDNFwdHUpdate; - using EpilogueGDNFwdHUpdate = Epilogue::Block::BlockEpilogue>; + using EpilogueGDNFwdHUpdate = Epilogue::Block::BlockEpilogue; using GDNFwdHOffsets = Catlass::Gemm::Block::GDNFwdHOffsets; @@ -134,6 +161,11 @@ class GDNFwdHKernel { using LayoutK = Catlass::layout::ColumnMajor; using LayoutVUpdate = typename VUpdateType::Layout; + static constexpr uint64_t DIRECT_UB_FREE_FLAG_BEGIN = 1; + static constexpr uint64_t DIRECT_UB_READY_FLAG_BEGIN = 6; + static constexpr uint64_t DIRECT_UB_FLAG_STRIDE = 16; + static constexpr uint32_t DIRECT_UB_STAGES = 2; + static constexpr uint32_t DIRECT_VEC_NUM = 2; uint32_t batch; uint32_t seqlen; @@ -153,6 +185,7 @@ class GDNFwdHKernel { uint32_t numSeqWorkspaceOffset; uint32_t numChunksWorkspaceOffset; uint32_t kDecayWorkspaceOffset; + bool useDirectFp32Ub; AscendC::GlobalTensor gmK; AscendC::GlobalTensor gmW; @@ -211,12 +244,18 @@ class GDNFwdHKernel { numSeqWorkspaceOffset = gdnFwdHTilingData->numSeqWorkspaceOffset; numChunksWorkspaceOffset = gdnFwdHTilingData->numChunksWorkspaceOffset; kDecayWorkspaceOffset = gdnFwdHTilingData->kDecayWorkspaceOffset; + uint64_t denseTaskCount = static_cast(shapeBatch) * vNumHead; + useDirectFp32Ub = std::is_same::value && + !isVariedLen && chunkSize <= 64 && + seqlen % chunkSize == 0 && + kHeadDim == 128 && vHeadDim == 128 && + denseTaskCount >= AscendC::GetBlockNum(); gmK.SetGlobalBuffer((__gm__ ElementK *)k); gmW.SetGlobalBuffer((__gm__ ElementW *)w); gmU.SetGlobalBuffer((__gm__ ElementU *)u); - gmG.SetGlobalBuffer((__gm__ ElementG *)g); - gmGk.SetGlobalBuffer((__gm__ ElementG *)gk); + gmG.SetGlobalBuffer((__gm__ ElementG *)(scalarGated ? g : gk)); + gmGk.SetGlobalBuffer((__gm__ ElementG *)(kGated ? gk : g)); gmInitialState.SetGlobalBuffer((__gm__ ElementInitialState *)inital_state); gmH.SetGlobalBuffer((__gm__ ElementH *)h); gmV.SetGlobalBuffer((__gm__ ElementV *)v_new); @@ -247,17 +286,239 @@ class GDNFwdHKernel { } } - __aicore__ inline void Process() { - if (isVariedLen) { - AscendC::SyncAll(); + template + __aicore__ inline void InitFromData( + GM_ADDR k, GM_ADDR w, GM_ADDR u, GM_ADDR g, GM_ADDR gk, GM_ADDR inital_state, + GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR h, GM_ADDR v_new, + GM_ADDR final_state, const TilingData& tilingData, GM_ADDR user) { + batch = tilingData.batch; + seqlen = tilingData.seqlen; + kNumHead = tilingData.kNumHead; + vNumHead = tilingData.vNumHead; + kHeadDim = tilingData.kHeadDim; + vHeadDim = tilingData.vHeadDim; + chunkSize = tilingData.chunkSize; + useInitialState = tilingData.useInitialState; + storeFinalState = tilingData.storeFinalState; + isVariedLen = tilingData.isVariedLen; + shapeBatch = tilingData.shapeBatch; + tokenBatch = tilingData.tokenBatch; + vWorkspaceOffset = tilingData.vWorkspaceOffset; + vUpdateWorkspaceOffset = tilingData.vUpdateWorkspaceOffset; + hWorkspaceOffset = tilingData.hWorkspaceOffset; + numSeqWorkspaceOffset = tilingData.numSeqWorkspaceOffset; + numChunksWorkspaceOffset = tilingData.numChunksWorkspaceOffset; + kDecayWorkspaceOffset = tilingData.kDecayWorkspaceOffset; + uint64_t denseTaskCount = static_cast(shapeBatch) * vNumHead; + useDirectFp32Ub = std::is_same::value && + !isVariedLen && chunkSize <= 64 && + seqlen % chunkSize == 0 && + kHeadDim == 128 && vHeadDim == 128 && + denseTaskCount >= AscendC::GetBlockNum(); + + gmK.SetGlobalBuffer((__gm__ ElementK *)k); + gmW.SetGlobalBuffer((__gm__ ElementW *)w); + gmU.SetGlobalBuffer((__gm__ ElementU *)u); + gmG.SetGlobalBuffer((__gm__ ElementG *)(scalarGated ? g : gk)); + gmInitialState.SetGlobalBuffer((__gm__ ElementInitialState *)inital_state); + gmH.SetGlobalBuffer((__gm__ ElementH *)h); + gmV.SetGlobalBuffer((__gm__ ElementV *)v_new); + gmFinalState.SetGlobalBuffer((__gm__ ElementFinalState *)final_state); + gmVWorkspace.SetGlobalBuffer((__gm__ ElementVWork *)(user + vWorkspaceOffset)); + gmVUpdateWorkspace.SetGlobalBuffer((__gm__ ElementV *)(user + vUpdateWorkspaceOffset)); + gmHWorkspace.SetGlobalBuffer((__gm__ ElementHWork *)(user + hWorkspaceOffset)); + gmGk.SetGlobalBuffer((__gm__ ElementG *)(kGated ? gk : g)); + gmKDecayWorkspace.SetGlobalBuffer((__gm__ ElementK *)(user + kDecayWorkspaceOffset)); + gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens); + gmNumSeq.SetGlobalBuffer((__gm__ int64_t *)(user + numSeqWorkspaceOffset)); + gmNumChunks.SetGlobalBuffer((__gm__ int64_t *)(user + numChunksWorkspaceOffset)); + + ubHUpdatePing = resource.ubBuf.template GetBufferByByte(32 * 1024); + ubHUpdatePong = resource.ubBuf.template GetBufferByByte(96 * 1024); + ubVWorkPing = resource.ubBuf.template GetBufferByByte(32 * 1024); + ubVWorkPong = resource.ubBuf.template GetBufferByByte(96 * 1024); + + l1VUpdatePing = resource.l1Buf.template GetBufferByByte(0); + l1VUpdatePong = resource.l1Buf.template GetBufferByByte(chunkSize * vHeadDim * sizeof(ElementV)); + + if ASCEND_IS_AIC { + cubeBlockScheduler.InitFromData(cu_seqlens, chunk_indices, tilingData, user); + } + if ASCEND_IS_AIV { + vecBlockScheduler.InitFromData(cu_seqlens, chunk_indices, tilingData, user); + } + } + + // Tail helpers borrow the stream's V_MTE2 free token and restore it before returning. + __aicore__ inline void ComputeTailVWorkspace( + const GDNFwdHOffsets& offsets, uint32_t tailEventId) + { + uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); + uint32_t subBlockNum = AscendC::GetSubBlockNum(); + uint32_t rowsPerSubBlock = CeilDiv(offsets.blockTokens, subBlockNum); + uint32_t rowBegin = subBlockIdx * rowsPerSubBlock; + uint32_t rowEnd = Min(rowBegin + rowsPerSubBlock, offsets.blockTokens); + if (rowBegin >= rowEnd) { + return; + } + AscendC::ResetMask(); + AscendC::WaitFlag(tailEventId); + + constexpr uint32_t TAIL_INPUT_OFFSET = 166 * 1024; + constexpr uint32_t TAIL_FLOAT_OFFSET = 167 * 1024; + constexpr uint32_t TAIL_ACCUM_OFFSET = 168 * 1024; + constexpr uint32_t TAIL_WEIGHT_INPUT_OFFSET = 169 * 1024; + constexpr uint32_t TAIL_WEIGHT_FLOAT_OFFSET = 170 * 1024; + AscendC::LocalTensor inputUb = + resource.ubBuf.template GetBufferByByte(TAIL_INPUT_OFFSET); + AscendC::LocalTensor floatUb = + resource.ubBuf.template GetBufferByByte(TAIL_FLOAT_OFFSET); + AscendC::LocalTensor accumUb = + resource.ubBuf.template GetBufferByByte(TAIL_ACCUM_OFFSET); + AscendC::LocalTensor weightInputUb = + resource.ubBuf.template GetBufferByByte(TAIL_WEIGHT_INPUT_OFFSET); + AscendC::LocalTensor weightFloatUb = + resource.ubBuf.template GetBufferByByte(TAIL_WEIGHT_FLOAT_OFFSET); + + for (uint32_t tokenRow = rowBegin; tokenRow < rowEnd; ++tokenRow) { + AscendC::DataCopy( + weightInputUb, + gmW[offsets.wOffset + tokenRow * kHeadDim], + kHeadDim); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::Cast( + weightFloatUb, weightInputUb, AscendC::RoundMode::CAST_NONE, + kHeadDim); + AscendC::PipeBarrier(); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + + AscendC::Duplicate(accumUb, 0.0f, offsets.vBlockDim); + AscendC::PipeBarrier(); + for (uint32_t kIdx = 0; kIdx < kHeadDim; ++kIdx) { + AscendC::DataCopy( + inputUb, gmH[offsets.hSrcOffset + kIdx * vHeadDim], + offsets.vBlockDim); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::Cast( + floatUb, inputUb, AscendC::RoundMode::CAST_NONE, + offsets.vBlockDim); + AscendC::PipeBarrier(); + float weight = weightFloatUb.GetValue(kIdx); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::Muls(floatUb, floatUb, weight, offsets.vBlockDim); + AscendC::PipeBarrier(); + AscendC::Add(accumUb, accumUb, floatUb, offsets.vBlockDim); + AscendC::PipeBarrier(); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + } + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::DataCopy( + gmVWorkspace[offsets.vWorkOffset + tokenRow * offsets.vBlockDim], + accumUb, offsets.vBlockDim); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + } + AscendC::SetFlag(tailEventId); + } + + __aicore__ inline void ComputeTailHWorkspace( + const GDNFwdHOffsets& offsets, uint32_t tailEventId) + { + uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); + uint32_t subBlockNum = AscendC::GetSubBlockNum(); + uint32_t rowsPerSubBlock = CeilDiv(kHeadDim, subBlockNum); + uint32_t rowBegin = subBlockIdx * rowsPerSubBlock; + uint32_t rowEnd = Min(rowBegin + rowsPerSubBlock, kHeadDim); + AscendC::ResetMask(); + AscendC::WaitFlag(tailEventId); + + constexpr uint32_t TAIL_INPUT_OFFSET = 166 * 1024; + constexpr uint32_t TAIL_FLOAT_OFFSET = 167 * 1024; + constexpr uint32_t TAIL_ACCUM_OFFSET = 168 * 1024; + constexpr uint32_t TAIL_WEIGHT_INPUT_OFFSET = 169 * 1024; + constexpr uint32_t TAIL_WEIGHT_FLOAT_OFFSET = 170 * 1024; + AscendC::LocalTensor inputUb = + resource.ubBuf.template GetBufferByByte(TAIL_INPUT_OFFSET); + AscendC::LocalTensor floatUb = + resource.ubBuf.template GetBufferByByte(TAIL_FLOAT_OFFSET); + AscendC::LocalTensor accumUb = + resource.ubBuf.template GetBufferByByte(TAIL_ACCUM_OFFSET); + AscendC::LocalTensor weightInputUb = + resource.ubBuf.template GetBufferByByte(TAIL_WEIGHT_INPUT_OFFSET); + AscendC::LocalTensor weightFloatUb = + resource.ubBuf.template GetBufferByByte(TAIL_WEIGHT_FLOAT_OFFSET); + + for (uint32_t kRow = rowBegin; kRow < rowEnd; ++kRow) { + AscendC::Duplicate(accumUb, 0.0f, offsets.vBlockDim); + AscendC::PipeBarrier(); + for (uint32_t tokenRow = 0; tokenRow < offsets.blockTokens; ++tokenRow) { + AscendC::DataCopy( + weightInputUb, + gmKDecayWorkspace[offsets.kDecayWorkOffset + tokenRow * kHeadDim], + kHeadDim); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::Cast( + weightFloatUb, weightInputUb, AscendC::RoundMode::CAST_NONE, + kHeadDim); + AscendC::PipeBarrier(); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::DataCopy( + inputUb, + gmVUpdateWorkspace[offsets.vWorkOffset + tokenRow * offsets.vBlockDim], + offsets.vBlockDim); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::Cast( + floatUb, inputUb, AscendC::RoundMode::CAST_NONE, + offsets.vBlockDim); + AscendC::PipeBarrier(); + float weight = weightFloatUb.GetValue(kRow); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::Muls(floatUb, floatUb, weight, offsets.vBlockDim); + AscendC::PipeBarrier(); + AscendC::Add(accumUb, accumUb, floatUb, offsets.vBlockDim); + AscendC::PipeBarrier(); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + } + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::DataCopy( + gmHWorkspace[offsets.hWorkOffset + kRow * offsets.vBlockDim], + accumUb, offsets.vBlockDim); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); } + AscendC::SetFlag(tailEventId); + } + + __aicore__ inline void Process() { + // FwdH can run after another stage in a megakernel. Start its AIC/AIV + // handshake only after every core has retired the preceding stage. + AscendC::SyncAll(); if ASCEND_IS_AIC { uint32_t coreIdx = AscendC::GetBlockIdx(); - uint32_t coreNum = AscendC::GetBlockNum(); + uint32_t coreNum = vecBlockScheduler.cubeCoreNum; BlockMmadWH blockMmadWH(resource, chunkSize * cubeBlockScheduler.vBlockSize * sizeof(ElementV) * PING_PONG_STAGES); BlockMmadKV blockMmadKV(resource, chunkSize * cubeBlockScheduler.vBlockSize * sizeof(ElementV) * PING_PONG_STAGES); + BlockMmadWHTail blockMmadWHTail(resource, chunkSize * cubeBlockScheduler.vBlockSize * sizeof(ElementV) * PING_PONG_STAGES); + BlockMmadKVTail blockMmadKVTail(resource, chunkSize * cubeBlockScheduler.vBlockSize * sizeof(ElementV) * PING_PONG_STAGES); + bool useBoundedMmad = isVariedLen || (seqlen % chunkSize != 0); auto wLayout = tla::MakeLayout(shapeBatch * kNumHead * cubeBlockScheduler.totalTokens, kHeadDim); auto hLayout = tla::MakeLayout(shapeBatch * vNumHead * cubeBlockScheduler.totalChunks * kHeadDim, vHeadDim); @@ -271,213 +532,380 @@ class GDNFwdHKernel { if (currStage == 0) { /* C1: v_work = w @ h[i] */ cubeBlockScheduler.InitTasks(); - for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { - uint32_t streamId = cubeBlockScheduler.GetStreamId(i); - const auto& stream = cubeBlockScheduler.GetStream(i); - if (cubeBlockScheduler.StreamIsDone(stream)) { - continue; + if (useDirectFp32Ub) { + BlockMmadWHDirectUb blockMmadWHDirectUb( + resource, chunkSize * cubeBlockScheduler.vBlockSize * sizeof(ElementV) * PING_PONG_STAGES); + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = cubeBlockScheduler.GetStreamId(i); + const auto& stream = cubeBlockScheduler.GetStream(i); + if (cubeBlockScheduler.StreamIsDone(stream)) { + continue; + } + + const GDNFwdHOffsets& cube1Offsets = cubeBlockScheduler.GetCurTaskOffsets(stream); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[streamId]); + if (cube1Offsets.blockTokens < 16) { + Arch::CrossCoreSetFlag<0x2, PIPE_MTE2>( + cubeBlockScheduler.cube1Done[streamId]); + continue; + } + int64_t cube1OffsetW = cube1Offsets.wOffset; + int64_t cube1OffsetH = cube1Offsets.hSrcOffset; + auto tensorW = tla::MakeTensor(gmW[cube1OffsetW], wLayout, Catlass::Arch::PositionGM{}); + auto tensorH = tla::MakeTensor(gmH[cube1OffsetH], hLayout, Catlass::Arch::PositionGM{}); + GemmCoord cube1Shape {cube1Offsets.blockTokens, cube1Offsets.vBlockDim, kHeadDim}; + auto tensorBlockW = GetTile(tensorW, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k())); + auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n())); + + auto ubLayout = tla::MakeLayout(cube1Shape.m(), cube1Shape.n()); + auto tensorUbPing = tla::MakeTensor(ubVWorkPing, ubLayout, Catlass::Arch::PositionUB{}); + auto tensorUbPong = tla::MakeTensor(ubVWorkPong, ubLayout, Catlass::Arch::PositionUB{}); + using UbTensor = decltype(tensorUbPing); + UbTensor tensorUbList[BlockMmadWHDirectUb::MAX_CUBE_VEC_SYNC_NUM]; + for (uint32_t ubIdx = 0; ubIdx < BlockMmadWHDirectUb::MAX_CUBE_VEC_SYNC_NUM; ++ubIdx) { + tensorUbList[ubIdx] = (ubIdx & 1U) ? tensorUbPong : tensorUbPing; + } + uint32_t ubListId = streamId; + uint32_t rowsPerSubBlock = CeilDiv(cube1Shape.m(), DIRECT_VEC_NUM); + blockMmadWHDirectUb( + tensorBlockW, tensorBlockH, tensorUbList, cube1Shape, rowsPerSubBlock, 0, + DIRECT_UB_FREE_FLAG_BEGIN, DIRECT_UB_READY_FLAG_BEGIN, ubListId, + DIRECT_VEC_NUM, DIRECT_UB_STAGES); } + } else if (useBoundedMmad) { + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = cubeBlockScheduler.GetStreamId(i); + const auto& stream = cubeBlockScheduler.GetStream(i); + if (cubeBlockScheduler.StreamIsDone(stream)) { + continue; + } - const GDNFwdHOffsets& cube1Offsets = cubeBlockScheduler.GetCurTaskOffsets(stream); - Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[streamId]); - auto vLayout = tla::MakeLayout(cube1Offsets.blockTokens, cube1Offsets.vBlockDim); - int64_t cube1OffsetW = cube1Offsets.wOffset; - int64_t cube1OffsetH = cube1Offsets.hSrcOffset; - int64_t cube1OffsetVwork = cube1Offsets.vWorkOffset; - auto tensorW = tla::MakeTensor(gmW[cube1OffsetW], wLayout, Catlass::Arch::PositionGM{}); - auto tensorH = tla::MakeTensor(gmH[cube1OffsetH], hLayout, Catlass::Arch::PositionGM{}); - auto tensorV = tla::MakeTensor(gmVWorkspace[cube1OffsetVwork], vLayout, Catlass::Arch::PositionGM{}); - GemmCoord cube1Shape {cube1Offsets.blockTokens, cube1Offsets.vBlockDim, kHeadDim}; - auto tensorBlockW = GetTile(tensorW, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k())); - auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n())); - auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n())); + const GDNFwdHOffsets& cube1Offsets = cubeBlockScheduler.GetCurTaskOffsets(stream); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[streamId]); + auto vLayout = tla::MakeLayout( + cube1Offsets.blockTokens, cube1Offsets.vBlockDim); + auto tensorW = tla::MakeTensor( + gmW[cube1Offsets.wOffset], wLayout, Catlass::Arch::PositionGM{}); + auto tensorH = tla::MakeTensor( + gmH[cube1Offsets.hSrcOffset], hLayout, Catlass::Arch::PositionGM{}); + auto tensorV = tla::MakeTensor( + gmVWorkspace[cube1Offsets.vWorkOffset], vLayout, + Catlass::Arch::PositionGM{}); + GemmCoord cube1Shape{ + cube1Offsets.blockTokens, cube1Offsets.vBlockDim, kHeadDim}; + auto tensorBlockW = GetTile( + tensorW, tla::MakeCoord(0, 0), + tla::MakeShape(cube1Shape.m(), cube1Shape.k())); + auto tensorBlockH = GetTile( + tensorH, tla::MakeCoord(0, 0), + tla::MakeShape(cube1Shape.k(), cube1Shape.n())); + auto tensorBlockV = GetTile( + tensorV, tla::MakeCoord(0, 0), + tla::MakeShape(cube1Shape.m(), cube1Shape.n())); + if (cube1Offsets.blockTokens < chunkSize) { + blockMmadWHTail.preSetFlags(); + blockMmadWHTail( + tensorBlockW, tensorBlockH, tensorBlockV, + cube1Shape, EmptyClass{}, true); + blockMmadWHTail.finalWaitFlags(); + } else { + blockMmadWH.preSetFlags(); + blockMmadWH( + tensorBlockW, tensorBlockH, tensorBlockV, cube1Shape); + blockMmadWH.finalWaitFlags(); + } + AscendC::PipeBarrier(); + Arch::CrossCoreSetFlag<0x2, PIPE_FIX>( + cubeBlockScheduler.cube1Done[streamId]); + } + } else { blockMmadWH.preSetFlags(); - blockMmadWH(tensorBlockW, tensorBlockH, tensorBlockV, cube1Shape); + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = cubeBlockScheduler.GetStreamId(i); + const auto& stream = cubeBlockScheduler.GetStream(i); + if (cubeBlockScheduler.StreamIsDone(stream)) { + continue; + } + + const GDNFwdHOffsets& cube1Offsets = cubeBlockScheduler.GetCurTaskOffsets(stream); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[streamId]); + if (cube1Offsets.blockTokens < 16) { + Arch::CrossCoreSetFlag<0x2, PIPE_MTE2>( + cubeBlockScheduler.cube1Done[streamId]); + continue; + } + auto vLayout = tla::MakeLayout(cube1Offsets.blockTokens, cube1Offsets.vBlockDim); + int64_t cube1OffsetW = cube1Offsets.wOffset; + int64_t cube1OffsetH = cube1Offsets.hSrcOffset; + int64_t cube1OffsetVwork = cube1Offsets.vWorkOffset; + auto tensorW = tla::MakeTensor(gmW[cube1OffsetW], wLayout, Catlass::Arch::PositionGM{}); + auto tensorH = tla::MakeTensor(gmH[cube1OffsetH], hLayout, Catlass::Arch::PositionGM{}); + auto tensorV = tla::MakeTensor(gmVWorkspace[cube1OffsetVwork], vLayout, Catlass::Arch::PositionGM{}); + GemmCoord cube1Shape {cube1Offsets.blockTokens, cube1Offsets.vBlockDim, kHeadDim}; + auto tensorBlockW = GetTile(tensorW, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k())); + auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n())); + auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n())); + + blockMmadWH(tensorBlockW, tensorBlockH, tensorBlockV, cube1Shape); + Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube1Done[streamId]); + } blockMmadWH.finalWaitFlags(); - Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube1Done[streamId]); } } else { /* C2: h[i+1] = k.T @ v_work */ - for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { - uint32_t streamId = cubeBlockScheduler.GetStreamId(i); - const auto& stream = cubeBlockScheduler.GetStream(i); - if (cubeBlockScheduler.StreamIsDone(stream)) { - continue; + if (useDirectFp32Ub) { + BlockMmadKVDirectUb blockMmadKVDirectUb( + resource, chunkSize * cubeBlockScheduler.vBlockSize * sizeof(ElementV) * PING_PONG_STAGES); + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = cubeBlockScheduler.GetStreamId(i); + const auto& stream = cubeBlockScheduler.GetStream(i); + if (cubeBlockScheduler.StreamIsDone(stream)) { + continue; + } + const GDNFwdHOffsets& cube2Offsets = cubeBlockScheduler.GetCurTaskOffsets(stream); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done[streamId]); + + if (cubeBlockScheduler.NeedProcessStage2(stream)) { + if (cube2Offsets.blockTokens < 16) { + Arch::CrossCoreSetFlag<0x2, PIPE_MTE2>( + cubeBlockScheduler.cube2Done[streamId]); + continue; + } + int64_t cube2OffsetK = kGated ? cube2Offsets.kDecayWorkOffset : cube2Offsets.wkOffset; + int64_t cube2OffsetVwork = cube2Offsets.vWorkOffset; + auto tensorK = kGated + ? tla::MakeTensor(gmKDecayWorkspace[cube2OffsetK], kLayout, Catlass::Arch::PositionGM{}) + : tla::MakeTensor(gmK[cube2OffsetK], kLayout, Catlass::Arch::PositionGM{}); + auto vUpdateLayout = tla::MakeLayout(cube2Offsets.blockTokens, cube2Offsets.vBlockDim); + auto tensorVwork = tla::MakeTensor(gmVUpdateWorkspace[cube2OffsetVwork], vUpdateLayout, Catlass::Arch::PositionGM{}); + GemmCoord cube2Shape{kHeadDim, cube2Offsets.vBlockDim, cube2Offsets.blockTokens}; + auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k())); + auto tensorBlockVwork = GetTile(tensorVwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n())); + + auto ubLayout = tla::MakeLayout(cube2Shape.m(), cube2Shape.n()); + auto tensorUbPing = tla::MakeTensor(ubHUpdatePing, ubLayout, Catlass::Arch::PositionUB{}); + auto tensorUbPong = tla::MakeTensor(ubHUpdatePong, ubLayout, Catlass::Arch::PositionUB{}); + using UbTensor = decltype(tensorUbPing); + UbTensor tensorUbList[BlockMmadKVDirectUb::MAX_CUBE_VEC_SYNC_NUM]; + for (uint32_t ubIdx = 0; ubIdx < BlockMmadKVDirectUb::MAX_CUBE_VEC_SYNC_NUM; ++ubIdx) { + tensorUbList[ubIdx] = (ubIdx & 1U) ? tensorUbPong : tensorUbPing; + } + uint32_t ubListId = streamId; + uint32_t rowsPerSubBlock = CeilDiv(cube2Shape.m(), DIRECT_VEC_NUM); + blockMmadKVDirectUb( + tensorBlockK, tensorBlockVwork, tensorUbList, cube2Shape, rowsPerSubBlock, 0, + DIRECT_UB_FREE_FLAG_BEGIN, DIRECT_UB_READY_FLAG_BEGIN, ubListId, + DIRECT_VEC_NUM, DIRECT_UB_STAGES); + } + } + } else if (useBoundedMmad) { + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = cubeBlockScheduler.GetStreamId(i); + const auto& stream = cubeBlockScheduler.GetStream(i); + if (cubeBlockScheduler.StreamIsDone(stream)) { + continue; + } + const GDNFwdHOffsets& cube2Offsets = + cubeBlockScheduler.GetCurTaskOffsets(stream); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done[streamId]); + + if (cubeBlockScheduler.NeedProcessStage2(stream)) { + int64_t cube2OffsetK = kGated + ? cube2Offsets.kDecayWorkOffset + : cube2Offsets.wkOffset; + auto tensorK = kGated + ? tla::MakeTensor( + gmKDecayWorkspace[cube2OffsetK], kLayout, + Catlass::Arch::PositionGM{}) + : tla::MakeTensor( + gmK[cube2OffsetK], kLayout, + Catlass::Arch::PositionGM{}); + auto vUpdateLayout = tla::MakeLayout( + cube2Offsets.blockTokens, cube2Offsets.vBlockDim); + auto tensorVwork = tla::MakeTensor( + gmVUpdateWorkspace[cube2Offsets.vWorkOffset], vUpdateLayout, + Catlass::Arch::PositionGM{}); + auto tensorHwork = tla::MakeTensor( + gmHWorkspace[cube2Offsets.hWorkOffset], hworkLayout, + Catlass::Arch::PositionGM{}); + GemmCoord cube2Shape{ + kHeadDim, cube2Offsets.vBlockDim, cube2Offsets.blockTokens}; + auto tensorBlockK = GetTile( + tensorK, tla::MakeCoord(0, 0), + tla::MakeShape(cube2Shape.m(), cube2Shape.k())); + auto tensorBlockVwork = GetTile( + tensorVwork, tla::MakeCoord(0, 0), + tla::MakeShape(cube2Shape.k(), cube2Shape.n())); + auto tensorBlockHwork = GetTile( + tensorHwork, tla::MakeCoord(0, 0), + tla::MakeShape(cube2Shape.m(), cube2Shape.n())); + + if (cube2Offsets.blockTokens < chunkSize) { + blockMmadKVTail.preSetFlags(); + blockMmadKVTail( + tensorBlockK, tensorBlockVwork, tensorBlockHwork, + cube2Shape, EmptyClass{}, true); + blockMmadKVTail.finalWaitFlags(); + } else { + blockMmadKV.preSetFlags(); + blockMmadKV( + tensorBlockK, tensorBlockVwork, tensorBlockHwork, + cube2Shape); + blockMmadKV.finalWaitFlags(); + } + AscendC::PipeBarrier(); + } + Arch::CrossCoreSetFlag<0x2, PIPE_FIX>( + cubeBlockScheduler.cube2Done[streamId]); } - const GDNFwdHOffsets& cube2Offsets = cubeBlockScheduler.GetCurTaskOffsets(stream); - Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done[streamId]); - - if (cubeBlockScheduler.NeedProcessStage2(stream)) { - // step 3: h[i+1] = k.T @ v_work - int64_t cube2OffsetK = kGated ? cube2Offsets.kDecayWorkOffset : cube2Offsets.wkOffset; - int64_t cube2OffsetVwork = cube2Offsets.vWorkOffset; - auto tensorK = kGated - ? tla::MakeTensor(gmKDecayWorkspace[cube2OffsetK], kLayout, Catlass::Arch::PositionGM{}) - : tla::MakeTensor(gmK[cube2OffsetK], kLayout, Catlass::Arch::PositionGM{}); - auto vUpdateLayout = tla::MakeLayout(cube2Offsets.blockTokens, cube2Offsets.vBlockDim); - auto tensorVwork = tla::MakeTensor(gmVUpdateWorkspace[cube2OffsetVwork], vUpdateLayout, Catlass::Arch::PositionGM{}); - auto tensorHwork = tla::MakeTensor(gmHWorkspace[cube2Offsets.hWorkOffset], hworkLayout, Catlass::Arch::PositionGM{}); - GemmCoord cube2Shape{kHeadDim, cube2Offsets.vBlockDim, cube2Offsets.blockTokens}; - auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k())); - auto tensorBlockVwork = GetTile(tensorVwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n())); - auto tensorBlockHwork = GetTile(tensorHwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n())); - - blockMmadKV.preSetFlags(); - blockMmadKV(tensorBlockK, tensorBlockVwork, tensorBlockHwork, cube2Shape); - blockMmadKV.finalWaitFlags(); + } else { + blockMmadKV.preSetFlags(); + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = cubeBlockScheduler.GetStreamId(i); + const auto& stream = cubeBlockScheduler.GetStream(i); + if (cubeBlockScheduler.StreamIsDone(stream)) { + continue; + } + const GDNFwdHOffsets& cube2Offsets = cubeBlockScheduler.GetCurTaskOffsets(stream); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done[streamId]); + + if (cubeBlockScheduler.NeedProcessStage2(stream)) { + if (cube2Offsets.blockTokens < 16) { + Arch::CrossCoreSetFlag<0x2, PIPE_MTE2>( + cubeBlockScheduler.cube2Done[streamId]); + continue; + } + // step 3: h[i+1] = k.T @ v_work + int64_t cube2OffsetK = kGated ? cube2Offsets.kDecayWorkOffset : cube2Offsets.wkOffset; + int64_t cube2OffsetVwork = cube2Offsets.vWorkOffset; + auto tensorK = kGated + ? tla::MakeTensor(gmKDecayWorkspace[cube2OffsetK], kLayout, Catlass::Arch::PositionGM{}) + : tla::MakeTensor(gmK[cube2OffsetK], kLayout, Catlass::Arch::PositionGM{}); + auto vUpdateLayout = tla::MakeLayout(cube2Offsets.blockTokens, cube2Offsets.vBlockDim); + auto tensorVwork = tla::MakeTensor(gmVUpdateWorkspace[cube2OffsetVwork], vUpdateLayout, Catlass::Arch::PositionGM{}); + auto tensorHwork = tla::MakeTensor(gmHWorkspace[cube2Offsets.hWorkOffset], hworkLayout, Catlass::Arch::PositionGM{}); + GemmCoord cube2Shape{kHeadDim, cube2Offsets.vBlockDim, cube2Offsets.blockTokens}; + auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k())); + auto tensorBlockVwork = GetTile(tensorVwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n())); + auto tensorBlockHwork = GetTile(tensorHwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n())); + + blockMmadKV(tensorBlockK, tensorBlockVwork, tensorBlockHwork, cube2Shape); + } + Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube2Done[streamId]); } - Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube2Done[streamId]); + blockMmadKV.finalWaitFlags(); } } currStage ^= 0x01; } Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[0]); Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[1]); + if (useDirectFp32Ub) { + for (uint32_t slot = 0; slot < DIRECT_UB_STAGES; ++slot) { + AscendC::CrossCoreWaitFlag<0x4, PIPE_FIX>(DIRECT_UB_FREE_FLAG_BEGIN + slot); + AscendC::CrossCoreWaitFlag<0x4, PIPE_FIX>( + DIRECT_UB_FREE_FLAG_BEGIN + DIRECT_UB_FLAG_STRIDE + slot); + } + } } if ASCEND_IS_AIV { - uint32_t coreIdx = AscendC::GetBlockIdx(); + bool useBoundedMmad = isVariedLen || (seqlen % chunkSize != 0); + uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); + uint32_t subBlockNum = AscendC::GetSubBlockNum(); + uint32_t coreIdx = AscendC::GetBlockIdx() / subBlockNum; uint32_t coreNum = AscendC::GetBlockNum(); - - if (useInitialState) { - AscendC::LocalTensor stateUbTensorPing = resource.ubBuf.template GetBufferByByte(0); - AscendC::LocalTensor stateUbTensorPong = resource.ubBuf.template GetBufferByByte(96 * 1024); - AscendC::LocalTensor hUbTensorPing = resource.ubBuf.template GetBufferByByte(64 * 1024); - AscendC::LocalTensor hUbTensorPong = resource.ubBuf.template GetBufferByByte(160 * 1024); - uint32_t totalChunks = isVariedLen ? vecBlockScheduler.totalChunks : ((seqlen + chunkSize - 1) / chunkSize); - uint32_t transferCount = isVariedLen ? (vecBlockScheduler.tokenBatch * vNumHead / coreNum) : (shapeBatch * vNumHead / coreNum); - uint32_t remainderFlag = isVariedLen ? (((vecBlockScheduler.tokenBatch * vNumHead) % coreNum) != 0): (((shapeBatch * vNumHead) % coreNum) != 0); - uint32_t step = transferCount + remainderFlag; - uint32_t stateBlockSize = kHeadDim * vHeadDim; - uint32_t pingpongFlag = 1; - uint32_t start = coreIdx * step; - uint32_t end = start + step; - uint32_t maxLimit = isVariedLen ? vecBlockScheduler.tokenBatch * vNumHead : shapeBatch * vNumHead; - uint32_t realEnd = min(end, maxLimit); - AscendC::SetFlag(EVENT_ID0); - AscendC::SetFlag(EVENT_ID1); - for (uint32_t initialStateBlockOffset = start; initialStateBlockOffset >= start && initialStateBlockOffset < realEnd; initialStateBlockOffset++) { - uint32_t batchIdx = initialStateBlockOffset / vNumHead; - uint32_t vHeadIdx = initialStateBlockOffset % vNumHead; - uint32_t chunkOffset = isVariedLen ? gmNumChunks.GetValue(batchIdx) : 0; - uint32_t initialStateBaseOffset = initialStateBlockOffset * stateBlockSize; - uint32_t shapeBatchIdx = isVariedLen ? 0 : batchIdx; - uint32_t hBaseOffset = (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset) * stateBlockSize; - if (vHeadDim <= 128) { - AscendC::LocalTensor stateUbTensor = pingpongFlag ? stateUbTensorPing : stateUbTensorPong; - AscendC::LocalTensor hUbTensor = pingpongFlag ? hUbTensorPing : hUbTensorPong; - auto event_id = pingpongFlag ? EVENT_ID1 : EVENT_ID0; - AscendC::WaitFlag(event_id); - if constexpr(!std::is_same::value) { - AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateBaseOffset], stateBlockSize); - AscendC::SetFlag(event_id); - AscendC::WaitFlag(event_id); - AscendC::Cast(hUbTensor, stateUbTensor, AscendC::RoundMode::CAST_RINT, stateBlockSize); - AscendC::SetFlag(event_id); - AscendC::WaitFlag(event_id); - AscendC::DataCopy(gmH[hBaseOffset], hUbTensor, stateBlockSize); - } else { - AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateBaseOffset], stateBlockSize); - AscendC::SetFlag(event_id); - AscendC::WaitFlag(event_id); - AscendC::DataCopy(gmH[hBaseOffset], stateUbTensor, stateBlockSize); - } - AscendC::SetFlag(event_id); - pingpongFlag = 1 - pingpongFlag; - } else { - uint32_t stateRowTile = (32 * 1024) / (vHeadDim * sizeof(ElementH)); - for (uint32_t rowOffset = 0; rowOffset < kHeadDim; rowOffset += stateRowTile) { - uint32_t rowsThisTile = Min(stateRowTile, kHeadDim - rowOffset); - uint32_t stateTileElems = rowsThisTile * vHeadDim; - uint32_t initialStateOffset = initialStateBaseOffset + rowOffset * vHeadDim; - uint32_t hOffset = hBaseOffset + rowOffset * vHeadDim; - AscendC::LocalTensor stateUbTensor = pingpongFlag ? stateUbTensorPing : stateUbTensorPong; - AscendC::LocalTensor hUbTensor = pingpongFlag ? hUbTensorPing : hUbTensorPong; - auto event_id = pingpongFlag ? EVENT_ID1 : EVENT_ID0; - AscendC::WaitFlag(event_id); - if constexpr(!std::is_same::value) { - AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateOffset], stateTileElems); - AscendC::SetFlag(event_id); - AscendC::WaitFlag(event_id); - AscendC::Cast(hUbTensor, stateUbTensor, AscendC::RoundMode::CAST_RINT, stateTileElems); - AscendC::SetFlag(event_id); - AscendC::WaitFlag(event_id); - AscendC::DataCopy(gmH[hOffset], hUbTensor, stateTileElems); - } else { - AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateOffset], stateTileElems); - AscendC::SetFlag(event_id); - AscendC::WaitFlag(event_id); - AscendC::DataCopy(gmH[hOffset], stateUbTensor, stateTileElems); - } - AscendC::SetFlag(event_id); - pingpongFlag = 1 - pingpongFlag; - } - } - } - - AscendC::WaitFlag(EVENT_ID0); - AscendC::WaitFlag(EVENT_ID1); - } else { - uint32_t stateRowTile = (32 * 1024) / (vHeadDim * sizeof(ElementH)); - AscendC::LocalTensor chunkOffsetsUb = - resource.ubBuf.template GetBufferByByte(0); - if (isVariedLen) { - uint32_t chunkOffsetBytes = (vecBlockScheduler.tokenBatch + 1) * sizeof(int64_t); - AscendC::DataCopyParams copyParams{ - 1, static_cast(chunkOffsetBytes), 0, 0}; - AscendC::DataCopyPadParams padParams{false, 0, 0, 0}; - AscendC::DataCopyPad(chunkOffsetsUb, gmNumChunks[0], copyParams, padParams); - AscendC::SetFlag(EVENT_ID3); - AscendC::WaitFlag(EVENT_ID3); - AscendC::SetFlag(EVENT_ID3); - AscendC::WaitFlag(EVENT_ID3); - } - auto chunkOffsets = reinterpret_cast<__ubuf__ int64_t *>(chunkOffsetsUb.GetPhyAddr()); - AscendC::LocalTensor hUbTensorPing = - resource.ubBuf.template GetBufferByByte(64 * 1024); - AscendC::LocalTensor hUbTensorPong = - resource.ubBuf.template GetBufferByByte(160 * 1024); - uint32_t totalChunks = - isVariedLen ? vecBlockScheduler.totalChunks : ((seqlen + chunkSize - 1) / chunkSize); - uint32_t taskCount = - (isVariedLen ? vecBlockScheduler.tokenBatch : shapeBatch) * vNumHead; - uint32_t step = taskCount / coreNum + ((taskCount % coreNum) != 0); - uint32_t start = coreIdx * step; - uint32_t realEnd = Min(start + step, taskCount); - uint32_t pingpongFlag = 1; - AscendC::SetFlag(EVENT_ID0); - AscendC::SetFlag(EVENT_ID1); - for (uint32_t taskIdx = start; taskIdx < realEnd; ++taskIdx) { + uint32_t taskCount = + (isVariedLen ? vecBlockScheduler.tokenBatch : shapeBatch) * vNumHead; + uint32_t tasksPerCore = taskCount > coreNum ? PING_PONG_STAGES : 1; + uint32_t taskStride = coreNum * tasksPerCore; + uint32_t rowsPerSubBlock = (kHeadDim + subBlockNum - 1) / subBlockNum; + uint32_t rowBegin = subBlockIdx * rowsPerSubBlock; + uint32_t rowEnd = Min(rowBegin + rowsPerSubBlock, kHeadDim); + uint32_t hRowsPerTile = (32 * 1024) / (vHeadDim * sizeof(ElementH)); + uint32_t stateRowsPerTile = + (64 * 1024) / (vHeadDim * sizeof(ElementInitialState)); + uint32_t rowsPerTile = Min(hRowsPerTile, stateRowsPerTile); + uint32_t totalChunks = + isVariedLen ? vecBlockScheduler.totalChunks : ((seqlen + chunkSize - 1) / chunkSize); + uint32_t stateBlockSize = kHeadDim * vHeadDim; + uint32_t pingpongFlag = 1; + AscendC::LocalTensor stateUbTensorPing = + resource.ubBuf.template GetBufferByByte(0); + AscendC::LocalTensor stateUbTensorPong = + resource.ubBuf.template GetBufferByByte(96 * 1024); + AscendC::LocalTensor hUbTensorPing = + resource.ubBuf.template GetBufferByByte(64 * 1024); + AscendC::LocalTensor hUbTensorPong = + resource.ubBuf.template GetBufferByByte(160 * 1024); + AscendC::SetFlag(EVENT_ID0); + AscendC::SetFlag(EVENT_ID1); + for (uint32_t slot = 0; slot < tasksPerCore; ++slot) { + for (uint32_t taskIdx = coreIdx * tasksPerCore + slot; + taskIdx < taskCount; taskIdx += taskStride) { uint32_t batchIdx = taskIdx / vNumHead; uint32_t vHeadIdx = taskIdx % vNumHead; - uint32_t chunkOffset = isVariedLen ? static_cast(chunkOffsets[batchIdx]) : 0; + uint32_t chunkOffset = + isVariedLen ? vecBlockScheduler.GetVarlenChunkOffset(batchIdx) : 0; uint32_t shapeBatchIdx = isVariedLen ? 0 : batchIdx; uint32_t hBaseOffset = (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset) * - kHeadDim * vHeadDim; - for (uint32_t rowOffset = 0; rowOffset < kHeadDim; rowOffset += stateRowTile) { - uint32_t rowsThisTile = Min(stateRowTile, kHeadDim - rowOffset); + stateBlockSize; + uint32_t initialStateBaseOffset = taskIdx * stateBlockSize; + for (uint32_t rowOffset = rowBegin; rowOffset < rowEnd; rowOffset += rowsPerTile) { + uint32_t rowsThisTile = Min(rowsPerTile, rowEnd - rowOffset); uint32_t stateTileElems = rowsThisTile * vHeadDim; + uint32_t hOffset = hBaseOffset + rowOffset * vHeadDim; + AscendC::LocalTensor stateUbTensor = + pingpongFlag ? stateUbTensorPing : stateUbTensorPong; AscendC::LocalTensor hUbTensor = pingpongFlag ? hUbTensorPing : hUbTensorPong; auto eventId = pingpongFlag ? EVENT_ID1 : EVENT_ID0; AscendC::WaitFlag(eventId); - AscendC::Duplicate(hUbTensor, static_cast(0), stateTileElems); - AscendC::SetFlag(eventId); - AscendC::WaitFlag(eventId); - AscendC::DataCopy(gmH[hBaseOffset + rowOffset * vHeadDim], hUbTensor, stateTileElems); + if (useInitialState) { + uint32_t initialStateOffset = + initialStateBaseOffset + rowOffset * vHeadDim; + if constexpr (!std::is_same::value) { + AscendC::DataCopy( + stateUbTensor, gmInitialState[initialStateOffset], stateTileElems); + AscendC::SetFlag(eventId); + AscendC::WaitFlag(eventId); + AscendC::Cast( + hUbTensor, stateUbTensor, AscendC::RoundMode::CAST_RINT, + stateTileElems); + AscendC::SetFlag(eventId); + AscendC::WaitFlag(eventId); + AscendC::DataCopy(gmH[hOffset], hUbTensor, stateTileElems); + } else { + AscendC::DataCopy( + stateUbTensor, gmInitialState[initialStateOffset], stateTileElems); + AscendC::SetFlag(eventId); + AscendC::WaitFlag(eventId); + AscendC::DataCopy(gmH[hOffset], stateUbTensor, stateTileElems); + } + } else { + AscendC::Duplicate(hUbTensor, static_cast(0), stateTileElems); + AscendC::SetFlag(eventId); + AscendC::WaitFlag(eventId); + AscendC::DataCopy(gmH[hOffset], hUbTensor, stateTileElems); + } AscendC::SetFlag(eventId); pingpongFlag = 1 - pingpongFlag; } } - AscendC::WaitFlag(EVENT_ID0); - AscendC::WaitFlag(EVENT_ID1); } + AscendC::WaitFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID1); AscendC::SyncAll(); + if (useDirectFp32Ub) { + for (uint32_t slot = 0; slot < DIRECT_UB_STAGES; ++slot) { + AscendC::CrossCoreSetFlag<0x4, PIPE_V>(DIRECT_UB_FREE_FLAG_BEGIN + slot); + } + } Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done[0]); Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done[1]); @@ -500,6 +928,10 @@ class GDNFwdHKernel { AscendC::SetFlag(EVENT_ID1 + pongBaseEvent); AscendC::SetFlag(EVENT_ID3); // preset g AscendC::SetFlag(EVENT_ID3 + pongBaseEvent); + AscendC::SetFlag(EVENT_ID0); // preset h_update + AscendC::SetFlag(EVENT_ID0 + pongBaseEvent); + AscendC::SetFlag(EVENT_ID2); // preset h + AscendC::SetFlag(EVENT_ID2 + pongBaseEvent); uint32_t currStage = 0; // 0: V1, 1: V2 bool event0FromMte3[PING_PONG_STAGES] = {false, false}; bool event2FromMte3[PING_PONG_STAGES] = {!(storeFinalState && std::is_same::value), @@ -521,8 +953,17 @@ class GDNFwdHKernel { } const GDNFwdHOffsets& vec1Offsets = vecBlockScheduler.GetCurTaskOffsets(stream); AscendC::LocalTensor l1VUpdate = (i == 0) ? l1VUpdatePing : l1VUpdatePong; + bool tailVectorPath = + vec1Offsets.blockTokens < 16 && !useBoundedMmad; + if (tailVectorPath) { + Arch::CrossCoreWaitFlag( + vecBlockScheduler.cube1Done[streamId]); + ComputeTailVWorkspace( + vec1Offsets, EVENT_ID3 + (i == 0 ? 0 : pongBaseEvent)); + } bool waitWsFromMte3 = storeFinalState && std::is_same::value && event0FromMte3[streamId]; + bool useDirectForTask = useDirectFp32Ub && !tailVectorPath; epilogueGDNFwdHVnew( gmV[vec1Offsets.uvOffset], gmVUpdateWorkspace[vec1Offsets.vWorkOffset], l1VUpdate, gmG[vec1Offsets.gOffset], gmU[vec1Offsets.uvOffset], gmVWorkspace[vec1Offsets.vWorkOffset], @@ -530,7 +971,8 @@ class GDNFwdHKernel { vec1Offsets.blockTokens, kHeadDim, vec1Offsets.vBlockDim, vHeadDim, vecBlockScheduler.cube1Done[streamId], vecBlockScheduler.vec1Done[streamId], vec1Offsets.isInitialState, vec1Offsets.isFinalState, storeFinalState, - waitWsFromMte3, (i == 0) + waitWsFromMte3, (i == 0), tailVectorPath, useDirectForTask, + DIRECT_UB_FREE_FLAG_BEGIN, DIRECT_UB_READY_FLAG_BEGIN ); if (storeFinalState && std::is_same::value) { event0FromMte3[streamId] = false; @@ -546,8 +988,18 @@ class GDNFwdHKernel { } const GDNFwdHOffsets& vec2Offsets = vecBlockScheduler.GetCurTaskOffsets(stream); if (vecBlockScheduler.NeedProcessStage2(stream)) { + bool tailVectorPath = + vec2Offsets.blockTokens < 16 && !useBoundedMmad; + if (tailVectorPath) { + Arch::CrossCoreWaitFlag( + vecBlockScheduler.cube2Done[streamId]); + ComputeTailHWorkspace( + vec2Offsets, EVENT_ID3 + (i == 0 ? 0 : pongBaseEvent)); + } if (storeFinalState && std::is_same::value) { - event0FromMte3[streamId] = vec2Offsets.isFinalState; + // Update always writes the FP32 state through MTE3, + // including intermediate chunks. + event0FromMte3[streamId] = true; event2FromMte3[streamId] = !vec2Offsets.isFinalState; } // step 4: h[i+1] += h_work if i < num_chunks - 1 else None @@ -557,11 +1009,17 @@ class GDNFwdHKernel { gmH[vec2Offsets.hSrcOffset], gmHWorkspace[vec2Offsets.hWorkOffset], gmGk[vec2Offsets.gkOffset], + gmInitialState[vec2Offsets.initialStateOffset], vec2Offsets.blockTokens, kHeadDim, vec2Offsets.vBlockDim, vHeadDim, vecBlockScheduler.cube2Done[streamId], - vec2Offsets.isInitialState, vec2Offsets.isFinalState, storeFinalState, (i == 0) + vec2Offsets.isInitialState, vec2Offsets.isFinalState, storeFinalState, + useInitialState, (i == 0), tailVectorPath, + useDirectFp32Ub && !tailVectorPath, + DIRECT_UB_FREE_FLAG_BEGIN, DIRECT_UB_READY_FLAG_BEGIN ); } else { - Arch::CrossCoreWaitFlag(vecBlockScheduler.cube2Done[streamId]); + if (!useDirectFp32Ub) { + Arch::CrossCoreWaitFlag(vecBlockScheduler.cube2Done[streamId]); + } } Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done[streamId]); } @@ -600,6 +1058,10 @@ class GDNFwdHKernel { AscendC::WaitFlag(EVENT_ID1 + pongBaseEvent); AscendC::WaitFlag(EVENT_ID3); // preset g AscendC::WaitFlag(EVENT_ID3 + pongBaseEvent); + AscendC::WaitFlag(EVENT_ID0); // drain h_update + AscendC::WaitFlag(EVENT_ID0 + pongBaseEvent); + AscendC::WaitFlag(EVENT_ID2); // drain h + AscendC::WaitFlag(EVENT_ID2 + pongBaseEvent); } } diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/epilogue/block/block_epilogue_gdn_fwdh_update.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/epilogue/block/block_epilogue_gdn_fwdh_update.hpp new file mode 100644 index 000000000000..973b02e361ff --- /dev/null +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/epilogue/block/block_epilogue_gdn_fwdh_update.hpp @@ -0,0 +1,571 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#ifndef CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_UPDATE_HPP +#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_UPDATE_HPP +#include "catlass/catlass.hpp" +#include "catlass/arch/resource.hpp" +#include "../gdn_fwd_h_epilogue_policies.hpp" +#include "catlass/gemm_coord.hpp" +#include "catlass/matrix_coord.hpp" +#include "catlass/epilogue/tile/tile_copy.hpp" + +namespace Catlass::Epilogue::Block { + +template < + class HOutputType_, + class GInputType_, + class HInputType_, + class HUpdateInputType_, + class FinalStateType_, + class KGatedTag +> +class BlockEpilogue < + EpilogueAtlasGDNFwdHUpdate, + HOutputType_, + GInputType_, + HInputType_, + HUpdateInputType_, + FinalStateType_, + KGatedTag +> { + static constexpr bool kGated = KGatedTag::value; + static constexpr bool scalarGated = KGatedTag::scalarGated; + static constexpr bool useExp2 = KGatedTag::useExp2; + static constexpr float LN2 = 0.6931471805599453f; +public: + using DispatchPolicy = EpilogueAtlasGDNFwdHUpdate; + using ArchTag = typename DispatchPolicy::ArchTag; + + using HElementOutput = typename HOutputType_::Element; + using GElementInput = typename GInputType_::Element; + using HElementInput = typename HInputType_::Element; + using HUpdateElementInput = typename HUpdateInputType_::Element; + using FinalStateElement = typename FinalStateType_::Element; + + CATLASS_DEVICE + BlockEpilogue(Arch::Resource &resource) + { + + constexpr uint32_t CALC_BUF_OFFSET = 0; + constexpr uint32_t PING_BUF_0_OFFSET = 32 * 1024; + constexpr uint32_t PING_BUF_1_OFFSET = 48 * 1024; + constexpr uint32_t PING_BUF_2_OFFSET = 64 * 1024; + constexpr uint32_t PING_BUF_3_OFFSET = 80 * 1024; + constexpr uint32_t PONG_BUF_0_OFFSET = 96 * 1024; + constexpr uint32_t PONG_BUF_1_OFFSET = 112 * 1024; + constexpr uint32_t PONG_BUF_2_OFFSET = 128 * 1024; + constexpr uint32_t PONG_BUF_3_OFFSET = 144 * 1024; + constexpr uint32_t PING_G_BUF_OFFSET = 160 * 1024; + constexpr uint32_t PONG_G_BUF_OFFSET = 161 * 1024; + constexpr uint32_t PING_G_SUB_BUF_OFFSET = 162 * 1024; + constexpr uint32_t PONG_G_SUB_BUF_OFFSET = 163 * 1024; + constexpr uint32_t PING_G_INPUT_BUF_OFFSET = 164 * 1024; + constexpr uint32_t PONG_G_INPUT_BUF_OFFSET = 165 * 1024; + constexpr uint32_t SHARE_BUF_OFFSET = 166 * 1024; + + + calcUbTensor = resource.ubBuf.template GetBufferByByte(CALC_BUF_OFFSET); + + hUpdateUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_0_OFFSET); + gkBroadcastUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_1_OFFSET); + hUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_3_OFFSET); + finalOutputUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_3_OFFSET); + glastUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_INPUT_BUF_OFFSET); + + hUpdateUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_0_OFFSET); + gkBroadcastUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_1_OFFSET); + hUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_3_OFFSET); + finalOutputUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_3_OFFSET); + glastUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_INPUT_BUF_OFFSET); + + if constexpr (kGated) { + gkLastUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_SUB_BUF_OFFSET); + gkLastUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_SUB_BUF_OFFSET); + gkInputUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_INPUT_BUF_OFFSET); + gkInputUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_INPUT_BUF_OFFSET); + shareBufferGk_ = resource.ubBuf.template GetBufferByByte(SHARE_BUF_OFFSET); + } + } + + CATLASS_DEVICE + ~BlockEpilogue() {} + + template + CATLASS_DEVICE + void CopyGmToUb( + AscendC::LocalTensor dst, + AscendC::GlobalTensor src, + uint32_t rows, + uint32_t cols, + uint32_t srcStride) + { + if (cols == srcStride) { + AscendC::DataCopy(dst, src, rows * cols); + return; + } + AscendC::DataCopyExtParams copyParams{ + static_cast(rows), + static_cast(cols * sizeof(Element)), + static_cast((srcStride - cols) * sizeof(Element)), + 0, + 0}; + AscendC::DataCopyPadExtParams padParams{false, 0, 0, 0}; + AscendC::DataCopyPad(dst, src, copyParams, padParams); + } + + template + CATLASS_DEVICE + void CopyUbToGm( + AscendC::GlobalTensor dst, + AscendC::LocalTensor src, + uint32_t rows, + uint32_t cols, + uint32_t dstStride) + { + if (cols == dstStride) { + AscendC::DataCopy(dst, src, rows * cols); + return; + } + AscendC::DataCopyExtParams copyParams{ + static_cast(rows), + static_cast(cols * sizeof(Element)), + 0, + static_cast((dstStride - cols) * sizeof(Element)), + 0}; + AscendC::DataCopyPad(dst, src, copyParams); + } + + CATLASS_DEVICE + void operator()( + AscendC::GlobalTensor hOutput, + AscendC::GlobalTensor finalState, + AscendC::GlobalTensor gInput, + AscendC::GlobalTensor hInput, + AscendC::GlobalTensor hUpdateInput, + AscendC::GlobalTensor gkInput, + AscendC::GlobalTensor initialState, + uint32_t chunkSize, + uint32_t kHeadDim, + uint32_t vBlockDim, + uint32_t vHeadDim, + Arch::CrossCoreFlag cube2Done, + bool isInitialState, + bool isFinalState, + bool storeFinalState, + bool useInitialState, + bool isPing, + bool cube2AlreadyWaited + ) + { + static constexpr uint32_t ROW_TILE = 16; + uint32_t mActual = kHeadDim; + uint32_t nActual = vBlockDim; + uint32_t outputStride = vHeadDim; + uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); + uint32_t subBlockNum = AscendC::GetSubBlockNum(); + uint32_t rowsPerSubBlock = CeilDiv(mActual, subBlockNum); + uint32_t rowBegin = subBlockIdx * rowsPerSubBlock; + uint32_t rowEnd = rowBegin + rowsPerSubBlock; + if (rowEnd > mActual) { + rowEnd = mActual; + } + if (rowBegin >= mActual) { + if (!cube2AlreadyWaited) { + Arch::CrossCoreWaitFlag(cube2Done); + } + return; + } + + AscendC::ResetMask(); + + AscendC::GlobalTensor gInputThisSubBlock = gInput; + + uint32_t pingpongFlag = isPing ? 0 : pongBaseEvent; + AscendC::LocalTensor hUpdateUbTensor = isPing ? hUpdateUbTensor_ping : hUpdateUbTensor_pong; + AscendC::LocalTensor gkBroadcastUbTensor = + isPing ? gkBroadcastUbTensor_ping : gkBroadcastUbTensor_pong; + AscendC::LocalTensor hUbTensor = isPing ? hUbTensor_ping : hUbTensor_pong; + AscendC::LocalTensor finalOutputUbTensor = isPing ? finalOutputUbTensor_ping : finalOutputUbTensor_pong; + AscendC::LocalTensor glastUbTensor = isPing ? glastUbTensor_ping : glastUbTensor_pong; + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + bool useFp32StateUpdate = storeFinalState && std::is_same::value && + (!isInitialState || useInitialState); + + float muls = 1.0f; + if constexpr (scalarGated) { + GElementInput gLastVal = gInputThisSubBlock.GetValue(chunkSize-1); + float gLastFloat = 0.0f; + if constexpr(std::is_same::value) { + gLastFloat = gLastVal; + } else if constexpr(std::is_same::value) { + gLastFloat = (float)gLastVal; + } else if constexpr(std::is_same::value) { + gLastFloat = AscendC::ToFloat(gLastVal); + } + glastUbTensor.SetValue(0, gLastFloat); + + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + if constexpr (useExp2) { + AscendC::Muls(glastUbTensor, glastUbTensor, LN2, 1); + AscendC::PipeBarrier(); + } + AscendC::Exp(glastUbTensor, glastUbTensor, 1); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + muls = glastUbTensor.GetValue(0); + } + if constexpr (kGated) { + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } + if constexpr (scalarGated) { + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + } + + if (nActual <= 128 && nActual == outputStride) { + uint32_t mActualThisSubBlock = rowEnd - rowBegin; + AscendC::GlobalTensor hOutputThisSubBlock = hOutput[rowBegin * outputStride]; + AscendC::GlobalTensor hInputThisSubBlock = hInput[rowBegin * outputStride]; + AscendC::GlobalTensor hUpdateInputThisSubBlock = hUpdateInput[rowBegin * nActual]; + AscendC::GlobalTensor finalStateThisSubBlock = finalState[rowBegin * outputStride]; + + if (storeFinalState && isInitialState && std::is_same::value) { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + } + if constexpr (std::is_same::value) { + if (useFp32StateUpdate) { + if (isInitialState) { + AscendC::DataCopy(calcUbTensor, initialState[rowBegin * outputStride], + mActualThisSubBlock * nActual); + } else { + AscendC::DataCopy(calcUbTensor, finalStateThisSubBlock, + mActualThisSubBlock * nActual); + } + } else { + AscendC::DataCopy(hUbTensor, hInputThisSubBlock, + mActualThisSubBlock * nActual); + } + } else { + AscendC::DataCopy(hUbTensor, hInputThisSubBlock, + mActualThisSubBlock * nActual); + } + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + + if (!useFp32StateUpdate) { + AscendC::Cast(calcUbTensor, hUbTensor, AscendC::RoundMode::CAST_NONE, + mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + } + AscendC::Muls(calcUbTensor, calcUbTensor, muls, mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + + if constexpr (kGated) { + AscendC::GlobalTensor gkLastInput = gkInput[(chunkSize - 1) * kHeadDim + rowBegin]; + AscendC::LocalTensor gkLastUbTensor = isPing ? gkLastUbTensor_ping : gkLastUbTensor_pong; + AscendC::LocalTensor gkInputUbTensor = isPing ? gkInputUbTensor_ping : gkInputUbTensor_pong; + + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + if constexpr(std::is_same::value) { + AscendC::DataCopy(gkLastUbTensor, gkLastInput, mActualThisSubBlock); + } else { + AscendC::DataCopy(gkInputUbTensor, gkLastInput, mActualThisSubBlock); + } + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + if constexpr(!std::is_same::value) { + AscendC::Cast(gkLastUbTensor, gkInputUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock); + } + AscendC::PipeBarrier(); + AscendC::Muls(gkLastUbTensor, gkLastUbTensor, LN2, mActualThisSubBlock); + AscendC::PipeBarrier(); + AscendC::Exp(gkLastUbTensor, gkLastUbTensor, mActualThisSubBlock); + AscendC::PipeBarrier(); + + uint32_t gkBrcReptime = (mActualThisSubBlock + 8 - 1) / 8; + uint32_t dstShapeGk[2] = {gkBrcReptime * 8, nActual}; + uint32_t srcShapeGk[2] = {gkBrcReptime * 8, 1}; + AscendC::Broadcast(hUpdateUbTensor, gkLastUbTensor, + dstShapeGk, srcShapeGk, shareBufferGk_); + AscendC::PipeBarrier(); + AscendC::Mul(calcUbTensor, calcUbTensor, hUpdateUbTensor, + mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + } + + if (!cube2AlreadyWaited) { + Arch::CrossCoreWaitFlag(cube2Done); + } + + if constexpr (kGated) { + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + } + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::DataCopy(hUpdateUbTensor, hUpdateInputThisSubBlock, mActualThisSubBlock * nActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::Add(hUpdateUbTensor, calcUbTensor, hUpdateUbTensor, mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + if (storeFinalState && isFinalState && std::is_same::value) { + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } + + if constexpr(std::is_same::value) { + if (storeFinalState) { + if (!isFinalState) { + AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, + mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + } + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::DataCopy(finalStateThisSubBlock, hUpdateUbTensor, mActualThisSubBlock * nActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + if (!isFinalState) { + AscendC::DataCopy(hOutputThisSubBlock, hUbTensor, mActualThisSubBlock * nActual); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } + } else { + AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + AscendC::DataCopy(hOutputThisSubBlock, hUbTensor, mActualThisSubBlock * nActual); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } + } else { + if (storeFinalState && isFinalState) { + AscendC::Cast(finalOutputUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + AscendC::DataCopy(finalStateThisSubBlock, finalOutputUbTensor, mActualThisSubBlock * nActual); + } else { + AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + AscendC::DataCopy(hOutputThisSubBlock, hUbTensor, mActualThisSubBlock * nActual); + } + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + return; + } + + if (!cube2AlreadyWaited) { + Arch::CrossCoreWaitFlag(cube2Done); + } + + bool waitHFromV = storeFinalState && isInitialState && std::is_same::value; + bool waitUpdateFromMte3 = false; + uint32_t updateReadyEvent = EVENT_ID3 + pingpongFlag; + bool waitGkFromScalar = true; + for (uint32_t rowStart = rowBegin; rowStart < rowEnd; rowStart += ROW_TILE) { + uint32_t rowsThisTile = rowEnd - rowStart; + if (rowsThisTile > ROW_TILE) { + rowsThisTile = ROW_TILE; + } + + AscendC::GlobalTensor hOutputThisTile = hOutput[rowStart * outputStride]; + AscendC::GlobalTensor hInputThisTile = hInput[rowStart * outputStride]; + AscendC::GlobalTensor hUpdateInputThisTile = hUpdateInput[rowStart * nActual]; + AscendC::GlobalTensor finalStateThisTile = finalState[rowStart * outputStride]; + + if (waitHFromV) { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + } + if constexpr (std::is_same::value) { + if (useFp32StateUpdate) { + if (isInitialState) { + CopyGmToUb(calcUbTensor, initialState[rowStart * outputStride], + rowsThisTile, nActual, outputStride); + } else { + CopyGmToUb(calcUbTensor, finalStateThisTile, rowsThisTile, nActual, + outputStride); + } + } else { + CopyGmToUb(hUbTensor, hInputThisTile, rowsThisTile, nActual, + outputStride); + } + } else { + CopyGmToUb(hUbTensor, hInputThisTile, rowsThisTile, nActual, outputStride); + } + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + + if (!useFp32StateUpdate) { + AscendC::Cast(calcUbTensor, hUbTensor, AscendC::RoundMode::CAST_NONE, rowsThisTile * nActual); + AscendC::PipeBarrier(); + } + AscendC::Muls(calcUbTensor, calcUbTensor, muls, rowsThisTile * nActual); + AscendC::PipeBarrier(); + + if constexpr (kGated) { + AscendC::GlobalTensor gkLastInput = gkInput[(chunkSize - 1) * kHeadDim + rowStart]; + AscendC::LocalTensor gkLastUbTensor = isPing ? gkLastUbTensor_ping : gkLastUbTensor_pong; + AscendC::LocalTensor gkInputUbTensor = isPing ? gkInputUbTensor_ping : gkInputUbTensor_pong; + + if (waitGkFromScalar) { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + waitGkFromScalar = false; + } else { + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + } + if constexpr(std::is_same::value) { + AscendC::DataCopy(gkLastUbTensor, gkLastInput, rowsThisTile); + } else { + AscendC::DataCopy(gkInputUbTensor, gkLastInput, rowsThisTile); + } + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + if constexpr(!std::is_same::value) { + AscendC::Cast(gkLastUbTensor, gkInputUbTensor, AscendC::RoundMode::CAST_NONE, rowsThisTile); + } + AscendC::PipeBarrier(); + AscendC::Muls(gkLastUbTensor, gkLastUbTensor, LN2, rowsThisTile); + AscendC::PipeBarrier(); + AscendC::Exp(gkLastUbTensor, gkLastUbTensor, rowsThisTile); + AscendC::PipeBarrier(); + + uint32_t gkBrcReptime = (rowsThisTile + 8 - 1) / 8; + uint32_t dstShapeGk[2] = {gkBrcReptime * 8, nActual}; + uint32_t srcShapeGk[2] = {gkBrcReptime * 8, 1}; + AscendC::Broadcast(gkBroadcastUbTensor, gkLastUbTensor, + dstShapeGk, srcShapeGk, shareBufferGk_); + AscendC::PipeBarrier(); + AscendC::Mul(calcUbTensor, calcUbTensor, gkBroadcastUbTensor, rowsThisTile * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + } + + if (waitUpdateFromMte3) { + AscendC::WaitFlag(updateReadyEvent); + } else { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } + AscendC::DataCopy(hUpdateUbTensor, hUpdateInputThisTile, rowsThisTile * nActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + AscendC::Add(hUpdateUbTensor, calcUbTensor, hUpdateUbTensor, rowsThisTile * nActual); + AscendC::PipeBarrier(); + if (storeFinalState && isFinalState && std::is_same::value) { + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + waitHFromV = true; + } else { + waitHFromV = false; + } + + if constexpr(std::is_same::value) { + if (storeFinalState) { + if (!isFinalState) { + AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, + rowsThisTile * nActual); + AscendC::PipeBarrier(); + } + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + CopyUbToGm(finalStateThisTile, hUpdateUbTensor, rowsThisTile, nActual, outputStride); + // A2 wide-V tiles reuse this UB in the next iteration; retire MTE3 before either writer advances. + AscendC::PipeBarrier(); + AscendC::SetFlag(updateReadyEvent); + AscendC::SetFlag(updateReadyEvent); + AscendC::WaitFlag(updateReadyEvent); + waitUpdateFromMte3 = true; + if (!isFinalState) { + CopyUbToGm(hOutputThisTile, hUbTensor, rowsThisTile, nActual, outputStride); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } + } else { + AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + CopyUbToGm(hOutputThisTile, hUbTensor, rowsThisTile, nActual, outputStride); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + waitUpdateFromMte3 = false; + } + } else { + if (storeFinalState && isFinalState) { + AscendC::Cast(finalOutputUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + CopyUbToGm(finalStateThisTile, finalOutputUbTensor, rowsThisTile, nActual, outputStride); + } else { + AscendC::Cast(hUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + CopyUbToGm(hOutputThisTile, hUbTensor, rowsThisTile, nActual, outputStride); + } + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + waitUpdateFromMte3 = false; + } + } + if (storeFinalState && std::is_same::value) { + AscendC::WaitFlag(updateReadyEvent); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + if (!isFinalState) { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } + } else { + AscendC::WaitFlag(EVENT_ID2 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + } + if constexpr (kGated) { + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + } + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::SetFlag(EVENT_ID2 + pingpongFlag); + + } + +private: + uint32_t pongBaseEvent = 4; + + AscendC::LocalTensor calcUbTensor; + + AscendC::LocalTensor hUpdateUbTensor_ping; + AscendC::LocalTensor gkBroadcastUbTensor_ping; + AscendC::LocalTensor hUbTensor_ping; + AscendC::LocalTensor finalOutputUbTensor_ping; + AscendC::LocalTensor glastUbTensor_ping; + + AscendC::LocalTensor hUpdateUbTensor_pong; + AscendC::LocalTensor gkBroadcastUbTensor_pong; + AscendC::LocalTensor hUbTensor_pong; + AscendC::LocalTensor finalOutputUbTensor_pong; + AscendC::LocalTensor glastUbTensor_pong; + + AscendC::LocalTensor gkLastUbTensor_ping; + AscendC::LocalTensor gkLastUbTensor_pong; + AscendC::LocalTensor gkInputUbTensor_ping; + AscendC::LocalTensor gkInputUbTensor_pong; + AscendC::LocalTensor shareBufferGk_; +}; +} + +#endif diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp new file mode 100644 index 000000000000..2883d100fb0d --- /dev/null +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp @@ -0,0 +1,480 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#ifndef CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_VNEW_HPP +#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_VNEW_HPP +#include "catlass/catlass.hpp" +#include "catlass/arch/resource.hpp" +#include "../gdn_fwd_h_epilogue_policies.hpp" +#include "catlass/gemm_coord.hpp" +#include "catlass/matrix_coord.hpp" +#include "catlass/epilogue/tile/tile_copy.hpp" + + + +namespace Catlass::Epilogue::Block { + +template < + class VOutputType_, + class GInputType_, + class UInputType_, + class WSInputType_, + class FinalStateType_, + class KGatedTag +> +class BlockEpilogue < + EpilogueAtlasGDNFwdHVnew, + VOutputType_, + GInputType_, + UInputType_, + WSInputType_, + FinalStateType_, + KGatedTag +> { + static constexpr bool kGated = KGatedTag::value; + static constexpr bool scalarGated = KGatedTag::scalarGated; + static constexpr bool useExp2 = KGatedTag::useExp2; + static constexpr float LN2 = 0.6931471805599453f; +public: + // Type aliases + using DispatchPolicy = EpilogueAtlasGDNFwdHVnew; + using ArchTag = typename DispatchPolicy::ArchTag; + + using VElementOutput = typename VOutputType_::Element; + using GElementInput = typename GInputType_::Element; + using UElementInput = typename UInputType_::Element; + using WSElementInput = typename WSInputType_::Element; + using FinalStateElement = typename FinalStateType_::Element; + + CATLASS_DEVICE + BlockEpilogue(Arch::Resource &resource) + { + + constexpr uint32_t CALC_BUF_OFFSET = 0; + constexpr uint32_t PING_BUF_0_OFFSET = 32 * 1024; + constexpr uint32_t PING_BUF_1_OFFSET = 48 * 1024; + constexpr uint32_t PING_BUF_2_OFFSET = 64 * 1024; + constexpr uint32_t PING_BUF_3_OFFSET = 80 * 1024; + constexpr uint32_t PONG_BUF_0_OFFSET = 96 * 1024; + constexpr uint32_t PONG_BUF_1_OFFSET = 112 * 1024; + constexpr uint32_t PONG_BUF_2_OFFSET = 128 * 1024; + constexpr uint32_t PONG_BUF_3_OFFSET = 144 * 1024; + constexpr uint32_t PING_G_BUF_OFFSET = 160 * 1024; + constexpr uint32_t PONG_G_BUF_OFFSET = 161 * 1024; + constexpr uint32_t PING_G_SUB_BUF_OFFSET = 162 * 1024; + constexpr uint32_t PONG_G_SUB_BUF_OFFSET = 163 * 1024; + constexpr uint32_t PING_G_INPUT_BUF_OFFSET = 164 * 1024; + constexpr uint32_t PONG_G_INPUT_BUF_OFFSET = 165 * 1024; + constexpr uint32_t SHARE_BUF_OFFSET = 166 * 1024; + + calcUbTensor = resource.ubBuf.template GetBufferByByte(CALC_BUF_OFFSET); + + uUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_2_OFFSET); + wsUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_0_OFFSET); + gUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_BUF_OFFSET); + gLastUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_SUB_BUF_OFFSET); + gInputUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_G_SUB_BUF_OFFSET); + vNewOutputUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_2_OFFSET); + vNewDecayUbTensor_ping = resource.ubBuf.template GetBufferByByte(PING_BUF_2_OFFSET); + + uUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_2_OFFSET); + wsUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_0_OFFSET); + gUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_BUF_OFFSET); + gLastUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_SUB_BUF_OFFSET); + gInputUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_G_SUB_BUF_OFFSET); + vNewOutputUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_2_OFFSET); + vNewDecayUbTensor_pong = resource.ubBuf.template GetBufferByByte(PONG_BUF_2_OFFSET); + + shareBuffer_ = resource.ubBuf.template GetBufferByByte(SHARE_BUF_OFFSET); + + } + + CATLASS_DEVICE + ~BlockEpilogue() {} + + template + CATLASS_DEVICE + void CopyGmToUb( + AscendC::LocalTensor dst, + AscendC::GlobalTensor src, + uint32_t rows, + uint32_t cols, + uint32_t srcStride) + { + if (cols == srcStride) { + AscendC::DataCopy(dst, src, rows * cols); + return; + } + AscendC::DataCopyExtParams copyParams{ + static_cast(rows), + static_cast(cols * sizeof(Element)), + static_cast((srcStride - cols) * sizeof(Element)), + 0, + 0}; + AscendC::DataCopyPadExtParams padParams{false, 0, 0, 0}; + AscendC::DataCopyPad(dst, src, copyParams, padParams); + } + + template + CATLASS_DEVICE + void CopyUbToGm( + AscendC::GlobalTensor dst, + AscendC::LocalTensor src, + uint32_t rows, + uint32_t cols, + uint32_t dstStride) + { + if (cols == dstStride) { + AscendC::DataCopy(dst, src, rows * cols); + return; + } + AscendC::DataCopyExtParams copyParams{ + static_cast(rows), + static_cast(cols * sizeof(Element)), + 0, + static_cast((dstStride - cols) * sizeof(Element)), + 0}; + AscendC::DataCopyPad(dst, src, copyParams); + } + + CATLASS_DEVICE + void PrepareG( + AscendC::LocalTensor gUbTensor, + AscendC::LocalTensor gLastUbTensor, + AscendC::LocalTensor gInputUbTensor, + AscendC::GlobalTensor gInputThisSubBlock, + uint32_t mActual, + uint32_t pingpongFlag) + { + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + if constexpr (!scalarGated) { + AscendC::Duplicate(gUbTensor, 1.0f, mActual); + AscendC::PipeBarrier(); + return; + } + if (mActual == 1) { + AscendC::Duplicate(gUbTensor, 1.0f, 1); + AscendC::PipeBarrier(); + return; + } + if constexpr(std::is_same::value) { + AscendC::DataCopyParams gUbParams{1, (uint16_t)(mActual * sizeof(float)), 0, 0}; + AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0}; + AscendC::DataCopyPad(gUbTensor, gInputThisSubBlock, gUbParams, gUbPadParams); + } else { + AscendC::DataCopyParams gUbParams{1, (uint16_t)(mActual * sizeof(GElementInput)), 0, 0}; + AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0}; + AscendC::DataCopyPad(gInputUbTensor, gInputThisSubBlock, gUbParams, gUbPadParams); + } + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + if constexpr(!std::is_same::value) { + AscendC::Cast(gUbTensor, gInputUbTensor, AscendC::RoundMode::CAST_NONE, mActual); + } + + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + float inputVal = gUbTensor.GetValue(mActual - 1); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + + AscendC::PipeBarrier(); + AscendC::Duplicate(gLastUbTensor, inputVal, mActual); + AscendC::PipeBarrier(); + + AscendC::Sub(gUbTensor, gLastUbTensor, gUbTensor, mActual); + AscendC::PipeBarrier(); + if constexpr (useExp2) { + AscendC::Muls(gUbTensor, gUbTensor, LN2, mActual); + AscendC::PipeBarrier(); + } + AscendC::Exp(gUbTensor, gUbTensor, mActual); + AscendC::PipeBarrier(); + } + + CATLASS_DEVICE + void operator()( + AscendC::GlobalTensor vnewOutput, + AscendC::GlobalTensor vnewdecayOutput, + AscendC::GlobalTensor gInput, + AscendC::GlobalTensor uInput, + AscendC::GlobalTensor wsInput, + AscendC::GlobalTensor gkInput, + AscendC::GlobalTensor kInput, + AscendC::GlobalTensor kDecayWorkspace, + uint32_t chunkSize, + uint32_t kHeadDim, + uint32_t vBlockDim, + uint32_t vHeadDim, + Arch::CrossCoreFlag cube1Done, + Arch::CrossCoreFlag vec1Done, + bool isInitialState, + bool isFinalState, + bool storeFinalState, + bool waitWsFromMte3, + bool isPing, + bool cube1AlreadyWaited + ) + { + static constexpr uint32_t ROW_TILE = 16; + uint32_t mActual = chunkSize; + uint32_t nvActual = vBlockDim; + uint32_t nkActual = kHeadDim; + uint32_t inputStride = vHeadDim; + + uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); + uint32_t subBlockNum = AscendC::GetSubBlockNum(); + uint32_t rowsPerSubBlock = CeilDiv(mActual, subBlockNum); + uint32_t rowBegin = subBlockIdx * rowsPerSubBlock; + uint32_t rowEnd = rowBegin + rowsPerSubBlock; + if (rowEnd > mActual) { + rowEnd = mActual; + } + uint32_t pingpongFlag = isPing ? 0 : pongBaseEvent; + if (rowBegin >= mActual) { + if (!cube1AlreadyWaited) { + Arch::CrossCoreWaitFlag(cube1Done); + } + // A zero-row AIV lane still owns the EVENT0 hand-off consumed by V2. + if (waitWsFromMte3) { + AscendC::WaitFlag( + EVENT_ID0 + pingpongFlag); + AscendC::SetFlag( + EVENT_ID0 + pingpongFlag); + } + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + return; + } + AscendC::ResetMask(); + + AscendC::GlobalTensor gInputThisSubBlock = gInput; + + AscendC::LocalTensor uUbTensor = isPing ? uUbTensor_ping : uUbTensor_pong; + AscendC::LocalTensor wsUbTensor = isPing ? wsUbTensor_ping : wsUbTensor_pong; + AscendC::LocalTensor gUbTensor = isPing ? gUbTensor_ping : gUbTensor_pong; + AscendC::LocalTensor gLastUbTensor = isPing ? gLastUbTensor_ping : gLastUbTensor_pong; + AscendC::LocalTensor gInputUbTensor = isPing ? gInputUbTensor_ping : gInputUbTensor_pong; + AscendC::LocalTensor vNewOutputUbTensor = isPing ? vNewOutputUbTensor_ping : vNewOutputUbTensor_pong; + AscendC::LocalTensor vNewDecayUbTensor = isPing ? vNewDecayUbTensor_ping : vNewDecayUbTensor_pong; + + if (nvActual <= 128 && nvActual == inputStride) { + uint32_t mActualThisSubBlock = rowEnd - rowBegin; + uint32_t gbrcRealStart = rowBegin & ~7; + uint32_t gbrcEffStart = rowBegin - gbrcRealStart; + uint32_t gbrcRealProcess = gbrcEffStart + mActualThisSubBlock; + uint32_t dstShape_[2] = {gbrcRealProcess, nvActual}; + uint32_t srcShape_[2] = {gbrcRealProcess, 1}; + + AscendC::GlobalTensor vnewOutputThisSubBlock = vnewOutput[rowBegin * inputStride]; + AscendC::GlobalTensor vnewdecayOutputThisSubBlock = vnewdecayOutput[rowBegin * nvActual]; + AscendC::GlobalTensor uInputThisSubBlock = uInput[rowBegin * inputStride]; + AscendC::GlobalTensor wsInputThisSubBlock = wsInput[rowBegin * nvActual]; + + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(uUbTensor, uInputThisSubBlock, mActualThisSubBlock * nvActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::Cast(calcUbTensor, uUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nvActual); + AscendC::PipeBarrier(); + + if constexpr (scalarGated) { + PrepareG(gUbTensor, gLastUbTensor, gInputUbTensor, gInputThisSubBlock, mActual, pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + } + if (!cube1AlreadyWaited) { + Arch::CrossCoreWaitFlag(cube1Done); + } + + if (waitWsFromMte3) { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } + AscendC::DataCopy(wsUbTensor, wsInputThisSubBlock, mActualThisSubBlock * nvActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + + AscendC::Sub(wsUbTensor, calcUbTensor, wsUbTensor, mActualThisSubBlock * nvActual); + AscendC::PipeBarrier(); + + uint32_t decayOffset = 0; + if constexpr (scalarGated) { + AscendC::Broadcast(calcUbTensor, gUbTensor[gbrcRealStart], dstShape_, srcShape_, shareBuffer_); + AscendC::PipeBarrier(); + AscendC::Mul(calcUbTensor[gbrcEffStart * nvActual], wsUbTensor, calcUbTensor[gbrcEffStart * nvActual], mActualThisSubBlock * nvActual); + decayOffset = gbrcEffStart * nvActual; + } else { + AscendC::Adds(calcUbTensor, wsUbTensor, 0.0f, mActualThisSubBlock * nvActual); + } + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + + AscendC::Cast(vNewDecayUbTensor, calcUbTensor[decayOffset], AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nvActual); + AscendC::PipeBarrier(); + + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(vnewdecayOutputThisSubBlock, vNewDecayUbTensor, mActualThisSubBlock * nvActual); + if constexpr (!kGated) { + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + } + + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::Cast(vNewOutputUbTensor, wsUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nvActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(vnewOutputThisSubBlock, vNewOutputUbTensor, mActualThisSubBlock * nvActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + + if constexpr (kGated) { + AscendC::GlobalTensor kInputThisSubBlock = kInput[rowBegin * nkActual]; + AscendC::GlobalTensor kDecayWorkspaceThisSubBlock = kDecayWorkspace[rowBegin * nkActual]; + // KDA passes kg = k * exp2(g_last - gk). Keep that decay exactly once. + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(vNewOutputUbTensor, kInputThisSubBlock, mActualThisSubBlock * nkActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(kDecayWorkspaceThisSubBlock, vNewOutputUbTensor, + mActualThisSubBlock * nkActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + } + return; + } + + if constexpr (scalarGated) { + PrepareG(gUbTensor, gLastUbTensor, gInputUbTensor, gInputThisSubBlock, mActual, pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID3 + pingpongFlag); + } + if (!cube1AlreadyWaited) { + Arch::CrossCoreWaitFlag(cube1Done); + } + + bool waitWsThisTileFromMte3 = waitWsFromMte3; + for (uint32_t rowStart = rowBegin; rowStart < rowEnd;) { + uint32_t alignExtra = rowStart & 7; + uint32_t maxRowsThisTile = ROW_TILE - alignExtra; + uint32_t rowsThisTile = rowEnd - rowStart; + if (rowsThisTile > maxRowsThisTile) { + rowsThisTile = maxRowsThisTile; + } + uint32_t gbrcRealStart = rowStart & ~7; + uint32_t gbrcRealProcess = alignExtra + rowsThisTile; + uint32_t dstShape_[2] = {gbrcRealProcess, nvActual}; + uint32_t srcShape_[2] = {gbrcRealProcess, 1}; + + AscendC::GlobalTensor vnewOutputThisTile = vnewOutput[rowStart * inputStride]; + AscendC::GlobalTensor vnewdecayOutputThisTile = vnewdecayOutput[rowStart * nvActual]; + AscendC::GlobalTensor uInputThisTile = uInput[rowStart * inputStride]; + AscendC::GlobalTensor wsInputThisTile = wsInput[rowStart * nvActual]; + + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + CopyGmToUb(uUbTensor, uInputThisTile, rowsThisTile, nvActual, inputStride); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::Cast(calcUbTensor, uUbTensor, AscendC::RoundMode::CAST_NONE, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + + if (waitWsThisTileFromMte3) { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } else { + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + } + AscendC::DataCopy(wsUbTensor, wsInputThisTile, rowsThisTile * nvActual); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID0 + pingpongFlag); + waitWsThisTileFromMte3 = false; + + AscendC::Sub(wsUbTensor, calcUbTensor, wsUbTensor, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + + uint32_t decayOffset = 0; + if constexpr (scalarGated) { + AscendC::Broadcast(calcUbTensor, gUbTensor[gbrcRealStart], dstShape_, srcShape_, shareBuffer_); + AscendC::PipeBarrier(); + AscendC::Mul(calcUbTensor[alignExtra * nvActual], wsUbTensor, calcUbTensor[alignExtra * nvActual], rowsThisTile * nvActual); + decayOffset = alignExtra * nvActual; + } else { + AscendC::Adds(calcUbTensor, wsUbTensor, 0.0f, rowsThisTile * nvActual); + } + AscendC::PipeBarrier(); + + AscendC::Cast(vNewDecayUbTensor, calcUbTensor[decayOffset], AscendC::RoundMode::CAST_RINT, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + AscendC::DataCopy(vnewdecayOutputThisTile, vNewDecayUbTensor, rowsThisTile * nvActual); + + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + if (rowStart + rowsThisTile >= rowEnd) { + if constexpr (!kGated) { + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + } + } + AscendC::Cast(vNewOutputUbTensor, wsUbTensor, AscendC::RoundMode::CAST_RINT, rowsThisTile * nvActual); + AscendC::PipeBarrier(); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::SetFlag(EVENT_ID0 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + CopyUbToGm(vnewOutputThisTile, vNewOutputUbTensor, rowsThisTile, nvActual, inputStride); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + rowStart += rowsThisTile; + } + + if constexpr (kGated) { + uint32_t mActualThisSubBlock = rowEnd - rowBegin; + AscendC::GlobalTensor kInputThisSubBlock = kInput[rowBegin * nkActual]; + AscendC::GlobalTensor kDecayWorkspaceThisSubBlock = kDecayWorkspace[rowBegin * nkActual]; + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + CopyGmToUb(vNewOutputUbTensor, kInputThisSubBlock, mActualThisSubBlock, nkActual, nkActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + AscendC::WaitFlag(EVENT_ID1 + pingpongFlag); + CopyUbToGm(kDecayWorkspaceThisSubBlock, vNewOutputUbTensor, mActualThisSubBlock, nkActual, nkActual); + AscendC::SetFlag(EVENT_ID1 + pingpongFlag); + + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vec1Done); + } + AscendC::SetFlag(EVENT_ID3 + pingpongFlag); + } + +private: + uint32_t pongBaseEvent = 4; + + AscendC::LocalTensor calcUbTensor; + + AscendC::LocalTensor uUbTensor_ping; + AscendC::LocalTensor wsUbTensor_ping; + AscendC::LocalTensor gUbTensor_ping; + AscendC::LocalTensor gLastUbTensor_ping; + AscendC::LocalTensor gInputUbTensor_ping; + AscendC::LocalTensor vNewOutputUbTensor_ping; + AscendC::LocalTensor vNewDecayUbTensor_ping; + + AscendC::LocalTensor uUbTensor_pong; + AscendC::LocalTensor wsUbTensor_pong; + AscendC::LocalTensor gUbTensor_pong; + AscendC::LocalTensor gLastUbTensor_pong; + AscendC::LocalTensor gInputUbTensor_pong; + AscendC::LocalTensor vNewOutputUbTensor_pong; + AscendC::LocalTensor vNewDecayUbTensor_pong; + + AscendC::LocalTensor shareBuffer_; + +}; +} + +#endif diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/epilogue/gdn_fwd_h_epilogue_policies.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/epilogue/gdn_fwd_h_epilogue_policies.hpp new file mode 100644 index 000000000000..d5287d55defd --- /dev/null +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/epilogue/gdn_fwd_h_epilogue_policies.hpp @@ -0,0 +1,27 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#ifndef CATLASS_EPILOGUE_GDN_FWD_H_EPILOGUE_POLICIES_HPP +#define CATLASS_EPILOGUE_GDN_FWD_H_EPILOGUE_POLICIES_HPP + +#include "catlass/catlass.hpp" + +namespace Catlass::Epilogue { + +struct EpilogueAtlasGDNFwdHVnew { + using ArchTag = Arch::AtlasA2; +}; + +struct EpilogueAtlasGDNFwdHUpdate { + using ArchTag = Arch::AtlasA2; +}; + +} // namespace Catlass::Epilogue + +#endif // CATLASS_EPILOGUE_GDN_FWD_H_EPILOGUE_POLICIES_HPP diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/gemm/block/block_scheduler_gdn_fwd_h.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/gemm/block/block_scheduler_gdn_fwd_h.hpp new file mode 100644 index 000000000000..cac9778e8b1f --- /dev/null +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/gemm/block/block_scheduler_gdn_fwd_h.hpp @@ -0,0 +1,433 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#include "catlass/gemm_coord.hpp" +using namespace Catlass; + +#ifndef CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP +#define CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP + +// constexpr uint32_t PING_PONG_STAGES = 1; +constexpr uint32_t PING_PONG_STAGES = 2; +constexpr uint32_t BYTE_SIZE_16_BIT = 2; +constexpr uint32_t BYTES_PER_C0 = 32; +constexpr uint32_t BYTE_SIZE_PER_REPEAT = 256; +constexpr uint32_t SIZE_16_NUM_PER_C0 = BYTES_PER_C0 / BYTE_SIZE_16_BIT; +constexpr uint32_t FLOAT_NUM_PER_REPEAT = BYTE_SIZE_PER_REPEAT / sizeof(float); +constexpr uint32_t NZ_BLOCK_SIZE = 16; + +template +CATLASS_DEVICE T AlignUp(T a, T b) { + return (b == 0) ? 0 : (a + b - 1) / b * b; +} + +template +CATLASS_DEVICE T Min(T a, T b) { + return (a > b) ? b : a; +} + +template +CATLASS_DEVICE T Max(T a, T b) { + return (a > b) ? a : b; +} + +namespace Catlass::Gemm::Block { + +struct GDNFwdHOffsets { + uint32_t hSrcOffset; + uint32_t hDstOffset; + uint32_t uvOffset; + uint32_t wkOffset; + uint32_t wOffset; + uint32_t gOffset; + uint32_t gkOffset; + uint32_t hWorkOffset; + uint32_t vWorkOffset; + uint32_t kDecayWorkOffset; + uint32_t vBlockOffset; + uint32_t vBlockDim; + uint32_t initialStateOffset; + uint32_t finalStateOffset; + bool isInitialState; + bool isFinalState; + uint32_t blockTokens; + uint32_t streamId; + // for debug + uint32_t batchIdx; + uint32_t headIdx; + uint32_t chunkIdx; + +}; + +struct GDNFwdHStream { + uint32_t batchIdx; + uint32_t chunkIdx{0}; + uint32_t vHeadIdx; + uint32_t kHeadIdx; + uint32_t shapeBatchIdx; + uint32_t tokenBatchIdx; + + uint32_t chunkOffset; + uint32_t tokenOffset; + uint32_t batchChunks{0}; + uint32_t batchTokens; + uint32_t nextTaskIdx{0}; + bool active{false}; + + GDNFwdHOffsets offset; +}; + +struct GDNFwdHRunningQ { + GDNFwdHStream streams[PING_PONG_STAGES]; +}; + +struct BlockSchedulerGdnFwdH { + uint32_t batch; + uint32_t seqlen; + uint32_t kNumHead; + uint32_t vNumHead; + uint32_t kHeadDim; + uint32_t vHeadDim; + uint32_t chunkSize; + uint32_t vBlockSize{128}; + uint32_t isVariedLen; + uint32_t shapeBatch; + uint32_t tokenBatch; + uint32_t inputTokenBatch; + bool useInitialState; + bool storeFinalState; + uint32_t numSeqWorkspaceOffset; + uint32_t numChunksWorkspaceOffset; + + uint32_t taskIdx; + uint32_t taskStride; + uint32_t cubeCoreIdx; + uint32_t cubeCoreNum; + uint32_t taskNum; + uint32_t headGroups; + uint32_t totalChunks; + uint32_t totalTokens; + + GDNFwdHRunningQ runningQ; + + bool isRunning; + + AscendC::GlobalTensor gmSeqlen; + AscendC::GlobalTensor gmNumSeq; + AscendC::GlobalTensor gmNumChunks; + + Arch::CrossCoreFlag cube1Done[PING_PONG_STAGES] = {0, 1}; + Arch::CrossCoreFlag vec1Done[PING_PONG_STAGES] = {2, 3}; + Arch::CrossCoreFlag cube2Done[PING_PONG_STAGES] = {4, 5}; + Arch::CrossCoreFlag vec2Done[PING_PONG_STAGES] = {6, 7}; + + CATLASS_DEVICE + BlockSchedulerGdnFwdH() {} + + CATLASS_DEVICE + void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR tiling, GM_ADDR user, uint32_t coreIdx, uint32_t coreNum) { + __gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict gdnFwdHTilingData = reinterpret_cast<__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict>(tiling); + + batch = gdnFwdHTilingData->batch; + seqlen = gdnFwdHTilingData->seqlen; + kNumHead = gdnFwdHTilingData->kNumHead; + vNumHead = gdnFwdHTilingData->vNumHead; + kHeadDim = gdnFwdHTilingData->kHeadDim; + vHeadDim = gdnFwdHTilingData->vHeadDim; + chunkSize = gdnFwdHTilingData->chunkSize; + isVariedLen = gdnFwdHTilingData->isVariedLen; + shapeBatch = gdnFwdHTilingData->shapeBatch; + tokenBatch = gdnFwdHTilingData->tokenBatch; + useInitialState = gdnFwdHTilingData->useInitialState; + storeFinalState = gdnFwdHTilingData->storeFinalState; + numSeqWorkspaceOffset = gdnFwdHTilingData->numSeqWorkspaceOffset; + numChunksWorkspaceOffset = gdnFwdHTilingData->numChunksWorkspaceOffset; + + InitRuntime(cu_seqlens, chunk_indices, user, coreIdx, coreNum); + } + + template + CATLASS_DEVICE + void InitFromData(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, const TilingData& tilingData, + GM_ADDR user, uint32_t coreIdx, uint32_t coreNum) { + batch = tilingData.batch; + seqlen = tilingData.seqlen; + kNumHead = tilingData.kNumHead; + vNumHead = tilingData.vNumHead; + kHeadDim = tilingData.kHeadDim; + vHeadDim = tilingData.vHeadDim; + chunkSize = tilingData.chunkSize; + isVariedLen = tilingData.isVariedLen; + shapeBatch = tilingData.shapeBatch; + tokenBatch = tilingData.tokenBatch; + useInitialState = tilingData.useInitialState; + storeFinalState = tilingData.storeFinalState; + numSeqWorkspaceOffset = tilingData.numSeqWorkspaceOffset; + numChunksWorkspaceOffset = tilingData.numChunksWorkspaceOffset; + + InitRuntime(cu_seqlens, chunk_indices, user, coreIdx, coreNum); + } + + CATLASS_DEVICE + void InitRuntime(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR user, + uint32_t coreIdx, uint32_t coreNum) { + + gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens); + gmNumSeq.SetGlobalBuffer((__gm__ int64_t *)(user + numSeqWorkspaceOffset)); + gmNumChunks.SetGlobalBuffer((__gm__ int64_t *)(user + numChunksWorkspaceOffset)); + + if (isVariedLen) { + inputTokenBatch = tokenBatch; + uint32_t actualBatch = 0; + int64_t chunkPrefix = 0; + int64_t prevSeq = 0, currSeq; + for (uint32_t b = 1; b <= inputTokenBatch; b++) { + currSeq = gmSeqlen.GetValue(b); + int64_t batchSeqLen = currSeq - prevSeq; + if (batchSeqLen > 0) { + actualBatch++; + int64_t batchChunk = (batchSeqLen + chunkSize - 1) / chunkSize; + chunkPrefix += batchChunk; + } + prevSeq = currSeq; + } + tokenBatch = actualBatch; + batch = actualBatch; + totalChunks = chunkPrefix; + totalTokens = prevSeq; + } else { + totalChunks = (seqlen + chunkSize - 1) / chunkSize; + totalTokens = seqlen; + } + + cubeCoreIdx = coreIdx; + cubeCoreNum = coreNum; + vBlockSize = vHeadDim; + taskNum = batch * vNumHead; + headGroups = vNumHead / kNumHead; + InitTaskWave(0); + + } + + CATLASS_DEVICE + void InitTaskWave(uint32_t waveIdx) { + uint32_t firstTaskIdx = waveIdx * cubeCoreNum + cubeCoreIdx; + taskStride = taskNum; + for (uint32_t streamId = 0; streamId < PING_PONG_STAGES; ++streamId) { + auto& stream = runningQ.streams[streamId]; + stream.nextTaskIdx = streamId == 0 ? firstTaskIdx : taskNum; + stream.chunkIdx = 0; + stream.batchChunks = 0; + stream.active = false; + } + isRunning = firstTaskIdx < taskNum; + } + + CATLASS_DEVICE + uint32_t GetTaskWaveCount() const { + return CeilDiv(taskNum, cubeCoreNum); + } + + CATLASS_DEVICE + void ResolveVarlenSequence(uint32_t compactBatchIdx, GDNFwdHStream& stream) { + uint32_t actualBatch = 0; + int64_t chunkPrefix = 0; + int64_t prevSeq = 0; + for (uint32_t b = 1; b <= inputTokenBatch; ++b) { + int64_t currSeq = gmSeqlen.GetValue(b); + int64_t batchTokens = currSeq - prevSeq; + if (batchTokens > 0) { + int64_t batchChunks = (batchTokens + chunkSize - 1) / chunkSize; + if (actualBatch == compactBatchIdx) { + stream.chunkOffset = static_cast(chunkPrefix); + stream.batchChunks = static_cast(batchChunks); + stream.tokenOffset = static_cast(prevSeq); + stream.batchTokens = static_cast(batchTokens); + return; + } + ++actualBatch; + chunkPrefix += batchChunks; + } + prevSeq = currSeq; + } + stream.chunkOffset = 0; + stream.batchChunks = 0; + stream.tokenOffset = 0; + stream.batchTokens = 0; + } + + CATLASS_DEVICE + uint32_t GetVarlenChunkOffset(uint32_t compactBatchIdx) { + GDNFwdHStream stream; + ResolveVarlenSequence(compactBatchIdx, stream); + return stream.chunkOffset; + } + + CATLASS_DEVICE + void InitNewStream(GDNFwdHStream& newStream) { + newStream.batchIdx = taskIdx / vNumHead; + newStream.vHeadIdx = taskIdx % vNumHead; + newStream.kHeadIdx = newStream.vHeadIdx / headGroups; + newStream.shapeBatchIdx = isVariedLen ? 0 : newStream.batchIdx; + newStream.tokenBatchIdx = isVariedLen ? newStream.batchIdx : 0; + if (isVariedLen) { + ResolveVarlenSequence(newStream.tokenBatchIdx, newStream); + } else { + newStream.chunkOffset = 0; + newStream.batchChunks = totalChunks; + newStream.tokenOffset = 0; + newStream.batchTokens = totalTokens; + } + newStream.chunkIdx = 0; + } + + CATLASS_DEVICE + void AssignNextStream(uint32_t streamId) { + auto& stream = runningQ.streams[streamId]; + taskIdx = stream.nextTaskIdx; + if (taskIdx >= taskNum) { + stream.active = false; + stream.batchChunks = 0; + return; + } + + stream.nextTaskIdx += taskStride; + InitNewStream(stream); + stream.active = stream.batchChunks > 0; + if (stream.active) { + UpdateTask(streamId); + } + } + + CATLASS_DEVICE + void UpdateTask(uint32_t streamId) { + auto& stream = runningQ.streams[streamId]; + auto& offset = stream.offset; + + offset.isInitialState = stream.chunkIdx == 0; + offset.isFinalState = stream.chunkIdx == (stream.batchChunks - 1); + uint32_t vBlockOffset = 0; + uint32_t vBlockDim = vBlockSize; + offset.initialStateOffset = (stream.batchIdx * vNumHead + stream.vHeadIdx) * kHeadDim * vHeadDim + vBlockOffset; + offset.finalStateOffset = (stream.batchIdx * vNumHead + stream.vHeadIdx) * kHeadDim * vHeadDim + vBlockOffset; + offset.hSrcOffset = (stream.shapeBatchIdx * vNumHead * totalChunks + stream.vHeadIdx * totalChunks + stream.chunkOffset + stream.chunkIdx) * kHeadDim * vHeadDim + vBlockOffset; + offset.hDstOffset = offset.hSrcOffset + kHeadDim * vHeadDim; + if (storeFinalState && offset.isFinalState) { + offset.hDstOffset = offset.hSrcOffset; + } + offset.uvOffset = (stream.shapeBatchIdx * vNumHead * totalTokens + stream.vHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize) * vHeadDim + vBlockOffset; + offset.wkOffset = (stream.shapeBatchIdx * kNumHead * totalTokens + stream.kHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize) * kHeadDim; + offset.wOffset = (stream.shapeBatchIdx * vNumHead * totalTokens + stream.vHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize) * kHeadDim; + offset.gOffset = stream.shapeBatchIdx * vNumHead * totalTokens + stream.vHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize; + offset.gkOffset = (stream.shapeBatchIdx * vNumHead * totalTokens + stream.vHeadIdx * totalTokens + stream.tokenOffset + stream.chunkIdx * chunkSize) * kHeadDim; + offset.hWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + streamId) * kHeadDim * vBlockSize; + offset.vWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + streamId) * chunkSize * vBlockSize; + offset.kDecayWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + streamId) * chunkSize * kHeadDim; + offset.vBlockOffset = vBlockOffset; + offset.vBlockDim = vBlockDim; + offset.blockTokens = offset.isFinalState ? (stream.batchTokens - stream.chunkIdx * chunkSize) : chunkSize; + offset.streamId = streamId; + offset.batchIdx = stream.batchIdx; + offset.headIdx = stream.vHeadIdx; + offset.chunkIdx = stream.chunkIdx; + } + + CATLASS_DEVICE + void InitTasks() { + isRunning = false; + for (uint32_t streamId = 0; streamId < PING_PONG_STAGES; ++streamId) { + auto& stream = runningQ.streams[streamId]; + if (stream.active) { + stream.chunkIdx += 1; + if (stream.chunkIdx >= stream.batchChunks) { + stream.active = false; + stream.batchChunks = 0; + } + } + if (!stream.active) { + AssignNextStream(streamId); + } else { + UpdateTask(streamId); + } + if (stream.active) { + isRunning = true; + } + } + } + + CATLASS_DEVICE + const GDNFwdHStream& GetStream(uint32_t i) const { + return runningQ.streams[i]; + } + + CATLASS_DEVICE + uint32_t GetStreamId(uint32_t i) const { + return i; + } + + CATLASS_DEVICE + const GDNFwdHOffsets& GetCurTaskOffsets(const GDNFwdHStream& stream) const { + return stream.offset; + } + + CATLASS_DEVICE + bool StreamIsDone(const GDNFwdHStream& stream) const { + return !stream.active; + } + + CATLASS_DEVICE + bool NeedProcessStage2(const GDNFwdHStream& stream) { + return storeFinalState || !stream.offset.isFinalState; + } +}; + +struct BlockSchedulerGdnFwdHCube : public BlockSchedulerGdnFwdH { + CATLASS_DEVICE + BlockSchedulerGdnFwdHCube() {} + + CATLASS_DEVICE + void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR tiling, GM_ADDR user) { + BlockSchedulerGdnFwdH::Init(cu_seqlens, chunk_indices, tiling, user, AscendC::GetBlockIdx(), AscendC::GetBlockNum()); + } + + template + CATLASS_DEVICE + void InitFromData(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, const TilingData& tilingData, GM_ADDR user) { + BlockSchedulerGdnFwdH::InitFromData( + cu_seqlens, chunk_indices, tilingData, user, AscendC::GetBlockIdx(), AscendC::GetBlockNum()); + } + +}; + +struct BlockSchedulerGdnFwdHVec : public BlockSchedulerGdnFwdH { + CATLASS_DEVICE + BlockSchedulerGdnFwdHVec() {} + + CATLASS_DEVICE + void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR tiling, GM_ADDR user) { + BlockSchedulerGdnFwdH::Init( + cu_seqlens, chunk_indices, tiling, user, + AscendC::GetBlockIdx() / AscendC::GetSubBlockNum(), + AscendC::GetBlockNum()); + } + + template + CATLASS_DEVICE + void InitFromData(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, const TilingData& tilingData, GM_ADDR user) { + BlockSchedulerGdnFwdH::InitFromData( + cu_seqlens, chunk_indices, tilingData, user, + AscendC::GetBlockIdx() / AscendC::GetSubBlockNum(), + AscendC::GetBlockNum()); + } + +}; + +} // namespace Catlass::Gemm::Block + +#endif // CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP diff --git a/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/gemm/kernel/gdn_fwd_h_kernel.hpp b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/gemm/kernel/gdn_fwd_h_kernel.hpp new file mode 100644 index 000000000000..4f2de83b1670 --- /dev/null +++ b/csrc/moe/chunk_gated_delta_rule_fwd_h/op_kernel/gemm/kernel/gdn_fwd_h_kernel.hpp @@ -0,0 +1,814 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * the BSD 3-Clause License (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + */ + +#define CATLASS_ARCH 2201 + +#include "catlass/arch/arch.hpp" +#include "catlass/arch/cross_core_sync.hpp" +#include "catlass/arch/resource.hpp" +#include "catlass/catlass.hpp" +#include "catlass/debug.hpp" +#include "catlass/epilogue/block/block_epilogue.hpp" +#include "../../epilogue/block/block_epilogue_gdn_fwdh_update.hpp" +#include "../../epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp" +#include "catlass/gemm/block/block_mmad.hpp" +#include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp" +#include "catlass/gemm/block/block_swizzle.hpp" +#include "../block/block_scheduler_gdn_fwd_h.hpp" +#include "catlass/gemm/dispatch_policy.hpp" +#include "catlass/gemm/gemm_type.hpp" +#include "catlass/layout/layout.hpp" +#include "catlass/gemm_coord.hpp" +#include "tla/tensor.hpp" +#include "tla/layout.hpp" +#include "tla/tensor.hpp" + + + +#include "kernel_operator.h" +using namespace Catlass; +using namespace tla; + +namespace Catlass::Gemm::Kernel { + +struct GDNFwdHTileShapes128 { + using L1TileShape = tla::Shape<_128, _128, _128>; + using L0TileShape = L1TileShape; +}; + +struct GDNFwdHTileShapes256 { + using L1TileShape = tla::Shape<_128, _256, _128>; + using L0TileShape = tla::Shape<_128, _256, _64>; +}; + +template +struct GDNFwdHGateTag { + static constexpr bool value = KGated; + static constexpr bool scalarGated = ScalarGated; + static constexpr bool useExp2 = UseExp2; +}; + +template< + typename INPUT_TYPE, + typename G_TYPE, + typename STATE_TYPE, + typename WORKSPACE_TYPE, + typename TileShapes = GDNFwdHTileShapes128, + bool kGated = false, + bool scalarGated = true, + bool useExp2 = false +> +class GDNFwdHKernel { +public: + + using ArchTag = Arch::AtlasA2; + using CubeScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdHCube; + using VecScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdHVec; + + using DispatchPolicyTla = Gemm::MmadPingpongTlaMulti; + using L1TileShapeVTla = typename TileShapes::L1TileShape; + using L0TileShapeVTla = typename TileShapes::L0TileShape; + + using WType = Gemm::GemmType; + using HType = Gemm::GemmType; + using VworkType = Gemm::GemmType; + using KType = Gemm::GemmType; + using HworkType = Gemm::GemmType; + using VType = Gemm::GemmType; + using GType = Gemm::GemmType; + using UType = Gemm::GemmType; + using FinalStateType = Gemm::GemmType; + + // cube 1 + using TileCopyWH = Catlass::Gemm::Tile::PackedTileCopyTla; + using BlockMmadWH = Gemm::Block::BlockMmadTla; + + // cube 2 + using TileCopyKV = Catlass::Gemm::Tile::PackedTileCopyTla; + using BlockMmadKV = Gemm::Block::BlockMmadTla; + + // vec 1 + using DispatchPolicyGDNFwdHVnew = Epilogue::EpilogueAtlasGDNFwdHVnew; + using GateTag = GDNFwdHGateTag; + using EpilogueGDNFwdHVnew = Epilogue::Block::BlockEpilogue; + + // vec 2 + using DispatchPolicyGDNFwdHUpdate = Epilogue::EpilogueAtlasGDNFwdHUpdate; + using EpilogueGDNFwdHUpdate = Epilogue::Block::BlockEpilogue; + + using GDNFwdHOffsets = Catlass::Gemm::Block::GDNFwdHOffsets; + + using ElementK = INPUT_TYPE; + using ElementW = INPUT_TYPE; + using ElementU = INPUT_TYPE; + using ElementG = G_TYPE; + using ElementH = INPUT_TYPE; + using ElementV = INPUT_TYPE; + using ElementVWork = WORKSPACE_TYPE; + using ElementHWork = WORKSPACE_TYPE; + using ElementInitialState = STATE_TYPE; + using ElementFinalState = STATE_TYPE; + + using LayoutW = Catlass::layout::RowMajor; + using LayoutH = Catlass::layout::RowMajor; + using LayoutV = Catlass::layout::RowMajor; + using LayoutK = Catlass::layout::ColumnMajor; + + + uint32_t batch; + uint32_t seqlen; + uint32_t kNumHead; + uint32_t vNumHead; + uint32_t kHeadDim; + uint32_t vHeadDim; + uint32_t chunkSize; + bool useInitialState; + bool storeFinalState; + uint32_t isVariedLen; + uint32_t shapeBatch; + uint32_t tokenBatch; + uint32_t vWorkspaceOffset; + uint32_t vUpdateWorkspaceOffset; + uint32_t hWorkspaceOffset; + uint32_t numSeqWorkspaceOffset; + uint32_t numChunksWorkspaceOffset; + uint32_t kDecayWorkspaceOffset; + + AscendC::GlobalTensor gmK; + AscendC::GlobalTensor gmW; + AscendC::GlobalTensor gmU; + AscendC::GlobalTensor gmG; + AscendC::GlobalTensor gmInitialState; + AscendC::GlobalTensor gmH; + AscendC::GlobalTensor gmV; + AscendC::GlobalTensor gmFinalState; + AscendC::GlobalTensor gmVWorkspace; + AscendC::GlobalTensor gmVUpdateWorkspace; + AscendC::GlobalTensor gmHWorkspace; + + AscendC::GlobalTensor gmGk; + AscendC::GlobalTensor gmKDecayWorkspace; + + AscendC::GlobalTensor gmSeqlen; + AscendC::GlobalTensor gmNumSeq; + AscendC::GlobalTensor gmNumChunks; + + CubeScheduler cubeBlockScheduler; + VecScheduler vecBlockScheduler; + + Arch::Resource resource; + + + __aicore__ inline GDNFwdHKernel() {} + + __aicore__ inline void Init(GM_ADDR k, GM_ADDR w, GM_ADDR u, GM_ADDR g, GM_ADDR gk, GM_ADDR inital_state, GM_ADDR cu_seqlens, GM_ADDR chunk_indices, + GM_ADDR h, GM_ADDR v_new, GM_ADDR final_state, GM_ADDR tiling, GM_ADDR user) { + + __gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict gdnFwdHTilingData = reinterpret_cast<__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict>(tiling); + + batch = gdnFwdHTilingData->batch; + seqlen = gdnFwdHTilingData->seqlen; + kNumHead = gdnFwdHTilingData->kNumHead; + vNumHead = gdnFwdHTilingData->vNumHead; + kHeadDim = gdnFwdHTilingData->kHeadDim; + vHeadDim = gdnFwdHTilingData->vHeadDim; + chunkSize = gdnFwdHTilingData->chunkSize; + useInitialState = gdnFwdHTilingData->useInitialState; + storeFinalState = gdnFwdHTilingData->storeFinalState; + isVariedLen = gdnFwdHTilingData->isVariedLen; + shapeBatch = gdnFwdHTilingData->shapeBatch; + tokenBatch = gdnFwdHTilingData->tokenBatch; + vWorkspaceOffset = gdnFwdHTilingData->vWorkspaceOffset; + vUpdateWorkspaceOffset = gdnFwdHTilingData->vUpdateWorkspaceOffset; + hWorkspaceOffset = gdnFwdHTilingData->hWorkspaceOffset; + numSeqWorkspaceOffset = gdnFwdHTilingData->numSeqWorkspaceOffset; + numChunksWorkspaceOffset = gdnFwdHTilingData->numChunksWorkspaceOffset; + kDecayWorkspaceOffset = gdnFwdHTilingData->kDecayWorkspaceOffset; + + gmK.SetGlobalBuffer((__gm__ ElementK *)k); + gmW.SetGlobalBuffer((__gm__ ElementW *)w); + gmU.SetGlobalBuffer((__gm__ ElementU *)u); + gmG.SetGlobalBuffer((__gm__ ElementG *)(scalarGated ? g : gk)); + gmInitialState.SetGlobalBuffer((__gm__ ElementInitialState *)inital_state); + gmH.SetGlobalBuffer((__gm__ ElementH *)h); + gmV.SetGlobalBuffer((__gm__ ElementV *)v_new); + gmFinalState.SetGlobalBuffer((__gm__ ElementFinalState *)final_state); + gmVWorkspace.SetGlobalBuffer((__gm__ ElementVWork *)(user + vWorkspaceOffset)); + gmVUpdateWorkspace.SetGlobalBuffer((__gm__ ElementV *)(user + vUpdateWorkspaceOffset)); + gmHWorkspace.SetGlobalBuffer((__gm__ ElementHWork *)(user + hWorkspaceOffset)); + gmGk.SetGlobalBuffer((__gm__ ElementG *)(kGated ? gk : g)); + gmKDecayWorkspace.SetGlobalBuffer((__gm__ ElementK *)(user + kDecayWorkspaceOffset)); + + gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens); + gmNumSeq.SetGlobalBuffer((__gm__ int64_t *)(user + numSeqWorkspaceOffset)); + gmNumChunks.SetGlobalBuffer((__gm__ int64_t *)(user + numChunksWorkspaceOffset)); + + if ASCEND_IS_AIC { + cubeBlockScheduler.Init(cu_seqlens, chunk_indices, tiling, user); + } + + if ASCEND_IS_AIV { + vecBlockScheduler.Init(cu_seqlens, chunk_indices, tiling, user); + } + } + + template + __aicore__ inline void InitFromData( + GM_ADDR k, GM_ADDR w, GM_ADDR u, GM_ADDR g, GM_ADDR gk, GM_ADDR inital_state, + GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR h, GM_ADDR v_new, + GM_ADDR final_state, const TilingData& tilingData, GM_ADDR user) { + batch = tilingData.batch; + seqlen = tilingData.seqlen; + kNumHead = tilingData.kNumHead; + vNumHead = tilingData.vNumHead; + kHeadDim = tilingData.kHeadDim; + vHeadDim = tilingData.vHeadDim; + chunkSize = tilingData.chunkSize; + useInitialState = tilingData.useInitialState; + storeFinalState = tilingData.storeFinalState; + isVariedLen = tilingData.isVariedLen; + shapeBatch = tilingData.shapeBatch; + tokenBatch = tilingData.tokenBatch; + vWorkspaceOffset = tilingData.vWorkspaceOffset; + vUpdateWorkspaceOffset = tilingData.vUpdateWorkspaceOffset; + hWorkspaceOffset = tilingData.hWorkspaceOffset; + numSeqWorkspaceOffset = tilingData.numSeqWorkspaceOffset; + numChunksWorkspaceOffset = tilingData.numChunksWorkspaceOffset; + kDecayWorkspaceOffset = tilingData.kDecayWorkspaceOffset; + + gmK.SetGlobalBuffer((__gm__ ElementK *)k); + gmW.SetGlobalBuffer((__gm__ ElementW *)w); + gmU.SetGlobalBuffer((__gm__ ElementU *)u); + gmG.SetGlobalBuffer((__gm__ ElementG *)(scalarGated ? g : gk)); + gmInitialState.SetGlobalBuffer((__gm__ ElementInitialState *)inital_state); + gmH.SetGlobalBuffer((__gm__ ElementH *)h); + gmV.SetGlobalBuffer((__gm__ ElementV *)v_new); + gmFinalState.SetGlobalBuffer((__gm__ ElementFinalState *)final_state); + gmVWorkspace.SetGlobalBuffer((__gm__ ElementVWork *)(user + vWorkspaceOffset)); + gmVUpdateWorkspace.SetGlobalBuffer((__gm__ ElementV *)(user + vUpdateWorkspaceOffset)); + gmHWorkspace.SetGlobalBuffer((__gm__ ElementHWork *)(user + hWorkspaceOffset)); + gmGk.SetGlobalBuffer((__gm__ ElementG *)(kGated ? gk : g)); + gmKDecayWorkspace.SetGlobalBuffer((__gm__ ElementK *)(user + kDecayWorkspaceOffset)); + gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens); + gmNumSeq.SetGlobalBuffer((__gm__ int64_t *)(user + numSeqWorkspaceOffset)); + gmNumChunks.SetGlobalBuffer((__gm__ int64_t *)(user + numChunksWorkspaceOffset)); + + if ASCEND_IS_AIC { + cubeBlockScheduler.InitFromData(cu_seqlens, chunk_indices, tilingData, user); + } + if ASCEND_IS_AIV { + vecBlockScheduler.InitFromData(cu_seqlens, chunk_indices, tilingData, user); + } + } + + template + __aicore__ inline float LoadScalarAsFloat( + AscendC::GlobalTensor tensor, uint32_t offset) const + { + Element value = tensor.GetValue(offset); + if constexpr (std::is_same::value) { + return AscendC::ToFloat(value); + } + return static_cast(value); + } + + __aicore__ inline void ComputeTailVWorkspace( + const GDNFwdHOffsets& offsets, uint32_t tailEventId) + { + uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); + uint32_t subBlockNum = AscendC::GetSubBlockNum(); + uint32_t rowsPerSubBlock = CeilDiv(offsets.blockTokens, subBlockNum); + uint32_t rowBegin = subBlockIdx * rowsPerSubBlock; + uint32_t rowEnd = Min(rowBegin + rowsPerSubBlock, offsets.blockTokens); + if (rowBegin >= rowEnd) { + return; + } + AscendC::ResetMask(); + AscendC::WaitFlag(tailEventId); + + constexpr uint32_t TAIL_INPUT_OFFSET = 166 * 1024; + constexpr uint32_t TAIL_FLOAT_OFFSET = 167 * 1024; + constexpr uint32_t TAIL_ACCUM_OFFSET = 168 * 1024; + AscendC::LocalTensor inputUb = + resource.ubBuf.template GetBufferByByte(TAIL_INPUT_OFFSET); + AscendC::LocalTensor floatUb = + resource.ubBuf.template GetBufferByByte(TAIL_FLOAT_OFFSET); + AscendC::LocalTensor accumUb = + resource.ubBuf.template GetBufferByByte(TAIL_ACCUM_OFFSET); + AscendC::LocalTensor weightInputUb = + resource.ubBuf.template GetBufferByByte(169 * 1024); + AscendC::LocalTensor weightFloatUb = + resource.ubBuf.template GetBufferByByte(170 * 1024); + + for (uint32_t tokenRow = rowBegin; tokenRow < rowEnd; ++tokenRow) { + AscendC::DataCopy( + weightInputUb, + gmW[offsets.wOffset + tokenRow * kHeadDim], + kHeadDim); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::Cast( + weightFloatUb, weightInputUb, AscendC::RoundMode::CAST_NONE, + kHeadDim); + AscendC::PipeBarrier(); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + + AscendC::Duplicate(accumUb, 0.0f, offsets.vBlockDim); + AscendC::PipeBarrier(); + for (uint32_t kIdx = 0; kIdx < kHeadDim; ++kIdx) { + AscendC::DataCopy( + inputUb, gmH[offsets.hSrcOffset + kIdx * vHeadDim], + offsets.vBlockDim); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::Cast( + floatUb, inputUb, AscendC::RoundMode::CAST_NONE, + offsets.vBlockDim); + AscendC::PipeBarrier(); + float weight = weightFloatUb.GetValue(kIdx); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::Muls(floatUb, floatUb, weight, offsets.vBlockDim); + AscendC::PipeBarrier(); + AscendC::Add(accumUb, accumUb, floatUb, offsets.vBlockDim); + AscendC::PipeBarrier(); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + } + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::DataCopy( + gmVWorkspace[offsets.vWorkOffset + tokenRow * offsets.vBlockDim], + accumUb, offsets.vBlockDim); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + } + AscendC::SetFlag(tailEventId); + } + + __aicore__ inline void ComputeTailHWorkspace( + const GDNFwdHOffsets& offsets, uint32_t tailEventId) + { + uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); + uint32_t subBlockNum = AscendC::GetSubBlockNum(); + uint32_t rowsPerSubBlock = CeilDiv(kHeadDim, subBlockNum); + uint32_t rowBegin = subBlockIdx * rowsPerSubBlock; + uint32_t rowEnd = Min(rowBegin + rowsPerSubBlock, kHeadDim); + AscendC::ResetMask(); + AscendC::WaitFlag(tailEventId); + + constexpr uint32_t TAIL_INPUT_OFFSET = 166 * 1024; + constexpr uint32_t TAIL_FLOAT_OFFSET = 167 * 1024; + constexpr uint32_t TAIL_ACCUM_OFFSET = 168 * 1024; + AscendC::LocalTensor inputUb = + resource.ubBuf.template GetBufferByByte(TAIL_INPUT_OFFSET); + AscendC::LocalTensor floatUb = + resource.ubBuf.template GetBufferByByte(TAIL_FLOAT_OFFSET); + AscendC::LocalTensor accumUb = + resource.ubBuf.template GetBufferByByte(TAIL_ACCUM_OFFSET); + AscendC::LocalTensor weightInputUb = + resource.ubBuf.template GetBufferByByte(169 * 1024); + AscendC::LocalTensor weightFloatUb = + resource.ubBuf.template GetBufferByByte(170 * 1024); + + for (uint32_t kRow = rowBegin; kRow < rowEnd; ++kRow) { + AscendC::Duplicate(accumUb, 0.0f, offsets.vBlockDim); + AscendC::PipeBarrier(); + for (uint32_t tokenRow = 0; tokenRow < offsets.blockTokens; ++tokenRow) { + AscendC::DataCopy( + weightInputUb, + gmKDecayWorkspace[offsets.kDecayWorkOffset + tokenRow * kHeadDim], + kHeadDim); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::Cast( + weightFloatUb, weightInputUb, AscendC::RoundMode::CAST_NONE, + kHeadDim); + AscendC::PipeBarrier(); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::DataCopy( + inputUb, + gmVUpdateWorkspace[offsets.vWorkOffset + tokenRow * offsets.vBlockDim], + offsets.vBlockDim); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::Cast( + floatUb, inputUb, AscendC::RoundMode::CAST_NONE, + offsets.vBlockDim); + AscendC::PipeBarrier(); + float weight = weightFloatUb.GetValue(kRow); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::Muls(floatUb, floatUb, weight, offsets.vBlockDim); + AscendC::PipeBarrier(); + AscendC::Add(accumUb, accumUb, floatUb, offsets.vBlockDim); + AscendC::PipeBarrier(); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + } + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + AscendC::DataCopy( + gmHWorkspace[offsets.hWorkOffset + kRow * offsets.vBlockDim], + accumUb, offsets.vBlockDim); + AscendC::SetFlag(tailEventId); + AscendC::WaitFlag(tailEventId); + } + AscendC::SetFlag(tailEventId); + } + + __aicore__ inline void PresetVectorPipelineEvents() + { + constexpr uint32_t pongBaseEvent = 4; + if (storeFinalState && std::is_same::value) { + AscendC::SetFlag(EVENT_ID0); + AscendC::SetFlag(EVENT_ID0 + pongBaseEvent); + AscendC::SetFlag(EVENT_ID2); + AscendC::SetFlag(EVENT_ID2 + pongBaseEvent); + } else { + AscendC::SetFlag(EVENT_ID0); + AscendC::SetFlag(EVENT_ID0 + pongBaseEvent); + AscendC::SetFlag(EVENT_ID2); + AscendC::SetFlag(EVENT_ID2 + pongBaseEvent); + } + AscendC::SetFlag(EVENT_ID1); + AscendC::SetFlag(EVENT_ID1 + pongBaseEvent); + AscendC::SetFlag(EVENT_ID3); + AscendC::SetFlag(EVENT_ID3 + pongBaseEvent); + AscendC::SetFlag(EVENT_ID0); + AscendC::SetFlag(EVENT_ID0 + pongBaseEvent); + AscendC::SetFlag(EVENT_ID2); + AscendC::SetFlag(EVENT_ID2 + pongBaseEvent); + } + + __aicore__ inline void DrainVectorPipelineEvents( + const bool (&event0FromMte3)[PING_PONG_STAGES], + const bool (&event2FromMte3)[PING_PONG_STAGES]) + { + constexpr uint32_t pongBaseEvent = 4; + if (storeFinalState && std::is_same::value) { + for (uint32_t streamId = 0; streamId < PING_PONG_STAGES; ++streamId) { + uint32_t eventOffset = streamId * pongBaseEvent; + if (event0FromMte3[streamId]) { + AscendC::WaitFlag(EVENT_ID0 + eventOffset); + } else { + AscendC::WaitFlag(EVENT_ID0 + eventOffset); + } + if (event2FromMte3[streamId]) { + AscendC::WaitFlag(EVENT_ID2 + eventOffset); + } else { + AscendC::WaitFlag(EVENT_ID2 + eventOffset); + } + } + } else { + AscendC::WaitFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID0 + pongBaseEvent); + AscendC::WaitFlag(EVENT_ID2); + AscendC::WaitFlag(EVENT_ID2 + pongBaseEvent); + } + AscendC::WaitFlag(EVENT_ID1); + AscendC::WaitFlag(EVENT_ID1 + pongBaseEvent); + AscendC::WaitFlag(EVENT_ID3); + AscendC::WaitFlag(EVENT_ID3 + pongBaseEvent); + AscendC::WaitFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID0 + pongBaseEvent); + AscendC::WaitFlag(EVENT_ID2); + AscendC::WaitFlag(EVENT_ID2 + pongBaseEvent); + } + + __aicore__ inline void Process() { + if (isVariedLen) { + AscendC::SyncAll(); + } + + if ASCEND_IS_AIC { + uint32_t coreIdx = AscendC::GetBlockIdx(); + uint32_t coreNum = AscendC::GetBlockNum(); + + auto wLayout = tla::MakeLayout(shapeBatch * kNumHead * cubeBlockScheduler.totalTokens, kHeadDim); + auto hLayout = tla::MakeLayout(shapeBatch * vNumHead * cubeBlockScheduler.totalChunks * kHeadDim, vHeadDim); + auto vLayout = tla::MakeLayout(coreNum * chunkSize * PING_PONG_STAGES, cubeBlockScheduler.vBlockSize); + + auto kLayout = tla::MakeLayout(kHeadDim, shapeBatch * kNumHead * cubeBlockScheduler.totalTokens); + auto vworkLayout = tla::MakeLayout(coreNum * chunkSize * PING_PONG_STAGES, cubeBlockScheduler.vBlockSize); + auto hworkLayout = tla::MakeLayout(coreNum * kHeadDim * PING_PONG_STAGES, cubeBlockScheduler.vBlockSize); + uint32_t taskWaveCount = cubeBlockScheduler.GetTaskWaveCount(); + for (uint32_t waveIdx = 0; waveIdx < taskWaveCount; ++waveIdx) { + BlockMmadWH blockMmadWH(resource); + BlockMmadKV blockMmadKV(resource); + BlockMmadWH blockMmadWHTail(resource); + BlockMmadKV blockMmadKVTail(resource); + AscendC::SyncAll(); + cubeBlockScheduler.InitTaskWave(waveIdx); + uint32_t currStage = 0; // 0: C1, 1: C2 + while (cubeBlockScheduler.isRunning) { + if (currStage == 0) { + /* C1: v_work = w @ h[i] */ + cubeBlockScheduler.InitTasks(); + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = cubeBlockScheduler.GetStreamId(i); + const auto& stream = cubeBlockScheduler.GetStream(i); + if (cubeBlockScheduler.StreamIsDone(stream)) { + continue; + } + + const GDNFwdHOffsets& cube1Offsets = cubeBlockScheduler.GetCurTaskOffsets(stream); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[streamId]); + if (cube1Offsets.blockTokens < 16) { + Arch::CrossCoreSetFlag<0x2, PIPE_MTE2>( + cubeBlockScheduler.cube1Done[streamId]); + continue; + } + int64_t cube1OffsetW = cube1Offsets.wOffset; + int64_t cube1OffsetH = cube1Offsets.hSrcOffset; + int64_t cube1OffsetVwork = cube1Offsets.vWorkOffset; + auto tensorW = tla::MakeTensor(gmW[cube1OffsetW], wLayout, Catlass::Arch::PositionGM{}); + auto tensorH = tla::MakeTensor(gmH[cube1OffsetH], hLayout, Catlass::Arch::PositionGM{}); + auto tensorV = tla::MakeTensor(gmVWorkspace[cube1OffsetVwork], vLayout, Catlass::Arch::PositionGM{}); + GemmCoord cube1Shape {cube1Offsets.blockTokens, cube1Offsets.vBlockDim, kHeadDim}; + auto tensorBlockW = GetTile(tensorW, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k())); + auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n())); + auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n())); + if (cube1Offsets.blockTokens < chunkSize) { + blockMmadWHTail.preSetFlags(); + blockMmadWHTail( + tensorBlockW, tensorBlockH, tensorBlockV, + cube1Shape, EmptyClass{}, true); + blockMmadWHTail.finalWaitFlags(); + } else { + blockMmadWH.preSetFlags(); + blockMmadWH( + tensorBlockW, tensorBlockH, tensorBlockV, cube1Shape); + blockMmadWH.finalWaitFlags(); + } + AscendC::PipeBarrier(); + Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube1Done[streamId]); + } + } else { + /* C2: h[i+1] = k.T @ v_work */ + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = cubeBlockScheduler.GetStreamId(i); + const auto& stream = cubeBlockScheduler.GetStream(i); + if (cubeBlockScheduler.StreamIsDone(stream)) { + continue; + } + const GDNFwdHOffsets& cube2Offsets = cubeBlockScheduler.GetCurTaskOffsets(stream); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done[streamId]); + + if (cubeBlockScheduler.NeedProcessStage2(stream)) { + if (cube2Offsets.blockTokens < 16) { + Arch::CrossCoreSetFlag<0x2, PIPE_MTE2>( + cubeBlockScheduler.cube2Done[streamId]); + continue; + } + // step 3: h[i+1] = k.T @ v_work + int64_t cube2OffsetKwork = kGated ? cube2Offsets.kDecayWorkOffset : cube2Offsets.wkOffset; + int64_t cube2OffsetVwork = cube2Offsets.vWorkOffset; + int64_t cube2OffsetH = cube2Offsets.hWorkOffset; + auto tensorK = kGated + ? tla::MakeTensor(gmKDecayWorkspace[cube2OffsetKwork], kLayout, Catlass::Arch::PositionGM{}) + : tla::MakeTensor(gmK[cube2OffsetKwork], kLayout, Catlass::Arch::PositionGM{}); + auto tensorVwork = tla::MakeTensor(gmVUpdateWorkspace[cube2OffsetVwork], vworkLayout, Catlass::Arch::PositionGM{}); + auto tensorHwork = tla::MakeTensor(gmHWorkspace[cube2OffsetH], hworkLayout, Catlass::Arch::PositionGM{}); + GemmCoord cube2Shape{kHeadDim, cube2Offsets.vBlockDim, cube2Offsets.blockTokens}; + auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k())); + auto tensorBlockVwork = GetTile(tensorVwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n())); + auto tensorBlockHwork = GetTile(tensorHwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n())); + if (cube2Offsets.blockTokens < chunkSize) { + blockMmadKVTail.preSetFlags(); + blockMmadKVTail( + tensorBlockK, tensorBlockVwork, tensorBlockHwork, + cube2Shape, EmptyClass{}, true); + blockMmadKVTail.finalWaitFlags(); + } else { + blockMmadKV.preSetFlags(); + blockMmadKV( + tensorBlockK, tensorBlockVwork, tensorBlockHwork, + cube2Shape); + blockMmadKV.finalWaitFlags(); + } + AscendC::PipeBarrier(); + } + Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube2Done[streamId]); + } + } + currStage ^= 0x01; + } + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[0]); + Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done[1]); + } + + } + + if ASCEND_IS_AIV { + uint32_t subBlockIdx = AscendC::GetSubBlockIdx(); + uint32_t subBlockNum = AscendC::GetSubBlockNum(); + uint32_t coreIdx = AscendC::GetBlockIdx() / subBlockNum; + uint32_t coreNum = AscendC::GetBlockNum(); + uint32_t taskCount = + (isVariedLen ? vecBlockScheduler.tokenBatch : shapeBatch) * vNumHead; + uint32_t rowsPerSubBlock = (kHeadDim + subBlockNum - 1) / subBlockNum; + uint32_t rowBegin = subBlockIdx * rowsPerSubBlock; + uint32_t rowEnd = Min(rowBegin + rowsPerSubBlock, kHeadDim); + uint32_t hRowsPerTile = (32 * 1024) / (vHeadDim * sizeof(ElementH)); + uint32_t stateRowsPerTile = + (64 * 1024) / (vHeadDim * sizeof(ElementInitialState)); + uint32_t rowsPerTile = Min(hRowsPerTile, stateRowsPerTile); + uint32_t totalChunks = + isVariedLen ? vecBlockScheduler.totalChunks : ((seqlen + chunkSize - 1) / chunkSize); + uint32_t stateBlockSize = kHeadDim * vHeadDim; + AscendC::LocalTensor stateUbTensorPing = + resource.ubBuf.template GetBufferByByte(0); + AscendC::LocalTensor stateUbTensorPong = + resource.ubBuf.template GetBufferByByte(96 * 1024); + AscendC::LocalTensor hUbTensorPing = + resource.ubBuf.template GetBufferByByte(64 * 1024); + AscendC::LocalTensor hUbTensorPong = + resource.ubBuf.template GetBufferByByte(160 * 1024); + uint32_t taskWaveCount = vecBlockScheduler.GetTaskWaveCount(); + for (uint32_t waveIdx = 0; waveIdx < taskWaveCount; ++waveIdx) { + EpilogueGDNFwdHVnew epilogueGDNFwdHVnew(resource); + EpilogueGDNFwdHUpdate epilogueGDNFwdHUpdate(resource); + uint32_t taskIdx = waveIdx * coreNum + coreIdx; + uint32_t pingpongFlag = 1; + AscendC::SetFlag(EVENT_ID0); + AscendC::SetFlag(EVENT_ID1); + if (taskIdx < taskCount) { + uint32_t batchIdx = taskIdx / vNumHead; + uint32_t vHeadIdx = taskIdx % vNumHead; + uint32_t chunkOffset = + isVariedLen ? vecBlockScheduler.GetVarlenChunkOffset(batchIdx) : 0; + uint32_t shapeBatchIdx = isVariedLen ? 0 : batchIdx; + uint32_t hBaseOffset = + (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset) * + stateBlockSize; + uint32_t initialStateBaseOffset = taskIdx * stateBlockSize; + for (uint32_t rowOffset = rowBegin; rowOffset < rowEnd; rowOffset += rowsPerTile) { + uint32_t rowsThisTile = Min(rowsPerTile, rowEnd - rowOffset); + uint32_t stateTileElems = rowsThisTile * vHeadDim; + uint32_t hOffset = hBaseOffset + rowOffset * vHeadDim; + AscendC::LocalTensor stateUbTensor = + pingpongFlag ? stateUbTensorPing : stateUbTensorPong; + AscendC::LocalTensor hUbTensor = + pingpongFlag ? hUbTensorPing : hUbTensorPong; + auto eventId = pingpongFlag ? EVENT_ID1 : EVENT_ID0; + AscendC::WaitFlag(eventId); + if (useInitialState) { + uint32_t initialStateOffset = + initialStateBaseOffset + rowOffset * vHeadDim; + if constexpr (!std::is_same::value) { + AscendC::DataCopy( + stateUbTensor, gmInitialState[initialStateOffset], stateTileElems); + AscendC::SetFlag(eventId); + AscendC::WaitFlag(eventId); + AscendC::Cast( + hUbTensor, stateUbTensor, AscendC::RoundMode::CAST_RINT, + stateTileElems); + AscendC::SetFlag(eventId); + AscendC::WaitFlag(eventId); + AscendC::DataCopy(gmH[hOffset], hUbTensor, stateTileElems); + } else { + AscendC::DataCopy( + stateUbTensor, gmInitialState[initialStateOffset], stateTileElems); + AscendC::SetFlag(eventId); + AscendC::WaitFlag(eventId); + AscendC::DataCopy(gmH[hOffset], stateUbTensor, stateTileElems); + } + } else { + AscendC::Duplicate(hUbTensor, static_cast(0), stateTileElems); + AscendC::SetFlag(eventId); + AscendC::WaitFlag(eventId); + AscendC::DataCopy(gmH[hOffset], hUbTensor, stateTileElems); + } + AscendC::SetFlag(eventId); + pingpongFlag = 1 - pingpongFlag; + } + } + AscendC::WaitFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID1); + + AscendC::SyncAll(); + vecBlockScheduler.InitTaskWave(waveIdx); + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done[0]); + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done[1]); + PresetVectorPipelineEvents(); + uint32_t currStage = 0; // 0: V1, 1: V2 + bool waitStageFence = false; + bool event0FromMte3[PING_PONG_STAGES] = {false, false}; + bool event2FromMte3[PING_PONG_STAGES] = { + !(storeFinalState && std::is_same::value), + !(storeFinalState && std::is_same::value)}; + while (vecBlockScheduler.isRunning) { + if (waitStageFence) { + AscendC::WaitFlag(EVENT_ID1); + } + if (currStage == 0) { + /* V1: + * gmV = gmU - gmVWorkspace + * g_buf = gmG[-1] - gmG + * g_buf = exp(g_buf) + * gmVWorkspace = g_buf * gmV + */ + vecBlockScheduler.InitTasks(); + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = vecBlockScheduler.GetStreamId(i); + const auto& stream = vecBlockScheduler.GetStream(i); + if (vecBlockScheduler.StreamIsDone(stream)) { + continue; + } + const GDNFwdHOffsets& vec1Offsets = vecBlockScheduler.GetCurTaskOffsets(stream); + bool tailVectorPath = vec1Offsets.blockTokens < 16; + if (tailVectorPath) { + Arch::CrossCoreWaitFlag( + vecBlockScheduler.cube1Done[streamId]); + ComputeTailVWorkspace( + vec1Offsets, + EVENT_ID3 + (streamId == 0 ? 0 : 4)); + } + bool waitWsFromMte3 = storeFinalState && std::is_same::value && + event0FromMte3[streamId]; + epilogueGDNFwdHVnew( + gmV[vec1Offsets.uvOffset], gmVUpdateWorkspace[vec1Offsets.vWorkOffset], + gmG[vec1Offsets.gOffset], gmU[vec1Offsets.uvOffset], gmVWorkspace[vec1Offsets.vWorkOffset], + gmGk[vec1Offsets.gkOffset], gmK[vec1Offsets.wkOffset], gmKDecayWorkspace[vec1Offsets.kDecayWorkOffset], + vec1Offsets.blockTokens, kHeadDim, vec1Offsets.vBlockDim, vHeadDim, + vecBlockScheduler.cube1Done[streamId], vecBlockScheduler.vec1Done[streamId], + vec1Offsets.isInitialState, vec1Offsets.isFinalState, storeFinalState, + waitWsFromMte3, (streamId == 0), tailVectorPath + ); + AscendC::SetFlag(EVENT_ID1); + AscendC::WaitFlag(EVENT_ID1); + if (storeFinalState && std::is_same::value) { + event0FromMte3[streamId] = false; + } + } + } else { + /* V2: h[i+1] += h_work if i < num_chunks - 1 else None */ + for (uint32_t i = 0; i < PING_PONG_STAGES; ++i) { + uint32_t streamId = vecBlockScheduler.GetStreamId(i); + const auto& stream = vecBlockScheduler.GetStream(i); + if (vecBlockScheduler.StreamIsDone(stream)) { + continue; + } + const GDNFwdHOffsets& vec2Offsets = vecBlockScheduler.GetCurTaskOffsets(stream); + if (vecBlockScheduler.NeedProcessStage2(stream)) { + bool tailVectorPath = vec2Offsets.blockTokens < 16; + if (tailVectorPath) { + Arch::CrossCoreWaitFlag( + vecBlockScheduler.cube2Done[streamId]); + ComputeTailHWorkspace( + vec2Offsets, + EVENT_ID3 + (streamId == 0 ? 0 : 4)); + } + if (storeFinalState && std::is_same::value) { + event0FromMte3[streamId] = true; + event2FromMte3[streamId] = !vec2Offsets.isFinalState; + } + // step 4: h[i+1] += h_work if i < num_chunks - 1 else None + epilogueGDNFwdHUpdate( + gmH[vec2Offsets.hDstOffset], gmFinalState[vec2Offsets.finalStateOffset], + gmG[vec2Offsets.gOffset], + gmH[vec2Offsets.hSrcOffset], + gmHWorkspace[vec2Offsets.hWorkOffset], + gmGk[vec2Offsets.gkOffset], + gmInitialState[vec2Offsets.initialStateOffset], + vec2Offsets.blockTokens, kHeadDim, vec2Offsets.vBlockDim, vHeadDim, vecBlockScheduler.cube2Done[streamId], + vec2Offsets.isInitialState, vec2Offsets.isFinalState, storeFinalState, + useInitialState, (streamId == 0), tailVectorPath + ); + AscendC::SetFlag(EVENT_ID1); + AscendC::WaitFlag(EVENT_ID1); + } else { + Arch::CrossCoreWaitFlag(vecBlockScheduler.cube2Done[streamId]); + } + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done[streamId]); + } + } + waitStageFence = vecBlockScheduler.isRunning; + if (waitStageFence) { + AscendC::SetFlag(EVENT_ID1); + } + currStage ^= 0x01; + } + + DrainVectorPipelineEvents(event0FromMte3, event2FromMte3); + } + + } + } + +}; + +} diff --git a/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla.hpp b/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla.hpp new file mode 100644 index 000000000000..803ceddf17b2 --- /dev/null +++ b/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla.hpp @@ -0,0 +1,1054 @@ +/** + * Copyright (c) 2025-2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef COMMON_BLOCK_MMAD_PINGPONG_TLA_HPP +#define COMMON_BLOCK_MMAD_PINGPONG_TLA_HPP + +#include "catlass/catlass.hpp" +#include "catlass/arch/resource.hpp" +#include "catlass/coord.hpp" +#include "catlass/gemm_coord.hpp" +#include "catlass/gemm/dispatch_policy.hpp" +#include "catlass/gemm/helper.hpp" +#include "catlass/gemm/tile/tile_copy.hpp" +#include "catlass/gemm/tile/tile_mmad.hpp" +#include "kernel_utils/tile/copy_l0c_to_ub.hpp" +#include "tla/layout.hpp" +#include "tla/tensor.hpp" + +namespace Common { + +template < + class DispatchPolicy, + class L1TileShape, + class L0TileShape, + class ElementA, + class ElementB, + class ElementC, + class ElementBias = void, + class TileCopy = Catlass::Gemm::Tile::PackedTileCopyTla, + class TileMmad = + Catlass::Gemm::Tile::TileMmadTla +> +struct BlockMmadTla { + static_assert(DEPENDENT_FALSE, "BlockMmadTla is not implemented for this DispatchPolicy"); +}; + +// Now ENABLE_UNIT_FLAG_ must be false when input element is int8 +template +struct MmadPingpong : public Catlass::Gemm::MmadBase { + static constexpr uint32_t L1A_STAGES = L1A_STAGES_; + static constexpr uint32_t L1B_STAGES = L1B_STAGES_; + static constexpr uint32_t L0A_STAGES = L0A_STAGES_; + static constexpr uint32_t L0B_STAGES = L0B_STAGES_; + static constexpr uint32_t L0C_STAGES = L0C_STAGES_; + static constexpr uint32_t UB_STAGES = UB_STAGES_; + static constexpr bool ENABLE_UNIT_FLAG = ENABLE_UNIT_FLAG_; + static constexpr bool USE_HF32_MODE = USE_HF32_MODE_; + static constexpr bool ENABLE_L1_RESIDENT = ENABLE_L1_RESIDENT_; +}; + +template < + class ArchTag_, + bool ENABLE_UNIT_FLAG_, + bool USE_HF32_MODE_, + uint32_t L0C_STAGES_, + bool ENABLE_L1_RESIDENT_, + uint32_t L1A_STAGES_, + uint32_t L1B_STAGES_, + uint32_t L0A_STAGES_, + uint32_t L0B_STAGES_, + uint32_t UB_STAGES_, + class L1TileShape_, + class L0TileShape_, + class ElementA_, + class ElementB_, + class ElementC_, + class ElementBias_, + class TileCopy_, + class TileMmad_ +> +struct BlockMmadTla < + MmadPingpong, + L1TileShape_, + L0TileShape_, + ElementA_, + ElementB_, + ElementC_, + ElementBias_, + TileCopy_, + TileMmad_ +> { +public: + // Type Aliases + using DispatchPolicy = MmadPingpong; + using ArchTag = typename DispatchPolicy::ArchTag; + using TileCopy = TileCopy_; + using L1TileShape = L1TileShape_; + using L0TileShape = L0TileShape_; + using ElementA = ElementA_; + using LayoutA = typename TileCopy::LayoutA; + using ElementB = ElementB_; + using LayoutB = typename TileCopy::LayoutB; + using ElementC = ElementC_; + using LayoutC = typename TileCopy::LayoutC; + using ElementBias = ElementBias_; + + using TileMmad = TileMmad_; + + using CopyL1ToL0A = typename TileCopy::CopyL1ToL0A; + using CopyL1ToL0B = typename TileCopy::CopyL1ToL0B; + using CopyL1ToBT = typename TileCopy::CopyL1ToBT; + + using ElementAccumulator = typename TileCopy::ElementAccumulator; + + static constexpr bool HAS_BIAS = TileCopy::HAS_BIAS; + + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; + using LayoutTagL0A = typename TileCopy::LayoutTagL0A; + using LayoutTagL0B = typename TileCopy::LayoutTagL0B; + + using L1AAlignHelper = typename TileCopy_::L1AAlignHelper; + using L1BAlignHelper = typename TileCopy_::L1BAlignHelper; + + static_assert(tla::is_tuple::value && tla::is_static::value, + "L1TileShape must be tla::tuple and static!"); + static_assert(tla::is_tuple::value && tla::is_static::value, + "L0TileShape must be tla::tuple and static!"); + + static constexpr uint64_t FLAG_ID_MAX = 16; + static constexpr bool ENABLE_UNIT_FLAG = DispatchPolicy::ENABLE_UNIT_FLAG; + static constexpr bool USE_HF32_MODE = DispatchPolicy::USE_HF32_MODE; + static constexpr bool ENABLE_L1_RESIDENT = DispatchPolicy::ENABLE_L1_RESIDENT; + static constexpr uint32_t L1A_STAGES = DispatchPolicy::L1A_STAGES; + static constexpr uint32_t L1B_STAGES = DispatchPolicy::L1B_STAGES; + static constexpr uint32_t L0A_STAGES = DispatchPolicy::L0A_STAGES; + static constexpr uint32_t L0B_STAGES = DispatchPolicy::L0B_STAGES; + static constexpr uint32_t L0C_STAGES = DispatchPolicy::L0C_STAGES; + static constexpr uint32_t L1_TILE_M = tla::get<0>(L1TileShape{}); + static constexpr uint32_t L1_TILE_N = tla::get<1>(L1TileShape{}); + static constexpr uint32_t L1_TILE_K = tla::get<2>(L1TileShape{}); + static constexpr uint32_t L0_TILE_M = tla::get<0>(L0TileShape{}); + static constexpr uint32_t L0_TILE_N = tla::get<1>(L0TileShape{}); + static constexpr uint32_t L0_TILE_K = tla::get<2>(L0TileShape{}); + static constexpr uint32_t UB_STAGES = UB_STAGES_; + static constexpr uint32_t MAX_CUBE_VEC_SYNC_NUM = 5; + + // L1 tile size + static constexpr uint32_t L1A_TILE_SIZE = L1_TILE_M * L1_TILE_K * sizeof(ElementA); + static constexpr uint32_t L1B_TILE_SIZE = L1_TILE_N * L1_TILE_K * sizeof(ElementB); + // L0 tile size + static constexpr uint32_t L0A_TILE_SIZE = L0_TILE_M * L0_TILE_K * sizeof(ElementA); + static constexpr uint32_t L0B_TILE_SIZE = L0_TILE_K * L0_TILE_N * sizeof(ElementB); + static constexpr uint32_t L0C_TILE_SIZE = L1_TILE_M * L1_TILE_N * sizeof(ElementAccumulator); + + // Check HF32_MODE + static_assert( + !USE_HF32_MODE || (USE_HF32_MODE && std::is_same_v && std::is_same_v), + "HF32 MODE only supports in float!" + ); + + // Check L0C_STAGES + static_assert(!(ENABLE_UNIT_FLAG && L0C_STAGES != 1), "L0C_STAGES must be 1 when UnitFlag is true!"); + + // Check LayoutC + static_assert(tla::detail::isRowMajor::value || + ((std::is_same_v || std::is_same_v || + std::is_same_v) && tla::detail::iszN::value), + "LayoutC only supports zN in half or bfloat16 or float, RowMajor in all dtype yet!"); + + // Check L1TileShape + static_assert(L1A_TILE_SIZE * L1A_STAGES + L1B_TILE_SIZE * L1B_STAGES <= ArchTag::L1_SIZE, + "L1TileShape exceeding the L1 space!"); + + // Check L0TileShape + static_assert(L0A_TILE_SIZE * L0A_STAGES <= ArchTag::L0A_SIZE, "L0TileShape exceeding the L0A space!"); + static_assert(L0B_TILE_SIZE * L0B_STAGES <= ArchTag::L0B_SIZE, "L0TileShape exceeding the L0B space!"); + static_assert(L0C_TILE_SIZE * L0C_STAGES <= ArchTag::L0C_SIZE, "L0TileShape exceeding the L0C space!"); + + static constexpr uint32_t _32B = 32*8; // in bits + static_assert(L1_TILE_M == L0_TILE_M && L1_TILE_N == L0_TILE_N, + "The situation where the basic blocks of L1 and L0 differ on the m and n axes is not supported yet"); + static_assert(L0_TILE_K <= L1_TILE_K, "L0TileShape::K cannot exceed L1TileShape::K"); +#if (defined (CATLASS_ARCH) && CATLASS_ARCH == 2201) + static_assert(L1_TILE_M * SizeOfBits::value % _32B == 0, "L1TileShape::M must be 32B aligned."); + static_assert(L1_TILE_K * SizeOfBits::value % _32B == 0, "L1TileShape::K must be 32B aligned."); + static_assert(L1_TILE_K * SizeOfBits::value % _32B == 0, "L1TileShape::K must be 32B aligned."); + static_assert(L1_TILE_N * SizeOfBits::value % _32B == 0, "L1TileShape::N must be 32B aligned."); + static_assert(L0_TILE_K * SizeOfBits::value % _32B == 0, "L0TileShape::K must be 32B aligned."); +#endif + + static_assert((!HAS_BIAS && (L1A_STAGES + L1B_STAGES) <= 8) || (HAS_BIAS && (L1A_STAGES + L1B_STAGES) <= 7), + "L1 Buffer overflow: Exceeds the supported range of EVENT(0~7)"); + + static_assert((!HAS_BIAS && (L0A_STAGES + L0B_STAGES) <= 8) || (HAS_BIAS && (L0A_STAGES + L0B_STAGES) <= 7), + "L0 Buffer overflow: Exceeds the supported range of EVENT_ID(0~7)"); + + static constexpr auto L1A_LAYOUT = + tla::MakeLayout(tla::Int{}, tla::Int{}); + static constexpr auto L1B_LAYOUT = + tla::MakeLayout(tla::Int{}, tla::Int{}); + static constexpr auto L1BIAS_LAYOUT = tla::MakeLayout(tla::Int{}); + static constexpr auto L0BIAS_LAYOUT = tla::MakeLayout(tla::Int{}); + + // When enabling L1 resident mode, restore the pointer and coordinates that record the last state + // to the initial state. if two blockmmad instances need to be consecutively invoked at the kernel layer, + // RestoreStatus() must be inserted between them. + CATLASS_DEVICE + void RestoreStatus() + { + for (int i = 0; i < L1A_STAGES; ++i) { + lastAddrA[i] = nullptr; + lastCoordA[i] = Catlass::MatrixCoord{0U, 0U}; + } + for (int i = 0; i < L1B_STAGES; ++i) { + lastAddrB[i] = nullptr; + lastCoordB[i] = Catlass::MatrixCoord{0U, 0U}; + } + } + + /// Construct + CATLASS_DEVICE + BlockMmadTla(Catlass::Arch::Resource &resource, uint32_t l1BufAddrStart = 0) + { + if ASCEND_IS_AIC { + // use HF32 when USE_HF32_MODE is true + if constexpr (USE_HF32_MODE) { + AscendC::SetHF32Mode(true); + } else { + AscendC::SetHF32Mode(false); + } + if constexpr (ENABLE_UNIT_FLAG && tla::detail::isRowMajor::value) { + AscendC::SetMMLayoutTransform(true); + } + uint32_t l1AOffset = l1BufAddrStart; + uint32_t l1BOffset = l1BufAddrStart + L1A_TILE_SIZE * L1A_STAGES; + // Init buffers + for (uint32_t i = 0; i < L1A_STAGES; i++) { + // Assign L1/L0A/L0B space for each stages + l1ATensorList[i] = resource.l1Buf.template GetBufferByByte(l1AOffset + L1A_TILE_SIZE * i); + // Assign event ID for each stages + l1AEventList[i] = i; + // The event id that needs to be set before the loop + AscendC::SetFlag(l1AEventList[i]); + } + for (uint32_t i = 0; i < L1B_STAGES; i++) { + // Assign L1/L0A/L0B space for each stages + l1BTensorList[i] = resource.l1Buf.template GetBufferByByte(l1BOffset + L1B_TILE_SIZE * i); + // Assign event ID for each stages + l1BEventList[i] = i + L1A_STAGES; + // The event id that needs to be set before the loop + AscendC::SetFlag(l1BEventList[i]); + } + for (uint32_t i = 0; i < L0A_STAGES; i++) { + // Assign L1/L0A/L0B space for each stages + l0ATensorList[i] = resource.l0ABuf.template GetBufferByByte(L0A_TILE_SIZE * i); + // Assign event ID for each stages + l0AEventList[i] = i; + // The event id that needs to be set before the loop + AscendC::SetFlag(l0AEventList[i]); + } + for (uint32_t i = 0; i < L0B_STAGES; i++) { + // Assign L1/L0A/L0B space for each stages + l0BTensorList[i] = resource.l0BBuf.template GetBufferByByte(L0B_TILE_SIZE * i); + // Assign event ID for each stages + l0BEventList[i] = i + L0A_STAGES; + // The event id that needs to be set before the loop + AscendC::SetFlag(l0BEventList[i]); + } + if constexpr(!ENABLE_UNIT_FLAG) { + for (uint32_t i = 0; i < L0C_STAGES; i++) { + l0CTensorList[i] = resource.l0CBuf.template GetBufferByByte(L0C_TILE_SIZE * i); + l0CEventList[i] = i; + AscendC::SetFlag(l0CEventList[i]); + } + } else { + l0CTensorList[0] = resource.l0CBuf.template GetBufferByByte(0); + } + if constexpr (HAS_BIAS) { + uint32_t l1BiasOffset = l1BOffset + L1B_TILE_SIZE * L1B_STAGES; + l1BiasTensor = resource.l1Buf.template GetBufferByByte(l1BiasOffset); + l0BiasTensor = resource.btBuf.template GetBufferByByte(0); + AscendC::SetFlag(L1A_STAGES + L1B_STAGES); + AscendC::SetFlag(L0A_STAGES + L0B_STAGES); + } + + if constexpr (ENABLE_L1_RESIDENT) { + RestoreStatus(); + } + } + } + + /// Destructor + CATLASS_DEVICE + ~BlockMmadTla() + { + if ASCEND_IS_AIC { + if constexpr (USE_HF32_MODE) { + AscendC::SetHF32Mode(false); + } + if constexpr (ENABLE_UNIT_FLAG && tla::detail::isRowMajor::value) { + AscendC::SetMMLayoutTransform(false); + } + for (uint32_t i = 0; i < L1A_STAGES; i++) { + AscendC::WaitFlag(l1AEventList[i]); + } + for (uint32_t i = 0; i < L1B_STAGES; i++) { + AscendC::WaitFlag(l1BEventList[i]); + } + for (uint32_t i = 0; i < L0A_STAGES; i++) { + AscendC::WaitFlag(l0AEventList[i]); + } + for (uint32_t i = 0; i < L0B_STAGES; i++) { + AscendC::WaitFlag(l0BEventList[i]); + } + if constexpr(!ENABLE_UNIT_FLAG) { + for (uint32_t i = 0; i < L0C_STAGES; i++) { + AscendC::WaitFlag(l0CEventList[i]); + } + } + if constexpr (HAS_BIAS) { + AscendC::WaitFlag(L1A_STAGES + L1B_STAGES); + AscendC::WaitFlag(L0A_STAGES + L0B_STAGES); + } + } + } + + /// Perform a block-scoped matrix multiply-accumulate + template + CATLASS_DEVICE void operator()(TensorA &tensorA, TensorB &tensorB, TensorC &tensorC, Catlass::GemmCoord const &actualShape, + TensorBias const &tensorBias = {}) + { + // Check L1TileShape + if constexpr (HAS_BIAS) { + static constexpr uint32_t BIAS_BUF_SIZE = L0_TILE_N * sizeof(ElementAccumulator); + static constexpr uint32_t L1BIAS_SIZE = L1_TILE_N * sizeof(ElementBias); + static_assert(BIAS_BUF_SIZE <= ArchTag::BIAS_SIZE, + "BIAS_BUF_SIZE exceeding the BT space! Reduce L0_TILE_N"); + static_assert(L1A_TILE_SIZE * L1A_STAGES + L1B_TILE_SIZE * L1B_STAGES + L1BIAS_SIZE <= ArchTag::L1_SIZE, + "L1TileShape exceeding the L1 space!"); + } + + using CopyGmToL1A = typename TileCopy_::template CopyGmToL1A; + using CopyGmToL1B = typename TileCopy_::template CopyGmToL1B; + CopyGmToL1A copyGmToL1A; + CopyGmToL1B copyGmToL1B; +#if (defined (CATLASS_ARCH) && CATLASS_ARCH == 2201) + using CopyL0CToGm = typename TileCopy_::template CopyL0CToGm; + CopyL0CToGm copyL0CToDst; +#endif +#if (defined (CATLASS_ARCH) && CATLASS_ARCH == 3510) + using CopyL0CToDst = typename TileCopy_::template CopyL0CToDst; + CopyL0CToDst copyL0CToDst; +#endif + + uint32_t mBlockActual = actualShape.m(); + uint32_t kBlockActual = actualShape.k(); + uint32_t nBlockActual = actualShape.n(); + + uint32_t mL1Actual = mBlockActual; + if constexpr (std::is_same_v) { + // Avoid using the gemv mode in mmad + if (mL1Actual == 1) { + mL1Actual = 16; + } + } + uint32_t nL1Actual = nBlockActual; + + auto layoutInL0C = tla::MakeLayoutL0C(mL1Actual, nL1Actual); + auto tensorL0C = tla::MakeTensor(l0CTensorList[l0CListId], layoutInL0C, Catlass::Arch::PositionL0C{}); + auto tensorL0Bias = tla::MakeTensor(l0BiasTensor, L0BIAS_LAYOUT, Catlass::Arch::PositionBias{}); + + uint32_t kL1Actual = min(kBlockActual, L1_TILE_K); + // load first matrix A tile from GM to L1 + AscendC::WaitFlag(l1AEventList[l1AListId]); + auto tensorL1A = tla::MakeTensor(l1ATensorList[l1AListId], L1A_LAYOUT, Catlass::Arch::PositionL1{}); + auto tensorTileA = GetTileA(tensorA, 0, 0, mBlockActual, kL1Actual); + if constexpr (ENABLE_L1_RESIDENT) { + // If the currently loaded GM pointer and block coordinates are the same as the last loaded ones, + // skip this loading. + if (lastAddrA[l1AListId] != tensorTileA.data().GetPhyAddr() + || tla::get<0>(tensorTileA.coord()) != lastCoordA[l1AListId].row() + || tla::get<1>(tensorTileA.coord()) != lastCoordA[l1AListId].column()) { + copyGmToL1A(tensorL1A, tensorTileA); + lastCoordA[l1AListId] = Catlass::MatrixCoord{tla::get<0>(tensorTileA.coord()), tla::get<1>(tensorTileA.coord())}; + lastAddrA[l1AListId] = const_cast<__gm__ typename AscendC::GlobalTensor::PrimType *>( + tensorTileA.data().GetPhyAddr() + ); + } + } else { + copyGmToL1A(tensorL1A, tensorTileA); + } + AscendC::SetFlag(l1AEventList[l1AListId]); + + // load first matrix B tile from GM to L1 + AscendC::WaitFlag(l1BEventList[l1BListId]); + auto tensorL1B = tla::MakeTensor(l1BTensorList[l1BListId], L1B_LAYOUT, Catlass::Arch::PositionL1{}); + auto tensorTileB = GetTile(tensorB, tla::MakeCoord(0, 0), tla::MakeShape(kL1Actual, nBlockActual)); + if constexpr (ENABLE_L1_RESIDENT) { + if (lastAddrB[l1BListId] != tensorTileB.data().GetPhyAddr() + || tla::get<0>(tensorTileB.coord()) != lastCoordB[l1BListId].row() + || tla::get<1>(tensorTileB.coord()) != lastCoordB[l1BListId].column()) { + copyGmToL1B(tensorL1B, tensorTileB); + lastCoordB[l1BListId] = Catlass::MatrixCoord{tla::get<0>(tensorTileB.coord()), tla::get<1>(tensorTileB.coord())}; + lastAddrB[l1BListId] = const_cast<__gm__ typename AscendC::GlobalTensor::PrimType *>( + tensorTileB.data().GetPhyAddr() + ); + } + } else { + copyGmToL1B(tensorL1B, tensorTileB); + } + AscendC::SetFlag(l1BEventList[l1BListId]); + + if constexpr (HAS_BIAS && !std::is_same_v) { + using CopyGmToL1Bias = typename TileCopy::template CopyGmToL1Bias; + CopyGmToL1Bias copyGmToL1Bias; + AscendC::WaitFlag(L1A_STAGES + L1B_STAGES); + auto l1Bias = l1BiasTensor.template ReinterpretCast(); + auto tensorL1Bias = tla::MakeTensor(l1Bias, L1BIAS_LAYOUT, Catlass::Arch::PositionL1{}); + copyGmToL1Bias(tensorL1Bias, tensorBias); + AscendC::SetFlag(L1A_STAGES + L1B_STAGES); + } + + if constexpr (!ENABLE_UNIT_FLAG) { + AscendC::WaitFlag(l0CEventList[l0CListId]); + } + + uint32_t mL0Loop = CeilDiv(mL1Actual); + uint32_t nL0Loop = CeilDiv(nL1Actual); + + // main loop + uint32_t kL1Loop = CeilDiv(kBlockActual); + for (uint32_t kL1Idx = 0; kL1Idx < kL1Loop; kL1Idx++) { + uint32_t l1AListIdNext = (l1AListId + 1 < L1A_STAGES) ? (l1AListId + 1) : 0; + uint32_t l1BListIdNext = (l1BListId + 1 < L1B_STAGES) ? (l1BListId + 1) : 0; + uint32_t kL1ActualNext{0}; + // preload next tile from GM to L1 + if (kL1Idx < kL1Loop - 1) { + uint32_t kL1IdxNext = kL1Idx + 1; + kL1ActualNext = (kL1IdxNext < kL1Loop - 1) ? L1_TILE_K : (kBlockActual - kL1IdxNext * L1_TILE_K); + + // Get L1 tensor for next stage + auto l1ATensor = l1ATensorList[l1AListIdNext]; + auto l1BTensor = l1BTensorList[l1BListIdNext]; + auto tensorL1A = tla::MakeTensor(l1ATensor, L1A_LAYOUT, Catlass::Arch::PositionL1{}); + auto tensorL1B = tla::MakeTensor(l1BTensor, L1B_LAYOUT, Catlass::Arch::PositionL1{}); + // Get GM tile for next stage + auto tensorTileA = GetTileA(tensorA, 0, kL1IdxNext * L1_TILE_K, mBlockActual, kL1ActualNext); + auto tensorTileB = GetTile(tensorB, tla::MakeCoord(kL1IdxNext * L1_TILE_K, 0), + tla::MakeShape(kL1ActualNext, nBlockActual)); + + // load next matrix A tile from GM to L1 + AscendC::WaitFlag(l1AEventList[l1AListIdNext]); + if constexpr (ENABLE_L1_RESIDENT) { + if (lastAddrA[l1AListIdNext] != tensorTileA.data().GetPhyAddr() + || tla::get<0>(tensorTileA.coord()) != lastCoordA[l1AListIdNext].row() + || tla::get<1>(tensorTileA.coord()) != lastCoordA[l1AListIdNext].column()) { + copyGmToL1A(tensorL1A, tensorTileA); + lastCoordA[l1AListIdNext] = + Catlass::MatrixCoord{tla::get<0>(tensorTileA.coord()), tla::get<1>(tensorTileA.coord())}; + lastAddrA[l1AListIdNext] = + const_cast<__gm__ typename AscendC::GlobalTensor::PrimType *>( + tensorTileA.data().GetPhyAddr() + ); + } + } else { + copyGmToL1A(tensorL1A, tensorTileA); + } + AscendC::SetFlag(l1AEventList[l1AListIdNext]); + + // load next matrix B tile from GM to L1 + AscendC::WaitFlag(l1BEventList[l1BListIdNext]); + if constexpr (ENABLE_L1_RESIDENT) { + if (lastAddrB[l1BListIdNext] != tensorTileB.data().GetPhyAddr() + || tla::get<0>(tensorTileB.coord()) != lastCoordB[l1BListIdNext].row() + || tla::get<1>(tensorTileB.coord()) != lastCoordB[l1BListIdNext].column()) { + copyGmToL1B(tensorL1B, tensorTileB); + lastCoordB[l1BListIdNext] = + Catlass::MatrixCoord{tla::get<0>(tensorTileB.coord()), tla::get<1>(tensorTileB.coord())}; + lastAddrB[l1BListIdNext] = + const_cast<__gm__ typename AscendC::GlobalTensor::PrimType *>( + tensorTileB.data().GetPhyAddr() + ); + } + } else { + copyGmToL1B(tensorL1B, tensorTileB); + } + AscendC::SetFlag(l1BEventList[l1BListIdNext]); + } + + // Get L1 tensor for current stage + auto l1ATensor = l1ATensorList[l1AListId]; + auto l1BTensor = l1BTensorList[l1BListId]; + tensorL1A = tla::MakeTensor(l1ATensor, L1A_LAYOUT, Catlass::Arch::PositionL1{}); + tensorL1B = tla::MakeTensor(l1BTensor, L1B_LAYOUT, Catlass::Arch::PositionL1{}); + // Get the loop nums on L0 + uint32_t kL0Loop = CeilDiv(kL1Actual); + + for (int mL0Idx = 0; mL0Idx < mL0Loop; mL0Idx++) { + uint32_t mL0Actual = (mL0Idx < mL0Loop - 1) ? L0_TILE_M : (mL1Actual - mL0Idx * L0_TILE_M); + + for (int kL0Idx = 0; kL0Idx < kL0Loop; kL0Idx++) { + uint32_t kL0Actual = (kL0Idx < kL0Loop - 1) ? L0_TILE_K : (kL1Actual - kL0Idx * L0_TILE_K); + + // Locate the current tile on L0A + auto l0ATile = l0ATensorList[l0AListId]; + auto layoutAInL0 = tla::MakeLayout(mL0Actual, kL0Actual); + auto tensorL0A = tla::MakeTensor(l0ATile, layoutAInL0, Catlass::Arch::PositionL0A{}); + // Locate the current tile of matrix A on L1 + auto tensorTileL1A = GetTileA(tensorL1A, mL0Idx * L0_TILE_M, kL0Idx * L0_TILE_K, mL0Actual, kL0Actual); + + AscendC::WaitFlag(l0AEventList[l0AListId]); + if ((mL0Idx == 0) && (kL0Idx == 0)) { + AscendC::WaitFlag(l1AEventList[l1AListId]); + } + + // Load current tile from L1 to L0A + copyL1ToL0A(tensorL0A, tensorTileL1A); + + if ((mL0Idx == mL0Loop - 1) && (kL0Idx == kL0Loop - 1)) { + AscendC::SetFlag(l1AEventList[l1AListId]); + } + + bool initC = ((kL1Idx == 0) && (kL0Idx == 0)); + for (int nL0Idx = 0; nL0Idx < nL0Loop; nL0Idx++) { + uint32_t nL0Actual = (nL0Idx < nL0Loop - 1) ? L0_TILE_N : (nL1Actual - nL0Idx * L0_TILE_N); + + // Locate the current tile on L0B + auto l0BTile = l0BTensorList[l0BListId]; + auto layoutBInL0 = tla::MakeLayout(kL0Actual, nL0Actual); + auto tensorL0B = tla::MakeTensor(l0BTile, layoutBInL0, Catlass::Arch::PositionL0B{}); + // Locate the current tile of matrix B on L1 + auto tensorTileL1B = GetTile(tensorL1B, + tla::MakeCoord(kL0Idx * L0_TILE_K, nL0Idx * L0_TILE_N), + tla::MakeShape(kL0Actual, nL0Actual)); + + // Wait for mmad finished + AscendC::WaitFlag(l0BEventList[l0BListId]); + // If the current tile is the first one on the k&n axis, wait for loading matrix B from GM to L1 + if ((mL0Idx == 0) && (kL0Idx == 0) && (nL0Idx == 0)) { + AscendC::WaitFlag(l1BEventList[l1BListId]); + } + + // Load current tile from L1 to L0B + copyL1ToL0B(tensorL0B, tensorTileL1B); + + // If the current tile is the last one on the k&n axis, notify to load matrix B from GM to L1 + if ((mL0Idx == mL0Loop - 1) && (kL0Idx == kL0Loop - 1) && (nL0Idx == nL0Loop - 1)) { + AscendC::SetFlag(l1BEventList[l1BListId]); + } + + if constexpr (HAS_BIAS && !std::is_same_v) { + if (initC) { + if (nL0Idx == 0) { + AscendC::WaitFlag(L1A_STAGES + L1B_STAGES); + } + AscendC::WaitFlag(L0A_STAGES + L0B_STAGES); + auto l1Bias = l1BiasTensor.template ReinterpretCast(); + auto tensorL1Bias = tla::MakeTensor(l1Bias, L1BIAS_LAYOUT, Catlass::Arch::PositionL1{}); + auto tensorTileL1Bias = GetTile(tensorL1Bias, + tla::MakeCoord(nL0Idx * L0_TILE_N), + tla::MakeShape(nL0Actual)); + // Load bias to l0 biasTable + copyL1ToBT(tensorL0Bias, tensorTileL1Bias); + if (nL0Idx == nL0Loop - 1) { + AscendC::SetFlag(L1A_STAGES + L1B_STAGES); + } + } + } + + // Notify to do mmad + AscendC::SetFlag(l0CEventList[l0CListId]); + + // Locate the current tile on L0C + auto tensorTileL0C = GetTile(tensorL0C, + tla::MakeCoord(mL0Idx * L0_TILE_M, nL0Idx * L0_TILE_N), + tla::MakeShape(mL0Actual, nL0Actual)); + + // Compute the matrix multiplication on L0A and L0B and write the result to the accumulator + // Wait for loading L0B + AscendC::WaitFlag(l0CEventList[l0CListId]); + + // If the unit flag is enabled, the unit flag is set according to the calculation progress + uint8_t unitFlag = 0b00; + if constexpr (ENABLE_UNIT_FLAG) { + if ((kL1Idx == kL1Loop - 1) && (mL0Idx == mL0Loop - 1) && + (kL0Idx == kL0Loop - 1) && (nL0Idx == nL0Loop - 1)) { + unitFlag = 0b11; + } else { + unitFlag = 0b10; + } + } + + if constexpr (HAS_BIAS && !std::is_same_v) { + if (initC) { + tileMmad(tensorTileL0C, tensorL0A, tensorL0B, tensorL0Bias, + mL0Actual, nL0Actual, kL0Actual, initC, unitFlag); + AscendC::SetFlag(L0A_STAGES + L0B_STAGES); + } else { + tileMmad(tensorTileL0C, tensorL0A, tensorL0B, + mL0Actual, nL0Actual, kL0Actual, initC, unitFlag); + } + } else { + tileMmad(tensorTileL0C, tensorL0A, tensorL0B, + mL0Actual, nL0Actual, kL0Actual, initC, unitFlag); + } + + // Notify to move the next L0B tile + AscendC::SetFlag(l0BEventList[l0BListId]); + l0BListId = (l0BListId + 1 < L0B_STAGES) ? (l0BListId + 1) : 0; + } + AscendC::SetFlag(l0AEventList[l0AListId]); + l0AListId = (l0AListId + 1 < L0A_STAGES) ? (l0AListId + 1) : 0; + } + } + l1AListId = l1AListIdNext; + l1BListId = l1BListIdNext; + kL1Actual = kL1ActualNext; + } + + // copy block out + if constexpr (!ENABLE_UNIT_FLAG) { + AscendC::SetFlag(l0CEventList[l0CListId]); + AscendC::WaitFlag(l0CEventList[l0CListId]); + copyL0CToDst(tensorC, tensorL0C); + AscendC::SetFlag(l0CEventList[l0CListId]); + l0CListId = (l0CListId + 1 < L0C_STAGES) ? (l0CListId + 1) : 0; + } else { + copyL0CToDst(tensorC, tensorL0C, 0b11); + } + } + + /// Perform a block-scoped matrix multiply-accumulate + template + CATLASS_DEVICE void operator()(TensorA &tensorA, TensorB &tensorB, TensorC tensorCList[MAX_CUBE_VEC_SYNC_NUM], + Catlass::GemmCoord const &actualShape, uint32_t ubRowNum, uint8_t beginSubBlockIdx, uint64_t l0C2UBAIVAICFlag, + uint64_t l0C2UBAICAIVFlag, uint32_t& ubListId, uint8_t sendVecNum, uint64_t ubNum = 2, TensorBias const &tensorBias = {}) + { + // Check L1TileShape + if constexpr (HAS_BIAS) { + static constexpr uint32_t BIAS_BUF_SIZE = L0_TILE_N * sizeof(ElementAccumulator); + static constexpr uint32_t L1BIAS_SIZE = L1_TILE_N * sizeof(ElementBias); + static_assert(BIAS_BUF_SIZE <= ArchTag::BIAS_SIZE, + "BIAS_BUF_SIZE exceeding the BT space! Reduce L0_TILE_N"); + static_assert(L1A_TILE_SIZE * L1A_STAGES + L1B_TILE_SIZE * L1B_STAGES + L1BIAS_SIZE <= ArchTag::L1_SIZE, + "L1TileShape exceeding the L1 space!"); + } + + using CopyGmToL1A = typename TileCopy_::template CopyGmToL1A; + using CopyGmToL1B = typename TileCopy_::template CopyGmToL1B; + using CopyL0CToDst = typename TileCopy_::template CopyL0CToDst; + CopyGmToL1A copyGmToL1A; + CopyGmToL1B copyGmToL1B; + CopyL0CToDst copyL0CToDst; + + uint32_t mBlockActual = actualShape.m(); + uint32_t kBlockActual = actualShape.k(); + uint32_t nBlockActual = actualShape.n(); + + uint32_t mL1Actual = mBlockActual; + if constexpr (std::is_same_v) { + // Avoid using the gemv mode in mmad + if (mL1Actual == 1) { + mL1Actual = 16; + } + } + uint32_t nL1Actual = nBlockActual; + + auto layoutInL0C = tla::MakeLayoutL0C(mL1Actual, nL1Actual); + auto tensorL0C = tla::MakeTensor(l0CTensorList[l0CListId], layoutInL0C, Catlass::Arch::PositionL0C{}); + auto tensorL0Bias = tla::MakeTensor(l0BiasTensor, L0BIAS_LAYOUT, Catlass::Arch::PositionBias{}); + + uint32_t kL1Actual = min(kBlockActual, L1_TILE_K); + // load first matrix A tile from GM to L1 + AscendC::WaitFlag(l1AEventList[l1AListId]); + auto tensorL1A = tla::MakeTensor(l1ATensorList[l1AListId], L1A_LAYOUT, Catlass::Arch::PositionL1{}); + auto tensorTileA = GetTileA(tensorA, 0, 0, mBlockActual, kL1Actual); + if constexpr (ENABLE_L1_RESIDENT) { + // If the currently loaded GM pointer and block coordinates are the same as the last loaded ones, + // skip this loadding. + if (lastAddrA[l1AListId] != tensorTileA.data().GetPhyAddr() + || tla::get<0>(tensorTileA.coord()) != lastCoordA[l1AListId].row() + || tla::get<1>(tensorTileA.coord()) != lastCoordA[l1AListId].column()) { + copyGmToL1A(tensorL1A, tensorTileA); + lastCoordA[l1AListId] = Catlass::MatrixCoord{tla::get<0>(tensorTileA.coord()), tla::get<1>(tensorTileA.coord())}; + lastAddrA[l1AListId] = const_cast<__gm__ typename AscendC::GlobalTensor::PrimType *>( + tensorTileA.data().GetPhyAddr() + ); + } + } else { + copyGmToL1A(tensorL1A, tensorTileA); + } + AscendC::SetFlag(l1AEventList[l1AListId]); + + // load first matrix B tile from GM to L1 + AscendC::WaitFlag(l1BEventList[l1BListId]); + auto tensorL1B = tla::MakeTensor(l1BTensorList[l1BListId], L1B_LAYOUT, Catlass::Arch::PositionL1{}); + auto tensorTileB = GetTile(tensorB, tla::MakeCoord(0, 0), tla::MakeShape(kL1Actual, nBlockActual)); + if constexpr (ENABLE_L1_RESIDENT) { + if (lastAddrB[l1BListId] != tensorTileB.data().GetPhyAddr() + || tla::get<0>(tensorTileB.coord()) != lastCoordB[l1BListId].row() + || tla::get<1>(tensorTileB.coord()) != lastCoordB[l1BListId].column()) { + copyGmToL1B(tensorL1B, tensorTileB); + lastCoordB[l1BListId] = Catlass::MatrixCoord{tla::get<0>(tensorTileB.coord()), tla::get<1>(tensorTileB.coord())}; + lastAddrB[l1BListId] = const_cast<__gm__ typename AscendC::GlobalTensor::PrimType *>( + tensorTileB.data().GetPhyAddr() + ); + } + } else { + copyGmToL1B(tensorL1B, tensorTileB); + } + AscendC::SetFlag(l1BEventList[l1BListId]); + + if constexpr (HAS_BIAS && !std::is_same_v) { + using CopyGmToL1Bias = typename TileCopy::template CopyGmToL1Bias; + CopyGmToL1Bias copyGmToL1Bias; + AscendC::WaitFlag(L1A_STAGES + L1B_STAGES); + auto l1Bias = l1BiasTensor.template ReinterpretCast(); + auto tensorL1Bias = tla::MakeTensor(l1Bias, L1BIAS_LAYOUT, Catlass::Arch::PositionL1{}); + copyGmToL1Bias(tensorL1Bias, tensorBias); + AscendC::SetFlag(L1A_STAGES + L1B_STAGES); + } + + if constexpr (!ENABLE_UNIT_FLAG) { + AscendC::WaitFlag(l0CEventList[l0CListId]); + } + + uint32_t mL0Loop = CeilDiv(mL1Actual); + uint32_t nL0Loop = CeilDiv(nL1Actual); + + // main loop + uint32_t kL1Loop = CeilDiv(kBlockActual); + for (uint32_t kL1Idx = 0; kL1Idx < kL1Loop; kL1Idx++) { + uint32_t l1AListIdNext = (l1AListId + 1 < L1A_STAGES) ? (l1AListId + 1) : 0; + uint32_t l1BListIdNext = (l1BListId + 1 < L1B_STAGES) ? (l1BListId + 1) : 0; + uint32_t kL1ActualNext{0}; + // preload next tile from GM to L1 + if (kL1Idx < kL1Loop - 1) { + uint32_t kL1IdxNext = kL1Idx + 1; + kL1ActualNext = (kL1IdxNext < kL1Loop - 1) ? L1_TILE_K : (kBlockActual - kL1IdxNext * L1_TILE_K); + + // Get L1 tensor for next stage + auto l1ATensor = l1ATensorList[l1AListIdNext]; + auto l1BTensor = l1BTensorList[l1BListIdNext]; + auto tensorL1A = tla::MakeTensor(l1ATensor, L1A_LAYOUT, Catlass::Arch::PositionL1{}); + auto tensorL1B = tla::MakeTensor(l1BTensor, L1B_LAYOUT, Catlass::Arch::PositionL1{}); + // Get GM tile for next stage + auto tensorTileA = GetTileA(tensorA, 0, kL1IdxNext * L1_TILE_K, mBlockActual, kL1ActualNext); + auto tensorTileB = GetTile(tensorB, tla::MakeCoord(kL1IdxNext * L1_TILE_K, 0), + tla::MakeShape(kL1ActualNext, nBlockActual)); + + // load next matrix A tile from GM to L1 + AscendC::WaitFlag(l1AEventList[l1AListIdNext]); + if constexpr (ENABLE_L1_RESIDENT) { + if (lastAddrA[l1AListIdNext] != tensorTileA.data().GetPhyAddr() + || tla::get<0>(tensorTileA.coord()) != lastCoordA[l1AListIdNext].row() + || tla::get<1>(tensorTileA.coord()) != lastCoordA[l1AListIdNext].column()) { + copyGmToL1A(tensorL1A, tensorTileA); + lastCoordA[l1AListIdNext] = + Catlass::MatrixCoord{tla::get<0>(tensorTileA.coord()), tla::get<1>(tensorTileA.coord())}; + lastAddrA[l1AListIdNext] = + const_cast<__gm__ typename AscendC::GlobalTensor::PrimType *>( + tensorTileA.data().GetPhyAddr() + ); + } + } else { + copyGmToL1A(tensorL1A, tensorTileA); + } + AscendC::SetFlag(l1AEventList[l1AListIdNext]); + + // load next matrix B tile from GM to L1 + AscendC::WaitFlag(l1BEventList[l1BListIdNext]); + if constexpr (ENABLE_L1_RESIDENT) { + if (lastAddrB[l1BListIdNext] != tensorTileB.data().GetPhyAddr() + || tla::get<0>(tensorTileB.coord()) != lastCoordB[l1BListIdNext].row() + || tla::get<1>(tensorTileB.coord()) != lastCoordB[l1BListIdNext].column()) { + copyGmToL1B(tensorL1B, tensorTileB); + lastCoordB[l1BListIdNext] = + Catlass::MatrixCoord{tla::get<0>(tensorTileB.coord()), tla::get<1>(tensorTileB.coord())}; + lastAddrB[l1BListIdNext] = + const_cast<__gm__ typename AscendC::GlobalTensor::PrimType *>( + tensorTileB.data().GetPhyAddr() + ); + } + } else { + copyGmToL1B(tensorL1B, tensorTileB); + } + AscendC::SetFlag(l1BEventList[l1BListIdNext]); + } + + // Get L1 tensor for current stage + auto l1ATensor = l1ATensorList[l1AListId]; + auto l1BTensor = l1BTensorList[l1BListId]; + tensorL1A = tla::MakeTensor(l1ATensor, L1A_LAYOUT, Catlass::Arch::PositionL1{}); + tensorL1B = tla::MakeTensor(l1BTensor, L1B_LAYOUT, Catlass::Arch::PositionL1{}); + // Get the loop nums on L0 + uint32_t kL0Loop = CeilDiv(kL1Actual); + + for (int mL0Idx = 0; mL0Idx < mL0Loop; mL0Idx++) { + uint32_t mL0Actual = (mL0Idx < mL0Loop - 1) ? L0_TILE_M : (mL1Actual - mL0Idx * L0_TILE_M); + + for (int kL0Idx = 0; kL0Idx < kL0Loop; kL0Idx++) { + uint32_t kL0Actual = (kL0Idx < kL0Loop - 1) ? L0_TILE_K : (kL1Actual - kL0Idx * L0_TILE_K); + + // Locate the current tile on L0A + auto l0ATile = l0ATensorList[l0AListId]; + auto layoutAInL0 = tla::MakeLayout(mL0Actual, kL0Actual); + auto tensorL0A = tla::MakeTensor(l0ATile, layoutAInL0, Catlass::Arch::PositionL0A{}); + // Locate the current tile of matrix A on L1 + auto tensorTileL1A = GetTileA(tensorL1A, mL0Idx * L0_TILE_M, kL0Idx * L0_TILE_K, mL0Actual, kL0Actual); + + AscendC::WaitFlag(l0AEventList[l0AListId]); + if ((mL0Idx == 0) && (kL0Idx == 0)) { + AscendC::WaitFlag(l1AEventList[l1AListId]); + } + + // Load current tile from L1 to L0A + copyL1ToL0A(tensorL0A, tensorTileL1A); + + if ((mL0Idx == mL0Loop - 1) && (kL0Idx == kL0Loop - 1)) { + AscendC::SetFlag(l1AEventList[l1AListId]); + } + + bool initC = ((kL1Idx == 0) && (kL0Idx == 0)); + for (int nL0Idx = 0; nL0Idx < nL0Loop; nL0Idx++) { + uint32_t nL0Actual = (nL0Idx < nL0Loop - 1) ? L0_TILE_N : (nL1Actual - nL0Idx * L0_TILE_N); + + // Locate the current tile on L0B + auto l0BTile = l0BTensorList[l0BListId]; + auto layoutBInL0 = tla::MakeLayout(kL0Actual, nL0Actual); + auto tensorL0B = tla::MakeTensor(l0BTile, layoutBInL0, Catlass::Arch::PositionL0B{}); + // Locate the current tile of matrix B on L1 + auto tensorTileL1B = GetTile(tensorL1B, + tla::MakeCoord(kL0Idx * L0_TILE_K, nL0Idx * L0_TILE_N), + tla::MakeShape(kL0Actual, nL0Actual)); + + // Wait for mmad finished + AscendC::WaitFlag(l0BEventList[l0BListId]); + // If the current tile is the first one on the k&n axis, wait for loading matrix B from GM to L1 + if ((mL0Idx == 0) && (kL0Idx == 0) && (nL0Idx == 0)) { + AscendC::WaitFlag(l1BEventList[l1BListId]); + } + + // Load current tile from L1 to L0B + copyL1ToL0B(tensorL0B, tensorTileL1B); + + // If the current tile is the last one on the k&n axis, notify to load matrix B from GM to L1 + if ((mL0Idx == mL0Loop - 1) && (kL0Idx == kL0Loop - 1) && (nL0Idx == nL0Loop - 1)) { + AscendC::SetFlag(l1BEventList[l1BListId]); + } + + if constexpr (HAS_BIAS && !std::is_same_v) { + if (initC) { + if (nL0Idx == 0) { + AscendC::WaitFlag(L1A_STAGES + L1B_STAGES); + } + AscendC::WaitFlag(L0A_STAGES + L0B_STAGES); + auto l1Bias = l1BiasTensor.template ReinterpretCast(); + auto tensorL1Bias = tla::MakeTensor(l1Bias, L1BIAS_LAYOUT, Catlass::Arch::PositionL1{}); + auto tensorTileL1Bias = GetTile(tensorL1Bias, + tla::MakeCoord(nL0Idx * L0_TILE_N), + tla::MakeShape(nL0Actual)); + // Load bias to l0 biastable + copyL1ToBT(tensorL0Bias, tensorTileL1Bias); + if (nL0Idx == nL0Loop - 1) { + AscendC::SetFlag(L1A_STAGES + L1B_STAGES); + } + } + } + + // Notify to do mmad + AscendC::SetFlag(l0CEventList[l0CListId]); + + // Locate the current tile on L0C + auto tensorTileL0C = GetTile(tensorL0C, + tla::MakeCoord(mL0Idx * L0_TILE_M, nL0Idx * L0_TILE_N), + tla::MakeShape(mL0Actual, nL0Actual)); + + // Compute the matrix multiplication on L0A and L0B and write the result to the accumulator + // Wait for loading L0B + AscendC::WaitFlag(l0CEventList[l0CListId]); + + // If the unit flag is enabled, the unit flag is set according to the calculation progress + uint8_t unitFlag = 0b00; + if constexpr (ENABLE_UNIT_FLAG) { + if ((kL1Idx == kL1Loop - 1) && (mL0Idx == mL0Loop - 1) && + (kL0Idx == kL0Loop - 1) && (nL0Idx == nL0Loop - 1)) { + unitFlag = 0b11; + } else { + unitFlag = 0b10; + } + } + + if constexpr (HAS_BIAS && !std::is_same_v) { + if (initC) { + tileMmad(tensorTileL0C, tensorL0A, tensorL0B, tensorL0Bias, + mL0Actual, nL0Actual, kL0Actual, initC, unitFlag); + AscendC::SetFlag(L0A_STAGES + L0B_STAGES); + } else { + tileMmad(tensorTileL0C, tensorL0A, tensorL0B, + mL0Actual, nL0Actual, kL0Actual, initC, unitFlag); + } + } else { + tileMmad(tensorTileL0C, tensorL0A, tensorL0B, + mL0Actual, nL0Actual, kL0Actual, initC, unitFlag); + } + + // Notify to move the next L0B tile + AscendC::SetFlag(l0BEventList[l0BListId]); + l0BListId = (l0BListId + 1 < L0B_STAGES) ? (l0BListId + 1) : 0; + } + AscendC::SetFlag(l0AEventList[l0AListId]); + l0AListId = (l0AListId + 1 < L0A_STAGES) ? (l0AListId + 1) : 0; + } + } + l1AListId = l1AListIdNext; + l1BListId = l1BListIdNext; + kL1Actual = kL1ActualNext; + } + + // copy block out + if constexpr (!ENABLE_UNIT_FLAG) { + AscendC::SetFlag(l0CEventList[l0CListId]); + AscendC::WaitFlag(l0CEventList[l0CListId]); + uint32_t leftRowNum = mL1Actual; + uint32_t twoVecRowNum = ubRowNum * sendVecNum; + uint32_t curRowNum = min(twoVecRowNum, leftRowNum); + uint32_t rowIdx = 0; + uint32_t tileNum = CeilDiv(mL1Actual, twoVecRowNum); + for (uint32_t i = 0; i < tileNum; i++) { + auto tensorTileL0C = GetTile(tensorL0C, + tla::MakeCoord(rowIdx, 0), + tla::MakeShape(curRowNum, nL1Actual)); + if (sendVecNum == 2) { + AscendC::CrossCoreWaitFlag<0x4, PIPE_FIX>(l0C2UBAIVAICFlag + ubListId); + AscendC::CrossCoreWaitFlag<0x4, PIPE_FIX>(l0C2UBAIVAICFlag + FLAG_ID_MAX + ubListId); + copyL0CToDst(tensorCList[ubListId], tensorTileL0C, ubRowNum, beginSubBlockIdx, sendVecNum); + AscendC::CrossCoreSetFlag<0x4, PIPE_FIX>(l0C2UBAICAIVFlag + ubListId); + AscendC::CrossCoreSetFlag<0x4, PIPE_FIX>(l0C2UBAICAIVFlag + FLAG_ID_MAX + ubListId); + rowIdx += curRowNum; + leftRowNum -= curRowNum; + curRowNum = min(twoVecRowNum, leftRowNum); + ubListId = (ubListId + 1 < ubNum) ? (ubListId + 1) : 0; + } else if (sendVecNum == 1) { + if (beginSubBlockIdx == 0) { + AscendC::CrossCoreWaitFlag<0x4, PIPE_FIX>(l0C2UBAIVAICFlag + ubListId); + } else if (beginSubBlockIdx == 1) { + AscendC::CrossCoreWaitFlag<0x4, PIPE_FIX>(l0C2UBAIVAICFlag + FLAG_ID_MAX + ubListId); + } + copyL0CToDst(tensorCList[ubListId], tensorTileL0C, ubRowNum, beginSubBlockIdx, sendVecNum); + if (beginSubBlockIdx == 0) { + AscendC::CrossCoreSetFlag<0x4, PIPE_FIX>(l0C2UBAICAIVFlag + ubListId); + } else if (beginSubBlockIdx == 1) { + AscendC::CrossCoreSetFlag<0x4, PIPE_FIX>(l0C2UBAICAIVFlag + FLAG_ID_MAX + ubListId); + } + rowIdx += curRowNum; + leftRowNum -= curRowNum; + curRowNum = min(twoVecRowNum, leftRowNum); + ubListId = (ubListId + 1 < ubNum) ? (ubListId + 1) : 0; + } + } + AscendC::SetFlag(l0CEventList[l0CListId]); + l0CListId = (l0CListId + 1 < L0C_STAGES) ? (l0CListId + 1) : 0; + } else { + uint32_t leftRowNum = mL1Actual; + uint32_t twoVecRowNum = ubRowNum * sendVecNum; + uint32_t curRowNum = min(twoVecRowNum, leftRowNum); + uint32_t rowIdx = 0; + uint32_t tileNum = CeilDiv(mL1Actual, twoVecRowNum); + for (uint32_t i = 0; i < tileNum; i++) { + auto tensorTileL0C = GetTile(tensorL0C, + tla::MakeCoord(rowIdx, 0), + tla::MakeShape(curRowNum, nL1Actual)); + if (sendVecNum == 2) { + AscendC::CrossCoreWaitFlag<0x4, PIPE_FIX>(l0C2UBAIVAICFlag + ubListId); + AscendC::CrossCoreWaitFlag<0x4, PIPE_FIX>(l0C2UBAIVAICFlag + FLAG_ID_MAX + ubListId); + copyL0CToDst(tensorCList[ubListId], tensorTileL0C, ubRowNum, beginSubBlockIdx, sendVecNum); + AscendC::CrossCoreSetFlag<0x4, PIPE_FIX>(l0C2UBAICAIVFlag + ubListId); + AscendC::CrossCoreSetFlag<0x4, PIPE_FIX>(l0C2UBAICAIVFlag + FLAG_ID_MAX + ubListId); + rowIdx += curRowNum; + leftRowNum -= curRowNum; + curRowNum = min(twoVecRowNum, leftRowNum); + ubListId = (ubListId + 1 < ubNum) ? (ubListId + 1) : 0; + } else if (sendVecNum == 1) { + if (beginSubBlockIdx == 0) { + AscendC::CrossCoreWaitFlag<0x4, PIPE_FIX>(l0C2UBAIVAICFlag + ubListId); + } else if (beginSubBlockIdx == 1) { + AscendC::CrossCoreWaitFlag<0x4, PIPE_FIX>(l0C2UBAIVAICFlag + FLAG_ID_MAX + ubListId); + } + copyL0CToDst(tensorCList[ubListId], tensorTileL0C, ubRowNum, beginSubBlockIdx, sendVecNum); + if (beginSubBlockIdx == 0) { + AscendC::CrossCoreSetFlag<0x4, PIPE_FIX>(l0C2UBAICAIVFlag + ubListId); + } else if (beginSubBlockIdx == 1) { + AscendC::CrossCoreSetFlag<0x4, PIPE_FIX>(l0C2UBAICAIVFlag + FLAG_ID_MAX + ubListId); + } + rowIdx += curRowNum; + leftRowNum -= curRowNum; + curRowNum = min(twoVecRowNum, leftRowNum); + ubListId = (ubListId + 1 < ubNum) ? (ubListId + 1) : 0; + } + } + } + } + +protected: + template + CATLASS_DEVICE auto GetTileA(TensorA &tensorA, uint32_t mIndex, uint32_t kIndex, uint32_t mSize, uint32_t kSize) + { + if constexpr(tla::detail::isVector::value) { + return GetTile(tensorA, tla::MakeCoord(kIndex), tla::MakeShape(kSize)); + } else { + return GetTile(tensorA, tla::MakeCoord(mIndex, kIndex), tla::MakeShape(mSize, kSize)); + } + } + + // Multi-stage tensors list + AscendC::LocalTensor l1ATensorList[L1A_STAGES]; + AscendC::LocalTensor l1BTensorList[L1B_STAGES]; + AscendC::LocalTensor l0ATensorList[L0A_STAGES]; + AscendC::LocalTensor l0BTensorList[L0B_STAGES]; + AscendC::LocalTensor l0CTensorList[L0C_STAGES]; + AscendC::LocalTensor l1BiasTensor; + AscendC::LocalTensor l0BiasTensor; + + // Multi-stage event id list + int32_t l1AEventList[L1A_STAGES]; + int32_t l1BEventList[L1B_STAGES]; + int32_t l0AEventList[L0A_STAGES]; + int32_t l0BEventList[L0B_STAGES]; + int32_t l0CEventList[L0C_STAGES]; + + __gm__ typename AscendC::GlobalTensor::PrimType* lastAddrA[L1A_STAGES]; + __gm__ typename AscendC::GlobalTensor::PrimType* lastAddrB[L1B_STAGES]; + Catlass::MatrixCoord lastCoordA[L1A_STAGES]; + Catlass::MatrixCoord lastCoordB[L1B_STAGES]; + + // The id of current stage + uint32_t l1AListId{0}; + uint32_t l1BListId{0}; + uint32_t l0AListId{0}; + uint32_t l0BListId{0}; + uint32_t l0CListId{0}; + + TileMmad tileMmad; + CopyL1ToL0A copyL1ToL0A; + CopyL1ToL0B copyL1ToL0B; + CopyL1ToBT copyL1ToBT; +}; + +} // namespace Common + +#endif // CATLASS_GEMM_BLOCK_BLOCK_MMAD_PINGPONG_TLA_HPP diff --git a/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla_multi.hpp b/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla_multi.hpp index 79e5666b19e3..a813e39fe473 100644 --- a/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla_multi.hpp +++ b/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla_multi.hpp @@ -61,7 +61,7 @@ template < class TileMmad_ > struct BlockMmadTla < - MmadPingpongTlaMulti, L1TileShape_, L0TileShape_, @@ -74,7 +74,7 @@ struct BlockMmadTla < > { public: // Type Aliases - using DispatchPolicy = MmadPingpongTlaMulti; using ArchTag = typename DispatchPolicy::ArchTag; using TileCopy = TileCopy_; @@ -170,10 +170,10 @@ struct BlockMmadTla < static_assert(L0_TILE_K * SizeOfBits::value % _32B == 0, "L0TileShape::K must be 32B aligned."); #endif - static_assert((!HAS_BIAS && (L1A_STAGES + L1B_STAGES) <= 8) || (HAS_BIAS && (L1A_STAGES + L1B_STAGES) <= 7), + static_assert((!HAS_BIAS && (L1A_STAGES + L1B_STAGES) <= 8) || (HAS_BIAS && (L1A_STAGES + L1B_STAGES) <= 7), "L1 Buffer overflow: Exceeds the supported range of EVENT(0~7)"); - static_assert((!HAS_BIAS && (L0A_STAGES + L0B_STAGES) <= 8) || (HAS_BIAS && (L0A_STAGES + L0B_STAGES) <= 7), + static_assert((!HAS_BIAS && (L0A_STAGES + L0B_STAGES) <= 8) || (HAS_BIAS && (L0A_STAGES + L0B_STAGES) <= 7), "L0 Buffer overflow: Exceeds the supported range of EVENT_ID(0~7)"); static constexpr auto L1A_LAYOUT = @@ -203,12 +203,7 @@ struct BlockMmadTla < CATLASS_DEVICE BlockMmadTla(Arch::Resource &resource, uint32_t l1BufAddrStart = 0) { -#ifdef CATLASS_UNIFIED_CORE - resourcePtr = &resource; - { -#else if ASCEND_IS_AIC { -#endif uint32_t l1AOffset = l1BufAddrStart; uint32_t l1BOffset = l1BufAddrStart + L1A_TILE_SIZE * L1A_STAGES; // Init buffers @@ -258,11 +253,8 @@ struct BlockMmadTla < CATLASS_DEVICE void preSetFlags() { -#ifdef CATLASS_UNIFIED_CORE - { -#else + if ASCEND_IS_AIC { -#endif // use HF32 when USE_HF32_MODE is true if constexpr (USE_HF32_MODE) { AscendC::SetHF32Mode(true); @@ -302,11 +294,7 @@ struct BlockMmadTla < CATLASS_DEVICE void finalWaitFlags() { -#ifdef CATLASS_UNIFIED_CORE - { -#else if ASCEND_IS_AIC { -#endif if constexpr (USE_HF32_MODE) { AscendC::SetHF32Mode(false); } @@ -340,7 +328,7 @@ struct BlockMmadTla < /// Perform a block-scoped matrix multiply-accumulate template CATLASS_DEVICE void operator()(TensorA &tensorA, TensorB &tensorB, TensorC &tensorC, GemmCoord const &actualShape, - TensorBias const &tensorBias = {}) + TensorBias const &tensorBias = {}, bool clearL1Padding = false) { // Check L1TileShape if constexpr (HAS_BIAS) { @@ -356,15 +344,14 @@ struct BlockMmadTla < using CopyGmToL1B = typename TileCopy_::template CopyGmToL1B; CopyGmToL1A copyGmToL1A; CopyGmToL1B copyGmToL1B; -#ifdef CATLASS_UNIFIED_CORE - // 310P: no Fixpipe, no DataCopyCO12Dst. L0C exits via DataCopy L0C→UB then UB→GM. -#elif (defined (CATLASS_ARCH) && CATLASS_ARCH == 2201) +#if (defined (CATLASS_ARCH) && CATLASS_ARCH == 2201) using CopyL0CToGm = typename TileCopy_::template CopyL0CToGm; CopyL0CToGm copyL0CToDst; -#elif (defined (CATLASS_ARCH) && CATLASS_ARCH == 3510) +#endif +#if (defined (CATLASS_ARCH) && CATLASS_ARCH == 3510) using CopyL0CToDst = typename TileCopy_::template CopyL0CToDst; CopyL0CToDst copyL0CToDst; -#endif +#endif uint32_t mBlockActual = actualShape.m(); uint32_t kBlockActual = actualShape.k(); @@ -386,6 +373,12 @@ struct BlockMmadTla < uint32_t kL1Actual = min(kBlockActual, L1_TILE_K); // load first matrix A tile from GM to L1 AscendC::WaitFlag(l1AEventList[l1AListId]); + if (clearL1Padding) { + AscendC::InitConstValueParams clearParams( + 1, static_cast(L1A_TILE_SIZE / 32), 0, + static_cast(0)); + AscendC::InitConstValue(l1ATensorList[l1AListId], clearParams); + } auto tensorL1A = tla::MakeTensor(l1ATensorList[l1AListId], L1A_LAYOUT, Arch::PositionL1{}); auto tensorTileA = GetTileA(tensorA, 0, 0, mBlockActual, kL1Actual); if constexpr (ENABLE_L1_RESIDENT) { @@ -395,7 +388,9 @@ struct BlockMmadTla < || tla::get<0>(tensorTileA.coord()) != lastCoordA[l1AListId].row() || tla::get<1>(tensorTileA.coord()) != lastCoordA[l1AListId].column()) { copyGmToL1A(tensorL1A, tensorTileA); - lastCoordA[l1AListId] = MatrixCoord{tla::get<0>(tensorTileA.coord()), tla::get<1>(tensorTileA.coord())}; + lastCoordA[l1AListId] = MatrixCoord{ + static_cast(tla::get<0>(tensorTileA.coord())), + static_cast(tla::get<1>(tensorTileA.coord()))}; lastAddrA[l1AListId] = const_cast<__gm__ typename AscendC::GlobalTensor::PrimType *>( tensorTileA.data().GetPhyAddr() ); @@ -407,6 +402,12 @@ struct BlockMmadTla < // load first matrix B tile from GM to L1 AscendC::WaitFlag(l1BEventList[l1BListId]); + if (clearL1Padding) { + AscendC::InitConstValueParams clearParams( + 1, static_cast(L1B_TILE_SIZE / 32), 0, + static_cast(0)); + AscendC::InitConstValue(l1BTensorList[l1BListId], clearParams); + } auto tensorL1B = tla::MakeTensor(l1BTensorList[l1BListId], L1B_LAYOUT, Arch::PositionL1{}); auto tensorTileB = GetTile(tensorB, tla::MakeCoord(0, 0), tla::MakeShape(kL1Actual, nBlockActual)); if constexpr (ENABLE_L1_RESIDENT) { @@ -414,7 +415,9 @@ struct BlockMmadTla < || tla::get<0>(tensorTileB.coord()) != lastCoordB[l1BListId].row() || tla::get<1>(tensorTileB.coord()) != lastCoordB[l1BListId].column()) { copyGmToL1B(tensorL1B, tensorTileB); - lastCoordB[l1BListId] = MatrixCoord{tla::get<0>(tensorTileB.coord()), tla::get<1>(tensorTileB.coord())}; + lastCoordB[l1BListId] = MatrixCoord{ + static_cast(tla::get<0>(tensorTileB.coord())), + static_cast(tla::get<1>(tensorTileB.coord()))}; lastAddrB[l1BListId] = const_cast<__gm__ typename AscendC::GlobalTensor::PrimType *>( tensorTileB.data().GetPhyAddr() ); @@ -464,13 +467,20 @@ struct BlockMmadTla < // load next matrix A tile from GM to L1 AscendC::WaitFlag(l1AEventList[l1AListIdNext]); + if (clearL1Padding) { + AscendC::InitConstValueParams clearParams( + 1, static_cast(L1A_TILE_SIZE / 32), 0, + static_cast(0)); + AscendC::InitConstValue(l1ATensorList[l1AListIdNext], clearParams); + } if constexpr (ENABLE_L1_RESIDENT) { if (lastAddrA[l1AListIdNext] != tensorTileA.data().GetPhyAddr() || tla::get<0>(tensorTileA.coord()) != lastCoordA[l1AListIdNext].row() || tla::get<1>(tensorTileA.coord()) != lastCoordA[l1AListIdNext].column()) { copyGmToL1A(tensorL1A, tensorTileA); - lastCoordA[l1AListIdNext] = - MatrixCoord{tla::get<0>(tensorTileA.coord()), tla::get<1>(tensorTileA.coord())}; + lastCoordA[l1AListIdNext] = MatrixCoord{ + static_cast(tla::get<0>(tensorTileA.coord())), + static_cast(tla::get<1>(tensorTileA.coord()))}; lastAddrA[l1AListIdNext] = const_cast<__gm__ typename AscendC::GlobalTensor::PrimType *>( tensorTileA.data().GetPhyAddr() @@ -483,13 +493,20 @@ struct BlockMmadTla < // load next matrix B tile from GM to L1 AscendC::WaitFlag(l1BEventList[l1BListIdNext]); + if (clearL1Padding) { + AscendC::InitConstValueParams clearParams( + 1, static_cast(L1B_TILE_SIZE / 32), 0, + static_cast(0)); + AscendC::InitConstValue(l1BTensorList[l1BListIdNext], clearParams); + } if constexpr (ENABLE_L1_RESIDENT) { if (lastAddrB[l1BListIdNext] != tensorTileB.data().GetPhyAddr() || tla::get<0>(tensorTileB.coord()) != lastCoordB[l1BListIdNext].row() || tla::get<1>(tensorTileB.coord()) != lastCoordB[l1BListIdNext].column()) { copyGmToL1B(tensorL1B, tensorTileB); - lastCoordB[l1BListIdNext] = - MatrixCoord{tla::get<0>(tensorTileB.coord()), tla::get<1>(tensorTileB.coord())}; + lastCoordB[l1BListIdNext] = MatrixCoord{ + static_cast(tla::get<0>(tensorTileB.coord())), + static_cast(tla::get<1>(tensorTileB.coord()))}; lastAddrB[l1BListIdNext] = const_cast<__gm__ typename AscendC::GlobalTensor::PrimType *>( tensorTileB.data().GetPhyAddr() @@ -632,60 +649,6 @@ struct BlockMmadTla < } // copy block out -#ifdef CATLASS_UNIFIED_CORE - { - // 310P unified core: L0C→UB via DataCopy, then UB→GM. - // No Fixpipe or DataCopyCO12Dst on dav_m200. - uint32_t mAligned = (mBlockActual + 15) / 16 * 16; - uint32_t nAligned = (nBlockActual + 15) / 16 * 16; - uint32_t tileElems = mAligned * nAligned; - uint32_t tileBytes = tileElems * sizeof(ElementAccumulator); - - // UB temp for L0C→UB transfer. Offset 0 is safe: on unified core, - // mmad and epilogue run sequentially so UB is not shared concurrently. - // The epilogue allocates its own UB regions at higher offsets (≥32KB). - AscendC::LocalTensor co2Temp = - resourcePtr->ubBuf.template GetBufferByByte(0); - - AscendC::PipeBarrier(); - - // L0C → UB: BLOCK_MODE_MATRIX copies raw NZ fractals to UB - // For float: blockLen unit = 1024B (one 16×16 fractal) - AscendC::DataCopyParams l0c2ubParams; - l0c2ubParams.blockCount = static_cast(nAligned / 16); - l0c2ubParams.blockLen = static_cast(mAligned / 16); - l0c2ubParams.srcStride = 0; - l0c2ubParams.dstStride = 0; - AscendC::DataCopyEnhancedParams enhParams; - enhParams.blockMode = AscendC::BlockMode::BLOCK_MODE_MATRIX; - AscendC::DataCopy(co2Temp, l0CTensorList[l0CListId], l0c2ubParams, enhParams); - AscendC::PipeBarrier(); - - // UB → GM: fractal-by-fractal with strided DataCopy (NZ→ND deformat) - // NZ in UB: [N/16 Z-cols][M/16 fractals][16 rows][16 cols] - // ND in GM: [M rows][N cols] - auto dstOffset = tensorC.layout()(tensorC.coord()); - uint32_t gmStride = tla::get<0>(tensorC.stride()); - uint32_t mFracs = mAligned / 16; - uint32_t nFracs = nAligned / 16; - for (uint32_t nf = 0; nf < nFracs; nf++) { - for (uint32_t mf = 0; mf < mFracs; mf++) { - uint32_t ubOff = (nf * mFracs + mf) * 256; - uint32_t gmRow = mf * 16; - uint32_t gmCol = nf * 16; - uint32_t gmOff = dstOffset + gmRow * gmStride + gmCol; - AscendC::DataCopyParams fracParams; - fracParams.blockCount = 16; - fracParams.blockLen = static_cast(16 * sizeof(ElementAccumulator) / 32); - fracParams.srcStride = 0; - fracParams.dstStride = static_cast((gmStride - 16) * sizeof(ElementAccumulator) / 32); - AscendC::DataCopy(tensorC.data()[gmOff], co2Temp[ubOff], fracParams); - } - } - AscendC::PipeBarrier(); - l0CListId = (l0CListId + 1 < L0C_STAGES) ? (l0CListId + 1) : 0; - } -#else if constexpr (!ENABLE_UNIT_FLAG) { AscendC::SetFlag(l0CEventList[l0CListId]); AscendC::WaitFlag(l0CEventList[l0CListId]); @@ -695,7 +658,6 @@ struct BlockMmadTla < } else { copyL0CToDst(tensorC, tensorL0C, 0b11); } -#endif } protected: @@ -717,9 +679,6 @@ struct BlockMmadTla < AscendC::LocalTensor l0CTensorList[L0C_STAGES]; AscendC::LocalTensor l1BiasTensor; AscendC::LocalTensor l0BiasTensor; -#ifdef CATLASS_UNIFIED_CORE - Arch::Resource* resourcePtr{nullptr}; -#endif // Multi-stage event id list int32_t l1AEventList[L1A_STAGES]; @@ -732,7 +691,7 @@ struct BlockMmadTla < __gm__ typename AscendC::GlobalTensor::PrimType* lastAddrB[L1B_STAGES]; MatrixCoord lastCoordA[L1A_STAGES]; MatrixCoord lastCoordB[L1B_STAGES]; - + // The id of current stage uint32_t l1AListId{0}; uint32_t l1BListId{0}; diff --git a/csrc/moe/common/kernel_utils/tile/copy_l0c_to_ub.hpp b/csrc/moe/common/kernel_utils/tile/copy_l0c_to_ub.hpp new file mode 100644 index 000000000000..99f0086d4c83 --- /dev/null +++ b/csrc/moe/common/kernel_utils/tile/copy_l0c_to_ub.hpp @@ -0,0 +1,414 @@ +/** + * Copyright (c) 2026 Tianjin University, Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef COMMON_TILE_ASCEND950_COPY_L0C_TO_UB_950_HPP +#define COMMON_TILE_ASCEND950_COPY_L0C_TO_UB_950_HPP + +#include "catlass/arch/arch.hpp" +#include "catlass/catlass.hpp" +#if defined(CATLASS_ARCH) && CATLASS_ARCH == 3510 +#include "catlass/gemm/tile/ascend950/copy_l0c_to_dst.hpp" +#endif +#include "tla/tensor.hpp" +#include "catlass/detail/tag_to_layout.hpp" + +namespace Common::Tile { + +#if defined(CATLASS_ARCH) && CATLASS_ARCH == 3510 + +template < + class ArchTag, + class TensorSrc, + class TensorDst, + Catlass::Gemm::Tile::CopyL0CToUBMode CopyMode = Catlass::Gemm::Tile::CopyL0CToUBMode::NO_SPLIT, + Catlass::Gemm::Tile::ScaleGranularity DEQUANT_GRANULARITY = Catlass::Gemm::Tile::ScaleGranularity::NO_QUANT, + bool ReluEnable = false, + class Enable = void +> +struct CopyL0CToUBTla { + static_assert(DEPENDENT_FALSE, "Unsupported copy l0c to ub, can not find the specialization."); +}; + +#endif + +#if defined(CATLASS_ARCH) && CATLASS_ARCH == 3510 + +template < + /// Tag indicating architecture + class ArchTag, + class ElementA_, + class LayoutTagA_, + class ElementB_, + class LayoutTagB_, + class ElementC_, + class LayoutTagC_, + class ElementBias = void, + bool ReluEnable_ = false, + Catlass::Gemm::Tile::ScaleGranularity DEQUANT_GRANULARITY_ = Catlass::Gemm::Tile::ScaleGranularity::NO_QUANT, + class L0CCopyMode = Catlass::Gemm::Tile::CopyToGM> +struct PackedTileCopyTla { + using ElementA = ElementA_; + using ElementB = ElementB_; + using LayoutTagA = LayoutTagA_; + using LayoutTagB = LayoutTagB_; + using LayoutTagC = LayoutTagC_; + using ElementAccumulator = typename Catlass::Gemm::helper::ElementAccumulatorSelector::ElementAccumulator; + static constexpr bool ReluEnable = ReluEnable_; + static constexpr Catlass::Gemm::Tile::ScaleGranularity DEQUANT_GRANULARITY = DEQUANT_GRANULARITY_; + + static constexpr bool HAS_BIAS = !std::is_void_v; + static constexpr bool HAS_QUANT_TENSOR = (DEQUANT_GRANULARITY == Catlass::Gemm::Tile::ScaleGranularity::PER_CHANNEL); + + using LayoutTagL1A = typename Catlass::Gemm::helper::L1ATypeSelector>::L1AType::Layout; + using LayoutTagL1B = typename Catlass::Gemm::helper::L1BTypeSelector>::L1BType::Layout; + using LayoutTagL0A = typename Catlass::Gemm::helper::L0ALayoutSelector::Layout; + using LayoutTagL0B = Catlass::layout::nZ; + using LayoutTagL0C = Catlass::layout::L0C; + + using LayoutA = Catlass::detail::TagToLayout_t; + using LayoutB = Catlass::detail::TagToLayout_t; + using LayoutC = Catlass::detail::TagToLayout_t; + + using LayoutL1A = Catlass::detail::TagToLayout_t; + using LayoutL1B = Catlass::detail::TagToLayout_t; + using LayoutL0A = Catlass::detail::TagToLayout_t; + using LayoutL0B = Catlass::detail::TagToLayout_t; + using LayoutL0C = typename Catlass::detail::LayoutL0C; + + using TensorL1AVectorLayout = + tla::Tensor, LayoutL1A, tla::Coord, AscendC::TPosition::A1>; + using TensorL1ALayout = + tla::Tensor, LayoutL1A, tla::Coord, AscendC::TPosition::A1>; + + using TensorL1A = + std::conditional_t::value, TensorL1AVectorLayout, TensorL1ALayout>; + using TensorL1B = + tla::Tensor, LayoutL1B, tla::Coord, AscendC::TPosition::A1>; + using TensorL0A = + tla::Tensor, LayoutL0A, tla::Coord, AscendC::TPosition::A2>; + using TensorL0B = + tla::Tensor, LayoutL0B, tla::Coord, AscendC::TPosition::B2>; + using TensorL0C = tla:: + Tensor, LayoutL0C, tla::Coord, AscendC::TPosition::CO1>; + using TensorL1Bias = std::conditional_t< + HAS_BIAS, + tla::Tensor< + AscendC::LocalTensor, + Catlass::detail::TagToLayout_t, + tla::Coord, + AscendC::TPosition::A1>, + Catlass::EmptyClass>; + using TensorL0Bias = tla::Tensor< + AscendC::LocalTensor, + Catlass::detail::TagToLayout_t, + tla::Coord, + AscendC::TPosition::C2>; + using TensorL1Quant = std::conditional_t< + HAS_QUANT_TENSOR, + tla::Tensor, + Catlass::detail::TagToLayout_t, + tla::Coord, + AscendC::TPosition::A1>, + Catlass::EmptyClass>; + + using L1AAlignHelper = Catlass::Gemm::helper::L1AlignHelper; + using L1BAlignHelper = Catlass::Gemm::helper::L1AlignHelper; + + template + using CopyGmToL1A = Catlass::Gemm::Tile::TileCopyTla; + + template + using CopyGmToL1B = Catlass::Gemm::Tile::TileCopyTla; + + template + using CopyGmToL1Bias = + std::conditional_t, Catlass::EmptyClass>; + + template + using CopyGmToL1Scale = + std::conditional_t, Catlass::EmptyClass>; + + using CopyL1ToL0A = Catlass::Gemm::Tile::TileCopyTla; + using CopyL1ToL0B = Catlass::Gemm::Tile::TileCopyTla; + using CopyL1ToBT = + std::conditional_t, Catlass::EmptyClass>; + +#if (defined(CATLASS_ARCH) && CATLASS_ARCH == 3510) + template + using CopyL0CToDst = Catlass::Gemm::Tile::CopyL0CToGmTla; +#endif +#if (defined(CATLASS_ARCH) && CATLASS_ARCH == 2201) + template + using CopyL0CToGm = Catlass::Gemm::Tile::CopyL0CToGmTla; +#endif +}; + +#if defined(CATLASS_ARCH) && CATLASS_ARCH == 3510 + +template +struct CopyL0CToUBTla< + Catlass::Arch::Ascend950, + TensorSrc_, + tla::Tensor, LayoutDst_, CoordDst_, AscendC::TPosition::VECCALC>, + Catlass::Gemm::Tile::CopyL0CToUBMode::NO_SPLIT, + Catlass::Gemm::Tile::ScaleGranularity::NO_QUANT, + ReluEnable_, + std::enable_if_t::value>> { + using ArchTag = Catlass::Arch::Ascend950; + using ElementDst = ElementDst_; + using ElementSrc = typename TensorSrc_::Element; + static constexpr auto quantPre = + Catlass::Gemm::Tile::CopyL0CToDstQuantMode::VALUE; + static constexpr auto reluEn = ReluEnable_; + + template + CATLASS_DEVICE void operator()(TensorDst const &dstTensor, TensorSrc const &srcTensor, uint8_t unitFlag = 0) + { + static_assert( + tla::detail::isRowMajor::value && TensorSrc::position == AscendC::TPosition::CO1 + && TensorDst::position == AscendC::TPosition::VECCALC, + "The input parameters do not match. TensorSrc must be L0C, while TensorDst must be UB and RowMajor" + ); + + AscendC::FixpipeParamsC310 intriParams; + + // Fixpipe layout information + intriParams.nSize = tla::get<1>(dstTensor.originShape()); + intriParams.mSize = tla::get<0>(dstTensor.originShape()); + intriParams.srcStride = tla::get<1, 1>(srcTensor.stride()) / tla::get<0, 0>(srcTensor.stride()); + intriParams.dstStride = tla::get<0>(dstTensor.stride()); + + // Fixpipe auxiliary arguments + intriParams.quantPre = quantPre; + intriParams.reluEn = reluEn; + intriParams.unitFlag = unitFlag; + + auto dstOffset = dstTensor.layout()(dstTensor.coord()); + auto srcOffset = srcTensor.layout()(srcTensor.coord()); + + // Call AscendC Fixpipe + AscendC::Fixpipe( + dstTensor.data()[dstOffset], srcTensor.data()[srcOffset], intriParams); + } + + template + CATLASS_DEVICE void operator()(TensorDst const &dstTensor, TensorSrc const &srcTensor, uint32_t ubRowNum, + uint8_t beginSubBlockIdx, uint8_t sendVecNum,uint8_t unitFlag = 0) + { + static_assert( + tla::detail::isRowMajor::value && TensorSrc::position == AscendC::TPosition::CO1 + && TensorDst::position == AscendC::TPosition::VECCALC, + "The input parameters do not match. TensorSrc must be L0C, while TensorDst must be UB and RowMajor" + ); + + AscendC::FixpipeParamsC310 intriParams; + + uint32_t dstTensorRowNum = tla::get<0>(dstTensor.originShape()); + // Fixpipe layout information + intriParams.nSize = tla::get<1>(dstTensor.originShape()); + intriParams.mSize = min(dstTensorRowNum, ubRowNum); + intriParams.srcStride = tla::get<1, 1>(srcTensor.stride()) / tla::get<0, 0>(srcTensor.stride()); + intriParams.dstStride = tla::get<0>(dstTensor.stride()); + + // Fixpipe auxiliary arguments + intriParams.quantPre = quantPre; + intriParams.reluEn = reluEn; + intriParams.unitFlag = unitFlag; + intriParams.dualDstCtl = 0; + intriParams.subBlockId = beginSubBlockIdx; + + auto dstOffset = dstTensor.layout()(dstTensor.coord()); + auto srcOffset = srcTensor.layout()(srcTensor.coord()); + + // Call AscendC Fixpipe + AscendC::Fixpipe( + dstTensor.data()[dstOffset], srcTensor.data()[srcOffset], intriParams); + if (sendVecNum > 1) { + if (dstTensorRowNum > ubRowNum) { + intriParams.mSize = dstTensorRowNum - ubRowNum; + intriParams.subBlockId = (beginSubBlockIdx + 1) % 2; + srcOffset += ubRowNum * tla::get<0, 0>(srcTensor.stride()); + AscendC::Fixpipe( + dstTensor.data()[dstOffset], srcTensor.data()[srcOffset], intriParams); + } + } + } + + template + CATLASS_DEVICE void operator()(TensorDst const &dstTensor, TensorSrc const &srcTensor, + uint8_t subBlockIdx, uint8_t unitFlag = 0) + { + static_assert( + tla::detail::isRowMajor::value && TensorSrc::position == AscendC::TPosition::CO1 + && TensorDst::position == AscendC::TPosition::VECCALC, + "The input parameters do not match. TensorSrc must be L0C, while TensorDst must be UB and RowMajor" + ); + + AscendC::FixpipeParamsC310 intriParams; + + // Fixpipe layout information + intriParams.nSize = tla::get<1>(dstTensor.shape()); + intriParams.mSize = tla::get<0>(dstTensor.shape()); + intriParams.srcStride = tla::get<1, 1>(srcTensor.stride()) / tla::get<0, 0>(srcTensor.stride()); + intriParams.dstStride = tla::get<0>(dstTensor.stride()); + + // Fixpipe auxiliary arguments + intriParams.quantPre = quantPre; + intriParams.reluEn = reluEn; + intriParams.unitFlag = unitFlag; + intriParams.dualDstCtl = 0; + intriParams.subBlockId = subBlockIdx; + + auto dstOffset = dstTensor.layout()(dstTensor.coord()); + auto srcOffset = srcTensor.layout()(srcTensor.coord()); + + // Call AscendC Fixpipe + AscendC::Fixpipe( + dstTensor.data()[dstOffset], srcTensor.data()[srcOffset], intriParams); + } +}; + +template +struct CopyL0CToUBTla< + Catlass::Arch::Ascend950, + TensorSrc_, + tla::Tensor, LayoutDst_, CoordDst_, AscendC::TPosition::VECCALC>, + Catlass::Gemm::Tile::CopyL0CToUBMode::SPLIT_M, + Catlass::Gemm::Tile::ScaleGranularity::NO_QUANT, + ReluEnable_, + std::enable_if_t::value>> { + using ArchTag = Catlass::Arch::Ascend950; + using ElementDst = ElementDst_; + using ElementSrc = typename TensorSrc_::Element; + static constexpr auto quantPre = + Catlass::Gemm::Tile::CopyL0CToDstQuantMode::VALUE; + static constexpr auto reluEn = ReluEnable_; + + template + CATLASS_DEVICE void operator()(TensorDst const &dstTensor, TensorSrc const &srcTensor, uint8_t unitFlag = 0) + { + static_assert( + tla::detail::isRowMajor::value && TensorSrc::position == AscendC::TPosition::CO1 + && TensorDst::position == AscendC::TPosition::VECCALC, + "The input parameters do not match. TensorSrc must be L0C, while TensorDst must be UB and RowMajor" + ); + + AscendC::FixpipeParamsC310 intriParams; + + // Fixpipe layout information + intriParams.nSize = tla::get<1>(dstTensor.originShape()); + intriParams.mSize = RoundUp(tla::get<0>(dstTensor.originShape()), 2); // m must be even when spilt m + intriParams.srcStride = tla::get<1, 1>(srcTensor.stride()) / tla::get<0, 0>(srcTensor.stride()); + intriParams.dstStride = tla::get<0>(dstTensor.stride()); + + // Fixpipe auxiliary arguments + intriParams.quantPre = quantPre; + intriParams.reluEn = reluEn; + intriParams.unitFlag = unitFlag; + intriParams.dualDstCtl = 1; + + auto dstOffset = dstTensor.layout()(dstTensor.coord()); + auto srcOffset = srcTensor.layout()(srcTensor.coord()); + + // Call AscendC Fixpipe + AscendC::Fixpipe( + dstTensor.data()[dstOffset], srcTensor.data()[srcOffset], intriParams); + } +}; + +template +struct CopyL0CToUBTla< + Catlass::Arch::Ascend950, + TensorSrc_, + tla::Tensor, LayoutDst_, CoordDst_, AscendC::TPosition::VECCALC>, + Catlass::Gemm::Tile::CopyL0CToUBMode::SPLIT_N, + Catlass::Gemm::Tile::ScaleGranularity::NO_QUANT, + ReluEnable_, + std::enable_if_t::value>> { + using ArchTag = Catlass::Arch::Ascend950; + using ElementDst = ElementDst_; + using ElementSrc = typename TensorSrc_::Element; + static constexpr auto quantPre = + Catlass::Gemm::Tile::CopyL0CToDstQuantMode::VALUE; + static constexpr auto reluEn = ReluEnable_; + + template + CATLASS_DEVICE void operator()(TensorDst const &dstTensor, TensorSrc const &srcTensor, uint8_t unitFlag = 0) + { + static_assert( + tla::detail::isRowMajor::value && TensorSrc::position == AscendC::TPosition::CO1 + && TensorDst::position == AscendC::TPosition::VECCALC, + "The input parameters do not match. TensorSrc must be L0C, while TensorDst must be UB and RowMajor" + ); + + AscendC::FixpipeParamsC310 intriParams; + + // Fixpipe layout information + intriParams.nSize = RoundUp(tla::get<1>(dstTensor.originShape()), 32); + intriParams.mSize = tla::get<0>(dstTensor.originShape()); // m must be even when spilt m + intriParams.srcStride = tla::get<1, 1>(srcTensor.stride()) / tla::get<0, 0>(srcTensor.stride()); + intriParams.dstStride = tla::get<0>(dstTensor.stride()); + + // Fixpipe auxiliary arguments + intriParams.quantPre = quantPre; + intriParams.reluEn = reluEn; + intriParams.unitFlag = unitFlag; + intriParams.dualDstCtl = 2; + + auto dstOffset = dstTensor.layout()(dstTensor.coord()); + auto srcOffset = srcTensor.layout()(srcTensor.coord()); + + // Call AscendC Fixpipe + AscendC::Fixpipe( + dstTensor.data()[dstOffset], srcTensor.data()[srcOffset], intriParams); + } +}; + +#endif + +template < + /// Tag indicating architecture + class ArchTag, + class ElementA_, + class LayoutTagA, + class ElementB_, + class LayoutTagB, + class ElementC_, + class LayoutTagC, + class ElementBias = void, + Catlass::Gemm::Tile::CopyL0CToUBMode CopyMode_ = Catlass::Gemm::Tile::CopyL0CToUBMode::NO_SPLIT, + bool ReluEnable = false, + Catlass::Gemm::Tile::ScaleGranularity DEQUANT_GRANULARITY = Catlass::Gemm::Tile::ScaleGranularity::NO_QUANT> +struct PackedTileCopyTlaToUB + : public PackedTileCopyTla { + static constexpr Catlass::Gemm::Tile::CopyL0CToUBMode CopyMode = CopyMode_; + // 重写 CopyL0CToDst + using TensorL0C = typename PackedTileCopyTla< + ArchTag, + ElementA_, + LayoutTagA, + ElementB_, + LayoutTagB, + ElementC_, + LayoutTagC, + ElementBias>::TensorL0C; + + template + using CopyL0CToDst = + CopyL0CToUBTla; +}; + +#endif + +///////////////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace Common::Tile + +#endif // CATLASS_GEMM_TILE_ASCEND950_COPY_L0C_TO_UB_950_HPP diff --git a/csrc/moe/common/kernel_utils/vector/regbase.hpp b/csrc/moe/common/kernel_utils/vector/regbase.hpp new file mode 100644 index 000000000000..96d9e3cbc328 --- /dev/null +++ b/csrc/moe/common/kernel_utils/vector/regbase.hpp @@ -0,0 +1,147 @@ + +using namespace AscendC::MicroAPI; + +#pragma once + +constexpr static CastTrait ctHalf2Fp32Zero = { + RegLayout::ZERO, + SatMode::SAT, + MaskMergeMode::ZEROING, + AscendC::RoundMode::CAST_NONE, +}; +constexpr static CastTrait ctHalf2Fp32One = { + RegLayout::ONE, + SatMode::SAT, + MaskMergeMode::ZEROING, + AscendC::RoundMode::CAST_NONE, +}; +constexpr static CastTrait ctFp322HalfZero = { + RegLayout::ZERO, + SatMode::NO_SAT, + MaskMergeMode::MERGING, + AscendC::RoundMode::CAST_ROUND +}; +constexpr static CastTrait ctFp322HalfOne = { + RegLayout::ONE, + SatMode::NO_SAT, + MaskMergeMode::ZEROING, + AscendC::RoundMode::CAST_ROUND +}; + +constexpr uint16_t PRELOAD_NUM = 2; + +template +__simd_callee__ inline void LoadIn(RegTensor& dstReg, __ubuf__ TType* srcUb) +{ + if constexpr (BroatCast) { + if constexpr (!std::is_same()) { + LoadAlign(dstReg, srcUb); + } else { + LoadAlign(dstReg, srcUb); + } + } else { + LoadAlign(dstReg, srcUb); + } +} + +template +__simd_callee__ inline void CastHalf2Float(RegTensor& dstZeroReg, RegTensor& dstOneReg, RegTensor& srcReg, MaskReg& mask) +{ + Cast(dstZeroReg, srcReg, mask); + Cast(dstOneReg, srcReg, mask); +} + +template +__simd_callee__ inline void CastFloat2Half(RegTensor& dstReg, RegTensor& srcZeroReg, RegTensor& srcOneReg, MaskReg& mask) +{ + Cast(dstReg, srcOneReg, mask); + Cast(dstReg, srcZeroReg, mask); +} + +template +__simd_callee__ inline void HalfOrFloat2Float(RegTensor& dstReg, RegTensor& srcReg, MaskReg& maskHalf, MaskReg& maskFloat) +{ + if constexpr (!std::is_same()) { + Cast(dstReg, srcReg, maskHalf); + } else { + Duplicate(dstReg, srcReg, maskFloat); + } +} + +__simd_callee__ inline void MulFloatTwoReg(RegTensor& dstZeroReg, RegTensor& dstOneReg, + RegTensor& src1ZeroReg, RegTensor& src1OneReg, + RegTensor& src2ZeroReg, RegTensor& src2OneReg, + MaskReg& maskFloat) +{ + Mul(dstZeroReg, src1ZeroReg, src2ZeroReg, maskFloat); + Mul(dstOneReg, src1OneReg, src2OneReg, maskFloat); +} + +__simd_callee__ inline void MinsFloatTwoReg(RegTensor& dstZeroReg, RegTensor& dstOneReg, + RegTensor& srcZeroReg, RegTensor& srcOneReg, + float scalarValue, + MaskReg& maskFloat) +{ + Mins(dstZeroReg, srcZeroReg, scalarValue, maskFloat); + Mins(dstOneReg, srcOneReg, scalarValue, maskFloat); +} + +__simd_callee__ inline void ExpFloatTwoReg(RegTensor& dstZeroReg, RegTensor& dstOneReg, + RegTensor& srcZeroReg, RegTensor& srcOneReg, + MaskReg& maskFloat) +{ + Exp(dstZeroReg, srcZeroReg, maskFloat); + Exp(dstOneReg, srcOneReg, maskFloat); +} + +__simd_callee__ inline void SubFloatTwoReg(RegTensor& dstZeroReg, RegTensor& dstOneReg, + RegTensor& src1ZeroReg, RegTensor& src1OneReg, + RegTensor& src2ZeroReg, RegTensor& src2OneReg, + MaskReg& maskFloat) +{ + Sub(dstZeroReg, src1ZeroReg, src2ZeroReg, maskFloat); + Sub(dstOneReg, src1OneReg, src2OneReg, maskFloat); +} + +__simd_callee__ inline void AddFloatTwoReg(RegTensor& dstZeroReg, RegTensor& dstOneReg, + RegTensor& src1ZeroReg, RegTensor& src1OneReg, + RegTensor& src2ZeroReg, RegTensor& src2OneReg, + MaskReg& maskFloat) +{ + Add(dstZeroReg, src1ZeroReg, src2ZeroReg, maskFloat); + Add(dstOneReg, src1OneReg, src2OneReg, maskFloat); +} + +template +__simd_callee__ inline void CompareTwoReg(MaskReg& maskZeroReg, MaskReg& maskOneReg, + RegTensor& cmp1ZeroReg, RegTensor& cmp1OneReg, + RegTensor& cmp2ZeroReg, RegTensor& cmp2OneReg, + MaskReg& mask) +{ + Compare(maskZeroReg, cmp1ZeroReg, cmp2ZeroReg, mask); + Compare(maskOneReg, cmp1OneReg, cmp2OneReg, mask); +} + +template +__simd_callee__ inline void SelectTwoReg(RegTensor& dstZeroReg, RegTensor& dstOneReg, + RegTensor& src1ZeroReg, RegTensor& src1OneReg, + RegTensor& src2ZeroReg, RegTensor& src2OneReg, + MaskReg& mask1, MaskReg& mask2) +{ + Select(dstZeroReg, src1ZeroReg, src2ZeroReg, mask1); + Select(dstOneReg, src1OneReg, src2OneReg, mask2); +} + +template +__simd_callee__ inline void StoreUnAlignOut(__ubuf__ TType* dstUb, RegTensor& srcReg, MaskReg& mask, UnalignRegForStore& uStore, uint32_t copyNum) +{ + RegTensor outZeroReg; + if constexpr (!std::is_same()) { + Cast(outZeroReg, srcReg, mask); + StoreUnAlign(dstUb, outZeroReg, uStore, copyNum); + StoreUnAlignPost(dstUb, uStore, 0); + } else { + StoreUnAlign(dstUb, srcReg, uStore, copyNum); + StoreUnAlignPost(dstUb, uStore, 0); + } +} diff --git a/csrc/torch_binding.cpp b/csrc/torch_binding.cpp index d651a11333d9..2b9b16a6b34c 100644 --- a/csrc/torch_binding.cpp +++ b/csrc/torch_binding.cpp @@ -1982,7 +1982,7 @@ TORCH_LIBRARY_EXPAND(CONCAT(_C, _ascend), ops) ops.impl("chunk_fwd_o", torch::kPrivateUse1, &vllm_ascend::chunk_fwd_o); ops.def( - "chunk_kda_fwd(Tensor q, Tensor k, Tensor v, Tensor gk, Tensor beta, float scale, int chunk_size, str layout=\"BSND\", *, Tensor? initial_state=None, bool? output_final_state=False, int[]? cu_seqlens=None, int[]? chunk_indices=None, bool? return_intermediate=False, bool? safe_gate=False, bool? transpose_state_layout=False) -> (Tensor o, Tensor final_state, Tensor g, Tensor aqk, Tensor akk, Tensor w, Tensor u, Tensor qg, Tensor kg, Tensor v_new, Tensor h, Tensor initial_state_out)" + "chunk_kda_fwd(Tensor q, Tensor k, Tensor v, Tensor g, Tensor beta, float scale, int chunk_size, str layout=\"BSND\", *, Tensor? initial_state=None, bool? output_final_state=False, int[]? cu_seqlens=None, int[]? chunk_indices=None, bool? safe_gate=False, float? lower_bound=None, bool? use_gate_in_kernel=False, Tensor? A_log=None, Tensor? dt_bias=None, bool? disable_recompute=False, bool? return_intermediate_states=False, bool? state_v_first=False) -> (Tensor o, Tensor? final_state, Tensor? gk, Tensor aqk, Tensor akk, Tensor? w, Tensor? u, Tensor? qg, Tensor? kg, Tensor? v_new, Tensor? h, Tensor? initial_state_out)" ); ops.impl("chunk_kda_fwd", torch::kPrivateUse1, &vllm_ascend::chunk_kda_fwd); @@ -2633,7 +2633,7 @@ TORCH_LIBRARY_EXPAND(CONCAT(_C, _ascend), ops) ops.impl("chunk_fwd_o", torch::kPrivateUse1, &vllm_ascend::chunk_fwd_o); ops.def( - "chunk_kda_fwd(Tensor q, Tensor k, Tensor v, Tensor gk, Tensor beta, float scale, int chunk_size, str layout=\"BSND\", *, Tensor? initial_state=None, bool? output_final_state=False, int[]? cu_seqlens=None, int[]? chunk_indices=None, bool? return_intermediate=False, bool? safe_gate=False, bool? transpose_state_layout=False) -> (Tensor o, Tensor final_state, Tensor g, Tensor aqk, Tensor akk, Tensor w, Tensor u, Tensor qg, Tensor kg, Tensor v_new, Tensor h, Tensor initial_state_out)" + "chunk_kda_fwd(Tensor q, Tensor k, Tensor v, Tensor g, Tensor beta, float scale, int chunk_size, str layout=\"BSND\", *, Tensor? initial_state=None, bool? output_final_state=False, int[]? cu_seqlens=None, int[]? chunk_indices=None, bool? safe_gate=False, float? lower_bound=None, bool? use_gate_in_kernel=False, Tensor? A_log=None, Tensor? dt_bias=None, bool? disable_recompute=False, bool? return_intermediate_states=False, bool? state_v_first=False) -> (Tensor o, Tensor? final_state, Tensor? gk, Tensor aqk, Tensor akk, Tensor? w, Tensor? u, Tensor? qg, Tensor? kg, Tensor? v_new, Tensor? h, Tensor? initial_state_out)" ); ops.impl("chunk_kda_fwd", torch::kPrivateUse1, &vllm_ascend::chunk_kda_fwd); diff --git a/csrc/torch_binding_meta.cpp b/csrc/torch_binding_meta.cpp index 6c6b88c025dd..b67905e87915 100644 --- a/csrc/torch_binding_meta.cpp +++ b/csrc/torch_binding_meta.cpp @@ -1504,13 +1504,15 @@ at::Tensor chunk_fwd_o_meta( return o; } -std::tuple +std::tuple, c10::optional, at::Tensor, at::Tensor, + c10::optional, c10::optional, c10::optional, + c10::optional, c10::optional, c10::optional, + c10::optional> chunk_kda_fwd_meta( const at::Tensor &q, const at::Tensor &k, const at::Tensor &v, - const at::Tensor &gk, + const at::Tensor &g, const at::Tensor &beta, double scale, int64_t chunk_size, @@ -1519,16 +1521,25 @@ chunk_kda_fwd_meta( c10::optional output_final_state, c10::optional cu_seqlens, c10::optional chunk_indices, - c10::optional return_intermediate, c10::optional safe_gate, - c10::optional transpose_state_layout) + c10::optional lower_bound, + c10::optional use_gate_in_kernel, + const c10::optional &A_log, + const c10::optional &dt_bias, + c10::optional disable_recompute, + c10::optional return_intermediate_states, + c10::optional state_v_first) { std::string layout_str = std::string(layout); bool is_tnd = layout_str == "TND"; bool is_ntd = layout_str == "NTD"; bool is_bnsd = layout_str == "BNSD"; bool is_rank3 = is_tnd || is_ntd; - bool is_internal_layout = is_bnsd || is_ntd; + bool output_final_state_ = output_final_state.value_or(false); + bool use_gate_in_kernel_ = use_gate_in_kernel.value_or(false); + bool disable_recompute_ = disable_recompute.value_or(false); + bool return_intermediate_states_ = return_intermediate_states.value_or(false); + bool state_v_first_ = state_v_first.value_or(false); c10::SymInt B = is_rank3 ? c10::SymInt(1) : q.sym_size(0); c10::SymInt T = is_tnd ? q.sym_size(0) : @@ -1555,54 +1566,60 @@ chunk_kda_fwd_meta( total_chunks = (T + c10::SymInt(chunk_size - 1)) / c10::SymInt(chunk_size); } - at::Tensor o = at::empty_like(v); - at::Tensor final_state_work = at::empty_symint( - c10::SymDimVector{seq_num, HV, K, V}, q.options().dtype(at::kFloat)); - at::Tensor final_state = output_final_state.value_or(false) ? - final_state_work : at::empty_symint(c10::SymDimVector{c10::SymInt(0)}, q.options().dtype(at::kFloat)); - at::Tensor g = gk.scalar_type() == at::kFloat ? - gk : at::empty_symint(gk.sym_sizes(), gk.options().dtype(at::kFloat)); - c10::SymInt chunk_size_sym(chunk_size); - c10::SymDimVector aqk_shape; - if (is_rank3) { - aqk_shape = is_internal_layout ? c10::SymDimVector{HV, T, chunk_size_sym} : - c10::SymDimVector{T, HV, chunk_size_sym}; - } else { - aqk_shape = is_internal_layout ? c10::SymDimVector{B, HV, T, chunk_size_sym} : - c10::SymDimVector{B, T, HV, chunk_size_sym}; + c10::SymDimVector attn_shape = is_rank3 ? c10::SymDimVector{T, HV, V} + : c10::SymDimVector{B, T, HV, V}; + c10::SymDimVector state_shape = state_v_first_ ? c10::SymDimVector{seq_num, HV, V, K} + : c10::SymDimVector{seq_num, HV, K, V}; + c10::SymDimVector matrix_shape = is_rank3 ? c10::SymDimVector{HV, T, c10::SymInt(chunk_size)} + : c10::SymDimVector{B, HV, T, c10::SymInt(chunk_size)}; + c10::SymDimVector k_shape = is_rank3 ? c10::SymDimVector{HV, T, K} + : c10::SymDimVector{B, HV, T, K}; + c10::SymDimVector v_shape = is_rank3 ? c10::SymDimVector{HV, T, V} + : c10::SymDimVector{B, HV, T, V}; + c10::SymDimVector h_shape = + is_rank3 ? (state_v_first_ ? c10::SymDimVector{total_chunks, HV, V, K} + : c10::SymDimVector{total_chunks, HV, K, V}) + : (state_v_first_ ? c10::SymDimVector{B, total_chunks, HV, V, K} + : c10::SymDimVector{B, total_chunks, HV, K, V}); + + at::Tensor o = at::empty_symint(attn_shape, v.options()); + c10::optional final_state; + if (output_final_state_) { + final_state = at::empty_symint(state_shape, q.options().dtype(at::kFloat)); } - at::Tensor aqk = at::empty_symint(aqk_shape, q.options()); + c10::optional gk; + if (!use_gate_in_kernel_ || disable_recompute_) { + gk = at::empty_symint(k_shape, q.options().dtype(at::kFloat)); + } + at::Tensor aqk = at::empty_symint(matrix_shape, q.options()); at::Tensor akk = at::empty_like(aqk); - c10::SymDimVector w_shape; - if (is_rank3) { - w_shape = is_internal_layout ? c10::SymDimVector{HV, T, K} : c10::SymDimVector{T, HV, K}; - } else { - w_shape = is_internal_layout ? c10::SymDimVector{B, HV, T, K} : c10::SymDimVector{B, T, HV, K}; + c10::optional w; + c10::optional u; + c10::optional qg; + c10::optional kg; + c10::optional v_new; + if (disable_recompute_) { + w = at::empty_symint(k_shape, q.options()); + u = at::empty_symint(v_shape, q.options()); + qg = at::empty_symint(k_shape, q.options()); + kg = at::empty_symint(k_shape, q.options()); + v_new = at::empty_symint(v_shape, q.options()); } - at::Tensor w = at::empty_symint(w_shape, q.options()); - at::Tensor u = at::empty_like(v); - at::Tensor qg = at::empty_like(w); - at::Tensor kg = at::empty_like(w); - at::Tensor v_new = at::empty_like(v); - c10::SymDimVector h_shape; - if (is_rank3) { - h_shape = is_internal_layout ? c10::SymDimVector{HV, total_chunks, K, V} : - c10::SymDimVector{total_chunks, HV, K, V}; - } else { - h_shape = is_internal_layout ? c10::SymDimVector{B, HV, total_chunks, K, V} : - c10::SymDimVector{B, total_chunks, HV, K, V}; + c10::optional h; + if (disable_recompute_ || return_intermediate_states_) { + h = at::empty_symint(h_shape, q.options()); } - at::Tensor h = at::empty_symint(h_shape, q.options()); - at::Tensor initial_state_tensor = initial_state.value_or(at::Tensor()); - at::Tensor initial_state_out = initial_state_tensor.defined() ? - initial_state_tensor : at::empty_symint(c10::SymDimVector{c10::SymInt(0)}, q.options()); + c10::optional initial_state_out = + initial_state.has_value() && initial_state->defined() ? initial_state : c10::nullopt; (void)k; + (void)g; (void)beta; (void)scale; - (void)return_intermediate; (void)safe_gate; - (void)transpose_state_layout; - return std::make_tuple(o, final_state, g, aqk, akk, w, u, qg, kg, v_new, h, initial_state_out); + (void)lower_bound; + (void)A_log; + (void)dt_bias; + return std::make_tuple(o, final_state, gk, aqk, akk, w, u, qg, kg, v_new, h, initial_state_out); } at::Tensor kda_gate_cumsum_meta( diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_chunk_kda_aclnn.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_chunk_kda_aclnn.py index f91b59188812..3b52f4278461 100644 --- a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_chunk_kda_aclnn.py +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_chunk_kda_aclnn.py @@ -15,6 +15,7 @@ # This file is a part of the vllm-ascend project. import gc +import math from dataclasses import dataclass import pytest @@ -126,18 +127,30 @@ def _cleanup_npu(): torch.npu.reset_peak_memory_stats() +def _is_ascend_950() -> bool: + try: + return "950" in torch.npu.get_device_name(0) + except Exception: + return False + + def _assert_close(name, actual, expected, rtol=5e-2, atol=5e-2): torch.testing.assert_close(actual.detach().cpu(), expected.detach().cpu(), rtol=rtol, atol=atol, msg=name) def _snapshot_outputs(outputs): torch.npu.synchronize() - return tuple(output.detach().cpu().contiguous() for output in outputs) + return tuple(None if output is None else output.detach().cpu().contiguous() for output in outputs) def _assert_outputs_bitwise_equal(reference, actual, repeat): assert len(reference) == len(actual) == len(CHUNK_KDA_OUTPUT_NAMES) for name, expected, current in zip(CHUNK_KDA_OUTPUT_NAMES, reference, actual): + if expected is None or current is None: + assert expected is None and current is None, ( + f"repeat={repeat} output={name} changed between None and Tensor" + ) + continue same_metadata = expected.shape == current.shape and expected.dtype == current.dtype same_bits = same_metadata and torch.equal(expected.view(torch.uint8), current.view(torch.uint8)) if same_bits: @@ -181,38 +194,45 @@ def test_kda_torch_bindings_have_shape_correct_meta_kernels(): v = torch.empty((1, 64, 2, 256), device="meta", dtype=torch.bfloat16) raw_gate = torch.empty((1, 64, 2, 128), device="meta", dtype=torch.bfloat16) beta = torch.empty((1, 64, 2), device="meta", dtype=torch.float32) + a_log = torch.empty((2,), device="meta", dtype=torch.float32) + dt_bias = torch.empty((2 * 128,), device="meta", dtype=torch.float32) gk = torch.ops._C_ascend.kda_gate_cumsum(raw_gate, 64, layout="BSND") outputs = torch.ops._C_ascend.chunk_kda_fwd( q, k, v, - gk, + raw_gate, beta, 128**-0.5, 64, layout="BSND", output_final_state=True, - return_intermediate=True, + safe_gate=True, + use_gate_in_kernel=True, + A_log=a_log, + dt_bias=dt_bias, + disable_recompute=True, + return_intermediate_states=True, ) swapped = torch.ops._C_ascend.kda_layout_swap12(raw_gate) assert gk.shape == raw_gate.shape assert gk.dtype == torch.float32 - assert [tuple(output.shape) for output in outputs] == [ + assert [tuple(output.shape) for output in outputs[:-1]] == [ (1, 64, 2, 256), (1, 2, 128, 256), - (1, 64, 2, 128), - (1, 64, 2, 64), - (1, 64, 2, 64), - (1, 64, 2, 128), - (1, 64, 2, 256), - (1, 64, 2, 128), - (1, 64, 2, 128), - (1, 64, 2, 256), + (1, 2, 64, 128), + (1, 2, 64, 64), + (1, 2, 64, 64), + (1, 2, 64, 128), + (1, 2, 64, 256), + (1, 2, 64, 128), + (1, 2, 64, 128), + (1, 2, 64, 256), (1, 1, 2, 128, 256), - (0,), ] + assert outputs[11] is None assert outputs[0].dtype == torch.bfloat16 assert outputs[1].dtype == torch.float32 assert swapped.shape == (1, 2, 64, 128) @@ -238,14 +258,15 @@ def test_chunk_kda_fwd_matches_reference_bsnd(): q, k, v, - gk, + g, beta, scale, 64, layout="BSND", initial_state=initial_state, output_final_state=True, - return_intermediate=True, + disable_recompute=True, + return_intermediate_states=True, ) ref = chunk_kda_forward_reference( q.cpu(), @@ -316,14 +337,15 @@ def run_chunk_kda_fwd(): q, k, v, - gk, + g, beta, scale, 64, layout="BSND", initial_state=initial_state, output_final_state=True, - return_intermediate=True, + disable_recompute=True, + return_intermediate_states=True, ) is_a5_determinism_case = ( @@ -358,6 +380,62 @@ def run_chunk_kda_fwd(): _cleanup_npu() +@pytest.mark.parametrize( + ("total_t", "disable_recompute"), + [ + pytest.param(15, False, id="single-tail-model-mode"), + pytest.param(15, True, id="single-tail-all-outputs"), + pytest.param(65, True, id="full-chunk-plus-tail-all-outputs"), + ], +) +@torch.inference_mode() +def test_chunk_kda_fwd_tail_is_bitwise_deterministic(total_t, disable_recompute): + torch.manual_seed(20260820 + total_t + int(disable_recompute)) + + shape = (1, total_t, 6, 128) + q = (torch.randn(shape) * 0.04).to(torch.bfloat16).npu() + k = (torch.randn(shape) * 0.04).to(torch.bfloat16).npu() + v = (torch.randn(shape) * 0.04).to(torch.bfloat16).npu() + raw_gate = (-7.0 + torch.randn(shape) * 0.03).to(torch.float32).npu() + beta = (torch.rand((1, total_t, 6)) * 0.2 + 0.05).to(torch.float32).npu() + initial_state = (torch.randn((1, 6, 128, 128)) * 0.01).to(torch.float32).npu() + a_log = torch.zeros(6, dtype=torch.float32, device="npu") + dt_bias = torch.zeros(6 * 128, dtype=torch.float32, device="npu") + chunk_indices = _canonical_chunk_indices([0, total_t], 64) + + def run_chunk_kda_fwd(): + return torch.ops._C_ascend.chunk_kda_fwd( + q, + k, + v, + raw_gate, + beta, + 128**-0.5, + 64, + layout="BSND", + initial_state=initial_state, + output_final_state=True, + cu_seqlens=[0, total_t], + chunk_indices=chunk_indices, + safe_gate=True, + lower_bound=-5.0, + use_gate_in_kernel=True, + A_log=a_log, + dt_bias=dt_bias, + disable_recompute=disable_recompute, + return_intermediate_states=True, + state_v_first=True, + ) + + run_chunk_kda_fwd() + reference_outputs = _snapshot_outputs(run_chunk_kda_fwd()) + for repeat in range(1, DETERMINISM_REPEATS): + current_outputs = _snapshot_outputs(run_chunk_kda_fwd()) + _assert_outputs_bitwise_equal(reference_outputs, current_outputs, repeat) + + _cleanup_npu() + + @torch.inference_mode() def test_kda_gate_cumsum_matches_reference(): torch.manual_seed(20260720) @@ -396,14 +474,15 @@ def test_chunk_kda_fwd_bnsd_layout_matches_reference(): q_bnsd, k_bnsd, v_bnsd, - gk_bnsd, + g_bnsd, beta_bns, scale, 64, layout="BNSD", initial_state=initial_state, output_final_state=True, - return_intermediate=True, + disable_recompute=True, + return_intermediate_states=True, ) gk_bsnd = gk_bnsd.transpose(1, 2).contiguous() ref = chunk_kda_forward_reference( @@ -418,9 +497,210 @@ def test_chunk_kda_fwd_bnsd_layout_matches_reference(): output_final_state=True, ) - out_bsnd = got[0].transpose(1, 2).contiguous() + out_bsnd = got[0] assert torch.isfinite(out_bsnd).all().item() assert torch.isfinite(got[1]).all().item() _assert_close("o", out_bsnd, ref.o) _assert_close("final_state", got[1], ref.final_state) _cleanup_npu() + + +def _canonical_chunk_indices(cu_seqlens, chunk_size): + if cu_seqlens is None: + return None + return [ + value + for seq_id, (start, end) in enumerate(zip(cu_seqlens[:-1], cu_seqlens[1:])) + for chunk_id in range((end - start + chunk_size - 1) // chunk_size) + for value in (seq_id, chunk_id) + ] + + +def _run_chunk_kda_fwd_a5_case( + layout, + tokens, + batch_size, + query_heads, + value_heads, + key_dim, + value_dim, + chunk_size, + dtype, + cu_seqlens, +): + torch.npu.set_device(0) + device = torch.device("npu:0") + if cu_seqlens is not None: + assert cu_seqlens[0] == 0 + assert cu_seqlens[-1] == tokens + is_tnd = layout == "TND" + q_shape = (tokens, query_heads, key_dim) if is_tnd else (batch_size, tokens, query_heads, key_dim) + v_shape = (tokens, value_heads, value_dim) if is_tnd else (batch_size, tokens, value_heads, value_dim) + g_shape = (tokens, value_heads, key_dim) if is_tnd else (batch_size, tokens, value_heads, key_dim) + beta_shape = (tokens, value_heads) if is_tnd else (batch_size, tokens, value_heads) + q = torch.full( + q_shape, + 1.0 / math.sqrt(key_dim), + dtype=dtype, + device=device, + ) + k = torch.full_like(q, 1.0 / math.sqrt(key_dim)) + v = torch.zeros(v_shape, dtype=dtype, device=device) + raw_gate = torch.full( + g_shape, + -0.005 * math.log(2.0), + dtype=torch.float32, + device=device, + ) + beta_dtype = torch.bfloat16 if dtype == torch.bfloat16 else torch.float32 + beta = torch.full(beta_shape, 0.5, dtype=beta_dtype, device=device) + chunk_indices = _canonical_chunk_indices(cu_seqlens, chunk_size) + + torch.npu.synchronize() + outputs = torch.ops._C_ascend.chunk_kda_fwd( + q, + k, + v, + raw_gate, + beta, + key_dim**-0.5, + chunk_size, + layout=layout, + output_final_state=True, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + safe_gate=True, + lower_bound=-5.0, + use_gate_in_kernel=False, + A_log=None, + dt_bias=None, + disable_recompute=False, + return_intermediate_states=False, + ) + torch.npu.synchronize() + + sequence_count = len(cu_seqlens) - 1 if cu_seqlens is not None else batch_size + expected_output_shape = (tokens, value_heads, value_dim) if is_tnd else (batch_size, tokens, value_heads, value_dim) + expected_gk_shape = (value_heads, tokens, key_dim) if is_tnd else (batch_size, value_heads, tokens, key_dim) + assert len(outputs) == len(CHUNK_KDA_OUTPUT_NAMES) + assert outputs[0].shape == expected_output_shape + assert outputs[1].shape == (sequence_count, value_heads, key_dim, value_dim) + assert outputs[2].shape == expected_gk_shape + assert torch.count_nonzero(outputs[0]).item() == 0 + assert torch.count_nonzero(outputs[1]).item() == 0 + assert torch.isfinite(outputs[2]).all().item() + + _cleanup_npu() + + +@pytest.mark.parametrize( + ("layout", "cu_seqlens"), + [ + pytest.param("BSND", None, id="BSND-dense"), + pytest.param("TND", [0, 2047, 4096, 8191], id="TND-varlen"), + ], +) +@pytest.mark.skip_global_cleanup +@torch.inference_mode() +def test_chunk_kda_fwd_a5_profile_t8191(layout, cu_seqlens): + if not _is_ascend_950(): + pytest.skip("requires an Ascend 950 device") + + _run_chunk_kda_fwd_a5_case( + layout=layout, + tokens=8191, + batch_size=1, + query_heads=16, + value_heads=32, + key_dim=128, + value_dim=128, + chunk_size=64, + dtype=torch.bfloat16, + cu_seqlens=cu_seqlens, + ) + + +@pytest.mark.parametrize( + ( + "layout", + "tokens", + "batch_size", + "query_heads", + "value_heads", + "key_dim", + "value_dim", + "chunk_size", + "dtype", + "cu_seqlens", + ), + [ + pytest.param( + "BSND", + 600, + 1, + 6, + 12, + 128, + 256, + 128, + torch.bfloat16, + [0, 127, 383, 600], + id="BSND-varlen-bf16", + ), + pytest.param( + "TND", + 257, + 1, + 6, + 12, + 128, + 256, + 128, + torch.bfloat16, + None, + id="TND-dense-bf16", + ), + pytest.param( + "TND", + 300, + 1, + 6, + 12, + 128, + 128, + 64, + torch.bfloat16, + [0, 63, 191, 300], + id="TND-varlen-bf16", + ), + ], +) +@pytest.mark.skip_global_cleanup +@torch.inference_mode() +def test_chunk_kda_fwd_a5_generalized_layouts( + layout, + tokens, + batch_size, + query_heads, + value_heads, + key_dim, + value_dim, + chunk_size, + dtype, + cu_seqlens, +): + if not _is_ascend_950(): + pytest.skip("requires an Ascend 950 device") + + _run_chunk_kda_fwd_a5_case( + layout=layout, + tokens=tokens, + batch_size=batch_size, + query_heads=query_heads, + value_heads=value_heads, + key_dim=key_dim, + value_dim=value_dim, + chunk_size=chunk_size, + dtype=dtype, + cu_seqlens=cu_seqlens, + ) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_k3_chunk_kda_tail_npu.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_k3_chunk_kda_tail_npu.py new file mode 100644 index 000000000000..5d1597de5e61 --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_k3_chunk_kda_tail_npu.py @@ -0,0 +1,280 @@ +# +# Copyright (c) 2026 Huawei Technologies Co., Ltd. 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. +# 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. +# This file is a part of the vllm-ascend project. + +"""Regression coverage for Kimi-K3 BF16 chunk-KDA tail execution.""" + +import math +from dataclasses import dataclass + +import pytest +import torch +import torch_npu + +from vllm_ascend.utils import enable_custom_op + +torch_npu.npu.config.allow_internal_format = True +enable_custom_op() + +_CHUNK_SIZE = 64 +_HEADS = 6 +_HEAD_DIM = 128 +_LOWER_BOUND = -5.0 +_MAX_REASONABLE_ABS = 1.0e6 +_TAIL_TEST_TOKENS = ( + 128, + 129, + 130, + 131, + 143, + 144, + 145, + 159, + 160, + 161, + 191, + 192, + 193, +) +_OUTPUT_NAMES = ( + "output", + "final_state", + "gk", + "aqk", + "akk", + "w", + "u", + "qg", + "kg", + "v_new", + "h", + "initial_state", +) + + +@dataclass +class _ChunkKdaInputs: + q: torch.Tensor + k: torch.Tensor + v: torch.Tensor + raw_gate: torch.Tensor + activated_gate: torch.Tensor + beta: torch.Tensor + a_log: torch.Tensor + dt_bias: torch.Tensor + initial_state: torch.Tensor + cu_seqlens: tuple[int, ...] + chunk_indices: tuple[int, ...] + + def clone(self) -> "_ChunkKdaInputs": + return _ChunkKdaInputs( + q=self.q.clone(), + k=self.k.clone(), + v=self.v.clone(), + raw_gate=self.raw_gate.clone(), + activated_gate=self.activated_gate.clone(), + beta=self.beta.clone(), + a_log=self.a_log.clone(), + dt_bias=self.dt_bias.clone(), + initial_state=self.initial_state.clone(), + cu_seqlens=self.cu_seqlens, + chunk_indices=self.chunk_indices, + ) + + +def _l2norm(value: torch.Tensor) -> torch.Tensor: + dtype = value.dtype + value_fp32 = value.float() + return (value_fp32 * torch.rsqrt((value_fp32 * value_fp32).sum(dim=-1, keepdim=True) + 1e-6)).to(dtype) + + +def _build_inputs(tokens: int) -> _ChunkKdaInputs: + seed = 20260820 + tokens + torch.manual_seed(seed) + torch.npu.manual_seed_all(seed) + + shape = (1, tokens, _HEADS, _HEAD_DIM) + q = _l2norm(torch.randn(shape, device="npu", dtype=torch.bfloat16)) + k = _l2norm(torch.randn(shape, device="npu", dtype=torch.bfloat16)) + v = torch.randn(shape, device="npu", dtype=torch.bfloat16) * 0.2 + raw_gate = torch.randn(shape, device="npu", dtype=torch.bfloat16) * 2.0 + beta = torch.sigmoid(torch.randn((1, tokens, _HEADS), device="npu", dtype=torch.float32)) + a_log = torch.empty((_HEADS,), device="npu", dtype=torch.float32).uniform_(-0.5, 0.8) + dt_bias = torch.empty((_HEADS * _HEAD_DIM,), device="npu", dtype=torch.float32).uniform_(-7.5, -1.5) + initial_state = torch.zeros( + (1, _HEADS, _HEAD_DIM, _HEAD_DIM), + device="npu", + dtype=torch.float32, + ) + activated_gate = _LOWER_BOUND * torch.sigmoid( + (raw_gate.float() + dt_bias.view(1, 1, _HEADS, _HEAD_DIM)) * a_log.exp().view(1, 1, _HEADS, 1) + ) + chunk_indices = tuple(value for chunk_index in range(math.ceil(tokens / _CHUNK_SIZE)) for value in (0, chunk_index)) + return _ChunkKdaInputs( + q=q, + k=k, + v=v, + raw_gate=raw_gate, + activated_gate=activated_gate, + beta=beta, + a_log=a_log, + dt_bias=dt_bias, + initial_state=initial_state, + cu_seqlens=(0, tokens), + chunk_indices=chunk_indices, + ) + + +def _final_state_reference(inputs: _ChunkKdaInputs) -> torch.Tensor: + k = inputs.k.detach().cpu() + v = inputs.v.detach().cpu() + gate = inputs.activated_gate.detach().cpu() + beta = inputs.beta.detach().cpu() + initial_state = inputs.initial_state.detach().cpu() + _, tokens, heads, head_dim = k.shape + final_state = torch.empty_like(initial_state, dtype=torch.float32) + + for head_index in range(heads): + state_kv = initial_state[0, head_index].float().transpose(-1, -2).contiguous() + for start in range(0, tokens, _CHUNK_SIZE): + end = min(start + _CHUNK_SIZE, tokens) + chunk_tokens = end - start + strict_causal = torch.ones((chunk_tokens, chunk_tokens), dtype=torch.bool).tril(diagonal=-1) + eye = torch.eye(chunk_tokens, dtype=torch.float32) + k_block = k[0, start:end, head_index].float() + v_block = v[0, start:end, head_index].float() + beta_block = beta[0, start:end, head_index].float() + gk_block = torch.cumsum(gate[0, start:end, head_index].float(), dim=0) / math.log(2.0) + relative_gate = gk_block[:, None, :] - gk_block[None, :, :] + gate_factor = torch.exp2(relative_gate.masked_fill(~strict_causal[:, :, None], 0.0)) + kk = torch.einsum("ik,jk,ijk->ij", k_block, k_block, gate_factor) + strict_kk = torch.where(strict_causal, kk * beta_block[:, None], 0.0) + akk_block = torch.linalg.solve_triangular(eye + strict_kk, eye, upper=False) + w_block = akk_block @ (k_block * beta_block[:, None] * torch.exp2(gk_block)) + u_block = akk_block @ (v_block * beta_block[:, None]) + kg_block = k_block * torch.exp2(gk_block[-1][None, :] - gk_block) + v_new_block = u_block - w_block @ state_kv + state_kv = torch.exp2(gk_block[-1])[:, None] * state_kv + kg_block.T @ v_new_block + final_state[0, head_index] = state_kv.transpose(-1, -2) + return final_state + + +def _run_chunk_kda(inputs: _ChunkKdaInputs, gate_mode: str, metadata_mode: str): + use_gate_in_kernel = gate_mode == "raw_gate" + gate = inputs.raw_gate if use_gate_in_kernel else inputs.activated_gate + use_varlen_metadata = metadata_mode == "varlen" + return torch.ops._C_ascend.chunk_kda_fwd( + inputs.q, + inputs.k, + inputs.v, + gate, + inputs.beta, + _HEAD_DIM**-0.5, + _CHUNK_SIZE, + layout="BSND", + initial_state=inputs.initial_state, + output_final_state=True, + cu_seqlens=inputs.cu_seqlens if use_varlen_metadata else None, + chunk_indices=inputs.chunk_indices if use_varlen_metadata else None, + safe_gate=True, + lower_bound=_LOWER_BOUND, + use_gate_in_kernel=use_gate_in_kernel, + A_log=inputs.a_log if use_gate_in_kernel else None, + dt_bias=inputs.dt_bias if use_gate_in_kernel else None, + disable_recompute=True, + return_intermediate_states=False, + state_v_first=True, + ) + + +def _snapshot_outputs(outputs) -> tuple[torch.Tensor | None, ...]: + torch.npu.synchronize() + return tuple(output.detach().cpu().contiguous() if isinstance(output, torch.Tensor) else None for output in outputs) + + +def _describe_difference( + name: str, + first: torch.Tensor | None, + second: torch.Tensor | None, +) -> str | None: + if first is None or second is None: + if first is None and second is None: + return None + return f"{name}: missing output first={type(first).__name__} second={type(second).__name__}" + same_metadata = first.shape == second.shape and first.dtype == second.dtype + first_fp32 = first.float() + second_fp32 = second.float() + values_are_finite = torch.isfinite(first_fp32).all().item() and torch.isfinite(second_fp32).all().item() + values_are_reasonable = ( + first_fp32.abs().max().item() <= _MAX_REASONABLE_ABS and second_fp32.abs().max().item() <= _MAX_REASONABLE_ABS + ) + same_bits = same_metadata and torch.equal(first.view(torch.uint8), second.view(torch.uint8)) + if same_bits and values_are_finite and values_are_reasonable: + return None + + if same_metadata: + changed = first.view(torch.uint8) != second.view(torch.uint8) + differing_elements = int(changed.reshape(-1, first.element_size()).any(dim=1).sum().item()) + max_abs_diff = (first.double() - second.double()).abs().max().item() + else: + differing_elements = -1 + max_abs_diff = float("nan") + return ( + f"{name}: shape_first={tuple(first.shape)} " + f"shape_second={tuple(second.shape)} dtype_first={first.dtype} " + f"dtype_second={second.dtype} max_abs_diff={max_abs_diff:.8e} " + f"differing_elements={differing_elements} " + f"values_are_finite={values_are_finite} " + f"values_are_reasonable={values_are_reasonable}" + ) + + +@pytest.mark.parametrize( + "tokens", + _TAIL_TEST_TOKENS, + ids=lambda tokens: f"tokens_{tokens}_remainder_{tokens % _CHUNK_SIZE}", +) +@pytest.mark.parametrize("gate_mode", ["external_gate", "raw_gate"]) +@pytest.mark.parametrize("metadata_mode", ["dense", "varlen"]) +@torch.inference_mode() +def test_kimi_k3_chunk_kda_bf16_tail_is_deterministic(tokens: int, gate_mode: str, metadata_mode: str): + if not hasattr(torch.ops._C_ascend, "chunk_kda_fwd"): + pytest.skip("requires the fused chunk KDA AscendC operator") + + inputs = _build_inputs(tokens) + # Allocate both sets before the first launch so an out-of-bounds write from + # that launch cannot change tensors allocated for the second invocation. + first_inputs = inputs.clone() + second_inputs = inputs.clone() + first = _snapshot_outputs(_run_chunk_kda(first_inputs, gate_mode, metadata_mode)) + second = _snapshot_outputs(_run_chunk_kda(second_inputs, gate_mode, metadata_mode)) + torch.testing.assert_close( + first[1], + _final_state_reference(inputs), + rtol=3e-2, + atol=3e-2, + msg=(f"final_state accuracy failed for tokens={tokens}, gate_mode={gate_mode}, metadata_mode={metadata_mode}"), + ) + + problems = [ + problem + for name, first_output, second_output in zip(_OUTPUT_NAMES, first, second) + if (problem := _describe_difference(name, first_output, second_output)) is not None + ] + assert not problems, ( + f"Kimi-K3 chunk KDA is not deterministic for tokens={tokens}, " + f"remainder={tokens % _CHUNK_SIZE}, gate_mode={gate_mode}, " + f"metadata_mode={metadata_mode}:\n" + "\n".join(problems) + ) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_kda_ascendc_npu.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_kda_ascendc_npu.py index 04486cd1ec27..6400373e68de 100644 --- a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_kda_ascendc_npu.py +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_kimi_kda_ascendc_npu.py @@ -14,7 +14,10 @@ # limitations under the License. # This file is a part of the vllm-ascend project. -"""Kimi K3 integration coverage for the AscendC prefill operators.""" +"""Kimi K3 integration coverage for the fused AscendC prefill operator.""" + +import importlib +import math import pytest import torch @@ -24,6 +27,31 @@ enable_custom_op() +CHUNK_KDA_OUTPUT_NAMES = ( + "o", + "final_state", + "gk", + "aqk", + "akk", + "w", + "u", + "qg", + "kg", + "v_new", + "h", + "initial_state", +) + + +def _has_chunk_kda_op() -> bool: + if hasattr(torch.ops._C_ascend, "chunk_kda_fwd"): + return True + try: + importlib.import_module("vllm_ascend.vllm_ascend_C") + except ImportError: + return False + return hasattr(torch.ops._C_ascend, "chunk_kda_fwd") + def _l2norm(x: torch.Tensor) -> torch.Tensor: dtype = x.dtype @@ -52,11 +80,109 @@ def _naive_kda( return out.to(dtype), state +def _chunked_kda_reference( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + gate: torch.Tensor, + beta: torch.Tensor, + initial_state_vk: torch.Tensor, +) -> tuple[torch.Tensor, ...]: + q, k, v, gate, beta, initial_state_vk = (x.detach().cpu() for x in (q, k, v, gate, beta, initial_state_vk)) + batch_size, tokens, heads, head_dim = q.shape + assert k.shape == q.shape == v.shape == gate.shape + assert beta.shape == (batch_size, tokens, heads) + + dtype = q.dtype + scale = head_dim**-0.5 + chunk_size = 64 + chunk_count = math.ceil(tokens / chunk_size) + gk = torch.empty((batch_size, heads, tokens, head_dim), dtype=torch.float32) + aqk = torch.zeros((batch_size, heads, tokens, chunk_size), dtype=dtype) + akk = torch.zeros_like(aqk) + w = torch.empty((batch_size, heads, tokens, head_dim), dtype=dtype) + u = torch.empty_like(w) + qg = torch.empty_like(w) + kg = torch.empty_like(w) + v_new = torch.empty_like(w) + h = torch.empty((batch_size, chunk_count, heads, head_dim, head_dim), dtype=dtype) + final_state_vk = torch.empty_like(initial_state_vk, dtype=torch.float32) + o = torch.empty_like(v) + + for batch_idx in range(batch_size): + for head_idx in range(heads): + state_kv = initial_state_vk[batch_idx, head_idx].float().transpose(-1, -2).contiguous() + for chunk_idx in range(chunk_count): + start = chunk_idx * chunk_size + end = min(start + chunk_size, tokens) + chunk_tokens = end - start + causal = torch.ones((chunk_tokens, chunk_tokens), dtype=torch.bool).tril() + strict_causal = torch.ones_like(causal).tril(diagonal=-1) + eye = torch.eye(chunk_tokens, dtype=torch.float32) + q_block = q[batch_idx, start:end, head_idx].float() + k_block = k[batch_idx, start:end, head_idx].float() + v_block = v[batch_idx, start:end, head_idx].float() + beta_block = beta[batch_idx, start:end, head_idx].float() + gk_block = torch.cumsum(gate[batch_idx, start:end, head_idx].float(), dim=0) / math.log(2.0) + relative_gate = gk_block[:, None, :] - gk_block[None, :, :] + gate_factor = torch.exp2(relative_gate.masked_fill(~causal[:, :, None], 0.0)) + qk = torch.einsum("ik,jk,ijk->ij", q_block, k_block, gate_factor) * scale + kk = torch.einsum("ik,jk,ijk->ij", k_block, k_block, gate_factor) + aqk_block = torch.where(causal, qk, 0.0) + strict_kk = torch.where(strict_causal, kk * beta_block[:, None], 0.0) + akk_block = torch.linalg.solve_triangular(eye + strict_kk, eye, upper=False) + + k_beta_g = k_block * beta_block[:, None] * torch.exp2(gk_block) + w_block = akk_block @ k_beta_g + u_block = akk_block @ (v_block * beta_block[:, None]) + qg_block = q_block * torch.exp2(gk_block) + kg_block = k_block * torch.exp2(gk_block[-1][None, :] - gk_block) + v_new_block = u_block - w_block @ state_kv + + gk[batch_idx, head_idx, start:end] = gk_block + aqk[batch_idx, head_idx, start:end, :chunk_tokens] = aqk_block.to(dtype) + akk[batch_idx, head_idx, start:end, :chunk_tokens] = akk_block.to(dtype) + w[batch_idx, head_idx, start:end] = w_block.to(dtype) + u[batch_idx, head_idx, start:end] = u_block.to(dtype) + qg[batch_idx, head_idx, start:end] = qg_block.to(dtype) + kg[batch_idx, head_idx, start:end] = kg_block.to(dtype) + v_new[batch_idx, head_idx, start:end] = v_new_block.to(dtype) + h[batch_idx, chunk_idx, head_idx] = state_kv.transpose(-1, -2).to(dtype) + o[batch_idx, start:end, head_idx] = (qg_block @ state_kv * scale + aqk_block @ v_new_block).to(dtype) + state_kv = torch.exp2(gk_block[-1])[:, None] * state_kv + kg_block.T @ v_new_block + final_state_vk[batch_idx, head_idx] = state_kv.transpose(-1, -2) + + return o, final_state_vk, gk, aqk, akk, w, u, qg, kg, v_new, h, initial_state_vk + + +def _assert_chunk_kda_outputs_close(actual, expected, retained_indices): + assert len(actual) == len(expected) == len(CHUNK_KDA_OUTPUT_NAMES) + for index, (name, expected_output) in enumerate(zip(CHUNK_KDA_OUTPUT_NAMES, expected)): + if index not in retained_indices: + assert actual[index] is None, f"{name} must be None" + continue + assert actual[index] is not None, f"{name} must be retained" + torch.testing.assert_close( + actual[index].detach().cpu(), + expected_output, + rtol=3e-2, + atol=3e-2, + msg=name, + ) + + +def _is_ascend_950() -> bool: + try: + return "950" in torch.npu.get_device_name(0) + except Exception: + return False + + @pytest.mark.skip_global_cleanup @torch.inference_mode() def test_kimi_k3_safe_gate_prefill_and_transposed_state_layout(): - if not hasattr(torch.ops._C_ascend, "kda_gate_cumsum") or not hasattr(torch.ops._C_ascend, "chunk_kda_fwd"): - pytest.skip("requires the KDA AscendC operators") + if not _has_chunk_kda_op(): + pytest.skip("requires the fused chunk KDA AscendC operator") torch.manual_seed(20260720) tokens, heads, head_dim = 64, 1, 128 @@ -64,7 +190,7 @@ def test_kimi_k3_safe_gate_prefill_and_transposed_state_layout(): q = _l2norm(torch.randn(1, tokens, heads, head_dim, dtype=dtype, device="npu")) k = _l2norm(torch.randn_like(q)) v = torch.randn_like(q) * 0.05 - raw_gate = torch.randn_like(q) * 0.1 + raw_gate = torch.randn(1, tokens, heads, head_dim, dtype=torch.float32, device="npu") * 0.1 beta = torch.rand(1, tokens, heads, dtype=torch.float32, device="npu").sigmoid() a_log = torch.randn(heads, dtype=torch.float32, device="npu") * 0.05 dt_bias = torch.randn(heads * head_dim, dtype=torch.float32, device="npu") * 0.05 @@ -73,41 +199,280 @@ def test_kimi_k3_safe_gate_prefill_and_transposed_state_layout(): cu_seqlens = (0, tokens) chunk_indices = (0, 0) - gate_cumsum = torch.ops._C_ascend.kda_gate_cumsum( + initial_state_kv = cache_vk.transpose(-1, -2).contiguous() + got = torch.ops._C_ascend.chunk_kda_fwd( + q, + k, + v, raw_gate, + beta, + head_dim**-0.5, 64, - A_log=a_log, - dt_bias=dt_bias, + layout="BSND", + initial_state=cache_vk, + output_final_state=True, cu_seqlens=cu_seqlens, - use_gate_in_kernel=True, + chunk_indices=chunk_indices, safe_gate=True, lower_bound=lower_bound, - layout="BSND", + use_gate_in_kernel=True, + A_log=a_log, + dt_bias=dt_bias, + disable_recompute=False, + return_intermediate_states=False, + state_v_first=True, ) - initial_state_kv = cache_vk.transpose(-1, -2).contiguous() - got = torch.ops._C_ascend.chunk_kda_fwd( + + retained = torch.ops._C_ascend.chunk_kda_fwd( q, k, - v, - gate_cumsum, - beta, + v.contiguous(), + raw_gate.contiguous(), + beta.contiguous(), head_dim**-0.5, 64, layout="BSND", - initial_state=initial_state_kv, + initial_state=cache_vk, output_final_state=True, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, - return_intermediate=False, + safe_gate=True, + lower_bound=lower_bound, + use_gate_in_kernel=True, + A_log=a_log, + dt_bias=dt_bias, + disable_recompute=True, + return_intermediate_states=True, + state_v_first=True, ) safe_gate = lower_bound * torch.sigmoid( (raw_gate.float() + dt_bias.view(1, 1, heads, head_dim)) * a_log.exp().view(1, 1, heads, 1) ) expected_out, expected_state_kv = _naive_kda(q, k, v, safe_gate, beta, initial_state_kv) + expected = _chunked_kda_reference(q, k, v, safe_gate, beta, cache_vk) torch.testing.assert_close(got[0], expected_out, rtol=3e-2, atol=3e-2) - torch.testing.assert_close(got[1], expected_state_kv, rtol=3e-2, atol=3e-2) + torch.testing.assert_close(got[1].transpose(-1, -2), expected_state_kv, rtol=3e-2, atol=3e-2) + _assert_chunk_kda_outputs_close(got, expected, retained_indices={0, 1, 3, 4, 11}) + _assert_chunk_kda_outputs_close( + retained, + expected, + retained_indices=set(range(len(CHUNK_KDA_OUTPUT_NAMES))), + ) + assert got[11] is cache_vk + assert retained[11] is cache_vk # The vLLM decode cache remains [H,V,K] after crossing the AscendC boundary. - cache_vk.copy_(got[1].transpose(-1, -2)) + cache_vk.copy_(got[1]) torch.testing.assert_close(cache_vk.transpose(-1, -2), expected_state_kv, rtol=3e-2, atol=3e-2) + + +@pytest.mark.skip_global_cleanup +@pytest.mark.parametrize(("tokens", "heads"), [(65, 1), (131, 6)]) +@torch.inference_mode() +def test_kimi_k3_a5_multichunk_all_outputs_match_reference(tokens, heads): + if not _has_chunk_kda_op(): + pytest.skip("requires the fused chunk KDA AscendC operator") + if not _is_ascend_950(): + pytest.skip("requires an Ascend 950 device") + + torch.manual_seed(20260819) + head_dim = 128 + dtype = torch.bfloat16 + q = _l2norm(torch.randn(1, tokens, heads, head_dim, dtype=dtype, device="npu")) + k = _l2norm(torch.randn_like(q)) + v = torch.randn_like(q) * 0.05 + raw_gate = torch.randn(1, tokens, heads, head_dim, dtype=torch.float32, device="npu") * 0.1 + beta = torch.rand(1, tokens, heads, dtype=torch.float32, device="npu").sigmoid() + a_log = torch.randn(heads, dtype=torch.float32, device="npu") * 0.05 + dt_bias = torch.randn(heads * head_dim, dtype=torch.float32, device="npu") * 0.05 + cache_vk = torch.randn(1, heads, head_dim, head_dim, dtype=torch.float32, device="npu") * 0.01 + chunk_indices = tuple(value for chunk_id in range(math.ceil(tokens / 64)) for value in (0, chunk_id)) + + result = torch.ops._C_ascend.chunk_kda_fwd( + q, + k, + v.contiguous(), + raw_gate.contiguous(), + beta.contiguous(), + head_dim**-0.5, + 64, + layout="BSND", + initial_state=cache_vk, + output_final_state=True, + cu_seqlens=(0, tokens), + chunk_indices=chunk_indices, + safe_gate=True, + lower_bound=-5.0, + use_gate_in_kernel=True, + A_log=a_log, + dt_bias=dt_bias, + disable_recompute=False, + return_intermediate_states=False, + state_v_first=True, + ) + retained = torch.ops._C_ascend.chunk_kda_fwd( + q, + k, + v.contiguous(), + raw_gate.contiguous(), + beta.contiguous(), + head_dim**-0.5, + 64, + layout="BSND", + initial_state=cache_vk, + output_final_state=True, + cu_seqlens=(0, tokens), + chunk_indices=chunk_indices, + safe_gate=True, + lower_bound=-5.0, + use_gate_in_kernel=True, + A_log=a_log, + dt_bias=dt_bias, + disable_recompute=True, + return_intermediate_states=True, + state_v_first=True, + ) + + safe_gate = -5.0 * torch.sigmoid( + (raw_gate.float() + dt_bias.view(1, 1, heads, head_dim)) * a_log.exp().view(1, 1, heads, 1) + ) + expected = _chunked_kda_reference(q, k, v, safe_gate, beta, cache_vk) + _assert_chunk_kda_outputs_close(result, expected, retained_indices={0, 1, 3, 4, 11}) + _assert_chunk_kda_outputs_close( + retained, + expected, + retained_indices=set(range(len(CHUNK_KDA_OUTPUT_NAMES))), + ) + + +@pytest.mark.skip_global_cleanup +@pytest.mark.parametrize("layout", ["BSND", "TND"]) +@torch.inference_mode() +def test_kimi_k3_a5_model_prefill_profile_shape(layout): + if not _has_chunk_kda_op(): + pytest.skip("requires the fused chunk KDA AscendC operator") + if not _is_ascend_950(): + pytest.skip("requires an Ascend 950 device") + + torch.manual_seed(20260819) + tokens, heads, head_dim = 8191, 12, 128 + dtype = torch.bfloat16 + q_bsnd = torch.full( + (1, tokens, heads, head_dim), + 1.0 / math.sqrt(head_dim), + dtype=dtype, + device="npu", + ) + k_bsnd = torch.full_like(q_bsnd, 1.0 / math.sqrt(head_dim)) + v_bsnd = torch.zeros_like(q_bsnd) + raw_gate_bsnd = torch.zeros((1, tokens, heads, head_dim), dtype=torch.float32, device="npu") + beta_bsnd = torch.full((1, tokens, heads), 0.5, dtype=torch.float32, device="npu") + output_shape: tuple[int, ...] + matrix_shape: tuple[int, ...] + if layout == "TND": + q, k, v = q_bsnd[0], k_bsnd[0], v_bsnd[0] + raw_gate, beta = raw_gate_bsnd[0], beta_bsnd[0] + output_shape = (tokens, heads, head_dim) + matrix_shape = (heads, tokens, 64) + else: + q, k, v = q_bsnd, k_bsnd, v_bsnd + raw_gate, beta = raw_gate_bsnd, beta_bsnd + output_shape = (1, tokens, heads, head_dim) + matrix_shape = (1, heads, tokens, 64) + a_log = torch.zeros(heads, dtype=torch.float32, device="npu") + dt_bias = torch.linspace( + -9.0, + -1.47, + heads * head_dim, + dtype=torch.float32, + device="npu", + ) + cache_vk = ( + torch.eye(head_dim, dtype=torch.float32, device="npu").view(1, 1, head_dim, head_dim).repeat(1, heads, 1, 1) + * 0.01 + ) + cu_seqlens = (0, tokens) + chunk_indices = tuple(value for chunk_id in range(math.ceil(tokens / 64)) for value in (0, chunk_id)) + + result = torch.ops._C_ascend.chunk_kda_fwd( + q, + k, + v.contiguous(), + raw_gate.contiguous(), + beta.contiguous(), + head_dim**-0.5, + 64, + layout=layout, + initial_state=cache_vk, + output_final_state=True, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + safe_gate=True, + lower_bound=-5.0, + use_gate_in_kernel=True, + A_log=a_log.reshape(-1).contiguous(), + dt_bias=dt_bias.contiguous(), + disable_recompute=False, + return_intermediate_states=False, + state_v_first=True, + ) + retained = torch.ops._C_ascend.chunk_kda_fwd( + q, + k, + v.contiguous(), + raw_gate.contiguous(), + beta.contiguous(), + head_dim**-0.5, + 64, + layout=layout, + initial_state=cache_vk, + output_final_state=True, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + safe_gate=True, + lower_bound=-5.0, + use_gate_in_kernel=True, + A_log=a_log.reshape(-1).contiguous(), + dt_bias=dt_bias.contiguous(), + disable_recompute=True, + return_intermediate_states=True, + state_v_first=True, + ) + + assert len(result) == len(CHUNK_KDA_OUTPUT_NAMES) + assert result[0].shape == output_shape + assert result[1].shape == cache_vk.shape + assert result[3].shape == result[4].shape == matrix_shape + assert result[11] is cache_vk + for index in (0, 1, 3, 4): + assert torch.isfinite(result[index]).all().item(), CHUNK_KDA_OUTPUT_NAMES[index] + for index in (2, 5, 6, 7, 8, 9, 10): + assert result[index] is None, CHUNK_KDA_OUTPUT_NAMES[index] + + token_head_shape = (heads, tokens, head_dim) + stored_token_head_shape = (1,) + token_head_shape if layout == "BSND" else token_head_shape + retained_shapes = ( + output_shape, + cache_vk.shape, + stored_token_head_shape, + matrix_shape, + matrix_shape, + stored_token_head_shape, + stored_token_head_shape, + stored_token_head_shape, + stored_token_head_shape, + stored_token_head_shape, + ( + (1, math.ceil(tokens / 64), heads, head_dim, head_dim) + if layout == "BSND" + else (math.ceil(tokens / 64), heads, head_dim, head_dim) + ), + cache_vk.shape, + ) + assert len(retained) == len(retained_shapes) == len(CHUNK_KDA_OUTPUT_NAMES) + for name, output, shape in zip(CHUNK_KDA_OUTPUT_NAMES, retained, retained_shapes): + assert output is not None, name + assert output.shape == shape, name + assert torch.isfinite(output).all().item(), name + assert retained[11] is cache_vk diff --git a/tests/ut/ops/test_kimi_kda.py b/tests/ut/ops/test_kimi_kda.py index ed9d845f628d..e25e55840a44 100644 --- a/tests/ut/ops/test_kimi_kda.py +++ b/tests/ut/ops/test_kimi_kda.py @@ -10,7 +10,6 @@ from vllm_ascend.ops.kimi_kda import ( _PACKED_CONV_WEIGHT_NAME, AscendKimiK3DeltaAttention, - AscendKimiK3MergedGateProjection, _prepare_beta, _zero_padded_output, _zero_padded_recurrent_output, @@ -91,12 +90,9 @@ def fake_upstream_init(attention, _config, _vllm_config, _prefix): enable_prompt_embeds=False, ) ) - with ( - patch( - "vllm_ascend.ops.kimi_kda.KimiK3DeltaAttention.__init__", - new=fake_upstream_init, - ), - patch("vllm_ascend.ops.kimi_kda.is_vl_model", return_value=False), + with patch( + "vllm_ascend.ops.kimi_kda.KimiK3DeltaAttention.__init__", + new=fake_upstream_init, ): attention = AscendKimiK3DeltaAttention(config, vllm_config) @@ -117,30 +113,7 @@ def test_prepare_beta_slices_and_applies_sigmoid_in_fp32(): assert torch.all((beta >= 0.0) & (beta <= 1.0)) -def test_recurrent_gate_uses_unbounded_kda_transform(): - attention = AscendKimiK3DeltaAttention.__new__(AscendKimiK3DeltaAttention) - nn.Module.__init__(attention) - attention.head_dim = 3 - attention.A_log = nn.Parameter(torch.randn(2)) - attention.dt_bias = nn.Parameter(torch.randn(6)) - raw_gate = torch.randn(1, 4, 2, 3) - expected = torch.randn(4, 2, 3) - - with patch( - "vllm_ascend.ops.kimi_kda.fused_kda_gate", - return_value=expected, - ) as fused_gate: - actual = attention._recurrent_gate(raw_gate) - - torch.testing.assert_close(actual, expected.unsqueeze(0)) - fused_gate.assert_called_once() - torch.testing.assert_close(fused_gate.call_args.args[0], raw_gate.reshape(4, 6)) - assert fused_gate.call_args.args[1] is attention.A_log - assert fused_gate.call_args.args[2] == attention.head_dim - assert fused_gate.call_args.kwargs["g_bias"] is attention.dt_bias - - -def test_prefill_accepts_unbounded_gate(): +def test_prefill_fuses_raw_gate_and_updates_v_first_state(): attention = AscendKimiK3DeltaAttention.__new__(AscendKimiK3DeltaAttention) nn.Module.__init__(attention) attention.head_dim = 2 @@ -162,27 +135,18 @@ def test_prefill_accepts_unbounded_gate(): keep_meta=None, chunk_indices_chunk64_host=(0, 0), ) - transformed_gate = torch.randn_like(raw_gate) - gate_cumsum = torch.randn_like(raw_gate, dtype=torch.float32) output = torch.randn_like(v) final_state = torch.randn(1, 1, 2, 2) with ( patch("vllm_ascend.ops.kimi_kda.clear_ssm_states"), patch("vllm_ascend.ops.kimi_kda.l2norm_fwd", side_effect=lambda x: x), - patch.object(attention, "_recurrent_gate", return_value=transformed_gate) as recurrent_gate, - patch.object( - torch.ops._C_ascend, - "kda_gate_cumsum", - return_value=gate_cumsum, - create=True, - ) as kda_gate_cumsum, patch.object( torch.ops._C_ascend, "chunk_kda_fwd", - return_value=(output, final_state), + return_value=(output, final_state, *([None] * 10)), create=True, - ), + ) as chunk_kda_fwd, ): actual = attention._run_prefill( q, @@ -197,32 +161,11 @@ def test_prefill_accepts_unbounded_gate(): ) assert actual is output - recurrent_gate.assert_called_once_with(raw_gate) - assert kda_gate_cumsum.call_args.args[0] is transformed_gate - assert kda_gate_cumsum.call_args.args[1] == 64 - assert "use_gate_in_kernel" not in kda_gate_cumsum.call_args.kwargs - - -def test_merged_gate_projection_uses_vllm_shard_loader(): - projection = AscendKimiK3MergedGateProjection.__new__( - AscendKimiK3MergedGateProjection, - ) - nn.Module.__init__(projection) - param = nn.Parameter(torch.empty(4, 3)) - loaded_weight = torch.empty(2, 3) - - with patch( - "vllm_ascend.ops.kimi_kda._KimiGDNMergedColumnParallelLinear.weight_loader", - autospec=True, - ) as weight_loader: - projection.load_shard_weight(param, loaded_weight, shard_id=2) - - weight_loader.assert_called_once_with( - projection, - param, - loaded_weight, - 2, - ) + assert chunk_kda_fwd.call_args.args[3] is raw_gate + assert chunk_kda_fwd.call_args.kwargs["use_gate_in_kernel"] is True + assert chunk_kda_fwd.call_args.kwargs["state_v_first"] is True + assert chunk_kda_fwd.call_args.kwargs["safe_gate"] is False + torch.testing.assert_close(recurrent_state[state_indices], final_state) def test_kda_empty_forward_context_clears_preallocated_output(): diff --git a/vllm_ascend/ops/kimi_kda.py b/vllm_ascend/ops/kimi_kda.py index 5ef2694c8d4e..a2055aa2e88c 100644 --- a/vllm_ascend/ops/kimi_kda.py +++ b/vllm_ascend/ops/kimi_kda.py @@ -29,7 +29,6 @@ from vllm_ascend.ops.gdn_attn_builder import AscendGDNAttentionBackend from vllm_ascend.ops.triton.fla.utils import clear_ssm_states -from vllm_ascend.ops.triton.kda.kda import fused_kda_gate _KDA_CHUNK_SIZE = 64 _PACKED_CONV_WEIGHT_NAME = "ascend_conv1d_weight" @@ -230,16 +229,6 @@ def _pack_conv_weights(self) -> None: prefer_copy=True, ) - def _recurrent_gate(self, raw_gate: torch.Tensor) -> torch.Tensor: - flat_gate = rearrange(raw_gate, "1 n h d -> n (h d)") - gate = fused_kda_gate( - flat_gate, - self.A_log, - self.head_dim, - g_bias=self.dt_bias, - ) - return gate.unsqueeze(0) - def _run_recurrent( self, q: torch.Tensor, @@ -296,50 +285,36 @@ def _run_prefill( state_indices = state_indices[keep] has_initial_state = has_initial_state[keep] - # The recurrent cache is [H, V, K], while chunk_kda_fwd consumes - # [H, K, V]. Keep the conversion at this operator boundary. + # The recurrent cache uses [H,V,K]. The fused prefill operator accepts + # that state layout directly through state_v_first. initial_state_vk = recurrent_state[state_indices].contiguous() clear_ssm_states(initial_state_vk, has_initial_state) - initial_state_kv = initial_state_vk.transpose(-1, -2).contiguous() q = l2norm_fwd(q.contiguous()) k = l2norm_fwd(k.contiguous()) - if self.gate_lower_bound is not None: - gate_cumsum = torch.ops._C_ascend.kda_gate_cumsum( - raw_gate.contiguous(), - _KDA_CHUNK_SIZE, - A_log=self.A_log.reshape(-1).contiguous(), - dt_bias=self.dt_bias.contiguous(), - cu_seqlens=cu_seqlens, - use_gate_in_kernel=True, - safe_gate=True, - lower_bound=self.gate_lower_bound, - layout="BSND", - ) - else: - gate = self._recurrent_gate(raw_gate) - gate_cumsum = torch.ops._C_ascend.kda_gate_cumsum( - gate.contiguous(), - _KDA_CHUNK_SIZE, - cu_seqlens=cu_seqlens, - layout="BSND", - ) result = torch.ops._C_ascend.chunk_kda_fwd( q, k, v.contiguous(), - gate_cumsum, + raw_gate.contiguous(), beta.contiguous(), self.head_dim**-0.5, _KDA_CHUNK_SIZE, layout="BSND", - initial_state=initial_state_kv, + initial_state=initial_state_vk, output_final_state=True, cu_seqlens=cu_seqlens, chunk_indices=prebuilt_metadata.chunk_indices_chunk64_host, - return_intermediate=False, + safe_gate=self.gate_lower_bound is not None, + lower_bound=self.gate_lower_bound if self.gate_lower_bound is not None else -5.0, + use_gate_in_kernel=True, + A_log=self.A_log.reshape(-1).contiguous(), + dt_bias=self.dt_bias.contiguous(), + disable_recompute=False, + return_intermediate_states=False, + state_v_first=True, ) - recurrent_state[state_indices] = result[1].transpose(-1, -2).contiguous().to(recurrent_state.dtype) + recurrent_state[state_indices] = result[1].to(recurrent_state.dtype) return result[0] @eager_break_during_capture From e8883909b5ab69784e060c16b5c9f919455925e5 Mon Sep 17 00:00:00 2001 From: zongersama <48584200+zongersama@users.noreply.github.com> Date: Tue, 4 Aug 2026 14:35:00 +0800 Subject: [PATCH 22/50] [Feature] Add mla_prolog_v3 and optional RoPE for MLA prolog (#13355) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit This PR adds the Ascend950 custom op `mla_prolog_v3` (Torch schema: `torch.ops._C_ascend.mla_prolog`) and enables a RoPE on/off switch for MLA prologue. **Changes** - Port `mla_prolog_v3` (host tiling / check / aclnn / kernel / torch adapter) into `csrc/attention/mla_prolog_v3`, and wire it into `build_aclnn.sh`, `torch_binding.cpp`, and `torch_binding_meta.cpp` (gated by `ASCEND_PLATFORM_950`). - **RoPE switch**: remove the `do_rope` bool attribute. `rope_sin` / `rope_cos` are optional inputs; both non-empty → RoPE on, both empty/null → RoPE off (`Dr` defaults to 64 when off). - Fix / extend head-num (`N`) validation to **[1, 128]** so large-head configs (e.g. `N=96`) are accepted. - Add operator API notes under `csrc/attention/mla_prolog_v3/docs/api.md`. **Why** MLA prologue needs a fused path on Ascend950, and some models/serving paths require turning RoPE off without a separate bool flag. Using empty `rope_sin`/`rope_cos` matches the optional-input contract and simplifies call sites. Yes (Ascend950 custom-op callers only). - New op: `torch.ops._C_ascend.mla_prolog(...)`. - **Breaking vs earlier draft API**: `do_rope` is removed from the schema. - RoPE control is now: - **ON**: pass valid `rope_sin` / `rope_cos` tensors - **OFF**: pass empty tensors for **both** `rope_sin` and `rope_cos` (must be consistent) Built and validated on Ascend950 (`SOC_VERSION=ascend950pr_9579`) with custom opp install + `vllm_ascend_C` extension. 1. **Smoke** (`script/test_mla_prolog_smoke.py`): schema has no `do_rope`; rope on/off both run; `query_rope` / `kr_cache` differ when RoPE is toggled. 2. **Functional with `N=96`**: decode/prefill, fused/unfused, rope on/off, `He∈{1024,7168}` — all passed (shape `(T, 96, 512)`). 3. **Accuracy** (`script/test_mla_prolog_accuracy.py --scenario bf16 --both --n 96 --he 7168`): - vs **CPU-REF (ATK)**: **PASS** for rope on/off (`query` fulfill ≈ 99.9%+; cache writes PASS) Example: ```bash source /set_env.sh source vllm_ascend/_cann_ops_custom/vendors/custom_transformer/bin/set_env.bash export SOC_VERSION=ascend950pr_9579 ASCEND_RT_VISIBLE_DEVICES=1 python3 -u script/test_mla_prolog_smoke.py python3 -u script/test_mla_prolog_accuracy.py --scenario bf16 --both --skip-rope-probe --n 96 --he 7168 - vLLM version: v0.26.0 - vLLM main: https://github.com/vllm-project/vllm/commit/d02df748bf9efd99022f1a062597dc3cb3808485 --------- Signed-off-by: zongersama <48584200+zongersama@users.noreply.github.com> Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> --- .gitleaks.toml | 3 +- csrc/attention/mla_prolog_v3/CMakeLists.txt | 20 + csrc/attention/mla_prolog_v3/docs/api.md | 276 ++ .../mla_prolog_v3/mla_prolog_v3_torch_adpt.h | 286 ++ .../op_api/aclnn_mla_prolog_v3_weight_nz.cpp | 274 ++ .../op_api/aclnn_mla_prolog_v3_weight_nz.h | 53 + .../mla_prolog_v3/op_host/CMakeLists.txt | 31 + .../op_host/mla_prolog_infershape.cpp | 138 + .../op_host/mla_prolog_infershape.h | 70 + .../op_host/mla_prolog_tiling.cpp | 877 ++++++ .../mla_prolog_v3/op_host/mla_prolog_tiling.h | 474 ++++ .../op_host/mla_prolog_tiling_check.cpp | 1285 +++++++++ .../op_host/mla_prolog_tiling_check.h | 215 ++ .../op_host/mla_prolog_v3_def.cpp | 309 +++ .../op_host/mla_prolog_v3_infershape.cpp | 285 ++ .../op_host/mla_prolog_v3_infershape.h | 58 + .../op_host/mla_prolog_v3_tiling.h | 26 + .../op_host/mla_prolog_v3_tiling_register.cpp | 26 + .../arch35/kernel_mla_prolog_split_m.h | 1785 +++++++++++++ .../arch35/kernel_mla_prolog_split_n.h | 2379 +++++++++++++++++ .../op_kernel/arch35/mla_prolog_comm.h | 409 +++ .../op_kernel/arch35/mla_prolog_vector_comm.h | 459 ++++ .../op_kernel/arch35/service_dequant.h | 195 ++ .../arch35/service_dynamic_quant_qn_mul_qr.h | 236 ++ .../op_kernel/arch35/service_gather_sin_cos.h | 66 + .../op_kernel/arch35/service_matmul.h | 987 +++++++ .../op_kernel/arch35/service_rms_norm.h | 159 ++ .../op_kernel/arch35/service_rope.h | 42 + .../service_rotary_position_embedding.h | 218 ++ .../op_kernel/arch35/service_scatter_cache.h | 135 + .../op_kernel/arch35/vf/vf_comm.h | 40 + .../op_kernel/arch35/vf/vf_dequant.h | 127 + .../op_kernel/arch35/vf/vf_dynamic_quant.h | 446 +++ .../op_kernel/arch35/vf/vf_mul_qr.h | 74 + .../op_kernel/arch35/vf/vf_quant_perchannel.h | 104 + .../op_kernel/arch35/vf/vf_quant_pertensor.h | 89 + .../op_kernel/arch35/vf/vf_rms_norm.h | 107 + .../op_kernel/arch35/vf/vf_rope.h | 91 + .../mla_prolog_template_tiling_key.h | 666 +++++ .../op_kernel/mla_prolog_tiling_data.h | 76 + .../mla_prolog_v3/op_kernel/mla_prolog_v3.cpp | 286 ++ csrc/build_aclnn.sh | 1 + csrc/torch_binding.cpp | 19 + csrc/torch_binding_meta.cpp | 146 + .../ops/singlecard_ops/test_mla_prolog_v3.py | 427 +++ 45 files changed, 14474 insertions(+), 1 deletion(-) create mode 100644 csrc/attention/mla_prolog_v3/CMakeLists.txt create mode 100644 csrc/attention/mla_prolog_v3/docs/api.md create mode 100644 csrc/attention/mla_prolog_v3/mla_prolog_v3_torch_adpt.h create mode 100644 csrc/attention/mla_prolog_v3/op_api/aclnn_mla_prolog_v3_weight_nz.cpp create mode 100644 csrc/attention/mla_prolog_v3/op_api/aclnn_mla_prolog_v3_weight_nz.h create mode 100644 csrc/attention/mla_prolog_v3/op_host/CMakeLists.txt create mode 100644 csrc/attention/mla_prolog_v3/op_host/mla_prolog_infershape.cpp create mode 100644 csrc/attention/mla_prolog_v3/op_host/mla_prolog_infershape.h create mode 100644 csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.cpp create mode 100644 csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.h create mode 100644 csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling_check.cpp create mode 100644 csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling_check.h create mode 100644 csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_def.cpp create mode 100644 csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_infershape.cpp create mode 100644 csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_infershape.h create mode 100644 csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_tiling.h create mode 100644 csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_tiling_register.cpp create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_m.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_n.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/mla_prolog_comm.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/mla_prolog_vector_comm.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/service_dequant.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/service_dynamic_quant_qn_mul_qr.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/service_gather_sin_cos.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/service_matmul.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/service_rms_norm.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/service_rope.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/service_rotary_position_embedding.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/service_scatter_cache.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_comm.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_dequant.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_dynamic_quant.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_mul_qr.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_quant_perchannel.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_quant_pertensor.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_rms_norm.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_rope.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_template_tiling_key.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_tiling_data.h create mode 100644 csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_v3.cpp create mode 100644 tests/e2e/nightly/single_node/ops/singlecard_ops/test_mla_prolog_v3.py diff --git a/.gitleaks.toml b/.gitleaks.toml index 476350aacf5e..7c9b4f9c40cb 100644 --- a/.gitleaks.toml +++ b/.gitleaks.toml @@ -56,7 +56,8 @@ paths = ["^vllm-empty/"] # Allow tilingKey generic-api-key false positives in the specified file [[allowlists]] paths = [ - "^csrc/notify_dispatch/op_host/notify_dispatch_tiling.cpp" + "^csrc/notify_dispatch/op_host/notify_dispatch_tiling.cpp", + "^csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_template_tiling_key.h" ] rules = ["generic-api-key"] diff --git a/csrc/attention/mla_prolog_v3/CMakeLists.txt b/csrc/attention/mla_prolog_v3/CMakeLists.txt new file mode 100644 index 000000000000..e3fb52cb5ca8 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/CMakeLists.txt @@ -0,0 +1,20 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2025 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- + +message(STATUS "=== Debug: start ops.mla_prolog_v3.CMakeLists.txt ") +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +if(NOT ENABLE_TEST AND NOT BENCHMARK) + list(REMOVE_ITEM CURRENT_DIRS tests) +endif() +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/docs/api.md b/csrc/attention/mla_prolog_v3/docs/api.md new file mode 100644 index 000000000000..3e54e7cfbf92 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/docs/api.md @@ -0,0 +1,276 @@ +# MlaPrologV3 API 与调用示例 + +## 1. API 总览 + +| 通路 | API/入口 | 支持情况 | +| --- | --- | --- | +| vllm-ascend 单算子入口 | `torch.ops._C_ascend.npu_mla_prolog_v3` | 支持 | +| aclnn | `aclnnMlaPrologV3WeightNzGetWorkspaceSize` / `aclnnMlaPrologV3WeightNz` | 支持 | +| Ascend C `<<<>>>` | `mla_prolog_v3<<>>` | 支持(诊断/直调;需自备 tiling) | + +各入口表达同一套 MLA 前处理融合语义:下采样 → RMSNorm → 上采样 / RoPE → 写入 KV/KR Cache(及可选量化)。 +底层算子名为 **MlaPrologV3**,权重 `weight_dq` / `weight_uq_qr` / `weight_dkv_kr` 需以 **FRACTAL_NZ** 格式传入。 + +## 2. 公共参数与约束 + +### 2.0 形状符号 + +| 符号 | 含义 | 典型/约束值 | +| --- | --- | --- | +| `B` / `S` / `T` | batch / seq / 合轴 token 数(`T=B*S`) | `T≤1M`;允许部分维为 0(空 Tensor) | +| `He` | 隐层宽度 | `{1024,2048,3072,4096,5120,6144,7168,7680,8192}` | +| `Hcq` | Query 压缩维 | `1536` | +| `Hckv` | KV 压缩维 | `512` | +| `D` | Qc 头维 | `128` | +| `Dr` | RoPE 维 | `64` | +| `N` | Query head 数 | `[1, 128]` | +| `Nkv` | KV head 数 | `1` | +| `BlockNum` / `BlockSize` | PA cache 页数/页长 | `BlockSize∈[16,1024]` 且为 16 的倍数 | + +### 2.1 输入 + +| 名称 | 必选/可选 | Shape | Dtype | Layout | 说明 | +| --- | --- | --- | --- | --- | --- | +| `token_x` | 必选 | 合轴 `(T,He)` 或非合轴 `(B,S,He)` | BF16 / INT8 / FP8_E4M3 / HIF8 | ND | 输入隐状态 | +| `weight_dq` | 必选 | `(He,Hcq)` | 同量化场景 | **FRACTAL_NZ** | \(W^{DQ}\) | +| `weight_uq_qr` | 必选 | `(Hcq,N*(D+Dr))` | 同量化场景 | **FRACTAL_NZ** | \(W^{UQ}\|W^{QR}\) | +| `weight_uk` | 必选 | `(N,D,Hckv)` | BF16 | ND | \(W^{UK}\) | +| `weight_dkv_kr` | 必选 | `(He,Hckv+Dr)` | 同量化场景 | **FRACTAL_NZ** | \(W^{DKV}\|W^{KR}\) | +| `rmsnorm_gamma_cq` | 必选 | `(Hcq,)` | BF16 | ND | Cq RMSNorm \(\gamma\) | +| `rmsnorm_gamma_ckv` | 必选 | `(Hckv,)` | BF16 | ND | Ckv RMSNorm \(\gamma\) | +| `rope_sin` / `rope_cos` | 条件必选 | 合轴 `(T,Dr)` 或非合轴 `(B,S,Dr)`;禁用 RoPE 时传 `nullptr` | BF16 | ND | RoPE 参数;同时非空时启用,同时为空时禁用,混合 null 返回错误 | +| `kv_cache` | 必选(可变) | 见 2.4 CacheMode | BF16 / INT8 / FP8… | ND | \(k^C\) 原地更新 | +| `kr_cache` | 必选(可变) | 见 2.4;`ckvkr_repo_mode=1` 时可为空 | BF16 / INT8 | ND | \(k^R\) 原地更新 | +| `cache_index` | 条件必选 | PA:`(T,)` 或 `(B,S)` 等 | INT64 | ND | PA 写 cache 槽位;取值见 2.4 | +| `dequant_scale_x` | 条件必选 | FULL/MXFP8/FP8/HIF8 必传 | FP32 / FP8_E8M0 | ND | `token_x` 反量化 | +| `dequant_scale_w_dq` | 条件必选 | 同上 | FP32 / FP8_E8M0 | ND | `weight_dq` 反量化 | +| `dequant_scale_w_uq_qr` | 条件必选 | PARTIAL 及以上必传 | FP32 / FP8_E8M0 | ND | `weight_uq_qr` 反量化 | +| `dequant_scale_w_dkv_kr` | 条件必选 | FULL 及以上必传 | FP32 / FP8_E8M0 | ND | `weight_dkv_kr` 反量化 | +| `quant_scale_ckv` / `quant_scale_ckr` | 条件必选 | KV per-channel / per-tensor 等 | FP32 | ND | cache 量化 scale | +| `smooth_scales_cq` | 可选 | `(Hcq,)` 等 | FP32 | ND | Cq 动态量化 smooth | +| `actual_seq_len` | 条件必选 | `(B,)` | INT32 | ND | `PA_BLK_*` 时必传 | +| `k_nope_clip_alpha` | 可选 | 标量/向量 | FP32 | ND | Ckv clip 缩放 | + +### 2.2 输出 + +| 名称 | Shape | Dtype | 说明 | +| --- | --- | --- | --- | +| `query` | 合轴 `(T,N,Hckv)` / 非合轴 `(B,S,N,Hckv)` | BF16 / INT8 / FP8… | \(q^N\) | +| `query_rope` | 合轴 `(T,N,Dr)` / 非合轴 `(B,S,N,Dr)` | BF16 | \(q^R\) | +| `dequant_scale_q_nope` | 全量化 + KV per-tensor 时非空,否则空 | FP32 | Query 动态量化 scale | +| `query_norm` | `query_norm_flag=True` 时非空 | BF16 / 量化 dtype | \(c^Q\) | +| `dequant_scale_q_norm` | `query_norm_flag` 且量化时非空 | FP32 / FP8_E8M0 | `query_norm` 反量化 scale | + +`kv_cache` / `kr_cache` 为可变输入:按 `cache_index` 原地写入,不作为独立 alias 输出返回。 + +### 2.3 属性 + +| 名称 | 类型 | 默认值 | 取值范围 | 说明 | +| --- | --- | --- | --- | --- | +| `rmsnorm_epsilon_cq` | float | `1e-5` | `>0` | Cq RMSNorm \(\epsilon\) | +| `rmsnorm_epsilon_ckv` | float | `1e-5` | `>0` | Ckv RMSNorm \(\epsilon\) | +| `cache_mode` | str | `"PA_BSND"` | 见 2.4 | cache 布局 | +| `query_norm_flag` | bool | `false` | `{false,true}` | 是否输出 `query_norm` | +| `weight_quant_mode` | int | `0` | `{0,1,2,3,4,5}` | 权重/激活量化模式 | +| `kv_cache_quant_mode` | int | `0` | `{0,1,2,3}` | KV cache 量化模式 | +| `query_quant_mode` | int | `0` | `{0,1}` | Query 量化;per-tensor KV 时需为 1 | +| `ckvkr_repo_mode` | int | `0` | `{0,1}` | 与 `quant_scale_repo_mode` 成对;pertile 必须为 1 | +| `quant_scale_repo_mode` | int | `0` | `{0,1}` | 同上 | +| `tile_size` | int | `128` | pertile 时必须为 `128` | per-token-per-group tile | +| `qc_qr_scale` | float | `1.0` | 有限浮点 | Query 尺度 \(\alpha_q\) | +| `kc_scale` | float | `1.0` | 有限浮点 | Key 尺度 \(\alpha_{kv}\) | + +RoPE 开关由 `ropeSin` / `ropeCos` 的 nullity 推导:同时非空 → 开启,同时为空 → 关闭;混合 null 返回参数错误。 + +#### 量化模式合法组合(`weight_quant_mode` × `kv_cache_quant_mode`) + +| wq | 含义 | 合法 kvq | +| --- | --- | --- | +| `0` | 非量化 | `{0}` | +| `1` | PARTIAL(仅 `weight_uq_qr` 量化) | `{0, 2, 3}` | +| `2` | FULL INT8 | `{0, 1, 3}` | +| `3` | MXFP8 | `{0, 1, 3}` | +| `4` | FP8 | `{0, 1, 3}` | +| `5` | HIF8 | `{0, 1, 3}` | + +`kvq`:`0` 非量化,`1` per-tensor,`2` per-channel,`3` per-tile。 + +### 2.4 CacheMode + +| `cache_mode` | `token_x` | `kv_cache` / `kr_cache`(非 pertile) | `cache_index` | +| --- | --- | --- | --- | +| `PA_BSND` / `PA_NZ` | `(T,He)` | `(BlockNum,BlockSize,Nkv,Hckv/Dr)` | `(T,)`,值 ∈ `[0, BlockNum*BlockSize)` | +| `PA_BLK_BSND` / `PA_BLK_NZ` | `(T,He)` | 同上 | block 级 index;需 `actual_seq_len` | +| `BSND` | `(B,S,He)` | `(B,S,Nkv,Hckv/Dr)` | `(B,S)` | +| `TND` | `(T,He)` | `(T,Nkv,Hckv/Dr)` | `(T,)` | + +pertile(`kvq=3`)时 `ckvkr_repo_mode=quant_scale_repo_mode=1`,`kv_cache` 末维为打包 `Dtile`,`kr_cache` 为空 Tensor。 + +## 3. aclnn API + +### 3.1 接口签名 + +```cpp +aclnnStatus aclnnMlaPrologV3WeightNzGetWorkspaceSize( + const aclTensor *tokenX, const aclTensor *weightDq, const aclTensor *weightUqQr, + const aclTensor *weightUk, const aclTensor *weightDkvKr, + const aclTensor *rmsnormGammaCq, const aclTensor *rmsnormGammaCkv, + const aclTensor *ropeSin, const aclTensor *ropeCos, + aclTensor *kvCacheRef, aclTensor *krCacheRef, + const aclTensor *cacheIndexOptional, + const aclTensor *dequantScaleXOptional, const aclTensor *dequantScaleWDqOptional, + const aclTensor *dequantScaleWUqQrOptional, const aclTensor *dequantScaleWDkvKrOptional, + const aclTensor *quantScaleCkvOptional, const aclTensor *quantScaleCkrOptional, + const aclTensor *smoothScalesCqOptional, const aclTensor *actualSeqLenOptional, + const aclTensor *kNopeClipAlphaOptional, + double rmsnormEpsilonCq, double rmsnormEpsilonCkv, char *cacheModeOptional, + int64_t weightQuantMode, int64_t kvCacheQuantMode, int64_t queryQuantMode, + int64_t ckvkrRepoMode, int64_t quantScaleRepoMode, int64_t tileSize, + double qcQrScale, double kcScale, + const aclTensor *queryOut, const aclTensor *queryRopeOut, + const aclTensor *dequantScaleQNopeOutOptional, + const aclTensor *queryNormOutOptional, const aclTensor *dequantScaleQNormOutOptional, + uint64_t *workspaceSize, aclOpExecutor **executor); + +aclnnStatus aclnnMlaPrologV3WeightNz( + void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream); +``` + +`GetWorkspaceSize` 完成参数校验与 executor 创建;第二段在传入 stream 上异步执行。 +`ropeSin` / `ropeCos` 同时非空时启用 RoPE,同时为空时禁用;一个空一个非空时返回参数错误。 +`kvCacheRef` / `krCacheRef` 同时是输入和输出。输入、输出、workspace 和 executor 必须保持有效直到 stream 完成。 + +### 3.2 调用示例 + +```cpp +// 按 2.1/2.2 创建 aclTensor;weightDq/UqQr/DkvKr 为 FRACTAL_NZ。 +uint64_t workspaceSize = 0; +aclOpExecutor *executor = nullptr; +ACLNN_CHECK(aclnnMlaPrologV3WeightNzGetWorkspaceSize( + tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, + gammaCq, gammaCkv, ropeSin, ropeCos, kvCache, krCache, + cacheIndex, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, + nullptr, nullptr, 1e-5, 1e-5, const_cast("PA_BSND"), + 0, 0, 0, 0, 0, 128, 1.0, 1.0, + queryOut, queryRopeOut, nullptr, nullptr, nullptr, + &workspaceSize, &executor)); +void *workspace = nullptr; +if (workspaceSize != 0) { + ACL_CHECK(aclrtMalloc(&workspace, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST)); +} +ACLNN_CHECK(aclnnMlaPrologV3WeightNz(workspace, workspaceSize, executor, stream)); +ACL_CHECK(aclrtSynchronizeStream(stream)); +``` + +## 4. `torch.ops._C_ascend` API + +### 4.1 接口签名 + +```python +query, query_rope, dequant_scale_q_nope, query_norm, dequant_scale_q_norm = ( + torch.ops._C_ascend.npu_mla_prolog_v3( + token_x, weight_dq, weight_uq_qr, weight_uk, weight_dkv_kr, + rmsnorm_gamma_cq, rmsnorm_gamma_ckv, rope_sin, rope_cos, + kv_cache, kr_cache, # mutable + *, + cache_index=None, + dequant_scale_x=None, dequant_scale_w_dq=None, + dequant_scale_w_uq_qr=None, dequant_scale_w_dkv_kr=None, + quant_scale_ckv=None, quant_scale_ckr=None, smooth_scales_cq=None, + actual_seq_len=None, k_nope_clip_alpha=None, + rmsnorm_epsilon_cq=1e-5, rmsnorm_epsilon_ckv=1e-5, + cache_mode="PA_BSND", query_norm_flag=False, + weight_quant_mode=0, kv_cache_quant_mode=0, query_quant_mode=0, + ckvkr_repo_mode=0, quant_scale_repo_mode=0, tile_size=128, + qc_qr_scale=1.0, kc_scale=1.0, + ) +) +``` + +仅在 Ascend950 构建且加载 `vllm_ascend_C` + 自定义 opp 后可用。 +`rope_sin` / `rope_cos` 为必传位置参数:同时非空启用 RoPE,同时为空(`numel()==0`)禁用;不允许一空一非空。 +`token_x` rank=2 为合轴 `(T,He)`,rank=3 为 `(B,S,He)`。 +`kv_cache` / `kr_cache` 原地更新;不需要的 optional 输出以空 Tensor 返回。 + +NZ 权重可用 `torch_npu.npu_format_cast(w.contiguous(), 29)` 转换。 + +### 4.2 调用示例(bf16 / PA_BSND) + +```python +import torch +import torch_npu + +# 需已加载 vllm_ascend_C,并 source 自定义 opp set_env.bash +torch_npu.npu.config.allow_internal_format = True +t, he, n = 2, 1024, 8 +hcq, hckv, d, dr = 1536, 512, 128, 64 +device, dtype = "npu:0", torch.bfloat16 + +def rnd(*shape): + return torch.randn(*shape, device=device, dtype=dtype) + +token_x = rnd(t, he) +weight_dq = torch_npu.npu_format_cast(rnd(he, hcq).contiguous(), 29) +weight_uq_qr = torch_npu.npu_format_cast(rnd(hcq, n * (d + dr)).contiguous(), 29) +weight_uk = rnd(n, d, hckv) +weight_dkv_kr = torch_npu.npu_format_cast(rnd(he, hckv + dr).contiguous(), 29) +gamma_cq = torch.ones(hcq, device=device, dtype=dtype) +gamma_ckv = torch.ones(hckv, device=device, dtype=dtype) +rope_cos = rnd(t, dr) +rope_sin = rnd(t, dr) +kv_cache = torch.zeros(2, 128, 1, hckv, device=device, dtype=dtype) +kr_cache = torch.zeros(2, 128, 1, dr, device=device, dtype=dtype) +cache_index = torch.arange(t, device=device, dtype=torch.int64) + +query, query_rope, *_ = torch.ops._C_ascend.npu_mla_prolog_v3( + token_x, weight_dq, weight_uq_qr, weight_uk, weight_dkv_kr, + gamma_cq, gamma_ckv, rope_sin, rope_cos, kv_cache, kr_cache, + cache_index=cache_index, cache_mode="PA_BSND") +# RoPE disabled: pass empty tensors for both rope inputs +empty_rope = torch.empty(0, device=device, dtype=dtype) +q_no_rope, qr_no_rope, *_ = torch.ops._C_ascend.npu_mla_prolog_v3( + token_x, weight_dq, weight_uq_qr, weight_uk, weight_dkv_kr, + gamma_cq, gamma_ckv, empty_rope, empty_rope, kv_cache.clone(), kr_cache.clone(), + cache_index=cache_index, cache_mode="PA_BSND") +torch.npu.synchronize() +# query: [T,N,Hckv], query_rope: [T,N,Dr];kv/kr_cache 已按 cache_index 写入 +``` + +## 5. Ascend C `<<<>>>` 直调 + +`blockDim`、workspace 与序列化 tiling data 必须来自同一组 host tiling。参数顺序与 kernel 定义一致: + +```cpp +mla_prolog_v3<<>>( + tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, + rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, + kvCache, krCache, cacheIndex, + dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, dequantScaleWDkvKr, + quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, kvCacheOut, krCacheOut, + dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, + workspace, tiling); +``` + +直调通路只作 route/诊断入口;公开 Python / aclnn 负责完整校验。GM 按连续物理布局解释。 + +## 6. 已知限制 + +- Torch schema 始终注册,实际可用性取决于 `csrc/build_aclnn.sh` 是否按 **Ascend950** 构建并安装了该自定义算子包。 +- `weight_dq` / `weight_uq_qr` / `weight_dkv_kr` 必须为 **FRACTAL_NZ**。 +- `Hcq=1536`,`Hckv=512`,`D=128`,`Dr=64`,`Nkv=1`;`He` 仅白名单集合;`N∈[1,128]`。 +- `weight_quant_mode` 与 `kv_cache_quant_mode` 必须落在 §2.3 合法表内。 +- pertile 要求 `ckvkr_repo_mode=quant_scale_repo_mode=1` 且 `tile_size=128`;`kr_cache` 为空。 +- KV per-tensor 时 `query_quant_mode` 必须为 `1`。 +- `PA_BLK_*` 需要 `actual_seq_len`;末项语义与合轴 `T` 一致。 +- RoPE 开关由 `rope_sin` / `rope_cos` 是否为空决定:同时非空启用,同为空(`numel()==0`)禁用。 +- B/S/T/Skv 允许为 0:空 query 时不更新 cache;Skv=0 时正常算 query 但不写 cache。 + +## 7. 异常与返回码 + +| 条件 | 返回码/异常 | +| --- | --- | +| 必选 tensor、workspaceSize 或 executor 为空 | `ACLNN_ERR_PARAM_NULLPTR` | +| rank/shape/dtype/layout、量化组合或 CacheMode 非法 | `ACLNN_ERR_PARAM_INVALID` / tiling `GRAPH_FAILED` | +| 内部 tensor 创建或 L0 调用失败 | `ACLNN_ERR_INNER_NULLPTR` | +| torch 侧未注册算子(非 950 构建)或输入非法 | `RuntimeError` / `AttributeError` | diff --git a/csrc/attention/mla_prolog_v3/mla_prolog_v3_torch_adpt.h b/csrc/attention/mla_prolog_v3/mla_prolog_v3_torch_adpt.h new file mode 100644 index 000000000000..e5bf3d0fc165 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/mla_prolog_v3_torch_adpt.h @@ -0,0 +1,286 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2026. 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. + * 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. + */ +#ifndef MLA_PROLOG_V3_TORCH_ADPT_H +#define MLA_PROLOG_V3_TORCH_ADPT_H + +namespace vllm_ascend { + +namespace { + +constexpr int64_t FP8_E4M3_BLOCK_SIZE = 32; +constexpr int64_t WEIGHT_QUANT_MODE_NO_QUANT = 0; +constexpr int64_t WEIGHT_QUANT_MODE_FULL_QUANT = 2; +constexpr int64_t WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT = 3; +constexpr int64_t WEIGHT_QUANT_MODE_FULL_QUANT_FP8 = 4; +constexpr int64_t WEIGHT_QUANT_MODE_FULL_QUANT_HIF8 = 5; +constexpr int64_t KV_CACHE_QUANT_MODE_PER_TENSOR = 1; + +bool NeedDequantScaleQNope(int64_t weight_quant_mode, int64_t kv_cache_quant_mode) +{ + return (weight_quant_mode == WEIGHT_QUANT_MODE_FULL_QUANT || + weight_quant_mode == WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT || + weight_quant_mode == WEIGHT_QUANT_MODE_FULL_QUANT_FP8 || + weight_quant_mode == WEIGHT_QUANT_MODE_FULL_QUANT_HIF8) && + kv_cache_quant_mode == KV_CACHE_QUANT_MODE_PER_TENSOR; +} + +at::ScalarType GetQueryDtype(const at::Tensor &rope_sin, int64_t weight_quant_mode, + int64_t kv_cache_quant_mode) +{ + if (weight_quant_mode == WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT && + kv_cache_quant_mode == KV_CACHE_QUANT_MODE_PER_TENSOR) { + return at::kFloat8_e4m3fn; + } + if (weight_quant_mode == WEIGHT_QUANT_MODE_FULL_QUANT && + kv_cache_quant_mode == KV_CACHE_QUANT_MODE_PER_TENSOR) { + return at::kChar; + } + // Empty rope means RoPE off; default query dtype to BF16. + if (!rope_sin.defined() || rope_sin.numel() == 0) { + return at::kBFloat16; + } + return rope_sin.scalar_type(); +} + +at::ScalarType GetQueryNormDtype(const at::Tensor &weight_uq_qr, int64_t weight_quant_mode) +{ + if (weight_quant_mode == WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT || + weight_quant_mode == WEIGHT_QUANT_MODE_FULL_QUANT_FP8) { + return at::kFloat8_e4m3fn; + } + if (weight_quant_mode == WEIGHT_QUANT_MODE_NO_QUANT) { + return at::kBFloat16; + } + return weight_uq_qr.scalar_type(); +} + +at::ScalarType GetDequantScaleQNormDtype(int64_t weight_quant_mode) +{ + if (weight_quant_mode == WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT) { + return at::kFloat8_e8m0fnu; + } + return at::kFloat; +} + +std::tuple +ConstructMlaPrologV3Outputs(const at::Tensor &token_x, const at::Tensor &weight_dq, + const at::Tensor &weight_uq_qr, const at::Tensor &weight_uk, + const at::Tensor &rope_sin, bool query_norm_flag, + int64_t weight_quant_mode, int64_t kv_cache_quant_mode) +{ + const int64_t token_x_dim = token_x.dim(); + TORCH_CHECK(token_x_dim == 2 || token_x_dim == 3, + "token_x dim num should be 2 or 3, but got ", token_x_dim); + TORCH_CHECK(weight_uk.dim() == 3, + "weight_uk dim num should be 3, but got ", weight_uk.dim()); + + std::vector query_shape; + std::vector query_rope_shape; + std::vector dequant_scale_q_nope_shape; + std::vector query_norm_shape; + std::vector dequant_scale_q_norm_shape; + + const bool rope_enabled = rope_sin.defined() && rope_sin.numel() > 0; + constexpr int64_t kDefaultRopeDim = 64; + + if (token_x_dim == 3) { + if (rope_enabled) { + TORCH_CHECK(rope_sin.dim() == 3, + "when token_x dim num is 3, rope_sin dim num should be 3, but got ", + rope_sin.dim()); + } + const int64_t rope_dim = rope_enabled ? rope_sin.size(2) : kDefaultRopeDim; + query_shape = {token_x.size(0), token_x.size(1), weight_uk.size(0), weight_uk.size(2)}; + query_rope_shape = {token_x.size(0), token_x.size(1), weight_uk.size(0), rope_dim}; + dequant_scale_q_nope_shape = {token_x.size(0) * token_x.size(1), weight_uk.size(0), 1}; + query_norm_shape = {token_x.size(0), token_x.size(1), weight_dq.size(1)}; + dequant_scale_q_norm_shape = {token_x.size(0) * token_x.size(1)}; + if (weight_quant_mode == WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT) { + dequant_scale_q_norm_shape.push_back(weight_dq.size(1) / FP8_E4M3_BLOCK_SIZE); + } else { + dequant_scale_q_norm_shape.push_back(1); + } + } else { + if (rope_enabled) { + TORCH_CHECK(rope_sin.dim() == 2, + "when token_x dim num is 2, rope_sin dim num should be 2, but got ", + rope_sin.dim()); + } + const int64_t rope_dim = rope_enabled ? rope_sin.size(1) : kDefaultRopeDim; + query_shape = {token_x.size(0), weight_uk.size(0), weight_uk.size(2)}; + query_rope_shape = {token_x.size(0), weight_uk.size(0), rope_dim}; + dequant_scale_q_nope_shape = {token_x.size(0), weight_uk.size(0), 1}; + query_norm_shape = {token_x.size(0), weight_dq.size(1)}; + dequant_scale_q_norm_shape = {token_x.size(0)}; + if (weight_quant_mode == WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT) { + dequant_scale_q_norm_shape.push_back(weight_dq.size(1) / FP8_E4M3_BLOCK_SIZE); + } else { + dequant_scale_q_norm_shape.push_back(1); + } + } + + const auto device_opts = token_x.options(); + at::Tensor query = at::empty( + query_shape, device_opts.dtype(GetQueryDtype(rope_sin, weight_quant_mode, kv_cache_quant_mode))); + at::Tensor query_rope = at::empty(query_rope_shape, device_opts.dtype(at::kBFloat16)); + + at::Tensor dequant_scale_q_nope; + if (NeedDequantScaleQNope(weight_quant_mode, kv_cache_quant_mode)) { + dequant_scale_q_nope = at::empty(dequant_scale_q_nope_shape, device_opts.dtype(at::kFloat)); + } else { + dequant_scale_q_nope = at::empty({0}, device_opts.dtype(at::kFloat)); + } + + at::Tensor query_norm; + at::Tensor dequant_scale_q_norm; + if (query_norm_flag) { + query_norm = at::empty(query_norm_shape, + device_opts.dtype(GetQueryNormDtype(weight_uq_qr, weight_quant_mode))); + if (weight_quant_mode != WEIGHT_QUANT_MODE_NO_QUANT) { + dequant_scale_q_norm = at::empty( + dequant_scale_q_norm_shape, + device_opts.dtype(GetDequantScaleQNormDtype(weight_quant_mode))); + } else { + dequant_scale_q_norm = at::empty( + {0}, device_opts.dtype(GetDequantScaleQNormDtype(weight_quant_mode))); + } + } else { + query_norm = at::empty({0}, device_opts.dtype(GetQueryNormDtype(weight_uq_qr, weight_quant_mode))); + dequant_scale_q_norm = at::empty( + {0}, device_opts.dtype(GetDequantScaleQNormDtype(weight_quant_mode))); + } + + return {query, query_rope, dequant_scale_q_nope, query_norm, dequant_scale_q_norm}; +} + +} // namespace + +// Torch schema name is npu_mla_prolog_v3 (aligned with torch_npu); underlying aclnn op is MlaPrologV3. +inline std::tuple npu_mla_prolog_v3( + const at::Tensor &token_x, + const at::Tensor &weight_dq, + const at::Tensor &weight_uq_qr, + const at::Tensor &weight_uk, + const at::Tensor &weight_dkv_kr, + const at::Tensor &rmsnorm_gamma_cq, + const at::Tensor &rmsnorm_gamma_ckv, + const at::Tensor &rope_sin, + const at::Tensor &rope_cos, + at::Tensor &kv_cache, + at::Tensor &kr_cache, + const c10::optional &cache_index, + const c10::optional &dequant_scale_x, + const c10::optional &dequant_scale_w_dq, + const c10::optional &dequant_scale_w_uq_qr, + const c10::optional &dequant_scale_w_dkv_kr, + const c10::optional &quant_scale_ckv, + const c10::optional &quant_scale_ckr, + const c10::optional &smooth_scales_cq, + const c10::optional &actual_seq_len, + const c10::optional &k_nope_clip_alpha, + double rmsnorm_epsilon_cq, + double rmsnorm_epsilon_ckv, + c10::string_view cache_mode, + bool query_norm_flag, + int64_t weight_quant_mode, + int64_t kv_cache_quant_mode, + int64_t query_quant_mode, + int64_t ckvkr_repo_mode, + int64_t quant_scale_repo_mode, + int64_t tile_size, + double qc_qr_scale, + double kc_scale) +{ + // Required args; empty (numel==0) means RoPE off. Both must be empty or both non-empty. + const bool rope_sin_empty = !rope_sin.defined() || rope_sin.numel() == 0; + const bool rope_cos_empty = !rope_cos.defined() || rope_cos.numel() == 0; + TORCH_CHECK(rope_sin_empty == rope_cos_empty, + "rope_sin and rope_cos must both be empty or both non-empty"); + + auto outputs = ConstructMlaPrologV3Outputs( + token_x, weight_dq, weight_uq_qr, weight_uk, rope_sin, query_norm_flag, + weight_quant_mode, kv_cache_quant_mode); + at::Tensor query = std::get<0>(outputs); + at::Tensor query_rope = std::get<1>(outputs); + at::Tensor dequant_scale_q_nope = std::get<2>(outputs); + at::Tensor query_norm = std::get<3>(outputs); + at::Tensor dequant_scale_q_norm = std::get<4>(outputs); + + // aclnnMlaPrologV3WeightNz derives queryNormFlag from whether optional outs are non-null. + c10::optional dequant_scale_q_nope_opt = + NeedDequantScaleQNope(weight_quant_mode, kv_cache_quant_mode) + ? c10::optional(dequant_scale_q_nope) + : c10::nullopt; + c10::optional query_norm_opt = + query_norm_flag ? c10::optional(query_norm) : c10::nullopt; + c10::optional dequant_scale_q_norm_opt = + (query_norm_flag && weight_quant_mode != WEIGHT_QUANT_MODE_NO_QUANT) + ? c10::optional(dequant_scale_q_norm) + : c10::nullopt; + + std::string cache_mode_str = std::string(cache_mode); + char *cache_mode_ptr = const_cast(cache_mode_str.c_str()); + + // Pass undefined tensors to aclnn when empty; tiling infers RoPE off from null rope. + at::Tensor rope_sin_aclnn = rope_sin_empty ? at::Tensor() : rope_sin; + at::Tensor rope_cos_aclnn = rope_cos_empty ? at::Tensor() : rope_cos; + + EXEC_NPU_CMD( + aclnnMlaPrologV3WeightNz, + token_x, + weight_dq, + weight_uq_qr, + weight_uk, + weight_dkv_kr, + rmsnorm_gamma_cq, + rmsnorm_gamma_ckv, + rope_sin_aclnn, + rope_cos_aclnn, + kv_cache, + kr_cache, + cache_index, + dequant_scale_x, + dequant_scale_w_dq, + dequant_scale_w_uq_qr, + dequant_scale_w_dkv_kr, + quant_scale_ckv, + quant_scale_ckr, + smooth_scales_cq, + actual_seq_len, + k_nope_clip_alpha, + rmsnorm_epsilon_cq, + rmsnorm_epsilon_ckv, + cache_mode_ptr, + weight_quant_mode, + kv_cache_quant_mode, + query_quant_mode, + ckvkr_repo_mode, + quant_scale_repo_mode, + tile_size, + qc_qr_scale, + kc_scale, + query, + query_rope, + dequant_scale_q_nope_opt, + query_norm_opt, + dequant_scale_q_norm_opt); + + return {query, query_rope, dequant_scale_q_nope, query_norm, dequant_scale_q_norm}; +} + +} // namespace vllm_ascend + +#endif // MLA_PROLOG_V3_TORCH_ADPT_H diff --git a/csrc/attention/mla_prolog_v3/op_api/aclnn_mla_prolog_v3_weight_nz.cpp b/csrc/attention/mla_prolog_v3/op_api/aclnn_mla_prolog_v3_weight_nz.cpp new file mode 100644 index 000000000000..2992bcd77f2e --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_api/aclnn_mla_prolog_v3_weight_nz.cpp @@ -0,0 +1,274 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ +#include +#include +#include +#include +#include "graph/types.h" +#include "aclnn_mla_prolog_v3_weight_nz.h" +#include "log/log.h" +#include "opdev/make_op_executor.h" +#include "opdev/op_dfx.h" +#include "opdev/op_executor.h" +#include "opdev/tensor_view_utils.h" +#include "opdev/op_def.h" +#include "opdev/op_log.h" +#include "opdev/common_types.h" +#include "opdev/data_type_utils.h" +#include "opdev/shape_utils.h" +#include "opdev/format_utils.h" + +using namespace op; + +#ifdef __cplusplus +extern "C" { +#endif + +namespace { + +extern aclnnStatus aclnnInnerMlaPrologV3GetWorkspaceSize( + const aclTensor *tokenX, const aclTensor *weightDq, const aclTensor *weightUqQr, const aclTensor *weightUk, + const aclTensor *weightDkvKr, const aclTensor *rmsnormGammaCq, const aclTensor *rmsnormGammaCkv, + const aclTensor *ropeSin, const aclTensor *ropeCos, aclTensor *kvCacheRef, aclTensor *krCacheRef, + const aclTensor *cacheIndexOptional, const aclTensor *dequantScaleXOptional, + const aclTensor *dequantScaleWDqOptional, const aclTensor *dequantScaleWUqQrOptional, + const aclTensor *dequantScaleWDkvKrOptional, const aclTensor *quantScaleCkvOptional, + const aclTensor *quantScaleCkrOptional, const aclTensor *smoothScalesCqOptional, + const aclTensor *actualSeqLenOptional, const aclTensor *kNopeClipAlphaOptional, double rmsnormEpsilonCq, + double rmsnormEpsilonCkv, char *cacheModeOptional, bool queryNormFlag, int64_t weightQuantMode, + int64_t kvCacheQuantMode, int64_t queryQuantMode, int64_t ckvkrRepoMode, int64_t quantScaleRepoMode, + int64_t tileSize, double qcQrScale, double kcScale, const aclTensor *queryOut, + const aclTensor *queryRopeOut, const aclTensor *dequantScaleQNopeOut, const aclTensor *queryNormOut, + const aclTensor *dequantScaleQNormOut, uint64_t *workspaceSize, aclOpExecutor **executor); + +extern aclnnStatus aclnnInnerMlaPrologV3(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, + const aclrtStream stream); + +class TensorHolder { +public: + TensorHolder(const aclTensor *&output, aclDataType dataType, std::string varName) + { + inner_ = nullptr; + name_ = varName; + if (output == nullptr) { + std::vector shape = {0}; + int64_t addr = 0xff; + inner_ = aclCreateTensor(shape.data(), shape.size(), dataType, shape.data(), 0, ACL_FORMAT_ND, shape.data(), + shape.size(), static_cast(&addr)); + output = inner_; + } + } + + ~TensorHolder() + { + if (inner_) { + aclDestroyTensor(inner_); + inner_ = nullptr; + } + } + + bool CheckTensorConditionalNotNull(bool conditional) const + { + if (inner_ && conditional) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON("MlaPrologV3", name_.c_str(), "null", + "this parameter is required under current configuration"); + return false; + } else if (!inner_ && !conditional) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON("MlaPrologV3", name_.c_str(), "not null", + "this parameter should be empty under current configuration"); + return false; + } + return true; + } + + bool IsTensorNotNull() const + { + return inner_ == nullptr; + } + +private: + const aclTensor *inner_; + std::string name_; +}; + +bool CheckWeightQuantModeValidity(int64_t weightQuantMode) +{ + std::set supportedWeightQuantMode; + if (op::GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510) { + supportedWeightQuantMode = {0LL, 1LL, 2LL, 3LL, 4LL, 5LL}; + } else { + supportedWeightQuantMode = {0LL, 1LL, 2LL}; + } + if (supportedWeightQuantMode.find(weightQuantMode) == supportedWeightQuantMode.end()) { + std::string supportedStr; + for (auto mode : supportedWeightQuantMode) { + supportedStr += std::to_string(mode) + ", "; + } + if (!supportedStr.empty()) { + supportedStr.pop_back(); + supportedStr.pop_back(); + } + OP_LOGE_FOR_INVALID_VALUE("MlaPrologV3", "weightQuantMode", std::to_string(weightQuantMode), supportedStr); + return false; + } + return true; +} + +bool CheckKvCacheQuantModeValidity(int64_t weightQuantMode, int64_t kvCacheQuantMode) +{ + std::map> supportedKvQuantMode; + if (op::GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510) { + supportedKvQuantMode = { + {0LL, {0LL}}, {1LL, {0LL, 2LL, 3LL}}, {2LL, {0LL, 1LL, 3LL}}, + {3LL, {0LL, 1LL, 3LL}}, {4LL, {0LL, 1LL, 3LL}}, {5LL, {0LL, 1LL, 3LL}}, + }; + } else { + supportedKvQuantMode = { + {0LL, {0LL}}, + {1LL, {0LL, 2LL, 3LL}}, + {2LL, {0LL, 1LL, 3LL}}, + }; + } + auto it = supportedKvQuantMode.find(weightQuantMode); + if (it == supportedKvQuantMode.end()) { + return true; // weightQuantMode itself is invalid, already checked by CheckWeightQuantModeValidity + } + if (it->second.find(kvCacheQuantMode) == it->second.end()) { + std::string supportedStr; + for (auto mode : it->second) { + supportedStr += std::to_string(mode) + ", "; + } + if (!supportedStr.empty()) { + supportedStr.pop_back(); + supportedStr.pop_back(); + } + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON("MlaPrologV3", "kvCacheQuantMode", std::to_string(kvCacheQuantMode), + "When weightQuantMode==" + std::to_string(weightQuantMode) + + ", must be within " + supportedStr); + return false; + } + return true; +} + +bool CheckQueryQuantModeValidity(int64_t queryQuantMode) +{ + std::set supportedQueryQuantMode = {0LL, 1LL}; + if (supportedQueryQuantMode.find(queryQuantMode) == supportedQueryQuantMode.end()) { + OP_LOGE_FOR_INVALID_VALUE("MlaPrologV3", "queryQuantMode", std::to_string(queryQuantMode), "0, 1"); + return false; + } + return true; +} + +aclnnStatus aclnnMlaPrologV3WeightNzGetWorkspaceSize( + const aclTensor *tokenX, const aclTensor *weightDq, const aclTensor *weightUqQr, const aclTensor *weightUk, + const aclTensor *weightDkvKr, const aclTensor *rmsnormGammaCq, const aclTensor *rmsnormGammaCkv, + const aclTensor *ropeSin, const aclTensor *ropeCos, aclTensor *kvCacheRef, aclTensor *krCacheRef, + const aclTensor *cacheIndexOptional, const aclTensor *dequantScaleXOptional, + const aclTensor *dequantScaleWDqOptional, const aclTensor *dequantScaleWUqQrOptional, + const aclTensor *dequantScaleWDkvKrOptional, const aclTensor *quantScaleCkvOptional, + const aclTensor *quantScaleCkrOptional, const aclTensor *smoothScalesCqOptional, + const aclTensor *actualSeqLenOptional, const aclTensor *kNopeClipAlphaOptional, double rmsnormEpsilonCq, + double rmsnormEpsilonCkv, char *cacheModeOptional, int64_t weightQuantMode, int64_t kvCacheQuantMode, + int64_t queryQuantMode, int64_t ckvkrRepoMode, int64_t quantScaleRepoMode, int64_t tileSize, double qcQrScale, + double kcScale, const aclTensor *queryOut, const aclTensor *queryRopeOut, + const aclTensor *dequantScaleQNopeOutOptional, const aclTensor *queryNormOutOptional, + const aclTensor *dequantScaleQNormOutOptional, uint64_t *workspaceSize, aclOpExecutor **executor) +{ + const int WEIGHT_QUANT_MODE_NO_QUANT = 0; + const int WEIGHT_QUANT_MODE_PARTIAL_QUANT = 1; + const int WEIGHT_QUANT_MODE_FULL_QUANT = 2; + const int WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT = 3; + const int WEIGHT_QUANT_MODE_FULL_QUANT_FP8 = 4; + const int WEIGHT_QUANT_MODE_FULL_QUANT_HIF8 = 5; + const int KV_CACHE_QUANT_MODE_NO_QUANT = 0; + const int KV_CACHE_QUANT_MODE_PER_TENSOR = 1; + const int KV_CACHE_QUANT_MODE_PER_CHANNEL = 2; + const int KV_CACHE_QUANT_MODE_PER_TILE = 3; + if (!CheckWeightQuantModeValidity(weightQuantMode)) { + return ge::GRAPH_FAILED; + }; + if (!CheckKvCacheQuantModeValidity(weightQuantMode, kvCacheQuantMode)) { + return ge::GRAPH_FAILED; + }; + if (!CheckQueryQuantModeValidity(queryQuantMode)) { + return ge::GRAPH_FAILED; + }; + + auto dequantScaleQNopeHolder = + TensorHolder(dequantScaleQNopeOutOptional, aclDataType::ACL_FLOAT, std::string("dequantScaleQNopeOut")); + aclDataType queryNormDataType = + weightQuantMode == WEIGHT_QUANT_MODE_NO_QUANT ? aclDataType::ACL_BF16 : aclDataType::ACL_INT8; + aclDataType dequantScaleQNormDataType = + weightQuantMode == WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT ? aclDataType::ACL_FLOAT8_E8M0 : aclDataType::ACL_FLOAT; + if (weightQuantMode == WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT || weightQuantMode == WEIGHT_QUANT_MODE_FULL_QUANT_FP8) { + queryNormDataType = aclDataType::ACL_FLOAT8_E4M3FN; + } else if (weightQuantMode == WEIGHT_QUANT_MODE_FULL_QUANT_HIF8) { + queryNormDataType = aclDataType::ACL_HIFLOAT8; + } + auto queryNormHolder = TensorHolder(queryNormOutOptional, queryNormDataType, std::string("queryNormOut")); + auto dequantScaleQNormHolder = + TensorHolder(dequantScaleQNormOutOptional, dequantScaleQNormDataType, std::string("dequantScaleQNormOut")); + if (dequantScaleQNopeOutOptional == nullptr) { + OP_LOGE_WITH_INVALID_INPUT("MlaPrologV3", "dequantScaleQNopeOut"); + return ge::GRAPH_FAILED; + } + if (queryNormOutOptional == nullptr) { + OP_LOGE_WITH_INVALID_INPUT("MlaPrologV3", "queryNormOut"); + return ge::GRAPH_FAILED; + } + if (dequantScaleQNormOutOptional == nullptr) { + OP_LOGE_WITH_INVALID_INPUT("MlaPrologV3", "dequantScaleQNormOut"); + return ge::GRAPH_FAILED; + } + // weightQuantMode == 2,4,5:全量化场景(int8,fp8,hif8) + // weightQuantMode == 3:mxfp8全量化场景 + // kvCacheQuantMode == 1:KV_PER_TENSOR量化场景 + if (!dequantScaleQNopeHolder.CheckTensorConditionalNotNull((weightQuantMode == WEIGHT_QUANT_MODE_FULL_QUANT || + weightQuantMode == WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT || + weightQuantMode == WEIGHT_QUANT_MODE_FULL_QUANT_FP8 || + weightQuantMode == WEIGHT_QUANT_MODE_FULL_QUANT_HIF8) && + kvCacheQuantMode == KV_CACHE_QUANT_MODE_PER_TENSOR)) { + return ge::GRAPH_FAILED; + } + bool queryNormFlag = queryNormHolder.IsTensorNotNull(); + // weightQuantMode != 0:量化场景 + if (!dequantScaleQNormHolder.CheckTensorConditionalNotNull(weightQuantMode != WEIGHT_QUANT_MODE_NO_QUANT && + queryNormFlag)) { + return ge::GRAPH_FAILED; + } + if ((ropeSin == nullptr) != (ropeCos == nullptr)) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON("MlaPrologV3", "ropeSin / ropeCos", + "(one null, one non-null)", + "ropeSin and ropeCos must both be non-null (RoPE enabled) or both be null (RoPE disabled)"); + return ge::GRAPH_FAILED; + } + return aclnnInnerMlaPrologV3GetWorkspaceSize( + tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, + kvCacheRef, krCacheRef, cacheIndexOptional, dequantScaleXOptional, dequantScaleWDqOptional, + dequantScaleWUqQrOptional, dequantScaleWDkvKrOptional, quantScaleCkvOptional, quantScaleCkrOptional, + smoothScalesCqOptional, actualSeqLenOptional, kNopeClipAlphaOptional, rmsnormEpsilonCq, rmsnormEpsilonCkv, + cacheModeOptional, queryNormFlag, weightQuantMode, kvCacheQuantMode, queryQuantMode, ckvkrRepoMode, + quantScaleRepoMode, tileSize, qcQrScale, kcScale, queryOut, queryRopeOut, + dequantScaleQNopeOutOptional, queryNormOutOptional, dequantScaleQNormOutOptional, workspaceSize, + executor); +} + +aclnnStatus aclnnMlaPrologV3WeightNz(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, + const aclrtStream stream) +{ + return aclnnInnerMlaPrologV3(workspace, workspaceSize, executor, stream); +} + +} // namespace + +#ifdef __cplusplus +} +#endif diff --git a/csrc/attention/mla_prolog_v3/op_api/aclnn_mla_prolog_v3_weight_nz.h b/csrc/attention/mla_prolog_v3/op_api/aclnn_mla_prolog_v3_weight_nz.h new file mode 100644 index 000000000000..dff6fd70ebf4 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_api/aclnn_mla_prolog_v3_weight_nz.h @@ -0,0 +1,53 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef ACLNN_MLA_PROLOG_V3_WEIGHT_NZ_H +#define ACLNN_MLA_PROLOG_V3_WEIGHT_NZ_H + +#include "aclnn/acl_meta.h" +#include "aclnn/aclnn_base.h" + +#ifdef __cplusplus +extern "C" { +#endif + +/** + * @brief The first interface of aclnnMlaPrologV3WeightNz calculates + * the workspace size based on the specific calculation process. + * @domain aclnn_ops_infer + */ +__attribute__((visibility("default"))) aclnnStatus aclnnMlaPrologV3WeightNzGetWorkspaceSize( + const aclTensor *tokenX, const aclTensor *weightDq, const aclTensor *weightUqQr, const aclTensor *weightUk, + const aclTensor *weightDkvKr, const aclTensor *rmsnormGammaCq, const aclTensor *rmsnormGammaCkv, + const aclTensor *ropeSin, const aclTensor *ropeCos, aclTensor *kvCacheRef, aclTensor *krCacheRef, + const aclTensor *cacheIndexOptional, const aclTensor *dequantScaleXOptional, + const aclTensor *dequantScaleWDqOptional, const aclTensor *dequantScaleWUqQrOptional, + const aclTensor *dequantScaleWDkvKrOptional, const aclTensor *quantScaleCkvOptional, + const aclTensor *quantScaleCkrOptional, const aclTensor *smoothScalesCqOptional, + const aclTensor *actualSeqLenOptional, const aclTensor *kNopeClipAlphaOptional, double rmsnormEpsilonCq, + double rmsnormEpsilonCkv, char *cacheModeOptional, int64_t weightQuantMode, int64_t kvCacheQuantMode, + int64_t queryQuantMode, int64_t ckvkrRepoMode, int64_t quantScaleRepoMode, int64_t tileSize, double qcQrScale, + double kcScale, const aclTensor *queryOut, const aclTensor *queryRopeOut, + const aclTensor *dequantScaleQNopeOutOptional, const aclTensor *queryNormOutOptional, + const aclTensor *dequantScaleQNormOutOptional, uint64_t *workspaceSize, aclOpExecutor **executor); + +/** + * @brief The second interface of aclnnMlaPrologV3WeightNz is used to perform calculations. + */ +__attribute__((visibility("default"))) aclnnStatus aclnnMlaPrologV3WeightNz(void *workspace, uint64_t workspaceSize, + aclOpExecutor *executor, + const aclrtStream stream); + + +#ifdef __cplusplus +} +#endif + +#endif // ACLNN_MLA_PROLOG_V3_WEIGHT_NZ_H diff --git a/csrc/attention/mla_prolog_v3/op_host/CMakeLists.txt b/csrc/attention/mla_prolog_v3/op_host/CMakeLists.txt new file mode 100644 index 000000000000..370002a93f9e --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_host/CMakeLists.txt @@ -0,0 +1,31 @@ +# ----------------------------------------------------------------------------------------------------------- +# Copyright (c) 2025 Huawei Technologies Co., Ltd. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +add_op_to_compiled_list() + +if (BUILD_OPEN_PROJECT) + target_sources(op_host_aclnnInner PRIVATE + mla_prolog_v3_def.cpp + ) + + add_ops_compile_options( + OP_NAME MlaPrologV3 + OPTIONS --cce-auto-sync=off + -Wno-deprecated-declarations + -Werror + -mllvm -cce-aicore-hoist-movemask=false + ) +endif() + +if (NOT BUILD_OPS_RTY_KERNEL) + add_modules_sources( + OP_API_INDEPENDENT ON + OP_API_DIR ${CMAKE_CURRENT_SOURCE_DIR}/../op_api + OPTYPE mla_prolog_v3 ACLNNTYPE aclnn_inner) +endif() \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_infershape.cpp b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_infershape.cpp new file mode 100644 index 000000000000..d375cba66295 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_infershape.cpp @@ -0,0 +1,138 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_proto.cpp + * \brief + */ + +#include "mla_prolog_infershape.h" + +using namespace ge; + +namespace ops { + +ge::graphStatus GetMlaPrologShapeDim(const gert::InferShapeContext *context, MlaPrologProtoShapeParam &shapeParam) +{ + auto tokenXShape = context->GetRequiredInputShape(TOKEN_X_INDEX); // (B, S, He) | (T, He) + OP_CHECK_NULL_WITH_CONTEXT(context, tokenXShape); + auto weightUkShape = context->GetRequiredInputShape(WEIGHT_UK_INDEX); // (N, D, Hckv) + OP_CHECK_NULL_WITH_CONTEXT(context, weightUkShape); + auto ropeSinShape = context->GetRequiredInputShape(ROPE_SIN_INDEX); // (B, S, Dr) | (T, Dr) + OP_CHECK_NULL_WITH_CONTEXT(context, ropeSinShape); + if (std::strcmp(context->GetNodeType(), "MlaPrologV3") == 0) { + auto kvCacheShape = context->GetRequiredInputShape(KV_CACHE_INDEX_V3); // (B, Nkv, Skv, Hckv) + OP_CHECK_NULL_WITH_CONTEXT(context, kvCacheShape); + auto krCacheShape = context->GetRequiredInputShape(KR_CACHE_INDEX_V3); // (B, Nkv, Skv, Dr) + OP_CHECK_NULL_WITH_CONTEXT(context, krCacheShape); + } else { + auto kvCacheShape = context->GetRequiredInputShape(KV_CACHE_INDEX); // (B, Nkv, Skv, Hckv) + OP_CHECK_NULL_WITH_CONTEXT(context, kvCacheShape); + auto krCacheShape = context->GetRequiredInputShape(KR_CACHE_INDEX); // (B, Nkv, Skv, Dr) + OP_CHECK_NULL_WITH_CONTEXT(context, krCacheShape); + } + + OP_CHECK_IF(((tokenXShape->GetDimNum() != DIM_NUM_3) && (tokenXShape->GetDimNum() != DIM_NUM_2)), + OP_LOGE(context->GetNodeName(), "tokenXShape is not 2 or 3, but %zu", tokenXShape->GetDimNum()), return ge::GRAPH_FAILED); + + if (tokenXShape->GetDimNum() == DIM_NUM_3) { // BS + shapeParam.isBsMerge = false; + shapeParam.B = tokenXShape->GetDim(DIM_INDEX_0); + shapeParam.S = tokenXShape->GetDim(DIM_INDEX_1); + shapeParam.Dr = ropeSinShape->GetDim(DIM_INDEX_2); + shapeParam.T = shapeParam.B * shapeParam.S; + } else { // T + shapeParam.isBsMerge = true; + shapeParam.T = tokenXShape->GetDim(DIM_INDEX_0); + shapeParam.Dr = ropeSinShape->GetDim(DIM_INDEX_1); + } + + shapeParam.N = weightUkShape->GetDim(DIM_INDEX_0); + shapeParam.Hckv = weightUkShape->GetDim(DIM_INDEX_2); + return GRAPH_SUCCESS; +} + +ge::graphStatus SetMlaPrologShapeDim(const MlaPrologProtoShapeParam &shapeParam, gert::InferShapeContext *context) +{ + auto queryShape = context->GetOutputShape(QUERY_INDEX); // query: (B, S, N, Hckv) | (T, N, Hckv) + OP_CHECK_NULL_WITH_CONTEXT(context, queryShape); + auto queryRopeShape = context->GetOutputShape(QUERY_ROPE_INDEX); // queryRope: (B, S, N, Dr) | (T, N, Dr) + OP_CHECK_NULL_WITH_CONTEXT(context, queryRopeShape); + auto kvCacheOutShape = context->GetOutputShape(KV_CACHE_OUT_INDEX); // kvCacheOut: (B, Nkv, Skv, Hckv) + OP_CHECK_NULL_WITH_CONTEXT(context, kvCacheOutShape); + auto krCacheOutShape = context->GetOutputShape(KR_CACHE_OUT_INDEX); // krCacheOut: (B, Nkv, Skv, Dr) + OP_CHECK_NULL_WITH_CONTEXT(context, krCacheOutShape); + + // Set output shape + if (!shapeParam.isBsMerge) { + queryShape->SetDimNum(DIM_NUM_4); // (B, S, N, Hckv) + queryShape->SetDim(DIM_INDEX_0, shapeParam.B); + queryShape->SetDim(DIM_INDEX_1, shapeParam.S); + queryShape->SetDim(DIM_INDEX_2, shapeParam.N); + queryShape->SetDim(DIM_INDEX_3, shapeParam.Hckv); + + queryRopeShape->SetDimNum(DIM_NUM_4); // (B, S, N, Dr) + queryRopeShape->SetDim(DIM_INDEX_0, shapeParam.B); + queryRopeShape->SetDim(DIM_INDEX_1, shapeParam.S); + queryRopeShape->SetDim(DIM_INDEX_2, shapeParam.N); + queryRopeShape->SetDim(DIM_INDEX_3, shapeParam.Dr); + } else { + queryShape->SetDimNum(DIM_NUM_3); // (T, N, Hckv) + queryShape->SetDim(DIM_INDEX_0, shapeParam.T); + queryShape->SetDim(DIM_INDEX_1, shapeParam.N); + queryShape->SetDim(DIM_INDEX_2, shapeParam.Hckv); + + queryRopeShape->SetDimNum(DIM_NUM_3); // (T, N, Dr) + queryRopeShape->SetDim(DIM_INDEX_0, shapeParam.T); + queryRopeShape->SetDim(DIM_INDEX_1, shapeParam.N); + queryRopeShape->SetDim(DIM_INDEX_2, shapeParam.Dr); + } + + if (std::strcmp(context->GetNodeType(), "MlaPrologV3") == 0) { + *kvCacheOutShape = *context->GetRequiredInputShape(KV_CACHE_INDEX_V3); + *krCacheOutShape = *context->GetRequiredInputShape(KR_CACHE_INDEX_V3); + } else { + *kvCacheOutShape = *context->GetRequiredInputShape(KV_CACHE_INDEX); + *krCacheOutShape = *context->GetRequiredInputShape(KR_CACHE_INDEX); + } + return GRAPH_SUCCESS; +} + +ge::graphStatus InferShapeMlaProlog(gert::InferShapeContext *context) { + OP_LOGI(context->GetNodeName(), "Enter MlaProlog infershape impl."); + + MlaPrologProtoShapeParam shapeParam {}; + auto apiRet = GetMlaPrologShapeDim(context, shapeParam); + OP_CHECK_IF((apiRet != GRAPH_SUCCESS), OP_LOGE(context->GetNodeName(), "Context get input shape failed"), return ge::GRAPH_FAILED); + + apiRet = SetMlaPrologShapeDim(shapeParam, context); + OP_CHECK_IF((apiRet != GRAPH_SUCCESS), OP_LOGE(context->GetNodeName(), "Context set output shape failed"), return ge::GRAPH_FAILED); + + OP_LOGI(context->GetNodeName(), "MlaProlog infershape end."); + return GRAPH_SUCCESS; +} + +ge::graphStatus InferDataTypeMlaProlog(gert::InferDataTypeContext *context) { + OP_LOGI(context->GetNodeName(), "Enter MlaProlog inferdatatype impl."); + + context->SetOutputDataType(QUERY_INDEX, context->GetRequiredInputDataType(WEIGHT_UK_INDEX)); + context->SetOutputDataType(QUERY_ROPE_INDEX, context->GetRequiredInputDataType(WEIGHT_UK_INDEX)); + if (std::strcmp(context->GetNodeType(), "MlaPrologV3") == 0) { + context->SetOutputDataType(KV_CACHE_OUT_INDEX, context->GetRequiredInputDataType(KV_CACHE_INDEX_V3)); + context->SetOutputDataType(KR_CACHE_OUT_INDEX, context->GetRequiredInputDataType(KR_CACHE_INDEX_V3)); + } else { + context->SetOutputDataType(KV_CACHE_OUT_INDEX, context->GetRequiredInputDataType(KV_CACHE_INDEX)); + context->SetOutputDataType(KR_CACHE_OUT_INDEX, context->GetRequiredInputDataType(KR_CACHE_INDEX)); + } + return GRAPH_SUCCESS; +} + +IMPL_OP_INFERSHAPE(MlaProlog).InferShape(InferShapeMlaProlog).InferDataType(InferDataTypeMlaProlog); +} // namespace ops \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_infershape.h b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_infershape.h new file mode 100644 index 000000000000..dcf22547cf5f --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_infershape.h @@ -0,0 +1,70 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_infershape.h + * \brief + */ + +#ifndef MLA_PROLOG_INFERSHAPE_H +#define MLA_PROLOG_INFERSHAPE_H + +#include +#include +#include "log/log.h" + +using namespace ge; + +namespace ops { +// INPUT +constexpr uint32_t TOKEN_X_INDEX = 0; +constexpr uint32_t WEIGHT_UK_INDEX = 3; +constexpr uint32_t ROPE_SIN_INDEX = 7; +constexpr uint32_t KV_CACHE_INDEX = 10; +constexpr uint32_t KR_CACHE_INDEX = 11; +constexpr uint32_t KV_CACHE_INDEX_V3 = 9; +constexpr uint32_t KR_CACHE_INDEX_V3 = 10; +// OUTPUT +constexpr uint32_t QUERY_INDEX = 0; +constexpr uint32_t QUERY_ROPE_INDEX = 1; +constexpr uint32_t KV_CACHE_OUT_INDEX = 2; +constexpr uint32_t KR_CACHE_OUT_INDEX = 3; +// TMP +constexpr uint32_t DIM_NUM_0 = 0; +constexpr uint32_t DIM_NUM_1 = 1; +constexpr uint32_t DIM_NUM_2 = 2; +constexpr uint32_t DIM_NUM_3 = 3; +constexpr uint32_t DIM_NUM_4 = 4; +constexpr uint32_t DIM_INDEX_0 = 0; +constexpr uint32_t DIM_INDEX_1 = 1; +constexpr uint32_t DIM_INDEX_2 = 2; +constexpr uint32_t DIM_INDEX_3 = 3; + +constexpr uint32_t FP8_E4M3_BLOCK_SIZE = 32; // Mxfp8全量化场景下 block_size = 32 + +struct MlaPrologProtoShapeParam { + bool isBsMerge { false }; + int64_t B { 0 }; + int64_t T { 0 }; + int64_t S { 0 }; + int64_t N { 0 }; + int64_t Hckv { 0 }; + int64_t He { 0 }; + int64_t Dr { 0 }; + int64_t Hcq { 0 }; +}; + +ge::graphStatus GetMlaPrologShapeDim(const gert::InferShapeContext *context, MlaPrologProtoShapeParam &shapeParam); +ge::graphStatus SetMlaPrologShapeDim(const MlaPrologProtoShapeParam &shapeParam, gert::InferShapeContext *context); +ge::graphStatus InferShapeMlaProlog(gert::InferShapeContext *context); +ge::graphStatus InferDataTypeMlaProlog(gert::InferDataTypeContext *context); +} // namespace ops + +#endif // MLA_PROLOG_INFERSHAPE_H \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.cpp b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.cpp new file mode 100644 index 000000000000..65da513d4c42 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.cpp @@ -0,0 +1,877 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_tiling.cpp + * \brief + */ + +#include +#include +#include +#include +#include "log/log.h" +#include "err/ops_err.h" +#include "register/op_def_registry.h" +#include "mla_prolog_tiling_check.h" +#include "mla_prolog_tiling.h" +using namespace ge; +namespace optiling { + +const std::unordered_map DTYPE_TO_SIZE{ + {ge::DT_BF16, 2}, {ge::DT_FLOAT16, 2}, {ge::DT_INT8, 1}, {ge::DT_FLOAT8_E4M3FN, 1}, + {ge::DT_FLOAT8_E8M0, 1}, {ge::DT_HIFLOAT8, 1}, {ge::DT_INT32, 4}, {ge::DT_FLOAT, 4}}; + +const std::unordered_map GE_TO_MM_DTYPE{ + {ge::DT_FLOAT16, matmul_tiling::DataType::DT_FLOAT16}, + {ge::DT_BF16, matmul_tiling::DataType::DT_BF16}, + {ge::DT_INT8, matmul_tiling::DataType::DT_INT8}, + {ge::DT_INT4, matmul_tiling::DataType::DT_INT4}, + {ge::DT_FLOAT, matmul_tiling::DataType::DT_FLOAT}, + {ge::DT_FLOAT8_E4M3FN, matmul_tiling::DataType::DT_FLOAT8_E4M3FN}, + {ge::DT_FLOAT8_E8M0, matmul_tiling::DataType::DT_FLOAT8_E8M0}, + {ge::DT_HIFLOAT8, matmul_tiling::DataType::DT_HIFLOAT8}}; + +template +inline auto CeilDiv(T a, T b) -> T +{ + if (b == 0) { + return b; + } + return (a + b - 1) / b; +} + +template +inline auto Align(T num, T rnd) -> T +{ + return (((rnd) == 0) ? 0 : (((num) + (rnd)-1) / (rnd) * (rnd))); +} + +NpuArch MlaPrologTiling::GetCurNpuArch() const +{ + auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_->platformInfo); + NpuArch npuArch = ascendcPlatform.GetCurNpuArch(); + return npuArch; +} + +ge::graphStatus MlaPrologTiling::GetNpuInfo() +{ + OP_CHECK_IF(context_->platformInfo == nullptr, + OPS_REPORT_VECTOR_INNER_ERR(context_->opName, "GetPlatformInfo is nullptr."), return ge::GRAPH_FAILED); + + auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_->platformInfo); + libapiSize_ = ascendcPlatform.GetLibApiWorkSpaceSize(); + + ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize_); + ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::L1, l1Size_); + ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::L0_C, l0cSize_); + ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::L0_B, l0bSize_); + + aivNum_ = ascendcPlatform.GetCoreNumAiv(); + aicNum_ = ascendcPlatform.GetCoreNumAic(); + + OP_CHECK_IF(aicNum_ == 0 || aivNum_ == 0, + OPS_REPORT_VECTOR_INNER_ERR(context_->opName, "num of core obtained is 0."), return GRAPH_FAILED); + + OP_CHECK_IF((aicNum_ != aivNum_) && (aicNum_ * 2 != aivNum_), + OPS_REPORT_VECTOR_INNER_ERR(context_->opName, "aicNum(%u):aivNum(%u) only support 1:1 or 1:2", aicNum_, + aivNum_), + return GRAPH_FAILED); + + return ge::GRAPH_SUCCESS; +} + +QUANT_MODE MlaPrologTiling::GetQuantizationModeV3() const +{ + if (*(context_->weightQuantMode) == static_cast(WEIGHT_QUANT_MODE::NO_QUANT)) { + if (*(context_->kvQuantMode) == static_cast(KV_QUANT_MODE::NO_QUANT)) { + return QUANT_MODE::NO_QUANT; + } else { + OP_LOGE(context_->opName, "When weightQuantMode == 0, kvQuantMode must be within {0}, actually is %ld.", + *(context_->kvQuantMode)); + } + } else if (*(context_->weightQuantMode) == static_cast(WEIGHT_QUANT_MODE::PARTIAL_QUANT)) { + if (*(context_->kvQuantMode) == static_cast(KV_QUANT_MODE::NO_QUANT)) { + return QUANT_MODE::PARTIAL_QUANT_KV_NO_QUANT; + } else if (*(context_->kvQuantMode) == static_cast(KV_QUANT_MODE::PER_CHANNEL)) { + return QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_CHANNEL; + } else if (*(context_->kvQuantMode) == static_cast(KV_QUANT_MODE::PER_TILE)) { + return QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_TILE; + } else { + OP_LOGE(context_->opName, + "When weightQuantMode == 1, kvQuantMode must be within {0, 2, 3}, actually is %ld.", + *(context_->kvQuantMode)); + } + } else if (*(context_->weightQuantMode) == static_cast(WEIGHT_QUANT_MODE::FULL_QUANT)) { + if (*(context_->kvQuantMode) == static_cast(KV_QUANT_MODE::NO_QUANT)) { + return QUANT_MODE::FULL_QUANT_KV_NO_QUANT; + } else if (*(context_->kvQuantMode) == static_cast(KV_QUANT_MODE::PER_TENSOR)) { + return QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TENSOR; + } else if (*(context_->kvQuantMode) == static_cast(KV_QUANT_MODE::PER_TILE)) { + return QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TILE; + } else { + OP_LOGE(context_->opName, + "When weightQuantMode == 2, kvQuantMode must be within {0, 1, 3}, actually is %ld.", + *(context_->kvQuantMode)); + } + } else { + OP_LOGE(context_->opName, "WeightQuantMode must be within {0, 1, 2}, actually is %ld.", + *(context_->weightQuantMode)); + } + return QUANT_MODE::ERROR_MODE; +} + +QUANT_MODE MlaPrologTiling::GetQuantizationModeV3Dav() const +{ + const int weightQuantMode = *(context_->weightQuantMode); + const int kvQuantMode = *(context_->kvQuantMode); + + // 卫语句1:weightQuantMode 越界(外层哈希找不到) + auto wqIt = QUANT_MODE_HASH_TABLE.find(weightQuantMode); + if (wqIt == QUANT_MODE_HASH_TABLE.end()) { + OP_LOGE_FOR_INVALID_VALUE(context_->opName, "weightQuantMode", std::to_string(weightQuantMode), + "{0, 1, 2, 3, 4, 5}"); + return QUANT_MODE::ERROR_MODE; + } + + // 卫语句2:kvQuantMode 越界或组合非法(内层哈希找不到) + auto kvIt = wqIt->second.find(kvQuantMode); + if (kvIt == wqIt->second.end()) { + auto reasonIt = VALID_KV_REASON_TABLE.find(weightQuantMode); + const char *reason = (reasonIt != VALID_KV_REASON_TABLE.end()) ? reasonIt->second : "invalid kvQuantMode"; + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context_->opName, "kvQuantMode", std::to_string(kvQuantMode), reason); + return QUANT_MODE::ERROR_MODE; + } + + // hash 命中,返回对应场景 + return kvIt->second; +} + +QUANT_MODE MlaPrologTiling::GetQuantizationMode() const +{ + if (std::strncmp(context_->opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + if (GetCurNpuArch() == NpuArch::DAV_3510) { + return GetQuantizationModeV3Dav(); + } else { + return GetQuantizationModeV3(); + } + } else { + if (context_->tokenX.desc->GetDataType() == ge::DT_INT8) { + if (context_->kvCache.desc->GetDataType() == ge::DT_INT8) { + return QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TENSOR; + } else { + return QUANT_MODE::FULL_QUANT_KV_NO_QUANT; + } + } + if (context_->tokenX.desc->GetDataType() == ge::DT_BF16 && + context_->weightUqQr.desc->GetDataType() == ge::DT_INT8) { + if (context_->kvCache.desc->GetDataType() == ge::DT_INT8) { + return QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_CHANNEL; + } else { + return QUANT_MODE::PARTIAL_QUANT_KV_NO_QUANT; + } + } + return QUANT_MODE::NO_QUANT; + } + return QUANT_MODE::ERROR_MODE; +} + +ge::graphStatus MlaPrologTiling::SetShapeInfo() +{ + if (context_->tokenX.shape->GetStorageShape().GetDimNum() == MLA_PROLOG_DIM_NUM_3) { + baseShapeInfo_.bSize = context_->tokenX.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_0); + baseShapeInfo_.s1Size = context_->tokenX.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_1); + baseShapeInfo_.heSize = context_->tokenX.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_2); + baseShapeInfo_.tSize = baseShapeInfo_.bSize * baseShapeInfo_.s1Size; + } else { + baseShapeInfo_.tSize = context_->tokenX.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_0); + baseShapeInfo_.heSize = context_->tokenX.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_1); + } + if (context_->weightDq.shape->GetStorageShape().GetDimNum() == MLA_PROLOG_DIM_NUM_2) { + baseShapeInfo_.hcqSize = context_->weightDq.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_1); + } else { + uint32_t weightDqAxisSize_ = 32U / ge::GetSizeByDataType(context_->weightDq.desc->GetDataType()); + // weightDq: [He, Hcq] -> [Hcq/16, He/16, 16, 16] || [Hcq/32, He/16, 16, 32] + baseShapeInfo_.hcqSize = + weightDqAxisSize_ * context_->weightDq.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_0); + } + baseShapeInfo_.nSize = context_->weightUk.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_0); + if (context_->ropeCos.shape != nullptr) { + baseShapeInfo_.drSize = context_->ropeCos.shape->GetStorageShape().GetDim( + context_->ropeCos.shape->GetStorageShape().GetDimNum() - 1); + } else { + // RoPE off: infer Dr from kr_cache last dim (or default 64). + OP_CHECK_IF(context_->krCache.shape == nullptr, + OP_LOGE_WITH_INVALID_INPUT(context_->opName, "krCache"), return ge::GRAPH_FAILED); + const auto &krShape = context_->krCache.shape->GetStorageShape(); + baseShapeInfo_.drSize = static_cast(krShape.GetDim(krShape.GetDimNum() - 1)); + } + baseShapeInfo_.dSize = context_->weightUk.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_1); + baseShapeInfo_.headSizeQc = baseShapeInfo_.dSize * baseShapeInfo_.nSize; + baseShapeInfo_.headSizeQr = baseShapeInfo_.drSize * baseShapeInfo_.nSize; + baseShapeInfo_.headSizeUqQr = baseShapeInfo_.headSizeQc + baseShapeInfo_.headSizeQr; + if (context_->kvCache.shape->GetStorageShape().GetDimNum() == MLA_PROLOG_DIM_NUM_3) { + baseShapeInfo_.blockNum = context_->kvCache.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_0); + baseShapeInfo_.nkvSize = context_->kvCache.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_1); + baseShapeInfo_.dtileSize = context_->kvCache.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_2); + } else { + baseShapeInfo_.blockNum = context_->kvCache.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_0); + baseShapeInfo_.blockSize = context_->kvCache.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_1); + baseShapeInfo_.nkvSize = context_->kvCache.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_2); + baseShapeInfo_.dtileSize = context_->kvCache.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_3); + } + if (context_->weightDkvKr.shape->GetStorageShape().GetDimNum() == MLA_PROLOG_DIM_NUM_4) { + baseShapeInfo_.hckvSize = context_->weightDkvKr.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_0) * + context_->weightDkvKr.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_3) - + baseShapeInfo_.drSize; + } else { + baseShapeInfo_.hckvSize = + context_->weightDkvKr.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_1) - baseShapeInfo_.drSize; + } + baseShapeInfo_.s2Size = baseShapeInfo_.nkvSize; + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTiling::SetScenarioInfo() +{ + scenarioInfo_.isV1Flag_ = (std::strncmp(context_->opType, V1_OP_NAME, OP_NAME_LEN) == 0); + scenarioInfo_.batchSeqFusedFlag_ = context_->tokenX.shape->GetStorageShape().GetDimNum() == MLA_PROLOG_DIM_NUM_2; + scenarioInfo_.quantMode_ = GetQuantizationMode(); + if (scenarioInfo_.quantMode_ == QUANT_MODE::ERROR_MODE) { + return ge::GRAPH_FAILED; + } + + // 由 quantMode_ 反向映射 wq/kvq(V1/V2 按 dtype 推断、V3 正向查表,均唯一可反推) + const auto &wqKvq = QUANT_MODE_REVERSE_TABLE.at(scenarioInfo_.quantMode_); + scenarioInfo_.weightQuantMode_ = wqKvq.first; + scenarioInfo_.kvQuantMode_ = wqKvq.second; + + // cacheMode 字符串 → 枚举,hash 查表(CheckCacheMode 已保证 cacheMode 合法,必命中) + scenarioInfo_.cacheMode_ = CACHE_MODE_HASH_TABLE.at(std::string(context_->cacheMode)); + + if ((scenarioInfo_.batchSeqFusedFlag_ && baseShapeInfo_.tSize == 0U) || + (!scenarioInfo_.batchSeqFusedFlag_ && (baseShapeInfo_.bSize * baseShapeInfo_.s1Size == 0U))) { + scenarioInfo_.emptyTensorMode_ = EMPTY_TENSOR_MODE::EMPTY_QUERY; + } else if (baseShapeInfo_.blockNum == 0U) { + scenarioInfo_.emptyTensorMode_ = EMPTY_TENSOR_MODE::EMPTY_CACHE; + } else { + scenarioInfo_.emptyTensorMode_ = EMPTY_TENSOR_MODE::NON_EMPTY; + } + + if (scenarioInfo_.batchSeqFusedFlag_ && + (scenarioInfo_.cacheMode_ == CACHE_MODE::PA_BLK_BSND || scenarioInfo_.cacheMode_ == CACHE_MODE::PA_BLK_NZ)) { + OP_CHECK_IF(context_->actualSeqLen.shape == nullptr, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context_->opName, "actualSeqLen", "null", + "When cacheMode in {PA_BLK_BSND, PA_BLK_NZ} and tokenX " + "shape dim num is 2, actualSeqLen should not be null"), + return GRAPH_FAILED); + baseShapeInfo_.bSize = context_->actualSeqLen.shape->GetStorageShape().GetDim(MLA_PROLOG_DIM_INDEX_0); + scenarioInfo_.actualSeqMode_ = ACTUAL_SEQ_MODE::EN_Q_LEN; + } else { + scenarioInfo_.actualSeqMode_ = ACTUAL_SEQ_MODE::DISABLED; + } + uint32_t cvRatio = aivNum_ / aicNum_; + // 当前仅在BS>=8K且数据类型为MXFP8时路由到切M模板,其他情况均路由到切N模板 + scenarioInfo_.splitMFlag_ = 0U; + if (scenarioInfo_.weightQuantMode_ == WEIGHT_QUANT_MODE::MXFP8_FULL_QUANT && + baseShapeInfo_.tSize >= 8192 && cvRatio == 2) { // 8192:BS >= 8K + if ((baseShapeInfo_.heSize == HEAD_SIZE1 || baseShapeInfo_.heSize == HEAD_SIZE2) && + baseShapeInfo_.nSize == 128) { // 128:N为128时路由到切M模板 + scenarioInfo_.splitMFlag_ = 1U; + } + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTiling::SetAttrInfo() +{ + reciprocalCq_ = 1.0f / baseShapeInfo_.hcqSize; + epsilonCq_ = *(context_->rmsNormEspilonCq); + reciprocalCkv_ = 1.0f / baseShapeInfo_.hckvSize; + epsilonCkv_ = *(context_->rmsNormEspilonCkv); + if (std::strncmp(context_->opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + queryNormFlag_ = *(context_->queryNormFlag); + weightQuantMode_ = static_cast(*(context_->weightQuantMode)); + kvQuantMode_ = static_cast(*(context_->kvQuantMode)); + queryQuantMode_ = static_cast(*(context_->queryQuantMode)); + ckvkrRepoMode_ = static_cast(*(context_->ckvkrRepoMode)); + quantSacleRepoMode_ = static_cast(*(context_->quantScaleRepoMode)); + tileSize_ = static_cast(*(context_->tileSize)); + qcQrScale_ = *(context_->qcQrScale); + kcScale_ = *(context_->kcScale); + } + + enableRope_ = true; + if (GetCurNpuArch() == NpuArch::DAV_3510 && context_->doRope != nullptr) { + enableRope_ = *(context_->doRope); + } + return ge::GRAPH_SUCCESS; +} + +bool MlaPrologTiling::GetMatmulType(ge::DataType getype, matmul_tiling::DataType *mmType) +{ + auto mmdt = GE_TO_MM_DTYPE.find(getype); + if (mmdt != GE_TO_MM_DTYPE.end()) { + *mmType = mmdt->second; + return true; + } + return false; +} + +uint32_t MlaPrologTiling::CalcSingleCoreN(uint32_t n, uint32_t coreNum, uint32_t alignNum) const +{ + return CeilDiv(n, alignNum * coreNum) * alignNum; +} + +// mm1.m = stepBatchSize // 32 +// mm1.n = singlecoreHeadSizeCq // 64 +// mm1.k = headSizeX // 7168 +// mm1.baseM = stepBatchSize // 32 +// mm1.baseN = singlecoreHeadSizeCq // 64 +// mm1.baseK = 256 +ge::graphStatus MlaPrologTiling::FillMatmul1Tiling() +{ + if (scenarioInfo_.splitMFlag_ == 1U) { + singlecoreHeadSizeCq_ = baseShapeInfo_.hcqSize; + mm1BlockNum_ = aicNum_; + } else { + auto dataType = context_->weightDq.desc->GetDataType(); + singlecoreHeadSizeCq_ = + CalcSingleCoreN(baseShapeInfo_.hcqSize, aicNum_, BLOCK_SIZE / DTYPE_TO_SIZE.at(dataType)); + singlecoreHeadSizeCq_ = std::max(singlecoreHeadSizeCq_, 64U); // 64:最大使用24核 + mm1BlockNum_ = CeilDiv(baseShapeInfo_.hcqSize, singlecoreHeadSizeCq_); + } + return ge::GRAPH_SUCCESS; +} + +// singlecoreHeadSizeCkvKr = HeadSizeCkvDr / mm2CoreNum // 576 / 9 == 64 +// mm2.m = stepBatchSize +// mm2.n = singlecoreHeadSizeCkvKr +// mm2.k = headSizeX // size of He +// mm2.baseN = n +// mm2.baseK = 256 +ge::graphStatus MlaPrologTiling::FillMatmul2Tiling() +{ + if (scenarioInfo_.emptyTensorMode_ == EMPTY_TENSOR_MODE::EMPTY_CACHE) { + return ge::GRAPH_SUCCESS; + } + if (scenarioInfo_.splitMFlag_ == 1U) { + singlecoreHeadSizeCkvKr_ = baseShapeInfo_.hckvSize + baseShapeInfo_.drSize; + mm2BlockNum_ = aicNum_; + } else if (aicNum_ >= 9U) { // 9是经验值 + uint32_t baseN = 64U; + mm2BlockNum_ = (baseShapeInfo_.hckvSize + baseShapeInfo_.drSize) / baseN; + singlecoreHeadSizeCkvKr_ = baseN; + } else { + auto dataType = context_->weightDkvKr.desc->GetDataType(); + singlecoreHeadSizeCkvKr_ = CalcSingleCoreN(baseShapeInfo_.hckvSize + baseShapeInfo_.drSize, aicNum_, + BLOCK_SIZE / DTYPE_TO_SIZE.at(dataType)); + mm2BlockNum_ = CeilDiv(baseShapeInfo_.hckvSize + baseShapeInfo_.drSize, singlecoreHeadSizeCkvKr_); + } + return ge::GRAPH_SUCCESS; +} + +// singlecoreHeadSizeQcQr = headNum * (dimHeadSizeQc + dimHeadRope) / mm3CoreNum = 32 * (128 + 64) / 24 +// mm3.m = stepBatchSize +// mm3.n = singlecoreHeadSizeQcQr // 256 +// mm3.k = headSizeCq // size of Hcq 1536 +// mm3.baseN = 64 // +// mm3.baseK = 256 // +ge::graphStatus MlaPrologTiling::FillMatmul3Tiling() +{ + auto dataType = context_->weightUqQr.desc->GetDataType(); + auto oriM = baseShapeInfo_.nSize * (baseShapeInfo_.dSize + baseShapeInfo_.drSize); + if (enableGroupComputeOpt_) { + // 算力分组场景下G=8,dimHeadSizeQc跨8核切,dimHeadSizeQr跨4核切;matmulQc和matmulQr的singleN都取128 + singlecoreHeadSizeQcQr_ = CalcSingleCoreN(baseShapeInfo_.nSize * baseShapeInfo_.dSize, + GROUP_COMPUTE_CUBE_NUM_PER_GROUP, baseShapeInfo_.dSize); + } else if (enableDequantOpt_) { + // dequant流水掩盖场景,dimHeadSizeQc + dimHeadRope不跨核 + singlecoreHeadSizeQcQr_ = CalcSingleCoreN(oriM, aicNum_, baseShapeInfo_.dSize + baseShapeInfo_.drSize); + } else { + // headnum * (dimHeadSizeQc + dimHeadRope) 合轴切 + singlecoreHeadSizeQcQr_ = CalcSingleCoreN(oriM, aicNum_, BLOCK_SIZE / DTYPE_TO_SIZE.at(dataType)); + } + mm3BlockNum_ = CeilDiv(oriM, singlecoreHeadSizeQcQr_); + + if (scenarioInfo_.splitMFlag_ == 1U) { + singlecoreHeadSizeQcQr_ = oriM; + mm3BlockNum_ = aicNum_; + } + + return ge::GRAPH_SUCCESS; +} + +// mm4.m = stepBatchSize +// mm4.n = headSizeCkv // 512 +// mm4.k = dimHeadSizeQc // size of Qc 128 +// mm4.baseN = 128 // +// mm4.baseK = 128 // +// mm4.Kstride = dimHeadSizeQc + dimHeadRope +ge::graphStatus MlaPrologTiling::FillMatmul4Tiling() +{ + if (scenarioInfo_.splitMFlag_ == 1U) { + singlecoreNumHeadSize_ = baseShapeInfo_.nSize; + mm4BlockNum_ = aicNum_; + } else { + singlecoreNumHeadSize_ = CeilDiv(baseShapeInfo_.nSize, aicNum_); + mm4BlockNum_ = CeilDiv(baseShapeInfo_.nSize, singlecoreNumHeadSize_); + } + + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTiling::ProcessBaseInputs() +{ + stepBatchSize_ = std::min(128U, baseShapeInfo_.tSize); + if (scenarioInfo_.splitMFlag_ == 1U && (stepBatchSize_ > 0U) && (aicNum_ > 0U) && (baseShapeInfo_.tSize > 0U)) { + mSubSize_ = (baseShapeInfo_.tSize + aicNum_ - 1U) / aicNum_; + // idx为[0, mSubCoreNum_]的核分到mSubSize_,其余核分到mSubSize_ - 1 + mSubCoreNum_ = baseShapeInfo_.tSize - (mSubSize_ - 1U) * aicNum_; + } + if (baseShapeInfo_.dSize == HIGH_THROUGHPUT__D_SIZE) { + stepNumHeadDequant_ = std::min(64U, baseShapeInfo_.nSize); + } else { + stepNumHeadDequant_ = std::min(16U, baseShapeInfo_.nSize); + } + vectorBlockNum_ = std::min(stepBatchSize_, aivNum_); + + uint32_t cvRatio = aivNum_ / aicNum_; + // 算力分组开关,仅当半量化场景,BS=1,G=8,可用核数大于等于16时进入分支 + // CV1:1时不支持分组计算场景 + if ((scenarioInfo_.quantMode_ == QUANT_MODE::PARTIAL_QUANT_KV_NO_QUANT || + scenarioInfo_.quantMode_ == QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_CHANNEL) && + baseShapeInfo_.tSize == GROUP_COMPUTE_T_SIZE && baseShapeInfo_.nkvSize == GROUP_COMPUTE_N_SIZE && + aivNum_ >= GROUP_COMPUTE_MIN_AIV_NUM && aicNum_ >= GROUP_COMPUTE_MIN_AIC_NUM && cvRatio != 1) { + enableGroupComputeOpt_ = true; + aivNum_ = 32U; + aicNum_ = 16U; + } else if ((context_->weightUqQr.desc->GetDataType() == ge::DT_INT8 && + baseShapeInfo_.nSize >= GROUP_COMPUTE_N_SIZE) || + context_->weightUqQr.desc->GetDataType() == ge::DT_FLOAT8_E4M3FN || + context_->weightUqQr.desc->GetDataType() == ge::DT_HIFLOAT8) { + // 场景1:INT8全量化且N大于等于8;场景2:MXFP8全量化场景 + // 通过切N处理MM3,MM4之后的操作例如Rope,DynamicQuant等会有性能收益 + enableDequantOpt_ = true; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTiling::FillTiling() +{ + baseParams_->tokenSize = baseShapeInfo_.tSize; + baseParams_->seq1Size = baseShapeInfo_.s1Size; + baseParams_->seq2Size = baseShapeInfo_.s2Size; + baseParams_->headSizeX = baseShapeInfo_.heSize; + baseParams_->headSizeCq = baseShapeInfo_.hcqSize; + baseParams_->headSizeCkv = baseShapeInfo_.hckvSize; + baseParams_->dtileSize = baseShapeInfo_.dtileSize; + baseParams_->headSizeQc = baseShapeInfo_.headSizeQc; + baseParams_->headSizeQr = baseShapeInfo_.headSizeQr; + baseParams_->headSizeKr = baseShapeInfo_.drSize; + baseParams_->numHeadSize = baseShapeInfo_.nSize; + baseParams_->numHeadKvSize = baseShapeInfo_.nkvSize; + baseParams_->dimHeadSizeQc = baseShapeInfo_.dSize; + baseParams_->dimHeadRope = baseShapeInfo_.drSize; + baseParams_->blockNum = baseShapeInfo_.blockNum; + baseParams_->blockSize = baseShapeInfo_.blockSize; + baseParams_->reciprocalCq = reciprocalCq_; + baseParams_->epsilonCq = epsilonCq_; + baseParams_->reciprocalCkv = reciprocalCkv_; + baseParams_->epsilonCkv = epsilonCkv_; + if (std::strncmp(context_->opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + baseParams_->queryNormFlag = queryNormFlag_ ? 1U : 0U; + baseParams_->kvQuantMode = static_cast(kvQuantMode_); + baseParams_->ckvkrRepoMode = static_cast(ckvkrRepoMode_); + baseParams_->quantScaleRepoMode = static_cast(quantSacleRepoMode_); + baseParams_->tileSize = tileSize_; + baseParams_->qcQrScale = qcQrScale_; + baseParams_->kcScale = kcScale_; + baseParams_->isQcQrScaleEnable = + static_cast(std::abs(qcQrScale_ - 1.0f) >= std::numeric_limits::epsilon()); + baseParams_->isKcScaleEnable = + static_cast(std::abs(kcScale_ - 1.0f) >= std::numeric_limits::epsilon()); + } else { + baseParams_->queryNormFlag = 0U; + baseParams_->kvQuantMode = 0U; + baseParams_->ckvkrRepoMode = 0U; + baseParams_->quantScaleRepoMode = 0U; + baseParams_->tileSize = 128U; + baseParams_->qcQrScale = 0; + baseParams_->kcScale = 0; + baseParams_->isQcQrScaleEnable = 0U; + baseParams_->isKcScaleEnable = 0U; + } + FillTilingCoreParams(); // 分核相关baseParams + return ge::GRAPH_SUCCESS; +} + +void MlaPrologTiling::FillTilingCoreParams() +{ + baseParams_->batchSize = baseShapeInfo_.bSize; + baseParams_->stepBatchSize = stepBatchSize_; + baseParams_->stepNumHeadDequant = stepNumHeadDequant_; + baseParams_->mSubSize = mSubSize_; + baseParams_->mSubCoreNum = mSubCoreNum_; + baseParams_->mm1BlockNum = mm1BlockNum_; + baseParams_->mm2BlockNum = mm2BlockNum_; + baseParams_->mm3BlockNum = mm3BlockNum_; + baseParams_->mm4BlockNum = mm4BlockNum_; + baseParams_->mm1SingleCoreN = singlecoreHeadSizeCq_; + baseParams_->mm2SingleCoreN = singlecoreHeadSizeCkvKr_; + baseParams_->mm3SingleCoreN = singlecoreHeadSizeQcQr_; + baseParams_->mm4SingleCoreBatch = singlecoreNumHeadSize_; + baseParams_->vectorBlockNum = vectorBlockNum_; +} + +ge::graphStatus MlaPrologTiling::CalcWorkSpace() +{ + workspaceSize_ = libapiSize_; + uint32_t mm1Mult = (scenarioInfo_.splitMFlag_ == 1U) ? mm1BlockNum_ : 1U; + uint32_t mm2Mult = (scenarioInfo_.splitMFlag_ == 1U) ? mm2BlockNum_ : 1U; + uint32_t mm3Mult = (scenarioInfo_.splitMFlag_ == 1U) ? mm3BlockNum_ : 1U; + uint32_t mm4Mult = (scenarioInfo_.splitMFlag_ == 1U) ? mm4BlockNum_ : 1U; + uint32_t dequantScaleMult = (scenarioInfo_.splitMFlag_ == 1U) ? static_cast(aicNum_) : 1U; + // 全量化场景:weightQuantMode 为 FULL/MXFP8/FP8/HIF8(即非 NO_QUANT 且非 PARTIAL) + if (scenarioInfo_.weightQuantMode_ != WEIGHT_QUANT_MODE::NO_QUANT && + scenarioInfo_.weightQuantMode_ != WEIGHT_QUANT_MODE::PARTIAL_QUANT) { + workspaceSize_ += static_cast(stepBatchSize_) * static_cast(baseShapeInfo_.hcqSize) * mm1Mult * + static_cast(NUM_BYTES_INT32); + workspaceSize_ += static_cast(stepBatchSize_) * static_cast(baseShapeInfo_.hcqSize) * mm1Mult * + static_cast(NUM_BYTES_BF16); + workspaceSize_ += static_cast(stepBatchSize_) * + static_cast(baseShapeInfo_.hckvSize + baseShapeInfo_.drSize) * mm2Mult * + static_cast(NUM_BYTES_INT32); + if (scenarioInfo_.kvQuantMode_ == KV_QUANT_MODE::PER_TENSOR) { + // 全量化场景mmQnRes输出到workspace, B, S1, N, Hckv, BF16 + workspaceSize_ += static_cast(stepBatchSize_) * static_cast(baseShapeInfo_.nSize) * + static_cast(baseShapeInfo_.hckvSize) * mm4Mult * + static_cast(NUM_BYTES_BF16); + } + } else { + workspaceSize_ += static_cast(stepBatchSize_) * static_cast(baseShapeInfo_.hcqSize) * mm1Mult * + static_cast(NUM_BYTES_BF16) * static_cast(2); // 2: double + workspaceSize_ += static_cast(stepBatchSize_) * + static_cast(baseShapeInfo_.hckvSize + baseShapeInfo_.drSize) * mm2Mult * + static_cast(NUM_BYTES_BF16); + } + workspaceSize_ += static_cast(stepBatchSize_) * + static_cast(baseShapeInfo_.headSizeQc + baseShapeInfo_.headSizeQr) * mm3Mult * + static_cast(NUM_BYTES_INT32); + workspaceSize_ += static_cast(stepBatchSize_) * static_cast(baseShapeInfo_.nSize) * + static_cast(baseShapeInfo_.dSize) * mm3Mult * static_cast(NUM_BYTES_BF16); + // SplitM非pertile量化场景下 mmQcQrResDequant使用两份workspace + if (scenarioInfo_.splitMFlag_ == 1U && + (scenarioInfo_.quantMode_ == QUANT_MODE::MXFP8_FULL_QUANT_KV_NO_QUANT || + scenarioInfo_.quantMode_ == QUANT_MODE::MXFP8_FULL_QUANT_KV_QUANT_PER_TENSOR)) { + workspaceSize_ += static_cast(stepBatchSize_) * static_cast(baseShapeInfo_.nSize) * + static_cast(baseShapeInfo_.dSize) * mm3Mult * static_cast(NUM_BYTES_BF16); + } + + if (enableGroupComputeOpt_ || enableDequantOpt_) { + workspaceSize_ += static_cast(stepBatchSize_) * dequantScaleMult * static_cast(BLOCK_SIZE); + } + if (context_->workSpaces) { + context_->workSpaces[0] = workspaceSize_; + } + OP_LOGI(context_->opName, "Tiling info: workspaceSize_ = %zu", workspaceSize_); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTiling::GenTilingKey() const +{ + uint8_t typeValue = 0; + uint8_t quantType = 0; + if (scenarioInfo_.quantMode_ == QUANT_MODE::NO_QUANT) { + typeValue = 1U; + } else { + typeValue = 2U; + // kvCache量化场景,对应tiling key为1(半量化:0 + kv量化:1)或3(全量化:2 + kv量化:1) + // 全量化场景,对应tiling key为2+0(全量化:2)或2+1(全量化:2+ kv量化:1) + // 非量化和半量化场景,对应tiling key为0 + quantType = static_cast(scenarioInfo_.quantMode_); + } + + uint8_t cvMode = ASCENDC_TPL_MIX_AIC_1_2; // 默认cv 1:2模式 + if (aivNum_ == aicNum_) { + cvMode = ASCENDC_TPL_MIX_AIC_1_1; // cv 1:1模式 + } + + if (cvMode == ASCENDC_TPL_MIX_AIC_1_1 && + (scenarioInfo_.quantMode_ != QUANT_MODE::NO_QUANT || + (scenarioInfo_.cacheMode_ != CACHE_MODE::PA_BSND && scenarioInfo_.cacheMode_ != CACHE_MODE::PA_NZ))) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON( + context_->opName, "quantMode", + std::to_string(static_cast(scenarioInfo_.quantMode_)) + ", cacheMode is " + + std::to_string(static_cast(scenarioInfo_.cacheMode_)), + "CV1:1 mode only support quantMode in {NO_QUANT} and cacheMode in {PA_BSND,PA_NZ}"); + return ge::GRAPH_FAILED; + } + + if (scenarioInfo_.emptyTensorMode_ == EMPTY_TENSOR_MODE::EMPTY_QUERY) { + context_->tilingKey = GET_TPL_TILING_KEY(0, 0, 0, false, false, + static_cast(scenarioInfo_.emptyTensorMode_), 0, 0, cvMode, + true); + } else { + uint8_t cacheMode = + scenarioInfo_.cacheMode_ == CACHE_MODE::TND ? 0 : static_cast(scenarioInfo_.cacheMode_); + context_->tilingKey = GET_TPL_TILING_KEY( + static_cast(cacheMode), typeValue, quantType, enableDequantOpt_, enableGroupComputeOpt_, + static_cast(scenarioInfo_.emptyTensorMode_), static_cast(scenarioInfo_.actualSeqMode_), + static_cast(scenarioInfo_.splitMFlag_), cvMode, enableRope_); + OP_LOGI( + context_->opName, + "MlaProlog tilingKey args: " + "CACHE_MODE:%u, SCENARIO:%u, QUANT_MODE:%u, ENABLE_DEQUANT_OPTIONAL:%u, ENABLE_GROUP_COMPUTE_OPTIONAL:%u, " + "EMPTY_TENSOR_MODE:%u, ACTUAL_SEQ_LEN_MODE:%u, SPLIT_M_MODE:%u, CV_MODE:%u, ENABLE_ROPE:%u", + static_cast(cacheMode), typeValue, quantType, enableDequantOpt_, enableGroupComputeOpt_, + static_cast(scenarioInfo_.emptyTensorMode_), static_cast(scenarioInfo_.actualSeqMode_), + static_cast(scenarioInfo_.splitMFlag_), cvMode, static_cast(enableRope_)); + } + OP_LOGI(context_->opName, "MlaProlog tilingKey:%lu", context_->tilingKey); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTiling::RunBigKernelTiling(MlaPrologContext &context, MlaPrologTilingData *tilingData) +{ + this->context_ = &context; + this->baseParams_ = &tilingData->baseParams; + MlaPrologTilingCheck tilingCheck_{*context_, baseShapeInfo_, scenarioInfo_}; + + using StatusFunction = std::function; + std::vector requiredTilingFuncs{ + std::bind(&MlaPrologTiling::GetNpuInfo, this), + std::bind(&MlaPrologTilingCheck::CheckAttrs, &tilingCheck_), + std::bind(&MlaPrologTilingCheck::CheckSingleRequiredParam, &tilingCheck_), + std::bind(&MlaPrologTilingCheck::CheckCacheMode, &tilingCheck_), + std::bind(&MlaPrologTiling::SetShapeInfo, this), + std::bind(&MlaPrologTiling::SetScenarioInfo, this), + std::bind(&MlaPrologTilingCheck::CheckScenarParam, &tilingCheck_), + std::bind(&MlaPrologTilingCheck::CheckQuantMode, &tilingCheck_), + std::bind(&MlaPrologTilingCheck::CheckDims, &tilingCheck_), + std::bind(&MlaPrologTilingCheck::CheckSpecialScenarioParamShape, &tilingCheck_), + std::bind(&MlaPrologTilingCheck::CheckParamByScenario, &tilingCheck_), + std::bind(&MlaPrologTiling::SetAttrInfo, this), + std::bind(&MlaPrologTiling::ProcessBaseInputs, this), + }; + for (const auto &func : requiredTilingFuncs) { + if (func() != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + } + + if (scenarioInfo_.emptyTensorMode_ == EMPTY_TENSOR_MODE::EMPTY_QUERY) { + FillTiling(); + if (context_->workSpaces) { + context_->workSpaces[0] = libapiSize_; + } + GenTilingKey(); + context_->blockDim = 1U; + return ge::GRAPH_SUCCESS; + } + + std::vector optionalTilingFuncs{ + std::bind(&MlaPrologTiling::FillMatmul1Tiling, this), std::bind(&MlaPrologTiling::FillMatmul2Tiling, this), + std::bind(&MlaPrologTiling::FillMatmul3Tiling, this), std::bind(&MlaPrologTiling::FillMatmul4Tiling, this), + std::bind(&MlaPrologTiling::FillTiling, this), std::bind(&MlaPrologTiling::CalcWorkSpace, this), + std::bind(&MlaPrologTiling::GenTilingKey, this)}; + for (const auto &func : optionalTilingFuncs) { + if (func() != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + } + + context_->blockDim = aicNum_; + + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTiling::ConvertContext(gert::TilingContext &context, MlaPrologContext &mlaPrologContext) +{ + OP_CHECK_IF(context.GetNodeName() == nullptr, OP_LOGE_WITH_INVALID_INPUT(V1_OP_NAME, "OpName"), + return ge::GRAPH_FAILED); + + mlaPrologContext.opName = context.GetNodeName(); + mlaPrologContext.opType = context.GetNodeType(); + mlaPrologContext.platformInfo = context.GetPlatformInfo(); + + ConvertRequiredParams(context, mlaPrologContext); + ConvertOptionalParams(context, mlaPrologContext); + + auto attrs = context.GetAttrs(); + OP_CHECK_IF(attrs == nullptr, OP_LOGE_WITH_INVALID_INPUT(context.GetNodeName(), "attrs"), return ge::GRAPH_FAILED); + mlaPrologContext.rmsNormEspilonCq = attrs->GetAttrPointer(RMS_NORM_EPSILON_CQ_ATTR_INDEX); + mlaPrologContext.rmsNormEspilonCkv = attrs->GetAttrPointer(RMS_NORM_EPSILON_CKV_ATTR_INDEX); + mlaPrologContext.cacheMode = attrs->GetStr(CACHE_MODE_ATTR_INDEX); + if (std::strncmp(mlaPrologContext.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + mlaPrologContext.queryNormFlag = attrs->GetAttrPointer(QUERY_NORM_FLAG_ATTR_INDEX); + mlaPrologContext.weightQuantMode = attrs->GetAttrPointer(WEIGHT_QUANT_MODE_ATTR_INDEX); + mlaPrologContext.kvQuantMode = attrs->GetAttrPointer(KV_CACHE_QUANT_MODE_ATTR_INDEX); + mlaPrologContext.queryQuantMode = attrs->GetAttrPointer(QUERY_QUANT_MODE_ATTR_INDEX); + mlaPrologContext.ckvkrRepoMode = attrs->GetAttrPointer(CKVKR_REPO_MODE_ATTR_INDEX); + mlaPrologContext.quantScaleRepoMode = attrs->GetAttrPointer(QUANT_SCALE_REPO_MODE_ATTR_INDEX); + mlaPrologContext.tileSize = attrs->GetAttrPointer(TILE_SIZE_ATTR_INDEX); + mlaPrologContext.qcQrScale = attrs->GetAttrPointer(QC_QR_SCALE_ATTR_INDEX); + mlaPrologContext.kcScale = attrs->GetAttrPointer(KC_SCALE_ATTR_INDEX); + // Infer RoPE switch from optional rope inputs (both null → off, both non-null → on). + const bool ropeSinNull = (mlaPrologContext.ropeSin.shape == nullptr); + const bool ropeCosNull = (mlaPrologContext.ropeCos.shape == nullptr); + OP_CHECK_IF(ropeSinNull != ropeCosNull, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON( + mlaPrologContext.opName, "ropeSin / ropeCos", "(one null, one non-null)", + "ropeSin and ropeCos must both be null (RoPE off) or both non-null (RoPE on)"), + return ge::GRAPH_FAILED); + mlaPrologContext.doRopeValue = !ropeSinNull; + mlaPrologContext.doRope = &mlaPrologContext.doRopeValue; + } else { + mlaPrologContext.queryNormFlag = nullptr; + mlaPrologContext.weightQuantMode = nullptr; + mlaPrologContext.kvQuantMode = nullptr; + mlaPrologContext.queryQuantMode = nullptr; + mlaPrologContext.ckvkrRepoMode = nullptr; + mlaPrologContext.quantScaleRepoMode = nullptr; + mlaPrologContext.tileSize = nullptr; + mlaPrologContext.qcQrScale = nullptr; + mlaPrologContext.kcScale = nullptr; + // 非 V3:RoPE 始终开启(rope 为必选输入) + mlaPrologContext.doRopeValue = true; + mlaPrologContext.doRope = nullptr; + } + + OP_CHECK_IF(context.GetWorkspaceSizes(1) == nullptr, + OPS_REPORT_VECTOR_INNER_ERR(context.GetNodeName(), "workSpaceSize got from ge is nullptr"), + return ge::GRAPH_FAILED); + mlaPrologContext.workSpaces = context.GetWorkspaceSizes(1); + return ge::GRAPH_SUCCESS; +} + +void MlaPrologTiling::ConvertRequiredParams(gert::TilingContext &context, MlaPrologContext &mlaPrologContext) +{ + mlaPrologContext.tokenX.desc = context.GetRequiredInputDesc(TOKEN_X_INPUT_INDEX); + mlaPrologContext.tokenX.shape = context.GetRequiredInputShape(TOKEN_X_INPUT_INDEX); + mlaPrologContext.weightDq.desc = context.GetRequiredInputDesc(WEIGHT_DQ_INPUT_INDEX); + mlaPrologContext.weightDq.shape = context.GetRequiredInputShape(WEIGHT_DQ_INPUT_INDEX); + mlaPrologContext.weightUqQr.desc = context.GetRequiredInputDesc(WEIGHT_UQ_QR_INPUT_INDEX); + mlaPrologContext.weightUqQr.shape = context.GetRequiredInputShape(WEIGHT_UQ_QR_INPUT_INDEX); + mlaPrologContext.weightUk.desc = context.GetRequiredInputDesc(WEIGHT_UK_INPUT_INDEX); + mlaPrologContext.weightUk.shape = context.GetRequiredInputShape(WEIGHT_UK_INPUT_INDEX); + mlaPrologContext.weightDkvKr.desc = context.GetRequiredInputDesc(WEIGHT_DKV_KR_INPUT_INDEX); + mlaPrologContext.weightDkvKr.shape = context.GetRequiredInputShape(WEIGHT_DKV_KR_INPUT_INDEX); + mlaPrologContext.rmsnormGammaCq.desc = context.GetRequiredInputDesc(RMSNORM_GAMMA_CQ_INPUT_INDEX); + mlaPrologContext.rmsnormGammaCq.shape = context.GetRequiredInputShape(RMSNORM_GAMMA_CQ_INPUT_INDEX); + mlaPrologContext.rmsnormGammaCkv.desc = context.GetRequiredInputDesc(RMS_NORM_GAMMA_CKV_INPUT_INDEX); + mlaPrologContext.rmsnormGammaCkv.shape = context.GetRequiredInputShape(RMS_NORM_GAMMA_CKV_INPUT_INDEX); + // V3: rope is optional (null when do_rope=false). Non-V3 keeps required. + if (std::strncmp(mlaPrologContext.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + mlaPrologContext.ropeSin.desc = context.GetOptionalInputDesc(ROPE_SIN_INPUT_INDEX); + mlaPrologContext.ropeSin.shape = context.GetOptionalInputShape(ROPE_SIN_INPUT_INDEX); + mlaPrologContext.ropeCos.desc = context.GetOptionalInputDesc(ROPE_COS_INPUT_INDEX); + mlaPrologContext.ropeCos.shape = context.GetOptionalInputShape(ROPE_COS_INPUT_INDEX); + } else { + mlaPrologContext.ropeSin.desc = context.GetRequiredInputDesc(ROPE_SIN_INPUT_INDEX); + mlaPrologContext.ropeSin.shape = context.GetRequiredInputShape(ROPE_SIN_INPUT_INDEX); + mlaPrologContext.ropeCos.desc = context.GetRequiredInputDesc(ROPE_COS_INPUT_INDEX); + mlaPrologContext.ropeCos.shape = context.GetRequiredInputShape(ROPE_COS_INPUT_INDEX); + } + + if (std::strncmp(mlaPrologContext.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + mlaPrologContext.kvCache.desc = context.GetRequiredInputDesc(KV_CACHE_INPUT_INDEX_V3); + mlaPrologContext.kvCache.shape = context.GetRequiredInputShape(KV_CACHE_INPUT_INDEX_V3); + mlaPrologContext.krCache.desc = context.GetRequiredInputDesc(KR_CACHE_INPUT_INDEX_V3); + mlaPrologContext.krCache.shape = context.GetRequiredInputShape(KR_CACHE_INPUT_INDEX_V3); + } else { + mlaPrologContext.cacheIndex.desc = context.GetRequiredInputDesc(CACHE_INDEX_INPUT_INDEX); + mlaPrologContext.cacheIndex.shape = context.GetRequiredInputShape(CACHE_INDEX_INPUT_INDEX); + mlaPrologContext.kvCache.desc = context.GetRequiredInputDesc(KV_CACHE_INPUT_INDEX); + mlaPrologContext.kvCache.shape = context.GetRequiredInputShape(KV_CACHE_INPUT_INDEX); + mlaPrologContext.krCache.desc = context.GetRequiredInputDesc(KR_CACHE_INPUT_INDEX); + mlaPrologContext.krCache.shape = context.GetRequiredInputShape(KR_CACHE_INPUT_INDEX); + } + + mlaPrologContext.query.desc = context.GetOutputDesc(QUERY_OUTPUT_INDEX); + mlaPrologContext.query.shape = context.GetOutputShape(QUERY_OUTPUT_INDEX); + mlaPrologContext.queryRope.desc = context.GetOutputDesc(QUERY_ROPE_OUTPUT_INDEX); + mlaPrologContext.queryRope.shape = context.GetOutputShape(QUERY_ROPE_OUTPUT_INDEX); + mlaPrologContext.kvCacheOut.desc = context.GetOutputDesc(KV_CACHE_OUT_OUTPUT_INDEX); + mlaPrologContext.kvCacheOut.shape = context.GetOutputShape(KV_CACHE_OUT_OUTPUT_INDEX); + mlaPrologContext.krCacheOut.desc = context.GetOutputDesc(KR_CACHE_OUT_OUTPUT_INDEX); + mlaPrologContext.krCacheOut.shape = context.GetOutputShape(KR_CACHE_OUT_OUTPUT_INDEX); +} + +void MlaPrologTiling::ConvertOptionalParams(gert::TilingContext &context, MlaPrologContext &mlaPrologContext) +{ + if (std::strncmp(mlaPrologContext.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + mlaPrologContext.cacheIndex.desc = context.GetRequiredInputDesc(CACHE_INDEX_INPUT_INDEX_V3); + mlaPrologContext.cacheIndex.shape = context.GetRequiredInputShape(CACHE_INDEX_INPUT_INDEX_V3); + } + mlaPrologContext.dequantScaleX.desc = context.GetOptionalInputDesc(DEQUANT_SCALE_X_INDEX); + mlaPrologContext.dequantScaleX.shape = context.GetOptionalInputShape(DEQUANT_SCALE_X_INDEX); + mlaPrologContext.dequantScaleWDq.desc = context.GetOptionalInputDesc(DEQUANT_SCALE_W_DQ_INDEX); + mlaPrologContext.dequantScaleWDq.shape = context.GetOptionalInputShape(DEQUANT_SCALE_W_DQ_INDEX); + mlaPrologContext.dequantScaleWUqQr.desc = context.GetOptionalInputDesc(DEQUANT_SCALE_W_UQ_QR_INDEX); + mlaPrologContext.dequantScaleWUqQr.shape = context.GetOptionalInputShape(DEQUANT_SCALE_W_UQ_QR_INDEX); + mlaPrologContext.dequantScaleWDkvKr.desc = context.GetOptionalInputDesc(DEQUANT_SCALE_W_DKV_KR_INDEX); + mlaPrologContext.dequantScaleWDkvKr.shape = context.GetOptionalInputShape(DEQUANT_SCALE_W_DKV_KR_INDEX); + mlaPrologContext.quantScaleCkv.desc = context.GetOptionalInputDesc(QUANT_SCALE_CKV_INDEX); + mlaPrologContext.quantScaleCkv.shape = context.GetOptionalInputShape(QUANT_SCALE_CKV_INDEX); + mlaPrologContext.quantScaleCkr.desc = context.GetOptionalInputDesc(QUANT_SCALE_CKR_INDEX); + mlaPrologContext.quantScaleCkr.shape = context.GetOptionalInputShape(QUANT_SCALE_CKR_INDEX); + mlaPrologContext.smoothScalesCq.desc = context.GetOptionalInputDesc(SMOOTH_SCALES_CQ_INDEX); + mlaPrologContext.smoothScalesCq.shape = context.GetOptionalInputShape(SMOOTH_SCALES_CQ_INDEX); + mlaPrologContext.actualSeqLen.desc = context.GetOptionalInputDesc(ACTUAL_SEQ_LEN_INDEX); + mlaPrologContext.actualSeqLen.shape = context.GetOptionalInputShape(ACTUAL_SEQ_LEN_INDEX); + mlaPrologContext.kNopeClipAlpha.desc = context.GetOptionalInputDesc(K_NOPE_CLIP_ALPHA_INDEX); + mlaPrologContext.kNopeClipAlpha.shape = context.GetOptionalInputShape(K_NOPE_CLIP_ALPHA_INDEX); + + // only v1 does not support dequantScaleQNope + if (std::strncmp(mlaPrologContext.opType, V1_OP_NAME, OP_NAME_LEN) == 0) { + mlaPrologContext.dequantScaleQNope.desc = nullptr; + mlaPrologContext.dequantScaleQNope.shape = nullptr; + } else { + mlaPrologContext.dequantScaleQNope.desc = context.GetOutputDesc(DEQUANT_SCALE_Q_NOPE_OUTPUT_INDEX); + mlaPrologContext.dequantScaleQNope.shape = context.GetOutputShape(DEQUANT_SCALE_Q_NOPE_OUTPUT_INDEX); + } + + if (std::strncmp(mlaPrologContext.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + mlaPrologContext.queryNorm.desc = context.GetOutputDesc(QUERY_NORM_OUTPUT_INDEX); + mlaPrologContext.queryNorm.shape = context.GetOutputShape(QUERY_NORM_OUTPUT_INDEX); + mlaPrologContext.dequantScaleQNorm.desc = context.GetOutputDesc(DEQUANT_SCALE_Q_NORM_OUTPUT_INDEX); + mlaPrologContext.dequantScaleQNorm.shape = context.GetOutputShape(DEQUANT_SCALE_Q_NORM_OUTPUT_INDEX); + } else { + mlaPrologContext.queryNorm.desc = nullptr; + mlaPrologContext.queryNorm.shape = nullptr; + mlaPrologContext.dequantScaleQNorm.desc = nullptr; + mlaPrologContext.dequantScaleQNorm.shape = nullptr; + } +} + +MLA_EXTERN_C ge::graphStatus TilingMlaProlog(gert::TilingContext *context) +{ + OP_CHECK_IF(context == nullptr, OPS_REPORT_VECTOR_INNER_ERR(V1_OP_NAME, "Context is nullptr."), + return ge::GRAPH_FAILED); + + MlaPrologContext mlaPrologContext{}; + OP_CHECK_IF(MlaPrologTiling::ConvertContext(*context, mlaPrologContext) != ge::GRAPH_SUCCESS, + OP_LOGE_WITHOUT_REPORT(context->GetNodeName(), + "Error occurred while converting tilingContext to MlaProlog context."), + return ge::GRAPH_FAILED); + + MlaPrologTiling mlaPrologTiling; + MlaPrologTilingData *tilingData = context->GetTilingData(); + OP_CHECK_IF(tilingData == nullptr, OPS_REPORT_VECTOR_INNER_ERR(context->GetNodeName(), "TilingData is nullptr."), + return ge::GRAPH_FAILED); + if (mlaPrologTiling.RunBigKernelTiling(mlaPrologContext, tilingData) == ge::SUCCESS) { + context->SetTilingKey(mlaPrologContext.tilingKey); + context->SetBlockDim(mlaPrologContext.blockDim); + return ge::GRAPH_SUCCESS; + } + return ge::GRAPH_FAILED; +} +} // namespace optiling diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.h b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.h new file mode 100644 index 000000000000..bfe22001205a --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.h @@ -0,0 +1,474 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_tiling.h + * \brief + */ + +#ifndef MLA_PROLOG_TILING_H +#define MLA_PROLOG_TILING_H + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include "register/tilingdata_base.h" +#include "tiling/tiling_api.h" +#include "tiling_base/data_copy_transpose_tiling.h" +#include "exe_graph/runtime/tiling_context.h" +#include "register/op_def_registry.h" +#include "../op_kernel/mla_prolog_template_tiling_key.h" +#include "../op_kernel/mla_prolog_tiling_data.h" +#include "platform/soc_spec.h" + +#ifdef ASCENDC_OP_TEST +#define MLA_EXTERN_C extern "C" +#else +#define MLA_EXTERN_C +#endif + +namespace optiling { + +// INPUT +constexpr uint32_t TOKEN_X_INPUT_INDEX = 0; +constexpr uint32_t WEIGHT_DQ_INPUT_INDEX = 1; +constexpr uint32_t WEIGHT_UQ_QR_INPUT_INDEX = 2; +constexpr uint32_t WEIGHT_UK_INPUT_INDEX = 3; +constexpr uint32_t WEIGHT_DKV_KR_INPUT_INDEX = 4; +constexpr uint32_t RMSNORM_GAMMA_CQ_INPUT_INDEX = 5; +constexpr uint32_t RMS_NORM_GAMMA_CKV_INPUT_INDEX = 6; +constexpr uint32_t ROPE_SIN_INPUT_INDEX = 7; +constexpr uint32_t ROPE_COS_INPUT_INDEX = 8; +constexpr uint32_t CACHE_INDEX_INPUT_INDEX = 9; +constexpr uint32_t KV_CACHE_INPUT_INDEX = 10; +constexpr uint32_t KR_CACHE_INPUT_INDEX = 11; + +constexpr uint32_t KV_CACHE_INPUT_INDEX_V3 = 9; +constexpr uint32_t KR_CACHE_INPUT_INDEX_V3 = 10; +constexpr uint32_t CACHE_INDEX_INPUT_INDEX_V3 = 11; + +// INPUT(OPTION) +constexpr uint32_t DEQUANT_SCALE_X_INDEX = 12; +constexpr uint32_t DEQUANT_SCALE_W_DQ_INDEX = 13; +constexpr uint32_t DEQUANT_SCALE_W_UQ_QR_INDEX = 14; +constexpr uint32_t DEQUANT_SCALE_W_DKV_KR_INDEX = 15; +constexpr uint32_t QUANT_SCALE_CKV_INDEX = 16; +constexpr uint32_t QUANT_SCALE_CKR_INDEX = 17; +constexpr uint32_t SMOOTH_SCALES_CQ_INDEX = 18; +constexpr uint32_t ACTUAL_SEQ_LEN_INDEX = 19; +constexpr uint32_t K_NOPE_CLIP_ALPHA_INDEX = 20; + +// OUTPUT +constexpr uint32_t QUERY_OUTPUT_INDEX = 0; +constexpr uint32_t QUERY_ROPE_OUTPUT_INDEX = 1; +constexpr uint32_t KV_CACHE_OUT_OUTPUT_INDEX = 2; +constexpr uint32_t KR_CACHE_OUT_OUTPUT_INDEX = 3; +constexpr uint32_t DEQUANT_SCALE_Q_NOPE_OUTPUT_INDEX = 4; +constexpr uint32_t QUERY_NORM_OUTPUT_INDEX = 5; +constexpr uint32_t DEQUANT_SCALE_Q_NORM_OUTPUT_INDEX = 6; + +// ATTR +constexpr uint32_t RMS_NORM_EPSILON_CQ_ATTR_INDEX = 0; +constexpr uint32_t RMS_NORM_EPSILON_CKV_ATTR_INDEX = 1; +constexpr uint32_t CACHE_MODE_ATTR_INDEX = 2; +constexpr uint32_t QUERY_NORM_FLAG_ATTR_INDEX = 3; +constexpr uint32_t WEIGHT_QUANT_MODE_ATTR_INDEX = 4; +constexpr uint32_t KV_CACHE_QUANT_MODE_ATTR_INDEX = 5; +constexpr uint32_t QUERY_QUANT_MODE_ATTR_INDEX = 6; +constexpr uint32_t CKVKR_REPO_MODE_ATTR_INDEX = 7; +constexpr uint32_t QUANT_SCALE_REPO_MODE_ATTR_INDEX = 8; +constexpr uint32_t TILE_SIZE_ATTR_INDEX = 9; +constexpr uint32_t QC_QR_SCALE_ATTR_INDEX = 10; +constexpr uint32_t KC_SCALE_ATTR_INDEX = 11; + +constexpr uint32_t MLA_PROLOG_DIM_INDEX_0 = 0; +constexpr uint32_t MLA_PROLOG_DIM_INDEX_1 = 1; +constexpr uint32_t MLA_PROLOG_DIM_INDEX_2 = 2; +constexpr uint32_t MLA_PROLOG_DIM_INDEX_3 = 3; + +constexpr uint32_t MLA_PROLOG_DIM_NUM_0 = 0; +constexpr uint32_t MLA_PROLOG_DIM_NUM_1 = 1; +constexpr uint32_t MLA_PROLOG_DIM_NUM_2 = 2; +constexpr uint32_t MLA_PROLOG_DIM_NUM_3 = 3; +constexpr uint32_t MLA_PROLOG_DIM_NUM_4 = 4; + +constexpr uint32_t HEAD_SIZE1 = 7168; +constexpr uint32_t HEAD_SIZE2 = 7680; + +constexpr char CACHE_MODE_BSND[]{"BSND"}; +constexpr char CACHE_MODE_TND[]{"TND"}; +constexpr char CACHE_MODE_PA_BSND[]{"PA_BSND"}; +constexpr char CACHE_MODE_PA_NZ[]{"PA_NZ"}; +constexpr char CACHE_MODE_PA_BLK_BSND[]{"PA_BLK_BSND"}; +constexpr char CACHE_MODE_PA_BLK_NZ[]{"PA_BLK_NZ"}; + +constexpr char V1_OP_NAME[]{"MlaProlog"}; +constexpr char V2_OP_NAME[]{"MlaPrologV2"}; +constexpr char V3_OP_NAME[]{"MlaPrologV3"}; + + +constexpr uint32_t CACHE_MODE_LEN = + std::max({sizeof(CACHE_MODE_BSND), sizeof(CACHE_MODE_TND), sizeof(CACHE_MODE_PA_BSND), sizeof(CACHE_MODE_PA_NZ), + sizeof(CACHE_MODE_PA_BLK_BSND), sizeof(CACHE_MODE_PA_BLK_NZ)}); + +constexpr uint32_t OP_NAME_LEN = std::max({sizeof(V1_OP_NAME), sizeof(V2_OP_NAME), sizeof(V3_OP_NAME)}); + +struct MlaPrologBaseShapeInfo { + uint32_t bSize = 0; // B + uint32_t s1Size = 0; // S1 + uint32_t tSize = 0; // T + uint32_t s2Size = 0; // S2 + uint32_t heSize = 0; // He + uint32_t hcqSize = 0; // Hcq + uint32_t hckvSize = 0; // Hckv + uint32_t headSizeQc = 0; // N * D + uint32_t headSizeQr = 0; // N * Dr + uint32_t headSizeUqQr = 0; // N * (D + Dr) + uint32_t nSize = 0; // N + uint32_t nkvSize = 0; // Nkv + uint32_t dSize = 0; // D + uint32_t blockNum = 0; + uint32_t blockSize = 0; + uint32_t drSize = 0; // Dr + uint32_t dtileSize = 0; // Dtile +}; + +enum class CACHE_MODE : uint8_t { + BSND = 0, + PA_BSND = 1, + PA_NZ = 2, + PA_BLK_BSND = 3, + PA_BLK_NZ = 4, + TND = 5 +}; + +// cacheMode 字符串 → CACHE_MODE 枚举 hash 查找表 +// CheckCacheMode 已保证执行到此处的 cacheMode 必为 6 个合法串之一 +inline const std::unordered_map CACHE_MODE_HASH_TABLE = { + {CACHE_MODE_BSND, CACHE_MODE::BSND}, + {CACHE_MODE_TND, CACHE_MODE::TND}, + {CACHE_MODE_PA_BSND, CACHE_MODE::PA_BSND}, + {CACHE_MODE_PA_NZ, CACHE_MODE::PA_NZ}, + {CACHE_MODE_PA_BLK_BSND, CACHE_MODE::PA_BLK_BSND}, + {CACHE_MODE_PA_BLK_NZ, CACHE_MODE::PA_BLK_NZ}}; + +enum class EMPTY_TENSOR_MODE : uint8_t { + NON_EMPTY = 0, + EMPTY_CACHE = 1, + EMPTY_QUERY = 2 +}; + +enum class ACTUAL_SEQ_MODE : uint8_t { + DISABLED = 0, + EN_Q_LEN = 1, +}; + +enum class QUANT_MODE : int8_t { + ERROR_MODE = -1, + NO_QUANT = 0, + PARTIAL_QUANT_KV_NO_QUANT = 1, + PARTIAL_QUANT_KV_QUANT_PER_CHANNEL = 2, + FULL_QUANT_KV_NO_QUANT = 3, + FULL_QUANT_KV_QUANT_PER_TENSOR = 4, + PARTIAL_QUANT_KV_QUANT_PER_TILE = 5, + FULL_QUANT_KV_QUANT_PER_TILE = 6, + MXFP8_FULL_QUANT_KV_NO_QUANT = 7, + MXFP8_FULL_QUANT_KV_QUANT_PER_TENSOR = 8, + MXFP8_FULL_QUANT_KV_QUANT_PER_TILE = 9, + FP8_FULL_QUANT_KV_NO_QUANT = 10, + FP8_FULL_QUANT_KV_QUANT_PER_TENSOR = 11, + HIF8_FULL_QUANT_KV_NO_QUANT = 12, + HIF8_FULL_QUANT_KV_QUANT_PER_TENSOR = 13, + FP8_FULL_QUANT_KV_QUANT_PER_TILE = 14, + HIF8_FULL_QUANT_KV_QUANT_PER_TILE = 15 +}; + +enum class WEIGHT_QUANT_MODE : uint8_t { + NO_QUANT = 0, + PARTIAL_QUANT = 1, + FULL_QUANT = 2, + MXFP8_FULL_QUANT = 3, + FP8_FULL_QUANT = 4, + HIF8_FULL_QUANT = 5 +}; + +enum class KV_QUANT_MODE : uint8_t { + NO_QUANT = 0, + PER_TENSOR = 1, + PER_CHANNEL = 2, + PER_TILE = 3 +}; + +// weightQuantMode(0-5) × kvQuantMode(0-3) → QUANT_MODE 双重哈希查找表 +// 外层 Key 是 weightQuantMode,内层 Key 是 kvQuantMode,Value 是 QUANT_MODE +// 只存放合法组合,查不到即非法(含越界与组合非法) +inline const std::unordered_map> QUANT_MODE_HASH_TABLE = { + // wq=0 NO_QUANT: 合法 kv 仅 {0} + {0, {{0, QUANT_MODE::NO_QUANT}}}, + // wq=1 PARTIAL: 合法 kv 为 {0, 2, 3} + {1, + {{0, QUANT_MODE::PARTIAL_QUANT_KV_NO_QUANT}, + {2, QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_CHANNEL}, + {3, QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_TILE}}}, + // wq=2 FULL: 合法 kv 为 {0, 1, 3} + {2, + {{0, QUANT_MODE::FULL_QUANT_KV_NO_QUANT}, + {1, QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TENSOR}, + {3, QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TILE}}}, + // wq=3 MXFP8: 合法 kv 为 {0, 1, 3} + {3, + {{0, QUANT_MODE::MXFP8_FULL_QUANT_KV_NO_QUANT}, + {1, QUANT_MODE::MXFP8_FULL_QUANT_KV_QUANT_PER_TENSOR}, + {3, QUANT_MODE::MXFP8_FULL_QUANT_KV_QUANT_PER_TILE}}}, + // wq=4 FP8: 合法 kv 为 {0, 1, 3} + {4, + {{0, QUANT_MODE::FP8_FULL_QUANT_KV_NO_QUANT}, + {1, QUANT_MODE::FP8_FULL_QUANT_KV_QUANT_PER_TENSOR}, + {3, QUANT_MODE::FP8_FULL_QUANT_KV_QUANT_PER_TILE}}}, + // wq=5 HIF8: 合法 kv 为 {0, 1, 3} + {5, + {{0, QUANT_MODE::HIF8_FULL_QUANT_KV_NO_QUANT}, + {1, QUANT_MODE::HIF8_FULL_QUANT_KV_QUANT_PER_TENSOR}, + {3, QUANT_MODE::HIF8_FULL_QUANT_KV_QUANT_PER_TILE}}}}; + +// 各 weightQuantMode 下合法 kvQuantMode 集合说明(用于错误日志 reason 文本,与原实现逐字一致) +inline const std::unordered_map VALID_KV_REASON_TABLE = { + {0, "When weightQuantMode==0, must be {0}"}, {1, "When weightQuantMode==1, must be {0, 2, 3}"}, + {2, "When weightQuantMode==2, must be {0, 1, 3}"}, {3, "When weightQuantMode==3, must be {0, 1, 3}"}, + {4, "When weightQuantMode==4, must be {0, 1, 3}"}, {5, "When weightQuantMode==5, must be {0, 1, 3}"}}; + +// QUANT_MODE → (weightQuantMode, kvQuantMode) 反向映射表 +// 用于从 quantMode_ 反推 wq/kvq,V1/V2(按 dtype 推断 quantMode_)与 V3(正向查表得 quantMode_)通用 +// 每个 QUANT_MODE 唯一对应一组 (wq, kvq),与正向表 QUANT_MODE_HASH_TABLE 互逆 +inline const std::unordered_map> QUANT_MODE_REVERSE_TABLE = { + {QUANT_MODE::NO_QUANT, {WEIGHT_QUANT_MODE::NO_QUANT, KV_QUANT_MODE::NO_QUANT}}, + {QUANT_MODE::PARTIAL_QUANT_KV_NO_QUANT, {WEIGHT_QUANT_MODE::PARTIAL_QUANT, KV_QUANT_MODE::NO_QUANT}}, + {QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_CHANNEL, {WEIGHT_QUANT_MODE::PARTIAL_QUANT, KV_QUANT_MODE::PER_CHANNEL}}, + {QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_TILE, {WEIGHT_QUANT_MODE::PARTIAL_QUANT, KV_QUANT_MODE::PER_TILE}}, + {QUANT_MODE::FULL_QUANT_KV_NO_QUANT, {WEIGHT_QUANT_MODE::FULL_QUANT, KV_QUANT_MODE::NO_QUANT}}, + {QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TENSOR, {WEIGHT_QUANT_MODE::FULL_QUANT, KV_QUANT_MODE::PER_TENSOR}}, + {QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TILE, {WEIGHT_QUANT_MODE::FULL_QUANT, KV_QUANT_MODE::PER_TILE}}, + {QUANT_MODE::MXFP8_FULL_QUANT_KV_NO_QUANT, {WEIGHT_QUANT_MODE::MXFP8_FULL_QUANT, KV_QUANT_MODE::NO_QUANT}}, + {QUANT_MODE::MXFP8_FULL_QUANT_KV_QUANT_PER_TENSOR, + {WEIGHT_QUANT_MODE::MXFP8_FULL_QUANT, KV_QUANT_MODE::PER_TENSOR}}, + {QUANT_MODE::MXFP8_FULL_QUANT_KV_QUANT_PER_TILE, {WEIGHT_QUANT_MODE::MXFP8_FULL_QUANT, KV_QUANT_MODE::PER_TILE}}, + {QUANT_MODE::FP8_FULL_QUANT_KV_NO_QUANT, {WEIGHT_QUANT_MODE::FP8_FULL_QUANT, KV_QUANT_MODE::NO_QUANT}}, + {QUANT_MODE::FP8_FULL_QUANT_KV_QUANT_PER_TENSOR, {WEIGHT_QUANT_MODE::FP8_FULL_QUANT, KV_QUANT_MODE::PER_TENSOR}}, + {QUANT_MODE::FP8_FULL_QUANT_KV_QUANT_PER_TILE, {WEIGHT_QUANT_MODE::FP8_FULL_QUANT, KV_QUANT_MODE::PER_TILE}}, + {QUANT_MODE::HIF8_FULL_QUANT_KV_NO_QUANT, {WEIGHT_QUANT_MODE::HIF8_FULL_QUANT, KV_QUANT_MODE::NO_QUANT}}, + {QUANT_MODE::HIF8_FULL_QUANT_KV_QUANT_PER_TENSOR, {WEIGHT_QUANT_MODE::HIF8_FULL_QUANT, KV_QUANT_MODE::PER_TENSOR}}, + {QUANT_MODE::HIF8_FULL_QUANT_KV_QUANT_PER_TILE, {WEIGHT_QUANT_MODE::HIF8_FULL_QUANT, KV_QUANT_MODE::PER_TILE}}}; + +enum class QUERY_QUANT_MODE : uint8_t { + NO_QUANT = 0, + PER_TOKEN_HEAD = 1 +}; + +enum class CKVKR_REPO_MODE : uint8_t { + DIVIDE = 0, + COMBINE = 1 +}; + +enum class QUANT_SCALE_REPO_MODE : uint8_t { + DIVIDE = 0, + COMBINE = 1 +}; + +struct MlaPrologScenarioInfo { + bool isV1Flag_; + bool batchSeqFusedFlag_; + uint8_t splitMFlag_; + QUANT_MODE quantMode_; + CACHE_MODE cacheMode_; + EMPTY_TENSOR_MODE emptyTensorMode_; + ACTUAL_SEQ_MODE actualSeqMode_; + WEIGHT_QUANT_MODE weightQuantMode_; + KV_QUANT_MODE kvQuantMode_; +}; + +struct MlaPrologCompileInfo { + int64_t core_num; +}; + +struct BaseParaInfo { + const gert::CompileTimeTensorDesc *desc; + const gert::StorageShape *shape; +}; + +struct RequiredParaInfo : BaseParaInfo {}; + +struct OptionalParaInfo : BaseParaInfo {}; + +constexpr uint32_t BLOCK_SIZE = 32; +constexpr uint32_t NUM_BYTES_BF16 = 2; +constexpr uint32_t NUM_BYTES_INT8 = 1; +constexpr uint32_t NUM_BYTES_INT32 = 4; +constexpr uint32_t NUM_BYTES_FP32 = 4; + +constexpr uint32_t GROUP_COMPUTE_CUBE_NUM_PER_GROUP = 8U; +constexpr uint32_t HIGH_THROUGHPUT__D_SIZE = 128U; +constexpr uint32_t GROUP_COMPUTE_T_SIZE = 1U; +constexpr uint32_t GROUP_COMPUTE_NKV_SIZE = 8U; +constexpr uint32_t GROUP_COMPUTE_MIN_AIC_NUM = 16U; +constexpr uint32_t GROUP_COMPUTE_MIN_AIV_NUM = 32U; +constexpr uint32_t GROUP_COMPUTE_N_SIZE = 8U; + +struct MlaPrologContext { + const char *opName; + const char *opType; + fe::PlatFormInfos *platformInfo; + RequiredParaInfo tokenX; + RequiredParaInfo weightDq; + RequiredParaInfo weightUqQr; + RequiredParaInfo weightUk; + RequiredParaInfo weightDkvKr; + RequiredParaInfo rmsnormGammaCq; + RequiredParaInfo rmsnormGammaCkv; + RequiredParaInfo ropeSin; + RequiredParaInfo ropeCos; + RequiredParaInfo kvCache; + RequiredParaInfo krCache; + OptionalParaInfo cacheIndex; + OptionalParaInfo dequantScaleX; + OptionalParaInfo dequantScaleWDq; + OptionalParaInfo dequantScaleWUqQr; + OptionalParaInfo dequantScaleWDkvKr; + OptionalParaInfo quantScaleCkv; + OptionalParaInfo quantScaleCkr; + OptionalParaInfo smoothScalesCq; + OptionalParaInfo actualSeqLen; + OptionalParaInfo kNopeClipAlpha; + RequiredParaInfo query; + RequiredParaInfo queryRope; + RequiredParaInfo kvCacheOut; + RequiredParaInfo krCacheOut; + OptionalParaInfo dequantScaleQNope; + OptionalParaInfo queryNorm; + OptionalParaInfo dequantScaleQNorm; + + const float *rmsNormEspilonCq; + const float *rmsNormEspilonCkv; + const char *cacheMode; + const bool *queryNormFlag; + + const int64_t *weightQuantMode; + const int64_t *kvQuantMode; + const int64_t *queryQuantMode; + const int64_t *ckvkrRepoMode; + const int64_t *quantScaleRepoMode; + const int64_t *tileSize; + + const float *qcQrScale; + const float *kcScale; + // Inferred in ConvertContext from ropeSin/ropeCos nullity (no OpDef Attr). + bool doRopeValue = true; + const bool *doRope = nullptr; + + size_t *workSpaces; + uint64_t tilingKey; + uint32_t blockDim; +}; + +class MlaPrologTiling { +public: + MlaPrologTiling() = default; + ~MlaPrologTiling() = default; + + ge::graphStatus RunBigKernelTiling(MlaPrologContext &context, MlaPrologTilingData *tilingData); + static ge::graphStatus ConvertContext(gert::TilingContext &context, MlaPrologContext &mlaPrologContext); + +private: + static void ConvertRequiredParams(gert::TilingContext &context, MlaPrologContext &mlaPrologContext); + static void ConvertOptionalParams(gert::TilingContext &context, MlaPrologContext &mlaPrologContext); + ge::graphStatus GetNpuInfo(); + ge::graphStatus SetScenarioInfo(); + ge::graphStatus SetAttrInfo(); + QUANT_MODE GetQuantizationMode() const; + QUANT_MODE GetQuantizationModeV3() const; + QUANT_MODE GetQuantizationModeV3Dav() const; + ge::graphStatus SetShapeInfo(); + ge::graphStatus ProcessBaseInputs(); + ge::graphStatus FillTiling(); + void FillTilingCoreParams(); + ge::graphStatus FillMatmul1Tiling(); + ge::graphStatus FillMatmul2Tiling(); + ge::graphStatus FillMatmul3Tiling(); + ge::graphStatus FillMatmul4Tiling(); + uint32_t CalcSingleCoreN(uint32_t n, uint32_t coreNum, uint32_t alignNum = 16) const; + bool GetMatmulType(ge::DataType getype, matmul_tiling::DataType *mmType); + ge::graphStatus CalcWorkSpace(); + ge::graphStatus GenTilingKey() const; + + NpuArch GetCurNpuArch() const; + + MlaPrologBaseShapeInfo baseShapeInfo_; + MlaPrologScenarioInfo scenarioInfo_; + + uint32_t stepBatchSize_ = 0; + uint32_t stepNumHeadDequant_ = 0; + uint32_t mSubSize_ = 0; + uint32_t mSubCoreNum_ = 0; + + uint32_t mm1BlockNum_ = 0; + uint32_t mm2BlockNum_ = 0; + uint32_t mm3BlockNum_ = 0; + uint32_t mm4BlockNum_ = 0; + uint32_t vectorBlockNum_ = 0; + + uint32_t singlecoreHeadSizeCq_ = 0; + uint32_t singlecoreHeadSizeQcQr_ = 0; + uint32_t singlecoreHeadSizeCkvKr_ = 0; + uint32_t singlecoreNumHeadSize_ = 0; + + float reciprocalCq_ = 0.00001f; + float epsilonCq_ = 1.0f; + float reciprocalCkv_ = 0.00001f; + float epsilonCkv_ = 1.0f; + bool queryNormFlag_ = false; + + WEIGHT_QUANT_MODE weightQuantMode_ = WEIGHT_QUANT_MODE::NO_QUANT; + KV_QUANT_MODE kvQuantMode_ = KV_QUANT_MODE::NO_QUANT; + QUERY_QUANT_MODE queryQuantMode_ = QUERY_QUANT_MODE::NO_QUANT; + CKVKR_REPO_MODE ckvkrRepoMode_ = CKVKR_REPO_MODE::DIVIDE; + QUANT_SCALE_REPO_MODE quantSacleRepoMode_ = QUANT_SCALE_REPO_MODE::DIVIDE; + uint32_t tileSize_ = 128; + float qcQrScale_ = 1.0f; + float kcScale_ = 1.0f; + + ge::DataType mmDateType_ = ge::DT_BF16; + bool enableDequantOpt_ = false; + bool enableGroupComputeOpt_ = false; // 低延时场景算例分组标记 + bool enableRope_ = true; // rope开关,仅 MlaPrologV3 在 DAV_3510 生效 + + size_t ubSize_ = 0; + size_t l1Size_ = 0; + size_t l0cSize_ = 0; + size_t l0bSize_ = 0; + uint32_t coreNum_ = 0; + uint32_t aicNum_ = 0; + uint32_t aivNum_ = 0; + size_t libapiSize_ = 0; + size_t workspaceSize_ = 0; + + MlaPrologContext *context_ = nullptr; + MlaPrologBaseParams *baseParams_ = nullptr; +}; + +ge::graphStatus TilingPrepareForMlaProlog(gert::TilingParseContext *context); +MLA_EXTERN_C ge::graphStatus TilingMlaProlog(gert::TilingContext *context); +} // namespace optiling + +#endif // MLA_PROLOG_TILING_H diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling_check.cpp b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling_check.cpp new file mode 100644 index 000000000000..31ab33f112cc --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling_check.cpp @@ -0,0 +1,1285 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_tiling_check.cpp + * \brief + */ + +#include "mla_prolog_tiling_check.h" +#include +#include "log/log.h" + +using namespace ge; + +namespace optiling { + +const std::unordered_map DTYPE_TO_SIZE{ + {ge::DT_BF16, 2}, {ge::DT_FLOAT16, 2}, {ge::DT_INT8, 1}, {ge::DT_FLOAT8_E4M3FN, 1}, + {ge::DT_FLOAT8_E8M0, 1}, {ge::DT_HIFLOAT8, 1}, {ge::DT_INT32, 4}, {ge::DT_FLOAT, 4}}; + +template +std::string ElemToString(const E &elem) +{ + return std::to_string(elem); +} + +std::string FormatToString(const ge::Format format) +{ + return std::string(ge::GetFormatName(format)); +} + +template +std::string ConvertContainerToString(const C &container, Func func = ElemToString) +{ + if (container.empty() || func == nullptr) { + return "[]"; + } + std::stringstream ss; + ss << "["; + bool isFirst = true; + for (const auto &elem : container) { + if (!isFirst) { + ss << ", "; + } + ss << func(elem); + isFirst = false; + } + ss << "]"; + return ss.str(); +} + +template +std::string ConvertContainerToStringV3(const C &container, Func func = ElemToString) +{ + if (container.empty() || func == nullptr) { + return "[]"; + } + std::stringstream ss; + bool isFirst = true; + for (const auto &elem : container) { + if (!isFirst) { + ss << ", "; + } + ss << func(elem); + isFirst = false; + } + return ss.str(); +} + +std::string GetShapeStr(const gert::Shape &aShape) +{ + std::string shapeStr = "["; + for (size_t i = 0; i < aShape.GetDimNum(); ++i) { + shapeStr += std::to_string(aShape.GetDim(i)) + (i + 1 < aShape.GetDimNum() ? ", " : ""); + } + return shapeStr + "]"; +} + +template +inline auto CeilDiv(T a, T b) -> T +{ + if (b == 0) { + return b; + } + return (a + b - 1) / b; +} + +NpuArch MlaPrologTilingCheck::GetCurNpuArch() const +{ + auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_.platformInfo); + NpuArch npuArch = ascendcPlatform.GetCurNpuArch(); + return npuArch; +} + +// =================================全量参数校验================================= +bool MlaPrologTilingCheck::CheckAttrsRange() const +{ + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + if (GetCurNpuArch() == NpuArch::DAV_3510) { + const std::unordered_set supportedWeightQuantMode{0U, 1U, 2U, 3U, 4U, 5U}; + OP_CHECK_IF(supportedWeightQuantMode.find(*context_.weightQuantMode) == supportedWeightQuantMode.end(), + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "WeightQuantMode", + std::to_string(*context_.weightQuantMode), "{0, 1, 2, 3, 4, 5}"), + return false); + } else { + const std::unordered_set supportedWeightQuantMode{0U, 1U, 2U}; + OP_CHECK_IF(supportedWeightQuantMode.find(*context_.weightQuantMode) == supportedWeightQuantMode.end(), + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "WeightQuantMode", + std::to_string(*context_.weightQuantMode), "{0, 1, 2}"), + return false); + } + + const std::unordered_set supportedKvQuantMode{0U, 1U, 2U, 3U}; + OP_CHECK_IF(supportedKvQuantMode.find(*context_.kvQuantMode) == supportedKvQuantMode.end(), + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "KvQuantMode", std::to_string(*context_.kvQuantMode), + "{0, 1, 2, 3}"), + return false); + + const std::unordered_set supportedQueryQuantMode{0U, 1U}; + OP_CHECK_IF(supportedQueryQuantMode.find(*context_.queryQuantMode) == supportedQueryQuantMode.end(), + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "QueryQuantMode", + std::to_string(*context_.queryQuantMode), "{0, 1}"), + return false); + + const std::unordered_set supportedCkvkrRepoMode{0U, 1U}; + OP_CHECK_IF(supportedCkvkrRepoMode.find(*context_.ckvkrRepoMode) == supportedCkvkrRepoMode.end(), + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "CkvkrRepoMode", std::to_string(*context_.ckvkrRepoMode), + "{0, 1}"), + return false); + + const std::unordered_set supportedQuantScaleRepoMode{0U, 1U}; + OP_CHECK_IF(supportedQuantScaleRepoMode.find(*context_.quantScaleRepoMode) == supportedQuantScaleRepoMode.end(), + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "QuantScaleRepoMode", + std::to_string(*context_.quantScaleRepoMode), "{0, 1}"), + return false); + + const std::unordered_set supportedTileSize{128U}; + OP_CHECK_IF(supportedTileSize.find(*context_.tileSize) == supportedTileSize.end(), + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "TileSize", std::to_string(*context_.tileSize), "{128}"), + return false); + } + return true; +} + +bool MlaPrologTilingCheck::CheckAttrsNotNull() const +{ + OP_CHECK_IF(context_.rmsNormEspilonCq == nullptr, OP_LOGE_WITH_INVALID_INPUT(context_.opName, "rmsNormEspilonCq"), + return false); + + OP_CHECK_IF(context_.rmsNormEspilonCkv == nullptr, OP_LOGE_WITH_INVALID_INPUT(context_.opName, "rmsNormEspilonCkv"), + return false); + + OP_CHECK_IF(context_.cacheMode == nullptr, OP_LOGE_WITH_INVALID_INPUT(context_.opName, "cacheMode"), return false); + + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + OP_CHECK_IF(context_.queryNormFlag == nullptr, OP_LOGE_WITH_INVALID_INPUT(context_.opName, "queryNormFlag"), + return false); + + OP_CHECK_IF(context_.weightQuantMode == nullptr, OP_LOGE_WITH_INVALID_INPUT(context_.opName, "weightQuantMode"), + return false); + + OP_CHECK_IF(context_.kvQuantMode == nullptr, OP_LOGE_WITH_INVALID_INPUT(context_.opName, "kvQuantMode"), + return false); + + OP_CHECK_IF(context_.queryQuantMode == nullptr, OP_LOGE_WITH_INVALID_INPUT(context_.opName, "queryQuantMode"), + return false); + + OP_CHECK_IF(context_.ckvkrRepoMode == nullptr, OP_LOGE_WITH_INVALID_INPUT(context_.opName, "ckvkrRepoMode"), + return false); + + OP_CHECK_IF(context_.quantScaleRepoMode == nullptr, + OP_LOGE_WITH_INVALID_INPUT(context_.opName, "quantScaleRepoMode"), return false); + + OP_CHECK_IF(context_.tileSize == nullptr, OP_LOGE_WITH_INVALID_INPUT(context_.opName, "tileSize"), + return false); + + OP_CHECK_IF(context_.qcQrScale == nullptr, OP_LOGE_WITH_INVALID_INPUT(context_.opName, "qcQrScale"), + return false); + + OP_CHECK_IF(context_.kcScale == nullptr, OP_LOGE_WITH_INVALID_INPUT(context_.opName, "kcScale"), return false); + } + return true; +} + +ge::graphStatus MlaPrologTilingCheck::CheckAttrs() const +{ + if (!CheckAttrsNotNull() || !CheckAttrsRange()) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTilingCheck::CheckDims() const +{ + OP_CHECK_IF( + baseShapeInfo_.bSize > MAX_B_SIZE, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context_.opName, "B", std::to_string(baseShapeInfo_.bSize), + "B size should not be greater than " + std::to_string(MAX_B_SIZE)), + return ge::GRAPH_FAILED); + const std::set supportedHeSize{1024U, 2048U, 3072U, 4096U, 5120U, 6144U, 7168U, 7680U, 8192U}; + OP_CHECK_IF(supportedHeSize.find(baseShapeInfo_.heSize) == supportedHeSize.end(), + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "He", std::to_string(baseShapeInfo_.heSize), + ConvertContainerToStringV3(supportedHeSize)), + return ge::GRAPH_FAILED); + + OP_CHECK_IF(baseShapeInfo_.nSize < 1U || baseShapeInfo_.nSize > 128U, + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "N", std::to_string(baseShapeInfo_.nSize), + "N size should be within [1, 128]"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(baseShapeInfo_.hckvSize != HCKV_SIZE, + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "Hckv", std::to_string(baseShapeInfo_.hckvSize), + std::to_string(HCKV_SIZE)), + return ge::GRAPH_FAILED); + OP_CHECK_IF(baseShapeInfo_.drSize != DR_SIZE, + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "Dr", std::to_string(baseShapeInfo_.drSize), + std::to_string(DR_SIZE)), + return ge::GRAPH_FAILED); + OP_CHECK_IF(baseShapeInfo_.nkvSize != NKV_SIZE, + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "Nkv", std::to_string(baseShapeInfo_.nkvSize), + std::to_string(NKV_SIZE)), + return ge::GRAPH_FAILED); + if (scenarioInfo_.cacheMode_ != CACHE_MODE::BSND && scenarioInfo_.cacheMode_ != CACHE_MODE::TND) { + OP_CHECK_IF(baseShapeInfo_.blockSize < MIN_BLOCK_SIZE || baseShapeInfo_.blockSize > MAX_BLOCK_SIZE || + baseShapeInfo_.blockSize % ALIGN_BLOCK_SIZE != 0, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON( + context_.opName, "blockSize", std::to_string(baseShapeInfo_.blockSize), + "BlockSize must be within [" + std::to_string(MIN_BLOCK_SIZE) + ", " + + std::to_string(MAX_BLOCK_SIZE) + "] and a multiple of " + std::to_string(ALIGN_BLOCK_SIZE)), + return ge::GRAPH_FAILED); + } + if (CheckHcqSize() != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + if (CheckDSize() != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + if (CheckDtileSize() != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTilingCheck::CheckQuantMode() const +{ + if (GetCurNpuArch() == NpuArch::DAV_3510) { + const std::set supportedQuantModes{ + static_cast(QUANT_MODE::NO_QUANT), + static_cast(QUANT_MODE::PARTIAL_QUANT_KV_NO_QUANT), + static_cast(QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_CHANNEL), + static_cast(QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_TILE), + static_cast(QUANT_MODE::FULL_QUANT_KV_NO_QUANT), + static_cast(QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TENSOR), + static_cast(QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TILE), + static_cast(QUANT_MODE::MXFP8_FULL_QUANT_KV_NO_QUANT), + static_cast(QUANT_MODE::MXFP8_FULL_QUANT_KV_QUANT_PER_TENSOR), + static_cast(QUANT_MODE::MXFP8_FULL_QUANT_KV_QUANT_PER_TILE), + static_cast(QUANT_MODE::FP8_FULL_QUANT_KV_NO_QUANT), + static_cast(QUANT_MODE::FP8_FULL_QUANT_KV_QUANT_PER_TENSOR), + static_cast(QUANT_MODE::HIF8_FULL_QUANT_KV_NO_QUANT), + static_cast(QUANT_MODE::HIF8_FULL_QUANT_KV_QUANT_PER_TENSOR), + static_cast(QUANT_MODE::FP8_FULL_QUANT_KV_QUANT_PER_TILE), + static_cast(QUANT_MODE::HIF8_FULL_QUANT_KV_QUANT_PER_TILE)}; + OP_CHECK_IF(supportedQuantModes.find(static_cast(scenarioInfo_.quantMode_)) == + supportedQuantModes.end(), + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON( + context_.opName, "quantMode", std::to_string(static_cast(scenarioInfo_.quantMode_)), + "On DAV3510, quantMode allows only " + ConvertContainerToStringV3(supportedQuantModes)), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTilingCheck::CheckHcqSize() const +{ + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + const std::set supportedHcqSize{1536U, 2048U}; + OP_CHECK_IF(supportedHcqSize.find(baseShapeInfo_.hcqSize) == supportedHcqSize.end(), + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "Hcq", std::to_string(baseShapeInfo_.hcqSize), + ConvertContainerToStringV3(supportedHcqSize)), + return ge::GRAPH_FAILED); + } else { + OP_CHECK_IF(baseShapeInfo_.hcqSize != HCQ_SIZE, + OP_LOGE(context_.opName, "Hcq allows only %u, got %u.", HCQ_SIZE, baseShapeInfo_.hcqSize), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTilingCheck::CheckDSize() const +{ + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + const std::set supportedDSize{128U, 192U}; + OP_CHECK_IF(supportedDSize.find(baseShapeInfo_.dSize) == supportedDSize.end(), + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "D", std::to_string(baseShapeInfo_.dSize), + ConvertContainerToStringV3(supportedDSize)), + return ge::GRAPH_FAILED); + } else { + OP_CHECK_IF(baseShapeInfo_.dSize != D_SIZE, + OP_LOGE(context_.opName, "D allows only %u, got %u.", D_SIZE, baseShapeInfo_.dSize), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTilingCheck::CheckDtileSize() const +{ + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + uint32_t supportedDtileSize = baseShapeInfo_.hckvSize; + if (*(context_.ckvkrRepoMode) == static_cast(CKVKR_REPO_MODE::COMBINE)) { + supportedDtileSize += + baseShapeInfo_.drSize * (DTYPE_TO_SIZE.at(ge::DT_BF16) / DTYPE_TO_SIZE.at(ge::DT_INT8)); + } + if (*(context_.quantScaleRepoMode) == static_cast(QUANT_SCALE_REPO_MODE::COMBINE)) { + supportedDtileSize += baseShapeInfo_.hckvSize / static_cast(*(context_.tileSize)) * + (DTYPE_TO_SIZE.at(ge::DT_FLOAT) / DTYPE_TO_SIZE.at(ge::DT_INT8)); + } + if (baseShapeInfo_.dtileSize != supportedDtileSize) { + if (*(context_.kvQuantMode) == static_cast(KV_QUANT_MODE::PER_TILE)) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON( + context_.opName, "dtileSize", std::to_string(baseShapeInfo_.dtileSize), + "when kvQuantMode is PER_TILE, dtileSize allows only " + std::to_string(supportedDtileSize)); + } else { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON( + context_.opName, "dtileSize", std::to_string(baseShapeInfo_.dtileSize), + "when kvQuantMode is in {NO_QUANT, PER_TENSOR, PER_CHANNEL}, dtileSize allows only " + + std::to_string(supportedDtileSize)); + } + return ge::GRAPH_FAILED; + } + } + return ge::GRAPH_SUCCESS; +} + +void MlaPrologTilingCheck::GenExpectedParamInfo() +{ + FillCommonParamInfo(); + FillScenarioParamInfo(); +} + +void MlaPrologTilingCheck::FillCommonParamInfo() +{ + FillRequiredParamShapeWithDims(); + FillOptionalOutputParamShapeWithDims(); + + if (context_.weightDq.shape->GetStorageShape().GetDimNum() == MLA_PROLOG_DIM_NUM_4) { + expectedParamInfo_[WEIGHT_DQ_NAME].dimNum = MLA_PROLOG_DIM_NUM_4; + int64_t weightAxisSize = 32L / ge::GetSizeByDataType(context_.weightDq.desc->GetDataType()); + expectedParamInfo_[WEIGHT_DQ_NAME].shape = + std::vector{static_cast(baseShapeInfo_.hcqSize) / weightAxisSize, + static_cast(baseShapeInfo_.heSize) / NZ_H0_SIZE, NZ_H0_SIZE, weightAxisSize}; + } + if (context_.weightUqQr.shape->GetStorageShape().GetDimNum() == MLA_PROLOG_DIM_NUM_4) { + expectedParamInfo_[WEIGHT_UQ_QR_NAME].dimNum = MLA_PROLOG_DIM_NUM_4; + int64_t weightAxisSize = 32L / ge::GetSizeByDataType(context_.weightUqQr.desc->GetDataType()); + expectedParamInfo_[WEIGHT_UQ_QR_NAME].shape = + std::vector{static_cast(baseShapeInfo_.headSizeUqQr) / weightAxisSize, + static_cast(baseShapeInfo_.hcqSize) / NZ_H0_SIZE, NZ_H0_SIZE, weightAxisSize}; + } + if (context_.weightDkvKr.shape->GetStorageShape().GetDimNum() == MLA_PROLOG_DIM_NUM_4) { + expectedParamInfo_[WEIGHT_DKV_KR_NAME].dimNum = MLA_PROLOG_DIM_NUM_4; + int64_t weightAxisSize = 32L / ge::GetSizeByDataType(context_.weightDkvKr.desc->GetDataType()); + expectedParamInfo_[WEIGHT_DKV_KR_NAME].shape = + std::vector{static_cast(baseShapeInfo_.hckvSize + baseShapeInfo_.drSize) / weightAxisSize, + static_cast(baseShapeInfo_.heSize) / NZ_H0_SIZE, NZ_H0_SIZE, weightAxisSize}; + } + + expectedParamInfo_[WEIGHT_DQ_NAME].format = ge::FORMAT_FRACTAL_NZ; + expectedParamInfo_[WEIGHT_UQ_QR_NAME].format = ge::FORMAT_FRACTAL_NZ; + expectedParamInfo_[WEIGHT_DKV_KR_NAME].format = ge::FORMAT_FRACTAL_NZ; + + if (scenarioInfo_.cacheMode_ == CACHE_MODE::PA_BLK_BSND || scenarioInfo_.cacheMode_ == CACHE_MODE::PA_BLK_NZ) { + if (scenarioInfo_.batchSeqFusedFlag_) { + expectedParamInfo_.emplace(ACTUAL_SEQ_LEN_NAME, std::vector{baseShapeInfo_.bSize}); + expectedParamInfo_[ACTUAL_SEQ_LEN_NAME].dtype = ge::DT_INT32; + expectedParamInfo_[ACTUAL_SEQ_LEN_NAME].format = ge::FORMAT_ND; + expectedParamInfo_[CACHE_INDEX_NAME].shape = actualParamInfo_[CACHE_INDEX_NAME].shape; + } else { + expectedParamInfo_[CACHE_INDEX_NAME].shape = + std::vector{baseShapeInfo_.bSize, CeilDiv(baseShapeInfo_.s1Size, baseShapeInfo_.blockSize)}; + } + } +} + +void MlaPrologTilingCheck::FillRequiredParamShapeWithDims() +{ + FillTokenAndQueryShapes(); + FillWeightAndNormShapes(); + FillCacheShapes(); + expectedParamInfo_.emplace(KV_CACHE_OUT_NAME, expectedParamInfo_[KV_CACHE_NAME]); + expectedParamInfo_.emplace(KR_CACHE_OUT_NAME, expectedParamInfo_[KR_CACHE_NAME]); +} + +void MlaPrologTilingCheck::FillTokenAndQueryShapes() +{ + const bool ropeEnabled = + !(std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0 && context_.doRope != nullptr && + !*(context_.doRope)); + if (scenarioInfo_.batchSeqFusedFlag_) { + expectedParamInfo_.emplace(TOKEN_X_NAME, std::vector{baseShapeInfo_.tSize, baseShapeInfo_.heSize}); + if (ropeEnabled) { + expectedParamInfo_.emplace(ROPE_SIN_NAME, + std::vector{baseShapeInfo_.tSize, baseShapeInfo_.drSize}); + expectedParamInfo_.emplace(ROPE_COS_NAME, + std::vector{baseShapeInfo_.tSize, baseShapeInfo_.drSize}); + } else { + expectedParamInfo_[ROPE_SIN_NAME].isValid = false; + expectedParamInfo_[ROPE_COS_NAME].isValid = false; + } + if ((scenarioInfo_.cacheMode_ != CACHE_MODE::BSND) && (scenarioInfo_.cacheMode_ != CACHE_MODE::TND)) { + expectedParamInfo_.emplace(CACHE_INDEX_NAME, std::vector{baseShapeInfo_.tSize}); + expectedParamInfo_[CACHE_INDEX_NAME].dtype = ge::DT_INT64; + } + expectedParamInfo_.emplace( + QUERY_NAME, std::vector{baseShapeInfo_.tSize, baseShapeInfo_.nSize, baseShapeInfo_.hckvSize}); + expectedParamInfo_.emplace( + QUERY_ROPE_NAME, std::vector{baseShapeInfo_.tSize, baseShapeInfo_.nSize, baseShapeInfo_.drSize}); + } else { + expectedParamInfo_.emplace( + TOKEN_X_NAME, std::vector{baseShapeInfo_.bSize, baseShapeInfo_.s1Size, baseShapeInfo_.heSize}); + if (ropeEnabled) { + expectedParamInfo_.emplace( + ROPE_SIN_NAME, std::vector{baseShapeInfo_.bSize, baseShapeInfo_.s1Size, baseShapeInfo_.drSize}); + expectedParamInfo_.emplace( + ROPE_COS_NAME, std::vector{baseShapeInfo_.bSize, baseShapeInfo_.s1Size, baseShapeInfo_.drSize}); + } else { + expectedParamInfo_[ROPE_SIN_NAME].isValid = false; + expectedParamInfo_[ROPE_COS_NAME].isValid = false; + } + if ((scenarioInfo_.cacheMode_ != CACHE_MODE::BSND) && (scenarioInfo_.cacheMode_ != CACHE_MODE::TND)) { + expectedParamInfo_.emplace(CACHE_INDEX_NAME, + std::vector{baseShapeInfo_.bSize, baseShapeInfo_.s1Size}); + expectedParamInfo_[CACHE_INDEX_NAME].dtype = ge::DT_INT64; + } + expectedParamInfo_.emplace(QUERY_NAME, std::vector{baseShapeInfo_.bSize, baseShapeInfo_.s1Size, + baseShapeInfo_.nSize, baseShapeInfo_.hckvSize}); + expectedParamInfo_.emplace(QUERY_ROPE_NAME, std::vector{baseShapeInfo_.bSize, baseShapeInfo_.s1Size, + baseShapeInfo_.nSize, baseShapeInfo_.drSize}); + } +} + +void MlaPrologTilingCheck::FillWeightAndNormShapes() +{ + expectedParamInfo_.emplace(WEIGHT_DQ_NAME, std::vector{baseShapeInfo_.heSize, baseShapeInfo_.hcqSize}); + expectedParamInfo_.emplace(WEIGHT_UQ_QR_NAME, + std::vector{baseShapeInfo_.hcqSize, baseShapeInfo_.headSizeUqQr}); + expectedParamInfo_.emplace( + WEIGHT_UK_NAME, std::vector{baseShapeInfo_.nSize, baseShapeInfo_.dSize, baseShapeInfo_.hckvSize}); + expectedParamInfo_.emplace( + WEIGHT_DKV_KR_NAME, + std::vector{baseShapeInfo_.heSize, baseShapeInfo_.hckvSize + baseShapeInfo_.drSize}); + expectedParamInfo_.emplace(RMSNORM_GAMMA_CQ_NAME, std::vector{baseShapeInfo_.hcqSize}); + expectedParamInfo_.emplace(RMSNORM_GAMMA_CKV_NAME, std::vector{baseShapeInfo_.hckvSize}); +} + +void MlaPrologTilingCheck::FillCacheShapes() +{ + if (scenarioInfo_.cacheMode_ == CACHE_MODE::TND) { + expectedParamInfo_.emplace(KV_CACHE_NAME, std::vector{baseShapeInfo_.tSize, baseShapeInfo_.nkvSize, + baseShapeInfo_.dtileSize}); + expectedParamInfo_.emplace( + KR_CACHE_NAME, std::vector{baseShapeInfo_.tSize, baseShapeInfo_.nkvSize, baseShapeInfo_.drSize}); + } else if (scenarioInfo_.cacheMode_ == CACHE_MODE::BSND) { + expectedParamInfo_.emplace(KV_CACHE_NAME, + std::vector{baseShapeInfo_.bSize, baseShapeInfo_.s1Size, + baseShapeInfo_.nkvSize, baseShapeInfo_.dtileSize}); + expectedParamInfo_.emplace(KR_CACHE_NAME, std::vector{baseShapeInfo_.bSize, baseShapeInfo_.s1Size, + baseShapeInfo_.nkvSize, baseShapeInfo_.drSize}); + } else { + expectedParamInfo_.emplace(KV_CACHE_NAME, + std::vector{baseShapeInfo_.blockNum, baseShapeInfo_.blockSize, + baseShapeInfo_.nkvSize, baseShapeInfo_.dtileSize}); + expectedParamInfo_.emplace(KR_CACHE_NAME, + std::vector{baseShapeInfo_.blockNum, baseShapeInfo_.blockSize, + baseShapeInfo_.nkvSize, baseShapeInfo_.drSize}); + } +} + +void MlaPrologTilingCheck::FillOptionalOutputParamShapeWithDims() +{ + if (std::strncmp(context_.opType, V2_OP_NAME, OP_NAME_LEN) == 0) { + FillOptionalOutputParamShapeWithDimsV2(); + } + + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + FillOptionalOutputParamShapeWithDimsV3(); + } +} + +void MlaPrologTilingCheck::FillOptionalOutputParamShapeWithDimsV2() +{ + if (scenarioInfo_.quantMode_ == QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TENSOR) { + expectedParamInfo_.emplace(DEQUANT_SCALE_Q_NOPE_NAME, + std::vector{baseShapeInfo_.tSize, baseShapeInfo_.nSize, 1}); + expectedParamInfo_[DEQUANT_SCALE_Q_NOPE_NAME].dtype = ge::DT_FLOAT; + } else { + // 仅校验dequantScaleQNope有传入 + expectedParamInfo_.emplace(DEQUANT_SCALE_Q_NOPE_NAME, context_.dequantScaleQNope); + } +} + +void MlaPrologTilingCheck::FillOptionalOutputParamShapeWithDimsV3() +{ + if (scenarioInfo_.kvQuantMode_ == KV_QUANT_MODE::PER_TENSOR) { + expectedParamInfo_.emplace(DEQUANT_SCALE_Q_NOPE_NAME, + std::vector{baseShapeInfo_.tSize, baseShapeInfo_.nSize, 1}); + } else { + expectedParamInfo_.emplace(DEQUANT_SCALE_Q_NOPE_NAME, std::vector{0}); + } + expectedParamInfo_[DEQUANT_SCALE_Q_NOPE_NAME].dtype = ge::DT_FLOAT; + + if (*(context_.queryNormFlag)) { + if (scenarioInfo_.batchSeqFusedFlag_) { + expectedParamInfo_.emplace(QUERY_NORM_NAME, + std::vector{baseShapeInfo_.tSize, baseShapeInfo_.hcqSize}); + } else { + expectedParamInfo_.emplace( + QUERY_NORM_NAME, + std::vector{baseShapeInfo_.bSize, baseShapeInfo_.s1Size, baseShapeInfo_.hcqSize}); + } + FillQueryNormScaleShape(); + } else { + expectedParamInfo_.emplace(QUERY_NORM_NAME, std::vector{0}); + expectedParamInfo_.emplace(DEQUANT_SCALE_Q_NORM_NAME, std::vector{0}); + } + FillQueryNormDtypes(); +} + +void MlaPrologTilingCheck::FillQueryNormScaleShape() +{ + if (scenarioInfo_.weightQuantMode_ == WEIGHT_QUANT_MODE::NO_QUANT) { + expectedParamInfo_.emplace(DEQUANT_SCALE_Q_NORM_NAME, std::vector{0}); + } else if (scenarioInfo_.weightQuantMode_ == WEIGHT_QUANT_MODE::MXFP8_FULL_QUANT) { + expectedParamInfo_.emplace( + DEQUANT_SCALE_Q_NORM_NAME, + std::vector{baseShapeInfo_.tSize, baseShapeInfo_.hcqSize / MXFP8_BLOCK_SIZE}); + } else { + expectedParamInfo_.emplace(DEQUANT_SCALE_Q_NORM_NAME, std::vector{baseShapeInfo_.tSize, 1}); + } +} + +void MlaPrologTilingCheck::FillQueryNormDtypes() +{ + if (scenarioInfo_.weightQuantMode_ == WEIGHT_QUANT_MODE::NO_QUANT) { + expectedParamInfo_[QUERY_NORM_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[DEQUANT_SCALE_Q_NORM_NAME].dtype = ge::DT_FLOAT; + } else if (scenarioInfo_.weightQuantMode_ == WEIGHT_QUANT_MODE::MXFP8_FULL_QUANT) { + expectedParamInfo_[QUERY_NORM_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[DEQUANT_SCALE_Q_NORM_NAME].dtype = ge::DT_FLOAT8_E8M0; + } else if (scenarioInfo_.weightQuantMode_ == WEIGHT_QUANT_MODE::FP8_FULL_QUANT) { + expectedParamInfo_[QUERY_NORM_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[DEQUANT_SCALE_Q_NORM_NAME].dtype = ge::DT_FLOAT; + } else if (scenarioInfo_.weightQuantMode_ == WEIGHT_QUANT_MODE::HIF8_FULL_QUANT) { + expectedParamInfo_[QUERY_NORM_NAME].dtype = ge::DT_HIFLOAT8; + expectedParamInfo_[DEQUANT_SCALE_Q_NORM_NAME].dtype = ge::DT_FLOAT; + } else { + expectedParamInfo_[QUERY_NORM_NAME].dtype = ge::DT_INT8; + expectedParamInfo_[DEQUANT_SCALE_Q_NORM_NAME].dtype = ge::DT_FLOAT; + } +} + +void MlaPrologTilingCheck::FillScenarioParamInfo() +{ + using FillFunc = void (MlaPrologTilingCheck::*)(); + static const std::unordered_map dispatchTable = { + {QUANT_MODE::NO_QUANT, &MlaPrologTilingCheck::FillNonQuantParamInfo}, + {QUANT_MODE::PARTIAL_QUANT_KV_NO_QUANT, &MlaPrologTilingCheck::FillPartialQuantParamInfo}, + {QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_CHANNEL, &MlaPrologTilingCheck::FillPartialKVQuantParamInfo}, + {QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_TILE, &MlaPrologTilingCheck::FillPartialKVPertileQuantParamInfo}, + {QUANT_MODE::FULL_QUANT_KV_NO_QUANT, &MlaPrologTilingCheck::FillFullQuantParamInfo}, + {QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TENSOR, &MlaPrologTilingCheck::FillFullKVQuantParamInfo}, + {QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TILE, &MlaPrologTilingCheck::FillFullKVPertileQuantParamInfo}, + {QUANT_MODE::MXFP8_FULL_QUANT_KV_NO_QUANT, &MlaPrologTilingCheck::FillMxfp8FullQuantParamInfo}, + {QUANT_MODE::MXFP8_FULL_QUANT_KV_QUANT_PER_TENSOR, &MlaPrologTilingCheck::FillMxfp8FullKVQuantParamInfo}, + {QUANT_MODE::MXFP8_FULL_QUANT_KV_QUANT_PER_TILE, &MlaPrologTilingCheck::FillMxfp8FullKVPertileParamInfo}, + {QUANT_MODE::FP8_FULL_QUANT_KV_NO_QUANT, &MlaPrologTilingCheck::FillFP8FullQuantParamInfo}, + {QUANT_MODE::FP8_FULL_QUANT_KV_QUANT_PER_TENSOR, &MlaPrologTilingCheck::FillFP8FullKVQuantParamInfo}, + {QUANT_MODE::HIF8_FULL_QUANT_KV_NO_QUANT, &MlaPrologTilingCheck::FillHIF8FullQuantParamInfo}, + {QUANT_MODE::HIF8_FULL_QUANT_KV_QUANT_PER_TENSOR, &MlaPrologTilingCheck::FillHIF8FullKVQuantParamInfo}, + {QUANT_MODE::FP8_FULL_QUANT_KV_QUANT_PER_TILE, &MlaPrologTilingCheck::FillFP8FullKVPertileQuantParamInfo}, + {QUANT_MODE::HIF8_FULL_QUANT_KV_QUANT_PER_TILE, &MlaPrologTilingCheck::FillHIF8FullKVPertileQuantParamInfo} + }; + auto it = dispatchTable.find(scenarioInfo_.quantMode_); + if (it != dispatchTable.end()) { + (this->*(it->second))(); + } +} + +void MlaPrologTilingCheck::FillNonQuantParamInfo() +{ + expectedParamInfo_[TOKEN_X_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[WEIGHT_DQ_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[WEIGHT_UQ_QR_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[WEIGHT_UK_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[WEIGHT_DKV_KR_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[RMSNORM_GAMMA_CQ_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[RMSNORM_GAMMA_CKV_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[ROPE_SIN_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[ROPE_COS_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[KV_CACHE_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[KR_CACHE_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[QUERY_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[QUERY_ROPE_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[KV_CACHE_OUT_NAME].dtype = ge::DT_BF16; + expectedParamInfo_[KR_CACHE_OUT_NAME].dtype = ge::DT_BF16; +} + +void MlaPrologTilingCheck::FillPartialQuantParamInfo() +{ + FillNonQuantParamInfo(); + + expectedParamInfo_.emplace(DEQUANT_SCALE_W_UQ_QR_NAME, std::vector{1, baseShapeInfo_.headSizeUqQr}); + expectedParamInfo_.emplace(SMOOTH_SCALES_CQ_NAME, std::vector{1, baseShapeInfo_.hcqSize}); + + expectedParamInfo_[WEIGHT_UQ_QR_NAME].dtype = ge::DT_INT8; + + expectedParamInfo_[DEQUANT_SCALE_W_UQ_QR_NAME].dtype = ge::DT_FLOAT; + expectedParamInfo_[SMOOTH_SCALES_CQ_NAME].dtype = ge::DT_FLOAT; + + expectedParamInfo_[SMOOTH_SCALES_CQ_NAME].isValid = (context_.smoothScalesCq.desc != nullptr); +} + +void MlaPrologTilingCheck::FillPartialKVQuantParamInfo() +{ + FillPartialQuantParamInfo(); + + expectedParamInfo_.emplace(QUANT_SCALE_CKV_NAME, std::vector{1, baseShapeInfo_.hckvSize}); + expectedParamInfo_.emplace(QUANT_SCALE_CKR_NAME, std::vector{1, baseShapeInfo_.drSize}); + + expectedParamInfo_[KV_CACHE_NAME].dtype = ge::DT_INT8; + expectedParamInfo_[KR_CACHE_NAME].dtype = ge::DT_INT8; + expectedParamInfo_[KV_CACHE_OUT_NAME].dtype = ge::DT_INT8; + expectedParamInfo_[KR_CACHE_OUT_NAME].dtype = ge::DT_INT8; + + expectedParamInfo_[QUANT_SCALE_CKV_NAME].dtype = ge::DT_FLOAT; + expectedParamInfo_[QUANT_SCALE_CKR_NAME].dtype = ge::DT_FLOAT; +} + +void MlaPrologTilingCheck::FillPartialKVPertileQuantParamInfo() +{ + FillPartialQuantParamInfo(); + + expectedParamInfo_.emplace(K_NOPE_CLIP_ALPHA_NAME, std::vector{1}); + expectedParamInfo_[KV_CACHE_NAME].dtype = ge::DT_INT8; + expectedParamInfo_[KV_CACHE_OUT_NAME].dtype = ge::DT_INT8; + expectedParamInfo_[K_NOPE_CLIP_ALPHA_NAME].dtype = ge::DT_FLOAT; +} + +void MlaPrologTilingCheck::FillFullQuantParamInfo() +{ + FillPartialQuantParamInfo(); + + expectedParamInfo_.emplace(DEQUANT_SCALE_X_NAME, std::vector{baseShapeInfo_.tSize, 1}); + expectedParamInfo_.emplace(DEQUANT_SCALE_W_DQ_NAME, std::vector{1, baseShapeInfo_.hcqSize}); + expectedParamInfo_.emplace(DEQUANT_SCALE_W_DKV_KR_NAME, + std::vector{1, baseShapeInfo_.hckvSize + baseShapeInfo_.drSize}); + + expectedParamInfo_[TOKEN_X_NAME].dtype = ge::DT_INT8; + expectedParamInfo_[WEIGHT_DQ_NAME].dtype = ge::DT_INT8; + expectedParamInfo_[WEIGHT_DKV_KR_NAME].dtype = ge::DT_INT8; + + if (GetCurNpuArch() == NpuArch::DAV_3510 && std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + if (*(context_.weightQuantMode) == static_cast(WEIGHT_QUANT_MODE::FP8_FULL_QUANT)) { + expectedParamInfo_[TOKEN_X_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[WEIGHT_DQ_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[WEIGHT_UQ_QR_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[WEIGHT_DKV_KR_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + } else if (*(context_.weightQuantMode) == static_cast(WEIGHT_QUANT_MODE::HIF8_FULL_QUANT)) { + expectedParamInfo_[TOKEN_X_NAME].dtype = ge::DT_HIFLOAT8; + expectedParamInfo_[WEIGHT_DQ_NAME].dtype = ge::DT_HIFLOAT8; + expectedParamInfo_[WEIGHT_UQ_QR_NAME].dtype = ge::DT_HIFLOAT8; + expectedParamInfo_[WEIGHT_DKV_KR_NAME].dtype = ge::DT_HIFLOAT8; + } + } + + expectedParamInfo_[DEQUANT_SCALE_X_NAME].dtype = ge::DT_FLOAT; + expectedParamInfo_[DEQUANT_SCALE_W_DQ_NAME].dtype = ge::DT_FLOAT; + expectedParamInfo_[DEQUANT_SCALE_W_DKV_KR_NAME].dtype = ge::DT_FLOAT; +} + +void MlaPrologTilingCheck::FillFullKVQuantParamInfo() +{ + FillFullQuantParamInfo(); + + expectedParamInfo_[QUERY_NAME].dtype = ge::DT_INT8; + expectedParamInfo_[KV_CACHE_NAME].dtype = ge::DT_INT8; + expectedParamInfo_[KV_CACHE_OUT_NAME].dtype = ge::DT_INT8; + if (GetCurNpuArch() == NpuArch::DAV_3510 && std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + if (*(context_.weightQuantMode) == static_cast(WEIGHT_QUANT_MODE::FP8_FULL_QUANT)) { + expectedParamInfo_[QUERY_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[KV_CACHE_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[KV_CACHE_OUT_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + } else if (*(context_.weightQuantMode) == static_cast(WEIGHT_QUANT_MODE::HIF8_FULL_QUANT)) { + expectedParamInfo_[QUERY_NAME].dtype = ge::DT_HIFLOAT8; + expectedParamInfo_[KV_CACHE_NAME].dtype = ge::DT_HIFLOAT8; + expectedParamInfo_[KV_CACHE_OUT_NAME].dtype = ge::DT_HIFLOAT8; + } + } + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + expectedParamInfo_.emplace(QUANT_SCALE_CKV_NAME, std::vector{1}); + } else { + expectedParamInfo_.emplace(QUANT_SCALE_CKV_NAME, std::vector{1, baseShapeInfo_.hckvSize}); + } + + expectedParamInfo_[QUANT_SCALE_CKV_NAME].dtype = ge::DT_FLOAT; +} + +void MlaPrologTilingCheck::FillFullKVPertileQuantParamInfo() +{ + FillFullQuantParamInfo(); + + expectedParamInfo_.emplace(K_NOPE_CLIP_ALPHA_NAME, std::vector{1}); + expectedParamInfo_[KV_CACHE_NAME].dtype = ge::DT_INT8; + expectedParamInfo_[KV_CACHE_OUT_NAME].dtype = ge::DT_INT8; + if (GetCurNpuArch() == NpuArch::DAV_3510 && std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0) { + if (*(context_.weightQuantMode) == static_cast(WEIGHT_QUANT_MODE::FP8_FULL_QUANT)) { + expectedParamInfo_[KV_CACHE_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[KV_CACHE_OUT_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + } else if (*(context_.weightQuantMode) == static_cast(WEIGHT_QUANT_MODE::HIF8_FULL_QUANT)) { + expectedParamInfo_[KV_CACHE_NAME].dtype = ge::DT_HIFLOAT8; + expectedParamInfo_[KV_CACHE_OUT_NAME].dtype = ge::DT_HIFLOAT8; + } + } + expectedParamInfo_[K_NOPE_CLIP_ALPHA_NAME].dtype = ge::DT_FLOAT; +} + +void MlaPrologTilingCheck::FillMxfp8FullQuantParamInfo() +{ + FillPartialQuantParamInfo(); + // dequantScaleX: (M, K / 32) [BS, He / 32] || [T, He / 32] + expectedParamInfo_.emplace(DEQUANT_SCALE_X_NAME, + std::vector{baseShapeInfo_.tSize, baseShapeInfo_.heSize / MXFP8_BLOCK_SIZE}); + // dequantScaleWDq: (N, K / 32) [Hcq, He / 32] + expectedParamInfo_.emplace(DEQUANT_SCALE_W_DQ_NAME, + std::vector{baseShapeInfo_.hcqSize, baseShapeInfo_.heSize / MXFP8_BLOCK_SIZE}); + // dequantScaleWDq: (N, K / 32) [Numhead * (D + Dr), Hcq / 32] + expectedParamInfo_.erase(DEQUANT_SCALE_W_UQ_QR_NAME); + expectedParamInfo_.emplace( + DEQUANT_SCALE_W_UQ_QR_NAME, + std::vector{baseShapeInfo_.headSizeUqQr, baseShapeInfo_.hcqSize / MXFP8_BLOCK_SIZE}); + // dequantScaleWDkvKr: (N, K / 32) [Hckv + Dr, He / 32] + expectedParamInfo_.emplace(DEQUANT_SCALE_W_DKV_KR_NAME, + std::vector{baseShapeInfo_.hckvSize + baseShapeInfo_.drSize, + baseShapeInfo_.heSize / MXFP8_BLOCK_SIZE}); + + expectedParamInfo_[TOKEN_X_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[WEIGHT_DQ_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[WEIGHT_UQ_QR_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[WEIGHT_DKV_KR_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[DEQUANT_SCALE_X_NAME].dtype = ge::DT_FLOAT8_E8M0; + expectedParamInfo_[DEQUANT_SCALE_W_DQ_NAME].dtype = ge::DT_FLOAT8_E8M0; + expectedParamInfo_[DEQUANT_SCALE_W_UQ_QR_NAME].dtype = ge::DT_FLOAT8_E8M0; + expectedParamInfo_[DEQUANT_SCALE_W_DKV_KR_NAME].dtype = ge::DT_FLOAT8_E8M0; + + expectedParamInfo_[QUANT_SCALE_CKR_NAME].isValid = false; + expectedParamInfo_[SMOOTH_SCALES_CQ_NAME].isValid = false; + expectedParamInfo_[K_NOPE_CLIP_ALPHA_NAME].isValid = false; +} + +void MlaPrologTilingCheck::FillMxfp8FullKVQuantParamInfo() +{ + FillMxfp8FullQuantParamInfo(); + + expectedParamInfo_.emplace(QUANT_SCALE_CKV_NAME, std::vector{1}); + + expectedParamInfo_[KV_CACHE_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[QUANT_SCALE_CKV_NAME].dtype = ge::DT_FLOAT; + expectedParamInfo_[QUERY_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[KV_CACHE_OUT_NAME].dtype = ge::DT_FLOAT8_E4M3FN; +} + +void MlaPrologTilingCheck::FillMxfp8FullKVPertileParamInfo() +{ + FillMxfp8FullQuantParamInfo(); + + expectedParamInfo_[KV_CACHE_NAME].dtype = ge::DT_FLOAT8_E4M3FN; + expectedParamInfo_[KV_CACHE_OUT_NAME].dtype = ge::DT_FLOAT8_E4M3FN; +} + +void MlaPrologTilingCheck::FillFP8FullQuantParamInfo() +{ + FillFullQuantParamInfo(); +} + +void MlaPrologTilingCheck::FillFP8FullKVQuantParamInfo() +{ + FillFullKVQuantParamInfo(); +} + +void MlaPrologTilingCheck::FillHIF8FullQuantParamInfo() +{ + FillFullQuantParamInfo(); +} + +void MlaPrologTilingCheck::FillHIF8FullKVQuantParamInfo() +{ + FillFullKVQuantParamInfo(); +} + +void MlaPrologTilingCheck::FillFP8FullKVPertileQuantParamInfo() +{ + FillFullKVPertileQuantParamInfo(); + expectedParamInfo_.erase(K_NOPE_CLIP_ALPHA_NAME); + expectedParamInfo_[K_NOPE_CLIP_ALPHA_NAME].isValid = false; +} + +void MlaPrologTilingCheck::FillHIF8FullKVPertileQuantParamInfo() +{ + FillFullKVPertileQuantParamInfo(); + expectedParamInfo_.erase(K_NOPE_CLIP_ALPHA_NAME); + expectedParamInfo_[K_NOPE_CLIP_ALPHA_NAME].isValid = false; +} + +void MlaPrologTilingCheck::GenActualParamInfo() +{ + actualParamInfo_.emplace(TOKEN_X_NAME, context_.tokenX); + actualParamInfo_.emplace(WEIGHT_DQ_NAME, context_.weightDq); + actualParamInfo_.emplace(WEIGHT_UQ_QR_NAME, context_.weightUqQr); + actualParamInfo_.emplace(WEIGHT_UK_NAME, context_.weightUk); + actualParamInfo_.emplace(WEIGHT_DKV_KR_NAME, context_.weightDkvKr); + actualParamInfo_.emplace(RMSNORM_GAMMA_CQ_NAME, context_.rmsnormGammaCq); + actualParamInfo_.emplace(RMSNORM_GAMMA_CKV_NAME, context_.rmsnormGammaCkv); + actualParamInfo_.emplace(ROPE_SIN_NAME, context_.ropeSin); + actualParamInfo_.emplace(ROPE_COS_NAME, context_.ropeCos); + actualParamInfo_.emplace(CACHE_INDEX_NAME, context_.cacheIndex); + actualParamInfo_.emplace(KV_CACHE_NAME, context_.kvCache); + actualParamInfo_.emplace(KR_CACHE_NAME, context_.krCache); + actualParamInfo_.emplace(DEQUANT_SCALE_X_NAME, context_.dequantScaleX); + actualParamInfo_.emplace(DEQUANT_SCALE_W_DQ_NAME, context_.dequantScaleWDq); + actualParamInfo_.emplace(DEQUANT_SCALE_W_UQ_QR_NAME, context_.dequantScaleWUqQr); + actualParamInfo_.emplace(DEQUANT_SCALE_W_DKV_KR_NAME, context_.dequantScaleWDkvKr); + actualParamInfo_.emplace(QUANT_SCALE_CKV_NAME, context_.quantScaleCkv); + actualParamInfo_.emplace(QUANT_SCALE_CKR_NAME, context_.quantScaleCkr); + actualParamInfo_.emplace(SMOOTH_SCALES_CQ_NAME, context_.smoothScalesCq); + actualParamInfo_.emplace(ACTUAL_SEQ_LEN_NAME, context_.actualSeqLen); + actualParamInfo_.emplace(K_NOPE_CLIP_ALPHA_NAME, context_.kNopeClipAlpha); + actualParamInfo_.emplace(QUERY_NAME, context_.query); + actualParamInfo_.emplace(QUERY_ROPE_NAME, context_.queryRope); + actualParamInfo_.emplace(KV_CACHE_OUT_NAME, context_.kvCacheOut); + actualParamInfo_.emplace(KR_CACHE_OUT_NAME, context_.krCacheOut); + actualParamInfo_.emplace(DEQUANT_SCALE_Q_NOPE_NAME, context_.dequantScaleQNope); + actualParamInfo_.emplace(QUERY_NORM_NAME, context_.queryNorm); + actualParamInfo_.emplace(DEQUANT_SCALE_Q_NORM_NAME, context_.dequantScaleQNorm); + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0 && + *(context_.ckvkrRepoMode) == static_cast(CKVKR_REPO_MODE::COMBINE)) { + actualParamInfo_.erase(KR_CACHE_NAME); + actualParamInfo_.erase(KR_CACHE_OUT_NAME); + } +} + +ge::graphStatus MlaPrologTilingCheck::CheckCkvkrRepoMode() +{ + ge::graphStatus isCorrect{ge::GRAPH_SUCCESS}; + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) != 0) { + return isCorrect; + } + if (*(context_.ckvkrRepoMode) == static_cast(CKVKR_REPO_MODE::COMBINE)) { + // 校验所有维度的乘积是否为0 + if (context_.krCache.shape->GetStorageShape().GetShapeSize() != 0) { + isCorrect = ge::GRAPH_FAILED; + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context_.opName, "krCache", + GetShapeStr(context_.krCache.shape->GetStorageShape()), + "When ckvkrRepoMode is COMBINE, krCache should be empty tensor"); + } + if (context_.krCacheOut.shape->GetStorageShape().GetShapeSize() != 0) { + isCorrect = ge::GRAPH_FAILED; + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context_.opName, "krCacheOut", + GetShapeStr(context_.krCacheOut.shape->GetStorageShape()), + "When ckvkrRepoMode is COMBINE, krCacheOut should be empty tensor"); + } + } + return isCorrect; +} + +ge::graphStatus MlaPrologTilingCheck::CheckCacheIndexDim() +{ + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) != 0) { + return ge::GRAPH_SUCCESS; + } + if (!scenarioInfo_.batchSeqFusedFlag_) { + return ge::GRAPH_SUCCESS; + } + if (scenarioInfo_.cacheMode_ != CACHE_MODE::PA_BLK_BSND && scenarioInfo_.cacheMode_ != CACHE_MODE::PA_BLK_NZ) { + return ge::GRAPH_SUCCESS; + } + OP_CHECK_IF(context_.cacheIndex.shape == nullptr, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON( + context_.opName, "cacheIndex", "null", + "When cacheMode is in {PA_BLK_BSND, PA_BLK_NZ}, cacheIndex should not be null"), + return ge::GRAPH_FAILED); + + OP_CHECK_IF(context_.cacheIndex.shape->GetStorageShape().GetDimNum() != MLA_PROLOG_DIM_NUM_1, + OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON( + context_.opName, "cacheIndex", + std::to_string(context_.cacheIndex.shape->GetStorageShape().GetDimNum()) + "D", + "When cacheMode in {PA_BLK_BSND, PA_BLK_NZ} and tokenX dim is 2, cacheIndex dim should be 1"), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTilingCheck::CheckSpecialScenarioParamShape() +{ + if (CheckCkvkrRepoMode() == ge::GRAPH_FAILED) { + return ge::GRAPH_FAILED; + } + if (CheckCacheIndexDim() == ge::GRAPH_FAILED) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus MlaPrologTilingCheck::CheckParamByScenario() +{ + GenActualParamInfo(); + GenExpectedParamInfo(); + ge::graphStatus isCorrect{ge::GRAPH_SUCCESS}; + for (const auto &it : actualParamInfo_) { + const auto &expectedParam{expectedParamInfo_[it.first]}; + if (__builtin_expect((expectedParam != it.second), 0)) { + isCorrect = ge::GRAPH_FAILED; + if (expectedParam.isValid != it.second.isValid) { + if (expectedParam.isValid && !it.second.isValid) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context_.opName, it.first, "null", + "this parameter is required under current configuration"); + } else { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context_.opName, it.first, "not null", + "this parameter is not required under current configuration"); + } + continue; + } + if (expectedParam.dtype != it.second.dtype) { + OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON( + context_.opName, it.first, TypeUtils::DataTypeToSerialString(it.second.dtype), + "this parameter requires dtype " + TypeUtils::DataTypeToSerialString(expectedParam.dtype) + + " under current configuration"); + } + if (expectedParam.format != it.second.format) { + OP_LOGE_FOR_INVALID_FORMATS_WITH_REASON( + context_.opName, it.first, std::string(ge::GetFormatName(it.second.format)), + "this parameter requires format " + std::string(ge::GetFormatName(expectedParam.format)) + + " under current configuration"); + } + if (expectedParam.shape != it.second.shape) { + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + context_.opName, it.first, ConvertContainerToStringV3(it.second.shape), + "this parameter requires shape " + ConvertContainerToString(expectedParam.shape) + + " under current configuration"); + } + } + } + return isCorrect; +} + +ge::graphStatus MlaPrologTilingCheck::CheckScenarParam() +{ + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) != 0) { + return ge::GRAPH_SUCCESS; + } + + ge::graphStatus isCorrect{ge::GRAPH_SUCCESS}; + CheckRepoMode(scenarioInfo_.kvQuantMode_ == KV_QUANT_MODE::PER_TILE, isCorrect); + CheckQueryQuantMode(scenarioInfo_.kvQuantMode_ == KV_QUANT_MODE::PER_TENSOR, isCorrect); + return isCorrect; +} + +void MlaPrologTilingCheck::CheckRepoMode(bool isPertile, ge::graphStatus &isCorrect) +{ + auto expectedCkvkr = isPertile ? CKVKR_REPO_MODE::COMBINE : CKVKR_REPO_MODE::DIVIDE; + auto expectedQuantScale = isPertile ? QUANT_SCALE_REPO_MODE::COMBINE : QUANT_SCALE_REPO_MODE::DIVIDE; + std::string desc = isPertile ? "pertile" : "non-pertile"; + std::string name = isPertile ? "COMBINE" : "DIVIDE"; + + if (*(context_.ckvkrRepoMode) != static_cast(expectedCkvkr)) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context_.opName, "ckvkrRepoMode", + std::to_string(*(context_.ckvkrRepoMode)), + "When " + desc + " quant mode, must be " + name + "(" + + std::to_string(static_cast(expectedCkvkr)) + ")"); + isCorrect = ge::GRAPH_FAILED; + } + if (*(context_.quantScaleRepoMode) != static_cast(expectedQuantScale)) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context_.opName, "quantScaleRepoMode", + std::to_string(*(context_.quantScaleRepoMode)), + "When " + desc + " quant mode, must be " + name + "(" + + std::to_string(static_cast(expectedQuantScale)) + ")"); + isCorrect = ge::GRAPH_FAILED; + } +} + +void MlaPrologTilingCheck::CheckQueryQuantMode(bool isPertensor, ge::graphStatus &isCorrect) +{ + auto expected = isPertensor ? QUERY_QUANT_MODE::PER_TOKEN_HEAD : QUERY_QUANT_MODE::NO_QUANT; + std::string desc = isPertensor ? "pertensor" : "non-pertensor"; + std::string name = isPertensor ? "PER_TOKEN_HEAD" : "NO_QUANT"; + + if (*(context_.queryQuantMode) != static_cast(expected)) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context_.opName, "queryQuantMode", + std::to_string(*(context_.queryQuantMode)), + "When " + desc + " quant mode, must be " + name + "(" + + std::to_string(static_cast(expected)) + ")"); + isCorrect = ge::GRAPH_FAILED; + } +} +// =================================全量参数校验================================= + +// ==================================单参数校验================================== +bool MlaPrologTilingCheck::IsSingleParamValid(const BaseParaInfo ¶m, const std::string ¶mName, + const std::set &expectedDtype, + const std::set &expectedFormat, + const std::set &expectedDimNum) const +{ + OP_CHECK_IF((param.shape == nullptr) || (param.desc == nullptr), + OP_LOGE_WITH_INVALID_INPUT(context_.opName, paramName), return false); + + ge::DataType dtype = param.desc->GetDataType(); + OP_CHECK_IF((expectedDtype.find(dtype) == expectedDtype.end()), + OP_LOGE_FOR_INVALID_DTYPE(context_.opName, paramName, TypeUtils::DataTypeToSerialString(dtype), + ConvertContainerToStringV3(expectedDtype, TypeUtils::DataTypeToSerialString)), + return false); + + ge::Format format = static_cast(ge::GetPrimaryFormat(param.desc->GetStorageFormat())); + OP_CHECK_IF((expectedFormat.find(format) == expectedFormat.end()), + OP_LOGE_FOR_INVALID_FORMAT(context_.opName, paramName, std::string(ge::GetFormatName(format)), + ConvertContainerToStringV3(expectedFormat, FormatToString)), + return false); + + size_t dimNum = param.shape->GetStorageShape().GetDimNum(); + OP_CHECK_IF((expectedDimNum.find(dimNum) == expectedDimNum.end()), + OP_LOGE_FOR_INVALID_SHAPEDIM(context_.opName, paramName, std::to_string(dimNum), + ConvertContainerToStringV3(expectedDimNum)), + return false); + return true; +} + +ge::graphStatus MlaPrologTilingCheck::CheckSingleRequiredParam() const +{ + if (!CheckTokenX() || !CheckWDq() || !CheckWUqQr() || !CheckWUk() || !CheckWDkvKr() || !CheckRmsnormGammaCq() || + !CheckRmsnormGammaCkv() || !CheckRopeSin() || !CheckRopeCos() || !CheckCacheIndex() || !CheckKvCache() || + !CheckKrCache() || !CheckActSeqLen()) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} + +bool MlaPrologTilingCheck::CheckTokenX() const +{ + if (GetCurNpuArch() == NpuArch::DAV_3510) { + return IsSingleParamValid(context_.tokenX, TOKEN_X_NAME, + {ge::DT_BF16, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8}, + {ge::FORMAT_ND, ge::FORMAT_NCHW}, {2, 3}); + } else { + return IsSingleParamValid(context_.tokenX, TOKEN_X_NAME, {ge::DT_BF16, ge::DT_INT8}, + {ge::FORMAT_ND, ge::FORMAT_NCHW}, {2, 3}); + } +} + +bool MlaPrologTilingCheck::CheckWDq() const +{ + if (GetCurNpuArch() == NpuArch::DAV_3510) { + return IsSingleParamValid(context_.weightDq, WEIGHT_DQ_NAME, + {ge::DT_BF16, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8}, + {ge::FORMAT_FRACTAL_NZ}, {2, 4}); + } else { + return IsSingleParamValid(context_.weightDq, WEIGHT_DQ_NAME, {ge::DT_BF16, ge::DT_INT8}, + {ge::FORMAT_FRACTAL_NZ}, {2, 4}); + } +} + +bool MlaPrologTilingCheck::CheckWUqQr() const +{ + if (GetCurNpuArch() == NpuArch::DAV_3510) { + return IsSingleParamValid(context_.weightUqQr, WEIGHT_UQ_QR_NAME, + {ge::DT_BF16, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8}, + {ge::FORMAT_FRACTAL_NZ}, {2, 4}); + } else { + return IsSingleParamValid(context_.weightUqQr, WEIGHT_UQ_QR_NAME, {ge::DT_BF16, ge::DT_INT8}, + {ge::FORMAT_FRACTAL_NZ}, {2, 4}); + } +} + +bool MlaPrologTilingCheck::CheckWUk() const +{ + if (GetCurNpuArch() == NpuArch::DAV_3510) { + return IsSingleParamValid(context_.weightUk, WEIGHT_UK_NAME, {ge::DT_BF16}, {ge::FORMAT_ND, ge::FORMAT_NCHW}, + {3}); + } else { + return IsSingleParamValid(context_.weightUk, WEIGHT_UK_NAME, {ge::DT_BF16, ge::DT_INT8}, + {ge::FORMAT_ND, ge::FORMAT_NCHW}, {3}); + } +} + +bool MlaPrologTilingCheck::CheckWDkvKr() const +{ + if (GetCurNpuArch() == NpuArch::DAV_3510) { + return IsSingleParamValid(context_.weightDkvKr, WEIGHT_DKV_KR_NAME, + {ge::DT_BF16, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8}, + {ge::FORMAT_FRACTAL_NZ}, {2, 4}); + } else { + return IsSingleParamValid(context_.weightDkvKr, WEIGHT_DKV_KR_NAME, {ge::DT_BF16, ge::DT_INT8}, + {ge::FORMAT_FRACTAL_NZ}, {2, 4}); + } +} + +bool MlaPrologTilingCheck::CheckRmsnormGammaCq() const +{ + return IsSingleParamValid(context_.rmsnormGammaCq, RMSNORM_GAMMA_CQ_NAME, {ge::DT_BF16}, + {ge::FORMAT_ND, ge::FORMAT_NCHW}, {1}); +} + +bool MlaPrologTilingCheck::CheckRmsnormGammaCkv() const +{ + return IsSingleParamValid(context_.rmsnormGammaCkv, RMSNORM_GAMMA_CKV_NAME, {ge::DT_BF16}, + {ge::FORMAT_ND, ge::FORMAT_NCHW}, {1}); +} + +bool MlaPrologTilingCheck::CheckRopeSin() const +{ + // V3: rope may be null when do_rope=false. + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0 && context_.doRope != nullptr && + !*(context_.doRope)) { + OP_CHECK_IF(context_.ropeSin.shape != nullptr || context_.ropeSin.desc != nullptr, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context_.opName, "ropeSin", "non-null", + "ropeSin must be null when do_rope is false"), + return false); + return true; + } + return IsSingleParamValid(context_.ropeSin, ROPE_SIN_NAME, {ge::DT_BF16}, {ge::FORMAT_ND, ge::FORMAT_NCHW}, {2, 3}); +} + +bool MlaPrologTilingCheck::CheckRopeCos() const +{ + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0 && context_.doRope != nullptr && + !*(context_.doRope)) { + OP_CHECK_IF(context_.ropeCos.shape != nullptr || context_.ropeCos.desc != nullptr, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context_.opName, "ropeCos", "non-null", + "ropeCos must be null when do_rope is false"), + return false); + return true; + } + return IsSingleParamValid(context_.ropeCos, ROPE_COS_NAME, {ge::DT_BF16}, {ge::FORMAT_ND, ge::FORMAT_NCHW}, {2, 3}); +} + +bool MlaPrologTilingCheck::CheckCacheIndex() const +{ + return std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0 || + IsSingleParamValid(context_.cacheIndex, CACHE_INDEX_NAME, {ge::DT_INT64}, {ge::FORMAT_ND, ge::FORMAT_NCHW}, + {1, 2}); +} + +bool MlaPrologTilingCheck::CheckKvCache() const +{ + if (GetCurNpuArch() == NpuArch::DAV_3510) { + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) != 0) { + return IsSingleParamValid(context_.kvCache, KV_CACHE_NAME, {ge::DT_BF16, ge::DT_INT8}, + {ge::FORMAT_ND, ge::FORMAT_NCHW}, {4}); + } else { + return IsSingleParamValid(context_.kvCache, KV_CACHE_NAME, + {ge::DT_BF16, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8}, + {ge::FORMAT_ND, ge::FORMAT_NCHW}, {3, 4}); + } + } else { + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) != 0) { + return IsSingleParamValid(context_.kvCache, KV_CACHE_NAME, {ge::DT_BF16, ge::DT_INT8}, + {ge::FORMAT_ND, ge::FORMAT_NCHW}, {4}); + } else { + return IsSingleParamValid(context_.kvCache, KV_CACHE_NAME, {ge::DT_BF16, ge::DT_INT8}, + {ge::FORMAT_ND, ge::FORMAT_NCHW}, {3, 4}); + } + } +} + +bool MlaPrologTilingCheck::CheckKrCache() const +{ + if (GetCurNpuArch() == NpuArch::DAV_3510) { + return IsSingleParamValid(context_.krCache, KR_CACHE_NAME, {ge::DT_BF16, ge::DT_INT8}, + {ge::FORMAT_ND, ge::FORMAT_NCHW}, {1, 3, 4}); + } else { + return (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) == 0 && + *(context_.ckvkrRepoMode) == static_cast(CKVKR_REPO_MODE::COMBINE)) || + IsSingleParamValid(context_.krCache, KR_CACHE_NAME, {ge::DT_BF16, ge::DT_INT8}, + {ge::FORMAT_ND, ge::FORMAT_NCHW}, {1, 3, 4}); + } +} + +bool MlaPrologTilingCheck::CheckActSeqLen() const +{ + if (context_.actualSeqLen.desc == nullptr) { + return true; + }; + ge::DataType dtype = context_.actualSeqLen.desc->GetDataType(); + OP_CHECK_IF((ge::DT_INT32 != dtype), + OP_LOGE_FOR_INVALID_DTYPE(context_.opName, "actualSeqLen", TypeUtils::DataTypeToSerialString(dtype), + TypeUtils::DataTypeToSerialString(ge::DT_INT32)), + return false); + return true; +} + +bool MlaPrologTilingCheck::CheckCacheModeParamShape() const +{ + if (std::strncmp(context_.cacheMode, CACHE_MODE_TND, CACHE_MODE_LEN) == 0) { + OP_CHECK_IF(context_.tokenX.shape->GetStorageShape().GetDimNum() != MLA_PROLOG_DIM_NUM_2, + OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON( + context_.opName, "tokenX", + std::to_string(context_.tokenX.shape->GetStorageShape().GetDimNum()) + "D", + "When cacheMode is TND, tokenX dim must be 2"), + return false); + OP_CHECK_IF(context_.kvCache.shape->GetStorageShape().GetDimNum() != MLA_PROLOG_DIM_NUM_3, + OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON( + context_.opName, "kvCache", + std::to_string(context_.kvCache.shape->GetStorageShape().GetDimNum()) + "D", + "When cacheMode is TND, kvCache dim must be 3"), + return false); + } else if (std::strncmp(context_.cacheMode, CACHE_MODE_BSND, CACHE_MODE_LEN) == 0) { + OP_CHECK_IF(context_.tokenX.shape->GetStorageShape().GetDimNum() != MLA_PROLOG_DIM_NUM_3, + OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON( + context_.opName, "tokenX", + std::to_string(context_.tokenX.shape->GetStorageShape().GetDimNum()) + "D", + "When cacheMode is BSND, tokenX dim must be 3"), + return false); + OP_CHECK_IF(context_.kvCache.shape->GetStorageShape().GetDimNum() != MLA_PROLOG_DIM_NUM_4, + OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON( + context_.opName, "kvCache", + std::to_string(context_.kvCache.shape->GetStorageShape().GetDimNum()) + "D", + "When cacheMode is BSND, kvCache dim must be 4"), + return false); + } else { + OP_CHECK_IF(context_.kvCache.shape->GetStorageShape().GetDimNum() != MLA_PROLOG_DIM_NUM_4, + OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON( + context_.opName, "kvCache", + std::to_string(context_.kvCache.shape->GetStorageShape().GetDimNum()) + "D", + "When cacheMode in {PA_BSND, PA_NZ, PA_BLK_BSND, PA_BLK_NZ}, kvCache dim must be 4"), + return false); + } + return true; +} + +ge::graphStatus MlaPrologTilingCheck::CheckCacheMode() const +{ + if (std::strncmp(context_.opType, V3_OP_NAME, OP_NAME_LEN) != 0) { + if (std::strncmp(context_.cacheMode, CACHE_MODE_PA_BSND, CACHE_MODE_LEN) != 0 && + std::strncmp(context_.cacheMode, CACHE_MODE_PA_NZ, CACHE_MODE_LEN) != 0) { + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "cacheMode", std::string(context_.cacheMode), + "{PA_BSND, PA_NZ}"); + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; + } + + if ((std::strncmp(context_.cacheMode, CACHE_MODE_BSND, CACHE_MODE_LEN) != 0) && + (std::strncmp(context_.cacheMode, CACHE_MODE_TND, CACHE_MODE_LEN) != 0) && + (std::strncmp(context_.cacheMode, CACHE_MODE_PA_BSND, CACHE_MODE_LEN) != 0) && + (std::strncmp(context_.cacheMode, CACHE_MODE_PA_NZ, CACHE_MODE_LEN) != 0) && + (std::strncmp(context_.cacheMode, CACHE_MODE_PA_BLK_BSND, CACHE_MODE_LEN) != 0) && + (std::strncmp(context_.cacheMode, CACHE_MODE_PA_BLK_NZ, CACHE_MODE_LEN) != 0)) { + OP_LOGE_FOR_INVALID_VALUE(context_.opName, "cacheMode", std::string(context_.cacheMode), + "{BSND, TND, PA_BSND, PA_NZ, PA_BLK_BSND, PA_BLK_NZ}"); + return ge::GRAPH_FAILED; + } + if (!CheckCacheModeParamShape()) { + return ge::GRAPH_FAILED; + } + + if (*(context_.kvQuantMode) != static_cast(KV_QUANT_MODE::PER_TILE)) { + return ge::GRAPH_SUCCESS; + } + + if ((std::strncmp(context_.cacheMode, CACHE_MODE_PA_NZ, CACHE_MODE_LEN) == 0) || + (std::strncmp(context_.cacheMode, CACHE_MODE_PA_BLK_BSND, CACHE_MODE_LEN) == 0) || + (std::strncmp(context_.cacheMode, CACHE_MODE_PA_BLK_NZ, CACHE_MODE_LEN) == 0)) { + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON( + context_.opName, "cacheMode", std::string(context_.cacheMode), + "When pertile quant mode, cacheMode cannot be {PA_NZ, PA_BLK_BSND, PA_BLK_NZ}"); + return ge::GRAPH_FAILED; + } + + return ge::GRAPH_SUCCESS; +} + +// ==================================单参数校验================================== +} // namespace optiling \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling_check.h b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling_check.h new file mode 100644 index 000000000000..82b6feedae47 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling_check.h @@ -0,0 +1,215 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_tiling_check.h + * \brief + */ + +#ifndef MLA_PROLOG_TILING_CHECK_H +#define MLA_PROLOG_TILING_CHECK_H + +#include "mla_prolog_tiling.h" + +namespace optiling { + +constexpr uint32_t MAX_B_SIZE = 65536U; +constexpr uint32_t MAX_S1_SIZE = 65536U; +constexpr uint32_t MAX_T_SIZE = 1024U * 1024U; +constexpr uint32_t HCQ_SIZE = 1536U; +constexpr uint32_t HCKV_SIZE = 512U; +constexpr uint32_t D_SIZE = 128U; +constexpr uint32_t DR_SIZE = 64U; +constexpr uint32_t NKV_SIZE = 1U; +constexpr uint32_t MIN_BLOCK_SIZE = 16U; +constexpr uint32_t MAX_BLOCK_SIZE = 1024U; +constexpr uint32_t ALIGN_BLOCK_SIZE = 16U; +constexpr uint32_t MXFP8_BLOCK_SIZE = 32U; + +constexpr int64_t NZ_H0_SIZE = 16U; + +constexpr char TOKEN_X_NAME[] {"tokenX"}; +constexpr char WEIGHT_DQ_NAME[] {"weightDq"}; +constexpr char WEIGHT_UQ_QR_NAME[] {"weightUqQr"}; +constexpr char WEIGHT_UK_NAME[] {"weightUk"}; +constexpr char WEIGHT_DKV_KR_NAME[] {"weightDkvKr"}; +constexpr char RMSNORM_GAMMA_CQ_NAME[] {"rmsnormGammaCq"}; +constexpr char RMSNORM_GAMMA_CKV_NAME[] {"rmsnormGammaCkv"}; +constexpr char ROPE_SIN_NAME[] {"ropeSin"}; +constexpr char ROPE_COS_NAME[] {"ropeCos"}; +constexpr char CACHE_INDEX_NAME[] {"cacheIndex"}; +constexpr char KV_CACHE_NAME[] {"kvCache"}; +constexpr char KR_CACHE_NAME[] {"krCache"}; +constexpr char DEQUANT_SCALE_X_NAME[] {"dequantScaleX"}; +constexpr char DEQUANT_SCALE_W_DQ_NAME[] {"dequantScaleWDq"}; +constexpr char DEQUANT_SCALE_W_UQ_QR_NAME[] {"dequantScaleWUqQr"}; +constexpr char DEQUANT_SCALE_W_DKV_KR_NAME[] {"dequantScaleWDkvKr"}; +constexpr char QUANT_SCALE_CKV_NAME[] {"quantScaleCkv"}; +constexpr char QUANT_SCALE_CKR_NAME[] {"quantScaleCkr"}; +constexpr char SMOOTH_SCALES_CQ_NAME[] {"smoothScalesCq"}; +constexpr char ACTUAL_SEQ_LEN_NAME[] {"actualSeqLen"}; +constexpr char K_NOPE_CLIP_ALPHA_NAME[] {"kNopeClipAlpha"}; +constexpr char QUERY_NAME[] {"query"}; +constexpr char QUERY_ROPE_NAME[] {"queryRope"}; +constexpr char KV_CACHE_OUT_NAME[] {"kvCacheOut"}; +constexpr char KR_CACHE_OUT_NAME[] {"krCacheOut"}; +constexpr char DEQUANT_SCALE_Q_NOPE_NAME[] {"dequantScaleQNope"}; +constexpr char QUERY_NORM_NAME[] {"queryNorm"}; +constexpr char DEQUANT_SCALE_Q_NORM_NAME[] {"dequantScaleQNorm"}; + +constexpr uint32_t PARAM_MAP_INIT_RESERVE_NUM = 28; // 预分配所有key的个数,避免使用时动态扩容 + +struct ParamInfo { + ParamInfo() = default; + explicit ParamInfo(const ParamInfo &) = default; + ParamInfo &operator=(const ParamInfo &) = default; + explicit ParamInfo(ParamInfo &&) = default; + ParamInfo &operator=(ParamInfo &&other) = default; + ~ParamInfo() = default; + ParamInfo(const gert::CompileTimeTensorDesc *actualDesc, const gert::StorageShape *actualShape) { + if (actualDesc != nullptr && actualShape != nullptr) { + isValid = true; + dtype = actualDesc->GetDataType(); + format = static_cast(ge::GetPrimaryFormat(actualDesc->GetStorageFormat())); + auto &&actualStorageShape = actualShape->GetStorageShape(); + dimNum = actualStorageShape.GetDimNum(); + this->shape.reserve(dimNum); + for (size_t i = 0; i < dimNum; i++) { + this->shape.emplace_back(actualStorageShape.GetDim(i)); + } + } + } + explicit ParamInfo(const BaseParaInfo &info) : ParamInfo(info.desc, info.shape) {} + explicit ParamInfo(const std::vector &expectedShape) + { + isValid = true; + format = ge::FORMAT_ND; + dimNum = expectedShape.size(); + this->shape.reserve(dimNum); + for (size_t i = 0; i < dimNum; i++) { + this->shape.emplace_back(static_cast(expectedShape[i])); + } + } + + bool operator == (const ParamInfo &other) const { + if (!isValid && !other.isValid) { + return true; + } + static const std::set ndFormats{ge::FORMAT_ND, ge::FORMAT_NCHW}; + if ((ndFormats.find(format) == ndFormats.end() || ndFormats.find(other.format) == ndFormats.end()) && + format != other.format) { + return false; + } + return (isValid == other.isValid && dtype == other.dtype && + dimNum == other.dimNum && shape == other.shape); + } + bool operator != (const ParamInfo &other) const { + return !(*this == other); + } + + bool isValid {}; + ge::DataType dtype {ge::DT_MAX}; + ge::Format format {ge::FORMAT_MAX}; + size_t dimNum {}; + std::vector shape; +}; + +using ParamInfoMap = std::unordered_map; + +class MlaPrologTilingCheck { +public: + MlaPrologTilingCheck(const MlaPrologContext &context, const MlaPrologBaseShapeInfo &baseShapeInfo, + const MlaPrologScenarioInfo &scenarioInfo) + : context_(context), baseShapeInfo_(baseShapeInfo), scenarioInfo_(scenarioInfo) {} + ge::graphStatus CheckSingleRequiredParam() const; + ge::graphStatus CheckCacheMode() const; + ge::graphStatus CheckQuantMode() const; + ge::graphStatus CheckDims() const; + ge::graphStatus CheckParamByScenario(); + ge::graphStatus CheckSpecialScenarioParamShape(); + ge::graphStatus CheckCkvkrRepoMode(); + ge::graphStatus CheckCacheIndexDim(); + ge::graphStatus CheckScenarParam(); + ge::graphStatus CheckAttrs() const; + + NpuArch GetCurNpuArch() const; + +private: + bool CheckAttrsNotNull() const; + bool CheckAttrsRange() const; + bool CheckCacheModeParamShape() const; + ge::graphStatus CheckHcqSize() const; + ge::graphStatus CheckDSize() const; + ge::graphStatus CheckDtileSize() const; + // ==================================单参数校验================================== + bool IsSingleParamValid(const BaseParaInfo ¶m, const std::string ¶mName, + const std::set &expectedDtype, + const std::set &expectedFormat, + const std::set &expectedDimNum) const; + bool CheckTokenX() const; + bool CheckWDq() const; + bool CheckWDkvKr() const; + bool CheckWUqQr() const; + bool CheckWUk() const; + bool CheckRmsnormGammaCkv() const; + bool CheckRmsnormGammaCq() const; + bool CheckRopeCos() const; + bool CheckRopeSin() const; + bool CheckCacheIndex() const; + bool CheckKvCache() const; + bool CheckKrCache() const; + bool CheckActSeqLen() const; + // ==================================单参数校验================================== + + // =================================全量参数校验================================= + void GenExpectedParamInfo(); + void FillCommonParamInfo(); + void FillRequiredParamShapeWithDims(); + void FillOptionalOutputParamShapeWithDims(); + void FillOptionalOutputParamShapeWithDimsV2(); + void FillOptionalOutputParamShapeWithDimsV3(); + void FillScenarioParamInfo(); + void FillQueryNormScaleShape(); + void FillQueryNormDtypes(); + void FillTokenAndQueryShapes(); + void FillWeightAndNormShapes(); + void FillCacheShapes(); + void CheckRepoMode(bool isPertile, ge::graphStatus &isCorrect); + void CheckQueryQuantMode(bool isPertensor, ge::graphStatus &isCorrect); + void FillNonQuantParamInfo(); + void FillPartialQuantParamInfo(); + void FillPartialKVQuantParamInfo(); + void FillPartialKVPertileQuantParamInfo(); + void FillFullQuantParamInfo(); + void FillFullKVQuantParamInfo(); + void FillFullKVPertileQuantParamInfo(); + void FillMxfp8FullQuantParamInfo(); + void FillMxfp8FullKVQuantParamInfo(); + void FillMxfp8FullKVPertileParamInfo(); + void FillFP8FullQuantParamInfo(); + void FillFP8FullKVQuantParamInfo(); + void FillHIF8FullQuantParamInfo(); + void FillHIF8FullKVQuantParamInfo(); + void FillFP8FullKVPertileQuantParamInfo(); + void FillHIF8FullKVPertileQuantParamInfo(); + + void GenActualParamInfo(); + // =================================全量参数校验================================= + + const MlaPrologContext &context_; + const MlaPrologBaseShapeInfo &baseShapeInfo_; + const MlaPrologScenarioInfo &scenarioInfo_; + ParamInfoMap expectedParamInfo_ = ParamInfoMap(PARAM_MAP_INIT_RESERVE_NUM); + ParamInfoMap actualParamInfo_ = ParamInfoMap(PARAM_MAP_INIT_RESERVE_NUM); +}; + +} // namespace optiling + +#endif \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_def.cpp b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_def.cpp new file mode 100644 index 000000000000..5b00fa6f34a7 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_def.cpp @@ -0,0 +1,309 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "register/op_def_registry.h" + +namespace ops { +class MlaPrologV3 : public OpDef { +public: + explicit MlaPrologV3(const char *name) : OpDef(name) + { + this->Input("token_x") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_BF16, ge::DT_INT8}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("weight_dq") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_BF16, ge::DT_INT8}) + .FormatList({ge::FORMAT_FRACTAL_NZ}) + .AutoContiguous(); + this->Input("weight_uq_qr") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8}) + .FormatList({ge::FORMAT_FRACTAL_NZ}) + .AutoContiguous(); + this->Input("weight_uk") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("weight_dkv_kr") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_BF16, ge::DT_INT8}) + .FormatList({ge::FORMAT_FRACTAL_NZ}) + .AutoContiguous(); + this->Input("rmsnorm_gamma_cq") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("rmsnorm_gamma_ckv") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("rope_sin") + .ParamType(OPTIONAL) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("rope_cos") + .ParamType(OPTIONAL) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("kv_cache") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_INT8, ge::DT_BF16, ge::DT_INT8, ge::DT_BF16, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("kr_cache") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_INT8, ge::DT_BF16, ge::DT_INT8, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("cache_index") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT64}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("dequant_scale_x") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("dequant_scale_w_dq") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("dequant_scale_w_uq_qr") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("dequant_scale_w_dkv_kr") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("quant_scale_ckv") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("quant_scale_ckr") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("smooth_scales_cq") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("actual_seq_len") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("k_nope_clip_alpha") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Output("query") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}); + this->Output("query_rope") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}); + this->Output("kv_cache") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_INT8, ge::DT_BF16, ge::DT_INT8, ge::DT_BF16, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8}) + .FormatList({ge::FORMAT_ND}); + this->Output("kr_cache") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_INT8, ge::DT_BF16, ge::DT_INT8, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}); + this->Output("dequant_scale_q_nope") + .ParamType(REQUIRED) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}); + this->Output("query_norm") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8}) + .FormatList({ge::FORMAT_ND}); + this->Output("dequant_scale_q_norm") + .ParamType(REQUIRED) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}); + this->Attr("rmsnorm_epsilon_cq").AttrType(OPTIONAL).Float(1e-05f); + this->Attr("rmsnorm_epsilon_ckv").AttrType(OPTIONAL).Float(1e-05f); + this->Attr("cache_mode").AttrType(OPTIONAL).String("PA_BSND"); + this->Attr("query_norm_flag").AttrType(OPTIONAL).Bool(false); + this->Attr("weight_quant_mode").AttrType(OPTIONAL).Int(0); + this->Attr("kv_cache_quant_mode").AttrType(OPTIONAL).Int(0); + this->Attr("query_quant_mode").AttrType(OPTIONAL).Int(0); + this->Attr("ckvkr_repo_mode").AttrType(OPTIONAL).Int(0); + this->Attr("quant_scale_repo_mode").AttrType(OPTIONAL).Int(0); + this->Attr("tile_size").AttrType(OPTIONAL).Int(128); // 128 : set value of tile size + this->Attr("qc_qr_scale").AttrType(OPTIONAL).Float(1.0f); + this->Attr("kc_scale").AttrType(OPTIONAL).Float(1.0f); + + OpAICoreConfig aicore_config_95; + aicore_config_95.Input("token_x") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("weight_dq") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN}) + .FormatList({ge::FORMAT_FRACTAL_NZ}) + .AutoContiguous(); + aicore_config_95.Input("weight_uq_qr") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN}) + .FormatList({ge::FORMAT_FRACTAL_NZ}) + .AutoContiguous(); + aicore_config_95.Input("weight_uk") + .ParamType(REQUIRED) + .DataTypeList({ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("weight_dkv_kr") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN}) + .FormatList({ge::FORMAT_FRACTAL_NZ}) + .AutoContiguous(); + aicore_config_95.Input("rmsnorm_gamma_cq") + .ParamType(REQUIRED) + .DataTypeList({ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("rmsnorm_gamma_ckv") + .ParamType(REQUIRED) + .DataTypeList({ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("rope_sin") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("rope_cos") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("kv_cache") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_INT8, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_BF16, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN}) + .FormatList({ge::FORMAT_ND}) + .IgnoreContiguous(); + aicore_config_95.Input("kr_cache") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .IgnoreContiguous(); + aicore_config_95.Input("cache_index") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT64}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("dequant_scale_x") + .ParamType(OPTIONAL) + .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("dequant_scale_w_dq") + .ParamType(OPTIONAL) + .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("dequant_scale_w_uq_qr") + .ParamType(OPTIONAL) + .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("dequant_scale_w_dkv_kr") + .ParamType(OPTIONAL) + .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("quant_scale_ckv") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("quant_scale_ckr") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("smooth_scales_cq") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("actual_seq_len") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT32}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Input("k_nope_clip_alpha") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + aicore_config_95.Output("query") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT8_E4M3FN, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}); + aicore_config_95.Output("query_rope") + .ParamType(REQUIRED) + .DataTypeList({ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}); + aicore_config_95.Output("kv_cache") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_INT8, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_BF16, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN}) + .FormatList({ge::FORMAT_ND}); + aicore_config_95.Output("kr_cache") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}); + aicore_config_95.Output("dequant_scale_q_nope") + .ParamType(REQUIRED) + .DataTypeList({ge::DT_FLOAT}) + .FormatList({ge::FORMAT_ND}); + aicore_config_95.Output("query_norm") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN}) + .FormatList({ge::FORMAT_ND}); + aicore_config_95.Output("dequant_scale_q_norm") + .ParamType(REQUIRED) + .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0}) + .FormatList({ge::FORMAT_ND}); + aicore_config_95.DynamicCompileStaticFlag(true) + .DynamicFormatFlag(true) + .DynamicRankSupportFlag(true) + .DynamicShapeSupportFlag(true) + .NeedCheckSupportFlag(false) + .PrecisionReduceFlag(true) + .ExtendCfgInfo("aclnnSupport.value", "support_aclnn"); + this->AICore().AddConfig("ascend950", aicore_config_95); + } +}; +OP_ADD(MlaPrologV3, optiling::MlaPrologCompileInfo); +} // namespace ops diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_infershape.cpp b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_infershape.cpp new file mode 100644 index 000000000000..b02cf007f149 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_infershape.cpp @@ -0,0 +1,285 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "mla_prolog_v3_infershape.h" + +using namespace ge; + +namespace ops { + +ge::graphStatus GetMlaPrologV3ShapeDim(const gert::InferShapeContext *context, MlaPrologProtoShapeParam &shapeParam) +{ + auto tokenXShape = context->GetRequiredInputShape(TOKEN_X_INDEX); // (B, S, He) | (T, He) + OP_CHECK_NULL_WITH_CONTEXT(context, tokenXShape); + auto weightUkShape = context->GetRequiredInputShape(WEIGHT_UK_INDEX); // (N, D, Hckv) + OP_CHECK_NULL_WITH_CONTEXT(context, weightUkShape); + auto ropeSinShape = context->GetRequiredInputShape(ROPE_SIN_INDEX); // (B, S, Dr) | (T, Dr) + OP_CHECK_NULL_WITH_CONTEXT(context, ropeSinShape); + auto weightDqShape = context->GetRequiredInputShape(WEIGHT_DQ_INDEX); // (He, Hcq) + OP_CHECK_NULL_WITH_CONTEXT(context, weightDqShape); + auto kvCacheShape = context->GetRequiredInputShape(KV_CACHE_INDEX_V3); // (B, Nkv, Skv, Hckv) + OP_CHECK_NULL_WITH_CONTEXT(context, kvCacheShape); + auto krCacheShape = context->GetRequiredInputShape(KR_CACHE_INDEX_V3); // (B, Nkv, Skv, Dr) + OP_CHECK_NULL_WITH_CONTEXT(context, krCacheShape); + + OP_CHECK_IF(((tokenXShape->GetDimNum() != DIM_NUM_2) && (tokenXShape->GetDimNum() != DIM_NUM_3)), + OP_LOGE_FOR_INVALID_SHAPEDIM(context->GetNodeName(), "tokenX", + std::to_string(tokenXShape->GetDimNum()) + "D", "2D or 3D"), + return ge::GRAPH_FAILED); + OP_CHECK_IF((weightUkShape->GetDimNum() != DIM_NUM_3), + OP_LOGE_FOR_INVALID_SHAPEDIM(context->GetNodeName(), "weightUk", + std::to_string(weightUkShape->GetDimNum()) + "D", "3D"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(((ropeSinShape->GetDimNum() != DIM_NUM_2) && (ropeSinShape->GetDimNum() != DIM_NUM_3)), + OP_LOGE_FOR_INVALID_SHAPEDIM(context->GetNodeName(), "ropeSin", + std::to_string(ropeSinShape->GetDimNum()) + "D", "2D or 3D"), + return ge::GRAPH_FAILED); + OP_CHECK_IF((weightDqShape->GetDimNum() != DIM_NUM_2), + OP_LOGE_FOR_INVALID_SHAPEDIM(context->GetNodeName(), "weightDq", + std::to_string(weightDqShape->GetDimNum()) + "D", "2D"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(((kvCacheShape->GetDimNum() != DIM_NUM_3) && (kvCacheShape->GetDimNum() != DIM_NUM_4)), + OP_LOGE_FOR_INVALID_SHAPEDIM(context->GetNodeName(), "kvCache", + std::to_string(kvCacheShape->GetDimNum()) + "D", "3D or 4D"), + return ge::GRAPH_FAILED); + OP_CHECK_IF(((krCacheShape->GetDimNum() != DIM_NUM_1) && (krCacheShape->GetDimNum() != DIM_NUM_3) && + (krCacheShape->GetDimNum() != DIM_NUM_4)), + OP_LOGE_FOR_INVALID_SHAPEDIM(context->GetNodeName(), "krCache", + std::to_string(krCacheShape->GetDimNum()) + "D", "1D or 3D or 4D"), + return ge::GRAPH_FAILED); + + if (tokenXShape->GetDimNum() == DIM_NUM_3) { // BS + shapeParam.isBsMerge = false; + shapeParam.B = tokenXShape->GetDim(DIM_INDEX_0); + shapeParam.S = tokenXShape->GetDim(DIM_INDEX_1); + shapeParam.Dr = ropeSinShape->GetDim(DIM_INDEX_2); + shapeParam.T = shapeParam.B * shapeParam.S; + } else { // T + shapeParam.isBsMerge = true; + shapeParam.T = tokenXShape->GetDim(DIM_INDEX_0); + shapeParam.Dr = ropeSinShape->GetDim(DIM_INDEX_1); + } + + shapeParam.N = weightUkShape->GetDim(DIM_INDEX_0); + shapeParam.Hckv = weightUkShape->GetDim(DIM_INDEX_2); + + shapeParam.Hcq = weightDqShape->GetDim(DIM_INDEX_1); + return GRAPH_SUCCESS; +} + +static bool IsWeightFullQuantFamily(int64_t weightQuantMode) +{ + return weightQuantMode == WEIGHT_QUANT_MODE_FULL_QUANT || + weightQuantMode == WEIGHT_QUANT_MODE_FULL_QUANT_FP8 || + weightQuantMode == WEIGHT_QUANT_MODE_FULL_QUANT_HIF8 || + weightQuantMode == WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT; +} + +static void SetQueryAndRopeShape(const MlaPrologProtoShapeParam &shapeParam, gert::Shape *queryShape, + gert::Shape *queryRopeShape) +{ + if (!shapeParam.isBsMerge) { + queryShape->SetDimNum(DIM_NUM_4); // (B, S, N, Hckv) + queryShape->SetDim(DIM_INDEX_0, shapeParam.B); + queryShape->SetDim(DIM_INDEX_1, shapeParam.S); + queryShape->SetDim(DIM_INDEX_2, shapeParam.N); + queryShape->SetDim(DIM_INDEX_3, shapeParam.Hckv); + + queryRopeShape->SetDimNum(DIM_NUM_4); // (B, S, N, Dr) + queryRopeShape->SetDim(DIM_INDEX_0, shapeParam.B); + queryRopeShape->SetDim(DIM_INDEX_1, shapeParam.S); + queryRopeShape->SetDim(DIM_INDEX_2, shapeParam.N); + queryRopeShape->SetDim(DIM_INDEX_3, shapeParam.Dr); + } else { + queryShape->SetDimNum(DIM_NUM_3); // (T, N, Hckv) + queryShape->SetDim(DIM_INDEX_0, shapeParam.T); + queryShape->SetDim(DIM_INDEX_1, shapeParam.N); + queryShape->SetDim(DIM_INDEX_2, shapeParam.Hckv); + + queryRopeShape->SetDimNum(DIM_NUM_3); // (T, N, Dr) + queryRopeShape->SetDim(DIM_INDEX_0, shapeParam.T); + queryRopeShape->SetDim(DIM_INDEX_1, shapeParam.N); + queryRopeShape->SetDim(DIM_INDEX_2, shapeParam.Dr); + } +} + +static void SetDequantScaleQNopeShape(const MlaPrologProtoShapeParam &shapeParam, gert::Shape *dequantScaleQNopeShape, + int64_t weightQuantMode, int64_t kvQuantMode) +{ + if (kvQuantMode == KV_QUANT_MODE_PER_TENSOR && IsWeightFullQuantFamily(weightQuantMode)) { + dequantScaleQNopeShape->SetDimNum(DIM_NUM_3); // (B*S, N, 1) | (T, N, 1) + dequantScaleQNopeShape->SetDim(DIM_INDEX_0, shapeParam.isBsMerge ? shapeParam.T : shapeParam.B * shapeParam.S); + dequantScaleQNopeShape->SetDim(DIM_INDEX_1, shapeParam.N); + dequantScaleQNopeShape->SetDim(DIM_INDEX_2, DIM_NUM_1); // 1: Fix dim 1 + } else { + dequantScaleQNopeShape->SetDimNum(DIM_NUM_1); + dequantScaleQNopeShape->SetDim(DIM_INDEX_0, DIM_NUM_0); + } +} + +static ge::graphStatus SetQueryNormShape(const MlaPrologProtoShapeParam &shapeParam, gert::InferShapeContext *context, + gert::Shape *queryNormShape, gert::Shape *dequantScaleQNormShape, + int64_t weightQuantMode, bool queryNormFlag) +{ + if (!queryNormFlag) { + queryNormShape->SetDimNum(DIM_NUM_1); + queryNormShape->SetDim(DIM_INDEX_0, DIM_NUM_0); + dequantScaleQNormShape->SetDimNum(DIM_NUM_1); + dequantScaleQNormShape->SetDim(DIM_INDEX_0, DIM_NUM_0); + return GRAPH_SUCCESS; + } + + auto weightUqQrDesc = context->GetInputDesc(WEIGHT_UQ_QR_INDEX); + OP_CHECK_NULL_WITH_CONTEXT(context, weightUqQrDesc); + (void)weightUqQrDesc; // 保持原有校验语义,desc 本身未被使用 + + int64_t tSize = shapeParam.isBsMerge ? shapeParam.T : shapeParam.B * shapeParam.S; + if (shapeParam.isBsMerge) { + // [T, Hcq] + queryNormShape->SetDimNum(DIM_NUM_2); + queryNormShape->SetDim(DIM_INDEX_0, shapeParam.T); + queryNormShape->SetDim(DIM_INDEX_1, shapeParam.Hcq); + } else { + // [B, S, Hcq] + queryNormShape->SetDimNum(DIM_NUM_3); + queryNormShape->SetDim(DIM_INDEX_0, shapeParam.B); + queryNormShape->SetDim(DIM_INDEX_1, shapeParam.S); + queryNormShape->SetDim(DIM_INDEX_2, shapeParam.Hcq); + } + + if (weightQuantMode == WEIGHT_QUANT_MODE_NO_QUANT) { + dequantScaleQNormShape->SetDimNum(DIM_NUM_1); + dequantScaleQNormShape->SetDim(DIM_INDEX_0, DIM_NUM_0); + } else if (weightQuantMode == WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT) { + dequantScaleQNormShape->SetDimNum(DIM_NUM_2); + dequantScaleQNormShape->SetDim(DIM_INDEX_0, tSize); + dequantScaleQNormShape->SetDim(DIM_INDEX_1, shapeParam.Hcq / FP8_E4M3_BLOCK_SIZE); + } else { + dequantScaleQNormShape->SetDimNum(DIM_NUM_2); + dequantScaleQNormShape->SetDim(DIM_INDEX_0, tSize); + dequantScaleQNormShape->SetDim(DIM_INDEX_1, DIM_NUM_1); + } + return GRAPH_SUCCESS; +} + +ge::graphStatus SetMlaPrologV3ShapeDim(const MlaPrologProtoShapeParam &shapeParam, gert::InferShapeContext *context) +{ + auto queryShape = context->GetOutputShape(QUERY_INDEX); // query: (B, S, N, Hckv) | (T, N, Hckv) + OP_CHECK_NULL_WITH_CONTEXT(context, queryShape); + auto queryRopeShape = context->GetOutputShape(QUERY_ROPE_INDEX); // queryRope: (B, S, N, Dr) | (T, N, Dr) + OP_CHECK_NULL_WITH_CONTEXT(context, queryRopeShape); + auto kvCacheOutShape = context->GetOutputShape(KV_CACHE_OUT_INDEX); // kvCacheOut: (B, Nkv, Skv, Hckv) + OP_CHECK_NULL_WITH_CONTEXT(context, kvCacheOutShape); + auto krCacheOutShape = context->GetOutputShape(KR_CACHE_OUT_INDEX); // krCacheOut: (B, Nkv, Skv, Dr) + OP_CHECK_NULL_WITH_CONTEXT(context, krCacheOutShape); + + // set output shape + auto attrs = context->GetAttrs(); + OP_CHECK_NULL_WITH_CONTEXT(context, attrs); + + // Get attribute pointers and dereference once + const int64_t *weightQuantModePtr = attrs->GetAttrPointer(ATTR_WEIGHT_QUANT_MODE_FLAG_INDEX); + const int64_t weightQuantMode = (weightQuantModePtr == nullptr) ? 0 : *weightQuantModePtr; + const int64_t *kvQuantModePtr = attrs->GetAttrPointer(ATTR_KV_QUANT_MODE_FLAG_INDEX); + const int64_t kvQuantMode = (kvQuantModePtr == nullptr) ? 0 : *kvQuantModePtr; + const bool *queryNormFlagPtr = attrs->GetAttrPointer(ATTR_QUERY_NORM_FLAG_INDEX); + const bool queryNormFlag = (queryNormFlagPtr == nullptr) ? 0 : *queryNormFlagPtr; + + // dequantScaleQNope: (B*S, N ,1) | (T, N, 1). (1) if not enabled + auto dequantScaleQNopeShape = context->GetOutputShape(DEQUANT_SCALE_Q_NOPE_INDEX); + OP_CHECK_NULL_WITH_CONTEXT(context, dequantScaleQNopeShape); + + // queryNorm + gert::Shape *queryNormShape = context->GetOutputShape(QUERY_NORM_INDEX); + OP_CHECK_NULL_WITH_CONTEXT(context, queryNormShape); + gert::Shape *dequantScaleQNormShape = context->GetOutputShape(DEQUANT_SCALE_Q_NORM_INDEX); + OP_CHECK_NULL_WITH_CONTEXT(context, dequantScaleQNormShape); + + SetQueryAndRopeShape(shapeParam, queryShape, queryRopeShape); + *kvCacheOutShape = *context->GetRequiredInputShape(KV_CACHE_INDEX_V3); + *krCacheOutShape = *context->GetRequiredInputShape(KR_CACHE_INDEX_V3); + SetDequantScaleQNopeShape(shapeParam, dequantScaleQNopeShape, weightQuantMode, kvQuantMode); + return SetQueryNormShape(shapeParam, context, queryNormShape, dequantScaleQNormShape, weightQuantMode, + queryNormFlag); +} + +ge::graphStatus InferShapeMlaPrologV3(gert::InferShapeContext *context) +{ + OP_LOGI(context->GetNodeName(), "Enter MlaPrologV3 infershape impl."); + + MlaPrologProtoShapeParam shapeParam{}; + auto apiRet = GetMlaPrologV3ShapeDim(context, shapeParam); + if (apiRet != GRAPH_SUCCESS) { + return GRAPH_FAILED; + } + + apiRet = SetMlaPrologV3ShapeDim(shapeParam, context); + if (apiRet != GRAPH_SUCCESS) { + return GRAPH_FAILED; + } + return GRAPH_SUCCESS; +} + +ge::graphStatus InferDataTypeMlaPrologV3(gert::InferDataTypeContext *context) +{ + OP_LOGI(context->GetNodeName(), "Enter MlaPrologV3 inferDataType impl."); + + auto attrs = context->GetAttrs(); + OP_CHECK_NULL_WITH_CONTEXT(context, attrs); + + // Get attribute pointers and dereference once + const int64_t *weightQuantModePtr = attrs->GetAttrPointer(ATTR_WEIGHT_QUANT_MODE_FLAG_INDEX); + const int weightQuantMode = (weightQuantModePtr == nullptr) ? 0 : *weightQuantModePtr; + const int64_t *kvQuantModePtr = attrs->GetAttrPointer(ATTR_KV_QUANT_MODE_FLAG_INDEX); + const int kvQuantMode = (kvQuantModePtr == nullptr) ? 0 : *kvQuantModePtr; + + // mxfp8 quant + if (weightQuantMode == WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT) { + bool isMxfp8FullQuant = (context->GetRequiredInputDataType(TOKEN_X_INDEX) == ge::DT_FLOAT8_E4M3FN && + context->GetOptionalInputDataType(QUANT_SCALE_CKV_INDEX) != ge::DT_UNDEFINED); + + context->SetOutputDataType(QUERY_INDEX, (isMxfp8FullQuant) ? + context->GetRequiredInputDataType(WEIGHT_DKV_KR_INDEX) : + context->GetRequiredInputDataType(WEIGHT_UK_INDEX)); + context->SetOutputDataType(QUERY_ROPE_INDEX, context->GetRequiredInputDataType(WEIGHT_UK_INDEX)); + context->SetOutputDataType(KV_CACHE_OUT_INDEX, context->GetRequiredInputDataType(KV_CACHE_INDEX_V3)); + context->SetOutputDataType(KR_CACHE_OUT_INDEX, context->GetRequiredInputDataType(KR_CACHE_INDEX_V3)); + context->SetOutputDataType(DEQUANT_SCALE_Q_NOPE_INDEX, ge::DT_FLOAT); + context->SetOutputDataType(QUERY_NORM_INDEX, context->GetRequiredInputDataType(WEIGHT_UQ_QR_INDEX)); + context->SetOutputDataType(DEQUANT_SCALE_Q_NORM_INDEX, ge::DT_FLOAT8_E8M0); + } else { + context->SetOutputDataType(QUERY_INDEX, context->GetRequiredInputDataType(WEIGHT_UK_INDEX)); + context->SetOutputDataType(QUERY_ROPE_INDEX, context->GetRequiredInputDataType(WEIGHT_UK_INDEX)); + context->SetOutputDataType(KV_CACHE_OUT_INDEX, context->GetRequiredInputDataType(KV_CACHE_INDEX_V3)); + context->SetOutputDataType(KR_CACHE_OUT_INDEX, context->GetRequiredInputDataType(KR_CACHE_INDEX_V3)); + + // full quant + bool isQuantQuery = + ((weightQuantMode == WEIGHT_QUANT_MODE_FULL_QUANT || weightQuantMode == WEIGHT_QUANT_MODE_FULL_QUANT_FP8 || + weightQuantMode == WEIGHT_QUANT_MODE_FULL_QUANT_HIF8) && + kvQuantMode == KV_QUANT_MODE_PER_TENSOR); + + context->SetOutputDataType(QUERY_INDEX, + isQuantQuery ? context->GetRequiredInputDataType(TOKEN_X_INDEX) : ge::DT_BF16); + context->SetOutputDataType(DEQUANT_SCALE_Q_NOPE_INDEX, ge::DT_FLOAT); + + if (weightQuantMode == WEIGHT_QUANT_MODE_NO_QUANT) { + context->SetOutputDataType(QUERY_NORM_INDEX, ge::DT_BF16); + } else { + context->SetOutputDataType(QUERY_NORM_INDEX, context->GetRequiredInputDataType(WEIGHT_UQ_QR_INDEX)); + } + context->SetOutputDataType(DEQUANT_SCALE_Q_NORM_INDEX, ge::DT_FLOAT); + } + + return GRAPH_SUCCESS; +} + +IMPL_OP_INFERSHAPE(MlaPrologV3).InferShape(InferShapeMlaPrologV3).InferDataType(InferDataTypeMlaPrologV3); +} // namespace ops \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_infershape.h b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_infershape.h new file mode 100644 index 000000000000..68bebaaee1b4 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_infershape.h @@ -0,0 +1,58 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_v3_infershape.h + * \brief + */ + +#ifndef MLA_PROLOG_V3_INFERSHAPE_H +#define MLA_PROLOG_V3_INFERSHAPE_H + +#include "mla_prolog_infershape.h" + +using namespace ge; + +namespace ops { +// INPUT +constexpr uint32_t WEIGHT_DQ_INDEX = 1; +constexpr uint32_t WEIGHT_UQ_QR_INDEX = 2; +constexpr uint32_t WEIGHT_DKV_KR_INDEX = 4; +constexpr uint32_t QUANT_SCALE_CKV_INDEX = 16; +// OUTPUT +constexpr uint32_t DEQUANT_SCALE_Q_NOPE_INDEX = 4; +constexpr uint32_t QUERY_NORM_INDEX = 5; +constexpr uint32_t DEQUANT_SCALE_Q_NORM_INDEX = 6; +// ATTRIBUTE +constexpr uint32_t ATTR_QUERY_NORM_FLAG_INDEX = 3; +constexpr uint32_t ATTR_WEIGHT_QUANT_MODE_FLAG_INDEX = 4; +constexpr uint32_t ATTR_KV_QUANT_MODE_FLAG_INDEX = 5; + +constexpr uint32_t WEIGHT_QUANT_MODE_NO_QUANT = 0; +constexpr uint32_t WEIGHT_QUANT_MODE_PARTIAL_QUANT = 1; +constexpr uint32_t WEIGHT_QUANT_MODE_FULL_QUANT = 2; +constexpr uint32_t WEIGHT_QUANT_MODE_MXFP8_FULL_QUANT = 3; +constexpr uint32_t WEIGHT_QUANT_MODE_FULL_QUANT_FP8 = 4; +constexpr uint32_t WEIGHT_QUANT_MODE_FULL_QUANT_HIF8 = 5; +constexpr uint32_t KV_QUANT_MODE_NO_QUANT = 0; +constexpr uint32_t KV_QUANT_MODE_PER_TENSOR = 1; +constexpr uint32_t KV_QUANT_MODE_PER_CHANNEL = 2; +constexpr uint32_t KV_QUANT_MODE_PER_TILE = 3; + +ge::graphStatus GetMlaPrologV3ShapeDim(const gert::InferShapeContext *context, MlaPrologProtoShapeParam &shapeParam); +ge::graphStatus SetMlaPrologV3ShapeDim(const MlaPrologProtoShapeParam &shapeParam, gert::InferShapeContext *context); + +ge::graphStatus InferShapeMlaPrologV3(gert::InferShapeContext *context); +ge::graphStatus InferDataTypeMlaPrologV3(gert::InferDataTypeContext *context); + + +} // namespace ops + +#endif // MLA_PROLOG_V3_INFERSHAPE_H \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_tiling.h b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_tiling.h new file mode 100644 index 000000000000..f0401114662a --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_tiling.h @@ -0,0 +1,26 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#ifndef MLA_PROLOG_V3_TILING_H +#define MLA_PROLOG_V3_TILING_H + +#include "register/tilingdata_base.h" +#include "mla_prolog_tiling.h" + +#ifdef ASCENDC_OP_TEST +#define MLA_EXTERN_C extern "C" +#else +#define MLA_EXTERN_C +#endif + +namespace optiling { +} // optiling + +#endif \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_tiling_register.cpp b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_tiling_register.cpp new file mode 100644 index 000000000000..e44fb089b15e --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_v3_tiling_register.cpp @@ -0,0 +1,26 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include "mla_prolog_v3_tiling.h" +#include "register/op_def_registry.h" + +using namespace ge; +using namespace AscendC; +namespace optiling { +ge::graphStatus TilingPrepareForMlaProlog(gert::TilingParseContext *context) +{ + (void)context; + return ge::GRAPH_SUCCESS; +} + +IMPL_OP_OPTILING(MlaPrologV3) + .Tiling(TilingMlaProlog) + .TilingParse(TilingPrepareForMlaProlog); +} // namespace optiling diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_m.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_m.h new file mode 100644 index 000000000000..7d73973f2698 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_m.h @@ -0,0 +1,1785 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file kernel_mla_prolog_split_m.h + * \brief + */ + +#ifndef KERNEL_MLA_PROLOG_SPLIT_M_H +#define KERNEL_MLA_PROLOG_SPLIT_M_H + +#include "mla_prolog_comm.h" +#include "mla_prolog_vector_comm.h" +#include "service_matmul.h" +#include "service_rms_norm.h" +#include "service_gather_sin_cos.h" +#include "service_rotary_position_embedding.h" +#include "service_scatter_cache.h" +#include "service_dequant.h" +#include "service_dynamic_quant_qn_mul_qr.h" +#include "../mla_prolog_tiling_data.h" +#include "../mla_prolog_template_tiling_key.h" + +namespace MlaProlog { +template +class MlaPrologV3SplitM { +public: + static constexpr bool isPertile = MLAPT::isPertile; + + using mmInputType = typename MLAPT::mmInputType; + using mmQcQrInputType = typename MLAPT::mmQcQrInputType; + using mmQnInputType = typename MLAPT::mmQnInputType; + using mmCqOutputType = typename MLAPT::mmCqOutputType; + using mmCkvKrOutputType = typename MLAPT::mmCkvKrOutputType; + using mmQcQrOutputType = typename MLAPT::mmQcQrOutputType; + using mmQnOutputType = typename MLAPT::mmQnOutputType; + using rmsNormGammaType = typename MLAPT::rmsNormGammaType; + using rmsNormComputType = typename MLAPT::rmsNormComputType; + using rmsNormCqOutputType = typename MLAPT::rmsNormCqOutputType; + using rmsNormCkvOutputType = typename MLAPT::rmsNormCkvOutputType; + using ropeSinCosType = typename MLAPT::ropeSinCosType; + using ropeComputType = typename MLAPT::ropeComputType; + using ropeOutputType = typename MLAPT::ropeOutputType; + using queryOutputType = typename std::conditional::type; + using kvCacheType = typename MLAPT::kvCacheType; + using krCacheType = typename MLAPT::krCacheType; + using dequantScaleQNopeType = typename MLAPT::dequantScaleQNopeType; + using dequantScaleQNormType = typename MLAPT::dequantScaleQNormType; + using dequantScaleType = typename MLAPT::dequantScaleType; + + MMParams mmCqParam_; + MMParams mmCkvKrParam_; + MMParams mmQcQrParam_; + MMParams mmQnParam_; + + __aicore__ inline MlaPrologV3SplitM(TPipe *pipe, const optiling::MlaPrologTilingData *__restrict tilingData, + const optiling::MlaPrologBaseParams *__restrict baseParams) + : pipe_(pipe), tilingData_(tilingData), baseParams_(baseParams) + { + } + + __aicore__ inline void Init(__gm__ uint8_t *tokenX, __gm__ uint8_t *weightDq, __gm__ uint8_t *weightUqQr, + __gm__ uint8_t *weightUk, __gm__ uint8_t *weightDkvKr, __gm__ uint8_t *rmsnormGammaCq, + __gm__ uint8_t *rmsnormGammaCkv, __gm__ uint8_t *ropeSin, __gm__ uint8_t *ropeCos, + __gm__ uint8_t *cacheIndex, __gm__ uint8_t *kvCache, __gm__ uint8_t *krCache, + __gm__ uint8_t *dequantScaleX, __gm__ uint8_t *dequantScaleWDq, + __gm__ uint8_t *deqScaleQcQrW, __gm__ uint8_t *dequantScaleWDkvkr, + __gm__ uint8_t *quantScaleCkv, __gm__ uint8_t *quantScaleCkr, + __gm__ uint8_t *smoothScaleCq, __gm__ uint8_t *actualSeqLen, + __gm__ uint8_t *kNopeClipAlpha, __gm__ uint8_t *queryOut, __gm__ uint8_t *queryRopeOut, + __gm__ uint8_t *dequantScaleQNopeOut, __gm__ uint8_t *queryNormOut, + __gm__ uint8_t *dequantScaleQNormOut, __gm__ uint8_t *workspace); + __aicore__ inline void Process(); + +private: + __aicore__ inline void CopyGlobalParams(); + __aicore__ inline void OutputInit(__gm__ uint8_t *actualSeqLen, __gm__ uint8_t *queryOut, + __gm__ uint8_t *queryRopeOut, __gm__ uint8_t *dequantScaleQNopeOut, + __gm__ uint8_t *queryNormOut, __gm__ uint8_t *dequantScaleQNormOut); + __aicore__ inline void ScaleInit(__gm__ uint8_t *dequantScaleX, __gm__ uint8_t *dequantScaleWDq, + __gm__ uint8_t *deqScaleQcQrW, __gm__ uint8_t *dequantScaleWDkvkr, + __gm__ uint8_t *quantScaleCkv, __gm__ uint8_t *quantScaleCkr, + __gm__ uint8_t *smoothScaleCq, __gm__ uint8_t *kNopeClipAlpha); + __aicore__ inline void WorkspaceInit(__gm__ uint8_t *workspace); + __aicore__ inline void MmParamInit(); + __aicore__ inline void MmCqParamInit(); + __aicore__ inline void MmCkvKrParamInit(); + __aicore__ inline void MmQcQrParamInit(); + __aicore__ inline void MmQnParamInit(); + __aicore__ inline void CubeBufferInit(); + __aicore__ inline void VectorBufferInit(); + __aicore__ inline void UpdateStepBatchParams(int64_t curMSize); + __aicore__ inline void ComputeBlkScatterOffsets(GlobalTensor indexGm, int64_t tokenIndex, int64_t rows, + CkvkrParams &rmsNormAndScatterCkvParams, + CkvkrParams &ropeAndScatterKrParams); + template + __aicore__ inline void AicProcess(AicOffset &aicOffset, int64_t batchOffset, int64_t mmQnLoops); + template + __aicore__ inline void AivProcess(AivOffset &aivOffset, int64_t batchOffset, int64_t curMSize, + int64_t numHeadOffset, int64_t mmQnLoops); + template + __aicore__ inline void + MatmulSplitM(const GlobalTensor &tensorResGm, const GlobalTensor &tensorAGm, const GlobalTensor &tensorBGm, + const MMParams &mmPara, const UsedBlockParams &mmBlockParams, + const GlobalTensor &tensorAScaleGm = {}, const GlobalTensor &tensorBScaleGm = {}); + __aicore__ inline void MatmulQcQr(AicOffset &aicOffset); + template + __aicore__ inline void MatmulQnSyncDynamicQuantAndMulQr(int64_t qcOffset, int64_t weightUkOffset, + int64_t qnResOffset, int64_t mmQnLoops); + __aicore__ inline void CopyInSinCos(int64_t tokenIndex, int64_t curVecToken, int64_t batchOffset, int64_t curMSize); + __aicore__ inline void RmsNormCq(int64_t tokenIndex, int64_t rmsNormCqOffset, int64_t rmsNormCqResOffset, + int64_t curVecToken, int64_t curBlockTokenOffset); + __aicore__ inline void RopeAndScatterKr(LocalTensor &dequantScaleXLocal, LocalTensor &shareTmpUb, + LocalTensor &cosLocalCkvKr, + LocalTensor &sinLocalCkvKr, + CkvkrParams ropeAndScatterKrParams); + __aicore__ inline void ScatterKr(LocalTensor &outputKrLocal, CkvkrParams ropeAndScatterKrParams); + __aicore__ inline void RmsNormAndScatterCkv(LocalTensor &dequantScaleXLocal, + LocalTensor &shareTmpUb, + LocalTensor &cosLocalCkvKr, + LocalTensor &sinLocalCkvKr, + CkvkrParams rmsNormAndScatterCkvParams); + __aicore__ inline void RmsNormAndQuantizeCkv(LocalTensor &outputLocal, + LocalTensor &rmsNormShareTmpUb, + LocalTensor &dequantScaleXLocal, RmsNormParam rmsNormParams, + CkvkrParams rmsNormAndScatterCkvParams); + __aicore__ inline void ScatterCkv(LocalTensor &outputLocal, CkvkrParams rmsNormAndScatterCkvParams); + __aicore__ inline void RmsNormRopeScatterCkvKr(int64_t tokenIndex, int64_t rmsNormCkvOffset, int64_t ropeKrOffset, + int64_t curVecToken); + // 低时延算力分组场景 + __aicore__ inline void QcQrSplit(int64_t curVecToken, int64_t curBlockTokenOffset, int64_t curMSize, + int64_t mmQnPreDequantOffset, int64_t mmQnPreDequantResOffset, + int64_t ropeQrOffset, int64_t ropeQrResOffset); + __aicore__ inline void DynamicQuantQnAndMulQrSyncMMQn(int64_t batchOffset, int64_t curMSize, int64_t numHeadOffset, + int64_t mmQnLoops); + +private: + TPipe *pipe_; + const optiling::MlaPrologTilingData *__restrict tilingData_; + const optiling::MlaPrologBaseParams *__restrict baseParams_; + uint32_t blockIdx_ = 0U; + uint32_t cubeBlockIdx_ = 0U; // AIV上使用AIC的blockIdx + int64_t vectorRow_ = 1; + int64_t curVectorBlockNum_; + int64_t vectorCoreNum_; + uint64_t dequantScaleCqSize_ = 1; + uint32_t curStepVecFrontToken_; + uint32_t curStepVecFrontListNum_; + uint32_t curStepVecBackToken_; + uint32_t curVecTokenMax_; + bool enableSmoothScalesCq_; + static constexpr uint32_t cvRatio_ = MLAPT::cvRatio; // 默认cv 1:2 + static constexpr bool isFp8E8m0 = std::is_same::value; + + uint32_t mSubSize_ = 0; + uint32_t mOffsetStart_ = 0; + + struct DequantTool { + GlobalTensor deQuantScaleCqGm_; + TBuf deQuantScaleCqBuffer_; // 用于临时存储每一行的Scale,以及汇总最终每一行的Scale参数 + LocalTensor deQuantScaleCqLocal_; + __aicore__ inline DequantTool() + { + } + }; + + // 算子分组开关 + DequantTool dequantTool_; + + // GM + GlobalTensor tokenXGm_; + GlobalTensor weightDqGm_; + GlobalTensor weightUqQrGm_; + GlobalTensor weightUkGm_; + GlobalTensor weightDkvKrGm_; + GlobalTensor rmsnormGammaCqGm_; + GlobalTensor rmsnormGammaCkvGm_; + GlobalTensor ropeSinGm_; + GlobalTensor ropeCosGm_; + GlobalTensor cacheIndexGm_; + GlobalTensor kvCacheGm_; + GlobalTensor krCacheGm_; + GlobalTensor qrOutGm_; + + GlobalTensor dequantScaleXGm_; + GlobalTensor dequantScaleWDqGm_; + GlobalTensor dequantScaleWDkvkrGm_; + GlobalTensor smoothScaleCqGm_; + GlobalTensor deqScaleQcQrW_; // per-channel反量化参数 + GlobalTensor quantScaleCkvGm_; + GlobalTensor quantScaleCkrGm_; + + GlobalTensor actualSeqLenGm_; + GlobalTensor kNopeClipAlphaGm_; + + GlobalTensor rmsNormCqResGm_; + GlobalTensor mmCqResGm_; + GlobalTensor mmCkvKrResGm_; + GlobalTensor mmQcQrResGm_; + GlobalTensor mmQcQrResDequantGm_; + GlobalTensor mmQnResGm_; + GlobalTensor dequantScaleQNopeGm_; + GlobalTensor queryOutGm_; + GlobalTensor dequantScaleQNormGm_; + + // UB + TBuf sincosBuffer_; + TBuf shareBuffer_; + TBuf dequantScaleWDqBuffer_; + TBuf dequantScaleWDkvKrBuffer_; + TBuf rmsnormGammaCqBuffer_; + TBuf rmsnormGammaCkvBuffer_; + TBuf smoothScaleCqBuffer_; + TBuf quantScaleCkvBuffer_; + TBuf quantScaleCkrBuffer_; + TBuf stepActualSeqBuffer_; + + LocalTensor cosLocal_; + LocalTensor sinLocal_; + LocalTensor dequantScaleWDqLocal_; + LocalTensor dequantScaleWDkvKrLocal_; + LocalTensor rmsnormGammaCqLocal_; + LocalTensor rmsnormGammaCkvLocal_; + LocalTensor smoothScaleCqLocal_; + LocalTensor quantScaleCkvLocal_; + LocalTensor quantScaleCkrLocal_; + LocalTensor stepActualSeqLocal_; + + struct ActSeqState { + bool inited = false; // whether state is initialized for the run + int64_t curBatch = 0; // current batch index for the running token index + int64_t prevPrefix = 0; // prefix sum of seq lengths up to (curBatch - 1) + int64_t curPrefix = 0; // prefix sum of seq lengths up to curBatch (exclusive high bound) + int64_t prevIndexOffset = 0; // prefix sum of blocks up to (curBatch - 1) + int64_t curIndexOffset = 0; // prefix sum of blocks up to curBatch + }; + ActSeqState actSeqState_; + int64_t blockSpanPerBatch_ = 0; // blockNum * blockSize + + TBuf aBufL1_; + TBuf bBufL1_; + LocalTensor aL1Tensor_; + LocalTensor bL1Tensor_; + MMBufParams bufParam_; + TBuf aBufL0_; + TBuf bBufL0_; + TBuf cBufL0_; + TBuf zeroBuf_; + LocalTensor zeroSrc_; +}; + + +template +__aicore__ inline void MlaPrologV3SplitM::Init( + __gm__ uint8_t *tokenX, __gm__ uint8_t *weightDq, __gm__ uint8_t *weightUqQr, __gm__ uint8_t *weightUk, + __gm__ uint8_t *weightDkvKr, __gm__ uint8_t *rmsnormGammaCq, __gm__ uint8_t *rmsnormGammaCkv, + __gm__ uint8_t *ropeSin, __gm__ uint8_t *ropeCos, __gm__ uint8_t *cacheIndex, __gm__ uint8_t *kvCache, + __gm__ uint8_t *krCache, __gm__ uint8_t *dequantScaleX, __gm__ uint8_t *dequantScaleWDq, + __gm__ uint8_t *deqScaleQcQrW, __gm__ uint8_t *dequantScaleWDkvkr, __gm__ uint8_t *quantScaleCkv, + __gm__ uint8_t *quantScaleCkr, __gm__ uint8_t *smoothScaleCq, __gm__ uint8_t *actualSeqLen, + __gm__ uint8_t *kNopeClipAlpha, __gm__ uint8_t *queryOut, __gm__ uint8_t *queryRopeOut, + __gm__ uint8_t *dequantScaleQNopeOut, __gm__ uint8_t *queryNormOut, __gm__ uint8_t *dequantScaleQNormOut, + __gm__ uint8_t *workspace) +{ + blockIdx_ = GetBlockIdx(); // cube:0-23 vec:0-47 + if ASCEND_IS_AIV { + cubeBlockIdx_ = blockIdx_ / cvRatio_; + } else { + cubeBlockIdx_ = blockIdx_; + } + curVectorBlockNum_ = static_cast(baseParams_->stepBatchSize); + vectorCoreNum_ = static_cast(baseParams_->vectorBlockNum); + curVecTokenMax_ = (curVectorBlockNum_ + vectorCoreNum_ - 1) / vectorCoreNum_; + enableSmoothScalesCq_ = smoothScaleCq == nullptr ? false : true; + // GM + tokenXGm_.SetGlobalBuffer((__gm__ mmInputType *)tokenX); + weightDqGm_.SetGlobalBuffer((__gm__ mmInputType *)weightDq); // NZ + weightUqQrGm_.SetGlobalBuffer((__gm__ mmQcQrInputType *)weightUqQr); // NZ + weightUkGm_.SetGlobalBuffer((__gm__ mmQnInputType *)weightUk); + weightDkvKrGm_.SetGlobalBuffer((__gm__ mmInputType *)weightDkvKr); // NZ + rmsnormGammaCqGm_.SetGlobalBuffer((__gm__ rmsNormGammaType *)rmsnormGammaCq); + rmsnormGammaCkvGm_.SetGlobalBuffer((__gm__ rmsNormGammaType *)rmsnormGammaCkv); + ropeSinGm_.SetGlobalBuffer((__gm__ ropeSinCosType *)ropeSin); + ropeCosGm_.SetGlobalBuffer((__gm__ ropeSinCosType *)ropeCos); + if constexpr (MLAPT::cacheMode != CACHE_MODE::ND) { + cacheIndexGm_.SetGlobalBuffer((__gm__ int64_t *)cacheIndex); + } + kvCacheGm_.SetGlobalBuffer((__gm__ kvCacheType *)kvCache); + krCacheGm_.SetGlobalBuffer((__gm__ krCacheType *)krCache); + + OutputInit(actualSeqLen, queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut); + ScaleInit(dequantScaleX, dequantScaleWDq, deqScaleQcQrW, dequantScaleWDkvkr, quantScaleCkv, quantScaleCkr, + smoothScaleCq, kNopeClipAlpha); + MmParamInit(); + WorkspaceInit(workspace); + if ASCEND_IS_AIV { + VectorBufferInit(); + } else { + CubeBufferInit(); + constexpr int64_t zeroChunk = 512; + pipe_->InitBuffer(zeroBuf_, zeroChunk * sizeof(float)); + zeroSrc_ = zeroBuf_.Get(); + Duplicate(zeroSrc_, static_cast(0), zeroChunk); + } +} + +template +__aicore__ inline void +MlaPrologV3SplitM::OutputInit(__gm__ uint8_t *actualSeqLen, __gm__ uint8_t *queryOut, + __gm__ uint8_t *queryRopeOut, __gm__ uint8_t *dequantScaleQNopeOut, + __gm__ uint8_t *queryNormOut, __gm__ uint8_t *dequantScaleQNormOut) +{ + qrOutGm_.SetGlobalBuffer((__gm__ ropeOutputType *)queryRopeOut); + if constexpr (((std::is_same::value && std::is_same::value) || + (std::is_same::value && std::is_same::value) || + (std::is_same::value && std::is_same::value)) && + !isPertile) { + dequantScaleQNopeGm_.SetGlobalBuffer((__gm__ dequantScaleQNopeType *)dequantScaleQNopeOut); + queryOutGm_.SetGlobalBuffer((__gm__ queryOutputType *)queryOut); + } else { + mmQnResGm_.SetGlobalBuffer((__gm__ mmQnOutputType *)queryOut); + } + if (baseParams_->queryNormFlag == 1U) { + rmsNormCqResGm_.SetGlobalBuffer((__gm__ mmQcQrInputType *)queryNormOut); + if constexpr (IsFullQuantMode()) { + dequantScaleQNormGm_.SetGlobalBuffer((__gm__ dequantScaleQNormType *)dequantScaleQNormOut); + } + } + if constexpr (MLAPT::actualSeqMode == ACTUAL_SEQ_MODE::EN_Q_LEN) { + actualSeqLenGm_.SetGlobalBuffer((__gm__ int32_t *)actualSeqLen); + } +} + +template +__aicore__ inline void +MlaPrologV3SplitM::ScaleInit(__gm__ uint8_t *dequantScaleX, __gm__ uint8_t *dequantScaleWDq, + __gm__ uint8_t *deqScaleQcQrW, __gm__ uint8_t *dequantScaleWDkvkr, + __gm__ uint8_t *quantScaleCkv, __gm__ uint8_t *quantScaleCkr, + __gm__ uint8_t *smoothScaleCq, __gm__ uint8_t *kNopeClipAlpha) +{ + if constexpr (IsFullQuantMode()) { + dequantScaleXGm_.SetGlobalBuffer((__gm__ dequantScaleType *)dequantScaleX); + dequantScaleWDqGm_.SetGlobalBuffer((__gm__ dequantScaleType *)dequantScaleWDq); + dequantScaleWDkvkrGm_.SetGlobalBuffer((__gm__ dequantScaleType *)dequantScaleWDkvkr); + } + if constexpr (std::is_same::value && isFp8E8m0) { + deqScaleQcQrW_.SetGlobalBuffer((__gm__ dequantScaleType *)deqScaleQcQrW); + quantScaleCkvGm_.SetGlobalBuffer((__gm__ float *)quantScaleCkv); + } + + if constexpr (isPertile) { + kNopeClipAlphaGm_.SetGlobalBuffer((__gm__ float *)kNopeClipAlpha); + } +} + +template +__aicore__ inline void MlaPrologV3SplitM::MmParamInit() +{ + MmCqParamInit(); + MmCkvKrParamInit(); + MmQcQrParamInit(); + MmQnParamInit(); +} + +template +__aicore__ inline void MlaPrologV3SplitM::MmCqParamInit() +{ + mmCqParam_.m = baseParams_->stepBatchSize; + mmCqParam_.n = baseParams_->mm1SingleCoreN; // 1536 / 24 = 64 + mmCqParam_.k = baseParams_->headSizeX; // 7168 + mmCqParam_.needSetOrgShape = 1; + mmCqParam_.orgM = mmCqParam_.m; + mmCqParam_.orgN = mmCqParam_.n; + mmCqParam_.orgKa = mmCqParam_.k; + mmCqParam_.orgKb = mmCqParam_.k; + mmCqParam_.orgKc = baseParams_->headSizeCq; // 1536 + mmCqParam_.baseK = + (sizeof(mmInputType) == ONE_BYTE_TYPE_SIZE) ? 256 : 128; // 128KB / (128 max baseN * 4 stepK * sizeof(type)) + mmCqParam_.baseN = (std::is_same::value && isFp8E8m0) ? 64 : 128; + mmCqParam_.stepK = 4; + if ((mmCqParam_.k / mmCqParam_.baseK) % mmCqParam_.stepK != 0) { + mmCqParam_.stepK = 3; // support k = 7680, mmInputType int8, no tail + } + mmCqParam_.kL1StepSize = mmCqParam_.baseK * mmCqParam_.stepK; + mmCqParam_.kScale = mmCqParam_.k / FP8_E4M3_BLOCK_SIZE; +} + +template +__aicore__ inline void MlaPrologV3SplitM::MmCkvKrParamInit() +{ + mmCkvKrParam_.m = baseParams_->stepBatchSize; + mmCkvKrParam_.n = baseParams_->mm2SingleCoreN; + mmCkvKrParam_.k = baseParams_->headSizeX; // 7168 + mmCkvKrParam_.needSetOrgShape = 1; + mmCkvKrParam_.orgM = mmCkvKrParam_.m; + mmCkvKrParam_.orgN = mmCkvKrParam_.n; + mmCkvKrParam_.orgKa = mmCkvKrParam_.k; + mmCkvKrParam_.orgKb = mmCkvKrParam_.k; + mmCkvKrParam_.orgKc = (baseParams_->headSizeCkv + baseParams_->dimHeadRope); // 576 + mmCkvKrParam_.baseK = + (sizeof(mmInputType) == sizeof(int8_t)) ? 256 : 128; // 128KB / (128 max baseN * 4 stepK * sizeof(type)) + mmCkvKrParam_.baseN = (std::is_same::value && isFp8E8m0) ? 64 : 128; + mmCkvKrParam_.stepK = 4; + if ((mmCkvKrParam_.k / mmCkvKrParam_.baseK) % mmCkvKrParam_.stepK != 0) { + mmCkvKrParam_.stepK = 3; // support k = 7680, mmInputType int8, no tail + } + mmCkvKrParam_.kL1StepSize = mmCkvKrParam_.baseK * mmCkvKrParam_.stepK; + mmCkvKrParam_.kScale = mmCkvKrParam_.k / FP8_E4M3_BLOCK_SIZE; +} + +template +__aicore__ inline void MlaPrologV3SplitM::MmQcQrParamInit() +{ + mmQcQrParam_.m = baseParams_->stepBatchSize; + mmQcQrParam_.n = baseParams_->mm3SingleCoreN; + mmQcQrParam_.k = baseParams_->headSizeCq; // 1536 + mmQcQrParam_.needSetOrgShape = 1; + mmQcQrParam_.orgM = mmQcQrParam_.m; + mmQcQrParam_.orgN = mmQcQrParam_.n; + mmQcQrParam_.orgKa = mmQcQrParam_.k; + mmQcQrParam_.orgKb = mmQcQrParam_.k; + mmQcQrParam_.orgKc = (baseParams_->headSizeQc + baseParams_->headSizeQr); // (128 * 32 + 64 * 32) + mmQcQrParam_.baseK = (sizeof(mmQcQrInputType) == sizeof(int8_t)) ? 128 : 64; + if constexpr (MLAPT::enableGroupComputeOpt) { + mmQcQrParam_.baseN = 128; + } else { + mmQcQrParam_.baseN = 128; + } + mmQcQrParam_.stepK = 4; + mmQcQrParam_.kL1StepSize = mmQcQrParam_.baseK * mmQcQrParam_.stepK; + mmQcQrParam_.kScale = mmQcQrParam_.k / FP8_E4M3_BLOCK_SIZE; +} + + +template +__aicore__ inline void MlaPrologV3SplitM::MmQnParamInit() +{ + mmQnParam_.m = baseParams_->stepBatchSize; + mmQnParam_.n = baseParams_->headSizeCkv; // 512 + mmQnParam_.k = baseParams_->dimHeadSizeQc; // 128, 这里numHeadSize被分核,matmul设置里不体现 + mmQnParam_.needSetOrgShape = 1; + mmQnParam_.orgM = mmQnParam_.m; + mmQnParam_.orgN = mmQnParam_.n; + if constexpr (std::is_same::value || std::is_same::value) { + mmQnParam_.orgKa = baseParams_->headSizeQc; + } else { + mmQnParam_.orgKa = baseParams_->headSizeQc + baseParams_->headSizeQr; + } + mmQnParam_.orgKb = baseParams_->dimHeadSizeQc; + mmQnParam_.orgKc = baseParams_->headSizeCkv * baseParams_->numHeadSize; + mmQnParam_.baseN = 128; + mmQnParam_.baseK = 128; + mmQnParam_.stepK = 1; + if ((mmQnParam_.k > mmQnParam_.baseK) && (mmQnParam_.k % mmQnParam_.baseK != 0)) { + mmQnParam_.baseK = 64; + mmQnParam_.stepK = 3; // support D = 192 + } + mmQnParam_.kL1StepSize = mmQnParam_.baseK * mmQnParam_.stepK; + mmQnParam_.kScale = mmQnParam_.k / FP8_E4M3_BLOCK_SIZE; +} + +template +__aicore__ inline void MlaPrologV3SplitM::VectorBufferInit() +{ + pipe_->InitBuffer(rmsnormGammaCqBuffer_, baseParams_->headSizeCq * sizeof(rmsNormGammaType)); // [1, 1536] bf16 + rmsnormGammaCqLocal_ = rmsnormGammaCqBuffer_.Get(); + + pipe_->InitBuffer(rmsnormGammaCkvBuffer_, baseParams_->headSizeCkv * sizeof(rmsNormGammaType)); // [1, 512] bf16 + rmsnormGammaCkvLocal_ = rmsnormGammaCkvBuffer_.Get(); + + if constexpr (std::is_same::value) { + pipe_->InitBuffer(quantScaleCkrBuffer_, baseParams_->dimHeadRope * sizeof(float)); // [1, 64] + quantScaleCkrLocal_ = quantScaleCkrBuffer_.Get(); + } + + if constexpr (IsFullQuantMode()) { + if constexpr (std::is_same::value || + (std::is_same::value && !isFp8E8m0)) { + pipe_->InitBuffer(quantScaleCkvBuffer_, ALIGN_BLOCK_SIZE); + } else { + pipe_->InitBuffer(quantScaleCkvBuffer_, baseParams_->headSizeCkv * sizeof(float)); // [1, 512] + } + quantScaleCkvLocal_ = quantScaleCkvBuffer_.Get(); + } + + // 预留brcb的空间 + pipe_->InitBuffer(dequantTool_.deQuantScaleCqBuffer_, (baseParams_->stepBatchSize + 7) * ALIGN_BLOCK_SIZE); + dequantTool_.deQuantScaleCqLocal_ = dequantTool_.deQuantScaleCqBuffer_.template Get(); + + if constexpr (MLAPT::enableDequantOpt) { + // 在ropeQr进行切M处理后,会复用shareBuffer的内存,不需要额外申请 + // 开启开关后会按照head切分rope qr,此时需要加载一半batchsize数量的sin和cos值 + // 需要2倍的空间分别存储sin和cos + pipe_->InitBuffer(sincosBuffer_, 2 * baseParams_->dimHeadRope * sizeof(ropeComputType) * + ((baseParams_->stepBatchSize + 1) >> 1)); + } else { + // 需要2倍的空间分别存储sin和cos + pipe_->InitBuffer(sincosBuffer_, + 2 * baseParams_->dimHeadRope * sizeof(ropeComputType) * curVecTokenMax_); // [2, 64] float + } + + uint64_t usedAddr; + if constexpr (MLAPT::enableDequantOpt) { + cosLocal_ = sincosBuffer_.Get(); + sinLocal_ = cosLocal_[baseParams_->dimHeadRope * ((baseParams_->stepBatchSize + 1) >> 1)]; + usedAddr = reinterpret_cast( + sinLocal_[baseParams_->dimHeadRope * ((baseParams_->stepBatchSize + 1) >> 1)].GetPhyAddr()); + } else { + cosLocal_ = sincosBuffer_.Get(); + sinLocal_ = cosLocal_[baseParams_->dimHeadRope * curVecTokenMax_]; + usedAddr = reinterpret_cast(sinLocal_[baseParams_->dimHeadRope * curVecTokenMax_].GetPhyAddr()); + } + + constexpr int64_t zeroChunk = 512; + pipe_->InitBuffer(zeroBuf_, zeroChunk * sizeof(float)); + zeroSrc_ = zeroBuf_.Get(); + Duplicate(zeroSrc_, static_cast(0), zeroChunk); + + // 由于shareBuffer属于各个vector操作临时申请内存的区域内存使用不固定,建议shareBuffer始终放在最后,防止写入shareBuffer越界导致前面固定申请的UB内存被踩。 + pipe_->InitBuffer(shareBuffer_, MAX_UB_SIZE - usedAddr - zeroChunk * sizeof(float)); +} + +template +__aicore__ inline void MlaPrologV3SplitM::CubeBufferInit() +{ + // cube相关Buffer初始化 + pipe_->InitBuffer(aBufL1_, L1_A_SIZE * 2); + pipe_->InitBuffer(bBufL1_, L1_B_SIZE * 2); + + SetFlag(A_EVENT0); + SetFlag(A_EVENT1); + SetFlag(B_EVENT0); + SetFlag(B_EVENT1); + aL1Tensor_ = aBufL1_.Get(); + bL1Tensor_ = bBufL1_.Get(); + bufParam_.aL1BufAddr = aBufL1_.GetBufferAddr(aL1Tensor_.GetBufferHandle()); + bufParam_.bL1BufAddr = bBufL1_.GetBufferAddr(bL1Tensor_.GetBufferHandle()); + + pipe_->InitBuffer(aBufL0_, L0A_PP_SIZE * 2); // 64K + pipe_->InitBuffer(bBufL0_, L0B_PP_SIZE * 2); // 64K + pipe_->InitBuffer(cBufL0_, L0C_PP_SIZE * 2); // 128K + + SetFlag(L0A_EVENT0); + SetFlag(L0A_EVENT1); + SetFlag(L0B_EVENT0); + SetFlag(L0B_EVENT1); + + SetFlag(L0C_EVENT0); + SetFlag(L0C_EVENT1); + + SetFlag(SCALE_EVENT); + + bufParam_.aL0BufAddr = aBufL0_.GetBufferAddr(aBufL0_.Get().GetBufferHandle()); + bufParam_.bL0BufAddr = bBufL0_.GetBufferAddr(bBufL0_.Get().GetBufferHandle()); + bufParam_.cL0BufAddr = cBufL0_.GetBufferAddr(cBufL0_.Get().GetBufferHandle()); +} + +/* + * workspace管理 + * 1. 常驻:dequantTool_.deQuantScaleCqGm_ stepBs * 32 Byte + * 2. 中间结果: + * tokenXGm_──────>mmCkvKrResGm_ + * | [stepBS, HCkv + Dr] + * | (bf16 | int32) + * └─────────>mmCqResGm_──────>rmsNormCqResGm_──────>mmQcQrResGm_──────>mmQcQrResDequantGm_──────>mmQnResGm_ + * [stepBS, HCq] [stepBS, HCq] [stepBS, N1, D + Dr] [stepBS, N1, D] [stepBS, N1, + * HCkv] (bf16 | int32) (bf16 | int8) (bf16 | int32) (bf16) (bf16) + */ +template +__aicore__ inline void MlaPrologV3SplitM::WorkspaceInit(__gm__ uint8_t *workspace) +{ + int64_t workspaceOffset = 0; + if constexpr (std::is_same::value && isFp8E8m0) { + dequantScaleCqSize_ = baseParams_->headSizeCq / FP8_E4M3_BLOCK_SIZE; + } + dequantScaleCqSize_ = Align(dequantScaleCqSize_, BYTE_BLOCK); + + if constexpr (MLAPT::enableGroupComputeOpt || MLAPT::enableDequantOpt) { + dequantTool_.deQuantScaleCqGm_.SetGlobalBuffer((__gm__ dequantScaleType *)(workspace + workspaceOffset)); + workspaceOffset += baseParams_->stepBatchSize * dequantScaleCqSize_ * baseParams_->mm1BlockNum; + } + // query_norm_flag有影响 + mmCqResGm_.SetGlobalBuffer((__gm__ mmCqOutputType *)(workspace + workspaceOffset)); // aicOffset.cqResOffset + if (baseParams_->queryNormFlag == 0U) { + // 全量化场景下`mmCqResGm_`与`rmsNormCqResGm_`的dtype不同,无法共用workspace,此处需要偏移`mmCqResGm_`占用的大小; + if constexpr (!std::is_same::value) { + workspaceOffset += static_cast(baseParams_->stepBatchSize) * + static_cast(baseParams_->headSizeCq) * baseParams_->mm1BlockNum * + static_cast(sizeof(mmCqOutputType)); + } + // 不返回queryNorm时,rmsNormCq结果放在workspace + rmsNormCqResGm_.SetGlobalBuffer((__gm__ rmsNormCqOutputType *)(workspace + workspaceOffset)); + workspaceOffset += static_cast(baseParams_->stepBatchSize) * baseParams_->mm1BlockNum * + static_cast(baseParams_->headSizeCq) * sizeof(rmsNormCqOutputType); + } else { + // 返回queryNorm时,rmsNormCq结果直接输出gm,此处需要偏移`mmCqResGm_`占用的大小 + workspaceOffset += static_cast(baseParams_->stepBatchSize) * + static_cast(baseParams_->headSizeCq) * baseParams_->mm1BlockNum * + static_cast(sizeof(mmCqOutputType)); + } + + mmCkvKrResGm_.SetGlobalBuffer((__gm__ mmCkvKrOutputType *)(workspace + workspaceOffset)); + workspaceOffset += static_cast(baseParams_->stepBatchSize) * + static_cast(baseParams_->headSizeCkv + baseParams_->dimHeadRope) * + baseParams_->mm2BlockNum * sizeof(mmCkvKrOutputType); + + mmQcQrResGm_.SetGlobalBuffer((__gm__ mmQcQrOutputType *)(workspace + workspaceOffset)); // aicOffset.qcQrResOffset + + if constexpr (IsFullQuantMode()) { + workspaceOffset += static_cast(baseParams_->stepBatchSize) * + static_cast(baseParams_->headSizeQc + baseParams_->headSizeQr) * + baseParams_->mm3BlockNum * sizeof(mmQcQrOutputType); + } + + mmQcQrResDequantGm_.SetGlobalBuffer( + (__gm__ mmQnInputType *)(workspace + workspaceOffset)); // qcOffset mmQnPreDequantResOffset + workspaceOffset += + static_cast(baseParams_->stepBatchSize) * static_cast(baseParams_->numHeadSize) * + static_cast(baseParams_->dimHeadSizeQc) * baseParams_->mm4BlockNum * sizeof(mmQnInputType); + if constexpr (((std::is_same::value && std::is_same::value) || + (std::is_same::value && std::is_same::value) || + (std::is_same::value && std::is_same::value)) && + !isPertile) { + workspaceOffset += + static_cast(baseParams_->stepBatchSize) * static_cast(baseParams_->numHeadSize) * + static_cast(baseParams_->dimHeadSizeQc) * baseParams_->mm4BlockNum * sizeof(mmQnInputType); + mmQnResGm_.SetGlobalBuffer((__gm__ mmQnOutputType *)(workspace + workspaceOffset)); // aicOffset.qnResOffset + } +} + +template +__aicore__ inline void MlaPrologV3SplitM::UpdateStepBatchParams(int64_t curMSize) +{ + mmCqParam_.m = curMSize; + mmCkvKrParam_.m = curMSize; + mmQcQrParam_.m = curMSize; + mmQnParam_.m = curMSize; + curVectorBlockNum_ = curMSize; +} + +/* + * MlaPrologV3算子计算&CV流水同步流程 + * ┌───────────────── token_x ─────────────────┐ + * | ▼ + * | MatmulCkvKr + * ▼ | wait mm CkvKr(0x1) + * MatmulCq ▼ + * | wait mm Cq(0x1) ┌─────────────────┐ + * ▼ ▼ ▼ + * RmsNorm(Cq) RmsNorm(Ckv) Rope(Kr) + * | wait rmsNorm cq(0x1) | | + * ▼ ▼ ▼ + * ┌───────MatmulQcQr───────┐ Scatter(Ckv) Scatter(Kr) + * | wait mm Qc(0x1) | | | + * ▼ | ▼ ▼ + * DequantQc | kv_cache_out kr_cache_out + * | wait dequant qc(0x1) | wait mm Qr(0x2) + * ▼ ▼ + * MatmulQn Rope(Qr) + * | wait mm Qn(0x1) | + * DynamicQuantQn──┐ | + * | ▼ | + * ▼ dequant_scale_out ▼ + * query_out query_rope_out + * 注:仅为表明基本计算与CV同步流程,仅包含了影响CV同步的量化分支,其余量化分支应参考设计文档。 + */ +template +__aicore__ inline void MlaPrologV3SplitM::Process() +{ + constexpr bool needQnDynamicQuant = + ((std::is_same::value && std::is_same::value) || + (std::is_same::value && std::is_same::value) || + (std::is_same::value && std::is_same::value)) && + !isPertile; + int64_t numHeadOffset = 0; + int64_t mmQnLoops = baseParams_->mm4SingleCoreBatch; + + // moffsetStart_ 当前核处理的token起始偏移 + // mSubSize_ 当前核处理的token数量 + // curMsize 当前step内处理的token数量 + // maxMOffset 当前核处理的token结束偏移 + // mSubCoreNum 前多少个核需要多处理1个token + if (cubeBlockIdx_ < baseParams_->mSubCoreNum) { + mSubSize_ = baseParams_->mSubSize; + mOffsetStart_ = mSubSize_ * cubeBlockIdx_; + } else { + mSubSize_ = baseParams_->mSubSize - 1; + mOffsetStart_ = + baseParams_->mSubSize * baseParams_->mSubCoreNum + mSubSize_ * (cubeBlockIdx_ - baseParams_->mSubCoreNum); + } + + // AIC的offset参数 + AicOffset aicOffset; + + // AIV的offset参数 + AivOffset aivOffset; + + // 需要考虑BS合轴的尾块情况 + int64_t maxMOffset = mOffsetStart_ + mSubSize_; + for (int64_t mOffset = mOffsetStart_; mOffset < maxMOffset; mOffset += baseParams_->stepBatchSize) { + int64_t curMSize = + (maxMOffset - mOffset) < baseParams_->stepBatchSize ? (maxMOffset - mOffset) : baseParams_->stepBatchSize; + + UpdateStepBatchParams(curMSize); + // AIV核用MTE3 DMA清零对应的AIC workspace (UB→GM仅AIV的MTE3支持) + if ASCEND_IS_AIV { + if (cubeBlockIdx_ < baseParams_->mm1BlockNum && (blockIdx_ % cvRatio_ == 0)) { + int64_t stepBS = static_cast(baseParams_->stepBatchSize); + int64_t cbIdx = static_cast(cubeBlockIdx_); + LocalTensor zeroSrc = zeroBuf_.Get(); + constexpr int64_t zeroChunk = 512; + Duplicate(zeroSrc, static_cast(0), zeroChunk); + + int64_t cqN = stepBS * static_cast(baseParams_->headSizeCq); + int64_t ckvN = stepBS * static_cast(baseParams_->headSizeCkv + baseParams_->dimHeadRope); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + for (int64_t off = 0; off < cqN; off += zeroChunk) { + int64_t len = (cqN - off < zeroChunk) ? (cqN - off) : zeroChunk; + DataCopy(mmCqResGm_[cbIdx * cqN + off], zeroSrc, static_cast(len)); + } + for (int64_t off = 0; off < ckvN; off += zeroChunk) { + int64_t len = (ckvN - off < zeroChunk) ? (ckvN - off) : zeroChunk; + DataCopy(mmCkvKrResGm_[cbIdx * ckvN + off], zeroSrc, static_cast(len)); + } + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + } + CrossCoreSetFlag(SYNC_MODE_CUBE_VEC); + } + + if ASCEND_IS_AIC { + AicProcess(aicOffset, mOffset, mmQnLoops); + } + if ASCEND_IS_AIV { + AivProcess(aivOffset, mOffset, curMSize, numHeadOffset, mmQnLoops); + } + } + + if ASCEND_IS_AIC { + WaitFlag(A_EVENT0); + WaitFlag(A_EVENT1); + WaitFlag(B_EVENT0); + WaitFlag(B_EVENT1); + + WaitFlag(L0A_EVENT0); + WaitFlag(L0A_EVENT1); + WaitFlag(L0B_EVENT0); + WaitFlag(L0B_EVENT1); + + WaitFlag(L0C_EVENT0); + WaitFlag(L0C_EVENT1); + WaitFlag(SCALE_EVENT); + } +} + +template +template +__aicore__ inline void MlaPrologV3SplitM::AicProcess(AicOffset &aicOffset, int64_t mOffset, int64_t mmQnLoops) +{ + int64_t tokenXOffset = mOffset * static_cast(baseParams_->headSizeX); + int64_t dequantScaleXOffset = mOffset * static_cast(baseParams_->headSizeX) / 32; // mxfp8需要除32 + aicOffset.weightDqOffset = 0; + aicOffset.dequantScaleWDqOffset = 0; + aicOffset.cqResOffset = baseParams_->stepBatchSize * baseParams_->headSizeCq * cubeBlockIdx_; + + CrossCoreWaitFlag(SYNC_MODE_CUBE_VEC); + + // MatmulCq ──> RmsNorm(Cq) [K-outer: kL1→nb→stepK] + if constexpr (std::is_same::value && isFp8E8m0) { + if (cubeBlockIdx_ < baseParams_->mm1BlockNum) { + MatmulSplitMKOuter( + mmCqResGm_[aicOffset.cqResOffset], tokenXGm_[tokenXOffset], weightDqGm_[aicOffset.weightDqOffset], + mmCqParam_, bufParam_, dequantScaleXGm_[dequantScaleXOffset], + dequantScaleWDqGm_[aicOffset.dequantScaleWDqOffset]); + } + } + + CrossCoreSetFlag(FINISH_MM_CQ); + + aicOffset.weightDkvKrOffset = 0; + aicOffset.dequantScaleWDkvKrOffset = 0; + aicOffset.ckvKrResOffset = + baseParams_->stepBatchSize * (baseParams_->headSizeCkv + baseParams_->dimHeadRope) * cubeBlockIdx_; + + // MatmulCkvKr ──> RmsNorm(Ckv) / Rope(Kr) [K-outer: kL1→nb→stepK] + if constexpr (std::is_same::value && isFp8E8m0) { + if (cubeBlockIdx_ < baseParams_->mm2BlockNum) { + MatmulSplitMKOuter( + mmCkvKrResGm_[aicOffset.ckvKrResOffset], tokenXGm_[tokenXOffset], + weightDkvKrGm_[aicOffset.weightDkvKrOffset], mmCkvKrParam_, bufParam_, + dequantScaleXGm_[dequantScaleXOffset], dequantScaleWDkvkrGm_[aicOffset.dequantScaleWDkvKrOffset]); + } + } + + PipeBarrier(); + CrossCoreSetFlag(FINISH_MM_CKVKR); + CrossCoreWaitFlag(FINISH_VEC_RMSNORM_CQ); + + // 设置 MatmulQcQr 需要的 offset + if constexpr (std::is_same::value && isFp8E8m0) { + if (baseParams_->queryNormFlag == 1U) { + aicOffset.rmsNormCqResOffset = mOffset * baseParams_->headSizeCq; + } else { + aicOffset.rmsNormCqResOffset = baseParams_->stepBatchSize * baseParams_->headSizeCq * cubeBlockIdx_; + } + aicOffset.dequantScaleCqOffset = baseParams_->stepBatchSize * cubeBlockIdx_ * dequantScaleCqSize_; + aicOffset.weightUqQrOffset = 0; + aicOffset.dequantScaleWuqqrOffset = 0; + aicOffset.qcQrResOffset = + baseParams_->stepBatchSize * (baseParams_->headSizeQc + baseParams_->headSizeQr) * cubeBlockIdx_; + MatmulQcQr(aicOffset); + } + + if constexpr (std::is_same::value && isFp8E8m0) { + WaitFlag(SCALE_EVENT); + } + + if constexpr (!needQnDynamicQuant) { + aicOffset.qnResOffset = mOffset * baseParams_->headSizeCkv * baseParams_->numHeadSize; + } else { + aicOffset.qnResOffset = + baseParams_->stepBatchSize * baseParams_->numHeadSize * baseParams_->headSizeCkv * cubeBlockIdx_; + } + + aicOffset.qcOffset = + baseParams_->stepBatchSize * baseParams_->numHeadSize * baseParams_->dimHeadSizeQc * cubeBlockIdx_; + aicOffset.weightUkOffset = 0; + MatmulQnSyncDynamicQuantAndMulQr(aicOffset.qcOffset, aicOffset.weightUkOffset, + aicOffset.qnResOffset, mmQnLoops); + if constexpr (std::is_same::value && isFp8E8m0) { + SetFlag(SCALE_EVENT); + } +} + +template +template +__aicore__ inline void MlaPrologV3SplitM::AivProcess(AivOffset &aivOffset, int64_t mOffset, int64_t curMSize, + int64_t numHeadOffset, int64_t mmQnLoops) +{ + if (mOffset == mOffsetStart_) { + // 只需要搬运一次 + CopyGlobalParams(); + } + + uint64_t isBackVec = blockIdx_ % cvRatio_; + curStepVecBackToken_ = curMSize / cvRatio_; + curStepVecFrontListNum_ = curMSize % cvRatio_; + + aivOffset.curVecToken = + curStepVecBackToken_ + isBackVec * curStepVecFrontListNum_; // 当前核中每个vec处理多少个token + aivOffset.curBlockTokenOffset = isBackVec * curStepVecBackToken_; // 当前核中每个vec在这个m中的偏移 + + int64_t tokenIndex = mOffset + aivOffset.curBlockTokenOffset; + int64_t curTokenStart = mOffset + aivOffset.curBlockTokenOffset; + int64_t curTokenEnd = mOffset + aivOffset.curBlockTokenOffset + aivOffset.curVecToken; + + CopyInSinCos(curTokenStart, aivOffset.curVecToken, mOffset, curMSize); + + CrossCoreWaitFlag(FINISH_MM_CQ); + + aivOffset.rmsNormCqOffset = baseParams_->stepBatchSize * baseParams_->headSizeCq * cubeBlockIdx_ + + baseParams_->headSizeCq * aivOffset.curBlockTokenOffset; + int64_t rmsNormCqResOffset; + if (baseParams_->queryNormFlag == 1U) { + rmsNormCqResOffset = tokenIndex * baseParams_->headSizeCq; + } else { + rmsNormCqResOffset = aivOffset.rmsNormCqOffset; + } + RmsNormCq(curTokenStart, aivOffset.rmsNormCqOffset, rmsNormCqResOffset, aivOffset.curVecToken, + aivOffset.curBlockTokenOffset); + + CrossCoreSetFlag(FINISH_VEC_RMSNORM_CQ); + CrossCoreWaitFlag(FINISH_MM_CKVKR); + + aivOffset.rmsNormCkvOffset = + baseParams_->stepBatchSize * (baseParams_->headSizeCkv + baseParams_->dimHeadRope) * cubeBlockIdx_ + + (baseParams_->headSizeCkv + baseParams_->dimHeadRope) * aivOffset.curBlockTokenOffset; + + aivOffset.ropeKrOffset = baseParams_->headSizeCkv + aivOffset.rmsNormCkvOffset; // 512 + (512 + 64) * idx + RmsNormRopeScatterCkvKr(tokenIndex, aivOffset.rmsNormCkvOffset, aivOffset.ropeKrOffset, aivOffset.curVecToken); + + aivOffset.mmQnPreDequantOffset = + ((baseParams_->stepBatchSize) * (baseParams_->headSizeQc + baseParams_->headSizeQr) * cubeBlockIdx_) + + (baseParams_->headSizeQc + baseParams_->headSizeQr) * aivOffset.curBlockTokenOffset; + + aivOffset.mmQnPreDequantResOffset = + baseParams_->stepBatchSize * baseParams_->numHeadSize * baseParams_->dimHeadSizeQc * cubeBlockIdx_ + + baseParams_->dimHeadSizeQc * baseParams_->numHeadSize * aivOffset.curBlockTokenOffset; + aivOffset.ropeQrOffset = + (baseParams_->stepBatchSize * (baseParams_->headSizeQc + baseParams_->headSizeQr) * cubeBlockIdx_) + + baseParams_->dimHeadSizeQc + + (baseParams_->headSizeQc + baseParams_->headSizeQr) * aivOffset.curBlockTokenOffset; + aivOffset.ropeQrResOffset = (baseParams_->headSizeQr) * (aivOffset.curBlockTokenOffset + mOffset); + + QcQrSplit(aivOffset.curVecToken, aivOffset.curBlockTokenOffset, curMSize, aivOffset.mmQnPreDequantOffset, + aivOffset.mmQnPreDequantResOffset, aivOffset.ropeQrOffset, aivOffset.ropeQrResOffset); + + if constexpr (needQnDynamicQuant) { + DynamicQuantQnAndMulQrSyncMMQn(mOffset, curMSize, numHeadOffset, mmQnLoops); + } +} + +template +template +__aicore__ inline void +MlaPrologV3SplitM::MatmulSplitM(const GlobalTensor &tensorResGm, const GlobalTensor &tensorAGm, + const GlobalTensor &tensorBGm, const MMParams &mmPara, + const UsedBlockParams &mmBlockParams, const GlobalTensor &tensorAScaleGm, + const GlobalTensor &tensorBScaleGm) +{ + if constexpr (needCheckEmptyTensor) { + if constexpr (MLAPT::emptyMode == EMPTY_TENSOR_MODE::EMPTY_CACHE) { + return; + } + } + if (blockIdx_ < mmBlockParams.blockStartIdx || blockIdx_ >= mmBlockParams.blockEndIdx) { + return; + } + // 用于enableGroupComputeOpt场景 + if constexpr (needCheckAFullLoad) { + constexpr uint32_t mSize = + (sizeof(mmQcQrInputType) == sizeof(int8_t)) ? INT8_AFULLLOAD_MAX_MSIZE : BF16_AFULLLOAD_MAX_MSIZE; + bool isAFullLoad = (mmQcQrParam_.m <= mSize) ? true : false; + if (isAFullLoad) { + MatmulGroupComputeAFullLoad(tensorResGm, tensorAGm, tensorBGm, mmPara, + bufParam_); + return; + } + } + + uint32_t nInput = mmPara.n; + uint32_t nL1SplitSize = mmPara.baseN; + uint32_t nL1loops = CeilDivT(nInput, nL1SplitSize); + uint32_t subNL1SplitSize = nL1SplitSize; + for (int64_t nL1 = 0; nL1 < nL1loops; nL1++) { + if (nL1 == nL1loops - 1) { + subNL1SplitSize = nInput - (nL1loops - 1) * nL1SplitSize; + } + MatmulSplitK(tensorResGm, tensorAGm, tensorBGm, mmPara, bufParam_, nL1 * nL1SplitSize, subNL1SplitSize, + tensorAScaleGm, tensorBScaleGm); + } +} + +template +__aicore__ inline void MlaPrologV3SplitM::MatmulQcQr(AicOffset &aicOffset) +{ + if (blockIdx_ >= baseParams_->mm3BlockNum) { + return; + } + // RmsNorm(Cq) ──> MatmulQcQr ──> MatmulQn + // └──> Rope(Qr) + // [32, 1536] * [1536, 32*(128+64)] = [32, 32*192] + constexpr uint32_t mSize = + (sizeof(mmQcQrInputType) == sizeof(int8_t)) ? INT8_AFULLLOAD_MAX_MSIZE : BF16_AFULLLOAD_MAX_MSIZE; + bool isAFullLoad = (mmQcQrParam_.m <= mSize) ? true : false; + + uint32_t nInput = mmQcQrParam_.n; + uint32_t nL1SplitSize = mmQcQrParam_.baseN; + uint32_t nL1loops = CeilDivT(nInput, nL1SplitSize); + uint32_t subNL1SplitSize = nL1SplitSize; + if (isAFullLoad) { + if constexpr (std::is_same::value && isFp8E8m0) { + uint32_t offsetL1B = L1_B_SIZE / 2 / sizeof(rmsNormCqOutputType); + LoadL1AAndScale( + rmsNormCqResGm_[aicOffset.rmsNormCqResOffset], + dequantTool_.deQuantScaleCqGm_[aicOffset.dequantScaleCqOffset], mmQcQrParam_.m, mmQcQrParam_.k, + mmQcQrParam_.k, mmQcQrParam_.kScale, offsetL1B, bufParam_); + } + WaitFlag(A_EVENT0 + (bufParam_.aL1BufIter & 1u)); + } + + for (int64_t nL1 = 0; nL1 < nL1loops; nL1++) { + if (nL1 == nL1loops - 1) { + subNL1SplitSize = nInput - (nL1loops - 1) * nL1SplitSize; + } + if constexpr (std::is_same::value && isFp8E8m0) { + if (isAFullLoad) { + MatmulSplitK( + mmQcQrResGm_[aicOffset.qcQrResOffset], rmsNormCqResGm_[aicOffset.rmsNormCqResOffset], + weightUqQrGm_[aicOffset.weightUqQrOffset], mmQcQrParam_, bufParam_, nL1 * nL1SplitSize, + subNL1SplitSize, dequantTool_.deQuantScaleCqGm_[aicOffset.dequantScaleCqOffset], + deqScaleQcQrW_[aicOffset.dequantScaleWuqqrOffset]); + } else { + MatmulSplitK( + mmQcQrResGm_[aicOffset.qcQrResOffset], rmsNormCqResGm_[aicOffset.rmsNormCqResOffset], + weightUqQrGm_[aicOffset.weightUqQrOffset], mmQcQrParam_, bufParam_, nL1 * nL1SplitSize, + subNL1SplitSize, dequantTool_.deQuantScaleCqGm_[aicOffset.dequantScaleCqOffset], + deqScaleQcQrW_[aicOffset.dequantScaleWuqqrOffset]); + } + } + if constexpr (MLAPT::enableDequantOpt) { + uint32_t spacing = CeilDivT(nL1loops, static_cast(MAX_SYNC_FLAG_COUNT - 1)); + if (nL1 % spacing == 0 || (nL1 + 1) == static_cast(nL1loops)) { + CrossCoreSetFlag<0x2, PIPE_FIX>(FINISH_MM_QCQR_SPLIT_N); + } + if ((nL1 + 1) % spacing == 0 || (nL1 + 1) == static_cast(nL1loops)) { + CrossCoreSetFlag<0x2, PIPE_FIX>(FINISH_MM_QCQR_SPLIT_BATCH); + } + } + } + if (isAFullLoad) { + SetFlag(A_EVENT0 + (bufParam_.aL1BufIter & 1u)); + bufParam_.aL1BufIter++; + } +} + +template +template +__aicore__ inline void +MlaPrologV3SplitM::MatmulQnSyncDynamicQuantAndMulQr(int64_t qcOffset, int64_t weightUkOffset, + int64_t qnResOffset, int64_t subLoopTimes) +{ + if (blockIdx_ >= baseParams_->mm4BlockNum) { + return; + } + + // MatmulQcQr ──> MatmulQn ──> query_out + // [32, 128] * [128, 512] = [32, 512] + // [32, 2, 128] * [2, 128, 512] = [32, 2, 512] + + for (int64_t i = 0; i < subLoopTimes; i++) { + uint32_t coarse = CeilDivT(static_cast(subLoopTimes), static_cast(MAX_SYNC_FLAG_COUNT - 1)); + if (i % coarse == 0) { + CrossCoreWaitFlag(FINISH_VEC_DEQUANT_QC_SPLIT_N); + } else if ((subLoopTimes - CeilDivT(static_cast(subLoopTimes), coarse)) <= MAX_SYNC_FLAG_COUNT) { + CrossCoreWaitFlag(FINISH_VEC_DEQUANT_QC_SPLIT_N_GAP); + } + + // Qn 用 MatmulSplitK 避免 L0A 溢出:每 L0 步 128×64=8192<16384,bf16 无 scale + for (int64_t nL1 = 0; nL1 < mmQnParam_.n; nL1 += mmQnParam_.baseN) { + uint32_t subN = static_cast((nL1 + mmQnParam_.baseN > mmQnParam_.n) ? (mmQnParam_.n - nL1) : + mmQnParam_.baseN); + MatmulSplitK( + mmQnResGm_[qnResOffset], mmQcQrResDequantGm_[qcOffset], weightUkGm_[weightUkOffset], mmQnParam_, + bufParam_, static_cast(nL1), subN); + } + + if constexpr (std::is_same::value || std::is_same::value || + std::is_same::value) { + qcOffset += static_cast(baseParams_->dimHeadSizeQc); + } + weightUkOffset += + static_cast(baseParams_->dimHeadSizeQc) * static_cast(baseParams_->headSizeCkv); + qnResOffset += static_cast(baseParams_->headSizeCkv); + } + PipeBarrier(); + if constexpr (needQnDynamicQuant) { + CrossCoreSetFlag(FINISH_MM_QN_SPLIT_N); + } +} + +template +__aicore__ inline void MlaPrologV3SplitM::CopyInSinCos(int64_t tokenIndex, int64_t curVecToken, + int64_t batchOffset, int64_t curMSize) +{ + if constexpr (!MLAPT::enableRope) { + return; + } + LocalTensor shareTmpUb = shareBuffer_.Get(); + if constexpr (MLAPT::enableDequantOpt) { + // 如果是切N场景,mm3的每个C核都会做rope + if (cubeBlockIdx_ >= baseParams_->mm3BlockNum) { + return; + } + // 如果curStepBatchSize是偶数,则两个核平分;如果curStepBatchSize是奇数,则奇数核比偶数核多分一个 + // >> 1 是将curStepBatchSize分到每个vec核上; + uint32_t subBlockIdx_ = blockIdx_ % cvRatio_; + int64_t offset = (curMSize / cvRatio_) * subBlockIdx_ + batchOffset; + GatherSinCos(cosLocal_, sinLocal_, ropeCosGm_, ropeSinGm_, offset, + (curMSize + cvRatio_ - 1) / cvRatio_, shareTmpUb, vectorRow_, + baseParams_->dimHeadRope); + } else { + if (cubeBlockIdx_ >= baseParams_->mm3BlockNum) { + return; + } + GatherSinCos(cosLocal_, sinLocal_, ropeCosGm_, ropeSinGm_, tokenIndex, + curVecToken, shareTmpUb, vectorRow_, baseParams_->dimHeadRope); + } +} + +template +__aicore__ inline void MlaPrologV3SplitM::CopyGlobalParams() +{ + // rmsnormGammaCq + DataCopy(rmsnormGammaCqLocal_, rmsnormGammaCqGm_, baseParams_->headSizeCq); + + // rmsnormGammaCkv + DataCopy(rmsnormGammaCkvLocal_, rmsnormGammaCkvGm_, baseParams_->headSizeCkv); + + // quantScaleCkv + + if constexpr (IsFullQuantMode() && !isPertile) { + if constexpr (std::is_same::value || + (std::is_same::value && !isFp8E8m0)) { + DataCopyExtParams quantCopyParams{1, sizeof(float), 0, 0, 0}; + DataCopyPadExtParams quantPadParams{false, 0, 0, 0}; + DataCopyPad(quantScaleCkvLocal_, quantScaleCkvGm_, quantCopyParams, quantPadParams); + } else { + DataCopy(quantScaleCkvLocal_, quantScaleCkvGm_, baseParams_->headSizeCkv); + } + } + + // quantScaleCkr + if constexpr (std::is_same::value) { + DataCopyExtParams quantCopyParams{1, static_cast(baseParams_->dimHeadRope * sizeof(float)), 0, 0, 0}; + DataCopyPadExtParams quantPadParams{false, 0, 0, 0}; + DataCopyPad(quantScaleCkrLocal_, quantScaleCkrGm_, quantCopyParams, quantPadParams); + } +} + + +/** + * @brief RmsNormCq流程,融合了dynamicquant + 内部所需空间约为 curVecToken(128) * 8*4 + 8*4 + (4*vectorRow_*baseParams_->headSizeCq + 8)*4 = 28.0625K + */ +template +__aicore__ inline void MlaPrologV3SplitM::RmsNormCq(int64_t tokenIndex, int64_t rmsNormCqOffset, + int64_t rmsNormCqResOffset, int64_t curVecToken, + int64_t curBlockTokenOffset) +{ + if (blockIdx_ < 0 || blockIdx_ >= cvRatio_ * baseParams_->mm1BlockNum) { + return; + } + + uint64_t dequantScaleCqElementNum = dequantScaleCqSize_ / sizeof(dequantScaleType); + LocalTensor outputLocal = shareBuffer_.Get(); + LocalTensor dequantScaleQcQr = + outputLocal[baseParams_->headSizeCq].template ReinterpretCast(); + LocalTensor dequantScaleXLocal = + dequantScaleQcQr[curVecToken * dequantScaleCqElementNum].template ReinterpretCast(); + + uint64_t dequantScaleXSize = 1; + if constexpr (std::is_same::value && isFp8E8m0) { + dequantScaleXSize = baseParams_->headSizeX / FP8_E4M3_BLOCK_SIZE; + } + uint64_t dequantScaleXElementNum = Align(dequantScaleXSize, BYTE_BLOCK) / sizeof(float); + LocalTensor shareTmpUb = dequantScaleXLocal[dequantScaleXElementNum].template ReinterpretCast(); + int64_t stepTokenIndex = tokenIndex; + for (int64_t curVecTokenIdx = 0; curVecTokenIdx < curVecToken; curVecTokenIdx++) { + // MatmulCq ──> RmsNorm(Cq) ──> MatmulQcQr + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); // wait for vector operations to finish + uint64_t scaleOffset; + if constexpr (std::is_same::value && isFp8E8m0) { + scaleOffset = curVecTokenIdx * dequantScaleCqElementNum; + } else { + scaleOffset = curVecTokenIdx * FP32_BLOCK_ELEMENT_NUM; + } + RmsNormParam rmsNormParams = {baseParams_->reciprocalCq, // reciprocal + baseParams_->epsilonCq, // epsilon + static_cast(vectorRow_), // row + baseParams_->headSizeCq, // col + baseParams_->qcQrScale, + baseParams_->isQcQrScaleEnable}; + + + if constexpr (IsFullQuantMode()) { + RmsNormDynamicQuant(outputLocal, dequantScaleQcQr[scaleOffset], + mmCqResGm_[rmsNormCqOffset], rmsnormGammaCqLocal_, + smoothScaleCqLocal_, dequantScaleWDqLocal_, dequantScaleXLocal, + shareTmpUb, rmsNormParams, enableSmoothScalesCq_); + } + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + // RmsNorm(Cq)的结果拷进mmCqResGm_中,用于MatmulQcQr的A矩阵 + DataCopy(rmsNormCqResGm_[rmsNormCqResOffset], outputLocal, baseParams_->headSizeCq); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + rmsNormCqOffset += baseParams_->headSizeCq; + rmsNormCqResOffset += baseParams_->headSizeCq; + tokenIndex++; + } + + if constexpr (std::is_same::value && isFp8E8m0) { + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + DataCopy(dequantTool_.deQuantScaleCqGm_[(curBlockTokenOffset + baseParams_->stepBatchSize * cubeBlockIdx_) * + dequantScaleCqElementNum], + dequantScaleQcQr, curVecToken * dequantScaleCqElementNum); + } + if (unlikely(baseParams_->queryNormFlag == 1U)) { + if constexpr (std::is_same::value && isFp8E8m0) { + DataCopyPad( + dequantScaleQNormGm_[stepTokenIndex * + static_cast((baseParams_->headSizeCq / FP8_E4M3_BLOCK_SIZE))], + dequantScaleQcQr, + {static_cast(curVecToken), + static_cast(sizeof(dequantScaleQNormType) * (baseParams_->headSizeCq / FP8_E4M3_BLOCK_SIZE)), + 0, 0}); + } + } +} + +template +__aicore__ inline void MlaPrologV3SplitM::ComputeBlkScatterOffsets(GlobalTensor indexGm, + int64_t tokenIndex, int64_t rows, + CkvkrParams &rmsNormAndScatterCkvParams, + CkvkrParams &ropeAndScatterKrParams) +{ + int64_t batchTokenIndex = 0; + int64_t batchIndexOffset = 0; + int64_t batchSeqSize = 0; + int64_t rowsThisStep = 0; + int64_t nextPageId = 0; + bool spill = false; + + if constexpr (MLAPT::actualSeqMode == ACTUAL_SEQ_MODE::DISABLED) { + const int64_t batchIndex = tokenIndex / baseParams_->seq1Size; + batchTokenIndex = tokenIndex % baseParams_->seq1Size; + batchSeqSize = baseParams_->seq1Size; + batchIndexOffset = batchIndex * CeilDivT(batchSeqSize, static_cast(baseParams_->blockSize)); + } else { + while (actSeqState_.curBatch + 1 < static_cast(baseParams_->batchSize) && + tokenIndex >= actSeqState_.curPrefix) { + actSeqState_.curBatch += 1; + actSeqState_.prevPrefix = actSeqState_.curPrefix; + actSeqState_.curPrefix = actualSeqLenGm_(actSeqState_.curBatch); + actSeqState_.prevIndexOffset = actSeqState_.curIndexOffset; + actSeqState_.curIndexOffset += CeilDivT(actSeqState_.curPrefix - actSeqState_.prevPrefix, + static_cast(baseParams_->blockSize)); + } + batchTokenIndex = tokenIndex - actSeqState_.prevPrefix; + batchSeqSize = actSeqState_.curPrefix - actSeqState_.prevPrefix; + batchIndexOffset = actSeqState_.prevIndexOffset; + } + + int64_t indexOffset = batchIndexOffset + batchTokenIndex / baseParams_->blockSize; + int64_t paBlkId = indexGm(indexOffset); + + int64_t pageTokenOffset = paBlkId * baseParams_->blockSize; + int64_t tokenOffsetInPage = batchTokenIndex % baseParams_->blockSize; + + int64_t leftRowsInPage = baseParams_->blockSize - tokenOffsetInPage; + if (leftRowsInPage >= rows) { + rowsThisStep = rows; + spill = false; + nextPageId = -1; + } else { + rowsThisStep = rows - leftRowsInPage; + spill = true; + nextPageId = indexGm(indexOffset + 1); + } + + MaterializeOffsetsWithHeadSize( + pageTokenOffset, tokenOffsetInPage, rowsThisStep, spill, nextPageId, baseParams_->headSizeCkv, + rmsNormAndScatterCkvParams); + + MaterializeOffsetsWithHeadSize( + pageTokenOffset, tokenOffsetInPage, rowsThisStep, spill, nextPageId, baseParams_->dimHeadRope, + ropeAndScatterKrParams); +} + +template +__aicore__ inline void MlaPrologV3SplitM::RmsNormAndScatterCkv(LocalTensor &dequantScaleXLocal, + LocalTensor &shareTmpUb, + LocalTensor &cosLocalCkvKr, + LocalTensor &sinLocalCkvKr, + CkvkrParams rmsNormAndScatterCkvParams) +{ + // MatmulCkvKr ──> RmsNorm(Ckv) ──> Scatter(Ckv) + LocalTensor outputLocal = shareTmpUb.ReinterpretCast(); + int32_t outputLocalSize = vectorRow_ * baseParams_->headSizeCkv * sizeof(kvCacheType); + if constexpr (isPertile) { // pertile量化场景,按照concat的最长长度申请内存 + int32_t tileNum = baseParams_->headSizeCkv / baseParams_->tileSize; + outputLocalSize = vectorRow_ * (baseParams_->headSizeCkv * sizeof(kvCacheType) + tileNum * sizeof(float)); + outputLocalSize = Align(outputLocalSize, static_cast(BYTE_BLOCK)); + } + LocalTensor rmsNormShareTmpUb = shareTmpUb[outputLocalSize].template ReinterpretCast(); + + RmsNormParam rmsNormParams = { + baseParams_->reciprocalCkv, // reciprocal + baseParams_->epsilonCkv, // epsilon + static_cast(vectorRow_), // row + baseParams_->headSizeCkv, // col + baseParams_->kcScale, // scale + baseParams_->isKcScaleEnable, // isScaleEnable + }; + + if constexpr (IsFullQuantMode()) { + RmsNormAndQuantizeCkv(outputLocal, rmsNormShareTmpUb, dequantScaleXLocal, rmsNormParams, + rmsNormAndScatterCkvParams); + } else { + RmsNormNormal( + outputLocal, mmCkvKrResGm_[rmsNormAndScatterCkvParams.offset], rmsnormGammaCkvLocal_, + dequantScaleWDkvKrLocal_, dequantScaleXLocal, rmsNormShareTmpUb, rmsNormParams); + } + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + // Scatter(Ckv) + // RmsNorm(Ckv) ──> Scatter(Ckv) ──> kv_cache_out + ScatterCkv(outputLocal, rmsNormAndScatterCkvParams); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); +} + +template +__aicore__ inline void MlaPrologV3SplitM::RmsNormAndQuantizeCkv(LocalTensor &outputLocal, + LocalTensor &rmsNormShareTmpUb, + LocalTensor &dequantScaleXLocal, + RmsNormParam rmsNormParams, + CkvkrParams rmsNormAndScatterCkvParams) +{ + // row = vectorRow_ = 1 col = baseParams_->headSizeCkv + LocalTensor inputLocal = rmsNormShareTmpUb.ReinterpretCast(); + LocalTensor sharedBuf = + inputLocal[vectorRow_ * baseParams_->headSizeCkv].template ReinterpretCast(); + RmsNormNormal( + inputLocal, mmCkvKrResGm_[rmsNormAndScatterCkvParams.offset], rmsnormGammaCkvLocal_, dequantScaleWDkvKrLocal_, + dequantScaleXLocal, sharedBuf, rmsNormParams); + + DataSyncBarrier(); + + if constexpr (isPertile) { + float kNopeClipAlpha = isFp8E8m0 ? 1.0f : kNopeClipAlphaGm_.GetValue(0); + PerTileQuantParams perTileQuantParams = { + static_cast(baseParams_->tileSize), // baseParams_->tileSize + static_cast(baseParams_->headSizeCkv / baseParams_->tileSize), // tileNum + kNopeClipAlpha, // alpha + static_cast(vectorRow_), // row + baseParams_->headSizeCkv // col + }; + if constexpr (std::is_same::value) { + QuantPerTile(outputLocal, inputLocal, sharedBuf, perTileQuantParams); + } else { + QuantPerTile8Bit(outputLocal, inputLocal, perTileQuantParams); + } + } else { + Rectangle rectangleParams{ + static_cast(vectorRow_), // row + static_cast(baseParams_->headSizeCkv), // col + static_cast(baseParams_->headSizeCkv) // columnStride + }; + if constexpr (std::is_same::value || + std::is_same::value) { + QuantPerTensor(outputLocal, inputLocal, quantScaleCkvLocal_, sharedBuf, rectangleParams); + } else { + QuantPerChannel(outputLocal, inputLocal, quantScaleCkvLocal_, sharedBuf, rectangleParams); + } + } + PipeBarrier(); +} + +template +__aicore__ inline void MlaPrologV3SplitM::ScatterCkv(LocalTensor &outputLocal, + CkvkrParams rmsNormAndScatterCkvParams) +{ + if constexpr ((MLAPT::cacheMode == CACHE_MODE::PA_NZ) || (MLAPT::cacheMode == CACHE_MODE::PA_BSND) || + (MLAPT::cacheMode == CACHE_MODE::ND)) { + int64_t paTokenIndex; + if constexpr (MLAPT::cacheMode == CACHE_MODE::ND) { + paTokenIndex = rmsNormAndScatterCkvParams.tokenIndex; + } else { + paTokenIndex = cacheIndexGm_(rmsNormAndScatterCkvParams.tokenIndex); + } + ScatterCache( + kvCacheGm_, outputLocal, + ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, baseParams_->headSizeCkv, + baseParams_->dtileSize}); + // 刷新量化scale + if (isPertile && baseParams_->quantScaleRepoMode == 1U) { + // BSND: + int32_t tileNum = baseParams_->headSizeCkv / baseParams_->tileSize; + LocalTensor quantScaleCkvInt8Tensor = + outputLocal[vectorRow_ * baseParams_->headSizeCkv].template ReinterpretCast(); + int64_t startOffset = 0; + int64_t startColOffset = baseParams_->headSizeCkv; + if (baseParams_->ckvkrRepoMode == 1U) { + startColOffset += baseParams_->headSizeKr * sizeof(krCacheType); + } + if constexpr ((MLAPT::cacheMode != CACHE_MODE::PA_NZ)) { + startOffset = startColOffset; + } + + ScatterCacheUnAligned( + kvCacheGm_[startOffset], quantScaleCkvInt8Tensor, + ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, + static_cast(tileNum * sizeof(float)), baseParams_->dtileSize}); + } + } else { + // 使用 已计算好的偏移参数 + ScatterCacheMultiRows( + kvCacheGm_, outputLocal, + ScatterCacheParams{baseParams_->blockSize, rmsNormAndScatterCkvParams.cacheOffset, vectorRow_, + baseParams_->headSizeCkv, baseParams_->seq1Size, rmsNormAndScatterCkvParams.tokenIndex}, + rmsNormAndScatterCkvParams.rowsInCurBatch, rmsNormAndScatterCkvParams.cacheOffset, + rmsNormAndScatterCkvParams.nextBatchOffset); + } +} + +template +__aicore__ inline void MlaPrologV3SplitM::RopeAndScatterKr(LocalTensor &dequantScaleXLocal, + LocalTensor &shareTmpUb, + LocalTensor &cosLocalCkvKr, + LocalTensor &sinLocalCkvKr, + CkvkrParams ropeAndScatterKrParams) +{ + // MatmulCkvKr ──> Rope(Ckv) ──> Scatter(Kr) + LocalTensor outputKrLocal = shareTmpUb.ReinterpretCast(); + LocalTensor ropeShareTmpUb = outputKrLocal[baseParams_->dimHeadRope].template ReinterpretCast(); + int64_t stride = baseParams_->headSizeCkv + baseParams_->headSizeKr; // 512 + 64 + + LocalTensor cosLocal; + LocalTensor sinLocal; + if constexpr (MLAPT::enableDequantOpt) { + cosLocal = cosLocalCkvKr[baseParams_->dimHeadRope * ropeAndScatterKrParams.curVecTokenIdx]; + sinLocal = sinLocalCkvKr[baseParams_->dimHeadRope * ropeAndScatterKrParams.curVecTokenIdx]; + } else { + cosLocal = cosLocal_[baseParams_->dimHeadRope * ropeAndScatterKrParams.curVecTokenIdx]; + sinLocal = sinLocal_[baseParams_->dimHeadRope * ropeAndScatterKrParams.curVecTokenIdx]; + } + Rectangle ropeParams{ + static_cast(vectorRow_), // row + static_cast(baseParams_->dimHeadRope), // col + static_cast(stride) // stride + }; + + if constexpr ((std::is_same::value || + (std::is_same::value && !isFp8E8m0)) && + std::is_same::value) { + // Use ropeShareTmpUb directly (same as RopeQr full-quant). cos/sin live in cosLocal/sinLocal, + // not in this temp buffer; the old dimHeadRope*sizeof(ropeSinCosType) skip caused + // ENABLE_ROPE=0 write-through to read/write the wrong UB region for Kr dequant. + LocalTensor sharedBuf = ropeShareTmpUb; + // input为int32_t需在rope中做反量化, intput为float根据模板参数判断是否做反量化 + RotaryPosEmbPerTensor( + outputKrLocal, mmCkvKrResGm_[ropeAndScatterKrParams.offset], cosLocal, sinLocal, sharedBuf, ropeParams, + dequantScaleWDkvKrLocal_[baseParams_->headSizeCkv], dequantScaleXLocal); + } else if constexpr (std::is_same::value) { + LocalTensor inputLocal = ropeShareTmpUb.ReinterpretCast(); + LocalTensor sharedBuf = + ropeShareTmpUb.ReinterpretCast()[baseParams_->dimHeadRope * sizeof(ropeSinCosType)]; + RotaryPosEmbPerTensor::value, MLAPT::enableRope>( + inputLocal, mmCkvKrResGm_[ropeAndScatterKrParams.offset], cosLocal, sinLocal, sharedBuf, ropeParams); + RopePostQuantPerChannel(outputKrLocal, inputLocal, quantScaleCkrLocal_, sharedBuf, + vectorRow_ * baseParams_->dimHeadRope); + } else { + RotaryPosEmbPerTensor( + outputKrLocal, mmCkvKrResGm_[ropeAndScatterKrParams.offset], cosLocal, sinLocal, ropeShareTmpUb, + ropeParams); + } + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + // scatter(Kr) + // Rope(Kr) ──> Scatter(Kr) ──> kr_cache_out + ScatterKr(outputKrLocal, ropeAndScatterKrParams); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); +} + +template +__aicore__ inline void MlaPrologV3SplitM::ScatterKr(LocalTensor &outputKrLocal, + CkvkrParams ropeAndScatterKrParams) +{ + int64_t paTokenIndex; + if constexpr ((MLAPT::cacheMode == CACHE_MODE::PA_NZ) || (MLAPT::cacheMode == CACHE_MODE::PA_BSND) || + (MLAPT::cacheMode == CACHE_MODE::ND)) { + paTokenIndex = cacheIndexGm_(ropeAndScatterKrParams.tokenIndex); + if constexpr (MLAPT::cacheMode == CACHE_MODE::ND) { + paTokenIndex = ropeAndScatterKrParams.tokenIndex; + } + if (isPertile && baseParams_->ckvkrRepoMode == 1) { + int64_t startOffset; + if constexpr ((MLAPT::cacheMode == CACHE_MODE::PA_NZ)) { + constexpr uint8_t col0 = ALIGN_BLOCK_SIZE / sizeof(kvCacheType); + startOffset = CeilDiv(baseParams_->headSizeCkv, col0) * col0 * baseParams_->blockSize; // 列方向的偏移 + } else { + startOffset = baseParams_->headSizeCkv; + } + LocalTensor outputKrInt8Tensor = outputKrLocal.template ReinterpretCast(); + ScatterCache( + kvCacheGm_[startOffset], outputKrInt8Tensor, + ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, + static_cast(baseParams_->dimHeadRope * sizeof(krCacheType)), + baseParams_->dtileSize}); + } else { + ScatterCache( + krCacheGm_, outputKrLocal, + ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, baseParams_->dimHeadRope, + baseParams_->dimHeadRope}); + } + } else if constexpr ((MLAPT::cacheMode == CACHE_MODE::PA_BLK_BSND) || (MLAPT::cacheMode == CACHE_MODE::PA_BLK_NZ)) { + // 使用 已计算好的偏移参数 + ScatterCacheMultiRows( + krCacheGm_, outputKrLocal, + ScatterCacheParams{baseParams_->blockSize, ropeAndScatterKrParams.cacheOffset, vectorRow_, + baseParams_->dimHeadRope, baseParams_->seq1Size, ropeAndScatterKrParams.tokenIndex}, + ropeAndScatterKrParams.rowsInCurBatch, ropeAndScatterKrParams.cacheOffset, + ropeAndScatterKrParams.nextBatchOffset); + } else { + int64_t batchTokenIndex = 0; + int64_t batchIndex = 0; + int64_t batchSeqSize = 0; + int64_t batchIndexOffset = 0; + int64_t rowsInCurBatch = 0; + int64_t cacheOffset = 0; + int64_t nextBatchOffset = 0; + } +} + +template +__aicore__ inline void MlaPrologV3SplitM::RmsNormRopeScatterCkvKr(int64_t tokenIndex, int64_t rmsNormCkvOffset, + int64_t ropeKrOffset, int64_t curVecToken) +{ + if (blockIdx_ < 0 || blockIdx_ >= cvRatio_ * baseParams_->mm2BlockNum) { + return; + } + LocalTensor dequantScaleXLocal = shareBuffer_.Get(); + LocalTensor cosLocalCkvKr = + dequantScaleXLocal[FP32_BLOCK_ELEMENT_NUM].template ReinterpretCast(); + LocalTensor sinLocalCkvKr = cosLocalCkvKr[baseParams_->dimHeadRope * curVecToken]; + LocalTensor shareTmpUb = + sinLocalCkvKr[baseParams_->dimHeadRope * curVecToken].template ReinterpretCast(); + if constexpr (MLAPT::enableDequantOpt && MLAPT::enableRope) { + GatherSinCos(cosLocalCkvKr, sinLocalCkvKr, ropeCosGm_, ropeSinGm_, tokenIndex, + curVecToken, shareTmpUb, vectorRow_, baseParams_->dimHeadRope); + } + if constexpr (MLAPT::emptyMode == EMPTY_TENSOR_MODE::EMPTY_CACHE) { + return; + } + + for (int64_t curVecTokenIdx = 0; curVecTokenIdx < curVecToken; curVecTokenIdx++) { + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + if constexpr (std::is_same::value || + (std::is_same::value && !isFp8E8m0)) { + DataCopyPad(dequantScaleXLocal, dequantScaleXGm_[tokenIndex], {1, sizeof(float), 0, 0}, {false, 0, 0, 0}); + } + + CkvkrParams rmsNormAndScatterCkvParams{tokenIndex, rmsNormCkvOffset, curVecTokenIdx, 0, 0, 0}; + + CkvkrParams ropeAndScatterKrParams{tokenIndex, ropeKrOffset, curVecTokenIdx, 0, 0, 0}; + + if constexpr ((MLAPT::cacheMode == CACHE_MODE::PA_BLK_BSND) || (MLAPT::cacheMode == CACHE_MODE::PA_BLK_NZ)) { + ComputeBlkScatterOffsets(cacheIndexGm_, tokenIndex, vectorRow_, rmsNormAndScatterCkvParams, + ropeAndScatterKrParams); + } + + RmsNormAndScatterCkv(dequantScaleXLocal, shareTmpUb, cosLocalCkvKr, sinLocalCkvKr, rmsNormAndScatterCkvParams); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + RopeAndScatterKr(dequantScaleXLocal, shareTmpUb, cosLocalCkvKr, sinLocalCkvKr, ropeAndScatterKrParams); + + tokenIndex += 1; + rmsNormCkvOffset += static_cast(baseParams_->headSizeCkv + baseParams_->dimHeadRope); + ropeKrOffset += static_cast(baseParams_->headSizeCkv + baseParams_->dimHeadRope); + } +} + +template +__aicore__ inline void MlaPrologV3SplitM::QcQrSplit(int64_t curVecToken, int64_t curBlockTokenOffset, + int64_t curMSize, int64_t mmQcQrOffset, + int64_t mmQnPreDequantResOffset, int64_t ropeQrOffset, + int64_t ropeQrResOffset) +{ + if (cubeBlockIdx_ >= baseParams_->mm3BlockNum) { + return; + } + + DataCopyParams inputCopyParams{ + static_cast(curVecToken), + static_cast(baseParams_->dimHeadSizeQc * sizeof(mmQcQrOutputType) / ALIGN_BLOCK_SIZE), + static_cast((baseParams_->headSizeQc + baseParams_->headSizeQr - baseParams_->dimHeadSizeQc) * + sizeof(mmQcQrOutputType) / ALIGN_BLOCK_SIZE), + 0}; + + LocalTensor shareTmpUb = shareBuffer_.Get(); + LocalTensor inputLocal = shareTmpUb.ReinterpretCast(); + LocalTensor outputLocal = inputLocal.template ReinterpretCast(); + + Rectangle dequantParams{ + static_cast(curVecToken), // row + baseParams_->dimHeadSizeQc, // col + baseParams_->dimHeadSizeQc // columnStride + }; + DataCopyParams outputCopyParams{ + static_cast(curVecToken), + static_cast(baseParams_->dimHeadSizeQc * sizeof(mmQnInputType) / ALIGN_BLOCK_SIZE), 0, + static_cast((baseParams_->headSizeQc - baseParams_->dimHeadSizeQc) * sizeof(mmQnInputType) / + ALIGN_BLOCK_SIZE)}; + + DataCopyParams outputRopeParams{ + static_cast(curVecToken), + static_cast(baseParams_->dimHeadRope * sizeof(ropeOutputType) / ALIGN_BLOCK_SIZE), 0, + static_cast((baseParams_->headSizeQr - baseParams_->dimHeadRope) * sizeof(ropeOutputType) / + ALIGN_BLOCK_SIZE)}; + + Rectangle ropeParams{ + static_cast(curVecToken), // row + baseParams_->dimHeadRope, // col + static_cast(baseParams_->numHeadSize * + (baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope)) // stride + }; + int64_t ropeStride = static_cast(baseParams_->numHeadSize) * + static_cast(baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope); + uint32_t deqScaleOffset = 0; + + int64_t deqScaleQcQrWOffset = 0; + uint32_t colOffsetCube = 0; + uint32_t colOffsetVec = 0; + uint32_t colOffsetVecEnd = 0; + uint32_t colOffsetRope = 0; + // cube一次处理row*colCube,对应的两个vec一次处理row*colQc,两vec之间切row + // 等cube生产足够数据了以后,vec开始消费 + uint32_t qcCount = 0; + uint32_t splitCount = 0; + uint32_t totalLoops = CeilDivT(mmQcQrParam_.n, mmQcQrParam_.baseN); + uint32_t totalQcLoops = CeilDivT(mmQcQrParam_.n, (baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope)); + + for (uint32_t colOffsetCube = 0; colOffsetCube < mmQcQrParam_.n; colOffsetCube += mmQcQrParam_.baseN) { + colOffsetVecEnd = colOffsetCube + mmQcQrParam_.baseN; + if (colOffsetVecEnd > mmQcQrParam_.n) { // 当oriCol不被colCube整除时,mm最后一个base块需要刷新col end + colOffsetVecEnd = mmQcQrParam_.n; + } + uint32_t qcQrSpacing = CeilDivT(totalLoops, static_cast(MAX_SYNC_FLAG_COUNT - 1)); + if (splitCount % qcQrSpacing == 0 || splitCount == totalLoops - 1) { + CrossCoreWaitFlag(FINISH_MM_QCQR_SPLIT_N); + } + if (splitCount % qcQrSpacing == 0) { + CrossCoreWaitFlag(FINISH_MM_QCQR_SPLIT_BATCH); + } + while (colOffsetVec + baseParams_->dimHeadSizeQc <= colOffsetVecEnd) { // 循环singleNumHeadSize次 + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + DataCopy(inputLocal, mmQcQrResGm_[mmQcQrOffset], inputCopyParams); + + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + + Cast(outputLocal, inputLocal, RoundMode::CAST_RINT, curVecToken * baseParams_->dimHeadSizeQc); + + SetFlag(EVENT_ID2); + WaitFlag(EVENT_ID2); + + DataCopy(mmQcQrResDequantGm_[mmQnPreDequantResOffset], outputLocal, outputCopyParams); + { + uint32_t qcSpacing = CeilDivT(totalQcLoops, static_cast(MAX_SYNC_FLAG_COUNT - 1)); + if ((qcCount + 1) % qcSpacing == 0 || qcCount + 1 == totalQcLoops) { + CrossCoreSetFlag(FINISH_VEC_DEQUANT_QC_SPLIT_N); + } else if ((totalQcLoops - CeilDivT(totalQcLoops, qcSpacing)) <= + static_cast(MAX_SYNC_FLAG_COUNT)) { + CrossCoreSetFlag(FINISH_VEC_DEQUANT_QC_SPLIT_N_GAP); + } + } + colOffsetVec += (baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope); + mmQcQrOffset += (baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope); + mmQnPreDequantResOffset += baseParams_->dimHeadSizeQc; + qcCount++; + } + + while (colOffsetRope + baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope <= + colOffsetVecEnd) { // 循环singleNumHeadSize次 + + GlobalTensor outputGmRope = qrOutGm_[ropeQrResOffset]; + + LocalTensor shareTmpUb = shareBuffer_.Get(); + LocalTensor outputLocalRope = shareTmpUb.ReinterpretCast(); + LocalTensor ropeShareTmpUb = + outputLocalRope[curVecToken * baseParams_->dimHeadRope].template ReinterpretCast(); + + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + + RotaryPosEmbPerHead( + outputLocalRope, mmQcQrResGm_[ropeQrOffset], cosLocal_, sinLocal_, ropeShareTmpUb, ropeParams, + ropeStride); + + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + + DataCopy(qrOutGm_[ropeQrResOffset], outputLocalRope, outputRopeParams); + + colOffsetRope += (baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope); + ropeQrOffset += (baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope); + ropeQrResOffset += baseParams_->dimHeadRope; + } + splitCount++; + } + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); +} + +template +__aicore__ inline void +MlaPrologV3SplitM::DynamicQuantQnAndMulQrSyncMMQn(int64_t batchOffset, int64_t curStepBatchSize, + int64_t numHeadOffset, int64_t mmQnLoops) +{ + // 如果curStepBatchSize是偶数,则两个核平分;如果curStepBatchSize是奇数,则奇数核比偶数核多分一个 + // >> 1 是将curStepBatchSize分到每个vec核上; + int64_t curStepBatchSizeVec = (curStepBatchSize + (blockIdx_ % cvRatio_)) >> 1; + if (blockIdx_ >= baseParams_->mm4BlockNum * cvRatio_) { + return; + } + + // 等待前面的Qr部分完成 + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + constexpr uint32_t DYNAMIC_QUANT_INPUT_READY = EVENT_ID0; + constexpr uint32_t MUL_QR_INPUT_COPY_READY = EVENT_ID3; + constexpr uint32_t DYNAMIC_QUANT_OUTPUT_READY = EVENT_ID3; + + int64_t blockBatchOffset = (blockIdx_ % cvRatio_) * (curStepBatchSize / cvRatio_); + int64_t totalSizeCkv = + static_cast(baseParams_->numHeadSize) * static_cast(baseParams_->headSizeCkv); + // 由于两个vec核分curStepBatchSize,各处理curStepBatchSize/2,blockIdx_ & 1 表示是否为第二个vec核 + int64_t dynamicQuantQueryOffset = + baseParams_->stepBatchSize * baseParams_->numHeadSize * baseParams_->headSizeCkv * cubeBlockIdx_ + + blockBatchOffset * totalSizeCkv + numHeadOffset * static_cast(baseParams_->headSizeCkv); + int64_t dynamicQuantQueryResOffset = batchOffset * totalSizeCkv + blockBatchOffset * totalSizeCkv + + numHeadOffset * static_cast(baseParams_->headSizeCkv); + int64_t scaleQueryNopeOffset = batchOffset * static_cast(baseParams_->numHeadSize) + + blockBatchOffset * static_cast(baseParams_->numHeadSize) + numHeadOffset; + int64_t queryOutStride = totalSizeCkv; + int64_t qrOutputStride = + static_cast(baseParams_->numHeadSize) * static_cast(baseParams_->dimHeadRope); + int64_t qrPostProcessResOffset = batchOffset * static_cast(baseParams_->headSizeQr) + + numHeadOffset * static_cast(baseParams_->dimHeadRope) + + blockBatchOffset * static_cast(baseParams_->headSizeQr); + + LocalTensor shareTmpUb = shareBuffer_.Get(); + + float quantScaleCkv = quantScaleCkvGm_.GetValue(0); + + // Dynamic Quant + SetFlag(DYNAMIC_QUANT_OUTPUT_READY); + SetFlag(DYNAMIC_QUANT_INPUT_READY); + + // Rope Post Process + SetFlag(MUL_QR_INPUT_COPY_READY); + CrossCoreWaitFlag(FINISH_MM_QN_SPLIT_N); + // per-head循环 + for (int64_t loopIdx = 0; loopIdx < mmQnLoops; loopIdx++) { + DynamicQuantQnWithMulQr( + dequantScaleQNopeGm_[scaleQueryNopeOffset], queryOutGm_[dynamicQuantQueryResOffset], + qrOutGm_[qrPostProcessResOffset], mmQnResGm_[dynamicQuantQueryOffset], shareTmpUb, curStepBatchSizeVec, + baseParams_->headSizeCkv, baseParams_->numHeadSize, queryOutStride, + // Rope Post Process + qrOutGm_[qrPostProcessResOffset], quantScaleCkv, baseParams_->dimHeadRope, qrOutputStride, cvRatio_); + dynamicQuantQueryOffset += baseParams_->headSizeCkv; + scaleQueryNopeOffset += 1; + dynamicQuantQueryResOffset += baseParams_->headSizeCkv; + qrPostProcessResOffset += baseParams_->dimHeadRope; + } + // Rope Post Process + WaitFlag(MUL_QR_INPUT_COPY_READY); + // Dynamic Quant + WaitFlag(DYNAMIC_QUANT_INPUT_READY); + WaitFlag(DYNAMIC_QUANT_OUTPUT_READY); +} + +} // namespace MlaProlog + +#endif // MLA_PROLOG_V3_SPLIT_M \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_n.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_n.h new file mode 100644 index 000000000000..711e629439ee --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_n.h @@ -0,0 +1,2379 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file kernel_mla_prolog_split_n.h + * \brief + */ + +#ifndef KERNEL_MLA_PROLOG_SPLIT_N_H +#define KERNEL_MLA_PROLOG_SPLIT_N_H + +#include "mla_prolog_comm.h" +#include "mla_prolog_vector_comm.h" +#include "service_matmul.h" +#include "service_rms_norm.h" +#include "service_gather_sin_cos.h" +#include "service_rotary_position_embedding.h" +#include "service_scatter_cache.h" +#include "service_dequant.h" +#include "service_dynamic_quant_qn_mul_qr.h" +#include "../mla_prolog_tiling_data.h" +#include "../mla_prolog_template_tiling_key.h" + +namespace MlaProlog { +template +class MlaPrologVecS1CubS2 { +public: + static constexpr bool isPertile = MLAPT::isPertile; + + using mmInputType = typename MLAPT::mmInputType; + using mmQcQrInputType = typename MLAPT::mmQcQrInputType; + using mmQnInputType = typename MLAPT::mmQnInputType; + using mmCqOutputType = typename MLAPT::mmCqOutputType; + using mmCkvKrOutputType = typename MLAPT::mmCkvKrOutputType; + using mmQcQrOutputType = typename MLAPT::mmQcQrOutputType; + using mmQnOutputType = typename MLAPT::mmQnOutputType; + using rmsNormGammaType = typename MLAPT::rmsNormGammaType; + using rmsNormComputType = typename MLAPT::rmsNormComputType; + using rmsNormCqOutputType = typename MLAPT::rmsNormCqOutputType; + using rmsNormCkvOutputType = typename MLAPT::rmsNormCkvOutputType; + using ropeSinCosType = typename MLAPT::ropeSinCosType; + using ropeComputType = typename MLAPT::ropeComputType; + using ropeOutputType = typename MLAPT::ropeOutputType; + using queryOutputType = typename std::conditional::type; + using kvCacheType = typename MLAPT::kvCacheType; + using krCacheType = typename MLAPT::krCacheType; + using dequantScaleQNopeType = typename MLAPT::dequantScaleQNopeType; + using dequantScaleQNormType = typename MLAPT::dequantScaleQNormType; + using dequantScaleType = typename MLAPT::dequantScaleType; + + MMParams mmCqParam_; + MMParams mmCkvKrParam_; + MMParams mmQcQrParam_; + MMParams mmQnParam_; + + __aicore__ inline MlaPrologVecS1CubS2(TPipe *pipe, const optiling::MlaPrologTilingData *__restrict tilingData, + const optiling::MlaPrologBaseParams *__restrict baseParams) + : pipe_(pipe), tilingData_(tilingData), baseParams_(baseParams) + { + } + + __aicore__ inline void Init(__gm__ uint8_t *tokenX, __gm__ uint8_t *weightDq, __gm__ uint8_t *weightUqQr, + __gm__ uint8_t *weightUk, __gm__ uint8_t *weightDkvKr, __gm__ uint8_t *rmsnormGammaCq, + __gm__ uint8_t *rmsnormGammaCkv, __gm__ uint8_t *ropeSin, __gm__ uint8_t *ropeCos, + __gm__ uint8_t *cacheIndex, __gm__ uint8_t *kvCache, __gm__ uint8_t *krCache, + __gm__ uint8_t *dequantScaleX, __gm__ uint8_t *dequantScaleWDq, + __gm__ uint8_t *deqScaleQcQrW, __gm__ uint8_t *dequantScaleWDkvkr, + __gm__ uint8_t *quantScaleCkv, __gm__ uint8_t *quantScaleCkr, + __gm__ uint8_t *smoothScaleCq, __gm__ uint8_t *actualSeqLen, + __gm__ uint8_t *kNopeClipAlpha, __gm__ uint8_t *queryOut, __gm__ uint8_t *queryRopeOut, + __gm__ uint8_t *dequantScaleQNopeOut, __gm__ uint8_t *queryNormOut, + __gm__ uint8_t *dequantScaleQNormOut, __gm__ uint8_t *workspace); + __aicore__ inline void Process(); + +private: + __aicore__ inline void CopyGlobalParams(); + __aicore__ inline void OutputInit(__gm__ uint8_t *actualSeqLen, __gm__ uint8_t *queryOut, + __gm__ uint8_t *queryRopeOut, __gm__ uint8_t *dequantScaleQNopeOut, + __gm__ uint8_t *queryNormOut, __gm__ uint8_t *dequantScaleQNormOut); + __aicore__ inline void ScaleInit(__gm__ uint8_t *dequantScaleX, __gm__ uint8_t *dequantScaleWDq, + __gm__ uint8_t *deqScaleQcQrW, __gm__ uint8_t *dequantScaleWDkvkr, + __gm__ uint8_t *quantScaleCkv, __gm__ uint8_t *quantScaleCkr, + __gm__ uint8_t *smoothScaleCq, __gm__ uint8_t *kNopeClipAlpha); + __aicore__ inline void WorkspaceInit(__gm__ uint8_t *workspace); + __aicore__ inline void MmParamInit(); + __aicore__ inline void MmCqParamInit(); + __aicore__ inline void MmCkvKrParamInit(); + __aicore__ inline void MmQcQrParamInit(); + __aicore__ inline void MmQnParamInit(); + __aicore__ inline void CubeBufferInit(); + __aicore__ inline void VectorBufferInit(); + __aicore__ inline void UpdateStepBatchParams(int64_t curStepBatchSize); + __aicore__ inline void ComputeAicOffset(AicOffset &aicOffset, int64_t numHeadOffset); + __aicore__ inline void ComputeBlkScatterOffsets(GlobalTensor indexGm, int64_t tokenIndex, int64_t rows, + CkvkrParams &rmsNormAndScatterCkvParams, + CkvkrParams &ropeAndScatterKrParams); + __aicore__ inline void ComputeAivOffset(AivOffset &aivOffset, int64_t batchOffset); + template + __aicore__ inline void AicProcess(AicOffset &aicOffset, int64_t batchOffset, int64_t mmQnLoops); + template + __aicore__ inline void AivProcess(AivOffset &aivOffset, int64_t batchOffset, int64_t curStepBatchSize, + int64_t numHeadOffset, int64_t mmQnLoops); + template + __aicore__ inline void + MatmulSplitN(const GlobalTensor &tensorResGm, const GlobalTensor &tensorAGm, const GlobalTensor &tensorBGm, + const MMParams &mmPara, const UsedBlockParams &mmBlockParams, + const GlobalTensor &tensorAScaleGm = {}, const GlobalTensor &tensorBScaleGm = {}); + __aicore__ inline void MatmulAndSyncQcQr(AicOffset &aicOffset); + __aicore__ inline void MatmulQcQr(AicOffset &aicOffset); + __aicore__ inline void PreloadQnAndSync(AicOffset &aicOffset, int64_t mmQnLoops); + __aicore__ inline void MatmulQnWeightPreload(int64_t weightUkOffset); + template + __aicore__ inline void MatmulQnSyncDynamicQuantAndMulQr(int64_t qcOffset, int64_t weightUkOffset, + int64_t qnResOffset, int64_t mmQnLoops); + __aicore__ inline void CopyInSinCos(int64_t tokenIndex, int64_t curVecToken, int64_t batchOffset, + int64_t curStepBatchSize); + __aicore__ inline void RmsNormCq(int64_t tokenIndex, int64_t rmsNormCqOffset, int64_t rmsNormCqResOffset, + int64_t curVecToken, int64_t curBlockTokenOffset); + __aicore__ inline void CopyDequantScaleCq(uint64_t dequantScaleCqElementNum, + LocalTensor &dequantScaleQcQr, int64_t curVecToken, + int64_t curBlockTokenOffset); + __aicore__ inline void CopyQueryNormScale(int64_t stepTokenIndex, LocalTensor &dequantScaleQcQr, + int64_t curVecToken); + __aicore__ inline void RopeAndScatterKr(LocalTensor &dequantScaleXLocal, LocalTensor &shareTmpUb, + LocalTensor &cosLocalCkvKr, + LocalTensor &sinLocalCkvKr, + CkvkrParams ropeAndScatterKrParams); + __aicore__ inline void ScatterKr(LocalTensor &outputKrLocal, CkvkrParams ropeAndScatterKrParams); + __aicore__ inline void RmsNormAndScatterCkv(LocalTensor &dequantScaleXLocal, + LocalTensor &shareTmpUb, + LocalTensor &cosLocalCkvKr, + LocalTensor &sinLocalCkvKr, + CkvkrParams rmsNormAndScatterCkvParams); + __aicore__ inline void RmsNormAndQuantizeCkv(LocalTensor &outputLocal, + LocalTensor &rmsNormShareTmpUb, + LocalTensor &dequantScaleXLocal, RmsNormParam rmsNormParams, + CkvkrParams rmsNormAndScatterCkvParams); + __aicore__ inline void ScatterCkv(LocalTensor &outputLocal, CkvkrParams rmsNormAndScatterCkvParams); + __aicore__ inline void RmsNormRopeScatterCkvKr(int64_t tokenIndex, int64_t rmsNormCkvOffset, int64_t ropeKrOffset, + int64_t curVecToken); + __aicore__ inline void RopeQr(int64_t ropeQrOffset, int64_t ropeQrResOffset, int64_t curVecToken, + int64_t curBlockTokenOffset); + __aicore__ inline void DequantQc(int64_t mmQnPreDequantOffset, int64_t mmQnPreDequantResOffset, int64_t curVecToken, + int64_t curBlockTokenOffset); + __aicore__ inline void CastQc(int64_t mmQnPreCastOffset, int64_t mmQnPreCastResOffset, int64_t curVecToken); + // 低时延算力分组场景 + template + __aicore__ inline void DequantQcAndRopeQc(AivOffset &aivOffset, int64_t batchOffset, int64_t curStepBatchSize, + int64_t numHeadOffset, int64_t mmQnLoops); + __aicore__ inline void DequantQcSplitNGroupCase(int64_t mmQnPreDequantOffset, int64_t mmQnPreDequantResOffset, + int64_t qcQrScaleOffset); + __aicore__ inline void RopeQrSplitNGroupCase(int64_t ropeQrOffset, int64_t ropeQrResOffset); + __aicore__ inline void DequantQcQrSplitN(const DequantQcQrSplitNParams &dequantQcQrSplitN); + __aicore__ inline void CastQcQrSplitN(const CastQcQrSplitNParams &castQcQrSplitN); + __aicore__ inline void RopeQrSplitN(const RopeQrSplitNParams &ropeQrSplitNParams); + __aicore__ inline void DequantAndRopeSplitNSyncMMQcQr(int64_t mmQnPreDequantOffset, int64_t mmQnPreDequantResOffset, + int64_t ropeQrOffset, int64_t ropeQrResOffset); + __aicore__ inline void DynamicQuantQnAndMulQrSyncMMQn(int64_t batchOffset, int64_t curStepBatchSize, + int64_t numHeadOffset, int64_t mmQnLoops); + + TPipe *pipe_; + const optiling::MlaPrologTilingData *__restrict tilingData_; + const optiling::MlaPrologBaseParams *__restrict baseParams_; + uint32_t blockIdx_ = 0U; + uint32_t cubeBlockIdx_ = 0U; // AIV上使用AIC的blockIdx + int64_t vectorRow_ = 1; + int64_t curVectorBlockNum_; + int64_t vectorCoreNum_; + uint64_t dequantScaleCqSize_ = 1; + uint32_t curStepVecFrontToken_; + uint32_t curStepVecFrontListNum_; + uint32_t curStepVecBackToken_; + uint32_t curVecTokenMax_; + bool enableSmoothScalesCq_; + static constexpr uint32_t cvMode = MLAPT::cvRatio; // 编译态,默认cv1:2 + static constexpr bool isFp8E8m0 = std::is_same::value; + uint32_t cvRatio_ = 2U; // 默认cv 1:2 + + struct DequantTool { + GlobalTensor deQuantScaleCqGm_; + TBuf deQuantScaleCqBuffer_; // 用于临时存储每一行的Scale,以及汇总最终每一行的Scale参数 + LocalTensor deQuantScaleCqLocal_; + __aicore__ inline DequantTool() + { + } + }; + + // 算子分组开关 + DequantTool dequantTool_; + + // GM + GlobalTensor tokenXGm_; + GlobalTensor weightDqGm_; + GlobalTensor weightUqQrGm_; + GlobalTensor weightUkGm_; + GlobalTensor weightDkvKrGm_; + GlobalTensor cacheIndexGm_; + GlobalTensor rmsnormGammaCqGm_; + GlobalTensor rmsnormGammaCkvGm_; + GlobalTensor ropeSinGm_; + GlobalTensor ropeCosGm_; + GlobalTensor kvCacheGm_; + GlobalTensor krCacheGm_; + GlobalTensor qrOutGm_; + + GlobalTensor dequantScaleXGm_; + GlobalTensor dequantScaleWDqGm_; + GlobalTensor dequantScaleWDkvkrGm_; + GlobalTensor smoothScaleCqGm_; + GlobalTensor deqScaleQcQrW_; // per-channel反量化参数 + GlobalTensor quantScaleCkvGm_; + GlobalTensor quantScaleCkrGm_; + + GlobalTensor actualSeqLenGm_; + GlobalTensor kNopeClipAlphaGm_; + + GlobalTensor rmsNormCqResGm_; + GlobalTensor mmCqResGm_; + GlobalTensor mmCkvKrResGm_; + GlobalTensor mmQcQrResGm_; + GlobalTensor mmQcQrResDequantGm_; + GlobalTensor mmQnResGm_; + GlobalTensor dequantScaleQNopeGm_; + GlobalTensor queryOutGm_; + GlobalTensor dequantScaleQNormGm_; + + // UB + TBuf sincosBuffer_; + TBuf shareBuffer_; + TBuf dequantScaleWDqBuffer_; + TBuf dequantScaleWDkvKrBuffer_; + TBuf rmsnormGammaCqBuffer_; + TBuf rmsnormGammaCkvBuffer_; + TBuf smoothScaleCqBuffer_; + TBuf quantScaleCkvBuffer_; + TBuf quantScaleCkrBuffer_; + TBuf stepActualSeqBuffer_; + + LocalTensor cosLocal_; + LocalTensor sinLocal_; + LocalTensor dequantScaleWDqLocal_; + LocalTensor dequantScaleWDkvKrLocal_; + LocalTensor rmsnormGammaCqLocal_; + LocalTensor rmsnormGammaCkvLocal_; + LocalTensor smoothScaleCqLocal_; + LocalTensor quantScaleCkvLocal_; + LocalTensor quantScaleCkrLocal_; + LocalTensor stepActualSeqLocal_; + + struct ActSeqState { + bool inited = false; // whether state is initialized for the run + int64_t curBatch = 0; // current batch index for the running token index + int64_t prevPrefix = 0; // prefix sum of seq lengths up to (curBatch - 1) + int64_t curPrefix = 0; // prefix sum of seq lengths up to curBatch (exclusive high bound) + int64_t prevIndexOffset = 0; // prefix sum of blocks up to (curBatch - 1) + int64_t curIndexOffset = 0; // prefix sum of blocks up to curBatch + }; + ActSeqState actSeqState_; + int64_t blockSpanPerBatch_ = 0; // blockNum * blockSize + + TBuf aBufL1_; + TBuf bBufL1_; + LocalTensor aL1Tensor_; + LocalTensor bL1Tensor_; + MMBufParams bufParam_; + TBuf aBufL0_; + TBuf bBufL0_; + TBuf cBufL0_; +}; + +template +__aicore__ inline void MlaPrologVecS1CubS2::Init( + __gm__ uint8_t *tokenX, __gm__ uint8_t *weightDq, __gm__ uint8_t *weightUqQr, __gm__ uint8_t *weightUk, + __gm__ uint8_t *weightDkvKr, __gm__ uint8_t *rmsnormGammaCq, __gm__ uint8_t *rmsnormGammaCkv, + __gm__ uint8_t *ropeSin, __gm__ uint8_t *ropeCos, __gm__ uint8_t *cacheIndex, __gm__ uint8_t *kvCache, + __gm__ uint8_t *krCache, __gm__ uint8_t *dequantScaleX, __gm__ uint8_t *dequantScaleWDq, + __gm__ uint8_t *deqScaleQcQrW, __gm__ uint8_t *dequantScaleWDkvkr, __gm__ uint8_t *quantScaleCkv, + __gm__ uint8_t *quantScaleCkr, __gm__ uint8_t *smoothScaleCq, __gm__ uint8_t *actualSeqLen, + __gm__ uint8_t *kNopeClipAlpha, __gm__ uint8_t *queryOut, __gm__ uint8_t *queryRopeOut, + __gm__ uint8_t *dequantScaleQNopeOut, __gm__ uint8_t *queryNormOut, __gm__ uint8_t *dequantScaleQNormOut, + __gm__ uint8_t *workspace) +{ + cvRatio_ = GetSubBlockNum(); // CV1:2场景返回2,其他场景返回1 + blockIdx_ = GetBlockIdx(); // cube:0-23 vec:0-47 + if ASCEND_IS_AIC { + cubeBlockIdx_ = blockIdx_; + } else { + cubeBlockIdx_ = blockIdx_ / cvRatio_; + } + curVectorBlockNum_ = static_cast(baseParams_->stepBatchSize); + vectorCoreNum_ = static_cast(baseParams_->vectorBlockNum); // aivNum 48 + if (cvMode == 1 && cvRatio_ == 2) { // 编译态cv1:1,运行态cv1:2 + if (vectorCoreNum_ < curVectorBlockNum_) { + vectorCoreNum_ = vectorCoreNum_ * 2; // 修正为运行态vector数目 + } + } + curVecTokenMax_ = (curVectorBlockNum_ + vectorCoreNum_ - 1) / vectorCoreNum_; + enableSmoothScalesCq_ = smoothScaleCq == nullptr ? false : true; + // GM + tokenXGm_.SetGlobalBuffer((__gm__ mmInputType *)tokenX); + weightDqGm_.SetGlobalBuffer((__gm__ mmInputType *)weightDq); // NZ + weightUqQrGm_.SetGlobalBuffer((__gm__ mmQcQrInputType *)weightUqQr); // NZ + weightUkGm_.SetGlobalBuffer((__gm__ mmQnInputType *)weightUk); + weightDkvKrGm_.SetGlobalBuffer((__gm__ mmInputType *)weightDkvKr); // NZ + rmsnormGammaCqGm_.SetGlobalBuffer((__gm__ rmsNormGammaType *)rmsnormGammaCq); + rmsnormGammaCkvGm_.SetGlobalBuffer((__gm__ rmsNormGammaType *)rmsnormGammaCkv); + ropeSinGm_.SetGlobalBuffer((__gm__ ropeSinCosType *)ropeSin); + ropeCosGm_.SetGlobalBuffer((__gm__ ropeSinCosType *)ropeCos); + if constexpr (MLAPT::cacheMode != CACHE_MODE::ND) { + cacheIndexGm_.SetGlobalBuffer((__gm__ int64_t *)cacheIndex); + } + kvCacheGm_.SetGlobalBuffer((__gm__ kvCacheType *)kvCache); + krCacheGm_.SetGlobalBuffer((__gm__ krCacheType *)krCache); + + OutputInit(actualSeqLen, queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut); + ScaleInit(dequantScaleX, dequantScaleWDq, deqScaleQcQrW, dequantScaleWDkvkr, quantScaleCkv, quantScaleCkr, + smoothScaleCq, kNopeClipAlpha); + MmParamInit(); + WorkspaceInit(workspace); + if ASCEND_IS_AIV { + VectorBufferInit(); + } else { + CubeBufferInit(); + } +} + +template +__aicore__ inline void +MlaPrologVecS1CubS2::OutputInit(__gm__ uint8_t *actualSeqLen, __gm__ uint8_t *queryOut, + __gm__ uint8_t *queryRopeOut, __gm__ uint8_t *dequantScaleQNopeOut, + __gm__ uint8_t *queryNormOut, __gm__ uint8_t *dequantScaleQNormOut) +{ + qrOutGm_.SetGlobalBuffer((__gm__ ropeOutputType *)queryRopeOut); + if constexpr (((std::is_same::value && std::is_same::value) || + (std::is_same::value && std::is_same::value) || + (std::is_same::value && std::is_same::value)) && + !isPertile) { + dequantScaleQNopeGm_.SetGlobalBuffer((__gm__ dequantScaleQNopeType *)dequantScaleQNopeOut); + queryOutGm_.SetGlobalBuffer((__gm__ queryOutputType *)queryOut); + } else { + mmQnResGm_.SetGlobalBuffer((__gm__ mmQnOutputType *)queryOut); + } + if (baseParams_->queryNormFlag == 1U) { + rmsNormCqResGm_.SetGlobalBuffer((__gm__ mmQcQrInputType *)queryNormOut); + + if constexpr (IsFullQuantMode()) { + dequantScaleQNormGm_.SetGlobalBuffer((__gm__ dequantScaleQNormType *)dequantScaleQNormOut); + } + } + if constexpr (MLAPT::actualSeqMode == ACTUAL_SEQ_MODE::EN_Q_LEN) { + actualSeqLenGm_.SetGlobalBuffer((__gm__ int32_t *)actualSeqLen); + } +} + +template +__aicore__ inline void +MlaPrologVecS1CubS2::ScaleInit(__gm__ uint8_t *dequantScaleX, __gm__ uint8_t *dequantScaleWDq, + __gm__ uint8_t *deqScaleQcQrW, __gm__ uint8_t *dequantScaleWDkvkr, + __gm__ uint8_t *quantScaleCkv, __gm__ uint8_t *quantScaleCkr, + __gm__ uint8_t *smoothScaleCq, __gm__ uint8_t *kNopeClipAlpha) +{ + if constexpr (IsFullQuantMode()) { + dequantScaleXGm_.SetGlobalBuffer((__gm__ dequantScaleType *)dequantScaleX); + dequantScaleWDqGm_.SetGlobalBuffer((__gm__ dequantScaleType *)dequantScaleWDq); + dequantScaleWDkvkrGm_.SetGlobalBuffer((__gm__ dequantScaleType *)dequantScaleWDkvkr); + } + + if constexpr (IsFullQuantMode()) { + smoothScaleCqGm_.SetGlobalBuffer((__gm__ float *)smoothScaleCq); + deqScaleQcQrW_.SetGlobalBuffer((__gm__ dequantScaleType *)deqScaleQcQrW); + quantScaleCkvGm_.SetGlobalBuffer((__gm__ float *)quantScaleCkv); + quantScaleCkrGm_.SetGlobalBuffer((__gm__ float *)quantScaleCkr); + } else if constexpr (std::is_same::value && isFp8E8m0) { + deqScaleQcQrW_.SetGlobalBuffer((__gm__ dequantScaleType *)deqScaleQcQrW); + quantScaleCkvGm_.SetGlobalBuffer((__gm__ float *)quantScaleCkv); + } + + if constexpr (isPertile) { + kNopeClipAlphaGm_.SetGlobalBuffer((__gm__ float *)kNopeClipAlpha); + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::MmParamInit() +{ + MmCqParamInit(); + MmCkvKrParamInit(); + MmQcQrParamInit(); + MmQnParamInit(); +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::MmCqParamInit() +{ + mmCqParam_.m = baseParams_->stepBatchSize; // 32 + if (cubeBlockIdx_ == baseParams_->mm1BlockNum - 1) { + mmCqParam_.n = baseParams_->headSizeCq - baseParams_->mm1SingleCoreN * cubeBlockIdx_; + } else { + mmCqParam_.n = baseParams_->mm1SingleCoreN; // 1536 / 24 = 64 + } + mmCqParam_.k = baseParams_->headSizeX; // 7168 + mmCqParam_.needSetOrgShape = 1; + mmCqParam_.orgM = mmCqParam_.m; + mmCqParam_.orgN = mmCqParam_.n; + mmCqParam_.orgKa = mmCqParam_.k; + mmCqParam_.orgKb = mmCqParam_.k; + mmCqParam_.orgKc = baseParams_->headSizeCq; // 1536 + mmCqParam_.baseK = + (sizeof(mmInputType) == ONE_BYTE_TYPE_SIZE) ? 256 : 128; // 128KB / (128 max baseN * 4 stepK * sizeof(type)) + mmCqParam_.baseN = (std::is_same::value && isFp8E8m0) ? 64 : 128; + mmCqParam_.stepK = 4; + if ((mmCqParam_.k / mmCqParam_.baseK) % mmCqParam_.stepK != 0) { + mmCqParam_.stepK = 3; // support k = 7680, mmInputType int8, no tail + } + mmCqParam_.kL1StepSize = mmCqParam_.baseK * mmCqParam_.stepK; + mmCqParam_.kScale = mmCqParam_.k / FP8_E4M3_BLOCK_SIZE; +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::MmCkvKrParamInit() +{ + mmCkvKrParam_.m = baseParams_->stepBatchSize; // 32 + if (cubeBlockIdx_ == baseParams_->mm2BlockNum - 1) { + mmCkvKrParam_.n = + baseParams_->headSizeCkv + baseParams_->headSizeKr - baseParams_->mm2SingleCoreN * cubeBlockIdx_; + } else { + mmCkvKrParam_.n = baseParams_->mm2SingleCoreN; + } + mmCkvKrParam_.k = baseParams_->headSizeX; // 7168 + mmCkvKrParam_.needSetOrgShape = 1; + mmCkvKrParam_.orgM = mmCkvKrParam_.m; + mmCkvKrParam_.orgN = mmCkvKrParam_.n; + mmCkvKrParam_.orgKa = mmCkvKrParam_.k; + mmCkvKrParam_.orgKb = mmCkvKrParam_.k; + mmCkvKrParam_.orgKc = (baseParams_->headSizeCkv + baseParams_->dimHeadRope); // 576 + mmCkvKrParam_.baseK = + (sizeof(mmInputType) == ONE_BYTE_TYPE_SIZE) ? 256 : 128; // 128KB / (128 max baseN * 4 stepK * sizeof(type)) + mmCkvKrParam_.baseN = 128; + mmCkvKrParam_.stepK = 4; + if ((mmCkvKrParam_.k / mmCkvKrParam_.baseK) % mmCkvKrParam_.stepK != 0) { + mmCkvKrParam_.stepK = 3; // support k = 7680, mmInputType int8, no tail + } + mmCkvKrParam_.kL1StepSize = mmCkvKrParam_.baseK * mmCkvKrParam_.stepK; + mmCkvKrParam_.kScale = mmCkvKrParam_.k / FP8_E4M3_BLOCK_SIZE; +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::MmQcQrParamInit() +{ + mmQcQrParam_.m = baseParams_->stepBatchSize; // 32 + if constexpr (MLAPT::enableGroupComputeOpt) { + // 算力分组仅考虑G8 Qc Qr(8+4核), n固定128, 这里不处理尾核 + mmQcQrParam_.n = baseParams_->mm3SingleCoreN; + } else { + if (cubeBlockIdx_ == baseParams_->mm3BlockNum - 1) { + mmQcQrParam_.n = + baseParams_->headSizeQc + baseParams_->headSizeQr - baseParams_->mm3SingleCoreN * cubeBlockIdx_; + } else { + mmQcQrParam_.n = baseParams_->mm3SingleCoreN; + } + } + + mmQcQrParam_.k = baseParams_->headSizeCq; // 1536 + mmQcQrParam_.needSetOrgShape = 1; + mmQcQrParam_.orgM = mmQcQrParam_.m; + mmQcQrParam_.orgN = mmQcQrParam_.n; + mmQcQrParam_.orgKa = mmQcQrParam_.k; + mmQcQrParam_.orgKb = mmQcQrParam_.k; + mmQcQrParam_.orgKc = (baseParams_->headSizeQc + baseParams_->headSizeQr); // (128 * 32 + 64 * 32) + mmQcQrParam_.baseK = (sizeof(mmQcQrInputType) == ONE_BYTE_TYPE_SIZE) ? 128 : 64; + if constexpr (MLAPT::enableGroupComputeOpt) { + mmQcQrParam_.baseN = 128; + } else { + if constexpr (std::is_same::value && + isFp8E8m0) { // FP8全量化场景下L1B用满,修改baseN会造成内存踩踏 + mmQcQrParam_.baseN = 128; + } else { + if (mmQcQrParam_.m <= 64) { // FP8全量化场景,scale需要额外占用L1,该优化不适用 + mmQcQrParam_.baseN = 256; + } else { + mmQcQrParam_.baseN = 128; + } + } + } + mmQcQrParam_.stepK = 4; + mmQcQrParam_.kL1StepSize = mmQcQrParam_.baseK * mmQcQrParam_.stepK; + mmQcQrParam_.kScale = mmQcQrParam_.k / FP8_E4M3_BLOCK_SIZE; +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::MmQnParamInit() +{ + mmQnParam_.m = baseParams_->stepBatchSize; // 32 + mmQnParam_.n = baseParams_->headSizeCkv; // 512 + mmQnParam_.k = baseParams_->dimHeadSizeQc; // 128, 这里numHeadSize被分核,matmul设置里不体现 + mmQnParam_.needSetOrgShape = 1; + mmQnParam_.orgM = mmQnParam_.m; + mmQnParam_.orgN = mmQnParam_.n; + if constexpr (std::is_same::value || std::is_same::value) { + mmQnParam_.orgKa = baseParams_->headSizeQc; + } else { + mmQnParam_.orgKa = baseParams_->headSizeQc + baseParams_->headSizeQr; + } + mmQnParam_.orgKb = baseParams_->dimHeadSizeQc; + mmQnParam_.orgKc = baseParams_->headSizeCkv * baseParams_->numHeadSize; + mmQnParam_.baseN = 128; + mmQnParam_.baseK = 128; + mmQnParam_.stepK = 1; + if ((mmQnParam_.k > mmQnParam_.baseK) && (mmQnParam_.k % mmQnParam_.baseK != 0)) { + mmQnParam_.baseK = 64; + mmQnParam_.stepK = 3; // support D = 192 + } + mmQnParam_.kL1StepSize = mmQnParam_.baseK * mmQnParam_.stepK; + mmQnParam_.kScale = mmQnParam_.k / FP8_E4M3_BLOCK_SIZE; +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::VectorBufferInit() +{ + // 暂存已分配UB大小,其余为shareBuffer + uint64_t usedBytes = 0; + if constexpr (IsFullQuantMode()) { + uint64_t dequantScaleWDqSize = baseParams_->headSizeCq * sizeof(float); + pipe_->InitBuffer(dequantScaleWDqBuffer_, dequantScaleWDqSize); // [1, 1536] + dequantScaleWDqLocal_ = dequantScaleWDqBuffer_.Get(); + usedBytes += dequantScaleWDqSize; + + uint64_t dequantScaleWDkvKrSize = (baseParams_->headSizeCkv + baseParams_->dimHeadRope) * sizeof(float); + pipe_->InitBuffer(dequantScaleWDkvKrBuffer_, dequantScaleWDkvKrSize); // [1, 512 + 64] + dequantScaleWDkvKrLocal_ = dequantScaleWDkvKrBuffer_.Get(); + usedBytes += dequantScaleWDkvKrSize; + } + uint64_t rmsnormGammaCqSize = baseParams_->headSizeCq * sizeof(rmsNormGammaType); + pipe_->InitBuffer(rmsnormGammaCqBuffer_, rmsnormGammaCqSize); // [1, 1536] bf16 + rmsnormGammaCqLocal_ = rmsnormGammaCqBuffer_.Get(); + usedBytes += rmsnormGammaCqSize; + + uint64_t rmsnormGammaCkvSize = baseParams_->headSizeCkv * sizeof(rmsNormGammaType); + pipe_->InitBuffer(rmsnormGammaCkvBuffer_, rmsnormGammaCkvSize); // [1, 512] bf16 + rmsnormGammaCkvLocal_ = rmsnormGammaCkvBuffer_.Get(); + usedBytes += rmsnormGammaCkvSize; + + if constexpr (IsFullQuantMode()) { + if (enableSmoothScalesCq_) { + uint64_t smoothScaleCqSize = baseParams_->headSizeCq * sizeof(float); + pipe_->InitBuffer(smoothScaleCqBuffer_, smoothScaleCqSize); // [1, 1536] + smoothScaleCqLocal_ = smoothScaleCqBuffer_.Get(); + usedBytes += smoothScaleCqSize; + } + } + + if constexpr (std::is_same::value) { + uint64_t quantScaleCkrSize = baseParams_->dimHeadRope * sizeof(float); + pipe_->InitBuffer(quantScaleCkrBuffer_, baseParams_->dimHeadRope * sizeof(float)); // [1, 64] + quantScaleCkrLocal_ = quantScaleCkrBuffer_.Get(); + usedBytes += quantScaleCkrSize; + } + + if constexpr (IsFullQuantMode()) { + uint64_t quantScaleCkvSize = 0; + if constexpr (std::is_same::value || + std::is_same::value) { + quantScaleCkvSize = ALIGN_BLOCK_SIZE; + pipe_->InitBuffer(quantScaleCkvBuffer_, quantScaleCkvSize); + } else { // per_channel量化场景走else分支 + quantScaleCkvSize = baseParams_->headSizeCkv * sizeof(float); + pipe_->InitBuffer(quantScaleCkvBuffer_, quantScaleCkvSize); // [1, 512] + } + quantScaleCkvLocal_ = quantScaleCkvBuffer_.Get(); + usedBytes += quantScaleCkvSize; + } + + // 预留brcb的空间 + uint64_t deQuantScaleCqSize = (baseParams_->stepBatchSize + 7) * ALIGN_BLOCK_SIZE; + pipe_->InitBuffer(dequantTool_.deQuantScaleCqBuffer_, deQuantScaleCqSize); + dequantTool_.deQuantScaleCqLocal_ = dequantTool_.deQuantScaleCqBuffer_.template Get(); + usedBytes += deQuantScaleCqSize; + + uint64_t sincosSize = 0; + if constexpr (MLAPT::enableDequantOpt) { + // 在ropeQr进行切N处理后,会复用shareBuffer的内存,不需要额外申请 + // 开启开关后会按照head切分rope qr,此时需要加载一半batchsize数量的sin和cos值 + // 需要2倍的空间分别存储sin和cos + sincosSize = 2 * baseParams_->dimHeadRope * sizeof(ropeComputType) * + ((baseParams_->stepBatchSize + cvRatio_ - 1) / cvRatio_); + pipe_->InitBuffer(sincosBuffer_, sincosSize); + } else { + // 需要2倍的空间分别存储sin和cos + sincosSize = 2 * baseParams_->dimHeadRope * sizeof(ropeComputType) * curVecTokenMax_; + pipe_->InitBuffer(sincosBuffer_, sincosSize); // [2, 64] float + } + usedBytes += sincosSize; + + if constexpr (MLAPT::enableDequantOpt) { + cosLocal_ = sincosBuffer_.Get(); + sinLocal_ = cosLocal_[baseParams_->dimHeadRope * ((baseParams_->stepBatchSize + cvRatio_ - 1) / cvRatio_)]; + } else { + cosLocal_ = sincosBuffer_.Get(); + sinLocal_ = cosLocal_[baseParams_->dimHeadRope * curVecTokenMax_]; + } + + if constexpr (MLAPT::actualSeqMode == ACTUAL_SEQ_MODE::EN_Q_LEN) { + uint64_t stepActualSeqSize = baseParams_->stepBatchSize * sizeof(int64_t); + pipe_->InitBuffer(stepActualSeqBuffer_, stepActualSeqSize); // [1, stepBatchSize] + stepActualSeqLocal_ = stepActualSeqBuffer_.Get(); + usedBytes += stepActualSeqSize; + } + + // 由于shareBuffer属于各个vector操作临时申请内存的区域内存使用不固定,建议shareBuffer始终放在最后 + // 防止写入shareBuffer越界导致前面固定申请的UB内存被踩。 + pipe_->InitBuffer(shareBuffer_, MAX_UB_SIZE - usedBytes); + CrossCoreSetFlag(FINISH_VEC_CKVKR); +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::CubeBufferInit() +{ + // cube相关Buffer初始化 + pipe_->InitBuffer(aBufL1_, L1_A_SIZE * 2); + pipe_->InitBuffer(bBufL1_, L1_B_SIZE * 2); + + SetFlag(A_EVENT0); + SetFlag(A_EVENT1); + SetFlag(B_EVENT0); + SetFlag(B_EVENT1); + aL1Tensor_ = aBufL1_.Get(); + bL1Tensor_ = bBufL1_.Get(); + bufParam_.aL1BufAddr = aBufL1_.GetBufferAddr(aL1Tensor_.GetBufferHandle()); + bufParam_.bL1BufAddr = bBufL1_.GetBufferAddr(bL1Tensor_.GetBufferHandle()); + + pipe_->InitBuffer(aBufL0_, L0A_PP_SIZE * 2); // 64K + pipe_->InitBuffer(bBufL0_, L0B_PP_SIZE * 2); // 64K + pipe_->InitBuffer(cBufL0_, L0C_PP_SIZE * 2); // 128K + + SetFlag(L0A_EVENT0); + SetFlag(L0A_EVENT1); + SetFlag(L0B_EVENT0); + SetFlag(L0B_EVENT1); + + SetFlag(L0C_EVENT0); + SetFlag(L0C_EVENT1); + + SetFlag(SCALE_EVENT); + + bufParam_.aL0BufAddr = aBufL0_.GetBufferAddr(aBufL0_.Get().GetBufferHandle()); + bufParam_.bL0BufAddr = bBufL0_.GetBufferAddr(bBufL0_.Get().GetBufferHandle()); + bufParam_.cL0BufAddr = cBufL0_.GetBufferAddr(cBufL0_.Get().GetBufferHandle()); +} + +/* + * workspace管理 + * 1. 常驻:dequantTool_.deQuantScaleCqGm_ stepBs * 32 Byte + * 2. 中间结果: + * tokenXGm_──────>mmCkvKrResGm_ + * | [stepBS, HCkv + Dr] + * | (bf16 | int32) + * └─────────>mmCqResGm_──────>rmsNormCqResGm_──────>mmQcQrResGm_──────>mmQcQrResDequantGm_──────>mmQnResGm_ + * [stepBS, HCq] [stepBS, HCq] [stepBS, N1, D + Dr] [stepBS, N1, D] [stepBS, N1, + * HCkv] (bf16 | int32) (bf16 | int8) (bf16 | int32) (bf16) (bf16) + */ +template +__aicore__ inline void MlaPrologVecS1CubS2::WorkspaceInit(__gm__ uint8_t *workspace) +{ + int64_t workspaceOffset = 0; + if constexpr (std::is_same::value && isFp8E8m0) { + dequantScaleCqSize_ = baseParams_->headSizeCq / FP8_E4M3_BLOCK_SIZE; + } + dequantScaleCqSize_ = Align(dequantScaleCqSize_, BYTE_BLOCK); + if constexpr (MLAPT::enableGroupComputeOpt || MLAPT::enableDequantOpt) { + dequantTool_.deQuantScaleCqGm_.SetGlobalBuffer((__gm__ dequantScaleType *)(workspace + workspaceOffset)); + workspaceOffset += baseParams_->stepBatchSize * dequantScaleCqSize_; + } + + mmCkvKrResGm_.SetGlobalBuffer((__gm__ mmCkvKrOutputType *)(workspace + workspaceOffset)); + workspaceOffset += static_cast(baseParams_->stepBatchSize) * + static_cast(baseParams_->headSizeCkv + baseParams_->dimHeadRope) * + static_cast(sizeof(mmCkvKrOutputType)); + + mmCqResGm_.SetGlobalBuffer((__gm__ mmCqOutputType *)(workspace + workspaceOffset)); + + if (baseParams_->queryNormFlag == 0U) { + // 全量化场景下`mmCqResGm_`与`rmsNormCqResGm_`的dtype不同,无法共用workspace,此处需要偏移`mmCqResGm_`占用的大小; + if constexpr (!std::is_same::value) { + workspaceOffset += static_cast(baseParams_->stepBatchSize) * + static_cast(baseParams_->headSizeCq) * + static_cast(sizeof(mmCqOutputType)); + } + // 不返回queryNorm时,rmsNormCq结果放在workspace + rmsNormCqResGm_.SetGlobalBuffer((__gm__ rmsNormCqOutputType *)(workspace + workspaceOffset)); + workspaceOffset += static_cast(baseParams_->stepBatchSize) * + static_cast(baseParams_->headSizeCq) * sizeof(rmsNormCqOutputType); + } else { + // 返回queryNorm时,rmsNormCq结果直接输出gm,此处需要偏移`mmCqResGm_`占用的大小 + workspaceOffset += static_cast(baseParams_->stepBatchSize) * + static_cast(baseParams_->headSizeCq) * static_cast(sizeof(mmCqOutputType)); + } + + mmQcQrResGm_.SetGlobalBuffer((__gm__ mmQcQrOutputType *)(workspace + workspaceOffset)); + if constexpr (IsFullQuantMode()) { + workspaceOffset += static_cast(baseParams_->stepBatchSize) * + static_cast(baseParams_->headSizeQc + baseParams_->headSizeQr) * + static_cast(sizeof(mmQcQrOutputType)); + } + + mmQcQrResDequantGm_.SetGlobalBuffer((__gm__ mmQnInputType *)(workspace + workspaceOffset)); + if constexpr (((std::is_same::value && std::is_same::value) || + (std::is_same::value && std::is_same::value) || + (std::is_same::value && std::is_same::value)) && + !isPertile) { + workspaceOffset += static_cast(baseParams_->stepBatchSize) * + static_cast(baseParams_->numHeadSize) * + static_cast(baseParams_->dimHeadSizeQc) * sizeof(mmQnInputType); + mmQnResGm_.SetGlobalBuffer((__gm__ mmQnOutputType *)(workspace + workspaceOffset)); + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::UpdateStepBatchParams(int64_t curStepBatchSize) +{ + mmCqParam_.m = curStepBatchSize; + mmCkvKrParam_.m = curStepBatchSize; + mmQcQrParam_.m = curStepBatchSize; + mmQnParam_.m = curStepBatchSize; + curVectorBlockNum_ = curStepBatchSize; +} + +/* + * MlaProlog算子计算&CV流水同步流程 + * ┌───────────────── token_x ─────────────────┐ + * | ▼ + * | MatmulCkvKr + * ▼ | wait mm CkvKr(0x1) + * MatmulCq ▼ + * | wait mm Cq(0x1) ┌─────────────────┐ + * ▼ ▼ ▼ + * RmsNorm(Cq) RmsNorm(Ckv) Rope(Kr) + * | wait rmsNorm cq(0x1) | | + * ▼ ▼ ▼ + * ┌───────MatmulQcQr───────┐ Scatter(Ckv) Scatter(Kr) + * | wait mm Qc(0x1) | | | + * ▼ | ▼ ▼ + * DequantQc | kv_cache_out kr_cache_out + * | wait dequant qc(0x1) | wait mm Qr(0x2) + * ▼ ▼ + * MatmulQn Rope(Qr) + * | wait mm Qn(0x1) | + * DynamicQuantQn──┐ | + * | ▼ | + * ▼ dequant_scale_out ▼ + * query_out query_rope_out + * 注:仅为表明基本计算与CV同步流程,仅包含了影响CV同步的量化分支,其余量化分支应参考设计文档。 + */ +template +__aicore__ inline void MlaPrologVecS1CubS2::Process() +{ + constexpr bool needQnDynamicQuant = + ((std::is_same::value && std::is_same::value) || + (std::is_same::value && std::is_same::value) || + (std::is_same::value && std::is_same::value)) && + !isPertile; + int64_t numHeadOffset = cubeBlockIdx_ * baseParams_->mm4SingleCoreBatch; + int64_t mmQnLoops; + if (cubeBlockIdx_ == baseParams_->mm4BlockNum - 1) { + mmQnLoops = static_cast(baseParams_->numHeadSize) - numHeadOffset; + } else { + mmQnLoops = static_cast(baseParams_->mm4SingleCoreBatch); + } + + // AIC的offset参数 + AicOffset aicOffset; + ComputeAicOffset(aicOffset, numHeadOffset); + + // AIV的offset参数 + AivOffset aivOffset; + ComputeAivOffset(aivOffset, 0); + + int64_t bsSize = static_cast(baseParams_->tokenSize); + // 需要考虑BS合轴的尾块情况 + for (int64_t batchOffset = 0; batchOffset < bsSize; + batchOffset += static_cast(baseParams_->stepBatchSize)) { + if constexpr (MLAPT::actualSeqMode == ACTUAL_SEQ_MODE::EN_Q_LEN) { + if (!actSeqState_.inited) { + actSeqState_.curBatch = 0; + actSeqState_.prevPrefix = 0; + actSeqState_.curPrefix = (baseParams_->batchSize > 0) ? actualSeqLenGm_(0) : 0; + actSeqState_.inited = true; + actSeqState_.prevIndexOffset = 0; + actSeqState_.curIndexOffset = + CeilDivT(actSeqState_.curPrefix, static_cast(baseParams_->blockSize)); + } + } + int64_t curStepBatchSize = bsSize - batchOffset; + if (curStepBatchSize < static_cast(baseParams_->stepBatchSize)) { + UpdateStepBatchParams(curStepBatchSize); // 320 - 256 + if (batchOffset != 0) { + ComputeAivOffset(aivOffset, batchOffset); + } + } else { + curStepBatchSize = static_cast(baseParams_->stepBatchSize); + } + if ASCEND_IS_AIC { + AicProcess(aicOffset, batchOffset, mmQnLoops); + } + if ASCEND_IS_AIV { + AivProcess(aivOffset, batchOffset, curStepBatchSize, numHeadOffset, mmQnLoops); + } + } + if ASCEND_IS_AIC { + WaitFlag(A_EVENT0); + WaitFlag(A_EVENT1); + WaitFlag(B_EVENT0); + WaitFlag(B_EVENT1); + + WaitFlag(L0A_EVENT0); + WaitFlag(L0A_EVENT1); + WaitFlag(L0B_EVENT0); + WaitFlag(L0B_EVENT1); + + WaitFlag(L0C_EVENT0); + WaitFlag(L0C_EVENT1); + + WaitFlag(SCALE_EVENT); + CrossCoreWaitFlag(FINISH_VEC_CKVKR); + } +} + +template +template +__aicore__ inline void MlaPrologVecS1CubS2::AicProcess(AicOffset &aicOffset, int64_t batchOffset, + int64_t mmQnLoops) +{ + int64_t tokenXOffset = batchOffset * static_cast(baseParams_->headSizeX); + int64_t dequantScaleXOffset = batchOffset * static_cast(baseParams_->headSizeX) / 32; + GlobalTensor scaleAGm{}; + GlobalTensor scaleBGmDq{}; + GlobalTensor scaleBGmDkvKr{}; + if constexpr (std::is_same::value && isFp8E8m0) { + scaleAGm = dequantScaleXGm_[dequantScaleXOffset]; + scaleBGmDq = dequantScaleWDqGm_[aicOffset.dequantScaleWDqOffset]; + scaleBGmDkvKr = dequantScaleWDkvkrGm_[aicOffset.dequantScaleWDkvKrOffset]; + } + // MatmulCq ──> RmsNorm(Cq) + // [32, 7168] * [7168, 1536] = [32, 1536] + MatmulSplitN( + mmCqResGm_[aicOffset.cqResOffset], tokenXGm_[tokenXOffset], weightDqGm_[aicOffset.weightDqOffset], mmCqParam_, + UsedBlockParams{0, baseParams_->mm1BlockNum}, scaleAGm, scaleBGmDq); + CrossCoreSetFlag(FINISH_MM_CQ); + // MatmulCkvKr ──> RmsNorm(Ckv) + // └──> Rope(Kr) + // [32, 7168] * [7168, 512+64] = [32, 576] + CrossCoreWaitFlag(FINISH_VEC_CKVKR); + MatmulSplitN( + mmCkvKrResGm_[aicOffset.ckvKrResOffset], tokenXGm_[tokenXOffset], weightDkvKrGm_[aicOffset.weightDkvKrOffset], + mmCkvKrParam_, UsedBlockParams{0, baseParams_->mm2BlockNum}, scaleAGm, scaleBGmDkvKr); + CrossCoreSetFlag(FINISH_MM_CKVKR); + CrossCoreWaitFlag(FINISH_VEC_RMSNORM_CQ); + + if constexpr (std::is_same::value && isFp8E8m0) { + MatmulQcQr(aicOffset); + } else { + MatmulAndSyncQcQr(aicOffset); + } + if constexpr (std::is_same::value && isFp8E8m0) { + WaitFlag(SCALE_EVENT); // FP8场景下Scale不做db,需要等scale用完才能做mmQn + } + PreloadQnAndSync(aicOffset, mmQnLoops); + MatmulQnSyncDynamicQuantAndMulQr(aicOffset.qcOffset, aicOffset.weightUkOffset, + aicOffset.qnResOffset, mmQnLoops); + if constexpr (std::is_same::value && isFp8E8m0) { + SetFlag(SCALE_EVENT); // FP8场景下Scale不做db,需要mmQn用完才能做下一轮 + } + if constexpr (!needQnDynamicQuant) { + // MatmulQn的结果直接输出到 queryOut, qnOffset需要按Batch轴偏移 + aicOffset.qnResOffset += static_cast(baseParams_->stepBatchSize) * + static_cast(baseParams_->headSizeCkv) * + static_cast(baseParams_->numHeadSize); + } + if (unlikely(baseParams_->queryNormFlag == 1U)) { + aicOffset.rmsNormCqResOffset += + static_cast(baseParams_->stepBatchSize) * static_cast(baseParams_->headSizeCq); + } +} + +template +template +__aicore__ inline void MlaPrologVecS1CubS2::AivProcess(AivOffset &aivOffset, int64_t batchOffset, + int64_t curStepBatchSize, int64_t numHeadOffset, + int64_t mmQnLoops) +{ + int64_t tokenIndex = batchOffset + aivOffset.curBlockTokenOffset; + int64_t rmsNormCqResOffset = + baseParams_->queryNormFlag == 1U ? tokenIndex * baseParams_->headSizeCq : aivOffset.rmsNormCqOffset; + + if (batchOffset == 0) { + // 只需要搬运一次 + CopyGlobalParams(); + } + CopyInSinCos(tokenIndex, aivOffset.curVecToken, batchOffset, curStepBatchSize); + CrossCoreWaitFlag(FINISH_MM_CQ); + WaitAllCore(FINISH_VEC_ALL); + RmsNormCq(tokenIndex, aivOffset.rmsNormCqOffset, rmsNormCqResOffset, aivOffset.curVecToken, + aivOffset.curBlockTokenOffset); + // 由于RmsNormCq和MatmulQcQr的分核策略不一样,需要等所有vector上的RmsNormCq执行完成后才能启动MatmulQcQr + // 需要所有vector核上的RmsNormCq执行完成后,才发起MatmulQcQr的执行 + WaitAllCore(FINISH_VEC_ALL); + + // 聚合全部scale结果 + if constexpr ((MLAPT::enableDequantOpt || MLAPT::enableGroupComputeOpt) && + IsFullQuantMode()) { + DataCopy(dequantTool_.deQuantScaleCqLocal_, dequantTool_.deQuantScaleCqGm_, + ALIGN_BLOCK_SIZE / sizeof(float) * baseParams_->stepBatchSize); + } + CrossCoreSetFlag(FINISH_VEC_RMSNORM_CQ); + CrossCoreWaitFlag(FINISH_MM_CKVKR); + WaitAllCore(FINISH_VEC_ALL); + RmsNormRopeScatterCkvKr(tokenIndex, aivOffset.rmsNormCkvOffset, aivOffset.ropeKrOffset, aivOffset.curVecToken); + CrossCoreSetFlag(FINISH_VEC_CKVKR); + + DequantQcAndRopeQc(aivOffset, batchOffset, curStepBatchSize, numHeadOffset, mmQnLoops); +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::ComputeAicOffset(AicOffset &aicOffset, int64_t numHeadOffset) +{ + aicOffset.cqResOffset = baseParams_->mm1SingleCoreN * blockIdx_; // 64 * idx + aicOffset.weightDqOffset = static_cast(baseParams_->headSizeX) * aicOffset.cqResOffset; // 7168 * 64 * idx + aicOffset.dequantScaleWDqOffset = + static_cast(baseParams_->headSizeX) / 32 * aicOffset.cqResOffset; // 7168 * 64 * idx + + aicOffset.ckvKrResOffset = baseParams_->mm2SingleCoreN * blockIdx_; // (512 + 64) / 9 * idx = 64 * idx + aicOffset.weightDkvKrOffset = static_cast(baseParams_->headSizeX) * + aicOffset.ckvKrResOffset; // 7168 * (512 + 64) / 9 * idx = 7168 * 64 * idx + aicOffset.dequantScaleWDkvKrOffset = static_cast(baseParams_->headSizeX) / 32 * + aicOffset.ckvKrResOffset; // 7168 * (512 + 64) / 9 * idx = 7168 * 64 * idx + + if constexpr (MLAPT::enableGroupComputeOpt) { + aicOffset.weightUqOffset = static_cast(baseParams_->headSizeCq) * + (baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope) * + blockIdx_; // 1536 * (32 * 128 + 32 * 64) / 24 * idx = 1536 * 256 * idx + aicOffset.weightQrOffset = static_cast(baseParams_->headSizeCq) * + (baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope) * 2 * + (blockIdx_ - QC_CORE_NUM) + + static_cast(baseParams_->headSizeCq) * baseParams_->dimHeadSizeQc; + aicOffset.qCResOffset = baseParams_->dimHeadSizeQc * blockIdx_; + aicOffset.qRResOffset = + baseParams_->stepBatchSize * baseParams_->dimHeadSizeQc * QC_CORE_NUM + + baseParams_->dimHeadRope * 2 * (blockIdx_ - QC_CORE_NUM); // m=baseParams_->stepBatchSize + } else { + aicOffset.qcQrResOffset = baseParams_->mm3SingleCoreN * blockIdx_; // 192 * idx + aicOffset.weightUqQrOffset = + static_cast(baseParams_->headSizeCq) * + aicOffset.qcQrResOffset; // 1536 * (32 * 128 + 32 * 64) / 24 * idx = 1536 * 256 * idx + aicOffset.dequantScaleWuqqrOffset = aicOffset.weightUqQrOffset / 32; + } + + if (cubeBlockIdx_ < baseParams_->mm4BlockNum) { + if constexpr (std::is_same::value || std::is_same::value) { + aicOffset.qcOffset = static_cast(baseParams_->dimHeadSizeQc) * numHeadOffset; // 128 * idx + } else { + aicOffset.qcOffset = static_cast(baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope) * + numHeadOffset; // (128 + 64) * idx + } + aicOffset.weightUkOffset = static_cast(baseParams_->dimHeadSizeQc) * + static_cast(baseParams_->headSizeCkv) * numHeadOffset; // (128 * 512) * idx + aicOffset.qnResOffset = static_cast(baseParams_->headSizeCkv) * numHeadOffset; // BS, N, Dq + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::ComputeAivOffset(AivOffset &aivOffset, int64_t batchOffset) +{ + curStepVecBackToken_ = curVectorBlockNum_ / static_cast(vectorCoreNum_); + curStepVecFrontListNum_ = curVectorBlockNum_ % static_cast(vectorCoreNum_); + curStepVecFrontToken_ = curStepVecFrontListNum_ == 0 ? curStepVecBackToken_ : curStepVecBackToken_ + 1; + + aivOffset.curVecToken = blockIdx_ < curStepVecFrontListNum_ ? curStepVecFrontToken_ : curStepVecBackToken_; + aivOffset.curBlockTokenOffset = blockIdx_ < curStepVecFrontListNum_ ? + blockIdx_ * aivOffset.curVecToken : + blockIdx_ * aivOffset.curVecToken + curStepVecFrontListNum_; + aivOffset.rmsNormCqOffset = baseParams_->headSizeCq * aivOffset.curBlockTokenOffset; // 1536 * idx + aivOffset.rmsNormCkvOffset = + (baseParams_->headSizeCkv + baseParams_->dimHeadRope) * aivOffset.curBlockTokenOffset; // (512 + 64) * idx + aivOffset.ropeKrOffset = baseParams_->headSizeCkv + aivOffset.rmsNormCkvOffset; // 512 + (512 + 64) * idx + if constexpr (!MLAPT::enableDequantOpt) { + aivOffset.mmQnPreDequantOffset = + (baseParams_->headSizeQc + baseParams_->headSizeQr) * aivOffset.curBlockTokenOffset; + aivOffset.mmQnPreDequantResOffset = baseParams_->headSizeQc * aivOffset.curBlockTokenOffset; + aivOffset.ropeQrOffset = static_cast(baseParams_->dimHeadSizeQc) + + static_cast(baseParams_->numHeadSize) * + static_cast(baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope) * + aivOffset.curBlockTokenOffset; // 128 + 32 * (128 + 64) * idx + aivOffset.ropeQrResOffset = + static_cast(baseParams_->headSizeQr) * + (aivOffset.curBlockTokenOffset + batchOffset); // 32 * 64 * idx; // 按BS合轴切分step + } + + // 以下的Offset均是与batchOffset无关的,仅初始化一次即可。 + if (batchOffset == 0) { + if constexpr (MLAPT::enableDequantOpt) { + aivOffset.mmQnPreDequantOffset = baseParams_->mm3SingleCoreN * cubeBlockIdx_; // qcQrResOffset + aivOffset.mmQnPreDequantResOffset = baseParams_->mm3SingleCoreN / + (baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope) * + baseParams_->dimHeadSizeQc * cubeBlockIdx_; + aivOffset.ropeQrOffset = aivOffset.mmQnPreDequantOffset + baseParams_->dimHeadSizeQc; + aivOffset.ropeQrResOffset = + (baseParams_->mm3SingleCoreN / (baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope)) * + baseParams_->dimHeadRope * cubeBlockIdx_; + } + + if constexpr (MLAPT::enableGroupComputeOpt) { + aivOffset.qcScaleOffsetSplitN = + (baseParams_->headSizeQc + baseParams_->headSizeQr) / QC_CORE_NUM * (blockIdx_ / cvRatio_); + aivOffset.mmQnPreDequantResOffset = (baseParams_->headSizeQc) / QC_CORE_NUM * (blockIdx_ / cvRatio_); + aivOffset.mmQnPreDequantOffset = aivOffset.mmQnPreDequantResOffset; + aivOffset.ropeQrResSplitNOffset = (blockIdx_ - QC_CORE_NUM * cvRatio_) * baseParams_->dimHeadRope; + aivOffset.ropeQrSplitNOffset = baseParams_->headSizeQc + aivOffset.ropeQrResSplitNOffset; + } + } +} + +// Mlaprolog 支持int8进int32出以及mxfp8进fp32出, 参考MatmulQcQr +template +template +__aicore__ inline void +MlaPrologVecS1CubS2::MatmulSplitN(const GlobalTensor &tensorResGm, const GlobalTensor &tensorAGm, + const GlobalTensor &tensorBGm, const MMParams &mmPara, + const UsedBlockParams &mmBlockParams, const GlobalTensor &tensorAScaleGm, + const GlobalTensor &tensorBScaleGm) +{ + if constexpr (needCheckEmptyTensor && MLAPT::emptyMode == EMPTY_TENSOR_MODE::EMPTY_CACHE) { + return; + } + if (blockIdx_ < mmBlockParams.blockStartIdx || blockIdx_ >= mmBlockParams.blockEndIdx) { + return; + } + // 用于enableGroupComputeOpt场景 + if constexpr (needCheckAFullLoad) { + constexpr uint32_t mSize = + (sizeof(mmQcQrInputType) == sizeof(int8_t)) ? INT8_AFULLLOAD_MAX_MSIZE : BF16_AFULLLOAD_MAX_MSIZE; + bool isAFullLoad = (mmQcQrParam_.m <= mSize) ? true : false; + if (isAFullLoad) { + MatmulGroupComputeAFullLoad(tensorResGm, tensorAGm, tensorBGm, mmPara, + bufParam_); + return; + } + } + uint32_t nInput = mmPara.n; + uint32_t nL1SplitSize = mmPara.baseN; + uint32_t nL1loops = CeilDivT(nInput, nL1SplitSize); + uint32_t subNL1SplitSize = nL1SplitSize; + for (int64_t nL1 = 0; nL1 < nL1loops; nL1++) { + if (nL1 == nL1loops - 1) { + subNL1SplitSize = nInput - (nL1loops - 1) * nL1SplitSize; + } + MatmulSplitK(tensorResGm, tensorAGm, tensorBGm, mmPara, bufParam_, nL1 * nL1SplitSize, subNL1SplitSize, + tensorAScaleGm, tensorBScaleGm); + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::MatmulAndSyncQcQr(AicOffset &aicOffset) +{ + if constexpr (MLAPT::enableGroupComputeOpt) { + // MatmulQc + // 复用mmCqResGm_ workspace + MatmulSplitN( + mmQcQrResGm_[aicOffset.qCResOffset], rmsNormCqResGm_[aicOffset.rmsNormCqResOffset], + weightUqQrGm_[aicOffset.weightUqOffset], mmQcQrParam_, UsedBlockParams{0, QC_CORE_NUM}); + CrossCoreSetFlag(FINISH_MM_QC); + // MatmulQr + MatmulSplitN( + mmQcQrResGm_[aicOffset.qRResOffset], rmsNormCqResGm_[aicOffset.rmsNormCqResOffset], + weightUqQrGm_[aicOffset.weightQrOffset], mmQcQrParam_, + UsedBlockParams{QC_CORE_NUM, QC_CORE_NUM + QR_CORE_NUM}); + } else { + MatmulQcQr(aicOffset); + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::MatmulQcQr(AicOffset &aicOffset) +{ + if (blockIdx_ >= baseParams_->mm3BlockNum) { + return; + } + // RmsNorm(Cq) ──> MatmulQcQr ──> MatmulQn + // └──> Rope(Qr) + // [32, 1536] * [1536, 32*(128+64)] = [32, 32*192] + constexpr uint32_t mSize = + (sizeof(mmQcQrInputType) == sizeof(int8_t)) ? INT8_AFULLLOAD_MAX_MSIZE : BF16_AFULLLOAD_MAX_MSIZE; + bool isAFullLoad = (mmQcQrParam_.m <= mSize) ? true : false; + + uint32_t nInput = mmQcQrParam_.n; + uint32_t nL1SplitSize = mmQcQrParam_.baseN; + uint32_t nL1loops = CeilDivT(nInput, nL1SplitSize); + uint32_t subNL1SplitSize = nL1SplitSize; + if (isAFullLoad) { + if constexpr (std::is_same::value && isFp8E8m0) { + uint32_t offsetL1B = + L1_B_SIZE / 2 / sizeof(rmsNormCqOutputType); // // 2表示scale起始地址固定从L1B上ping的64k开始 + WaitFlag(SCALE_EVENT); + LoadL1AAndScale( + rmsNormCqResGm_[aicOffset.rmsNormCqResOffset], + dequantTool_.deQuantScaleCqGm_[aicOffset.dequantScaleCqOffset], mmQcQrParam_.m, mmQcQrParam_.k, + mmQcQrParam_.k, mmQcQrParam_.kScale, offsetL1B, bufParam_); + SetFlag(SCALE_EVENT); + } else { + LoadL1A(rmsNormCqResGm_[aicOffset.rmsNormCqResOffset], mmQcQrParam_.m, mmQcQrParam_.k, mmQcQrParam_.k, + bufParam_); + } + WaitFlag(A_EVENT0 + (bufParam_.aL1BufIter & 1u)); + } + + for (int64_t nL1 = 0; nL1 < nL1loops; nL1++) { + if (nL1 == nL1loops - 1) { + subNL1SplitSize = nInput - (nL1loops - 1) * nL1SplitSize; + } + if constexpr (std::is_same::value && isFp8E8m0) { + if (isAFullLoad) { + MatmulSplitK( + mmQcQrResGm_[aicOffset.qcQrResOffset], rmsNormCqResGm_[aicOffset.rmsNormCqResOffset], + weightUqQrGm_[aicOffset.weightUqQrOffset], mmQcQrParam_, bufParam_, nL1 * nL1SplitSize, + subNL1SplitSize, dequantTool_.deQuantScaleCqGm_[aicOffset.dequantScaleCqOffset], + deqScaleQcQrW_[aicOffset.dequantScaleWuqqrOffset]); + } else { + MatmulSplitK( + mmQcQrResGm_[aicOffset.qcQrResOffset], rmsNormCqResGm_[aicOffset.rmsNormCqResOffset], + weightUqQrGm_[aicOffset.weightUqQrOffset], mmQcQrParam_, bufParam_, nL1 * nL1SplitSize, + subNL1SplitSize, dequantTool_.deQuantScaleCqGm_[aicOffset.dequantScaleCqOffset], + deqScaleQcQrW_[aicOffset.dequantScaleWuqqrOffset]); + } + } else { + if (isAFullLoad) { + MatmulSplitK( + mmQcQrResGm_[aicOffset.qcQrResOffset], rmsNormCqResGm_[aicOffset.rmsNormCqResOffset], + weightUqQrGm_[aicOffset.weightUqQrOffset], mmQcQrParam_, bufParam_, nL1 * nL1SplitSize, + subNL1SplitSize); + } else { + MatmulSplitK( + mmQcQrResGm_[aicOffset.qcQrResOffset], rmsNormCqResGm_[aicOffset.rmsNormCqResOffset], + weightUqQrGm_[aicOffset.weightUqQrOffset], mmQcQrParam_, bufParam_, nL1 * nL1SplitSize, + subNL1SplitSize); + } + } + if constexpr (MLAPT::enableDequantOpt) { + CrossCoreSetFlag<0x2, PIPE_FIX>(FINISH_MM_QCQR_SPLIT_N); + } + } + if (isAFullLoad) { + SetFlag(A_EVENT0 + (bufParam_.aL1BufIter & 1u)); + bufParam_.aL1BufIter++; + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::PreloadQnAndSync(AicOffset &aicOffset, int64_t mmQnLoops) +{ + MatmulQnWeightPreload(aicOffset.weightUkOffset); + if constexpr (MLAPT::enableGroupComputeOpt) { + CrossCoreSetFlag(FINISH_MM_QR); + } else if constexpr (MLAPT::enableDequantOpt) { + // enableDequantOpt分支的CV同步由更细粒度的子函数控制 + return; + } else { + CrossCoreSetFlag(FINISH_MM_QCQR); + } + + if constexpr (IsFullQuantMode()) { + CrossCoreWaitFlag(FINISH_VEC_DEQUANT_QC); + } else { + WaitAllCore(FINISH_MM_ALL); + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::MatmulQnWeightPreload(int64_t weightUkOffset) +{ + uint32_t maxBlockIdx = MLAPT::enableGroupComputeOpt ? QC_CORE_NUM : baseParams_->mm4BlockNum; + if (blockIdx_ >= maxBlockIdx) { + return; + } + if (mmQnParam_.k > mmQnParam_.baseK) { + return; + } + int64_t weightOffset = weightUkOffset; + LoadL1B(weightUkGm_[weightOffset], mmQnParam_.n, mmQnParam_.n, mmQnParam_.k, + mmQnParam_.k, bufParam_); + WaitFlag(B_EVENT0 + (bufParam_.bL1BufIter & 1u)); + weightOffset += static_cast(baseParams_->dimHeadSizeQc) * static_cast(baseParams_->headSizeCkv); +} + +template +template +__aicore__ inline void +MlaPrologVecS1CubS2::MatmulQnSyncDynamicQuantAndMulQr(int64_t qcOffset, int64_t weightUkOffset, + int64_t qnResOffset, int64_t subLoopTimes) +{ + uint32_t maxBlockIdx = MLAPT::enableGroupComputeOpt ? QC_CORE_NUM : baseParams_->mm4BlockNum; + if (blockIdx_ >= maxBlockIdx) { + return; + } + // MatmulQcQr ──> MatmulQn ──> query_out + // [32, 128] * [128, 512] = [32, 512] + // [32, 2, 128] * [2, 128, 512] = [32, 2, 512] + bool needSparseSync = subLoopTimes > MAX_SYNC_FLAG_COUNT; + for (int64_t i = 0; i < subLoopTimes; i++) { + if constexpr (MLAPT::enableDequantOpt) { + if (!needSparseSync || i % 2 == 0 || i == subLoopTimes - 1) { + CrossCoreWaitFlag(FINISH_VEC_DEQUANT_QC_SPLIT_N); + } + } + if (unlikely(mmQnParam_.baseK < mmQnParam_.k)) { + uint32_t nInput = baseParams_->headSizeCkv; + uint32_t nL1SplitSize = mmQnParam_.baseN; + uint32_t nL1loops = CeilDivT(nInput, nL1SplitSize); + uint32_t subNL1SplitSize = nL1SplitSize; + for (int64_t nL1 = 0; nL1 < nL1loops; nL1++) { + if (nL1 == nL1loops - 1) { + subNL1SplitSize = nInput - (nL1loops - 1) * nL1SplitSize; + } + MatmulSplitK( + mmQnResGm_[qnResOffset], mmQcQrResDequantGm_[qcOffset], weightUkGm_[weightUkOffset], mmQnParam_, + bufParam_, nL1 * nL1SplitSize, subNL1SplitSize); + } + } else { + if (i < 1) { + MatmulFullLoad( + mmQnResGm_[qnResOffset], mmQcQrResDequantGm_[qcOffset], weightUkGm_[weightUkOffset], mmQnParam_, + bufParam_); + } else { + MatmulFullLoad( + mmQnResGm_[qnResOffset], mmQcQrResDequantGm_[qcOffset], weightUkGm_[weightUkOffset], mmQnParam_, + bufParam_); + } + } + if constexpr (std::is_same::value || std::is_same::value || + std::is_same::value) { + qcOffset += static_cast(baseParams_->dimHeadSizeQc); + } else { + qcOffset += static_cast(baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope); + } + weightUkOffset += + static_cast(baseParams_->dimHeadSizeQc) * static_cast(baseParams_->headSizeCkv); + qnResOffset += static_cast(baseParams_->headSizeCkv); + + if constexpr (needQnDynamicQuant) { + CrossCoreSetFlag(FINISH_MM_QN_SPLIT_N); + } + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::CopyInSinCos(int64_t tokenIndex, int64_t curVecToken, + int64_t batchOffset, int64_t curStepBatchSize) +{ + if constexpr (!MLAPT::enableRope) { + return; + } + LocalTensor shareTmpUb = shareBuffer_.Get(); + if constexpr (MLAPT::enableDequantOpt) { + // 如果是切N场景,mm3的每个C核都会做rope + if (cubeBlockIdx_ >= baseParams_->mm3BlockNum) { + return; + } + // 如果curStepBatchSize是偶数,则两个核平分;如果curStepBatchSize是奇数,则奇数核比偶数核多分一个 + // >> 1 是将curStepBatchSize分到每个vec核上; + uint32_t subBlockIdx_ = blockIdx_ % cvRatio_; + int64_t offset = (curStepBatchSize / cvRatio_) * subBlockIdx_ + batchOffset; + GatherSinCos(cosLocal_, sinLocal_, ropeCosGm_, ropeSinGm_, offset, + (curStepBatchSize + cvRatio_ - 1) / cvRatio_, shareTmpUb, + vectorRow_, baseParams_->dimHeadRope); + } else { + if (blockIdx_ >= curVectorBlockNum_) { + return; + } + GatherSinCos(cosLocal_, sinLocal_, ropeCosGm_, ropeSinGm_, tokenIndex, + curVecToken, shareTmpUb, vectorRow_, baseParams_->dimHeadRope); + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::CopyGlobalParams() +{ + // dequantScaleWDq + if constexpr (IsFullQuantMode()) { + DataCopyExtParams dequantCopyParams{1, static_cast(baseParams_->headSizeCq * sizeof(float)), 0, 0, 0}; + DataCopyPadExtParams dequantPadParams{false, 0, 0, 0}; + DataCopyPad(dequantScaleWDqLocal_, dequantScaleWDqGm_, dequantCopyParams, dequantPadParams); + } + + // rmsnormGammaCq + DataCopy(rmsnormGammaCqLocal_, rmsnormGammaCqGm_, baseParams_->headSizeCq); + + // rmsnormGammaCkv + DataCopy(rmsnormGammaCkvLocal_, rmsnormGammaCkvGm_, baseParams_->headSizeCkv); + + // smoothScaleCq + + if constexpr (IsFullQuantMode()) { + if (enableSmoothScalesCq_) { + DataCopyExtParams smoothCopyParams{1, static_cast(baseParams_->headSizeCq * sizeof(float)), 0, 0, + 0}; + DataCopyPadExtParams smoothPadParams{false, 0, 0, 0}; + DataCopyPad(smoothScaleCqLocal_, smoothScaleCqGm_, smoothCopyParams, smoothPadParams); + } + } + + // dequantScaleWDkvKr + if constexpr (IsFullQuantMode()) { + DataCopyExtParams dequantCopyParams{ + 1, static_cast((baseParams_->headSizeCkv + baseParams_->dimHeadRope) * sizeof(float)), 0, 0, 0}; + DataCopyPadExtParams dequantPadParams{false, 0, 0, 0}; + DataCopyPad(dequantScaleWDkvKrLocal_, dequantScaleWDkvkrGm_, dequantCopyParams, dequantPadParams); + } + + // quantScaleCkv + + if constexpr (IsFullQuantMode() && !isPertile) { + if constexpr (std::is_same::value || + std::is_same::value) { + DataCopyExtParams quantCopyParams{1, sizeof(float), 0, 0, 0}; + DataCopyPadExtParams quantPadParams{false, 0, 0, 0}; + DataCopyPad(quantScaleCkvLocal_, quantScaleCkvGm_, quantCopyParams, quantPadParams); + } else { // per_channel量化场景走else分支 + DataCopy(quantScaleCkvLocal_, quantScaleCkvGm_, baseParams_->headSizeCkv); + } + } + + // quantScaleCkr + if constexpr (std::is_same::value) { + DataCopyExtParams quantCopyParams{1, static_cast(baseParams_->dimHeadRope * sizeof(float)), 0, 0, 0}; + DataCopyPadExtParams quantPadParams{false, 0, 0, 0}; + DataCopyPad(quantScaleCkrLocal_, quantScaleCkrGm_, quantCopyParams, quantPadParams); + } +} + + +/** + * @brief RmsNormCq流程,融合了dynamicquant + 内部所需空间约为 curVecToken(128) * 8*4 + 8*4 + (4*vectorRow_*baseParams_->headSizeCq + 8)*4 = 28.0625K + */ +template +__aicore__ inline void MlaPrologVecS1CubS2::RmsNormCq(int64_t tokenIndex, int64_t rmsNormCqOffset, + int64_t rmsNormCqResOffset, int64_t curVecToken, + int64_t curBlockTokenOffset) +{ + if (blockIdx_ >= curVectorBlockNum_) { + return; + } + uint64_t dequantScaleXSize = 1; + if constexpr (std::is_same::value && isFp8E8m0) { + dequantScaleXSize = baseParams_->headSizeX / FP8_E4M3_BLOCK_SIZE; + } + uint64_t dequantScaleCqElementNum = dequantScaleCqSize_ / sizeof(dequantScaleType); + uint64_t dequantScaleXElementNum = Align(dequantScaleXSize, BYTE_BLOCK) / sizeof(float); + int64_t stepTokenIndex = tokenIndex; + LocalTensor outputLocal = shareBuffer_.Get(); + LocalTensor dequantScaleQcQr = + outputLocal[baseParams_->headSizeCq].template ReinterpretCast(); + LocalTensor dequantScaleXLocal = + dequantScaleQcQr[curVecToken * dequantScaleCqElementNum].template ReinterpretCast(); + LocalTensor shareTmpUb = dequantScaleXLocal[dequantScaleXElementNum].template ReinterpretCast(); + for (int64_t curVecTokenIdx = 0; curVecTokenIdx < curVecToken; curVecTokenIdx++) { + // MatmulCq ──> RmsNorm(Cq) ──> MatmulQcQr + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); // wait for vector operations to finish + + // dequantScaleXGm_ [BS , 1] 每个每个token对应一个系数,此处扩展为一个DataBlock + if constexpr (std::is_same::value || + (std::is_same::value && !isFp8E8m0)) { + DataCopyPad(dequantScaleXLocal, dequantScaleXGm_[tokenIndex], {1, sizeof(float), 0, 0}, {false, 0, 0, 0}); + } + + uint64_t scaleOffset = curVecTokenIdx * dequantScaleCqElementNum; + RmsNormParam rmsNormParams = { + baseParams_->reciprocalCq, // reciprocal + baseParams_->epsilonCq, // epsilon + static_cast(vectorRow_), // row + baseParams_->headSizeCq, // col + baseParams_->qcQrScale, // scale + baseParams_->isQcQrScaleEnable, // isScaleEnable + }; + + if constexpr (IsFullQuantMode()) { + RmsNormDynamicQuant(outputLocal, dequantScaleQcQr[scaleOffset], + mmCqResGm_[rmsNormCqOffset], rmsnormGammaCqLocal_, + smoothScaleCqLocal_, dequantScaleWDqLocal_, dequantScaleXLocal, + shareTmpUb, rmsNormParams, enableSmoothScalesCq_); + } else { + RmsNormNormal( + outputLocal, mmCqResGm_[rmsNormCqOffset], rmsnormGammaCqLocal_, dequantScaleWDqLocal_, + dequantScaleXLocal, shareTmpUb, rmsNormParams); + } + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + // RmsNorm(Cq)的结果拷进mmCqResGm_中,用于MatmulQcQr的A矩阵 + DataCopy(rmsNormCqResGm_[rmsNormCqResOffset], outputLocal, baseParams_->headSizeCq); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + rmsNormCqOffset += static_cast(baseParams_->headSizeCq); + rmsNormCqResOffset += static_cast(baseParams_->headSizeCq); + tokenIndex++; + } + + CopyDequantScaleCq(dequantScaleCqElementNum, dequantScaleQcQr, curVecToken, curBlockTokenOffset); + CopyQueryNormScale(stepTokenIndex, dequantScaleQcQr, curVecToken); +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::CopyDequantScaleCq(uint64_t dequantScaleCqElementNum, + LocalTensor &dequantScaleQcQr, + int64_t curVecToken, int64_t curBlockTokenOffset) +{ + if constexpr (std::is_same::value && isFp8E8m0) { + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + DataCopy(dequantTool_.deQuantScaleCqGm_[curBlockTokenOffset * dequantScaleCqElementNum], dequantScaleQcQr, + curVecToken * dequantScaleCqElementNum); + } else { + Brcb(dequantTool_.deQuantScaleCqLocal_[curBlockTokenOffset * FP32_BLOCK_ELEMENT_NUM], dequantScaleQcQr, + curVecToken, {1, 1}); + if constexpr (MLAPT::enableDequantOpt || MLAPT::enableGroupComputeOpt) { + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + DataCopy(dequantTool_.deQuantScaleCqGm_[curBlockTokenOffset * FP32_BLOCK_ELEMENT_NUM], + dequantTool_.deQuantScaleCqLocal_[curBlockTokenOffset * FP32_BLOCK_ELEMENT_NUM], + FP32_BLOCK_ELEMENT_NUM * curVecToken); + } + } +} +template +__aicore__ inline void MlaPrologVecS1CubS2::CopyQueryNormScale(int64_t stepTokenIndex, + LocalTensor &dequantScaleQcQr, + int64_t curVecToken) +{ + if (unlikely(baseParams_->queryNormFlag == 1U)) { + if constexpr (IsFullQuantMode()) { + DataCopyPad(dequantScaleQNormGm_[stepTokenIndex], dequantScaleQcQr, + {static_cast(curVecToken), sizeof(dequantScaleQNormType), 0, 0}); + } else if constexpr (std::is_same::value && isFp8E8m0) { + DataCopyPad( + dequantScaleQNormGm_[stepTokenIndex * + static_cast((baseParams_->headSizeCq / FP8_E4M3_BLOCK_SIZE))], + dequantScaleQcQr, + {static_cast(curVecToken), + static_cast(sizeof(dequantScaleQNormType) * (baseParams_->headSizeCq / FP8_E4M3_BLOCK_SIZE)), + 0, 0}); + } + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::ComputeBlkScatterOffsets(GlobalTensor indexGm, + int64_t tokenIndex, int64_t rows, + CkvkrParams &rmsNormAndScatterCkvParams, + CkvkrParams &ropeAndScatterKrParams) +{ + int64_t batchTokenIndex = 0; + int64_t batchSeqSize = 0; + int64_t batchIndexOffset = 0; + int64_t rowsThisStep = 0; + int64_t nextPageId = 0; + bool spill = false; + + // Advance batch cursor until the token index falls in [prevPrefix, curPrefix) + // curPrefix is an exclusive upper bound for the current batch + if constexpr (MLAPT::actualSeqMode == ACTUAL_SEQ_MODE::DISABLED) { + const int64_t batchIndex = tokenIndex / baseParams_->seq1Size; + batchTokenIndex = tokenIndex % baseParams_->seq1Size; + batchSeqSize = baseParams_->seq1Size; + batchIndexOffset = + batchIndex * CeilDivT(batchSeqSize, static_cast(baseParams_->blockSize)); // cacheIndex的offset + } else { + while (actSeqState_.curBatch + 1 < static_cast(baseParams_->batchSize) && + tokenIndex >= actSeqState_.curPrefix) { + actSeqState_.curBatch += 1; + actSeqState_.prevPrefix = actSeqState_.curPrefix; + actSeqState_.curPrefix = actualSeqLenGm_(actSeqState_.curBatch); + actSeqState_.prevIndexOffset = actSeqState_.curIndexOffset; + actSeqState_.curIndexOffset += CeilDivT(actSeqState_.curPrefix - actSeqState_.prevPrefix, + static_cast(baseParams_->blockSize)); + } + batchTokenIndex = tokenIndex - actSeqState_.prevPrefix; + batchSeqSize = actSeqState_.curPrefix - actSeqState_.prevPrefix; + batchIndexOffset = actSeqState_.prevIndexOffset; + } + + int64_t indexOffset = batchIndexOffset + batchTokenIndex / baseParams_->blockSize; + int64_t paBlkId = indexGm(indexOffset); // 取cacheIdx + + int64_t pageTokenOffset = paBlkId * baseParams_->blockSize; + int64_t tokenOffsetInPage = batchTokenIndex % baseParams_->blockSize; + + int64_t leftRowsInPage = baseParams_->blockSize - tokenOffsetInPage; + if (leftRowsInPage >= rows) { + rowsThisStep = rows; + spill = false; + nextPageId = -1; + } else { + rowsThisStep = rows - leftRowsInPage; + spill = true; + nextPageId = indexGm(indexOffset + 1); + } + + // --- Materialize for RMSNorm/CKV --- + MaterializeOffsetsWithHeadSize( + pageTokenOffset, tokenOffsetInPage, rowsThisStep, spill, nextPageId, baseParams_->headSizeCkv, + rmsNormAndScatterCkvParams); + + MaterializeOffsetsWithHeadSize( + pageTokenOffset, tokenOffsetInPage, rowsThisStep, spill, nextPageId, baseParams_->dimHeadRope, + ropeAndScatterKrParams); +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::RmsNormAndScatterCkv(LocalTensor &dequantScaleXLocal, + LocalTensor &shareTmpUb, + LocalTensor &cosLocalCkvKr, + LocalTensor &sinLocalCkvKr, + CkvkrParams rmsNormAndScatterCkvParams) +{ + // MatmulCkvKr ──> RmsNorm(Ckv) ──> Scatter(Ckv) + LocalTensor outputLocal = shareTmpUb.ReinterpretCast(); + uint32_t outputLocalSize = vectorRow_ * baseParams_->headSizeCkv * sizeof(kvCacheType); + if constexpr (isPertile) { // pertile量化场景,按照concat的最长长度申请内存 + uint32_t tileNum = baseParams_->headSizeCkv / baseParams_->tileSize; + outputLocalSize = vectorRow_ * (baseParams_->headSizeCkv * sizeof(kvCacheType) + tileNum * sizeof(float)); + outputLocalSize = Align(outputLocalSize, static_cast(BYTE_BLOCK)); + } + LocalTensor rmsNormShareTmpUb = shareTmpUb[outputLocalSize].template ReinterpretCast(); + + RmsNormParam rmsNormParams = { + baseParams_->reciprocalCkv, // reciprocal + baseParams_->epsilonCkv, // epsilon + static_cast(vectorRow_), // row + baseParams_->headSizeCkv, // col + baseParams_->kcScale, // scale + baseParams_->isKcScaleEnable, // isScaleEnable + }; + if constexpr (IsFullQuantMode()) { + RmsNormAndQuantizeCkv(outputLocal, rmsNormShareTmpUb, dequantScaleXLocal, rmsNormParams, + rmsNormAndScatterCkvParams); + } else { + RmsNormNormal( + outputLocal, mmCkvKrResGm_[rmsNormAndScatterCkvParams.offset], rmsnormGammaCkvLocal_, + dequantScaleWDkvKrLocal_, dequantScaleXLocal, rmsNormShareTmpUb, rmsNormParams); + } + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + // Scatter(Ckv) + // RmsNorm(Ckv) ──> Scatter(Ckv) ──> kv_cache_out + ScatterCkv(outputLocal, rmsNormAndScatterCkvParams); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::RmsNormAndQuantizeCkv(LocalTensor &outputLocal, + LocalTensor &rmsNormShareTmpUb, + LocalTensor &dequantScaleXLocal, + RmsNormParam rmsNormParams, + CkvkrParams rmsNormAndScatterCkvParams) +{ + // row = vectorRow_ = 1 col = baseParams_->headSizeCkv + LocalTensor inputLocal = rmsNormShareTmpUb.ReinterpretCast(); + LocalTensor sharedBuf = + inputLocal[vectorRow_ * baseParams_->headSizeCkv].template ReinterpretCast(); + RmsNormNormal( + inputLocal, mmCkvKrResGm_[rmsNormAndScatterCkvParams.offset], rmsnormGammaCkvLocal_, dequantScaleWDkvKrLocal_, + dequantScaleXLocal, sharedBuf, rmsNormParams); + + DataSyncBarrier(); + + if constexpr (isPertile) { + float kNopeClipAlpha = isFp8E8m0 ? 1.0f : kNopeClipAlphaGm_.GetValue(0); + PerTileQuantParams perTileQuantParams = { + static_cast(baseParams_->tileSize), // baseParams_->tileSize + static_cast(baseParams_->headSizeCkv / baseParams_->tileSize), // tileNum + kNopeClipAlpha, // alpha + static_cast(vectorRow_), // row + baseParams_->headSizeCkv // col + }; + if constexpr (std::is_same::value) { + QuantPerTile(outputLocal, inputLocal, sharedBuf, perTileQuantParams); + } else { + QuantPerTile8Bit(outputLocal, inputLocal, perTileQuantParams); + } + } else { + Rectangle rectangleParams{ + static_cast(vectorRow_), // row + static_cast(baseParams_->headSizeCkv), // col + static_cast(baseParams_->headSizeCkv) // columnStride + }; + if constexpr (std::is_same::value || + std::is_same::value) { + QuantPerTensor(outputLocal, inputLocal, quantScaleCkvLocal_, sharedBuf, rectangleParams); + } else { + QuantPerChannel(outputLocal, inputLocal, quantScaleCkvLocal_, sharedBuf, rectangleParams); + } + } + PipeBarrier(); +} + + +template +__aicore__ inline void MlaPrologVecS1CubS2::ScatterCkv(LocalTensor &outputLocal, + CkvkrParams rmsNormAndScatterCkvParams) +{ + if constexpr ((MLAPT::cacheMode == CACHE_MODE::PA_NZ) || (MLAPT::cacheMode == CACHE_MODE::PA_BSND) || + (MLAPT::cacheMode == CACHE_MODE::ND)) { + int64_t paTokenIndex; + if constexpr (MLAPT::cacheMode == CACHE_MODE::ND) { + paTokenIndex = rmsNormAndScatterCkvParams.tokenIndex; + } else { + paTokenIndex = cacheIndexGm_(rmsNormAndScatterCkvParams.tokenIndex); + } + ScatterCache( + kvCacheGm_, outputLocal, + ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, baseParams_->headSizeCkv, + baseParams_->dtileSize}); + // 刷新量化scale + if (isPertile && baseParams_->quantScaleRepoMode == 1U) { + // BSND: + uint32_t tileNum = baseParams_->headSizeCkv / baseParams_->tileSize; + LocalTensor quantScaleCkvInt8Tensor = + outputLocal[vectorRow_ * baseParams_->headSizeCkv].template ReinterpretCast(); + int64_t startOffset = 0; + int64_t startColOffset = baseParams_->headSizeCkv; + if (baseParams_->ckvkrRepoMode == 1U) { + startColOffset += baseParams_->headSizeKr * sizeof(krCacheType); + } + if constexpr ((MLAPT::cacheMode != CACHE_MODE::PA_NZ)) { + startOffset = startColOffset; + } + + ScatterCacheUnAligned( + kvCacheGm_[startOffset], quantScaleCkvInt8Tensor, + ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, + static_cast(tileNum * sizeof(float)), baseParams_->dtileSize}); + } + } else { + ScatterCacheMultiRows( + kvCacheGm_, outputLocal, + ScatterCacheParams{baseParams_->blockSize, rmsNormAndScatterCkvParams.cacheOffset, vectorRow_, + baseParams_->headSizeCkv, baseParams_->headSizeCkv, baseParams_->seq1Size, + rmsNormAndScatterCkvParams.tokenIndex}, + rmsNormAndScatterCkvParams.rowsInCurBatch, rmsNormAndScatterCkvParams.cacheOffset, + rmsNormAndScatterCkvParams.nextBatchOffset); + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::RopeAndScatterKr(LocalTensor &dequantScaleXLocal, + LocalTensor &shareTmpUb, + LocalTensor &cosLocalCkvKr, + LocalTensor &sinLocalCkvKr, + CkvkrParams ropeAndScatterKrParams) +{ + // MatmulCkvKr ──> Rope(Ckv) ──> Scatter(Kr) + LocalTensor outputKrLocal = shareTmpUb.ReinterpretCast(); + LocalTensor ropeShareTmpUb = outputKrLocal[baseParams_->dimHeadRope].template ReinterpretCast(); + int64_t stride = static_cast(baseParams_->headSizeCkv + baseParams_->headSizeKr); // 512 + 64 + + LocalTensor cosLocal; + LocalTensor sinLocal; + if constexpr (MLAPT::enableDequantOpt) { + cosLocal = cosLocalCkvKr[baseParams_->dimHeadRope * ropeAndScatterKrParams.curVecTokenIdx]; + sinLocal = sinLocalCkvKr[baseParams_->dimHeadRope * ropeAndScatterKrParams.curVecTokenIdx]; + } else { + cosLocal = cosLocal_[baseParams_->dimHeadRope * ropeAndScatterKrParams.curVecTokenIdx]; + sinLocal = sinLocal_[baseParams_->dimHeadRope * ropeAndScatterKrParams.curVecTokenIdx]; + } + Rectangle ropeParams{ + static_cast(vectorRow_), // row + static_cast(baseParams_->dimHeadRope), // col + static_cast(stride) // stride + }; + if constexpr ((std::is_same::value || + (std::is_same::value && !isFp8E8m0)) && + std::is_same::value) { + // Use ropeShareTmpUb directly (same as RopeQr full-quant). cos/sin live in cosLocal/sinLocal, + // not in this temp buffer; the old dimHeadRope*sizeof(ropeSinCosType) skip caused + // ENABLE_ROPE=0 write-through to read/write the wrong UB region for Kr dequant. + LocalTensor sharedBuf = ropeShareTmpUb; + // input为int32_t需在rope中做反量化, intput为float根据模板参数判断是否做反量化 + RotaryPosEmbPerTensor( + outputKrLocal, mmCkvKrResGm_[ropeAndScatterKrParams.offset], cosLocal, sinLocal, sharedBuf, ropeParams, + dequantScaleWDkvKrLocal_[baseParams_->headSizeCkv], dequantScaleXLocal); + } else if constexpr (std::is_same::value) { + LocalTensor inputLocal = ropeShareTmpUb.ReinterpretCast(); + LocalTensor sharedBuf = + ropeShareTmpUb.ReinterpretCast()[baseParams_->dimHeadRope * sizeof(ropeSinCosType)]; + RotaryPosEmbPerTensor::value, MLAPT::enableRope>( + inputLocal, mmCkvKrResGm_[ropeAndScatterKrParams.offset], cosLocal, sinLocal, sharedBuf, ropeParams); + RopePostQuantPerChannel(outputKrLocal, inputLocal, quantScaleCkrLocal_, sharedBuf, + vectorRow_ * baseParams_->dimHeadRope); + } else { + RotaryPosEmbPerTensor( + outputKrLocal, mmCkvKrResGm_[ropeAndScatterKrParams.offset], cosLocal, sinLocal, ropeShareTmpUb, + ropeParams); + } + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + // scatter(Kr) + // Rope(Kr) ──> Scatter(Kr) ──> kr_cache_out + ScatterKr(outputKrLocal, ropeAndScatterKrParams); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::ScatterKr(LocalTensor &outputKrLocal, + CkvkrParams ropeAndScatterKrParams) +{ + int64_t paTokenIndex; + if constexpr ((MLAPT::cacheMode == CACHE_MODE::PA_NZ) || (MLAPT::cacheMode == CACHE_MODE::PA_BSND) || + (MLAPT::cacheMode == CACHE_MODE::ND)) { + if constexpr (MLAPT::cacheMode == CACHE_MODE::ND) { + paTokenIndex = ropeAndScatterKrParams.tokenIndex; + } else { + paTokenIndex = cacheIndexGm_(ropeAndScatterKrParams.tokenIndex); + } + if (isPertile && baseParams_->ckvkrRepoMode == 1U) { + int64_t startOffset; + if constexpr ((MLAPT::cacheMode == CACHE_MODE::PA_NZ)) { + constexpr uint8_t col0 = ALIGN_BLOCK_SIZE / sizeof(kvCacheType); + startOffset = CeilDiv(baseParams_->headSizeCkv, col0) * col0 * baseParams_->blockSize; // 列方向的偏移 + } else { + startOffset = baseParams_->headSizeCkv; + } + LocalTensor outputKrInt8Tensor = outputKrLocal.template ReinterpretCast(); + ScatterCache( + kvCacheGm_[startOffset], outputKrInt8Tensor, + ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, + static_cast(baseParams_->dimHeadRope * sizeof(krCacheType)), + baseParams_->dtileSize}); + } else { + ScatterCache( + krCacheGm_, outputKrLocal, + ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, baseParams_->dimHeadRope, + baseParams_->dimHeadRope}); + } + } else { + ScatterCacheMultiRows( + krCacheGm_, outputKrLocal, + ScatterCacheParams{baseParams_->blockSize, ropeAndScatterKrParams.cacheOffset, vectorRow_, + baseParams_->dimHeadRope, baseParams_->dimHeadRope, baseParams_->seq1Size, + ropeAndScatterKrParams.tokenIndex}, + ropeAndScatterKrParams.rowsInCurBatch, ropeAndScatterKrParams.cacheOffset, + ropeAndScatterKrParams.nextBatchOffset); + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::RmsNormRopeScatterCkvKr(int64_t tokenIndex, int64_t rmsNormCkvOffset, + int64_t ropeKrOffset, int64_t curVecToken) +{ + if (blockIdx_ >= curVectorBlockNum_) { + return; + } + LocalTensor dequantScaleXLocal = shareBuffer_.Get(); + LocalTensor cosLocalCkvKr = + dequantScaleXLocal[FP32_BLOCK_ELEMENT_NUM].template ReinterpretCast(); + LocalTensor sinLocalCkvKr = cosLocalCkvKr[baseParams_->dimHeadRope * curVecToken]; + LocalTensor shareTmpUb = + sinLocalCkvKr[baseParams_->dimHeadRope * curVecToken].template ReinterpretCast(); + if constexpr (MLAPT::enableDequantOpt && MLAPT::enableRope) { + GatherSinCos(cosLocalCkvKr, sinLocalCkvKr, ropeCosGm_, ropeSinGm_, tokenIndex, + curVecToken, shareTmpUb, vectorRow_, baseParams_->dimHeadRope); + } + if constexpr (MLAPT::emptyMode == EMPTY_TENSOR_MODE::EMPTY_CACHE) { + return; + } + + for (int64_t curVecTokenIdx = 0; curVecTokenIdx < curVecToken; curVecTokenIdx++) { + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + if constexpr (std::is_same::value || + (std::is_same::value && !isFp8E8m0)) { + DataCopyPad(dequantScaleXLocal, dequantScaleXGm_[tokenIndex], {1, sizeof(float), 0, 0}, {false, 0, 0, 0}); + } + + CkvkrParams rmsNormAndScatterCkvParams{tokenIndex, rmsNormCkvOffset, curVecTokenIdx, 0, 0, 0}; + + CkvkrParams ropeAndScatterKrParams{tokenIndex, ropeKrOffset, curVecTokenIdx, 0, 0, 0}; + + if constexpr ((MLAPT::cacheMode == CACHE_MODE::PA_BLK_BSND) || (MLAPT::cacheMode == CACHE_MODE::PA_BLK_NZ)) { + // --- Compute headSize-independent variables once --- + ComputeBlkScatterOffsets(cacheIndexGm_, tokenIndex, vectorRow_, rmsNormAndScatterCkvParams, + ropeAndScatterKrParams); + } + + RmsNormAndScatterCkv(dequantScaleXLocal, shareTmpUb, cosLocalCkvKr, sinLocalCkvKr, rmsNormAndScatterCkvParams); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + RopeAndScatterKr(dequantScaleXLocal, shareTmpUb, cosLocalCkvKr, sinLocalCkvKr, ropeAndScatterKrParams); + + tokenIndex += 1; + rmsNormCkvOffset += static_cast(baseParams_->headSizeCkv + baseParams_->dimHeadRope); + ropeKrOffset += static_cast(baseParams_->headSizeCkv + baseParams_->dimHeadRope); + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::RopeQr(int64_t ropeQrOffset, int64_t ropeQrResOffset, + int64_t curVecToken, int64_t curBlockTokenOffset) +{ + if (blockIdx_ >= curVectorBlockNum_) { + return; + } + uint64_t stride = static_cast(baseParams_->dimHeadRope + baseParams_->dimHeadSizeQc); + + LocalTensor outputLocal; + LocalTensor channelDeqScaleLocal = shareBuffer_.Get(); + if constexpr (IsFullQuantMode()) { + uint64_t row = baseParams_->numHeadSize; + uint64_t col = baseParams_->dimHeadRope; + DataCopyExtParams copyParams{static_cast(row), static_cast(col * sizeof(float)), + static_cast((stride - col) * sizeof(float)), 0, 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + DataCopyPad(channelDeqScaleLocal, deqScaleQcQrW_[baseParams_->dimHeadSizeQc], copyParams, + padParams); // 复用内存 + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + + outputLocal = channelDeqScaleLocal[row * col].template ReinterpretCast(); + } else { + outputLocal = shareBuffer_.Get(); + } + LocalTensor ropeShareTmpUb = outputLocal[baseParams_->headSizeQr].template ReinterpretCast(); + + // row, col, stride + Rectangle ropeParams{baseParams_->numHeadSize, baseParams_->dimHeadRope, static_cast(stride)}; + + for (int64_t curVecTokenIdx = 0; curVecTokenIdx < curVecToken; curVecTokenIdx++) { + // MatmulQcQr ──> Rope(Qr) ──> query_rope_out + if constexpr (IsFullQuantMode()) { + RotaryPosEmbPerTensor( + outputLocal, mmQcQrResGm_[ropeQrOffset], cosLocal_[baseParams_->dimHeadRope * curVecTokenIdx], + sinLocal_[baseParams_->dimHeadRope * curVecTokenIdx], ropeShareTmpUb, ropeParams, channelDeqScaleLocal, + dequantTool_.deQuantScaleCqLocal_[(curBlockTokenOffset + curVecTokenIdx) * FP32_BLOCK_ELEMENT_NUM]); + } else { + RotaryPosEmbPerTensor( + outputLocal, mmQcQrResGm_[ropeQrOffset], cosLocal_[baseParams_->dimHeadRope * curVecTokenIdx], + sinLocal_[baseParams_->dimHeadRope * curVecTokenIdx], ropeShareTmpUb, ropeParams); + } + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + DataCopy(qrOutGm_[ropeQrResOffset], outputLocal, baseParams_->headSizeQr); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + ropeQrOffset += static_cast(baseParams_->numHeadSize) * + static_cast(baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope); + ropeQrResOffset += static_cast(baseParams_->headSizeQr); + } +} + +// 用于enableGroupComputeOpt场景 BS = 1 , G = 8 +template +__aicore__ inline void MlaPrologVecS1CubS2::RopeQrSplitNGroupCase(int64_t ropeQrOffset, int64_t ropeQrResOffset) +{ + if (blockIdx_ < QC_CORE_NUM * cvRatio_ || blockIdx_ >= (QC_CORE_NUM + QR_CORE_NUM) * cvRatio_) { + return; + } + + uint32_t row = baseParams_->numHeadSize / QC_CORE_NUM; + uint32_t col = baseParams_->dimHeadRope; + int64_t stride = col; + int64_t strideScale = static_cast(baseParams_->dimHeadRope + baseParams_->dimHeadSizeQc); + int64_t deqScaleOffset = baseParams_->dimHeadSizeQc + row * + (baseParams_->dimHeadSizeQc + baseParams_->dimHeadRope) * + (blockIdx_ - QC_CORE_NUM * cvRatio_); + LocalTensor channelDeqScaleLocal = shareBuffer_.Get(); + DataCopyExtParams copyParams{static_cast(row), static_cast(col * sizeof(float)), + static_cast((strideScale - col) * sizeof(float)), 0, 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + DataCopyPad(channelDeqScaleLocal, deqScaleQcQrW_[deqScaleOffset], copyParams, padParams); // 复用内存 + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + + LocalTensor outputLocal = + channelDeqScaleLocal[row * col].template ReinterpretCast(); + LocalTensor ropeShareTmpUb = outputLocal[baseParams_->headSizeQr].template ReinterpretCast(); + + Rectangle ropeParams{ + static_cast(col), // row + static_cast(col), // col + static_cast(stride) // stride + }; + + for (int64_t curVecTokenIdx = 0; curVecTokenIdx < curVectorBlockNum_; curVecTokenIdx++) { + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + // sharetmp共用,待V算完再执行MTE2搬运规避风险 + if constexpr (MLAPT::enableRope) { + GatherSinCos(cosLocal_, sinLocal_, ropeCosGm_, ropeSinGm_, + curVecTokenIdx * baseParams_->dimHeadRope, 1, ropeShareTmpUb, + vectorRow_, baseParams_->dimHeadRope); + } + + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + + RotaryPosEmbPerTensor( + outputLocal, mmQcQrResGm_[ropeQrOffset], cosLocal_, sinLocal_, ropeShareTmpUb, ropeParams, + channelDeqScaleLocal, dequantTool_.deQuantScaleCqLocal_); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + DataCopy(qrOutGm_[ropeQrResOffset], outputLocal, row * col); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + ropeQrOffset += static_cast(baseParams_->numHeadSize) * static_cast(baseParams_->dimHeadRope); + ropeQrResOffset += static_cast(baseParams_->headSizeQr); + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::RopeQrSplitN(const RopeQrSplitNParams &ropeQrSplitNParams) +{ + uint32_t ropeCnt = ropeQrSplitNParams.ropeCnt; + uint32_t colQr = baseParams_->dimHeadRope; + uint32_t deQuantScaleCqOffset = ropeQrSplitNParams.deQuantScaleCqOffset; + DataCopyParams outputRopeParams{ + static_cast(ropeCnt), static_cast(colQr * sizeof(ropeOutputType) / ALIGN_BLOCK_SIZE), 0, + static_cast(ropeQrSplitNParams.ropeDstStride * sizeof(ropeOutputType) / ALIGN_BLOCK_SIZE)}; + + Rectangle ropeParams{ + ropeCnt, // row + colQr, // col + static_cast(ropeQrSplitNParams.ropeStride) // stride + }; + + GlobalTensor inputGmRope = mmQcQrResGm_[ropeQrSplitNParams.ropeQrOffset]; + GlobalTensor outputGmRope = qrOutGm_[ropeQrSplitNParams.ropeQrResOffset]; + + LocalTensor shareTmpUb = shareBuffer_.Get(); + LocalTensor outputLocalRope = shareTmpUb.ReinterpretCast(); + LocalTensor ropeShareTmpUb = outputLocalRope[ropeCnt * colQr].template ReinterpretCast(); + + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + if constexpr (IsFullQuantMode()) { + GlobalTensor deqScaleRope = deqScaleQcQrW_[ropeQrSplitNParams.ropeQrOffset]; + RotaryPosEmbPerHead( + outputLocalRope, inputGmRope[ropeQrSplitNParams.inputOffsetRope], + cosLocal_[ropeQrSplitNParams.sinCosOffset], sinLocal_[ropeQrSplitNParams.sinCosOffset], ropeShareTmpUb, + ropeParams, ropeQrSplitNParams.ropeStride, deqScaleRope[ropeQrSplitNParams.deqScaleOffset], + dequantTool_.deQuantScaleCqLocal_[deQuantScaleCqOffset]); + } else { + RotaryPosEmbPerHead( + outputLocalRope, inputGmRope[ropeQrSplitNParams.inputOffsetRope], + cosLocal_[ropeQrSplitNParams.sinCosOffset], sinLocal_[ropeQrSplitNParams.sinCosOffset], ropeShareTmpUb, + ropeParams, ropeQrSplitNParams.ropeStride); + } + + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + + DataCopy(outputGmRope[ropeQrSplitNParams.outputOffsetRope], outputLocalRope, outputRopeParams); +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::CastQcQrSplitN(const CastQcQrSplitNParams &castQcQrSplitN) +{ + uint32_t row = mmQcQrParam_.m; + uint32_t colQc = baseParams_->dimHeadSizeQc; + uint32_t colQcSingle = colQc / cvRatio_; + uint32_t count = row * colQcSingle; + + DataCopyParams inputCopyParams{ + static_cast(row), static_cast(colQcSingle * sizeof(mmQcQrOutputType) / ALIGN_BLOCK_SIZE), + static_cast(castQcQrSplitN.srcStride * sizeof(mmQcQrOutputType) / ALIGN_BLOCK_SIZE), 0}; + + DataCopyParams outputCopyParams{ + static_cast(row), static_cast(colQcSingle * sizeof(mmQnInputType) / ALIGN_BLOCK_SIZE), 0, + static_cast(castQcQrSplitN.dstStride * sizeof(mmQnInputType) / ALIGN_BLOCK_SIZE)}; + + + GlobalTensor inputGm = mmQcQrResGm_[castQcQrSplitN.mmQnPreCastOffset]; + GlobalTensor outputGm = mmQcQrResDequantGm_[castQcQrSplitN.mmQnPreCastResOffset]; + + LocalTensor shareTmpUb = shareBuffer_.Get(); + LocalTensor inputLocal = shareTmpUb.ReinterpretCast(); + LocalTensor outputLocal = inputLocal.template ReinterpretCast(); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + DataCopy(inputLocal, inputGm[castQcQrSplitN.inputOffset], inputCopyParams); + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + // cast + Cast(outputLocal, inputLocal, RoundMode::CAST_RINT, count); + SetFlag(EVENT_ID2); + // copy out + WaitFlag(EVENT_ID2); + DataCopy(outputGm[castQcQrSplitN.outputOffset], outputLocal, outputCopyParams); +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::DequantQcQrSplitN(const DequantQcQrSplitNParams &dequantQcQrSplitN) +{ + uint32_t row = mmQcQrParam_.m; + uint32_t colQc = baseParams_->dimHeadSizeQc; + uint32_t colQcSingle = colQc / cvRatio_; + uint32_t count = row * colQcSingle; + + DataCopyParams inputCopyParams{ + static_cast(row), static_cast(colQcSingle * sizeof(mmQcQrOutputType) / ALIGN_BLOCK_SIZE), + static_cast(dequantQcQrSplitN.srcStride * sizeof(mmQcQrOutputType) / ALIGN_BLOCK_SIZE), 0}; + + DataCopyParams outputCopyParams{ + static_cast(row), static_cast(colQcSingle * sizeof(mmQnInputType) / ALIGN_BLOCK_SIZE), 0, + static_cast(dequantQcQrSplitN.dstStride * sizeof(mmQnInputType) / ALIGN_BLOCK_SIZE)}; + + + GlobalTensor inputGm = mmQcQrResGm_[dequantQcQrSplitN.mmQnPreDequantOffset]; + GlobalTensor scale1Gm = deqScaleQcQrW_[dequantQcQrSplitN.mmQnPreDequantOffset]; + GlobalTensor outputGm = mmQcQrResDequantGm_[dequantQcQrSplitN.mmQnPreDequantResOffset]; + + LocalTensor shareTmpUb = shareBuffer_.Get(); + LocalTensor scale2Local = dequantTool_.deQuantScaleCqLocal_; + LocalTensor inputLocal = shareTmpUb.ReinterpretCast(); + // 可以和inputLocal共享内存地址,减少UB使用 + LocalTensor computeLocal = shareTmpUb.ReinterpretCast(); + LocalTensor scaleLocal = inputLocal[count + FP32_BLOCK_ELEMENT_NUM].template ReinterpretCast(); + // outputLocal比scaleLocal占用UB少,且不会同时使用,故可以复用UB内存 + LocalTensor outputLocal = scaleLocal.template ReinterpretCast(); + + Rectangle dequantParams{ + row, // row + colQcSingle, // col + colQcSingle // columnStride + }; + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + DataCopy(inputLocal, inputGm[dequantQcQrSplitN.inputOffset], inputCopyParams); + DataCopy(scaleLocal, scale1Gm[dequantQcQrSplitN.inputOffset], colQcSingle); + SetFlag(EVENT_ID1); + // cast + WaitFlag(EVENT_ID1); + Dequant(computeLocal, inputLocal, scaleLocal, scale2Local, dequantParams); + AscendC::PipeBarrier(); + // cast + Cast(outputLocal, computeLocal, RoundMode::CAST_RINT, count); + SetFlag(EVENT_ID2); + // copy out + WaitFlag(EVENT_ID2); + DataCopy(outputGm[dequantQcQrSplitN.outputOffset], outputLocal, outputCopyParams); +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::DequantAndRopeSplitNSyncMMQcQr(int64_t mmQnPreDequantOffset, + int64_t mmQnPreDequantResOffset, + int64_t ropeQrOffset, + int64_t ropeQrResOffset) +{ + if (cubeBlockIdx_ >= baseParams_->mm3BlockNum) { + return; + } + // mmQcQr一个C核算stepBatchSize * singleN,每两个V核做对应C核输出的Qc部分的dequant + // singleN一定整除(dimHeadSizeQc + dimHeadSizeQc),保证dequant能找到dimHeadSizeQc部分的起始位置 + uint32_t subBlockIdx_ = blockIdx_ % cvRatio_; + uint32_t colCube = mmQcQrParam_.baseN; + uint32_t colQc = baseParams_->dimHeadSizeQc; + uint32_t colQr = baseParams_->dimHeadRope; + uint32_t oriCol = mmQcQrParam_.n; + uint32_t colOffsetCube = 0; + + // DequantSplitN参数 + uint32_t srcStride = baseParams_->headSizeQc + baseParams_->headSizeQr - baseParams_->dimHeadSizeQc / cvRatio_; + uint32_t dstStride = baseParams_->headSizeQc - baseParams_->dimHeadSizeQc / cvRatio_; + uint32_t colQcSingle = colQc / cvRatio_; + uint32_t colOffsetVec = 0; + uint32_t inputOffset = subBlockIdx_ * colQcSingle; + uint32_t outputOffset = inputOffset; + + // RopeQrSplitN参数 + uint32_t ropeDstStride = baseParams_->headSizeQr - colQr; + // CV1:2和1:1,M轴方向切成2块处理 + uint32_t ropeCntDown = mmQcQrParam_.m / 2; + // cv1:1时,为了防止UB内存溢出,每次处理一半的数量 + uint32_t ropeCnt = cvRatio_ == 1 ? ropeCntDown : (mmQcQrParam_.m + subBlockIdx_) / cvRatio_; + int64_t ropeStride = static_cast(baseParams_->numHeadSize) * static_cast(colQc + colQr); + uint32_t inputOffsetRope = ropeCntDown * subBlockIdx_ * ropeStride; + uint32_t outputOffsetRope = ropeCntDown * subBlockIdx_ * baseParams_->headSizeQr; + uint32_t deqScaleOffset = 0; + uint32_t colOffsetRope = 0; + uint32_t deQuantScaleCqOffset = ropeCntDown * subBlockIdx_ * FP32_BLOCK_ELEMENT_NUM; + // cube一次处理row*colCube,对应的两个vec一次处理row*colQc,两vec之间切colQc + // 等cube生产足够数据了以后,vec开始消费 + uint32_t dequantLoopCount = 0; + uint32_t totalDequantLoops = CeilDiv(oriCol, (colQc + colQr)); + bool needSparseSync = totalDequantLoops > MAX_SYNC_FLAG_COUNT; + while (colOffsetCube < oriCol) { // 循环CeilDiv(oriCol, colCube)次 + colOffsetCube += colCube; + if (colOffsetCube > oriCol) { // 当oriCol不被colCube整除时,mm最后一个base块需要刷新col end + colOffsetCube = oriCol; + } + CrossCoreWaitFlag(FINISH_MM_QCQR_SPLIT_N); + // DequantSplitN + while (colOffsetVec + colQc <= colOffsetCube) { // 循环singleNumHeadSize次 + if constexpr (IsFullQuantMode()) { + DequantQcQrSplitN(DequantQcQrSplitNParams{mmQnPreDequantOffset, mmQnPreDequantResOffset, inputOffset, + outputOffset, srcStride, dstStride}); + } else if constexpr (std::is_same::value && isFp8E8m0) { + CastQcQrSplitN(CastQcQrSplitNParams{mmQnPreDequantOffset, mmQnPreDequantResOffset, inputOffset, + outputOffset, srcStride, dstStride}); + } + if (!needSparseSync || dequantLoopCount % 2 == 0 || dequantLoopCount == totalDequantLoops - 1) { + CrossCoreSetFlag(FINISH_VEC_DEQUANT_QC_SPLIT_N); + } + dequantLoopCount++; + colOffsetVec += (colQc + colQr); + inputOffset += (colQc + colQr); + outputOffset += colQc; + } + // RopeQrSplitN + while ((colOffsetRope + colQc + colQr) <= colOffsetCube) { + RopeQrSplitN(RopeQrSplitNParams{ropeQrOffset, ropeQrResOffset, inputOffsetRope, deqScaleOffset, + outputOffsetRope, ropeStride, ropeDstStride, deQuantScaleCqOffset, 0, + ropeCnt}); + // cv1:1时,还需处理第二次,第二次在第一次的基础上计算偏移等 + if (cvRatio_ == 1) { + RopeQrSplitN(RopeQrSplitNParams{ + ropeQrOffset, ropeQrResOffset, static_cast(inputOffsetRope + ropeCntDown * ropeStride), + deqScaleOffset, static_cast(outputOffsetRope + ropeCntDown * baseParams_->headSizeQr), + ropeStride, ropeDstStride, ropeCntDown * FP32_BLOCK_ELEMENT_NUM, + ropeCntDown * baseParams_->dimHeadRope, (mmQcQrParam_.m + 1) / 2}); + } + + colOffsetRope += (colQc + colQr); + inputOffsetRope += (colQc + colQr); + deqScaleOffset += (colQc + colQr); + outputOffsetRope += colQr; + } + } +} + +template +template +__aicore__ inline void MlaPrologVecS1CubS2::DequantQcAndRopeQc(AivOffset &aivOffset, int64_t batchOffset, + int64_t curStepBatchSize, int64_t numHeadOffset, + int64_t mmQnLoops) +{ + // 根据不同分支条件处理 + if constexpr (MLAPT::enableGroupComputeOpt) { + CrossCoreWaitFlag(FINISH_MM_QC); + DequantQcSplitNGroupCase(aivOffset.mmQnPreDequantOffset, aivOffset.mmQnPreDequantResOffset, + aivOffset.qcScaleOffsetSplitN); + CrossCoreSetFlag(FINISH_VEC_DEQUANT_QC); + CrossCoreWaitFlag(FINISH_MM_QR); + RopeQrSplitNGroupCase(aivOffset.ropeQrSplitNOffset, aivOffset.ropeQrResSplitNOffset); + } else { + if constexpr (MLAPT::enableDequantOpt) { + DequantAndRopeSplitNSyncMMQcQr(aivOffset.mmQnPreDequantOffset, aivOffset.mmQnPreDequantResOffset, + aivOffset.ropeQrOffset, aivOffset.ropeQrResOffset); + } else if constexpr (IsFullQuantMode()) { + CrossCoreWaitFlag(FINISH_MM_QCQR); + WaitAllCore(FINISH_VEC_ALL); + DequantQc(aivOffset.mmQnPreDequantOffset, aivOffset.mmQnPreDequantResOffset, aivOffset.curVecToken, + aivOffset.curBlockTokenOffset); + WaitAllCore(FINISH_VEC_ALL); + CrossCoreSetFlag(FINISH_VEC_DEQUANT_QC); + RopeQr(aivOffset.ropeQrOffset, aivOffset.ropeQrResOffset, aivOffset.curVecToken, + aivOffset.curBlockTokenOffset); + } else if constexpr (std::is_same::value && isFp8E8m0) { + CrossCoreWaitFlag(FINISH_MM_QCQR); + WaitAllCore(FINISH_VEC_ALL); + CastQc(aivOffset.mmQnPreDequantOffset, aivOffset.mmQnPreDequantResOffset, aivOffset.curVecToken); + WaitAllCore(FINISH_VEC_ALL); + CrossCoreSetFlag(FINISH_VEC_DEQUANT_QC); + RopeQr(aivOffset.ropeQrOffset, aivOffset.ropeQrResOffset, aivOffset.curVecToken, + aivOffset.curBlockTokenOffset); + } else { + CrossCoreWaitFlag(FINISH_MM_QCQR); + WaitAllCore(FINISH_VEC_ALL); + RopeQr(aivOffset.ropeQrOffset, aivOffset.ropeQrResOffset, aivOffset.curVecToken, + aivOffset.curBlockTokenOffset); + } + + if constexpr (needQnDynamicQuant && !(MLAPT::enableDequantOpt)) { + // 非切N场景,需要等待全部rope的结果做完并搬到GM + WaitAllCore(FINISH_VEC_ALL); + } + if constexpr (needQnDynamicQuant) { + DynamicQuantQnAndMulQrSyncMMQn(batchOffset, curStepBatchSize, numHeadOffset, mmQnLoops); + } + aivOffset.ropeQrResOffset += + static_cast(baseParams_->stepBatchSize) * static_cast(baseParams_->headSizeQr); + } +} + +// 用于算力切分dequant切N场景 +template +__aicore__ inline void MlaPrologVecS1CubS2::DequantQcSplitNGroupCase(int64_t mmQnPreDequantOffset, + int64_t mmQnPreDequantResOffset, + int64_t qcQrScaleOffset) +{ + if (blockIdx_ >= QC_CORE_NUM * cvRatio_) { + return; + } + uint32_t subBlockIdx_ = blockIdx_ % cvRatio_; + uint32_t oriCol = (baseParams_->headSizeQc) / QC_CORE_NUM; + uint32_t curCol = (baseParams_->headSizeQc) / (QC_CORE_NUM * cvRatio_); + uint32_t srcStride = baseParams_->headSizeQc - curCol; + uint32_t dstStride = baseParams_->headSizeQc - curCol; + LocalTensor shareTmpUb = shareBuffer_.Get(); + Rectangle rectangleParams{ + static_cast(mmQcQrParam_.m), // row + static_cast(curCol), // col + static_cast(baseParams_->headSizeQc) // columnStride + }; + DequantSplitNQc(mmQcQrResDequantGm_[mmQnPreDequantResOffset], mmQcQrResGm_[mmQnPreDequantOffset], + deqScaleQcQrW_[qcQrScaleOffset], dequantTool_.deQuantScaleCqLocal_, shareTmpUb, rectangleParams, + oriCol, dstStride, subBlockIdx_); +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::DequantQc(int64_t mmQnPreDequantOffset, + int64_t mmQnPreDequantResOffset, int64_t curVecToken, + int64_t curBlockTokenOffset) +{ + if (blockIdx_ >= curVectorBlockNum_) { + return; + } + + Rectangle rectangleParams{ + static_cast(baseParams_->stepNumHeadDequant), // row + static_cast(baseParams_->dimHeadSizeQc), // col + static_cast(baseParams_->dimHeadRope + baseParams_->dimHeadSizeQc) // columnStride + }; + + for (int64_t curVecTokenIdx = 0; curVecTokenIdx < curVecToken; curVecTokenIdx++) { + LocalTensor shareTmpUb = shareBuffer_.Get(); + DequantPerTokenQc( + mmQcQrResDequantGm_[mmQnPreDequantResOffset], mmQcQrResGm_[mmQnPreDequantOffset], deqScaleQcQrW_, + dequantTool_.deQuantScaleCqLocal_[(curBlockTokenOffset + curVecTokenIdx) * FP32_BLOCK_ELEMENT_NUM], + shareTmpUb, rectangleParams, baseParams_->numHeadSize); + mmQnPreDequantOffset += baseParams_->headSizeQc + baseParams_->headSizeQr; + mmQnPreDequantResOffset += baseParams_->headSizeQc; + } +} + +template +__aicore__ inline void MlaPrologVecS1CubS2::CastQc(int64_t mmQnPreCastOffset, int64_t mmQnPreCastResOffset, + int64_t curVecToken) +{ + if (blockIdx_ >= curVectorBlockNum_) { + return; + } + + Rectangle rectangleCastParams{ + static_cast(baseParams_->stepNumHeadDequant), // row + static_cast(baseParams_->dimHeadSizeQc), // col + static_cast(baseParams_->dimHeadRope + baseParams_->dimHeadSizeQc) // columnStride + }; + + for (int64_t curVecTokenIdx = 0; curVecTokenIdx < curVecToken; curVecTokenIdx++) { + LocalTensor shareTmpUb = shareBuffer_.Get(); + CastPerTokenQc(mmQcQrResDequantGm_[mmQnPreCastResOffset], mmQcQrResGm_[mmQnPreCastOffset], shareTmpUb, + rectangleCastParams, baseParams_->numHeadSize); + + mmQnPreCastOffset += baseParams_->headSizeQc + baseParams_->headSizeQr; + mmQnPreCastResOffset += baseParams_->headSizeQc; + } +} + +template +__aicore__ inline void +MlaPrologVecS1CubS2::DynamicQuantQnAndMulQrSyncMMQn(int64_t batchOffset, int64_t curStepBatchSize, + int64_t numHeadOffset, int64_t mmQnLoops) +{ + // 如果curStepBatchSize是偶数,则两个核平分;如果curStepBatchSize是奇数,则奇数核比偶数核多分一个 + // >> 1 是将curStepBatchSize分到每个vec核上; + int64_t curStepBatchSizeVec = (curStepBatchSize + (blockIdx_ % cvRatio_)) / cvRatio_; + if (blockIdx_ >= baseParams_->mm4BlockNum * cvRatio_) { + return; + } + // 等待前面的Qr部分完成 + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + + constexpr uint32_t DYNAMIC_QUANT_INPUT_READY = EVENT_ID0; + constexpr uint32_t MUL_QR_INPUT_COPY_READY = EVENT_ID3; + constexpr uint32_t DYNAMIC_QUANT_OUTPUT_READY = EVENT_ID3; + + int64_t blockBatchOffset = (blockIdx_ % cvRatio_) * (curStepBatchSize / cvRatio_); + int64_t totalSizeCkv = + static_cast(baseParams_->numHeadSize) * static_cast(baseParams_->headSizeCkv); + // 由于两个vec核分curStepBatchSize,各处理curStepBatchSize/2,blockIdx_ & 1 表示是否为第二个vec核 + int64_t dynamicQuantQueryOffset = + blockBatchOffset * totalSizeCkv + numHeadOffset * static_cast(baseParams_->headSizeCkv); + int64_t dynamicQuantQueryResOffset = batchOffset * totalSizeCkv + blockBatchOffset * totalSizeCkv + + numHeadOffset * static_cast(baseParams_->headSizeCkv); + int64_t scaleQueryNopeOffset = batchOffset * static_cast(baseParams_->numHeadSize) + + blockBatchOffset * static_cast(baseParams_->numHeadSize) + numHeadOffset; + int64_t queryOutStride = totalSizeCkv; + int64_t qrOutputStride = + static_cast(baseParams_->numHeadSize) * static_cast(baseParams_->dimHeadRope); + int64_t qrPostProcessResOffset = batchOffset * static_cast(baseParams_->headSizeQr) + + numHeadOffset * static_cast(baseParams_->dimHeadRope) + + blockBatchOffset * static_cast(baseParams_->headSizeQr); + + + LocalTensor shareTmpUb = shareBuffer_.Get(); + + float quantScaleCkv = quantScaleCkvGm_.GetValue(0); + + // Dynamic Quant + SetFlag(DYNAMIC_QUANT_OUTPUT_READY); + SetFlag(DYNAMIC_QUANT_INPUT_READY); + + // Rope Post Process + SetFlag(MUL_QR_INPUT_COPY_READY); + // per-head循环 + for (int64_t loopIdx = 0; loopIdx < mmQnLoops; loopIdx++) { + CrossCoreWaitFlag(FINISH_MM_QN_SPLIT_N); + DynamicQuantQnWithMulQr( + dequantScaleQNopeGm_[scaleQueryNopeOffset], queryOutGm_[dynamicQuantQueryResOffset], + qrOutGm_[qrPostProcessResOffset], mmQnResGm_[dynamicQuantQueryOffset], shareTmpUb, curStepBatchSizeVec, + baseParams_->headSizeCkv, baseParams_->numHeadSize, queryOutStride, + // Rope Post Process + qrOutGm_[qrPostProcessResOffset], quantScaleCkv, baseParams_->dimHeadRope, qrOutputStride, cvRatio_); + + dynamicQuantQueryOffset += static_cast(baseParams_->headSizeCkv); + scaleQueryNopeOffset += 1; + dynamicQuantQueryResOffset += static_cast(baseParams_->headSizeCkv); + qrPostProcessResOffset += static_cast(baseParams_->dimHeadRope); + } + // Rope Post Process + WaitFlag(MUL_QR_INPUT_COPY_READY); + // Dynamic Quant + WaitFlag(DYNAMIC_QUANT_INPUT_READY); + WaitFlag(DYNAMIC_QUANT_OUTPUT_READY); +} + +} // namespace MlaProlog + +#endif // MLA_PROLOG_VEC_S1_CUB_S2_H diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/mla_prolog_comm.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/mla_prolog_comm.h new file mode 100644 index 000000000000..b7ca0ccf90af --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/mla_prolog_comm.h @@ -0,0 +1,409 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_comm.h + * \brief + */ + +#ifndef MLA_PROLOG_COMM_H +#define MLA_PROLOG_COMM_H + +#include "kernel_operator.h" +#include "kernel_operator_list_tensor_intf.h" +#include "kernel_tiling/kernel_tiling.h" +#include "lib/matmul_intf.h" +#include "lib/matrix/matmul/tiling.h" + +using namespace AscendC; + +namespace MlaProlog { + +template +__aicore__ inline T CeilDivT(T num1, T num2) +{ + if (num2 == 0) { + return static_cast(0); + } + return (num1 + num2 - 1) / num2; +} + +template +__aicore__ inline T Align(T num, T rnd) +{ + return (((rnd) == 0) ? 0 : (((num) + (rnd)-1) / (rnd) * (rnd))); +} + +enum class CACHE_MODE : uint8_t { + ND = static_cast(0), // BSND or TND + PA_BSND = static_cast(1), + PA_NZ = static_cast(2), + PA_BLK_BSND = static_cast(3), + PA_BLK_NZ = static_cast(4), + PA_BS = static_cast(5) +}; + +enum class SCENARIO : uint8_t { + RESERVED = static_cast(0), + NO_QUANT = static_cast(1), + QUANT = static_cast(2) +}; + +enum class QUANT_MODE : uint8_t { + NO_QUANT = static_cast(0), + PARTIAL_QUANT_KV_NO_QUANT = static_cast(1), + PARTIAL_QUANT_KV_QUANT_PER_CHANNEL = static_cast(2), + FULL_QUANT_KV_NO_QUANT = static_cast(3), + FULL_QUANT_KV_QUANT_PER_TENSOR = static_cast(4), + PARTIAL_QUANT_KV_QUANT_PERTILE = static_cast(5), + FULL_QUANT_KV_QUANT_PERTILE = static_cast(6), + MXFP8_FULL_QUANT_KV_NO_QUANT = static_cast(7), + MXFP8_FULL_QUANT_KV_QUANT_PER_TENSOR = static_cast(8), + MXFP8_FULL_QUANT_KV_QUANT_PER_TILE = static_cast(9), + FP8_FULL_QUANT_KV_NO_QUANT = static_cast(10), + FP8_FULL_QUANT_KV_QUANT_PER_TENSOR = static_cast(11), + HIF8_FULL_QUANT_KV_NO_QUANT = static_cast(12), + HIF8_FULL_QUANT_KV_QUANT_PER_TENSOR = static_cast(13), + FP8_FULL_QUANT_KV_QUANT_PER_TILE = static_cast(14), + HIF8_FULL_QUANT_KV_QUANT_PER_TILE = static_cast(15) +}; + +enum class EMPTY_TENSOR_MODE : uint8_t { + NON_EMPTY = static_cast(0), + EMPTY_CACHE = static_cast(1), + EMPTY_QUERY = static_cast(2), +}; + +enum class ACTUAL_SEQ_MODE : uint8_t { + DISABLED = static_cast(0), + EN_Q_LEN = static_cast(1), +}; + +enum class SPLIT_M_MODE : uint8_t { + DISABLED = static_cast(0), + ENABLED = static_cast(1), +}; + +enum class ROPE_MODE : uint8_t { + INTERLEAVE_HALF = static_cast(0), + HALF = static_cast(1), +}; + +constexpr uint64_t BYTE_BLOCK = 32UL; +constexpr uint8_t ALIGN_BLOCK_SIZE = 32; // 32B对齐 +constexpr uint32_t BLOCK_CUBE_SIZE = 16; // L1上m轴16对齐 +constexpr uint32_t FP8_TWO = 2; // L1上scale的存储方式为2个fp8类型存成1个bf16 +constexpr uint32_t REPEAT_BLOCK_BYTE = 256; +constexpr uint32_t REPEAT_STRIDE_UP_BOUND = 256; // repeat stride 不能超过256 +constexpr uint32_t FP32_BLOCK_ELEMENT_NUM = ALIGN_BLOCK_SIZE / sizeof(float); +constexpr uint32_t FP16_BLOCK_ELEMENT_NUM = ALIGN_BLOCK_SIZE / sizeof(half); +constexpr uint32_t FP32_REPEAT_ELEMENT_NUM = REPEAT_BLOCK_BYTE / sizeof(float); +constexpr uint32_t MAX_UB_SIZE = 192 * 1024; // 最大的UB大小 +constexpr uint32_t DIM_HEAD_SIZE_QCQR = 192; // 算力分组方案D + Dr = 192 +constexpr uint32_t QC_CORE_NUM = 8; // 算力分组方案QC占用8核 +constexpr uint32_t QR_CORE_NUM = 4; // 算力分组方案QR占用4核 +constexpr uint32_t INT8_AFULLLOAD_MAX_MSIZE = 64; // 计算mmQcQr时,int8类型的A矩阵在msize小于等于64可以全载L1 +constexpr uint32_t BF16_AFULLLOAD_MAX_MSIZE = 32; // 计算mmQcQr时,bf16类型的A矩阵在msize小于等于32可以全载L1 +constexpr uint32_t ONE_BYTE_TYPE_SIZE = 1; // 数据类型int8_t fp8大小为1字节 +constexpr uint32_t FP8_E4M3_BLOCK_SIZE = 32; +constexpr uint32_t K_STEP_SIZE_32 = 32; // for move left or right +constexpr uint32_t SHIFTS_UNIT = 4; +constexpr uint32_t UNIT_SIZE = 512; +constexpr uint32_t ROUND_UP_UNIT = 15; // for round up +constexpr uint32_t MAX_SYNC_FLAG_COUNT = 15; // 同一个flagId的计数器最多设置15次 + +constexpr int SYNC_MODE_ALL_CUBE = 0x0; +constexpr int SYNC_MODE_CUBE_VEC = 0x2; +constexpr int SYNC_MODE_ALL_VEC = 0x0; + +constexpr int FINISH_MM_CQ = 0x6; +constexpr int FINISH_MM_CKVKR = 0x6; +constexpr int FINISH_MM_QCQR = 0x6; +constexpr int FINISH_MM_QR = 0x8; // 算力分组场景 +constexpr int FINISH_MM_QC = 0x6; // 算力分组场景 +constexpr int FINISH_MM_ALL = 0x7; + +constexpr int FINISH_VEC_RMSNORM_CQ = 0x6; +constexpr int FINISH_VEC_DEQUANT_QC = 0x6; +constexpr int FINISH_VEC_CKVKR = 0x9; +constexpr int FINISH_VEC_ALL = 0x7; + +constexpr int FINISH_MM_QCQR_SPLIT_N = 0XA; +constexpr int FINISH_MM_QCQR_SPLIT_BATCH = 0x3; +constexpr int FINISH_VEC_DEQUANT_QC_SPLIT_N = 0X7; +constexpr int FINISH_VEC_DEQUANT_QC_SPLIT_N_GAP = 0X1; +constexpr int FINISH_MM_QN_SPLIT_N = 0X8; + +#ifdef ENABLE_DUMP_DATA +#define DO_DUMP_DATA(srcTensor, id, len) DumpTensor(srcTensor, id, len) +#else +#define DO_DUMP_DATA(srcTensor, id, len) +#endif + +class NoneType {}; + +using FP8E4M3 = fp8_e4m3fn_t; + +using FP8E8M0 = fp8_e8m0_t; + +using HIF8 = hifloat8_t; + +// mte2 <> mte1 +#define SCALE_EVENT EVENT_ID3 +#define A_EVENT0 EVENT_ID4 +#define A_EVENT1 EVENT_ID5 +#define B_EVENT0 EVENT_ID6 +#define B_EVENT1 EVENT_ID7 + +// m <> mte1 +#define L0A_EVENT0 EVENT_ID3 +#define L0A_EVENT1 EVENT_ID4 +#define L0B_EVENT0 EVENT_ID5 +#define L0B_EVENT1 EVENT_ID6 + +// fix <> m +#define L0C_EVENT0 EVENT_ID3 +#define L0C_EVENT1 EVENT_ID4 + +constexpr uint32_t L1_A_SIZE = 128 * 1024; // 512 / 4 +constexpr uint32_t L1_B_SIZE = 128 * 1024; // 512 / 4 +constexpr uint32_t L0A_PP_SIZE = 32 * 1024; +constexpr uint32_t L0B_PP_SIZE = 32 * 1024; +constexpr uint32_t L0C_PP_SIZE = 64 * 1024; + + +/* + 非量化 半量化(kv非量化) 半量化(kv量化) int8全量化(kv非量化) int8全量化(kv量化) 半量化(kv per-tile量化) int8全量化(kv per-tile量化) Mxfp8量化(kv非量化) Mxfp8量化(kv量化) Mxfp8量化(kv per-tile量化) fp8全量化(kv非量化) fp8全量化(kv量化) hif8全量化(kv非量化) hif8全量化(kv量化) + cacheMode PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/BSND/TND PA_BSND/BSND/TND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/BSND/TND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND + /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ + /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND + enableDequantOpt false true/false true/false true/false true/false true true true true true true/false true/false true/false true/false + enableGroupDequantOpt false true/false true/false false false false false false false false false false false false + quantMode 0 1 2 3 4 5 6 7 8 9 10 11 12 13 + tokenXType(复用mmInputType) bfloat16_t bfloat16_t bfloat16_t int8_t int8_t bfloat16_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + WdqType(复用mmInputType) bfloat16_t bfloat16_t bfloat16_t int8_t int8_t bfloat16_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + WuqqrType(复用mmQcQrInputType) bfloat16_t int8_t int8_t int8_t int8_t int8_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + WukType(复用mmQnInputType) bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + WdkvkrType(复用mmInputType) bfloat16_t bfloat16_t bfloat16_t int8_t int8_t bfloat16_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + rmsNormGammaType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + gammaCkvType(复用rmsNormGammaType)bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + ropeSinCosType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + cosType(复用ropeSinCosType) bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + cacheIndexType int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t + kvCacheType bfloat16_t bfloat16_t int8_t bfloat16_t int8_t int8_t int8_t bfloat16_t fp8_e4m3fn_t fp8_e4m3fn_t bfloat16_t fp8_e4m3fn_t bfloat16_t hifloat8_t + krCacheType bfloat16_t bfloat16_t int8_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + deqScaleXType / / / float float / float fp8_e8m0_t fp8_e8m0_t fp8_e8m0_t float float float float + deqScaleWdqType / / / float float / float fp8_e8m0_t fp8_e8m0_t fp8_e8m0_t float float float float + deqScaleWuqqrType / float float float float float float fp8_e8m0_t fp8_e8m0_t fp8_e8m0_t float float float float + deqScaleWdkvkrType / / / float float / float fp8_e8m0_t fp8_e8m0_t fp8_e8m0_t float float float float + quantScaleCkvType / / float / float / / / float / / float / float + quantScaleCkrType / / float / / / / / / / / / / / + smoothScaleCqType / float float float float float float / / / float float float float + queryOutputType bfloat16_t bfloat16_t int8_t bfloat16_t int8_t bfloat16_t bfloat16_t bfloat16_t fp8_e4m3fn_t bfloat16_t bfloat16_t fp8_e4m3fn_t bfloat16_t hifloat8_t + ropeOutputType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + dequantScaleQNopeType / / / / float / / / float / / float / float + queryNormType(复用mmQcQrInputType)bfloat16_t int8_t int8_t int8_t int8_t int8_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + dequantScaleQNormType / float float float float float float fp8_e8m0_t fp8_e8m0_t fp8_e8m0_t float float float float + mmInputType bfloat16_t bfloat16_t bfloat16_t int8_t int8_t bfloat16_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + mmCqOutputType bfloat16_t bfloat16_t bfloat16_t int32_t int32_t bfloat16_t int32_t float float float float float float float + mmCkvKrInputType(复用mmInputType) bfloat16_t bfloat16_t bfloat16_t int8_t int8_t bfloat16_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + mmCkvKrOutputType bfloat16_t bfloat16_t bfloat16_t int32_t int32_t bfloat16_t int32_t float float float float float float float + mmQcQrInputType bfloat16_t int8_t int8_t int8_t int8_t int8_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + mmQcQrOutputType bfloat16_t bfloat16_t bfloat16_t int32_t int32_t bfloat16_t int32_t float float float float float float float + mmQnInputType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + mmQnOutputType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + rmsNormComputType float float float float float float float float float float float float float float + rmsNormCqOutputType bfloat16_t int8_t int8_t int8_t int8_t int8_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + rmsNormCkvOutputType bfloat16_t bfloat16_t int8_t bfloat16_t int8_t int8_t int8_t bfloat16_t fp8_e4m3fn_t fp8_e4m3fn_t bfloat16_t fp8_e4m3fn_t bfloat16_t hifloat8_t + ropeComputType float float float float float float float float float float float float float float +*/ + +template +struct MLAPType { + // 如果是 FP8 或 HIF8,输出 float;如果是 int8,输出 int32_t;否则输出 bfloat16_t + template + using GetMatmulOutType_t = typename std::conditional< + std::is_same::value || std::is_same::value, float, + typename std::conditional::value, int32_t, bfloat16_t>::type>::type; + + using mmInputType = X_T; // tokenX的类型与weight的类型一致 + using mmQcQrInputType = W_T; + using mmQnInputType = bfloat16_t; // matmul计算Qn的输入类型 + using mmCqOutputType = GetMatmulOutType_t; // matmul计算Cq的输出类型 + using mmCkvKrOutputType = GetMatmulOutType_t; // matmul计算QcQr的输出类型 + using mmQcQrOutputType = GetMatmulOutType_t; // matmul计算Qn的输出类型 + using mmQnOutputType = bfloat16_t; // matmul计算Qn的输出类型 + using rmsNormGammaType = bfloat16_t; // gamma的输入类型 + using rmsNormComputType = float; + using rmsNormCqOutputType = W_T; + using rmsNormCkvOutputType = C_T; + using ropeSinCosType = bfloat16_t; // sin cos的输入类型 + using ropeComputType = float; + using ropeOutputType = bfloat16_t; + using kvCacheType = C_T; // kvcache的类型 + // 如果是 X_T 为 bfloat16_t && C_T 为 int8_t 且不是 pertile (半量化kv perchannel量化), 则为 int8_t;否则为 + // bfloat16_t + using krCacheType = typename std::conditional::value && + std::is_same::value && !IS_PERTILE, + int8_t, bfloat16_t>::type; + using dequantScaleQNopeType = float; // dequantScaleQNope的类型 + using dequantScaleQNormType = D_S; // dequantScaleQNorm的类型 + using dequantScaleType = D_S; + + static constexpr CACHE_MODE cacheMode = C_M; + static constexpr bool enableDequantOpt = ENABLE_DEQUANT_OPT; + static constexpr bool enableGroupComputeOpt = ENABLE_GROUP_COMPUTE_OPT; + static constexpr bool enableRope = ENABLE_ROPE; + static constexpr EMPTY_TENSOR_MODE emptyMode = EMPTY_MODE; + static constexpr ACTUAL_SEQ_MODE actualSeqMode = SEQ_MODE; + static constexpr bool isPertile = IS_PERTILE; + static constexpr uint32_t cvRatio = CV_RATIO; // 默认C:V 1:2 +}; + +struct MMParams { + uint32_t m; + uint32_t n; + uint32_t k; + uint32_t orgM; + uint32_t orgN; + uint32_t orgKa; + uint32_t orgKb; + uint32_t orgKc; + uint32_t baseM; + uint32_t baseN; + uint32_t baseK; + uint32_t stepK; + uint32_t needSetOrgShape; + uint32_t kL1StepSize; + uint32_t kScale; +}; + +struct MMBufParams { + uint32_t aL1BufIter = 0; + uint32_t bL1BufIter = 0; + TBuffAddr aL1BufAddr; + TBuffAddr bL1BufAddr; + uint32_t aL0BufIter = 0; + uint32_t bL0BufIter = 0; + uint32_t cL0BufIter = 0; + TBuffAddr aL0BufAddr; + TBuffAddr bL0BufAddr; + TBuffAddr cL0BufAddr; +}; + +struct AicOffset { + int64_t weightDqOffset = 0; + int64_t weightUqQrOffset = 0; + int64_t weightUqOffset = 0; + int64_t weightQrOffset = 0; + int64_t weightUkOffset = 0; + int64_t weightDkvKrOffset = 0; + int64_t dequantScaleWDqOffset = 0; + int64_t dequantScaleWDkvKrOffset = 0; + int64_t dequantScaleCqOffset = 0; + int64_t dequantScaleWuqqrOffset = 0; + int64_t cqResOffset = 0; + int64_t rmsNormCqResOffset = 0; + int64_t ckvKrResOffset = 0; + int64_t qcQrResOffset = 0; + int64_t qCResOffset = 0; + int64_t qRResOffset = 0; + int64_t qcOffset = 0; + int64_t qnResOffset = 0; +}; + +struct AivOffset { + int64_t curVecToken = 0; + int64_t curBlockTokenOffset = 0; + int64_t rmsNormCqOffset = 0; + int64_t rmsNormCqResOffset = 0; + int64_t rmsNormCkvOffset = 0; + int64_t mmQnPreDequantOffset = 0; + int64_t mmQnPreDequantResOffset = 0; + int64_t ropeKrOffset = 0; + int64_t ropeQrOffset = 0; + int64_t ropeQrResOffset = 0; + int64_t ropeQrSplitNOffset = 0; + int64_t ropeQrResSplitNOffset = 0; + int64_t qcScaleOffsetSplitN = 0; +}; + +struct UsedBlockParams { + uint32_t blockStartIdx; + uint32_t blockEndIdx; +}; + +struct CkvkrParams { + int64_t tokenIndex; + int64_t offset; + int64_t curVecTokenIdx; + int64_t rowsInCurBatch; + int64_t cacheOffset; + int64_t nextBatchOffset; +}; + + +struct RopeQrSplitNParams { + int64_t ropeQrOffset; + int64_t ropeQrResOffset; + uint32_t inputOffsetRope; + uint32_t deqScaleOffset; + uint32_t outputOffsetRope; + int64_t ropeStride; + uint32_t ropeDstStride; + uint32_t deQuantScaleCqOffset; + uint32_t sinCosOffset; + uint32_t ropeCnt; +}; + +struct DequantQcQrSplitNParams { + int64_t mmQnPreDequantOffset; + int64_t mmQnPreDequantResOffset; + uint32_t inputOffset; + uint32_t outputOffset; + uint32_t srcStride; + uint32_t dstStride; +}; + +struct CastQcQrSplitNParams { + int64_t mmQnPreCastOffset; + int64_t mmQnPreCastResOffset; + uint32_t inputOffset; + uint32_t outputOffset; + uint32_t srcStride; + uint32_t dstStride; +}; + +template +__aicore__ inline void WaitAllCore(uint16_t flagId) +{ + CrossCoreSetFlag(flagId); + CrossCoreWaitFlag(flagId); +} + +template +__aicore__ constexpr bool IsFullQuantMode() +{ + constexpr bool IS_FP8_E8M0 = std::is_same::value; + if constexpr (EXCLUDE_MXFP8) { + // Pattern A: INT8 ‖ HIF8 ‖ (FP8E4M3 && !IS_FP8_E8M0) + return std::is_same::value || std::is_same::value || + (std::is_same::value && !IS_FP8_E8M0); + } else { + // Pattern B: INT8 ‖ FP8E4M3 ‖ HIF8 + return std::is_same::value || std::is_same::value || + std::is_same::value; + } +} + +} +#endif diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/mla_prolog_vector_comm.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/mla_prolog_vector_comm.h new file mode 100644 index 000000000000..8f39fafc0664 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/mla_prolog_vector_comm.h @@ -0,0 +1,459 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_comm.h + * \brief 存放各种vector的公共组件 + */ + +#ifndef MLA_PROLOG_VECTOR_COMM_H +#define MLA_PROLOG_VECTOR_COMM_H + +#include "vf/vf_quant_pertensor.h" +#include "vf/vf_quant_perchannel.h" +#include "vf/vf_dynamic_quant.h" +#include "vf/vf_dequant.h" + +#include "mla_prolog_comm.h" +namespace MlaProlog { + +struct Rectangle { + uint32_t row; + uint32_t col; + uint32_t stride; +}; + +struct RmsNormParam { + float reciprocal; + float epsilon; + uint32_t row; + uint32_t col; + float scale; + uint16_t isScaleEnable; +}; + +struct PerTileQuantParams { + uint32_t tileSize; + uint32_t tileNum; + float alpha; + uint32_t row; + uint32_t col; +}; + +/** + * @brief RowMuls muls by row, 每行的元素乘以相同的元素,该元素需要扩展到一个数据块; + * dstUb[i, (j * 8) : (j * 8 + 7)] = src0Ub[i, (j * 8) : (j * 8 + 7)] * src1Ub[i, 0 : 7] + * @param dstUb 输出tensor [row, columnStride] + * @param src0Ub 输入tensor [row, columnStride] + * @param src1Ub 输入tensor [row, FP32_BLOCK_ELEMENT_NUM] + * @param rectangleParams 描述待处理数据的排布,包括 + row 行数 + col 列数 + stride 一行的真实长度 + */ +template +__aicore__ inline void RowMuls(LocalTensor dstUb, LocalTensor src0Ub, LocalTensor src1Ub, + const Rectangle &rectangleParams) +{ + uint32_t repeatElementNum = FP32_REPEAT_ELEMENT_NUM; + uint32_t blockElementNum = FP32_BLOCK_ELEMENT_NUM; + + if constexpr (std::is_same::value) { + // 此限制由于每个repeat至多连续读取256B数据 + repeatElementNum = FP32_REPEAT_ELEMENT_NUM * 2; // 256/4 * 2 = 128 + blockElementNum = FP32_BLOCK_ELEMENT_NUM * 2; // 32/4 * 2 = 16 + } + + // 每次只能连续读取256B的数据进行计算,故每次只能处理256B/sizeof(dType)= + // 列方向分dLoop次,每次处理8列数据 + uint32_t dLoop = rectangleParams.col / repeatElementNum; + uint32_t dRemain = rectangleParams.col % repeatElementNum; + // REPEAT_STRIDE_UP_BOUND=256,此限制由于src0RepStride数据类型为uint8至多256个datablock间距 + if (rectangleParams.stride < REPEAT_STRIDE_UP_BOUND * blockElementNum) { + // dstBlkStrideIn src0BlkStrideIn src1BlkStrideIn dstRepStrideIn src0RepStrideIn src1RepStrideIn + BinaryRepeatParams repeatParams{1, + 1, + 0, + static_cast(rectangleParams.stride / blockElementNum), + static_cast(rectangleParams.stride / blockElementNum), + 1}; + + // 如果以列为repeat所处理的次数小于行处理次数,则以列方式处理。反之则以行进行repeat处理 + if (dLoop <= rectangleParams.row) { + uint32_t offset = 0; + for (uint32_t i = 0; i < dLoop; i++) { + Mul(dstUb[offset], src0Ub[offset], src1Ub, repeatElementNum, rectangleParams.row, repeatParams); + offset += repeatElementNum; + } + } else { + repeatParams.src0RepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block + repeatParams.src1RepStride = 0; + repeatParams.dstRepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block + for (uint32_t i = 0; i < rectangleParams.row; i++) { + Mul(dstUb[i * rectangleParams.stride], src0Ub[i * rectangleParams.stride], src1Ub[i * blockElementNum], + repeatElementNum, dLoop, repeatParams); + } + } + + // 最后一次完成[row, dRemain] * [row, blockElementNum] 只计算有效部分 + if (dRemain > 0) { + Mul(dstUb[dLoop * repeatElementNum], src0Ub[dLoop * repeatElementNum], src1Ub, dRemain, rectangleParams.row, + repeatParams); + } + } else { + // dstBlkStrideIn src0BlkStrideIn src1BlkStrideIn dstRepStrideIn src0RepStrideIn src1RepStrideIn + // 8 : 每个repeat为256B数据,正好8个datablock + BinaryRepeatParams repeatParams{1, 1, 0, 8, 8, 0}; + + // 每次计算一行,共计算dealRowCount行 + for (uint32_t i = 0; i < rectangleParams.row; i++) { + // 计算一行中的dLoop个repeat,每个repeat计算256/block_size个data_block + Mul(dstUb[i * rectangleParams.stride], src0Ub[i * rectangleParams.stride], src1Ub[i * blockElementNum], + repeatElementNum, dLoop, repeatParams); + // 计算一行中的尾块 + if (dRemain > 0) { + Mul(dstUb[i * rectangleParams.stride + dLoop * repeatElementNum], + src0Ub[i * rectangleParams.stride + dLoop * repeatElementNum], src1Ub[i * blockElementNum], dRemain, + 1, repeatParams); + } + } + } +} + +/** + * @brief RowMax max by row, 按行求最大值 + * @param dstUb 输出tensor [row, 1] + dstUb[i] = max(srcUb[i, :]) + * @param srcUb 输入tensor [row, columnStride] + * @param rectangleParams 描述待处理数据的排布,包括 + row 行数 + col 列数 + stride 一行的真实长度 + */ +__aicore__ inline void RowMax(LocalTensor &dstUb, LocalTensor &srcUb, const Rectangle rectangleParams) +{ + uint32_t dtypeMask = FP32_REPEAT_ELEMENT_NUM; + uint32_t blockCount = rectangleParams.col / dtypeMask; + uint32_t remain = rectangleParams.col % dtypeMask; + + BinaryRepeatParams repeatParamsMax; + repeatParamsMax.src0BlkStride = 1; + repeatParamsMax.src1BlkStride = 1; + repeatParamsMax.dstBlkStride = 1; + repeatParamsMax.src0RepStride = rectangleParams.stride / FP32_BLOCK_ELEMENT_NUM; + repeatParamsMax.src1RepStride = rectangleParams.stride / FP32_BLOCK_ELEMENT_NUM; + repeatParamsMax.dstRepStride = rectangleParams.stride / FP32_BLOCK_ELEMENT_NUM; + if (blockCount > 0 && remain > 0) { + Max(srcUb, srcUb, srcUb[blockCount * dtypeMask], remain, rectangleParams.row, repeatParamsMax); + PipeBarrier(); + } + + for (uint32_t columnLoopCount = blockCount >> 1; columnLoopCount > 0; + columnLoopCount = blockCount >> 1) { // 2: 每次处理2个block + blockCount = (blockCount + 1) >> 1; // 2: 每次处理2个block + for (uint32_t j = 0; j < columnLoopCount; j++) { + Max(srcUb[j * dtypeMask], srcUb[j * dtypeMask], srcUb[(j + blockCount) * dtypeMask], dtypeMask, + rectangleParams.row, repeatParamsMax); + } + PipeBarrier(); + } + + WholeReduceMax(dstUb, srcUb, (rectangleParams.col < dtypeMask) ? rectangleParams.col : dtypeMask, + rectangleParams.row, 1, 1, rectangleParams.stride / FP32_BLOCK_ELEMENT_NUM, + ReduceOrder::ORDER_ONLY_VALUE); +} + +/** + * @brief Dequant 对[row * col]的tensor进行反量化。需要输入一个行向量和列向量,共同构成反量化的矩阵。 + outputLocal [i,j] = inputLocal[i,j] * scaleLocal[j] * scale2Local [i] + * @param outputLocal 输出tensor [row , col] + * @param inputLocal 输入tensor [row , col] + * @param scaleLocal [1,col] 量化系数;行向量 + * @param scale2Local [row,8] 量化系数;列向量,8为float扩充为32Bytes + * @param rectangleParams 描述待处理数据的排布,包括 + row 行数 + col 列数 + stride 一行的真实长度 + */ +template +__aicore__ inline void Dequant(const LocalTensor &outputLocal, const LocalTensor &inputLocal, + const LocalTensor &scaleLocal, const LocalTensor &scale2Local, + const Rectangle &rectangleParams) +{ + DequantVf(outputLocal, inputLocal, scaleLocal, scale2Local, rectangleParams.row, rectangleParams.col, + rectangleParams.stride); +} + +/** + * @brief CastFP32ToINT8 将float类型cast为int8,路径为float--------->int32--------->half---------->int8 + CAST_RINT CAST_ROUND CAST_TRUNC + * @param outputLocal 输出tensor [cnt] + * @param inputLocal 输入tensor [cnt] + * @param shareTmpUb 临时buffer 内部需要的空间为 [cnt * 4] 4 : sizeof(int32) + * @param cnt tensor长度 + */ +template +__aicore__ inline void CastFP32ToINT8(const LocalTensor outLocal, const LocalTensor &inputLocal, + const LocalTensor &shareTmpUb, uint64_t cnt) +{ + LocalTensor int32 = shareTmpUb.ReinterpretCast(); + LocalTensor tmpHalf = shareTmpUb.ReinterpretCast(); + Cast(int32, inputLocal, RoundMode::CAST_RINT, cnt); + PipeBarrier(); + SetDeqScale(static_cast(1.0)); + PipeBarrier(); + Cast(tmpHalf, int32, RoundMode::CAST_ROUND, cnt); + PipeBarrier(); + Cast(outLocal, tmpHalf, RoundMode::CAST_TRUNC, cnt); +} + +// divs by row, 每行的元素除以相同的元素 +// dstUb[i, (j * 8) : (j * 8 + 7)] = src0Ub[i, (j * 8) : (j * 8 + 7)] / src1Ub[i, 0 : 7] +// src0Ub:[dealRowCount, columnCount], src1Ub:[dealRowCount, FP32_BLOCK_ELEMENT_NUM] dstUb:[dealRowCount, +// columnCount] +__aicore__ inline void RowClips(LocalTensor &dstUb, const LocalTensor &src0Ub, + const LocalTensor &maxUb, const LocalTensor &minUb, uint32_t dealRowCount, + uint32_t columnCount, uint32_t actualColumnCount) +{ + uint32_t dtypeMask = FP32_REPEAT_ELEMENT_NUM; + uint32_t dLoop = actualColumnCount / dtypeMask; + uint32_t dRemain = actualColumnCount % dtypeMask; + + BinaryRepeatParams repeatParams; + repeatParams.src0BlkStride = 1; + repeatParams.src1BlkStride = 0; + repeatParams.dstBlkStride = 1; + repeatParams.src0RepStride = columnCount / FP32_BLOCK_ELEMENT_NUM; + repeatParams.src1RepStride = 1; + repeatParams.dstRepStride = columnCount / FP32_BLOCK_ELEMENT_NUM; + uint32_t offset = 0; + + for (uint32_t i = 0; i < dLoop; i++) { + Min(dstUb[offset], src0Ub[offset], maxUb, dtypeMask, dealRowCount, repeatParams); + PipeBarrier(); + Max(dstUb[offset], dstUb[offset], minUb, dtypeMask, dealRowCount, repeatParams); + offset += dtypeMask; + } + + if (dRemain > 0) { + Min(dstUb[dLoop * dtypeMask], src0Ub[dLoop * dtypeMask], maxUb, dRemain, dealRowCount, repeatParams); + PipeBarrier(); + Max(dstUb[dLoop * dtypeMask], dstUb[dLoop * dtypeMask], minUb, dRemain, dealRowCount, repeatParams); + } +} + +__aicore__ inline void PerTileClipWithAlpha(LocalTensor &dstClipUb, LocalTensor &aMax, + LocalTensor &aMaxBrcb, const LocalTensor &srcUb, + const LocalTensor &shareTmpUb, + const PerTileQuantParams &perTileQuantParams) +{ + constexpr uint32_t brcnNum = 8; // brcn一次性处理8个数据 + uint32_t srcSize = perTileQuantParams.row * perTileQuantParams.col; + uint32_t maxAlphaSize = perTileQuantParams.row * perTileQuantParams.tileNum * FP32_BLOCK_ELEMENT_NUM; + + LocalTensor absSrcUb = shareTmpUb.template ReinterpretCast(); + LocalTensor maxAlpha = absSrcUb[Align(srcSize, FP32_BLOCK_ELEMENT_NUM)]; + LocalTensor minAlpha = maxAlpha[Align(maxAlphaSize, FP32_BLOCK_ELEMENT_NUM)]; + + Abs(absSrcUb, srcUb, srcSize); + PipeBarrier(); + + Rectangle rectangleParams{ + static_cast(perTileQuantParams.row * perTileQuantParams.tileNum), + static_cast(perTileQuantParams.tileSize), + static_cast(perTileQuantParams.tileSize) // columnStride + }; + RowMax(aMax, absSrcUb, rectangleParams); + PipeBarrier(); + Brcb(aMaxBrcb, aMax, (CeilDivT((perTileQuantParams.row * perTileQuantParams.tileNum), brcnNum)), {1, brcnNum}); + PipeBarrier(); + + Muls(maxAlpha, aMaxBrcb, perTileQuantParams.alpha, maxAlphaSize); + Muls(minAlpha, aMaxBrcb, -perTileQuantParams.alpha, maxAlphaSize); + PipeBarrier(); + + RowClips(dstClipUb, srcUb, maxAlpha, minAlpha, perTileQuantParams.row * perTileQuantParams.tileNum, + perTileQuantParams.tileSize, perTileQuantParams.tileSize); + PipeBarrier(); +} + +/** + * @brief DynamicQuant 对row行进行dynamicquant, float ---> int8, 每一行出一个系数。 + * @param outputLocal 输出tensor [row , col],支持和inputLocal是同一块空间 + * @param scale 输出每行的反量化系数 [row] + * @param inputLocal 输入tensor [row , col] + * @param rowMax [row] 一行的最大值 + * @param rowMaxBrcb [row, 8] 一行的最大值,brcb扩充为8的倍数 + * @param shareTmpUb 临时buffer 内部需要的空间为 [row * col * sizeof(float)] + * @param row 待处理的行数 + * @param col 待处理的列数 + */ +__aicore__ inline void DynamicQuant(const LocalTensor &outputLocal, const LocalTensor &scale, + const LocalTensor &inputLocal, const LocalTensor &aMax, + const LocalTensor &rowMaxBrcb, const LocalTensor &shareTmpUb, + float alpha, uint64_t row, uint64_t col) +{ + constexpr float maxInt8 = 127.0f; + int32_t aMaxSizeAlign = Align(static_cast(row), FP32_BLOCK_ELEMENT_NUM); + LocalTensor aMaxAlpha = shareTmpUb.ReinterpretCast(); + LocalTensor dupTensor = aMaxAlpha[aMaxSizeAlign]; + LocalTensor scaleReciprocal = dupTensor[aMaxSizeAlign]; + LocalTensor scaleReciprocalBrcb = scaleReciprocal[aMaxSizeAlign]; + + Duplicate(dupTensor, 1.0f, row); + Muls(aMaxAlpha, aMax, alpha, row); + PipeBarrier(); + Muls(scale, aMaxAlpha, 1.0f / maxInt8, row); + PipeBarrier(); + Div(scaleReciprocal, dupTensor, scale, row); + PipeBarrier(); + Brcb(scaleReciprocalBrcb, scaleReciprocal, CeilDivT(row, 8UL), {1, 8}); + PipeBarrier(); + RowMuls(outputLocal, inputLocal, scaleReciprocalBrcb, + Rectangle{static_cast(row), static_cast(col), static_cast(col)}); + PipeBarrier(); +} + +/** + * @brief QuantPerChannel 同时对row行进行FP32到int8的per-channel量化操作。一行中的每一列用不同的量化参数。 + outLocal[i, j] = inputLocal[i, j] * quantScaleLocal[j] + * @param outLocal 输出tensor [row , col] + * @param inputLocal 输入tensor [row , col] + * @param quantScaleLocal quant系数 [1 , col] + * @param shareTmpUb 临时buffer 内部需要的空间为 [row * col * 4],源自CastFP32ToINT8 + * @param rectangleParams 描述待处理数据的排布,包括 + row 行数 + col 列数 + stride 一行的真实长度 + */ +template +__aicore__ inline void QuantPerChannel(const LocalTensor &outLocal, const LocalTensor &inputLocal, + const LocalTensor &quantScaleLocal, const LocalTensor &shareTmpUb, + const Rectangle &rectangleParams) +{ + QuantPerChannelVf(outLocal, inputLocal, quantScaleLocal, rectangleParams.row, rectangleParams.col, + rectangleParams.stride); +} + +/** + * @brief QuantPerTensor 同时对row行进行FP32到int8的per-tensor量化操作。一行内共用同一个量化系数。 + outLocal[i , j] = inputLocal[i , j] * quantScaleLocal[i] + * @param outLocal 输出tensor [row , col] + * @param inputLocal 输入tensor [row , col] + * @param quantScaleLocal quant系数 [row , 8]; 8 : 32Bytes对齐 + * @param shareTmpUb 临时buffer 内部需要的空间为 [row * col * 4],源自CastFP32ToINT8 + * @param rectangleParams 描述待处理数据的排布,包括 + row 行数 + col 列数 + stride 一行的真实长度 + */ +template +__aicore__ inline void QuantPerTensor(const LocalTensor &outLocal, const LocalTensor &inputLocal, + const LocalTensor &quantScaleLocal, const LocalTensor &shareTmpUb, + const Rectangle &rectangleParams) +{ + QuantPerTensorVF(outLocal, inputLocal, quantScaleLocal, rectangleParams.row, rectangleParams.col); +} + +/** + * @brief QuantPerTile 对输入tensor进行per-tile量化操作,FP32->fp8 e4m3/hifloat8,每个tile出一个量化系数。 + 量化流程:先对输入的每行的每个tile做动态量化,最后转换为fp8/hif8. + * @param outLocal 输出tensor [row * col],量化后的fp8/hif8数据,后续跟随scale数据 + * @param inputLocal 输入tensor [row * col] + * @param shareTmpUb 临时buffer,内部需要空间为: + * [Align(row * tileNum, 8) + row * tileNum * 8 + 其他中间计算所需空间] * sizeof(float) + * @param perTileQuantParams 描述pertile量化的参数,包括 + * tileSize 每个tile的大小 + * tileNum 每行的tile数量 + * alpha clip的alpha系数 + * row 待处理的行数 + * col 待处理的列数 (col = tileSize * tileNum) + */ +template +__aicore__ inline void QuantPerTile8Bit(const LocalTensor &outLocal, const LocalTensor &inputLocal, + const PerTileQuantParams &perTileQuantParams) +{ + LocalTensor quantScaleLocal = + outLocal[perTileQuantParams.row * perTileQuantParams.col].template ReinterpretCast(); + QuantPerTileVF(outLocal, inputLocal, quantScaleLocal, perTileQuantParams.row, + perTileQuantParams.col, perTileQuantParams.tileSize); +} + +__aicore__ inline void QuantPerTile(const LocalTensor &outLocal, const LocalTensor &inputLocal, + const LocalTensor &shareTmpUb, + const PerTileQuantParams &perTileQuantParams) +{ + LocalTensor scale = + outLocal[perTileQuantParams.row * perTileQuantParams.col].template ReinterpretCast(); + + LocalTensor clipOut = inputLocal; + uint32_t aMaxSizeAlign = + Align(static_cast(perTileQuantParams.row * perTileQuantParams.tileNum), FP32_BLOCK_ELEMENT_NUM); + uint32_t aMaxBrcbSize = perTileQuantParams.row * perTileQuantParams.tileNum * FP32_BLOCK_ELEMENT_NUM; + LocalTensor aMax = shareTmpUb.template ReinterpretCast(); + LocalTensor aMaxBrcb = aMax[aMaxSizeAlign]; + LocalTensor sharedBuf = aMaxBrcb[aMaxBrcbSize].template ReinterpretCast(); + + PerTileClipWithAlpha(clipOut, aMax, aMaxBrcb, inputLocal, sharedBuf, perTileQuantParams); + + LocalTensor quantOut = clipOut; + DynamicQuant(quantOut, scale, clipOut, aMax, aMaxBrcb, sharedBuf, perTileQuantParams.alpha, + perTileQuantParams.row * perTileQuantParams.tileNum, perTileQuantParams.tileSize); + PipeBarrier(); + CastFP32ToINT8(outLocal, quantOut, sharedBuf, perTileQuantParams.row * perTileQuantParams.col); +} + +/** + * @brief DynamicQuant 同时对row行进行dynamicquant, float ---> int8, 每一行出一个系数。 + * @param outLocal 输出tensor [row , col],支持和inputLocal是同一块空间 + * @param inputLocal 输入tensor [row , col] + * @param scale 输出每行的反量化系数 [1 , row] + * @param maxInt8Tensor [1 , row] 元素均为int8的最大值127 + * @param shareTmpUb 临时buffer 内部需要的空间为 [(Align(row * col, 8) + Align(row , 8) * ALIGN_BLOCK_SIZE) * + * sizeof(float)] + * @param row 待处理的行数 + * @param col 待处理的列数 + */ +__aicore__ inline void DynamicQuant(const LocalTensor &outputLocal, const LocalTensor &scale, + const LocalTensor &inputLocal, const LocalTensor &maxInt8Tensor, + const LocalTensor &shareTmpUb, uint64_t row, uint64_t col) +{ + constexpr uint64_t brcnNum = 8; // brcb一次处理8个数据 + uint64_t computeSize = row * col; + LocalTensor inputCopy = shareTmpUb.ReinterpretCast(); + LocalTensor rowMaxBrcb = inputCopy[Align(computeSize, static_cast(ALIGN_BLOCK_SIZE))]; + // abs(x) + Abs(inputCopy, inputLocal, computeSize); + PipeBarrier(); + Rectangle rectangleParams{ + static_cast(row), static_cast(col), + static_cast(col) // columnStride + }; + // rowMax(abs(x)) + RowMax(inputCopy, inputCopy, rectangleParams); + PipeBarrier(); + + // scaleOut = rowMax(abs(x)) / 127 + Div(scale, inputCopy, maxInt8Tensor, row); + PipeBarrier(); + + // 1 / scaleOut = 127 / rowMax(abs(x)) + Div(inputCopy, maxInt8Tensor, inputCopy, row); + PipeBarrier(); + + Brcb(rowMaxBrcb, inputCopy, static_cast(CeilDivT(row, brcnNum)), {1, brcnNum}); + PipeBarrier(); + + // x * 1 / scaleOut + RowMuls(outputLocal, inputLocal, rowMaxBrcb, rectangleParams); +} + +} // namespace MlaProlog +#endif // MLA_PROLOG_VECTOR_COMM_H \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_dequant.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_dequant.h new file mode 100644 index 000000000000..cc0571151836 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_dequant.h @@ -0,0 +1,195 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file service_dequant.h + * \brief + */ + +#ifndef SERVICE_DEQUANT_H +#define SERVICE_DEQUANT_H + +#include "mla_prolog_comm.h" +#include "mla_prolog_vector_comm.h" + +namespace MlaProlog { + +/** + * @brief DequantPerTokenQc 用于对Qc做反量化;按行做dequant流程,给oriRow * col的数据做反量化 + * @param outputGm 输出tensor + * @param inputGm 输入tensor + * @param deqScaleQcQrWGm deqScaleQcQr反量化系数,原shape[1,ND] + * @param deQuantScaleQcQrLocal deQuantScaleQcQr反量化系数,原shape[BS,1] + * @param shareTmpUb 临时buffer + * @param dequantRowColStrideParams 描述待处理数据的排布,包括 + row 行数 + col 列数 + stride 一行的真实长度 + * @param oriRow 一共有多少行 +*/ +template +__aicore__ inline void +DequantPerTokenQc(const GlobalTensor &outputGm, const GlobalTensor &inputGm, + const GlobalTensor &deqScaleQcQrWGm, const LocalTensor deQuantScaleQcQrLocal, + const LocalTensor &shareTmpUb, Rectangle dequantRowColStrideParams, uint32_t oriRow) +{ + int64_t count = dequantRowColStrideParams.row * dequantRowColStrideParams.col; + + LocalTensor inputLocal = shareTmpUb.ReinterpretCast(); // count * sizeof(T) + LocalTensor scaleLocal = inputLocal[count + 16].template ReinterpretCast(); // count * sizeof(C) + LocalTensor computeLocal = scaleLocal[count + 16].template ReinterpretCast(); // count * sizeof(C) + LocalTensor outputLocal = computeLocal[count + 16].template ReinterpretCast(); // count * sizeof(O) + + DataCopyParams copyParams{ + static_cast(dequantRowColStrideParams.row), + static_cast(dequantRowColStrideParams.col * sizeof(T) / 32U), + static_cast((dequantRowColStrideParams.stride - dequantRowColStrideParams.col) * sizeof(T) / 32U), 0}; + + Rectangle rectangleParams{ + static_cast(1), // row + static_cast(count), // col + static_cast(count) // columnStride + }; + + for (int64_t rowOffset = 0; rowOffset < oriRow; rowOffset += dequantRowColStrideParams.row) { + int64_t inputOffset = rowOffset * dequantRowColStrideParams.stride; + int64_t outputOffset = rowOffset * dequantRowColStrideParams.col; + // copy in + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + DataCopy(inputLocal, inputGm[inputOffset], copyParams); + DataCopy(scaleLocal, deqScaleQcQrWGm[inputOffset], copyParams); + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + // compute + Dequant(computeLocal, inputLocal, scaleLocal, deQuantScaleQcQrLocal, rectangleParams); + AscendC::PipeBarrier(); + // cast + Cast(outputLocal, computeLocal, RoundMode::CAST_RINT, count); + SetFlag(EVENT_ID2); + // copy out + WaitFlag(EVENT_ID2); + DataCopy(outputGm[outputOffset], outputLocal, count); + } +} + +/** + * @brief CastPerTokenQc 用于对Qc做类型转换;按行做cast流程,给oriRow * col的数据做类型转换 + * @param outputGm 输出tensor + * @param inputGm 输入tensor + * @param shareTmpUb 临时buffer + * @param castRowColStrideParams 描述待处理数据的排布,包括 + row 行数 + col 列数 + stride 一行的真实长度 + * @param oriRow 一共有多少行 +*/ +template +__aicore__ inline void CastPerTokenQc(const GlobalTensor &outputGm, const GlobalTensor &inputGm, + const LocalTensor &shareTmpUb, Rectangle castRowColStrideParams, + uint32_t oriRow) +{ + int64_t count = castRowColStrideParams.row * castRowColStrideParams.col; + + LocalTensor inputLocal = shareTmpUb.ReinterpretCast(); // count * sizeof(T) + LocalTensor outputLocal = inputLocal[count + 16].template ReinterpretCast(); // count * sizeof(O) + + DataCopyParams copyParams{ + static_cast(castRowColStrideParams.row), + static_cast(castRowColStrideParams.col * sizeof(T) / 32U), + static_cast((castRowColStrideParams.stride - castRowColStrideParams.col) * sizeof(T) / 32U), 0}; + + for (int64_t rowOffset = 0; rowOffset < oriRow; rowOffset += castRowColStrideParams.row) { + int64_t inputOffset = rowOffset * castRowColStrideParams.stride; + int64_t outputOffset = rowOffset * castRowColStrideParams.col; + // copy in + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + DataCopy(inputLocal, inputGm[inputOffset], copyParams); + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + // cast + Cast(outputLocal, inputLocal, RoundMode::CAST_RINT, count); + SetFlag(EVENT_ID2); + // copy out + WaitFlag(EVENT_ID2); + DataCopy(outputGm[outputOffset], outputLocal, count); + } +} + +/** + DequantSplitNQc 用于enableGroupComputeOpt场景 + 场景特征:半量化,BS = 1, headsize = 8 + * @brief 用于对Qc做反量化;按N方向切分的dequant,理解为做一个head的 + * @param outputGm 输出位置 + * @param inputGm 输入tensor + * @param deqScaleQcQrWGm deqScaleQcQr反量化系数,原shape[1,ND] + * @param deQuantScaleQcQrLocal deQuantScaleQcQr反量化系数,原shape[BS,1] + * @param shareTmpUb 临时buffer + * @param dequantRowColStrideParams 描述待处理数据的排布,包括 + row 行数 + col 列数 + stride 一行的真实长度 + * @param oriCol 一列真实的长度 + * @param dstStride 目的数据行间的偏移 + * @param subBlockIdx_ 这是奇数核还是偶数核 + */ +// 用于enableGroupComputeOpt场景 +template +__aicore__ inline void +DequantSplitNQc(const GlobalTensor &outputGm, const GlobalTensor &inputGm, const GlobalTensor &deqScaleQcQrWGm, + const LocalTensor &deQuantScaleQcQrLocal, const LocalTensor &shareTmpUb, + Rectangle dequantRowColStrideParams, uint32_t oriCol, uint32_t dstStride, uint32_t subBlockIdx_) +{ + int64_t count = dequantRowColStrideParams.col * 1; + LocalTensor inputLocal = shareTmpUb.ReinterpretCast(); // count * sizeof(T) + LocalTensor scaleLocal = inputLocal[count + 16].template ReinterpretCast(); // count * sizeof(C) + LocalTensor computeLocal = scaleLocal[count + 16].template ReinterpretCast(); // count * sizeof(C) + LocalTensor outputLocal = computeLocal[count + 16].template ReinterpretCast(); // count * sizeof(O) + + DataCopyParams inputCopyParams{ + static_cast(1), static_cast(count * sizeof(T) / 32U), + static_cast((dequantRowColStrideParams.stride - dequantRowColStrideParams.col) * sizeof(T) / 32U), 0}; + + DataCopyParams outputCopyParams{static_cast(1), static_cast(count * sizeof(O) / 32U), 0, + static_cast(dstStride * sizeof(O) / 32U)}; + + Rectangle rectangleParams{ + static_cast(1), // row + static_cast(count), // col + static_cast(count) // columnStride + }; + + SetFlag(EVENT_ID0); + int64_t scaleOffset = subBlockIdx_ * count; + DataCopy(scaleLocal, deqScaleQcQrWGm[scaleOffset], dequantRowColStrideParams.col); + for (int64_t rowOffset = 0; rowOffset < dequantRowColStrideParams.row; rowOffset++) { + int64_t inputOffset = rowOffset * oriCol + subBlockIdx_ * count; + int64_t outputOffset = rowOffset * oriCol + subBlockIdx_ * count; + WaitFlag(EVENT_ID0); + DataCopy(inputLocal, inputGm[inputOffset], inputCopyParams); + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + // compute + Dequant(computeLocal, inputLocal, scaleLocal, deQuantScaleQcQrLocal, rectangleParams); + AscendC::PipeBarrier(); + // cast + Cast(outputLocal, computeLocal, RoundMode::CAST_RINT, count); + SetFlag(EVENT_ID2); + // copy out + WaitFlag(EVENT_ID2); + DataCopy(outputGm[outputOffset], outputLocal, count); + SetFlag(EVENT_ID0); + } + WaitFlag(EVENT_ID0); +} + +} // namespace MlaProlog +#endif diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_dynamic_quant_qn_mul_qr.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_dynamic_quant_qn_mul_qr.h new file mode 100644 index 000000000000..42a4fa89eb28 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_dynamic_quant_qn_mul_qr.h @@ -0,0 +1,236 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file service_dynamic_quant_qn_mul_qr.h + * \brief + */ + +#ifndef SERVICE_DYNAMIC_QUANT_H +#define SERVICE_DYNAMIC_QUANT_H + +#include "mla_prolog_comm.h" +#include "mla_prolog_vector_comm.h" +#include "vf/vf_mul_qr.h" +#include "vf/vf_dynamic_quant.h" + +namespace MlaProlog { + +template +__aicore__ inline void DynamicQuantMultiRow(const GlobalTensor &outputGm, const LocalTensor &scaleOutputLocal, + const GlobalTensor &inputGm, const LocalTensor outputLocal, + const LocalTensor &inputHalf, const LocalTensor &inputLocal, + const LocalTensor &maxInt8Tensor, const LocalTensor &shareTmpUb, + uint64_t row, uint64_t col, uint64_t subRow, uint64_t queryOutStride, + uint32_t DYNAMIC_QUANT_INPUT_READY, uint32_t DYNAMIC_QUANT_OUTPUT_READY) +{ + constexpr uint64_t inputBlockAlign = (ALIGN_BLOCK_SIZE / sizeof(T)); + constexpr uint64_t outputBlockAlign = (ALIGN_BLOCK_SIZE / sizeof(O)); + // DataCopyParams : count len srcStrideIn dstStrideIn + DataCopyParams outParams{static_cast(subRow), static_cast(col / outputBlockAlign), 0, + static_cast((queryOutStride - col) / outputBlockAlign)}; + DataCopyParams inputParams{static_cast(subRow), static_cast(col / inputBlockAlign), + static_cast((queryOutStride - col) / inputBlockAlign), 0}; + + uint32_t computeSize = subRow * col; + uint64_t loopCnt = CeilDivT(row, subRow); + uint64_t lastSubRow = row - (loopCnt - 1) * subRow; + uint64_t inputGmOffset = 0; + uint64_t scaleOffset = 0; + + for (uint64_t loopIdx = 0; loopIdx < loopCnt; loopIdx++) { + if ((loopIdx == loopCnt - 1) && (lastSubRow != subRow)) { + subRow = lastSubRow; + computeSize = subRow * col; + outParams.blockCount = subRow; + inputParams.blockCount = subRow; + } + + WaitFlag(DYNAMIC_QUANT_INPUT_READY); // 计算是否已经完成可以搬运 + DataCopy(inputHalf, inputGm[inputGmOffset], inputParams); + SetFlag(DYNAMIC_QUANT_INPUT_READY); + WaitFlag(DYNAMIC_QUANT_INPUT_READY); // 搬运是否已经完成可以计算 + + LocalTensor output = outputLocal.template ReinterpretCast(); + WaitFlag(DYNAMIC_QUANT_OUTPUT_READY); + PipeBarrier(); + DynamicQuantPerTokenVf(output, scaleOutputLocal[scaleOffset], inputHalf, subRow, col); + PipeBarrier(); + SetFlag(DYNAMIC_QUANT_INPUT_READY); + SetFlag(DYNAMIC_QUANT_OUTPUT_READY); + WaitFlag(DYNAMIC_QUANT_OUTPUT_READY); // 计算是否已经完成可以搬运 + DataCopy(outputGm[inputGmOffset], output, outParams); + SetFlag(DYNAMIC_QUANT_OUTPUT_READY); + inputGmOffset += subRow * queryOutStride; + scaleOffset += subRow; + } +} + +template +__aicore__ inline void +MulQr(const GlobalTensor &outputGmRope, const GlobalTensor &inputGmRope, LocalTensor outputLocalRope, + const LocalTensor &qrInputLocal, const LocalTensor &qrFp32Local, const LocalTensor &reciprocalLocal, + const LocalTensor &dequantScaleBrcbLocal, uint64_t row, uint64_t colRope, uint64_t subRowRope, + uint64_t qrOutputStrideRope, float quantScaleCkvRope, uint32_t MUL_QR_INPUT_COPY_READY, uint32_t MUL_QR) +{ + constexpr uint64_t inputBlockAlign = (ALIGN_BLOCK_SIZE / sizeof(T)); + constexpr uint64_t computeBlockAlign = (ALIGN_BLOCK_SIZE / sizeof(C)); + + uint64_t loopCntRope = CeilDivT(row, subRowRope); + uint64_t lastSubRowRope = row - (loopCntRope - 1) * subRowRope; + uint32_t computeSizeRope = subRowRope * colRope; + + // DataCopyParams : count len srcStrideIn dstStrideIn + DataCopyParams inputParamsRope{static_cast(row), static_cast(colRope / inputBlockAlign), + static_cast((qrOutputStrideRope - colRope) / inputBlockAlign), 0}; + DataCopyParams outputParamsRope{static_cast(subRowRope), static_cast(colRope / inputBlockAlign), + 0, static_cast((qrOutputStrideRope - colRope) / inputBlockAlign)}; + + uint64_t inputGmRopeOffset = 0; + uint64_t dequantScaleOffset = 0; + uint64_t inputLocalRopeOffset = 0; + + + WaitFlag(MUL_QR_INPUT_COPY_READY); + DataCopy(qrInputLocal, inputGmRope, inputParamsRope); + for (uint64_t loopIdx = 0; loopIdx < loopCntRope; loopIdx++) { + if (loopIdx == (loopCntRope - 1) && subRowRope != lastSubRowRope) { + subRowRope = lastSubRowRope; + computeSizeRope = subRowRope * colRope; + outputParamsRope.blockCount = subRowRope; + } + SetFlag(MUL_QR); + WaitFlag(MUL_QR); + + SetFlag(MUL_QR); + WaitFlag(MUL_QR); + + Cast(qrFp32Local, qrInputLocal[inputLocalRopeOffset], RoundMode::CAST_NONE, computeSizeRope); + PipeBarrier(); + MulQrVF(qrFp32Local, qrFp32Local, dequantScaleBrcbLocal, quantScaleCkvRope, computeSizeRope, computeBlockAlign); + Cast(outputLocalRope, qrFp32Local, RoundMode::CAST_RINT, computeSizeRope); + PipeBarrier(); + + SetFlag(MUL_QR); + WaitFlag(MUL_QR); + + DataCopy(outputGmRope[inputGmRopeOffset], outputLocalRope, outputParamsRope); + inputLocalRopeOffset += subRowRope * colRope; + inputGmRopeOffset += subRowRope * qrOutputStrideRope; + dequantScaleOffset += subRowRope * computeBlockAlign; + } + SetFlag(MUL_QR_INPUT_COPY_READY); +} +/** + * @brief DynamicQuantQnWithMulQr 流程中dynamicquant和ropeMul融合流程 + UB,流水,中间产物统筹规划 + * @tparam T mmQnOutputType -> bf16 + * @tparam C dequantScaleQNopeType -> float + * @tparam O queryOutputType -> int8 + * @param scaleOutputGm 动态量化参数输出GM位置 + * @param outputGm 动态量化结果输出GM位置 + * @param outputGmRope rope处理后的输出位置 + * @param inputGm dynamicQuant部分的输入 + * @param shareTmpUb 临时buffer空间 + * @param row + * @param col + * @param scaleOutStride 描述动态量化参数输出GM位置的偏移关系 + * @param queryOutStride 描述动态量化结果输出GM位置 + * @param inputGmRope + * @param quantScaleCkvRope + * @param colRope + * @param qrOutputStrideRope 描述rope处理后的输出位置 + */ +template +__aicore__ inline void DynamicQuantQnWithMulQr( + // Dynamic Quant With MulQr 输出 + const GlobalTensor &scaleOutputGm, const GlobalTensor &outputGm, const GlobalTensor &outputGmRope, + // Dynamic Quant 入参 + const GlobalTensor &inputGm, LocalTensor &shareTmpUb, uint64_t row, uint64_t col, + uint64_t scaleOutStride, uint64_t queryOutStride, + // Mul Qr 入参 + const GlobalTensor &inputGmRope, float quantScaleCkvRope, uint64_t colRope, uint64_t qrOutputStrideRope, + uint32_t cvRatio) +{ + if (row == 0 || col == 0) { + return; + } + // 常量 + constexpr uint32_t MUL_QR = EVENT_ID1; // 用于控制Mul_Qr的同步 + // dynamicquant输入的计算/搬运 是否已经完成可以开始下一轮的 搬运/计算 + constexpr uint32_t DYNAMIC_QUANT_INPUT_READY = EVENT_ID0; + constexpr uint32_t MUL_QR_INPUT_COPY_READY = EVENT_ID3; // 是否可以开始下一轮的MUL_QR_INPUT_COPY + constexpr uint32_t CALC_SCALE_FINISH = EVENT_ID0; // dynamicquant的scale是否计算完成可以开始搬运 + // dynamicquant输出的计算/搬运 是否已经完成可以开始下一轮的 搬运/计算 + constexpr uint32_t DYNAMIC_QUANT_OUTPUT_READY = EVENT_ID3; + + constexpr uint64_t computeBlockAlign = (ALIGN_BLOCK_SIZE / sizeof(C)); + constexpr uint64_t inputBlockAlign = (ALIGN_BLOCK_SIZE / sizeof(T)); + constexpr float maxInt8 = 127.0; + // Dynamic Quant 局部变量 + uint64_t rowStepSize = 8 * cvRatio; // 单次处理最大行数, cv1:1场景,改为8行进行计算,降低UB使用 + + uint64_t subRow = row < rowStepSize ? row : rowStepSize; + uint32_t computeSize = subRow * col; + uint32_t brcbCnt = Align(subRow, computeBlockAlign); + + // Rope Post Process 局部变量 + uint64_t rowStepSizeRope = 64; // 单次处理最大行数 + uint64_t subRowRope = row < rowStepSizeRope ? row : rowStepSizeRope; + uint32_t computeSizeRope = subRowRope * colRope; + uint32_t loadRopeGmSize = row * colRope; + + // Dynamic Quant & Rope Post Process UB Buffer分配 + LocalTensor inputHalf = shareTmpUb.ReinterpretCast(); + LocalTensor qrInputLocal = inputHalf[computeSize]; + LocalTensor outputLocalRope = qrInputLocal[loadRopeGmSize + inputBlockAlign]; + + LocalTensor inputLocal = outputLocalRope[computeSizeRope + inputBlockAlign].template ReinterpretCast(); + LocalTensor outputLocal = inputLocal[computeSize]; + LocalTensor maxInt8Tensor = outputLocal[computeSize]; + LocalTensor dynamicQuantUb = maxInt8Tensor[brcbCnt]; + LocalTensor scaleOutputLocal = dynamicQuantUb[computeSize + brcbCnt * computeBlockAlign]; + LocalTensor scaleBrcb = scaleOutputLocal[Align(row, computeBlockAlign)]; + + // Rope Post Process流程在Dynamic Quant流程结束,两者UB Buffer不会同时使用,故可以从起始位置开始重新计算,减少UB + // Buffer使用 + LocalTensor qrFp32Local = inputLocal; + LocalTensor reciprocalLocal = qrFp32Local[computeSizeRope + computeBlockAlign]; + + // Dynamic Quant + Duplicate(maxInt8Tensor, static_cast(maxInt8), brcbCnt); + PipeBarrier(); + + DynamicQuantMultiRow(outputGm, scaleOutputLocal, inputGm, outputLocal, inputHalf, inputLocal, maxInt8Tensor, + dynamicQuantUb.template ReinterpretCast(), row, col, subRow, queryOutStride, + DYNAMIC_QUANT_INPUT_READY, DYNAMIC_QUANT_OUTPUT_READY); + + Brcb(scaleBrcb, scaleOutputLocal, CeilDivT(row, computeBlockAlign), {1, computeBlockAlign}); + PipeBarrier(); + + // DataCopyParams : count len srcStrideIn dstStrideIn + DataCopyParams scaleOutCopyParams{(uint16_t)row, (uint16_t)sizeof(C), 0, + (uint16_t)((scaleOutStride - 1) * sizeof(C))}; + + SetFlag(CALC_SCALE_FINISH); + WaitFlag(CALC_SCALE_FINISH); + DataCopyPad(scaleOutputGm, scaleBrcb, scaleOutCopyParams); + + // qr rope 后的乘法 + MulQr(outputGmRope, inputGmRope, outputLocalRope, qrInputLocal, qrFp32Local, reciprocalLocal, scaleBrcb, row, + colRope, subRowRope, qrOutputStrideRope, quantScaleCkvRope, MUL_QR_INPUT_COPY_READY, MUL_QR); + + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); +} + +} // namespace MlaProlog + +#endif diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_gather_sin_cos.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_gather_sin_cos.h new file mode 100644 index 000000000000..aca6647a0755 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_gather_sin_cos.h @@ -0,0 +1,66 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file service_gather_sin_cos.h + * \brief + */ + +#ifndef SERVICE_GATHER_SIN_COS_H +#define SERVICE_GATHER_SIN_COS_H + +#include "mla_prolog_comm.h" + +namespace MlaProlog { + +/** + * @brief GatherSinCos 读取rope所需的sin和cos系数,并进行预处理,rope中sin的-1系数会在此处融合进sin中; + * @param cosLocal cos搬运的位置; + * @param sinLocal sin搬运的位置;rope中sin的-1系数会在此处融合进sin中; + * @param cosGm cos在GM的位置,可以认为一个token有col个系数;整体连续排布 + * @param sinGm sin在GM的位置,同上 + * @param tokenIndex 从哪个token开始读取 + * @param curVecToken 读取多少个token的系数 + * @param shareTmpUb 临时buffer 内部需要的空间为 [2 * curVecToken * col * sizeof(bfloat16_t)] + * @param row 行数;预留参数 + * @param col 列数 + */ +template +__aicore__ inline void GatherSinCos(LocalTensor &cosLocal, LocalTensor &sinLocal, const GlobalTensor &cosGm, + const GlobalTensor &sinGm, int64_t tokenIndex, int64_t curVecToken, + LocalTensor &shareTmpUb, int64_t row, int64_t col) +{ + int64_t offset = col * tokenIndex; + int64_t curDataSize = col * curVecToken; + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + if constexpr (IsSameType::value) { + auto tmpUbBf16 = shareTmpUb.ReinterpretCast(); + DataCopy(tmpUbBf16, cosGm[offset], curDataSize); + DataCopy(tmpUbBf16[curDataSize], sinGm[offset], curDataSize); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + Cast(cosLocal, tmpUbBf16, RoundMode::CAST_NONE, curDataSize); + Cast(sinLocal, tmpUbBf16[curDataSize], RoundMode::CAST_NONE, curDataSize); + PipeBarrier(); + } else { + DataCopy(cosLocal, cosGm[offset], curDataSize); + DataCopy(sinLocal, sinGm[offset], curDataSize); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + } + uint8_t blockNumPerRow = col / (ALIGN_BLOCK_SIZE / sizeof(O)); + Muls(sinLocal, sinLocal, -1.0f, col >> 1, curVecToken, {1, 1, blockNumPerRow, blockNumPerRow}); + PipeBarrier(); +} + +} // namespace MlaProlog + +#endif \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_matmul.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_matmul.h new file mode 100644 index 000000000000..7a5d35efd87b --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_matmul.h @@ -0,0 +1,987 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/* ! + * \file service_matmul.h + * \brief + */ + +#ifndef SERVICE_MATMUL_H +#define SERVICE_MATMUL_H + +#include "mla_prolog_comm.h" + +namespace MlaProlog { + +constexpr uint8_t UNIT_FLAG_DISABLE = 0; // 0: disable: 不配置unitFlag +constexpr uint8_t UNIT_FLAG_CHECK = 0b10; // 2: enable: 将mmadParams.unitFlag设置为 0b10 +constexpr uint8_t UNIT_FLAG_SET = 0b11; // 3: enable: 在k的最后一轮循环,会将mmadParams.unitFlag设置为 0b11 + + +/** + * @brief Struct to encapsulate all local tensor objects for matrix multiplication + * @tparam T Data type for A, B, and L0 tensors + * @tparam O_L0C Data type for C L0 tensor + */ +template +struct mmLocalTensors { + LocalTensor aL1Tensor; + LocalTensor bL1Tensor; + LocalTensor aL0Tensor; + LocalTensor bL0Tensor; + LocalTensor cL0Tensor; + + /** + * @brief Initialize all local tensors with buffer addresses from bufParam + * @param bufParam Buffer parameter object containing all buffer addresses + */ + __aicore__ inline void Init(const MMBufParams &bufParam) + { + aL1Tensor.SetAddr(bufParam.aL1BufAddr); + bL1Tensor.SetAddr(bufParam.bL1BufAddr); + aL0Tensor.SetAddr(bufParam.aL0BufAddr); + bL0Tensor.SetAddr(bufParam.bL0BufAddr); + cL0Tensor.SetAddr(bufParam.cL0BufAddr); + } +}; + +template +__aicore__ inline constexpr uint32_t GetC0Num() +{ + if (sizeof(SrcT) == sizeof(float)) { + return 8; + } else if (sizeof(SrcT) == sizeof(int8_t)) { + return 32; + } + return 16; +} + +template +__aicore__ inline void CopyNDGmToL1(LocalTensor &l1Tensor, const GlobalTensor &gmSrcTensor, uint32_t srcN, + uint32_t srcD, uint32_t srcDstride) +{ + Nd2NzParams nd2nzPara; + nd2nzPara.ndNum = 1; + nd2nzPara.nValue = srcN; // 行数 + nd2nzPara.dValue = srcD; + nd2nzPara.srcDValue = srcDstride; + nd2nzPara.dstNzC0Stride = CeilDivT(srcN, BLOCK_CUBE_SIZE) * BLOCK_CUBE_SIZE; // 对齐到16 单位 block + nd2nzPara.dstNzNStride = 1; + nd2nzPara.srcNdMatrixStride = 0; + nd2nzPara.dstNzMatrixStride = 0; + DataCopy(l1Tensor, gmSrcTensor, nd2nzPara); +} + +template +__aicore__ inline void CopyNZGmToL1(LocalTensor &l1Tensor, const GlobalTensor &gmSrcTensor, uint32_t srcN, + uint32_t srcD, uint32_t srcNstride) +{ + DataCopyParams param; + param.blockCount = CeilDivT(srcD, GetC0Num()); + param.blockLen = srcN; // 单位为32B srcN*16/16 + param.srcStride = (srcNstride - srcN); // 单位为32B (srcNstride - srcN)*16/16 + param.dstStride = 0; + DataCopy(l1Tensor, gmSrcTensor, param); +} + +template +__aicore__ inline void LoadL1A(const GlobalTensor &tensorAGm, const uint32_t mInput, const uint32_t kL1StepSize, + const uint32_t kInput, MMBufParams &bufParam) +{ + uint32_t bufIter = bufParam.aL1BufIter; + if constexpr (preload) { + bufIter++; + } + LocalTensor aL1Tensor; + aL1Tensor.SetAddr(bufParam.aL1BufAddr); + uint32_t aL1Offset = L1_A_SIZE / sizeof(T); + + WaitFlag(A_EVENT0 + (bufIter & 1u)); + LocalTensor aL1 = aL1Tensor[(bufIter & 1u) * aL1Offset]; + CopyNDGmToL1(aL1, tensorAGm, mInput, kL1StepSize, kInput); + SetFlag(A_EVENT0 + (bufIter & 1u)); +} + +template +__aicore__ inline void LoadL1B(const GlobalTensor &tensorBGm, const uint32_t nL1Size, const uint32_t nInput, + const uint32_t kL1StepSize, const uint32_t kInput, MMBufParams &bufParam) +{ + uint32_t bufIter = bufParam.bL1BufIter; + if constexpr (preload) { + bufIter++; + } + LocalTensor bL1Tensor; + bL1Tensor.SetAddr(bufParam.bL1BufAddr); + uint32_t bL1Offset = L1_B_SIZE / sizeof(T); + + WaitFlag(B_EVENT0 + (bufIter & 1u)); + LocalTensor bL1 = bL1Tensor[(bufIter & 1u) * bL1Offset]; + if constexpr (loadFormat == DataFormat::NZ) { + CopyNZGmToL1(bL1, tensorBGm, kL1StepSize, nL1Size, kInput); + } else { + CopyNDGmToL1(bL1, tensorBGm, kInput, nL1Size, nInput); + } + SetFlag(B_EVENT0 + (bufIter & 1u)); +} + +template +__aicore__ inline void LoadL1AAndScale(const GlobalTensor &tensorAGm, const GlobalTensor &tensorAScaleGm, + const uint32_t mInput, const uint32_t kL1StepSize, const uint32_t kInput, + const uint32_t nInput, const uint64_t scaleAInL1BOffset, MMBufParams &bufParam) +{ + uint32_t bufIter = bufParam.aL1BufIter; + if constexpr (preload) { + bufIter++; + } + LocalTensor aL1Tensor; + aL1Tensor.SetAddr(bufParam.aL1BufAddr); + uint32_t aL1Offset = L1_A_SIZE / sizeof(T); + + LocalTensor bL1Tensor; + bL1Tensor.SetAddr(bufParam.bL1BufAddr); + uint32_t srcDValue = nInput / FP8_TWO; + if constexpr (scaleSrcPadFlag) { + srcDValue = Align(nInput / FP8_TWO, BLOCK_CUBE_SIZE); + } + GlobalTensor tensorScaleGmCast; + tensorScaleGmCast.SetGlobalBuffer(((__gm__ bfloat16_t *)(tensorAScaleGm.GetPhyAddr()))); + + WaitFlag(A_EVENT0 + (bufIter & 1u)); + LocalTensor aL1 = aL1Tensor[(bufIter & 1u) * aL1Offset]; + CopyNDGmToL1(aL1, tensorAGm, mInput, kL1StepSize, kInput); + + Dn2NzParams dn2Nzparam; + dn2Nzparam.dnNum = 1; + dn2Nzparam.nValue = nInput / FP8_TWO; // 单个DN矩阵的列数N, 单位为元素个数 + dn2Nzparam.dValue = mInput; // 单个DN矩阵的实际行数D, 单位为元素个数 + dn2Nzparam.srcDnMatrixStride = 0; // 相邻DN矩阵起始地址之间的偏移, 单位为元素个数 + dn2Nzparam.srcDValue = srcDValue; // 同一个DN矩阵中相邻行起始地址之间的偏移, 单位为元素个数 + dn2Nzparam.dstNzC0Stride = + nInput / FP8_TWO; // 来自DN矩阵中同一列的相邻两个block在NZ矩阵中起始地址间的偏移,单位为block(32B)个数 + dn2Nzparam.dstNzNStride = 1; // 转换为NZ矩阵后,DN中相邻两行在NZ矩阵中起始地址之间的偏移,单位为元素个数 + dn2Nzparam.dstNzMatrixStride = dn2Nzparam.nValue; // 两个NZ矩阵,起始地址之间的偏移,单位为block个数 + DataCopy(bL1Tensor[scaleAInL1BOffset / FP8_TWO], tensorScaleGmCast, dn2Nzparam); + + SetFlag(A_EVENT0 + (bufIter & 1u)); +} + +template +__aicore__ inline void LoadL1BAndScale(const GlobalTensor &tensorBGm, const GlobalTensor &tensorBScaleGm, + const uint32_t nL1Size, const uint32_t kL1StepSize, const uint32_t kInput, + const uint32_t nInput, const uint64_t scaleBInL1BOffset, MMBufParams &bufParam) +{ + uint32_t bufIter = bufParam.bL1BufIter; + if constexpr (preload) { + bufIter++; + } + LocalTensor bL1Tensor; + bL1Tensor.SetAddr(bufParam.bL1BufAddr); + uint32_t bL1Offset = L1_B_SIZE / sizeof(T); + + LocalTensor bL1ScaleTensor; + bL1ScaleTensor.SetAddr(bufParam.bL1BufAddr); + GlobalTensor tensorScaleGmCast; + tensorScaleGmCast.SetGlobalBuffer(((__gm__ bfloat16_t *)(tensorBScaleGm.GetPhyAddr()))); + + WaitFlag(B_EVENT0 + (bufIter & 1u)); + + LocalTensor bL1 = bL1Tensor[(bufIter & 1u) * bL1Offset]; + if constexpr (loadFormat == DataFormat::NZ) { + CopyNZGmToL1(bL1, tensorBGm, kL1StepSize, nL1Size, kInput); + } else { + CopyNDGmToL1(bL1, tensorBGm, kInput, nL1Size, nL1Size); + } + + Dn2NzParams dn2Nzparam; + dn2Nzparam.dnNum = 1; + dn2Nzparam.nValue = nInput / FP8_TWO; // 单个DN矩阵的列数N, 单位为元素个数 + dn2Nzparam.dValue = nL1Size; // 单个DN矩阵的实际行数D, 单位为元素个数 + dn2Nzparam.srcDnMatrixStride = 0; // 相邻DN矩阵起始地址之间的偏移, 单位为元素个数 + dn2Nzparam.srcDValue = nInput / FP8_TWO; // 同一个DN矩阵中相邻行起始地址之间的偏移, 单位为元素个数 + dn2Nzparam.dstNzC0Stride = + dn2Nzparam.srcDValue; // 来自DN矩阵中同一列的相邻两个block在NZ矩阵中起始地址间的偏移,单位为block(32B)个数 + dn2Nzparam.dstNzNStride = 1; // 转换为NZ矩阵后,DN中相邻两行在NZ矩阵中起始地址之间的偏移,单位为元素个数 + dn2Nzparam.dstNzMatrixStride = dn2Nzparam.dstNzC0Stride; // 两个NZ矩阵,起始地址之间的偏移,单位为block个数 + DataCopy(bL1ScaleTensor[scaleBInL1BOffset / FP8_TWO], tensorScaleGmCast, dn2Nzparam); + + SetFlag(B_EVENT0 + (bufIter & 1u)); +} + +/** + * @brief 使用3D Pro模式将数据从L1加载到L0,可选择是否进行转置操作 + * @tparam T 张量的数据类型 + * @tparam enTranspose 是否启用转置操作(用于B矩阵) + * @param l0Tensor L0缓冲区中的目标张量 + * @param l1Tensor L1缓冲区中的源张量 + * @param mSize 行维度(对应A矩阵的m,B矩阵的k) + * @param kSize 列维度(对应A矩阵的k,B矩阵的n) + * @param l1StepSize L1缓冲区中的步长大小,用于B矩阵转置 + * @param kl1Size L1缓冲区中的大小,用于B矩阵转置 + */ +template +__aicore__ inline void LoadDataL1ToL0(const LocalTensor &l0Tensor, const LocalTensor &l1Tensor, + const uint32_t mSize, const uint32_t kSize, const uint32_t l1StepSize = 0) +{ + if constexpr (enTranspose) { // B + LoadData2DParamsV2 loadData2DV2; + loadData2DV2.mStartPosition = 0; + loadData2DV2.kStartPosition = 0; + loadData2DV2.mStep = ((mSize + ROUND_UP_UNIT) >> SHIFTS_UNIT << SHIFTS_UNIT) / BLOCK_CUBE_SIZE; + loadData2DV2.kStep = ((kSize + ROUND_UP_UNIT) >> SHIFTS_UNIT << SHIFTS_UNIT) / BLOCK_CUBE_SIZE; + loadData2DV2.srcStride = ((l1StepSize + ROUND_UP_UNIT) >> SHIFTS_UNIT << SHIFTS_UNIT) / BLOCK_CUBE_SIZE; + loadData2DV2.dstStride = ((kSize + ROUND_UP_UNIT) >> SHIFTS_UNIT << SHIFTS_UNIT) / BLOCK_CUBE_SIZE; + loadData2DV2.ifTranspose = true; + LoadData(l0Tensor, l1Tensor, loadData2DV2); + } else { // A + LoadData3DParamsV2 loadDataAParams; + loadDataAParams.l1W = 1; + loadDataAParams.l1H = mSize; + loadDataAParams.channelSize = kSize; + loadDataAParams.kExtension = kSize; + loadDataAParams.mExtension = mSize; + loadDataAParams.kStartPt = 0; + loadDataAParams.mStartPt = 0; + loadDataAParams.strideW = 1; + loadDataAParams.strideH = 1; + loadDataAParams.filterW = 1; + loadDataAParams.filterH = 1; + loadDataAParams.dilationFilterW = 1; + loadDataAParams.dilationFilterH = 1; + loadDataAParams.enTranspose = false; + loadDataAParams.enSmallK = false; + loadDataAParams.padValue = 0; + loadDataAParams.filterSizeW = 0; + loadDataAParams.filterSizeH = 0; + loadDataAParams.fMatrixCtrl = false; + uint16_t dstStride = CeilDivT(mSize, BLOCK_CUBE_SIZE); +#if defined(ASC_DEVKIT_VERSION_NUM) && (ASC_DEVKIT_VERSION_NUM >= 90000000) + SetLoadDataRepeatWithStride({0, 1, 0, dstStride}); // >= 9.0.0 release 新 API + LoadDataWithStride(l0Tensor, l1Tensor, loadDataAParams); +#else + SetLoadDataRepeat({0, 1, 0, dstStride}); // < 9.0.0 (beta.2) 旧 API + LoadData(l0Tensor, l1Tensor, loadDataAParams); +#endif + } +} + +template +__aicore__ inline void LoadDataL1ToL0Mxfp8(const LocalTensor &l0Tensor, const LocalTensor &l1Tensor, + const LocalTensor &l1AMxTensor, const uint32_t mSize, + const uint32_t kSize, const uint32_t kL1Size = 0) +{ + // LoadData and A Scale + LocalTensor dstTensor = l0Tensor.template ReinterpretCast(); + LocalTensor l1MxTensor = l1AMxTensor.template ReinterpretCast(); + + LoadData2DParamsV2 loadDataA2DParams; + loadDataA2DParams.mStartPosition = 0; + loadDataA2DParams.kStartPosition = 0; + loadDataA2DParams.ifTranspose = false; + loadDataA2DParams.mStep = CeilDivT(mSize, BLOCK_CUBE_SIZE); // m轴分型大小为16 + loadDataA2DParams.kStep = CeilDivT(kSize, K_STEP_SIZE_32); // k轴分型大小为32B + // 配合ub->L1使用256 * 32 / 256 + // 64搬运 + loadDataA2DParams.srcStride = loadDataA2DParams.mStep; // m轴全搬,无pad + loadDataA2DParams.dstStride = loadDataA2DParams.mStep; // 搬运前后不转置,m轴不变 + + LoadData2DMxParams loadAScaleParam; + loadAScaleParam.xStartPosition = 0; + loadAScaleParam.yStartPosition = 0; + loadAScaleParam.xStep = CeilDivT(mSize, BLOCK_CUBE_SIZE); + loadAScaleParam.yStep = CeilDivT(kSize, K_STEP_SIZE_32 * FP8_TWO); // ksize对应baseK + loadAScaleParam.srcStride = CeilDivT(kL1Size, K_STEP_SIZE_32 * FP8_TWO); // kL1Size对应baseK*stepK + loadAScaleParam.dstStride = loadAScaleParam.yStep; + LoadData(dstTensor, l1Tensor, l1MxTensor, loadDataA2DParams, loadAScaleParam); +} + +template +__aicore__ inline void LoadDataL1ToL0B(const LocalTensor &bL0Tensor, const LocalTensor &bL1Tensor, + const LocalTensor &bscaleL1, const uint32_t kSize, const uint32_t nSize, + const uint32_t kL1StepSize = 0, const uint32_t kL1Size = 0) +{ + LoadData2DParamsV2 loadData2DV2; + loadData2DV2.mStartPosition = 0; + loadData2DV2.kStartPosition = 0; + loadData2DV2.mStep = CeilDivT(kSize, BLOCK_CUBE_SIZE); + loadData2DV2.kStep = CeilDivT(nSize, K_STEP_SIZE_32); + loadData2DV2.srcStride = CeilDivT(kL1StepSize, BLOCK_CUBE_SIZE); // k轴切分,stride为完整k的长度 + loadData2DV2.dstStride = + CeilDivT(nSize, UNIT_SIZE / K_STEP_SIZE_32); // n轴切分,B矩阵搬运后转置,dst中M轴变成src中的K轴 + loadData2DV2.ifTranspose = true; // 搬运后转置 + + if (std::is_same::value) { + LocalTensor dstTensor = bL0Tensor.template ReinterpretCast(); + LocalTensor srcTensor = bL1Tensor.template ReinterpretCast(); + LocalTensor l1MxTensor = bscaleL1.template ReinterpretCast(); + LoadData2DMxParams loadDataMxParams; + loadDataMxParams.xStartPosition = 0; + loadDataMxParams.yStartPosition = 0; + loadDataMxParams.xStep = CeilDivT(nSize, BLOCK_CUBE_SIZE); + loadDataMxParams.yStep = CeilDivT(kSize, K_STEP_SIZE_32 * FP8_TWO); // ksize对应baseK + loadDataMxParams.srcStride = CeilDivT(kL1Size, K_STEP_SIZE_32 * FP8_TWO); // kL1Size对应baseK*stepK + loadDataMxParams.dstStride = loadDataMxParams.yStep; + LoadData(dstTensor, srcTensor, l1MxTensor, loadData2DV2, loadDataMxParams); + } else { + LoadData(bL0Tensor, bL1Tensor, loadData2DV2); + } +} + +template +__aicore__ inline void MatmulL0(MMBufParams &bufParam, const LocalTensor &aL1, const LocalTensor &bL1, + const LocalTensor &aL0Tensor, const LocalTensor &bL0Tensor, + const LocalTensor &cL0Tensor, const MmadParams &mmadParams, + const uint32_t kL1StepSize, const uint32_t kL1Size = 0, + const LocalTensor &aScaleL1 = {}, const LocalTensor &bscaleL1 = {}) +{ + WaitFlag(L0A_EVENT0 + (bufParam.aL0BufIter & 1u)); + LocalTensor aL0 = aL0Tensor[(bufParam.aL0BufIter & 1u) * (L0A_PP_SIZE / sizeof(T))]; + if constexpr (std::is_same::value && std::is_same::value) { + LoadDataL1ToL0Mxfp8(aL0, aL1, aScaleL1, mmadParams.m, mmadParams.k, kL1Size); + } else { + LoadDataL1ToL0(aL0, aL1, mmadParams.m, mmadParams.k); + } + SetFlag(L0A_EVENT0 + (bufParam.aL0BufIter & 1u)); + WaitFlag(L0A_EVENT0 + (bufParam.aL0BufIter & 1u)); + WaitFlag(L0B_EVENT0 + (bufParam.bL0BufIter & 1u)); + + LocalTensor bL0 = bL0Tensor[(bufParam.bL0BufIter & 1u) * (L0B_PP_SIZE / sizeof(T))]; + if constexpr (std::is_same::value || std::is_same::value || + (std::is_same::value && std::is_same::value)) { + LoadDataL1ToL0B(bL0, bL1, bscaleL1, mmadParams.k, mmadParams.n, kL1StepSize); + } else if constexpr (std::is_same::value && std::is_same::value) { + LoadDataL1ToL0B(bL0, bL1, bscaleL1, mmadParams.k, mmadParams.n, kL1StepSize, kL1Size); + } else { + LoadDataL1ToL0(bL0, bL1, mmadParams.k, mmadParams.n, kL1StepSize); + } + + SetFlag(L0B_EVENT0 + (bufParam.bL0BufIter & 1u)); + WaitFlag(L0B_EVENT0 + (bufParam.bL0BufIter & 1u)); + if constexpr (std::is_same::value && std::is_same::value) { + LocalTensor aL0Tmp = aL0.template ReinterpretCast(); + LocalTensor bL0Tmp = bL0.template ReinterpretCast(); + Mmad(cL0Tensor, aL0Tmp, bL0Tmp, mmadParams); + } else { + Mmad(cL0Tensor, aL0, bL0, mmadParams); + } + PipeBarrier(); + SetFlag(L0B_EVENT0 + (bufParam.bL0BufIter & 1u)); + bufParam.bL0BufIter++; + SetFlag(L0A_EVENT0 + (bufParam.aL0BufIter & 1u)); + bufParam.aL0BufIter++; +} + +template +__aicore__ inline void GetTensorC(const GlobalTensor &tensorCGm, const LocalTensor &cL0, const uint32_t mSize, + const uint32_t nSize, const uint32_t srcStride, const uint32_t dstStride, + MMBufParams &bufParam) +{ + FixpipeParamsV220 fixParams; + fixParams.nSize = nSize; // 实现切片大小 + fixParams.mSize = mSize; // msdIterNum * gSize; // 有效数据不足16行,只需要输出部分行即可 + fixParams.srcStride = srcStride; // ((fixParams.mSize + 15) / 16) * 16 + fixParams.dstStride = dstStride; + fixParams.ndNum = 1; + if constexpr (enUnitFlag) { + fixParams.unitFlag = UNIT_FLAG_SET; + } else { + fixParams.unitFlag = UNIT_FLAG_DISABLE; + } + if constexpr (std::is_same::value && std::is_same::value) { + fixParams.quantPre = QuantMode_t::F322BF16; + } + if constexpr (!enUnitFlag) { + SetFlag(L0C_EVENT0 + (bufParam.cL0BufIter & 1u)); + WaitFlag(L0C_EVENT0 + (bufParam.cL0BufIter & 1u)); + } + Fixpipe(tensorCGm, cL0, fixParams); +} + +/** + * @brief 将AB矩阵从GM加载到L1缓存以进行矩阵乘法运算 + * @tparam T 张量的数据类型 + * @tparam hasL1ALoaded A矩阵是否已经加载到L1中 + * @param tensorAGm GM中的A矩阵 + * @param tensorBGm GM中的B矩阵 + * @param kL1 当前K L1循环迭代次数 + * @param kL1Loops K L1循环的总次数 + * @param para 包含维度和步长信息的矩阵乘法参数 + * @param nL1Offset L1缓冲区中的N维度偏移 + * @param nL1Size L1缓冲区中的N维度大小 + * @param kOffesetUnit K维度寻址的偏移单位 + * @param bufParam L1缓冲区管理的参数对象 + */ +template +__aicore__ inline void LoadL1AB(const GlobalTensor &tensorAGm, const GlobalTensor &tensorBGm, uint32_t kL1, + uint32_t kL1Loops, const MMParams ¶, uint32_t nL1Offset, uint32_t nL1Size, + uint32_t kOffesetUnit, MMBufParams &bufParam) +{ + if (kL1 == 0) { + if constexpr (!hasL1ALoaded) { + LoadL1A(tensorAGm[kL1 * para.kL1StepSize], para.m, para.kL1StepSize, para.orgKa, bufParam); + } + if constexpr (bLoadFormat == DataFormat::NZ) { + LoadL1B(tensorBGm[para.k * nL1Offset + kL1 * kOffesetUnit], nL1Size, para.n, para.kL1StepSize, para.k, + bufParam); + } else { + LoadL1B(tensorBGm[kL1 * para.kL1StepSize * para.n + nL1Offset], nL1Size, para.n, + para.kL1StepSize, para.kL1StepSize, bufParam); + } + } + + if (kL1 + 1 < kL1Loops) { + if constexpr (!hasL1ALoaded) { + LoadL1A(tensorAGm[(kL1 + 1) * para.kL1StepSize], para.m, para.kL1StepSize, para.orgKa, bufParam); + } + if constexpr (bLoadFormat == DataFormat::NZ) { + LoadL1B(tensorBGm[para.k * nL1Offset + (kL1 + 1) * kOffesetUnit], nL1Size, para.n, + para.kL1StepSize, para.k, bufParam); + } else { + LoadL1B(tensorBGm[(kL1 + 1) * para.kL1StepSize * para.n + nL1Offset], nL1Size, + para.n, para.kL1StepSize, para.kL1StepSize, bufParam); + } + } +} + +/** + * @brief 将AB矩阵和用于伪量化的Scale矩阵从GM加载到L1缓存以进行矩阵乘法运算 + * @tparam T 张量的数据类型 + * @tparam S Scale矩阵的数据类型 + * @tparam hasL1ALoaded A矩阵是否已经加载到L1中 + * @tparam scaleSrcPadFlag Scale矩阵是否需要对齐搬运,在mm3中AScale为True + * @param tensorAGm GM中的A矩阵 + * @param tensorBGm GM中的B矩阵 + * @param tensorAScaleGm GM中的AScale矩阵 + * @param tensorBScaleGm GM中的BScale矩阵 + * @param kL1 当前K L1循环迭代次数 + * @param kL1Loops K L1循环的总次数 + * @param para 包含维度和步长信息的矩阵乘法参数 + * @param nL1Offset L1缓冲区中的N维度偏移 + * @param nL1Size L1缓冲区中的N维度大小 + * @param kOffesetUnit K维度寻址的偏移单位 + * @param bufParam L1缓冲区管理的参数对象 + */ +template +__aicore__ inline void LoadL1ABAndScale(const GlobalTensor &tensorAGm, const GlobalTensor &tensorBGm, + const GlobalTensor &tensorAScaleGm, const GlobalTensor &tensorBScaleGm, + uint32_t kL1, uint32_t kL1Loops, const MMParams ¶, uint32_t nL1Offset, + uint32_t nL1Size, uint32_t kOffesetUnit, MMBufParams &bufParam) +{ + uint64_t offsetL1B = L1_B_SIZE / 2 / sizeof(T); // // 2表示scale起始地址固定从L1B上ping的64k开始 + if (kL1 == 0) { + if constexpr (!hasL1ALoaded) { + LoadL1AAndScale(tensorAGm, tensorAScaleGm, para.m, para.kL1StepSize, + para.k, para.kScale, offsetL1B, bufParam); + } + uint64_t scaleOffsetL1B = offsetL1B + para.kScale * Align(para.m, BLOCK_CUBE_SIZE); + LoadL1BAndScale(tensorBGm[para.k * nL1Offset], tensorBScaleGm[para.kScale * nL1Offset], nL1Size, + para.kL1StepSize, para.k, para.kScale, scaleOffsetL1B, bufParam); + } + + if (kL1 + 1 < kL1Loops) { + if constexpr (!hasL1ALoaded) { + LoadL1A(tensorAGm[(kL1 + 1) * para.kL1StepSize], para.m, para.kL1StepSize, para.k, bufParam); + } + LoadL1B(tensorBGm[para.k * nL1Offset + (kL1 + 1) * kOffesetUnit], nL1Size, para.n, + para.kL1StepSize, para.k, bufParam); + } +} + +/** + * @brief B矩阵数据加载到L1缓冲区的辅助函数 + * @tparam T B张量的数据类型 + * @tparam isContinuousCopy 是否使用连续复制模式 + * @param bL1 B矩阵L1张量 + * @param tensorBGm B矩阵GM张量 + * @param kL1StepSize L1中K维度的步长大小 + * @param subNL1SplitSize L1中N维度的分割大小 + * @param para MMParams参数 + * @param bufParam 缓冲区管理参数 + * @param nL1 当前N L1迭代 + * @param kL1 当前K L1迭代 + * @param kOffesetUnit K维度偏移单位 + */ +template +__aicore__ inline void LoadL1BGroupCompute(LocalTensor &bL1, const GlobalTensor &tensorBGm, + uint32_t subNL1SplitSize, const MMParams ¶, int64_t nL1, uint32_t kL1, + uint32_t kOffesetUnit) +{ + int64_t tensorBGmOffset = para.k * nL1 * para.baseN + kL1 * kOffesetUnit; + auto tensorBGmForL1 = tensorBGm[tensorBGmOffset]; + if constexpr (isContinuousCopy) { + CopyNZGmToL1(bL1, tensorBGmForL1, para.kL1StepSize, subNL1SplitSize, para.k); + } else { + // 每次搬运两块K*64拼接成K*128 + CopyNZGmToL1(bL1, tensorBGmForL1, para.kL1StepSize, subNL1SplitSize >> 1, para.k); + LocalTensor b2L1 = bL1[(para.kL1StepSize * subNL1SplitSize) >> 1]; + auto tensorB2GmForL1 = tensorBGmForL1[para.k * DIM_HEAD_SIZE_QCQR]; + CopyNZGmToL1(b2L1, tensorB2GmForL1, para.kL1StepSize, subNL1SplitSize >> 1, para.k); + } +} + +/** + * @brief 执行矩阵乘法的K L0循环计算 + * @tparam T A和B矩阵的数据类型 + * @tparam O_L0C L0C缓存中的数据类型 + * @tparam enUnitFlag 是否启用unitFlag功能 + * @param cL0 L0缓存中的C矩阵张量 + * @param aL1 L1缓存中的A矩阵张量 + * @param bL1 L1缓存中的B矩阵张量 + * @param localTensors 矩阵乘法所需的本地张量对象 + * @param bufParam L1缓存管理的缓冲区参数对象 + * @param kL1 当前K L1循环迭代次数 + * @param kL1Loops K L1循环的总次数 + * @param stepK K L0步数 + * @param aOffsetUnit A矩阵L0寻址的偏移单位 + * @param bOffsetUnit B矩阵L0寻址的偏移单位 + * @param mmadParams 矩阵乘法参数 + * @param aOffset A矩阵计算的起始偏移 + * @param bOffset B矩阵计算的起始偏移 + */ +template +__aicore__ inline void MatmulL1(const LocalTensor &cL0, const LocalTensor &aL1, const LocalTensor &bL1, + const mmLocalTensors &localTensors, MMBufParams &bufParam, + const MMParams ¶, uint32_t kL1, uint32_t kL1Loops, MmadParams &mmadParams, + int64_t &aOffset) +{ + uint32_t mSize = Align(para.m, BLOCK_CUBE_SIZE); + uint64_t weightSizeL1B = L1_B_SIZE / 2 / sizeof(T); // 2表示scale起始地址固定从L1B上ping的64k开始 + uint64_t scaleASize = mSize * para.kScale; // BS*224 + uint32_t bOffset = 0; + uint32_t aOffsetUnit = mSize * para.baseK; + uint32_t bOffsetUnit = GetC0Num() * para.baseK; + + uint32_t aScaleOffset = kL1 * para.kScale / kL1Loops * BLOCK_CUBE_SIZE; + uint32_t bScaleOffset = kL1 * para.kScale / kL1Loops * BLOCK_CUBE_SIZE; + uint32_t aScaleOffsetUnit = para.baseK / BYTE_BLOCK * BLOCK_CUBE_SIZE; + uint32_t bScaleOffsetUnit = para.baseK / BYTE_BLOCK * BLOCK_CUBE_SIZE; + const LocalTensor scaleALocalTensor = localTensors.bL1Tensor[weightSizeL1B]; + const LocalTensor scaleBLocalTensor = scaleALocalTensor[scaleASize]; + + for (int64_t kL0Loops = 0; kL0Loops < para.stepK; kL0Loops++) { + mmadParams.cmatrixInitVal = ((kL1 == 0) && (kL0Loops == 0)); + if constexpr (enUnitFlag) { + mmadParams.unitFlag = + (kL1 == kL1Loops - 1) && (kL0Loops == para.stepK - 1) ? UNIT_FLAG_SET : UNIT_FLAG_CHECK; + } + MatmulL0(bufParam, aL1[aOffset], bL1[bOffset], localTensors.aL0Tensor, localTensors.bL0Tensor, cL0, + mmadParams, para.kL1StepSize, para.k, scaleALocalTensor[aScaleOffset], + scaleBLocalTensor[bScaleOffset]); + aOffset += aOffsetUnit; // 16(BS)*256=4096 + bOffset += bOffsetUnit; // 32(Get)*256=8192 + if constexpr (std::is_same::value && std::is_same::value) { + aScaleOffset += aScaleOffsetUnit; // 16(BS)*baseK/32=16(BS)*8=128 + bScaleOffset += bScaleOffsetUnit; // baseK/32*baseN=8*64 + } + } +} + +/** + * @brief MatmulSplitK 通过切分K轴,进行L1的数据管理和矩阵乘运算; 用于mmCq, mmCKvKr和mmQcQr。 + * @param tensorAGm A矩阵在GM的位置 + * @param tensorBGm B矩阵在GM的位置 + * @param tensorCGm C矩阵在GM的位置 + * @param para 表示matmul形状信息的结构体参数 + * @param bufParam 管理L1 buffer地址和同步计数的结构体参数 + * @param nL1Offset 本次计算中,L1 buffer内B矩阵在n轴上的偏移 + * @param nL1Size 本次计算中,L1 buffer内B矩阵在n轴上的计算量 + * @param tensorAScaleGm AScale矩阵在GM的位置 + * @param tensorBScaleGm BScale矩阵在GM的位置 + */ +template +__aicore__ inline void +MatmulSplitK(const GlobalTensor &tensorCGm, const GlobalTensor &tensorAGm, const GlobalTensor &tensorBGm, + const MMParams ¶, MMBufParams &bufParam, const uint32_t nL1Offset, const uint32_t nL1Size, + const GlobalTensor &tensorAScaleGm = {}, const GlobalTensor &tensorBScaleGm = {}) +{ + using O_L0C = typename std::conditional::value, int32_t, float>::type; + + // 全局L1管理 + uint32_t kOffesetUnit = para.kL1StepSize * GetC0Num(); + uint32_t mSize = Align(para.m, BLOCK_CUBE_SIZE); + uint32_t kL1Loops = CeilDivT(para.k, para.kL1StepSize); + + mmLocalTensors localTensors; + localTensors.Init(bufParam); + + LocalTensor aL1, bL1; + + if constexpr (hasL1ALoaded) { + aL1 = localTensors.aL1Tensor[(bufParam.aL1BufIter & 1u) * L1_A_SIZE / sizeof(T)]; + } + + WaitFlag(L0C_EVENT0 + (bufParam.cL0BufIter & 1u)); + LocalTensor cL0 = localTensors.cL0Tensor[(bufParam.cL0BufIter & 1u) * (L0C_PP_SIZE / sizeof(O_L0C))]; + + MmadParams mmadParams = MmadParams(mSize, nL1Size, para.baseK, UNIT_FLAG_DISABLE, false, true); + + int64_t aOffset = 0; + if constexpr (std::is_same::value && std::is_same::value) { + WaitFlag(SCALE_EVENT); + } + + for (uint32_t kL1 = 0; kL1 < kL1Loops; kL1++) { + // Load data from global memory to L1 cache + if constexpr (std::is_same::value && std::is_same::value) { + LoadL1ABAndScale(tensorAGm, tensorBGm, tensorAScaleGm, tensorBScaleGm, + kL1, kL1Loops, para, nL1Offset, nL1Size, kOffesetUnit, + bufParam); + } else { + LoadL1AB(tensorAGm, tensorBGm, kL1, kL1Loops, para, nL1Offset, nL1Size, + kOffesetUnit, bufParam); + } + + WaitFlag(B_EVENT0 + (bufParam.bL1BufIter & 1u)); + bL1 = localTensors.bL1Tensor[(bufParam.bL1BufIter & 1u) * L1_B_SIZE / sizeof(T)]; + if constexpr (!hasL1ALoaded) { + WaitFlag(A_EVENT0 + (bufParam.aL1BufIter & 1u)); + aL1 = localTensors.aL1Tensor[(bufParam.aL1BufIter & 1u) * L1_A_SIZE / sizeof(T)]; + aOffset = 0; + } + + // Perform core K L0 loops computation + MatmulL1(cL0, aL1, bL1, localTensors, bufParam, para, kL1, kL1Loops, mmadParams, + aOffset); + + SetFlag(B_EVENT0 + (bufParam.bL1BufIter & 1u)); + bufParam.bL1BufIter++; + if constexpr (!hasL1ALoaded) { + SetFlag(A_EVENT0 + (bufParam.aL1BufIter & 1u)); + bufParam.aL1BufIter++; + } + } + if constexpr (std::is_same::value && std::is_same::value) { + SetFlag(SCALE_EVENT); + } + GetTensorC(tensorCGm[nL1Offset], cL0, para.m, nL1Size, mSize, para.orgKc, bufParam); + SetFlag(L0C_EVENT0 + (bufParam.cL0BufIter & 1u)); + bufParam.cL0BufIter++; +} + +/** + * @brief MatmulFullLoad A,B矩阵能够全载L1时的矩阵运算,用于mmQn; + * @param tensorAGm A矩阵在GM的位置 + * @param tensorBGm B矩阵在GM的位置 + * @param tensorCGm C矩阵在GM的位置 + * @param para 表示matmul形状信息的结构体参数 + * @param bufParam 管理L1 buffer地址和同步计数的结构体参数 + */ +template +__aicore__ inline void MatmulFullLoad(const GlobalTensor &tensorCGm, const GlobalTensor &tensorAGm, + const GlobalTensor &tensorBGm, const MMParams ¶, MMBufParams &bufParam) +{ + using O_L0C = typename std::conditional::value, int32_t, float>::type; + uint32_t mSize = Align(para.m, BLOCK_CUBE_SIZE); + constexpr uint32_t aL1Offset = L1_A_SIZE / sizeof(T); + constexpr uint32_t bL1Offset = L1_B_SIZE / sizeof(T); + + mmLocalTensors localTensors; + localTensors.Init(bufParam); + + LoadL1A(tensorAGm, para.m, para.k, para.orgKa, bufParam); + WaitFlag(A_EVENT0 + (bufParam.aL1BufIter & 1u)); + LocalTensor aL1 = localTensors.aL1Tensor[(bufParam.aL1BufIter & 1u) * aL1Offset]; + + if constexpr (!hasL1BLoaded) { + LoadL1B(tensorBGm, para.n, para.n, para.k, para.k, bufParam); + WaitFlag(B_EVENT0 + (bufParam.bL1BufIter & 1u)); + } + LocalTensor bL1 = localTensors.bL1Tensor[(bufParam.bL1BufIter & 1u) * bL1Offset]; + + uint32_t nSplitSize = 128; + uint32_t nSplitSizeAct = nSplitSize; + uint32_t nloops = CeilDivT(para.n, nSplitSize); + + MmadParams mmadParams = MmadParams(mSize, nSplitSizeAct, para.k, UNIT_FLAG_DISABLE, false, true); + + for (uint32_t n = 0; n < nloops; n++) { + if (n == nloops - 1) { + nSplitSizeAct = para.n - (nloops - 1) * nSplitSize; + mmadParams.n = nSplitSizeAct; + } + // Perform matrix multiplication computation for current N-dimension split + if constexpr (enUnitFlag) { + mmadParams.unitFlag = UNIT_FLAG_SET; + } + WaitFlag(L0C_EVENT0 + (bufParam.cL0BufIter & 1u)); + LocalTensor cL0 = localTensors.cL0Tensor[(bufParam.cL0BufIter & 1u) * (L0C_PP_SIZE / sizeof(O_L0C))]; + MatmulL0(bufParam, aL1, bL1[para.k * nSplitSize * n], localTensors.aL0Tensor, + localTensors.bL0Tensor, cL0, mmadParams, para.kL1StepSize); + GetTensorC(tensorCGm[n * nSplitSize], cL0, para.m, nSplitSizeAct, mSize, para.orgKc, + bufParam); + SetFlag(L0C_EVENT0 + (bufParam.cL0BufIter & 1u)); + bufParam.cL0BufIter++; + } + + SetFlag(B_EVENT0 + (bufParam.bL1BufIter & 1u)); + bufParam.bL1BufIter++; + SetFlag(A_EVENT0 + (bufParam.aL1BufIter & 1u)); + bufParam.aL1BufIter++; +} + +/** + * @brief MatmulGroupComputeAFullLoad + * A矩阵能够全载,通过切分K轴,进行L1的数据管理和矩阵乘运算,仅适用于算力分组场景下的mmQc和mmQr; + * @param tensorAGm A矩阵在GM的位置 + * @param tensorBGm B矩阵在GM的位置 + * @param tensorCGm C矩阵在GM的位置 + * @param para 表示matmul形状信息的结构体参数 + * @param bufParam 管理L1 buffer地址和同步计数的结构体参数 + */ +template +__aicore__ inline void MatmulGroupComputeAFullLoad(const GlobalTensor &tensorCGm, const GlobalTensor &tensorAGm, + const GlobalTensor &tensorBGm, const MMParams ¶, + MMBufParams &bufParam) +{ + using O_L0C = typename std::conditional::value, int32_t, float>::type; + // 全局L1管理 + uint32_t kOffesetUnit = para.kL1StepSize * GetC0Num(); + uint32_t mSize = Align(para.m, BLOCK_CUBE_SIZE); + uint32_t kL1Loops = CeilDivT(para.k, para.kL1StepSize); + uint32_t nL1loops = CeilDivT(para.n, para.baseN); + + mmLocalTensors localTensors; + localTensors.Init(bufParam); + + WaitFlag(A_EVENT0 + (bufParam.aL1BufIter & 1u)); + LocalTensor aL1 = localTensors.aL1Tensor[(bufParam.aL1BufIter & 1u) * (L1_A_SIZE / sizeof(T))]; + CopyNDGmToL1(aL1, tensorAGm, para.m, para.k, para.k); + SetFlag(A_EVENT0 + (bufParam.aL1BufIter & 1u)); + WaitFlag(A_EVENT0 + (bufParam.aL1BufIter & 1u)); + + uint32_t subNL1SplitSize = para.baseN; + + MmadParams mmadParams = MmadParams(mSize, subNL1SplitSize, para.baseK, UNIT_FLAG_DISABLE, false, true); + + for (int64_t nL1 = 0; nL1 < nL1loops; nL1++) { + if (nL1 == nL1loops - 1) { + subNL1SplitSize = para.n - (nL1loops - 1) * para.baseN; + mmadParams.n = subNL1SplitSize; + } + + WaitFlag(L0C_EVENT0 + (bufParam.cL0BufIter & 1u)); + LocalTensor cL0 = localTensors.cL0Tensor[(bufParam.cL0BufIter & 1u) * (L0C_PP_SIZE / sizeof(O_L0C))]; + + int64_t aOffset = 0; + for (uint32_t kL1 = 0; kL1 < kL1Loops; kL1++) { + WaitFlag(B_EVENT0 + (bufParam.bL1BufIter & 1u)); + LocalTensor bL1 = localTensors.bL1Tensor[(bufParam.bL1BufIter & 1u) * (L1_B_SIZE / sizeof(T))]; + LoadL1BGroupCompute(bL1, tensorBGm, subNL1SplitSize, para, nL1, kL1, kOffesetUnit); + SetFlag(B_EVENT0 + (bufParam.bL1BufIter & 1u)); + WaitFlag(B_EVENT0 + (bufParam.bL1BufIter & 1u)); + + // Perform core K L0 loops computation using shared function + MatmulL1(cL0, aL1, bL1, localTensors, bufParam, para, kL1, kL1Loops, mmadParams, + aOffset); + SetFlag(B_EVENT0 + (bufParam.bL1BufIter & 1u)); + bufParam.bL1BufIter++; + } + GetTensorC(tensorCGm[nL1 * para.baseN], cL0, para.m, subNL1SplitSize, mSize, + para.orgKc, bufParam); + SetFlag(L0C_EVENT0 + (bufParam.cL0BufIter & 1u)); + bufParam.cL0BufIter++; + } + SetFlag(A_EVENT0 + (bufParam.aL1BufIter & 1u)); + bufParam.aL1BufIter++; +} + +/** + * @brief K-outer: 加载B矩阵到L1B(支持双缓冲slot选择) + * bScaleL1Offset: B-scale在slot内的bf16偏移 (需放在A-scale之后避免冲突) + */ +template +__aicore__ inline void LoadL1BOuter(const GlobalTensor &tensorBGm, const GlobalTensor &tensorBScaleGm, + uint32_t nL1Size, uint32_t kL1StepSize, uint32_t kInput, uint32_t kScale, + uint32_t bEventIdx, uint32_t bScaleL1Offset, MMBufParams &bufParam) +{ + LocalTensor bL1Tensor; + bL1Tensor.SetAddr(bufParam.bL1BufAddr); + uint32_t bL1SlotSize = L1_B_SIZE / sizeof(T); + + WaitFlag(B_EVENT0 + bEventIdx); + LocalTensor bL1 = bL1Tensor[bEventIdx * bL1SlotSize]; + + if constexpr (std::is_same::value && std::is_same::value) { + CopyNZGmToL1(bL1, tensorBGm, kL1StepSize, nL1Size, kInput); + + LocalTensor bL1ScaleTensor = bL1.template ReinterpretCast(); + GlobalTensor tensorScaleGmCast; + tensorScaleGmCast.SetGlobalBuffer(((__gm__ bfloat16_t *)(tensorBScaleGm.GetPhyAddr()))); + Dn2NzParams dn2Nzparam; + dn2Nzparam.dnNum = 1; + uint32_t nValueB = CeilDivT(kScale, static_cast(FP8_TWO)); + dn2Nzparam.nValue = nValueB; + dn2Nzparam.dValue = nL1Size; + dn2Nzparam.srcDnMatrixStride = 0; + dn2Nzparam.srcDValue = nValueB; // B-scale无padding, 保持原值 + dn2Nzparam.dstNzC0Stride = nValueB; // = nValue, NZ输出布局不变 + dn2Nzparam.dstNzNStride = 1; + dn2Nzparam.dstNzMatrixStride = nValueB; + DataCopy(bL1ScaleTensor[bScaleL1Offset], tensorScaleGmCast, dn2Nzparam); + } else if constexpr (loadFormat == DataFormat::NZ) { + CopyNZGmToL1(bL1, tensorBGm, kL1StepSize, nL1Size, kInput); + } else { + CopyNDGmToL1(bL1, tensorBGm, kInput, nL1Size, nL1Size); + } + SetFlag(B_EVENT0 + bEventIdx); +} + + +/** + * @brief K-outer: 主Matmul函数。A数据每kL1加载到L1A,A-scale复用原LoadL1AAndScale放L1B一次加载全量。 + * B每nb双缓冲加载,B-scale紧随B数据之后。FixPipe使用AtomicAdd累加到GM。 + */ +template +__aicore__ inline void MatmulSplitMKOuter(const GlobalTensor &tensorCGm, const GlobalTensor &tensorAGm, + const GlobalTensor &tensorBGm, const MMParams ¶, MMBufParams &bufParam, + const GlobalTensor &tensorAScaleGm, const GlobalTensor &tensorBScaleGm) +{ + using O_L0C = typename std::conditional::value, int32_t, float>::type; + + uint32_t nBlocks = CeilDivT(para.n, para.baseN); + uint32_t kL1Loops = CeilDivT(para.k, para.kL1StepSize); + uint32_t kOffesetUnit = para.kL1StepSize * GetC0Num(); + uint32_t mSizeAligned = Align(para.m, BLOCK_CUBE_SIZE); + uint32_t aOffsetUnit = mSizeAligned * para.baseK; + uint32_t bOffsetUnit = GetC0Num() * para.baseK; + uint32_t bL1SlotSize = L1_B_SIZE / sizeof(T); + uint32_t scaleOffUnit = para.baseK / static_cast(BYTE_BLOCK) * BLOCK_CUBE_SIZE; + + mmLocalTensors localTensors; + localTensors.Init(bufParam); + + // ──── 一次加载全量A-scale到L1B, 复用原LoadL1AAndScale的DN2NZ参数 ──── + LocalTensor aScaleL1Base; + uint32_t bScaleBf16Off = 0; + uint32_t bScaleL1BaseOff = 0; + if constexpr (std::is_same::value && std::is_same::value) { + uint64_t scaleAOff = L1_B_SIZE / 2 / sizeof(T); + LocalTensor bL1Bf16; + bL1Bf16.SetAddr(bufParam.bL1BufAddr); + GlobalTensor scaleGmCast; + scaleGmCast.SetGlobalBuffer(((__gm__ bfloat16_t *)(tensorAScaleGm.GetPhyAddr()))); + Dn2NzParams dn2nz; + dn2nz.dnNum = 1; + uint32_t nValue = CeilDivT(para.kScale, static_cast(FP8_TWO)); + dn2nz.nValue = nValue; + dn2nz.dValue = para.m; + dn2nz.srcDnMatrixStride = 0; + dn2nz.srcDValue = nValue; + dn2nz.dstNzC0Stride = nValue; + dn2nz.dstNzNStride = 1; + dn2nz.dstNzMatrixStride = nValue; + DataCopy(bL1Bf16[scaleAOff / FP8_TWO], scaleGmCast, dn2nz); + aScaleL1Base = localTensors.bL1Tensor[scaleAOff]; + + uint32_t aScaleTotalSize = mSizeAligned * para.kScale; + bScaleL1BaseOff = static_cast(scaleAOff) + aScaleTotalSize; + bScaleBf16Off = bScaleL1BaseOff / FP8_TWO; + } + + for (uint32_t kL1 = 0; kL1 < kL1Loops; ++kL1) { + // ── 加载A数据到L1A slot0 ── + LocalTensor aL1; + WaitFlag(A_EVENT0); + aL1 = localTensors.aL1Tensor; + CopyNDGmToL1(aL1, tensorAGm[kL1 * para.kL1StepSize], para.m, para.kL1StepSize, para.k); + SetFlag(A_EVENT0); + WaitFlag(A_EVENT0); + + // A-scale起始偏移: 原MatmulL1公式 + uint32_t kL1ScaleOff = kL1 * para.kScale / kL1Loops * BLOCK_CUBE_SIZE; + + // ── 预加载 B[0] ── + uint32_t subN0 = (nBlocks == 1) ? para.n : para.baseN; + int64_t bGmOff0 = static_cast(kL1) * kOffesetUnit; + if constexpr (std::is_same::value && std::is_same::value) { + int64_t bScaleOff0 = static_cast(0); // 全量K-step从0开始 + LoadL1BOuter(tensorBGm[bGmOff0], tensorBScaleGm[bScaleOff0], subN0, para.kL1StepSize, para.k, + para.kScale, 0, bScaleBf16Off, bufParam); + } else { + LoadL1BOuter(tensorBGm[bGmOff0], tensorBScaleGm, subN0, para.kL1StepSize, para.k, 0, 0, bScaleBf16Off, + bufParam); + } + + // ── N分块: B双缓冲ping-pong ── + for (uint32_t nb = 0; nb < nBlocks; ++nb) { + uint32_t subN = (nb == nBlocks - 1) ? (para.n - nb * para.baseN) : para.baseN; + uint32_t curSlot = nb & 1u; + + WaitFlag(B_EVENT0 + curSlot); + LocalTensor bL1 = localTensors.bL1Tensor[curSlot * bL1SlotSize]; + LocalTensor bScaleL1; + if constexpr (std::is_same::value && std::is_same::value) { + bScaleL1 = localTensors.bL1Tensor[curSlot * bL1SlotSize + bScaleL1BaseOff]; + } + + // 预取 B[nb+1] + if (nb + 1 < nBlocks) { + uint32_t nextSlot = (nb + 1) & 1u; + uint32_t subNx = ((nb + 1) == nBlocks - 1) ? (para.n - (nb + 1) * para.baseN) : para.baseN; + int64_t bOffNext = + static_cast(para.k) * (nb + 1) * para.baseN + static_cast(kL1) * kOffesetUnit; + if constexpr (std::is_same::value && std::is_same::value) { + int64_t bScOffNext = static_cast((nb + 1) * para.baseN * para.kScale); + LoadL1BOuter(tensorBGm[bOffNext], tensorBScaleGm[bScOffNext], subNx, para.kL1StepSize, para.k, + para.kScale, nextSlot, bScaleBf16Off, bufParam); + } else { + LoadL1BOuter(tensorBGm[bOffNext], tensorBScaleGm, subNx, para.kL1StepSize, para.k, 0, + nextSlot, bScaleBf16Off, bufParam); + } + } + + // ── L0C ── + WaitFlag(L0C_EVENT0 + (bufParam.cL0BufIter & 1u)); + LocalTensor cL0 = localTensors.cL0Tensor[(bufParam.cL0BufIter & 1u) * (L0C_PP_SIZE / sizeof(O_L0C))]; + + // ── stepK L0循环 ── + MmadParams mmadParams = MmadParams(mSizeAligned, subN, para.baseK, UNIT_FLAG_DISABLE, false, true); + int64_t aOffset = 0; + uint32_t bOffset = 0; + uint32_t aScaleOff = kL1ScaleOff; + uint32_t bScaleOff = kL1ScaleOff; + for (int64_t sk = 0; sk < para.stepK; sk++) { + mmadParams.cmatrixInitVal = (sk == 0); + MatmulL0(bufParam, aL1[aOffset], bL1[bOffset], localTensors.aL0Tensor, + localTensors.bL0Tensor, cL0, mmadParams, para.kL1StepSize, para.k, + aScaleL1Base[aScaleOff], bScaleL1[bScaleOff]); + aOffset += aOffsetUnit; + bOffset += bOffsetUnit; + aScaleOff += scaleOffUnit; + bScaleOff += scaleOffUnit; + } + + // ── FixPipe AtomicAdd ── + SetFlag(L0C_EVENT0 + (bufParam.cL0BufIter & 1u)); + WaitFlag(L0C_EVENT0 + (bufParam.cL0BufIter & 1u)); + FixpipeParamsV220 fixParams; + fixParams.nSize = subN; + fixParams.mSize = para.m; + fixParams.srcStride = mSizeAligned; + fixParams.dstStride = para.orgKc; + fixParams.ndNum = 1; + fixParams.unitFlag = UNIT_FLAG_DISABLE; + SetAtomicAdd(); + Fixpipe(tensorCGm[nb * para.baseN], cL0, fixParams); + SetAtomicNone(); + + SetFlag(L0C_EVENT0 + (bufParam.cL0BufIter & 1u)); + bufParam.cL0BufIter++; + + SetFlag(B_EVENT0 + curSlot); + } + SetFlag(A_EVENT0); + } +} + +} // namespace MlaProlog + + +#endif diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_rms_norm.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_rms_norm.h new file mode 100644 index 000000000000..aaeab1838cb7 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_rms_norm.h @@ -0,0 +1,159 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file service_rms_norm.h + * \brief + */ + +#ifndef SERVICE_RMS_NORM_H +#define SERVICE_RMS_NORM_H + +#include "mla_prolog_comm.h" +#include "mla_prolog_vector_comm.h" + +#include "vf/vf_rms_norm.h" +#include "vf/vf_dynamic_quant.h" + +namespace MlaProlog { + +/** + * @brief RmsNormNormal rmsNorm流程 + * @param outputLocal 输出tensor,[B, S1, H] + * @param inputGm 输入tensor,[B, S1, H],dtype支持bf16,fp32,int32 + * @param gammaLocal 系数gamma [H],dtype只支持bf16 + * @param dequantScaleWDqLocal 权重w反量化系数 + * @param dequantScaleXLocal x反量化系数 + * @param shareTmpUb 临时buffer 所需空间[2 * cnt * sizeof(float) + ALIGN_BLOCK_SIZE] + ALIGN_BLOCK_SIZE = 32Bytes, cnt = row * col + * @param rmsNormParams rms所需系数,包括 + reciprocal rmsnorm系数reciprocal + epsilon rmsnorm系数epsilon + row 处理的行数;预留参数,当前仅支持单个batch的处理,row为1,对应S1 + col 列数,对应H + */ +template +__aicore__ inline void RmsNormNormal(const LocalTensor &outputLocal, const GlobalTensor &inputGm, + const LocalTensor &gammaLocal, + const LocalTensor &dequantScaleWDqLocal, + const LocalTensor &dequantScaleXLocal, + const LocalTensor &shareTmpUb, RmsNormParam &rmsNormParams) +{ + int64_t cnt = rmsNormParams.row * rmsNormParams.col; + LocalTensor xFp32Local = shareTmpUb.ReinterpretCast(); + + // load input [1, col] + DataCopyExtParams copyParams{1, static_cast(rmsNormParams.col * sizeof(T)), 0, 0, 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + + if constexpr (std::is_same::value && std::is_same::value) { + DataCopyPad(xFp32Local, inputGm, copyParams, padParams); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + } else if constexpr (std::is_same::value && std::is_same::value) { + DataCopyPad(xFp32Local, inputGm, copyParams, padParams); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + Rectangle rectangleParams{ + static_cast(rmsNormParams.row), static_cast(rmsNormParams.col), + static_cast(rmsNormParams.col) // columnStride + }; + Dequant(xFp32Local, xFp32Local, dequantScaleWDqLocal, dequantScaleXLocal, rectangleParams); + PipeBarrier(); + } else if constexpr (std::is_same::value) { + LocalTensor xInt32Local = shareTmpUb.ReinterpretCast(); + DataCopyPad(xInt32Local, inputGm, copyParams, padParams); + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + Rectangle rectangleParams{ + static_cast(rmsNormParams.row), static_cast(rmsNormParams.col), + static_cast(rmsNormParams.col) // columnStride + }; + Dequant(xFp32Local, xInt32Local, dequantScaleWDqLocal, dequantScaleXLocal, rectangleParams); + PipeBarrier(); + } else { + LocalTensor inputLocal = xFp32Local[rmsNormParams.col].template ReinterpretCast(); + DataCopyPad(inputLocal, inputGm, copyParams, padParams); + + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + // Cast input to fp32 [1, col] + Cast(xFp32Local, inputLocal, RoundMode::CAST_NONE, cnt); + PipeBarrier(); + } + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + + if constexpr (std::is_same::value) { + RmsNormVF(outputLocal, xFp32Local, gammaLocal, rmsNormParams); + PipeBarrier(); + } else { + RmsNormVF(xFp32Local, xFp32Local, gammaLocal, rmsNormParams); + // Cast xFp32 to outputLocal + PipeBarrier(); + Cast(outputLocal, xFp32Local, RoundMode::CAST_RINT, cnt); + PipeBarrier(); + } +} + +/** + * @brief RmsNormDynamicQuant rmsNorm + dynamicQuant融合流程 + 目前C均为float + * @param outputLocal 输出tensor,[B, S1, H], 量化完后为int8 + * @param outputScales dynamicQuant计算系数输出 + * @param inputGm 输入tensor,[B, S1, H],dtype支持bf16,fp32,int32 + * @param gammaLocal 输入tensor,[H, ], dtype只支持bf16 + * @param smoothLocal dynamicQuant平滑参数 + * @param dequantScaleWDqLocal 权重w反量化系数 + * @param dequantScaleXLocal x反量化系数 + * @param shareTmpUb 临时buffer 所需空间[(4 * cnt + 8) * sizeof(float)] cnt = row * col + * @param rmsNormParams rms所需系数,包括 + reciprocal rmsnorm系数reciprocal + epsilon rmsnorm系数epsilon + row 处理的行数;预留参数,当前仅支持单个batch的处理,row为1,对应S1 + col 列数,对应H + * @param enableSmoothScalesCq 表示是否有smoothGm的需求 + */ +template +__aicore__ inline void +RmsNormDynamicQuant(const LocalTensor &outputLocal, const LocalTensor &outputScales, + const GlobalTensor &inputGm, const LocalTensor &gammaLocal, + const LocalTensor &smoothLocal, const LocalTensor &dequantScaleWDqLocal, + const LocalTensor &dequantScaleXLocal, const LocalTensor &shareTmpUb, + RmsNormParam &rmsNormParams, bool enableSmoothScalesCq = true) +{ + int64_t cnt = rmsNormParams.row * rmsNormParams.col; + LocalTensor xFp32Local = shareTmpUb.ReinterpretCast(); + RmsNormNormal(xFp32Local, inputGm, gammaLocal, dequantScaleWDqLocal, dequantScaleXLocal, + shareTmpUb[rmsNormParams.col * sizeof(C)], rmsNormParams); + PipeBarrier(); + if constexpr (std::is_same::value && std::is_same::value) { + LocalTensor xBf16Local = xFp32Local[cnt].template ReinterpretCast(); + Cast(xBf16Local, xFp32Local, RoundMode::CAST_ROUND, cnt); + PipeBarrier(); + LocalTensor outputScalesLocal = outputScales.template ReinterpretCast(); + LocalTensor outLocal = outputLocal.template ReinterpretCast(); + LocalTensor tmpLocal = xBf16Local[cnt].template ReinterpretCast(); + DynamicQuantPerBlockMxfp8Vf(outLocal, outputScalesLocal, xBf16Local, tmpLocal, + rmsNormParams.row, rmsNormParams.col); + LocalTensor scale = outputScalesLocal.template ReinterpretCast(); + PipeBarrier(); + } else { + if (enableSmoothScalesCq) { + Mul(xFp32Local, xFp32Local, smoothLocal, cnt); + } + DynamicQuantPerTokenVf(outputLocal, outputScales, xFp32Local, rmsNormParams.row, rmsNormParams.col); + PipeBarrier(); + } +} + +} // namespace MlaProlog + +#endif \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_rope.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_rope.h new file mode 100644 index 000000000000..3afe4b333b79 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_rope.h @@ -0,0 +1,42 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file service_rope.h + * \brief + */ + +#ifndef SERVICE_ROPE_H +#define SERVICE_ROPE_H + +#include "vf/vf_rope.h" + +namespace MlaProlog { +/** + * @brief RotaryPosEmb, 同时做row行的RotaryPosEmb,每一行的元素为col + * @param outputLocal 输出tensor [row * col],支持和inputLocal是同一块空间 + * @param inputLocal 输入tensor [row * col] + * @param cosLocal cos系数tensor [(row - 1) * sinCosRepStride + col] + * @param sinLocal sin系数tensor [(row - 1) * sinCosRepStride + col] - 1 应已在sin中 + * @param row 待处理的行数 + * @param col 待处理的列数 col <= 512 / sizeof(C) + * @param sinCosRepStride 行与行之间sin/cos系数的偏移,单位为元素个数。 + */ +template +__aicore__ inline void RotaryPosEmb(const LocalTensor &outputLocal, const LocalTensor &inputLocal, + const LocalTensor &cosLocal, const LocalTensor &sinLocal, + uint64_t row, uint64_t col, uint8_t sinCosRepStride) +{ + DataSyncBarrier(); + RotaryPosEmbVF(outputLocal, inputLocal, cosLocal, sinLocal, + static_cast(row), col, col, static_cast(sinCosRepStride)); +} +} // namespace MlaProlog +#endif \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_rotary_position_embedding.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_rotary_position_embedding.h new file mode 100644 index 000000000000..6f4f3766876c --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_rotary_position_embedding.h @@ -0,0 +1,218 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file service_rotary_position_embedding.h + * \brief + */ + +#ifndef SERVICE_ROTARY_POSITION_EMBEDDING_H +#define SERVICE_ROTARY_POSITION_EMBEDDING_H + +#include "mla_prolog_comm.h" +#include "mla_prolog_vector_comm.h" +#include "service_rope.h" + +namespace MlaProlog { + +template +__aicore__ inline void PreprocessRopeInput(const GlobalTensor &inputGm, LocalTensor &shareTmpUb, + const Rectangle &ropeParams, LocalTensor &kFp32Local, + LocalTensor &kLocal, int64_t cnt) +{ + int64_t baseOffset; + kLocal = shareTmpUb.ReinterpretCast(); + + if constexpr (std::is_same::value) { + baseOffset = cnt; + } else { + baseOffset = cnt >> 1; + } + kFp32Local = shareTmpUb.ReinterpretCast()[baseOffset]; + + DataCopyExtParams copyParams{static_cast(ropeParams.row), + static_cast(ropeParams.col * sizeof(T)), + static_cast((ropeParams.stride - ropeParams.col) * sizeof(T)), 0, 0}; + DataCopyPadExtParams padParams{false, 0, 0, 0}; + + if constexpr (std::is_same::value) { + DataCopyPad(kFp32Local, inputGm, copyParams, padParams); + } else { + DataCopyPad(kLocal, inputGm, copyParams, padParams); + } +} + +/** + * @brief RotaryPosEmbPerTensor 对一个tensor进行RotartPosEmb,tensor的维度为[row * col] + 行与行之间sin/cos公用; + C:ropeComputType, float + * @param outputLocal 输出tensor + * @param inputGm 输入tensor + * @param cosLocal cos系数 + * @param sinLocal sin系数 + * @param shareTmpUb 临时buffer,需要大小为 cnt * 5 * sizeof(float) + * @param ropeParams 描述待处理数据的排布,包括 + row 行数 + col 列数 + stride 一行的真实长度 + * @param channelDeqScaleGm 量化参数;该tensor的每个元素不同 + * @param scale 量化参数;该tensor共用 + */ +template +__aicore__ inline void RotaryPosEmbPerTensor(LocalTensor &outputLocal, const GlobalTensor &inputGm, + const LocalTensor &cosLocal, const LocalTensor &sinLocal, + LocalTensor &shareTmpUb, Rectangle ropeParams, + LocalTensor channelDeqScaleLocal = LocalTensor(), + LocalTensor scale = LocalTensor()) +{ + int64_t cnt = ropeParams.row * ropeParams.col; + LocalTensor kFp32Local; + LocalTensor kLocal; + PreprocessRopeInput(inputGm, shareTmpUb, ropeParams, kFp32Local, kLocal, cnt); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + LocalTensor kFp32OutputLocal = kFp32Local[cnt]; + + if constexpr (std::is_same::value || enableDequant) { // 反量化 + Rectangle rectangleParams{ + static_cast(1), // row + static_cast(cnt), // col + static_cast(cnt) // columnStride + }; + if constexpr (std::is_same::value) { + Dequant(kFp32Local, kFp32Local, channelDeqScaleLocal, scale, rectangleParams); + } else { + Dequant(kFp32Local, kLocal, channelDeqScaleLocal, scale, rectangleParams); + } + PipeBarrier(); + } else if constexpr (std::is_same::value) { + Cast(kFp32Local, kLocal, RoundMode::CAST_NONE, cnt); + PipeBarrier(); + } + if constexpr (enableRope) { + if constexpr (std::is_same::value) { + RotaryPosEmb(outputLocal, kFp32Local, cosLocal, sinLocal, ropeParams.row, ropeParams.col, 0); + PipeBarrier(); + } else { + RotaryPosEmb(kFp32OutputLocal, kFp32Local, cosLocal, sinLocal, ropeParams.row, ropeParams.col, 0); + PipeBarrier(); + Cast(outputLocal, kFp32OutputLocal, RoundMode::CAST_RINT, cnt); + PipeBarrier(); + } + } else { + DataSyncBarrier(); + if constexpr (std::is_same::value) { + DataCopy(outputLocal, kFp32Local, cnt); + PipeBarrier(); + } else { + Cast(outputLocal, kFp32Local, RoundMode::CAST_RINT, cnt); + PipeBarrier(); + } + } +} + +/** + * @brief RotaryPosEmbPerHead 进行row行col列的RotaryPosEmb + 每行的量化系数,sin/cos均不同 + C:ropeComputType,float + * @param outputLocal 输出tensor + * @param inputGm 输入tensor + * @param cosLocal cos系数 + * @param sinLocal sin系数 + * @param shareTmpUb 临时buffer + * @param ropeParams 描述待处理数据的排布,包括 + row 行数 + col 列数 + stride 一行的真实长度 + * @param strideScale 一段的真实长度,描述channelDeqScaledGm数据排布 + * @param channelDeqScaleGm 量化参数:最终使用shape[1,col] + * @param deQuantScale 量化参数;最终使用shape[row,8] + */ +template +__aicore__ inline void RotaryPosEmbPerHead(LocalTensor &outputLocal, const GlobalTensor &inputGm, + const LocalTensor &cosLocal, const LocalTensor &sinLocal, + LocalTensor &shareTmpUb, Rectangle ropeParams, int64_t strideScale, + GlobalTensor channelDeqScaleGm = GlobalTensor(), + LocalTensor deQuantScale = LocalTensor()) +{ + // 在 BS = 1 场景可能存在有row为零的情况,提前返回减少运算 + if (ropeParams.row == 0) { + return; + } + + int64_t cnt = ropeParams.row * ropeParams.col; + LocalTensor kFp32Local; + LocalTensor kLocal; + PreprocessRopeInput(inputGm, shareTmpUb, ropeParams, kFp32Local, kLocal, cnt); + // scale参数可以和rope使用的空间复用 + LocalTensor scaleLocal = kFp32Local[cnt]; + LocalTensor kFp32OutputLocal = scaleLocal[ropeParams.col]; + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + if constexpr (std::is_same::value || enableDequant) { // 反量化 + DataCopyExtParams copyParams1{static_cast(1), static_cast(ropeParams.col * sizeof(C)), + static_cast((strideScale - ropeParams.col) * sizeof(C)), 0, 0}; + DataCopyPadExtParams padParams1{false, 0, 0, 0}; + DataCopyPad(scaleLocal, channelDeqScaleGm, copyParams1, padParams1); // 复用内存 + SetFlag(EVENT_ID1); + WaitFlag(EVENT_ID1); + + uint8_t blockNumPerRow = ropeParams.col / (ALIGN_BLOCK_SIZE / sizeof(C)); + // row col stride + Rectangle rectangleParams{static_cast(ropeParams.row), static_cast(ropeParams.col), + static_cast(ropeParams.col)}; + if constexpr (std::is_same::value) { + Dequant(kFp32Local, kFp32Local, scaleLocal, deQuantScale, rectangleParams); + } else { + Dequant(kFp32Local, kLocal, scaleLocal, deQuantScale, rectangleParams); + } + } else if constexpr (std::is_same::value) { + Cast(kFp32Local, kLocal, RoundMode::CAST_NONE, cnt); + } + PipeBarrier(); + if constexpr (enableRope) { + RotaryPosEmb(kFp32OutputLocal, kFp32Local, cosLocal, sinLocal, ropeParams.row, ropeParams.col, ropeParams.col); + PipeBarrier(); + if constexpr (std::is_same::value) { + DataCopy(outputLocal, kFp32OutputLocal, cnt); + } else { + Cast(outputLocal, kFp32OutputLocal, RoundMode::CAST_RINT, cnt); + } + } else { + DataSyncBarrier(); + if constexpr (std::is_same::value) { + DataCopy(outputLocal, kFp32Local, cnt); + } else { + Cast(outputLocal, kFp32Local, RoundMode::CAST_RINT, cnt); + } + } + PipeBarrier(); +} + +template +__aicore__ inline void RopePostQuantPerChannel(LocalTensor &outputLocal, LocalTensor &inputLocal, + LocalTensor &quantScaleLocal, LocalTensor &shareTmpUb, + int64_t cnt) +{ + LocalTensor inFp32; + if constexpr (std::is_same::value) { + inFp32 = inputLocal; + } else { + inFp32 = shareTmpUb.ReinterpretCast()[cnt]; + Cast(inFp32, inputLocal, RoundMode::CAST_NONE, cnt); + PipeBarrier(); + } + Mul(inFp32, inFp32, quantScaleLocal, cnt); + PipeBarrier(); + CastFP32ToINT8(outputLocal, inFp32, shareTmpUb, cnt); + PipeBarrier(); +} +} // namespace MlaProlog +#endif \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_scatter_cache.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_scatter_cache.h new file mode 100644 index 000000000000..c79165668223 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_scatter_cache.h @@ -0,0 +1,135 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file service_scatter_cache.h + * \brief + */ + +#ifndef SERVICE_SCATTER_CACHE_H +#define SERVICE_SCATTER_CACHE_H + +#include "mla_prolog_comm.h" + +namespace MlaProlog { + +__aicore__ inline int64_t CeilDiv64(int64_t x, int64_t y) +{ + if (y == 0) { + return static_cast(0); + } + return (x + y - 1) / y; +}; + +/** + * @brief PA场景,将inputLocal中的数据scatter到cacheGm,支持ND和Nz cache + * @param cacheGm 输出tensor + * ND [blockNum, blockSize, col] + * Nz [blockNum, ceil(col/col0), blockSize, col0] + * @param inputLocal 输入tensor,[row, col],一行对应一个token,只支持单行数据处理,row为1 + * @param scatterCacheParams 所需相关参数,包括 + blockSize KV blocks的大小 + paTokenIndex 待处理token在cache中的全局index,取值[0, blockNum*blockSize) + row 待处理的行数 + col 待处理的列数,需满足32 bytes对齐 + */ + +struct ScatterCacheParams { + int64_t blockSize; + int64_t paTokenIndex; + int64_t row; + int64_t col; + int64_t stride; + int64_t seqLength; + int64_t tokenIndex; +}; + +template +__aicore__ inline void ScatterCache(const GlobalTensor &cacheGm, const LocalTensor &inputLocal, + const ScatterCacheParams &scatterCacheParams) +{ + if (scatterCacheParams.paTokenIndex < 0) { + return; + } + if constexpr (!IS_NZ) { + DataCopy(cacheGm[scatterCacheParams.paTokenIndex * scatterCacheParams.stride], inputLocal, + scatterCacheParams.col); + } else { + constexpr uint8_t col0 = ALIGN_BLOCK_SIZE / sizeof(T); + int64_t cacheOffset = scatterCacheParams.paTokenIndex / scatterCacheParams.blockSize * + scatterCacheParams.blockSize * scatterCacheParams.stride + + scatterCacheParams.paTokenIndex % scatterCacheParams.blockSize * col0; + DataCopyParams copyParams{static_cast(scatterCacheParams.col / col0), 1, 0, + static_cast(scatterCacheParams.blockSize - 1)}; + DataCopy(cacheGm[cacheOffset], inputLocal, copyParams); + } +} + +template +__aicore__ inline void ScatterCacheUnAligned(const GlobalTensor &cacheGm, const LocalTensor &inputLocal, + const ScatterCacheParams &scatterCacheParams) +{ + if (scatterCacheParams.paTokenIndex < 0) { + return; + } + if constexpr (!IS_NZ) { + // blockCount, blockLen, srcStride, dstStride + DataCopyParams dataCopyParams{1, static_cast(scatterCacheParams.col * sizeof(T)), 0, 0}; + DataCopyPad(cacheGm[scatterCacheParams.paTokenIndex * scatterCacheParams.stride], inputLocal, dataCopyParams); + } +} + +template +__aicore__ inline void ScatterCacheMultiRows(GlobalTensor &cacheGm, const LocalTensor &inputLocal, + const ScatterCacheParams &scatterCacheParams, int64_t rowsInCurBatch, + int64_t cacheOffset, int64_t nextBatchOffset) +{ + if (cacheOffset < 0) { + return; + } + int64_t copyCnt = scatterCacheParams.col * rowsInCurBatch; + + if constexpr (!IS_NZ) { + DataCopy(cacheGm[cacheOffset], inputLocal, copyCnt); + if (rowsInCurBatch != scatterCacheParams.row) { + DataCopy(cacheGm[nextBatchOffset], inputLocal[copyCnt], + (scatterCacheParams.row - rowsInCurBatch) * scatterCacheParams.col); + } + } else { + constexpr uint8_t col0 = ALIGN_BLOCK_SIZE / sizeof(T); + DataCopyParams copyParams{static_cast(scatterCacheParams.col / col0), 1, 0, + static_cast(scatterCacheParams.blockSize - 1)}; + DataCopy(cacheGm[cacheOffset], inputLocal, copyParams); + if (rowsInCurBatch != scatterCacheParams.row) { + for (int row = 0; row < scatterCacheParams.row - rowsInCurBatch; ++row) { + DataCopy(cacheGm[nextBatchOffset + row * col0], inputLocal[copyCnt + row * col0], copyParams); + } + } + } +} + +template +__aicore__ inline void MaterializeOffsetsWithHeadSize(int64_t pageTokenOffset, int64_t tokenOffsetInPage, + int64_t rowsThisStep, bool spill, int64_t nextPageId, + int64_t headSize, CkvkrParams &ckvkrParams) +{ + ckvkrParams.rowsInCurBatch = rowsThisStep; + if constexpr (IS_NZ) { + constexpr uint8_t col0 = ALIGN_BLOCK_SIZE / sizeof(T); + ckvkrParams.cacheOffset = pageTokenOffset * headSize + tokenOffsetInPage * col0; + } else { + ckvkrParams.cacheOffset = (pageTokenOffset + tokenOffsetInPage) * headSize; + } + ckvkrParams.nextBatchOffset = (spill && nextPageId >= 0) ? nextPageId * headSize : 0; +} + +} // namespace MlaProlog + +#endif \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_comm.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_comm.h new file mode 100644 index 000000000000..10e38dc09f64 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_comm.h @@ -0,0 +1,40 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file vf_comm.h + * \brief VF comm + */ + +#ifndef VF_COMM_H +#define VF_COMM_H + +#include "kernel_tensor.h" + +namespace MlaProlog { +constexpr uint32_t ROPE_VF_COL = 64; + +constexpr MicroAPI::CastTrait castTraitB162B32 = { + MicroAPI::RegLayout::ZERO, + MicroAPI::SatMode::UNKNOWN, + MicroAPI::MaskMergeMode::ZEROING, + RoundMode::UNKNOWN, +}; + +constexpr MicroAPI::CastTrait castTraitB322B16 = { + MicroAPI::RegLayout::ZERO, + MicroAPI::SatMode::NO_SAT, + MicroAPI::MaskMergeMode::ZEROING, + RoundMode::CAST_RINT, +}; + +} // namespace MlaProlog + +#endif // VF_COMM_H diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_dequant.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_dequant.h new file mode 100644 index 000000000000..61c65f370bac --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_dequant.h @@ -0,0 +1,127 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file vf_dequant.h + * \brief + */ + +#ifndef VF_DEQUANT_H +#define VF_DEQUANT_H +#include "kernel_tensor.h" + +namespace MlaProlog { + +template +__simd_vf__ void DequantVFImpl(__ubuf__ float *yAddr, __ubuf__ T *xAddr, __ubuf__ float *scalePerChannelAddr, + __ubuf__ float *scalePerTokenAddr, uint32_t floatRepSize, uint32_t fp32BlockElementNum, + uint32_t dLoops, uint32_t dTail, uint32_t dTailLoop, uint32_t row, uint32_t col, + uint32_t stride) +{ + constexpr static AscendC::MicroAPI::CastTrait castTraitInt32ToFp32 = { + AscendC::MicroAPI::RegLayout::UNKNOWN, AscendC::MicroAPI::SatMode::NO_SAT, + AscendC::MicroAPI::MaskMergeMode::ZEROING, AscendC::RoundMode::CAST_RINT}; + + AscendC::MicroAPI::RegTensor vregInput; + AscendC::MicroAPI::RegTensor vregScalePerChannel; + AscendC::MicroAPI::RegTensor vregScalePerToken; + AscendC::MicroAPI::RegTensor vregInputFp32; // cast成float之后的vregInput + AscendC::MicroAPI::MaskReg fullMask = AscendC::MicroAPI::CreateMask(); + AscendC::MicroAPI::MaskReg tailMask; + tailMask = AscendC::MicroAPI::UpdateMask(dTail); + + uint32_t colOffset = 0; + uint32_t rowOffset = 0; + uint32_t scaleOffset = 0; + for (uint32_t j = 0; j < dLoops; j++) { + AscendC::MicroAPI::LoadAlign(vregScalePerChannel, + scalePerChannelAddr + colOffset); + rowOffset = 0; + scaleOffset = 0; + for (uint32_t i = 0; i < row; i++) { + AscendC::MicroAPI::LoadAlign( + vregScalePerToken, scalePerTokenAddr + scaleOffset); + if constexpr (!std::is_same::value) { + AscendC::MicroAPI::LoadAlign( + vregInput, xAddr + colOffset + rowOffset); + AscendC::MicroAPI::Cast( + vregInputFp32, vregInput, fullMask); + } else { + AscendC::MicroAPI::LoadAlign(vregInputFp32, + xAddr + colOffset + rowOffset); + } + AscendC::MicroAPI::Mul(vregInputFp32, vregInputFp32, vregScalePerChannel, fullMask); + AscendC::MicroAPI::Mul(vregInputFp32, vregInputFp32, vregScalePerToken, fullMask); + AscendC::MicroAPI::StoreAlign( + yAddr + colOffset + rowOffset, vregInputFp32, fullMask); + rowOffset += stride; + scaleOffset += fp32BlockElementNum; + } + colOffset += floatRepSize; + } + + if (dTailLoop > 0) { + rowOffset = 0; + scaleOffset = 0; + AscendC::MicroAPI::LoadAlign( + vregScalePerChannel, scalePerChannelAddr + dLoops * floatRepSize); + for (uint32_t i = 0; i < row; i++) { + AscendC::MicroAPI::LoadAlign( + vregScalePerToken, scalePerTokenAddr + scaleOffset); + if constexpr (!std::is_same::value) { + AscendC::MicroAPI::LoadAlign( + vregInput, xAddr + dLoops * floatRepSize + rowOffset); + AscendC::MicroAPI::Cast( + vregInputFp32, vregInput, tailMask); + } else { + AscendC::MicroAPI::LoadAlign( + vregInputFp32, xAddr + dLoops * floatRepSize + rowOffset); + } + AscendC::MicroAPI::Mul(vregInputFp32, vregInputFp32, vregScalePerChannel, tailMask); + AscendC::MicroAPI::Mul(vregInputFp32, vregInputFp32, vregScalePerToken, tailMask); + AscendC::MicroAPI::StoreAlign( + yAddr + dLoops * floatRepSize + rowOffset, vregInputFp32, tailMask); + rowOffset += stride; + scaleOffset += fp32BlockElementNum; + } + } +} + +/** + * @brief DequantVf 对输入做per-token叠加per-channel的反量化, INT32 ---> FP32. + * @param outputLocal 输出tensor [row, col] + * @param inputLocal 输入tensor [row, col] + * @param scalePerChannelLocal 输入tensor [1, col] + * @param scalePerTokenLocal 输入tensor [row, 8] + * @param row 待处理的行数 + * @param col 待处理的列数 + * @param stride 待处理数据一行的真实长度 + */ +template +__aicore__ inline void DequantVf(const LocalTensor &outputLocal, const LocalTensor &inputLocal, + const LocalTensor &scalePerChannelLocal, + const LocalTensor &scalePerTokenLocal, + uint32_t row, uint32_t col, uint32_t stride) +{ + __ubuf__ float *outputUb = (__ubuf__ float *)outputLocal.GetPhyAddr(); + __ubuf__ T *inputUb = (__ubuf__ T *)inputLocal.GetPhyAddr(); + __ubuf__ float *scalePerChannelLocalUb = (__ubuf__ float *)scalePerChannelLocal.GetPhyAddr(); + __ubuf__ float *scalePerTokenLocalUb = (__ubuf__ float *)scalePerTokenLocal.GetPhyAddr(); + + const uint32_t floatRepSize = 64; // 一个寄存器能够存放64个FP32 + const uint32_t fp32BlockElementNum = 8; + uint32_t dLoops = col / floatRepSize; + uint32_t dTail = col % floatRepSize; + uint32_t dTailLoop = dTail > 0 ? 1 : 0; + DequantVFImpl(outputUb, inputUb, scalePerChannelLocalUb, scalePerTokenLocalUb, floatRepSize, fp32BlockElementNum, + dLoops, dTail, dTailLoop, row, col, stride); +} +} // namespace MlaProlog +#endif diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_dynamic_quant.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_dynamic_quant.h new file mode 100644 index 000000000000..f6f115f32fae --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_dynamic_quant.h @@ -0,0 +1,446 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file vf_dynamic_quant.h + * \brief + */ + +#ifndef VF_DYNAMIC_QUANT_H +#define VF_DYNAMIC_QUANT_H + +#include "kernel_tensor.h" + +namespace MlaProlog { +constexpr float INT8_MAX_VALUE = 127.0f; +constexpr float FP8_E4M3FN_MAX_VALUE = 448.0f; +constexpr float FP8_E4M3FN_MIN_VALUE = -448.0f; +constexpr float HIFLOAT8_MAX_VALUE = 32768.0f; +constexpr uint32_t FP8_E4M3FN_BLOCK_SIZE = 32; +constexpr uint16_t FP8_E4M3_MAX_EXP = 0x0400; +constexpr uint16_t MAX_EXP_FOR_BF16 = 0x7f80; +constexpr uint16_t BF16_EXP_BIAS = 0x7f00; +constexpr uint16_t MAX_EXP_FOR_FP8 = 0x00ff; +constexpr uint16_t NAN_CUSTOMIZATION = 0x7f81; +constexpr uint16_t SPECIAL_EXP_THRESHOLD = 0x0040; +constexpr int64_t OUT_ELE_NUM_ONE_BLK = 64; +constexpr int64_t DIGIT_TWO = 2; +constexpr int16_t SHR_NUM_FOR_BF16 = 7; +constexpr uint8_t PER_TILE_QUANT_MODE = 1; +#ifndef INFINITY +#define INFINITY (__builtin_inff()) +#endif +constexpr float NEG_INFINITY = -INFINITY; +constexpr uint16_t REDUCE_SIZE = 8; + +template // M=0为fp8全量化pertoken量化;M=1为fp8全量化pertile量化 +__simd_vf__ void ComputeVFImpl(__ubuf__ T *xAddr, __ubuf__ O *yAddr, __ubuf__ float *scaleAddr, uint32_t rowIndex, + uint32_t rowCount, uint32_t dtypeSize, uint16_t VL, uint16_t vfLoop, + const float alphaValue) +{ + constexpr static AscendC::MicroAPI::CastTrait castTraitB16ToF32 = { + AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::UNKNOWN, + AscendC::MicroAPI::MaskMergeMode::ZEROING, AscendC::RoundMode::UNKNOWN}; + constexpr static AscendC::MicroAPI::CastTrait castTraitPack2 = { + AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::SAT, AscendC::MicroAPI::MaskMergeMode::ZEROING, + RoundMode::CAST_RINT}; + constexpr static AscendC::MicroAPI::CastTrait castTraitF32ToHalf = { + AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::NO_SAT, + AscendC::MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ODD}; + constexpr static AscendC::MicroAPI::CastTrait castTraitF32ToHif8 = { + AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::SAT, AscendC::MicroAPI::MaskMergeMode::ZEROING, + RoundMode::CAST_ROUND}; + static constexpr AscendC::MicroAPI::DivSpecificMode mode = {AscendC::MicroAPI::MaskMergeMode::ZEROING, true}; + AscendC::MicroAPI::RegTensor xInput; // 搬入的x + AscendC::MicroAPI::RegTensor xFp32; // cast成float之后的x + AscendC::MicroAPI::RegTensor xFp32Abs; // x的绝对值absvalue + AscendC::MicroAPI::RegTensor xMaxAbs; // x的abs与-inf比较的结果,可以认为是x的absvalue的一些最大值 + AscendC::MicroAPI::RegTensor xReduceMax; // reduceMax + AscendC::MicroAPI::RegTensor xScale; // scale + AscendC::MicroAPI::RegTensor xScaleDup; // Duplicate之后的scale,为了和input一起得到y + AscendC::MicroAPI::RegTensor xNorm; // input/scale + AscendC::MicroAPI::RegTensor yHalf; // float-->half-->int8 + AscendC::MicroAPI::RegTensor yOutput; // 最终y + + AscendC::MicroAPI::MaskReg validMask0; // 有效掩码 + AscendC::MicroAPI::MaskReg fullMask1 = + AscendC::MicroAPI::CreateMask(); // 启用所有通道,全掩码 + AscendC::MicroAPI::MaskReg validMask2; + + AscendC::MicroAPI::UnalignRegForStore ureg0; + AscendC::MicroAPI::Duplicate(xMaxAbs, NEG_INFINITY, fullMask1); + uint32_t sreg0 = rowCount; + + // 计算量化参数 + for (uint16_t j = 0; j < vfLoop; j++) { + validMask0 = AscendC::MicroAPI::UpdateMask(sreg0); // 有效元素 + if constexpr (!std::is_same::value) { + AscendC::MicroAPI::LoadAlign( + xInput, xAddr + rowIndex * rowCount + j * VL); + AscendC::MicroAPI::Cast(xFp32, xInput, validMask0); + } else { + AscendC::MicroAPI::LoadAlign(xFp32, xAddr + rowIndex * rowCount + + j * VL); + } + AscendC::MicroAPI::Abs(xFp32Abs, xFp32, validMask0); + AscendC::MicroAPI::Max(xMaxAbs, xFp32Abs, xMaxAbs, fullMask1); + } + AscendC::MicroAPI::Reduce( + xReduceMax, xMaxAbs, fullMask1); + if constexpr (M == PER_TILE_QUANT_MODE) { + constexpr float epsilonValue = 1e-4f; + AscendC::MicroAPI::RegTensor epsilonReg; + AscendC::MicroAPI::Duplicate(epsilonReg, epsilonValue, fullMask1); + AscendC::MicroAPI::Max(xReduceMax, xReduceMax, epsilonReg, fullMask1); // regtensor类型 + } + AscendC::MicroAPI::Muls(xScale, xReduceMax, alphaValue, fullMask1); + AscendC::MicroAPI::Duplicate(xScaleDup, xScale, fullMask1); + AscendC::MicroAPI::StoreUnAlign(scaleAddr, xScale, ureg0, + 1); + + uint32_t sreg1 = rowCount; + for (uint16_t j = 0; j < vfLoop; j++) { + auto addr = yAddr + rowIndex * rowCount + j * VL; + validMask2 = AscendC::MicroAPI::UpdateMask(sreg1); + if constexpr (!std::is_same::value) { + AscendC::MicroAPI::LoadAlign( + xInput, xAddr + rowIndex * rowCount + j * VL); + AscendC::MicroAPI::Cast(xFp32, xInput, validMask2); + } else { + AscendC::MicroAPI::LoadAlign(xFp32, xAddr + rowIndex * rowCount + + j * VL); + } + AscendC::MicroAPI::Div(xNorm, xFp32, xScaleDup, validMask2); + if constexpr (std::is_same::value) { + AscendC::MicroAPI::Cast(yOutput, xNorm, validMask2); + } else if constexpr (std::is_same::value) { + AscendC::MicroAPI::Cast(yOutput, xNorm, validMask2); + } else { + AscendC::MicroAPI::Cast(yHalf, xNorm, validMask2); + AscendC::MicroAPI::Cast(yOutput, yHalf, validMask2); + } + AscendC::MicroAPI::StoreAlign(addr, yOutput, validMask2); + } + AscendC::MicroAPI::StoreUnAlignPost(scaleAddr, ureg0, 0); +} + +/** + * @brief DynamicQuantPerTokenVf 对row行进行dynamicquant, BF16 ---> int8/FP8E4M3, 每一行出一个系数。 + * @param outputLocal 输出tensor [row , col] + * @param scale 输出每行的反量化系数 [row] + * @param inputLocal 输入tensor [row , col] + * @param row 待处理的行数 + * @param col 待处理的列数 + */ +template +__aicore__ inline void DynamicQuantPerTokenVf(const LocalTensor &outputLocal, const LocalTensor &scale, + const LocalTensor &inputLocal, uint64_t row, uint64_t col) +{ + auto xAddr = (__local_mem__ T *)inputLocal.GetPhyAddr(); + auto yAddr = (__local_mem__ O *)outputLocal.GetPhyAddr(); + auto scaleAddr = (__local_mem__ C *)scale.GetPhyAddr(); + uint32_t dtypeSize = sizeof(float); + uint16_t VL = AscendC::VECTOR_REG_WIDTH / dtypeSize; + uint32_t rowCount = col; + uint16_t vfLoop = (rowCount + VL - 1) / VL; + + constexpr float maxValue = std::is_same::value ? FP8_E4M3FN_MAX_VALUE : + std::is_same::value ? HIFLOAT8_MAX_VALUE : + INT8_MAX_VALUE; + const float alphaValue = static_cast(1.0) / maxValue; + for (int32_t i = 0; i < row; i++) { + ComputeVFImpl(xAddr, yAddr, scaleAddr + i, i, rowCount, dtypeSize, VL, vfLoop, alphaValue); + } +} + +template +__simd_vf__ void ComputeMaxExpVF(__ubuf__ T *srcAddr, __ubuf__ uint16_t *maxExpAddr, uint32_t totalCountInUB, + uint16_t loopNum, uint16_t vecLen) +{ + AscendC::MicroAPI::RegTensor vdExp0; + AscendC::MicroAPI::RegTensor vdExp1; + AscendC::MicroAPI::RegTensor vdExpExtract0; + AscendC::MicroAPI::RegTensor vdExpExtract1; + + AscendC::MicroAPI::RegTensor expMaskBF16; + AscendC::MicroAPI::Duplicate(expMaskBF16, MAX_EXP_FOR_BF16); + + AscendC::MicroAPI::RegTensor vdMaxExp; + AscendC::MicroAPI::MaskReg scaleMask1; + AscendC::MicroAPI::MaskReg scaleMask2; + AscendC::MicroAPI::UnalignRegForStore u1; + + for (uint16_t i = 0; i < loopNum; i++) { + scaleMask1 = AscendC::MicroAPI::UpdateMask(totalCountInUB); + scaleMask2 = AscendC::MicroAPI::UpdateMask(totalCountInUB); + AscendC::MicroAPI::LoadAlign(vdExp0, vdExp1, srcAddr, + vecLen * DIGIT_TWO); + // 通过位与运算得到bf16的指数位保留,尾数位置0所对应的值, 0x7f80是bf16 8个指数位为1,7个尾数位为0对应的值 + AscendC::MicroAPI::And(vdExpExtract0, (AscendC::MicroAPI::RegTensor &)vdExp0, expMaskBF16, + scaleMask1); + + AscendC::MicroAPI::And(vdExpExtract1, (AscendC::MicroAPI::RegTensor &)vdExp1, expMaskBF16, + scaleMask1); + // 得到指数位最大的值 + AscendC::MicroAPI::Max(vdMaxExp, vdExpExtract0, vdExpExtract1, scaleMask1); + // 得到每个block(32个元素)中的指数位最大的值 + AscendC::MicroAPI::ReduceDataBlock( + vdMaxExp, vdMaxExp, scaleMask1); + AscendC::MicroAPI::StoreUnAlign( + maxExpAddr, vdMaxExp, u1, REDUCE_SIZE); + } + AscendC::MicroAPI::StoreUnAlignPost(maxExpAddr, u1, 0); +} + + +template +__simd_vf__ void ComputeScaleVF(__ubuf__ uint16_t *maxExpAddr, __ubuf__ uint16_t *mxScaleLocalAddr, + __ubuf__ uint16_t *halfScaleLocalAddr, uint32_t totalScaleInUB, uint16_t loopNumScale, + uint16_t vecLen) +{ + AscendC::MicroAPI::RegTensor expMask; + AscendC::MicroAPI::Duplicate(expMask, MAX_EXP_FOR_BF16); + AscendC::MicroAPI::RegTensor vdMaxExp; + + AscendC::MicroAPI::MaskReg cmpResult; + AscendC::MicroAPI::MaskReg zeroMask; + AscendC::MicroAPI::MaskReg preMaskScale; + AscendC::MicroAPI::MaskReg invalidDataMask; + AscendC::MicroAPI::MaskReg specialDataMask; + + AscendC::MicroAPI::RegTensor maxExpValue; + AscendC::MicroAPI::Duplicate(maxExpValue, FP8_E4M3_MAX_EXP); + AscendC::MicroAPI::RegTensor sharedExp; + AscendC::MicroAPI::RegTensor scaleValue; + AscendC::MicroAPI::RegTensor scaleBias; + AscendC::MicroAPI::Duplicate(scaleBias, BF16_EXP_BIAS); + AscendC::MicroAPI::RegTensor halfScale; + AscendC::MicroAPI::RegTensor fp8NanRegTensor; + AscendC::MicroAPI::Duplicate(fp8NanRegTensor, MAX_EXP_FOR_FP8); + AscendC::MicroAPI::RegTensor zeroRegTensor; + AscendC::MicroAPI::Duplicate(zeroRegTensor, 0); + AscendC::MicroAPI::RegTensor nanRegTensor; + AscendC::MicroAPI::Duplicate(nanRegTensor, NAN_CUSTOMIZATION); + AscendC::MicroAPI::RegTensor specialExpRegTensor; + AscendC::MicroAPI::Duplicate(specialExpRegTensor, SPECIAL_EXP_THRESHOLD); + + for (uint16_t i = 0; i < loopNumScale; i++) { + preMaskScale = AscendC::MicroAPI::UpdateMask(totalScaleInUB); + AscendC::MicroAPI::LoadAlign(vdMaxExp, maxExpAddr, + vecLen); + AscendC::MicroAPI::Compare(cmpResult, vdMaxExp, expMask, preMaskScale); // INF/NAN + AscendC::MicroAPI::Compare(zeroMask, vdMaxExp, zeroRegTensor, preMaskScale); + // 如果vdMaxExp小于等于maxExpValue, 则置为maxExpValue, maxExpValue为FP8E4M3最大正整数的指数位8左移7位是0x400 + AscendC::MicroAPI::Compare(invalidDataMask, vdMaxExp, maxExpValue, preMaskScale); + AscendC::MicroAPI::Select(vdMaxExp, maxExpValue, vdMaxExp, invalidDataMask); + // vdMaxExp - maxExpValue后右移7位得到FP8E8M0的值,右移7位是因为bf16的指数位从第7位开始 + AscendC::MicroAPI::Sub(sharedExp, vdMaxExp, maxExpValue, preMaskScale); + AscendC::MicroAPI::ShiftRights(scaleValue, sharedExp, SHR_NUM_FOR_BF16, preMaskScale); + + AscendC::MicroAPI::Select(scaleValue, scaleValue, fp8NanRegTensor, cmpResult); + AscendC::MicroAPI::Select(scaleValue, scaleValue, zeroRegTensor, zeroMask); + + AscendC::MicroAPI::StoreAlign(mxScaleLocalAddr, scaleValue, + vecLen / DIGIT_TWO, preMaskScale); + + AscendC::MicroAPI::Compare(specialDataMask, sharedExp, scaleBias, preMaskScale); + // 0x7f00 - sharedExp得到1/sharedExp + AscendC::MicroAPI::Sub(halfScale, scaleBias, sharedExp, preMaskScale); + AscendC::MicroAPI::Select(halfScale, halfScale, nanRegTensor, cmpResult); + AscendC::MicroAPI::Select(halfScale, halfScale, zeroRegTensor, zeroMask); + AscendC::MicroAPI::Select(halfScale, specialExpRegTensor, halfScale, specialDataMask); + + AscendC::MicroAPI::StoreAlign( + halfScaleLocalAddr, halfScale, vecLen, preMaskScale); + } +} + +template +__simd_vf__ void ComputeDataVF(__ubuf__ T *srcAddr, __ubuf__ uint16_t *halfScaleLocalAddr, + __ubuf__ int8_t *outLocalAddr, uint32_t totalCountInUB, uint32_t totalCountInUB2, + uint16_t loopNum, uint16_t vecLen) +{ + AscendC::MicroAPI::MaskReg dataMask1; + AscendC::MicroAPI::MaskReg dataMask2; + AscendC::MicroAPI::MaskReg dataMask3; + AscendC::MicroAPI::MaskReg dataMask4; + AscendC::MicroAPI::MaskReg nanResult; + + AscendC::MicroAPI::RegTensor halfScaleForMul; + AscendC::MicroAPI::RegTensor floatScaleForMul; + AscendC::MicroAPI::RegTensor vdExp0; + AscendC::MicroAPI::RegTensor vdExp1; + AscendC::MicroAPI::RegTensor vdExp0Convert; + AscendC::MicroAPI::RegTensor vdExp1Convert; + AscendC::MicroAPI::RegTensor vdExp0FP32Zero; + AscendC::MicroAPI::RegTensor vdExp0FP32One; + AscendC::MicroAPI::RegTensor vdExp1FP32Zero; + AscendC::MicroAPI::RegTensor vdExp1FP32One; + AscendC::MicroAPI::RegTensor maxFp8Value; + AscendC::MicroAPI::Duplicate(maxFp8Value, FP8_E4M3FN_MAX_VALUE); + AscendC::MicroAPI::RegTensor minFp8Value; + AscendC::MicroAPI::Duplicate(minFp8Value, FP8_E4M3FN_MIN_VALUE); + AscendC::MicroAPI::RegTensor vdExp0FP8Zero; + AscendC::MicroAPI::RegTensor vdExp0FP8One; + AscendC::MicroAPI::RegTensor vdExp1FP8Zero; + AscendC::MicroAPI::RegTensor vdExp1FP8One; + + static constexpr AscendC::MicroAPI::CastTrait castTraitZero = { + AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::UNKNOWN, + AscendC::MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + static constexpr AscendC::MicroAPI::CastTrait castTraitOne = { + AscendC::MicroAPI::RegLayout::ONE, AscendC::MicroAPI::SatMode::UNKNOWN, + AscendC::MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + static constexpr AscendC::MicroAPI::CastTrait castTrait32to8 = { + AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::SAT, AscendC::MicroAPI::MaskMergeMode::ZEROING, + RoundMode::CAST_RINT}; + for (uint16_t i = 0; i < loopNum; i++) { + dataMask1 = AscendC::MicroAPI::UpdateMask(totalCountInUB); + dataMask2 = AscendC::MicroAPI::UpdateMask(totalCountInUB); + dataMask3 = AscendC::MicroAPI::UpdateMask(totalCountInUB2); + dataMask4 = AscendC::MicroAPI::UpdateMask(totalCountInUB2); + AscendC::MicroAPI::LoadAlign(vdExp0, vdExp1, srcAddr, + vecLen * DIGIT_TWO); + AscendC::MicroAPI::LoadAlign(halfScaleForMul, halfScaleLocalAddr, + REDUCE_SIZE); + + // X / mxscale + AscendC::MicroAPI::Mul(vdExp0, vdExp0, (AscendC::MicroAPI::RegTensor &)halfScaleForMul, dataMask1); + AscendC::MicroAPI::Mul(vdExp1, vdExp1, (AscendC::MicroAPI::RegTensor &)halfScaleForMul, dataMask1); + AscendC::MicroAPI::Interleave(vdExp0, vdExp1, vdExp0, vdExp1); + AscendC::MicroAPI::Cast(vdExp0FP32Zero, vdExp0, dataMask1); + AscendC::MicroAPI::Cast(vdExp0FP32One, vdExp0, dataMask1); + AscendC::MicroAPI::Interleave(vdExp0FP32Zero, vdExp0FP32One, vdExp0FP32Zero, vdExp0FP32One); + // 大于448.0的值设为448.0 + AscendC::MicroAPI::Compare(nanResult, vdExp0FP32Zero, maxFp8Value, dataMask1); + AscendC::MicroAPI::Select(vdExp0FP32Zero, maxFp8Value, vdExp0FP32Zero, nanResult); + AscendC::MicroAPI::Compare(nanResult, vdExp0FP32One, maxFp8Value, dataMask1); + AscendC::MicroAPI::Select(vdExp0FP32One, maxFp8Value, vdExp0FP32One, nanResult); + // 小于-448.0的值设为-448。0 + AscendC::MicroAPI::Compare(nanResult, vdExp0FP32Zero, minFp8Value, dataMask1); + AscendC::MicroAPI::Select(vdExp0FP32Zero, minFp8Value, vdExp0FP32Zero, nanResult); + AscendC::MicroAPI::Compare(nanResult, vdExp0FP32One, minFp8Value, dataMask1); + AscendC::MicroAPI::Select(vdExp0FP32One, minFp8Value, vdExp0FP32One, nanResult); + // 将结果转为FP8E4M3 + AscendC::MicroAPI::Cast(vdExp0FP8Zero, vdExp0FP32Zero, dataMask3); + AscendC::MicroAPI::Cast(vdExp0FP8One, vdExp0FP32One, dataMask3); + + AscendC::MicroAPI::Cast(vdExp1FP32Zero, vdExp1, dataMask2); + AscendC::MicroAPI::Cast(vdExp1FP32One, vdExp1, dataMask2); + AscendC::MicroAPI::Interleave(vdExp1FP32Zero, vdExp1FP32One, vdExp1FP32Zero, vdExp1FP32One); + // 大于448.0的值设为448.0 + AscendC::MicroAPI::Compare(nanResult, vdExp1FP32Zero, maxFp8Value, dataMask2); + AscendC::MicroAPI::Select(vdExp1FP32Zero, maxFp8Value, vdExp1FP32Zero, nanResult); + AscendC::MicroAPI::Compare(nanResult, vdExp1FP32One, maxFp8Value, dataMask2); + AscendC::MicroAPI::Select(vdExp1FP32One, maxFp8Value, vdExp1FP32One, nanResult); + // 小于-448.0的值设为-448。0 + AscendC::MicroAPI::Compare(nanResult, vdExp1FP32Zero, minFp8Value, dataMask2); + AscendC::MicroAPI::Select(vdExp1FP32Zero, minFp8Value, vdExp1FP32Zero, nanResult); + AscendC::MicroAPI::Compare(nanResult, vdExp1FP32One, minFp8Value, dataMask2); + AscendC::MicroAPI::Select(vdExp1FP32One, minFp8Value, vdExp1FP32One, nanResult); + // 将结果转为FP8E4M3 + AscendC::MicroAPI::Cast(vdExp1FP8Zero, vdExp1FP32Zero, dataMask4); + AscendC::MicroAPI::Cast(vdExp1FP8One, vdExp1FP32One, dataMask4); + AscendC::MicroAPI::StoreAlign( + outLocalAddr, (AscendC::MicroAPI::RegTensor &)vdExp0FP8Zero, OUT_ELE_NUM_ONE_BLK, dataMask3); + AscendC::MicroAPI::StoreAlign( + outLocalAddr, (AscendC::MicroAPI::RegTensor &)vdExp0FP8One, OUT_ELE_NUM_ONE_BLK, dataMask3); + AscendC::MicroAPI::StoreAlign( + outLocalAddr, (AscendC::MicroAPI::RegTensor &)vdExp1FP8Zero, OUT_ELE_NUM_ONE_BLK, dataMask4); + AscendC::MicroAPI::StoreAlign( + outLocalAddr, (AscendC::MicroAPI::RegTensor &)vdExp1FP8One, OUT_ELE_NUM_ONE_BLK, dataMask4); + } +} + +/** + * @brief DynamicQuantPerBlockMxfp8Vf 对row进行dynamicquant, BF16 ---> FP8E4M3, 每个BLOCK出一个系数。 + * @param outputLocal 输出tensor [row, col] + * @param outputScale 输出每行的反量化系数 [row, col/32] + * @param inputLocal 输入tensor [row, col] + * @param tmpLocal 临时buffer,所需空间大小row * col * 2 + * @param row 待处理的行数 + * @param col 待处理的列数 + */ +/** + shared_exp = floor(log2(max(|Vi|))) - emax + mxscale = 2^shared_exp + Pi = cast_to_dst_type(Vi/mxscale, round_mode) +**/ +template +__aicore__ inline void DynamicQuantPerBlockMxfp8Vf(const LocalTensor &outputLocal, + const LocalTensor &outputScale, + const LocalTensor &inputLocal, + const LocalTensor &tmpLocal, uint32_t row, uint32_t col) +{ + LocalTensor maxExpLocal = tmpLocal.ReinterpretCast(); + uint32_t totalScaleInUB = row * col / FP8_E4M3FN_BLOCK_SIZE; + uint32_t totalCountInUB = row * col; + uint16_t vecLen = AscendC::VECTOR_REG_WIDTH / sizeof(T); + uint16_t loopNum = (totalCountInUB + vecLen * DIGIT_TWO - 1) / (vecLen * DIGIT_TWO); + uint16_t loopNumScale = (totalScaleInUB + vecLen - 1) / vecLen; + + auto srcAddr = reinterpret_cast<__ubuf__ T *>(inputLocal.GetPhyAddr()); + auto maxExpAddr = reinterpret_cast<__ubuf__ uint16_t *>(maxExpLocal.GetPhyAddr()); + auto mxScaleLocalAddr = reinterpret_cast<__ubuf__ uint16_t *>(outputScale.GetPhyAddr()); + LocalTensor halfScaleLocal = maxExpLocal[totalCountInUB].template ReinterpretCast(); + auto halfScaleLocalAddr = reinterpret_cast<__ubuf__ uint16_t *>(halfScaleLocal.GetPhyAddr()); + auto outLocalAddr = reinterpret_cast<__ubuf__ int8_t *>(outputLocal.GetPhyAddr()); + ComputeMaxExpVF(srcAddr, maxExpAddr, totalCountInUB, loopNum, vecLen); + ComputeScaleVF(maxExpAddr, mxScaleLocalAddr, halfScaleLocalAddr, totalScaleInUB, loopNumScale, vecLen); + + srcAddr = reinterpret_cast<__ubuf__ T *>(inputLocal.GetPhyAddr()); + halfScaleLocalAddr = reinterpret_cast<__ubuf__ uint16_t *>(halfScaleLocal.GetPhyAddr()); + + uint32_t totalCountInUB2 = totalCountInUB * DIGIT_TWO; + ComputeDataVF(srcAddr, halfScaleLocalAddr, outLocalAddr, totalCountInUB, totalCountInUB2, loopNum, vecLen); +} + +/** + * @brief QuantPerTileVF 计算一行中每个tile的最大值,并计算量化参数和量化后的激活 + * @param outputLocal 输出tensor [row * col],row为rmsnorm输出的结果,均为1 + * @param inputLocal 输入tensor [row * col] + * @param quantScaleLocal 量化参数tensor [row, col / tileSize] + * @param row 处理数据的行数,默认为1 后续可拓展 + * @param col 处理数据的列数 + * @param tileSize tile的大小 当前只支持128,且col可被tileSize整除 + */ +template +__aicore__ inline void QuantPerTileVF(const LocalTensor &outputLocal, const LocalTensor &inputLocal, + const LocalTensor &quantScaleLocal, const uint32_t row, const uint32_t col, + const uint32_t tileSize) +{ + uint32_t cnt = row * col; + __ubuf__ T *inputBuf = (__ubuf__ T *)inputLocal.GetPhyAddr(); + __ubuf__ T *quantScaleBuf = (__ubuf__ T *)quantScaleLocal.GetPhyAddr(); + __ubuf__ O *outputBuf = (__ubuf__ O *)outputLocal.GetPhyAddr(); + uint32_t dtypeSize = sizeof(float); + uint16_t VL = AscendC::VECTOR_REG_WIDTH / dtypeSize; + uint16_t vfLoop = (tileSize + VL - 1) / VL; + constexpr float maxValue = std::is_same::value ? FP8_E4M3FN_MAX_VALUE : + std::is_same::value ? HIFLOAT8_MAX_VALUE : + INT8_MAX_VALUE; + const float alphaValue = static_cast(1.0) / maxValue; + uint32_t loopCount = cnt / tileSize; + for (uint32_t rowIndex = 0; rowIndex < loopCount; rowIndex++) { + ComputeVFImpl(inputBuf, outputBuf, quantScaleBuf + rowIndex, rowIndex, tileSize, + dtypeSize, VL, vfLoop, alphaValue); + } +} + +} // namespace MlaProlog +#endif \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_mul_qr.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_mul_qr.h new file mode 100644 index 000000000000..b2375018e5dc --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_mul_qr.h @@ -0,0 +1,74 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file vf_mul_qr.h + * \brief + */ + +#ifndef VF_MUL_QR_H +#define VF_MUL_QR_H + +namespace MlaProlog { +template +__simd_vf__ void MulQrVFImpl(__ubuf__ T *inputBuf, __ubuf__ T *outputBuf, __ubuf__ T *dequantScaleBrcbBuf, + const uint16_t floatRepSize, uint16_t repeatTimes, float quantScaleCkvRope, + uint64_t computeBlockAlign) +{ + MicroAPI::MaskReg pregAll = MicroAPI::CreateMask(); + + for (uint16_t i = 0; i < repeatTimes; i++) { + MicroAPI::RegTensor vregSrc; + MicroAPI::RegTensor vregQuantScale; + MicroAPI::RegTensor vregDequantScaleVrcb; + MicroAPI::RegTensor vregMulScale; + MicroAPI::RegTensor vregRes; + uint16_t loopOffset = i * floatRepSize; // 计算数据偏移 + uint16_t dequantLoopOffset = i * computeBlockAlign; // 动态量化参数偏移 + + MicroAPI::LoadAlign(vregSrc, inputBuf + loopOffset); + // broadcast量化系数 + MicroAPI::LoadAlign(vregDequantScaleVrcb, + dequantScaleBrcbBuf + dequantLoopOffset); + MicroAPI::Duplicate(vregQuantScale, quantScaleCkvRope); + + MicroAPI::Div(vregMulScale, vregQuantScale, vregDequantScaleVrcb, pregAll); + + MicroAPI::Mul(vregRes, vregSrc, vregMulScale, pregAll); + + MicroAPI::StoreAlign(outputBuf + loopOffset, vregRes, pregAll); + } +} + +/** + * @brief MulQrVF 对于RopeQr输出的结果进行mul系数计算 + * @param outputLocal 输出tensor [subRow, colRope],row目前均为1 + * @param inputLocal 输入tensor [subRow, colRope] + * @param dequantScaleBrcbLocal 动态量化参数tensor [row, computeBlockAlign] + * @param quantScaleCkvRope 量化系数 + * @param computeSizeRope 输入参数的大小 subRow * colRope + * @param computeBlockAlign 输入动态量化参数的单个数据块存放多少个动态量化参数 + */ +template +__aicore__ inline void MulQrVF(const LocalTensor &outputLocal, const LocalTensor &inputLocal, + const LocalTensor &dequantScaleBrcbLocal, float quantScaleCkvRope, + uint64_t computeSizeRope, uint64_t computeBlockAlign) +{ + const uint16_t floatRepSize = 64; // 一个寄存器能够存放64个FP32 + uint16_t repeatTimes = (computeSizeRope + floatRepSize - 1) / floatRepSize; // 对尾块处理的扩展,循环处理的次数 + + __ubuf__ T *inputBuf = (__ubuf__ T *)inputLocal.GetPhyAddr(); + __ubuf__ T *outputBuf = (__ubuf__ T *)outputLocal.GetPhyAddr(); + __ubuf__ T *dequantScaleBrcbBuf = (__ubuf__ T *)dequantScaleBrcbLocal.GetPhyAddr(); + MulQrVFImpl(inputBuf, outputBuf, dequantScaleBrcbBuf, floatRepSize, repeatTimes, quantScaleCkvRope, + computeBlockAlign); +} +} // namespace MlaProlog +#endif \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_quant_perchannel.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_quant_perchannel.h new file mode 100644 index 000000000000..2b9a5a973063 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_quant_perchannel.h @@ -0,0 +1,104 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file vf_quant_perchannel.h + * \brief + */ + +#ifndef VF_QUANT_PERCHANNEL_H +#define VF_QUANT_PERCHANNEL_H +#include "kernel_tensor.h" + +namespace MlaProlog { + +template +__simd_vf__ void QuantChannelVFImpl(__ubuf__ O *yAddr, __ubuf__ T *xAddr, __ubuf__ C *quantScaleAddr, + const uint32_t floatRepSize, uint32_t dLoops, uint32_t dTail, uint32_t dTailLoop, + uint32_t row, uint32_t col, uint32_t stride) +{ + AscendC::MicroAPI::RegTensor vregInput; + AscendC::MicroAPI::RegTensor vregQuantScale; + AscendC::MicroAPI::RegTensor vregOutput; + AscendC::MicroAPI::RegTensor vregOutputHalf; // float-->half-->int8 + AscendC::MicroAPI::MaskReg fullMask = AscendC::MicroAPI::CreateMask(); + AscendC::MicroAPI::MaskReg tailMask; + tailMask = AscendC::MicroAPI::UpdateMask(dTail); + constexpr static AscendC::MicroAPI::CastTrait castTraitPack2 = { + AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::SAT, AscendC::MicroAPI::MaskMergeMode::ZEROING, + AscendC::RoundMode::CAST_RINT}; + constexpr static AscendC::MicroAPI::CastTrait castTraitF32ToHalf = { + AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::NO_SAT, + AscendC::MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ODD}; + + uint32_t colOffset = 0; + uint32_t rowOffset = 0; + for (uint32_t j = 0; j < dLoops; j++) { + AscendC::MicroAPI::LoadAlign(vregQuantScale, + quantScaleAddr + colOffset); + rowOffset = 0; + for (uint32_t i = 0; i < row; i++) { + AscendC::MicroAPI::LoadAlign(vregInput, + xAddr + colOffset + rowOffset); + AscendC::MicroAPI::Mul(vregInput, vregInput, vregQuantScale, fullMask); + AscendC::MicroAPI::Cast(vregOutputHalf, vregInput, fullMask); + AscendC::MicroAPI::Cast(vregOutput, vregOutputHalf, fullMask); + AscendC::MicroAPI::StoreAlign( + yAddr + colOffset + rowOffset, vregOutput, fullMask); + rowOffset += stride; + } + colOffset += floatRepSize; + } + + if (dTailLoop > 0) { + rowOffset = 0; + AscendC::MicroAPI::LoadAlign(vregQuantScale, + quantScaleAddr + dLoops * floatRepSize); + for (uint32_t i = 0; i < row; i++) { + AscendC::MicroAPI::LoadAlign( + vregInput, xAddr + dLoops * floatRepSize + rowOffset); + AscendC::MicroAPI::Mul(vregInput, vregInput, vregQuantScale, tailMask); + AscendC::MicroAPI::Cast(vregOutputHalf, vregInput, tailMask); + AscendC::MicroAPI::Cast(vregOutput, vregOutputHalf, tailMask); + AscendC::MicroAPI::StoreAlign( + yAddr + dLoops * floatRepSize + rowOffset, vregOutput, tailMask); + rowOffset += stride; + } + } +} + +/** + * @brief QuantPerChannelVF 同时对row进行FP32到int8的per-channel量化操作。一行中的每一列用不同的量化参数。 + outLocal[i , j] = inputLocal[i , j] * quantScaleLocal[j] + * @param outputLocal 输出tensor [row, col] + * @param inputLocal 输入tensor [row, col] + * @param quantScaleLocal 量化参数tensor [1, col] + * @param row 处理数据的行数 默认为1 后续可扩展 + * @param col 处理数据的列数 + * @param stride 待处理数据一行的真实长度 + */ +template +__aicore__ inline void QuantPerChannelVf(const LocalTensor &outputLocal, const LocalTensor &inputLocal, + const LocalTensor &quantScaleLocal, const uint32_t row, const uint32_t col, + const uint32_t stride) +{ + __ubuf__ O *outputUb = (__ubuf__ O *)outputLocal.GetPhyAddr(); + __ubuf__ T *inputUb = (__ubuf__ T *)inputLocal.GetPhyAddr(); + __ubuf__ C *quantScaleUb = (__ubuf__ C *)quantScaleLocal.GetPhyAddr(); + + const uint32_t floatRepSize = 64; // 一个寄存器能够存放64个FP32 + uint32_t dLoops = col / floatRepSize; + uint32_t dTail = col % floatRepSize; + uint32_t dTailLoop = dTail > 0 ? 1 : 0; + + QuantChannelVFImpl(outputUb, inputUb, quantScaleUb, floatRepSize, dLoops, dTail, dTailLoop, row, col, stride); +} +} // namespace MlaProlog +#endif \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_quant_pertensor.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_quant_pertensor.h new file mode 100644 index 000000000000..9339b00712b6 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_quant_pertensor.h @@ -0,0 +1,89 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file vf_quant_pertensor.h + * \brief + */ + +#ifndef VF_QUANT_PERTENSOR_H +#define VF_QUANT_PERTENSOR_H +#include "kernel_tensor.h" + +namespace MlaProlog { +template +__simd_vf__ void QuantPerTensorVFImpl(__ubuf__ T *inputBuf, __ubuf__ T *quantScaleBuf, __ubuf__ U *outputBuf, + uint32_t cnt, const uint16_t floatRepSize, uint16_t repeatTimes) +{ + MicroAPI::MaskReg pregAll = MicroAPI::CreateMask(); + + // float -> fp8e4m3 类型转换模式结构体 + static constexpr MicroAPI::CastTrait CAST_TRAITF322FP8E4M3 = { + MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::SAT, MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_RINT}; + static constexpr MicroAPI::CastTrait CAST_TRAIT = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::SAT, + MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_RINT}; + static constexpr MicroAPI::CastTrait CAST_TRAITB162F32 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, + MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + static constexpr MicroAPI::CastTrait CAST_TRAITF322HIF8 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::SAT, + MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND}; + static constexpr MicroAPI::CastTrait castTraitF32ToHalf = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT, + MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ODD}; + MicroAPI::RegTensor vregSrc; + MicroAPI::RegTensor vregQuantScale; + MicroAPI::RegTensor vregFloat; + MicroAPI::RegTensor vregRes; + MicroAPI::RegTensor yHalf; + // 量化系数broadcast到寄存器所有位置 + MicroAPI::LoadAlign(vregQuantScale, quantScaleBuf); + for (uint16_t i = 0; i < uint16_t(repeatTimes); i++) { + uint16_t loopOffset = i * floatRepSize; + if constexpr (std::is_same::value) { + MicroAPI::LoadAlign(vregFloat, inputBuf + loopOffset); + } else if constexpr (std::is_same::value) { + MicroAPI::LoadAlign(vregSrc, inputBuf + loopOffset); + MicroAPI::Cast(vregFloat, vregSrc, pregAll); + } + MicroAPI::Mul(vregFloat, vregFloat, vregQuantScale, pregAll); + if constexpr (std::is_same::value) { + MicroAPI::Cast(vregRes, vregFloat, pregAll); + } else if constexpr (std::is_same::value) { + MicroAPI::Cast(yHalf, vregFloat, pregAll); + MicroAPI::Cast(vregRes, yHalf, pregAll); + } else { + MicroAPI::Cast(vregRes, vregFloat, pregAll); + } + MicroAPI::StoreAlign(outputBuf + loopOffset, vregRes, pregAll); + } +} + +/** + * @brief QuantPerTensorVF 对一行进行mul,并量化到fp8e4m3 T float U fp8e4m3 可根据不同量化结果扩展 + * @param outputLocal 输出tensor [row, col],row为rmsnorm输出的结果,均为1 + * @param inputLocal 输入tensor [row, col] + * @param quantScaleLocal 量化参数tensor [row, 1] + * @param row 处理数据的行数 默认为1 后续可扩展 + * @param col 处理数据的列数 + */ +template +__aicore__ inline void QuantPerTensorVF(const LocalTensor &outputLocal, const LocalTensor &inputLocal, + const LocalTensor &quantScaleLocal, const uint32_t row, const uint32_t col) +{ + uint32_t cnt = row * col; + const uint16_t floatRepSize = 64; // 一个寄存器能够存放64个FP32 + uint16_t repeatTimes = (cnt + floatRepSize - 1) / floatRepSize; // 对尾块处理的扩展,循环处理的次数 + + __ubuf__ T *inputBuf = (__ubuf__ T *)inputLocal.GetPhyAddr(); + __ubuf__ T *quantScaleBuf = (__ubuf__ T *)quantScaleLocal.GetPhyAddr(); + __ubuf__ U *outputBuf = (__ubuf__ U *)outputLocal.GetPhyAddr(); + + QuantPerTensorVFImpl(inputBuf, quantScaleBuf, outputBuf, cnt, floatRepSize, repeatTimes); +} +} // namespace MlaProlog +#endif \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_rms_norm.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_rms_norm.h new file mode 100644 index 000000000000..171cd9d5eb0e --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_rms_norm.h @@ -0,0 +1,107 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file vf_rms_norm.h + * \brief + */ + +#ifndef VF_RMS_NORM_H +#define VF_RMS_NORM_H +#include "kernel_tensor.h" + +namespace MlaProlog { +constexpr uint64_t FLOAT_REP_SIZE = 64; + +template +__simd_vf__ void RmsNormVFImpl(__ubuf__ InType *inputBuf, __ubuf__ GammaType *gammaBuf, __ubuf__ OutType *outputBuf, + uint32_t cnt, uint32_t repeatTimes, const RmsNormParam rmsNormParams) +{ + MicroAPI::RegTensor vregSum; + MicroAPI::RegTensor vregSumReduce; + MicroAPI::RegTensor vregDiv; + MicroAPI::RegTensor vregSquareRoot; + + MicroAPI::MaskReg pregAll = MicroAPI::CreateMask(); + MicroAPI::MaskReg pregFirst = MicroAPI::CreateMask(); + + static constexpr MicroAPI::CastTrait castTraitB162B32 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, + MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN}; + + MicroAPI::Duplicate(vregSum, 0.0); + + for (uint16_t i = 0; i < uint16_t(repeatTimes); ++i) { + MicroAPI::RegTensor vregXCast; + MicroAPI::RegTensor vregXSquare; + uint64_t loopOffset = i * FLOAT_REP_SIZE; + + MicroAPI::LoadAlign(vregXCast, inputBuf + loopOffset); + MicroAPI::Mul(vregXSquare, vregXCast, vregXCast, pregAll); + MicroAPI::Add(vregSum, vregXSquare, vregSum, pregAll); + } + + MicroAPI::Reduce(vregSumReduce, vregSum, + pregAll); + MicroAPI::Muls(vregSumReduce, vregSumReduce, rmsNormParams.reciprocal, + pregFirst); + MicroAPI::Adds(vregSumReduce, vregSumReduce, rmsNormParams.epsilon, + pregFirst); + MicroAPI::Sqrt(vregSquareRoot, vregSumReduce, pregFirst); + MicroAPI::Duplicate(vregDiv, vregSquareRoot, + pregAll); + + for (uint16_t i = 0; i < uint16_t(repeatTimes); ++i) { + MicroAPI::RegTensor vregXCast; + MicroAPI::RegTensor vregGamma; + MicroAPI::RegTensor vregGammaCast; + uint16_t loopOffset = i * FLOAT_REP_SIZE; + + MicroAPI::LoadAlign(vregXCast, inputBuf + loopOffset); + MicroAPI::LoadAlign(vregGamma, gammaBuf + loopOffset); + MicroAPI::Cast(vregGammaCast, vregGamma, pregAll); + + MicroAPI::Div(vregXCast, vregXCast, vregDiv, pregAll); + MicroAPI::Mul(vregXCast, vregXCast, vregGammaCast, pregAll); + + MicroAPI::StoreAlign(outputBuf + loopOffset, vregXCast, pregAll); + } +} + +/** + * @brief RmsNormVF 对一行进行rmsnorm + * @param outputLocal 输出tensor [row, col],row目前均为1 + * @param inputLocal 输入tensor [row, col] + * @param gammaLocal gamma参数tensor [row, col] + * @param rmsNormParams rmsNrom计算所需系数,包括 + row 行数 + col 列数,对应headSizeCq或headSizeCkv + reciprocal ,1/N + epsilon,防止除零极小数 + */ +template +__aicore__ inline void RmsNormVF(const LocalTensor &outputLocal, const LocalTensor &inputLocal, + const LocalTensor &gammaLocal, const RmsNormParam rmsNormParams) +{ + uint32_t cnt = rmsNormParams.row * rmsNormParams.col; + uint32_t repeatTimes = (cnt + FLOAT_REP_SIZE - 1) / FLOAT_REP_SIZE; + + __ubuf__ InType *inputBuf = (__ubuf__ InType *)inputLocal.GetPhyAddr(); + __ubuf__ GammaType *gammaBuf = (__ubuf__ GammaType *)gammaLocal.GetPhyAddr(); + __ubuf__ OutType *outputBuf = (__ubuf__ OutType *)outputLocal.GetPhyAddr(); + + RmsNormVFImpl(inputBuf, gammaBuf, outputBuf, cnt, repeatTimes, rmsNormParams); + + if (unlikely(rmsNormParams.isScaleEnable)) { + AscendC::PipeBarrier(); + Muls(outputLocal, outputLocal, rmsNormParams.scale, cnt); + } +} +} // namespace MlaProlog +#endif \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_rope.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_rope.h new file mode 100644 index 000000000000..c1c4481f75c0 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_rope.h @@ -0,0 +1,91 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file vf_rope.h + * \brief VF implementation for mla_prolog rope path (interleave-half mode). + */ + +#ifndef VF_ROPE_H +#define VF_ROPE_H + +#include "kernel_tensor.h" +#include "vf_comm.h" + +namespace MlaProlog { + +template +__simd_vf__ inline void RopeVFImpl(__ubuf__ C *outputUb, __ubuf__ C *inputUb, __ubuf__ C *sinUb, __ubuf__ C *cosUb, + uint32_t row, uint64_t srcStride, uint64_t dstStride, uint64_t sinCosStride) +{ + MicroAPI::RegTensor vregX; + MicroAPI::RegTensor vregSin; + MicroAPI::RegTensor vregCos; + MicroAPI::RegTensor vregEven; + MicroAPI::RegTensor vregOdd; + MicroAPI::RegTensor vregHigh; + MicroAPI::RegTensor vregLow; + MicroAPI::RegTensor vregTemp; + MicroAPI::RegTensor vregRes; + + MicroAPI::MaskReg maskAll = MicroAPI::CreateMask(); + MicroAPI::MaskReg maskLowHalf = MicroAPI::CreateMask(); + MicroAPI::MaskReg maskHighHalf; + MicroAPI::Not(maskHighHalf, maskLowHalf, maskAll); + + for (uint16_t i = 0; i < row; ++i) { + __ubuf__ C *curXUb = inputUb + i * srcStride; + __ubuf__ C *curResUb = outputUb + i * dstStride; + __ubuf__ C *curSinUb = sinUb + i * sinCosStride; + __ubuf__ C *curCosUb = cosUb + i * sinCosStride; + + MicroAPI::LoadAlign(vregX, curXUb); + MicroAPI::LoadAlign(vregSin, curSinUb); + MicroAPI::LoadAlign(vregCos, curCosUb); + + // vregEven = [evens(0..31), evens(32..63)] + // vregOdd = [odds(0..31), odds(32..63)] + MicroAPI::DeInterleave(vregEven, vregOdd, vregX, vregX); + + // Part1 low = cos * evens, Part1 high preserved + MicroAPI::Mul(vregRes, vregCos, vregEven, maskLowHalf); + // Part1 high = sin * evens, Part1 low preserved + MicroAPI::Mul(vregTemp, vregSin, vregEven, maskHighHalf); + // Part2 low = sin(-) * odds, Part2 high preserved + MicroAPI::Mul(vregLow, vregSin, vregOdd, maskLowHalf); + // Part2 high = cos * odds, Part2 low preserved + MicroAPI::Mul(vregHigh, vregCos, vregOdd, maskHighHalf); + + // Part1 = [cos_l*even + sin_l*odd, cos_u*odd + sin_u*even] = [y_lower, y_upper] + MicroAPI::Add(vregRes, vregRes, vregLow, maskLowHalf); + MicroAPI::Add(vregTemp, vregTemp, vregHigh, maskHighHalf); + MicroAPI::Move(vregRes, vregTemp, maskHighHalf); + + MicroAPI::StoreAlign(curResUb, vregRes, maskAll); + } +} + + +// col == 64 +template +__aicore__ inline void RotaryPosEmbVF(const LocalTensor &outputLocal, const LocalTensor &inputLocal, + const LocalTensor &cosLocal, const LocalTensor &sinLocal, uint32_t row, + uint64_t srcStride, uint64_t dstStride, uint64_t sinCosStride) +{ + __ubuf__ C *inputUb = (__ubuf__ C *)inputLocal.GetPhyAddr(); + __ubuf__ C *sinUb = (__ubuf__ C *)sinLocal.GetPhyAddr(); + __ubuf__ C *cosUb = (__ubuf__ C *)cosLocal.GetPhyAddr(); + __ubuf__ C *outputUb = (__ubuf__ C *)outputLocal.GetPhyAddr(); + + RopeVFImpl(outputUb, inputUb, sinUb, cosUb, row, srcStride, dstStride, sinCosStride); +} +} // namespace MlaProlog + +#endif // VF_ROPE_H diff --git a/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_template_tiling_key.h b/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_template_tiling_key.h new file mode 100644 index 000000000000..e6b367cba869 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_template_tiling_key.h @@ -0,0 +1,666 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_template_tiling_key.h + * \brief + */ + +#ifndef MLA_PROLOG_TEMPLATE_TILING_KEY_H +#define MLA_PROLOG_TEMPLATE_TILING_KEY_H + +#ifndef ORIG_DTYPE_TOKEN_X +#define ORIG_DTYPE_TOKEN_X (-1) +#endif + +#ifndef ORIG_DTYPE_WEIGHT_UQ_QR +#define ORIG_DTYPE_WEIGHT_UQ_QR (-1) +#endif + +#ifndef ORIG_DTYPE_KV_CACHE +#define ORIG_DTYPE_KV_CACHE (-1) +#endif + +#ifndef ORIG_DTYPE_KR_CACHE +#define ORIG_DTYPE_KR_CACHE (-1) +#endif + +#ifndef ORIG_DTYPE_QUERY +#define ORIG_DTYPE_QUERY (-1) +#endif + +#ifndef ORIG_DTYPE_DEQUANT_SCALE_X +#define ORIG_DTYPE_DEQUANT_SCALE_X (-1) +#endif + +#ifndef MLA_PROLOG_VERSION +#define MLA_PROLOG_VERSION (-1) +#endif + +#include "ascendc/host_api/tiling/template_argument.h" + + +#define ASCENDC_TPL_2_BW 2 // 每个参数占用2个bit位 +#define ASCENDC_TPL_4_BW 4 // 每个参数占用4个bit位 +#define ASCENDC_TPL_6_BW 6 // 每个参数占用6个bit位 + +// 可表示的tilingkey范围为64bit,注意不可超过限制 +ASCENDC_TPL_ARGS_DECL( + mla_prolog_v3, // 算子唯一标识,必须与 OPTYPE / opc --main_func=mla_prolog_v3 一致 + // bit:0-3 CACHE_MODE:0-ND 1-PA_BSND 2-PA_NZ 3-PA_BLK_BSND 4-PA_BLK_NZ + ASCENDC_TPL_UINT_DECL(CACHE_MODE, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, 0, 1, 2, 3, 4), + // bit:4-5 场景标识:0-FP16(预留) 1-BF16 2-量化场景 + ASCENDC_TPL_UINT_DECL(SCENARIO, ASCENDC_TPL_2_BW, ASCENDC_TPL_UI_LIST, 0, 1, 2), + // bit:6-11 量化场景:0-非量化 1-MMQcQr量化 2-MMQcQr量化+KVcache量化 3-MMcqCkvKr量化+MMQcQr量化 + // 4-MMCqCkvkr量化+MMQcQr量化+KVcache量化 5-MMQcQr量化+KVcache pertile量化 + // 6-MMCqCkvkr量化+MMQcQr量化+KVcache pertile量化 + // 7-Mxfp8量化+MMCqCkvkr量化+MMQcQr量化 8-Mxfp8量化+MMCqCkvkr量化+MMQcQr量化+KVcache量化 + // 9-Mxfp8量化+MMCqCkvkr量化+MMQcQr量化+KVcache pertile量化 + // 10-fp8量化+MMcqCkvKr量化+MMQcQr量化 11-fp8量化+MMCqCkvkr量化+MMQcQr量化+KVcache量化 + // 12-hif8量化+MMcqCkvKr量化+MMQcQr量化 13-hif8量化+MMCqCkvkr量化+MMQcQr量化+KVcache量化 + // 14-fp8量化+MMcqCkvKr量化+MMQcQr量化+KVcache pertile量化 15-hif8量化+MMCqCkvkr量化+MMQcQr量化+KVcache pertile量化 + ASCENDC_TPL_UINT_DECL(QUANT_MODE, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, + 13, 14, 15), + // bit:12 反量化使能:0-关闭 1-开启 + ASCENDC_TPL_BOOL_DECL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + // bit:13 量化算力分组:0-关闭 1-开启 + ASCENDC_TPL_BOOL_DECL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0, 1), + // bit:14-15 空tensor场景:0-无空tensor 1-kv_cache/kr_cache为空 2-query为空且不更新cache + ASCENDC_TPL_UINT_DECL(EMPTY_TENSOR_MODE, ASCENDC_TPL_2_BW, ASCENDC_TPL_UI_LIST, 0, 1, 2), + // bit:16-17 actualSeqLen使能场景 0-关闭 1-使能actualSeqLen + ASCENDC_TPL_UINT_DECL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_2_BW, ASCENDC_TPL_UI_LIST, 0, 1), + // bit:18-19 切M模式 0-关闭(切N) 1-使能(切M) + ASCENDC_TPL_UINT_DECL(SPLIT_M_MODE, ASCENDC_TPL_2_BW, ASCENDC_TPL_UI_LIST, 0, 1), + // bit:20-25 cv分核模式 ASCENDC_TPL_MIX_AIC_1_1(6)-1:1 ASCENDC_TPL_MIX_AIC_1_2(7)-1:2 + ASCENDC_TPL_KERNEL_TYPE_DECL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + // bit:26 rope使能:0-关闭 1-开启(仅A5/arch35生效) + ASCENDC_TPL_BOOL_DECL(ENABLE_ROPE, 0, 1)); + +ASCENDC_TPL_SEL( + +#if MLA_PROLOG_VERSION == -1 || MLA_PROLOG_VERSION >= 1 +// -------------------------- 非量化场景 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_BF16 && ORIG_DTYPE_WEIGHT_UQ_QR == DT_BF16) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1, 2, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif + +// -------------------------- 半量化kv非量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_BF16 && ORIG_DTYPE_WEIGHT_UQ_QR == DT_INT8 && ORIG_DTYPE_KV_CACHE == DT_BF16) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1, 2, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0, 1), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0, 1), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1, 2, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0, 1), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif + +// -------------------------- 半量化kv perchannel量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + ORIG_DTYPE_KR_CACHE == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_BF16 && ORIG_DTYPE_WEIGHT_UQ_QR == DT_INT8 && ORIG_DTYPE_KV_CACHE == DT_INT8 && \ + ORIG_DTYPE_KR_CACHE == DT_INT8) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1, 2, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0, 1), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0, 1), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1, 2, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0, 1), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif +#endif + +#if MLA_PROLOG_VERSION == -1 || MLA_PROLOG_VERSION >= 2 +// -------------------------- 全量化kv非量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_INT8 && ORIG_DTYPE_WEIGHT_UQ_QR == DT_INT8 && ORIG_DTYPE_KV_CACHE == DT_BF16) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1, 2, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 3), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 3), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 3), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif + +// -------------------------- 全量化kv pertensor量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + ORIG_DTYPE_KR_CACHE == -1 || ORIG_DTYPE_QUERY == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_INT8 && ORIG_DTYPE_WEIGHT_UQ_QR == DT_INT8 && ORIG_DTYPE_KV_CACHE == DT_INT8 && \ + ORIG_DTYPE_KR_CACHE == DT_BF16 && ORIG_DTYPE_QUERY == DT_INT8) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1, 2, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 4), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 4), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 4), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif +#endif + +#if MLA_PROLOG_VERSION == -1 || MLA_PROLOG_VERSION == 3 +// -------------------------- 半量化kv pertile量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_BF16 && ORIG_DTYPE_WEIGHT_UQ_QR == DT_INT8 && ORIG_DTYPE_KV_CACHE == DT_INT8) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 5), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 5), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif + +// -------------------------- 全量化kv pertile量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + ORIG_DTYPE_QUERY == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_INT8 && ORIG_DTYPE_WEIGHT_UQ_QR == DT_INT8 && ORIG_DTYPE_KV_CACHE == DT_INT8 && \ + ORIG_DTYPE_QUERY == DT_BF16) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 6), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif + +// -------------------------- Mxfp8全量化kv非量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + ORIG_DTYPE_DEQUANT_SCALE_X == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_FLOAT8_E4M3FN && ORIG_DTYPE_WEIGHT_UQ_QR == DT_FLOAT8_E4M3FN && \ + ORIG_DTYPE_KV_CACHE == DT_BF16 && ORIG_DTYPE_DEQUANT_SCALE_X == DT_FLOAT8_E8M0) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0, 1, 2, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 7), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif + +// -------------------------- Mxfp8全量化kv量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + ORIG_DTYPE_KR_CACHE == -1 || ORIG_DTYPE_DEQUANT_SCALE_X == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_FLOAT8_E4M3FN && ORIG_DTYPE_WEIGHT_UQ_QR == DT_FLOAT8_E4M3FN && \ + ORIG_DTYPE_KV_CACHE == DT_FLOAT8_E4M3FN && ORIG_DTYPE_KR_CACHE == DT_BF16 && \ + ORIG_DTYPE_DEQUANT_SCALE_X == DT_FLOAT8_E8M0) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0, 1, 2, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 8), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif + +// -------------------------- Mxfp8全量化kv pertile量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + ORIG_DTYPE_DEQUANT_SCALE_X == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_FLOAT8_E4M3FN && ORIG_DTYPE_WEIGHT_UQ_QR == DT_FLOAT8_E4M3FN && \ + ORIG_DTYPE_KV_CACHE == DT_FLOAT8_E4M3FN && ORIG_DTYPE_DEQUANT_SCALE_X == DT_FLOAT8_E8M0) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 9), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif + +// -------------------------- fp8全量化kv非量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + ORIG_DTYPE_DEQUANT_SCALE_X == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_FLOAT8_E4M3FN && ORIG_DTYPE_WEIGHT_UQ_QR == DT_FLOAT8_E4M3FN && \ + ORIG_DTYPE_KV_CACHE == DT_BF16 && ORIG_DTYPE_DEQUANT_SCALE_X == DT_FLOAT) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1, 2, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 10), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 10), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 10), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif + +// -------------------------- fp8全量化kv pertensor量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + ORIG_DTYPE_KR_CACHE == -1 || ORIG_DTYPE_QUERY == -1 || ORIG_DTYPE_DEQUANT_SCALE_X == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_FLOAT8_E4M3FN && ORIG_DTYPE_WEIGHT_UQ_QR == DT_FLOAT8_E4M3FN && \ + ORIG_DTYPE_KV_CACHE == DT_FLOAT8_E4M3FN && ORIG_DTYPE_KR_CACHE == DT_BF16 && \ + ORIG_DTYPE_QUERY == DT_FLOAT8_E4M3FN && ORIG_DTYPE_DEQUANT_SCALE_X == DT_FLOAT) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1, 2, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 11), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 11), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 11), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif +// -------------------------- hif8全量化kv非量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_HIFLOAT8 && ORIG_DTYPE_WEIGHT_UQ_QR == DT_HIFLOAT8 && ORIG_DTYPE_KV_CACHE == DT_BF16) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1, 2, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 12), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 12), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 12), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif + +// -------------------------- hif8全量化kv pertensor量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + ORIG_DTYPE_KR_CACHE == -1 || ORIG_DTYPE_QUERY == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_HIFLOAT8 && ORIG_DTYPE_WEIGHT_UQ_QR == DT_HIFLOAT8 && \ + ORIG_DTYPE_KV_CACHE == DT_HIFLOAT8 && ORIG_DTYPE_KR_CACHE == DT_BF16 && ORIG_DTYPE_QUERY == DT_HIFLOAT8) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1, 2, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 13), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 13), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 3, 4), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 13), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif + +// -------------------------- fp8全量化kv pertile量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + ORIG_DTYPE_QUERY == -1 || ORIG_DTYPE_DEQUANT_SCALE_X == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_FLOAT8_E4M3FN && ORIG_DTYPE_WEIGHT_UQ_QR == DT_FLOAT8_E4M3FN && \ + ORIG_DTYPE_KV_CACHE == DT_FLOAT8_E4M3FN && ORIG_DTYPE_QUERY == DT_BF16 && ORIG_DTYPE_DEQUANT_SCALE_X == DT_FLOAT) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 14), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 14), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif + +// -------------------------- hif8全量化kv pertile量化 -------------------------- +#if ORIG_DTYPE_TOKEN_X == -1 || ORIG_DTYPE_WEIGHT_UQ_QR == -1 || ORIG_DTYPE_KV_CACHE == -1 || \ + ORIG_DTYPE_QUERY == -1 || ORIG_DTYPE_DEQUANT_SCALE_X == -1 || \ + (ORIG_DTYPE_TOKEN_X == DT_HIFLOAT8 && ORIG_DTYPE_WEIGHT_UQ_QR == DT_HIFLOAT8 && \ + ORIG_DTYPE_KV_CACHE == DT_HIFLOAT8 && ORIG_DTYPE_QUERY == DT_BF16 && ORIG_DTYPE_DEQUANT_SCALE_X == DT_FLOAT) + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 1), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 15), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), + + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 15), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0, 1), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0, 1), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1)), +#endif +#endif + + // -------------------------- 空tensor场景 -------------------------- + ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(CACHE_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SCENARIO, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(QUANT_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_DEQUANT_OPTIONAL, 0), + ASCENDC_TPL_BOOL_SEL(ENABLE_GROUP_COMPUTE_OPTIONAL, 0), + ASCENDC_TPL_UINT_SEL(EMPTY_TENSOR_MODE, ASCENDC_TPL_UI_LIST, 2), + ASCENDC_TPL_UINT_SEL(ACTUAL_SEQ_LEN_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_UINT_SEL(SPLIT_M_MODE, ASCENDC_TPL_UI_LIST, 0), + ASCENDC_TPL_SHARED_KERNEL_TYPE_SEL(CV_MODE, ASCENDC_TPL_MIX_AIC_1_1, ASCENDC_TPL_MIX_AIC_1_2), + ASCENDC_TPL_BOOL_SEL(ENABLE_ROPE, 0, 1))); + +#endif // MLA_PROLOG_TEMPLATE_TILING_KEY_H \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_tiling_data.h b/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_tiling_data.h new file mode 100644 index 000000000000..5a963232bde3 --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_tiling_data.h @@ -0,0 +1,76 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_tiling_datay.h + * \brief + */ + +#ifndef MLA_PROLOG_TILING_DATA_H +#define MLA_PROLOG_TILING_DATA_H +#include +#include "kernel_tiling/kernel_tiling.h" + +namespace optiling { +// 1. 基础参数结构体(对应 MlaPrologBaseParams 宏定义) +struct MlaPrologBaseParams { + uint32_t batchSize; // batch size(批大小) + uint32_t stepBatchSize; // batch size per step 32(每步批大小) + uint32_t mSubSize; + uint32_t mSubCoreNum; + uint32_t stepNumHeadDequant; // head size per step when dequant before mmqn(反量化前每步头大小) + uint32_t tokenSize; // token size = batchSize * seq1Size(token总数:批大小×序列1长度) + uint32_t seq1Size; // seq1(序列1长度,通常为query序列长度) + uint32_t seq2Size; // seq2(序列2长度,通常为key/value序列长度) + uint32_t headSizeX; // head size of Input Hidden 7168(输入隐藏层的头维度) + uint32_t headSizeCq; // head size of Latent Query 1536(潜在Query的头维度) + uint32_t headSizeCkv; // head size of Latent KeyValue 512(潜在KeyValue的头维度) + uint32_t headSizeQc; // head size of Query = dimHeadSizeQc * numHeadSize = 128 * 32(Query总头维度) + uint32_t headSizeQr; // head size of Query Rope = dimHeadRope * numHeadSize = 64 * 32(带RoPE的Query头维度) + uint32_t headSizeKr; // head size of Key Rope 64(带RoPE的Key头维度) + uint32_t numHeadSize; // number of head 32(头数量) + uint32_t numHeadKvSize; // number of headkv(KeyValue的头数量) + uint32_t dimHeadSizeQc; // dim size per query head 128(单个Query头的维度) + uint32_t dimHeadRope; // dim size per rope head 64(单个带RoPE头的维度) + uint32_t blockNum; // pa block num(PA格式的块数量) + uint32_t blockSize; // pa block size 128(PA格式的块大小) + uint32_t mm1BlockNum; // 24 Cq(矩阵乘1的块数量,对应Cq计算) + uint32_t mm2BlockNum; // 9 Ckv(矩阵乘2的块数量,对应Ckv计算) + uint32_t mm3BlockNum; // 24 QcQr(矩阵乘3的块数量,对应QcQr计算) + uint32_t mm4BlockNum; // 24 Qn(矩阵乘4的块数量,对应Qn计算) + uint32_t vectorBlockNum; // 32(向量计算的块数量) + uint32_t mm1SingleCoreN; // single headSizeCq(单核心矩阵乘1的N维度大小,对应单个Cq头维度) + uint32_t mm2SingleCoreN; // single headSizeCkv+headSizeKr(单核心矩阵乘2的N维度大小,Ckv+Kr头维度之和) + uint32_t mm3SingleCoreN; // single headSizeQc+headSizeQr(单核心矩阵乘3的N维度大小,Qc+Qr头维度之和) + uint32_t mm4SingleCoreBatch; // single numHeadSize(单核心矩阵乘4的批大小,对应单个头数量) + uint32_t dtileSize; + uint32_t kvQuantMode; + uint32_t tileSize; + uint32_t ckvkrRepoMode; + uint32_t quantScaleRepoMode; + uint32_t queryNormFlag; + float reciprocalCq; // 1 / headSizeCq(headSizeCq的倒数,用于快速计算) + float epsilonCq; // Cq计算的epsilon(数值稳定性参数) + float reciprocalCkv; // 1 / headSizeCkv(headSizeCkv的倒数,用于快速计算) + float epsilonCkv; // Ckv计算的epsilon(数值稳定性参数) + float kNopeClipAlpha; + float qcQrScale; // query 的尺度矫正因子 + float kcScale; // kv 的尺度矫正因子 + uint16_t isQcQrScaleEnable; // query 的尺度矫正因子是否生效(默认是1.0的时候不生效) + uint16_t isKcScaleEnable; // kv 的尺度矫正因子是否生效(默认是1.0的时候不生效) +}; + +// 2. 完整分块数据结构体(对应 MlaPrologTilingData 宏定义,嵌套基础参数) +struct MlaPrologTilingData { + MlaPrologBaseParams baseParams; // 嵌套基础参数结构体(包含维度、头信息等核心配置) +}; +} // namespace optiling + +#endif // MLA_PROLOG_TILING_DATA_H \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_v3.cpp b/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_v3.cpp new file mode 100644 index 000000000000..4a806673815c --- /dev/null +++ b/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_v3.cpp @@ -0,0 +1,286 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_v3.cpp + * \brief + */ + +#define MLA_PROLOG_VERSION 3 +#define GLOBAL_OVERFLOW_MODE_CTRL 60 + +#if __has_include("arch35/kernel_mla_prolog_split_n.h") +#include "arch35/kernel_mla_prolog_split_n.h" +#include "arch35/kernel_mla_prolog_split_m.h" +#else +#include "arch35/kernel_mla_prolog_split_n.h" +#include "arch35/kernel_mla_prolog_split_m.h" +#endif +using namespace MlaProlog; + +template +__global__ __aicore__ void mla_prolog_v3( + __gm__ uint8_t *tokenX, + __gm__ uint8_t *weightDq, + __gm__ uint8_t *weightUqQr, + __gm__ uint8_t *weightUk, + __gm__ uint8_t *weightDkvKr, + __gm__ uint8_t *rmsnormGammaCq, + __gm__ uint8_t *rmsnormGammaCkv, + __gm__ uint8_t *ropeSin, + __gm__ uint8_t *ropeCos, + __gm__ uint8_t *kvCache, + __gm__ uint8_t *krCache, + __gm__ uint8_t *cacheIndex, + __gm__ uint8_t *dequantScaleX, + __gm__ uint8_t *dequantScaleWDq, + __gm__ uint8_t *dequantScaleWUqQr, + __gm__ uint8_t *dequantScaleWDkvKr, + __gm__ uint8_t *quantScaleCkv, + __gm__ uint8_t *quantScaleCkr, + __gm__ uint8_t *smoothScalesCq, + __gm__ uint8_t *actualSeqLen, + __gm__ uint8_t *kNopeClipAlpha, + __gm__ uint8_t *queryOut, + __gm__ uint8_t *queryRopeOut, + __gm__ uint8_t *kvCacheOut, + __gm__ uint8_t *krCacheOut, + __gm__ uint8_t *dequantScaleQNopeOut, + __gm__ uint8_t *queryNormOut, + __gm__ uint8_t *dequantScaleQNormOut, + __gm__ uint8_t *workspace, + __gm__ uint8_t *tiling) +{ +#if (__NPU_ARCH__ == 3510) + int64_t globalOriOverflowMode = AscendC::GetCtrlSpr(); +#endif + + REGISTER_TILING_DEFAULT(optiling::MlaPrologTilingData); + KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2); + + constexpr auto emptyMode = static_cast(EmptyTensorMode); + if constexpr (emptyMode == EMPTY_TENSOR_MODE::EMPTY_QUERY) { + return; + } + constexpr auto cacheMode = static_cast(CacheMode); + constexpr auto actualSeqLenMode = static_cast(ActualSeqLenMode); + constexpr auto splitMMode = static_cast(SplitMMode); + constexpr uint32_t cvRatio = CvMode == ASCENDC_TPL_MIX_AIC_1_1 ? 1 : 2; + + GET_TILING_DATA_WITH_STRUCT(optiling::MlaPrologTilingData, tilingDataIn, tiling); + const optiling::MlaPrologTilingData *__restrict tilingData = nullptr; + const optiling::MlaPrologBaseParams *__restrict tilingDataBaseParams = &tilingDataIn.baseParams; + + TPipe pipe; +#if (__NPU_ARCH__ == 3510) + AscendC::SetCtrlSpr(0); +#endif + if constexpr (static_cast(Scenario) == SCENARIO::NO_QUANT) { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::PARTIAL_QUANT_KV_NO_QUANT) { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_CHANNEL) { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::FULL_QUANT_KV_NO_QUANT) { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TENSOR) { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PERTILE) { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::FULL_QUANT_KV_QUANT_PERTILE) { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } +#if __CCE_AICORE__ == 310 + else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::MXFP8_FULL_QUANT_KV_NO_QUANT) { + if constexpr (splitMMode == SPLIT_M_MODE::ENABLED) { + MlaPrologV3SplitM> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::MXFP8_FULL_QUANT_KV_QUANT_PER_TENSOR) { + if constexpr (splitMMode == SPLIT_M_MODE::ENABLED) { + MlaPrologV3SplitM> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::MXFP8_FULL_QUANT_KV_QUANT_PER_TILE) { + if constexpr (splitMMode == SPLIT_M_MODE::ENABLED) { + MlaPrologV3SplitM> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::FP8_FULL_QUANT_KV_NO_QUANT) { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::FP8_FULL_QUANT_KV_QUANT_PER_TENSOR) { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::HIF8_FULL_QUANT_KV_NO_QUANT) { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::HIF8_FULL_QUANT_KV_QUANT_PER_TENSOR) { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::FP8_FULL_QUANT_KV_QUANT_PER_TILE) { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + static_cast(QuantMode) == QUANT_MODE::HIF8_FULL_QUANT_KV_QUANT_PER_TILE) { + MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); + op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, + ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, + dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, + queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); + op.Process(); + } +#endif +#if (__NPU_ARCH__ == 3510) + AscendC::SetCtrlSpr(globalOriOverflowMode); +#endif +} diff --git a/csrc/build_aclnn.sh b/csrc/build_aclnn.sh index b3e7d3aace90..cc32781add6b 100755 --- a/csrc/build_aclnn.sh +++ b/csrc/build_aclnn.sh @@ -224,6 +224,7 @@ elif [[ "$SOC_VERSION" =~ ^ascend950 ]]; then "store_kv_block_metadata" "k2q_csr" "sparse_attention_score" + "mla_prolog_v3" ) CUSTOM_OPS=$(IFS=';'; echo "${CUSTOM_OPS_ARRAY[*]}") diff --git a/csrc/torch_binding.cpp b/csrc/torch_binding.cpp index 2b9b16a6b34c..f961908cbed3 100644 --- a/csrc/torch_binding.cpp +++ b/csrc/torch_binding.cpp @@ -56,6 +56,7 @@ #include "attention/store_kv_block_metadata/store_kv_block_metadata_torch_adpt.cpp" #include "moe/dequant_situ_quant/dequant_situ_quant_torch_adpt.h" #include "moe/situ_mx_quant/situ_mx_quant_torch_adpt.h" +#include "attention/mla_prolog_v3/mla_prolog_v3_torch_adpt.h" #include #include #include @@ -2097,6 +2098,24 @@ TORCH_LIBRARY_EXPAND(CONCAT(_C, _ascend), ops) ops.impl("swap_blocks", torch::kPrivateUse1, &vllm_ascend::swap_blocks); #endif + // npu_mla_prolog_v3: aligned with torch_npu; underlying aclnn op is MlaPrologV3 (950-only). + ops.def( + "npu_mla_prolog_v3(Tensor token_x, Tensor weight_dq, Tensor weight_uq_qr, Tensor weight_uk," + " Tensor weight_dkv_kr, Tensor rmsnorm_gamma_cq, Tensor rmsnorm_gamma_ckv," + " Tensor rope_sin, Tensor rope_cos, Tensor(a!) kv_cache, Tensor(b!) kr_cache, *," + " Tensor? cache_index=None, Tensor? dequant_scale_x=None," + " Tensor? dequant_scale_w_dq=None, Tensor? dequant_scale_w_uq_qr=None," + " Tensor? dequant_scale_w_dkv_kr=None, Tensor? quant_scale_ckv=None," + " Tensor? quant_scale_ckr=None, Tensor? smooth_scales_cq=None," + " Tensor? actual_seq_len=None, Tensor? k_nope_clip_alpha=None," + " float rmsnorm_epsilon_cq=1e-05, float rmsnorm_epsilon_ckv=1e-05," + " str cache_mode=\"PA_BSND\", bool query_norm_flag=False," + " int weight_quant_mode=0, int kv_cache_quant_mode=0, int query_quant_mode=0," + " int ckvkr_repo_mode=0, int quant_scale_repo_mode=0, int tile_size=128," + " float qc_qr_scale=1.0, float kc_scale=1.0)" + " -> (Tensor, Tensor, Tensor, Tensor, Tensor)"); + ops.impl("npu_mla_prolog_v3", torch::kPrivateUse1, &vllm_ascend::npu_mla_prolog_v3); + // swap_blocks_batch takes CPU tensors (int64 pointer/size arrays), not NPU // tensors, so dispatch must be registered on the CPU backend. The function // internally submits async memcpy on the current NPU stream. diff --git a/csrc/torch_binding_meta.cpp b/csrc/torch_binding_meta.cpp index b67905e87915..158e5bc8ba36 100644 --- a/csrc/torch_binding_meta.cpp +++ b/csrc/torch_binding_meta.cpp @@ -1424,6 +1424,150 @@ void npu_scatter_nd_update_v2_meta( } +std::tuple npu_mla_prolog_v3_meta( + const at::Tensor &token_x, + const at::Tensor &weight_dq, + const at::Tensor &weight_uq_qr, + const at::Tensor &weight_uk, + const at::Tensor &weight_dkv_kr, + const at::Tensor &rmsnorm_gamma_cq, + const at::Tensor &rmsnorm_gamma_ckv, + const at::Tensor &rope_sin, + const at::Tensor &rope_cos, + at::Tensor &kv_cache, + at::Tensor &kr_cache, + const c10::optional &cache_index, + const c10::optional &dequant_scale_x, + const c10::optional &dequant_scale_w_dq, + const c10::optional &dequant_scale_w_uq_qr, + const c10::optional &dequant_scale_w_dkv_kr, + const c10::optional &quant_scale_ckv, + const c10::optional &quant_scale_ckr, + const c10::optional &smooth_scales_cq, + const c10::optional &actual_seq_len, + const c10::optional &k_nope_clip_alpha, + double rmsnorm_epsilon_cq, + double rmsnorm_epsilon_ckv, + c10::string_view cache_mode, + bool query_norm_flag, + int64_t weight_quant_mode, + int64_t kv_cache_quant_mode, + int64_t query_quant_mode, + int64_t ckvkr_repo_mode, + int64_t quant_scale_repo_mode, + int64_t tile_size, + double qc_qr_scale, + double kc_scale) +{ + constexpr int64_t FP8_E4M3_BLOCK_SIZE = 32; + const bool need_dequant_scale_q_nope = + (weight_quant_mode == 2 || weight_quant_mode == 3 || weight_quant_mode == 4 || + weight_quant_mode == 5) && + kv_cache_quant_mode == 1; + + // rope_sin/rope_cos are required args; empty tensors mean RoPE off (Dr defaults to 64). + // symbolic-meta-ok: empty rope_sin (numel==0) is the concrete RoPE-off runtime sentinel. + const bool rope_enabled = rope_sin.defined() && rope_sin.numel() > 0; + at::ScalarType query_dtype = rope_enabled ? rope_sin.scalar_type() : at::kBFloat16; + if (weight_quant_mode == 3 && kv_cache_quant_mode == 1) { + query_dtype = at::kFloat8_e4m3fn; + } else if (weight_quant_mode == 2 && kv_cache_quant_mode == 1) { + query_dtype = at::kChar; + } + + at::ScalarType query_norm_dtype = at::kBFloat16; + if (weight_quant_mode == 3 || weight_quant_mode == 4) { + query_norm_dtype = at::kFloat8_e4m3fn; + } else if (weight_quant_mode != 0) { + query_norm_dtype = weight_uq_qr.scalar_type(); + } + + at::ScalarType dequant_scale_q_norm_dtype = + weight_quant_mode == 3 ? at::kFloat8_e8m0fnu : at::kFloat; + + c10::SymDimVector query_shape; + c10::SymDimVector query_rope_shape; + c10::SymDimVector dequant_scale_q_nope_shape; + c10::SymDimVector query_norm_shape; + c10::SymDimVector dequant_scale_q_norm_shape; + + if (token_x.dim() == 3) { + c10::SymInt rope_dim = rope_enabled ? rope_sin.sym_size(2) : c10::SymInt(64); + query_shape = {token_x.sym_size(0), token_x.sym_size(1), weight_uk.sym_size(0), + weight_uk.sym_size(2)}; + query_rope_shape = {token_x.sym_size(0), token_x.sym_size(1), weight_uk.sym_size(0), + rope_dim}; + dequant_scale_q_nope_shape = {token_x.sym_size(0) * token_x.sym_size(1), + weight_uk.sym_size(0), c10::SymInt(1)}; + query_norm_shape = {token_x.sym_size(0), token_x.sym_size(1), weight_dq.sym_size(1)}; + dequant_scale_q_norm_shape = {token_x.sym_size(0) * token_x.sym_size(1)}; + if (weight_quant_mode == 3) { + dequant_scale_q_norm_shape.push_back(weight_dq.sym_size(1) / c10::SymInt(FP8_E4M3_BLOCK_SIZE)); + } else { + dequant_scale_q_norm_shape.push_back(c10::SymInt(1)); + } + } else { + c10::SymInt rope_dim = rope_enabled ? rope_sin.sym_size(1) : c10::SymInt(64); + query_shape = {token_x.sym_size(0), weight_uk.sym_size(0), weight_uk.sym_size(2)}; + query_rope_shape = {token_x.sym_size(0), weight_uk.sym_size(0), rope_dim}; + dequant_scale_q_nope_shape = {token_x.sym_size(0), weight_uk.sym_size(0), c10::SymInt(1)}; + query_norm_shape = {token_x.sym_size(0), weight_dq.sym_size(1)}; + dequant_scale_q_norm_shape = {token_x.sym_size(0)}; + if (weight_quant_mode == 3) { + dequant_scale_q_norm_shape.push_back(weight_dq.sym_size(1) / c10::SymInt(FP8_E4M3_BLOCK_SIZE)); + } else { + dequant_scale_q_norm_shape.push_back(c10::SymInt(1)); + } + } + + at::Tensor query = at::empty_symint(query_shape, token_x.options().dtype(query_dtype)); + at::Tensor query_rope = at::empty_symint(query_rope_shape, token_x.options().dtype(at::kBFloat16)); + at::Tensor dequant_scale_q_nope = + need_dequant_scale_q_nope + ? at::empty_symint(dequant_scale_q_nope_shape, token_x.options().dtype(at::kFloat)) + : at::empty_symint(c10::SymDimVector{c10::SymInt(0)}, + token_x.options().dtype(at::kFloat)); + at::Tensor query_norm = + query_norm_flag + ? at::empty_symint(query_norm_shape, token_x.options().dtype(query_norm_dtype)) + : at::empty_symint(c10::SymDimVector{c10::SymInt(0)}, + token_x.options().dtype(query_norm_dtype)); + at::Tensor dequant_scale_q_norm = + (query_norm_flag && weight_quant_mode != 0) + ? at::empty_symint(dequant_scale_q_norm_shape, + token_x.options().dtype(dequant_scale_q_norm_dtype)) + : at::empty_symint(c10::SymDimVector{c10::SymInt(0)}, + token_x.options().dtype(dequant_scale_q_norm_dtype)); + + (void)weight_dkv_kr; + (void)rmsnorm_gamma_cq; + (void)rmsnorm_gamma_ckv; + (void)rope_cos; + (void)kv_cache; + (void)kr_cache; + (void)cache_index; + (void)dequant_scale_x; + (void)dequant_scale_w_dq; + (void)dequant_scale_w_uq_qr; + (void)dequant_scale_w_dkv_kr; + (void)quant_scale_ckv; + (void)quant_scale_ckr; + (void)smooth_scales_cq; + (void)actual_seq_len; + (void)k_nope_clip_alpha; + (void)rmsnorm_epsilon_cq; + (void)rmsnorm_epsilon_ckv; + (void)cache_mode; + (void)query_quant_mode; + (void)ckvkr_repo_mode; + (void)quant_scale_repo_mode; + (void)tile_size; + (void)qc_qr_scale; + (void)kc_scale; + + return {query, query_rope, dequant_scale_q_nope, query_norm, dequant_scale_q_norm}; +} + std::tuple chunk_gated_delta_rule_fwd_h_meta( const at::Tensor & k, const at::Tensor & w, @@ -1866,6 +2010,8 @@ TORCH_LIBRARY_IMPL_EXPAND(CONCAT(_C, _ascend), Meta, ops) { ops.impl("npu_scatter_nd_update_v2", &vllm_ascend::meta::npu_scatter_nd_update_v2_meta); // Lightning indexer quant ops.impl("npu_lightning_indexer_quant", &vllm_ascend::meta::npu_lightning_indexer_quant_meta); + // MLA prolog (MlaPrologV3), Ascend950-only; name aligned with torch_npu + ops.impl("npu_mla_prolog_v3", &vllm_ascend::meta::npu_mla_prolog_v3_meta); // chunk_gated_delta_rule_fwd_h ops.impl("chunk_gated_delta_rule_fwd_h", &vllm_ascend::meta::chunk_gated_delta_rule_fwd_h_meta); // chunk_fwd_o diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_mla_prolog_v3.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_mla_prolog_v3.py new file mode 100644 index 000000000000..96981fdde2ed --- /dev/null +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_mla_prolog_v3.py @@ -0,0 +1,427 @@ +import gc + +import pytest +import torch +import torch_npu + +from vllm_ascend.utils import enable_custom_op + +enable_custom_op() + + +def _skip_if_mla_prolog_v3_unavailable(): + if not hasattr(torch.ops, "_C_ascend") or not hasattr(torch.ops._C_ascend, "npu_mla_prolog_v3"): + pytest.skip("requires the npu_mla_prolog_v3 custom operator") + + +@torch.inference_mode() +def test_mla_prolog_v3_native_bf16_head96(): + """Kimi K3 native bf16: head_num=96, q_lora=1536, kv_lora=512, D=128, Dr=64.""" + _skip_if_mla_prolog_v3_unavailable() + + token_num = 1 + head_num = 96 + he = 7168 + hcq = 1536 + hckv = 512 + d = 128 + dr = 64 + block_num = 2 + block_size = 128 + dtype = torch.bfloat16 + + token_x = torch.randn((token_num, he), dtype=dtype).npu() + weight_dq = torch_npu.npu_format_cast(torch.randn((he, hcq), dtype=dtype).npu().contiguous(), 29) + weight_uq_qr = torch_npu.npu_format_cast( + torch.randn((hcq, head_num * (d + dr)), dtype=dtype).npu().contiguous(), 29 + ) + weight_uk = torch.randn((head_num, d, hckv), dtype=dtype).npu() + weight_dkv_kr = torch_npu.npu_format_cast(torch.randn((he, hckv + dr), dtype=dtype).npu().contiguous(), 29) + rmsnorm_gamma_cq = torch.ones((hcq,), dtype=dtype).npu() + rmsnorm_gamma_ckv = torch.ones((hckv,), dtype=dtype).npu() + rope_sin = torch.randn((token_num, dr), dtype=dtype).npu() + rope_cos = torch.randn((token_num, dr), dtype=dtype).npu() + kv_cache = torch.zeros((block_num, block_size, 1, hckv), dtype=dtype).npu() + kr_cache = torch.zeros((block_num, block_size, 1, dr), dtype=dtype).npu() + cache_index = torch.arange(token_num, dtype=torch.int64).npu() + + kv_old = kv_cache.clone() + kr_old = kr_cache.clone() + + query, query_rope, dequant_scale_q_nope, query_norm, dequant_scale_q_norm = torch.ops._C_ascend.npu_mla_prolog_v3( + token_x, + weight_dq, + weight_uq_qr, + weight_uk, + weight_dkv_kr, + rmsnorm_gamma_cq, + rmsnorm_gamma_ckv, + rope_sin, + rope_cos, + kv_cache, + kr_cache, + cache_index=cache_index, + cache_mode="PA_BSND", + ) + + assert query.shape == (token_num, head_num, hckv) + assert query_rope.shape == (token_num, head_num, dr) + assert query.dtype == dtype + assert query_rope.dtype == dtype + assert dequant_scale_q_nope.numel() == 0 + assert query_norm.numel() == 0 + assert dequant_scale_q_norm.numel() == 0 + assert not torch.equal(kv_cache, kv_old) + assert not torch.equal(kr_cache, kr_old) + + gc.collect() + torch.npu.empty_cache() + torch.npu.reset_peak_memory_stats() + + +@torch.inference_mode() +def test_mla_prolog_v3_rope_disabled(): + _skip_if_mla_prolog_v3_unavailable() + + token_num = 1 + head_num = 96 + he = 7168 + hcq = 1536 + hckv = 512 + d = 128 + dr = 64 + block_num = 2 + block_size = 128 + dtype = torch.bfloat16 + + token_x = torch.randn((token_num, he), dtype=dtype).npu() + weight_dq = torch_npu.npu_format_cast(torch.randn((he, hcq), dtype=dtype).npu().contiguous(), 29) + weight_uq_qr = torch_npu.npu_format_cast( + torch.randn((hcq, head_num * (d + dr)), dtype=dtype).npu().contiguous(), 29 + ) + weight_uk = torch.randn((head_num, d, hckv), dtype=dtype).npu() + weight_dkv_kr = torch_npu.npu_format_cast(torch.randn((he, hckv + dr), dtype=dtype).npu().contiguous(), 29) + rmsnorm_gamma_cq = torch.ones((hcq,), dtype=dtype).npu() + rmsnorm_gamma_ckv = torch.ones((hckv,), dtype=dtype).npu() + rope_sin = torch.empty((0, dr), dtype=dtype).npu() + rope_cos = torch.empty((0, dr), dtype=dtype).npu() + kv_cache = torch.zeros((block_num, block_size, 1, hckv), dtype=dtype).npu() + kr_cache = torch.zeros((block_num, block_size, 1, dr), dtype=dtype).npu() + cache_index = torch.arange(token_num, dtype=torch.int64).npu() + + query, query_rope, *_ = torch.ops._C_ascend.npu_mla_prolog_v3( + token_x, + weight_dq, + weight_uq_qr, + weight_uk, + weight_dkv_kr, + rmsnorm_gamma_cq, + rmsnorm_gamma_ckv, + rope_sin, + rope_cos, + kv_cache, + kr_cache, + cache_index=cache_index, + cache_mode="PA_BSND", + ) + + assert query.shape == (token_num, head_num, hckv) + assert query_rope.shape == (token_num, head_num, dr) + + gc.collect() + torch.npu.empty_cache() + torch.npu.reset_peak_memory_stats() + + +@torch.inference_mode() +def test_mla_prolog_v3_query_norm_flag(): + _skip_if_mla_prolog_v3_unavailable() + + token_num = 1 + head_num = 96 + he = 7168 + hcq = 1536 + hckv = 512 + d = 128 + dr = 64 + block_num = 2 + block_size = 128 + dtype = torch.bfloat16 + + token_x = torch.randn((token_num, he), dtype=dtype).npu() + weight_dq = torch_npu.npu_format_cast(torch.randn((he, hcq), dtype=dtype).npu().contiguous(), 29) + weight_uq_qr = torch_npu.npu_format_cast( + torch.randn((hcq, head_num * (d + dr)), dtype=dtype).npu().contiguous(), 29 + ) + weight_uk = torch.randn((head_num, d, hckv), dtype=dtype).npu() + weight_dkv_kr = torch_npu.npu_format_cast(torch.randn((he, hckv + dr), dtype=dtype).npu().contiguous(), 29) + rmsnorm_gamma_cq = torch.ones((hcq,), dtype=dtype).npu() + rmsnorm_gamma_ckv = torch.ones((hckv,), dtype=dtype).npu() + rope_sin = torch.randn((token_num, dr), dtype=dtype).npu() + rope_cos = torch.randn((token_num, dr), dtype=dtype).npu() + kv_cache = torch.zeros((block_num, block_size, 1, hckv), dtype=dtype).npu() + kr_cache = torch.zeros((block_num, block_size, 1, dr), dtype=dtype).npu() + cache_index = torch.arange(token_num, dtype=torch.int64).npu() + + query, query_rope, _, query_norm, dequant_scale_q_norm = torch.ops._C_ascend.npu_mla_prolog_v3( + token_x, + weight_dq, + weight_uq_qr, + weight_uk, + weight_dkv_kr, + rmsnorm_gamma_cq, + rmsnorm_gamma_ckv, + rope_sin, + rope_cos, + kv_cache, + kr_cache, + cache_index=cache_index, + cache_mode="PA_BSND", + query_norm_flag=True, + ) + + assert query.shape == (token_num, head_num, hckv) + assert query_rope.shape == (token_num, head_num, dr) + assert query_norm.shape == (token_num, hcq) + assert query_norm.dtype == dtype + assert dequant_scale_q_norm.numel() == 0 + + gc.collect() + torch.npu.empty_cache() + torch.npu.reset_peak_memory_stats() + + +@pytest.mark.parametrize("cache_mode", ["PA_NZ", "TND"]) +@torch.inference_mode() +def test_mla_prolog_v3_cache_mode(cache_mode: str): + _skip_if_mla_prolog_v3_unavailable() + + token_num = 1 + head_num = 96 + he = 7168 + hcq = 1536 + hckv = 512 + d = 128 + dr = 64 + block_num = 2 + block_size = 128 + dtype = torch.bfloat16 + + token_x = torch.randn((token_num, he), dtype=dtype).npu() + weight_dq = torch_npu.npu_format_cast(torch.randn((he, hcq), dtype=dtype).npu().contiguous(), 29) + weight_uq_qr = torch_npu.npu_format_cast( + torch.randn((hcq, head_num * (d + dr)), dtype=dtype).npu().contiguous(), 29 + ) + weight_uk = torch.randn((head_num, d, hckv), dtype=dtype).npu() + weight_dkv_kr = torch_npu.npu_format_cast(torch.randn((he, hckv + dr), dtype=dtype).npu().contiguous(), 29) + rmsnorm_gamma_cq = torch.ones((hcq,), dtype=dtype).npu() + rmsnorm_gamma_ckv = torch.ones((hckv,), dtype=dtype).npu() + rope_sin = torch.randn((token_num, dr), dtype=dtype).npu() + rope_cos = torch.randn((token_num, dr), dtype=dtype).npu() + + if cache_mode == "TND": + kv_cache = torch.zeros((token_num, 1, hckv), dtype=dtype).npu() + kr_cache = torch.zeros((token_num, 1, dr), dtype=dtype).npu() + cache_index = None + else: + kv_cache = torch.zeros((block_num, block_size, 1, hckv), dtype=dtype).npu() + kr_cache = torch.zeros((block_num, block_size, 1, dr), dtype=dtype).npu() + cache_index = torch.arange(token_num, dtype=torch.int64).npu() + + query, query_rope, *_ = torch.ops._C_ascend.npu_mla_prolog_v3( + token_x, + weight_dq, + weight_uq_qr, + weight_uk, + weight_dkv_kr, + rmsnorm_gamma_cq, + rmsnorm_gamma_ckv, + rope_sin, + rope_cos, + kv_cache, + kr_cache, + cache_index=cache_index, + cache_mode=cache_mode, + ) + + assert query.shape == (token_num, head_num, hckv) + assert query_rope.shape == (token_num, head_num, dr) + + gc.collect() + torch.npu.empty_cache() + torch.npu.reset_peak_memory_stats() + + +@torch.inference_mode() +def test_mla_prolog_v3_cache_mode_bsnd(): + _skip_if_mla_prolog_v3_unavailable() + + batch = 1 + seq = 2 + head_num = 96 + he = 7168 + hcq = 1536 + hckv = 512 + d = 128 + dr = 64 + dtype = torch.bfloat16 + + token_x = torch.randn((batch, seq, he), dtype=dtype).npu() + weight_dq = torch_npu.npu_format_cast(torch.randn((he, hcq), dtype=dtype).npu().contiguous(), 29) + weight_uq_qr = torch_npu.npu_format_cast( + torch.randn((hcq, head_num * (d + dr)), dtype=dtype).npu().contiguous(), 29 + ) + weight_uk = torch.randn((head_num, d, hckv), dtype=dtype).npu() + weight_dkv_kr = torch_npu.npu_format_cast(torch.randn((he, hckv + dr), dtype=dtype).npu().contiguous(), 29) + rmsnorm_gamma_cq = torch.ones((hcq,), dtype=dtype).npu() + rmsnorm_gamma_ckv = torch.ones((hckv,), dtype=dtype).npu() + rope_sin = torch.randn((batch, seq, dr), dtype=dtype).npu() + rope_cos = torch.randn((batch, seq, dr), dtype=dtype).npu() + kv_cache = torch.zeros((batch, seq, 1, hckv), dtype=dtype).npu() + kr_cache = torch.zeros((batch, seq, 1, dr), dtype=dtype).npu() + + query, query_rope, *_ = torch.ops._C_ascend.npu_mla_prolog_v3( + token_x, + weight_dq, + weight_uq_qr, + weight_uk, + weight_dkv_kr, + rmsnorm_gamma_cq, + rmsnorm_gamma_ckv, + rope_sin, + rope_cos, + kv_cache, + kr_cache, + cache_mode="BSND", + ) + + assert query.shape == (batch, seq, head_num, hckv) + assert query_rope.shape == (batch, seq, head_num, dr) + + gc.collect() + torch.npu.empty_cache() + torch.npu.reset_peak_memory_stats() + + +@torch.inference_mode() +def test_mla_prolog_v3_partial_quant(): + _skip_if_mla_prolog_v3_unavailable() + + token_num = 1 + head_num = 96 + he = 7168 + hcq = 1536 + hckv = 512 + d = 128 + dr = 64 + block_num = 2 + block_size = 128 + dtype = torch.bfloat16 + + token_x = torch.randn((token_num, he), dtype=dtype).npu() + weight_dq = torch_npu.npu_format_cast(torch.randn((he, hcq), dtype=dtype).npu().contiguous(), 29) + weight_uq_qr = torch_npu.npu_format_cast( + torch.randint(-7, 8, (hcq, head_num * (d + dr)), dtype=torch.int8).npu().contiguous(), 29 + ) + weight_uk = torch.randn((head_num, d, hckv), dtype=dtype).npu() + weight_dkv_kr = torch_npu.npu_format_cast(torch.randn((he, hckv + dr), dtype=dtype).npu().contiguous(), 29) + rmsnorm_gamma_cq = torch.ones((hcq,), dtype=dtype).npu() + rmsnorm_gamma_ckv = torch.ones((hckv,), dtype=dtype).npu() + rope_sin = torch.randn((token_num, dr), dtype=dtype).npu() + rope_cos = torch.randn((token_num, dr), dtype=dtype).npu() + kv_cache = torch.zeros((block_num, block_size, 1, hckv), dtype=dtype).npu() + kr_cache = torch.zeros((block_num, block_size, 1, dr), dtype=dtype).npu() + cache_index = torch.arange(token_num, dtype=torch.int64).npu() + dequant_scale_w_uq_qr = torch.rand((1, head_num * (d + dr)), dtype=torch.float).npu() + + query, query_rope, *_ = torch.ops._C_ascend.npu_mla_prolog_v3( + token_x, + weight_dq, + weight_uq_qr, + weight_uk, + weight_dkv_kr, + rmsnorm_gamma_cq, + rmsnorm_gamma_ckv, + rope_sin, + rope_cos, + kv_cache, + kr_cache, + cache_index=cache_index, + dequant_scale_w_uq_qr=dequant_scale_w_uq_qr, + cache_mode="PA_BSND", + weight_quant_mode=1, + kv_cache_quant_mode=0, + ) + + assert query.shape == (token_num, head_num, hckv) + assert query_rope.shape == (token_num, head_num, dr) + + gc.collect() + torch.npu.empty_cache() + torch.npu.reset_peak_memory_stats() + + +@torch.inference_mode() +def test_mla_prolog_v3_full_int8_quant(): + _skip_if_mla_prolog_v3_unavailable() + + token_num = 1 + head_num = 96 + he = 7168 + hcq = 1536 + hckv = 512 + d = 128 + dr = 64 + block_num = 2 + block_size = 128 + dtype = torch.bfloat16 + + token_x = torch.randint(-7, 8, (token_num, he), dtype=torch.int8).npu() + weight_dq = torch_npu.npu_format_cast(torch.randint(-7, 8, (he, hcq), dtype=torch.int8).npu().contiguous(), 29) + weight_uq_qr = torch_npu.npu_format_cast( + torch.randint(-7, 8, (hcq, head_num * (d + dr)), dtype=torch.int8).npu().contiguous(), 29 + ) + weight_uk = torch.randn((head_num, d, hckv), dtype=dtype).npu() + weight_dkv_kr = torch_npu.npu_format_cast( + torch.randint(-7, 8, (he, hckv + dr), dtype=torch.int8).npu().contiguous(), 29 + ) + rmsnorm_gamma_cq = torch.ones((hcq,), dtype=dtype).npu() + rmsnorm_gamma_ckv = torch.ones((hckv,), dtype=dtype).npu() + rope_sin = torch.randn((token_num, dr), dtype=dtype).npu() + rope_cos = torch.randn((token_num, dr), dtype=dtype).npu() + kv_cache = torch.zeros((block_num, block_size, 1, hckv), dtype=dtype).npu() + kr_cache = torch.zeros((block_num, block_size, 1, dr), dtype=dtype).npu() + cache_index = torch.arange(token_num, dtype=torch.int64).npu() + dequant_scale_x = torch.rand((token_num, 1), dtype=torch.float).npu() + dequant_scale_w_dq = torch.rand((1, hcq), dtype=torch.float).npu() + dequant_scale_w_uq_qr = torch.rand((1, head_num * (d + dr)), dtype=torch.float).npu() + dequant_scale_w_dkv_kr = torch.rand((1, hckv + dr), dtype=torch.float).npu() + + query, query_rope, *_ = torch.ops._C_ascend.npu_mla_prolog_v3( + token_x, + weight_dq, + weight_uq_qr, + weight_uk, + weight_dkv_kr, + rmsnorm_gamma_cq, + rmsnorm_gamma_ckv, + rope_sin, + rope_cos, + kv_cache, + kr_cache, + cache_index=cache_index, + dequant_scale_x=dequant_scale_x, + dequant_scale_w_dq=dequant_scale_w_dq, + dequant_scale_w_uq_qr=dequant_scale_w_uq_qr, + dequant_scale_w_dkv_kr=dequant_scale_w_dkv_kr, + cache_mode="PA_BSND", + weight_quant_mode=2, + kv_cache_quant_mode=0, + ) + + assert query.shape == (token_num, head_num, hckv) + assert query_rope.shape == (token_num, head_num, dr) + + gc.collect() + torch.npu.empty_cache() + torch.npu.reset_peak_memory_stats() From 20bd5a7cb14cd645b67eb7f1ab67293da10ef80a Mon Sep 17 00:00:00 2001 From: yolic66 <747731294@qq.com> Date: Wed, 5 Aug 2026 09:43:10 +0800 Subject: [PATCH 23/50] [Ops][Feature] Support first-axis non-contiguous kv/kr cache for mla_prolog_v3 (#13422) ### What this PR does / why we need it? This PR enables `mla_prolog_v3` to write into `kv_cache` / `kr_cache` when the **first axis is non-contiguous** (common in Kimi K3 / PA cache layouts on Ascend950). Other axes remain required to be contiguous. **Changes** - Derive `kv_cache_stride0` / `kr_cache_stride0` from the cache tensor view in the aclnn path, and pass them into host tiling. - Propagate dim0 strides into kernel scatter (`ScatterCache` / `ScatterCacheMultiRows` / PA_BLK offset materialization) so writes use `blockIndex * stride0 + tokenOffset * tokenStride` instead of assuming a contiguous first axis. - Document the Ascend950 first-axis non-contiguous constraint in `csrc/attention/mla_prolog_v3/docs/api.md`. **Why** Serving paths may allocate KV/KR caches with a non-contiguous dim0 (e.g. strided / shared buffer layouts). Without stride-aware scatter, cache writes land at wrong offsets and produce incorrect results. Depends on / builds on #13355 (`mla_prolog_v3` + optional RoPE). After #13355 lands, this PR should rebase cleanly to the stride-only delta. Refs #13355 ### Does this PR introduce _any_ user-facing change? Yes (Ascend950 `mla_prolog_v3` callers only). - `kv_cache` / `kr_cache` may be **first-axis non-contiguous**; dims other than dim0 must still be contiguous. - Optional attrs `kv_cache_stride0` / `kr_cache_stride0` are filled automatically by the aclnn wrapper from the tensor view (callers normally do not set them). ### How was this patch tested? - Built on Ascend950 against the #13355 `mla_prolog_v3` baseline. - Functional / accuracy checks for contiguous and first-axis non-contiguous `kv_cache` / `kr_cache` (PA layouts used by Kimi K3). - Existing `tests/e2e/nightly/single_node/ops/singlecard_ops/test_mla_prolog_v3.py` still applicable for the base op path. - vLLM version: v0.26.0 - vLLM main: https://github.com/vllm-project/vllm/commit/d02df748bf9efd99022f1a062597dc3cb3808485 --------- Signed-off-by: yolic66 <747731294@qq.com> --- csrc/attention/mla_prolog_v3/docs/api.md | 2 + .../op_host/mla_prolog_tiling.cpp | 75 +++++++++++++ .../mla_prolog_v3/op_host/mla_prolog_tiling.h | 3 + .../arch35/kernel_mla_prolog_split_m.h | 24 ++-- .../arch35/kernel_mla_prolog_split_n.h | 22 ++-- .../op_kernel/arch35/service_scatter_cache.h | 40 +++++-- .../op_kernel/mla_prolog_tiling_data.h | 2 + .../ops/singlecard_ops/test_mla_prolog_v3.py | 103 ++++++++++++++++++ 8 files changed, 242 insertions(+), 29 deletions(-) diff --git a/csrc/attention/mla_prolog_v3/docs/api.md b/csrc/attention/mla_prolog_v3/docs/api.md index 3e54e7cfbf92..679dca4524ec 100644 --- a/csrc/attention/mla_prolog_v3/docs/api.md +++ b/csrc/attention/mla_prolog_v3/docs/api.md @@ -82,6 +82,8 @@ RoPE 开关由 `ropeSin` / `ropeCos` 的 nullity 推导:同时非空 → 开启,同时为空 → 关闭;混合 null 返回参数错误。 +`kv_cache` / `kr_cache` 在 Ascend 950PR/Ascend 950DT 上支持首轴非连续;除首轴外的其余轴必须连续。 + #### 量化模式合法组合(`weight_quant_mode` × `kv_cache_quant_mode`) | wq | 含义 | 合法 kvq | diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.cpp b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.cpp index 65da513d4c42..ae56d3a5aa85 100644 --- a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.cpp +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.cpp @@ -14,6 +14,7 @@ */ #include +#include #include #include #include @@ -54,6 +55,61 @@ inline auto Align(T num, T rnd) -> T return (((rnd) == 0) ? 0 : (((num) + (rnd)-1) / (rnd) * (rnd))); } +inline bool GetDefaultStride0(const gert::Shape &shape, uint64_t &stride0) +{ + stride0 = 1U; + for (size_t dim = 1U; dim < shape.GetDimNum(); ++dim) { + const int64_t dimSize = shape.GetDim(dim); + if (dimSize < 0 || + (dimSize != 0 && stride0 > static_cast(std::numeric_limits::max()) / + static_cast(dimSize))) { + return false; + } + stride0 *= static_cast(dimSize); + } + return true; +} + +inline ge::graphStatus GetCacheStride0(gert::TilingContext &context, uint32_t inputIndex, const gert::Shape &shape, + const char *tensorName, uint64_t &stride0) +{ + OP_CHECK_IF(shape.GetDimNum() == 0U, + OP_LOGE(context.GetNodeName(), "%s rank must be greater than 0.", tensorName), + return ge::GRAPH_FAILED); + uint64_t defaultStride0 = 0U; + OP_CHECK_IF(!GetDefaultStride0(shape, defaultStride0), + OP_LOGE(context.GetNodeName(), "%s shape cannot be represented by an int64 stride.", tensorName), + return ge::GRAPH_FAILED); + auto *stride = context.GetInputStride(inputIndex); + if (stride == nullptr || stride->GetDimNum() != shape.GetDimNum()) { + stride0 = defaultStride0; + OP_LOGD(context.GetNodeName(), "%s has no valid stride descriptor, use contiguous stride0=%lu.", tensorName, + stride0); + return ge::GRAPH_SUCCESS; + } + + uint64_t expectedStride = 1U; + for (int64_t dim = static_cast(shape.GetDimNum()) - 1; dim >= 1; --dim) { + const uint64_t actualStride = static_cast(stride->GetStride(static_cast(dim))); + OP_CHECK_IF(actualStride != expectedStride, + OP_LOGE(context.GetNodeName(), + "%s dim%ld must be contiguous, actual stride is %lu, expected stride is %lu. " + "Only dim0 may be non-contiguous.", + tensorName, dim, actualStride, expectedStride), + return ge::GRAPH_FAILED); + expectedStride *= static_cast(shape.GetDim(static_cast(dim))); + } + + const int64_t actualStride0 = stride->GetStride(MLA_PROLOG_DIM_INDEX_0); + OP_CHECK_IF(actualStride0 < 0 || static_cast(actualStride0) < defaultStride0, + OP_LOGE(context.GetNodeName(), "%s dim0 stride must be at least %lu, but got %ld.", tensorName, + defaultStride0, actualStride0), + return ge::GRAPH_FAILED); + stride0 = static_cast(actualStride0); + OP_LOGD(context.GetNodeName(), "%s stride0=%lu, contiguous stride0=%lu.", tensorName, stride0, defaultStride0); + return ge::GRAPH_SUCCESS; +} + NpuArch MlaPrologTiling::GetCurNpuArch() const { auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_->platformInfo); @@ -482,6 +538,8 @@ ge::graphStatus MlaPrologTiling::FillTiling() baseParams_->dimHeadRope = baseShapeInfo_.drSize; baseParams_->blockNum = baseShapeInfo_.blockNum; baseParams_->blockSize = baseShapeInfo_.blockSize; + baseParams_->kvCacheStride0 = context_->kvCacheStride0; + baseParams_->krCacheStride0 = context_->krCacheStride0; baseParams_->reciprocalCq = reciprocalCq_; baseParams_->epsilonCq = epsilonCq_; baseParams_->reciprocalCkv = reciprocalCkv_; @@ -705,6 +763,23 @@ ge::graphStatus MlaPrologTiling::ConvertContext(gert::TilingContext &context, Ml ConvertRequiredParams(context, mlaPrologContext); ConvertOptionalParams(context, mlaPrologContext); + OP_CHECK_IF(mlaPrologContext.kvCache.shape == nullptr || mlaPrologContext.krCache.shape == nullptr, + OP_LOGE(context.GetNodeName(), "kvCache or krCache shape is nullptr."), + return ge::GRAPH_FAILED); + const gert::Shape &kvCacheShape = mlaPrologContext.kvCache.shape->GetStorageShape(); + const gert::Shape &krCacheShape = mlaPrologContext.krCache.shape->GetStorageShape(); + const bool isV3 = std::strncmp(mlaPrologContext.opType, V3_OP_NAME, OP_NAME_LEN) == 0; + const uint32_t kvCacheIndex = isV3 ? KV_CACHE_INPUT_INDEX_V3 : KV_CACHE_INPUT_INDEX; + const uint32_t krCacheIndex = isV3 ? KR_CACHE_INPUT_INDEX_V3 : KR_CACHE_INPUT_INDEX; + OP_CHECK_IF(GetCacheStride0(context, kvCacheIndex, kvCacheShape, KV_CACHE_NAME, + mlaPrologContext.kvCacheStride0) != ge::GRAPH_SUCCESS, + OP_LOGE(context.GetNodeName(), "Failed to get or validate kvCache strides."), + return ge::GRAPH_FAILED); + OP_CHECK_IF(GetCacheStride0(context, krCacheIndex, krCacheShape, KR_CACHE_NAME, + mlaPrologContext.krCacheStride0) != ge::GRAPH_SUCCESS, + OP_LOGE(context.GetNodeName(), "Failed to get or validate krCache strides."), + return ge::GRAPH_FAILED); + auto attrs = context.GetAttrs(); OP_CHECK_IF(attrs == nullptr, OP_LOGE_WITH_INVALID_INPUT(context.GetNodeName(), "attrs"), return ge::GRAPH_FAILED); mlaPrologContext.rmsNormEspilonCq = attrs->GetAttrPointer(RMS_NORM_EPSILON_CQ_ATTR_INDEX); diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.h b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.h index bfe22001205a..71e779759a8a 100644 --- a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.h +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling.h @@ -377,6 +377,9 @@ struct MlaPrologContext { bool doRopeValue = true; const bool *doRope = nullptr; + uint64_t kvCacheStride0 = 0U; + uint64_t krCacheStride0 = 0U; + size_t *workSpaces; uint64_t tilingKey; uint32_t blockDim; diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_m.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_m.h index 7d73973f2698..30bd778bf879 100644 --- a/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_m.h +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_m.h @@ -1241,7 +1241,6 @@ __aicore__ inline void MlaPrologV3SplitM::ComputeBlkScatterOffsets(Global int64_t indexOffset = batchIndexOffset + batchTokenIndex / baseParams_->blockSize; int64_t paBlkId = indexGm(indexOffset); - int64_t pageTokenOffset = paBlkId * baseParams_->blockSize; int64_t tokenOffsetInPage = batchTokenIndex % baseParams_->blockSize; int64_t leftRowsInPage = baseParams_->blockSize - tokenOffsetInPage; @@ -1256,11 +1255,13 @@ __aicore__ inline void MlaPrologV3SplitM::ComputeBlkScatterOffsets(Global } MaterializeOffsetsWithHeadSize( - pageTokenOffset, tokenOffsetInPage, rowsThisStep, spill, nextPageId, baseParams_->headSizeCkv, + paBlkId, tokenOffsetInPage, rowsThisStep, spill, nextPageId, baseParams_->headSizeCkv, + baseParams_->kvCacheStride0, rmsNormAndScatterCkvParams); MaterializeOffsetsWithHeadSize( - pageTokenOffset, tokenOffsetInPage, rowsThisStep, spill, nextPageId, baseParams_->dimHeadRope, + paBlkId, tokenOffsetInPage, rowsThisStep, spill, nextPageId, baseParams_->dimHeadRope, + baseParams_->krCacheStride0, ropeAndScatterKrParams); } @@ -1373,7 +1374,7 @@ __aicore__ inline void MlaPrologV3SplitM::ScatterCkv(LocalTensor( kvCacheGm_, outputLocal, ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, baseParams_->headSizeCkv, - baseParams_->dtileSize}); + baseParams_->dtileSize, static_cast(baseParams_->kvCacheStride0)}); // 刷新量化scale if (isPertile && baseParams_->quantScaleRepoMode == 1U) { // BSND: @@ -1392,14 +1393,17 @@ __aicore__ inline void MlaPrologV3SplitM::ScatterCkv(LocalTensor( kvCacheGm_[startOffset], quantScaleCkvInt8Tensor, ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, - static_cast(tileNum * sizeof(float)), baseParams_->dtileSize}); + static_cast(tileNum * sizeof(float)), baseParams_->dtileSize, + static_cast(baseParams_->kvCacheStride0)}); } } else { // 使用 已计算好的偏移参数 ScatterCacheMultiRows( kvCacheGm_, outputLocal, ScatterCacheParams{baseParams_->blockSize, rmsNormAndScatterCkvParams.cacheOffset, vectorRow_, - baseParams_->headSizeCkv, baseParams_->seq1Size, rmsNormAndScatterCkvParams.tokenIndex}, + baseParams_->headSizeCkv, baseParams_->headSizeCkv, + static_cast(baseParams_->kvCacheStride0), baseParams_->seq1Size, + rmsNormAndScatterCkvParams.tokenIndex}, rmsNormAndScatterCkvParams.rowsInCurBatch, rmsNormAndScatterCkvParams.cacheOffset, rmsNormAndScatterCkvParams.nextBatchOffset); } @@ -1494,19 +1498,21 @@ __aicore__ inline void MlaPrologV3SplitM::ScatterKr(LocalTensorblockSize, paTokenIndex, vectorRow_, static_cast(baseParams_->dimHeadRope * sizeof(krCacheType)), - baseParams_->dtileSize}); + baseParams_->dtileSize, static_cast(baseParams_->kvCacheStride0)}); } else { ScatterCache( krCacheGm_, outputKrLocal, ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, baseParams_->dimHeadRope, - baseParams_->dimHeadRope}); + baseParams_->dimHeadRope, static_cast(baseParams_->krCacheStride0)}); } } else if constexpr ((MLAPT::cacheMode == CACHE_MODE::PA_BLK_BSND) || (MLAPT::cacheMode == CACHE_MODE::PA_BLK_NZ)) { // 使用 已计算好的偏移参数 ScatterCacheMultiRows( krCacheGm_, outputKrLocal, ScatterCacheParams{baseParams_->blockSize, ropeAndScatterKrParams.cacheOffset, vectorRow_, - baseParams_->dimHeadRope, baseParams_->seq1Size, ropeAndScatterKrParams.tokenIndex}, + baseParams_->dimHeadRope, baseParams_->dimHeadRope, + static_cast(baseParams_->krCacheStride0), baseParams_->seq1Size, + ropeAndScatterKrParams.tokenIndex}, ropeAndScatterKrParams.rowsInCurBatch, ropeAndScatterKrParams.cacheOffset, ropeAndScatterKrParams.nextBatchOffset); } else { diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_n.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_n.h index 711e629439ee..f6d6dbda5354 100644 --- a/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_n.h +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_n.h @@ -1500,7 +1500,6 @@ __aicore__ inline void MlaPrologVecS1CubS2::ComputeBlkScatterOffsets(Glob int64_t indexOffset = batchIndexOffset + batchTokenIndex / baseParams_->blockSize; int64_t paBlkId = indexGm(indexOffset); // 取cacheIdx - int64_t pageTokenOffset = paBlkId * baseParams_->blockSize; int64_t tokenOffsetInPage = batchTokenIndex % baseParams_->blockSize; int64_t leftRowsInPage = baseParams_->blockSize - tokenOffsetInPage; @@ -1516,11 +1515,13 @@ __aicore__ inline void MlaPrologVecS1CubS2::ComputeBlkScatterOffsets(Glob // --- Materialize for RMSNorm/CKV --- MaterializeOffsetsWithHeadSize( - pageTokenOffset, tokenOffsetInPage, rowsThisStep, spill, nextPageId, baseParams_->headSizeCkv, + paBlkId, tokenOffsetInPage, rowsThisStep, spill, nextPageId, baseParams_->headSizeCkv, + baseParams_->kvCacheStride0, rmsNormAndScatterCkvParams); MaterializeOffsetsWithHeadSize( - pageTokenOffset, tokenOffsetInPage, rowsThisStep, spill, nextPageId, baseParams_->dimHeadRope, + paBlkId, tokenOffsetInPage, rowsThisStep, spill, nextPageId, baseParams_->dimHeadRope, + baseParams_->krCacheStride0, ropeAndScatterKrParams); } @@ -1633,7 +1634,7 @@ __aicore__ inline void MlaPrologVecS1CubS2::ScatterCkv(LocalTensor( kvCacheGm_, outputLocal, ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, baseParams_->headSizeCkv, - baseParams_->dtileSize}); + baseParams_->dtileSize, static_cast(baseParams_->kvCacheStride0)}); // 刷新量化scale if (isPertile && baseParams_->quantScaleRepoMode == 1U) { // BSND: @@ -1652,13 +1653,15 @@ __aicore__ inline void MlaPrologVecS1CubS2::ScatterCkv(LocalTensor( kvCacheGm_[startOffset], quantScaleCkvInt8Tensor, ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, - static_cast(tileNum * sizeof(float)), baseParams_->dtileSize}); + static_cast(tileNum * sizeof(float)), baseParams_->dtileSize, + static_cast(baseParams_->kvCacheStride0)}); } } else { ScatterCacheMultiRows( kvCacheGm_, outputLocal, ScatterCacheParams{baseParams_->blockSize, rmsNormAndScatterCkvParams.cacheOffset, vectorRow_, - baseParams_->headSizeCkv, baseParams_->headSizeCkv, baseParams_->seq1Size, + baseParams_->headSizeCkv, baseParams_->headSizeCkv, + static_cast(baseParams_->kvCacheStride0), baseParams_->seq1Size, rmsNormAndScatterCkvParams.tokenIndex}, rmsNormAndScatterCkvParams.rowsInCurBatch, rmsNormAndScatterCkvParams.cacheOffset, rmsNormAndScatterCkvParams.nextBatchOffset); @@ -1753,18 +1756,19 @@ __aicore__ inline void MlaPrologVecS1CubS2::ScatterKr(LocalTensorblockSize, paTokenIndex, vectorRow_, static_cast(baseParams_->dimHeadRope * sizeof(krCacheType)), - baseParams_->dtileSize}); + baseParams_->dtileSize, static_cast(baseParams_->kvCacheStride0)}); } else { ScatterCache( krCacheGm_, outputKrLocal, ScatterCacheParams{baseParams_->blockSize, paTokenIndex, vectorRow_, baseParams_->dimHeadRope, - baseParams_->dimHeadRope}); + baseParams_->dimHeadRope, static_cast(baseParams_->krCacheStride0)}); } } else { ScatterCacheMultiRows( krCacheGm_, outputKrLocal, ScatterCacheParams{baseParams_->blockSize, ropeAndScatterKrParams.cacheOffset, vectorRow_, - baseParams_->dimHeadRope, baseParams_->dimHeadRope, baseParams_->seq1Size, + baseParams_->dimHeadRope, baseParams_->dimHeadRope, + static_cast(baseParams_->krCacheStride0), baseParams_->seq1Size, ropeAndScatterKrParams.tokenIndex}, ropeAndScatterKrParams.rowsInCurBatch, ropeAndScatterKrParams.cacheOffset, ropeAndScatterKrParams.nextBatchOffset); diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_scatter_cache.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_scatter_cache.h index c79165668223..869a44f05c7f 100644 --- a/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_scatter_cache.h +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/service_scatter_cache.h @@ -47,10 +47,22 @@ struct ScatterCacheParams { int64_t row; int64_t col; int64_t stride; + int64_t cacheStride0; int64_t seqLength; int64_t tokenIndex; }; +__aicore__ inline int64_t GetCacheOffset(int64_t paTokenIndex, int64_t blockSize, int64_t tokenStride, + int64_t cacheStride0) +{ + if (blockSize <= 0) { + return paTokenIndex * cacheStride0; + } + int64_t blockIndex = paTokenIndex / blockSize; + int64_t tokenIndexInBlock = paTokenIndex % blockSize; + return blockIndex * cacheStride0 + tokenIndexInBlock * tokenStride; +} + template __aicore__ inline void ScatterCache(const GlobalTensor &cacheGm, const LocalTensor &inputLocal, const ScatterCacheParams &scatterCacheParams) @@ -59,13 +71,15 @@ __aicore__ inline void ScatterCache(const GlobalTensor &cacheGm, const LocalT return; } if constexpr (!IS_NZ) { - DataCopy(cacheGm[scatterCacheParams.paTokenIndex * scatterCacheParams.stride], inputLocal, - scatterCacheParams.col); + int64_t cacheOffset = + GetCacheOffset(scatterCacheParams.paTokenIndex, scatterCacheParams.blockSize, + scatterCacheParams.stride, scatterCacheParams.cacheStride0); + DataCopy(cacheGm[cacheOffset], inputLocal, scatterCacheParams.col); } else { constexpr uint8_t col0 = ALIGN_BLOCK_SIZE / sizeof(T); - int64_t cacheOffset = scatterCacheParams.paTokenIndex / scatterCacheParams.blockSize * - scatterCacheParams.blockSize * scatterCacheParams.stride + - scatterCacheParams.paTokenIndex % scatterCacheParams.blockSize * col0; + int64_t cacheOffset = + GetCacheOffset(scatterCacheParams.paTokenIndex, scatterCacheParams.blockSize, col0, + scatterCacheParams.cacheStride0); DataCopyParams copyParams{static_cast(scatterCacheParams.col / col0), 1, 0, static_cast(scatterCacheParams.blockSize - 1)}; DataCopy(cacheGm[cacheOffset], inputLocal, copyParams); @@ -80,9 +94,12 @@ __aicore__ inline void ScatterCacheUnAligned(const GlobalTensor &cacheGm, con return; } if constexpr (!IS_NZ) { + int64_t cacheOffset = + GetCacheOffset(scatterCacheParams.paTokenIndex, scatterCacheParams.blockSize, + scatterCacheParams.stride, scatterCacheParams.cacheStride0); // blockCount, blockLen, srcStride, dstStride DataCopyParams dataCopyParams{1, static_cast(scatterCacheParams.col * sizeof(T)), 0, 0}; - DataCopyPad(cacheGm[scatterCacheParams.paTokenIndex * scatterCacheParams.stride], inputLocal, dataCopyParams); + DataCopyPad(cacheGm[cacheOffset], inputLocal, dataCopyParams); } } @@ -116,18 +133,19 @@ __aicore__ inline void ScatterCacheMultiRows(GlobalTensor &cacheGm, const Loc } template -__aicore__ inline void MaterializeOffsetsWithHeadSize(int64_t pageTokenOffset, int64_t tokenOffsetInPage, +__aicore__ inline void MaterializeOffsetsWithHeadSize(int64_t pageId, int64_t tokenOffsetInPage, int64_t rowsThisStep, bool spill, int64_t nextPageId, - int64_t headSize, CkvkrParams &ckvkrParams) + int64_t headSize, int64_t cacheStride0, + CkvkrParams &ckvkrParams) { ckvkrParams.rowsInCurBatch = rowsThisStep; if constexpr (IS_NZ) { constexpr uint8_t col0 = ALIGN_BLOCK_SIZE / sizeof(T); - ckvkrParams.cacheOffset = pageTokenOffset * headSize + tokenOffsetInPage * col0; + ckvkrParams.cacheOffset = pageId * cacheStride0 + tokenOffsetInPage * col0; } else { - ckvkrParams.cacheOffset = (pageTokenOffset + tokenOffsetInPage) * headSize; + ckvkrParams.cacheOffset = pageId * cacheStride0 + tokenOffsetInPage * headSize; } - ckvkrParams.nextBatchOffset = (spill && nextPageId >= 0) ? nextPageId * headSize : 0; + ckvkrParams.nextBatchOffset = (spill && nextPageId >= 0) ? nextPageId * cacheStride0 : 0; } } // namespace MlaProlog diff --git a/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_tiling_data.h b/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_tiling_data.h index 5a963232bde3..5ac674341290 100644 --- a/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_tiling_data.h +++ b/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_tiling_data.h @@ -41,6 +41,8 @@ struct MlaPrologBaseParams { uint32_t dimHeadRope; // dim size per rope head 64(单个带RoPE头的维度) uint32_t blockNum; // pa block num(PA格式的块数量) uint32_t blockSize; // pa block size 128(PA格式的块大小) + uint64_t kvCacheStride0; // kv cache first-axis stride in elements + uint64_t krCacheStride0; // kr cache first-axis stride in elements uint32_t mm1BlockNum; // 24 Cq(矩阵乘1的块数量,对应Cq计算) uint32_t mm2BlockNum; // 9 Ckv(矩阵乘2的块数量,对应Ckv计算) uint32_t mm3BlockNum; // 24 QcQr(矩阵乘3的块数量,对应QcQr计算) diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_mla_prolog_v3.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_mla_prolog_v3.py index 96981fdde2ed..b762bcb670ac 100644 --- a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_mla_prolog_v3.py +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_mla_prolog_v3.py @@ -252,6 +252,109 @@ def test_mla_prolog_v3_cache_mode(cache_mode: str): torch.npu.reset_peak_memory_stats() +@pytest.mark.parametrize("cache_mode", ["PA_BSND", "PA_NZ"]) +@torch.inference_mode() +def test_mla_prolog_v3_noncontiguous_cache_dim0(cache_mode: str): + """KV/KR cache views may have padding between physical blocks.""" + _skip_if_mla_prolog_v3_unavailable() + + token_num = 2 + head_num = 96 + he = 7168 + hcq = 1536 + hckv = 512 + d = 128 + dr = 64 + block_num = 4 + block_size = 128 + dtype = torch.bfloat16 + + token_x = torch.randn((token_num, he), dtype=dtype).npu() + weight_dq = torch_npu.npu_format_cast(torch.randn((he, hcq), dtype=dtype).npu().contiguous(), 29) + weight_uq_qr = torch_npu.npu_format_cast( + torch.randn((hcq, head_num * (d + dr)), dtype=dtype).npu().contiguous(), 29 + ) + weight_uk = torch.randn((head_num, d, hckv), dtype=dtype).npu() + weight_dkv_kr = torch_npu.npu_format_cast(torch.randn((he, hckv + dr), dtype=dtype).npu().contiguous(), 29) + rmsnorm_gamma_cq = torch.ones((hcq,), dtype=dtype).npu() + rmsnorm_gamma_ckv = torch.ones((hckv,), dtype=dtype).npu() + rope_sin = torch.randn((token_num, dr), dtype=dtype).npu() + rope_cos = torch.randn((token_num, dr), dtype=dtype).npu() + cache_index = torch.tensor([block_size + 1, 2 * block_size + 3], dtype=torch.int64).npu() + + kv_contiguous = torch.zeros((block_num, block_size, 1, hckv), dtype=dtype).npu() + kr_contiguous = torch.zeros((block_num, block_size, 1, dr), dtype=dtype).npu() + query_ref, query_rope_ref, *_ = torch.ops._C_ascend.npu_mla_prolog_v3( + token_x, + weight_dq, + weight_uq_qr, + weight_uk, + weight_dkv_kr, + rmsnorm_gamma_cq, + rmsnorm_gamma_ckv, + rope_sin, + rope_cos, + kv_contiguous, + kr_contiguous, + cache_index=cache_index, + cache_mode=cache_mode, + ) + + # Keep one physical padding block between adjacent logical KV blocks and + # two physical padding blocks between adjacent logical KR blocks. + kv_storage = torch.full((block_num * 2, block_size, 1, hckv), 7.0, dtype=dtype).npu() + kr_storage = torch.full((block_num * 3, block_size, 1, dr), -5.0, dtype=dtype).npu() + kv_cache = kv_storage[::2] + kr_cache = kr_storage[::3] + kv_cache.zero_() + kr_cache.zero_() + + assert not kv_cache.is_contiguous() + assert not kr_cache.is_contiguous() + assert kv_cache.stride(0) == 2 * block_size * hckv + assert kr_cache.stride(0) == 3 * block_size * dr + + kv_before = kv_cache.clone() + kr_before = kr_cache.clone() + kv_padding_before = kv_storage[1::2].clone() + kr_padding_1_before = kr_storage[1::3].clone() + kr_padding_2_before = kr_storage[2::3].clone() + + # Both tokens target non-zero logical blocks, so an implementation that + # ignores dim0 stride would overwrite physical padding blocks. + query, query_rope, *_ = torch.ops._C_ascend.npu_mla_prolog_v3( + token_x, + weight_dq, + weight_uq_qr, + weight_uk, + weight_dkv_kr, + rmsnorm_gamma_cq, + rmsnorm_gamma_ckv, + rope_sin, + rope_cos, + kv_cache, + kr_cache, + cache_index=cache_index, + cache_mode=cache_mode, + ) + + assert query.shape == (token_num, head_num, hckv) + assert query_rope.shape == (token_num, head_num, dr) + assert not torch.equal(kv_cache, kv_before) + assert not torch.equal(kr_cache, kr_before) + torch.testing.assert_close(query, query_ref, rtol=0, atol=0) + torch.testing.assert_close(query_rope, query_rope_ref, rtol=0, atol=0) + assert torch.equal(kv_cache, kv_contiguous) + assert torch.equal(kr_cache, kr_contiguous) + assert torch.equal(kv_storage[1::2], kv_padding_before) + assert torch.equal(kr_storage[1::3], kr_padding_1_before) + assert torch.equal(kr_storage[2::3], kr_padding_2_before) + + gc.collect() + torch.npu.empty_cache() + torch.npu.reset_peak_memory_stats() + + @torch.inference_mode() def test_mla_prolog_v3_cache_mode_bsnd(): _skip_if_mla_prolog_v3_unavailable() From cc0048791eb6515dc03ccd722b088745b911ae81 Mon Sep 17 00:00:00 2001 From: Dawn952 Date: Thu, 13 Aug 2026 15:50:32 +0800 Subject: [PATCH 24/50] [Feature][MLA] Support Kimi K3 no-RoPE MLAPO on A5 (#13507) Kimi K3 uses MLA without RoPE. This PR keeps the fused A5 decode preprocess available when MLA RoPE is disabled and adds the minimum native BF16 weight adaptation needed by the unquantized Kimi K3 checkpoint. - pass `use_mla_rope` through `DeviceOperator.mla_preprocess_only_decode`; - preserve the existing RoPE-enabled behavior; - pass `rope_sin=None` and `rope_cos=None` through the AscendC prolog wrapper for no-RoPE A5 decode, while keeping the RoPE path on `torch_npu`; - accept `UnquantizedLinearMethod` only on A5 instead of falling back to `enable_mlapo=False`; - prepare native BF16 prolog weights in `[in_features, out_features]` NZ format and use `weight_quant_mode=0` without dequant scales; - pad Kimi K3's 12 local query heads to the A5 prolog-supported 16 heads, then remove the padding before the existing attention path; - keep the W8A8 and W8A8-MXFP8 path unchanged. Yes. An unquantized Kimi K3 checkpoint can keep MLAPO enabled for A5 decode while MLA RoPE is disabled. - Targeted unit tests: `9 passed, 82 deselected`. - `git diff --check` passed. Four-node Kimi K3 service validation is intentionally pending code review of the BF16 adaptation. - vLLM version: v0.26.0 - vLLM main: https://github.com/vllm-project/vllm/commit/d02df748bf9efd99022f1a062597dc3cb3808485 --------- Signed-off-by: Dawn952 --- tests/ut/attention/a2/test_mla_v1.py | 1 + vllm_ascend/attention/mla_v1.py | 24 +++++++++++++++--------- 2 files changed, 16 insertions(+), 9 deletions(-) diff --git a/tests/ut/attention/a2/test_mla_v1.py b/tests/ut/attention/a2/test_mla_v1.py index 0d11c879ba35..ab87f4fe7245 100644 --- a/tests/ut/attention/a2/test_mla_v1.py +++ b/tests/ut/attention/a2/test_mla_v1.py @@ -1651,6 +1651,7 @@ def test_process_weights_for_fused_mlapo_a5(self, mock_format_cast, mock_get_asc self.impl.q_proj.weight_scale.data = torch.randn(128, 128, 128) self.impl.q_lora_rank = 32 self.impl._mlapo_quant_type = object + self.impl._mlapo_uses_native_weights = False from vllm_ascend.attention.mla_v1 import AscendDeviceType diff --git a/vllm_ascend/attention/mla_v1.py b/vllm_ascend/attention/mla_v1.py index 72f10ab275ba..da3e74ce817f 100644 --- a/vllm_ascend/attention/mla_v1.py +++ b/vllm_ascend/attention/mla_v1.py @@ -79,6 +79,13 @@ MLAPO_MAX_SUPPORTED_TOKENS = 1024 +def _npu_mla_prolog_v3_no_rope(**kwargs): + """Call the AscendC MLA prolog with optional RoPE inputs omitted.""" + import vllm_ascend.vllm_ascend_C # type: ignore[import-untyped] # noqa: F401, PLC0415 + + return torch.ops._C_ascend.npu_mla_prolog_v3(**kwargs) + + class AscendMLABackend(AttentionBackend): accept_output_buffer: bool = True @@ -1726,9 +1733,7 @@ def mla_preprocess_only_decode(self, hidden_states, kv_cache, attn_metadata): dequant_scale_w_dkv_kr = None else: hidden_states = hidden_states.unsqueeze(1) - quantized_x, dynamic_scale = torch_npu.npu_dynamic_mx_quant( - hidden_states, dst_type=torch.float8_e4m3fn - ) + quantized_x, dynamic_scale = torch_npu.npu_dynamic_mx_quant(hidden_states, dst_type=torch.float8_e4m3fn) dequant_scale_x = dynamic_scale.reshape(quantized_x.shape[0] * quantized_x.shape[1], -1).view( torch.float8_e8m0fnu ) @@ -1738,15 +1743,15 @@ def mla_preprocess_only_decode(self, hidden_states, kv_cache, attn_metadata): if self.use_mla_rope: cos_shape = attn_metadata.decode.cos.shape rope_shape = ( - (cos_shape[0], 1, cos_shape[-1]) - if quantized_x.dim() == 3 - else (cos_shape[0], cos_shape[-1]) + (cos_shape[0], 1, cos_shape[-1]) if quantized_x.dim() == 3 else (cos_shape[0], cos_shape[-1]) ) cos = attn_metadata.decode.cos.view(rope_shape) sin = attn_metadata.decode.sin.view(rope_shape) + prolog_op = torch_npu.npu_mla_prolog_v3 else: - cos = quantized_x.new_empty((0,), dtype=torch.bfloat16) - sin = quantized_x.new_empty((0,), dtype=torch.bfloat16) + cos = None + sin = None + prolog_op = _npu_mla_prolog_v3_no_rope cache_index = cache_index.view(bsz, -1) if quantized_x.dim() == 3 else cache_index.view(-1) cache_mode = "PA_BSND" weight_quant_mode = self.mlapo_weight_quant_mode @@ -1760,13 +1765,14 @@ def mla_preprocess_only_decode(self, hidden_states, kv_cache, attn_metadata): dequant_scale_w_dkv_kr = self.dequant_scale_w_dkv_kr cos = attn_metadata.decode.cos.view(cos_shape[0], cos_shape[-1]) sin = attn_metadata.decode.sin.view(cos_shape[0], cos_shape[-1]) + prolog_op = torch_npu.npu_mla_prolog_v3 cache_mode = "PA_NZ" if (self.fa_quant_layer or self.enable_kv_nz) else "PA_BSND" weight_quant_mode = 2 # v3 full-quant uses a per-tensor kv scale; quant_kscale is one scalar # broadcast to (1, Hckv), so slice out the single per-tensor value. quant_scale_ckv = self.quant_kscale[:, :1] if self.fa_quant_layer else None - decode_q_nope, decode_q_pe, dequant_scale_q_nope, _, _ = torch_npu.npu_mla_prolog_v3( + decode_q_nope, decode_q_pe, dequant_scale_q_nope, _, _ = prolog_op( kv_cache=decode_k_nope, kr_cache=decode_k_pe, token_x=quantized_x, From 1dc36309fceaa996c31bcecf9abc295e3b50cd83 Mon Sep 17 00:00:00 2001 From: Dawn952 Date: Tue, 25 Aug 2026 14:57:24 +0800 Subject: [PATCH 25/50] chore(ops): normalize MLA prolog source line endings Normalize the imported AscendC MLA prolog headers to LF and remove the remaining mixed indentation so the main-branch diff passes whitespace checks. Signed-off-by: Dawn952 --- csrc/attention/mla_prolog_v3/docs/api.md | 12 +- .../op_host/mla_prolog_tiling_check.h | 430 +++++++++--------- .../arch35/kernel_mla_prolog_split_m.h | 6 +- .../op_kernel/arch35/mla_prolog_comm.h | 50 +- .../op_kernel/arch35/vf/vf_comm.h | 80 ++-- .../op_kernel/arch35/vf/vf_dequant.h | 254 +++++------ .../op_kernel/arch35/vf/vf_quant_perchannel.h | 206 ++++----- .../op_kernel/mla_prolog_tiling_data.h | 150 +++--- .../mla_prolog_v3/op_kernel/mla_prolog_v3.cpp | 56 +-- 9 files changed, 622 insertions(+), 622 deletions(-) diff --git a/csrc/attention/mla_prolog_v3/docs/api.md b/csrc/attention/mla_prolog_v3/docs/api.md index 679dca4524ec..b5354208e22f 100644 --- a/csrc/attention/mla_prolog_v3/docs/api.md +++ b/csrc/attention/mla_prolog_v3/docs/api.md @@ -8,7 +8,7 @@ | aclnn | `aclnnMlaPrologV3WeightNzGetWorkspaceSize` / `aclnnMlaPrologV3WeightNz` | 支持 | | Ascend C `<<<>>>` | `mla_prolog_v3<<>>` | 支持(诊断/直调;需自备 tiling) | -各入口表达同一套 MLA 前处理融合语义:下采样 → RMSNorm → 上采样 / RoPE → 写入 KV/KR Cache(及可选量化)。 +各入口表达同一套 MLA 前处理融合语义:下采样 → RMSNorm → 上采样 / RoPE → 写入 KV/KR Cache(及可选量化)。 底层算子名为 **MlaPrologV3**,权重 `weight_dq` / `weight_uq_qr` / `weight_dkv_kr` 需以 **FRACTAL_NZ** 格式传入。 ## 2. 公共参数与约束 @@ -138,8 +138,8 @@ aclnnStatus aclnnMlaPrologV3WeightNz( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream); ``` -`GetWorkspaceSize` 完成参数校验与 executor 创建;第二段在传入 stream 上异步执行。 -`ropeSin` / `ropeCos` 同时非空时启用 RoPE,同时为空时禁用;一个空一个非空时返回参数错误。 +`GetWorkspaceSize` 完成参数校验与 executor 创建;第二段在传入 stream 上异步执行。 +`ropeSin` / `ropeCos` 同时非空时启用 RoPE,同时为空时禁用;一个空一个非空时返回参数错误。 `kvCacheRef` / `krCacheRef` 同时是输入和输出。输入、输出、workspace 和 executor 必须保持有效直到 stream 完成。 ### 3.2 调用示例 @@ -189,9 +189,9 @@ query, query_rope, dequant_scale_q_nope, query_norm, dequant_scale_q_norm = ( ) ``` -仅在 Ascend950 构建且加载 `vllm_ascend_C` + 自定义 opp 后可用。 -`rope_sin` / `rope_cos` 为必传位置参数:同时非空启用 RoPE,同时为空(`numel()==0`)禁用;不允许一空一非空。 -`token_x` rank=2 为合轴 `(T,He)`,rank=3 为 `(B,S,He)`。 +仅在 Ascend950 构建且加载 `vllm_ascend_C` + 自定义 opp 后可用。 +`rope_sin` / `rope_cos` 为必传位置参数:同时非空启用 RoPE,同时为空(`numel()==0`)禁用;不允许一空一非空。 +`token_x` rank=2 为合轴 `(T,He)`,rank=3 为 `(B,S,He)`。 `kv_cache` / `kr_cache` 原地更新;不需要的 optional 输出以空 Tensor 返回。 NZ 权重可用 `torch_npu.npu_format_cast(w.contiguous(), 29)` 转换。 diff --git a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling_check.h b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling_check.h index 82b6feedae47..ae025b4a6a46 100644 --- a/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling_check.h +++ b/csrc/attention/mla_prolog_v3/op_host/mla_prolog_tiling_check.h @@ -1,215 +1,215 @@ -/** - * Copyright (c) 2025 Huawei Technologies Co., Ltd. - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of - * CANN Open Software License Agreement Version 2.0 (the "License"). - * Please refer to the License for details. You may not use this file except in compliance with the License. - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. - * See LICENSE in the root of the software repository for the full text of the License. - */ - -/*! - * \file mla_prolog_tiling_check.h - * \brief - */ - -#ifndef MLA_PROLOG_TILING_CHECK_H -#define MLA_PROLOG_TILING_CHECK_H - -#include "mla_prolog_tiling.h" - -namespace optiling { - -constexpr uint32_t MAX_B_SIZE = 65536U; -constexpr uint32_t MAX_S1_SIZE = 65536U; -constexpr uint32_t MAX_T_SIZE = 1024U * 1024U; -constexpr uint32_t HCQ_SIZE = 1536U; -constexpr uint32_t HCKV_SIZE = 512U; -constexpr uint32_t D_SIZE = 128U; -constexpr uint32_t DR_SIZE = 64U; -constexpr uint32_t NKV_SIZE = 1U; -constexpr uint32_t MIN_BLOCK_SIZE = 16U; -constexpr uint32_t MAX_BLOCK_SIZE = 1024U; -constexpr uint32_t ALIGN_BLOCK_SIZE = 16U; -constexpr uint32_t MXFP8_BLOCK_SIZE = 32U; - -constexpr int64_t NZ_H0_SIZE = 16U; - -constexpr char TOKEN_X_NAME[] {"tokenX"}; -constexpr char WEIGHT_DQ_NAME[] {"weightDq"}; -constexpr char WEIGHT_UQ_QR_NAME[] {"weightUqQr"}; -constexpr char WEIGHT_UK_NAME[] {"weightUk"}; -constexpr char WEIGHT_DKV_KR_NAME[] {"weightDkvKr"}; -constexpr char RMSNORM_GAMMA_CQ_NAME[] {"rmsnormGammaCq"}; -constexpr char RMSNORM_GAMMA_CKV_NAME[] {"rmsnormGammaCkv"}; -constexpr char ROPE_SIN_NAME[] {"ropeSin"}; -constexpr char ROPE_COS_NAME[] {"ropeCos"}; -constexpr char CACHE_INDEX_NAME[] {"cacheIndex"}; -constexpr char KV_CACHE_NAME[] {"kvCache"}; -constexpr char KR_CACHE_NAME[] {"krCache"}; -constexpr char DEQUANT_SCALE_X_NAME[] {"dequantScaleX"}; -constexpr char DEQUANT_SCALE_W_DQ_NAME[] {"dequantScaleWDq"}; -constexpr char DEQUANT_SCALE_W_UQ_QR_NAME[] {"dequantScaleWUqQr"}; -constexpr char DEQUANT_SCALE_W_DKV_KR_NAME[] {"dequantScaleWDkvKr"}; -constexpr char QUANT_SCALE_CKV_NAME[] {"quantScaleCkv"}; -constexpr char QUANT_SCALE_CKR_NAME[] {"quantScaleCkr"}; -constexpr char SMOOTH_SCALES_CQ_NAME[] {"smoothScalesCq"}; -constexpr char ACTUAL_SEQ_LEN_NAME[] {"actualSeqLen"}; -constexpr char K_NOPE_CLIP_ALPHA_NAME[] {"kNopeClipAlpha"}; -constexpr char QUERY_NAME[] {"query"}; -constexpr char QUERY_ROPE_NAME[] {"queryRope"}; -constexpr char KV_CACHE_OUT_NAME[] {"kvCacheOut"}; -constexpr char KR_CACHE_OUT_NAME[] {"krCacheOut"}; -constexpr char DEQUANT_SCALE_Q_NOPE_NAME[] {"dequantScaleQNope"}; -constexpr char QUERY_NORM_NAME[] {"queryNorm"}; -constexpr char DEQUANT_SCALE_Q_NORM_NAME[] {"dequantScaleQNorm"}; - -constexpr uint32_t PARAM_MAP_INIT_RESERVE_NUM = 28; // 预分配所有key的个数,避免使用时动态扩容 - -struct ParamInfo { - ParamInfo() = default; - explicit ParamInfo(const ParamInfo &) = default; - ParamInfo &operator=(const ParamInfo &) = default; - explicit ParamInfo(ParamInfo &&) = default; - ParamInfo &operator=(ParamInfo &&other) = default; - ~ParamInfo() = default; - ParamInfo(const gert::CompileTimeTensorDesc *actualDesc, const gert::StorageShape *actualShape) { - if (actualDesc != nullptr && actualShape != nullptr) { - isValid = true; - dtype = actualDesc->GetDataType(); - format = static_cast(ge::GetPrimaryFormat(actualDesc->GetStorageFormat())); - auto &&actualStorageShape = actualShape->GetStorageShape(); - dimNum = actualStorageShape.GetDimNum(); - this->shape.reserve(dimNum); - for (size_t i = 0; i < dimNum; i++) { - this->shape.emplace_back(actualStorageShape.GetDim(i)); - } - } - } - explicit ParamInfo(const BaseParaInfo &info) : ParamInfo(info.desc, info.shape) {} - explicit ParamInfo(const std::vector &expectedShape) - { - isValid = true; - format = ge::FORMAT_ND; - dimNum = expectedShape.size(); - this->shape.reserve(dimNum); - for (size_t i = 0; i < dimNum; i++) { - this->shape.emplace_back(static_cast(expectedShape[i])); - } - } - - bool operator == (const ParamInfo &other) const { - if (!isValid && !other.isValid) { - return true; - } - static const std::set ndFormats{ge::FORMAT_ND, ge::FORMAT_NCHW}; - if ((ndFormats.find(format) == ndFormats.end() || ndFormats.find(other.format) == ndFormats.end()) && - format != other.format) { - return false; - } - return (isValid == other.isValid && dtype == other.dtype && - dimNum == other.dimNum && shape == other.shape); - } - bool operator != (const ParamInfo &other) const { - return !(*this == other); - } - - bool isValid {}; - ge::DataType dtype {ge::DT_MAX}; - ge::Format format {ge::FORMAT_MAX}; - size_t dimNum {}; - std::vector shape; -}; - -using ParamInfoMap = std::unordered_map; - -class MlaPrologTilingCheck { -public: - MlaPrologTilingCheck(const MlaPrologContext &context, const MlaPrologBaseShapeInfo &baseShapeInfo, - const MlaPrologScenarioInfo &scenarioInfo) - : context_(context), baseShapeInfo_(baseShapeInfo), scenarioInfo_(scenarioInfo) {} - ge::graphStatus CheckSingleRequiredParam() const; - ge::graphStatus CheckCacheMode() const; - ge::graphStatus CheckQuantMode() const; - ge::graphStatus CheckDims() const; - ge::graphStatus CheckParamByScenario(); - ge::graphStatus CheckSpecialScenarioParamShape(); - ge::graphStatus CheckCkvkrRepoMode(); - ge::graphStatus CheckCacheIndexDim(); - ge::graphStatus CheckScenarParam(); - ge::graphStatus CheckAttrs() const; - - NpuArch GetCurNpuArch() const; - -private: - bool CheckAttrsNotNull() const; - bool CheckAttrsRange() const; - bool CheckCacheModeParamShape() const; - ge::graphStatus CheckHcqSize() const; - ge::graphStatus CheckDSize() const; - ge::graphStatus CheckDtileSize() const; - // ==================================单参数校验================================== - bool IsSingleParamValid(const BaseParaInfo ¶m, const std::string ¶mName, - const std::set &expectedDtype, - const std::set &expectedFormat, - const std::set &expectedDimNum) const; - bool CheckTokenX() const; - bool CheckWDq() const; - bool CheckWDkvKr() const; - bool CheckWUqQr() const; - bool CheckWUk() const; - bool CheckRmsnormGammaCkv() const; - bool CheckRmsnormGammaCq() const; - bool CheckRopeCos() const; - bool CheckRopeSin() const; - bool CheckCacheIndex() const; - bool CheckKvCache() const; - bool CheckKrCache() const; - bool CheckActSeqLen() const; - // ==================================单参数校验================================== - - // =================================全量参数校验================================= - void GenExpectedParamInfo(); - void FillCommonParamInfo(); - void FillRequiredParamShapeWithDims(); - void FillOptionalOutputParamShapeWithDims(); - void FillOptionalOutputParamShapeWithDimsV2(); - void FillOptionalOutputParamShapeWithDimsV3(); - void FillScenarioParamInfo(); - void FillQueryNormScaleShape(); - void FillQueryNormDtypes(); - void FillTokenAndQueryShapes(); - void FillWeightAndNormShapes(); - void FillCacheShapes(); - void CheckRepoMode(bool isPertile, ge::graphStatus &isCorrect); - void CheckQueryQuantMode(bool isPertensor, ge::graphStatus &isCorrect); - void FillNonQuantParamInfo(); - void FillPartialQuantParamInfo(); - void FillPartialKVQuantParamInfo(); - void FillPartialKVPertileQuantParamInfo(); - void FillFullQuantParamInfo(); - void FillFullKVQuantParamInfo(); - void FillFullKVPertileQuantParamInfo(); - void FillMxfp8FullQuantParamInfo(); - void FillMxfp8FullKVQuantParamInfo(); - void FillMxfp8FullKVPertileParamInfo(); - void FillFP8FullQuantParamInfo(); - void FillFP8FullKVQuantParamInfo(); - void FillHIF8FullQuantParamInfo(); - void FillHIF8FullKVQuantParamInfo(); - void FillFP8FullKVPertileQuantParamInfo(); - void FillHIF8FullKVPertileQuantParamInfo(); - - void GenActualParamInfo(); - // =================================全量参数校验================================= - - const MlaPrologContext &context_; - const MlaPrologBaseShapeInfo &baseShapeInfo_; - const MlaPrologScenarioInfo &scenarioInfo_; - ParamInfoMap expectedParamInfo_ = ParamInfoMap(PARAM_MAP_INIT_RESERVE_NUM); - ParamInfoMap actualParamInfo_ = ParamInfoMap(PARAM_MAP_INIT_RESERVE_NUM); -}; - -} // namespace optiling - -#endif \ No newline at end of file +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_tiling_check.h + * \brief + */ + +#ifndef MLA_PROLOG_TILING_CHECK_H +#define MLA_PROLOG_TILING_CHECK_H + +#include "mla_prolog_tiling.h" + +namespace optiling { + +constexpr uint32_t MAX_B_SIZE = 65536U; +constexpr uint32_t MAX_S1_SIZE = 65536U; +constexpr uint32_t MAX_T_SIZE = 1024U * 1024U; +constexpr uint32_t HCQ_SIZE = 1536U; +constexpr uint32_t HCKV_SIZE = 512U; +constexpr uint32_t D_SIZE = 128U; +constexpr uint32_t DR_SIZE = 64U; +constexpr uint32_t NKV_SIZE = 1U; +constexpr uint32_t MIN_BLOCK_SIZE = 16U; +constexpr uint32_t MAX_BLOCK_SIZE = 1024U; +constexpr uint32_t ALIGN_BLOCK_SIZE = 16U; +constexpr uint32_t MXFP8_BLOCK_SIZE = 32U; + +constexpr int64_t NZ_H0_SIZE = 16U; + +constexpr char TOKEN_X_NAME[] {"tokenX"}; +constexpr char WEIGHT_DQ_NAME[] {"weightDq"}; +constexpr char WEIGHT_UQ_QR_NAME[] {"weightUqQr"}; +constexpr char WEIGHT_UK_NAME[] {"weightUk"}; +constexpr char WEIGHT_DKV_KR_NAME[] {"weightDkvKr"}; +constexpr char RMSNORM_GAMMA_CQ_NAME[] {"rmsnormGammaCq"}; +constexpr char RMSNORM_GAMMA_CKV_NAME[] {"rmsnormGammaCkv"}; +constexpr char ROPE_SIN_NAME[] {"ropeSin"}; +constexpr char ROPE_COS_NAME[] {"ropeCos"}; +constexpr char CACHE_INDEX_NAME[] {"cacheIndex"}; +constexpr char KV_CACHE_NAME[] {"kvCache"}; +constexpr char KR_CACHE_NAME[] {"krCache"}; +constexpr char DEQUANT_SCALE_X_NAME[] {"dequantScaleX"}; +constexpr char DEQUANT_SCALE_W_DQ_NAME[] {"dequantScaleWDq"}; +constexpr char DEQUANT_SCALE_W_UQ_QR_NAME[] {"dequantScaleWUqQr"}; +constexpr char DEQUANT_SCALE_W_DKV_KR_NAME[] {"dequantScaleWDkvKr"}; +constexpr char QUANT_SCALE_CKV_NAME[] {"quantScaleCkv"}; +constexpr char QUANT_SCALE_CKR_NAME[] {"quantScaleCkr"}; +constexpr char SMOOTH_SCALES_CQ_NAME[] {"smoothScalesCq"}; +constexpr char ACTUAL_SEQ_LEN_NAME[] {"actualSeqLen"}; +constexpr char K_NOPE_CLIP_ALPHA_NAME[] {"kNopeClipAlpha"}; +constexpr char QUERY_NAME[] {"query"}; +constexpr char QUERY_ROPE_NAME[] {"queryRope"}; +constexpr char KV_CACHE_OUT_NAME[] {"kvCacheOut"}; +constexpr char KR_CACHE_OUT_NAME[] {"krCacheOut"}; +constexpr char DEQUANT_SCALE_Q_NOPE_NAME[] {"dequantScaleQNope"}; +constexpr char QUERY_NORM_NAME[] {"queryNorm"}; +constexpr char DEQUANT_SCALE_Q_NORM_NAME[] {"dequantScaleQNorm"}; + +constexpr uint32_t PARAM_MAP_INIT_RESERVE_NUM = 28; // 预分配所有key的个数,避免使用时动态扩容 + +struct ParamInfo { + ParamInfo() = default; + explicit ParamInfo(const ParamInfo &) = default; + ParamInfo &operator=(const ParamInfo &) = default; + explicit ParamInfo(ParamInfo &&) = default; + ParamInfo &operator=(ParamInfo &&other) = default; + ~ParamInfo() = default; + ParamInfo(const gert::CompileTimeTensorDesc *actualDesc, const gert::StorageShape *actualShape) { + if (actualDesc != nullptr && actualShape != nullptr) { + isValid = true; + dtype = actualDesc->GetDataType(); + format = static_cast(ge::GetPrimaryFormat(actualDesc->GetStorageFormat())); + auto &&actualStorageShape = actualShape->GetStorageShape(); + dimNum = actualStorageShape.GetDimNum(); + this->shape.reserve(dimNum); + for (size_t i = 0; i < dimNum; i++) { + this->shape.emplace_back(actualStorageShape.GetDim(i)); + } + } + } + explicit ParamInfo(const BaseParaInfo &info) : ParamInfo(info.desc, info.shape) {} + explicit ParamInfo(const std::vector &expectedShape) + { + isValid = true; + format = ge::FORMAT_ND; + dimNum = expectedShape.size(); + this->shape.reserve(dimNum); + for (size_t i = 0; i < dimNum; i++) { + this->shape.emplace_back(static_cast(expectedShape[i])); + } + } + + bool operator == (const ParamInfo &other) const { + if (!isValid && !other.isValid) { + return true; + } + static const std::set ndFormats{ge::FORMAT_ND, ge::FORMAT_NCHW}; + if ((ndFormats.find(format) == ndFormats.end() || ndFormats.find(other.format) == ndFormats.end()) && + format != other.format) { + return false; + } + return (isValid == other.isValid && dtype == other.dtype && + dimNum == other.dimNum && shape == other.shape); + } + bool operator != (const ParamInfo &other) const { + return !(*this == other); + } + + bool isValid {}; + ge::DataType dtype {ge::DT_MAX}; + ge::Format format {ge::FORMAT_MAX}; + size_t dimNum {}; + std::vector shape; +}; + +using ParamInfoMap = std::unordered_map; + +class MlaPrologTilingCheck { +public: + MlaPrologTilingCheck(const MlaPrologContext &context, const MlaPrologBaseShapeInfo &baseShapeInfo, + const MlaPrologScenarioInfo &scenarioInfo) + : context_(context), baseShapeInfo_(baseShapeInfo), scenarioInfo_(scenarioInfo) {} + ge::graphStatus CheckSingleRequiredParam() const; + ge::graphStatus CheckCacheMode() const; + ge::graphStatus CheckQuantMode() const; + ge::graphStatus CheckDims() const; + ge::graphStatus CheckParamByScenario(); + ge::graphStatus CheckSpecialScenarioParamShape(); + ge::graphStatus CheckCkvkrRepoMode(); + ge::graphStatus CheckCacheIndexDim(); + ge::graphStatus CheckScenarParam(); + ge::graphStatus CheckAttrs() const; + + NpuArch GetCurNpuArch() const; + +private: + bool CheckAttrsNotNull() const; + bool CheckAttrsRange() const; + bool CheckCacheModeParamShape() const; + ge::graphStatus CheckHcqSize() const; + ge::graphStatus CheckDSize() const; + ge::graphStatus CheckDtileSize() const; + // ==================================单参数校验================================== + bool IsSingleParamValid(const BaseParaInfo ¶m, const std::string ¶mName, + const std::set &expectedDtype, + const std::set &expectedFormat, + const std::set &expectedDimNum) const; + bool CheckTokenX() const; + bool CheckWDq() const; + bool CheckWDkvKr() const; + bool CheckWUqQr() const; + bool CheckWUk() const; + bool CheckRmsnormGammaCkv() const; + bool CheckRmsnormGammaCq() const; + bool CheckRopeCos() const; + bool CheckRopeSin() const; + bool CheckCacheIndex() const; + bool CheckKvCache() const; + bool CheckKrCache() const; + bool CheckActSeqLen() const; + // ==================================单参数校验================================== + + // =================================全量参数校验================================= + void GenExpectedParamInfo(); + void FillCommonParamInfo(); + void FillRequiredParamShapeWithDims(); + void FillOptionalOutputParamShapeWithDims(); + void FillOptionalOutputParamShapeWithDimsV2(); + void FillOptionalOutputParamShapeWithDimsV3(); + void FillScenarioParamInfo(); + void FillQueryNormScaleShape(); + void FillQueryNormDtypes(); + void FillTokenAndQueryShapes(); + void FillWeightAndNormShapes(); + void FillCacheShapes(); + void CheckRepoMode(bool isPertile, ge::graphStatus &isCorrect); + void CheckQueryQuantMode(bool isPertensor, ge::graphStatus &isCorrect); + void FillNonQuantParamInfo(); + void FillPartialQuantParamInfo(); + void FillPartialKVQuantParamInfo(); + void FillPartialKVPertileQuantParamInfo(); + void FillFullQuantParamInfo(); + void FillFullKVQuantParamInfo(); + void FillFullKVPertileQuantParamInfo(); + void FillMxfp8FullQuantParamInfo(); + void FillMxfp8FullKVQuantParamInfo(); + void FillMxfp8FullKVPertileParamInfo(); + void FillFP8FullQuantParamInfo(); + void FillFP8FullKVQuantParamInfo(); + void FillHIF8FullQuantParamInfo(); + void FillHIF8FullKVQuantParamInfo(); + void FillFP8FullKVPertileQuantParamInfo(); + void FillHIF8FullKVPertileQuantParamInfo(); + + void GenActualParamInfo(); + // =================================全量参数校验================================= + + const MlaPrologContext &context_; + const MlaPrologBaseShapeInfo &baseShapeInfo_; + const MlaPrologScenarioInfo &scenarioInfo_; + ParamInfoMap expectedParamInfo_ = ParamInfoMap(PARAM_MAP_INIT_RESERVE_NUM); + ParamInfoMap actualParamInfo_ = ParamInfoMap(PARAM_MAP_INIT_RESERVE_NUM); +}; + +} // namespace optiling + +#endif diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_m.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_m.h index 30bd778bf879..77c2797140e6 100644 --- a/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_m.h +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/kernel_mla_prolog_split_m.h @@ -613,7 +613,7 @@ __aicore__ inline void MlaPrologV3SplitM::WorkspaceInit(__gm__ uint8_t *w baseParams_->mm2BlockNum * sizeof(mmCkvKrOutputType); mmQcQrResGm_.SetGlobalBuffer((__gm__ mmQcQrOutputType *)(workspace + workspaceOffset)); // aicOffset.qcQrResOffset - + if constexpr (IsFullQuantMode()) { workspaceOffset += static_cast(baseParams_->stepBatchSize) * static_cast(baseParams_->headSizeQc + baseParams_->headSizeQr) * @@ -1098,7 +1098,7 @@ __aicore__ inline void MlaPrologV3SplitM::CopyGlobalParams() DataCopy(rmsnormGammaCkvLocal_, rmsnormGammaCkvGm_, baseParams_->headSizeCkv); // quantScaleCkv - + if constexpr (IsFullQuantMode() && !isPertile) { if constexpr (std::is_same::value || (std::is_same::value && !isFp8E8m0)) { @@ -1163,7 +1163,7 @@ __aicore__ inline void MlaPrologV3SplitM::RmsNormCq(int64_t tokenIndex, i baseParams_->qcQrScale, baseParams_->isQcQrScaleEnable}; - + if constexpr (IsFullQuantMode()) { RmsNormDynamicQuant(outputLocal, dequantScaleQcQr[scaleOffset], diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/mla_prolog_comm.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/mla_prolog_comm.h index b7ca0ccf90af..9fe1c5fd9d5f 100644 --- a/csrc/attention/mla_prolog_v3/op_kernel/arch35/mla_prolog_comm.h +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/mla_prolog_comm.h @@ -180,7 +180,7 @@ constexpr uint32_t L0C_PP_SIZE = 64 * 1024; /* - 非量化 半量化(kv非量化) 半量化(kv量化) int8全量化(kv非量化) int8全量化(kv量化) 半量化(kv per-tile量化) int8全量化(kv per-tile量化) Mxfp8量化(kv非量化) Mxfp8量化(kv量化) Mxfp8量化(kv per-tile量化) fp8全量化(kv非量化) fp8全量化(kv量化) hif8全量化(kv非量化) hif8全量化(kv量化) + 非量化 半量化(kv非量化) 半量化(kv量化) int8全量化(kv非量化) int8全量化(kv量化) 半量化(kv per-tile量化) int8全量化(kv per-tile量化) Mxfp8量化(kv非量化) Mxfp8量化(kv量化) Mxfp8量化(kv per-tile量化) fp8全量化(kv非量化) fp8全量化(kv量化) hif8全量化(kv非量化) hif8全量化(kv量化) cacheMode PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/BSND/TND PA_BSND/BSND/TND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/BSND/TND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND PA_BSND/PA_BLK_BSND /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /PA_NZ/PA_BLK_NZ /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND /BSND/TND @@ -191,37 +191,37 @@ constexpr uint32_t L0C_PP_SIZE = 64 * 1024; WdqType(复用mmInputType) bfloat16_t bfloat16_t bfloat16_t int8_t int8_t bfloat16_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t WuqqrType(复用mmQcQrInputType) bfloat16_t int8_t int8_t int8_t int8_t int8_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t WukType(复用mmQnInputType) bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t - WdkvkrType(复用mmInputType) bfloat16_t bfloat16_t bfloat16_t int8_t int8_t bfloat16_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t - rmsNormGammaType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t - gammaCkvType(复用rmsNormGammaType)bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t - ropeSinCosType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t - cosType(复用ropeSinCosType) bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + WdkvkrType(复用mmInputType) bfloat16_t bfloat16_t bfloat16_t int8_t int8_t bfloat16_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + rmsNormGammaType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + gammaCkvType(复用rmsNormGammaType)bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + ropeSinCosType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + cosType(复用ropeSinCosType) bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t cacheIndexType int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t int64_t - kvCacheType bfloat16_t bfloat16_t int8_t bfloat16_t int8_t int8_t int8_t bfloat16_t fp8_e4m3fn_t fp8_e4m3fn_t bfloat16_t fp8_e4m3fn_t bfloat16_t hifloat8_t - krCacheType bfloat16_t bfloat16_t int8_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + kvCacheType bfloat16_t bfloat16_t int8_t bfloat16_t int8_t int8_t int8_t bfloat16_t fp8_e4m3fn_t fp8_e4m3fn_t bfloat16_t fp8_e4m3fn_t bfloat16_t hifloat8_t + krCacheType bfloat16_t bfloat16_t int8_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t deqScaleXType / / / float float / float fp8_e8m0_t fp8_e8m0_t fp8_e8m0_t float float float float - deqScaleWdqType / / / float float / float fp8_e8m0_t fp8_e8m0_t fp8_e8m0_t float float float float - deqScaleWuqqrType / float float float float float float fp8_e8m0_t fp8_e8m0_t fp8_e8m0_t float float float float + deqScaleWdqType / / / float float / float fp8_e8m0_t fp8_e8m0_t fp8_e8m0_t float float float float + deqScaleWuqqrType / float float float float float float fp8_e8m0_t fp8_e8m0_t fp8_e8m0_t float float float float deqScaleWdkvkrType / / / float float / float fp8_e8m0_t fp8_e8m0_t fp8_e8m0_t float float float float quantScaleCkvType / / float / float / / / float / / float / float quantScaleCkrType / / float / / / / / / / / / / / smoothScaleCqType / float float float float float float / / / float float float float - queryOutputType bfloat16_t bfloat16_t int8_t bfloat16_t int8_t bfloat16_t bfloat16_t bfloat16_t fp8_e4m3fn_t bfloat16_t bfloat16_t fp8_e4m3fn_t bfloat16_t hifloat8_t - ropeOutputType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + queryOutputType bfloat16_t bfloat16_t int8_t bfloat16_t int8_t bfloat16_t bfloat16_t bfloat16_t fp8_e4m3fn_t bfloat16_t bfloat16_t fp8_e4m3fn_t bfloat16_t hifloat8_t + ropeOutputType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t dequantScaleQNopeType / / / / float / / / float / / float / float - queryNormType(复用mmQcQrInputType)bfloat16_t int8_t int8_t int8_t int8_t int8_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t - dequantScaleQNormType / float float float float float float fp8_e8m0_t fp8_e8m0_t fp8_e8m0_t float float float float - mmInputType bfloat16_t bfloat16_t bfloat16_t int8_t int8_t bfloat16_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t - mmCqOutputType bfloat16_t bfloat16_t bfloat16_t int32_t int32_t bfloat16_t int32_t float float float float float float float - mmCkvKrInputType(复用mmInputType) bfloat16_t bfloat16_t bfloat16_t int8_t int8_t bfloat16_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t - mmCkvKrOutputType bfloat16_t bfloat16_t bfloat16_t int32_t int32_t bfloat16_t int32_t float float float float float float float - mmQcQrInputType bfloat16_t int8_t int8_t int8_t int8_t int8_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t - mmQcQrOutputType bfloat16_t bfloat16_t bfloat16_t int32_t int32_t bfloat16_t int32_t float float float float float float float - mmQnInputType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t - mmQnOutputType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t - rmsNormComputType float float float float float float float float float float float float float float - rmsNormCqOutputType bfloat16_t int8_t int8_t int8_t int8_t int8_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t - rmsNormCkvOutputType bfloat16_t bfloat16_t int8_t bfloat16_t int8_t int8_t int8_t bfloat16_t fp8_e4m3fn_t fp8_e4m3fn_t bfloat16_t fp8_e4m3fn_t bfloat16_t hifloat8_t + queryNormType(复用mmQcQrInputType)bfloat16_t int8_t int8_t int8_t int8_t int8_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + dequantScaleQNormType / float float float float float float fp8_e8m0_t fp8_e8m0_t fp8_e8m0_t float float float float + mmInputType bfloat16_t bfloat16_t bfloat16_t int8_t int8_t bfloat16_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + mmCqOutputType bfloat16_t bfloat16_t bfloat16_t int32_t int32_t bfloat16_t int32_t float float float float float float float + mmCkvKrInputType(复用mmInputType) bfloat16_t bfloat16_t bfloat16_t int8_t int8_t bfloat16_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + mmCkvKrOutputType bfloat16_t bfloat16_t bfloat16_t int32_t int32_t bfloat16_t int32_t float float float float float float float + mmQcQrInputType bfloat16_t int8_t int8_t int8_t int8_t int8_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + mmQcQrOutputType bfloat16_t bfloat16_t bfloat16_t int32_t int32_t bfloat16_t int32_t float float float float float float float + mmQnInputType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + mmQnOutputType bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t bfloat16_t + rmsNormComputType float float float float float float float float float float float float float float + rmsNormCqOutputType bfloat16_t int8_t int8_t int8_t int8_t int8_t int8_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t fp8_e4m3fn_t hifloat8_t hifloat8_t + rmsNormCkvOutputType bfloat16_t bfloat16_t int8_t bfloat16_t int8_t int8_t int8_t bfloat16_t fp8_e4m3fn_t fp8_e4m3fn_t bfloat16_t fp8_e4m3fn_t bfloat16_t hifloat8_t ropeComputType float float float float float float float float float float float float float float */ diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_comm.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_comm.h index 10e38dc09f64..be21b1046531 100644 --- a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_comm.h +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_comm.h @@ -1,40 +1,40 @@ -/** - * Copyright (c) 2025 Huawei Technologies Co., Ltd. - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of - * CANN Open Software License Agreement Version 2.0 (the "License"). - * Please refer to the License for details. You may not use this file except in compliance with the License. - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. - * See LICENSE in the root of the software repository for the full text of the License. - */ - -/*! - * \file vf_comm.h - * \brief VF comm - */ - -#ifndef VF_COMM_H -#define VF_COMM_H - -#include "kernel_tensor.h" - -namespace MlaProlog { -constexpr uint32_t ROPE_VF_COL = 64; - -constexpr MicroAPI::CastTrait castTraitB162B32 = { - MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::UNKNOWN, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::UNKNOWN, -}; - -constexpr MicroAPI::CastTrait castTraitB322B16 = { - MicroAPI::RegLayout::ZERO, - MicroAPI::SatMode::NO_SAT, - MicroAPI::MaskMergeMode::ZEROING, - RoundMode::CAST_RINT, -}; - -} // namespace MlaProlog - -#endif // VF_COMM_H +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file vf_comm.h + * \brief VF comm + */ + +#ifndef VF_COMM_H +#define VF_COMM_H + +#include "kernel_tensor.h" + +namespace MlaProlog { +constexpr uint32_t ROPE_VF_COL = 64; + +constexpr MicroAPI::CastTrait castTraitB162B32 = { + MicroAPI::RegLayout::ZERO, + MicroAPI::SatMode::UNKNOWN, + MicroAPI::MaskMergeMode::ZEROING, + RoundMode::UNKNOWN, +}; + +constexpr MicroAPI::CastTrait castTraitB322B16 = { + MicroAPI::RegLayout::ZERO, + MicroAPI::SatMode::NO_SAT, + MicroAPI::MaskMergeMode::ZEROING, + RoundMode::CAST_RINT, +}; + +} // namespace MlaProlog + +#endif // VF_COMM_H diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_dequant.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_dequant.h index 61c65f370bac..024fc6aa8820 100644 --- a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_dequant.h +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_dequant.h @@ -1,127 +1,127 @@ -/** - * Copyright (c) 2026 Huawei Technologies Co., Ltd. - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of - * CANN Open Software License Agreement Version 2.0 (the "License"). - * Please refer to the License for details. You may not use this file except in compliance with the License. - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. - * See LICENSE in the root of the software repository for the full text of the License. - */ - -/*! - * \file vf_dequant.h - * \brief - */ - -#ifndef VF_DEQUANT_H -#define VF_DEQUANT_H -#include "kernel_tensor.h" - -namespace MlaProlog { - -template -__simd_vf__ void DequantVFImpl(__ubuf__ float *yAddr, __ubuf__ T *xAddr, __ubuf__ float *scalePerChannelAddr, - __ubuf__ float *scalePerTokenAddr, uint32_t floatRepSize, uint32_t fp32BlockElementNum, - uint32_t dLoops, uint32_t dTail, uint32_t dTailLoop, uint32_t row, uint32_t col, - uint32_t stride) -{ - constexpr static AscendC::MicroAPI::CastTrait castTraitInt32ToFp32 = { - AscendC::MicroAPI::RegLayout::UNKNOWN, AscendC::MicroAPI::SatMode::NO_SAT, - AscendC::MicroAPI::MaskMergeMode::ZEROING, AscendC::RoundMode::CAST_RINT}; - - AscendC::MicroAPI::RegTensor vregInput; - AscendC::MicroAPI::RegTensor vregScalePerChannel; - AscendC::MicroAPI::RegTensor vregScalePerToken; - AscendC::MicroAPI::RegTensor vregInputFp32; // cast成float之后的vregInput - AscendC::MicroAPI::MaskReg fullMask = AscendC::MicroAPI::CreateMask(); - AscendC::MicroAPI::MaskReg tailMask; - tailMask = AscendC::MicroAPI::UpdateMask(dTail); - - uint32_t colOffset = 0; - uint32_t rowOffset = 0; - uint32_t scaleOffset = 0; - for (uint32_t j = 0; j < dLoops; j++) { - AscendC::MicroAPI::LoadAlign(vregScalePerChannel, - scalePerChannelAddr + colOffset); - rowOffset = 0; - scaleOffset = 0; - for (uint32_t i = 0; i < row; i++) { - AscendC::MicroAPI::LoadAlign( - vregScalePerToken, scalePerTokenAddr + scaleOffset); - if constexpr (!std::is_same::value) { - AscendC::MicroAPI::LoadAlign( - vregInput, xAddr + colOffset + rowOffset); - AscendC::MicroAPI::Cast( - vregInputFp32, vregInput, fullMask); - } else { - AscendC::MicroAPI::LoadAlign(vregInputFp32, - xAddr + colOffset + rowOffset); - } - AscendC::MicroAPI::Mul(vregInputFp32, vregInputFp32, vregScalePerChannel, fullMask); - AscendC::MicroAPI::Mul(vregInputFp32, vregInputFp32, vregScalePerToken, fullMask); - AscendC::MicroAPI::StoreAlign( - yAddr + colOffset + rowOffset, vregInputFp32, fullMask); - rowOffset += stride; - scaleOffset += fp32BlockElementNum; - } - colOffset += floatRepSize; - } - - if (dTailLoop > 0) { - rowOffset = 0; - scaleOffset = 0; - AscendC::MicroAPI::LoadAlign( - vregScalePerChannel, scalePerChannelAddr + dLoops * floatRepSize); - for (uint32_t i = 0; i < row; i++) { - AscendC::MicroAPI::LoadAlign( - vregScalePerToken, scalePerTokenAddr + scaleOffset); - if constexpr (!std::is_same::value) { - AscendC::MicroAPI::LoadAlign( - vregInput, xAddr + dLoops * floatRepSize + rowOffset); - AscendC::MicroAPI::Cast( - vregInputFp32, vregInput, tailMask); - } else { - AscendC::MicroAPI::LoadAlign( - vregInputFp32, xAddr + dLoops * floatRepSize + rowOffset); - } - AscendC::MicroAPI::Mul(vregInputFp32, vregInputFp32, vregScalePerChannel, tailMask); - AscendC::MicroAPI::Mul(vregInputFp32, vregInputFp32, vregScalePerToken, tailMask); - AscendC::MicroAPI::StoreAlign( - yAddr + dLoops * floatRepSize + rowOffset, vregInputFp32, tailMask); - rowOffset += stride; - scaleOffset += fp32BlockElementNum; - } - } -} - -/** - * @brief DequantVf 对输入做per-token叠加per-channel的反量化, INT32 ---> FP32. - * @param outputLocal 输出tensor [row, col] - * @param inputLocal 输入tensor [row, col] - * @param scalePerChannelLocal 输入tensor [1, col] - * @param scalePerTokenLocal 输入tensor [row, 8] - * @param row 待处理的行数 - * @param col 待处理的列数 - * @param stride 待处理数据一行的真实长度 - */ -template -__aicore__ inline void DequantVf(const LocalTensor &outputLocal, const LocalTensor &inputLocal, - const LocalTensor &scalePerChannelLocal, - const LocalTensor &scalePerTokenLocal, - uint32_t row, uint32_t col, uint32_t stride) -{ - __ubuf__ float *outputUb = (__ubuf__ float *)outputLocal.GetPhyAddr(); - __ubuf__ T *inputUb = (__ubuf__ T *)inputLocal.GetPhyAddr(); - __ubuf__ float *scalePerChannelLocalUb = (__ubuf__ float *)scalePerChannelLocal.GetPhyAddr(); - __ubuf__ float *scalePerTokenLocalUb = (__ubuf__ float *)scalePerTokenLocal.GetPhyAddr(); - - const uint32_t floatRepSize = 64; // 一个寄存器能够存放64个FP32 - const uint32_t fp32BlockElementNum = 8; - uint32_t dLoops = col / floatRepSize; - uint32_t dTail = col % floatRepSize; - uint32_t dTailLoop = dTail > 0 ? 1 : 0; - DequantVFImpl(outputUb, inputUb, scalePerChannelLocalUb, scalePerTokenLocalUb, floatRepSize, fp32BlockElementNum, - dLoops, dTail, dTailLoop, row, col, stride); -} -} // namespace MlaProlog -#endif +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file vf_dequant.h + * \brief + */ + +#ifndef VF_DEQUANT_H +#define VF_DEQUANT_H +#include "kernel_tensor.h" + +namespace MlaProlog { + +template +__simd_vf__ void DequantVFImpl(__ubuf__ float *yAddr, __ubuf__ T *xAddr, __ubuf__ float *scalePerChannelAddr, + __ubuf__ float *scalePerTokenAddr, uint32_t floatRepSize, uint32_t fp32BlockElementNum, + uint32_t dLoops, uint32_t dTail, uint32_t dTailLoop, uint32_t row, uint32_t col, + uint32_t stride) +{ + constexpr static AscendC::MicroAPI::CastTrait castTraitInt32ToFp32 = { + AscendC::MicroAPI::RegLayout::UNKNOWN, AscendC::MicroAPI::SatMode::NO_SAT, + AscendC::MicroAPI::MaskMergeMode::ZEROING, AscendC::RoundMode::CAST_RINT}; + + AscendC::MicroAPI::RegTensor vregInput; + AscendC::MicroAPI::RegTensor vregScalePerChannel; + AscendC::MicroAPI::RegTensor vregScalePerToken; + AscendC::MicroAPI::RegTensor vregInputFp32; // cast成float之后的vregInput + AscendC::MicroAPI::MaskReg fullMask = AscendC::MicroAPI::CreateMask(); + AscendC::MicroAPI::MaskReg tailMask; + tailMask = AscendC::MicroAPI::UpdateMask(dTail); + + uint32_t colOffset = 0; + uint32_t rowOffset = 0; + uint32_t scaleOffset = 0; + for (uint32_t j = 0; j < dLoops; j++) { + AscendC::MicroAPI::LoadAlign(vregScalePerChannel, + scalePerChannelAddr + colOffset); + rowOffset = 0; + scaleOffset = 0; + for (uint32_t i = 0; i < row; i++) { + AscendC::MicroAPI::LoadAlign( + vregScalePerToken, scalePerTokenAddr + scaleOffset); + if constexpr (!std::is_same::value) { + AscendC::MicroAPI::LoadAlign( + vregInput, xAddr + colOffset + rowOffset); + AscendC::MicroAPI::Cast( + vregInputFp32, vregInput, fullMask); + } else { + AscendC::MicroAPI::LoadAlign(vregInputFp32, + xAddr + colOffset + rowOffset); + } + AscendC::MicroAPI::Mul(vregInputFp32, vregInputFp32, vregScalePerChannel, fullMask); + AscendC::MicroAPI::Mul(vregInputFp32, vregInputFp32, vregScalePerToken, fullMask); + AscendC::MicroAPI::StoreAlign( + yAddr + colOffset + rowOffset, vregInputFp32, fullMask); + rowOffset += stride; + scaleOffset += fp32BlockElementNum; + } + colOffset += floatRepSize; + } + + if (dTailLoop > 0) { + rowOffset = 0; + scaleOffset = 0; + AscendC::MicroAPI::LoadAlign( + vregScalePerChannel, scalePerChannelAddr + dLoops * floatRepSize); + for (uint32_t i = 0; i < row; i++) { + AscendC::MicroAPI::LoadAlign( + vregScalePerToken, scalePerTokenAddr + scaleOffset); + if constexpr (!std::is_same::value) { + AscendC::MicroAPI::LoadAlign( + vregInput, xAddr + dLoops * floatRepSize + rowOffset); + AscendC::MicroAPI::Cast( + vregInputFp32, vregInput, tailMask); + } else { + AscendC::MicroAPI::LoadAlign( + vregInputFp32, xAddr + dLoops * floatRepSize + rowOffset); + } + AscendC::MicroAPI::Mul(vregInputFp32, vregInputFp32, vregScalePerChannel, tailMask); + AscendC::MicroAPI::Mul(vregInputFp32, vregInputFp32, vregScalePerToken, tailMask); + AscendC::MicroAPI::StoreAlign( + yAddr + dLoops * floatRepSize + rowOffset, vregInputFp32, tailMask); + rowOffset += stride; + scaleOffset += fp32BlockElementNum; + } + } +} + +/** + * @brief DequantVf 对输入做per-token叠加per-channel的反量化, INT32 ---> FP32. + * @param outputLocal 输出tensor [row, col] + * @param inputLocal 输入tensor [row, col] + * @param scalePerChannelLocal 输入tensor [1, col] + * @param scalePerTokenLocal 输入tensor [row, 8] + * @param row 待处理的行数 + * @param col 待处理的列数 + * @param stride 待处理数据一行的真实长度 + */ +template +__aicore__ inline void DequantVf(const LocalTensor &outputLocal, const LocalTensor &inputLocal, + const LocalTensor &scalePerChannelLocal, + const LocalTensor &scalePerTokenLocal, + uint32_t row, uint32_t col, uint32_t stride) +{ + __ubuf__ float *outputUb = (__ubuf__ float *)outputLocal.GetPhyAddr(); + __ubuf__ T *inputUb = (__ubuf__ T *)inputLocal.GetPhyAddr(); + __ubuf__ float *scalePerChannelLocalUb = (__ubuf__ float *)scalePerChannelLocal.GetPhyAddr(); + __ubuf__ float *scalePerTokenLocalUb = (__ubuf__ float *)scalePerTokenLocal.GetPhyAddr(); + + const uint32_t floatRepSize = 64; // 一个寄存器能够存放64个FP32 + const uint32_t fp32BlockElementNum = 8; + uint32_t dLoops = col / floatRepSize; + uint32_t dTail = col % floatRepSize; + uint32_t dTailLoop = dTail > 0 ? 1 : 0; + DequantVFImpl(outputUb, inputUb, scalePerChannelLocalUb, scalePerTokenLocalUb, floatRepSize, fp32BlockElementNum, + dLoops, dTail, dTailLoop, row, col, stride); +} +} // namespace MlaProlog +#endif diff --git a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_quant_perchannel.h b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_quant_perchannel.h index 2b9a5a973063..cec65fc22ee0 100644 --- a/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_quant_perchannel.h +++ b/csrc/attention/mla_prolog_v3/op_kernel/arch35/vf/vf_quant_perchannel.h @@ -1,104 +1,104 @@ -/** - * Copyright (c) 2026 Huawei Technologies Co., Ltd. - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of - * CANN Open Software License Agreement Version 2.0 (the "License"). - * Please refer to the License for details. You may not use this file except in compliance with the License. - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. - * See LICENSE in the root of the software repository for the full text of the License. - */ - -/*! - * \file vf_quant_perchannel.h - * \brief - */ - -#ifndef VF_QUANT_PERCHANNEL_H -#define VF_QUANT_PERCHANNEL_H -#include "kernel_tensor.h" - -namespace MlaProlog { - -template -__simd_vf__ void QuantChannelVFImpl(__ubuf__ O *yAddr, __ubuf__ T *xAddr, __ubuf__ C *quantScaleAddr, - const uint32_t floatRepSize, uint32_t dLoops, uint32_t dTail, uint32_t dTailLoop, - uint32_t row, uint32_t col, uint32_t stride) -{ - AscendC::MicroAPI::RegTensor vregInput; - AscendC::MicroAPI::RegTensor vregQuantScale; - AscendC::MicroAPI::RegTensor vregOutput; - AscendC::MicroAPI::RegTensor vregOutputHalf; // float-->half-->int8 - AscendC::MicroAPI::MaskReg fullMask = AscendC::MicroAPI::CreateMask(); - AscendC::MicroAPI::MaskReg tailMask; - tailMask = AscendC::MicroAPI::UpdateMask(dTail); - constexpr static AscendC::MicroAPI::CastTrait castTraitPack2 = { - AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::SAT, AscendC::MicroAPI::MaskMergeMode::ZEROING, - AscendC::RoundMode::CAST_RINT}; - constexpr static AscendC::MicroAPI::CastTrait castTraitF32ToHalf = { - AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::NO_SAT, - AscendC::MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ODD}; - - uint32_t colOffset = 0; - uint32_t rowOffset = 0; - for (uint32_t j = 0; j < dLoops; j++) { - AscendC::MicroAPI::LoadAlign(vregQuantScale, - quantScaleAddr + colOffset); - rowOffset = 0; - for (uint32_t i = 0; i < row; i++) { - AscendC::MicroAPI::LoadAlign(vregInput, - xAddr + colOffset + rowOffset); - AscendC::MicroAPI::Mul(vregInput, vregInput, vregQuantScale, fullMask); - AscendC::MicroAPI::Cast(vregOutputHalf, vregInput, fullMask); - AscendC::MicroAPI::Cast(vregOutput, vregOutputHalf, fullMask); - AscendC::MicroAPI::StoreAlign( - yAddr + colOffset + rowOffset, vregOutput, fullMask); - rowOffset += stride; - } - colOffset += floatRepSize; - } - - if (dTailLoop > 0) { - rowOffset = 0; - AscendC::MicroAPI::LoadAlign(vregQuantScale, - quantScaleAddr + dLoops * floatRepSize); - for (uint32_t i = 0; i < row; i++) { - AscendC::MicroAPI::LoadAlign( - vregInput, xAddr + dLoops * floatRepSize + rowOffset); - AscendC::MicroAPI::Mul(vregInput, vregInput, vregQuantScale, tailMask); - AscendC::MicroAPI::Cast(vregOutputHalf, vregInput, tailMask); - AscendC::MicroAPI::Cast(vregOutput, vregOutputHalf, tailMask); - AscendC::MicroAPI::StoreAlign( - yAddr + dLoops * floatRepSize + rowOffset, vregOutput, tailMask); - rowOffset += stride; - } - } -} - -/** - * @brief QuantPerChannelVF 同时对row进行FP32到int8的per-channel量化操作。一行中的每一列用不同的量化参数。 - outLocal[i , j] = inputLocal[i , j] * quantScaleLocal[j] - * @param outputLocal 输出tensor [row, col] - * @param inputLocal 输入tensor [row, col] - * @param quantScaleLocal 量化参数tensor [1, col] - * @param row 处理数据的行数 默认为1 后续可扩展 - * @param col 处理数据的列数 - * @param stride 待处理数据一行的真实长度 - */ -template -__aicore__ inline void QuantPerChannelVf(const LocalTensor &outputLocal, const LocalTensor &inputLocal, - const LocalTensor &quantScaleLocal, const uint32_t row, const uint32_t col, - const uint32_t stride) -{ - __ubuf__ O *outputUb = (__ubuf__ O *)outputLocal.GetPhyAddr(); - __ubuf__ T *inputUb = (__ubuf__ T *)inputLocal.GetPhyAddr(); - __ubuf__ C *quantScaleUb = (__ubuf__ C *)quantScaleLocal.GetPhyAddr(); - - const uint32_t floatRepSize = 64; // 一个寄存器能够存放64个FP32 - uint32_t dLoops = col / floatRepSize; - uint32_t dTail = col % floatRepSize; - uint32_t dTailLoop = dTail > 0 ? 1 : 0; - - QuantChannelVFImpl(outputUb, inputUb, quantScaleUb, floatRepSize, dLoops, dTail, dTailLoop, row, col, stride); -} -} // namespace MlaProlog +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file vf_quant_perchannel.h + * \brief + */ + +#ifndef VF_QUANT_PERCHANNEL_H +#define VF_QUANT_PERCHANNEL_H +#include "kernel_tensor.h" + +namespace MlaProlog { + +template +__simd_vf__ void QuantChannelVFImpl(__ubuf__ O *yAddr, __ubuf__ T *xAddr, __ubuf__ C *quantScaleAddr, + const uint32_t floatRepSize, uint32_t dLoops, uint32_t dTail, uint32_t dTailLoop, + uint32_t row, uint32_t col, uint32_t stride) +{ + AscendC::MicroAPI::RegTensor vregInput; + AscendC::MicroAPI::RegTensor vregQuantScale; + AscendC::MicroAPI::RegTensor vregOutput; + AscendC::MicroAPI::RegTensor vregOutputHalf; // float-->half-->int8 + AscendC::MicroAPI::MaskReg fullMask = AscendC::MicroAPI::CreateMask(); + AscendC::MicroAPI::MaskReg tailMask; + tailMask = AscendC::MicroAPI::UpdateMask(dTail); + constexpr static AscendC::MicroAPI::CastTrait castTraitPack2 = { + AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::SAT, AscendC::MicroAPI::MaskMergeMode::ZEROING, + AscendC::RoundMode::CAST_RINT}; + constexpr static AscendC::MicroAPI::CastTrait castTraitF32ToHalf = { + AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::NO_SAT, + AscendC::MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ODD}; + + uint32_t colOffset = 0; + uint32_t rowOffset = 0; + for (uint32_t j = 0; j < dLoops; j++) { + AscendC::MicroAPI::LoadAlign(vregQuantScale, + quantScaleAddr + colOffset); + rowOffset = 0; + for (uint32_t i = 0; i < row; i++) { + AscendC::MicroAPI::LoadAlign(vregInput, + xAddr + colOffset + rowOffset); + AscendC::MicroAPI::Mul(vregInput, vregInput, vregQuantScale, fullMask); + AscendC::MicroAPI::Cast(vregOutputHalf, vregInput, fullMask); + AscendC::MicroAPI::Cast(vregOutput, vregOutputHalf, fullMask); + AscendC::MicroAPI::StoreAlign( + yAddr + colOffset + rowOffset, vregOutput, fullMask); + rowOffset += stride; + } + colOffset += floatRepSize; + } + + if (dTailLoop > 0) { + rowOffset = 0; + AscendC::MicroAPI::LoadAlign(vregQuantScale, + quantScaleAddr + dLoops * floatRepSize); + for (uint32_t i = 0; i < row; i++) { + AscendC::MicroAPI::LoadAlign( + vregInput, xAddr + dLoops * floatRepSize + rowOffset); + AscendC::MicroAPI::Mul(vregInput, vregInput, vregQuantScale, tailMask); + AscendC::MicroAPI::Cast(vregOutputHalf, vregInput, tailMask); + AscendC::MicroAPI::Cast(vregOutput, vregOutputHalf, tailMask); + AscendC::MicroAPI::StoreAlign( + yAddr + dLoops * floatRepSize + rowOffset, vregOutput, tailMask); + rowOffset += stride; + } + } +} + +/** + * @brief QuantPerChannelVF 同时对row进行FP32到int8的per-channel量化操作。一行中的每一列用不同的量化参数。 + outLocal[i , j] = inputLocal[i , j] * quantScaleLocal[j] + * @param outputLocal 输出tensor [row, col] + * @param inputLocal 输入tensor [row, col] + * @param quantScaleLocal 量化参数tensor [1, col] + * @param row 处理数据的行数 默认为1 后续可扩展 + * @param col 处理数据的列数 + * @param stride 待处理数据一行的真实长度 + */ +template +__aicore__ inline void QuantPerChannelVf(const LocalTensor &outputLocal, const LocalTensor &inputLocal, + const LocalTensor &quantScaleLocal, const uint32_t row, const uint32_t col, + const uint32_t stride) +{ + __ubuf__ O *outputUb = (__ubuf__ O *)outputLocal.GetPhyAddr(); + __ubuf__ T *inputUb = (__ubuf__ T *)inputLocal.GetPhyAddr(); + __ubuf__ C *quantScaleUb = (__ubuf__ C *)quantScaleLocal.GetPhyAddr(); + + const uint32_t floatRepSize = 64; // 一个寄存器能够存放64个FP32 + uint32_t dLoops = col / floatRepSize; + uint32_t dTail = col % floatRepSize; + uint32_t dTailLoop = dTail > 0 ? 1 : 0; + + QuantChannelVFImpl(outputUb, inputUb, quantScaleUb, floatRepSize, dLoops, dTail, dTailLoop, row, col, stride); +} +} // namespace MlaProlog #endif \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_tiling_data.h b/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_tiling_data.h index 5ac674341290..00fe4445c8f2 100644 --- a/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_tiling_data.h +++ b/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_tiling_data.h @@ -1,78 +1,78 @@ -/** - * Copyright (c) 2025 Huawei Technologies Co., Ltd. - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of - * CANN Open Software License Agreement Version 2.0 (the "License"). - * Please refer to the License for details. You may not use this file except in compliance with the License. - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. - * See LICENSE in the root of the software repository for the full text of the License. - */ - -/*! - * \file mla_prolog_tiling_datay.h - * \brief - */ - -#ifndef MLA_PROLOG_TILING_DATA_H -#define MLA_PROLOG_TILING_DATA_H -#include -#include "kernel_tiling/kernel_tiling.h" - -namespace optiling { -// 1. 基础参数结构体(对应 MlaPrologBaseParams 宏定义) -struct MlaPrologBaseParams { - uint32_t batchSize; // batch size(批大小) - uint32_t stepBatchSize; // batch size per step 32(每步批大小) - uint32_t mSubSize; - uint32_t mSubCoreNum; - uint32_t stepNumHeadDequant; // head size per step when dequant before mmqn(反量化前每步头大小) - uint32_t tokenSize; // token size = batchSize * seq1Size(token总数:批大小×序列1长度) - uint32_t seq1Size; // seq1(序列1长度,通常为query序列长度) - uint32_t seq2Size; // seq2(序列2长度,通常为key/value序列长度) - uint32_t headSizeX; // head size of Input Hidden 7168(输入隐藏层的头维度) - uint32_t headSizeCq; // head size of Latent Query 1536(潜在Query的头维度) - uint32_t headSizeCkv; // head size of Latent KeyValue 512(潜在KeyValue的头维度) - uint32_t headSizeQc; // head size of Query = dimHeadSizeQc * numHeadSize = 128 * 32(Query总头维度) - uint32_t headSizeQr; // head size of Query Rope = dimHeadRope * numHeadSize = 64 * 32(带RoPE的Query头维度) - uint32_t headSizeKr; // head size of Key Rope 64(带RoPE的Key头维度) - uint32_t numHeadSize; // number of head 32(头数量) - uint32_t numHeadKvSize; // number of headkv(KeyValue的头数量) - uint32_t dimHeadSizeQc; // dim size per query head 128(单个Query头的维度) - uint32_t dimHeadRope; // dim size per rope head 64(单个带RoPE头的维度) - uint32_t blockNum; // pa block num(PA格式的块数量) - uint32_t blockSize; // pa block size 128(PA格式的块大小) +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +/*! + * \file mla_prolog_tiling_datay.h + * \brief + */ + +#ifndef MLA_PROLOG_TILING_DATA_H +#define MLA_PROLOG_TILING_DATA_H +#include +#include "kernel_tiling/kernel_tiling.h" + +namespace optiling { +// 1. 基础参数结构体(对应 MlaPrologBaseParams 宏定义) +struct MlaPrologBaseParams { + uint32_t batchSize; // batch size(批大小) + uint32_t stepBatchSize; // batch size per step 32(每步批大小) + uint32_t mSubSize; + uint32_t mSubCoreNum; + uint32_t stepNumHeadDequant; // head size per step when dequant before mmqn(反量化前每步头大小) + uint32_t tokenSize; // token size = batchSize * seq1Size(token总数:批大小×序列1长度) + uint32_t seq1Size; // seq1(序列1长度,通常为query序列长度) + uint32_t seq2Size; // seq2(序列2长度,通常为key/value序列长度) + uint32_t headSizeX; // head size of Input Hidden 7168(输入隐藏层的头维度) + uint32_t headSizeCq; // head size of Latent Query 1536(潜在Query的头维度) + uint32_t headSizeCkv; // head size of Latent KeyValue 512(潜在KeyValue的头维度) + uint32_t headSizeQc; // head size of Query = dimHeadSizeQc * numHeadSize = 128 * 32(Query总头维度) + uint32_t headSizeQr; // head size of Query Rope = dimHeadRope * numHeadSize = 64 * 32(带RoPE的Query头维度) + uint32_t headSizeKr; // head size of Key Rope 64(带RoPE的Key头维度) + uint32_t numHeadSize; // number of head 32(头数量) + uint32_t numHeadKvSize; // number of headkv(KeyValue的头数量) + uint32_t dimHeadSizeQc; // dim size per query head 128(单个Query头的维度) + uint32_t dimHeadRope; // dim size per rope head 64(单个带RoPE头的维度) + uint32_t blockNum; // pa block num(PA格式的块数量) + uint32_t blockSize; // pa block size 128(PA格式的块大小) uint64_t kvCacheStride0; // kv cache first-axis stride in elements uint64_t krCacheStride0; // kr cache first-axis stride in elements - uint32_t mm1BlockNum; // 24 Cq(矩阵乘1的块数量,对应Cq计算) - uint32_t mm2BlockNum; // 9 Ckv(矩阵乘2的块数量,对应Ckv计算) - uint32_t mm3BlockNum; // 24 QcQr(矩阵乘3的块数量,对应QcQr计算) - uint32_t mm4BlockNum; // 24 Qn(矩阵乘4的块数量,对应Qn计算) - uint32_t vectorBlockNum; // 32(向量计算的块数量) - uint32_t mm1SingleCoreN; // single headSizeCq(单核心矩阵乘1的N维度大小,对应单个Cq头维度) - uint32_t mm2SingleCoreN; // single headSizeCkv+headSizeKr(单核心矩阵乘2的N维度大小,Ckv+Kr头维度之和) - uint32_t mm3SingleCoreN; // single headSizeQc+headSizeQr(单核心矩阵乘3的N维度大小,Qc+Qr头维度之和) - uint32_t mm4SingleCoreBatch; // single numHeadSize(单核心矩阵乘4的批大小,对应单个头数量) - uint32_t dtileSize; - uint32_t kvQuantMode; - uint32_t tileSize; - uint32_t ckvkrRepoMode; - uint32_t quantScaleRepoMode; - uint32_t queryNormFlag; - float reciprocalCq; // 1 / headSizeCq(headSizeCq的倒数,用于快速计算) - float epsilonCq; // Cq计算的epsilon(数值稳定性参数) - float reciprocalCkv; // 1 / headSizeCkv(headSizeCkv的倒数,用于快速计算) - float epsilonCkv; // Ckv计算的epsilon(数值稳定性参数) - float kNopeClipAlpha; - float qcQrScale; // query 的尺度矫正因子 - float kcScale; // kv 的尺度矫正因子 - uint16_t isQcQrScaleEnable; // query 的尺度矫正因子是否生效(默认是1.0的时候不生效) - uint16_t isKcScaleEnable; // kv 的尺度矫正因子是否生效(默认是1.0的时候不生效) -}; - -// 2. 完整分块数据结构体(对应 MlaPrologTilingData 宏定义,嵌套基础参数) -struct MlaPrologTilingData { - MlaPrologBaseParams baseParams; // 嵌套基础参数结构体(包含维度、头信息等核心配置) -}; -} // namespace optiling - + uint32_t mm1BlockNum; // 24 Cq(矩阵乘1的块数量,对应Cq计算) + uint32_t mm2BlockNum; // 9 Ckv(矩阵乘2的块数量,对应Ckv计算) + uint32_t mm3BlockNum; // 24 QcQr(矩阵乘3的块数量,对应QcQr计算) + uint32_t mm4BlockNum; // 24 Qn(矩阵乘4的块数量,对应Qn计算) + uint32_t vectorBlockNum; // 32(向量计算的块数量) + uint32_t mm1SingleCoreN; // single headSizeCq(单核心矩阵乘1的N维度大小,对应单个Cq头维度) + uint32_t mm2SingleCoreN; // single headSizeCkv+headSizeKr(单核心矩阵乘2的N维度大小,Ckv+Kr头维度之和) + uint32_t mm3SingleCoreN; // single headSizeQc+headSizeQr(单核心矩阵乘3的N维度大小,Qc+Qr头维度之和) + uint32_t mm4SingleCoreBatch; // single numHeadSize(单核心矩阵乘4的批大小,对应单个头数量) + uint32_t dtileSize; + uint32_t kvQuantMode; + uint32_t tileSize; + uint32_t ckvkrRepoMode; + uint32_t quantScaleRepoMode; + uint32_t queryNormFlag; + float reciprocalCq; // 1 / headSizeCq(headSizeCq的倒数,用于快速计算) + float epsilonCq; // Cq计算的epsilon(数值稳定性参数) + float reciprocalCkv; // 1 / headSizeCkv(headSizeCkv的倒数,用于快速计算) + float epsilonCkv; // Ckv计算的epsilon(数值稳定性参数) + float kNopeClipAlpha; + float qcQrScale; // query 的尺度矫正因子 + float kcScale; // kv 的尺度矫正因子 + uint16_t isQcQrScaleEnable; // query 的尺度矫正因子是否生效(默认是1.0的时候不生效) + uint16_t isKcScaleEnable; // kv 的尺度矫正因子是否生效(默认是1.0的时候不生效) +}; + +// 2. 完整分块数据结构体(对应 MlaPrologTilingData 宏定义,嵌套基础参数) +struct MlaPrologTilingData { + MlaPrologBaseParams baseParams; // 嵌套基础参数结构体(包含维度、头信息等核心配置) +}; +} // namespace optiling + #endif // MLA_PROLOG_TILING_DATA_H \ No newline at end of file diff --git a/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_v3.cpp b/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_v3.cpp index 4a806673815c..e2d14364aeb4 100644 --- a/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_v3.cpp +++ b/csrc/attention/mla_prolog_v3/op_kernel/mla_prolog_v3.cpp @@ -58,7 +58,7 @@ __global__ __aicore__ void mla_prolog_v3( __gm__ uint8_t *queryNormOut, __gm__ uint8_t *dequantScaleQNormOut, __gm__ uint8_t *workspace, - __gm__ uint8_t *tiling) + __gm__ uint8_t *tiling) { #if (__NPU_ARCH__ == 3510) int64_t globalOriOverflowMode = AscendC::GetCtrlSpr(); @@ -86,67 +86,67 @@ __global__ __aicore__ void mla_prolog_v3( #endif if constexpr (static_cast(Scenario) == SCENARIO::NO_QUANT) { MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); op.Process(); - } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::PARTIAL_QUANT_KV_NO_QUANT) { MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); op.Process(); - } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PER_CHANNEL) { MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); op.Process(); - } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::FULL_QUANT_KV_NO_QUANT) { MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); op.Process(); - } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::FULL_QUANT_KV_QUANT_PER_TENSOR) { MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); op.Process(); - } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::PARTIAL_QUANT_KV_QUANT_PERTILE) { MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); op.Process(); - } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::FULL_QUANT_KV_QUANT_PERTILE) { MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, @@ -176,7 +176,7 @@ __global__ __aicore__ void mla_prolog_v3( queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); op.Process(); } - } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::MXFP8_FULL_QUANT_KV_QUANT_PER_TENSOR) { if constexpr (splitMMode == SPLIT_M_MODE::ENABLED) { MlaPrologV3SplitM(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::MXFP8_FULL_QUANT_KV_QUANT_PER_TILE) { if constexpr (splitMMode == SPLIT_M_MODE::ENABLED) { MlaPrologV3SplitM(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::FP8_FULL_QUANT_KV_NO_QUANT) { MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); op.Process(); - } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::FP8_FULL_QUANT_KV_QUANT_PER_TENSOR) { MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); op.Process(); - } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::HIF8_FULL_QUANT_KV_NO_QUANT) { MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); op.Process(); - } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::HIF8_FULL_QUANT_KV_QUANT_PER_TENSOR) { MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); op.Process(); - } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::FP8_FULL_QUANT_KV_QUANT_PER_TILE) { MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, dequantScaleWDkvKr, quantScaleCkv, quantScaleCkr, smoothScalesCq, actualSeqLen, kNopeClipAlpha, queryOut, queryRopeOut, dequantScaleQNopeOut, queryNormOut, dequantScaleQNormOut, workspace); op.Process(); - } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && + } else if constexpr (static_cast(Scenario) == SCENARIO::QUANT && static_cast(QuantMode) == QUANT_MODE::HIF8_FULL_QUANT_KV_QUANT_PER_TILE) { MlaPrologVecS1CubS2> op(&pipe, tilingData, tilingDataBaseParams); op.Init(tokenX, weightDq, weightUqQr, weightUk, weightDkvKr, rmsnormGammaCq, rmsnormGammaCkv, ropeSin, ropeCos, cacheIndex, kvCacheOut, krCacheOut, dequantScaleX, dequantScaleWDq, dequantScaleWUqQr, From adf50d7f3e65685d57447d23372f638613b2be00 Mon Sep 17 00:00:00 2001 From: MQ Date: Wed, 26 Aug 2026 08:24:33 +0800 Subject: [PATCH 26/50] fix(quantization): use forward config for MX scale selection Runtime MoE quantization executes under the forward context after the model initialization config context has been restored. Resolve the model config from ForwardContext during model execution while retaining the initialization fallback. Signed-off-by: MQ --- tests/ut/quantization/test_utils.py | 14 ++++++++++++++ vllm_ascend/quantization/utils.py | 9 +++++++-- 2 files changed, 21 insertions(+), 2 deletions(-) diff --git a/tests/ut/quantization/test_utils.py b/tests/ut/quantization/test_utils.py index 0cdcb371a070..5d0eaa136bad 100644 --- a/tests/ut/quantization/test_utils.py +++ b/tests/ut/quantization/test_utils.py @@ -49,6 +49,20 @@ def test_uses_current_vllm_config_when_config_is_omitted(self, mock_current_conf self.assertEqual(get_dynamic_mx_quant_scale_alg(), 1) + @patch("vllm.forward_context.get_forward_context") + @patch("vllm.forward_context.is_forward_context_available", return_value=True) + @patch("vllm_ascend.quantization.utils.get_ascend_device_type", return_value=AscendDeviceType.A5) + def test_uses_forward_vllm_config_when_available( + self, + _mock_device_type, + _mock_forward_context_available, + mock_forward_context, + ): + minimax_config = self._config(None, model_type="minimax_m3") + mock_forward_context.return_value.vllm_config = minimax_config + + self.assertEqual(get_dynamic_mx_quant_scale_alg(), 1) + class TestDetectQuantizationMethod(TestBase): def test_returns_none_for_non_existent_path(self): diff --git a/vllm_ascend/quantization/utils.py b/vllm_ascend/quantization/utils.py index 59638965d738..d30874554151 100644 --- a/vllm_ascend/quantization/utils.py +++ b/vllm_ascend/quantization/utils.py @@ -49,9 +49,14 @@ def get_dynamic_mx_quant_scale_alg(vllm_config=None) -> int: return 0 if vllm_config is None: - from vllm.config import get_current_vllm_config + from vllm.forward_context import get_forward_context, is_forward_context_available - vllm_config = get_current_vllm_config() + if is_forward_context_available(): + vllm_config = get_forward_context().vllm_config + else: + from vllm.config import get_current_vllm_config + + vllm_config = get_current_vllm_config() model_config = vllm_config.model_config architectures = getattr(model_config, "architectures", None) or () From da6e07123a052ee816a929bfe05e30b9092bde9f Mon Sep 17 00:00:00 2001 From: MQ Date: Wed, 26 Aug 2026 09:46:36 +0800 Subject: [PATCH 27/50] fix(quantization): carry MX scale algorithm into forward context Signed-off-by: MQ --- tests/ut/quantization/test_utils.py | 18 +++++++++++++++++- tests/ut/test_platform.py | 18 ++++++++++++++++++ vllm_ascend/platform.py | 5 ++++- vllm_ascend/quantization/utils.py | 12 ++++++++---- 4 files changed, 47 insertions(+), 6 deletions(-) diff --git a/tests/ut/quantization/test_utils.py b/tests/ut/quantization/test_utils.py index 5d0eaa136bad..2ee23446deb0 100644 --- a/tests/ut/quantization/test_utils.py +++ b/tests/ut/quantization/test_utils.py @@ -57,9 +57,25 @@ def test_uses_forward_vllm_config_when_available( _mock_device_type, _mock_forward_context_available, mock_forward_context, + ): + mock_forward_context.return_value.additional_kwargs = {"dynamic_mx_quant_scale_alg": 1} + + self.assertEqual(get_dynamic_mx_quant_scale_alg(), 1) + + @patch("vllm.config.get_current_vllm_config") + @patch("vllm.forward_context.get_forward_context") + @patch("vllm.forward_context.is_forward_context_available", return_value=True) + @patch("vllm_ascend.quantization.utils.get_ascend_device_type", return_value=AscendDeviceType.A5) + def test_falls_back_to_current_config_when_forward_scale_is_missing( + self, + _mock_device_type, + _mock_forward_context_available, + mock_forward_context, + mock_current_config, ): minimax_config = self._config(None, model_type="minimax_m3") - mock_forward_context.return_value.vllm_config = minimax_config + mock_forward_context.return_value.additional_kwargs = {} + mock_current_config.return_value = minimax_config self.assertEqual(get_dynamic_mx_quant_scale_alg(), 1) diff --git a/tests/ut/test_platform.py b/tests/ut/test_platform.py index 6a366840a27c..d344b69bb1eb 100644 --- a/tests/ut/test_platform.py +++ b/tests/ut/test_platform.py @@ -602,6 +602,24 @@ def test_set_additional_forward_context_v2_includes_required_moe_fields(self): self.assertFalse(kwargs["in_profile_run"]) self.assertEqual(kwargs["padded_num_tokens"], 8) self.assertIs(kwargs["moe_comm_method"], dummy_comm_method) + self.assertEqual(kwargs["dynamic_mx_quant_scale_alg"], 0) + + def test_set_additional_forward_context_v1_includes_dynamic_mx_scale_alg(self): + vllm_config = TestNPUPlatform.mock_vllm_config() + vllm_config.use_v2_model_runner = False + + with patch( + "vllm_ascend.quantization.utils.get_dynamic_mx_quant_scale_alg", + return_value=1, + ): + kwargs = self.platform.set_additional_forward_context( + attn_metadata=None, + vllm_config=vllm_config, + dp_metadata=None, + num_tokens=5, + ) + + self.assertEqual(kwargs, {"dynamic_mx_quant_scale_alg": 1}) def test_set_additional_forward_context_reads_v2_profile_override(self): vllm_config = TestNPUPlatform.mock_vllm_config() diff --git a/vllm_ascend/platform.py b/vllm_ascend/platform.py index cf58351c5aee..129d957b7b83 100644 --- a/vllm_ascend/platform.py +++ b/vllm_ascend/platform.py @@ -502,6 +502,7 @@ def set_additional_forward_context( select_moe_comm_method, ) from vllm_ascend.ops.fused_moe.moe_comm_method import get_moe_comm_method + from vllm_ascend.quantization.utils import get_dynamic_mx_quant_scale_alg from vllm.distributed import get_dp_group, get_tensor_model_parallel_world_size # NOTE(Ronald1995): avoid circular import, cudagraph_runtime_mode is @@ -518,8 +519,9 @@ def set_additional_forward_context( # compared to v1, v2's forward context lacks some fields, such as: # is_first_layer, prefetch_mlp_gate_up_proj, prefetch_mlp_gate_down_proj, # prefetch_mlp_enabled, model_instance, is_draft_model. + dynamic_mx_quant_scale_alg = get_dynamic_mx_quant_scale_alg(vllm_config) if not vllm_config.use_v2_model_runner: - return {} + return {"dynamic_mx_quant_scale_alg": dynamic_mx_quant_scale_alg} # is_draft_model will be removed later, so we set it to False temporarily. is_draft_model = False @@ -581,6 +583,7 @@ def set_additional_forward_context( "in_profile_run": in_profile_run, "padded_num_tokens": padded_num_tokens, "sinks": sinks, + "dynamic_mx_quant_scale_alg": dynamic_mx_quant_scale_alg, } diff --git a/vllm_ascend/quantization/utils.py b/vllm_ascend/quantization/utils.py index d30874554151..7f4c02b03a80 100644 --- a/vllm_ascend/quantization/utils.py +++ b/vllm_ascend/quantization/utils.py @@ -52,11 +52,15 @@ def get_dynamic_mx_quant_scale_alg(vllm_config=None) -> int: from vllm.forward_context import get_forward_context, is_forward_context_available if is_forward_context_available(): - vllm_config = get_forward_context().vllm_config - else: - from vllm.config import get_current_vllm_config + scale_alg = get_forward_context().additional_kwargs.get( + "dynamic_mx_quant_scale_alg" + ) + if scale_alg is not None: + return scale_alg + + from vllm.config import get_current_vllm_config - vllm_config = get_current_vllm_config() + vllm_config = get_current_vllm_config() model_config = vllm_config.model_config architectures = getattr(model_config, "architectures", None) or () From 9e8b8f795deb7af0e62bdc84a92d790755b35e1a Mon Sep 17 00:00:00 2001 From: MQ Date: Wed, 26 Aug 2026 10:08:40 +0800 Subject: [PATCH 28/50] fix(ci): cover Kimi KDA and format MX context helper Signed-off-by: MQ --- .github/workflows/scripts/test_config.yaml | 1 + vllm_ascend/quantization/utils.py | 4 +--- 2 files changed, 2 insertions(+), 3 deletions(-) diff --git a/.github/workflows/scripts/test_config.yaml b/.github/workflows/scripts/test_config.yaml index ea08d20fd912..6eba786b6ecb 100644 --- a/.github/workflows/scripts/test_config.yaml +++ b/.github/workflows/scripts/test_config.yaml @@ -392,6 +392,7 @@ source_file_dependencies: - vllm_ascend/ops/gdn.py - vllm_ascend/ops/gdn_attn_builder.py + - vllm_ascend/ops/kimi_kda.py - vllm_ascend/ops/triton/gdn_chunk_meta.py tests: - tests/ut/ops diff --git a/vllm_ascend/quantization/utils.py b/vllm_ascend/quantization/utils.py index 7f4c02b03a80..a4fa3c25a5b0 100644 --- a/vllm_ascend/quantization/utils.py +++ b/vllm_ascend/quantization/utils.py @@ -52,9 +52,7 @@ def get_dynamic_mx_quant_scale_alg(vllm_config=None) -> int: from vllm.forward_context import get_forward_context, is_forward_context_available if is_forward_context_available(): - scale_alg = get_forward_context().additional_kwargs.get( - "dynamic_mx_quant_scale_alg" - ) + scale_alg = get_forward_context().additional_kwargs.get("dynamic_mx_quant_scale_alg") if scale_alg is not None: return scale_alg From 8a3fdcd766f9360d271e911f399dbed9f4e394f5 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Tue, 25 Aug 2026 21:46:26 -0500 Subject: [PATCH 29/50] refactor(kimi-k3): remove operator-integration detours Restore the established K3 KV-cache and proposer paths, disable MLAPO only at the K3 draft construction boundary, and keep the Block5 RoPE/causality integration explicit. Remove mock-heavy tests that only covered discarded workaround code. Signed-off-by: maoxx241 --- tests/ut/models/test_kimi_k3_adapter.py | 26 ++++- tests/ut/worker/a2/test_model_runner_v1.py | 130 +++++++++++++++------ vllm_ascend/models/kimi_k3.py | 5 +- vllm_ascend/models/kimi_k3_dspark.py | 19 ++- vllm_ascend/worker/model_runner_v1.py | 45 ++++++- 5 files changed, 177 insertions(+), 48 deletions(-) diff --git a/tests/ut/models/test_kimi_k3_adapter.py b/tests/ut/models/test_kimi_k3_adapter.py index 8f466c54ccf8..51d91820e33f 100644 --- a/tests/ut/models/test_kimi_k3_adapter.py +++ b/tests/ut/models/test_kimi_k3_adapter.py @@ -114,16 +114,18 @@ def test_dspark_decoder_uses_upstream_mlp_activation_contract( intermediate_size=16, hidden_act="silu", rms_norm_eps=1e-6, + full_attention_causal=True, ) vllm_config = SimpleNamespace(cache_config=None) mlp_factory = MagicMock(return_value=nn.Identity()) + attention_factory = MagicMock(return_value=nn.Identity()) monkeypatch.setattr( "vllm_ascend.models.kimi_k3_dspark.get_draft_quant_config", lambda _: None, ) monkeypatch.setattr( "vllm_ascend.models.kimi_k3_dspark.AscendKimiMLAAttention", - lambda **_: nn.Identity(), + attention_factory, ) monkeypatch.setattr( "vllm_ascend.models.kimi_k3_dspark.KimiMLP", @@ -142,6 +144,22 @@ def test_dspark_decoder_uses_upstream_mlp_activation_contract( assert mlp_factory.call_args.kwargs["hidden_act"] == "silu" assert "activation_situ_beta" not in mlp_factory.call_args.kwargs assert "activation_situ_linear_beta" not in mlp_factory.call_args.kwargs + assert attention_factory.call_args.kwargs["non_causal_multi_token_decode"] is False + + +def test_k3_dspark_reports_draft_attention_causality(): + model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) + nn.Module.__init__(model) + model.model = SimpleNamespace(layers=[object(), object(), object()]) + + model.config = SimpleNamespace(dflash_config={"causal": True}) + assert model.get_draft_attn_causal() == [True, True, True] + + model.config = SimpleNamespace(full_attention_causal=True) + assert model.get_draft_attn_causal() == [True, True, True] + + model.config = SimpleNamespace() + assert model.get_draft_attn_causal() == [False, False, False] def test_ascend_kimi_moe_quantizes_modelslim_latent_projections(monkeypatch): @@ -657,8 +675,9 @@ def test_k3_dspark_load_weights_keeps_per_layer_context_kv(monkeypatch): seen_names: list[str] = [] class CapturingLoader: - def __init__(self, loaded_model): + def __init__(self, loaded_model, *, skip_substrs): assert loaded_model is model + assert skip_substrs == list(model.checkpoint_skip_substrs) def load_weights(self, weights, *, mapper): assert mapper is model.hf_to_vllm_mapper @@ -691,8 +710,9 @@ def test_k3_dspark_reuses_modelslim_rotation_loader(monkeypatch): seen_weights: list[tuple[str, torch.Tensor]] = [] class CapturingLoader: - def __init__(self, loaded_model): + def __init__(self, loaded_model, *, skip_substrs): assert loaded_model is model + assert skip_substrs == list(model.checkpoint_skip_substrs) def load_weights(self, weights, *, mapper): assert mapper is model.hf_to_vllm_mapper diff --git a/tests/ut/worker/a2/test_model_runner_v1.py b/tests/ut/worker/a2/test_model_runner_v1.py index 170daf9ef525..81a812e9e621 100644 --- a/tests/ut/worker/a2/test_model_runner_v1.py +++ b/tests/ut/worker/a2/test_model_runner_v1.py @@ -4,6 +4,7 @@ import numpy as np import torch +from vllm.config import CUDAGraphMode from vllm.model_executor.layers.attention import MLAAttention from vllm.model_executor.models.deepseek_v2 import DeepseekV32IndexerCache from vllm.v1.kv_cache_interface import ( @@ -11,7 +12,6 @@ KVCacheConfig, KVCacheGroupSpec, KVCacheTensor, - MLAAttentionSpec, UniformTypeKVCacheSpecs, ) @@ -134,7 +134,13 @@ def reinitialize_input_batch(_kv_cache_config): runner.may_reinitialize_input_batch = MagicMock(side_effect=reinitialize_input_batch) runner.initialize_kv_cache_tensors = MagicMock(return_value={}) - runner.initialize_kv_cache(SimpleNamespace(kv_cache_groups=[])) + runner.initialize_kv_cache( + KVCacheConfig( + num_blocks=0, + kv_cache_tensors=[], + kv_cache_groups=[], + ) + ) drafter.initialize_attn_backend.assert_called_once() self.assertEqual( @@ -164,52 +170,102 @@ def test_allocate_kv_cache_uses_layer_spec_for_draft_gqa(self): self.assertEqual(k_cache_raw.numel(), kv_cache_spec.page_size_bytes) self.assertEqual(v_cache_raw.numel(), kv_cache_spec.page_size_bytes) - @patch("vllm_ascend.worker.model_runner_v1.has_ec_transfer", return_value=False) @patch("vllm_ascend.worker.model_runner_v1.get_layers_from_vllm_config") - def test_draft_mla_uses_separate_target_kv_cache_group( - self, - mock_get_layers, - _mock_has_ec_transfer, - ): + def test_mla_rope_modes_use_separate_metadata_groups(self, mock_get_layers): + class FakeBuilder: + def __init__(self, _spec, layer_names, _config, _device): + self.layer_names = layer_names + + class FakeBackend: + @classmethod + def full_cls_name(cls): + return "test.FakeBackend" + + @classmethod + def get_builder_cls(cls): + return FakeBuilder + runner = self._build_runner() - runner.block_size = 16 - runner.shared_kv_cache_layers = {} + runner.attn_groups = [] + runner._check_and_update_cudagraph_mode = MagicMock() + runner.calculate_reorder_batch_threshold = MagicMock() - draft_attn = MLAAttention.__new__(MLAAttention) - torch.nn.Module.__init__(draft_attn) - draft_attn.impl = SimpleNamespace(fa_quant_layer=False) - draft_attn.get_kv_cache_spec = MagicMock( - return_value=MLAAttentionSpec( + target_layer = "language_model.model.layers.0.self_attn.attn" + draft_layer = "model.layers.0.self_attn.attn" + mock_get_layers.return_value = { + target_layer: SimpleNamespace( + impl=SimpleNamespace(use_mla_rope=False), + get_attn_backend=lambda: FakeBackend, + ), + draft_layer: SimpleNamespace( + impl=SimpleNamespace(use_mla_rope=True), + get_attn_backend=lambda: FakeBackend, + ), + } + specs = { + target_layer: AscendMLAAttentionSpec( block_size=16, num_kv_heads=1, head_size=576, dtype=torch.bfloat16, - non_causal_multi_token_decode=True, - ) + ), + draft_layer: AscendMLAAttentionSpec( + block_size=16, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + ), + } + group_spec = UniformTypeKVCacheSpecs.from_specs(specs) + self.assertIsNotNone(group_spec) + assert group_spec is not None + kv_cache_config = KVCacheConfig( + num_blocks=2, + kv_cache_tensors=[], + kv_cache_groups=[ + KVCacheGroupSpec( + layer_names=[target_layer, draft_layer], + kv_cache_spec=group_spec, + ) + ], ) - mock_get_layers.return_value = {"draft.self_attn": draft_attn} - draft_spec = runner.get_kv_cache_spec()["draft.self_attn"] - target_spec = AscendMLAAttentionSpec( - block_size=16, - num_kv_heads=1, - head_size=576, - dtype=torch.bfloat16, - ) + runner.initialize_attn_backend(kv_cache_config) - self.assertTrue(draft_spec.non_causal_multi_token_decode) - uniform_spec = UniformTypeKVCacheSpecs.from_specs( - { - "target.self_attn": target_spec, - "draft.self_attn": draft_spec, - } + self.assertEqual(len(runner.attn_groups), 1) + self.assertEqual( + {tuple(group.layer_names) for group in runner.attn_groups[0]}, + {(target_layer,), (draft_layer,)}, ) - self.assertIsNone(uniform_spec) - with self.assertRaisesRegex( - AssertionError, - "Causal target layers and non-causal multi-token draft layers", - ): - AscendMLAAttentionSpec.merge([target_spec, draft_spec]) + + def test_explicit_capture_sizes_must_align_spec_decode_and_sp(self): + for capture_sizes, expected_tp_size in (([48, 96], 1), ([16, 32], 16)): + with self.subTest(capture_sizes=capture_sizes): + runner = self._build_runner() + compilation_config = SimpleNamespace( + pass_config=SimpleNamespace(enable_sp=True), + cudagraph_capture_sizes=capture_sizes, + resolve_cudagraph_mode_and_sizes=MagicMock(return_value=CUDAGraphMode.FULL_DECODE_ONLY), + ) + runner.compilation_config = compilation_config + runner.vllm_config.compilation_config = compilation_config + runner.parallel_config = SimpleNamespace(tensor_parallel_size=16) + runner.uniform_decode_query_len = 6 + runner.kv_cache_config = SimpleNamespace() + runner.max_num_reqs = 16 + runner.cudagraph_dispatcher = MagicMock() + runner.cudagraph_dispatcher.get_capture_descs.return_value = [] + runner.speculative_config = None + runner.drafter = None + runner.use_aclgraph = False + + runner._check_and_update_cudagraph_mode([], []) + + call_kwargs = compilation_config.resolve_cudagraph_mode_and_sizes.call_args.kwargs + self.assertEqual( + call_kwargs["tensor_parallel_size"], + expected_tp_size, + ) def test_sparse_c8_indexer_reuses_raw_cache_from_shared_descriptor(self): runner = self._build_runner() diff --git a/vllm_ascend/models/kimi_k3.py b/vllm_ascend/models/kimi_k3.py index b9569d095556..727b503188ba 100644 --- a/vllm_ascend/models/kimi_k3.py +++ b/vllm_ascend/models/kimi_k3.py @@ -259,6 +259,7 @@ def __init__( quant_config: QuantizationConfig | None = None, prefix: str = "", non_causal_multi_token_decode: bool = False, + disable_mlapo: bool = False, ) -> None: upstream_config = copy(config) upstream_config.mla_use_output_gate = use_output_gate @@ -276,6 +277,9 @@ def __init__( quant_config=quant_config, prefix=prefix, ) + attention_layer = self._attention_layer + if disable_mlapo: + attention_layer.impl.enable_mlapo = False if not use_rope and not non_causal_multi_token_decode: return @@ -303,7 +307,6 @@ def __init__( # registered MLA wrapper, including all projections and weight loaders. # Configure that existing Ascend attention layer for DSpark instead of # constructing and registering a second wrapper with the same prefix. - attention_layer = self._attention_layer attention_layer.scale = self.scaling attention_layer.non_causal_multi_token_decode = non_causal_multi_token_decode attention_layer.impl.scale = float(self.scaling) diff --git a/vllm_ascend/models/kimi_k3_dspark.py b/vllm_ascend/models/kimi_k3_dspark.py index 5754bbac535d..64e51bc70e86 100644 --- a/vllm_ascend/models/kimi_k3_dspark.py +++ b/vllm_ascend/models/kimi_k3_dspark.py @@ -49,6 +49,13 @@ from vllm_ascend.ops.rotary_embedding import get_cos_and_sin_mla +def _uses_causal_draft_attention(config) -> bool: + dflash_config = getattr(config, "dflash_config", None) + if isinstance(dflash_config, dict) and "causal" in dflash_config: + return bool(dflash_config["causal"]) + return bool(getattr(config, "full_attention_causal", False)) + + class AscendK3DSparkDecoderLayer(UpstreamK3DSparkDecoderLayer): def __init__( self, @@ -82,7 +89,8 @@ def __init__( cache_config=vllm_config.cache_config, quant_config=quant_config, prefix=f"{layer_prefix}.self_attn", - non_causal_multi_token_decode=True, + non_causal_multi_token_decode=not _uses_causal_draft_attention(config), + disable_mlapo=True, ) self.mlp = KimiMLP( hidden_size=config.hidden_size, @@ -264,6 +272,10 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: prefix=maybe_prefix(prefix, "lm_head"), ) + def get_draft_attn_causal(self) -> list[bool]: + causal = _uses_causal_draft_attention(self.config) + return [causal] * len(self.model.layers) + def load_weights( self, weights: Iterable[tuple[str, torch.Tensor]], @@ -275,7 +287,10 @@ def load_weights( quantization-aware per-layer projections, so use vLLM's public loader interface without creating that extra packed parameter. """ - loader = AutoWeightsLoader(self) + loader = AutoWeightsLoader( + self, + skip_substrs=list(self.checkpoint_skip_substrs), + ) rotation_weight = None if self.rotation_path is not None: rotation_weight = get_rotation_matrix(self.rotation_path) diff --git a/vllm_ascend/worker/model_runner_v1.py b/vllm_ascend/worker/model_runner_v1.py index 5938d1d6c3b8..5b4937902b98 100644 --- a/vllm_ascend/worker/model_runner_v1.py +++ b/vllm_ascend/worker/model_runner_v1.py @@ -4608,6 +4608,7 @@ def initialize_attn_backend(self, kv_cache_config: KVCacheConfig) -> None: class AttentionGroupKey(NamedTuple): attn_backend: type[AttentionBackend] kv_cache_spec: KVCacheSpec + use_mla_rope: bool | None def get_attn_backends_for_group( kv_cache_group_spec: KVCacheGroupSpec, @@ -4654,8 +4655,17 @@ def backend_supports_kernel_block_size( attn_backend = AscendSFAIndexerBackend full_cls_name = attn_backend.full_cls_name() - key = (full_cls_name, layer_kv_cache_spec) - attn_backends[key] = AttentionGroupKey(attn_backend, layer_kv_cache_spec) + use_mla_rope = ( + getattr(layers[layer_name].impl, "use_mla_rope", None) + if isinstance(layer_kv_cache_spec, AscendMLAAttentionSpec) + else None + ) + key = (full_cls_name, layer_kv_cache_spec, use_mla_rope) + attn_backends[key] = AttentionGroupKey( + attn_backend, + layer_kv_cache_spec, + use_mla_rope, + ) attn_backend_layers[key].append(layer_name) return ( {attn_backends[k]: v for k, v in attn_backend_layers.items()}, @@ -4663,10 +4673,13 @@ def backend_supports_kernel_block_size( ) def create_attn_groups( - attn_backends_map: dict[AttentionBackend, list[str]], kv_cache_group_id: int + attn_backends_map: dict[AttentionGroupKey, list[str]], + kv_cache_group_id: int, ) -> list[AttentionGroup]: attn_groups: list[AttentionGroup] = [] - for (attn_backend, kv_cache_spec), layer_names in attn_backends_map.items(): + for group_key, layer_names in attn_backends_map.items(): + attn_backend = group_key.attn_backend + kv_cache_spec = group_key.kv_cache_spec attn_metadata_builders = [] attn_metadata_builders.append( attn_backend.get_builder_cls()( @@ -4865,12 +4878,34 @@ def _check_and_update_cudagraph_mode( min_cg_attn_backend = attn_backend.__name__ with update_pass_config(self): + tensor_parallel_size = self.parallel_config.tensor_parallel_size + resolver_tensor_parallel_size = tensor_parallel_size + if ( + self.compilation_config.pass_config.enable_sp + and self.uniform_decode_query_len > 1 + and tensor_parallel_size > 1 + ): + graph_alignment = math.lcm( + self.uniform_decode_query_len, + tensor_parallel_size, + ) + capture_sizes = self.compilation_config.cudagraph_capture_sizes + # vLLM 0.27 has no path for explicit capture sizes that are + # already aligned to both speculative steps and SP. Skip its + # redundant TP adjustment only for that exact case. + if ( + graph_alignment + > max(self.uniform_decode_query_len, tensor_parallel_size) + and capture_sizes + and all(size % graph_alignment == 0 for size in capture_sizes) + ): + resolver_tensor_parallel_size = 1 cudagraph_mode = self.compilation_config.resolve_cudagraph_mode_and_sizes( min_cg_support=min_cg_support, min_cg_attn_backend=min_cg_attn_backend, uniform_decode_query_len=self.uniform_decode_query_len, use_v2_model_runner=False, - tensor_parallel_size=self.parallel_config.tensor_parallel_size, + tensor_parallel_size=resolver_tensor_parallel_size, kv_cache_config=self.kv_cache_config, max_num_reqs=self.max_num_reqs, ) From ac63c664a4bd84bd4ea7eac3f6af849055fba334 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Wed, 26 Aug 2026 00:17:08 -0500 Subject: [PATCH 30/50] fix(quantization): preserve ModelSlim optional metadata Keep checkpoint-level optional metadata out of model weight name mapping so QuaRot rotation descriptors remain available while constructing DSpark draft models. Signed-off-by: maoxx241 --- .../ut/quantization/test_modelslim_config.py | 34 +++++++++++++++++++ vllm_ascend/quantization/modelslim_config.py | 15 ++++++-- 2 files changed, 46 insertions(+), 3 deletions(-) diff --git a/tests/ut/quantization/test_modelslim_config.py b/tests/ut/quantization/test_modelslim_config.py index e56bb0c463db..8584d383d021 100644 --- a/tests/ut/quantization/test_modelslim_config.py +++ b/tests/ut/quantization/test_modelslim_config.py @@ -1,6 +1,8 @@ import json import os import tempfile +from pathlib import Path +from types import SimpleNamespace from unittest.mock import MagicMock, patch import torch @@ -8,8 +10,10 @@ from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase from vllm.model_executor.layers.fused_moe import RoutedExperts from vllm.model_executor.layers.linear import LinearBase +from vllm.model_executor.models.utils import WeightsMapper from tests.ut.base import TestBase +from vllm_ascend.models.llama_eagle3 import get_rotation_path from vllm_ascend.ops.linear import AscendUnquantizedLinearMethod from vllm_ascend.quantization.modelslim_config import ( MODELSLIM_CONFIG_FILENAME, @@ -520,6 +524,36 @@ def test_apply_mapper_with_populated_quant_description(self): self.assertEqual(config.quant_description, {"new_key.weight": "INT8"}) mock_mapper.apply_dict.assert_called_once_with({"old_key.weight": "INT8"}) + def test_apply_mapper_preserves_optional_metadata(self): + optional_metadata = { + "quarot": { + "rotation_map": { + "global_rotation": "optional/quarot.safetensors", + } + } + } + config = AscendModelSlimConfig( + { + "context_proj.weight": "W8A8", + "optional": optional_metadata, + } + ) + draft_mapper = WeightsMapper(orig_to_new_prefix={"": "model."}) + + config.apply_vllm_mapper(draft_mapper) + + self.assertEqual(config.quant_description["model.context_proj.weight"], "W8A8") + self.assertEqual(config.quant_description["optional"], optional_metadata) + self.assertNotIn("model.optional", config.quant_description) + vllm_config = SimpleNamespace( + quant_config=config, + model_config=SimpleNamespace(model="/target"), + ) + self.assertEqual( + get_rotation_path(vllm_config), + Path("/target/optional/quarot.safetensors"), + ) + class TestQuantPrefixMapper(TestBase): def test_qwen3_5_text_backbones_use_packed_module_mappings(self): diff --git a/vllm_ascend/quantization/modelslim_config.py b/vllm_ascend/quantization/modelslim_config.py index e8f1f83e16df..717c5d9da6d8 100644 --- a/vllm_ascend/quantization/modelslim_config.py +++ b/vllm_ascend/quantization/modelslim_config.py @@ -635,8 +635,9 @@ def apply_vllm_mapper(self, hf_to_vllm_mapper: "WeightsMapper"): This method is called by vLLM to apply the model-specific weight mapper to the quantization configuration. It directly uses the forward mapping - (HF -> vLLM) to transform keys in quant_description from HF format to - vLLM format. + (HF -> vLLM) to transform layer keys in quant_description from HF format + to vLLM format. Checkpoint-level metadata under ``optional`` does not + name model parameters and must retain its schema. Args: hf_to_vllm_mapper: The WeightsMapper instance provided by vLLM @@ -649,7 +650,15 @@ def apply_vllm_mapper(self, hf_to_vllm_mapper: "WeightsMapper"): self._mapper_applied = True if self.quant_description: - self.quant_description = hf_to_vllm_mapper.apply_dict(self.quant_description) + optional_metadata = self.quant_description.get("optional") + layer_descriptions = { + name: description + for name, description in self.quant_description.items() + if name != "optional" + } + self.quant_description = hf_to_vllm_mapper.apply_dict(layer_descriptions) + if optional_metadata is not None: + self.quant_description["optional"] = optional_metadata self._add_kvcache_quant_metadata() logger.info("Applied hf_to_vllm_mapper to quant_description keys") From eb118171cbc8c5bddc0d866124a35d3fcb2bd7d4 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Wed, 26 Aug 2026 05:17:34 -0500 Subject: [PATCH 31/50] fix(ci): restore 310P unified-core output path Preserve the existing L0C-to-UB-to-GM fallback for dav_m200 when integrating the fused chunk KDA kernel changes, avoiding unsupported Fixpipe instantiation on 310P. Also apply the required Ruff formatting to the ModelSlim metadata mapping. Signed-off-by: maoxx241 --- .../block/block_mmad_pingpong_tla_multi.hpp | 79 ++++++++++++++++++- vllm_ascend/quantization/modelslim_config.py | 4 +- 2 files changed, 76 insertions(+), 7 deletions(-) diff --git a/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla_multi.hpp b/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla_multi.hpp index a813e39fe473..794a0ee63b35 100644 --- a/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla_multi.hpp +++ b/csrc/moe/common/kernel_utils/block/block_mmad_pingpong_tla_multi.hpp @@ -203,7 +203,12 @@ struct BlockMmadTla < CATLASS_DEVICE BlockMmadTla(Arch::Resource &resource, uint32_t l1BufAddrStart = 0) { +#ifdef CATLASS_UNIFIED_CORE + resourcePtr = &resource; + { +#else if ASCEND_IS_AIC { +#endif uint32_t l1AOffset = l1BufAddrStart; uint32_t l1BOffset = l1BufAddrStart + L1A_TILE_SIZE * L1A_STAGES; // Init buffers @@ -253,8 +258,11 @@ struct BlockMmadTla < CATLASS_DEVICE void preSetFlags() { - +#ifdef CATLASS_UNIFIED_CORE + { +#else if ASCEND_IS_AIC { +#endif // use HF32 when USE_HF32_MODE is true if constexpr (USE_HF32_MODE) { AscendC::SetHF32Mode(true); @@ -294,7 +302,11 @@ struct BlockMmadTla < CATLASS_DEVICE void finalWaitFlags() { +#ifdef CATLASS_UNIFIED_CORE + { +#else if ASCEND_IS_AIC { +#endif if constexpr (USE_HF32_MODE) { AscendC::SetHF32Mode(false); } @@ -344,11 +356,12 @@ struct BlockMmadTla < using CopyGmToL1B = typename TileCopy_::template CopyGmToL1B; CopyGmToL1A copyGmToL1A; CopyGmToL1B copyGmToL1B; -#if (defined (CATLASS_ARCH) && CATLASS_ARCH == 2201) +#ifdef CATLASS_UNIFIED_CORE + // 310P: no Fixpipe, no DataCopyCO12Dst. L0C exits via DataCopy L0C→UB then UB→GM. +#elif (defined (CATLASS_ARCH) && CATLASS_ARCH == 2201) using CopyL0CToGm = typename TileCopy_::template CopyL0CToGm; CopyL0CToGm copyL0CToDst; -#endif -#if (defined (CATLASS_ARCH) && CATLASS_ARCH == 3510) +#elif (defined (CATLASS_ARCH) && CATLASS_ARCH == 3510) using CopyL0CToDst = typename TileCopy_::template CopyL0CToDst; CopyL0CToDst copyL0CToDst; #endif @@ -649,6 +662,60 @@ struct BlockMmadTla < } // copy block out +#ifdef CATLASS_UNIFIED_CORE + { + // 310P unified core: L0C→UB via DataCopy, then UB→GM. + // No Fixpipe or DataCopyCO12Dst on dav_m200. + uint32_t mAligned = (mBlockActual + 15) / 16 * 16; + uint32_t nAligned = (nBlockActual + 15) / 16 * 16; + uint32_t tileElems = mAligned * nAligned; + uint32_t tileBytes = tileElems * sizeof(ElementAccumulator); + + // UB temp for L0C→UB transfer. Offset 0 is safe: on unified core, + // the matrix multiply and epilogue run sequentially so UB is not shared concurrently. + // The epilogue allocates its own UB regions at higher offsets (≥32KB). + AscendC::LocalTensor co2Temp = + resourcePtr->ubBuf.template GetBufferByByte(0); + + AscendC::PipeBarrier(); + + // L0C → UB: BLOCK_MODE_MATRIX copies raw NZ fractals to UB + // For float: blockLen unit = 1024B (one 16×16 fractal) + AscendC::DataCopyParams l0c2ubParams; + l0c2ubParams.blockCount = static_cast(nAligned / 16); + l0c2ubParams.blockLen = static_cast(mAligned / 16); + l0c2ubParams.srcStride = 0; + l0c2ubParams.dstStride = 0; + AscendC::DataCopyEnhancedParams enhParams; + enhParams.blockMode = AscendC::BlockMode::BLOCK_MODE_MATRIX; + AscendC::DataCopy(co2Temp, l0CTensorList[l0CListId], l0c2ubParams, enhParams); + AscendC::PipeBarrier(); + + // UB → GM: fractal-by-fractal with strided DataCopy (NZ→ND deformat) + // NZ in UB: [N/16 Z-cols][M/16 fractals][16 rows][16 cols] + // ND in GM: [M rows][N cols] + auto dstOffset = tensorC.layout()(tensorC.coord()); + uint32_t gmStride = tla::get<0>(tensorC.stride()); + uint32_t mFracs = mAligned / 16; + uint32_t nFracs = nAligned / 16; + for (uint32_t nf = 0; nf < nFracs; nf++) { + for (uint32_t mf = 0; mf < mFracs; mf++) { + uint32_t ubOff = (nf * mFracs + mf) * 256; + uint32_t gmRow = mf * 16; + uint32_t gmCol = nf * 16; + uint32_t gmOff = dstOffset + gmRow * gmStride + gmCol; + AscendC::DataCopyParams fracParams; + fracParams.blockCount = 16; + fracParams.blockLen = static_cast(16 * sizeof(ElementAccumulator) / 32); + fracParams.srcStride = 0; + fracParams.dstStride = static_cast((gmStride - 16) * sizeof(ElementAccumulator) / 32); + AscendC::DataCopy(tensorC.data()[gmOff], co2Temp[ubOff], fracParams); + } + } + AscendC::PipeBarrier(); + l0CListId = (l0CListId + 1 < L0C_STAGES) ? (l0CListId + 1) : 0; + } +#else if constexpr (!ENABLE_UNIT_FLAG) { AscendC::SetFlag(l0CEventList[l0CListId]); AscendC::WaitFlag(l0CEventList[l0CListId]); @@ -658,6 +725,7 @@ struct BlockMmadTla < } else { copyL0CToDst(tensorC, tensorL0C, 0b11); } +#endif } protected: @@ -679,6 +747,9 @@ struct BlockMmadTla < AscendC::LocalTensor l0CTensorList[L0C_STAGES]; AscendC::LocalTensor l1BiasTensor; AscendC::LocalTensor l0BiasTensor; +#ifdef CATLASS_UNIFIED_CORE + Arch::Resource* resourcePtr{nullptr}; +#endif // Multi-stage event id list int32_t l1AEventList[L1A_STAGES]; diff --git a/vllm_ascend/quantization/modelslim_config.py b/vllm_ascend/quantization/modelslim_config.py index 717c5d9da6d8..0fb56952affb 100644 --- a/vllm_ascend/quantization/modelslim_config.py +++ b/vllm_ascend/quantization/modelslim_config.py @@ -652,9 +652,7 @@ def apply_vllm_mapper(self, hf_to_vllm_mapper: "WeightsMapper"): if self.quant_description: optional_metadata = self.quant_description.get("optional") layer_descriptions = { - name: description - for name, description in self.quant_description.items() - if name != "optional" + name: description for name, description in self.quant_description.items() if name != "optional" } self.quant_description = hf_to_vllm_mapper.apply_dict(layer_descriptions) if optional_metadata is not None: From 530884a3ea8be4d75e58676d8f5e2fa4b8cfc1b5 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Wed, 26 Aug 2026 06:27:49 -0500 Subject: [PATCH 32/50] fix(kimi-k3): align CI coverage with vLLM versions Restore Ascend MLA scale storage accounting and the v0.27.1 UniformType block-table contract. Select the appropriate DSpark checkpoint filtering interface for pinned and current vLLM, and update CPU test fixtures to match the runtime metadata and shared-expert APIs. Signed-off-by: maoxx241 --- .../fused_moe/test_shared_fused_moe_310.py | 5 +++- tests/ut/models/test_kimi_k3_adapter.py | 15 ++++++++--- tests/ut/ops/test_fused_moe.py | 2 ++ tests/ut/ops/test_moe_mlp.py | 11 +++++--- tests/ut/patch/platform/test_patch_eplb.py | 8 ++++-- .../platform/test_prefix_cache_cp_patches.py | 10 ++++--- tests/ut/spec_decode/test_dspark_proposer.py | 2 ++ vllm_ascend/core/kv_cache_interface.py | 4 +++ vllm_ascend/models/kimi_k3_dspark.py | 14 +++++++--- .../patch/platform/patch_kv_cache_utils.py | 26 +++++++++++++++++-- 10 files changed, 78 insertions(+), 19 deletions(-) diff --git a/tests/ut/_310p/fused_moe/test_shared_fused_moe_310.py b/tests/ut/_310p/fused_moe/test_shared_fused_moe_310.py index 815fe04ec060..2f7cc721ee23 100644 --- a/tests/ut/_310p/fused_moe/test_shared_fused_moe_310.py +++ b/tests/ut/_310p/fused_moe/test_shared_fused_moe_310.py @@ -229,7 +229,10 @@ def test_forward_impl_310_returns_current_runner_contract(monkeypatch, has_share router_logits = torch.randn(2, 3) routed_out = torch.randn(2, 4) shared_out = torch.randn(2, 4) - ascend_shared_experts = SimpleNamespace(forward=MagicMock(return_value=shared_out)) + ascend_shared_experts = SimpleNamespace( + prepare_input_before_routed_experts=MagicMock(return_value=(hidden_states, None)), + forward=MagicMock(return_value=shared_out), + ) routed_events = FusedMoEEvents( before_routed_experts=None, after_routed_experts=None, diff --git a/tests/ut/models/test_kimi_k3_adapter.py b/tests/ut/models/test_kimi_k3_adapter.py index 51d91820e33f..98697d733437 100644 --- a/tests/ut/models/test_kimi_k3_adapter.py +++ b/tests/ut/models/test_kimi_k3_adapter.py @@ -22,6 +22,7 @@ AscendK3DSparkDecoderLayer, AscendK3DSparkForCausalLM, ) +from vllm_ascend.utils import vllm_version_is def test_ascend_attn_res_matches_canonical_k3_math(): @@ -675,9 +676,12 @@ def test_k3_dspark_load_weights_keeps_per_layer_context_kv(monkeypatch): seen_names: list[str] = [] class CapturingLoader: - def __init__(self, loaded_model, *, skip_substrs): + def __init__(self, loaded_model, **kwargs): assert loaded_model is model - assert skip_substrs == list(model.checkpoint_skip_substrs) + if vllm_version_is("0.27.1"): + assert kwargs["skip_substrs"] == list(model.checkpoint_skip_substrs) + else: + assert kwargs == {} def load_weights(self, weights, *, mapper): assert mapper is model.hf_to_vllm_mapper @@ -710,9 +714,12 @@ def test_k3_dspark_reuses_modelslim_rotation_loader(monkeypatch): seen_weights: list[tuple[str, torch.Tensor]] = [] class CapturingLoader: - def __init__(self, loaded_model, *, skip_substrs): + def __init__(self, loaded_model, **kwargs): assert loaded_model is model - assert skip_substrs == list(model.checkpoint_skip_substrs) + if vllm_version_is("0.27.1"): + assert kwargs["skip_substrs"] == list(model.checkpoint_skip_substrs) + else: + assert kwargs == {} def load_weights(self, weights, *, mapper): assert mapper is model.hf_to_vllm_mapper diff --git a/tests/ut/ops/test_fused_moe.py b/tests/ut/ops/test_fused_moe.py index a34a74b36d37..bed3f7be40d2 100644 --- a/tests/ut/ops/test_fused_moe.py +++ b/tests/ut/ops/test_fused_moe.py @@ -941,6 +941,7 @@ def test_w8a8_shared_situ_uses_dequant_situ_quant(monkeypatch): shared_experts_module.torch_npu, "npu_dynamic_quant", MagicMock(return_value=(quantized_input, input_scale)), + raising=False, ) monkeypatch.setattr( shared_experts_module.torch_npu, @@ -984,6 +985,7 @@ def test_w4a8_mxfp_shared_situ_uses_situ_mx_quant(monkeypatch): shared_experts_module.torch_npu, "npu_dynamic_mx_quant", MagicMock(return_value=(quantized_input, input_scale)), + raising=False, ) monkeypatch.setattr( shared_experts_module.torch.ops, diff --git a/tests/ut/ops/test_moe_mlp.py b/tests/ut/ops/test_moe_mlp.py index 643d4fb327c1..61c3c01a026c 100644 --- a/tests/ut/ops/test_moe_mlp.py +++ b/tests/ut/ops/test_moe_mlp.py @@ -7,7 +7,7 @@ import torch import torch_npu # noqa: F401 -- registers torch.npu used by the module under test from torch.nn import functional as F -from vllm.config import VllmConfig, set_current_vllm_config +from vllm.config import CompilationConfig, VllmConfig, set_current_vllm_config from vllm.model_executor.layers.activation import SituAndMul from vllm.model_executor.layers.fused_moe.activation import MoEActivation @@ -30,6 +30,10 @@ MXFP4_TEST_DTYPE = getattr(torch, "float4_e2m1fn_x2", torch.float16) +def _custom_op_vllm_config(): + return SimpleNamespace(compilation_config=CompilationConfig(custom_ops=["none"])) + + class TestCumsumGroupList(unittest.TestCase): glist_dict: ClassVar[dict[int, torch.Tensor]] @@ -596,6 +600,7 @@ def test_dynamic_eplb_tensor_lists_reach_both_grouped_matmuls(self): stream_patch, evt = _patch_npu_stream() with ( + patch(f"{MOE_MLP}._EXTRA_CTX", MagicMock(moe_comm_type=-1)), stream_patch, patch("torch_npu.npu_grouped_matmul", return_value=[gate_up_out], create=True) as mock_gmm1, patch( @@ -679,7 +684,7 @@ def test_antiquant_weights_use_native_situ_between_grouped_matmuls(self): stream_patch, evt = _patch_npu_stream() with ( - set_current_vllm_config(VllmConfig()), + set_current_vllm_config(_custom_op_vllm_config()), stream_patch, patch(f"{MOE_MLP}._EXTRA_CTX", MagicMock(moe_comm_type=-1)), patch( @@ -718,7 +723,7 @@ def test_w4a16_mxfp_uses_native_situ_without_activation_requant(self): stream_patch, evt = _patch_npu_stream() with ( - set_current_vllm_config(VllmConfig()), + set_current_vllm_config(_custom_op_vllm_config()), stream_patch, patch(f"{MOE_MLP}._EXTRA_CTX", MagicMock(moe_comm_type=-1)), patch("torch_npu.npu_grouped_matmul", return_value=[gate_up_out], create=True), diff --git a/tests/ut/patch/platform/test_patch_eplb.py b/tests/ut/patch/platform/test_patch_eplb.py index b77d98e88fad..e9fa0374c552 100644 --- a/tests/ut/patch/platform/test_patch_eplb.py +++ b/tests/ut/patch/platform/test_patch_eplb.py @@ -3,7 +3,7 @@ from contextlib import contextmanager from types import SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch from vllm.config import EPLBConfig, ParallelConfig, VllmConfig from vllm.config import parallel as parallel_module @@ -32,7 +32,11 @@ def _npu_parallel_config_platform(): def test_parallel_and_vllm_config_keep_upstream_validation(): - with _npu_parallel_config_platform(): + with ( + _npu_parallel_config_platform(), + patch("vllm_ascend.logger.configure_ascend_file_logging"), + patch("vllm_ascend.logger.configure_ascend_logging"), + ): parallel_config = ParallelConfig( tensor_parallel_size=2, enable_expert_parallel=True, diff --git a/tests/ut/patch/platform/test_prefix_cache_cp_patches.py b/tests/ut/patch/platform/test_prefix_cache_cp_patches.py index a992b9dc757b..fcec9e286752 100644 --- a/tests/ut/patch/platform/test_prefix_cache_cp_patches.py +++ b/tests/ut/patch/platform/test_prefix_cache_cp_patches.py @@ -38,6 +38,7 @@ group_and_unify_kv_cache_specs, ) from vllm_ascend.patch.platform.patch_mamba_manager import AscendMambaManager +from vllm_ascend.utils import vllm_version_is def _make_hybrid_kv_cache_config( @@ -399,10 +400,9 @@ def test_kimi_k3_gqa_mixed_groups_use_expected_physical_layout(monkeypatch) -> N assert sum(tensor.size for tensor in tensors) == available_memory -def test_kimi_k3_gqa_mixed_grouping_falls_back_on_partial_signature() -> None: +def test_kimi_k3_gqa_mixed_grouping_falls_back_on_unrecognized_layer() -> None: specs = _make_kimi_k3_dspark_kv_cache_specs() - draft_layer = "model.layers.93.self_attn.attn" - specs.pop(draft_layer) + specs["unrecognized.layer"] = next(iter(specs.values())) assert _get_kimi_k3_dspark_mixed_kv_cache_groups(specs) is None @@ -560,6 +560,10 @@ def _fake_orig(*args, **kwargs): @pytest.mark.parametrize("num_prefill_lookahead", [0, 8]) +@pytest.mark.skipif( + vllm_version_is("0.27.1"), + reason="num_prefill_lookahead was added to the coordinator after v0.27.1", +) def test_get_kv_cache_coordinator_forwards_prefill_lookahead( monkeypatch, num_prefill_lookahead: int, diff --git a/tests/ut/spec_decode/test_dspark_proposer.py b/tests/ut/spec_decode/test_dspark_proposer.py index bed7ab8d69f0..2e88b243f351 100644 --- a/tests/ut/spec_decode/test_dspark_proposer.py +++ b/tests/ut/spec_decode/test_dspark_proposer.py @@ -219,6 +219,8 @@ def _call_set_inputs_first_pass(proposer, *, num_reqs, block_size): query_start_loc=torch.arange(num_reqs + 1, dtype=torch.int32) * block_size, query_start_loc_cpu=torch.zeros(num_reqs + 1, dtype=torch.int32), seq_lens=torch.full((num_reqs,), 128, dtype=torch.int32), + _seq_lens_cpu=torch.full((num_reqs,), 128, dtype=torch.int32), + seq_lens_cpu=torch.full((num_reqs,), 128, dtype=torch.int32), max_seq_len=128, ) proposer.set_inputs_first_pass( diff --git a/vllm_ascend/core/kv_cache_interface.py b/vllm_ascend/core/kv_cache_interface.py index cab6c85bae84..fb43662779ba 100644 --- a/vllm_ascend/core/kv_cache_interface.py +++ b/vllm_ascend/core/kv_cache_interface.py @@ -60,6 +60,10 @@ def real_page_size_bytes(self) -> int: * (self.head_size * get_dtype_size(self.dtype) + self.scale_dim * get_dtype_size(self.scale_dtype)) ) + @property + def unpadded_page_size_bytes(self) -> int: + return self.real_page_size_bytes + @classmethod def merge(cls, specs: list[Self]) -> Self: assert all(isinstance(spec, MLAAttentionSpec) for spec in specs), ( diff --git a/vllm_ascend/models/kimi_k3_dspark.py b/vllm_ascend/models/kimi_k3_dspark.py index 64e51bc70e86..1609543f43b8 100644 --- a/vllm_ascend/models/kimi_k3_dspark.py +++ b/vllm_ascend/models/kimi_k3_dspark.py @@ -47,6 +47,7 @@ process_weight, ) from vllm_ascend.ops.rotary_embedding import get_cos_and_sin_mla +from vllm_ascend.utils import vllm_version_is def _uses_causal_draft_attention(config) -> bool: @@ -287,10 +288,15 @@ def load_weights( quantization-aware per-layer projections, so use vLLM's public loader interface without creating that extra packed parameter. """ - loader = AutoWeightsLoader( - self, - skip_substrs=list(self.checkpoint_skip_substrs), - ) + if vllm_version_is("0.27.1"): + loader = AutoWeightsLoader( + self, + skip_substrs=list(self.checkpoint_skip_substrs), + ) + else: + # Current vLLM drops the training-only and shared checkpoint + # weights in hf_to_vllm_mapper instead of AutoWeightsLoader. + loader = AutoWeightsLoader(self) rotation_weight = None if self.rotation_path is not None: rotation_weight = get_rotation_matrix(self.rotation_path) diff --git a/vllm_ascend/patch/platform/patch_kv_cache_utils.py b/vllm_ascend/patch/platform/patch_kv_cache_utils.py index 6dc24f26e2d3..6aa468dbc268 100644 --- a/vllm_ascend/patch/platform/patch_kv_cache_utils.py +++ b/vllm_ascend/patch/platform/patch_kv_cache_utils.py @@ -22,12 +22,34 @@ get_kv_cache_spec_kind, ) +from vllm_ascend.utils import vllm_version_is + _KIMI_K3_TARGET_LAYER_PREFIX = "language_model.model.layers." _KIMI_K3_DRAFT_LAYER_PREFIX = "model.layers." _orig_resolve_kv_cache_block_sizes = vllm.v1.core.kv_cache_utils.resolve_kv_cache_block_sizes _orig_get_kv_cache_groups_uniform_page_size = vllm.v1.core.kv_cache_utils._get_kv_cache_groups_uniform_page_size +if vllm_version_is("0.27.1"): + + def _uniform_type_max_num_blocks_per_req( + self: UniformTypeKVCacheSpecs, + vllm_config: VllmConfig, + max_len: int, + ) -> int: + """Use the per-layer block-table width, matching current vLLM.""" + widths = {spec.max_num_blocks_per_req(vllm_config, max_len) for spec in self.kv_cache_specs.values()} + assert len(widths) == 1, ( + "All layers in the same KV cache group must need the same number " + f"of block table entries, got {sorted(widths)}." + ) + return next(iter(widths)) + + UniformTypeKVCacheSpecs.max_num_blocks_per_req = ( # type: ignore[method-assign] + _uniform_type_max_num_blocks_per_req + ) + + def _ascend_resolve_kv_cache_block_sizes( kv_cache_config: KVCacheConfig, vllm_config: VllmConfig, @@ -78,8 +100,8 @@ def _get_kimi_k3_dspark_mixed_kv_cache_groups( Block and page sizes are resolved by the runtime and intentionally not fixed here: TP8 and TP16 produce different sizes but the same ownership - relation. A partial or incompatible signature falls back to vLLM's generic - hybrid grouping. + relation. An unrecognized or incompatible signature falls back to vLLM's + generic hybrid grouping. """ target_attention_specs = { name: spec From fb5c390dbadbb11cf5ace6b25828e6396a301b2c Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Wed, 26 Aug 2026 07:22:25 -0500 Subject: [PATCH 33/50] fix(ci): execute SiTU custom op eagerly in CPU tests Enable the custom-op dispatch in the two CPU-only SiTU flow tests so the out-of-tree native implementation runs eagerly instead of sending forward_native through torch.compile, which requires torch.npu on the CPU runner. Signed-off-by: maoxx241 --- tests/ut/ops/test_moe_mlp.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/ut/ops/test_moe_mlp.py b/tests/ut/ops/test_moe_mlp.py index 61c3c01a026c..0583da394524 100644 --- a/tests/ut/ops/test_moe_mlp.py +++ b/tests/ut/ops/test_moe_mlp.py @@ -31,7 +31,7 @@ def _custom_op_vllm_config(): - return SimpleNamespace(compilation_config=CompilationConfig(custom_ops=["none"])) + return SimpleNamespace(compilation_config=CompilationConfig(custom_ops=["all"])) class TestCumsumGroupList(unittest.TestCase): From 45804ed3cd1e171ad80b317c35bcffc8acd9d808 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Wed, 26 Aug 2026 08:50:38 -0500 Subject: [PATCH 34/50] fix(mla): preserve defaults for existing attention models Keep the output gate optional for models such as DeepSeek and initialize the MLAPO native-weight mode before either FA quantization or MLAPO weight processing uses it. Signed-off-by: maoxx241 --- vllm_ascend/attention/mla_v1.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/vllm_ascend/attention/mla_v1.py b/vllm_ascend/attention/mla_v1.py index da3e74ce817f..d67f4d24d564 100644 --- a/vllm_ascend/attention/mla_v1.py +++ b/vllm_ascend/attention/mla_v1.py @@ -844,7 +844,7 @@ def __init__( self.q_proj = kwargs["q_proj"] if self.q_lora_rank is None else kwargs["q_b_proj"] self.kv_b_proj = kwargs["kv_b_proj"] self.o_proj = kwargs["o_proj"] - self.g_proj = kwargs["g_proj"] + self.g_proj = kwargs.get("g_proj") self.use_output_gate = self.g_proj is not None self.use_mla_rope = kwargs["use_mla_rope"] self.vllm_config = get_current_vllm_config() @@ -874,6 +874,7 @@ def __init__( self.head_padding = self.num_heads_padded - self.num_heads self.mlapo_num_heads = self.num_heads self.mlapo_weight_quant_mode = 3 + self._mlapo_uses_native_weights = False @staticmethod def update_graph_params( From d7a93b20e6a4bd34a27943cb16a42e438bcfb406 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Wed, 26 Aug 2026 10:21:12 -0500 Subject: [PATCH 35/50] fix(mla): default existing models to RoPE Keep MLA RoPE enabled when a model does not provide the Kimi-specific use_mla_rope override. Kimi K3 still explicitly disables it. Signed-off-by: maoxx241 --- vllm_ascend/attention/mla_v1.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/vllm_ascend/attention/mla_v1.py b/vllm_ascend/attention/mla_v1.py index d67f4d24d564..fd92e95df1cc 100644 --- a/vllm_ascend/attention/mla_v1.py +++ b/vllm_ascend/attention/mla_v1.py @@ -846,7 +846,7 @@ def __init__( self.o_proj = kwargs["o_proj"] self.g_proj = kwargs.get("g_proj") self.use_output_gate = self.g_proj is not None - self.use_mla_rope = kwargs["use_mla_rope"] + self.use_mla_rope = kwargs.get("use_mla_rope", True) self.vllm_config = get_current_vllm_config() self.kv_a_proj_with_mqa = kwargs.get("kv_a_proj_with_mqa") self.kv_a_layernorm = kwargs.get("kv_a_layernorm") From 342ee0c51f640d01ff347f77ff4f6b54497c3f33 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Wed, 26 Aug 2026 17:54:09 -0500 Subject: [PATCH 36/50] fix(kv-cache): scope RoPE mode to MLA backends Only the Ascend MLA backend consumes the per-layer RoPE mode. Keep SFA, cache-only, and DeepSeek indexer layers grouped by their own backends without reading an unrelated implementation property. Signed-off-by: maoxx241 --- tests/ut/worker/a2/test_model_runner_v1.py | 44 +++++++++++++++------- vllm_ascend/worker/model_runner_v1.py | 7 ++-- 2 files changed, 35 insertions(+), 16 deletions(-) diff --git a/tests/ut/worker/a2/test_model_runner_v1.py b/tests/ut/worker/a2/test_model_runner_v1.py index 81a812e9e621..8c48316668cd 100644 --- a/tests/ut/worker/a2/test_model_runner_v1.py +++ b/tests/ut/worker/a2/test_model_runner_v1.py @@ -15,6 +15,7 @@ UniformTypeKVCacheSpecs, ) +from vllm_ascend.attention.mla_v1 import AscendMLABackend from vllm_ascend.attention.utils import get_sfa_qsfa_packed_head_dim from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec, AscendSFAIndexerCacheSpec from vllm_ascend.spec_decode.dspark_proposer import AscendDSparkProposer @@ -171,15 +172,24 @@ def test_allocate_kv_cache_uses_layer_spec_for_draft_gqa(self): self.assertEqual(v_cache_raw.numel(), kv_cache_spec.page_size_bytes) @patch("vllm_ascend.worker.model_runner_v1.get_layers_from_vllm_config") - def test_mla_rope_modes_use_separate_metadata_groups(self, mock_get_layers): + def test_mla_rope_modes_and_cache_layers_use_separate_metadata_groups(self, mock_get_layers): class FakeBuilder: def __init__(self, _spec, layer_names, _config, _device): self.layer_names = layer_names - class FakeBackend: + class FakeMLABackend(AscendMLABackend): @classmethod def full_cls_name(cls): - return "test.FakeBackend" + return "test.FakeMLABackend" + + @classmethod + def get_builder_cls(cls): + return FakeBuilder + + class FakeCacheBackend: + @classmethod + def full_cls_name(cls): + return "test.FakeCacheBackend" @classmethod def get_builder_cls(cls): @@ -192,15 +202,17 @@ def get_builder_cls(cls): target_layer = "language_model.model.layers.0.self_attn.attn" draft_layer = "model.layers.0.self_attn.attn" + cache_layer = "language_model.model.layers.0.self_attn.indexer.k_cache" + target_attn = MagicMock(spec=MLAAttention) + target_attn.impl = SimpleNamespace(use_mla_rope=False) + target_attn.get_attn_backend.return_value = FakeMLABackend + draft_attn = MagicMock(spec=MLAAttention) + draft_attn.impl = SimpleNamespace(use_mla_rope=True) + draft_attn.get_attn_backend.return_value = FakeMLABackend mock_get_layers.return_value = { - target_layer: SimpleNamespace( - impl=SimpleNamespace(use_mla_rope=False), - get_attn_backend=lambda: FakeBackend, - ), - draft_layer: SimpleNamespace( - impl=SimpleNamespace(use_mla_rope=True), - get_attn_backend=lambda: FakeBackend, - ), + target_layer: target_attn, + draft_layer: draft_attn, + cache_layer: SimpleNamespace(get_attn_backend=lambda: FakeCacheBackend), } specs = { target_layer: AscendMLAAttentionSpec( @@ -215,6 +227,12 @@ def get_builder_cls(cls): head_size=576, dtype=torch.bfloat16, ), + cache_layer: AscendMLAAttentionSpec( + block_size=16, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + ), } group_spec = UniformTypeKVCacheSpecs.from_specs(specs) self.assertIsNotNone(group_spec) @@ -224,7 +242,7 @@ def get_builder_cls(cls): kv_cache_tensors=[], kv_cache_groups=[ KVCacheGroupSpec( - layer_names=[target_layer, draft_layer], + layer_names=[target_layer, draft_layer, cache_layer], kv_cache_spec=group_spec, ) ], @@ -235,7 +253,7 @@ def get_builder_cls(cls): self.assertEqual(len(runner.attn_groups), 1) self.assertEqual( {tuple(group.layer_names) for group in runner.attn_groups[0]}, - {(target_layer,), (draft_layer,)}, + {(target_layer,), (draft_layer,), (cache_layer,)}, ) def test_explicit_capture_sizes_must_align_spec_decode_and_sp(self): diff --git a/vllm_ascend/worker/model_runner_v1.py b/vllm_ascend/worker/model_runner_v1.py index 5b4937902b98..1faadb90a8de 100644 --- a/vllm_ascend/worker/model_runner_v1.py +++ b/vllm_ascend/worker/model_runner_v1.py @@ -4638,12 +4638,13 @@ def backend_supports_kernel_block_size( # they are cached correctly, there will be different objects per # layer. for layer_name in kv_cache_group_spec.layer_names: + layer = layers[layer_name] layer_kv_cache_spec = kv_cache_group_spec.kv_cache_spec if isinstance(layer_kv_cache_spec, UniformTypeKVCacheSpecs): layer_kv_cache_spec = layer_kv_cache_spec.kv_cache_specs[layer_name] # Prefer the backend declared by the layer itself. Some # indexer-cache layers require their own metadata builder. - attn_backend = layers[layer_name].get_attn_backend() + attn_backend = layer.get_attn_backend() if ( isinstance(layer_kv_cache_spec, AscendSFAIndexerCacheSpec) and not backend_supports_kernel_block_size( @@ -4656,8 +4657,8 @@ def backend_supports_kernel_block_size( attn_backend = AscendSFAIndexerBackend full_cls_name = attn_backend.full_cls_name() use_mla_rope = ( - getattr(layers[layer_name].impl, "use_mla_rope", None) - if isinstance(layer_kv_cache_spec, AscendMLAAttentionSpec) + layer.impl.use_mla_rope + if issubclass(attn_backend, AscendMLABackend) else None ) key = (full_cls_name, layer_kv_cache_spec, use_mla_rope) From 2c83f016e0ccb1a80ec7a9e7d4ca7756507d1dc4 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Thu, 27 Aug 2026 10:04:02 +0800 Subject: [PATCH 37/50] refactor(mla): remove no-rope identity metadata Represent disabled MLA RoPE with absent cos/sin metadata and rely on the existing explicit no-RoPE execution paths. Keep target and draft separation in the model-runner grouping layer instead of duplicating a mixed-mode assertion in the metadata builder. Signed-off-by: maoxx241 --- tests/ut/attention/a2/test_mla_v1.py | 36 --------------------------- tests/ut/ops/test_rotary_embedding.py | 20 --------------- vllm_ascend/attention/mla_v1.py | 28 +++++++++++---------- vllm_ascend/ops/rotary_embedding.py | 9 ------- 4 files changed, 15 insertions(+), 78 deletions(-) diff --git a/tests/ut/attention/a2/test_mla_v1.py b/tests/ut/attention/a2/test_mla_v1.py index ab87f4fe7245..9977ca2943fa 100644 --- a/tests/ut/attention/a2/test_mla_v1.py +++ b/tests/ut/attention/a2/test_mla_v1.py @@ -543,42 +543,6 @@ def test_metadata_builder_uses_target_layer_nope_mode(self): self.assertFalse(builder.use_mla_rope) - def test_metadata_builder_rejects_mixed_rope_modes(self): - mock_vllm_config = MagicMock() - mock_vllm_config.model_config.max_model_len = 1024 - mock_vllm_config.model_config.get_head_size.return_value = 64 - mock_vllm_config.model_config.dtype = torch.float16 - mock_vllm_config.model_config.hf_text_config = SimpleNamespace( - qk_rope_head_dim=64, - mla_use_nope=False, - ) - mock_vllm_config.cache_config.block_size = 16 - mock_vllm_config.scheduler_config.max_num_seqs = 4 - mock_vllm_config.scheduler_config.enable_chunked_prefill = False - mock_vllm_config.speculative_config = None - mock_vllm_config.compilation_config.static_forward_context = { - "rope.self_attn": SimpleNamespace( - impl=SimpleNamespace(use_mla_rope=True), - ), - "nope.self_attn": SimpleNamespace( - impl=SimpleNamespace(use_mla_rope=False), - ), - } - - with ( - patch( - "vllm_ascend.attention.mla_v1.get_ascend_config", - return_value=MagicMock(), - ), - self.assertRaisesRegex(AssertionError, "separate KV cache groups"), - ): - AscendMLAMetadataBuilder( - None, - ["rope.self_attn", "nope.self_attn"], - mock_vllm_config, - "cpu", - ) - def test_ascend_mla_metadata_builder_spec_decode(self): mock_vllm_config = MagicMock() mock_vllm_config.model_config.max_model_len = 1024 diff --git a/tests/ut/ops/test_rotary_embedding.py b/tests/ut/ops/test_rotary_embedding.py index ec6484701b06..37db24ec55bc 100644 --- a/tests/ut/ops/test_rotary_embedding.py +++ b/tests/ut/ops/test_rotary_embedding.py @@ -21,7 +21,6 @@ import torch from vllm.model_executor.layers.rotary_embedding import RotaryEmbedding, YaRNScalingRotaryEmbedding -from vllm_ascend.ops import rotary_embedding as rotary_embedding_ops from vllm_ascend.ops.rotary_embedding import AscendRotaryEmbedding, AscendYaRNRotaryEmbedding HEAD_SIZE = 64 @@ -33,25 +32,6 @@ NUM_HEADS = 2 -def test_get_identity_cos_and_sin_mla_skips_rotary_cache(monkeypatch): - cos_buffer = torch.ones(4, 1, 1, ROTARY_DIM) - sin_buffer = torch.zeros(4, 1, 1, ROTARY_DIM) - monkeypatch.setattr(rotary_embedding_ops, "_cos_mla", cos_buffer) - monkeypatch.setattr(rotary_embedding_ops, "_sin_mla", sin_buffer) - monkeypatch.setattr(rotary_embedding_ops, "_cos_cache", None) - monkeypatch.setattr(rotary_embedding_ops, "_sin_cache", None) - - cos, sin = rotary_embedding_ops.get_identity_cos_and_sin_mla( - torch.tensor([1, 3]), - use_cache=True, - ) - - assert cos.data_ptr() == cos_buffer.data_ptr() - assert sin.data_ptr() == sin_buffer.data_ptr() - torch.testing.assert_close(cos, torch.ones_like(cos)) - torch.testing.assert_close(sin, torch.zeros_like(sin)) - - def _make_tensors(seq_len=SEQ_LEN, num_heads=NUM_HEADS, head_size=HEAD_SIZE): positions = torch.arange(seq_len, dtype=torch.long) query = torch.randn(seq_len, num_heads * head_size) diff --git a/vllm_ascend/attention/mla_v1.py b/vllm_ascend/attention/mla_v1.py index fd92e95df1cc..0304b59745a5 100644 --- a/vllm_ascend/attention/mla_v1.py +++ b/vllm_ascend/attention/mla_v1.py @@ -46,10 +46,7 @@ from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec from vllm_ascend.device.device_op import DeviceOperator from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.attention_fence import record_attention_compute_start -from vllm_ascend.ops.rotary_embedding import ( - get_cos_and_sin_mla, - get_identity_cos_and_sin_mla, -) +from vllm_ascend.ops.rotary_embedding import get_cos_and_sin_mla from vllm_ascend.quantization.methods.w8a8_mxfp8 import AscendW8A8MXFP8DynamicLinearMethod from vllm_ascend.quantization.methods.w8a8_static import AscendW8A8LinearMethod from vllm_ascend.quantization.utils import enable_fa_quant @@ -301,9 +298,8 @@ def __init__( self.reorder_batch_threshold = self.decode_threshold self.rope_dim = self.model_config.hf_text_config.qk_rope_head_dim static_forward_context = vllm_config.compilation_config.static_forward_context - layer_rope_modes = {static_forward_context[layer_name].impl.use_mla_rope for layer_name in layer_names} - assert len(layer_rope_modes) <= 1, "MLA layers with and without RoPE must use separate KV cache groups." - self.use_mla_rope = layer_rope_modes.pop() if layer_rope_modes else True + # MLA layers are grouped by RoPE mode before metadata builders are created. + self.use_mla_rope = static_forward_context[layer_names[0]].impl.use_mla_rope if layer_names else True self.cos_cache = None self.sin_cache = None @@ -604,8 +600,10 @@ def build_prefill_metadata( prefill_query_start_loc = query_start_loc[reqs_start:] - query_start_loc[reqs_start] prefill_input_positions = input_positions[tokens_start:] - cos_sin_getter = get_cos_and_sin_mla if self.use_mla_rope else get_identity_cos_and_sin_mla - cos, sin = cos_sin_getter(prefill_input_positions) + if self.use_mla_rope: + cos, sin = get_cos_and_sin_mla(prefill_input_positions) + else: + cos = sin = None prefill_query_lens = self.query_lens[reqs_start:].to(torch.int32) actual_seq_lengths_q = torch.cumsum(prefill_query_lens, dim=0).tolist() return AscendMLAPrefillMetadata( @@ -689,8 +687,12 @@ def build_decode_metadata( num_reqs_pad_size, num_reqs, actual_seq_lengths_q, common_attn_metadata ) - cos_sin_getter = get_cos_and_sin_mla if self.use_mla_rope else get_identity_cos_and_sin_mla - cos, sin = cos_sin_getter(input_positions, use_cache=True) + if self.use_mla_rope: + cos, sin = get_cos_and_sin_mla(input_positions, use_cache=True) + cos = cos[: self.num_decode_tokens, ...] + sin = sin[: self.num_decode_tokens, ...] + else: + cos = sin = None decode_metadata = self.decode_metadata_cls( input_positions=input_positions, block_table=self.block_table, @@ -699,8 +701,8 @@ def build_decode_metadata( max_seq_lens=max_seq_lens, attn_mask=self.attn_mask_builder.get_splitfuse_attn_mask(), actual_seq_lengths_q=actual_seq_lengths_q, - sin=sin[: self.num_decode_tokens, ...], - cos=cos[: self.num_decode_tokens, ...], + sin=sin, + cos=cos, ) return decode_metadata diff --git a/vllm_ascend/ops/rotary_embedding.py b/vllm_ascend/ops/rotary_embedding.py index 11243724ad20..92776d186b83 100644 --- a/vllm_ascend/ops/rotary_embedding.py +++ b/vllm_ascend/ops/rotary_embedding.py @@ -103,15 +103,6 @@ def get_cos_and_sin_mla(positions, use_cache=False): return _cos_mla[:num_tokens, ...], _sin_mla[:num_tokens, ...] -def get_identity_cos_and_sin_mla(positions, use_cache=False): - """Return the existing stable-shape identity rotation for no-RoPE MLA.""" - del use_cache - global _cos_mla - global _sin_mla - num_tokens = positions.size(0) - return _cos_mla[:num_tokens, ...], _sin_mla[:num_tokens, ...] - - def _record_cos_sin_cache(cos_sin_cache): global _cos_sin_cache if _cos_sin_cache is not None: From 11bc0aa3ef14f0bbc5b93265ede4bb258993b749 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Thu, 27 Aug 2026 10:20:22 +0800 Subject: [PATCH 38/50] fix(kv-cache): preserve grouped Mamba block capacity Delegate UniformTypeKVCacheSpecs block-table sizing to its inner specs when the installed vLLM lacks the upstream implementation. This preserves Mamba speculative blocks for grouped Kimi K3 caches and keeps scheduler and worker table widths aligned near the context limit. Signed-off-by: maoxx241 --- .../platform/test_prefix_cache_cp_patches.py | 18 +++++++++++------- .../patch/platform/patch_kv_cache_utils.py | 6 ++---- 2 files changed, 13 insertions(+), 11 deletions(-) diff --git a/tests/ut/patch/platform/test_prefix_cache_cp_patches.py b/tests/ut/patch/platform/test_prefix_cache_cp_patches.py index fcec9e286752..aa1f3f6d4c19 100644 --- a/tests/ut/patch/platform/test_prefix_cache_cp_patches.py +++ b/tests/ut/patch/platform/test_prefix_cache_cp_patches.py @@ -354,6 +354,8 @@ def test_kimi_k3_mla_dspark_uses_four_groups_and_five_builders() -> None: def test_kimi_k3_gqa_mixed_groups_preserve_scheduler_and_mamba_contracts() -> None: groups = _get_kimi_k3_dspark_mixed_kv_cache_groups(_make_kimi_k3_dspark_kv_cache_specs()) assert groups is not None + max_model_len = 133120 + vllm_config = _make_vllm_config(enable_prefix_caching=True, dcp=1, block_size=384) worker_config = KVCacheConfig( num_blocks=100, kv_cache_tensors=[], @@ -362,17 +364,19 @@ def test_kimi_k3_gqa_mixed_groups_preserve_scheduler_and_mamba_contracts() -> No assert worker_config.has_mamba_layers assert worker_config.needs_kv_cache_zeroing - assert ( - groups[1].kv_cache_spec.max_num_blocks_per_req( - _make_vllm_config(enable_prefix_caching=True, dcp=1, block_size=384), - 3840, - ) - == 17 - ) + worker_mamba_widths = { + group.kv_cache_spec.max_num_blocks_per_req(vllm_config, max_model_len) for group in groups[1:] + } + assert worker_mamba_widths == {354} scheduler_config = generate_scheduler_kv_cache_config([worker_config]) assert isinstance(scheduler_config.kv_cache_groups[0].kv_cache_spec, MLAAttentionSpec) assert all(isinstance(group.kv_cache_spec, MambaSpec) for group in scheduler_config.kv_cache_groups[1:]) + scheduler_mamba_widths = { + group.kv_cache_spec.max_num_blocks_per_req(vllm_config, max_model_len) + for group in scheduler_config.kv_cache_groups[1:] + } + assert scheduler_mamba_widths == worker_mamba_widths assert scheduler_config.needs_kv_cache_zeroing diff --git a/vllm_ascend/patch/platform/patch_kv_cache_utils.py b/vllm_ascend/patch/platform/patch_kv_cache_utils.py index 6aa468dbc268..4f6f4134c620 100644 --- a/vllm_ascend/patch/platform/patch_kv_cache_utils.py +++ b/vllm_ascend/patch/platform/patch_kv_cache_utils.py @@ -22,22 +22,20 @@ get_kv_cache_spec_kind, ) -from vllm_ascend.utils import vllm_version_is - _KIMI_K3_TARGET_LAYER_PREFIX = "language_model.model.layers." _KIMI_K3_DRAFT_LAYER_PREFIX = "model.layers." _orig_resolve_kv_cache_block_sizes = vllm.v1.core.kv_cache_utils.resolve_kv_cache_block_sizes _orig_get_kv_cache_groups_uniform_page_size = vllm.v1.core.kv_cache_utils._get_kv_cache_groups_uniform_page_size -if vllm_version_is("0.27.1"): +if UniformTypeKVCacheSpecs.max_num_blocks_per_req is KVCacheSpec.max_num_blocks_per_req: def _uniform_type_max_num_blocks_per_req( self: UniformTypeKVCacheSpecs, vllm_config: VllmConfig, max_len: int, ) -> int: - """Use the per-layer block-table width, matching current vLLM.""" + """Preserve the inner spec's block-table width.""" widths = {spec.max_num_blocks_per_req(vllm_config, max_len) for spec in self.kv_cache_specs.values()} assert len(widths) == 1, ( "All layers in the same KV cache group must need the same number " From 548c0dc634742b6777fbee3925f2affe3a0c1e23 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Thu, 27 Aug 2026 16:40:07 +0800 Subject: [PATCH 39/50] test(kimi-k3): remove dummy execution parity nightly Remove the Kimi K3 dummy nightly case, its dedicated fixture and workflow entry. Clean up the matching documentation claims while retaining real-checkpoint validation guidance. Signed-off-by: maoxx241 --- .github/workflows/configs/nightly_config.yaml | 4 - docs/source/tutorials/models/Kimi-K3.md | 30 ++--- .../kimi_k3_5layers_16experts/config.json | 63 ---------- .../models/test_kimi_k3_execution_parity.py | 110 ------------------ 4 files changed, 10 insertions(+), 197 deletions(-) delete mode 100644 tests/e2e/nightly/single_node/models/fixtures/kimi_k3_5layers_16experts/config.json delete mode 100644 tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py diff --git a/.github/workflows/configs/nightly_config.yaml b/.github/workflows/configs/nightly_config.yaml index 8f5adeaff85e..716adc58dab9 100644 --- a/.github/workflows/configs/nightly_config.yaml +++ b/.github/workflows/configs/nightly_config.yaml @@ -231,10 +231,6 @@ a3: multi_card: test_config: # pytest-driven tests - - name: kimi-k3-execution-parity - os: linux-aarch64-nightly-a3-16 - tests: tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py - testcase_timeout: 180 - name: qwen3-30b-acc os: linux-aarch64-nightly-a3-4 tests: tests/e2e/weekly/single_node/models/test_qwen3_30b_acc.py diff --git a/docs/source/tutorials/models/Kimi-K3.md b/docs/source/tutorials/models/Kimi-K3.md index 56aca600db56..4b882a9c9711 100644 --- a/docs/source/tutorials/models/Kimi-K3.md +++ b/docs/source/tutorials/models/Kimi-K3.md @@ -33,7 +33,7 @@ for the project-wide feature matrix. ## 3 Choosing a Checkpoint -Kimi K3 checkpoints serve different validation purposes. A reduced or dummy +Kimi K3 checkpoints serve different validation purposes. A reduced checkpoint must not be used to report GPQA or other semantic accuracy. | Checkpoint | Typical storage | What it can validate | What it cannot validate | @@ -41,14 +41,8 @@ checkpoint must not be used to report GPQA or other semantic accuracy. | Full 93-layer, 896-expert W4A8 | About 1.49 TB | Deployment and benchmark accuracy | Not applicable | | Full 93-layer, 16-expert derivative | About 113 GB | Single-node integration and long-context execution | Full-model semantics and full expert routing | | Five-layer, 16-expert W4A8 derivative | About 12.1 GiB | Real quantized loading, KDA/MLA/MoE execution, graph and cache parity | Full-depth behavior and benchmark accuracy | -| Five-layer, 16-expert dummy | Configuration only | Model construction, TP16/EP, graph replay, cache-state parity, finite outputs | Real weight loading, W4A8 numerics, or semantic accuracy | -The nightly test uses the last option so it does not depend on a large model -cache. Its configuration preserves the production hidden size, head geometry, -mixed KDA/MLA layout, attention residuals, and top-16 expert routing. It reduces -only the layer and expert counts, then initializes BF16 dummy weights. - -For a storage-limited real-weight nightly artifact, derive the five-layer, +For a storage-limited real-weight validation checkpoint, derive the five-layer, 16-expert checkpoint from the W4A8 checkpoint as follows: 1. Keep tokenizer and configuration metadata. @@ -61,7 +55,7 @@ For a storage-limited real-weight nightly artifact, derive the five-layer, and log-probability parity. Do not create a task-accuracy baseline from it. :::{note} -The reduced checkpoint is a CI fixture, not a model release. Keep the source +The reduced checkpoint is a validation artifact, not a model release. Keep the source checkpoint revision and a manifest of retained tensors with the artifact so that it can be reproduced when the quantized weights change. ::: @@ -395,22 +389,18 @@ Run concurrency and benchmark requests from a host inside the trusted serving network. A developer workstation or VPN should be used only for bounded smoke requests. -## 9 Accuracy and Nightly Validation +## 9 Accuracy Validation Use the following validation ladder: -1. The storage-free nightly guard loads the committed five-layer, 16-expert - configuration with dummy BF16 weights on one 16-NPU A3 node. In one model - instance it compares cold prefill, a Prefix Cache hit that leaves exactly - one token to prefill, and a post-reset cold run. It requires identical token - IDs, close and finite chosen-token log probabilities, complete outputs, and - `FULL_DECODE_ONLY` replay. -1. A hosted five-layer, 16-expert W4A8 fixture should run the same parity case - with real weights. This adds quantized loader and W4A8 numerical coverage - while staying small enough for a nightly worker. +1. Use a reduced W4A8 checkpoint for execution validation. Compare cold + prefill, a Prefix Cache hit that leaves exactly one token to prefill, and + a post-reset cold run. Check deterministic output tokens, close and finite + chosen-token log probabilities, complete outputs, and `FULL_DECODE_ONLY` + replay. 1. Run GPQA with the full 93-layer, 896-expert checkpoint on the four-node deployment. This is the semantic accuracy gate and cannot be replaced by - either reduced fixture. + a reduced checkpoint. Keep the model revision, tokenizer, chat rendering, reasoning mode, sampling parameters, dataset revision, and evaluator revision fixed when comparing diff --git a/tests/e2e/nightly/single_node/models/fixtures/kimi_k3_5layers_16experts/config.json b/tests/e2e/nightly/single_node/models/fixtures/kimi_k3_5layers_16experts/config.json deleted file mode 100644 index 0dd1a5e42c0e..000000000000 --- a/tests/e2e/nightly/single_node/models/fixtures/kimi_k3_5layers_16experts/config.json +++ /dev/null @@ -1,63 +0,0 @@ -{ - "activation_situ_beta": 4, - "activation_situ_linear_beta": 25, - "architectures": [ - "KimiLinearForCausalLM" - ], - "attn_res_block_size": 12, - "bos_token_id": 163584, - "dtype": "bfloat16", - "eos_token_id": 163586, - "first_k_dense_replace": 1, - "hidden_act": "situ", - "hidden_size": 7168, - "intermediate_size": 33792, - "kv_lora_rank": 512, - "latent_moe_use_norm": true, - "linear_attn_config": { - "full_attn_layers": [ - 4 - ], - "gate_lower_bound": -5, - "head_dim": 128, - "kda_layers": [ - 1, - 2, - 3, - 5 - ], - "num_heads": 96, - "short_conv_kernel_size": 4, - "use_full_rank_gate": true - }, - "max_position_embeddings": 1048576, - "mla_use_nope": true, - "mla_use_output_gate": true, - "model_type": "kimi_linear", - "moe_intermediate_size": 3072, - "moe_layer_freq": 1, - "moe_renormalize": true, - "moe_router_activation_func": "sigmoid", - "num_attention_heads": 96, - "num_expert_group": 1, - "num_experts": 16, - "num_experts_per_token": 16, - "num_hidden_layers": 5, - "num_key_value_heads": 96, - "num_nextn_predict_layers": 0, - "num_shared_experts": 2, - "pad_token_id": 163839, - "q_lora_rank": 1536, - "qk_nope_head_dim": 128, - "qk_rope_head_dim": 64, - "rms_norm_eps": 1e-05, - "routed_expert_hidden_size": 3584, - "routed_scaling_factor": 1, - "tie_word_embeddings": false, - "topk_group": 1, - "topk_method": "noaux_tc", - "use_cache": true, - "use_grouped_topk": true, - "v_head_dim": 128, - "vocab_size": 163840 -} diff --git a/tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py b/tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py deleted file mode 100644 index a28c793923b5..000000000000 --- a/tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py +++ /dev/null @@ -1,110 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project -# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved. -"""Storage-light Kimi K3 execution parity guard. - -The committed fixture keeps Kimi K3's production dimensions and its mixed -KDA/MLA layout, but limits the model to five layers and sixteen experts. Dummy -weights deliberately make this an execution-parity test, not a semantic -accuracy test. Full-checkpoint GPQA remains a separate release gate. -""" - -from pathlib import Path - -import pytest -import torch -from vllm import SamplingParams -from vllm.inputs import TokensPrompt - -from tests.e2e.conftest import VllmRunner, wait_until_npu_memory_free - -MODEL_CONFIG = Path(__file__).parent / "fixtures" / "kimi_k3_5layers_16experts" -SCHEDULER_BLOCK_SIZE = 16 -PROMPT_TOKEN_IDS = [163584, *range(100, 100 + SCHEDULER_BLOCK_SIZE)] -MAX_TOKENS = 4 - - -def _assert_complete_output(request_output): - assert request_output is not None - assert request_output.finished - assert request_output.outputs is not None - assert len(request_output.outputs) == 1 - - completion = request_output.outputs[0] - assert completion is not None - assert completion.token_ids is not None - assert len(completion.token_ids) == MAX_TOKENS - assert completion.logprobs is not None - assert len(completion.logprobs) == MAX_TOKENS - - chosen_logprobs = [] - for token_id, step_logprobs in zip(completion.token_ids, completion.logprobs): - assert step_logprobs is not None - assert token_id in step_logprobs - logprob = step_logprobs[token_id].logprob - assert logprob is not None - assert torch.isfinite(torch.tensor(logprob)) - chosen_logprobs.append(logprob) - - return list(completion.token_ids), torch.tensor(chosen_logprobs, dtype=torch.float32) - - -@pytest.mark.e2e_model("sgl-npu/Kimi-K3-W4A8") -@pytest.mark.e2e_coverage( - arch="moe", - feature="aclgraph,prefix_caching,logprobs", - parallel="TP,EP", - deploy="pd_mix", - hardware="A3", - quantization="BF16", - graph_mode="full_decode_only", -) -@wait_until_npu_memory_free() -def test_kimi_k3_dummy_prefix_cache_one_token_prefill_parity(): - """Compare cold prefill with the cached block-size-plus-one path.""" - sampling_params = SamplingParams( - temperature=0, - max_tokens=MAX_TOKENS, - logprobs=1, - ignore_eos=True, - seed=0, - ) - prompt = TokensPrompt(prompt_token_ids=PROMPT_TOKEN_IDS) - - with VllmRunner( - str(MODEL_CONFIG), - skip_tokenizer_init=True, - load_format="dummy", - dtype="bfloat16", - seed=0, - block_size=SCHEDULER_BLOCK_SIZE, - max_model_len=64, - max_num_seqs=1, - max_num_batched_tokens=64, - tensor_parallel_size=16, - enable_expert_parallel=True, - enable_prefix_caching=True, - gpu_memory_utilization=0.75, - compilation_config={ - "cudagraph_mode": "FULL_DECODE_ONLY", - "cudagraph_capture_sizes": [1], - }, - ) as vllm_model: - cold = vllm_model.model.generate([prompt], sampling_params, use_tqdm=False)[0] - hit = vllm_model.model.generate([prompt], sampling_params, use_tqdm=False)[0] - - assert cold.num_cached_tokens in (None, 0) - assert hit.num_cached_tokens == SCHEDULER_BLOCK_SIZE - assert len(PROMPT_TOKEN_IDS) - hit.num_cached_tokens == 1 - - cold_tokens, cold_logprobs = _assert_complete_output(cold) - hit_tokens, hit_logprobs = _assert_complete_output(hit) - assert hit_tokens == cold_tokens - torch.testing.assert_close(hit_logprobs, cold_logprobs, rtol=5e-3, atol=5e-3) - - assert vllm_model.model.reset_prefix_cache() - reset = vllm_model.model.generate([prompt], sampling_params, use_tqdm=False)[0] - assert reset.num_cached_tokens in (None, 0) - reset_tokens, reset_logprobs = _assert_complete_output(reset) - assert reset_tokens == cold_tokens - torch.testing.assert_close(reset_logprobs, cold_logprobs, rtol=5e-3, atol=5e-3) From 4c6924ca33ebe04ccae266658b6bdfae2093788b Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Thu, 27 Aug 2026 16:55:59 +0800 Subject: [PATCH 40/50] fix(moe): compute SiTU without a runtime CustomOp Keep SiTU in the existing MoE activation dispatch and evaluate the upstream native formula directly. This avoids constructing a CustomOp after the model-init vLLM config context has exited in the unquantized, antiquant and W4A16 paths. Exercise the existing SiTU regression tests without a forward config context and compare all supported floating dtypes with the upstream native activation. The quantized fused SiTU kernels are unchanged. Signed-off-by: maoxx241 --- tests/ut/ops/test_moe_mlp.py | 62 ++++++++++++++-------------- vllm_ascend/ops/fused_moe/moe_mlp.py | 32 ++++++++++---- 2 files changed, 55 insertions(+), 39 deletions(-) diff --git a/tests/ut/ops/test_moe_mlp.py b/tests/ut/ops/test_moe_mlp.py index 0583da394524..841819d3c10b 100644 --- a/tests/ut/ops/test_moe_mlp.py +++ b/tests/ut/ops/test_moe_mlp.py @@ -7,7 +7,7 @@ import torch import torch_npu # noqa: F401 -- registers torch.npu used by the module under test from torch.nn import functional as F -from vllm.config import CompilationConfig, VllmConfig, set_current_vllm_config +from vllm.config import VllmConfig, set_current_vllm_config from vllm.model_executor.layers.activation import SituAndMul from vllm.model_executor.layers.fused_moe.activation import MoEActivation @@ -30,10 +30,6 @@ MXFP4_TEST_DTYPE = getattr(torch, "float4_e2m1fn_x2", torch.float16) -def _custom_op_vllm_config(): - return SimpleNamespace(compilation_config=CompilationConfig(custom_ops=["all"])) - - class TestCumsumGroupList(unittest.TestCase): glist_dict: ClassVar[dict[int, torch.Tensor]] @@ -123,31 +119,35 @@ def test_uses_small_op_activation_and_dynamic_mx_quant(self): class TestUnifiedApplyMlpRequest(unittest.TestCase): - def test_unquant_situ_uses_upstream_activation_contract(self): - hidden_states = torch.randn(2, 8, dtype=torch.bfloat16) - gate_up_out = torch.randn(2, 16, dtype=torch.bfloat16) - expected_output = torch.randn(2, 8, dtype=torch.bfloat16) - with set_current_vllm_config(VllmConfig()): - expected_activation = SituAndMul(beta=4.0, linear_beta=25.0)(gate_up_out) - - with patch( - f"{MOE_MLP}.torch_npu.npu_grouped_matmul", - side_effect=[[gate_up_out], [expected_output]], - create=True, - ) as grouped_matmul: - output, _ = unquant_apply_mlp( - hidden_states=hidden_states, - w1=torch.randn(1, 8, 16), - w2=torch.randn(1, 8, 8), - group_list=torch.tensor([1, 1]), - activation=MoEActivation.SITU, - activation_situ_beta=4.0, - activation_situ_linear_beta=25.0, - need_trans=False, - ) - - self.assertIs(output, expected_output) - torch.testing.assert_close(grouped_matmul.call_args_list[1].kwargs["x"][0], expected_activation) + def test_unquant_situ_matches_upstream_without_config_context(self): + for dtype in (torch.bfloat16, torch.float16, torch.float32): + for linear_beta in (None, 25.0): + with self.subTest(dtype=dtype, linear_beta=linear_beta): + hidden_states = torch.randn(2, 8, dtype=dtype) + gate_up_out = torch.linspace(-40, 40, 32).reshape(2, 16).to(dtype) + expected_output = torch.randn(2, 8, dtype=dtype) + with set_current_vllm_config(VllmConfig()): + expected_activation = SituAndMul(beta=4.0, linear_beta=linear_beta).forward_native(gate_up_out) + + # The worker forward runs outside the model-init config context. + with patch( + f"{MOE_MLP}.torch_npu.npu_grouped_matmul", + side_effect=[[gate_up_out], [expected_output]], + create=True, + ) as grouped_matmul: + output, _ = unquant_apply_mlp( + hidden_states=hidden_states, + w1=torch.randn(1, 8, 16), + w2=torch.randn(1, 8, 8), + group_list=torch.tensor([1, 1]), + activation=MoEActivation.SITU, + activation_situ_beta=4.0, + activation_situ_linear_beta=linear_beta, + need_trans=False, + ) + + self.assertIs(output, expected_output) + torch.testing.assert_close(grouped_matmul.call_args_list[1].kwargs["x"][0], expected_activation) def test_unquant_swigluoai_uninterleave_falls_back_on_a5(self): hidden_states = torch.randn(2, 8, dtype=torch.bfloat16) @@ -684,7 +684,6 @@ def test_antiquant_weights_use_native_situ_between_grouped_matmuls(self): stream_patch, evt = _patch_npu_stream() with ( - set_current_vllm_config(_custom_op_vllm_config()), stream_patch, patch(f"{MOE_MLP}._EXTRA_CTX", MagicMock(moe_comm_type=-1)), patch( @@ -723,7 +722,6 @@ def test_w4a16_mxfp_uses_native_situ_without_activation_requant(self): stream_patch, evt = _patch_npu_stream() with ( - set_current_vllm_config(_custom_op_vllm_config()), stream_patch, patch(f"{MOE_MLP}._EXTRA_CTX", MagicMock(moe_comm_type=-1)), patch("torch_npu.npu_grouped_matmul", return_value=[gate_up_out], create=True), diff --git a/vllm_ascend/ops/fused_moe/moe_mlp.py b/vllm_ascend/ops/fused_moe/moe_mlp.py index 51160ed12dd6..61e4ad3de5f9 100644 --- a/vllm_ascend/ops/fused_moe/moe_mlp.py +++ b/vllm_ascend/ops/fused_moe/moe_mlp.py @@ -18,7 +18,6 @@ import torch import torch_npu from torch.nn.functional import pad -from vllm.model_executor.layers.activation import SituAndMul from vllm.model_executor.layers.fused_moe.activation import MoEActivation from vllm.triton_utils import HAS_TRITON @@ -121,6 +120,22 @@ def _prepare_swigluoai_grouped_matmul_scales( return [scale.to(output_dtype) if scale.dtype != output_dtype else scale for scale in scales] +def _apply_situ( + hidden_states: torch.Tensor, + *, + beta: float, + linear_beta: float | None, +) -> torch.Tensor: + """Match SituAndMul.forward_native without constructing a CustomOp in forward.""" + gate, up = hidden_states.chunk(2, dim=-1) + gate = gate.float() + up = up.float() + gate = beta * torch.tanh(gate / beta) * torch.sigmoid(gate) + if linear_beta is not None: + up = linear_beta * torch.tanh(up / linear_beta) + return (gate * up).to(hidden_states.dtype) + + def _apply_clipped_swiglu( hidden_states: torch.Tensor, *, @@ -406,10 +421,11 @@ def quant_apply_mlp( approximate = "tanh" if activation == MoEActivation.GELU_TANH else "none" hidden_states = torch.nn.functional.gelu(gate, approximate=approximate) * up elif activation == MoEActivation.SITU: - hidden_states = SituAndMul( + hidden_states = _apply_situ( + hidden_states, beta=situ_beta, linear_beta=activation_situ_linear_beta, - )(hidden_states) + ) elif is_swigluoai_uninterleave: hidden_states = _apply_clipped_swiglu( hidden_states, @@ -573,10 +589,11 @@ def quant_apply_mlp( quant_mode="dynamic", ) else: - hidden_states = SituAndMul( + hidden_states = _apply_situ( + hidden_states, beta=situ_beta, linear_beta=activation_situ_linear_beta, - )(hidden_states) + ) swiglu_out_scale = None elif is_swigluoai_uninterleave: if use_mxfp_quant: @@ -703,10 +720,11 @@ def unquant_apply_mlp( act_name = getattr(activation, "value", activation) if activation == MoEActivation.SITU: - gate_up_out = SituAndMul( + gate_up_out = _apply_situ( + gate_up_out, beta=activation_situ_beta, linear_beta=activation_situ_linear_beta, - )(gate_up_out) + ) elif activation == MoEActivation.SWIGLUOAI: num_experts, _, hidden_size = w1.shape gate_up_out = AscendSwigluOAIAndMul.swiglu_oai_forward(gate_up_out.view(-1, hidden_size)) From 6da8cfa2d9f717aa36179e38fb2231b35cbc09c5 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Thu, 27 Aug 2026 17:33:47 +0800 Subject: [PATCH 41/50] fix(pd): preserve decode graphs for stateful handoffs Allow uniform handoff batches to use decode graphs once initial state is available, including prompts at N-1 with speculative padding. Keep first-token prefills and nonuniform batches on their existing path without changing DP mode synchronization. Replace the completed-prompt assertion test with CPU behavior checks using the real dispatcher and DP synchronization for DP1 and DP4. Signed-off-by: maoxx241 --- .../a2/test_model_runner_v1_with_device.py | 101 +++++++++++++----- vllm_ascend/worker/model_runner_v1.py | 11 +- 2 files changed, 80 insertions(+), 32 deletions(-) diff --git a/tests/ut/worker/a2/test_model_runner_v1_with_device.py b/tests/ut/worker/a2/test_model_runner_v1_with_device.py index bd7875ba728e..e76986ed690f 100644 --- a/tests/ut/worker/a2/test_model_runner_v1_with_device.py +++ b/tests/ut/worker/a2/test_model_runner_v1_with_device.py @@ -1,8 +1,10 @@ import os +from types import SimpleNamespace from unittest.mock import MagicMock, patch import numpy as np import pytest +import torch from vllm.config import ( CacheConfig, CUDAGraphMode, @@ -15,6 +17,7 @@ from vllm.distributed.parallel_state import GroupCoordinator from vllm.model_executor.layers.attention import Attention from vllm.platforms import current_platform +from vllm.v1.cudagraph_dispatcher import CudagraphDispatcher from vllm.v1.kv_cache_interface import ( FullAttentionSpec, KVCacheConfig, @@ -392,34 +395,82 @@ def test_determine_batch_execution_and_padding( @pytest.mark.parametrize( - ("num_prompt_tokens", "expected_uniform_decode"), + ("num_spec_tokens", "computed", "prompts", "scheduled", "expected_mode"), [ - pytest.param(8, False, id="one_token_pd_prefill"), - pytest.param(7, True, id="decode_after_prompt"), + pytest.param(0, [7], [8], [1], CUDAGraphMode.FULL, id="stateful_one_token_handoff"), + pytest.param(0, [0], [1], [1], CUDAGraphMode.NONE, id="first_token_without_state"), + pytest.param(7, [16, 24], [8, 8], [8, 8], CUDAGraphMode.FULL, id="steady_spec_decode"), + pytest.param(7, [16, 7], [8, 8], [8, 8], CUDAGraphMode.FULL, id="handoff_padded_to_spec_width"), + pytest.param(7, [16, 0], [8, 8], [8, 8], CUDAGraphMode.NONE, id="spec_width_prefill_without_state"), + pytest.param(7, [16, 7], [8, 8], [8, 1], CUDAGraphMode.NONE, id="nonuniform_handoff"), ], ) -def test_uniform_decode_requires_completed_prompt_without_spec_decode( - model_runner, - num_prompt_tokens: int, - expected_uniform_decode: bool, +@pytest.mark.parametrize("dp_size", [1, 4]) +def test_stateful_handoff_preserves_decode_graph( + monkeypatch, + num_spec_tokens, + computed, + prompts, + scheduled, + expected_mode, + dp_size, ): - runner = model_runner - runner.speculative_config = None - runner.uniform_decode_query_len = 1 - runner.input_batch.num_computed_tokens_cpu[0] = 7 - runner.input_batch.num_prompt_tokens[0] = num_prompt_tokens + # Exercise the real dispatcher and DP synchronization using CPU metadata only. + runner = NPUModelRunner.__new__(NPUModelRunner) + runner.dp_size = dp_size + runner.dp_rank = 0 + runner.parallel_config = SimpleNamespace( + data_parallel_size=dp_size, + data_parallel_rank=0, + tensor_parallel_size=8, + use_sequence_parallel_moe=True, + ) + runner.compilation_config = SimpleNamespace( + pass_config=SimpleNamespace(enable_sp=True), + cudagraph_mode=CUDAGraphMode.FULL_DECODE_ONLY, + cudagraph_capture_sizes=[8, 16, 24, 32], + max_cudagraph_capture_size=32, + compile_sizes=[], + ) + runner.vllm_config = SimpleNamespace( + parallel_config=runner.parallel_config, + compilation_config=runner.compilation_config, + scheduler_config=SimpleNamespace(max_num_seqs=32), + observability_config=SimpleNamespace(cudagraph_metrics=False), + num_speculative_tokens=num_spec_tokens, + lora_config=None, + ) + runner.model_config = SimpleNamespace(is_encoder_decoder=False) + runner.uniform_decode_query_len = 1 + num_spec_tokens + runner.input_batch = SimpleNamespace( + num_computed_tokens_cpu=np.array(computed), + num_prompt_tokens=np.array(prompts), + lora_id_to_lora_request={}, + ) + runner.cudagraph_dispatcher = CudagraphDispatcher(runner.vllm_config) + runner.cudagraph_dispatcher.initialize_cudagraph_keys( + CUDAGraphMode.FULL_DECODE_ONLY, runner.uniform_decode_query_len + ) - with patch.object( - runner.cudagraph_dispatcher, - "dispatch", - wraps=runner.cudagraph_dispatcher.dispatch, - ) as dispatch: - runner._determine_batch_execution_and_padding( - num_tokens=1, - num_reqs=1, - num_scheduled_tokens_np=np.array([1], dtype=np.int32), - max_num_scheduled_tokens=1, - use_cascade_attn=False, - ) + def all_reduce(packed_tensor, group): + # Other DP replicas are already decoding at the largest captured size. + packed_tensor[0, 1:] = 32 + packed_tensor[1, 1:] = CUDAGraphMode.FULL.value + + module = "vllm_ascend.worker.model_runner_v1" + monkeypatch.setattr(f"{module}.should_skip_allreduce_across_dp_group", lambda *args: False) + monkeypatch.setattr(f"{module}.get_dp_group", lambda: SimpleNamespace(cpu_group=None)) + monkeypatch.setattr(f"{module}.dist.all_reduce", all_reduce) + mode, descriptor, _, tokens_across_dp, _ = runner._determine_batch_execution_and_padding( + num_tokens=sum(scheduled), + num_reqs=len(scheduled), + num_scheduled_tokens_np=np.array(scheduled, dtype=np.int32), + max_num_scheduled_tokens=max(scheduled), + use_cascade_attn=False, + ) - assert bool(dispatch.call_args.kwargs["uniform_decode"]) is expected_uniform_decode + assert mode == expected_mode + assert descriptor.uniform == (expected_mode == CUDAGraphMode.FULL) + if dp_size > 1: + assert descriptor.num_tokens == 32 + torch.testing.assert_close(tokens_across_dp, torch.full((dp_size,), 32, dtype=torch.int32)) diff --git a/vllm_ascend/worker/model_runner_v1.py b/vllm_ascend/worker/model_runner_v1.py index 1faadb90a8de..b83fd6ddf0de 100644 --- a/vllm_ascend/worker/model_runner_v1.py +++ b/vllm_ascend/worker/model_runner_v1.py @@ -2746,15 +2746,12 @@ def _determine_batch_execution_and_padding( num_encoder_reqs: int = 0, ) -> tuple[CUDAGraphMode, BatchDescriptor, bool, torch.Tensor | None, CUDAGraphStat | None]: num_tokens_padded = self._pad_for_sequence_parallelism(num_tokens) - # A one-token chunk can still be prefill at a P/D handoff. Decode graph - # replay is valid only after every prompt has been fully computed. - is_all_decode = np.all( - self.input_batch.num_computed_tokens_cpu[:num_reqs] - >= self.input_batch.num_prompt_tokens[:num_reqs] - ) + # A stateful P/D handoff can use a uniform decode graph even at + # prompt_len - 1 computed tokens. Keep first-token prefills out. + has_initial_state = np.all(self.input_batch.num_computed_tokens_cpu[:num_reqs] > 0) uniform_decode = ( ( - is_all_decode + has_initial_state and (max_num_scheduled_tokens == self.uniform_decode_query_len) and (num_tokens == max_num_scheduled_tokens * num_reqs) ) From 5e327cb7f6afe8dc81aefdfc706ec7400d7b3af5 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Thu, 27 Aug 2026 18:08:19 +0800 Subject: [PATCH 42/50] fix(mamba): preserve accepted snapshots and per-rank table capacity Keep asynchronous accepted-token D2H results outside the mutable InputBatch until they are remapped into current request order, following the ownership fix in #13864. Reuse the existing buffers and event, and preserve synchronous scheduling behavior. Trust the per-rank capacity returned by each KV spec instead of multiplying replicated Mamba tables by DCP again. Cover real request replacement and backend reordering, scheduling modes, and final speculative slots for raw and grouped Mamba specs at DCP1/DCP4. Validation: 89 targeted tests passed, 2 upstream-version cases skipped; bash format.sh ci passed. No full-model NPU serving run was performed for this change. Signed-off-by: maoxx241 --- tests/ut/worker/a2/test_block_table.py | 66 ++++++---- tests/ut/worker/a2/test_model_runner_v1.py | 145 +++++++++++++++++++++ vllm_ascend/worker/block_table.py | 4 +- vllm_ascend/worker/model_runner_v1.py | 81 ++++++++---- 4 files changed, 246 insertions(+), 50 deletions(-) diff --git a/tests/ut/worker/a2/test_block_table.py b/tests/ut/worker/a2/test_block_table.py index a746c25b249e..d1e940cc4ce1 100644 --- a/tests/ut/worker/a2/test_block_table.py +++ b/tests/ut/worker/a2/test_block_table.py @@ -104,9 +104,11 @@ def test_compute_slot_mapping_draft_reserves_mtp_slots(self): self.assertEqual(block_table.slot_mapping.cpu.numel(), 128) self.assertEqual(block_table.slot_mapping.cpu[: req_indices.size].numel(), 110) - def test_uniform_mamba_group_is_recognized_as_mamba(self): + def test_mamba_table_preserves_speculative_capacity_with_dcp(self): + from vllm_ascend.worker.block_table import BlockTable + mamba_spec = MambaSpec( - block_size=self.block_size, + block_size=384, shapes=((4, 8),), dtypes=(torch.float32,), page_size_padded=128, @@ -116,30 +118,42 @@ def test_uniform_mamba_group_is_recognized_as_mamba(self): layer_specs = {f"mamba.{i}": mamba_spec for i in range(3)} uniform_spec = UniformTypeKVCacheSpecs.from_specs(layer_specs) self.assertIsNotNone(uniform_spec) - kv_cache_group = KVCacheGroupSpec( - layer_names=list(layer_specs), - kv_cache_spec=uniform_spec, - ) - - with patch("vllm_ascend.worker.block_table.get_dcp_group") as mock_get_dcp_group: - mock_get_dcp_group.return_value = SimpleNamespace( - world_size=1, - rank_in_group=0, - ) - from vllm_ascend.worker.block_table import BlockTable - - block_table = BlockTable( - block_size=self.block_size, - max_num_reqs=self.max_num_reqs, - max_num_blocks_per_req=self.max_num_blocks_per_req, - max_num_batched_tokens=self.max_num_batched_tokens, - pin_memory=self.pin_memory, - device=self.device, - kernel_sizes=[0], - kv_cache_group=kv_cache_group, - ) - - self.assertTrue(block_table.is_mamba_group) + # ceil(133120 / 384) + 7 speculative state slots, on every DCP rank. + expected_width = 354 + for spec in (mamba_spec, uniform_spec): + for dcp_size in (1, 4): + with self.subTest(spec=type(spec).__name__, dcp_size=dcp_size): + config = SimpleNamespace( + cache_config=SimpleNamespace(mamba_cache_mode="align"), + parallel_config=SimpleNamespace(decode_context_parallel_size=dcp_size), + ) + width = spec.max_num_blocks_per_req(config, 133120) + self.assertEqual(width, expected_width) + group = KVCacheGroupSpec(layer_names=list(layer_specs), kv_cache_spec=spec) + with patch( + "vllm_ascend.worker.block_table.get_dcp_group", + return_value=SimpleNamespace(world_size=dcp_size, rank_in_group=0), + ): + table = BlockTable( + block_size=spec.block_size, + max_num_reqs=self.max_num_reqs, + max_num_blocks_per_req=width, + max_num_batched_tokens=self.max_num_batched_tokens, + pin_memory=self.pin_memory, + device=self.device, + kernel_sizes=[0], + num_speculative_tokens=7, + kv_cache_group=group, + ) + + self.assertTrue(table.is_mamba_group) + self.assertEqual(table.block_table.cpu.shape[1], expected_width) + self.assertEqual(table.max_num_blocks_per_req, expected_width) + # Exercise the last speculative slot, not just the formula. + block_ids = list(range(expected_width)) + table.add_row(block_ids[:-2], 0) + table.append_row(block_ids[-2:], 0) + np.testing.assert_array_equal(table.block_table.np[0], block_ids) def setup_block_table_data(self, block_table, num_reqs=2): """Helper method to populate block table with test data""" diff --git a/tests/ut/worker/a2/test_model_runner_v1.py b/tests/ut/worker/a2/test_model_runner_v1.py index 8c48316668cd..403daee403fd 100644 --- a/tests/ut/worker/a2/test_model_runner_v1.py +++ b/tests/ut/worker/a2/test_model_runner_v1.py @@ -7,6 +7,8 @@ from vllm.config import CUDAGraphMode from vllm.model_executor.layers.attention import MLAAttention from vllm.model_executor.models.deepseek_v2 import DeepseekV32IndexerCache +from vllm.sampling_params import SamplingParams +from vllm.v1.attention.backends.utils import reorder_batch_to_split_decodes_and_prefills from vllm.v1.kv_cache_interface import ( FullAttentionSpec, KVCacheConfig, @@ -14,6 +16,8 @@ KVCacheTensor, UniformTypeKVCacheSpecs, ) +from vllm.v1.utils import CpuGpuBuffer +from vllm.v1.worker.gpu_input_batch import CachedRequestState, InputBatch from vllm_ascend.attention.mla_v1 import AscendMLABackend from vllm_ascend.attention.utils import get_sfa_qsfa_packed_head_dim @@ -69,6 +73,147 @@ def test_non_dspark_keeps_raw_stream(self): self.assertFalse(runner._draft_uses_qwen3_gqa_dspark()) +class TestAcceptedTokenSnapshot(unittest.TestCase): + def _build_runner(self): + runner = NPUModelRunner.__new__(NPUModelRunner) + runner.use_async_scheduling = True + runner.speculative_config = object() + runner.model_config = SimpleNamespace(is_hybrid=True) + runner.cache_config = SimpleNamespace(mamba_cache_mode="align") + runner.num_accepted_tokens = CpuGpuBuffer(12, dtype=torch.int32, device=torch.device("cpu"), pin_memory=False) + runner.prev_positions = CpuGpuBuffer(12, dtype=torch.int32, device=torch.device("cpu"), pin_memory=False) + batch_counts = torch.ones(12, dtype=torch.int32) + runner.input_batch = SimpleNamespace( + num_accepted_tokens_cpu=batch_counts.numpy(), + num_accepted_tokens_cpu_tensor=batch_counts, + ) + runner.num_accepted_tokens_event = MagicMock() + return runner + + def test_snapshot_survives_request_replacement_and_backend_reorder(self): + runner = self._build_runner() + with patch("vllm.v1.worker.gpu_input_batch.PIN_MEMORY", False): + batch = InputBatch( + max_num_reqs=12, + max_model_len=32, + max_num_batched_tokens=32, + device=torch.device("cpu"), + vocab_size=128, + block_sizes=[4], + kernel_block_sizes=[4], + max_num_blocks_per_req=[8], + num_spec_tokens=3, + ) + runner.input_batch = batch + for req_id in ("A", "B", "C"): + batch.add_request( + CachedRequestState( + req_id=req_id, + prompt_token_ids=[10, 11], + mm_features=[], + sampling_params=SamplingParams(temperature=0), + generator=None, + block_ids=([1],), + num_computed_tokens=2, + output_token_ids=[], + ) + ) + batch.refresh_metadata() + batch.prev_req_id_to_index = dict(batch.req_id_to_index) + runner._get_mamba_bufs = MagicMock() + runner.kv_cache_config = object() + runner.compilation_config = SimpleNamespace(static_forward_context={}) + runner.model = MagicMock() + + # Only replace the device kernel boundary; use real InputBatch mutations. + def postprocess(**kwargs): + kwargs["num_accepted_tokens_cpu_tensor"][:3].copy_(kwargs["num_accepted_tokens_gpu"][:3]) + + with patch( + "vllm_ascend.worker.model_runner_v1.mamba_utils.postprocess_mamba_align_gpu", + side_effect=postprocess, + ): + runner._update_states_after_model_execute( + torch.tensor([[10, 11, -1, -1], [10, 11, 12, -1], [10, 11, 12, 13]]), + SimpleNamespace(), + ) + runner.num_accepted_tokens_event.record.assert_called_once() + np.testing.assert_array_equal(runner.num_accepted_tokens.np[:3], [2, 3, 4]) + + batch.remove_request("A") + batch.add_request( + CachedRequestState( + req_id="D", + prompt_token_ids=[10] * 16, + mm_features=[], + sampling_params=SamplingParams(temperature=0), + generator=None, + block_ids=([2],), + num_computed_tokens=0, + output_token_ids=[], + ) + ) + batch.condense() + self.assertTrue( + reorder_batch_to_split_decodes_and_prefills( + batch, + SimpleNamespace(num_scheduled_tokens={"B": 4, "C": 4, "D": 16}), + decode_threshold=4, + ) + ) + batch.refresh_metadata() + runner._compute_prev_positions(batch.num_reqs) + runner._sync_num_accepted_tokens(batch.num_reqs, has_prev_mapping=True) + + expected = [{"B": 3, "C": 4, "D": 1}[req_id] for req_id in batch.req_ids] + np.testing.assert_array_equal(runner.num_accepted_tokens.np[:3], expected) + np.testing.assert_array_equal(batch.num_accepted_tokens_cpu[:3], expected) + + def test_sync_respects_snapshot_and_current_batch_ownership(self): + for async_scheduling, has_prev_mapping, expected in ( + (True, True, [4, 1, 3]), + (True, False, [1, 1, 1]), + (False, True, [5, 6, 7]), + ): + with self.subTest(async_scheduling=async_scheduling, has_prev_mapping=has_prev_mapping): + runner = self._build_runner() + runner.use_async_scheduling = async_scheduling + runner.num_accepted_tokens.np[:] = np.arange(12) + runner.num_accepted_tokens.np[[11, 4]] = [4, 3] + runner.prev_positions.np[:3] = [11, -1, 4] + runner.input_batch.num_accepted_tokens_cpu[:3] = [5, 6, 7] + + runner._sync_num_accepted_tokens(3, has_prev_mapping=has_prev_mapping) + + np.testing.assert_array_equal(runner.num_accepted_tokens.np[:3], expected) + np.testing.assert_array_equal(runner.input_batch.num_accepted_tokens_cpu[:3], expected) + + def test_non_align_postprocess_keeps_an_independent_snapshot(self): + for mode in ("none", "all"): + with self.subTest(mode=mode): + runner = self._build_runner() + runner.cache_config.mamba_cache_mode = mode + runner.kv_cache_config = object() + runner.requests = {} + runner.mamba_state_idx = {} + runner.num_spec_tokens = 3 + with patch("vllm_ascend.worker.model_runner_v1.mamba_utils.postprocess_mamba_all") as postprocess_all: + runner._update_states_after_model_execute(torch.tensor([[10, -1], [11, 12]]), SimpleNamespace()) + np.testing.assert_array_equal(runner.num_accepted_tokens.np[:2], [1, 2]) + np.testing.assert_array_equal(runner.input_batch.num_accepted_tokens_cpu[:2], [1, 1]) + self.assertEqual(postprocess_all.call_count, int(mode == "all")) + runner.num_accepted_tokens_event.record.assert_called_once() + + def test_non_async_postprocess_keeps_upstream_behavior(self): + runner = self._build_runner() + runner.use_async_scheduling = False + output_token_ids = torch.tensor([[10, -1]]) + scheduler_output = SimpleNamespace() + with patch("vllm.v1.worker.gpu_model_runner.GPUModelRunner._update_states_after_model_execute") as upstream: + runner._update_states_after_model_execute(output_token_ids, scheduler_output) + upstream.assert_called_once_with(output_token_ids, scheduler_output) + + class TestNPUModelRunnerKVCache(unittest.TestCase): def _build_runner(self): runner = NPUModelRunner.__new__(NPUModelRunner) diff --git a/vllm_ascend/worker/block_table.py b/vllm_ascend/worker/block_table.py index b200c1bbeb70..dcf8b4fa0e2d 100644 --- a/vllm_ascend/worker/block_table.py +++ b/vllm_ascend/worker/block_table.py @@ -39,8 +39,8 @@ def __init__( and hasattr(kv_cache_group, "kv_cache_spec") and get_kv_cache_spec_kind(kv_cache_group.kv_cache_spec) == KVCacheSpecKind.MAMBA ) - if self.dcp_world_size > 1 and is_mamba_group: - max_num_blocks_per_req = max_num_blocks_per_req * self.dcp_world_size + # The KV cache spec already provides the per-rank table capacity. + # Mamba state is replicated across DCP ranks, not sharded then expanded. self.max_num_blocks_per_req = max_num_blocks_per_req self.max_num_batched_tokens = max_num_batched_tokens self.pin_memory = pin_memory diff --git a/vllm_ascend/worker/model_runner_v1.py b/vllm_ascend/worker/model_runner_v1.py index b83fd6ddf0de..8294de55df03 100644 --- a/vllm_ascend/worker/model_runner_v1.py +++ b/vllm_ascend/worker/model_runner_v1.py @@ -817,6 +817,63 @@ def _update_states(self, scheduler_output: "SchedulerOutput") -> Callable | None self._apply_pp_sampled_tokens_from_scheduler_output(scheduler_output) return super()._update_states(scheduler_output) + def _update_states_after_model_execute( + self, output_token_ids: torch.Tensor, scheduler_output: "SchedulerOutput" + ) -> None: + if not self.use_async_scheduling: + return super()._update_states_after_model_execute(output_token_ids, scheduler_output) + if not self.speculative_config or not self.model_config.is_hybrid: + return + + # InputBatch can be condensed/reordered before the next input preparation. + # Keep the asynchronous D2H result in the previous iteration's row order, + # independently of InputBatch, until the existing event is synchronized. + num_reqs = output_token_ids.size(0) + self.num_accepted_tokens.gpu[:num_reqs] = (output_token_ids != -1).sum(dim=1) + if self.cache_config.mamba_cache_mode == "align": + mamba_utils.postprocess_mamba_align_gpu( + bufs=self._get_mamba_bufs(), + num_reqs=num_reqs, + num_accepted_tokens_gpu=self.num_accepted_tokens.gpu, + num_accepted_tokens_cpu_tensor=self.num_accepted_tokens.cpu, + input_batch=self.input_batch, + kv_cache_config=self.kv_cache_config, + forward_context=self.compilation_config.static_forward_context, + mamba_state_copy_funcs=self.model.get_mamba_state_copy_func(), + ) + else: + self.num_accepted_tokens.copy_to_cpu(num_reqs) + if self.cache_config.mamba_cache_mode == "all": + mamba_utils.postprocess_mamba_all( + scheduler_output, + self.kv_cache_config, + self.input_batch, + self.requests, + self.mamba_state_idx, + self.num_spec_tokens, + num_reqs, + ) + assert self.num_accepted_tokens_event is not None + self.num_accepted_tokens_event.record() + + def _sync_num_accepted_tokens(self, num_reqs: int, has_prev_mapping: bool) -> None: + """Publish accepted counts in current request order after the D2H event.""" + accepted = self.num_accepted_tokens.np + if not self.use_async_scheduling: + # Synchronous scheduling already updates counts in current batch order. + accepted[:num_reqs] = self.input_batch.num_accepted_tokens_cpu[:num_reqs] + return + + if has_prev_mapping: + prev_idx = self.prev_positions.np[:num_reqs] + new_mask = prev_idx < 0 + # Advanced indexing reads the snapshot before overwriting its rows. + accepted[:num_reqs] = accepted[np.where(new_mask, 0, prev_idx)] + accepted[:num_reqs][new_mask] = 1 + else: + accepted[:num_reqs].fill(1) + self.input_batch.num_accepted_tokens_cpu[:num_reqs] = accepted[:num_reqs] + def _pad_query_start_loc_for_fia( self, query_start_loc: torch.Tensor, @@ -1088,24 +1145,7 @@ def _prepare_inputs( # _update_states_after_model_execute for hybrid models). if self.num_accepted_tokens_event is not None: self.num_accepted_tokens_event.synchronize() - # Async mode: condense() reordered indices, use prev_positions mapping - if self.use_async_scheduling and prev_req_id_to_index: - prev_idx = self.prev_positions.np[:num_reqs] - new_mask = prev_idx < 0 - self.num_accepted_tokens.np[:num_reqs] = ( - self.input_batch.num_accepted_tokens_cpu[ - np.where(new_mask, 0, prev_idx) - ] - ) - self.num_accepted_tokens.np[:num_reqs][new_mask] = 1 - self.input_batch.num_accepted_tokens_cpu[:num_reqs] = ( - self.num_accepted_tokens.np[:num_reqs] - ) - else: - # Non-async mode: use values directly - self.num_accepted_tokens.np[:num_reqs] = ( - self.input_batch.num_accepted_tokens_cpu[:num_reqs] - ) + self._sync_num_accepted_tokens(num_reqs, has_prev_mapping=bool(prev_req_id_to_index)) self.num_accepted_tokens.np[num_reqs:].fill(1) self.num_accepted_tokens.copy_to_gpu() else: @@ -1988,10 +2028,7 @@ def execute_model( pad_attn = cudagraph_mode == CUDAGraphMode.FULL - # NOTE(Angazenn): According to https://github.com/vllm-project/vllm/pull/30877, - # there should be a corresponding 'postprocess_mamba'. However, it is called inside - # '_update_states_after_model_execute', which is not overridden in vLLM-Ascend. - # We simply utilize the implementation in vLLM. + # postprocess_mamba runs later in _update_states_after_model_execute. if self.cache_config.mamba_cache_mode == "align": # preprocess_mamba reads req_state.num_computed_tokens (CPU) # to decide copy operations, so we must apply deferred From 4a96af9420b90ace5fc9fff7c01e3982f66ee2d5 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Thu, 27 Aug 2026 18:43:51 +0800 Subject: [PATCH 43/50] fix(moe): normalize optional SiTU beta before activation Use the same beta=1.0 default as the existing quantized MoE path before calling native SiTU from the unquantized path. Keep configured beta values unchanged and avoid constructing a runtime CustomOp. Extend the existing upstream parity test to cover an omitted beta across floating dtypes and optional linear clipping. Signed-off-by: maoxx241 --- tests/ut/ops/test_moe_mlp.py | 10 ++++++---- vllm_ascend/ops/fused_moe/moe_mlp.py | 2 +- 2 files changed, 7 insertions(+), 5 deletions(-) diff --git a/tests/ut/ops/test_moe_mlp.py b/tests/ut/ops/test_moe_mlp.py index 841819d3c10b..ee1ef422b31e 100644 --- a/tests/ut/ops/test_moe_mlp.py +++ b/tests/ut/ops/test_moe_mlp.py @@ -121,13 +121,15 @@ def test_uses_small_op_activation_and_dynamic_mx_quant(self): class TestUnifiedApplyMlpRequest(unittest.TestCase): def test_unquant_situ_matches_upstream_without_config_context(self): for dtype in (torch.bfloat16, torch.float16, torch.float32): - for linear_beta in (None, 25.0): - with self.subTest(dtype=dtype, linear_beta=linear_beta): + for beta, linear_beta in ((None, None), (None, 25.0), (4.0, None), (4.0, 25.0)): + with self.subTest(dtype=dtype, beta=beta, linear_beta=linear_beta): hidden_states = torch.randn(2, 8, dtype=dtype) gate_up_out = torch.linspace(-40, 40, 32).reshape(2, 16).to(dtype) expected_output = torch.randn(2, 8, dtype=dtype) with set_current_vllm_config(VllmConfig()): - expected_activation = SituAndMul(beta=4.0, linear_beta=linear_beta).forward_native(gate_up_out) + expected_activation = SituAndMul( + beta=1.0 if beta is None else beta, linear_beta=linear_beta + ).forward_native(gate_up_out) # The worker forward runs outside the model-init config context. with patch( @@ -141,7 +143,7 @@ def test_unquant_situ_matches_upstream_without_config_context(self): w2=torch.randn(1, 8, 8), group_list=torch.tensor([1, 1]), activation=MoEActivation.SITU, - activation_situ_beta=4.0, + activation_situ_beta=beta, activation_situ_linear_beta=linear_beta, need_trans=False, ) diff --git a/vllm_ascend/ops/fused_moe/moe_mlp.py b/vllm_ascend/ops/fused_moe/moe_mlp.py index 61e4ad3de5f9..998dbcdcc69a 100644 --- a/vllm_ascend/ops/fused_moe/moe_mlp.py +++ b/vllm_ascend/ops/fused_moe/moe_mlp.py @@ -722,7 +722,7 @@ def unquant_apply_mlp( if activation == MoEActivation.SITU: gate_up_out = _apply_situ( gate_up_out, - beta=activation_situ_beta, + beta=1.0 if activation_situ_beta is None else activation_situ_beta, linear_beta=activation_situ_linear_beta, ) elif activation == MoEActivation.SWIGLUOAI: From 17d249171c0029add94dd0c0aa34c9209c8e5c5e Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Thu, 27 Aug 2026 19:36:56 +0800 Subject: [PATCH 44/50] test(moe): isolate SiTU reference configuration Give the upstream reference CustomOp only its compilation configuration instead of constructing a full VllmConfig. This avoids platform logging initialization against Ascend mock state left by preceding unit tests. Keep the native upstream numerical comparison and run the MoE call outside the config context. Validate alongside linear, encoder-attention and MoE communication tests. Signed-off-by: maoxx241 --- tests/ut/ops/test_moe_mlp.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/tests/ut/ops/test_moe_mlp.py b/tests/ut/ops/test_moe_mlp.py index ee1ef422b31e..939d98efa2d5 100644 --- a/tests/ut/ops/test_moe_mlp.py +++ b/tests/ut/ops/test_moe_mlp.py @@ -7,7 +7,7 @@ import torch import torch_npu # noqa: F401 -- registers torch.npu used by the module under test from torch.nn import functional as F -from vllm.config import VllmConfig, set_current_vllm_config +from vllm.config import CompilationConfig, VllmConfig, set_current_vllm_config from vllm.model_executor.layers.activation import SituAndMul from vllm.model_executor.layers.fused_moe.activation import MoEActivation @@ -120,15 +120,18 @@ def test_uses_small_op_activation_and_dynamic_mx_quant(self): class TestUnifiedApplyMlpRequest(unittest.TestCase): def test_unquant_situ_matches_upstream_without_config_context(self): + # The reference CustomOp needs dispatch config, not platform/logging setup. + reference_config = MagicMock(spec=VllmConfig) + reference_config.compilation_config = CompilationConfig(custom_ops=["none"]) for dtype in (torch.bfloat16, torch.float16, torch.float32): for beta, linear_beta in ((None, None), (None, 25.0), (4.0, None), (4.0, 25.0)): with self.subTest(dtype=dtype, beta=beta, linear_beta=linear_beta): hidden_states = torch.randn(2, 8, dtype=dtype) gate_up_out = torch.linspace(-40, 40, 32).reshape(2, 16).to(dtype) expected_output = torch.randn(2, 8, dtype=dtype) - with set_current_vllm_config(VllmConfig()): + with set_current_vllm_config(reference_config): expected_activation = SituAndMul( - beta=1.0 if beta is None else beta, linear_beta=linear_beta + beta=1.0 if beta is None else beta, linear_beta=linear_beta, compile_native=False ).forward_native(gate_up_out) # The worker forward runs outside the model-init config context. From b106ea1b48dec9cd1e4b7d3107e2dca24b242df0 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Thu, 27 Aug 2026 23:18:54 +0800 Subject: [PATCH 45/50] fix(mla): preserve absent RoPE metadata during prefill Keep absent cos/sin metadata intact before calling the existing no-RoPE query and KV paths. The PCP prefill changes introduced unconditional slicing before those paths could bypass rotation. Preserve the distinct actual-query and padded-KV lengths for RoPE layers, and cover pure and mixed prefill with numeric Q/K/V and cache-write assertions. Signed-off-by: maoxx241 --- tests/ut/attention/a2/test_mla_v1.py | 43 ++++++++++++++++++++++++++++ vllm_ascend/attention/mla_v1.py | 8 +++--- 2 files changed, 47 insertions(+), 4 deletions(-) diff --git a/tests/ut/attention/a2/test_mla_v1.py b/tests/ut/attention/a2/test_mla_v1.py index 9977ca2943fa..0c9e33bf42d5 100644 --- a/tests/ut/attention/a2/test_mla_v1.py +++ b/tests/ut/attention/a2/test_mla_v1.py @@ -2194,6 +2194,49 @@ def test_mla_preprocess(self): self.assertIsNotNone(decode_res) self.assertIsNotNone(prefill_res) + @patch("vllm_ascend.attention.mla_v1.DeviceOperator.reshape_and_cache") + @patch("torch_npu.npu_interleave_rope") + @patch("torch_npu.npu_kv_rmsnorm_rope_cache") + def test_mla_preprocess_prefill_without_rope(self, mock_rope_cache, mock_rope, mock_cache): + self.impl.use_mla_rope = False + self.impl.num_heads = self.impl.num_kv_heads = 1 + self.impl.qk_nope_head_dim = self.impl.qk_rope_head_dim = 2 + self.impl.qk_head_dim = 4 + self.impl.kv_lora_rank = self.impl.v_head_dim = 2 + self.impl.q_proj.side_effect = lambda x: (x,) + self.impl.kv_a_layernorm = torch.nn.Identity() + self.impl.kv_b_proj.side_effect = lambda x: (torch.cat((x, x), dim=-1),) + + q_c = torch.arange(16, dtype=torch.float32).view(4, 4) + kv_no_split = q_c + 100 + kv_cache = (torch.empty(4, 1, 2), torch.empty(4, 1, 2)) + for num_decode_tokens in (0, 1): + with self.subTest(num_decode_tokens=num_decode_tokens): + num_actual_tokens = num_decode_tokens + 2 + metadata = SimpleNamespace( + num_decode_tokens=num_decode_tokens, + num_actual_tokens=num_actual_tokens, + slot_mapping=torch.arange(4), + prefill=SimpleNamespace(cos=None, sin=None), + ) + result = self.impl.mla_preprocess_prefill(q_c, kv_no_split, kv_cache, metadata) + + q = q_c[num_decode_tokens:num_actual_tokens].unsqueeze(1) + kv = kv_no_split[num_decode_tokens:num_actual_tokens].unsqueeze(1) + torch.testing.assert_close(result.q_nope, q[..., :2]) + torch.testing.assert_close(result.q_pe, q[..., 2:]) + torch.testing.assert_close(result.k_nope, kv[..., :2]) + torch.testing.assert_close(result.k_pe, kv[..., 2:]) + torch.testing.assert_close(result.value, kv[..., :2]) + torch.testing.assert_close(mock_cache.call_args.kwargs["key"], kv[..., :2]) + torch.testing.assert_close(mock_cache.call_args.kwargs["value"], kv[..., 2:]) + torch.testing.assert_close( + mock_cache.call_args.kwargs["slot_mapping"], + metadata.slot_mapping[num_decode_tokens:num_actual_tokens], + ) + mock_rope.assert_not_called() + mock_rope_cache.assert_not_called() + @patch("torch_npu.npu_kv_rmsnorm_rope_cache") def test_exec_kv_prefill(self, mock_kv_rmsnorm_rope_cache): B = 2 diff --git a/vllm_ascend/attention/mla_v1.py b/vllm_ascend/attention/mla_v1.py index 0304b59745a5..ec2e65b542c6 100644 --- a/vllm_ascend/attention/mla_v1.py +++ b/vllm_ascend/attention/mla_v1.py @@ -1829,13 +1829,13 @@ def mla_preprocess_prefill(self, q_c, kv_no_split, kv_cache, attn_metadata): prefill_slots = attn_metadata.slot_mapping[num_decode_tokens:num_actual_tokens] prefill_q_pe = self.rope_single( prefill_q_pe, - cos[:num_actual_prefill_tokens], - sin[:num_actual_prefill_tokens], + cos[:num_actual_prefill_tokens] if cos is not None else None, + sin[:num_actual_prefill_tokens] if sin is not None else None, ) prefill_k_pe, prefill_k_c_normed = self.exec_kv_prefill( prefill_kv_no_split, - cos[:num_prefill_kv_tokens], - sin[:num_prefill_kv_tokens], + cos[:num_prefill_kv_tokens] if cos is not None else None, + sin[:num_prefill_kv_tokens] if sin is not None else None, kv_cache, prefill_slots, attn_metadata=attn_metadata, From 5dc5cb302cfea8d03c2786d0007f92695ddd0e29 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Fri, 28 Aug 2026 00:10:06 +0800 Subject: [PATCH 46/50] perf(kda): restore Ascend fused RMSNorm gate dispatch Register the upstream FLA FusedRMSNormGated class with the Ascend backend so K3 output normalization reaches the existing fused Triton kernel instead of decomposed native operations inside kda_attention. Reuse the v0.26 kernel arithmetic and tiling, preserve the loaded parameters and epsilon, and allocate a separate result as in the v0.26 K3 adapter. Cover CustomOp dispatch plus NPU numerics for packed gate strides, FP16/BF16, sigmoid/SiLU, affine and residual/prenorm paths. Signed-off-by: maoxx241 --- .../ops/a3_2/test_kimi_kda_fused_norm_gate.py | 67 +++++++++++++++++++ tests/ut/ops/test_layernorm.py | 29 ++++++++ vllm_ascend/ops/kimi_kda.py | 5 +- vllm_ascend/ops/layernorm.py | 19 ++++++ vllm_ascend/ops/triton/kda/kda.py | 1 + vllm_ascend/utils.py | 3 +- 6 files changed, 120 insertions(+), 4 deletions(-) create mode 100644 tests/ut/ops/a3_2/test_kimi_kda_fused_norm_gate.py diff --git a/tests/ut/ops/a3_2/test_kimi_kda_fused_norm_gate.py b/tests/ut/ops/a3_2/test_kimi_kda_fused_norm_gate.py new file mode 100644 index 000000000000..c9a46fd8f24c --- /dev/null +++ b/tests/ut/ops/a3_2/test_kimi_kda_fused_norm_gate.py @@ -0,0 +1,67 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest +import torch + +from vllm_ascend.ops.triton.kda.kda import rms_norm_gated + + +@torch.inference_mode() +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize("tokens, heads, head_dim", [(1, 1, 128), (16, 4, 128), (37, 4, 128), (2, 1, 1024)]) +@pytest.mark.parametrize("strided_gate", [False, True]) +def test_kimi_kda_fused_rms_norm_sigmoid_gate(dtype, tokens, heads, head_dim, strided_gate): + torch.manual_seed(20260801) + eps = 1e-6 + weight = torch.randn(head_dim, dtype=dtype, device="npu") + core_attn_out = torch.randn(1, tokens, heads, head_dim, dtype=dtype, device="npu") + output_gate = torch.randn(tokens, heads, head_dim, dtype=dtype, device="npu") + if strided_gate: + # K3's packed projection leaves gaps between consecutive gate rows. + packed_gate = torch.full((tokens, 2 * heads, head_dim), torch.nan, dtype=dtype, device="npu") + packed_gate[:, :heads].copy_(output_gate) + output_gate = packed_gate[:, :heads] + core_attn_out_before = core_attn_out.clone() + output_gate_before = output_gate.clone() + + actual = rms_norm_gated(core_attn_out, output_gate, weight, None, activation="sigmoid", eps=eps) + + x_float = core_attn_out_before.float() + variance = x_float.square().mean(dim=-1, keepdim=True) + expected = x_float * torch.rsqrt(variance + eps) + expected = expected * weight.float() + expected = expected * output_gate.float().sigmoid().unsqueeze(0) + + torch.testing.assert_close(actual, expected.to(dtype), rtol=2e-3, atol=2e-3) + torch.testing.assert_close(core_attn_out, core_attn_out_before, rtol=0, atol=0) + torch.testing.assert_close(output_gate, output_gate_before, rtol=0, atol=0) + + +@torch.inference_mode() +@pytest.mark.parametrize("residual_dtype", [None, torch.bfloat16, torch.float32]) +@pytest.mark.parametrize("elementwise_affine", [False, True]) +def test_fused_rms_norm_silu_gate_preserves_prenorm_contract(residual_dtype, elementwise_affine): + torch.manual_seed(20260801) + x = torch.randn(1, 3, 2, 128, dtype=torch.bfloat16, device="npu") + gate = torch.randn_like(x) + residual = torch.randn_like(x, dtype=residual_dtype) if residual_dtype is not None else None + weight = torch.randn(128, dtype=x.dtype, device="npu") if elementwise_affine else None + before = x.clone() + eps = 1e-6 + + actual, residual_out = rms_norm_gated( + x, gate, weight, None, activation="silu", residual=residual, prenorm=True, residual_in_fp32=True, eps=eps + ) + + summed = before.float() if residual is None else before.float() + residual.float() + expected = summed * torch.rsqrt(summed.square().mean(-1, keepdim=True) + eps) + if weight is not None: + expected *= weight.float() + expected *= gate.float() * gate.float().sigmoid() + expected_residual_dtype = torch.float32 if residual is None else residual.dtype + + torch.testing.assert_close(actual, expected.to(x.dtype), rtol=2e-3, atol=2e-3) + assert residual_out.dtype == expected_residual_dtype + torch.testing.assert_close(residual_out, summed.to(expected_residual_dtype), rtol=0, atol=0) + torch.testing.assert_close(x, before, rtol=0, atol=0) diff --git a/tests/ut/ops/test_layernorm.py b/tests/ut/ops/test_layernorm.py index a296ffa8b5bb..4fa8b72bceaf 100644 --- a/tests/ut/ops/test_layernorm.py +++ b/tests/ut/ops/test_layernorm.py @@ -4,7 +4,9 @@ import torch from vllm.config import set_current_vllm_config from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.third_party.flash_linear_attention.ops.kda import FusedRMSNormGated +from vllm_ascend.ops.layernorm import AscendFusedRMSNormGated from vllm_ascend.utils import enable_custom_op from vllm_ascend.utils import is_310p as is_310p_hw @@ -83,6 +85,33 @@ def test_RMSNorm_creates_bias_from_quant_description(default_vllm_config): assert not layer.bias.requires_grad +@pytest.mark.parametrize("activation", ["sigmoid", "swish"]) +@pytest.mark.parametrize("prenorm", [False, True]) +def test_FusedRMSNormGated_dispatches_to_ascend_kernel(default_vllm_config, activation, prenorm): + layer = FusedRMSNormGated(hidden_size=8, eps=1e-6, activation=activation) + x = torch.randn(1, 4, 2, 8) + gate = torch.randn(4, 2, 8) + residual = torch.randn_like(x) if prenorm else None + expected = (torch.empty_like(x), torch.empty_like(x)) if prenorm else torch.empty_like(x) + + with patch("vllm_ascend.ops.layernorm.rms_norm_gated", return_value=expected) as fused_norm_gate: + actual = layer(x, gate, residual=residual, prenorm=prenorm, residual_in_fp32=prenorm) + + assert isinstance(layer, AscendFusedRMSNormGated) + assert actual is expected + fused_norm_gate.assert_called_once_with( + x, + gate, + layer.weight, + layer.bias, + activation, + residual=residual, + eps=1e-6, + prenorm=prenorm, + residual_in_fp32=prenorm, + ) + + @pytest.mark.skipif(not is_310p_hw(), reason="310P device unittest case.") @pytest.mark.parametrize("residual", [None, torch.randn(4, 8, dtype=torch.float16)]) @patch("torch_npu.npu_rms_norm", side_effect=mock_rms_norm) diff --git a/vllm_ascend/ops/kimi_kda.py b/vllm_ascend/ops/kimi_kda.py index a2055aa2e88c..344189cc8655 100644 --- a/vllm_ascend/ops/kimi_kda.py +++ b/vllm_ascend/ops/kimi_kda.py @@ -535,9 +535,8 @@ def _forward( elif core_non_spec is not None: core_attn_out[:, :num_actual_tokens] = core_non_spec - # Let vLLM's CustomOp dispatch select FusedRMSNormGated.forward_native - # on Ascend. Calling the CUDA/Triton helper directly bypasses platform - # dispatch. + # The registered Ascend FusedRMSNormGated uses the fused norm-gate + # kernel while preserving the upstream parameter/loading contract. normalized = self.o_norm(core_attn_out[:, :num_actual_tokens], g2) # Mask again after the norm gate: zero * sigmoid(NaN) is still NaN in # static padding rows whose captured gate values are not live. diff --git a/vllm_ascend/ops/layernorm.py b/vllm_ascend/ops/layernorm.py index d88246743ed4..cbe73d0b06c0 100644 --- a/vllm_ascend/ops/layernorm.py +++ b/vllm_ascend/ops/layernorm.py @@ -19,8 +19,10 @@ from torch import nn from vllm.config import get_current_vllm_config from vllm.model_executor.layers.layernorm import GemmaRMSNorm, RMSNorm, RMSNormGated +from vllm.third_party.flash_linear_attention.ops.kda import FusedRMSNormGated from vllm_ascend.device.device_op import DeviceOperator +from vllm_ascend.ops.triton.kda.kda import rms_norm_gated from vllm_ascend.ops.triton.layernorm_gated import layer_norm_fwd_npu from vllm_ascend.utils import enable_custom_op @@ -194,3 +196,20 @@ def reset_parameters(self): def forward_oot(self, x, z=None): """If z is not None, we do norm(x) * silu(z) if norm_before_gate, else norm(x * silu(z))""" return LayerNormFn.apply(x, self.weight, self.bias, z, self.eps, self.group_size, self.norm_before_gate, True) + + +class AscendFusedRMSNormGated(FusedRMSNormGated): + """Use Ascend's fused kernel at the upstream FLA CustomOp boundary.""" + + def forward_oot(self, x, g, residual=None, prenorm=False, residual_in_fp32=False): + return rms_norm_gated( + x, + g, + self.weight, + self.bias, + self.activation, + residual=residual, + eps=self.eps, + prenorm=prenorm, + residual_in_fp32=residual_in_fp32, + ) diff --git a/vllm_ascend/ops/triton/kda/kda.py b/vllm_ascend/ops/triton/kda/kda.py index fc09fde39677..313b3b8949b7 100644 --- a/vllm_ascend/ops/triton/kda/kda.py +++ b/vllm_ascend/ops/triton/kda/kda.py @@ -411,6 +411,7 @@ def rms_norm_gated( activation=activation, eps=eps, residual=residual, + out_dtype=x.dtype, # Preserve the input, as in the v0.26 K3 fused norm gate. residual_dtype=residual_dtype, is_rms_norm=True, ) diff --git a/vllm_ascend/utils.py b/vllm_ascend/utils.py index 9555f27774dc..d6518d82fcae 100644 --- a/vllm_ascend/utils.py +++ b/vllm_ascend/utils.py @@ -674,7 +674,7 @@ def register_ascend_customop(vllm_config: VllmConfig | None = None): from vllm_ascend.ops.fused_moe.fused_moe import AscendMoERunner from vllm_ascend.ops.fused_moe.routed_experts import AscendRoutedExperts from vllm_ascend.ops.gdn import AscendGatedDeltaNetAttention - from vllm_ascend.ops.layernorm import AscendGemmaRMSNorm, AscendRMSNorm, AscendRMSNormGated + from vllm_ascend.ops.layernorm import AscendFusedRMSNormGated, AscendGemmaRMSNorm, AscendRMSNorm, AscendRMSNormGated from vllm_ascend.ops.linear import ( AscendColumnParallelLinear, AscendMergedColumnParallelLinear, @@ -722,6 +722,7 @@ def register_ascend_customop(vllm_config: VllmConfig | None = None): "MMEncoderAttention": AscendMMEncoderAttention, "ApplyRotaryEmb": AscendApplyRotaryEmb, "RMSNormGated": AscendRMSNormGated, + "FusedRMSNormGated": AscendFusedRMSNormGated, "Conv3dLayer": AscendConv3dLayer, "RelPosAttention": AscendRelPosAttention, "CustomQwen2Decoder": AscendCustomQwen2Decoder, From 3b6c845bbed21e786196272b2bb25250b67d1a99 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Fri, 28 Aug 2026 00:25:47 +0800 Subject: [PATCH 47/50] fix(ci): register fused norm gate test duration Add the measured NPU norm-gate test to estimated_times so selective CI coverage validation can schedule the existing test. Signed-off-by: maoxx241 --- .github/workflows/scripts/test_config.yaml | 1 + 1 file changed, 1 insertion(+) diff --git a/.github/workflows/scripts/test_config.yaml b/.github/workflows/scripts/test_config.yaml index 6eba786b6ecb..4b555bfd54d0 100644 --- a/.github/workflows/scripts/test_config.yaml +++ b/.github/workflows/scripts/test_config.yaml @@ -915,6 +915,7 @@ estimated_times: tests/ut/ops/a2/test_gdn_layerwise_kv.py: 70 tests/ut/ops/a2/test_token_dispatcher.py: 30 tests/ut/ops/a3_2/test_activation.py: 50 + tests/ut/ops/a3_2/test_kimi_kda_fused_norm_gate.py: 60 tests/ut/ops/a3_2/test_select_experts.py: 20 tests/ut/quantization/methods/a2/test_w4a16.py: 30 tests/ut/quantization/methods/a2/test_w4a4_flatquant.py: 40 From 8c8deeddd853742dd297c5c360972ed48fec5760 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Fri, 28 Aug 2026 02:10:20 +0800 Subject: [PATCH 48/50] test(kimi-k3): cover A3 dummy deployment variants Build reduced 16-layer/16-expert K3 configs in one A3 16-device PR test, covering MLA and GQA DSpark, W4A8, DP2/TP8, MTP images, prefix-cache boundaries and P/D transfer with local prefill fallback. Restore DP-wide MC2 padding around local SP shards while preserving rank-local masks and hash token IDs. Match upstream text-only draft handling, use the target language-model head for MTP, and adapt its decoder attention return convention. Exclude K3 MLA from Transformers' inherited MHA divisibility check. Validated all six A3 functional cases on vLLM 0.27.1, MC2 unit tests, targeted mypy for Python 3.10/3.11/3.12, CI routing/coverage and format.sh ci. Dummy results are not accuracy or QuaRot checkpoint validation. Signed-off-by: maoxx241 --- .github/workflows/scripts/runner_label.json | 6 + .github/workflows/scripts/test_config.yaml | 30 + .../pull_request/sixteen_card/test_kimi_k3.py | 539 ++++++++++++++++++ tests/ut/ops/test_prepare_finalize.py | 49 ++ vllm_ascend/models/kimi_k3.py | 8 + vllm_ascend/ops/fused_moe/prepare_finalize.py | 42 +- .../platform/patch_speculative_config.py | 19 +- vllm_ascend/spec_decode/llm_base_proposer.py | 11 +- 8 files changed, 697 insertions(+), 7 deletions(-) create mode 100644 tests/e2e/pull_request/sixteen_card/test_kimi_k3.py diff --git a/.github/workflows/scripts/runner_label.json b/.github/workflows/scripts/runner_label.json index e6e5ece3b714..f57abbace1e2 100644 --- a/.github/workflows/scripts/runner_label.json +++ b/.github/workflows/scripts/runner_label.json @@ -35,6 +35,12 @@ "image_tag": "9.1.0-a3-ubuntu22.04-py3.12", "csrc_cache_target": "a3-arm64-ubuntu" }, + "linux-aarch64-a3-16": { + "chip": "a3", + "npu_num": 16, + "image_tag": "9.1.0-a3-ubuntu22.04-py3.12", + "csrc_cache_target": "a3-arm64-ubuntu" + }, "linux-aarch64-a3-2-": { "chip": "a3", "npu_num": 2, diff --git a/.github/workflows/scripts/test_config.yaml b/.github/workflows/scripts/test_config.yaml index 4b555bfd54d0..f83f3984e8fd 100644 --- a/.github/workflows/scripts/test_config.yaml +++ b/.github/workflows/scripts/test_config.yaml @@ -322,6 +322,30 @@ - tests/e2e/pull_request/one_card/test_minimax_m3_sparse_attn.py - tests/e2e/pull_request/eight_card/test_minimax_m3.py +- name: models_kimi_k3 + optional: true + source_file_dependencies: + - vllm_ascend/models/kimi_k3.py + - vllm_ascend/models/kimi_k3_dspark.py + - vllm_ascend/models/kimi_k3_mtp.py + - vllm_ascend/models/qwen3_dspark.py + - vllm_ascend/worker/model_runner_v1.py + - vllm_ascend/worker/worker.py + - vllm_ascend/attention/mla_v1.py + - vllm_ascend/core/kv_cache_interface.py + - vllm_ascend/ops/kimi_kda.py + - vllm_ascend/ops/fused_moe + - vllm_ascend/ops/triton/kimi_k3 + - vllm_ascend/spec_decode/dspark_proposer.py + - vllm_ascend/spec_decode/llm_base_proposer.py + - vllm_ascend/quantization/modelslim_config.py + - vllm_ascend/patch/platform/patch_kv_cache_utils.py + - vllm_ascend/patch/platform/patch_kv_cache_coordinator.py + - vllm_ascend/patch/platform/patch_speculative_config.py + - vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py + tests: + - tests/e2e/pull_request/sixteen_card/test_kimi_k3.py + # Ops - name: ops_basic optional: false @@ -894,6 +918,7 @@ estimated_times: tests/e2e/pull_request/four_card/test_qwen3_32b_bf16_ll_performance.py: 1200 tests/e2e/pull_request/two_card/test_gemma4.py: 420 tests/e2e/pull_request/eight_card/test_minimax_m3.py: 630 + tests/e2e/pull_request/sixteen_card/test_kimi_k3.py: 1200 tests/ut/attention/a2/test_attention_cp.py: 30 tests/ut/attention/a2/test_attention_cp_precision.py: 30 tests/ut/attention/a2/test_attention_v1.py: 30 @@ -974,6 +999,8 @@ runner_mapping: 310p: 310p-4 tests/e2e/pull_request/eight_card: default: a3-8 + tests/e2e/pull_request/sixteen_card: + default: a3-16 tests/ut/.+/a3_2: default: a3-2 tests/ut/.+/a2: @@ -1007,6 +1034,9 @@ partition: a3-8: runner_label: linux-aarch64-a3-8- count: 1 + a3-16: + runner_label: linux-aarch64-a3-16 + count: 1 cpu-0: runner_label: linux-amd64-cpu-8-hk count: 1 diff --git a/tests/e2e/pull_request/sixteen_card/test_kimi_k3.py b/tests/e2e/pull_request/sixteen_card/test_kimi_k3.py new file mode 100644 index 000000000000..45a9d9a6eb9f --- /dev/null +++ b/tests/e2e/pull_request/sixteen_card/test_kimi_k3.py @@ -0,0 +1,539 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Single-node A3 (16 logical NPUs) K3 functional tests, not accuracy tests. + +Build local configs and initialize dummy weights; no full checkpoint is needed. +Keep production tensor widths while reducing the target to 16 layers/experts. +Random weights cannot validate checkpoint loading, QuaRot or acceptance rates. +""" + +import json +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path + +import pytest +import requests +from PIL import Image +from prometheus_client.parser import text_string_to_metric_families +from tokenizers import Tokenizer # type: ignore[import-untyped] +from tokenizers.models import WordLevel # type: ignore[import-untyped] +from tokenizers.pre_tokenizers import Whitespace # type: ignore[import-untyped] +from transformers import PreTrainedTokenizerFast +from vllm import SamplingParams +from vllm.utils.network_utils import get_open_port +from vllm.v1.metrics.reader import Counter + +from tests.e2e.conftest import RemoteOpenAIServer, RemotePDServer, VllmRunner + +NUM_LAYERS = 16 +NUM_EXPERTS = 16 +NUM_VISION_LAYERS = 2 +VOCAB_SIZE = 163840 +MAX_MODEL_LEN = 2048 +OUTPUT_TOKENS = 2 +# Keep the final consecutive MLA layers as well as the regular KDA/MLA pattern. +# Layer 13 also crosses the production attention-residual block boundary (12). +FULL_ATTN_LAYERS = (4, 8, 12, 15, 16) + + +def _text_config() -> dict: + return { + "architectures": ["KimiLinearForCausalLM"], + "model_type": "kimi_linear", + "torch_dtype": "bfloat16", + "hidden_size": 7168, + "intermediate_size": 33792, + "num_hidden_layers": NUM_LAYERS, + "num_experts": NUM_EXPERTS, + "num_experts_per_token": 16, + "num_shared_experts": 2, + "moe_intermediate_size": 3072, + "routed_expert_hidden_size": 3584, + "first_k_dense_replace": 1, + "moe_layer_freq": 1, + "hidden_act": "situ", + "activation_situ_beta": 4.0, + "activation_situ_linear_beta": 25.0, + "latent_moe_use_norm": True, + "moe_router_activation_func": "sigmoid", + "use_grouped_topk": True, + "num_expert_group": 1, + "topk_group": 1, + "topk_method": "noaux_tc", + "moe_renormalize": True, + "attn_res_block_size": 12, + "num_attention_heads": 96, + "num_key_value_heads": 96, + "q_lora_rank": 1536, + "kv_lora_rank": 512, + "qk_nope_head_dim": 128, + "qk_rope_head_dim": 64, + "v_head_dim": 128, + "mla_use_nope": True, + "mla_use_output_gate": True, + "rms_norm_eps": 1e-5, + "vocab_size": VOCAB_SIZE, + "bos_token_id": 163584, + "eos_token_id": 163586, + "pad_token_id": 163839, + "tie_word_embeddings": False, + "max_position_embeddings": 8192, + "num_nextn_predict_layers": 0, + "linear_attn_config": { + "head_dim": 128, + "num_heads": 96, + "short_conv_kernel_size": 4, + "use_full_rank_gate": True, + "gate_lower_bound": -5.0, + "full_attn_layers": list(FULL_ATTN_LAYERS), + "kda_layers": [i for i in range(1, NUM_LAYERS + 1) if i not in FULL_ATTN_LAYERS], + }, + } + + +def _draft_config(variant: str) -> dict: + config: dict = { + "hidden_size": 7168, + "intermediate_size": 14336, + "hidden_act": "silu", + "num_hidden_layers": 5, + "num_attention_heads": 64, + "num_key_value_heads": 64, + "rms_norm_eps": 1e-5, + "vocab_size": VOCAB_SIZE, + "draft_vocab_size": VOCAB_SIZE, + "bos_token_id": 163584, + "eos_token_id": 163586, + "pad_token_id": 163839, + "torch_dtype": "bfloat16", + "tie_word_embeddings": False, + "max_position_embeddings": 8192, + "markov_rank": 256, + "markov_head_type": "vanilla", + "enable_confidence_head": True, + "confidence_head_with_markov": True, + } + if variant == "gqa": + config.update( + architectures=["DSparkDraftModel"], + model_type="qwen3", + num_key_value_heads=16, + head_dim=64, + layer_types=["full_attention"] * 5, + block_size=7, + num_target_layers=NUM_LAYERS, + # The drafter consumes intermediate layer outputs, before the final + # target layer. Remap all five taps to the reduced target. + dflash_config={"mask_token_id": 163824, "target_layer_ids": [1, 3, 7, 11, 14]}, + rope_parameters={ + "rope_type": "yarn", + "factor": 16.0, + "original_max_position_embeddings": 65536, + "rope_theta": 10000.0, + }, + ) + return config + + config.update( + architectures=["K3DSparkModel"], + model_type="k3_dspark", + q_lora_rank=1536, + kv_lora_rank=512, + qk_nope_head_dim=128, + qk_rope_head_dim=64, + v_head_dim=128, + target_hidden_size=7168, + target_num_hidden_layers=NUM_LAYERS, + num_target_layers=5, + target_layer_ids=[1, 3, 7, 11, 14], + mask_token_id=163837, + rope_parameters={ + "rope_type": "yarn", + "factor": 32.0, + "original_max_position_embeddings": 32768, + "rope_theta": 50000.0, + "beta_fast": 32, + "beta_slow": 1, + "mscale": 1.0, + "mscale_all_dim": 1.0, + }, + ) + if variant == "mla_block5": + config.update( + num_hidden_layers=3, + num_attention_heads=96, + num_key_value_heads=96, + head_dim=64, + qk_head_dim=192, + num_target_layers=3, + target_layer_ids=[7, 11, 14], + layer_types=["full_attention"] * 3, + mask_token_id=163839, + markov_rank=512, + sample_from_anchor=True, + block_size=5, + full_attention_causal=True, + dflash_config={"causal": True}, + rope_interleave=True, + ) + config["rope_parameters"].update(rope_theta=1000000.0, mscale_all_dim=0.0, attn_factor=0.7426255848312643) + return config + + +def _write_config(path: Path, config: dict) -> str: + path.mkdir() + (path / "config.json").write_text(json.dumps(config), encoding="utf-8") + return str(path) + + +def _write_target(path: Path, *, mtp: bool = False) -> str: + text_config = _text_config() + text_config["num_nextn_predict_layers"] = int(mtp) + _write_config( + path, + { + "architectures": ["KimiK3ForConditionalGeneration"], + "model_type": "kimi_k3", + "text_config": text_config, + "vision_config": {"vt_num_hidden_layers": NUM_VISION_LAYERS, "text_hidden_size": 7168}, + "media_placeholder_token_id": 163605, + }, + ) + # A small local tokenizer/processor description keeps the multimodal wrapper + # offline too. Token IDs and vocabulary size still match the real K3 model. + special_tokens = { + 0: "", + 163584: "", + 163586: "", + 163600: "<|kimi_image_placeholder|>", + 163601: "<|media_begin|>", + 163602: "<|media_content|>", + 163603: "<|media_end|>", + 163605: "<|media_pad|>", + 163839: "", + } + vocabulary = {special_tokens.get(i, f"token_{i}"): i for i in range(VOCAB_SIZE)} + tokenizer = Tokenizer(WordLevel(vocabulary, unk_token="")) + tokenizer.pre_tokenizer = Whitespace() + PreTrainedTokenizerFast( + tokenizer_object=tokenizer, + unk_token="", + bos_token="", + eos_token="", + pad_token="", + additional_special_tokens=list(special_tokens.values()), + chat_template="{% for message in messages %}{{ message['content'] }}{% endfor %}", + ).save_pretrained(path) + # K3 and K2.5 use the same MoonViT patch format. Reuse vLLM's native image + # preprocessor rather than copying checkpoint Python code into the test. + (path / "image_processing_k3_dummy.py").write_text( + "from vllm.transformers_utils.processors.kimi_k25_vision_fused import KimiK25FusedVisionProcessor\n", + encoding="utf-8", + ) + (path / "preprocessor_config.json").write_text( + json.dumps( + { + "auto_map": {"AutoImageProcessor": "image_processing_k3_dummy.KimiK25FusedVisionProcessor"}, + "media_proc_cfg": { + "patch_size": 14, + "merge_kernel_size": 2, + "temporal_merge_kernel_size": 4, + "in_patch_limit": 256, + "patch_limit_on_one_side": 16, + "fixed_output_tokens": None, + "image_mean": [0.5, 0.5, 0.5], + "image_std": [0.5, 0.5, 0.5], + }, + } + ), + encoding="utf-8", + ) + return str(path) + + +def _write_w4a8_description(path: Path) -> None: + # Use the released mixed KDA precision: W8A8 q/k/v, floating-point gates; + # routed experts use W4A8. No rotation file means this is NOT a QuaRot test. + description = {"model.embed_tokens.weight": "FLOAT", "lm_head.weight": "FLOAT", "optional": {}} + for layer in range(NUM_LAYERS): + prefix = f"model.layers.{layer}" + for projection in ("self_attention_res_proj", "mlp_res_proj"): + description[f"{prefix}.{projection}.weight"] = "FLOAT" + if layer + 1 in FULL_ATTN_LAYERS: + quantized = ("q_a_proj", "q_b_proj", "kv_a_proj_with_mqa") + floating: tuple[str, ...] = ("kv_b_proj", "o_proj", "g_proj") + else: + quantized = ("q_proj", "k_proj", "v_proj") + floating = ("g_proj", "f_a_proj", "f_b_proj", "b_proj", "o_proj", "q_conv1d", "k_conv1d", "v_conv1d") + for projection in quantized: + description[f"{prefix}.self_attn.{projection}.weight"] = "W8A8_DYNAMIC" + for projection in floating: + description[f"{prefix}.self_attn.{projection}.weight"] = "FLOAT" + mlp_prefix = f"{prefix}.block_sparse_moe.shared_experts" if layer else f"{prefix}.mlp" + for projection in ("gate_proj", "up_proj", "down_proj"): + description[f"{mlp_prefix}.{projection}.weight"] = "W8A8_DYNAMIC" + if layer: + for projection in ("routed_expert_down_proj", "routed_expert_up_proj"): + description[f"{prefix}.block_sparse_moe.{projection}.weight"] = "W8A8_DYNAMIC" + for expert in range(NUM_EXPERTS): + for projection in ("w1", "w2", "w3"): + description[f"{prefix}.block_sparse_moe.experts.{expert}.{projection}.weight"] = "W4A8_DYNAMIC" + description = {f"language_model.{key}" if key != "optional" else key: value for key, value in description.items()} + for layer in range(NUM_VISION_LAYERS): + for projection in ("mlp.fc0", "mlp.fc1", "wqkv", "wo"): + description[f"vision_tower.encoder.blocks.{layer}.{projection}.weight"] = "FLOAT" + for projection in ("proj.0", "proj.2", "rot_proj"): + description[f"mm_projector.{projection}.weight"] = "FLOAT" + (path / "quant_model_description.json").write_text(json.dumps(description), encoding="utf-8") + + +@pytest.fixture +def k3_models(tmp_path: Path) -> dict[str, str]: + models = {"target": _write_target(tmp_path / "target")} + models["w4a8"] = _write_target(tmp_path / "w4a8") + _write_w4a8_description(tmp_path / "w4a8") + for variant in ("mla", "mla_block5", "gqa"): + models[variant] = _write_config(tmp_path / variant, _draft_config(variant)) + return models + + +@pytest.fixture(autouse=True) +def k3_runtime(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("VLLM_USE_V2_MODEL_RUNNER", "0") + monkeypatch.setenv("VLLM_WORKER_MULTIPROC_METHOD", "spawn") + monkeypatch.setenv("HCCL_OP_EXPANSION_MODE", "AIV") + monkeypatch.setenv("HCCL_BUFFSIZE", "512") + + +def _engine_args(models: dict[str, str], variant: str, tp: int = 16) -> dict: + steps = 5 if variant == "mla_block5" else 7 + # Respect upstream LCM(TP, speculative query width), including Block5's 48. + graph_sizes = [48, 96] if steps == 5 else [16, 32] + return { + "load_format": "dummy", + "dtype": "bfloat16", + "tensor_parallel_size": tp, + "enable_expert_parallel": True, + "distributed_executor_backend": "mp", + "max_model_len": MAX_MODEL_LEN, + "max_num_seqs": 4, + "max_num_batched_tokens": 512, + "block_size": 128, + "kv_cache_memory_bytes": 512 * 1024**2, + "gpu_memory_utilization": 0.8, + "enable_prefix_caching": True, + "enable_chunked_prefill": True, + "mamba_cache_mode": "align", + "async_scheduling": True, + "disable_log_stats": False, + "seed": 0, + "limit_mm_per_prompt": {"image": 0}, + "mm_encoder_tp_mode": "data", + "speculative_config": { + "method": "dspark", + "model": models[variant], + "num_speculative_tokens": steps, + "draft_sample_method": "greedy", + "enforce_eager": True, + "draft_load_config": {"load_format": "dummy"}, + }, + "compilation_config": {"cudagraph_mode": "FULL_DECODE_ONLY", "cudagraph_capture_sizes": graph_sizes}, + "additional_config": { + "enable_cpu_binding": False, + "enable_shared_expert_dp": False, + "multistream_overlap_shared_expert": True, + "enable_fused_mc2": 0, + }, + } + + +def _prompt(length: int, salt: int = 0) -> dict: + return {"prompt_token_ids": [10 + (i + salt) % 1000 for i in range(length)]} + + +def _generate(llm, prompts: list[dict]): + print(f"K3 smoke: generating {OUTPUT_TOKENS} tokens for {len(prompts)} requests", flush=True) + outputs = llm.generate( + prompts, + SamplingParams(temperature=0, max_tokens=OUTPUT_TOKENS, ignore_eos=True, detokenize=False), + use_tqdm=False, + ) + assert len(outputs) == len(prompts) + for output in outputs: + assert output.finished + assert len(output.outputs) == 1 + assert len(output.outputs[0].token_ids) == OUTPUT_TOKENS + assert all(0 <= token < VOCAB_SIZE for token in output.outputs[0].token_ids) + print(f"K3 smoke completed; cached tokens: {[output.num_cached_tokens for output in outputs]}", flush=True) + return outputs + + +@pytest.mark.parametrize("variant", ["mla", "mla_block5", "gqa"]) +def test_k3_dspark_tp16(k3_models: dict[str, str], variant: str) -> None: + args = _engine_args(k3_models, variant) + target = k3_models["target"] + if variant == "gqa": + args["quantization"] = "ascend" + target = k3_models["w4a8"] + with VllmRunner(target, **args) as runner: + llm = runner.model + _generate(llm, [_prompt(1)]) + # Kernel-block and block+1 prefill, mixed lengths, and chunked prefill. + _generate(llm, [_prompt(length, salt=i * 137) for i, length in enumerate((127, 128, 129, 769))]) + _generate(llm, [_prompt(length, salt=311 + i * 137) for i, length in enumerate((383, 384, 385, 1537))]) + + prefix = _prompt(1153, salt=421) + assert llm.reset_prefix_cache() + cold = _generate(llm, [prefix])[0] + warm = _generate(llm, [prefix])[0] + assert cold.num_cached_tokens == 0 + assert warm.num_cached_tokens > 0, "Repeated prompt did not reuse the hybrid prefix cache" + assert llm.reset_prefix_cache() + assert _generate(llm, [prefix])[0].num_cached_tokens == 0 + + # Exercise the final block-table columns plus speculative lookahead. + _generate(llm, [_prompt(MAX_MODEL_LEN - OUTPUT_TOKENS, salt=713)]) + drafts = [m for m in llm.get_metrics() if m.name == "vllm:spec_decode_num_drafts"] + assert drafts and all(isinstance(m, Counter) for m in drafts) + assert sum(m.value for m in drafts) > 0, "Requests bypassed speculative decoding" + + +def _serve_args(args: dict) -> list[str]: + result = [ + "--served-model-name", + "k3-dummy", + "--host", + "127.0.0.1", + "--trust-remote-code", + "--enable-prompt-tokens-details", + ] + for name, value in args.items(): + option = "--" + name.replace("_", "-") + if isinstance(value, bool): + if value: + result.append(option) + else: + result.extend([option, json.dumps(value) if isinstance(value, dict) else str(value)]) + return result + + +def _completion(url: str, prompt: list[int], *, max_tokens: int = OUTPUT_TOKENS, **kwargs) -> dict: + response = requests.post( + url, + json={ + "model": "k3-dummy", + "prompt": prompt, + "max_tokens": max_tokens, + "ignore_eos": True, + "temperature": 0, + "return_token_ids": True, + **kwargs, + }, + timeout=180, + ) + response.raise_for_status() + output = response.json() + assert output["usage"]["completion_tokens"] == max_tokens + assert len(output["choices"][0]["token_ids"]) == max_tokens + assert all(0 <= token < VOCAB_SIZE for token in output["choices"][0]["token_ids"]) + return output + + +def _draft_counts(url: str) -> dict[str, float]: + response = requests.get(url, timeout=10) + response.raise_for_status() + return { + sample.labels["engine"]: sample.value + for family in text_string_to_metric_families(response.text) + for sample in family.samples + if sample.name == "vllm:spec_decode_num_drafts_total" + } + + +def test_k3_gqa_dp2_tp8(k3_models: dict[str, str]) -> None: + args = _engine_args(k3_models, "gqa", tp=8) + args["data_parallel_size"] = 2 + args["additional_config"]["enable_shared_expert_dp"] = True + port = get_open_port() + with RemoteOpenAIServer( + k3_models["target"], + [*_serve_args(args), "--port", str(port)], + server_host="127.0.0.1", + server_port=port, + auto_port=False, + max_wait_seconds=600, + ) as server: + prompts = [_prompt(length, salt=i * 137)["prompt_token_ids"] for i, length in enumerate((1, 129, 769, 1153))] + with ThreadPoolExecutor(max_workers=4) as pool: + outputs = list(pool.map(lambda prompt: _completion(server.url_for("v1", "completions"), prompt), prompts)) + assert len(outputs) == len(prompts) + counts = _draft_counts(server.url_for("metrics")) + assert len(counts) == 2 and all(count > 0 for count in counts.values()), counts + + +def test_k3_mtp_image_tp16(k3_models: dict[str, str], tmp_path: Path) -> None: + target = _write_target(tmp_path / "mtp", mtp=True) + args = _engine_args(k3_models, "gqa") + args["speculative_config"] = { + "method": "mtp", + "num_speculative_tokens": 1, + "enforce_eager": True, + "draft_load_config": {"load_format": "dummy"}, + } + args["limit_mm_per_prompt"] = {"image": 1} + with VllmRunner(target, **args) as runner: + _generate(runner.model, [_prompt(1)]) + image_prompt = { + "prompt": "<|kimi_image_placeholder|> describe", + "multi_modal_data": {"image": Image.new("RGB", (56, 56), (32, 64, 128))}, + } + _generate(runner.model, [image_prompt]) + drafts = [m for m in runner.model.get_metrics() if m.name == "vllm:spec_decode_num_drafts"] + assert drafts and sum(m.value for m in drafts) > 0, "Requests bypassed MTP" + + +def test_k3_gqa_pd_tp8(k3_models: dict[str, str]) -> None: + prefill_port, decode_port = get_open_port(), get_open_port() + transfer_config = { + "kv_connector": "MooncakeConnectorV1", + "kv_connector_extra_config": { + "prefill": {"dp_size": 1, "tp_size": 8}, + "decode": {"dp_size": 1, "tp_size": 8}, + }, + } + prefill_args = _engine_args(k3_models, "gqa", tp=8) + # P also builds the draft KV that D consumes; both peers need the same + # target/draft layer layout even though P only generates one token. + prefill_args.pop("compilation_config") + prefill_args["enforce_eager"] = True + prefill_args["kv_transfer_config"] = dict(transfer_config, kv_role="kv_producer", kv_port=get_open_port()) + decode_args = _engine_args(k3_models, "gqa", tp=8) + decode_args["kv_transfer_config"] = dict(transfer_config, kv_role="kv_consumer", kv_port=get_open_port()) + servers = [ + [k3_models["target"], "--port", str(prefill_port), *_serve_args(prefill_args)], + [k3_models["target"], "--port", str(decode_port), *_serve_args(decode_args)], + ] + # Use the normal P/D transfer protocol directly so missing transfer metadata + # cannot silently fall back to local prefill and make this test pass. + with RemotePDServer(servers): + prefill_url = f"http://127.0.0.1:{prefill_port}/v1/completions" + decode_url = f"http://127.0.0.1:{decode_port}/v1/completions" + for length in (1, 129, 769): + prompt = _prompt(length, salt=length)["prompt_token_ids"] + prefill = _completion( + prefill_url, + prompt, + max_tokens=1, + kv_transfer_params={"do_remote_decode": True, "do_remote_prefill": False}, + ) + transfer = prefill["kv_transfer_params"] + assert transfer["do_remote_prefill"] + assert any(transfer["remote_block_ids"]) + decoded = _completion(decode_url, prompt, kv_transfer_params=transfer) + if length > 1: + assert decoded["usage"]["prompt_tokens_details"]["cached_tokens"] > 0 + # A D worker can also receive a request without remote KV. Its MLA + # prefill weights must remain usable (the previous P/D fallback bug). + _completion(decode_url, _prompt(513, salt=911)["prompt_token_ids"]) + counts = _draft_counts(f"http://127.0.0.1:{decode_port}/metrics") + assert counts and sum(counts.values()) > 0 diff --git a/tests/ut/ops/test_prepare_finalize.py b/tests/ut/ops/test_prepare_finalize.py index eb4afa465527..54548ae309c8 100644 --- a/tests/ut/ops/test_prepare_finalize.py +++ b/tests/ut/ops/test_prepare_finalize.py @@ -3,6 +3,7 @@ import torch from vllm.model_executor.layers.fused_moe import FusedMoEConfig +from vllm.model_executor.models.utils import sequence_parallel_chunk_impl from vllm_ascend.ops.fused_moe.prepare_finalize import ( PrepareAndFinalizeWithAll2All, @@ -60,6 +61,54 @@ def test_mc2_prepare_finalize(self, mock_get_forward_context, mock_tp_rank, mock result = layer.finalize(h_out, reduce_results=False, padded_hidden_states_shape=padded_hidden_states_shape) self.assertEqual(result.shape[0], 3) + @patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_tensor_model_parallel_world_size", return_value=4) + @patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_tensor_model_parallel_rank") + @patch("vllm_ascend.ascend_forward_context.get_forward_context") + def test_mc2_sp_preserves_local_mask_and_unpads(self, mock_context, mock_tp_rank, mock_tp_size): + # DP peers can have different local SP lengths. Valid bits follow the + # local TP shard, not the larger DP-wide communication stride. + for num_tokens, padded_num_tokens in ((3, 8), (7, 8), (8, 8), (9, 16)): + shard_size = (num_tokens + 3) // 4 + hidden = torch.arange(shard_size * 4 * 8, dtype=torch.float32).reshape(-1, 8) + context = MagicMock() + context.mc2_mask = torch.arange(padded_num_tokens) < num_tokens + context.padded_num_tokens = padded_num_tokens + mock_context.return_value = context + for rank in range(4): + with self.subTest(num_tokens=num_tokens, rank=rank): + mock_tp_rank.return_value = rank + layer = PrepareAndFinalizeWithMC2(self.moe_config) + local = hidden[rank * shard_size : (rank + 1) * shard_size] + logits = local[:, :2].clone() + prepared = layer.prepare(local, logits, replace_allreduce=True) + expected_mask = torch.zeros(padded_num_tokens // 4, dtype=torch.bool) + expected_mask[:shard_size] = torch.arange(rank * shard_size, (rank + 1) * shard_size) < num_tokens + torch.testing.assert_close(prepared.mc2_mask, expected_mask) + torch.testing.assert_close(prepared.hidden_states[:shard_size], local) + torch.testing.assert_close(prepared.router_logits[:shard_size], logits) + self.assertEqual(prepared.hidden_states.shape[0], len(expected_mask)) + self.assertEqual(prepared.router_logits.shape[0], len(expected_mask)) + input_ids = torch.arange(rank * shard_size, (rank + 1) * shard_size) + prepared_ids = layer.pad_and_split_input_ids(input_ids) + torch.testing.assert_close(prepared_ids[:shard_size], input_ids) + self.assertEqual(len(prepared_ids), len(expected_mask)) + full_ids = torch.arange(num_tokens) + 1 + local_ids = torch.nn.functional.pad(full_ids, (0, shard_size * 4 - num_tokens)).chunk(4)[rank] + with ( + patch("vllm.model_executor.models.utils.get_tensor_model_parallel_world_size", return_value=4), + patch("vllm.model_executor.models.utils.get_tensor_model_parallel_rank", return_value=rank), + # Execute the upstream implementation without NPU-only + # custom-op dispatch in this CPU unit test. + patch( + "vllm_ascend.ops.fused_moe.prepare_finalize.sequence_parallel_chunk", + side_effect=sequence_parallel_chunk_impl, + ), + ): + prepared_ids = layer.pad_and_split_input_ids(full_ids) + torch.testing.assert_close(prepared_ids[:shard_size], local_ids) + self.assertEqual(len(prepared_ids), len(expected_mask)) + torch.testing.assert_close(layer.finalize(prepared.hidden_states, reduce_results=False), local) + @patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_tensor_model_parallel_world_size", return_value=2) @patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_tensor_model_parallel_rank", return_value=0) @patch("vllm_ascend.ascend_forward_context.get_forward_context") diff --git a/vllm_ascend/models/kimi_k3.py b/vllm_ascend/models/kimi_k3.py index 727b503188ba..d1eca77215a0 100644 --- a/vllm_ascend/models/kimi_k3.py +++ b/vllm_ascend/models/kimi_k3.py @@ -467,6 +467,14 @@ def __init__( if self.use_sequence_parallel: self.self_attn.o_proj.reduce_results = False + def _run_self_attn( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + ) -> torch.Tensor: + # Ascend attention returns its output instead of filling an AMD buffer. + return self.self_attn(positions=positions, hidden_states=hidden_states) + def forward_attn_residual( self, positions: torch.Tensor, diff --git a/vllm_ascend/ops/fused_moe/prepare_finalize.py b/vllm_ascend/ops/fused_moe/prepare_finalize.py index 1f1c7bdddd19..f72fcaf2d1e9 100644 --- a/vllm_ascend/ops/fused_moe/prepare_finalize.py +++ b/vllm_ascend/ops/fused_moe/prepare_finalize.py @@ -27,6 +27,7 @@ get_tensor_model_parallel_world_size, ) from vllm.model_executor.layers.fused_moe import FusedMoEConfig +from vllm.model_executor.models.utils import sequence_parallel_chunk from vllm_ascend.ascend_forward_context import _EXTRA_CTX from vllm_ascend.lora.fused_moe import prepare_lora_indices @@ -228,8 +229,8 @@ def finalize( class PrepareAndFinalizeWithMC2(PrepareAndFinalizeWithAll2All): """ - MoE communication strategy using MC2, which is based on All2All. Hence, it inherits - All2All and share the same finalize method. + MoE communication strategy using MC2, based on All2All with additional + DP-wide padding and unpadding for sequence-parallel inputs. Designed for Ascend or environments requiring explicit padding and slicing control. Relies on `mc2_mask` and `padded_num_tokens` from forward_context for alignment. """ @@ -261,14 +262,26 @@ def prepare( 3. If TP > 1, split tensors along token dimension and select current TP rank's slice. 4. Split and return corresponding `mc2_mask`. - Skips padding/slicing if `replace_allreduce` is True. + With `replace_allreduce`, inputs are already TP-sharded. Pad only the + local shard to the DP-wide MC2 length, preserving its original mask. Returns: MoEPrepareOutput, possibly sliced/padded. """ self.replace_allreduce = replace_allreduce mc2_mask = _EXTRA_CTX.mc2_mask - if self.tp_size > 1: + if self.replace_allreduce: + # SP shards use the local token count, not the largest DP batch. + # Select valid bits before adding padding for uniform MC2 batches. + self.num_tokens = hidden_states.shape[0] + start = self.tp_rank * self.num_tokens + mc2_mask = mc2_mask[start : start + self.num_tokens] + pad_size = _EXTRA_CTX.padded_num_tokens // self.tp_size - self.num_tokens + if pad_size > 0: + hidden_states = nn.functional.pad(hidden_states, (0, 0, 0, pad_size)) + router_logits = nn.functional.pad(router_logits, (0, 0, 0, pad_size)) + mc2_mask = nn.functional.pad(mc2_mask, (0, pad_size), value=False) + elif self.tp_size > 1: # Also slice mc2_mask split_mc2_mask = torch.tensor_split(mc2_mask, self.tp_size, dim=0) mc2_mask = split_mc2_mask[self.tp_rank] @@ -299,11 +312,30 @@ def prepare( pertoken_scale=None, ) + def finalize( + self, + hidden_states: torch.Tensor, + reduce_results: bool, + padded_hidden_states_shape: torch.Size | None = None, + ) -> torch.Tensor: + if self.replace_allreduce: + # Return the original SP shard to the residual/shared-expert path. + return hidden_states[: self.num_tokens] + return super().finalize(hidden_states, reduce_results, padded_hidden_states_shape) + def pad_and_split_input_ids( self, input_ids, ): - if not self.replace_allreduce: + if self.replace_allreduce: + # MoE-only SP retains full token IDs, while model-level SP may + # already shard them. Align to the local hidden states first. + if input_ids.numel() != self.num_tokens: + input_ids = sequence_parallel_chunk(input_ids.reshape(-1, 1)).reshape(-1) + pad_size = _EXTRA_CTX.padded_num_tokens // self.tp_size - self.num_tokens + if pad_size > 0: + input_ids = nn.functional.pad(input_ids, (0, pad_size)) + else: target_pad_length = _EXTRA_CTX.padded_num_tokens pad_size = target_pad_length - self.num_tokens if pad_size > 0: diff --git a/vllm_ascend/patch/platform/patch_speculative_config.py b/vllm_ascend/patch/platform/patch_speculative_config.py index 269607a336de..0e35e889c3e2 100644 --- a/vllm_ascend/patch/platform/patch_speculative_config.py +++ b/vllm_ascend/patch/platform/patch_speculative_config.py @@ -1,10 +1,27 @@ -from transformers import PretrainedConfig +from transformers import DeepseekV2Config, PretrainedConfig from vllm.config.speculative import SpeculativeConfig _orig_post_init = SpeculativeConfig.__post_init__ _orig_hf_config_override = SpeculativeConfig.hf_config_override +# Transformers 5.14 inherited a hidden_size % num_heads check from Llama in +# DeepseekV2Config. K3 MLA has independent projection/head dimensions (e.g. +# hidden_size=7168, num_heads=96), so that MHA constraint does not apply. +# strict stores unbound validators; patch that entry, not all config validation. +if hasattr(DeepseekV2Config, "__class_validators__"): + _orig_validate_architecture = DeepseekV2Config.validate_architecture + + def _validate_dspark_architecture(config): + if config.model_type != "k3_dspark": + _orig_validate_architecture(config) + + DeepseekV2Config.__class_validators__ = [ + _validate_dspark_architecture if validator is _orig_validate_architecture else validator + for validator in DeepseekV2Config.__class_validators__ + ] + + def _normalize_legacy_qwen3_dspark_config(hf_config: PretrainedConfig) -> PretrainedConfig: hf_config = _orig_hf_config_override(hf_config) architectures = hf_config.architectures or () diff --git a/vllm_ascend/spec_decode/llm_base_proposer.py b/vllm_ascend/spec_decode/llm_base_proposer.py index d6bfd42f7131..cb95895b5f9b 100644 --- a/vllm_ascend/spec_decode/llm_base_proposer.py +++ b/vllm_ascend/spec_decode/llm_base_proposer.py @@ -326,6 +326,15 @@ def load_model(self, model: nn.Module) -> None: with self.maybe_eager_context: self.model = self._get_model() + if self.supports_mm_inputs: + # Match upstream: a multimodal target can use a text-only drafter. + try: + dummy_input_ids = torch.tensor([[1]], device=self.input_ids.device) + self.model.embed_input_ids(dummy_input_ids, multimodal_embeddings=None) + except (NotImplementedError, AttributeError, TypeError): + logger.warning("Draft model does not support multimodal inputs, falling back to text-only mode") + self.supports_mm_inputs: bool = False + # Find draft layers (attention layers added by draft model) all_attn_layers = get_layers_from_vllm_config( self.vllm_config, @@ -377,7 +386,7 @@ def load_model(self, model: nn.Module) -> None: # share embed_tokens with the target model if needed self._maybe_share_embeddings(target_language_model) self._maybe_share_topk_indices(target_language_model) - self._maybe_share_lm_head(model) + self._maybe_share_lm_head(target_language_model) if ( self.parallel_drafting From b6dce8cb2beb269604d4c5af5abae3aa29337877 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Fri, 28 Aug 2026 02:35:39 +0800 Subject: [PATCH 49/50] test(kimi-k3): streamline A3 smoke coverage Use four complementary deployment cases: Block5 MLA at TP16, W4A8 GQA at DP2/TP8, legacy MLA at P8/D8, and MTP with an image. Keep production widths and 16 experts in a six-layer target with a reduced attention-residual block. Share local model fixtures and omit redundant fused-expert quantization descriptions. Initialize the text-only multimodal capability in the bare proposer UT fixture to match the base constructor contract. No production runtime changes. Validated all four cases on one A3-16 with vLLM 0.27.1 in 357.92 seconds, 23 proposer unit tests, targeted mypy, CI routing/coverage, and format.sh ci. Set the suite estimate to 360 seconds. Dummy smoke is not checkpoint accuracy or QuaRot validation. Signed-off-by: maoxx241 --- .github/workflows/scripts/test_config.yaml | 2 +- .../pull_request/sixteen_card/test_kimi_k3.py | 66 ++++++++++--------- .../ut/spec_decode/test_llm_base_proposer.py | 1 + 3 files changed, 36 insertions(+), 33 deletions(-) diff --git a/.github/workflows/scripts/test_config.yaml b/.github/workflows/scripts/test_config.yaml index f83f3984e8fd..35923912f6f6 100644 --- a/.github/workflows/scripts/test_config.yaml +++ b/.github/workflows/scripts/test_config.yaml @@ -918,7 +918,7 @@ estimated_times: tests/e2e/pull_request/four_card/test_qwen3_32b_bf16_ll_performance.py: 1200 tests/e2e/pull_request/two_card/test_gemma4.py: 420 tests/e2e/pull_request/eight_card/test_minimax_m3.py: 630 - tests/e2e/pull_request/sixteen_card/test_kimi_k3.py: 1200 + tests/e2e/pull_request/sixteen_card/test_kimi_k3.py: 360 tests/ut/attention/a2/test_attention_cp.py: 30 tests/ut/attention/a2/test_attention_cp_precision.py: 30 tests/ut/attention/a2/test_attention_v1.py: 30 diff --git a/tests/e2e/pull_request/sixteen_card/test_kimi_k3.py b/tests/e2e/pull_request/sixteen_card/test_kimi_k3.py index 45a9d9a6eb9f..2e2020d98385 100644 --- a/tests/e2e/pull_request/sixteen_card/test_kimi_k3.py +++ b/tests/e2e/pull_request/sixteen_card/test_kimi_k3.py @@ -3,7 +3,8 @@ """Single-node A3 (16 logical NPUs) K3 functional tests, not accuracy tests. Build local configs and initialize dummy weights; no full checkpoint is needed. -Keep production tensor widths while reducing the target to 16 layers/experts. +Keep production tensor widths with a six-layer, 16-expert target. Cover Block5 +at TP16, quantized GQA at DP2/TP8, legacy MLA at P8/D8, and MTP with an image. Random weights cannot validate checkpoint loading, QuaRot or acceptance rates. """ @@ -25,15 +26,16 @@ from tests.e2e.conftest import RemoteOpenAIServer, RemotePDServer, VllmRunner -NUM_LAYERS = 16 +NUM_LAYERS = 6 NUM_EXPERTS = 16 NUM_VISION_LAYERS = 2 VOCAB_SIZE = 163840 MAX_MODEL_LEN = 2048 OUTPUT_TOKENS = 2 -# Keep the final consecutive MLA layers as well as the regular KDA/MLA pattern. -# Layer 13 also crosses the production attention-residual block boundary (12). -FULL_ATTN_LAYERS = (4, 8, 12, 15, 16) +# Keep both KDA and MLA, including the final consecutive MLA layers. A smaller +# attention-residual block retains cross-block execution in the reduced model. +FULL_ATTN_LAYERS = (2, 5, 6) +ATTN_RES_BLOCK_SIZE = 4 def _text_config() -> dict: @@ -61,7 +63,7 @@ def _text_config() -> dict: "topk_group": 1, "topk_method": "noaux_tc", "moe_renormalize": True, - "attn_res_block_size": 12, + "attn_res_block_size": ATTN_RES_BLOCK_SIZE, "num_attention_heads": 96, "num_key_value_heads": 96, "q_lora_rank": 1536, @@ -124,7 +126,7 @@ def _draft_config(variant: str) -> dict: num_target_layers=NUM_LAYERS, # The drafter consumes intermediate layer outputs, before the final # target layer. Remap all five taps to the reduced target. - dflash_config={"mask_token_id": 163824, "target_layer_ids": [1, 3, 7, 11, 14]}, + dflash_config={"mask_token_id": 163824, "target_layer_ids": [0, 1, 2, 3, 4]}, rope_parameters={ "rope_type": "yarn", "factor": 16.0, @@ -145,7 +147,7 @@ def _draft_config(variant: str) -> dict: target_hidden_size=7168, target_num_hidden_layers=NUM_LAYERS, num_target_layers=5, - target_layer_ids=[1, 3, 7, 11, 14], + target_layer_ids=[0, 1, 2, 3, 4], mask_token_id=163837, rope_parameters={ "rope_type": "yarn", @@ -166,7 +168,7 @@ def _draft_config(variant: str) -> dict: head_dim=64, qk_head_dim=192, num_target_layers=3, - target_layer_ids=[7, 11, 14], + target_layer_ids=[1, 3, 4], layer_types=["full_attention"] * 3, mask_token_id=163839, markov_rank=512, @@ -275,9 +277,11 @@ def _write_w4a8_description(path: Path) -> None: if layer: for projection in ("routed_expert_down_proj", "routed_expert_up_proj"): description[f"{prefix}.block_sparse_moe.{projection}.weight"] = "W8A8_DYNAMIC" - for expert in range(NUM_EXPERTS): - for projection in ("w1", "w2", "w3"): - description[f"{prefix}.block_sparse_moe.experts.{expert}.{projection}.weight"] = "W4A8_DYNAMIC" + # ModelSlim selects the fused expert group's scheme from expert 0. + # All 16 experts still run; duplicating their identical descriptions + # only bloats the config sent to each spawned worker. + for projection in ("w1", "w2", "w3"): + description[f"{prefix}.block_sparse_moe.experts.0.{projection}.weight"] = "W4A8_DYNAMIC" description = {f"language_model.{key}" if key != "optional" else key: value for key, value in description.items()} for layer in range(NUM_VISION_LAYERS): for projection in ("mlp.fc0", "mlp.fc1", "wqkv", "wo"): @@ -287,11 +291,13 @@ def _write_w4a8_description(path: Path) -> None: (path / "quant_model_description.json").write_text(json.dumps(description), encoding="utf-8") -@pytest.fixture -def k3_models(tmp_path: Path) -> dict[str, str]: +@pytest.fixture(scope="module") +def k3_models(tmp_path_factory: pytest.TempPathFactory) -> dict[str, str]: + tmp_path = tmp_path_factory.mktemp("k3-dummy") models = {"target": _write_target(tmp_path / "target")} models["w4a8"] = _write_target(tmp_path / "w4a8") _write_w4a8_description(tmp_path / "w4a8") + models["mtp"] = _write_target(tmp_path / "mtp", mtp=True) for variant in ("mla", "mla_block5", "gqa"): models[variant] = _write_config(tmp_path / variant, _draft_config(variant)) return models @@ -368,14 +374,9 @@ def _generate(llm, prompts: list[dict]): return outputs -@pytest.mark.parametrize("variant", ["mla", "mla_block5", "gqa"]) -def test_k3_dspark_tp16(k3_models: dict[str, str], variant: str) -> None: - args = _engine_args(k3_models, variant) - target = k3_models["target"] - if variant == "gqa": - args["quantization"] = "ascend" - target = k3_models["w4a8"] - with VllmRunner(target, **args) as runner: +def test_k3_mla_block5_tp16(k3_models: dict[str, str]) -> None: + args = _engine_args(k3_models, "mla_block5") + with VllmRunner(k3_models["target"], **args) as runner: llm = runner.model _generate(llm, [_prompt(1)]) # Kernel-block and block+1 prefill, mixed lengths, and chunked prefill. @@ -450,20 +451,22 @@ def _draft_counts(url: str) -> dict[str, float]: } -def test_k3_gqa_dp2_tp8(k3_models: dict[str, str]) -> None: +def test_k3_gqa_w4a8_dp2_tp8(k3_models: dict[str, str]) -> None: args = _engine_args(k3_models, "gqa", tp=8) args["data_parallel_size"] = 2 + args["quantization"] = "ascend" args["additional_config"]["enable_shared_expert_dp"] = True port = get_open_port() with RemoteOpenAIServer( - k3_models["target"], - [*_serve_args(args), "--port", str(port)], + k3_models["w4a8"], + [*_serve_args(args), "--port", str(port), "--api-server-count", "1"], server_host="127.0.0.1", server_port=port, auto_port=False, max_wait_seconds=600, ) as server: - prompts = [_prompt(length, salt=i * 137)["prompt_token_ids"] for i, length in enumerate((1, 129, 769, 1153))] + lengths = (1, 129, 769, MAX_MODEL_LEN - OUTPUT_TOKENS) + prompts = [_prompt(length, salt=i * 137)["prompt_token_ids"] for i, length in enumerate(lengths)] with ThreadPoolExecutor(max_workers=4) as pool: outputs = list(pool.map(lambda prompt: _completion(server.url_for("v1", "completions"), prompt), prompts)) assert len(outputs) == len(prompts) @@ -471,8 +474,7 @@ def test_k3_gqa_dp2_tp8(k3_models: dict[str, str]) -> None: assert len(counts) == 2 and all(count > 0 for count in counts.values()), counts -def test_k3_mtp_image_tp16(k3_models: dict[str, str], tmp_path: Path) -> None: - target = _write_target(tmp_path / "mtp", mtp=True) +def test_k3_mtp_image_tp16(k3_models: dict[str, str]) -> None: args = _engine_args(k3_models, "gqa") args["speculative_config"] = { "method": "mtp", @@ -481,7 +483,7 @@ def test_k3_mtp_image_tp16(k3_models: dict[str, str], tmp_path: Path) -> None: "draft_load_config": {"load_format": "dummy"}, } args["limit_mm_per_prompt"] = {"image": 1} - with VllmRunner(target, **args) as runner: + with VllmRunner(k3_models["mtp"], **args) as runner: _generate(runner.model, [_prompt(1)]) image_prompt = { "prompt": "<|kimi_image_placeholder|> describe", @@ -492,7 +494,7 @@ def test_k3_mtp_image_tp16(k3_models: dict[str, str], tmp_path: Path) -> None: assert drafts and sum(m.value for m in drafts) > 0, "Requests bypassed MTP" -def test_k3_gqa_pd_tp8(k3_models: dict[str, str]) -> None: +def test_k3_mla_pd_tp8(k3_models: dict[str, str]) -> None: prefill_port, decode_port = get_open_port(), get_open_port() transfer_config = { "kv_connector": "MooncakeConnectorV1", @@ -501,13 +503,13 @@ def test_k3_gqa_pd_tp8(k3_models: dict[str, str]) -> None: "decode": {"dp_size": 1, "tp_size": 8}, }, } - prefill_args = _engine_args(k3_models, "gqa", tp=8) + prefill_args = _engine_args(k3_models, "mla", tp=8) # P also builds the draft KV that D consumes; both peers need the same # target/draft layer layout even though P only generates one token. prefill_args.pop("compilation_config") prefill_args["enforce_eager"] = True prefill_args["kv_transfer_config"] = dict(transfer_config, kv_role="kv_producer", kv_port=get_open_port()) - decode_args = _engine_args(k3_models, "gqa", tp=8) + decode_args = _engine_args(k3_models, "mla", tp=8) decode_args["kv_transfer_config"] = dict(transfer_config, kv_role="kv_consumer", kv_port=get_open_port()) servers = [ [k3_models["target"], "--port", str(prefill_port), *_serve_args(prefill_args)], diff --git a/tests/ut/spec_decode/test_llm_base_proposer.py b/tests/ut/spec_decode/test_llm_base_proposer.py index 2c8be1bfcd38..8f54384935ce 100644 --- a/tests/ut/spec_decode/test_llm_base_proposer.py +++ b/tests/ut/spec_decode/test_llm_base_proposer.py @@ -116,6 +116,7 @@ def test_load_model_reads_validated_draft_window_size(): proposer.runner = SimpleNamespace(max_num_reqs=8) proposer.device = "cpu" proposer.parallel_drafting = False + proposer.supports_mm_inputs = False proposer._maybe_share_embeddings = MagicMock() proposer._maybe_share_topk_indices = MagicMock() proposer._maybe_share_lm_head = MagicMock() From ba14b852dedb0fbeb1bdec0124c9eb6db6ba7a96 Mon Sep 17 00:00:00 2001 From: maoxx241 Date: Fri, 28 Aug 2026 02:59:13 +0800 Subject: [PATCH 50/50] test(kimi-k3): prune redundant unit test scaffolding Remove source-text assertions, duplicated constructor and forwarding checks, and redundant mock scaffolding from the Kimi K3 unit tests. Consolidate rotation loading and Mamba copy checks around actual tensor results while retaining cache-capacity, P/D transfer, accepted-token and operator regressions. Signed-off-by: maoxx241 --- tests/ut/attention/a2/test_mla_v1.py | 67 +-- .../ut/kv_offload/test_mooncake_connector.py | 165 ++---- tests/ut/model_executor/test_qwen3_dspark.py | 108 +--- tests/ut/models/test_kimi_k3_adapter.py | 528 ++---------------- tests/ut/ops/test_fused_moe.py | 28 - tests/ut/ops/test_gdn_attn_builder.py | 34 -- tests/ut/ops/test_kimi_kda.py | 31 - tests/ut/ops/test_layernorm.py | 18 +- .../platform/test_prefix_cache_cp_patches.py | 86 +-- .../ut/patch/worker/test_patch_mamba_utils.py | 78 +-- .../worker/test_patch_mamba_utils_source.py | 61 -- .../ut/quantization/test_modelslim_config.py | 58 -- tests/ut/spec_decode/test_dspark_proposer.py | 223 +------- tests/ut/test_utils.py | 12 +- tests/ut/worker/a2/test_model_runner_v1.py | 59 -- 15 files changed, 197 insertions(+), 1359 deletions(-) delete mode 100644 tests/ut/patch/worker/test_patch_mamba_utils_source.py diff --git a/tests/ut/attention/a2/test_mla_v1.py b/tests/ut/attention/a2/test_mla_v1.py index 0c9e33bf42d5..6dedcd9efec5 100644 --- a/tests/ut/attention/a2/test_mla_v1.py +++ b/tests/ut/attention/a2/test_mla_v1.py @@ -479,69 +479,30 @@ def test_ascend_mla_metadata_builder_default(self): self.assertEqual(builder.block_size, mock_vllm_config.cache_config.block_size) self.assertEqual(builder.chunked_prefill_enabled, mock_vllm_config.scheduler_config.enable_chunked_prefill) - def test_metadata_builder_uses_draft_layer_rope_mode(self): + def test_metadata_builder_uses_layer_rope_mode(self): mock_vllm_config = MagicMock() mock_vllm_config.model_config.max_model_len = 1024 mock_vllm_config.model_config.get_head_size.return_value = 64 mock_vllm_config.model_config.dtype = torch.float16 - mock_vllm_config.model_config.hf_text_config = SimpleNamespace( - qk_rope_head_dim=64, - mla_use_nope=True, - ) - mock_vllm_config.cache_config.block_size = 16 - mock_vllm_config.scheduler_config.max_num_seqs = 4 - mock_vllm_config.scheduler_config.enable_chunked_prefill = False - mock_vllm_config.speculative_config = None - mock_vllm_config.compilation_config.static_forward_context = { - "draft.self_attn": SimpleNamespace( - impl=SimpleNamespace(use_mla_rope=True), - ), - } - - with patch( - "vllm_ascend.attention.mla_v1.get_ascend_config", - return_value=MagicMock(), - ): - builder = AscendMLAMetadataBuilder( - None, - ["draft.self_attn"], - mock_vllm_config, - "cpu", - ) - - self.assertTrue(builder.use_mla_rope) - - def test_metadata_builder_uses_target_layer_nope_mode(self): - mock_vllm_config = MagicMock() - mock_vllm_config.model_config.max_model_len = 1024 - mock_vllm_config.model_config.get_head_size.return_value = 64 - mock_vllm_config.model_config.dtype = torch.float16 - mock_vllm_config.model_config.hf_text_config = SimpleNamespace( - qk_rope_head_dim=64, - mla_use_nope=False, - ) mock_vllm_config.cache_config.block_size = 16 mock_vllm_config.scheduler_config.max_num_seqs = 4 mock_vllm_config.scheduler_config.enable_chunked_prefill = False mock_vllm_config.speculative_config = None - mock_vllm_config.compilation_config.static_forward_context = { - "target.self_attn": SimpleNamespace( - impl=SimpleNamespace(use_mla_rope=False), - ), - } - with patch( - "vllm_ascend.attention.mla_v1.get_ascend_config", - return_value=MagicMock(), - ): - builder = AscendMLAMetadataBuilder( - None, - ["target.self_attn"], - mock_vllm_config, - "cpu", - ) + for layer_uses_rope in (True, False): + with self.subTest(layer_uses_rope=layer_uses_rope): + # Deliberately disagree with the target config: the layer wins. + mock_vllm_config.model_config.hf_text_config = SimpleNamespace( + qk_rope_head_dim=64, + mla_use_nope=layer_uses_rope, + ) + mock_vllm_config.compilation_config.static_forward_context = { + "self_attn": SimpleNamespace(impl=SimpleNamespace(use_mla_rope=layer_uses_rope)), + } + with patch("vllm_ascend.attention.mla_v1.get_ascend_config", return_value=MagicMock()): + builder = AscendMLAMetadataBuilder(None, ["self_attn"], mock_vllm_config, "cpu") - self.assertFalse(builder.use_mla_rope) + self.assertEqual(builder.use_mla_rope, layer_uses_rope) def test_ascend_mla_metadata_builder_spec_decode(self): mock_vllm_config = MagicMock() diff --git a/tests/ut/kv_offload/test_mooncake_connector.py b/tests/ut/kv_offload/test_mooncake_connector.py index 2fa3774f43ef..d60b6aa5a899 100644 --- a/tests/ut/kv_offload/test_mooncake_connector.py +++ b/tests/ut/kv_offload/test_mooncake_connector.py @@ -7,7 +7,7 @@ import types import unittest from collections import OrderedDict, defaultdict, deque -from typing import Any, TypedDict, cast +from typing import Any, cast from unittest.mock import MagicMock, patch import msgspec @@ -95,19 +95,6 @@ DONE_RECVING_MSG = b"done_recving_msg" -class KimiMambaTransferCase(TypedDict): - decode_tp: int - pulls: int - conv_shape: list[int] - ssm_shape: list[int] - local_conv_len: int - local_ssm_len: int - remote_conv_stride: int - remote_ssm_stride: int - remote_tp_offset: int - expected_segment: int - - def make_mock_kv_caches() -> dict[str, Any]: kv_cache = MagicMock(device=torch.device("npu:0")) return {"layer_0": (kv_cache, kv_cache)} @@ -1358,114 +1345,52 @@ def test_append_mamba_transfer_meta_uses_block_stride_for_block_offsets(self): return_value=False, ) def test_append_mamba_transfer_meta_kimi_k3_sd_unequal_tp(self, _mock_layout): - """Cover K3 conv/SSM address slicing when prefill and decode TP differ.""" - cases: list[KimiMambaTransferCase] = [ - { - "decode_tp": 8, - "pulls": 2, - "conv_shape": [3, 4608], - "ssm_shape": [12, 128, 128], - "local_conv_len": 27648, - "local_ssm_len": 786432, - "remote_conv_stride": 15000, - "remote_ssm_stride": 400000, - "remote_tp_offset": 1, - "expected_segment": 768, - }, - { - "decode_tp": 8, - "pulls": 4, - "conv_shape": [3, 4608], - "ssm_shape": [12, 128, 128], - "local_conv_len": 27648, - "local_ssm_len": 786432, - "remote_conv_stride": 6912, - "remote_ssm_stride": 196608, - "remote_tp_offset": 3, - "expected_segment": 384, - }, - { - "decode_tp": 16, - "pulls": 2, - "conv_shape": [3, 2304], - "ssm_shape": [6, 128, 128], - "local_conv_len": 13824, - "local_ssm_len": 393216, - "remote_conv_stride": 6912, - "remote_ssm_stride": 196608, - "remote_tp_offset": 1, - "expected_segment": 384, - }, - ] - - for case in cases: - with self.subTest(case=case): - self.thread.tp_size = case["decode_tp"] - self.thread.vllm_config.model_config.hf_text_config = types.SimpleNamespace( - linear_attn_config={"num_heads": 96, "head_dim": 128} - ) - src_list: list[int] = [] - dst_list: list[int] = [] - length_list: list[int] = [] - local_bases = [0x100000, 0x300000] - remote_bases = [0x200000, 0x400000] - local_strides = [case["local_conv_len"] + 4096, case["local_ssm_len"] + 8192] - local_block_id = 2 - remote_block_id = 3 - - self.thread._append_mamba_transfer_meta( - src_list, - dst_list, - length_list, - group_spec={ - "kv_cache_spec_type": "MambaSpec", - "shapes": [case["conv_shape"], case["ssm_shape"]], - "dtype_sizes": [2, 4], - }, - src_layer_base_addr=local_bases, - dst_layer_base_addr=remote_bases, - block_len=[case["local_conv_len"], case["local_ssm_len"]], - block_stride=local_strides, - remote_block_stride=[case["remote_conv_stride"], case["remote_ssm_stride"]], - remote_block_id=remote_block_id, - local_block_id=local_block_id, - tp_num_need_pulls=case["pulls"], - remote_tp_offset=case["remote_tp_offset"], - ) + """P16 -> D8 slices Q/K/V per state row and honors padded page strides.""" + self.thread.tp_size = 8 + self.thread.vllm_config.model_config.hf_text_config = types.SimpleNamespace( + linear_attn_config={"num_heads": 96, "head_dim": 128} + ) + src_list: list[int] = [] + dst_list: list[int] = [] + length_list: list[int] = [] + local_strides = [27648 + 4096, 786432 + 8192] + remote_strides = [15000, 400000] - remote_segment = case["expected_segment"] - local_segment = remote_segment * case["pulls"] - local_conv_base = local_bases[0] + local_block_id * local_strides[0] - remote_conv_base = remote_bases[0] + remote_block_id * case["remote_conv_stride"] - expected_src: list[int] = [] - expected_dst: list[int] = [] - for state_idx in range(3): - for segment_idx in range(3): - expected_src.append( - local_conv_base - + ( - state_idx * case["conv_shape"][1] - + segment_idx * local_segment - + case["remote_tp_offset"] * remote_segment - ) - * 2 - ) - expected_dst.append( - remote_conv_base + (state_idx * remote_segment * 3 + segment_idx * remote_segment) * 2 - ) - expected_src.append( - local_bases[1] - + local_block_id * local_strides[1] - + case["remote_tp_offset"] * case["local_ssm_len"] // case["pulls"] - ) - expected_dst.append(remote_bases[1] + remote_block_id * case["remote_ssm_stride"]) + self.thread._append_mamba_transfer_meta( + src_list, + dst_list, + length_list, + group_spec={ + "kv_cache_spec_type": "MambaSpec", + "shapes": [[3, 4608], [12, 128, 128]], + "dtype_sizes": [2, 4], + }, + src_layer_base_addr=[0x100000, 0x300000], + dst_layer_base_addr=[0x200000, 0x400000], + block_len=[27648, 786432], + block_stride=local_strides, + remote_block_stride=remote_strides, + local_block_id=2, + remote_block_id=3, + tp_num_need_pulls=2, + remote_tp_offset=1, + ) - self.assertEqual(src_list, expected_src) - self.assertEqual(dst_list, expected_dst) - self.assertEqual( - length_list, - [remote_segment * 2] * 9 + [case["local_ssm_len"] // case["pulls"]], - ) + # The second P shard occupies the upper 768 elements of each 1536-wide + # Q/K/V segment in each D row. The remote rows have no inter-segment gaps. + local_offsets = [1536, 4608, 7680, 10752, 13824, 16896, 19968, 23040, 26112] + remote_offsets = [0, 1536, 3072, 4608, 6144, 7680, 9216, 10752, 12288] + self.assertEqual( + src_list, + [0x100000 + 2 * local_strides[0] + offset for offset in local_offsets] + + [0x300000 + 2 * local_strides[1] + 393216], + ) + self.assertEqual( + dst_list, + [0x200000 + 3 * remote_strides[0] + offset for offset in remote_offsets] + + [0x400000 + 3 * remote_strides[1]], + ) + self.assertEqual(length_list, [1536] * 9 + [393216]) @patch( "vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.is_conv_state_dim_first", diff --git a/tests/ut/model_executor/test_qwen3_dspark.py b/tests/ut/model_executor/test_qwen3_dspark.py index 02c0f9bc31b8..e14f66f1c43f 100644 --- a/tests/ut/model_executor/test_qwen3_dspark.py +++ b/tests/ut/model_executor/test_qwen3_dspark.py @@ -28,7 +28,6 @@ from torch import nn import vllm_ascend.models.qwen3_dspark as qwen3_dspark -from vllm_ascend.models.llama_eagle3 import load_quarot_target_layer class TestQwen3DSparkWeightLoading: @@ -74,84 +73,39 @@ def test_rotates_only_fc_weights(self) -> None: torch.testing.assert_close(processed_weights[1][1], non_fc_weight) torch.testing.assert_close(processed_weights[2][1], non_fc_weight) - def test_quarot_loads_missing_boundaries_in_modeling(self) -> None: - model_cls = qwen3_dspark.AscendQwen3DSparkForCausalLM - model = model_cls.__new__(model_cls) - nn.Module.__init__(model) - model.rotation_path = "quarot.safetensors" - model.target_model_path = "/target" - model.enable_confidence_head = False - model.model = SimpleNamespace(embed_tokens=object()) - model.lm_head = object() - rotation = torch.eye(2) - with ( - patch.object( - qwen3_dspark, - "get_rotation_matrix", - return_value=rotation, - ), - patch.object( - qwen3_dspark, - "load_quarot_target_layer", - ) as load_target_layer, - patch.object( - qwen3_dspark.Qwen3DSparkForCausalLM, - "load_weights", - ), - ): - model.load_weights([("model.fc.weight", torch.eye(2))]) - - assert load_target_layer.call_count == 2 - assert load_target_layer.call_args_list[0].args[:2] == ( - model.model.embed_tokens, - model.target_model_path, - ) - assert load_target_layer.call_args_list[1].args[:2] == ( - model.lm_head, - model.target_model_path, - ) - assert model.has_own_embed_tokens - assert model.has_own_lm_head - - -def test_load_quarot_target_layer_reads_local_vocab_shard(tmp_path) -> None: - weight_name = "language_model.model.embed_tokens.weight" +def test_quarot_loads_missing_target_vocab_shards(tmp_path) -> None: + embed_name = "language_model.model.embed_tokens.weight" + head_name = "language_model.lm_head.weight" shard_name = "model-00001-of-00001.safetensors" - target_weight = torch.tensor( - [ - [1.0, 2.0], - [3.0, 4.0], - [5.0, 6.0], - [7.0, 8.0], - ] - ) - save_file({weight_name: target_weight}, tmp_path / shard_name) + target_weight = torch.arange(8, dtype=torch.float32).view(4, 2) + save_file({embed_name: target_weight, head_name: target_weight + 10}, tmp_path / shard_name) (tmp_path / "model.safetensors.index.json").write_text( - json.dumps({"weight_map": {weight_name: shard_name}}), + json.dumps({"weight_map": {embed_name: shard_name, head_name: shard_name}}), encoding="utf-8", ) - - layer = nn.Linear(2, 3, bias=False) - layer.weight.data.fill_(99) - layer.shard_indices = SimpleNamespace( - org_vocab_start_index=1, - org_vocab_end_index=3, - ) - rotation = torch.tensor([[0.0, 1.0], [1.0, 0.0]]) - - load_quarot_target_layer( - layer, - tmp_path, - (weight_name,), - rotation, - "test embedding", - ) - - expected = torch.cat( - ( - target_weight[1:3] @ rotation.T, - torch.zeros(1, 2), - ) - ) - torch.testing.assert_close(layer.weight, expected) + rotation = torch.tensor([[0.0, -1.0], [1.0, 0.0]]) + rotation_path = tmp_path / "rotation.safetensors" + save_file({"global_rotation": rotation}, rotation_path) + + model_cls = qwen3_dspark.AscendQwen3DSparkForCausalLM + model = model_cls.__new__(model_cls) + nn.Module.__init__(model) + model.rotation_path = rotation_path + model.target_model_path = tmp_path + model.enable_confidence_head = False + model.model = SimpleNamespace(embed_tokens=nn.Linear(2, 3, bias=False)) + model.lm_head = nn.Linear(2, 3, bias=False) + for layer in (model.model.embed_tokens, model.lm_head): + layer.weight.data.fill_(99) + layer.shard_indices = SimpleNamespace(org_vocab_start_index=1, org_vocab_end_index=3) + + # Draft omits its vocabulary weights; load and align the target's local TP shard. + with patch.object(qwen3_dspark.Qwen3DSparkForCausalLM, "load_weights"): + model.load_weights(iter([])) + + for layer, weight in ((model.model.embed_tokens, target_weight), (model.lm_head, target_weight + 10)): + expected = torch.cat((weight[1:3] @ rotation.T, torch.zeros(1, 2))) + torch.testing.assert_close(layer.weight, expected) + assert model.has_own_embed_tokens + assert model.has_own_lm_head diff --git a/tests/ut/models/test_kimi_k3_adapter.py b/tests/ut/models/test_kimi_k3_adapter.py index 98697d733437..0627ac012a3d 100644 --- a/tests/ut/models/test_kimi_k3_adapter.py +++ b/tests/ut/models/test_kimi_k3_adapter.py @@ -2,27 +2,20 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from types import MethodType, SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import patch import torch +from safetensors.torch import save_file from torch import nn -from vllm.config import VllmConfig, set_current_vllm_config -from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec from vllm_ascend.models import kimi_k3 from vllm_ascend.models.kimi_k3 import ( - AscendKimiK3ForConditionalGeneration, AscendKimiK3MultiModalProjector, - AscendKimiLinearForCausalLM, AscendKimiLinearModel, - AscendKimiMLAAttention, - AscendKimiMoE, ) from vllm_ascend.models.kimi_k3_dspark import ( - AscendK3DSparkDecoderLayer, AscendK3DSparkForCausalLM, ) -from vllm_ascend.utils import vllm_version_is def test_ascend_attn_res_matches_canonical_k3_math(): @@ -56,98 +49,6 @@ def test_ascend_attn_res_matches_canonical_k3_math(): torch.testing.assert_close(output, expected) -def _make_moe_config(**overrides): - values = { - "hidden_size": 16, - "moe_intermediate_size": 32, - "num_experts": 8, - "num_experts_per_token": 2, - "moe_renormalize": True, - "routed_expert_hidden_size": None, - "latent_moe_use_norm": False, - "routed_scaling_factor": 1.0, - "num_shared_experts": None, - "hidden_act": "silu", - "activation_situ_beta": None, - "activation_situ_linear_beta": None, - "use_grouped_topk": False, - "num_expert_group": None, - "topk_group": None, - "moe_router_activation_func": "softmax", - "rms_norm_eps": 1e-6, - } - values.update(overrides) - return SimpleNamespace(**values) - - -def test_ascend_kimi_moe_uses_standard_runner_dispatch(monkeypatch): - class FakeGate(nn.Module): - def __init__(self, **kwargs): - super().__init__() - self.kwargs = kwargs - - factory = MagicMock(return_value=nn.Identity()) - monkeypatch.setattr(kimi_k3, "GateLinear", FakeGate) - monkeypatch.setattr(kimi_k3, "FusedMoEFactory", factory) - - AscendKimiMoE( - config=_make_moe_config(), - prefix="model.layers.1.block_sparse_moe", - use_sequence_parallel=True, - ) - - assert factory.call_args.kwargs["intermediate_size"] == 32 - assert factory.call_args.kwargs["is_sequence_parallel"] is True - assert "runner_cls" not in factory.call_args.kwargs - - -def test_dspark_decoder_uses_upstream_mlp_activation_contract( - monkeypatch, -): - config = SimpleNamespace( - hidden_size=8, - num_attention_heads=2, - qk_nope_head_dim=2, - qk_rope_head_dim=2, - v_head_dim=2, - q_lora_rank=4, - kv_lora_rank=4, - intermediate_size=16, - hidden_act="silu", - rms_norm_eps=1e-6, - full_attention_causal=True, - ) - vllm_config = SimpleNamespace(cache_config=None) - mlp_factory = MagicMock(return_value=nn.Identity()) - attention_factory = MagicMock(return_value=nn.Identity()) - monkeypatch.setattr( - "vllm_ascend.models.kimi_k3_dspark.get_draft_quant_config", - lambda _: None, - ) - monkeypatch.setattr( - "vllm_ascend.models.kimi_k3_dspark.AscendKimiMLAAttention", - attention_factory, - ) - monkeypatch.setattr( - "vllm_ascend.models.kimi_k3_dspark.KimiMLP", - mlp_factory, - ) - - with set_current_vllm_config(VllmConfig()): - AscendK3DSparkDecoderLayer( - vllm_config=vllm_config, - config=config, - layer_idx=0, - start_layer_id=4, - prefix="model", - ) - - assert mlp_factory.call_args.kwargs["hidden_act"] == "silu" - assert "activation_situ_beta" not in mlp_factory.call_args.kwargs - assert "activation_situ_linear_beta" not in mlp_factory.call_args.kwargs - assert attention_factory.call_args.kwargs["non_causal_multi_token_decode"] is False - - def test_k3_dspark_reports_draft_attention_causality(): model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) nn.Module.__init__(model) @@ -163,51 +64,6 @@ def test_k3_dspark_reports_draft_attention_causality(): assert model.get_draft_attn_causal() == [False, False, False] -def test_ascend_kimi_moe_quantizes_modelslim_latent_projections(monkeypatch): - class FakeGate(nn.Module): - def __init__(self, **kwargs): - super().__init__() - - class FakeLinear(nn.Module): - def __init__(self, input_size, output_size, **kwargs): - super().__init__() - self.input_size = input_size - self.output_size = output_size - self.kwargs = kwargs - - config = _make_moe_config(routed_expert_hidden_size=8) - quant_config = MagicMock() - quant_config.get_name.return_value = "ascend" - factory = MagicMock(return_value=nn.Identity()) - monkeypatch.setattr(kimi_k3, "GateLinear", FakeGate) - monkeypatch.setattr(kimi_k3, "ReplicatedLinear", FakeLinear) - monkeypatch.setattr(kimi_k3, "FusedMoEFactory", factory) - - moe = AscendKimiMoE( - config=config, - quant_config=quant_config, - prefix="model.layers.1.block_sparse_moe", - ) - - assert moe.routed_expert_down_proj.input_size == 16 - assert moe.routed_expert_down_proj.output_size == 8 - assert moe.routed_expert_down_proj.kwargs == { - "bias": False, - "quant_config": quant_config, - "prefix": "model.layers.1.block_sparse_moe.routed_expert_down_proj", - } - assert moe.routed_expert_up_proj.input_size == 8 - assert moe.routed_expert_up_proj.output_size == 16 - assert moe.routed_expert_up_proj.kwargs == { - "bias": False, - "quant_config": quant_config, - "prefix": "model.layers.1.block_sparse_moe.routed_expert_up_proj", - } - assert factory.call_args.kwargs["routed_input_transform"] is moe.routed_expert_down_proj - assert factory.call_args.kwargs["routed_output_transform"] is moe.routed_output_transform - assert "runner_cls" not in factory.call_args.kwargs - - def test_kimi_mixed_kda_gate_weights_use_upstream_packed_loader(monkeypatch): model = AscendKimiLinearModel.__new__(AscendKimiLinearModel) nn.Module.__init__(model) @@ -254,82 +110,6 @@ def fake_upstream_load_weights(_self, weights): } -def test_kimi_text_model_layer_factory_accepts_prefix_keyword(monkeypatch): - config = SimpleNamespace( - vocab_size=64, - hidden_size=16, - num_hidden_layers=1, - rms_norm_eps=1e-5, - attn_res_block_size=None, - num_attention_heads=1, - ) - vllm_config = MagicMock() - vllm_config.model_config.hf_text_config = config - vllm_config.parallel_config = SimpleNamespace( - pipeline_parallel_size=1, - enable_expert_parallel=True, - tensor_parallel_size=2, - ) - pp_group = SimpleNamespace(is_first_rank=False, is_last_rank=False) - decoder_layer = nn.Identity() - decoder_layer_factory = MagicMock(return_value=decoder_layer) - - def fake_make_layers(num_hidden_layers, layer_fn, *, prefix): - assert num_hidden_layers == 1 - assert layer_fn(prefix=f"{prefix}.0") is decoder_layer - return 0, 1, nn.ModuleList([decoder_layer]) - - monkeypatch.setattr(kimi_k3, "get_pp_group", lambda: pp_group) - monkeypatch.setattr(kimi_k3, "get_tensor_model_parallel_world_size", lambda: 1) - monkeypatch.setattr(kimi_k3, "AscendKimiDecoderLayer", decoder_layer_factory) - monkeypatch.setattr(kimi_k3, "make_layers", fake_make_layers) - - model = AscendKimiLinearModel(vllm_config=vllm_config, prefix="model") - - assert model.start_layer == 0 - assert model.end_layer == 1 - decoder_layer_factory.assert_called_once_with( - config, - vllm_config, - "model.layers.0", - use_sequence_parallel=True, - ) - - -def test_kimi_mla_cache_spec_preserves_hybrid_page_padding(): - real_page_size = 128 * 576 * torch.bfloat16.itemsize - padded_page_size = real_page_size + 128 - spec = AscendMLAAttentionSpec( - block_size=128, - num_kv_heads=1, - head_size=576, - dtype=torch.bfloat16, - page_size_padded=padded_page_size, - ) - - assert spec.real_page_size_bytes == real_page_size - assert spec.page_size_bytes == padded_page_size - assert AscendMLAAttentionSpec.merge([spec, spec]).page_size_bytes == padded_page_size - - -def test_ascend_mla_exposes_layer_and_cache_contract(): - attention = AscendKimiMLAAttention.__new__(AscendKimiMLAAttention) - layer = MagicMock() - layer.layer_name = "model.layers.1.self_attn.attn" - layer.impl = object() - layer.kv_cache = (object(), object()) - layer.kv_cache_dtype = "auto" - layer._k_scale = 1.0 - attention.mla_attn = MagicMock() - attention.mla_attn.mla_attn = layer - - assert attention.layer_name == layer.layer_name - assert attention.impl is layer.impl - assert attention.kv_cache is layer.kv_cache - assert attention.kv_cache_dtype == layer.kv_cache_dtype - assert attention._k_scale == layer._k_scale - - def test_kimi_attention_residual_stays_sequence_sharded(monkeypatch): class IdentityAttention(nn.Module): def forward(self, *, hidden_states, positions): @@ -443,17 +223,12 @@ def forward(self, *, positions, hidden_states, residual): def test_kimi_model_selects_materialized_or_raw_dspark_aux_stream(monkeypatch): - class Marker(nn.Module): - def __init__(self, value: int) -> None: - super().__init__() - self.value = value - class RecordingLayer(nn.Module): def __init__(self, layer_idx: int) -> None: super().__init__() self.layer_idx = layer_idx self.prev_valid_blocks = layer_idx - self.self_attention_res_proj = Marker(layer_idx) + self.self_attention_res_proj = nn.Identity() self.self_attention_res_norm = nn.Identity() def forward(self, *, positions, hidden_states, residual): @@ -467,10 +242,7 @@ def forward(self, *, positions, hidden_states, residual): ) return materialized + 10, residual - residual_calls: list[int] = [] - - def fake_attn_res(prefix_sum, _residual, projection, _norm, num_valid_blocks): - residual_calls.append(projection.value) + def fake_attn_res(prefix_sum, _residual, _projection, _norm, num_valid_blocks): return prefix_sum + 100 * num_valid_blocks monkeypatch.setattr(kimi_k3, "_apply_ascend_attn_res", fake_attn_res) @@ -487,7 +259,7 @@ def fake_attn_res(prefix_sum, _residual, projection, _norm, num_valid_blocks): model.end_layer = 2 model.layers = nn.ModuleList([RecordingLayer(0), RecordingLayer(1)]) model.use_sequence_parallel = False - model.output_attn_res_proj = Marker(2) + model.output_attn_res_proj = nn.Identity() model.output_attn_res_norm = nn.Identity() model._set_aux_hidden_state_layers((1,)) @@ -500,7 +272,6 @@ def fake_attn_res(prefix_sum, _residual, projection, _norm, num_valid_blocks): ) torch.testing.assert_close(materialized_aux[0], torch.tensor([[111.0]])) - residual_calls.clear() model.dspark_aux_capture_materialized = False _, raw_aux = model( input_ids=None, @@ -511,76 +282,6 @@ def fake_attn_res(prefix_sum, _residual, projection, _norm, num_valid_blocks): torch.testing.assert_close(raw_aux[0], torch.tensor([[11.0]])) -def test_kimi_dspark_aux_capture_mode_is_forwarded(): - causal_model = AscendKimiLinearForCausalLM.__new__(AscendKimiLinearForCausalLM) - nn.Module.__init__(causal_model) - causal_model.model = SimpleNamespace(dspark_aux_capture_materialized=False) - - causal_model.set_dspark_aux_capture_materialized(True) - - assert causal_model.model.dspark_aux_capture_materialized is True - - wrapper = AscendKimiK3ForConditionalGeneration.__new__(AscendKimiK3ForConditionalGeneration) - nn.Module.__init__(wrapper) - wrapper.language_model = MagicMock() - - wrapper.set_dspark_aux_capture_materialized(True) - - wrapper.language_model.set_dspark_aux_capture_materialized.assert_called_once_with(True) - - -def test_dspark_configures_upstream_mla_without_rebuilding(monkeypatch): - impl = SimpleNamespace( - scale=0.0, - rotary_emb=None, - use_mla_rope=False, - ) - layer = SimpleNamespace( - scale=0.0, - non_causal_multi_token_decode=False, - impl=impl, - ) - upstream_wrapper = SimpleNamespace(mla_attn=layer) - - def fake_upstream_init(self, **_kwargs): - nn.Module.__init__(self) - self.scaling = 0.125 - self.mla_attn = upstream_wrapper - - rotary_emb = object() - monkeypatch.setattr( - kimi_k3.UpstreamKimiMLAAttention, - "__init__", - fake_upstream_init, - ) - monkeypatch.setattr(kimi_k3, "get_rope", lambda *_args, **_kwargs: rotary_emb) - - attention = AscendKimiMLAAttention( - config=SimpleNamespace( - rope_parameters={"rope_type": "default"}, - max_position_embeddings=4096, - ), - hidden_size=16, - num_heads=2, - qk_nope_head_dim=4, - qk_rope_head_dim=4, - v_head_dim=4, - q_lora_rank=8, - kv_lora_rank=8, - use_output_gate=False, - use_rope=True, - prefix="model.layers.1.self_attn", - non_causal_multi_token_decode=True, - ) - - assert attention.mla_attn is upstream_wrapper - assert layer.scale == attention.scaling - assert layer.non_causal_multi_token_decode is True - assert impl.scale == attention.scaling - assert impl.rotary_emb is rotary_emb - assert impl.use_mla_rope is True - - def test_projector_applies_optional_modelslim_rotation(): class ScaleLinear(nn.Module): def forward(self, hidden_states): @@ -604,183 +305,59 @@ def forward(self, hidden_states): torch.testing.assert_close(projector(image_features), image_features) -def test_projector_creates_rotation_only_when_enabled(monkeypatch): - def fake_upstream_init(self, *_args, **_kwargs): - nn.Module.__init__(self) - - monkeypatch.setattr( - kimi_k3.KimiK25MultiModalProjector, - "__init__", - fake_upstream_init, - ) - rotation = nn.Linear(1, 1, bias=False) - rotation_factory = MagicMock(return_value=rotation) - monkeypatch.setattr(kimi_k3, "ReplicatedLinear", rotation_factory) - config = SimpleNamespace(text_hidden_size=16) - - plain_projector = AscendKimiK3MultiModalProjector(config, prefix="mm_projector") - - assert plain_projector.rot_proj is None - rotation_factory.assert_not_called() - - rotated_projector = AscendKimiK3MultiModalProjector( - config, - prefix="mm_projector", - enable_rotation=True, - ) - - assert rotated_projector.rot_proj is rotation - rotation_factory.assert_called_once_with( - 16, - 16, - bias=False, - quant_config=None, - prefix="mm_projector.rot_proj", - ) - - -class _DraftTokenEmbedder(nn.Module): - def __init__(self) -> None: - super().__init__() - self.embedding = nn.Embedding.from_pretrained( - torch.tensor( - [ - [0.0, 0.0], - [1.0, 2.0], - [3.0, 4.0], - ] - ) - ) - - def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: - return self.embedding(input_ids) - - -def _make_k3_dspark_for_embedding_test() -> AscendK3DSparkForCausalLM: - model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) - nn.Module.__init__(model) - model.model = _DraftTokenEmbedder() - return model - - -def test_k3_dspark_load_weights_keeps_per_layer_context_kv(monkeypatch): +def test_k3_dspark_load_weights_rotates_projection_and_target_boundaries(tmp_path): model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) nn.Module.__init__(model) - model.rotation_path = None - source_weights = [ - ( - "layers.0.self_attn.kv_a_proj_with_mqa.weight", - torch.ones(1, 1), + model.model = nn.Module() + model.model.context_proj = nn.Linear(4, 2, bias=False) + model.model.context_norm = nn.LayerNorm(2) + model.model.embed_tokens = nn.Linear(2, 3, bias=False) + model.lm_head = nn.Linear(2, 3, bias=False) + model.rotation_path = tmp_path / "rotation.safetensors" + model.target_model_path = tmp_path + + # A non-symmetric rotation distinguishes projection R from vocabulary R.T. + rotation = torch.tensor([[0.0, -1.0], [1.0, 0.0]]) + embed_weight = torch.arange(6, dtype=torch.float32).view(3, 2) + head_weight = embed_weight + 10 + save_file({"global_rotation": rotation}, model.rotation_path) + save_file( + { + "language_model.model.embed_tokens.weight": embed_weight, + "language_model.lm_head.weight": head_weight, + }, + tmp_path / "model.safetensors", + ) + projection = torch.arange(8, dtype=torch.float32).view(2, 4) + norm_weight = torch.tensor([2.0, 3.0]) + + # Load the draft projection plus vocabulary weights from the target checkpoint. + model.load_weights( + iter( + [ + ("context_proj.weight", projection), + ("context_norm.weight", norm_weight), + ] ) - ] - seen_names: list[str] = [] - - class CapturingLoader: - def __init__(self, loaded_model, **kwargs): - assert loaded_model is model - if vllm_version_is("0.27.1"): - assert kwargs["skip_substrs"] == list(model.checkpoint_skip_substrs) - else: - assert kwargs == {} - - def load_weights(self, weights, *, mapper): - assert mapper is model.hf_to_vllm_mapper - seen_names.extend(name for name, _ in weights) - return {"model.layers.0.self_attn.fused_qkv_a_proj.weight"} - - monkeypatch.setattr( - "vllm_ascend.models.kimi_k3_dspark.AutoWeightsLoader", - CapturingLoader, - ) - - loaded = model.load_weights(iter(source_weights)) - - assert seen_names == [source_weights[0][0]] - assert loaded == {"model.layers.0.self_attn.fused_qkv_a_proj.weight"} - - -def test_k3_dspark_reuses_modelslim_rotation_loader(monkeypatch): - model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) - nn.Module.__init__(model) - model.rotation_path = "rotation.safetensors" - model.target_model_path = "/target" - model.model = SimpleNamespace(embed_tokens=object()) - model.lm_head = object() - source_weights = [ - ("context_proj.weight", torch.ones(2, 4)), - ("context_norm.weight", torch.ones(2)), - ] - rotated_weight = torch.full((2, 4), 2.0) - seen_weights: list[tuple[str, torch.Tensor]] = [] - - class CapturingLoader: - def __init__(self, loaded_model, **kwargs): - assert loaded_model is model - if vllm_version_is("0.27.1"): - assert kwargs["skip_substrs"] == list(model.checkpoint_skip_substrs) - else: - assert kwargs == {} - - def load_weights(self, weights, *, mapper): - assert mapper is model.hf_to_vllm_mapper - seen_weights.extend(weights) - return {name for name, _ in seen_weights} - - monkeypatch.setattr( - "vllm_ascend.models.kimi_k3_dspark.AutoWeightsLoader", - CapturingLoader, - ) - rotation = torch.eye(4) - monkeypatch.setattr( - "vllm_ascend.models.kimi_k3_dspark.get_rotation_matrix", - lambda path: rotation if path == model.rotation_path else None, - ) - process_weight = MagicMock(return_value=rotated_weight) - monkeypatch.setattr( - "vllm_ascend.models.kimi_k3_dspark.process_weight", - process_weight, - ) - load_target_layer = MagicMock() - monkeypatch.setattr( - "vllm_ascend.models.kimi_k3_dspark.load_quarot_target_layer", - load_target_layer, ) - model.load_weights(iter(source_weights)) - - process_weight.assert_called_once() - torch.testing.assert_close(process_weight.call_args.args[0], source_weights[0][1]) - torch.testing.assert_close(process_weight.call_args.args[1], rotation) - assert load_target_layer.call_count == 2 - assert load_target_layer.call_args_list[0].args[:2] == ( - model.model.embed_tokens, - model.target_model_path, - ) - assert load_target_layer.call_args_list[1].args[:2] == ( - model.lm_head, - model.target_model_path, + torch.testing.assert_close( + model.model.context_proj.weight, + projection @ torch.block_diag(rotation, rotation), ) + torch.testing.assert_close(model.model.context_norm.weight, norm_weight) + torch.testing.assert_close(model.model.embed_tokens.weight, embed_weight @ rotation.T) + torch.testing.assert_close(model.lm_head.weight, head_weight @ rotation.T) assert model.has_own_embed_tokens assert model.has_own_lm_head - assert seen_weights[0][0] == "context_proj.weight" - assert seen_weights[0][1] is rotated_weight - assert seen_weights[1][0] == source_weights[1][0] - assert seen_weights[1][1] is source_weights[1][1] - - -def test_k3_dspark_embed_input_ids_keeps_text_only_path(): - model = _make_k3_dspark_for_embedding_test() - - output = model.embed_input_ids(torch.tensor([1, 2])) - - torch.testing.assert_close( - output, - torch.tensor([[1.0, 2.0], [3.0, 4.0]]), - ) def test_k3_dspark_embed_input_ids_merges_multimodal_embeddings(): - model = _make_k3_dspark_for_embedding_test() + model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) + nn.Module.__init__(model) + model.model = SimpleNamespace( + embed_input_ids=nn.Embedding.from_pretrained(torch.tensor([[0.0, 0.0], [1.0, 2.0], [3.0, 4.0]])), + ) input_ids = torch.tensor([1, 999, 2]) is_multimodal = torch.tensor([False, True, False]) image_embedding = torch.tensor([[9.0, 10.0]]) @@ -801,14 +378,3 @@ def test_k3_dspark_embed_input_ids_merges_multimodal_embeddings(): ] ), ) - - -def test_k3_dspark_embed_input_ids_without_multimodal_mask_uses_text_path(): - model = _make_k3_dspark_for_embedding_test() - - output = model.embed_input_ids( - torch.tensor([1]), - multimodal_embeddings=(torch.tensor([[9.0, 10.0]]),), - ) - - torch.testing.assert_close(output, torch.tensor([[1.0, 2.0]])) diff --git a/tests/ut/ops/test_fused_moe.py b/tests/ut/ops/test_fused_moe.py index bed3f7be40d2..3855d0981605 100644 --- a/tests/ut/ops/test_fused_moe.py +++ b/tests/ut/ops/test_fused_moe.py @@ -857,34 +857,6 @@ def test_shared_experts_part2_applies_optional_gate(with_gate): torch.testing.assert_close(output, expected) -def test_shared_expert_consistency_uses_projection_input_width(monkeypatch): - shared_experts = AscendSharedExperts.__new__(AscendSharedExperts) - shared_experts.hidden_size = 3584 - shared_experts.shared_expert_input_size = 7168 - shared_experts.in_dtype = torch.float16 - output = torch.ones(10, 7168) - shared_experts.layer = MagicMock(return_value=output) - shared_experts.part1 = MagicMock(return_value=output) - shared_experts.part2 = MagicMock(return_value=output) - random_input = torch.ones(10, 7168) - random = MagicMock(return_value=random_input) - monkeypatch.setattr(shared_experts_module.torch, "rand", random) - - shared_experts.validate_consistency() - - random.assert_called_once_with( - 10, - 7168, - device="npu", - dtype=torch.float16, - ) - shared_experts.layer.assert_called_once() - torch.testing.assert_close( - shared_experts.layer.call_args.args[0], - random_input, - ) - - def _make_quantized_situ_shared_experts(quant_type, gate_up_proj, down_proj): shared_experts = AscendSharedExperts.__new__(AscendSharedExperts) shared_experts.layer = SimpleNamespace( diff --git a/tests/ut/ops/test_gdn_attn_builder.py b/tests/ut/ops/test_gdn_attn_builder.py index 0b0a28191b6c..efa615721f30 100644 --- a/tests/ut/ops/test_gdn_attn_builder.py +++ b/tests/ut/ops/test_gdn_attn_builder.py @@ -591,40 +591,6 @@ def test_full_graph_spec_actual_seq_lengths_use_padded_builder_buffer(): ) -def test_full_graph_non_spec_actual_seq_lengths_use_padded_builder_buffer(): - batch_spec = BatchSpec( - seq_lens=[1, 1, 0, 0], - query_lens=[1, 1, 0, 0], - name="full_graph_padded_non_spec_actual_seq_lengths", - ) - common_attn_metadata = create_common_attn_metadata( - batch_spec=batch_spec, - block_size=16, - device=torch.device("cpu"), - ) - builder = _make_builder( - device=torch.device("cpu"), - num_heads=32, - num_speculative_tokens=0, - cudagraph_mode=CUDAGraphMode.FULL_DECODE_ONLY, - ) - - attn_metadata = builder.build(0, common_attn_metadata) - - assert torch.equal( - attn_metadata.non_spec_query_start_loc, - torch.tensor([0, 1, 2, 2, 2], dtype=torch.int32), - ) - assert ( - attn_metadata.non_spec_decode_metadata.actual_seq_lengths.data_ptr() - == builder.non_spec_actual_seq_lengths.data_ptr() - ) - assert torch.equal( - attn_metadata.non_spec_decode_metadata.actual_seq_lengths, - torch.tensor([0, 1, 1, 0, 0], dtype=torch.int32), - ) - - def test_causal_conv1d_cache_indices_use_device_block_table(monkeypatch: pytest.MonkeyPatch): _patch_missing_runtime_cdiv(monkeypatch) batch_spec = BatchSpec( diff --git a/tests/ut/ops/test_kimi_kda.py b/tests/ut/ops/test_kimi_kda.py index e25e55840a44..fc1ae8a096d3 100644 --- a/tests/ut/ops/test_kimi_kda.py +++ b/tests/ut/ops/test_kimi_kda.py @@ -41,37 +41,6 @@ def test_zero_padded_output_uses_combined_live_token_count(): assert torch.equal(actual[:, 6:], torch.zeros_like(actual[:, 6:])) -def test_run_causal_conv1d_returns_declared_output_alias(): - mixed_qkv = torch.randn(3, 8) - conv_weights = torch.randn(4, 8) - conv_state = torch.randn(2, 8, 4) - query_start_loc = torch.tensor([0, 3], dtype=torch.int32) - cache_indices = torch.tensor([1], dtype=torch.int32) - returned_alias = torch.full_like(mixed_qkv, 7) - - with patch.object( - torch.ops._C_ascend, - "npu_causal_conv1d_custom", - return_value=returned_alias, - create=True, - ) as causal_conv: - actual = AscendKimiK3DeltaAttention._run_causal_conv1d( - mixed_qkv, - conv_weights, - conv_state, - query_start_loc, - cache_indices, - None, - run_mode=1, - num_accepted_tokens=torch.tensor([3], dtype=torch.int32), - ) - - assert actual is returned_alias - assert causal_conv.call_args.kwargs["query_start_loc_opt"] is query_start_loc - assert causal_conv.call_args.kwargs["cache_indices_opt"] is cache_indices - assert causal_conv.call_args.kwargs["initial_state_mode_opt"] is None - - def test_kda_output_norm_uses_checkpoint_epsilon(): def fake_upstream_init(attention, _config, _vllm_config, _prefix): nn.Module.__init__(attention) diff --git a/tests/ut/ops/test_layernorm.py b/tests/ut/ops/test_layernorm.py index 4fa8b72bceaf..4e207e6dab33 100644 --- a/tests/ut/ops/test_layernorm.py +++ b/tests/ut/ops/test_layernorm.py @@ -85,17 +85,15 @@ def test_RMSNorm_creates_bias_from_quant_description(default_vllm_config): assert not layer.bias.requires_grad -@pytest.mark.parametrize("activation", ["sigmoid", "swish"]) -@pytest.mark.parametrize("prenorm", [False, True]) -def test_FusedRMSNormGated_dispatches_to_ascend_kernel(default_vllm_config, activation, prenorm): - layer = FusedRMSNormGated(hidden_size=8, eps=1e-6, activation=activation) +def test_FusedRMSNormGated_dispatches_to_ascend_kernel(default_vllm_config): + layer = FusedRMSNormGated(hidden_size=8, eps=1e-6, activation="sigmoid") x = torch.randn(1, 4, 2, 8) gate = torch.randn(4, 2, 8) - residual = torch.randn_like(x) if prenorm else None - expected = (torch.empty_like(x), torch.empty_like(x)) if prenorm else torch.empty_like(x) + residual = torch.randn_like(x) + expected = (torch.empty_like(x), torch.empty_like(x)) with patch("vllm_ascend.ops.layernorm.rms_norm_gated", return_value=expected) as fused_norm_gate: - actual = layer(x, gate, residual=residual, prenorm=prenorm, residual_in_fp32=prenorm) + actual = layer(x, gate, residual=residual, prenorm=True, residual_in_fp32=True) assert isinstance(layer, AscendFusedRMSNormGated) assert actual is expected @@ -104,11 +102,11 @@ def test_FusedRMSNormGated_dispatches_to_ascend_kernel(default_vllm_config, acti gate, layer.weight, layer.bias, - activation, + "sigmoid", residual=residual, eps=1e-6, - prenorm=prenorm, - residual_in_fp32=prenorm, + prenorm=True, + residual_in_fp32=True, ) diff --git a/tests/ut/patch/platform/test_prefix_cache_cp_patches.py b/tests/ut/patch/platform/test_prefix_cache_cp_patches.py index aa1f3f6d4c19..545b9f76fa10 100644 --- a/tests/ut/patch/platform/test_prefix_cache_cp_patches.py +++ b/tests/ut/patch/platform/test_prefix_cache_cp_patches.py @@ -38,7 +38,6 @@ group_and_unify_kv_cache_specs, ) from vllm_ascend.patch.platform.patch_mamba_manager import AscendMambaManager -from vllm_ascend.utils import vllm_version_is def _make_hybrid_kv_cache_config( @@ -217,6 +216,7 @@ def test_ascend_mla_merge_preserves_upstream_layout_fields() -> None: head_size=128, dtype=torch.bfloat16, cache_dtype_str="fp8_ds_mla", + page_size_padded=(512 // 4) * (128 * 2 + 2) + 128, compress_ratio=4, model_version="deepseek_v4", indexes_kv_by_block_stride=True, @@ -227,6 +227,8 @@ def test_ascend_mla_merge_preserves_upstream_layout_fields() -> None: merged = AscendMLAAttentionSpec.merge([spec, replace(spec)]) assert merged.block_size == spec.block_size + assert merged.real_page_size_bytes == (512 // 4) * (128 * 2 + 2) + assert merged.page_size_bytes == spec.page_size_padded assert merged.compress_ratio == spec.compress_ratio assert merged.model_version == spec.model_version assert merged.indexes_kv_by_block_stride == spec.indexes_kv_by_block_stride @@ -302,53 +304,29 @@ def test_deepseek_v4_groups_use_logical_sizes_and_full_attention_manager() -> No @pytest.mark.parametrize( - ("block_size", "page_size"), + ("block_size", "page_size", "draft_uses_mla"), [ - pytest.param(384, 488448, id="tp16"), - pytest.param(768, 976896, id="tp8"), + pytest.param(384, 488448, False, id="gqa-tp16"), + pytest.param(768, 976896, False, id="gqa-tp8"), + pytest.param(384, 488448, True, id="mla-tp16"), ], ) -def test_kimi_k3_gqa_uses_four_mixed_kv_groups_and_five_builders( - block_size: int, - page_size: int, -) -> None: +def test_kimi_k3_dspark_uses_four_mixed_kv_groups(block_size, page_size, draft_uses_mla) -> None: groups = _get_kimi_k3_dspark_mixed_kv_cache_groups( _make_kimi_k3_dspark_kv_cache_specs( block_size=block_size, page_size=page_size, + draft_uses_mla=draft_uses_mla, ) ) assert groups is not None assert [len(group.layer_names) for group in groups] == [29, 23, 23, 23] assert all(isinstance(group.kv_cache_spec, UniformTypeKVCacheSpecs) for group in groups) - mixed_specs = groups[0].kv_cache_spec.kv_cache_specs - assert sum(isinstance(spec, MLAAttentionSpec) for spec in mixed_specs.values()) == 24 - assert ( - sum( - isinstance(spec, FullAttentionSpec) and not isinstance(spec, MLAAttentionSpec) - for spec in mixed_specs.values() - ) - == 5 - ) - - # The model runner keys builders by backend and exact inner spec. The - # mixed attention group contributes MLA + GQA, and each Mamba group one. - metadata_builder_count = sum(len(set(group.kv_cache_spec.kv_cache_specs.values())) for group in groups) - assert metadata_builder_count == 5 - - -def test_kimi_k3_mla_dspark_uses_four_groups_and_five_builders() -> None: - groups = _get_kimi_k3_dspark_mixed_kv_cache_groups(_make_kimi_k3_dspark_kv_cache_specs(draft_uses_mla=True)) - - assert groups is not None - assert [len(group.layer_names) for group in groups] == [29, 23, 23, 23] - # Target MLA is causal while draft MLA enables non-causal multi-token - # decode, so the mixed group needs two exact-spec builders. The three - # recurrent groups each contribute one more builder. - metadata_builder_count = sum(len(set(group.kv_cache_spec.kv_cache_specs.values())) for group in groups) - assert metadata_builder_count == 5 + expected_mla_count = 29 if draft_uses_mla else 24 + assert sum(isinstance(spec, MLAAttentionSpec) for spec in mixed_specs.values()) == expected_mla_count + assert all(isinstance(spec, FullAttentionSpec) for spec in mixed_specs.values()) def test_kimi_k3_gqa_mixed_groups_preserve_scheduler_and_mamba_contracts() -> None: @@ -563,46 +541,6 @@ def _fake_orig(*args, **kwargs): assert coordinator is sentinel -@pytest.mark.parametrize("num_prefill_lookahead", [0, 8]) -@pytest.mark.skipif( - vllm_version_is("0.27.1"), - reason="num_prefill_lookahead was added to the coordinator after v0.27.1", -) -def test_get_kv_cache_coordinator_forwards_prefill_lookahead( - monkeypatch, - num_prefill_lookahead: int, -) -> None: - kv_cache_config = _make_hybrid_kv_cache_config( - full_block_size=16, - mamba_block_size=16, - ) - captured_kwargs = {} - - def _fake_ascend_coordinator(*args, **kwargs): - captured_kwargs.update(kwargs) - return object() - - monkeypatch.setattr( - "vllm_ascend.patch.platform.patch_kv_cache_coordinator.AscendHybridKVCacheCoordinator", - _fake_ascend_coordinator, - ) - - get_kv_cache_coordinator( - kv_cache_config, - max_model_len=1024, - max_num_batched_tokens=1024, - use_eagle=False, - enable_caching=True, - enable_kv_cache_events=False, - dcp_world_size=1, - pcp_world_size=1, - hash_block_size=16, - num_prefill_lookahead=num_prefill_lookahead, - ) - - assert captured_kwargs["num_prefill_lookahead"] == num_prefill_lookahead - - def test_get_kv_cache_coordinator_uses_ascend_for_deepseek_v4(monkeypatch) -> None: sentinel = object() kv_cache_config = _make_deepseek_v4_kv_cache_config() diff --git a/tests/ut/patch/worker/test_patch_mamba_utils.py b/tests/ut/patch/worker/test_patch_mamba_utils.py index f43b5b9e0454..cf47c2151c19 100644 --- a/tests/ut/patch/worker/test_patch_mamba_utils.py +++ b/tests/ut/patch/worker/test_patch_mamba_utils.py @@ -1,66 +1,31 @@ # SPDX-License-Identifier: Apache-2.0 from types import SimpleNamespace -from unittest.mock import MagicMock, call, patch +from unittest.mock import patch import numpy as np +import torch +from vllm.v1.utils import CpuGpuBuffer from vllm_ascend.patch.worker.patch_mamba_utils import ( _do_mamba_copy_block_npu, - _stage_mamba_copy_metadata, preprocess_mamba, ) -def _copy_buffer(): - gpu = MagicMock() - gpu_view = MagicMock() - gpu.__getitem__.return_value = gpu_view - return SimpleNamespace(copy_to_gpu=MagicMock(), gpu=gpu), gpu_view - - -def test_mamba_copy_metadata_is_staged_asynchronously_during_preprocess(): - src_ptrs, _ = _copy_buffer() - dst_ptrs, _ = _copy_buffer() - sizes, _ = _copy_buffer() - copy_bufs = SimpleNamespace( - offset=2, - src_ptrs=src_ptrs, - dst_ptrs=dst_ptrs, - sizes=sizes, - ) - - _stage_mamba_copy_metadata(copy_bufs) - - for buffer in (src_ptrs, dst_ptrs, sizes): - buffer.copy_to_gpu.assert_called_once_with(2) - - -def test_mamba_state_copy_uses_previously_staged_metadata(): - src_ptrs, src_view = _copy_buffer() - dst_ptrs, dst_view = _copy_buffer() - sizes, sizes_view = _copy_buffer() - copy_bufs = SimpleNamespace( - offset=2, - src_ptrs=src_ptrs, - dst_ptrs=dst_ptrs, - sizes=sizes, - ) - - with patch("vllm_ascend.patch.worker.patch_mamba_utils._batch_memcpy_triton") as batch_memcpy: - _do_mamba_copy_block_npu(copy_bufs) - - for buffer in (src_ptrs, dst_ptrs, sizes): - buffer.copy_to_gpu.assert_not_called() - assert buffer.gpu.__getitem__.call_args_list == [call(slice(None, 2))] - batch_memcpy.assert_called_once_with(src_view, dst_view, sizes_view) - - def test_preprocess_stages_metadata_but_defers_state_copy(): + # Separate CPU-backed buffers let us check staging without an NPU. + buffers = [ + CpuGpuBuffer(2, dtype=dtype, device=torch.device("cpu"), pin_memory=False) + for dtype in (torch.int64, torch.int64, torch.int32) + ] copy_bufs = SimpleNamespace( offset=0, mamba_group_ids=[0], mamba_spec=SimpleNamespace(num_speculative_blocks=1, block_size=7), + src_ptrs=buffers[0], + dst_ptrs=buffers[1], + sizes=buffers[2], ) scheduler_output = SimpleNamespace( finished_req_ids=[], @@ -76,16 +41,17 @@ def test_preprocess_stages_metadata_but_defers_state_copy(): mamba_state_idx = {"req": 0} def collect_metadata(copy_buffers, *_args): + for buffer, value in zip(buffers, (100, 200, 32)): + buffer.np[0] = value copy_buffers.offset = 1 with ( patch( "vllm_ascend.patch.worker.patch_mamba_utils.mamba_utils.collect_mamba_copy_meta", side_effect=collect_metadata, - ) as collect, + ), patch("vllm_ascend.patch.worker.patch_mamba_utils._can_launch_triton_batch_memcpy", return_value=True), - patch("vllm_ascend.patch.worker.patch_mamba_utils._stage_mamba_copy_metadata") as stage, - patch("vllm_ascend.patch.worker.patch_mamba_utils._do_mamba_copy_block_npu") as state_copy, + patch("vllm_ascend.patch.worker.patch_mamba_utils._batch_memcpy_triton") as state_copy, ): preprocess_mamba( scheduler_output, @@ -99,9 +65,17 @@ def collect_metadata(copy_buffers, *_args): copy_bufs, ) - collect.assert_called_once() - stage.assert_called_once_with(copy_bufs) - state_copy.assert_not_called() + state_copy.assert_not_called() + for buffer, value in zip(buffers, (100, 200, 32)): + torch.testing.assert_close(buffer.gpu, torch.tensor([value, 0], dtype=buffer.gpu.dtype)) + # Later host reuse must not overwrite metadata staged for this step. + buffer.cpu.fill_(-1) + + _do_mamba_copy_block_npu(copy_bufs) + + state_copy.assert_called_once() + for actual, value in zip(state_copy.call_args.args, (100, 200, 32)): + torch.testing.assert_close(actual, torch.tensor([value], dtype=actual.dtype)) assert input_batch.num_accepted_tokens_cpu.tolist() == [1] diff --git a/tests/ut/patch/worker/test_patch_mamba_utils_source.py b/tests/ut/patch/worker/test_patch_mamba_utils_source.py deleted file mode 100644 index 88dc71282be5..000000000000 --- a/tests/ut/patch/worker/test_patch_mamba_utils_source.py +++ /dev/null @@ -1,61 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -"""Source-level checks for the Ascend Mamba precision-kernel overridess.""" - -from __future__ import annotations - -import ast -from pathlib import Path - -from vllm_ascend.utils import vllm_version_is - -ROOT = Path(__file__).resolve().parents[4] -POSTPROCESS = ROOT / "vllm_ascend" / "ops" / "triton" / "mamba" / "postprocess.py" - - -def _top_level_functions(path: Path) -> dict[str, ast.FunctionDef]: - return {node.name: node for node in ast.parse(path.read_text()).body if isinstance(node, ast.FunctionDef)} - - -def _postprocess_kernels(path: Path) -> list[ast.FunctionDef]: - """Collect both postprocess_mamba_fused_kernel definitions nested under - the vllm_version_is('0.27.1') if/else gate. Index 0 = v0.27.1 branch, - index 1 = main (else) branch.""" - kernels = [] - - def _walk(node: ast.AST) -> None: - for child in ast.iter_child_nodes(node): - if isinstance(child, ast.FunctionDef) and child.name == "postprocess_mamba_fused_kernel": - kernels.append(child) - else: - _walk(child) - - _walk(ast.parse(path.read_text())) - return kernels - - -def _selected_kernel_source(path: Path) -> str: - """Return the kernel source that is active for the current vllm version.""" - kernels = _postprocess_kernels(path) - assert len(kernels) == 2, "expected one kernel per vllm_version_is branch" - idx = 0 if vllm_version_is("0.27.1") else 1 - return ast.unparse(kernels[idx]) - - -def test_postprocess_keeps_only_existing_ascend_precision_kernel() -> None: - functions = _top_level_functions(POSTPROCESS) - - assert set(functions) == set() - postprocess_source = _selected_kernel_source(POSTPROCESS) - assert "src_ptr = src_addr.to(tl.pointer_type(tl.uint8))" in postprocess_source - assert "dst_ptr = dst_addr.to(tl.pointer_type(tl.uint8))" in postprocess_source - assert "PRECOMPUTED_NEW_COMPUTED" in postprocess_source - assert "tl.store(num_accepted_tokens_ptr + req_idx, 1)" in postprocess_source - - if vllm_version_is("0.27.1"): - assert "TEMPORAL_TILES" not in postprocess_source - assert "tile_idx" not in postprocess_source - else: - assert "TEMPORAL_TILES" in postprocess_source - assert "tile_idx" in postprocess_source - assert "if tile_idx == 0:" in postprocess_source - assert "and state_idx == 0 and tile_idx == 0" not in postprocess_source diff --git a/tests/ut/quantization/test_modelslim_config.py b/tests/ut/quantization/test_modelslim_config.py index 8584d383d021..745024416cbf 100644 --- a/tests/ut/quantization/test_modelslim_config.py +++ b/tests/ut/quantization/test_modelslim_config.py @@ -199,42 +199,6 @@ def test_get_quant_method_for_moe_installs_modelslim_weight_loader(self): return_success=True, ) - def test_get_quant_method_for_kimi_linear_moe_uses_kimi_k3_mapping(self): - prefix = "language_model.model.layers.1.block_sparse_moe.experts" - quant_description = {f"{prefix}.0.{name}.weight": "W4A8_DYNAMIC" for name in ("w1", "w2", "w3")} - config = AscendModelSlimConfig(quant_description) - layer = RoutedExperts.__new__(RoutedExperts) - torch.nn.Module.__init__(layer) - layer.moe_config = MagicMock() - layer.weight_loader = MagicMock(return_value=True) - mock_vllm_config = MagicMock() - mock_vllm_config.model_config.hf_config.model_type = "kimi_linear" - mock_scheme = MagicMock() - - with ( - patch( - "vllm_ascend.quantization.modelslim_config.get_current_vllm_config", - return_value=mock_vllm_config, - ), - patch( - "vllm_ascend.quantization.modelslim_config.create_scheme_for_layer", - return_value=mock_scheme, - ) as create_scheme, - patch( - "vllm_ascend.quantization.method_adapters.AscendFusedMoEMethod", - return_value=MagicMock(), - ), - ): - config.get_quant_method(layer, prefix) - - self.assertEqual(config.packed_modules_mapping, get_packed_modules_mapping("kimi_k3")) - create_scheme.assert_called_once_with( - quant_description, - prefix, - "moe", - get_packed_modules_mapping("kimi_k3"), - ) - def test_get_quant_method_for_c8_kv_cache_attention(self): c8_config = AscendModelSlimConfig( { @@ -642,28 +606,6 @@ def test_gemma4_packed_modules_mapping_covers_attention_mlp_and_moe(self): with self.subTest(model_type=model_type): self.assertEqual(get_packed_modules_mapping(model_type), expected_mapping) - def test_kimi_k3_packed_modules_mapping_covers_kda_and_moe(self): - expected_mapping = { - "gate_up_proj": ["gate_proj", "up_proj"], - "experts": ["experts.0.w1", "experts.0.w2", "experts.0.w3"], - "in_proj_qkvgfab": [ - "q_proj", - "k_proj", - "v_proj", - "g_proj", - "f_a_proj", - "b_proj", - ], - "in_proj_qkv": ["q_proj", "k_proj", "v_proj"], - "in_proj_gfab": ["g_proj", "f_a_proj", "b_proj"], - "conv1d": ["q_conv1d", "k_conv1d", "v_conv1d"], - "fused_qkv_a_proj": ["q_a_proj", "kv_a_proj_with_mqa"], - } - - for model_type in ("kimi_k3", "kimi_linear"): - with self.subTest(model_type=model_type): - self.assertEqual(get_packed_modules_mapping(model_type), expected_mapping) - def test_kimi_k3_modelslim_resolves_fused_kda_and_moe_types(self): layer_prefix = "language_model.model.layers.1" quant_description = { diff --git a/tests/ut/spec_decode/test_dspark_proposer.py b/tests/ut/spec_decode/test_dspark_proposer.py index 2e88b243f351..617be3b3be8d 100644 --- a/tests/ut/spec_decode/test_dspark_proposer.py +++ b/tests/ut/spec_decode/test_dspark_proposer.py @@ -19,7 +19,6 @@ from __future__ import annotations -import inspect from types import SimpleNamespace from unittest.mock import MagicMock, patch @@ -34,9 +33,7 @@ from vllm.v1.worker.utils import AttentionGroup from vllm_ascend.attention.attention_v1 import AscendAttentionState -from vllm_ascend.spec_decode.dflash_proposer import AscendDflashProposer from vllm_ascend.spec_decode.dspark_proposer import AscendDSparkProposer -from vllm_ascend.spec_decode.llm_base_proposer import AscendSpecDecodeBaseProposer # 0 = single-DP (no padding); >0 = multi-DP where num_input_tokens > # num_query_total, the out-of-bounds regime. @@ -64,12 +61,12 @@ class _DSparkProposerTestBase: """Shared helpers for ``AscendDSparkProposer`` tests.""" @staticmethod - def _make_vllm_config(hf_config: SimpleNamespace) -> SimpleNamespace: + def _make_vllm_config(hf_config: SimpleNamespace, draft_sample_method: str) -> SimpleNamespace: """Build the minimal config consumed by the DSpark initializer.""" draft_model_config = SimpleNamespace(hf_config=hf_config, get_hidden_size=lambda: _HIDDEN_SIZE) return SimpleNamespace( speculative_config=SimpleNamespace( - draft_sample_method="greedy", + draft_sample_method=draft_sample_method, draft_model_config=draft_model_config, ) ) @@ -83,9 +80,10 @@ def _make_proposer( block_size: int, hf_config: SimpleNamespace | None = None, draft_attn_causal: bool | None = None, + draft_sample_method: str = "greedy", ): device = torch.device("cpu") - vllm_config = cls._make_vllm_config(hf_config or SimpleNamespace()) + vllm_config = cls._make_vllm_config(hf_config or SimpleNamespace(), draft_sample_method) def mock_parent_init( proposer: AscendDSparkProposer, @@ -313,62 +311,17 @@ def test_noop_without_dp_padding(self): proposer._pad_draft_buffers(num_actual, num_actual) assert torch.equal(proposer.positions, snapshot) - def test_must_precede_build(self): - """build_draft_attn_metadata reads positions but does not zero it, so - _pad_draft_buffers must run first.""" - num_reqs, block_size, max_num_tokens = 4, 5, 256 - num_actual = num_reqs * block_size - num_input = num_actual + 16 - - def capture_build(): - captured = {} - - def fake_build(common_attn_metadata, num_input_tokens, num_actual_tokens): - captured["region"] = common_attn_metadata.positions[num_actual:num_input].clone() - return None, common_attn_metadata - - return captured, fake_build - - ok = self._make_proposer(max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size) - ok.positions[num_actual:num_input] = -999 - cap_ok, build_ok = capture_build() - ok.build_draft_attn_metadata = build_ok - ok._pad_draft_buffers(num_actual, num_input) - ok.build_draft_attn_metadata(SimpleNamespace(positions=ok.positions), num_input, num_actual) - assert torch.all(cap_ok["region"] == 0) - - bug = self._make_proposer(max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size) - bug.positions[num_actual:num_input] = -999 - cap_bug, build_bug = capture_build() - bug.build_draft_attn_metadata = build_bug - bug.build_draft_attn_metadata(SimpleNamespace(positions=bug.positions), num_input, num_actual) - bug._pad_draft_buffers(num_actual, num_input) - assert torch.all(cap_bug["region"] == -999) - - def test_called_before_build_in_propose(self): - """In ``_propose`` the ``_pad_draft_buffers`` call must precede - ``build_draft_attn_metadata``.""" - src = inspect.getsource(AscendSpecDecodeBaseProposer._propose) - pad_idx = src.find("self._pad_draft_buffers(") - build_idx = src.find("self.build_draft_attn_metadata(") - # Only assert when both calls live directly in _propose; a refactor that - # extracts them elsewhere leaves this guard inert rather than brittle. - if pad_idx != -1 and build_idx != -1: - assert pad_idx < build_idx, ( - "_pad_draft_buffers must be called before build_draft_attn_metadata " - "in _propose, otherwise the attention backend reads un-zeroed " - "positions in the DP-padding region." - ) - class TestDSparkInitialization(_DSparkProposerTestBase): """Tests for DSpark initialization configuration.""" @pytest.mark.parametrize( - ("hf_config", "expected_sample_from_anchor", "expected_num_query_per_req"), + ("hf_config", "expected_sample_from_anchor", "expected_num_query_per_req", "draft_sample_method"), [ - pytest.param(SimpleNamespace(), True, _NUM_SPECULATIVE_TOKENS), - pytest.param(SimpleNamespace(sample_from_anchor=False), False, 1 + _NUM_SPECULATIVE_TOKENS), + pytest.param(SimpleNamespace(), True, _NUM_SPECULATIVE_TOKENS, "greedy"), + pytest.param( + SimpleNamespace(sample_from_anchor=False), False, 1 + _NUM_SPECULATIVE_TOKENS, "probabilistic" + ), ], ) def test_configures_anchor_sampling( @@ -376,6 +329,7 @@ def test_configures_anchor_sampling( hf_config: SimpleNamespace, expected_sample_from_anchor: bool, expected_num_query_per_req: int, + draft_sample_method: str, ) -> None: """Verify the bonus-anchor flag selects the expected query layout.""" proposer = self._make_proposer( @@ -383,166 +337,13 @@ def test_configures_anchor_sampling( num_reqs=_MAX_BATCH_SIZE, block_size=_NUM_SPECULATIVE_TOKENS, hf_config=hf_config, + draft_sample_method=draft_sample_method, ) expected_max_query_tokens = _MAX_BATCH_SIZE * expected_num_query_per_req assert proposer.sample_from_anchor is expected_sample_from_anchor assert proposer.num_query_per_req == expected_num_query_per_req assert proposer.max_query_tokens == expected_max_query_tokens - - -class TestSetPerGroupAttnMetadata(_DSparkProposerTestBase): - """``set_per_group_attn_metadata`` stores the runner-provided per-group - block table / slot mapping into the read-only dicts the proposer consults - during ``set_inputs_first_pass``.""" - - def test_stores_block_table_and_slot_mapping(self): - num_reqs, block_size, max_num_tokens = 4, 5, 256 - proposer = self._make_proposer(max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size) - # a gid not pre-populated by _make_proposer (which only seeds gid=0) - gid = 7 - block_table = torch.zeros((num_reqs, 16), dtype=torch.int32) - slot_mapping = torch.full((max_num_tokens,), 42, dtype=torch.int32) - - proposer.set_per_group_attn_metadata(gid, block_table, slot_mapping) - - assert proposer._per_group_block_tables[gid] is block_table - assert proposer._per_group_slot_mappings[gid] is slot_mapping - - def test_overwrites_existing_gid(self): - num_reqs, block_size, max_num_tokens = 2, 5, 256 - proposer = self._make_proposer(max_num_tokens=max_num_tokens, num_reqs=num_reqs, block_size=block_size) - gid = 0 # already populated by _make_proposer - old_block_table = proposer._per_group_block_tables[gid] - new_block_table = torch.ones((num_reqs, 16), dtype=torch.int32) - new_slot_mapping = torch.ones(max_num_tokens, dtype=torch.int32) - - proposer.set_per_group_attn_metadata(gid, new_block_table, new_slot_mapping) - - assert proposer._per_group_block_tables[gid] is new_block_table - assert proposer._per_group_slot_mappings[gid] is new_slot_mapping - assert proposer._per_group_block_tables[gid] is not old_block_table - - -class TestDSparkInitValidation: - """``AscendDSparkProposer.__init__`` accepts probabilistic draft sampling - (supported via the per-step probability collection in - ``llm_base_proposer``), allocates the DSpark-specific draft/seed buffers - and overrides the DFlash query-token / cudagraph defaults.""" - - @staticmethod - def _make_vllm_config( - *, - num_speculative_tokens, - max_batch_size, - max_num_tokens, - draft_sample_method, - hidden_size=8, - ): - speculative_config = SimpleNamespace( - num_speculative_tokens=num_speculative_tokens, - draft_sample_method=draft_sample_method, - draft_model_config=SimpleNamespace(hf_config=SimpleNamespace(), get_hidden_size=lambda: hidden_size), - ) - return SimpleNamespace(speculative_config=speculative_config) - - @staticmethod - def _stub_dflash_init( - monkeypatch, - *, - num_speculative_tokens, - max_batch_size, - max_num_tokens, - dtype, - device, - ): - """Replace the heavy DFlash/Eagle base init with a stub that only sets - the attributes DSpark's ``__init__`` subsequently reads.""" - - def _stub(self, vllm_config, device, runner=None): - self.num_speculative_tokens = num_speculative_tokens - self.max_batch_size = max_batch_size - self.max_num_tokens = max_num_tokens - self.dtype = dtype - self.device = device - self.draft_model_config = vllm_config.speculative_config.draft_model_config - self.hidden_size = 0 - self.hidden_states = None - self._dflash_hidden_states = None - - monkeypatch.setattr(AscendDflashProposer, "__init__", _stub) - - def test_probabilistic_accepted(self, monkeypatch): - device = torch.device("cpu") - self._stub_dflash_init( - monkeypatch, - num_speculative_tokens=5, - max_batch_size=16, - max_num_tokens=256, - dtype=torch.float32, - device=device, - ) - vllm_config = self._make_vllm_config( - num_speculative_tokens=5, - max_batch_size=16, - max_num_tokens=256, - draft_sample_method="probabilistic", - ) - # Probabilistic draft sampling is supported (per-step probabilities - # are collected in llm_base_proposer); init must not raise. - proposer = AscendDSparkProposer(vllm_config, device) - assert proposer._dspark_draft_buffer.shape == (16, 6) - assert proposer._dspark_seed_buffer.shape == (16,) - - def test_greedy_allocates_dspark_buffers(self, monkeypatch): - device = torch.device("cpu") - num_spec, max_batch, max_num_tokens, hidden = 5, 16, 256, 8 - self._stub_dflash_init( - monkeypatch, - num_speculative_tokens=num_spec, - max_batch_size=max_batch, - max_num_tokens=max_num_tokens, - dtype=torch.float32, - device=device, - ) - vllm_config = self._make_vllm_config( - num_speculative_tokens=num_spec, - max_batch_size=max_batch, - max_num_tokens=max_num_tokens, - draft_sample_method="greedy", - hidden_size=hidden, - ) - dynamic_spec_config = SimpleNamespace(method="", method_params={}) - with patch( - "vllm_ascend.spec_decode.dspark_proposer.get_ascend_config", - return_value=SimpleNamespace( - dynamic_spec_config=dynamic_spec_config, - ), - ): - proposer = AscendDSparkProposer(vllm_config, device) - - blk = 1 + num_spec - max_query_tokens = max_batch * num_spec - # DSpark-specific draft / seed buffers. - assert proposer._dspark_draft_buffer.shape == (max_batch, blk) - assert proposer._dspark_draft_buffer.dtype == torch.int64 - assert proposer._dspark_seed_buffer.shape == (max_batch,) - assert proposer._dspark_seed_buffer.dtype == torch.int64 - # hidden_size / hidden states come from the draft model config. - assert proposer.hidden_size == hidden - assert proposer.hidden_states.shape == (max_num_tokens, hidden) - assert proposer._dflash_hidden_states.shape == (max_num_tokens, hidden) - # DSpark runs eager only (Ascend cudagraph unsupported on this path). - assert proposer.use_cuda_graph is False - # anchor-first: N query tokens per request, no bonus token (unlike - # DFlash's 1+N). - assert proposer.max_query_tokens == max_query_tokens - assert proposer.positions.shape == (max_query_tokens,) - assert proposer.positions.dtype == torch.int32 - assert proposer._slot_mapping_buffer.shape == (max_query_tokens,) - # per-group bookkeeping dicts start empty / None. - assert proposer._per_group_block_tables == {} - assert proposer._per_group_slot_mappings == {} - assert proposer._context_slot_mapping_buffers is None + assert proposer._dspark_draft_buffer.shape == (_MAX_BATCH_SIZE, 1 + _NUM_SPECULATIVE_TOKENS) class TestSetInputsFirstPassOutputs(_DSparkProposerTestBase): diff --git a/tests/ut/test_utils.py b/tests/ut/test_utils.py index e4e30e040a68..aceef39fa0b0 100644 --- a/tests/ut/test_utils.py +++ b/tests/ut/test_utils.py @@ -526,16 +526,8 @@ def test_check_gdn_layer_supports_kimi_linear_config_property(): assert utils.check_gdn_layer(vllm_config) is True -@pytest.mark.parametrize( - "hf_config", - [ - pytest.param( - SimpleNamespace(text_config=SimpleNamespace(layer_types=["linear_attention"])), - id="qwen3-5-nested-text-config", - ), - ], -) -def test_check_gdn_layer_supports_layer_types(hf_config): +def test_check_gdn_layer_supports_nested_layer_types(): + hf_config = SimpleNamespace(text_config=SimpleNamespace(layer_types=["linear_attention"])) vllm_config = SimpleNamespace(model_config=SimpleNamespace(hf_config=hf_config)) assert utils.check_gdn_layer(vllm_config) is True diff --git a/tests/ut/worker/a2/test_model_runner_v1.py b/tests/ut/worker/a2/test_model_runner_v1.py index 403daee403fd..aa4d2a92eb96 100644 --- a/tests/ut/worker/a2/test_model_runner_v1.py +++ b/tests/ut/worker/a2/test_model_runner_v1.py @@ -22,7 +22,6 @@ from vllm_ascend.attention.mla_v1 import AscendMLABackend from vllm_ascend.attention.utils import get_sfa_qsfa_packed_head_dim from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec, AscendSFAIndexerCacheSpec -from vllm_ascend.spec_decode.dspark_proposer import AscendDSparkProposer from vllm_ascend.utils import AscendDeviceType from vllm_ascend.worker.model_runner_v1 import NPUModelRunner @@ -204,15 +203,6 @@ def test_non_align_postprocess_keeps_an_independent_snapshot(self): self.assertEqual(postprocess_all.call_count, int(mode == "all")) runner.num_accepted_tokens_event.record.assert_called_once() - def test_non_async_postprocess_keeps_upstream_behavior(self): - runner = self._build_runner() - runner.use_async_scheduling = False - output_token_ids = torch.tensor([[10, -1]]) - scheduler_output = SimpleNamespace() - with patch("vllm.v1.worker.gpu_model_runner.GPUModelRunner._update_states_after_model_execute") as upstream: - runner._update_states_after_model_execute(output_token_ids, scheduler_output) - upstream.assert_called_once_with(output_token_ids, scheduler_output) - class TestNPUModelRunnerKVCache(unittest.TestCase): def _build_runner(self): @@ -245,55 +235,6 @@ def _build_runner(self): runner.attn_backend = backend return runner - @patch("vllm_ascend.worker.model_runner_v1.has_kv_transfer_group", return_value=False) - @patch("vllm_ascend.worker.model_runner_v1.apply_layerwise_kv_cache_plan") - def test_drafter_receives_logical_block_size_for_every_cache_group( - self, - _mock_apply_layerwise_plan, - _mock_has_kv_transfer_group, - ): - runner = self._build_runner() - runner.attn_groups = [] - runner.model_config.enable_return_routed_experts = False - runner.speculative_config = SimpleNamespace( - use_eagle=lambda: False, - uses_draft_model=lambda: True, - uses_extract_hidden_states=lambda: False, - ) - drafter = AscendDSparkProposer.__new__(AscendDSparkProposer) - drafter.initialize_attn_backend = MagicMock() - runner.drafter = drafter - runner.may_add_encoder_only_layers_to_kv_cache_config = MagicMock() - runner.maybe_add_kv_sharing_layers_to_kv_cache_groups = MagicMock() - - def initialize_attn_backend(_kv_cache_config): - runner.attn_groups = [ - [SimpleNamespace(kv_cache_spec=object())], - [SimpleNamespace(kv_cache_spec=object())], - ] - - runner.initialize_attn_backend = MagicMock(side_effect=initialize_attn_backend) - - def reinitialize_input_batch(_kv_cache_config): - runner.kernel_block_sizes = [[0], [128]] - - runner.may_reinitialize_input_batch = MagicMock(side_effect=reinitialize_input_batch) - runner.initialize_kv_cache_tensors = MagicMock(return_value={}) - - runner.initialize_kv_cache( - KVCacheConfig( - num_blocks=0, - kv_cache_tensors=[], - kv_cache_groups=[], - ) - ) - - drafter.initialize_attn_backend.assert_called_once() - self.assertEqual( - drafter.initialize_attn_backend.call_args.args[1], - [0, 128], - ) - def test_allocate_kv_cache_uses_layer_spec_for_draft_gqa(self): runner = self._build_runner() runner.sparse_kv_offload_enabled = False