已合并
add new operator matmul_abft_verify #94
add new operator matmul_abft_verify #94
已合并
starfican创建于 9月4日
共 90 个文件变更+18846-12
@@ -77,6 +77,24 @@ function(kernel_src_copy)
77 endif()77 endif()
78endfunction()78endfunction()
79 79 
80+# ######################################################################################################################
81+# get op_type from *_def.cpp
82+# ######################################################################################################################
83+function(get_op_type_from_op_name OP_NAME OP_TYPE)
84+ execute_process(
85+ COMMAND
86+ find ${CMAKE_CURRENT_SOURCE_DIR} -name ${OP_NAME}_def.cpp -exec grep OP_ADD {} \;
87+ OUTPUT_VARIABLE op_type
88+ )
89+ if(NOT op_type)
90+ set(op_type "")
91+ else()
92+ string(REGEX REPLACE "[\t ]*OP_ADD\\([\t ]*" "" op_type "${op_type}")
93+ string(REGEX REPLACE "[\t ]*\\).*$" "" op_type "${op_type}")
94+ endif()
95+ set(${OP_TYPE} "${op_type}" PARENT_SCOPE)
96+endfunction()
97+ 
80function(get_op_type_and_validate OP_DIR compute_unit op_name_var op_type_var is_valid_var)98function(get_op_type_and_validate OP_DIR compute_unit op_name_var op_type_var is_valid_var)
81 get_filename_component(op_name "${OP_DIR}" NAME)99 get_filename_component(op_name "${OP_DIR}" NAME)
82 set(${op_name_var} "${op_name}" PARENT_SCOPE)100 set(${op_name_var} "${op_name}" PARENT_SCOPE)
@@ -95,7 +113,7 @@ function(get_op_type_and_validate OP_DIR compute_unit op_name_var op_type_var is
95 endif()113 endif()
96 114 
97 set(op_type "")115 set(op_type "")
98- set(binary_json ${OP_DIR}/op_host/config/${compute_unit}/${op_name}_binary.json)116+ set(binary_json "${OP_DIR}/op_host/config/${compute_unit}/${op_name}_binary.json")
99 117 
100 if(NOT EXISTS "${OP_DIR}/op_kernel")118 if(NOT EXISTS "${OP_DIR}/op_kernel")
101 message(STATUS "[INFO] The op_kernel folder does not exist, [${op_name}] not need to compile.")119 message(STATUS "[INFO] The op_kernel folder does not exist, [${op_name}] not need to compile.")
@@ -104,7 +122,7 @@ function(get_op_type_and_validate OP_DIR compute_unit op_name_var op_type_var is
104 return()122 return()
105 endif()123 endif()
106 124 
107- if(EXISTS ${binary_json})125+ if(EXISTS "${binary_json}")
108 get_op_type_from_binary_json("${binary_json}" op_type)126 get_op_type_from_binary_json("${binary_json}" op_type)
109 message(STATUS "[INFO] On [${compute_unit}], [${op_name}] compile binary with self config.")127 message(STATUS "[INFO] On [${compute_unit}], [${op_name}] compile binary with self config.")
110 if(NOT op_type)128 if(NOT op_type)
@@ -125,7 +143,7 @@ function(get_op_type_and_validate OP_DIR compute_unit op_name_var op_type_var is
125 check_op_supported("${op_name}" "${compute_unit}" check_op_supported_result)143 check_op_supported("${op_name}" "${compute_unit}" check_op_supported_result)
126 if(NOT check_op_supported_result)144 if(NOT check_op_supported_result)
127 message(STATUS "[INFO] On [${compute_unit}], [${op_name}] not supported.")145 message(STATUS "[INFO] On [${compute_unit}], [${op_name}] not supported.")
128- set(${op_type_var} ${op_type} PARENT_SCOPE)146+ set(${op_type_var} "${op_type}" PARENT_SCOPE)
129 set(${is_valid_var} FALSE PARENT_SCOPE)147 set(${is_valid_var} FALSE PARENT_SCOPE)
130 return()148 return()
131 endif()149 endif()
@@ -391,6 +409,8 @@ function(prepare_compile_from_config)
391 COMMAND cp ${CMAKE_BINARY_DIR}/binary/${CONFCMP_COMPUTE_UNIT}/gen/${CONFCMP_OP_NAME}/${CONFCMP_OP_NAME}_binary.json ${ASCEND_KERNEL_CONF_DST}/${CONFCMP_COMPUTE_UNIT}/${CONFCMP_OP_NAME}409 COMMAND cp ${CMAKE_BINARY_DIR}/binary/${CONFCMP_COMPUTE_UNIT}/gen/${CONFCMP_OP_NAME}/${CONFCMP_OP_NAME}_binary.json ${ASCEND_KERNEL_CONF_DST}/${CONFCMP_COMPUTE_UNIT}/${CONFCMP_OP_NAME}
392 COMMENT "cp ${CMAKE_BINARY_DIR}/binary/${CONFCMP_COMPUTE_UNIT}/gen/${CONFCMP_OP_NAME}/${CONFCMP_OP_NAME}_binary.json ${ASCEND_KERNEL_CONF_DST}/${CONFCMP_COMPUTE_UNIT}/${CONFCMP_OP_NAME}"410 COMMENT "cp ${CMAKE_BINARY_DIR}/binary/${CONFCMP_COMPUTE_UNIT}/gen/${CONFCMP_OP_NAME}/${CONFCMP_OP_NAME}_binary.json ${ASCEND_KERNEL_CONF_DST}/${CONFCMP_COMPUTE_UNIT}/${CONFCMP_OP_NAME}"
393 )411 )
412+ add_dependencies(bin_conf_${CONFCMP_OP_NAME}_${CONFCMP_COMPUTE_UNIT}_copy
413+ generate_bin_scripts_${CONFCMP_COMPUTE_UNIT}_${CONFCMP_OP_NAME})
394 endif()414 endif()
395 415 
396 if(NOT TARGET gen_opc_info_${CONFCMP_COMPUTE_UNIT})416 if(NOT TARGET gen_opc_info_${CONFCMP_COMPUTE_UNIT})
@@ -418,9 +438,10 @@ function(prepare_compile_from_config)
418 add_custom_target(prepare_binary_compile_${CONFCMP_COMPUTE_UNIT})438 add_custom_target(prepare_binary_compile_${CONFCMP_COMPUTE_UNIT})
419 endif()439 endif()
420 440 
441+ file(MAKE_DIRECTORY ${CONFCMP_OP_PYTHON_DIR})
421 add_custom_target(${CONFCMP_TARGET}442 add_custom_target(${CONFCMP_TARGET}
422 COMMAND cp -r ${CONFCMP_IMPL_DIR}/*.* ${CONFCMP_OUT_DIR}/src443 COMMAND cp -r ${CONFCMP_IMPL_DIR}/*.* ${CONFCMP_OUT_DIR}/src
423- COMMAND cp ${CONFCMP_OP_PYTHON_DIR}/*.py ${CONFCMP_OUT_DIR}/src444+ COMMAND ${CMAKE_COMMAND} -E copy_directory ${CONFCMP_OP_PYTHON_DIR} ${CONFCMP_OUT_DIR}/src
424 )445 )
425 add_dependencies(prepare_binary_compile_${CONFCMP_COMPUTE_UNIT} config_compile_${CONFCMP_COMPUTE_UNIT}_${CONFCMP_OP_NAME} ${CONFCMP_TARGET})446 add_dependencies(prepare_binary_compile_${CONFCMP_COMPUTE_UNIT} config_compile_${CONFCMP_COMPUTE_UNIT}_${CONFCMP_OP_NAME} ${CONFCMP_TARGET})
426 447 
@@ -612,7 +633,7 @@ function(gen_ops_info_and_python)
612 set(HAS_OP_COMPILE_OF_COMPUTE_UNIT FALSE)633 set(HAS_OP_COMPILE_OF_COMPUTE_UNIT FALSE)
613 foreach(OP_DIR ${COMPILED_OP_DIRS})634 foreach(OP_DIR ${COMPILED_OP_DIRS})
614 get_op_type_and_validate("${OP_DIR}" "${compute_unit}" op_name op_type is_valid)635 get_op_type_and_validate("${OP_DIR}" "${compute_unit}" op_name op_type is_valid)
615- set(binary_json ${OP_DIR}/op_host/config/${compute_unit}/${op_name}_binary.json)636+ set(binary_json "${OP_DIR}/op_host/config/${compute_unit}/${op_name}_binary.json")
616 if(NOT is_valid)637 if(NOT is_valid)
617 continue()638 continue()
618 endif()639 endif()
@@ -629,7 +650,7 @@ function(gen_ops_info_and_python)
629 generate_bin_scripts(650 generate_bin_scripts(
630 TARGET gen_bin_scripts651 TARGET gen_bin_scripts
631 OP_NAME ${op_name}652 OP_NAME ${op_name}
632- OP_TYPE ${op_type}653+ OP_TYPE "${op_type}"
633 OPS_INFO_DIR ${ASCEND_AUTOGEN_PATH}654 OPS_INFO_DIR ${ASCEND_AUTOGEN_PATH}
634 COMPUTE_UNIT ${compute_unit}655 COMPUTE_UNIT ${compute_unit}
635 OUT_DIR ${CMAKE_BINARY_DIR}/binary/${compute_unit}656 OUT_DIR ${CMAKE_BINARY_DIR}/binary/${compute_unit}
@@ -639,7 +660,7 @@ function(gen_ops_info_and_python)
639 prepare_compile_from_config(660 prepare_compile_from_config(
640 TARGET ascendc_bin_${compute_unit}_${op_name}661 TARGET ascendc_bin_${compute_unit}_${op_name}
641 OP_NAME ${op_name}662 OP_NAME ${op_name}
642- OP_TYPE ${op_type}663+ OP_TYPE "${op_type}"
643 BINARY_JSON ${binary_json}664 BINARY_JSON ${binary_json}
644 OPS_INFO_DIR ${ASCEND_AUTOGEN_PATH}665 OPS_INFO_DIR ${ASCEND_AUTOGEN_PATH}
645 IMPL_DIR ${OP_DIR}/op_kernel666 IMPL_DIR ${OP_DIR}/op_kernel
@@ -2,9 +2,10 @@
2 2 
3- [RAS aclnn Interface List](op_api_list.md)3- [RAS aclnn Interface List](op_api_list.md)
4- [aclnnCrypto](../../reliability/crypto/docs/aclnnCrypto.md)4- [aclnnCrypto](../../reliability/crypto/docs/aclnnCrypto.md)
5+- [aclnnMatmulAbftVerify](../../reliability/matmul_abft_verify/docs/aclnnMatmulAbftVerify_en.md)
5- [aclnnObfuscationCalculate](../../reliability/obfuscation_calculate/docs/aclnnObfuscationCalculate.md)6- [aclnnObfuscationCalculate](../../reliability/obfuscation_calculate/docs/aclnnObfuscationCalculate.md)
6- [aclnnObfuscationCalculateV2](../../reliability/obfuscation_calculate/docs/aclnnObfuscationCalculateV2.md)7- [aclnnObfuscationCalculateV2](../../reliability/obfuscation_calculate/docs/aclnnObfuscationCalculateV2.md)
7- [aclnnObfuscationSetup](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetup.md)8- [aclnnObfuscationSetup](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetup.md)
8- [aclnnObfuscationSetupV2](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetupV2.md)9- [aclnnObfuscationSetupV2](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetupV2.md)
9- [Compile and Run Sample](context/compile_and_run_sample.md)10- [Compile and Run Sample](context/compile_and_run_sample.md)
10-- [aclnn Return Codes](context/aclnn_return_code.md)11+- [aclnn Return Codes](context/aclnn_return_code.md)
@@ -33,4 +33,5 @@ The operator interface list is as follows:
33| [aclnnObfuscationCalculateV2](../../reliability/obfuscation_calculate/docs/aclnnObfuscationCalculateV2.md) | Sends the tensor x and configuration parameters (such as param, cmd) to the PMCC obfuscation engine. The CA module of the engine calls the TA module to perform tensor obfuscation processing, and finally returns an obfuscated tensor y with the same shape as x. |Default deterministic implementation| - |33| [aclnnObfuscationCalculateV2](../../reliability/obfuscation_calculate/docs/aclnnObfuscationCalculateV2.md) | Sends the tensor x and configuration parameters (such as param, cmd) to the PMCC obfuscation engine. The CA module of the engine calls the TA module to perform tensor obfuscation processing, and finally returns an obfuscated tensor y with the same shape as x. |Default deterministic implementation| - |
34| [aclnnObfuscationSetup](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetup.md) | Completes the resource initialization and release of the PMCC model obfuscation engine. |Default deterministic implementation| - |34| [aclnnObfuscationSetup](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetup.md) | Completes the resource initialization and release of the PMCC model obfuscation engine. |Default deterministic implementation| - |
35| [aclnnObfuscationSetupV2](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetupV2.md) | Completes the resource initialization and release of the PMCC model obfuscation engine. |Default deterministic implementation| - |35| [aclnnObfuscationSetupV2](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetupV2.md) | Completes the resource initialization and release of the PMCC model obfuscation engine. |Default deterministic implementation| - |
36-| [aclnnCrypto](../../reliability/crypto/docs/aclnnCrypto.md) | Calls the aicpu encryption/decryption operator and executes according to the input parameters. |Default deterministic implementation| - |36+| [aclnnCrypto](../../reliability/crypto/docs/aclnnCrypto.md) | Calls the aicpu encryption/decryption operator and executes according to the input parameters. |Default deterministic implementation| - |
37+| [aclnnMatmulAbftVerify](../../reliability/matmul_abft_verify/docs/MatmulAbftVerify_en.md) | Performs V-ABFT fault detection on a precomputed matrix multiplication result and outputs per-row detection results. |Default deterministic implementation| - |
@@ -27,6 +27,15 @@ All operator categories and operator lists provided by the project are as follow
27 <th>op_graph</th>27 <th>op_graph</th>
28 </tr></thead>28 </tr></thead>
29<tbody>29<tbody>
30- 30+ <tr>
31+ <td>Reliability</td>
32+ <td><a href="../../reliability/matmul_abft_verify">matmul_abft_verify</a></td>
33+ <td>√</td>
34+ <td>√</td>
35+ <td>√</td>
36+ <td>-</td>
37+ <td>AI Core</td>
38+ <td>Performs fault detection on a precomputed matrix multiplication result based on V-ABFT.</td>
39+ </tr>
31</tbody>40</tbody>
32-</table>41+</table>
@@ -2,6 +2,7 @@
2 2 
3- [Ras类aclnn接口列表](op_api_list.md)3- [Ras类aclnn接口列表](op_api_list.md)
4- [aclnnCrypto](../../reliability/crypto/docs/aclnnCrypto.md)4- [aclnnCrypto](../../reliability/crypto/docs/aclnnCrypto.md)
5+- [aclnnMatmulAbftVerify](../../reliability/matmul_abft_verify/docs/aclnnMatmulAbftVerify.md)
5- [aclnnObfuscationCalculate](../../reliability/obfuscation_calculate/docs/aclnnObfuscationCalculate.md)6- [aclnnObfuscationCalculate](../../reliability/obfuscation_calculate/docs/aclnnObfuscationCalculate.md)
6- [aclnnObfuscationCalculateV2](../../reliability/obfuscation_calculate/docs/aclnnObfuscationCalculateV2.md)7- [aclnnObfuscationCalculateV2](../../reliability/obfuscation_calculate/docs/aclnnObfuscationCalculateV2.md)
7- [aclnnObfuscationSetup](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetup.md)8- [aclnnObfuscationSetup](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetup.md)
@@ -34,3 +34,4 @@
34| [aclnnObfuscationSetup](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetup.md) | 完成PMCC模型混淆引擎的资源初始化和释放。 |默认确定性实现| - |34| [aclnnObfuscationSetup](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetup.md) | 完成PMCC模型混淆引擎的资源初始化和释放。 |默认确定性实现| - |
35| [aclnnObfuscationSetupV2](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetupV2.md) | 完成PMCC模型混淆引擎的资源初始化和释放。 |默认确定性实现| - |35| [aclnnObfuscationSetupV2](../../reliability/obfuscation_setup/docs/aclnnObfuscationSetupV2.md) | 完成PMCC模型混淆引擎的资源初始化和释放。 |默认确定性实现| - |
36| [aclnnCrypto](../../reliability/crypto/docs/aclnnCrypto.md) | 调用aicpu加解密算子,按输入参数执行。 |默认确定性实现| - |36| [aclnnCrypto](../../reliability/crypto/docs/aclnnCrypto.md) | 调用aicpu加解密算子,按输入参数执行。 |默认确定性实现| - |
37+| [aclnnMatmulAbftVerify](../../reliability/matmul_abft_verify/docs/MatmulAbftVerify.md) | 基于V-ABFT对预先计算的矩阵乘结果进行容错检测,并输出逐行检测结果。 |默认确定性实现| - |
@@ -27,6 +27,15 @@
27 <th>op_graph</th>27 <th>op_graph</th>
28 </tr></thead>28 </tr></thead>
29<tbody>29<tbody>
30- 30+ <tr>
31+ <td>可靠性</td>
32+ <td><a href="../../reliability/matmul_abft_verify">matmul_abft_verify</a></td>
33+ <td>√</td>
34+ <td>√</td>
35+ <td>√</td>
36+ <td>-</td>
37+ <td>AI Core</td>
38+ <td>基于V-ABFT对预先计算的矩阵乘结果进行容错检测。</td>
39+ </tr>
31</tbody>40</tbody>
32</table>41</table>
@@ -0,0 +1,19 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2025 Huawei Technologies Co., Ltd.
3+# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+ 
11+file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
12+if(NOT ENABLE_TEST AND NOT BENCHMARK)
13+ list(REMOVE_ITEM CURRENT_DIRS tests)
14+endif()
15+foreach(SUB_DIR ${CURRENT_DIRS})
16+ if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
17+ add_subdirectory(${SUB_DIR})
18+ endif()
19+endforeach()
@@ -0,0 +1,529 @@
1+# aclnnMatmulAbftVerify
2+ 
3+## 产品支持情况
4+ 
5+<!-- npu="950" id1 -->
6+- <term>Ascend 950PR/Ascend 950DT</term>:不支持
7+<!-- end id1 -->
8+<!-- npu="A3" id2 -->
9+- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持
10+<!-- end id2 -->
11+<!-- npu="910b" id3 -->
12+- <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:支持
13+<!-- end id3 -->
14+<!-- npu="310b" id4 -->
15+- <term>Atlas 200I/500 A2 推理产品</term>:不支持
16+<!-- end id4 -->
17+<!-- npu="310p" id5 -->
18+- <term>Atlas 推理系列产品</term>:不支持
19+<!-- end id5 -->
20+<!-- npu="910" id6 -->
21+- <term>Atlas 训练系列产品</term>:不支持
22+<!-- end id6 -->
23+ 
24+## 功能说明
25+ 
26+- 接口功能:实现基于方差估计自适应门限算法(V-ABFT)的GEMM容错检测算子。算子接收矩阵A、B以及预先计算的矩阵乘结果C,对C进行分块ABFT校验,检测静默计算错误并输出逐行检测结果。
27+ 
28+- 特点
29+ 
30+ - 自适应阈值: 该算子能够根据矩阵大小与值域自动确定用于比对的阈值,能够在保证检出率的同时避免误检。
31+ - 自适应阈值算法与推导见https://gitee.com/yihenggao/v-abft的文档
32+ 
33+ - 计算量显著小于基于重新计算的容错方案。在矩阵维度m=n=k=a时,本算子只需8a^2次浮点计算,而重算则需要2a^3次浮点计算。且容错阈值也与该算法相匹配
34+ 
35+ 
36+- 计算公式:
37+ 
38+ $$
39+ C = A \times B, \quad A \in \mathbb{R}^{M \times K},\; B \in \mathbb{R}^{K \times N}
40+ $$
41+ 
42+ $$
43+ C^r = C\times r, B^r=B\times r
44+ $$
45+ 
46+ 
47+ 其中阈值$Threshold_i$由输入矩阵A、B的局部统计特征(均值、标准差)动态估计,无需依赖C矩阵输出结果。
48+ 
49+- 算子功能说明:
50+ 
51+ - 输入矩阵A、B和预先计算的C经过checksum编码、阈值估计和校验比对流程,输出压缩后的逐行故障检测结果张量`comp_row`,1表示正确,0表示检测到错误。
52+ 
53+## 函数原型
54+ 
55+每个算子分为[两段式接口](../../../docs/zh/context/two_phase_api.md),必须先调用“aclnnMatmulAbftVerifyGetWorkspaceSize”接口获取入参并根据计算流程计算所需workspace大小,再调用“aclnnMatmulAbftVerify”接口执行计算。
56+ 
57+```cpp
58+aclnnStatus aclnnMatmulAbftVerifyGetWorkspaceSize(
59+ const aclTensor *a,
60+ const aclTensor *b,
61+ const aclTensor *c,
62+ const aclTensor *checksumWeight,
63+ double eMax,
64+ const aclTensor *compRow,
65+ uint64_t *workspaceSize,
66+ aclOpExecutor **executor);
67+```
C
Ccaiwenwen21 天前

星号对齐

likedislike
68+ 
69+```cpp
70+aclnnStatus aclnnMatmulAbftVerify(
71+ void *workspace,
72+ uint64_t workspaceSize,
73+ aclOpExecutor *executor,
74+ aclrtStream stream);
75+```
76+ 
77+## aclnnMatmulAbftVerifyGetWorkspaceSize
78+ 
79+- **参数说明:**
80+ 
81+ <table style="undefined;table-layout: fixed;width: 1540px"><colgroup>
82+ <col style="width: 170px">
83+ <col style="width: 120px">
84+ <col style="width: 300px">
85+ <col style="width: 330px">
86+ <col style="width: 212px">
87+ <col style="width: 100px">
88+ <col style="width: 190px">
89+ <col style="width: 118px">
90+ </colgroup>
91+ <thead>
92+ <tr>
93+ <th>参数名</th>
94+ <th style="white-space: nowrap">输入/输出</th>
95+ <th>描述</th>
96+ <th>使用说明</th>
97+ <th>数据类型</th>
98+ <th><a href="../../../docs/zh/context/数据格式.md" target="_blank">数据格式</a></th>
99+ <th style="white-space: nowrap">维度(shape)</th>
100+ <th><a href="../../../docs/zh/context/非连续的Tensor.md" target="_blank">非连续的Tensor</a></th>
101+ </tr>
102+ </thead>
103+ <tbody>
104+ <tr>
105+ <td>a(aclTensor)</td>
106+ <td>输入</td>
107+ <td>矩阵乘法输入A。</td>
108+ <td>
109+ <ul>
110+ <li>维度为2,shape为[M, K]。</li>
111+ </ul>
112+ </td>
113+ <td>FLOAT16、BFLOAT16、FLOAT32</td>
114+ <td>ND</td>
115+ <td>[M, K]</td>
116+ <td>-</td>
117+ </tr>
118+ <tr>
119+ <td>b(aclTensor)</td>
120+ <td>输入</td>
121+ <td>矩阵乘法输入B。</td>
122+ <td>
123+ <ul>
124+ <li>维度为2,shape为[K, N]。</li>
125+ </ul>
126+ </td>
127+ <td>FLOAT16、BFLOAT16、FLOAT32</td>
128+ <td>ND</td>
129+ <td>[K, N]</td>
130+ <td>-</td>
131+ </tr>
132+ <tr>
133+ <td>c(aclTensor)</td>
134+ <td>输入</td>
135+ <td>预先计算的矩阵乘结果C = A × B,作为容错检测的数据输入。</td>
136+ <td>
137+ <ul>
138+ <li>维度为2,shape为[M, N]。</li>
139+ </ul>
140+ </td>
141+ <td>FLOAT32</td>
142+ <td>ND</td>
143+ <td>[M, N]</td>
144+ <td>-</td>
145+ </tr>
146+ <tr>
147+ <td>checksumWeight(aclTensor)</td>
148+ <td>输入</td>
149+ <td>列编码值向量,对应ABFT中的行校验和编码向量$r$(加权向量)。</td>
150+ <td>
151+ <ul>
152+ <li>维度为1,shape为[N]。</li>
153+ </ul>
154+ </td>
155+ <td>FLOAT16、BFLOAT16、FLOAT32</td>
156+ <td>ND</td>
157+ <td>[N]</td>
158+ <td>-</td>
159+ </tr>
160+ <tr>
161+ <td>eMax(double)</td>
162+ <td>输入</td>
163+ <td>误差阈值系数,控制故障检测的灵敏度。</td>
164+ <td>
165+ <ul>
166+ <li>默认值为0.001。</li>
167+ <li>A,B为BF16精度,推荐值:0.001</li>
168+ <li>A,B为FP32精度, 推荐值: 0.00002</li>
169+ <li>使用小于推荐值的eMax会增大检出率,但是也可能会出现误检情况。</li>
170+ <li>在出现误报的时候,可以将eMax调大。</li>
171+ </ul>
172+ </td>
173+ <td>-</td>
174+ <td>-</td>
175+ <td>-</td>
176+ <td>-</td>
177+ </tr>
178+ <tr>
179+ <td>compRow(aclTensor)</td>
180+ <td>输出</td>
181+ <td>压缩后的行方向故障检测位流输出。</td>
182+ <td>
183+ <ul>
184+ <li>每个bit表示一行一段的检测结果,1表示正确,0表示检测到错误。</li>
185+ </ul>
186+ </td>
187+ <td>UINT8</td>
188+ <td>ND</td>
189+ <td>[ceil(M/8) * splitN]</td>
190+ <td>-</td>
191+ </tr>
192+ <tr>
193+ <td>workspaceSize(uint64_t)</td>
194+ <td>输出</td>
195+ <td>返回需要在Device侧申请的workspace大小。</td>
196+ <td>-</td>
197+ <td>-</td>
198+ <td>-</td>
199+ <td>-</td>
200+ <td>-</td>
201+ </tr>
202+ <tr>
203+ <td>executor(aclOpExecutor)</td>
204+ <td>输出</td>
205+ <td>返回op执行器,包含了算子计算流程。</td>
206+ <td>-</td>
207+ <td>-</td>
208+ <td>-</td>
209+ <td>-</td>
210+ <td>-</td>
211+ </tr>
212+ </tbody></table>
213+ 
214+ 其中 $splitN = \lceil N / 256\rceil$。
215+ 
216+- **返回值:**
217+ 
218+ 返回aclnnStatus状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。
219+ 
220+ 第一阶段接口完成入参校验,出现以下场景时报错:
221+ 
222+ <table style="undefined;table-layout: fixed;width: 1030px"><colgroup>
223+ <col style="width: 250px">
224+ <col style="width: 130px">
225+ <col style="width: 650px">
226+ </colgroup>
227+ <thead>
228+ <tr>
229+ <th>返回值</th>
230+ <th>错误码</th>
231+ <th>描述</th>
232+ </tr>
233+ </thead>
234+ <tbody>
235+ <tr>
236+ <td>ACLNN_ERR_PARAM_NULLPTR</td>
237+ <td>161001</td>
238+ <td>必选输入、输出或者必选属性是空指针。</td>
239+ </tr>
240+ <tr>
241+ <td rowspan="5">ACLNN_ERR_PARAM_INVALID</td>
242+ <td rowspan="5">161002</td>
243+ <td>a、b、c、weight、eMax、beSplitFactor、compRow的数据类型和数据格式不在支持的范围内。</td>
244+ </tr>
245+ <tr>
246+ <td>a和b的维度不为2。</td>
247+ </tr>
248+ <tr>
249+ <td>a的第1维(K)与b的第0维(K)不相等。</td>
250+ </tr>
251+ <tr>
252+ <td>eMax为负数。</td>
253+ </tr>
254+ <tr>
255+ <td>reduceCores小于0。</td>
256+ </tr>
257+ </tbody></table>
258+ 
259+## aclnnMatmulAbftVerify
260+ 
261+- **参数说明:**
262+ <table>
263+ <thead>
264+ <tr><th>参数名</th><th>输入/输出</th><th>描述</th></tr>
265+ </thead>
266+ <tbody>
267+ <tr><td>workspace</td><td>输入</td><td>在Device侧申请的workspace内存地址。</td></tr>
268+ <tr><td>workspaceSize</td><td>输入</td><td>在Device侧申请的workspace大小,由第一段接口aclnnMatmulAbftVerifyGetWorkspaceSize获取。</td></tr>
269+ <tr><td>executor</td><td>输入</td><td>op执行器,包含了算子计算流程。</td></tr>
270+ <tr><td>stream</td><td>输入</td><td>指定执行任务的Stream。</td></tr>
271+ </tbody>
272+ </table>
273+ 
274+- **返回值:**
275+ 
276+ 返回aclnnStatus状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。
277+ 
278+## 约束说明
279+ 
280+- 确定性计算:
281+ - aclnnMatmulAbftVerify默认确定性实现。
282+- 输入矩阵a、b和c必须为2维,shape分别为[M, K]、[K, N]和[M, N],且a的第1维(K)与b的第0维(K)必须相等。
283+- 输入向量weight的shape必须为[N]。
284+- 支持的数据类型组合为:
285+ 
286+ | a | b | c | checksumWeight |
287+ |:-------:|:-------:|:-------:|:-------:|
288+ | FLOAT16 | FLOAT16 | FLOAT32 | FLOAT16 |
289+ | BFLOAT16| BFLOAT16| FLOAT32 | BFLOAT16|
290+ | FLOAT32 | FLOAT32 | FLOAT32 | FLOAT32 |
291+ 
292+- 不支持的场景:
293+ - 不支持ND格式以外的数据格式。
294+ - 不支持非连续的Tensor。
295+ - 不支持a、b和c的维度不为2的场景。
296+ 
297+## Workspace使用设计
298+ 
299+workspace由两部分组成:
300+ 
301+1. 固定的16 MiB系统workspace。
302+2. user workspace,其中依次放置13个内部中间张量及FT内部临时区。
303+ 
304+算子tiling阶段根据M、N、K和输入精度计算完整大小,并通过`workspaceSize`返回。调用者必须申请不少于该大小的连续device内存,不能只按`compRow`大小申请。所有中间张量的起始地址按32字节向上对齐。
305+ 
306+定义:
307+ 
308+```text
309+splitN = ceil(N / 256)
310+rowSplitLen = M * splitN
311+bStatLen = ceil(splitN / 8) * 8 + 8
312+beLen = K * splitN
313+align32(x) = ceil(x / 32) * 32
314+```
315+ 
316+设`W`为user workspace首地址。在kernel内,`W = AscendC::GetUserWorkspace(workspace)`;对于算子外部的device地址计算,当前实现等价于`W = (uint8_t *)workspace + 16 MiB`。
317+ 
318+以下偏移均相对`W`。令`S0 = 0`,每个张量的起始偏移为`Oi = align32(Si)`,结束位置为`Si+1 = Oi + 元素数 × 元素字节数`:
319+ 
320+| 顺序 | 中间结果 | 起始地址 | 元素类型 | 元素数 |
321+|:--:|:--|:--|:--|--:|
322+| 0 | z_row | `W + O0` | FLOAT32 | `rowSplitLen` |
323+| 1 | d_row | `W + O1` | FLOAT32 | `rowSplitLen` |
324+| 2 | threshold | `W + O2` | FLOAT32 | `rowSplitLen` |
325+| 3 | b_mean_abs | `W + O3` | FLOAT32 | `bStatLen` |
326+| 4 | b_mean_square | `W + O4` | FLOAT32 | `bStatLen` |
327+| 5 | b_var | `W + O5` | FLOAT32 | `bStatLen` |
328+| 6 | be | `W + O6` | 与a相同 | `beLen` |
329+| 7 | be_for_aiv | `W + O7` | a为FLOAT16时是FLOAT16,否则是FLOAT32 | `beLen` |
330+| 8 | b_max_slice | `W + O8` | FLOAT32 | `beLen` |
331+| 9 | b_min_slice | `W + O9` | FLOAT32 | `beLen` |
332+| 10 | a_max | `W + O10` | FLOAT32 | `M` |
333+| 11 | a_mean | `W + O11` | FLOAT32 | `M` |
334+| 12 | a_min | `W + O12` | FLOAT32 | `M` |
335+其中FLOAT16和BFLOAT16元素占2字节,FLOAT32元素占4字节。第13段之后再次按32字节对齐,剩余区域是FT内部临时区,其逻辑大小为:
336+ 
337+```text
338+M * (splitN + 1) * sizeof(float)
339+```
340+ 
341+AMean计算使用的常量因子`1/K`不占用workspace。tiling阶段按输入精度生成该scalar,kernel在首次计算AMean时直接用它初始化FT的L1内部缓冲区,因此该因子不能作为workspace中间结果拷出。
342+ 
343+如果需要调试并拷出某个中间结果,应在算子执行完成且workspace尚未释放或复用时,从上表对应的`W + Oi`开始执行device-to-host拷贝,拷贝字节数为“元素数 × 元素字节数”。对外读取workspace属于调试能力,不是稳定的公开输出接口;布局发生变更时,以[matmul_abft_verify.cpp](../op_kernel/matmul_abft_verify.cpp)中`Workspace ABI`注释下的地址切分代码为准,该代码也是偏移计算的示例实现。
344+ 
345+## 调用示例
346+ 
347+下面以FLOAT32输入为例。示例先调用`aclnnGemm`计算`C = A × B`,再将矩阵乘结果`C`传给`aclnnMatmulAbftVerify`进行校验。
348+ 
349+```cpp
350+#include <cstdint>
351+#include <cstdio>
352+#include <vector>
353+ 
354+#include "acl/acl.h"
355+#include "aclnnop/aclnn_gemm.h"
356+#include "aclnnop/aclnn_matmul_abft_verify.h"
357+ 
358+#define CHECK_RET(cond, action) \
359+ do { \
360+ if (!(cond)) { \
361+ action; \
362+ } \
363+ } while (0)
364+ 
365+int64_t GetElementCount(const std::vector<int64_t>& shape)
366+{
367+ int64_t count = 1;
368+ for (int64_t dim : shape) {
369+ count *= dim;
370+ }
371+ return count;
372+}
373+ 
374+template <typename T>
375+int CreateAclTensor(const std::vector<T>& hostData,
376+ const std::vector<int64_t>& shape,
377+ aclDataType dataType,
378+ void** deviceAddr,
379+ aclTensor** tensor)
380+{
381+ const uint64_t bytes = static_cast<uint64_t>(hostData.size()) * sizeof(T);
382+ auto ret = aclrtMalloc(deviceAddr, bytes, ACL_MEM_MALLOC_HUGE_FIRST);
383+ CHECK_RET(ret == ACL_SUCCESS, return ret);
384+ 
385+ ret = aclrtMemcpy(*deviceAddr, bytes, hostData.data(), bytes,
386+ ACL_MEMCPY_HOST_TO_DEVICE);
387+ CHECK_RET(ret == ACL_SUCCESS, return ret);
388+ 
389+ std::vector<int64_t> strides(shape.size(), 1);
390+ for (int64_t i = static_cast<int64_t>(shape.size()) - 2; i >= 0; --i) {
391+ strides[i] = shape[i + 1] * strides[i + 1];
392+ }
393+ 
394+ *tensor = aclCreateTensor(shape.data(), shape.size(), dataType,
395+ strides.data(), 0, ACL_FORMAT_ND,
396+ shape.data(), shape.size(), *deviceAddr);
397+ CHECK_RET(*tensor != nullptr, return ACL_ERROR_FAILURE);
398+ return ACL_SUCCESS;
399+}
400+ 
401+int main()
402+{
403+ constexpr int32_t deviceId = 0;
404+ constexpr int64_t M = 4096;
405+ constexpr int64_t N = 4096;
406+ constexpr int64_t K = 4096;
407+ constexpr double eMax = 0.00002;
408+ 
409+ const int64_t splitN = (N + 255) / 256;
410+ const std::vector<int64_t> aShape{M, K};
411+ const std::vector<int64_t> bShape{K, N};
412+ const std::vector<int64_t> cShape{M, N};
413+ const std::vector<int64_t> checksumWeightShape{N};
414+ const std::vector<int64_t> compRowShape{((M + 7) / 8) * splitN};
415+ 
416+ // checksumWeight全为1时,对应普通的分块行校验和。
417+ std::vector<float> hostA(GetElementCount(aShape), 0.5F);
418+ std::vector<float> hostB(GetElementCount(bShape), 0.25F);
419+ std::vector<float> hostC(GetElementCount(cShape), 0.0F);
420+ std::vector<float> hostChecksumWeight(GetElementCount(checksumWeightShape), 1.0F);
421+ std::vector<uint8_t> hostCompRow(GetElementCount(compRowShape), 0);
422+ 
423+ aclrtContext context = nullptr;
424+ aclrtStream stream = nullptr;
425+ auto ret = aclInit(nullptr);
426+ CHECK_RET(ret == ACL_SUCCESS, return ret);
427+ ret = aclrtSetDevice(deviceId);
428+ CHECK_RET(ret == ACL_SUCCESS, return ret);
429+ ret = aclrtCreateContext(&context, deviceId);
430+ CHECK_RET(ret == ACL_SUCCESS, return ret);
431+ ret = aclrtSetCurrentContext(context);
432+ CHECK_RET(ret == ACL_SUCCESS, return ret);
433+ ret = aclrtCreateStream(&stream);
434+ CHECK_RET(ret == ACL_SUCCESS, return ret);
435+ 
436+ void* devA = nullptr;
437+ void* devB = nullptr;
438+ void* devC = nullptr;
439+ void* devGemmAddend = nullptr;
440+ void* devChecksumWeight = nullptr;
441+ void* devCompRow = nullptr;
442+ aclTensor* tensorA = nullptr;
443+ aclTensor* tensorB = nullptr;
444+ aclTensor* tensorC = nullptr;
445+ aclTensor* tensorGemmAddend = nullptr;
446+ aclTensor* tensorChecksumWeight = nullptr;
447+ aclTensor* tensorCompRow = nullptr;
448+ 
449+ ret = CreateAclTensor(hostA, aShape, ACL_FLOAT, &devA, &tensorA);
450+ CHECK_RET(ret == ACL_SUCCESS, return ret);
451+ ret = CreateAclTensor(hostB, bShape, ACL_FLOAT, &devB, &tensorB);
452+ CHECK_RET(ret == ACL_SUCCESS, return ret);
453+ ret = CreateAclTensor(hostC, cShape, ACL_FLOAT, &devC, &tensorC);
454+ CHECK_RET(ret == ACL_SUCCESS, return ret);
455+ ret = CreateAclTensor(hostC, cShape, ACL_FLOAT,
456+ &devGemmAddend, &tensorGemmAddend);
457+ CHECK_RET(ret == ACL_SUCCESS, return ret);
458+ ret = CreateAclTensor(hostChecksumWeight, checksumWeightShape, ACL_FLOAT,
459+ &devChecksumWeight, &tensorChecksumWeight);
460+ CHECK_RET(ret == ACL_SUCCESS, return ret);
461+ ret = CreateAclTensor(hostCompRow, compRowShape, ACL_UINT8,
462+ &devCompRow, &tensorCompRow);
463+ CHECK_RET(ret == ACL_SUCCESS, return ret);
464+ 
465+ // 前序矩阵乘法:C = 1.0 * A * B + 0.0 * gemmAddend。
466+ uint64_t gemmWorkspaceSize = 0;
467+ aclOpExecutor* gemmExecutor = nullptr;
468+ ret = aclnnGemmGetWorkspaceSize(
469+ tensorA, tensorB, tensorGemmAddend,
470+ 1.0F, 0.0F, 0, 0, tensorC, 0,
471+ &gemmWorkspaceSize, &gemmExecutor);
472+ CHECK_RET(ret == ACL_SUCCESS, return ret);
473+ 
474+ void* gemmWorkspace = nullptr;
475+ if (gemmWorkspaceSize > 0) {
476+ ret = aclrtMalloc(&gemmWorkspace, gemmWorkspaceSize,
477+ ACL_MEM_MALLOC_HUGE_FIRST);
478+ CHECK_RET(ret == ACL_SUCCESS, return ret);
479+ }
480+ ret = aclnnGemm(gemmWorkspace, gemmWorkspaceSize, gemmExecutor, stream);
481+ CHECK_RET(ret == ACL_SUCCESS, return ret);
482+ ret = aclrtSynchronizeStream(stream);
483+ CHECK_RET(ret == ACL_SUCCESS, return ret);
484+ 
485+ // 对前序GEMM的结果进行ABFT校验。
486+ uint64_t workspaceSize = 0;
487+ aclOpExecutor* executor = nullptr;
488+ ret = aclnnMatmulAbftVerifyGetWorkspaceSize(
489+ tensorA, tensorB, tensorC, tensorChecksumWeight,
490+ eMax, tensorCompRow, &workspaceSize, &executor);
491+ CHECK_RET(ret == ACL_SUCCESS, return ret);
492+ 
493+ void* workspace = nullptr;
494+ if (workspaceSize > 0) {
495+ ret = aclrtMalloc(&workspace, workspaceSize,
496+ ACL_MEM_MALLOC_HUGE_FIRST);
497+ CHECK_RET(ret == ACL_SUCCESS, return ret);
498+ }
499+ ret = aclnnMatmulAbftVerify(workspace, workspaceSize, executor, stream);
500+ CHECK_RET(ret == ACL_SUCCESS, return ret);
501+ ret = aclrtSynchronizeStream(stream);
502+ CHECK_RET(ret == ACL_SUCCESS, return ret);
503+ 
504+ aclDestroyTensor(tensorA);
505+ aclDestroyTensor(tensorB);
506+ aclDestroyTensor(tensorC);
507+ aclDestroyTensor(tensorGemmAddend);
508+ aclDestroyTensor(tensorChecksumWeight);
509+ aclDestroyTensor(tensorCompRow);
510+ aclrtFree(devA);
511+ aclrtFree(devB);
512+ aclrtFree(devC);
513+ aclrtFree(devGemmAddend);
514+ aclrtFree(devChecksumWeight);
515+ aclrtFree(devCompRow);
516+ if (gemmWorkspace != nullptr) {
517+ aclrtFree(gemmWorkspace);
518+ }
519+ if (workspace != nullptr) {
520+ aclrtFree(workspace);
521+ }
522+ aclrtDestroyStream(stream);
523+ aclrtDestroyContext(context);
524+ aclrtResetDevice(deviceId);
525+ aclFinalize();
526+ return ACL_SUCCESS;
527+}
528+ 
529+```
@@ -0,0 +1,188 @@
1+# aclnnMatmulAbftVerify
2+ 
3+## Product Support
4+ 
5+<!-- npu="950" id1 -->
6+- <term>Ascend 950PR/Ascend 950DT</term>: Not supported
7+<!-- end id1 -->
8+<!-- npu="A3" id2 -->
9+- <term>Atlas A3 Training Series/Atlas A3 Inference Series</term>: Supported
10+<!-- end id2 -->
11+<!-- npu="910b" id3 -->
12+- <term>Atlas A2 Training Series/Atlas A2 Inference Series</term>: Supported
13+<!-- end id3 -->
14+<!-- npu="310b" id4 -->
15+- <term>Atlas 200I/500 A2 Inference Products</term>: Not supported
16+<!-- end id4 -->
17+<!-- npu="310p" id5 -->
18+- <term>Atlas Inference Series</term>: Not supported
19+<!-- end id5 -->
20+<!-- npu="910" id6 -->
21+- <term>Atlas Training Series</term>: Not supported
22+<!-- end id6 -->
23+ 
24+## Overview
25+ 
26+- Function: Implements a GEMM fault-detection operator based on the variance-estimation adaptive-threshold algorithm (V-ABFT). The operator accepts matrices A and B and the precomputed matrix multiplication result C, performs block-wise ABFT verification on C, detects silent data corruption, and outputs a per-row detection result.
27+ 
28+- Features:
29+ 
30+ - Adaptive threshold: The operator automatically determines the comparison threshold based on the matrix dimensions and value range. This maintains the detection rate while avoiding false positives. For details about the adaptive-threshold algorithm and its derivation, see the [V-ABFT documentation](https://gitee.com/yihenggao/v-abft).
31+ - Significantly less computation than fault-tolerance schemes based on recomputation. When `M = N = K = a`, this operator requires only `8a^2` floating-point operations, whereas recomputation requires `2a^3` floating-point operations. The fault-detection threshold is also designed for this algorithm.
32+ 
33+- Formulas:
34+ 
35+ $$
36+ C = A \times B, \quad A \in \mathbb{R}^{M \times K},\; B \in \mathbb{R}^{K \times N}
37+ $$
38+ 
39+ $$
40+ C^r = C \times r, \quad B^r = B \times r
41+ $$
42+ 
43+ The threshold $Threshold_i$ is dynamically estimated from local statistics (mean and standard deviation) of input matrices A and B. It does not depend on the output values of matrix C.
44+ 
45+- Operator behavior:
46+ 
47+ - Matrices A and B and the precomputed matrix C pass through checksum encoding, threshold estimation, and verification. The operator outputs the compressed per-row fault-detection tensor `compRow`. A bit value of 1 indicates a correct result, and 0 indicates that an error has been detected.
48+ 
49+## Function Prototypes
50+ 
51+This operator uses a [two-stage interface](../../../docs/en/context/two_phase_api.md). First call `aclnnMatmulAbftVerifyGetWorkspaceSize` to validate the inputs and calculate the required workspace size, and then call `aclnnMatmulAbftVerify` to execute the operator.
52+ 
53+```cpp
54+aclnnStatus aclnnMatmulAbftVerifyGetWorkspaceSize(
55+ const aclTensor *a,
56+ const aclTensor *b,
57+ const aclTensor *c,
58+ const aclTensor *checksumWeight,
59+ double eMax,
60+ const aclTensor *compRow,
61+ uint64_t *workspaceSize,
62+ aclOpExecutor **executor);
63+```
64+ 
65+```cpp
66+aclnnStatus aclnnMatmulAbftVerify(
67+ void *workspace,
68+ uint64_t workspaceSize,
69+ aclOpExecutor *executor,
70+ aclrtStream stream);
71+```
72+ 
73+## aclnnMatmulAbftVerifyGetWorkspaceSize
74+ 
75+- **Parameters:**
76+ 
77+ | Parameter | Input/Output | Description | Usage | Data Type | Format | Shape | Non-contiguous Tensor |
78+ |:--|:--:|:--|:--|:--|:--:|:--|:--:|
79+ | `a` (`aclTensor`) | Input | Input matrix A. | Must be a 2D tensor with shape `[M, K]`. | FLOAT16, BFLOAT16, FLOAT32 | ND | `[M, K]` | No |
80+ | `b` (`aclTensor`) | Input | Input matrix B. | Must be a 2D tensor with shape `[K, N]`. | FLOAT16, BFLOAT16, FLOAT32 | ND | `[K, N]` | No |
81+ | `c` (`aclTensor`) | Input | Precomputed matrix multiplication result `C = A × B`, used as the data to be verified. | Must be a 2D tensor with shape `[M, N]`. | FLOAT32 | ND | `[M, N]` | No |
82+ | `checksumWeight` (`aclTensor`) | Input | Column-encoding vector corresponding to the row-checksum encoding vector $r$ (weight vector) in ABFT. | Must be a 1D tensor with shape `[N]`. | FLOAT16, BFLOAT16, FLOAT32 | ND | `[N]` | No |
83+ | `eMax` (`double`) | Input | Error-threshold coefficient that controls fault-detection sensitivity. | The default value is `0.001`. The recommended value is `0.001` for BF16 inputs and `0.00002` for FP32 inputs. A smaller value increases the detection rate but may cause false positives. Increase this value if false positives occur. | - | - | - | - |
84+ | `compRow` (`aclTensor`) | Output | Compressed row-direction fault-detection bitstream. | Each bit represents the detection result of one row segment. A value of 1 indicates a correct result, and 0 indicates an error. | UINT8 | ND | `[ceil(M/8) * splitN]` | No |
85+ | `workspaceSize` (`uint64_t`) | Output | Size of the workspace that must be allocated on the Device. | - | - | - | - | - |
86+ | `executor` (`aclOpExecutor`) | Output | Operator executor containing the operator execution process. | - | - | - | - | - |
87+ 
88+ Here, $splitN = \lceil N / 256 \rceil$.
89+ 
90+- **Return Value:**
91+ 
92+ Returns an `aclnnStatus` status code. For details, see [aclnn Return Codes](../../../docs/en/context/aclnn_return_code.md). The first-stage interface validates the input arguments and reports an error in the following cases:
93+ 
94+ | Return Value | Error Code | Description |
95+ |:--|:--:|:--|
96+ | ACLNN_ERR_PARAM_NULLPTR | 161001 | A required input, output, or attribute is a null pointer. |
97+ | ACLNN_ERR_PARAM_INVALID | 161002 | The data type or format of `a`, `b`, `c`, `weight`, `eMax`, `beSplitFactor`, or `compRow` is unsupported. |
98+ | ACLNN_ERR_PARAM_INVALID | 161002 | `a` or `b` is not a 2D tensor. |
99+ | ACLNN_ERR_PARAM_INVALID | 161002 | Dimension 1 (K) of `a` differs from dimension 0 (K) of `b`. |
100+ | ACLNN_ERR_PARAM_INVALID | 161002 | `eMax` is negative. |
101+ | ACLNN_ERR_PARAM_INVALID | 161002 | `reduceCores` is less than 0. |
102+ 
103+## aclnnMatmulAbftVerify
104+ 
105+- **Parameters:**
106+ 
107+ | Parameter | Input/Output | Description |
108+ |:--|:--:|:--|
109+ | `workspace` | Input | Address of the workspace allocated on the Device. |
110+ | `workspaceSize` | Input | Size of the workspace allocated on the Device, obtained from `aclnnMatmulAbftVerifyGetWorkspaceSize`. |
111+ | `executor` | Input | Operator executor containing the operator execution process. |
112+ | `stream` | Input | Stream on which the task is executed. |
113+ 
114+- **Return Value:**
115+ 
116+ Returns an `aclnnStatus` status code. For details, see [aclnn Return Codes](../../../docs/en/context/aclnn_return_code.md).
117+ 
118+## Constraints
119+ 
120+- Deterministic computation: `aclnnMatmulAbftVerify` uses a deterministic implementation by default.
121+- Input matrices `a`, `b`, and `c` must be 2D tensors with shapes `[M, K]`, `[K, N]`, and `[M, N]`, respectively. Dimension 1 (K) of `a` must equal dimension 0 (K) of `b`.
122+- The shape of input vector `checksumWeight` must be `[N]`.
123+- The following data type combinations are supported:
124+ 
125+ | a | b | c | checksumWeight |
126+ |:--:|:--:|:--:|:--:|
127+ | FLOAT16 | FLOAT16 | FLOAT32 | FLOAT16 |
128+ | BFLOAT16 | BFLOAT16 | FLOAT32 | BFLOAT16 |
129+ | FLOAT32 | FLOAT32 | FLOAT32 | FLOAT32 |
130+ 
131+- Unsupported scenarios:
132+ 
133+ - Formats other than ND are not supported.
134+ - Non-contiguous tensors are not supported.
135+ - Non-2D `a`, `b`, or `c` tensors are not supported.
136+ 
137+## Workspace Design
138+ 
139+The workspace consists of two parts:
140+ 
141+1. A fixed 16 MiB system workspace.
142+2. A user workspace containing 13 internal intermediate tensors followed by the FT internal temporary area.
143+ 
144+During tiling, the operator calculates the complete size from M, N, K, and the input precision and returns it through `workspaceSize`. The caller must allocate at least this amount of contiguous Device memory. Allocating only enough memory for `compRow` is insufficient. The starting address of every intermediate tensor is aligned upwards to 32 bytes.
145+ 
146+Definitions:
147+ 
148+```text
149+splitN = ceil(N / 256)
150+rowSplitLen = M * splitN
151+bStatLen = ceil(splitN / 8) * 8 + 8
152+beLen = K * splitN
153+align32(x) = ceil(x / 32) * 32
154+```
155+ 
156+Let `W` be the start address of the user workspace. In the kernel, `W = AscendC::GetUserWorkspace(workspace)`. For Device-address calculations outside the operator, the current implementation is equivalent to `W = (uint8_t *)workspace + 16 MiB`.
157+ 
158+All offsets below are relative to `W`. Let `S0 = 0`. The starting offset of each tensor is `Oi = align32(Si)`, and its end position is `Si+1 = Oi + element count × element size`:
159+ 
160+| Order | Intermediate Result | Start Address | Element Type | Element Count |
161+|:--:|:--|:--|:--|--:|
162+| 0 | z_row | `W + O0` | FLOAT32 | `rowSplitLen` |
163+| 1 | d_row | `W + O1` | FLOAT32 | `rowSplitLen` |
164+| 2 | threshold | `W + O2` | FLOAT32 | `rowSplitLen` |
165+| 3 | b_mean_abs | `W + O3` | FLOAT32 | `bStatLen` |
166+| 4 | b_mean_square | `W + O4` | FLOAT32 | `bStatLen` |
167+| 5 | b_var | `W + O5` | FLOAT32 | `bStatLen` |
168+| 6 | be | `W + O6` | Same as `a` | `beLen` |
169+| 7 | be_for_aiv | `W + O7` | FLOAT16 when `a` is FLOAT16; otherwise FLOAT32 | `beLen` |
170+| 8 | b_max_slice | `W + O8` | FLOAT32 | `beLen` |
171+| 9 | b_min_slice | `W + O9` | FLOAT32 | `beLen` |
172+| 10 | a_max | `W + O10` | FLOAT32 | `M` |
173+| 11 | a_mean | `W + O11` | FLOAT32 | `M` |
174+| 12 | a_min | `W + O12` | FLOAT32 | `M` |
175+ 
176+FLOAT16 and BFLOAT16 elements occupy 2 bytes, and FLOAT32 elements occupy 4 bytes. After the 13th segment, the address is aligned upwards to 32 bytes again. The remaining region is the FT internal temporary area, whose logical size is:
177+ 
178+```text
179+M * (splitN + 1) * sizeof(float)
180+```
181+ 
182+The constant factor `1/K` used to calculate AMean does not occupy workspace. During tiling, the scalar is generated according to the input precision. When AMean is first calculated, the kernel uses the scalar directly to initialize the FT L1 internal buffer. Therefore, this factor cannot be copied out as a workspace intermediate result.
183+ 
184+To debug and copy an intermediate result, perform a Device-to-Host copy starting at the corresponding `W + Oi` address in the table after the operator has finished and before the workspace is released or reused. The number of bytes to copy is `element count × element size`. Reading the workspace externally is a debugging capability and is not a stable public output interface. If the layout changes, refer to the address-partitioning code below the `Workspace ABI` comment in [matmul_abft_verify.cpp](../op_kernel/matmul_abft_verify.cpp). That code is also an example of the offset calculation.
185+ 
186+## Example
187+ 
188+The following example uses FLOAT32 inputs. It first calls `aclnnGemm` to calculate `C = A × B`, and then passes C to `aclnnMatmulAbftVerify` for verification. A complete runnable example is provided in [test_aclnn_matmul_abft_verify_bf16.cpp](../examples/test_aclnn_matmul_abft_verify_bf16.cpp).
@@ -0,0 +1,423 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+#include <algorithm>
12+#include <cstdint>
13+#include <cstdio>
14+#include <cstring>
15+#include <random>
16+#include <vector>
17+#include <iostream>
18+ 
19+#include "acl/acl.h"
20+#include "aclnnop/aclnn_matmul_abft_verify.h"
21+ 
22+#define CHECK_RET(cond, return_expr) \
23+ do { \
24+ if (!(cond)) { \
25+ return_expr; \
26+ } \
27+ } while (0)
28+ 
29+#define LOG_PRINT(message, ...) \
30+ do { \
31+ std::printf(message, ##__VA_ARGS__); \
32+ } while (0)
33+ 
34+using Bf16 = uint16_t;
35+ 
36+Bf16 FloatToBf16(float value)
37+{
38+ uint32_t bits = 0;
39+ std::copy_n(reinterpret_cast<const unsigned char *>(&value), sizeof(value),
40+ reinterpret_cast<unsigned char *>(&bits));
41+ bits += 0x7FFFU + ((bits >> 16U) & 1U);
42+ return static_cast<Bf16>(bits >> 16U);
43+}
44+ 
45+float Bf16ToFloat(Bf16 value)
46+{
47+ uint32_t bits = static_cast<uint32_t>(value) << 16U;
48+ float result = 0.0F;
49+ std::copy_n(reinterpret_cast<const unsigned char *>(&bits), sizeof(bits),
50+ reinterpret_cast<unsigned char *>(&result));
51+ return result;
52+}
53+ 
54+struct TestArgs {
55+ int64_t m = 4096;
56+ int64_t n = 4096;
57+ int64_t k = 4096;
58+ int32_t deviceId = 0;
59+ double eMax = 0.001;
60+ int64_t reduceCores = 8;
61+ int32_t verbose = 1;
62+ 
63+ void Print() const
64+ {
65+ LOG_PRINT("Args: m=%ld n=%ld k=%ld eMax=%f reduceCores=%ld "
66+ "verbose=%d deviceId=%d\n",
67+ m, n, k, eMax, reduceCores, verbose, deviceId);
68+ }
69+ 
70+ int Parse(int argc, const char** argv)
71+ {
72+ const std::string helper =
73+ "test_aclnn_matmul_abft_verify_bf16 m n k eMax reduceCores verbose [deviceId]";
74+ enum {
75+ M_IDX = 1,
76+ N_IDX,
77+ K_IDX,
78+ EMAX_IDX,
79+ RED_CORES_IDX,
80+ VERBOSE_IDX,
81+ DEVICE_ID_IDX,
82+ ARGS_MAX
83+ };
84+ if (argc < DEVICE_ID_IDX || argc > ARGS_MAX) {
85+ std::cerr << helper << std::endl;
86+ return -1;
87+ }
88+ m = std::atol(argv[M_IDX]);
89+ n = std::atol(argv[N_IDX]);
90+ k = std::atol(argv[K_IDX]);
91+ eMax = std::stod(argv[EMAX_IDX]);
92+ reduceCores = std::atol(argv[RED_CORES_IDX]);
93+ verbose = std::atoi(argv[VERBOSE_IDX]);
94+ if (argc == ARGS_MAX) {
95+ deviceId = std::atoi(argv[DEVICE_ID_IDX]);
96+ }
97+ return 0;
98+ }
99+};
100+ 
101+int64_t GetShapeSize(const std::vector<int64_t>& shape)
102+{
103+ int64_t size = 1;
104+ for (auto dim : shape) {
105+ size *= dim;
106+ }
107+ return size;
108+}
109+ 
110+template <typename T>
111+int CreateAclTensor(const std::vector<T>& hostData, const std::vector<int64_t>& shape,
112+ void** deviceAddr, aclDataType dataType, aclTensor** tensor)
113+{
114+ auto size = GetShapeSize(shape) * static_cast<int64_t>(sizeof(T));
115+ auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST);
116+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMalloc failed. ERROR: %d\n", ret); return ret);
117+ ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE);
118+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMemcpy failed. ERROR: %d\n", ret); return ret);
119+ 
120+ std::vector<int64_t> strides(shape.size(), 1);
121+ for (int64_t i = static_cast<int64_t>(shape.size()) - 2; i >= 0; --i) {
122+ strides[i] = shape[i + 1] * strides[i + 1];
123+ }
124+ *tensor = aclCreateTensor(shape.data(), shape.size(), dataType,
125+ strides.data(), 0, ACL_FORMAT_ND,
126+ shape.data(), shape.size(), *deviceAddr);
127+ CHECK_RET(*tensor != nullptr, LOG_PRINT("aclCreateTensor failed.\n"); return ACL_ERROR_FAILURE);
128+ return ACL_SUCCESS;
129+}
130+ 
131+template <typename T>
132+int CopyDeviceToHost(std::vector<T>& hostData, const void* deviceAddr, const char* name)
133+{
134+ const size_t bytes = hostData.size() * sizeof(T);
135+ const auto ret = aclrtMemcpy(hostData.data(), bytes, deviceAddr, bytes,
136+ ACL_MEMCPY_DEVICE_TO_HOST);
137+ CHECK_RET(ret == ACL_SUCCESS,
138+ LOG_PRINT("copy %s from device to host failed. ERROR: %d\n", name, ret);
139+ return ret);
140+ return ACL_SUCCESS;
141+}
142+ 
143+ 
144+int InitAcl(int32_t deviceId, aclrtContext* context, aclrtStream* stream)
145+{
146+ auto ret = aclInit(nullptr);
147+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclInit failed. ERROR: %d\n", ret); return ret);
148+ ret = aclrtSetDevice(deviceId);
149+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetDevice failed. ERROR: %d\n", ret); return ret);
150+ ret = aclrtCreateContext(context, deviceId);
151+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateContext failed. ERROR: %d\n", ret); return ret);
152+ ret = aclrtSetCurrentContext(*context);
153+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetCurrentContext failed. ERROR: %d\n", ret); return ret);
154+ ret = aclrtCreateStream(stream);
155+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateStream failed. ERROR: %d\n", ret); return ret);
156+ return ACL_SUCCESS;
157+}
158+ 
159+void PrintChecksum(const std::vector<int64_t>& shape, void* deviceAddr, aclDataType dataType,
160+ const char* name)
161+{
162+ auto elementCount = GetShapeSize(shape);
163+ if (dataType == ACL_FLOAT) {
164+ std::vector<float> hostData(elementCount, 0.0F);
165+ auto ret = aclrtMemcpy(hostData.data(), hostData.size() * sizeof(float),
166+ deviceAddr, hostData.size() * sizeof(float), ACL_MEMCPY_DEVICE_TO_HOST);
167+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("copy %s from device to host failed. ERROR: %d\n", name, ret); return);
168+ 
169+ double sum = 0.0;
170+ for (auto v : hostData) {
171+ sum += static_cast<double>(v);
172+ }
173+ LOG_PRINT("%s (FP32) first 8:", name);
174+ int64_t limit = elementCount < 8 ? elementCount : 8;
175+ for (int64_t i = 0; i < limit; ++i) {
176+ LOG_PRINT(" %.6f", hostData[i]);
177+ }
178+ LOG_PRINT(" checksum: %.6f\n", sum);
179+ 
180+ } else if (dataType == ACL_BF16) {
181+ std::vector<Bf16> hostData(elementCount, 0);
182+ auto ret = aclrtMemcpy(hostData.data(), hostData.size() * sizeof(Bf16),
183+ deviceAddr, hostData.size() * sizeof(Bf16), ACL_MEMCPY_DEVICE_TO_HOST);
184+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("copy %s from device to host failed. ERROR: %d\n", name, ret); return);
185+ 
186+ double sum = 0.0;
187+ for (auto v : hostData) {
188+ sum += static_cast<double>(Bf16ToFloat(v));
189+ }
190+ LOG_PRINT("%s (BF16) first 8:", name);
191+ int64_t limit = elementCount < 8 ? elementCount : 8;
192+ for (int64_t i = 0; i < limit; ++i) {
193+ LOG_PRINT(" %.6f", Bf16ToFloat(hostData[i]));
194+ }
195+ LOG_PRINT(" checksum: %.6f\n", sum);
196+ 
197+ } else if (dataType == ACL_UINT8) {
198+ std::vector<uint8_t> hostData(elementCount, 0);
199+ auto ret = aclrtMemcpy(hostData.data(), hostData.size() * sizeof(uint8_t),
200+ deviceAddr, hostData.size() * sizeof(uint8_t), ACL_MEMCPY_DEVICE_TO_HOST);
201+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("copy %s from device to host failed. ERROR: %d\n", name, ret); return);
202+ 
203+ uint64_t sum = 0;
204+ for (auto v : hostData) {
205+ sum += static_cast<uint64_t>(v);
206+ }
207+ LOG_PRINT("%s (U8) first 8:", name);
208+ int64_t limit = elementCount < 8 ? elementCount : 8;
209+ for (int64_t i = 0; i < limit; ++i) {
210+ LOG_PRINT(" %u", static_cast<uint32_t>(hostData[i]));
211+ }
212+ LOG_PRINT(" checksum: %lu\n", sum);
213+ }
214+}
215+ 
216+int main(int argc, const char** argv)
217+{
218+ TestArgs args;
219+ if (args.Parse(argc, argv) != 0) {
220+ return -1;
221+ }
222+ args.Print();
223+ 
224+ const int64_t M = args.m;
225+ const int64_t N = args.n;
226+ const int64_t K = args.k;
227+ const int64_t splitN = (N + 256 - 1) / 256; // L1_TILE_N = 256
228+ const int64_t bStatBlock = (splitN + 8 - 1) / 8; // FLOAT_ELEMENTS_PER_BLOCK = 8
229+ const int64_t bStatLen = bStatBlock * 8 + 8; // aligned length for B stats
230+ const int64_t rowSplitLen = M * splitN; // length of per-row split outputs
231+ const int64_t beLen = K * splitN; // length of BE / BMaxSlice etc.
232+ 
233+ // ── Input shapes ──
234+ const std::vector<int64_t> aShape = {M, K};
235+ const std::vector<int64_t> bShape = {K, N};
236+ const std::vector<int64_t> xvShape = {N};
237+ const std::vector<int64_t> vxForAeShape = {K};
238+ 
239+ // ── Output shapes (from infershape) ──
240+ const std::vector<int64_t> zRowColShape = {rowSplitLen};
241+ const std::vector<int64_t> compRowShape = {((M + 7) / 8) * splitN};
242+ const std::vector<int64_t> cShape = {M, N};
243+ const std::vector<int64_t> bStatShape = {bStatLen};
244+ const std::vector<int64_t> beShape = {beLen};
245+ const std::vector<int64_t> aRedShape = {M};
246+ 
247+ // ── Init ACL ──
248+ aclrtContext context = nullptr;
249+ aclrtStream stream = nullptr;
250+ auto ret = InitAcl(args.deviceId, &context, &stream);
251+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret);
252+ 
253+ // ── Fill host data ──
254+ std::mt19937 rng(42);
255+ std::uniform_real_distribution<float> dist(-1.0F, 1.0F);
256+ 
257+ std::vector<Bf16> hostA(GetShapeSize(aShape), 0);
258+ std::vector<Bf16> hostB(GetShapeSize(bShape), 0);
259+ std::vector<Bf16> hostXv(GetShapeSize(xvShape), 0);
260+ std::vector<Bf16> hostVxForAe(GetShapeSize(vxForAeShape), 0);
261+ 
262+ for (auto& v : hostA) {
263+ v = FloatToBf16(dist(rng));
264+ }
265+ for (auto& v : hostB) {
266+ v = FloatToBf16(dist(rng));
267+ }
268+ for (auto& v : hostXv) {
269+ v = FloatToBf16(1.0F);
270+ }
271+ for (auto& v : hostVxForAe) {
272+ v = FloatToBf16(1.0F / static_cast<float>(K));
273+ }
274+ 
275+ // Output data (zero-initialised)
276+ std::vector<float> hostZRow(GetShapeSize(zRowColShape), 0.0F);
277+ std::vector<float> hostDRow(GetShapeSize(zRowColShape), 0.0F);
278+ std::vector<uint8_t> hostCompRow(GetShapeSize(compRowShape), 0);
279+ std::vector<float> hostC(GetShapeSize(cShape), 0.0F);
280+ std::vector<float> hostThreshold(GetShapeSize(zRowColShape), 0.0F);
281+ std::vector<float> hostBMeanAbs(GetShapeSize(bStatShape), 0.0F);
282+ std::vector<float> hostBMeanSquare(GetShapeSize(bStatShape), 0.0F);
283+ std::vector<float> hostBVar(GetShapeSize(bStatShape), 0.0F);
284+ std::vector<Bf16> hostBe(GetShapeSize(beShape), 0);
285+ std::vector<float> hostBeForAiv(GetShapeSize(beShape), 0.0F);
286+ std::vector<float> hostBMaxSlice(GetShapeSize(beShape), 0.0F);
287+ std::vector<float> hostBMinSlice(GetShapeSize(beShape), 0.0F);
288+ std::vector<float> hostAMax(GetShapeSize(aRedShape), 0.0F);
289+ std::vector<float> hostAMean(GetShapeSize(aRedShape), 0.0F);
290+ std::vector<float> hostAMin(GetShapeSize(aRedShape), 0.0F);
291+ 
292+ // ── Create device tensors ──
293+ void *devA = nullptr, *devB = nullptr, *devXv = nullptr, *devVxForAe = nullptr;
294+ void *devZRow = nullptr, *devDRow = nullptr;
295+ void *devCompRow = nullptr, *devGemmAddend = nullptr, *devC = nullptr;
296+ void *devThreshold = nullptr;
297+ void *devBMeanAbs = nullptr, *devBMeanSquare = nullptr, *devBVar = nullptr;
298+ void *devBe = nullptr, *devBeForAiv = nullptr;
299+ void *devBMaxSlice = nullptr, *devBMinSlice = nullptr;
300+ void *devAMax = nullptr, *devAMean = nullptr, *devAMin = nullptr;
301+ 
302+ aclTensor *tA = nullptr, *tB = nullptr, *tXv = nullptr, *tVxForAe = nullptr;
303+ aclTensor *tZRow = nullptr, *tDRow = nullptr;
304+ aclTensor *tCompRow = nullptr, *tGemmAddend = nullptr, *tC = nullptr;
305+ aclTensor *tThreshold = nullptr;
306+ aclTensor *tBMeanAbs = nullptr, *tBMeanSquare = nullptr, *tBVar = nullptr;
307+ aclTensor *tBe = nullptr, *tBeForAiv = nullptr;
308+ aclTensor *tBMaxSlice = nullptr, *tBMinSlice = nullptr;
309+ aclTensor *tAMax = nullptr, *tAMean = nullptr, *tAMin = nullptr;
310+ 
311+ CHECK_RET(CreateAclTensor(hostA, aShape, &devA, ACL_BF16, &tA) == ACL_SUCCESS, return ret);
312+ CHECK_RET(CreateAclTensor(hostB, bShape, &devB, ACL_BF16, &tB) == ACL_SUCCESS, return ret);
313+ CHECK_RET(CreateAclTensor(hostXv, xvShape, &devXv, ACL_BF16, &tXv) == ACL_SUCCESS, return ret);
314+ CHECK_RET(CreateAclTensor(hostVxForAe, vxForAeShape, &devVxForAe, ACL_BF16, &tVxForAe) == ACL_SUCCESS, return ret);
315+ 
316+ CHECK_RET(CreateAclTensor(hostZRow, zRowColShape, &devZRow, ACL_FLOAT, &tZRow) == ACL_SUCCESS, return ret);
317+ CHECK_RET(CreateAclTensor(hostDRow, zRowColShape, &devDRow, ACL_FLOAT, &tDRow) == ACL_SUCCESS, return ret);
318+ CHECK_RET(CreateAclTensor(hostCompRow, compRowShape, &devCompRow, ACL_UINT8, &tCompRow) == ACL_SUCCESS, return ret);
319+ CHECK_RET(CreateAclTensor(hostC, cShape, &devGemmAddend, ACL_FLOAT, &tGemmAddend) == ACL_SUCCESS, return ret);
320+ CHECK_RET(CreateAclTensor(hostC, cShape, &devC, ACL_FLOAT, &tC) == ACL_SUCCESS, return ret);
321+ CHECK_RET(CreateAclTensor(hostThreshold, zRowColShape, &devThreshold, ACL_FLOAT, &tThreshold) == ACL_SUCCESS, return ret);
322+ CHECK_RET(CreateAclTensor(hostBMeanAbs, bStatShape, &devBMeanAbs, ACL_FLOAT, &tBMeanAbs) == ACL_SUCCESS, return ret);
323+ CHECK_RET(CreateAclTensor(hostBMeanSquare, bStatShape, &devBMeanSquare, ACL_FLOAT, &tBMeanSquare) == ACL_SUCCESS, return ret);
324+ CHECK_RET(CreateAclTensor(hostBVar, bStatShape, &devBVar, ACL_FLOAT, &tBVar) == ACL_SUCCESS, return ret);
325+ CHECK_RET(CreateAclTensor(hostBe, beShape, &devBe, ACL_BF16, &tBe) == ACL_SUCCESS, return ret);
326+ CHECK_RET(CreateAclTensor(hostBeForAiv, beShape, &devBeForAiv, ACL_FLOAT, &tBeForAiv) == ACL_SUCCESS, return ret);
327+ CHECK_RET(CreateAclTensor(hostBMaxSlice, beShape, &devBMaxSlice, ACL_FLOAT, &tBMaxSlice) == ACL_SUCCESS, return ret);
328+ CHECK_RET(CreateAclTensor(hostBMinSlice, beShape, &devBMinSlice, ACL_FLOAT, &tBMinSlice) == ACL_SUCCESS, return ret);
329+ CHECK_RET(CreateAclTensor(hostAMax, aRedShape, &devAMax, ACL_FLOAT, &tAMax) == ACL_SUCCESS, return ret);
330+ CHECK_RET(CreateAclTensor(hostAMean, aRedShape, &devAMean, ACL_FLOAT, &tAMean) == ACL_SUCCESS, return ret);
331+ CHECK_RET(CreateAclTensor(hostAMin, aRedShape, &devAMin, ACL_FLOAT, &tAMin) == ACL_SUCCESS, return ret);
332+ 
333+ // ── Launch MatmulAbftVerify ──
334+ uint64_t workspaceSize = 0;
335+ aclOpExecutor* executor = nullptr;
336+ 
337+ ret = aclnnMatmulAbftVerifyGetWorkspaceSize(
338+ tA, tB, tC, tXv, tVxForAe,
339+ args.eMax,
340+ args.reduceCores,
341+ tZRow, tDRow, tCompRow,
342+ tThreshold,
343+ tBMeanAbs, tBMeanSquare, tBVar,
344+ tBe, tBeForAiv,
345+ tBMaxSlice, tBMinSlice,
346+ tAMax, tAMean, tAMin,
347+ &workspaceSize, &executor);
348+ CHECK_RET(ret == ACL_SUCCESS,
349+ LOG_PRINT("aclnnMatmulAbftVerifyGetWorkspaceSize failed. ERROR: %d.\n[ERROR msg]%s\n",
350+ ret, aclGetRecentErrMsg());
351+ return ret);
352+ 
353+ void* workspaceAddr = nullptr;
354+ if (workspaceSize > 0) {
355+ ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST);
356+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("allocate workspace failed. ERROR: %d\n", ret); return ret);
357+ }
358+ 
359+ ret = aclnnMatmulAbftVerify(workspaceAddr, workspaceSize, executor, stream);
360+ CHECK_RET(ret == ACL_SUCCESS,
361+ LOG_PRINT("aclnnMatmulAbftVerify failed. ERROR: %d.\n[ERROR msg]%s\n",
362+ ret, aclGetRecentErrMsg());
363+ return ret);
364+ 
365+ ret = aclrtSynchronizeStream(stream);
366+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSynchronizeStream failed. ERROR: %d\n", ret); return ret);
367+ 
368+ LOG_PRINT("MatmulAbftVerify run finished successfully.\n\n");
369+ // ── Print results ──
370+ if (args.verbose) {
371+ PrintChecksum(cShape, devC, ACL_FLOAT, "c");
372+ PrintChecksum(zRowColShape, devZRow, ACL_FLOAT, "zRow");
373+ PrintChecksum(zRowColShape, devDRow, ACL_FLOAT, "dRow");
374+ PrintChecksum(compRowShape, devCompRow, ACL_UINT8, "compRow");
375+ PrintChecksum(zRowColShape, devThreshold, ACL_FLOAT, "threshold");
376+ PrintChecksum(bStatShape, devBMeanAbs, ACL_FLOAT, "bMeanAbs");
377+ PrintChecksum(bStatShape, devBMeanSquare, ACL_FLOAT, "bMeanSquare");
378+ PrintChecksum(bStatShape, devBVar, ACL_FLOAT, "bVar");
379+ PrintChecksum(beShape, devBe, ACL_BF16, "be");
380+ PrintChecksum(beShape, devBeForAiv, ACL_FLOAT, "beForAiv");
381+ PrintChecksum(beShape, devBMaxSlice, ACL_FLOAT, "bMaxSlice");
382+ PrintChecksum(beShape, devBMinSlice, ACL_FLOAT, "bMinSlice");
383+ PrintChecksum(aRedShape, devAMax, ACL_FLOAT, "aMax");
384+ PrintChecksum(aRedShape, devAMean, ACL_FLOAT, "aMean");
385+ PrintChecksum(aRedShape, devAMin, ACL_FLOAT, "aMin");
386+ }
387+ 
388+ ret = aclDestroyAclOpExecutor(executor);
389+ CHECK_RET(ret == ACL_SUCCESS,
390+ LOG_PRINT("aclDestroyAclOpExecutor failed. ERROR: %d\n", ret);
391+ return ret);
392+ executor = nullptr;
393+ 
394+ // ── Cleanup ──
395+ auto Destroy = [](aclTensor*& t) { if (t) { aclDestroyTensor(t); t = nullptr; } };
396+ auto Free = [](void*& p) { if (p) { aclrtFree(p); p = nullptr; } };
397+ 
398+ Destroy(tA); Destroy(tB); Destroy(tXv); Destroy(tVxForAe);
399+ Destroy(tZRow); Destroy(tDRow);
400+ Destroy(tCompRow); Destroy(tGemmAddend); Destroy(tC);
401+ Destroy(tThreshold);
402+ Destroy(tBMeanAbs); Destroy(tBMeanSquare); Destroy(tBVar);
403+ Destroy(tBe); Destroy(tBeForAiv);
404+ Destroy(tBMaxSlice); Destroy(tBMinSlice);
405+ Destroy(tAMax); Destroy(tAMean); Destroy(tAMin);
406+ 
407+ Free(devA); Free(devB); Free(devXv); Free(devVxForAe);
408+ Free(devZRow); Free(devDRow);
409+ Free(devCompRow); Free(devGemmAddend); Free(devC);
410+ Free(devThreshold);
411+ Free(devBMeanAbs); Free(devBMeanSquare); Free(devBVar);
412+ Free(devBe); Free(devBeForAiv);
413+ Free(devBMaxSlice); Free(devBMinSlice);
414+ Free(devAMax); Free(devAMean); Free(devAMin);
415+ Free(workspaceAddr);
416+ 
417+ aclrtDestroyStream(stream);
418+ aclrtDestroyContext(context);
419+ aclrtResetDevice(args.deviceId);
420+ aclFinalize();
421+ 
422+ return 0;
423+}
@@ -0,0 +1,51 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2025 Huawei Technologies Co., Ltd.
3+# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+set(ASCEND_OP_NAME "" CACHE STRING "Ascend op names to compile")
11+set(OP_TYPE "matmul_abft_verify")
12+ 
13+# ACLNN code generation is disabled unless this switch is visible from the
14+# top-level generation stage. Keep the setting local to this operator's build
15+# configuration instead of changing the repository root CMakeLists.txt.
16+set(ENABLE_GEN_ACLNN ON CACHE BOOL "Generate ACLNN APIs from operator definitions" FORCE)
17+ 
18+if(ASCEND_OP_NAME)
19+ set(_ASCEND_OP_NAME_TMP "${ASCEND_OP_NAME}")
20+ separate_arguments(_ASCEND_OP_NAME_TMP)
21+ 
22+ list(FIND _ASCEND_OP_NAME_TMP "${OP_TYPE}" _index)
23+ if(_index EQUAL -1)
24+ message(STATUS "[${OP_TYPE}] skipped, not in ASCEND_OP_NAME list: ${ASCEND_OP_NAME}")
25+ return()
26+ else()
27+ message(STATUS "[${OP_TYPE}] selected, in ASCEND_OP_NAME list: ${ASCEND_OP_NAME}")
28+ endif()
29+endif()
30+ 
31+file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
32+if(NOT ENABLE_TEST AND NOT BENCHMARK)
33+ list(REMOVE_ITEM CURRENT_DIRS tests)
34+endif()
35+foreach(SUB_DIR ${CURRENT_DIRS})
36+ if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
37+ add_subdirectory(${SUB_DIR})
38+ endif()
39+endforeach()
40+ 
41+add_ops_compile_options(
42+ MatmulAbftVerify
43+ OPTIONS -I${CMAKE_CURRENT_SOURCE_DIR}/../op_kernel
44+)
45+add_modules_sources(
46+ HOSTNAME ${OPHOST_NAME}
47+ MODE PRIVATE
48+ DIR ${CMAKE_CURRENT_SOURCE_DIR}
49+ OPTYPE matmul_abft_verify
50+ ACLNNTYPE aclnn
51+)
@@ -0,0 +1,76 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file matmul_abft_verify_def.cpp
13+ * \brief MatmulAbftVerify operator definition.
14+ */
15+#include "register/op_def_registry.h"
16+ 
17+namespace ops {
18+namespace {
19+const std::vector<ge::DataType> kPrecisionDtype = {ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT};
20+const std::vector<ge::DataType> kFloatDtype = {ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT};
21+const std::vector<ge::DataType> kUint8Dtype = {ge::DT_UINT8, ge::DT_UINT8, ge::DT_UINT8};
22+const std::vector<ge::Format> kNdFormat = {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND};
23+} // namespace
24+ 
25+class MatmulAbftVerify : public OpDef {
26+public:
27+ explicit MatmulAbftVerify(const char* name) : OpDef(name)
28+ {
29+ this->Input("a")
30+ .ParamType(REQUIRED)
31+ .DataType(kPrecisionDtype)
32+ .Format(kNdFormat)
33+ .UnknownShapeFormat(kNdFormat)
34+ .AutoContiguous();
35+ this->Input("b")
36+ .ParamType(REQUIRED)
37+ .DataType(kPrecisionDtype)
38+ .Format(kNdFormat)
39+ .UnknownShapeFormat(kNdFormat)
40+ .AutoContiguous();
41+ this->Input("c")
42+ .ParamType(REQUIRED)
43+ .DataType(kFloatDtype)
44+ .Format(kNdFormat)
45+ .UnknownShapeFormat(kNdFormat)
46+ .AutoContiguous();
47+ this->Input("checksum_weight")
48+ .ParamType(REQUIRED)
49+ .DataType(kPrecisionDtype)
50+ .Format(kNdFormat)
51+ .UnknownShapeFormat(kNdFormat)
52+ .AutoContiguous();
53+ 
54+ this->Output("comp_row")
55+ .ParamType(REQUIRED)
56+ .DataType(kUint8Dtype)
57+ .Format(kNdFormat)
58+ .UnknownShapeFormat(kNdFormat)
59+ .AutoContiguous();
60+ 
61+ this->Attr("e_max").AttrType(OPTIONAL).Float(0.001);
62+ OpAICoreConfig aicoreConfig;
63+ aicoreConfig.DynamicCompileStaticFlag(true)
64+ .DynamicFormatFlag(false)
65+ .DynamicRankSupportFlag(true)
66+ .NeedCheckSupportFlag(false)
67+ .DynamicShapeSupportFlag(true)
68+ .PrecisionReduceFlag(true)
69+ .ExtendCfgInfo("opFile.value", "matmul_abft_verify");
70+ this->AICore().AddConfig("ascend910b", aicoreConfig);
71+ this->AICore().AddConfig("ascend910_93", aicoreConfig);
72+ }
73+};
74+ 
75+OP_ADD(MatmulAbftVerify);
76+} // namespace ops
@@ -0,0 +1,62 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file matmul_abft_verify_infershape.cpp
13+ * \brief MatmulAbftVerify shape and dtype inference.
14+ */
15+#include "register/op_impl_registry.h"
16+#include "log/log.h"
17+ 
18+namespace ops {
19+namespace {
20+constexpr int64_t L1_TILE_N = 256;
21+constexpr int64_t COMP_PACK = 8;
22+ 
23+ge::graphStatus SetVectorShape(gert::InferShapeContext* context, size_t index, int64_t length)
24+{
25+ gert::Shape* shape = context->GetOutputShape(index);
26+ OP_CHECK_NULL_WITH_CONTEXT(context, shape);
27+ shape->SetDimNum(1);
28+ shape->SetDim(0, length);
29+ return ge::GRAPH_SUCCESS;
30+}
31+} // namespace
32+ 
33+static ge::graphStatus InferShapeMatmulAbftVerify(gert::InferShapeContext* context)
34+{
35+ const gert::Shape* aShape = context->GetInputShape(0);
36+ const gert::Shape* bShape = context->GetInputShape(1);
37+ OP_CHECK_NULL_WITH_CONTEXT(context, aShape);
38+ OP_CHECK_NULL_WITH_CONTEXT(context, bShape);
39+ OP_CHECK_IF(aShape->GetDimNum() != 2 || bShape->GetDimNum() != 2,
40+ OP_LOGE(context->GetNodeName(), "MatmulAbftVerify expects rank-2 A and B"), return ge::GRAPH_FAILED);
41+ 
42+ const int64_t m = aShape->GetDim(0);
43+ const int64_t k = aShape->GetDim(1);
44+ const int64_t bK = bShape->GetDim(0);
45+ const int64_t n = bShape->GetDim(1);
46+ OP_CHECK_IF(k != bK, OP_LOGE(context->GetNodeName(), "A.K must equal B.K"), return ge::GRAPH_FAILED);
47+ 
48+ const int64_t splitN = (n + L1_TILE_N - 1) / L1_TILE_N;
49+ OP_CHECK_IF(SetVectorShape(context, 0, ((m + COMP_PACK - 1) / COMP_PACK) * splitN) !=
50+ ge::GRAPH_SUCCESS,
51+ OP_LOGE(context->GetNodeName(), "failed to set comp_row shape"), return ge::GRAPH_FAILED);
52+ return ge::GRAPH_SUCCESS;
53+}
54+ 
55+static ge::graphStatus InferDataTypeMatmulAbftVerify(gert::InferDataTypeContext* context)
56+{
57+ context->SetOutputDataType(0, ge::DT_UINT8);
58+ return ge::GRAPH_SUCCESS;
59+}
60+ 
61+IMPL_OP_INFERSHAPE(MatmulAbftVerify).InferShape(InferShapeMatmulAbftVerify).InferDataType(InferDataTypeMatmulAbftVerify);
62+} // namespace ops
@@ -0,0 +1,288 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file matmul_abft_verify_tiling.cpp
13+ * \brief Host-side tiling for MatmulAbftVerify.
14+ */
15+#include <algorithm>
16+#include <cstring>
17+#include <cmath>
18+#include <limits>
19+ 
20+#include "log/log.h"
21+#include "tiling/platform/platform_ascendc.h"
22+#include "op_host/tiling_util.h"
23+#include "op_host/tiling_templates_registry.h"
24+#include "../op_kernel/matmul_abft_verify_tiling_data.h"
25+#include "../op_kernel/matmul_abft_verify_tiling_key.h"
26+ 
27+namespace optiling {
28+namespace {
29+constexpr size_t SYSTEM_WORKSPACE_SIZE = 16UL * 1024UL * 1024UL;
30+constexpr uint64_t WORKSPACE_ALIGNMENT = 32;
31+constexpr uint64_t FLOAT_BYTES = sizeof(float);
32+ 
33+bool AddWorkspaceTensor(uint64_t elements, uint64_t elementBytes, uint64_t& workspaceBytes)
34+{
35+ if (workspaceBytes > std::numeric_limits<uint64_t>::max() - (WORKSPACE_ALIGNMENT - 1)) {
36+ return false;
37+ }
38+ workspaceBytes = (workspaceBytes + WORKSPACE_ALIGNMENT - 1) / WORKSPACE_ALIGNMENT * WORKSPACE_ALIGNMENT;
39+ if (elements != 0 && elementBytes > std::numeric_limits<uint64_t>::max() / elements) {
40+ return false;
41+ }
42+ const uint64_t tensorBytes = elements * elementBytes;
43+ if (workspaceBytes > std::numeric_limits<uint64_t>::max() - tensorBytes) {
44+ return false;
45+ }
46+ workspaceBytes += tensorBytes;
47+ return true;
48+}
49+ 
50+uint32_t FloatBits(float value)
51+{
52+ uint32_t bits;
53+ std::copy_n(reinterpret_cast<const unsigned char *>(&value), sizeof(value),
54+ reinterpret_cast<unsigned char *>(&bits));
55+ return bits;
56+}
57+ 
58+uint16_t FloatToBf16Bits(float value)
59+{
60+ const uint32_t bits = FloatBits(value);
61+ return static_cast<uint16_t>((bits + 0x7FFFU + ((bits >> 16U) & 1U)) >> 16U);
62+}
63+ 
64+uint16_t FloatToFp16Bits(float value)
65+{
66+ const uint32_t bits = FloatBits(value);
67+ const int32_t exponent = static_cast<int32_t>((bits >> 23U) & 0xFFU) - 127;
68+ uint32_t mantissa = bits & 0x7FFFFFU;
69+ if (exponent < -24) {
70+ return 0;
71+ }
72+ if (exponent < -14) {
73+ mantissa |= 0x800000U;
74+ const uint32_t shift = static_cast<uint32_t>(-exponent - 1);
75+ const uint32_t rounded = mantissa + ((1U << (shift - 1U)) - 1U) + ((mantissa >> shift) & 1U);
76+ return static_cast<uint16_t>(rounded >> shift);
77+ }
78+ uint32_t halfExponent = static_cast<uint32_t>(exponent + 15);
79+ mantissa += 0xFFFU + ((mantissa >> 13U) & 1U);
80+ if ((mantissa & 0x800000U) != 0U) {
81+ ++halfExponent;
82+ mantissa = 0;
83+ }
84+ return static_cast<uint16_t>((halfExponent << 10U) | (mantissa >> 13U));
85+}
86+}
87+ 
88+template <typename Policy>
89+MatmulAbftVerifyComputeTilingData BuildComputeTiling(uint32_t m, uint32_t n, uint32_t k,
90+ float eMax, uint32_t reduceCores, uint32_t splitKs)
91+{
92+ // DIVISOR_GUARDED[585]: keep the reduction divisor non-zero for all caller inputs.
93+ const uint32_t safeReduceCores = reduceCores == 0U ? 1U : reduceCores;
94+ const uint32_t splitNNum = (n + Policy::L1TileShape_N - 1) / Policy::L1TileShape_N;
95+ const uint32_t firstBlockN = splitNNum;
96+ const uint32_t remainBlockN = n;
97+ 
98+ const uint32_t totalInputElements = m + n;
99+ const uint32_t totalOutputElements = (m + 7) / 8 + (n + 7) / 8;
100+ const uint32_t rowOutputElements = (m + 7) / 8;
101+ const uint32_t colOutputElements = (n + 7) / 8;
102+ const uint32_t xLen = m > n ? m : n;
103+ 
104+ uint32_t splitReduceM = (splitNNum + safeReduceCores - 1U) / safeReduceCores;
105+ splitReduceM = splitReduceM >= Policy::UBTileShapeforBRed_M ? Policy::UBTileShapeforBRed_M : splitReduceM;
106+ splitReduceM = splitReduceM < 2 ? 2 : splitReduceM;
107+ 
108+ const uint32_t splitReduceNNum = (k + Policy::UBTileShapeforBRed_N - 1) / Policy::UBTileShapeforBRed_N;
109+ uint32_t splitReduceN = Policy::UBTileShapeforBRed_N;
110+ if (splitReduceNNum < 2) {
111+ splitReduceN = (k + 1) / 2;
112+ }
113+ 
114+ const uint32_t nRemainSplit = n % Policy::L1TileShape_N;
115+ const float commonSize = static_cast<float>(k) * Policy::L1TileShape_N;
116+ const float remainSize = static_cast<float>(k) * nRemainSplit;
117+ 
118+ const float commonStdFactor = std::sqrt(2.0f * std::log(commonSize));
119+ float remainStdFactor = commonStdFactor;
120+ float remainKnRatio = commonSize;
121+ float remainKnSqrtRatio = std::sqrt(commonSize);
122+ float remainKSqrtNRatio = std::sqrt(static_cast<float>(k)) * Policy::L1TileShape_N;
123+ if (nRemainSplit > 0) {
124+ remainStdFactor = std::sqrt(2.0f * std::log(remainSize));
125+ remainKnRatio = remainSize;
126+ remainKnSqrtRatio = std::sqrt(static_cast<float>(nRemainSplit));
127+ remainKSqrtNRatio = std::sqrt(static_cast<float>(k)) * nRemainSplit;
128+ }
129+ 
130+ MatmulAbftVerifyComputeTilingData tiling{};
131+ tiling.problemGemmShape = {m, n, k};
132+ tiling.problemGemmShapeFirst = {m, firstBlockN, k};
133+ tiling.problemGemmShapeRemain = {m, remainBlockN, k};
134+ tiling.problemShape = {m, n};
135+ tiling.problemShapeCol = {n, m};
136+ tiling.problemCompShape = {1, totalInputElements};
137+ tiling.problemSliceShape = {splitNNum, m};
138+ tiling.totalInputElements = totalInputElements;
139+ tiling.totalOutputElements = totalOutputElements;
140+ tiling.rowOutputElements = rowOutputElements;
141+ tiling.colOutputElements = colOutputElements;
142+ tiling.xLen = xLen;
143+ tiling.layoutThreLen = m;
144+ tiling.roundingAlpha = 1.0f;
145+ tiling.eMax = eMax * 6 * std::sqrt(static_cast<float>(k) / 1024.0f);
146+ tiling.stdEstARowRatio = 1.0f / std::sqrt(2.0f * std::log(static_cast<float>(k)));
147+ tiling.aRowScaleRatio = 1.0f / static_cast<float>(k);
148+ 
149+ tiling.stdEstRatios[0] = commonStdFactor == 0.0f ? 0.0f : 1.0f / commonStdFactor;
150+ tiling.stdEstRatios[1] = remainStdFactor == 0.0f ? 0.0f : 1.0f / remainStdFactor;
151+ tiling.knRatios[0] = commonSize;
152+ tiling.knRatios[1] = remainKnRatio;
153+ tiling.knScaleRatios[0] = 1.0f / commonSize;
154+ tiling.knScaleRatios[1] = 1.0f / remainKnRatio;
155+ tiling.knSqrtRatios[0] = std::sqrt(commonSize);
156+ tiling.knSqrtRatios[1] = remainKnSqrtRatio;
157+ tiling.kSqrtNRatios[0] = std::sqrt(static_cast<float>(k)) * Policy::L1TileShape_N;
158+ tiling.kSqrtNRatios[1] = remainKSqrtNRatio;
159+ 
160+ tiling.nScaleRatios[0] = 1.0f / Policy::L1TileShape_N;
161+ tiling.nRatios[0] = static_cast<float>(Policy::L1TileShape_N);
162+ tiling.nSqrtRatios[0] = std::sqrt(static_cast<float>(Policy::L1TileShape_N));
163+ tiling.nSquareRatios[0] = static_cast<float>(Policy::L1TileShape_N * Policy::L1TileShape_N);
164+ const uint32_t remainN = nRemainSplit > 0 ? nRemainSplit : Policy::L1TileShape_N;
165+ tiling.nScaleRatios[1] = 1.0f / remainN;
166+ tiling.nRatios[1] = static_cast<float>(remainN);
167+ tiling.nSqrtRatios[1] = std::sqrt(static_cast<float>(remainN));
168+ tiling.nSquareRatios[1] = static_cast<float>(remainN * remainN);
169+ 
170+ tiling.splitNNum = splitNNum;
171+ tiling.splitReduceM = splitReduceM;
172+ tiling.splitReduceN = splitReduceN;
173+ tiling.splitKNum = splitKs;
174+ tiling.outputThre = 1;
175+ tiling.outputCE = 1;
176+ tiling.outputWorkspace = 0;
177+ return tiling;
178+}
179+ 
180+ 
181+static ge::graphStatus MatmulAbftVerifyTilingFunc(gert::TilingContext* context)
182+{
183+ const gert::StorageShape* aShape = context->GetInputShape(0);
184+ const gert::StorageShape* bShape = context->GetInputShape(1);
185+ OP_CHECK_NULL_WITH_CONTEXT(context, aShape);
186+ OP_CHECK_NULL_WITH_CONTEXT(context, bShape);
187+ 
188+ const gert::Shape& aStorageShape = aShape->GetStorageShape();
189+ const gert::Shape& bStorageShape = bShape->GetStorageShape();
190+ OP_CHECK_IF(aStorageShape.GetDimNum() != 2 || bStorageShape.GetDimNum() != 2,
191+ OP_LOGE(context, "MatmulAbftVerify expects rank-2 A and B"), return ge::GRAPH_FAILED);
192+ 
193+ const int64_t m = aStorageShape.GetDim(0);
194+ const int64_t k = aStorageShape.GetDim(1);
195+ const int64_t bK = bStorageShape.GetDim(0);
196+ const int64_t n = bStorageShape.GetDim(1);
197+ OP_CHECK_IF(m <= 0 || n <= 0 || k <= 0 || k != bK,
198+ OP_LOGE(context, "invalid MatmulAbftVerify problem shape"), return ge::GRAPH_FAILED);
199+ OP_CHECK_IF(m > UINT32_MAX || n > UINT32_MAX || k > UINT32_MAX,
200+ OP_LOGE(context, "MatmulAbftVerify shape exceeds uint32 range"), return ge::GRAPH_FAILED);
201+ 
202+ auto attrs = context->GetAttrs();
203+ OP_CHECK_NULL_WITH_CONTEXT(context, attrs);
204+ uint32_t idx = 0;
205+ const float* eMaxAttr = attrs->GetAttrPointer<float>(idx++);
206+ OP_CHECK_NULL_WITH_CONTEXT(context, eMaxAttr);
207+ 
208+ const float eMax = *eMaxAttr;
209+ const uint32_t splitKs = 1;
210+ const auto* aDesc = context->GetInputDesc(0);
211+ OP_CHECK_NULL_WITH_CONTEXT(context, aDesc);
212+ const ge::DataType aDtype = aDesc->GetDataType();
213+ 
214+ auto platformInfoPtr = context->GetPlatformInfo();
215+ OP_CHECK_IF(platformInfoPtr == nullptr,
216+ OP_LOGE(context, "platformInfoPtr is null"), return ge::GRAPH_FAILED);
217+ auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
218+ const uint32_t aicNum = ascendcPlatform.GetCoreNumAic();
219+ OP_CHECK_IF(aicNum == 0,
220+ OP_LOGE(context, "invalid AIC core count"), return ge::GRAPH_FAILED);
221+ const uint32_t reduceCores = std::min(aicNum, 8U);
222+ 
223+ MatmulAbftVerifyTilingData* tiling = context->GetTilingData<MatmulAbftVerifyTilingData>();
224+ OP_CHECK_NULL_WITH_CONTEXT(context, tiling);
225+ if (aDtype == ge::DT_FLOAT16) {
226+ tiling->compute = BuildComputeTiling<MatmulAbftVerifyFp16TilingPolicy>(
227+ static_cast<uint32_t>(m), static_cast<uint32_t>(n), static_cast<uint32_t>(k),
228+ eMax, reduceCores, splitKs);
229+ } else if (aDtype == ge::DT_FLOAT) {
230+ tiling->compute = BuildComputeTiling<MatmulAbftVerifyFp32TilingPolicy>(
231+ static_cast<uint32_t>(m), static_cast<uint32_t>(n), static_cast<uint32_t>(k),
232+ eMax, reduceCores, splitKs);
233+ } else {
234+ tiling->compute = BuildComputeTiling<MatmulAbftVerifyBf16TilingPolicy>(
235+ static_cast<uint32_t>(m), static_cast<uint32_t>(n), static_cast<uint32_t>(k),
236+ eMax, reduceCores, splitKs);
237+ }
238+ const float aMeanFactor = 1.0f / static_cast<float>(k);
239+ if (aDtype == ge::DT_FLOAT16) {
240+ tiling->compute.aMeanFactorBits = FloatToFp16Bits(aMeanFactor);
241+ } else if (aDtype == ge::DT_FLOAT) {
242+ tiling->compute.aMeanFactorBits = FloatBits(aMeanFactor);
243+ } else {
244+ tiling->compute.aMeanFactorBits = FloatToBf16Bits(aMeanFactor);
245+ }
246+ 
247+ const uint64_t splitN = tiling->compute.splitNNum;
248+ const uint64_t rowSplitElements = static_cast<uint64_t>(m) * splitN;
249+ const uint64_t bStatElements = ((splitN + 7) / 8) * 8 + 8;
250+ const uint64_t beElements = static_cast<uint64_t>(k) * splitN;
251+ const uint64_t precisionBytes = aDtype == ge::DT_FLOAT ? sizeof(float) : sizeof(uint16_t);
252+ const uint64_t ceBytes = aDtype == ge::DT_FLOAT16 ? sizeof(uint16_t) : sizeof(float);
253+ 
254+ uint64_t userWorkspaceBytes = 0;
255+ bool workspaceValid = true;
256+ // The order is part of the workspace ABI and must match matmul_abft_verify.cpp.
257+ for (uint32_t i = 0; i < 3; ++i) { // z_row, d_row, threshold
258+ workspaceValid &= AddWorkspaceTensor(rowSplitElements, FLOAT_BYTES, userWorkspaceBytes);
259+ }
260+ for (uint32_t i = 0; i < 3; ++i) { // b_mean_abs, b_mean_square, b_var
261+ workspaceValid &= AddWorkspaceTensor(bStatElements, FLOAT_BYTES, userWorkspaceBytes);
262+ }
263+ workspaceValid &= AddWorkspaceTensor(beElements, precisionBytes, userWorkspaceBytes); // be
264+ workspaceValid &= AddWorkspaceTensor(beElements, ceBytes, userWorkspaceBytes); // be_for_aiv
265+ for (uint32_t i = 0; i < 2; ++i) { // b_max_slice, b_min_slice
266+ workspaceValid &= AddWorkspaceTensor(beElements, FLOAT_BYTES, userWorkspaceBytes);
267+ }
268+ for (uint32_t i = 0; i < 3; ++i) { // a_max, a_mean, a_min
269+ workspaceValid &= AddWorkspaceTensor(m, FLOAT_BYTES, userWorkspaceBytes);
270+ }
271+ const uint64_t internalWorkspaceElements =
272+ static_cast<uint64_t>(m) * (splitN + 1) * tiling->compute.splitKNum;
273+ workspaceValid &= AddWorkspaceTensor(internalWorkspaceElements, FLOAT_BYTES, userWorkspaceBytes);
274+ 
275+ OP_CHECK_IF(!workspaceValid || userWorkspaceBytes >
276+ std::numeric_limits<size_t>::max() - SYSTEM_WORKSPACE_SIZE,
277+ OP_LOGE(context, "MatmulAbftVerify workspace size overflow"), return ge::GRAPH_FAILED);
278+ size_t* workspaceSizes = context->GetWorkspaceSizes(1);
279+ OP_CHECK_NULL_WITH_CONTEXT(context, workspaceSizes);
280+ workspaceSizes[0] = SYSTEM_WORKSPACE_SIZE + static_cast<size_t>(userWorkspaceBytes);
281+ 
282+ context->SetBlockDim(aicNum);
283+ context->SetTilingKey(GET_TPL_TILING_KEY(MATMUL_ABFT_VERIFY_TPL_SCH_MODE_BF16));
284+ return ge::GRAPH_SUCCESS;
285+}
286+ 
287+IMPL_OP_OPTILING(MatmulAbftVerify).Tiling(MatmulAbftVerifyTilingFunc);
288+} // namespace optiling
@@ -0,0 +1,75 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+#ifndef MATMUL_FT_FTSELF_ARCH_CROSS_CORE_SYNC_AIV_HPP
12+#define MATMUL_FT_FTSELF_ARCH_CROSS_CORE_SYNC_AIV_HPP
13+ 
14+#include "../core/macros.hpp"
15+ 
16+namespace FTSelf {
17+ 
18+namespace Arch {
19+ 
20+constexpr uint32_t MAX_REVERSE_DEPTH = 15;
21+using FlagID = uint16_t;
22+ 
23+template <uint32_t ReverseDepth = MAX_REVERSE_DEPTH>
24+struct CrossCoreFlagWithReverse {
25+ FlagID id;
26+ FlagID reverseId;
27+ uint32_t count{0};
28+ 
29+ FTSELF_DEVICE CrossCoreFlagWithReverse(FlagID id = 0, FlagID reverseId = 0)
30+ : id(id), reverseId(reverseId) {}
31+};
32+ 
33+template <uint8_t Mode, pipe_t Pipe, uint32_t ReverseDepth>
34+FTSELF_DEVICE void CrossCoreSetFlagWithReverse(CrossCoreFlagWithReverse<ReverseDepth> &flag)
35+{
36+ AscendC::CrossCoreSetFlag<Mode, Pipe>(flag.id);
37+ if (++flag.count >= ReverseDepth) {
38+ AscendC::CrossCoreWaitFlag(flag.reverseId);
39+ flag.count = 0;
40+ }
41+}
42+ 
43+template <uint8_t Mode, pipe_t Pipe, uint32_t ReverseDepth>
44+FTSELF_DEVICE void CrossCoreWaitFlagWithReverse(CrossCoreFlagWithReverse<ReverseDepth> &flag)
45+{
46+ AscendC::CrossCoreWaitFlag(flag.id);
47+ if (++flag.count >= ReverseDepth) {
48+ AscendC::CrossCoreSetFlag<Mode, Pipe>(flag.reverseId);
49+ flag.count = 0;
50+ }
51+}
52+ 
53+} // namespace Arch
54+ 
55+template <uint8_t MODE, pipe_t PIPE>
56+FTSELF_DEVICE
57+void CrossCoreBarrierAIC()
58+{
59+ constexpr Arch::FlagID flagId = 9;
60+ AscendC::CrossCoreSetFlag<MODE, PIPE>(flagId);
61+ AscendC::CrossCoreWaitFlag(flagId);
62+}
63+ 
64+template <uint8_t MODE, pipe_t PIPE>
65+FTSELF_DEVICE
66+void CrossCoreBarrierAIV()
67+{
68+ constexpr Arch::FlagID flagId = MODE == 0x1 ? 10 : 8;
69+ AscendC::CrossCoreSetFlag<MODE, PIPE>(flagId);
70+ AscendC::CrossCoreWaitFlag(flagId);
71+}
72+ 
73+} // namespace FTSelf
74+ 
75+#endif // MATMUL_FT_FTSELF_ARCH_CROSS_CORE_SYNC_AIV_HPP
@@ -0,0 +1,35 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+#ifndef MATMUL_FT_FTSELF_ARCH_RESOURCE_AIV_HPP
12+#define MATMUL_FT_FTSELF_ARCH_RESOURCE_AIV_HPP
13+ 
14+#include "../core/arch.hpp"
15+#include "../core/macros.hpp"
16+ 
17+namespace FTSelf {
18+ 
19+template <class ArchTag>
20+struct ResourceAIV {
21+public:
22+ AscendC::TPipe pipe;
23+ FTSelf::Arch::LocalTensorBuffer<ArchTag, AscendC::TPosition::VECCALC> ubBuf;
24+ 
25+ FTSELF_DEVICE
26+ ResourceAIV()
27+ {
28+ // TPipe initialization inserts synchronization that may conflict with callers.
29+ pipe.Destroy();
30+ }
31+};
32+ 
33+} // namespace FTSelf
34+ 
35+#endif // MATMUL_FT_FTSELF_ARCH_RESOURCE_AIV_HPP
@@ -0,0 +1,90 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+#ifndef FTSELF_CORE_ARCH_HPP
12+#define FTSELF_CORE_ARCH_HPP
13+ 
14+#include <cstdint>
15+#include <type_traits>
16+ 
17+#if defined(__CCE__)
18+#include <kernel_operator.h>
19+#endif
20+#include "../core/macros.hpp"
21+ 
22+namespace FTSelf {
23+ 
24+namespace Arch {
25+ 
26+struct AtlasA2 {
27+ static constexpr uint32_t BIAS_SIZE = 1024;
28+ static constexpr uint32_t FIXBUF_SIZE = 7 * 1024;
29+ static constexpr uint32_t UB_SIZE = 192 * 1024;
30+ static constexpr uint32_t L1_SIZE = 512 * 1024;
31+ static constexpr uint32_t L0A_SIZE = 64 * 1024;
32+ static constexpr uint32_t L0B_SIZE = 64 * 1024;
33+ static constexpr uint32_t L0C_SIZE = 128 * 1024;
34+};
35+ 
36+template <AscendC::TPosition Position>
37+using PositionType = std::integral_constant<AscendC::TPosition, Position>;
38+ 
39+using PositionGM = PositionType<AscendC::TPosition::GM>;
40+using PositionL1 = PositionType<AscendC::TPosition::A1>;
41+using PositionL0A = PositionType<AscendC::TPosition::A2>;
42+using PositionL0B = PositionType<AscendC::TPosition::B2>;
43+using PositionL0C = PositionType<AscendC::TPosition::CO1>;
44+using PositionBias = PositionType<AscendC::TPosition::C2>;
45+using PositionUB = PositionType<AscendC::TPosition::VECCALC>;
46+ 
47+template <class ArchTag, AscendC::TPosition Position>
48+struct LocalTensorBuffer {
49+ AscendC::LocalTensor<uint8_t> tensor;
50+ 
51+ FTSELF_DEVICE LocalTensorBuffer()
52+ {
53+ AscendC::TBuf<Position> buffer;
54+ constexpr uint32_t size = Position == AscendC::TPosition::A1 || Position == AscendC::TPosition::B1 ||
55+ Position == AscendC::TPosition::C1 ? ArchTag::L1_SIZE :
56+ Position == AscendC::TPosition::A2 ? ArchTag::L0A_SIZE :
57+ Position == AscendC::TPosition::B2 ? ArchTag::L0B_SIZE :
58+ Position == AscendC::TPosition::C2 ? ArchTag::BIAS_SIZE :
59+ Position == AscendC::TPosition::CO1 ? ArchTag::L0C_SIZE :
60+ Position == AscendC::TPosition::C2PIPE2GM ? ArchTag::FIXBUF_SIZE : ArchTag::UB_SIZE;
61+ GetTPipePtr()->InitBuffer(buffer, size);
62+ tensor = buffer.template Get<uint8_t>();
63+ }
64+ 
65+ template <class Element = half>
66+ FTSELF_DEVICE AscendC::LocalTensor<Element> GetBufferByByte(uint32_t offset) const
67+ {
68+ return tensor[offset].template ReinterpretCast<Element>();
69+ }
70+};
71+ 
72+template <class ArchTag>
73+struct Resource {
74+ AscendC::TPipe pipe;
75+ LocalTensorBuffer<ArchTag, AscendC::TPosition::A1> l1Buf;
76+ LocalTensorBuffer<ArchTag, AscendC::TPosition::A2> l0ABuf;
77+ LocalTensorBuffer<ArchTag, AscendC::TPosition::B2> l0BBuf;
78+ LocalTensorBuffer<ArchTag, AscendC::TPosition::C2> btBuf;
79+ LocalTensorBuffer<ArchTag, AscendC::TPosition::CO1> l0CBuf;
80+ LocalTensorBuffer<ArchTag, AscendC::TPosition::VECCALC> ubBuf;
81+ LocalTensorBuffer<ArchTag, AscendC::TPosition::C2PIPE2GM> fpBuf;
82+ 
83+ FTSELF_DEVICE Resource() { pipe.Destroy(); }
84+};
85+ 
86+} // namespace Arch
87+ 
88+} // namespace FTSelf
89+ 
90+#endif // FTSELF_CORE_ARCH_HPP
@@ -0,0 +1,116 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+#ifndef FTSELF_CORE_COORD_HPP
12+#define FTSELF_CORE_COORD_HPP
13+ 
14+#include <cstdint>
15+#include "../core/macros.hpp"
16+ 
17+namespace FTSelf {
18+ 
19+constexpr uint32_t BYTE_PER_BLK = 32;
20+constexpr uint32_t BYTE_PER_C0 = 32;
21+constexpr uint32_t BYTE_PER_C2 = 64;
22+constexpr uint32_t C0_NUM_PER_FRACTAL = 16;
23+constexpr uint32_t BYTE_PER_FRACTAL = BYTE_PER_C0 * C0_NUM_PER_FRACTAL;
24+constexpr uint32_t BLK_NUM_PER_VECTOR_FRACTAL = 8;
25+constexpr uint32_t BYTE_PER_VECTOR_FRACTAL = BYTE_PER_BLK * BLK_NUM_PER_VECTOR_FRACTAL;
26+constexpr uint64_t L2_OFFSET = 0;
27+constexpr uint32_t STRIDE_LIMIT = 65536;
28+ 
29+enum class Status { kSuccess, kInvalid };
30+ 
31+template <int Rank, class IndexType = uint32_t>
32+struct Coord {
33+ using Index = IndexType;
34+ static constexpr int RANK = Rank;
35+ Index data[Rank];
36+ 
37+ FTSELF_HOST_DEVICE constexpr Coord() : data{} {}
38+ template <class... Values>
39+ FTSELF_HOST_DEVICE constexpr explicit Coord(Values... values) : data{Index(values)...} {}
40+ FTSELF_HOST_DEVICE constexpr Index const &operator[](int index) const { return data[index]; }
41+ FTSELF_HOST_DEVICE constexpr Index &operator[](int index) { return data[index]; }
42+ FTSELF_HOST_DEVICE constexpr Index const &At(int index) const { return data[index]; }
43+ FTSELF_HOST_DEVICE constexpr Index &At(int index) { return data[index]; }
44+};
45+ 
46+template <class... Values>
47+FTSELF_HOST_DEVICE constexpr auto MakeCoord(Values... values)
48+{
49+ return Coord<sizeof...(Values)>{values...};
50+}
51+ 
52+struct MatrixCoord : Coord<2> {
53+ using Base = Coord<2>;
54+ FTSELF_HOST_DEVICE constexpr MatrixCoord(uint32_t row = 0, uint32_t column = 0) : Base(row, column) {}
55+ FTSELF_HOST_DEVICE constexpr MatrixCoord(Base const &coord) : Base(coord) {}
56+ FTSELF_HOST_DEVICE constexpr uint32_t const &row() const { return data[0]; }
57+ FTSELF_HOST_DEVICE constexpr uint32_t &row() { return data[0]; }
58+ FTSELF_HOST_DEVICE constexpr uint32_t const &column() const { return data[1]; }
59+ FTSELF_HOST_DEVICE constexpr uint32_t &column() { return data[1]; }
60+};
61+ 
62+struct GemvCoord : Coord<2> {
63+ using Base = Coord<2>;
64+ FTSELF_HOST_DEVICE constexpr GemvCoord(uint32_t m = 0, uint32_t n = 0) : Base(m, n) {}
65+ FTSELF_HOST_DEVICE constexpr GemvCoord(Base const &coord) : Base(coord) {}
66+ FTSELF_HOST_DEVICE constexpr uint32_t const &m() const { return data[0]; }
67+ FTSELF_HOST_DEVICE constexpr uint32_t &m() { return data[0]; }
68+ FTSELF_HOST_DEVICE constexpr uint32_t const &n() const { return data[1]; }
69+ FTSELF_HOST_DEVICE constexpr uint32_t &n() { return data[1]; }
70+};
71+ 
72+struct GemmCoord : Coord<3> {
73+ using Base = Coord<3>;
74+ FTSELF_HOST_DEVICE constexpr GemmCoord(uint32_t m = 0, uint32_t n = 0, uint32_t k = 0) : Base(m, n, k) {}
75+ FTSELF_HOST_DEVICE constexpr GemmCoord(Base const &coord) : Base(coord) {}
76+ FTSELF_HOST_DEVICE constexpr uint32_t const &m() const { return data[0]; }
77+ FTSELF_HOST_DEVICE constexpr uint32_t &m() { return data[0]; }
78+ FTSELF_HOST_DEVICE constexpr uint32_t const &n() const { return data[1]; }
79+ FTSELF_HOST_DEVICE constexpr uint32_t &n() { return data[1]; }
80+ FTSELF_HOST_DEVICE constexpr uint32_t const &k() const { return data[2]; }
81+ FTSELF_HOST_DEVICE constexpr uint32_t &k() { return data[2]; }
82+};
83+ 
84+template <uint32_t M_, uint32_t N_, uint32_t K_>
85+struct GemmShape {
86+ static constexpr uint32_t M = M_;
87+ static constexpr uint32_t N = N_;
88+ static constexpr uint32_t K = K_;
89+ static constexpr uint64_t MN = uint64_t(M) * N;
90+ static constexpr uint64_t MK = uint64_t(M) * K;
91+ static constexpr uint64_t NK = uint64_t(N) * K;
92+ static constexpr uint64_t MNK = MN * K;
93+ static constexpr uint64_t COUNT = MNK;
94+ FTSELF_HOST_DEVICE static constexpr GemmCoord ToCoord() { return GemmCoord(M, N, K); }
95+};
96+ 
97+template <uint32_t M_, uint32_t N_>
98+struct GemvShape {
99+ static constexpr uint32_t M = M_;
100+ static constexpr uint32_t N = N_;
101+ static constexpr uint64_t MN = uint64_t(M) * N;
102+ static constexpr uint64_t COUNT = MN;
103+ FTSELF_HOST_DEVICE static constexpr GemvCoord ToCoord() { return GemvCoord(M, N); }
104+};
105+ 
106+template <uint32_t Rows, uint32_t Columns>
107+struct MatrixShape {
108+ static constexpr uint32_t ROW = Rows;
109+ static constexpr uint32_t COLUMN = Columns;
110+ static constexpr uint64_t COUNT = uint64_t(Rows) * Columns;
111+ FTSELF_HOST_DEVICE static constexpr MatrixCoord ToCoord() { return MatrixCoord(Rows, Columns); }
112+};
113+ 
114+} // namespace FTSelf
115+ 
116+#endif // FTSELF_CORE_COORD_HPP
@@ -0,0 +1,61 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+#ifndef FTSELF_CORE_GEMM_TYPE_HPP
12+#define FTSELF_CORE_GEMM_TYPE_HPP
13+ 
14+#include "../core/arch.hpp"
15+#include "../helper/type_helper.hpp"
16+ 
17+namespace FTSelf {
18+ 
19+namespace Epilogue {}
20+ 
21+// FTSelf extends both Gemm and Gemv with its own Block/Tile/helper namespaces,
22+// so these must be real namespaces rather than namespace aliases.
23+namespace Gemm {
24+template <class ArchTag_, bool Async_ = false>
25+struct MmadBase {
26+ using ArchTag = ArchTag_;
27+ static constexpr bool ASYNC = Async_;
28+};
29+ 
30+using MmadAtlasA2 = MmadBase<Arch::AtlasA2, false>;
31+ 
32+template <bool EnableUnitFlag = false>
33+struct MmadAtlasA2Pingpong : MmadAtlasA2 {
34+ static constexpr uint32_t STAGES = 2;
35+ static constexpr bool ENABLE_UNIT_FLAG = EnableUnitFlag;
36+};
37+} // namespace Gemm
38+ 
39+template <class Element_, class Layout_, AscendC::TPosition Position_ = AscendC::TPosition::GM>
40+struct GemmType {
41+ using Element = Element_;
42+ using Layout = Layout_;
43+ static constexpr AscendC::TPosition POSITION = Position_;
44+};
45+ 
46+namespace Gemm {
47+using FTSelf::GemmType;
48+} // namespace Gemm
49+ 
50+namespace Gemv {
51+using FTSelf::GemmType;
52+} // namespace Gemv
53+ 
54+namespace core {
55+template <class ElementA, class ElementB>
56+using ElementAccumulatorSelector = FTSelf::helper::ElementAccumulatorSelector<ElementA, ElementB>;
57+} // namespace core
58+ 
59+} // namespace FTSelf
60+ 
61+#endif // FTSELF_CORE_GEMM_TYPE_HPP
@@ -0,0 +1,22 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+#ifndef FTSELF_CORE_LAYOUT_HPP
12+#define FTSELF_CORE_LAYOUT_HPP
13+ 
14+#include "../helper/layout_helper.hpp"
15+ 
16+namespace FTSelf {
17+ 
18+namespace layout = Gemv::layout;
19+ 
20+} // namespace FTSelf
21+ 
22+#endif // FTSELF_CORE_LAYOUT_HPP
@@ -0,0 +1,26 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+#ifndef FTSELF_CORE_MACROS_HPP
12+#define FTSELF_CORE_MACROS_HPP
13+ 
14+#if defined(__CCE__)
15+#include <kernel_operator.h>
16+#endif
17+ 
18+#define FTSELF_DEVICE __forceinline__ __aicore__
19+#ifdef __CCE__
20+#define FTSELF_HOST_DEVICE __forceinline__ [host, aicore]
21+#else
22+#define FTSELF_HOST_DEVICE
23+#endif
24+#define FTSELF_KERNEL __global__ __aicore__
25+ 
26+#endif // FTSELF_CORE_MACROS_HPP
@@ -0,0 +1,45 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+#ifndef FTSELF_DEBUG_HPP
12+#define FTSELF_DEBUG_HPP
13+ 
14+#undef inline
15+#include <iostream>
16+#include <sstream>
17+#include <functional>
18+#define inline __inline__ __attribute__((always_inline))
19+ 
20+#include <acl/acl.h>
21+ 
22+#define SINGLE_CORE_DUMPSIZE (1024 * 1024)
23+// 75 is from AscendC host stub
24+#define ALL_DUMPSIZE (75 * SINGLE_CORE_DUMPSIZE)
25+ 
26+namespace FTSelf {
27+ 
28+using LogFuncType = std::function<void(const char *)>;
29+inline void aclCheck(aclError status, LogFuncType logFunc = [](const char *logStrPtr) { std::cerr << logStrPtr; })
30+{
31+ if (status != ACL_SUCCESS) {
32+ std::stringstream ss;
33+ ss << "AclError: " << status;
34+ logFunc(ss.str().c_str());
35+ }
36+}
37+ 
38+void AdumpPrintWorkSpace(const void *dumpBufferAddr,
39+ const size_t dumpBufferSize,
40+ aclrtStream stream,
41+ const char *opType);
42+ 
43+} // namespace FTSelf
44+ 
45+#endif // FTSELF_DEBUG_HPP
@@ -0,0 +1,44 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+#ifndef FTSELF_DEVICE_KERNEL_ADAPTER_HPP
12+#define FTSELF_DEVICE_KERNEL_ADAPTER_HPP
13+ 
14+#if defined(ENABLE_ASCENDC_DUMP)
15+#include "../debug.hpp"
16+#endif
17+ 
18+namespace FTSelf {
19+ 
20+template <class Operator>
21+FTSELF_KERNEL void KernelAdapter(typename Operator::Params params, GM_ADDR ptrDump = nullptr)
22+{
23+ Operator op;
24+#if defined(ENABLE_ASCENDC_DUMP)
25+ AscendC::InitDump(false, ptrDump, ALL_DUMPSIZE);
26+#endif
27+ op(params);
28+}
29+ 
30+template <class Operator>
31+FTSELF_KERNEL void KernelAdapter(typename Operator::Params params, uint64_t fftsAddr,
32+ GM_ADDR ptrDump = nullptr)
33+{
34+ AscendC::SetSyncBaseAddr(fftsAddr);
35+ Operator op;
36+#if defined(ENABLE_ASCENDC_DUMP)
37+ AscendC::InitDump(false, ptrDump, ALL_DUMPSIZE);
38+#endif
39+ op(params);
40+}
41+ 
42+} // namespace FTSelf
43+ 
44+#endif // FTSELF_DEVICE_KERNEL_ADAPTER_HPP
@@ -0,0 +1,207 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+#ifndef FTSELF_FTSELF_GEMM_BLOCK_BLOCK_MMAD_HPP
Y
Yyue-ma25 天前

文件无copy right

likedislike
starfican
19 天前 评论:
12+#define FTSELF_FTSELF_GEMM_BLOCK_BLOCK_MMAD_HPP
13+ 
14+#include "../../gemm/tile/gemm_tile_copy.hpp"
15+#include "../../gemm/tile/tile_mmad.hpp"
16+#include "../../gemv/helper.hpp"
17+ 
18+namespace FTSelf::Gemm::Block {
19+ 
20+#define FTSELF_BLOCK_MMAD_BASIC_PARAMS \
21+ class DispatchPolicy, class L1TileShape, class L0TileShape, class AType, class BType, class CType, \
22+ class BiasType = void, \
23+ class TileCopy = FTSelf::Gemm::Tile::TileCopy<typename DispatchPolicy::ArchTag, AType, BType, CType, BiasType>, \
24+ class TileMmad = FTSelf::Gemm::Tile::TileMmad<typename DispatchPolicy::ArchTag, AType, BType, BiasType>
25+ 
26+#define FTSELF_BLOCK_MMAD_FT_PARAMS \
27+ class DispatchPolicy, FTSelf::Gemv::helper::FT_ENC_TYPE ENC_TYPE_, \
28+ FTSelf::Gemv::helper::FT_L02L1_TYPE COPY_TYPE_, class L1TileShape, class L0TileShape, \
29+ class L0TileShapeforFT, class AType, class BType, class CType, class XType, class YType, \
30+ class BiasType = void, \
31+ class TileCopyFT = FTSelf::Gemm::Tile::TileCopyFT<typename DispatchPolicy::ArchTag, AType, BType, CType, XType, YType, BiasType, COPY_TYPE_>, \
32+ class TileMmad = FTSelf::Gemm::Tile::TileMmad<typename DispatchPolicy::ArchTag, AType, BType, BiasType>
33+ 
34+#define FTSELF_BLOCK_MMAD_AB_PARAMS \
35+ class DispatchPolicy, class L1TileShape, class L1TileShapeforFT, class L0TileShape, \
36+ class L0TileShapeforFT, class AType, class BType, class CType, class XType, class YType, \
37+ class BiasType = void, \
38+ class TileCopyFTABonAic = FTSelf::Gemm::Tile::TileCopyFTABonAic<typename DispatchPolicy::ArchTag, AType, BType, CType, XType, YType, BiasType>, \
39+ class TileMmad = FTSelf::Gemm::Tile::TileMmad<typename DispatchPolicy::ArchTag, AType, BType, BiasType>
40+ 
41+#define FTSELF_BLOCK_MMAD_AUGED_PARAMS \
42+ class DispatchPolicy, class L1TileShape, class L1TileShapeforFT, class L0TileShape, \
43+ class L0TileShapeforFT, class AType, class BType, class CType, class XType, class XColType, \
44+ class YType, class BiasType = void, \
45+ class TileCopyFTABonAicAuged = FTSelf::Gemm::Tile::TileCopyFTABonAicAuged<typename DispatchPolicy::ArchTag, AType, BType, CType, XType, XColType, YType, BiasType>, \
46+ class TileMmad = FTSelf::Gemm::Tile::TileMmad<typename DispatchPolicy::ArchTag, AType, BType, BiasType>
47+ 
48+#define FTSELF_BLOCK_MMAD_SPEC_PARAMS \
49+ class DispatchPolicy, class L1TileShapeforFT, class L0TileShapeforFT, class AType, class BType, \
50+ class CType, class XType, class YType, class BiasType, \
51+ class TileCopyFTABonAic = FTSelf::Gemm::Tile::TileCopyFTABonAic<typename DispatchPolicy::ArchTag, AType, BType, CType, XType, YType, BiasType>, \
52+ class TileMmad = FTSelf::Gemm::Tile::TileMmad<typename DispatchPolicy::ArchTag, AType, XType, BiasType>
53+ 
54+ 
55+ 
56+template <FTSELF_BLOCK_MMAD_BASIC_PARAMS>
57+struct BlockMmadPreload {
58+ static_assert(FTSelf::helper::DEPENDENT_FALSE<DispatchPolicy>, "BlockMmadPreload is not implemented for this DispatchPolicy");
59+};
60+ 
61+ 
62+ 
63+// using TileMmadAIC = Gemm::Tile::TileMmad<typename GEMVAICDispatchPolicy::ArchTag, XType, CType, BiasType>;
64+// class TileMmadforFT = FTSelf::Gemm::Tile::TileMmad<typename DispatchPolicy::ArchTag,>
65+// = FTSelf::Gemv::helper::FT_L02L1_TYPE::FIX_PIPE
66+ 
67+ 
68+ 
69+template <FTSELF_BLOCK_MMAD_FT_PARAMS>
70+struct BlockMmadFTNOSPLIT {
71+ static_assert(FTSelf::helper::DEPENDENT_FALSE<DispatchPolicy>, "BlockMmad is not implemented for this DispatchPolicy");
72+};
73+ 
74+/*
75+struct BlockMmadFTABeNoSplitK<
76+ FTSelf::Gemm::MmadAtlasA2Pingpong<ENABLE_UNIT_FLAG_>,
77+ L1TileShape_,
78+ L1TileShapeforFT_,
79+ L0TileShape_,
80+ L0TileShapeforFT_,
81+ AType_,
82+ BType_,
83+ CType_,
84+ XType_,
85+ YType_,
86+ BiasType_,
87+ TileCopyFTABonAic_,
88+ TileMmad_
89+>
90+*/
91+template <FTSELF_BLOCK_MMAD_AB_PARAMS>
92+struct BlockMmadFTABeNoSplitK {
93+ /*
94+ L1TileShape_,
95+ L1TileShapeforFT_,
96+ L0TileShape_,
97+ L0TileShapeforFT_,
98+ */
99+ static_assert(FTSelf::helper::DEPENDENT_FALSE<DispatchPolicy>, "BlockMmad is not implemented for this DispatchPolicy");
100+};
101+ 
102+template <FTSELF_BLOCK_MMAD_AUGED_PARAMS>
103+struct BlockMmadFTABeAugedNoSplitK{
104+ 
105+ /*
106+ FTSelf::Gemm::MmadAtlasA2Pingpong<ENABLE_UNIT_FLAG_>,
107+ L1TileShape_,
108+ L1TileShapeforFT_,
109+ L0TileShape_,
110+ L0TileShapeforFT_,
111+ AType_,
112+ BType_,
113+ CType_,
114+ XType_,
115+ XColType_,
116+ YType_,
117+ BiasType_,
118+ TileCopyFTABonAicAuged_,
119+ TileMmad_
120+ */
121+ static_assert(FTSelf::helper::DEPENDENT_FALSE<DispatchPolicy>, "BlockMmadFTABeAugedNoSplitK is not implemented for this DispatchPolicy");
122+};
123+ 
124+ 
125+template <FTSELF_BLOCK_MMAD_AUGED_PARAMS>
126+struct BlockMmadFTABeAugedNoSplitKGemv{
127+ 
128+ /*
129+ FTSelf::Gemm::MmadAtlasA2Pingpong<ENABLE_UNIT_FLAG_>,
130+ L1TileShape_,
131+ L1TileShapeforFT_,
132+ L0TileShape_,
133+ L0TileShapeforFT_,
134+ AType_,
135+ BType_,
136+ CType_,
137+ XType_,
138+ XColType_,
139+ YType_,
140+ BiasType_,
141+ TileCopyFTABonAicAuged_,
142+ TileMmad_
143+ */
144+ static_assert(FTSelf::helper::DEPENDENT_FALSE<DispatchPolicy>, "BlockMmadFTABeAugedNoSplitKGemv is not implemented for this DispatchPolicy");
145+};
146+ 
147+template <FTSELF_BLOCK_MMAD_AUGED_PARAMS>
148+struct BlockMmadFTABeAugedNoSplitKRobust{
149+ 
150+ /*
151+ FTSelf::Gemm::MmadAtlasA2Pingpong<ENABLE_UNIT_FLAG_>,
152+ L1TileShape_,
153+ L1TileShapeforFT_,
154+ L0TileShape_,
155+ L0TileShapeforFT_,
156+ AType_,
157+ BType_,
158+ CType_,
159+ XType_,
160+ XColType_,
161+ YType_,
162+ BiasType_,
163+ TileCopyFTABonAicAuged_,
164+ TileMmad_
165+ */
166+ static_assert(FTSelf::helper::DEPENDENT_FALSE<DispatchPolicy>, "BlockMmadFTABeAugedNoSplitKRobust is not implemented for this DispatchPolicy");
167+};
168+ 
169+template <FTSELF_BLOCK_MMAD_FT_PARAMS>
170+struct BlockMmadFTSpiltK {
171+ static_assert(FTSelf::helper::DEPENDENT_FALSE<DispatchPolicy>, "BlockMmadFTSpiltK is not implemented for this DispatchPolicy");
172+};
173+ 
174+template <FTSELF_BLOCK_MMAD_SPEC_PARAMS>
175+struct BlockMmadSpecABeNoSplitK {
176+ static_assert(FTSelf::helper::DEPENDENT_FALSE<DispatchPolicy>, "BlockMmadSpecABeNoSplitK is not implemented for this DispatchPolicy");
177+};
178+ 
179+template <FTSELF_BLOCK_MMAD_SPEC_PARAMS>
180+struct BlockMmadSpecABeNoSplitKRobust {
181+ static_assert(FTSelf::helper::DEPENDENT_FALSE<DispatchPolicy>, "BlockMmadSpecABeNoSplitK is not implemented for this DispatchPolicy");
182+};
183+ 
184+template <FTSELF_BLOCK_MMAD_AB_PARAMS>
185+struct BlockMmadFTABeNoSplitKRobust {
186+ static_assert(FTSelf::helper::DEPENDENT_FALSE<DispatchPolicy>, "BlockMmad is not implemented for this DispatchPolicy");
187+};
188+ 
189+template <FTSELF_BLOCK_MMAD_BASIC_PARAMS>
190+struct BlockMmadFault {
191+ static_assert(FTSelf::helper::DEPENDENT_FALSE<DispatchPolicy>, "BlockMmadFault is not implemented for this DispatchPolicy");
192+};
193+ 
194+ 
195+#undef FTSELF_BLOCK_MMAD_SPEC_PARAMS
196+#undef FTSELF_BLOCK_MMAD_AUGED_PARAMS
197+#undef FTSELF_BLOCK_MMAD_AB_PARAMS
198+#undef FTSELF_BLOCK_MMAD_FT_PARAMS
199+#undef FTSELF_BLOCK_MMAD_BASIC_PARAMS
200+ 
201+} // namespace FTSelf::Gemm::Block
202+ 
203+#include "../../gemm/block/block_mmad_pingpong_fault_abe_spec_no_splitk_robust.hpp"
204+#include "../../gemm/block/block_mmad_pingpong_preload.hpp"
205+ 
206+ 
207+#endif // FTSELF_FTSELF_GEMM_BLOCK_BLOCK_MMAD_HPP