已开启
【代码侦探Challenge08】赛诺信致(北京)软件技术有限公司 - challenge08_tanhcustom #4054
【代码侦探Challenge08】赛诺信致(北京)软件技术有限公司 - challenge08_tanhcustom #4054
已开启
小东创建于 20 天前
共 36 个文件变更+3036-0
@@ -0,0 +1,15 @@
1+# 赛诺信致(北京)软件技术有限公司
2+ 
3+## 团队信息
4+ 
5+- 提交者: ciknife
6+- 身份: 企业员工
7+- 单位: 赛诺信致(北京)软件技术有限公司
8+ 
9+## 成员
10+ 
11+- ciknife (ciknife): 提交者
12+ 
13+## 算子: challenge08_tanhcustom
14+ 
15+多核实现,UB 双缓冲流水线。Tiling 策略:总元素数按 AIV 核数均分,blockFactor 向上对齐到 16 元素(32B,兼容 float32/float16);数据量小于 1024×核数 时自动减少启动核数;ubFactor 按 UB 容量自适应(输入/输出 2 队列 × BUFFER_NUM=2),尾核取剩余元素防越界。Kernel 侧 CopyIn → AscendC::Tanh → CopyOut 三段式,逐元素计算 y = (exp(x) - exp(-x)) / (exp(x) + exp(-x)),最后一个 tile 取剩余量。
@@ -0,0 +1,29 @@
1+cmake_minimum_required(VERSION 3.16.0)
2+project(tanh_op_prj)
3+find_package(ASC REQUIRED)
4+set(CMAKE_CXX_STANDARD 17)
5+set(CMAKE_CXX_STANDARD_REQUIRED ON)
6+ 
7+set(ARCH32_COMPUTE_UNITS ascend910b ascend910_93)
8+set(ARCH35_COMPUTE_UNITS ascend950)
9+ 
10+if(NOT DEFINED ASCEND_COMPUTE_UNIT OR ASCEND_COMPUTE_UNIT STREQUAL "")
11+ set(ASCEND_COMPUTE_UNIT ${ARCH32_COMPUTE_UNITS} ${ARCH35_COMPUTE_UNITS})
12+endif()
13+set(package_name tanh_custom)
14+ 
15+ npu_op_package(${package_name}
16+ TYPE RUN
17+ CONFIG
18+ INSTALL_PATH ${CMAKE_BINARY_DIR}
19+)
20+ 
21+if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/op_host")
22+ add_subdirectory(op_host)
23+endif()
24+ 
25+if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/op_kernel")
26+ add_subdirectory(op_kernel)
27+endif()
28+ 
29+message(WARNING "cmake 'make' does NOT build kernel binary by default. Use: bash build.sh --soc=<soc>")
@@ -0,0 +1,4 @@
1+# challenge08_tanhcustom
2+ 
3+提交团队: 赛诺信致(北京)软件技术有限公司
4+提交者: ciknife
@@ -0,0 +1,158 @@
1+#!/bin/bash
2+set -e
3+ 
4+export BASE_PATH=$(
5+ cd "$(dirname $0)"
6+ pwd
7+)
8+export BUILD_PATH="${BASE_PATH}/build"
9+export BUILD_OUT_PATH="${BASE_PATH}/build_out"
10+ 
11+CORE_NUMS=$(cat /proc/cpuinfo | grep "processor" | wc -l)
12+if [ ${CORE_NUMS} -gt 8 ]; then
13+ CORE_NUMS=8
14+fi
15+ 
16+usage() {
17+ echo "Build script for tanh operator"
18+ echo "Usage: bash build.sh [OPTIONS]"
19+ echo ""
20+ echo "Options:"
21+ echo " -h, --help Print this help message"
22+ echo " -j[n] Compile thread nums, default is ${CORE_NUMS}, eg: -j8"
23+ echo " --make_clean Clean build artifacts"
24+ echo " -u, --ut Run UT (Unit Tests)"
25+ echo " -e, --example Run examples (requires NPU)"
26+ echo ""
27+ echo "Examples:"
28+ echo " bash build.sh # Build with default soc (ascend910b)"
29+ echo " bash build.sh -j8 # Build with 8 threads"
30+ echo " bash build.sh --make_clean"
31+ echo " bash build.sh -u # Run UT tests"
32+ echo " bash build.sh -e # Run aclnn example (requires NPU)"
33+}
34+ 
35+clean_build() {
36+ if [ -d "${BUILD_PATH}" ]; then
37+ echo "Cleaning build directory..."
38+ rm -rf ${BUILD_PATH}/*
39+ fi
40+}
41+ 
42+clean_build_out() {
43+ if [ -d "${BUILD_OUT_PATH}" ]; then
44+ echo "Cleaning build_out directory..."
45+ rm -rf ${BUILD_OUT_PATH}/*
46+ fi
47+}
48+ 
49+THREAD_NUM=${CORE_NUMS}
50+COMPUTE_UNIT="ascend910b"
51+ENABLE_CLEAN=FALSE
52+RUN_UT=FALSE
53+RUN_EXAMPLE=FALSE
54+ 
55+while [[ $# -gt 0 ]]; do
56+ case "$1" in
57+ -h|--help)
58+ usage
59+ exit 0
60+ ;;
61+ -j*)
62+ THREAD_NUM="${1:2}"
63+ if [ -z "$THREAD_NUM" ]; then
64+ THREAD_NUM=${CORE_NUMS}
65+ fi
66+ shift
67+ ;;
68+ -u|--ut)
69+ RUN_UT=true
70+ shift
71+ ;;
72+ -e|--example)
73+ RUN_EXAMPLE=true
74+ shift
75+ ;;
76+ --make_clean)
77+ ENABLE_CLEAN=TRUE
78+ shift
79+ ;;
80+ -*)
81+ echo "[ERROR] Invalid option: $1"
82+ usage
83+ exit 1
84+ ;;
85+ *)
86+ echo "[ERROR] Unexpected argument: $1"
87+ usage
88+ exit 1
89+ ;;
90+ esac
91+done
92+ 
93+if [ "$ENABLE_CLEAN" = "TRUE" ]; then
94+ clean_build
95+ clean_build_out
96+ exit 0
97+fi
98+ 
99+if [ "$RUN_UT" = true ]; then
100+ echo "[INFO] Running UT tests..."
101+ cd "${BASE_PATH}/tests/ut"
102+ ./run.sh
103+ UT_RESULT=$?
104+ if [ $UT_RESULT -ne 0 ]; then
105+ echo "[ERROR] UT tests failed"
106+ exit 1
107+ fi
108+ echo "[INFO] UT tests passed!"
109+ exit 0
110+fi
111+ 
112+CMAKE_ARGS="-DASCEND_COMPUTE_UNIT=$COMPUTE_UNIT"
113+ 
114+if [ ! -d "${BUILD_PATH}" ]; then
115+ mkdir -p "${BUILD_PATH}"
116+fi
117+ 
118+[ -f "${BUILD_PATH}/CMakeCache.txt" ] && rm -f ${BUILD_PATH}/CMakeCache.txt
119+ 
120+echo "----------------------------------------------------------------"
121+echo "[INFO] Configuring project..."
122+echo "[INFO] CMAKE_ARGS: ${CMAKE_ARGS}"
123+cd "${BUILD_PATH}" && cmake ${CMAKE_ARGS} ..
124+ 
125+echo "----------------------------------------------------------------"
126+echo "[INFO] Building project with ${THREAD_NUM} threads..."
127+cmake --build . --target all binary package install -- -j ${THREAD_NUM}
128+ 
129+KERNEL_O=$(find ${BUILD_PATH}/op_kernel/ascendc_kernels/binary/${COMPUTE_UNIT} -name "*.o" 2>/dev/null | head -1)
130+if [ -z "$KERNEL_O" ]; then
131+ echo "[ERROR] Kernel binary not found"
132+ exit 1
133+fi
134+ 
135+PKG_PATH=$(ls "${BUILD_PATH}"/custom_opp_*.run 2>/dev/null | head -n 1)
136+if [ -z "$PKG_PATH" ] || [ ! -f "$PKG_PATH" ] || [ ! -s "$PKG_PATH" ]; then
137+ echo "[ERROR] Package not found or empty"
138+ exit 1
139+fi
140+ 
141+echo "----------------------------------------------------------------"
142+echo "[INFO] Build completed successfully!"
143+echo "[INFO] Kernel binary: ${KERNEL_O}"
144+echo "[INFO] Package: ${PKG_PATH}"
145+ 
146+if [ "$RUN_EXAMPLE" = true ]; then
147+ echo "----------------------------------------------------------------"
148+ echo "[INFO] Running examples..."
149+ cd "${BASE_PATH}/examples"
150+ ./run.sh
151+ EXAMPLE_RESULT=$?
152+ cd - > /dev/null
153+ if [ $EXAMPLE_RESULT -ne 0 ]; then
154+ echo "[ERROR] Example execution failed"
155+ exit 1
156+ fi
157+ echo "[INFO] Example completed successfully!"
158+fi
@@ -0,0 +1,58 @@
1+cmake_minimum_required(VERSION 3.14)
2+project(ACLNN_EXAMPLE)
3+ 
4+add_compile_options(-std=c++17)
5+set(CMAKE_RUNTIME_OUTPUT_DIRECTORY "./bin")
6+set(CMAKE_CXX_FLAGS_DEBUG "-fPIC -O0 -g -Wall")
7+set(CMAKE_CXX_FLAGS_RELEASE "-fPIC -O2 -Wall")
8+ 
9+add_executable(test_aclnn_tanh
10+test_aclnn_tanh.cpp)
11+ 
12+if(NOT "$ENV{ASCEND_HOME_PATH}" STREQUAL "")
13+ set(ASCEND_PATH $ENV{ASCEND_HOME_PATH})
14+else()
15+ set(ASCEND_PATH "/usr/local/Ascend/cann")
16+endif()
17+ 
18+find_path(CUSTOM_OP_INCLUDE_DIR
19+ NAMES aclnn_tanh.h
20+ PATHS
21+ ${ASCEND_PATH}/opp/vendors/tanh_custom/op_api/include
22+ /usr/local/Ascend/opp/vendors/tanh_custom/op_api/include
23+ $ENV{HOME}/Ascend/opp/vendors/tanh_custom/op_api/include
24+)
25+ 
26+if(NOT CUSTOM_OP_INCLUDE_DIR)
27+ message(FATAL_ERROR "未找到自定义算子头文件 aclnn_tanh.h,请先安装算子包")
28+endif()
29+message(STATUS "自定义算子头文件目录: ${CUSTOM_OP_INCLUDE_DIR}")
30+ 
31+find_library(CUSTOM_OP_LIBRARY cust_opapi
32+ PATHS
33+ ${ASCEND_PATH}/opp/vendors/tanh_custom/op_api/lib
34+ /usr/local/Ascend/opp/vendors/tanh_custom/op_api/lib
35+ $ENV{HOME}/Ascend/opp/vendors/tanh_custom/op_api/lib
36+)
37+ 
38+if(NOT CUSTOM_OP_LIBRARY)
39+ message(FATAL_ERROR "未找到自定义算子库 libcust_opapi.so,请先安装算子包")
40+endif()
41+ 
42+include_directories(
43+ ${ASCEND_PATH}/include
44+ ${CUSTOM_OP_INCLUDE_DIR}
45+)
46+ 
47+target_link_libraries(test_aclnn_tanh PRIVATE
48+ ${CUSTOM_OP_LIBRARY}
49+ ${ASCEND_PATH}/lib64/libascendcl.so
50+ ${ASCEND_PATH}/lib64/libnnopbase.so
51+ ${ASCEND_PATH}/lib64/libopapi.so
52+)
53+get_filename_component(CUSTOM_OP_LIB_DIR ${CUSTOM_OP_LIBRARY} DIRECTORY)
54+target_link_options(test_aclnn_tanh PRIVATE
55+ "-Wl,-rpath,${CUSTOM_OP_LIB_DIR}"
56+)
57+ 
58+install(TARGETS test_aclnn_tanh DESTINATION ${CMAKE_RUNTIME_OUTPUT_DIRECTORY})
@@ -0,0 +1,30 @@
1+#!/bin/bash
2+# tanh 算子调用示例执行脚本
3+ 
4+set -e
5+ 
6+SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
7+BUILD_DIR="${SCRIPT_DIR}/build"
8+ 
9+echo "========================================"
10+echo "tanh 算子调用示例"
11+echo "========================================"
12+ 
13+if [ -z "$ASCEND_HOME_PATH" ]; then
14+ export ASCEND_HOME_PATH=/usr/local/Ascend/cann
15+fi
16+ 
17+export LD_LIBRARY_PATH=${ASCEND_HOME_PATH}/lib64:${LD_LIBRARY_PATH}
18+ 
19+mkdir -p "${BUILD_DIR}"
20+cd "${BUILD_DIR}"
21+cmake ..
22+make -j$(nproc)
23+ 
24+echo "执行调用示例..."
25+cd bin
26+./test_aclnn_tanh
27+ 
28+echo "========================================"
29+echo "执行完成"
30+echo "========================================"
@@ -0,0 +1,184 @@
1+#include <iostream>
2+#include <vector>
3+#include <cstring>
4+#include <cstdint>
5+#include <algorithm>
6+#include "acl/acl.h"
7+#include "aclnn_tanh.h"
8+ 
9+#define CHECK_RET(cond, return_expr) \
10+ do { \
11+ if (!(cond)) { \
12+ return_expr; \
13+ } \
14+ } while (0)
15+ 
16+#define LOG_PRINT(message, ...) \
17+ do { \
18+ printf(message, ##__VA_ARGS__); \
19+ } while (0)
20+ 
21+int64_t GetShapeSize(const std::vector<int64_t>& shape)
22+{
23+ int64_t shapeSize = 1;
24+ for (auto i : shape) {
25+ shapeSize *= i;
26+ }
27+ return shapeSize;
28+}
29+ 
30+int Init(int32_t deviceId, aclrtStream* stream)
31+{
32+ auto ret = aclInit(nullptr);
33+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclInit failed. ERROR: %d\n", ret); return ret);
34+ ret = aclrtSetDevice(deviceId);
35+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetDevice failed. ERROR: %d\n", ret); return ret);
36+ ret = aclrtCreateStream(stream);
37+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateStream failed. ERROR: %d\n", ret); return ret);
38+ return 0;
39+}
40+ 
41+ 
42+static uint16_t FloatToHalf(float f) {
43+ uint32_t bits;
44+ memcpy(&bits, &f, sizeof(float));
45+ uint32_t sign = (bits >> 16) & 0x8000;
46+ int32_t exp = ((bits >> 23) & 0xff) - 127 + 15;
47+ uint32_t mant = (bits >> 13) & 0x3ff;
48+ if (exp <= 0) return sign;
49+ if (exp >= 31) return sign | 0x7c00;
50+ return sign | (exp << 10) | mant;
51+}
52+ 
53+static uint16_t FloatToBFloat16(float f) {
54+ uint32_t bits;
55+ memcpy(&bits, &f, sizeof(float));
56+ return (uint16_t)(bits >> 16);
57+}
58+ 
59+template <typename T>
60+int CreateAclTensor(
61+ const std::vector<T>& hostData, const std::vector<int64_t>& shape, void** deviceAddr, aclDataType dataType,
62+ aclTensor** tensor)
63+{
64+ auto elemCount = GetShapeSize(shape);
65+ int64_t elemSize = sizeof(T);
66+ switch (dataType) {
67+ case aclDataType::ACL_FLOAT16:
68+ case aclDataType::ACL_BF16:
69+ case aclDataType::ACL_INT16:
70+ case aclDataType::ACL_UINT16:
71+ elemSize = 2;
72+ break;
73+ case aclDataType::ACL_INT8:
74+ case aclDataType::ACL_UINT8:
75+ case aclDataType::ACL_BOOL:
76+ elemSize = 1;
77+ break;
78+ case aclDataType::ACL_INT64:
79+ case aclDataType::ACL_UINT64:
80+ case aclDataType::ACL_DOUBLE:
81+ elemSize = 8;
82+ break;
83+ default:
84+ break;
85+ }
86+ auto size = elemCount * elemSize;
87+ auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST);
88+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMalloc failed. ERROR: %d\n", ret); return ret);
89+ 
90+ std::vector<uint8_t> convBuf(size);
91+ if (dataType == aclDataType::ACL_FLOAT16) {
92+ for (int64_t i = 0; i < elemCount; i++) {
93+ uint16_t h = FloatToHalf(static_cast<float>(hostData[i]));
94+ memcpy(convBuf.data() + i * 2, &h, 2);
95+ }
96+ } else if (dataType == aclDataType::ACL_BF16) {
97+ for (int64_t i = 0; i < elemCount; i++) {
98+ uint16_t b = FloatToBFloat16(static_cast<float>(hostData[i]));
99+ memcpy(convBuf.data() + i * 2, &b, 2);
100+ }
101+ } else if (dataType == aclDataType::ACL_DOUBLE) {
102+ for (int64_t i = 0; i < elemCount; i++) {
103+ double d = static_cast<double>(hostData[i]);
104+ memcpy(convBuf.data() + i * 8, &d, 8);
105+ }
106+ } else {
107+ auto copySize = std::min((int64_t)(elemCount * sizeof(T)), size);
108+ memcpy(convBuf.data(), hostData.data(), copySize);
109+ }
110+ ret = aclrtMemcpy(*deviceAddr, size, convBuf.data(), size, ACL_MEMCPY_HOST_TO_DEVICE);
111+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMemcpy failed. ERROR: %d\n", ret); return ret);
112+ 
113+ std::vector<int64_t> strides(shape.size(), 1);
114+ for (int64_t i = shape.size() - 2; i >= 0; i--) {
115+ strides[i] = shape[i + 1] * strides[i + 1];
116+ }
117+ 
118+ *tensor = aclCreateTensor(
119+ shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(),
120+ *deviceAddr);
121+ return 0;
122+}
123+ 
124+int main()
125+{
126+ int32_t deviceId = 0;
127+ aclrtStream stream;
128+ auto ret = Init(deviceId, &stream);
129+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret);
130+ 
131+ // 构造输入 tensor
132+ aclTensor* x = nullptr;
133+ void* xDeviceAddr = nullptr;
134+ std::vector<int64_t> xShape = {8, 2048};
135+ std::vector<float> xHostData(16384, 1);
136+ ret = CreateAclTensor(xHostData, xShape, &xDeviceAddr, aclDataType::ACL_FLOAT, &x);
137+ CHECK_RET(ret == ACL_SUCCESS, return ret);
138+ 
139+ // 构造输出 tensor
140+ aclTensor* y = nullptr;
141+ void* yDeviceAddr = nullptr;
142+ std::vector<int64_t> yShape = {8, 2048};
143+ std::vector<float> yHostData(16384, 0);
144+ ret = CreateAclTensor(yHostData, yShape, &yDeviceAddr, aclDataType::ACL_FLOAT, &y);
145+ CHECK_RET(ret == ACL_SUCCESS, return ret);
146+
147+ 
148+ // 调用 aclnnTanhGetWorkspaceSize 第一段接口
149+ uint64_t workspaceSize = 0;
150+ aclOpExecutor* executor = nullptr;
151+ ret = aclnnTanhGetWorkspaceSize(x, y, &workspaceSize, &executor);
152+ CHECK_RET(ret == 0, LOG_PRINT("aclnnTanhGetWorkspaceSize failed. ERROR: %d\n", ret); return ret);
153+ 
154+ // 申请 workspace
155+ void* workspaceAddr = nullptr;
156+ if (workspaceSize > 0) {
157+ ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST);
158+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("allocate workspace failed. ERROR: %d\n", ret); return ret);
159+ }
160+ 
161+ // 调用 aclnnTanh 第二段接口
162+ ret = aclnnTanh(workspaceAddr, workspaceSize, executor, stream);
163+ CHECK_RET(ret == 0, LOG_PRINT("aclnnTanh failed. ERROR: %d\n", ret); return ret);
164+ 
165+ // 同步等待
166+ ret = aclrtSynchronizeStream(stream);
167+ CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSynchronizeStream failed. ERROR: %d\n", ret); return ret);
168+ 
169+ // 释放资源
170+ aclDestroyTensor(x);
171+ aclrtFree(xDeviceAddr);
172+
173+ aclDestroyTensor(y);
174+ aclrtFree(yDeviceAddr);
175+ if (workspaceSize > 0) {
176+ aclrtFree(workspaceAddr);
177+ }
178+ 
179+ aclrtDestroyStream(stream);
180+ aclrtResetDevice(deviceId);
181+ aclFinalize();
182+ 
183+ return 0;
184+}
@@ -0,0 +1,101 @@
1+file(GLOB host_ops_def_srcs
2+ ${CMAKE_CURRENT_SOURCE_DIR}/*def.cpp
3+)
4+ 
5+file(GLOB host_ops_infershape_srcs
6+ ${CMAKE_CURRENT_SOURCE_DIR}/*_infershape.cpp
7+)
8+ 
9+set(host_ops_tiling_srcs)
10+file(GLOB TILING_FILES ${CMAKE_CURRENT_SOURCE_DIR}/*tiling.cpp)
11+list(APPEND host_ops_tiling_srcs ${TILING_FILES})
12+ 
13+set(host_ops_srcs
14+ ${host_ops_def_srcs}
15+ ${host_ops_infershape_srcs}
16+ ${host_ops_tiling_srcs}
17+)
18+ 
19+npu_op_code_gen(
20+ SRC ${host_ops_srcs}
21+ PACKAGE ${package_name}
22+ OUT_DIR ${ASCEND_AUTOGEN_PATH}
23+ COMPILE_OPTIONS
24+ -I$ENV{ASCEND_HOME_PATH}/aarch64-linux/include
25+ -I$ENV{ASCEND_HOME_PATH}/aarch64-linux/asc/include/tiling
26+ -I$ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc
27+ -I$ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/op_common
28+ -I$ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/base
29+ -I$ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/exe_graph
30+ -I$ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/graph
31+)
32+ 
33+npu_op_library(cust_optiling TILING
34+ ${host_ops_srcs}
35+)
36+ 
37+target_include_directories(cust_optiling PRIVATE
38+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/include
39+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/asc/include/tiling
40+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc
41+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/op_common
42+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/base
43+)
44+ 
45+set(op_api_dir ${CMAKE_CURRENT_SOURCE_DIR}/../op_api)
46+if(EXISTS ${op_api_dir} AND IS_DIRECTORY ${op_api_dir})
47+ file(GLOB op_api_srcs ${op_api_dir}/*.cpp)
48+else()
49+ file(GLOB op_api_srcs "${CMAKE_BINARY_DIR}/autogen/aclnn_*.cpp")
50+endif()
51+ 
52+npu_op_library(cust_opapi ACLNN
53+ ${op_api_srcs}
54+)
55+ 
56+target_include_directories(cust_opapi PRIVATE
57+ ${op_api_dir}
58+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/include
59+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/include/aclnn
60+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/asc/include
61+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc
62+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/op_common
63+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/base
64+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/aicpu
65+)
66+ 
67+target_compile_options(cust_opapi PRIVATE -UACLNN_WITH_BINARY)
68+ 
69+file(GLOB proto_src ${ASCEND_AUTOGEN_PATH}/op_proto.cc)
70+set_source_files_properties(${proto_src} PROPERTIES GENERATED TRUE)
71+ 
72+npu_op_library(cust_op_proto GRAPH
73+ ${host_ops_srcs}
74+ ${proto_src}
75+)
76+ 
77+target_include_directories(cust_op_proto PRIVATE
78+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/include
79+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/asc/include/tiling
80+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc
81+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/op_common
82+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/base
83+)
84+ 
85+npu_op_package_add(${package_name}
86+ LIBRARY
87+ cust_optiling
88+ cust_op_proto
89+ cust_opapi
90+)
91+ 
92+# Install hand-written op_api headers into RUN package
93+if(EXISTS ${op_api_dir} AND IS_DIRECTORY ${op_api_dir})
94+ file(GLOB op_api_aclnn_headers ${op_api_dir}/aclnn_*.h ${op_api_dir}/acl_*.h)
95+ if(op_api_aclnn_headers)
96+ npu_op_package_add(${package_name}
97+ FILES ${op_api_aclnn_headers}
98+ TYPE ACLNN
99+ )
100+ endif()
101+endif()
@@ -0,0 +1,29 @@
1+/*!
2+ * \file tanh_def.cpp
3+ * \brief Tanh 算子定义
4+ */
5+#include "register/op_def_registry.h"
6+ 
7+namespace ops {
8+class Tanh : public OpDef {
9+public:
10+ explicit Tanh(const char* name) : OpDef(name)
11+ {
12+ this->Input("x")
13+ .ParamType(REQUIRED)
14+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT16})
15+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
16+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND})
17+ .AutoContiguous();
18+ this->Output("y")
19+ .ParamType(REQUIRED)
20+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT16})
21+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
22+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND})
23+ .AutoContiguous();
24+ 
25+ this->AICore().AddConfig("ascend910b");
26+ }
27+};
28+OP_ADD(Tanh);
29+} // namespace ops
@@ -0,0 +1,37 @@
1+/*!
2+ * \file tanh_infershape.cpp
3+ * \brief Tanh 算子形状推导实现
4+ */
5+ 
6+#include "register/op_impl_registry.h"
7+#include "exe_graph/runtime/infer_shape_context.h"
8+ 
9+using namespace ge;
10+ 
11+namespace ops {
12+ 
13+static ge::graphStatus InferShapeTanh(gert::InferShapeContext* context)
14+{
15+ // TODO: 实现形状推导逻辑
16+ const gert::Shape* input_shape = (1 > 0) ? context->GetInputShape(0) : nullptr;
17+ // 注意:无输入算子时 input_shape 为 nullptr,需在此处手动设置输出 shape
18+ 
19+ for (size_t i = 0; i < 1; i++) {
20+ gert::Shape* output_shape = context->GetOutputShape(i);
21+ if (output_shape == nullptr) {
22+ return ge::GRAPH_FAILED;
23+ }
24+ const gert::Shape* in_shape = (i < 1) ? context->GetInputShape(i) : input_shape;
25+ if (in_shape == nullptr) {
26+ in_shape = input_shape;
27+ }
28+ if (in_shape != nullptr) {
29+ *output_shape = *in_shape;
30+ }
31+ }
32+ return ge::GRAPH_SUCCESS;
33+}
34+ 
35+IMPL_OP_INFERSHAPE(Tanh).InferShape(InferShapeTanh);
36+ 
37+} // namespace ops
@@ -0,0 +1,125 @@
1+/*!
2+ * \file tanh_tiling.cpp
3+ * \brief Tanh 算子 Tiling 实现
4+ */
5+ 
6+#include <algorithm>
7+#include "register/op_def_registry.h"
8+#include "op_common/log/log.h"
9+#include "op_common/op_host/util/math_util.h"
10+#include "op_common/op_host/util/platform_util.h"
11+#include "../op_kernel/tanh_tiling_data.h"
12+#include "../op_kernel/tanh_tiling_key.h"
13+ 
14+namespace optiling {
15+ 
16+using Ops::Base::CeilDiv;
17+using Ops::Base::CeilAlign;
18+using Ops::Base::FloorDiv;
19+using Ops::Base::FloorAlign;
20+using Ops::Base::GetUbBlockSize;
21+ 
22+constexpr uint32_t WS_SYS_SIZE = 0U;
23+constexpr int64_t TYPE_SIZE = 4;
24+constexpr int64_t MIN_SPLIT_THRESHOLD = 1024;
25+constexpr int64_t BUFFER_NUM = 2; // 与 kernel 侧 Double Buffer 保持一致
26+constexpr int64_t QUEUE_NUM = 2; // 输入队列 + 输出队列
27+ 
28+static const gert::Shape g_vec_1_shape = {1};
29+ 
30+static inline const gert::Shape EnsureNotScalar(const gert::Shape& in_shape) {
31+ if (in_shape.GetDimNum() == 0) {
32+ return g_vec_1_shape;
33+ }
34+ return in_shape;
35+}
36+ 
37+static ge::graphStatus GetPlatformInfo(gert::TilingContext* context, uint64_t& ubSize, int64_t& coreNum)
38+{
39+ fe::PlatFormInfos* platformInfoPtr = context->GetPlatformInfo();
40+ OP_CHECK_NULL_WITH_CONTEXT(context, platformInfoPtr);
41+ auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
42+ coreNum = ascendcPlatform.GetCoreNumAiv();
43+ OP_CHECK_IF(coreNum == 0, OP_LOGE(context, "coreNum is 0"), return ge::GRAPH_FAILED);
44+ ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
45+ OP_CHECK_IF(ubSize == 0, OP_LOGE(context, "ubSize is 0"), return ge::GRAPH_FAILED);
46+ return ge::GRAPH_SUCCESS;
47+}
48+ 
49+static ge::graphStatus GetWorkspaceSize(gert::TilingContext* context)
50+{
51+ size_t* currentWorkspace = context->GetWorkspaceSizes(1);
52+ OP_CHECK_NULL_WITH_CONTEXT(context, currentWorkspace);
53+ currentWorkspace[0] = WS_SYS_SIZE;
54+ return ge::GRAPH_SUCCESS;
55+}
56+ 
57+static ge::graphStatus TanhTilingFunc(gert::TilingContext* context)
58+{
59+ uint64_t ubSize;
60+ int64_t coreNum;
61+ OP_CHECK_IF(
62+ GetPlatformInfo(context, ubSize, coreNum) != ge::GRAPH_SUCCESS,
63+ OP_LOGE(context, "GetPlatformInfo error"),
64+ return ge::GRAPH_FAILED);
65+ 
66+ OP_CHECK_IF(
67+ GetWorkspaceSize(context) != ge::GRAPH_SUCCESS,
68+ OP_LOGE(context, "GetWorkspaceSize error"),
69+ return ge::GRAPH_FAILED);
70+ 
71+ TanhTilingData* tiling = context->GetTilingData<TanhTilingData>();
72+ OP_CHECK_NULL_WITH_CONTEXT(context, tiling);
73+ 
74+ // 1. 计算总元素数:输入 shape 各维相乘(标量按 1 个元素处理)
75+ const gert::Shape& inputShape = EnsureNotScalar(context->GetInputShape(0)->GetStorageShape());
76+ int64_t totalNum = 1;
77+ for (size_t i = 0; i < inputShape.GetDimNum(); i++) {
78+ totalNum *= inputShape.GetDim(i);
79+ }
80+ 
81+ // 2. 对齐粒度:DataCopy 要求 32B 对齐,float16 为 16 元素、float32 为 8 元素,取 16 兼容两种 dtype
82+ const int64_t alignNum = static_cast<int64_t>(GetUbBlockSize(context)) / 2; // 32B / 2B = 16
83+ 
84+ // 3. 确定使用核数:数据量过小时减少核数,保证每核至少处理 MIN_SPLIT_THRESHOLD 个元素
85+ int64_t usedCores = coreNum;
86+ if (totalNum < MIN_SPLIT_THRESHOLD * coreNum) {
87+ usedCores = CeilDiv(totalNum, MIN_SPLIT_THRESHOLD);
88+ }
89+ 
90+ // 4. 每核处理元素数,向上对齐到 alignNum,并重算实际启动核数
91+ int64_t blockFactor = CeilAlign(CeilDiv(totalNum, usedCores), alignNum);
92+ usedCores = CeilDiv(totalNum, blockFactor);
93+ 
94+ // 5. UB 单次循环处理量:输入/输出 2 个队列 × BUFFER_NUM × ubFactor × 4 字节 ≤ UB 容量
95+ int64_t ubFactor = FloorAlign(static_cast<int64_t>(ubSize) / (QUEUE_NUM * BUFFER_NUM * TYPE_SIZE), alignNum);
96+ ubFactor = std::min(ubFactor, blockFactor);
97+ 
98+ tiling->totalNum = totalNum;
99+ tiling->blockFactor = blockFactor;
100+ tiling->ubFactor = ubFactor;
101+ 
102+ context->SetBlockDim(static_cast<uint32_t>(usedCores));
103+ 
104+ // 根据输入 dtype 选择 tilingKey
105+ uint64_t tilingKey;
106+ auto inputDesc = context->GetInputDesc(0);
107+ if (inputDesc != nullptr && (inputDesc->GetDataType() == ge::DT_FLOAT16 || inputDesc->GetDataType() == ge::DT_BF16)) {
108+ tilingKey = GET_TPL_TILING_KEY(TANH_TPL_SCH_MODE_0);
109+ } else {
110+ tilingKey = GET_TPL_TILING_KEY(TANH_TPL_SCH_MODE_1);
111+ }
112+ context->SetTilingKey(tilingKey);
113+ return ge::GRAPH_SUCCESS;
114+}
115+ 
116+static ge::graphStatus TilingParseForTanh([[maybe_unused]] gert::TilingParseContext* context)
117+{
118+ return ge::GRAPH_SUCCESS;
119+}
120+ 
121+struct TanhCompileInfo {};
122+ 
123+IMPL_OP_OPTILING(Tanh).Tiling(TanhTilingFunc).TilingParse<TanhCompileInfo>(TilingParseForTanh);
124+ 
125+} // namespace optiling
@@ -0,0 +1,17 @@
1+file(GLOB_RECURSE ALL_KERNEL_FILES RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} *.cpp)
2+ 
3+npu_op_kernel_sources(ascendc_kernels
4+ OP_TYPE OP
5+ KERNEL_DIR .
6+ KERNEL_FILE ${ALL_KERNEL_FILES}
7+)
8+ 
9+npu_op_kernel_library(ascendc_kernels
10+ SRC_BASE ${CMAKE_CURRENT_SOURCE_DIR}
11+ TILING_LIBRARY cust_optiling
12+)
13+ 
14+npu_op_package_add(${package_name}
15+ LIBRARY
16+ ascendc_kernels
17+)
@@ -0,0 +1,29 @@
1+/*!
2+ * \file tanh.cpp
3+ * \brief Tanh 算子 kernel 入口
4+ */
5+ 
6+#include "tanh.h"
7+ 
8+enum class TanhTilingKey : uint32_t
9+{
10+ TILING_KEY_TANH_MODE_0 = 0,
11+ TILING_KEY_TANH_MODE_1 = 1,
12+};
13+ 
14+template <uint32_t schMode>
15+__global__ __aicore__ void tanh(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
16+{
17+ REGISTER_TILING_DEFAULT(TanhTilingData);
18+ GET_TILING_DATA_WITH_STRUCT(TanhTilingData, tilingData, tiling);
19+ if constexpr (schMode == static_cast<uint32_t>(TanhTilingKey::TILING_KEY_TANH_MODE_0)) {
20+ NsTanh::Tanh<half> op;
21+ op.Init(x, y, &tilingData);
22+ op.Process();
23+ }
24+ if constexpr (schMode == static_cast<uint32_t>(TanhTilingKey::TILING_KEY_TANH_MODE_1)) {
25+ NsTanh::Tanh<float> op;
26+ op.Init(x, y, &tilingData);
27+ op.Process();
28+ }
29+}
@@ -0,0 +1,111 @@
1+/*!
2+ * \file tanh.h
3+ * \brief Tanh 算子 kernel 类定义
4+ */
5+ 
6+#ifndef TANH_H
7+#define TANH_H
8+ 
9+#include "kernel_operator.h"
10+#include "kernel_tiling/kernel_tiling.h"
11+#include "tanh_tiling_data.h"
12+#include "tanh_tiling_key.h"
13+ 
14+namespace NsTanh {
15+ 
16+using namespace AscendC;
17+ 
18+constexpr int32_t BUFFER_NUM = 2;
19+ 
20+template <typename T>
21+class Tanh {
22+public:
23+ __aicore__ inline Tanh(){};
24+ 
25+ __aicore__ inline void Init(GM_ADDR x, GM_ADDR y, const TanhTilingData* tilingData);
26+ __aicore__ inline void Process();
27+ 
28+private:
29+ __aicore__ inline void CopyIn(int64_t progress, int64_t currentNum);
30+ __aicore__ inline void CopyOut(int64_t progress, int64_t currentNum);
31+ __aicore__ inline void Compute(int64_t currentNum);
32+ 
33+private:
34+ TPipe pipe;
35+ TQue<QuePosition::VECIN, BUFFER_NUM> inputQueueX;
36+ TQue<QuePosition::VECOUT, BUFFER_NUM> outputQueueY;
37+ 
38+ GlobalTensor<T> inputGMX;
39+ GlobalTensor<T> outputGMY;
40+ 
41+ int64_t blockLength_ = 0;
42+ int64_t ubLength_ = 0;
43+};
44+ 
45+template <typename T>
46+__aicore__ inline void Tanh<T>::Init(GM_ADDR x, GM_ADDR y, const TanhTilingData* tilingData)
47+{
48+ // 本核处理的 GM 起始偏移 = 核号 × 每核元素数
49+ int64_t start = static_cast<int64_t>(GetBlockIdx()) * tilingData->blockFactor;
50+ // 最后一核可能不足 blockFactor,取剩余元素数,防止越界读写
51+ blockLength_ = tilingData->blockFactor;
52+ if (start + blockLength_ > tilingData->totalNum) {
53+ blockLength_ = tilingData->totalNum - start;
54+ }
55+ ubLength_ = tilingData->ubFactor;
56+ 
57+ inputGMX.SetGlobalBuffer(reinterpret_cast<__gm__ T*>(x) + start, blockLength_);
58+ outputGMY.SetGlobalBuffer(reinterpret_cast<__gm__ T*>(y) + start, blockLength_);
59+ 
60+ // 为输入/输出队列分配 UB 内存(Double Buffer,共 BUFFER_NUM 块)
61+ pipe.InitBuffer(inputQueueX, BUFFER_NUM, ubLength_ * sizeof(T));
62+ pipe.InitBuffer(outputQueueY, BUFFER_NUM, ubLength_ * sizeof(T));
63+}
64+ 
65+template <typename T>
66+__aicore__ inline void Tanh<T>::CopyIn(int64_t progress, int64_t currentNum)
67+{
68+ // 分配 UB 空间 → 从 GM 搬入当前 tile → 入队
69+ LocalTensor<T> xLocal = inputQueueX.AllocTensor<T>();
70+ DataCopy(xLocal, inputGMX[progress * ubLength_], currentNum);
71+ inputQueueX.EnQue(xLocal);
72+}
73+ 
74+template <typename T>
75+__aicore__ inline void Tanh<T>::Compute(int64_t currentNum)
76+{
77+ // 输入出队 → 逐元素 Tanh(y = (exp(x) - exp(-x)) / (exp(x) + exp(-x)))→ 结果入队 → 释放输入
78+ LocalTensor<T> xLocal = inputQueueX.DeQue<T>();
79+ LocalTensor<T> yLocal = outputQueueY.AllocTensor<T>();
80+ AscendC::Tanh(yLocal, xLocal, currentNum);
81+ outputQueueY.EnQue(yLocal);
82+ inputQueueX.FreeTensor(xLocal);
83+}
84+ 
85+template <typename T>
86+__aicore__ inline void Tanh<T>::CopyOut(int64_t progress, int64_t currentNum)
87+{
88+ // 结果出队 → 搬回 GM → 释放 UB 空间
89+ LocalTensor<T> yLocal = outputQueueY.DeQue<T>();
90+ DataCopy(outputGMY[progress * ubLength_], yLocal, currentNum);
91+ outputQueueY.FreeTensor(yLocal);
92+}
93+ 
94+template <typename T>
95+__aicore__ inline void Tanh<T>::Process()
96+{
97+ // 按 ubLength_ 将本核数据切成若干 tile,流水线执行 CopyIn → Compute → CopyOut
98+ int64_t loopCount = (blockLength_ + ubLength_ - 1) / ubLength_;
99+ for (int64_t i = 0; i < loopCount; i++) {
100+ int64_t currentNum = ubLength_;
101+ if ((i + 1) * ubLength_ > blockLength_) {
102+ currentNum = blockLength_ - i * ubLength_; // 最后一个 tile 取剩余量
103+ }
104+ CopyIn(i, currentNum);
105+ Compute(currentNum);
106+ CopyOut(i, currentNum);
107+ }
108+}
109+ 
110+} // namespace NsTanh
111+#endif // TANH_H
@@ -0,0 +1,14 @@
1+/*!
2+ * \file tanh_tiling_data.h
3+ * \brief tiling data struct
4+ */
5+ 
6+#ifndef _TANH_TILING_DATA_H_
7+#define _TANH_TILING_DATA_H_
8+ 
9+struct TanhTilingData {
10+ int64_t totalNum = 0; // 总元素数量
11+ int64_t blockFactor = 1; // 每个核处理的元素数量
12+ int64_t ubFactor = 0; // 每次 UB 循环处理的元素数量
13+};
14+#endif
@@ -0,0 +1,21 @@
1+/*!
2+ * \file tanh_tiling_key.h
3+ * \brief Tiling 模板参数定义
4+ */
5+ 
6+#ifndef __TANH_TILING_KEY_H__
7+#define __TANH_TILING_KEY_H__
8+ 
9+#include "ascendc/host_api/tiling/template_argument.h"
10+ 
11+#define TANH_TPL_SCH_MODE_0 0
12+#define TANH_TPL_SCH_MODE_1 1
13+ 
14+ASCENDC_TPL_ARGS_DECL(
15+ Tanh,
16+ ASCENDC_TPL_UINT_DECL(schMode, 1, ASCENDC_TPL_UI_LIST, TANH_TPL_SCH_MODE_0, TANH_TPL_SCH_MODE_1));
17+ 
18+ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(
19+ ASCENDC_TPL_UINT_SEL(schMode, ASCENDC_TPL_UI_LIST, TANH_TPL_SCH_MODE_0, TANH_TPL_SCH_MODE_1)));
20+ 
21+#endif
@@ -0,0 +1,68 @@
1+cmake_minimum_required(VERSION 3.16.0)
2+project(tanh_ut CXX C)
3+ 
4+set(CMAKE_CXX_STANDARD 17)
5+set(CMAKE_CXX_STANDARD_REQUIRED ON)
6+set(CMAKE_POSITION_INDEPENDENT_CODE ON)
7+ 
8+if(NOT DEFINED ENV{ASCEND_HOME_PATH})
9+ message(FATAL_ERROR "ASCEND_HOME_PATH environment variable is not set!")
10+endif()
11+ 
12+message(STATUS "ASCEND_HOME_PATH: $ENV{ASCEND_HOME_PATH}")
13+ 
14+include(${CMAKE_CURRENT_SOURCE_DIR}/cmake/BuildGoogleTest.cmake)
15+ 
16+# DTYPE 宏定义(根据算子原型和测试用例生成)
17+add_definitions(
18+ -DDTYPE_X=float
19+ -DDTYPE_Y=float
20+)
21+ 
22+set(UT_COMMON_DIR ${CMAKE_CURRENT_SOURCE_DIR}/common)
23+set(UT_COMMON_SRCS
24+ ${UT_COMMON_DIR}/infershape_case_executor.cpp
25+ ${UT_COMMON_DIR}/infershape_context_faker.cpp
26+ ${UT_COMMON_DIR}/tiling_case_executor.cpp
27+ ${UT_COMMON_DIR}/tiling_context_faker.cpp
28+)
29+ 
30+set(UT_COMMON_INCLUDE_DIRS
31+ ${UT_COMMON_DIR}
32+ $ENV{ASCEND_HOME_PATH}/include
33+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/include
34+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/include/base
35+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/include/base/context_builder
36+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/asc/include
37+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/asc/include/tiling
38+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc
39+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/op_common
40+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/base
41+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/exe_graph
42+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/pkg_inc/graph
43+)
44+ 
45+set(UT_COMMON_LIB_DIRS
46+ $ENV{ASCEND_HOME_PATH}/lib64
47+ $ENV{ASCEND_HOME_PATH}/aarch64-linux/lib64
48+)
49+ 
50+add_compile_options(
51+ -Wall
52+ -Wno-deprecated-declarations
53+ -Wno-unused-variable
54+ -fno-access-control
55+)
56+ 
57+if(NOT CMAKE_BUILD_TYPE)
58+ set(CMAKE_BUILD_TYPE Debug)
59+endif()
60+ 
61+message(STATUS "Build Type: ${CMAKE_BUILD_TYPE}")
62+ 
63+add_subdirectory(op_host)
64+ 
65+if(EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/../../op_api)
66+ add_subdirectory(op_api)
67+endif()
68+add_subdirectory(op_kernel)
@@ -0,0 +1,132 @@
1+#
2+# BuildGoogleTest.cmake
3+#
4+# Purpose: Build Google Test from source with OLD ABI to match CANN libraries
5+#
6+# Reference: ops-math/cmake/third_party/gtest.cmake
7+# ----------------------------------------------------------------------------------------------------------
8+ 
9+include_guard(GLOBAL)
10+ 
11+# CRITICAL: Force build from source with OLD ABI to match CANN libraries
12+# We cannot use system Google Test because it uses new ABI by default
13+message(STATUS "")
14+message(STATUS "=== Building Google Test from source with OLD ABI ===")
15+message(STATUS " Reason: System gtest uses new ABI, incompatible with libplatform.so")
16+message(STATUS " Target ABI: _GLIBCXX_USE_CXX11_ABI=0 (old ABI)")
17+message(STATUS "")
18+ 
19+set(GTEST_VERSION "1.14.0")
20+set(GTEST_INSTALL_DIR ${CMAKE_BINARY_DIR}/3rd_party/gtest)
21+set(GTEST_SOURCE_DIR ${CMAKE_BINARY_DIR}/3rd_party/gtest-src)
22+ 
23+# Download URL (using gitcode mirror for China access)
24+set(GTEST_URL "https://gitcode.com/cann-src-third-party/googletest/releases/download/v${GTEST_VERSION}/googletest-${GTEST_VERSION}.tar.gz")
25+ 
26+# Compiler flags - CRITICAL: Use OLD ABI
27+set(GTEST_CXX_FLAGS "-D_GLIBCXX_USE_CXX11_ABI=0 -O2 -D_FORTIFY_SOURCE=2 -fPIC -fstack-protector-all -w")
28+set(GTEST_C_FLAGS "-D_GLIBCXX_USE_CXX11_ABI=0 -O2 -D_FORTIFY_SOURCE=2 -fPIC -fstack-protector-all -w")
29+ 
30+include(ExternalProject)
31+ExternalProject_Add(
32+ third_party_gtest
33+ URL ${GTEST_URL}
34+ TLS_VERIFY OFF
35+ DOWNLOAD_DIR ${CMAKE_BINARY_DIR}/downloads
36+ SOURCE_DIR ${GTEST_SOURCE_DIR}
37+ INSTALL_DIR ${GTEST_INSTALL_DIR}
38+
39+ CMAKE_ARGS
40+ -DCMAKE_CXX_COMPILER=${CMAKE_CXX_COMPILER}
41+ -DCMAKE_C_COMPILER=${CMAKE_C_COMPILER}
42+ -DCMAKE_CXX_FLAGS=${GTEST_CXX_FLAGS}
43+ -DCMAKE_C_FLAGS=${GTEST_C_FLAGS}
44+ -DCMAKE_INSTALL_PREFIX=<INSTALL_DIR>
45+ -DCMAKE_INSTALL_LIBDIR=lib
46+ -DBUILD_SHARED_LIBS=OFF
47+ -Dgtest_build_tests=OFF
48+ -Dgtest_build_samples=OFF
49+ -Dgmock_build_tests=OFF
50+
51+ BUILD_COMMAND $(MAKE)
52+ INSTALL_COMMAND $(MAKE) install
53+
54+ LOG_DOWNLOAD ON
55+ LOG_CONFIGURE ON
56+ LOG_BUILD ON
57+ LOG_INSTALL ON
58+)
59+ 
60+# Create imported targets (matching ops-math)
61+set(GTEST_INCLUDE_DIR ${GTEST_INSTALL_DIR}/include)
62+ 
63+# Ensure include directory exists for target_link_libraries
64+file(MAKE_DIRECTORY ${GTEST_INCLUDE_DIR})
65+ 
66+# gtest
67+add_library(gtest STATIC IMPORTED GLOBAL)
68+set_target_properties(gtest PROPERTIES
69+ IMPORTED_LOCATION ${GTEST_INSTALL_DIR}/lib/libgtest.a
70+ INTERFACE_INCLUDE_DIRECTORIES ${GTEST_INCLUDE_DIR}
71+)
72+add_dependencies(gtest third_party_gtest)
73+ 
74+# gtest_main
75+add_library(gtest_main STATIC IMPORTED GLOBAL)
76+set_target_properties(gtest_main PROPERTIES
77+ IMPORTED_LOCATION ${GTEST_INSTALL_DIR}/lib/libgtest_main.a
78+ INTERFACE_INCLUDE_DIRECTORIES ${GTEST_INCLUDE_DIR}
79+)
80+add_dependencies(gtest_main third_party_gtest)
81+ 
82+# gmock
83+add_library(gmock STATIC IMPORTED GLOBAL)
84+set_target_properties(gmock PROPERTIES
85+ IMPORTED_LOCATION ${GTEST_INSTALL_DIR}/lib/libgmock.a
86+ INTERFACE_INCLUDE_DIRECTORIES ${GTEST_INCLUDE_DIR}
87+)
88+add_dependencies(gmock third_party_gtest)
89+ 
90+# gmock_main
91+add_library(gmock_main STATIC IMPORTED GLOBAL)
92+set_target_properties(gmock_main PROPERTIES
93+ IMPORTED_LOCATION ${GTEST_INSTALL_DIR}/lib/libgmock_main.a
94+ INTERFACE_INCLUDE_DIRECTORIES ${GTEST_INCLUDE_DIR}
95+)
96+add_dependencies(gmock_main third_party_gtest)
97+ 
98+# Create interface library for UT (matching ops-math intf_llt_pub_asan_cxx17)
99+add_library(intf_llt_pub_asan_cxx17 INTERFACE)
100+target_include_directories(intf_llt_pub_asan_cxx17 INTERFACE
101+ ${GTEST_INCLUDE_DIR}
102+)
103+target_compile_definitions(intf_llt_pub_asan_cxx17 INTERFACE
104+ _GLIBCXX_USE_CXX11_ABI=0
105+ CFG_BUILD_DEBUG
106+)
107+target_compile_options(intf_llt_pub_asan_cxx17 INTERFACE
108+ -g
109+ --coverage
110+ -fprofile-arcs
111+ -ftest-coverage
112+ -w
113+ -std=c++17
114+ -fPIC
115+)
116+target_link_options(intf_llt_pub_asan_cxx17 INTERFACE
117+ -fprofile-arcs
118+ -ftest-coverage
119+)
120+target_link_libraries(intf_llt_pub_asan_cxx17 INTERFACE
121+ gcov
122+ pthread
123+)
124+ 
125+message(STATUS "")
126+message(STATUS "=== Google Test Build Configuration (ops-math mode) ===")
127+message(STATUS " Version: ${GTEST_VERSION}")
128+message(STATUS " Install Dir: ${GTEST_INSTALL_DIR}")
129+message(STATUS " ABI: OLD (_GLIBCXX_USE_CXX11_ABI=0)")
130+message(STATUS " Reason: Match CANN libraries (libplatform.so, libtiling_api.a)")
131+message(STATUS "=====================================================")
132+message(STATUS "")
@@ -0,0 +1,106 @@
1+#ifndef OPS_MATH_DEV_TESTS_UT_COMMON_ANY_VALUE_H
2+#define OPS_MATH_DEV_TESTS_UT_COMMON_ANY_VALUE_H
3+ 
4+#include <memory>
5+#include <cstdint>
6+#include <string>
7+#include <vector>
8+ 
9+namespace Ops {
10+namespace Math {
11+class AnyValue {
12+public:
13+ enum ValueType
14+ {
15+ VT_STRING = 1,
16+ VT_FLOAT = 2,
17+ VT_BOOL = 3,
18+ VT_INT = 4,
19+ VT_LIST_LIST_INT = 10,
20+ VT_LIST_BASE = 1000,
21+ 
22+ VT_LIST_FLOAT = static_cast<int32_t>(VT_LIST_BASE) + static_cast<int32_t>(VT_FLOAT),
23+ VT_LIST_BOOL = static_cast<int32_t>(VT_LIST_BASE) + static_cast<int32_t>(VT_BOOL),
24+ VT_LIST_INT = static_cast<int32_t>(VT_LIST_BASE) + static_cast<int32_t>(VT_INT),
25+ };
26+ 
27+ AnyValue(ValueType type, const std::shared_ptr<void>& valuePtr) : type_(type), valuePtr_(valuePtr)
28+ {}
29+ ~AnyValue() = default;
30+ AnyValue(const AnyValue& anyValue) : type_(anyValue.type_), valuePtr_(anyValue.valuePtr_)
31+ {}
32+ 
33+ template<typename T>
34+ static inline AnyValue CreateFrom(const T& value);
35+ 
36+ ValueType type_;
37+ std::shared_ptr<void> valuePtr_;
38+};
39+ 
40+template <>
41+inline AnyValue AnyValue::CreateFrom<std::string>(const std::string& value)
42+{
43+ auto valuePtr = new std::string;
44+ *valuePtr = value;
45+ return AnyValue(VT_STRING, std::shared_ptr<void>(valuePtr));
46+}
47+ 
48+template <>
49+inline AnyValue AnyValue::CreateFrom<float>(const float& value)
50+{
51+ auto valuePtr = new float;
52+ *valuePtr = value;
53+ return AnyValue(VT_FLOAT, std::shared_ptr<void>(valuePtr));
54+}
55+ 
56+template <>
57+inline AnyValue AnyValue::CreateFrom<bool>(const bool& value)
58+{
59+ auto valuePtr = new bool;
60+ *valuePtr = value;
61+ return AnyValue(VT_BOOL, std::shared_ptr<void>(valuePtr));
62+}
63+ 
64+template <>
65+inline AnyValue AnyValue::CreateFrom<int64_t>(const int64_t& value)
66+{
67+ auto valuePtr = new int64_t;
68+ *valuePtr = value;
69+ return AnyValue(VT_INT, std::shared_ptr<void>(valuePtr));
70+}
71+ 
72+template <>
73+inline AnyValue AnyValue::CreateFrom<std::vector<float>>(const std::vector<float>& value)
74+{
75+ auto valuePtr = new std::vector<float>;
76+ *valuePtr = value;
77+ return AnyValue(VT_LIST_FLOAT, std::shared_ptr<void>(valuePtr));
78+}
79+ 
80+template <>
81+inline AnyValue AnyValue::CreateFrom<std::vector<bool>>(const std::vector<bool>& value)
82+{
83+ auto valuePtr = new std::vector<bool>;
84+ *valuePtr = value;
85+ return AnyValue(VT_LIST_BOOL, std::shared_ptr<void>(valuePtr));
86+}
87+ 
88+template <>
89+inline AnyValue AnyValue::CreateFrom<std::vector<int64_t>>(const std::vector<int64_t>& value)
90+{
91+ auto valuePtr = new std::vector<int64_t>;
92+ *valuePtr = value;
93+ return AnyValue(VT_LIST_INT, std::shared_ptr<void>(valuePtr));
94+}
95+ 
96+template <>
97+inline AnyValue AnyValue::CreateFrom<std::vector<std::vector<int64_t>>>(const std::vector<std::vector<int64_t>>& value)
98+{
99+ auto valuePtr = new std::vector<std::vector<int64_t>>;
100+ *valuePtr = value;
101+ return AnyValue(VT_LIST_LIST_INT, std::shared_ptr<void>(valuePtr));
102+}
103+} // namespace Math
104+} // namespace Ops
105+ 
106+#endif // OPS_MATH_DEV_TESTS_UT_COMMON_ANY_VALUE_H
@@ -0,0 +1,101 @@
1+#include "infershape_case_executor.h"
2+#include <gtest/gtest.h>
3+#include "base/registry/op_impl_space_registry_v2.h"
4+ 
5+#define DO_INFERSHAPE(infershapeContextPara) \
6+ auto contextFaker = gert::InferShapeContextFaker(); \
7+ /* 1. input/output information */ \
8+ size_t inputNum = infershapeContextPara.inputTensorDesc_.size(); \
9+ size_t outputNum = infershapeContextPara.outputTensorDesc_.size(); \
10+ if (infershapeContextPara.inputInstanceNum_.size() != 0 || infershapeContextPara.outputInstanceNum_.size() != 0) { \
11+ contextFaker.IrInstanceNum(infershapeContextPara.inputInstanceNum_, infershapeContextPara.outputInstanceNum_); \
12+ } else { \
13+ contextFaker.NodeIoNum(inputNum, outputNum); \
14+ } \
15+ std::vector<gert::Tensor *> inputTensors = {}; \
16+ std::vector<std::unique_ptr<gert::Tensor>> inputTensorsKeepAlive = {}; \
17+ for (size_t index = 0; index < inputNum; index++) { \
18+ std::unique_ptr<gert::Tensor> curTensor = std::make_unique<gert::Tensor>( \
19+ infershapeContextPara.inputTensorDesc_[index].shape_, \
20+ gert::StorageFormat(infershapeContextPara.inputTensorDesc_[index].format_, \
21+ infershapeContextPara.inputTensorDesc_[index].format_, \
22+ gert::ExpandDimsType()), \
23+ gert::TensorPlacement::kOnHost, \
24+ infershapeContextPara.inputTensorDesc_[index].dtype_, \
25+ infershapeContextPara.inputTensorDesc_[index].isConst_ ? \
26+ infershapeContextPara.inputTensorDesc_[index].constValue_: \
27+ nullptr); \
28+ inputTensors.push_back(curTensor.get()); \
29+ inputTensorsKeepAlive.push_back(std::move(curTensor)); \
30+ } \
31+ for (size_t index = 0; index < outputNum; index++) { \
32+ contextFaker.NodeOutputTd(index, \
33+ infershapeContextPara.outputTensorDesc_[index].dtype_, \
34+ infershapeContextPara.outputTensorDesc_[index].format_, \
35+ infershapeContextPara.outputTensorDesc_[index].format_); \
36+ } \
37+ contextFaker.InputTensors(inputTensors); \
38+ for (auto& attrInfo : infershapeContextPara.attrs_) { \
39+ switch (attrInfo.attr_.type_) { \
40+ case Ops::Math::AnyValue::ValueType::VT_BOOL: { \
41+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<bool*>(attrInfo.attr_.valuePtr_.get())); \
42+ break;} \
43+ case Ops::Math::AnyValue::ValueType::VT_INT: { \
44+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<int64_t*>(attrInfo.attr_.valuePtr_.get())); \
45+ break;} \
46+ case Ops::Math::AnyValue::ValueType::VT_FLOAT: { \
47+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<float*>(attrInfo.attr_.valuePtr_.get())); \
48+ break;} \
49+ case Ops::Math::AnyValue::ValueType::VT_STRING: { \
50+ contextFaker.Attr(attrInfo.attrName_, ge::AscendString(reinterpret_cast<std::string*>(attrInfo.attr_.valuePtr_.get())->c_str()));\
51+ break;} \
52+ case Ops::Math::AnyValue::ValueType::VT_LIST_BOOL: { \
53+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<std::vector<bool>*>(attrInfo.attr_.valuePtr_.get()));\
54+ break;} \
55+ case Ops::Math::AnyValue::ValueType::VT_LIST_INT: { \
56+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<std::vector<int64_t>*>(attrInfo.attr_.valuePtr_.get()));\
57+ break;} \
58+ case Ops::Math::AnyValue::ValueType::VT_LIST_LIST_INT: { \
59+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<std::vector<std::vector<int64_t>>*>(attrInfo.attr_.valuePtr_.get()));\
60+ break;} \
61+ case Ops::Math::AnyValue::ValueType::VT_LIST_FLOAT: { \
62+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<std::vector<float>*>(attrInfo.attr_.valuePtr_.get()));\
63+ break;} \
64+ default: \
65+ std::cout << "[ERROR]" << __FILE__ << ":" << __LINE__ << "The ValueType " << attrInfo.attr_.type_ << "is not supported!" << std::endl;\
66+ } \
67+ } \
68+ auto contextHolder = contextFaker.SetOpType(infershapeContextPara.opName_.c_str()).Build(); \
69+ /* 2. get infershape func */ \
70+ auto spaceRegistry = gert::DefaultOpImplSpaceRegistryV2::GetInstance().GetSpaceRegistry(); \
71+ auto infershapeFunc = spaceRegistry->GetOpImpl(infershapeContextPara.opName_.c_str())->infer_shape; \
72+ /* 3. check infershape func */ \
73+ auto infershapeRet = infershapeFunc(contextHolder.GetContext());
74+ 
75+static std::vector<int64_t> ToVector(const gert::Shape& shape) {
76+ size_t shapeSize = shape.GetDimNum();
77+ std::vector<int64_t> shapeVec(shapeSize, 0);
78+ 
79+ for (size_t i = 0; i < shapeSize; i++) {
80+ shapeVec[i] = shape.GetDim(i);
81+ }
82+ return shapeVec;
83+}
84+ 
85+void ExecuteTestCase(gert::InfershapeContextPara& infershapeContextPara,
86+ ge::graphStatus expectResult,
87+ const std::vector<std::vector<int64_t>>& expectOutputShape)
88+{
89+ DO_INFERSHAPE(infershapeContextPara);
90+ 
91+ // check infershape func
92+ EXPECT_EQ(infershapeRet, expectResult);
93+ if (expectResult == ge::GRAPH_FAILED) {
94+ return;
95+ }
96+ 
97+ // check output shape
98+ for (size_t i = 0; i < expectOutputShape.size(); i++) {
99+ EXPECT_EQ(ToVector(*contextHolder.GetContext()->GetOutputShape(i)), expectOutputShape[i]);
100+ }
101+}
@@ -0,0 +1,10 @@
1+#ifndef OPS_MATH_DEV_TESTS_UT_COMMON_INFERSHAPE_CASE_EXECUTOR_H
2+#define OPS_MATH_DEV_TESTS_UT_COMMON_INFERSHAPE_CASE_EXECUTOR_H
3+ 
4+#include "infershape_context_faker.h"
5+ 
6+void ExecuteTestCase(gert::InfershapeContextPara& infershapeContextPara,
7+ ge::graphStatus expectResult = ge::GRAPH_FAILED,
8+ const std::vector<std::vector<int64_t>>& expectOutputShape = {});
9+ 
10+#endif // OPS_MATH_DEV_TESTS_UT_COMMON_INFERSHAPE_CASE_EXECUTOR_H
@@ -0,0 +1,42 @@
1+#include "infershape_context_faker.h"
2+ 
3+namespace gert {
4+ 
5+InferShapeContextFaker& InferShapeContextFaker::SetOpType(const std::string opType)
6+{
7+ OpInferShapeContextBuilder::OpType(opType.c_str()).OpName(opType.c_str());
8+ return *this;
9+}
10+ 
11+InferShapeContextFaker& InferShapeContextFaker::NodeIoNum(size_t inputNum, size_t outputNum)
12+{
13+ OpInferShapeContextBuilder::IONum(inputNum, outputNum);
14+ return *this;
15+}
16+ 
17+InferShapeContextFaker& InferShapeContextFaker::IrInstanceNum(const std::vector<uint32_t>& inputInstanceNum,
18+ const std::vector<uint32_t>& outputInstanceNum)
19+{
20+ OpInferShapeContextBuilder::IOInstanceNum(inputInstanceNum, outputInstanceNum);
21+ return *this;
22+}
23+ 
24+InferShapeContextFaker& InferShapeContextFaker::NodeOutputTd(int32_t index, ge::DataType dtype, ge::Format originFormat,
25+ ge::Format storageFormat)
26+{
27+ OpInferShapeContextBuilder::OutputTensorDesc(index, dtype, originFormat, storageFormat);
28+ return *this;
29+}
30+ 
31+InferShapeContextFaker& InferShapeContextFaker::InputTensors(const std::vector<Tensor *>& inputTensors)
32+{
33+ OpInferShapeContextBuilder::InputTensors(inputTensors);
34+ return *this;
35+}
36+ 
37+ContextHolder<InferShapeContext> InferShapeContextFaker::Build()
38+{
39+ return OpInferShapeContextBuilder::Build();
40+}
41+ 
42+} // namespace gert
@@ -0,0 +1,124 @@
1+#ifndef OPS_MATH_DEV_TESTS_UT_COMMON_INFERSHAPE_CONTEXT_FAKER_H
2+#define OPS_MATH_DEV_TESTS_UT_COMMON_INFERSHAPE_CONTEXT_FAKER_H
3+ 
4+#include "op_infer_shape_context_builder.h"
5+#include "any_value.h"
6+ 
7+namespace gert {
8+ 
9+class InfershapeContextPara {
10+public:
11+ class TensorDescription {
12+ public:
13+ TensorDescription(const gert::StorageShape& shape, ge::DataType dtype, ge::Format format, bool isConst = false,
14+ void* constValue = nullptr) :
15+ shape_(shape), dtype_(dtype), format_(format), isConst_(isConst), constValue_(constValue) {}
16+ public:
17+ gert::StorageShape shape_;
18+ ge::DataType dtype_ = ge::DT_FLOAT;
19+ ge::Format format_ = ge::FORMAT_ND;
20+ bool isConst_ = false;
21+ void* constValue_ = nullptr;
22+ };
23+ 
24+ class OpAttr {
25+ public:
26+ OpAttr(const std::string& attrName, const Ops::Math::AnyValue& attr) : attrName_(attrName), attr_(attr) {}
27+ public:
28+ std::string attrName_;
29+ Ops::Math::AnyValue attr_;
30+ };
31+public:
32+ InfershapeContextPara(const std::string& opName,
33+ const std::vector<TensorDescription>& inputTensorDesc,
34+ const std::vector<TensorDescription>& outputTensorDesc,
35+ const std::vector<OpAttr>& attrs,
36+ const std::vector<uint32_t>& inputInstanceNum = {},
37+ const std::vector<uint32_t>& outputInstanceNum = {}) :
38+ opName_(opName),
39+ inputInstanceNum_(inputInstanceNum),
40+ outputInstanceNum_(outputInstanceNum),
41+ inputTensorDesc_(inputTensorDesc),
42+ outputTensorDesc_(outputTensorDesc),
43+ attrs_(attrs) {}
44+ 
45+ InfershapeContextPara(const std::string& opName,
46+ const std::vector<TensorDescription>& inputTensorDesc,
47+ const std::vector<TensorDescription>& outputTensorDesc,
48+ const std::vector<uint32_t>& inputInstanceNum = {},
49+ const std::vector<uint32_t>& outputInstanceNum = {}) :
50+ opName_(opName),
51+ inputInstanceNum_(inputInstanceNum),
52+ outputInstanceNum_(outputInstanceNum),
53+ inputTensorDesc_(inputTensorDesc),
54+ outputTensorDesc_(outputTensorDesc),
55+ attrs_() {}
56+ 
57+public:
58+ std::string opName_;
59+ std::vector<uint32_t> inputInstanceNum_;
60+ std::vector<uint32_t> outputInstanceNum_;
61+ std::vector<TensorDescription> inputTensorDesc_;
62+ std::vector<TensorDescription> outputTensorDesc_;
63+ std::vector<OpAttr> attrs_;
64+};
65+ 
66+class InferShapeContextFaker : public OpInferShapeContextBuilder {
67+public:
68+ InferShapeContextFaker& SetOpType(const std::string opType);
69+ 
70+ /* only one can be choosed from IrInstanceNum */
71+ InferShapeContextFaker& NodeIoNum(size_t inputNum, size_t outputNum);
72+ 
73+ /* can be used for dynamic inputs/outputs
74+ * only one can be choosed from NodeIoNum */
75+ InferShapeContextFaker& IrInstanceNum(const std::vector<uint32_t>& inputInstanceNum,
76+ const std::vector<uint32_t>& outputInstanceNum);
77+ 
78+ InferShapeContextFaker& NodeOutputTd(int32_t index, ge::DataType dtype, ge::Format originFormat,
79+ ge::Format storageFormat);
80+ 
81+ InferShapeContextFaker& Attr(const std::string& attrName, bool attr) {
82+ OpInferShapeContextBuilder::AppendAttr(attr);
83+ return *this;
84+ }
85+ InferShapeContextFaker& Attr(const std::string& attrName, int64_t attr) {
86+ OpInferShapeContextBuilder::AppendAttr(attr);
87+ return *this;
88+ }
89+ InferShapeContextFaker& Attr(const std::string& attrName, float attr) {
90+ OpInferShapeContextBuilder::AppendAttr(attr);
91+ return *this;
92+ }
93+ InferShapeContextFaker& Attr(const std::string& attrName, const ge::AscendString& attr) {
94+ OpInferShapeContextBuilder::AppendAttr(attr);
95+ return *this;
96+ }
97+ InferShapeContextFaker& Attr(const std::string& attrName, const std::vector<bool>& attr) {
98+ OpInferShapeContextBuilder::AppendAttr(attr);
99+ return *this;
100+ }
101+ InferShapeContextFaker& Attr(const std::string& attrName, const std::vector<int64_t>& attr) {
102+ OpInferShapeContextBuilder::AppendAttr(attr);
103+ return *this;
104+ }
105+ InferShapeContextFaker& Attr(const std::string& attrName, const std::vector<float>& attr) {
106+ OpInferShapeContextBuilder::AppendAttr(attr);
107+ return *this;
108+ }
109+ InferShapeContextFaker& Attr(const std::string& attrName, const std::vector<ge::AscendString>& attr) {
110+ OpInferShapeContextBuilder::AppendAttr(attr);
111+ return *this;
112+ }
113+ InferShapeContextFaker& Attr(const std::string& attrName, const std::vector<std::vector<int64_t>>& attr) {
114+ OpInferShapeContextBuilder::AppendAttr(attr);
115+ return *this;
116+ }
117+ 
118+ InferShapeContextFaker& InputTensors(const std::vector<Tensor *>& inputTensors);
119+ 
120+ ContextHolder<InferShapeContext> Build();
121+};
122+ 
123+} // namespace gert
124+#endif // OPS_MATH_DEV_TESTS_UT_COMMON_INFERSHAPE_CONTEXT_FAKER_H
@@ -0,0 +1,295 @@
1+#include "tiling_case_executor.h"
2+#include <gtest/gtest.h>
3+#include <nlohmann/json.hpp>
4+#include "platform/platform_infos_def.h"
5+#include "base/registry/op_impl_space_registry_v2.h"
6+ 
7+#define STR_IMPL(x) #x
8+#define STR(x) STR_IMPL(x)
9+#define DO_TILING(tilingContextPara) \
10+ auto contextFaker = gert::TilingContextFaker(); \
11+ /* 1. input/output information */ \
12+ size_t inputNum = tilingContextPara.inputTensorDesc_.size(); \
13+ size_t outputNum = tilingContextPara.outputTensorDesc_.size(); \
14+ if (tilingContextPara.inputInstanceNum_.size() != 0 || tilingContextPara.outputInstanceNum_.size() != 0) { \
15+ contextFaker.IrInstanceNum(tilingContextPara.inputInstanceNum_, tilingContextPara.outputInstanceNum_); \
16+ } else { \
17+ contextFaker.NodeIoNum(inputNum, outputNum); \
18+ } \
19+ std::vector<gert::Tensor *> inputTensors = {}; \
20+ std::vector<gert::Tensor *> outputTensors = {}; \
21+ std::vector<std::unique_ptr<gert::Tensor>> inputTensorsKeepAlive = {}; \
22+ std::vector<std::unique_ptr<gert::Tensor>> outputTensorsKeepAlive = {}; \
23+ for (size_t index = 0; index < inputNum; index++) { \
24+ std::unique_ptr<gert::Tensor> curTensor = std::make_unique<gert::Tensor>( \
25+ tilingContextPara.inputTensorDesc_[index].shape_, \
26+ gert::StorageFormat(tilingContextPara.inputTensorDesc_[index].format_, \
27+ tilingContextPara.inputTensorDesc_[index].format_, \
28+ gert::ExpandDimsType()), \
29+ gert::TensorPlacement::kOnHost, \
30+ tilingContextPara.inputTensorDesc_[index].dtype_, \
31+ tilingContextPara.inputTensorDesc_[index].isConst_ ? \
32+ tilingContextPara.inputTensorDesc_[index].constValue_: \
33+ nullptr); \
34+ inputTensors.push_back(curTensor.get()); \
35+ inputTensorsKeepAlive.push_back(std::move(curTensor)); \
36+ } \
37+ for (size_t index = 0; index < outputNum; index++) { \
38+ std::unique_ptr<gert::Tensor> curTensor = std::make_unique<gert::Tensor>( \
39+ tilingContextPara.outputTensorDesc_[index].shape_, \
40+ gert::StorageFormat(tilingContextPara.outputTensorDesc_[index].format_, \
41+ tilingContextPara.outputTensorDesc_[index].format_, \
42+ gert::ExpandDimsType()), \
43+ gert::TensorPlacement::kOnHost, \
44+ tilingContextPara.outputTensorDesc_[index].dtype_, \
45+ tilingContextPara.outputTensorDesc_[index].isConst_ ? \
46+ tilingContextPara.outputTensorDesc_[index].constValue_: \
47+ nullptr); \
48+ outputTensors.push_back(curTensor.get()); \
49+ outputTensorsKeepAlive.push_back(std::move(curTensor)); \
50+ } \
51+ contextFaker.InputTensors(inputTensors).OutputTensors(outputTensors); \
52+ for (auto& attrInfo : tilingContextPara.attrs_) { \
53+ switch (attrInfo.attr_.type_) { \
54+ case Ops::Math::AnyValue::ValueType::VT_BOOL: { \
55+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<bool*>(attrInfo.attr_.valuePtr_.get())); \
56+ break;} \
57+ case Ops::Math::AnyValue::ValueType::VT_INT: { \
58+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<int64_t*>(attrInfo.attr_.valuePtr_.get())); \
59+ break;} \
60+ case Ops::Math::AnyValue::ValueType::VT_FLOAT: { \
61+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<float*>(attrInfo.attr_.valuePtr_.get())); \
62+ break;} \
63+ case Ops::Math::AnyValue::ValueType::VT_STRING: { \
64+ contextFaker.Attr(attrInfo.attrName_, ge::AscendString(reinterpret_cast<std::string*>(attrInfo.attr_.valuePtr_.get())->c_str()));\
65+ break;} \
66+ case Ops::Math::AnyValue::ValueType::VT_LIST_BOOL: { \
67+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<std::vector<bool>*>(attrInfo.attr_.valuePtr_.get()));\
68+ break;} \
69+ case Ops::Math::AnyValue::ValueType::VT_LIST_INT: { \
70+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<std::vector<int64_t>*>(attrInfo.attr_.valuePtr_.get()));\
71+ break;} \
72+ case Ops::Math::AnyValue::ValueType::VT_LIST_LIST_INT: { \
73+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<std::vector<std::vector<int64_t>>*>(attrInfo.attr_.valuePtr_.get()));\
74+ break;} \
75+ case Ops::Math::AnyValue::ValueType::VT_LIST_FLOAT: { \
76+ contextFaker.Attr(attrInfo.attrName_, *reinterpret_cast<std::vector<float>*>(attrInfo.attr_.valuePtr_.get()));\
77+ break;} \
78+ default: \
79+ std::cout << "[ERROR]" << __FILE__ << ":" << __LINE__ << "The ValueType " << attrInfo.attr_.type_ << "is not supported!" << std::endl;\
80+ } \
81+ } \
82+ /* 2. base information */ \
83+ fe::PlatFormInfos platformInfo; \
84+ platformInfo.Init(); \
85+ auto tilingData = gert::TilingData::CreateCap(tilingContextPara.tilingDataSize_); \
86+ auto workspace = gert::ContinuousVector::Create<size_t>(4096); \
87+ auto contextHolder = contextFaker.SetOpType(tilingContextPara.opName_.c_str()) \
88+ .CompileInfo(tilingContextPara.compileInfo_) \
89+ .PlatformInfo(reinterpret_cast<char*>(&platformInfo)) \
90+ .TilingData(tilingData.get()) \
91+ .Workspace(reinterpret_cast<gert::ContinuousVector *>(workspace.get())) \
92+ .Build(); \
93+ string compileInfoStringPrefix = R"({"hardware_info": {"BT_SIZE": 0, "load3d_constraints": "1", "Intrinsic_fix_pipe_l0c2out": false, "Intrinsic_data_move_l12ub": true, "Intrinsic_data_move_l0c2ub": true, "Intrinsic_data_move_out2l1_nd2nz": false, "UB_SIZE": )";\
94+ string compileInfoStringMiddle = R"(, "L2_SIZE": 33554432, "L1_SIZE": 524288, "L0A_SIZE": 65536, "L0B_SIZE": 65536, "L0C_SIZE": 131072, "CORE_NUM": )";\
95+ map<string, string> socToUpper = { \
96+ {"ascend910b", "Ascend910B"}, \
97+ {"ascend910_93", "Ascend910_93"}, \
98+ {"ascend950", "Ascend950"}, \
99+ {"ascend310p", "Ascend310P"}, \
100+ {"ascend910", "Ascend910"}, \
101+ {"ascend310b", "Ascend310B"}, \
102+ {"ascend610lite", "Ascend610Lite"}, \
103+ {"ascend031", "Ascend031"}, \
104+ {"ascend035", "Ascend035"}, \
105+ {"kirinx90", "KrinX90"}, \
106+ {"kirin9030", "Kirin9030"}, \
107+ {"mc62cm12a", "MC62CM12A"} \
108+ }; \
109+ std::string buildSocVersion = STR(BUILD_SOC_VERSION); \
110+ if (!buildSocVersion.empty()) \
111+ { \
112+ buildSocVersion = socToUpper[buildSocVersion]; \
113+ } \
114+ string compileInfoStringSuffix = R"(, "socVersion":)" R"(")" + buildSocVersion + R"("} })"; \
115+ string compileInfoString = compileInfoStringPrefix + \
116+ std::to_string(tilingContextPara.ubSize_) + \
117+ compileInfoStringMiddle + \
118+ std::to_string(tilingContextPara.coreNum_) + \
119+ compileInfoStringSuffix; \
120+ map<string, string> socToArch = { \
121+ {"Ascend310P", "2002"}, \
122+ {"Ascend910B", "2201"}, \
123+ {"Ascend910_93", "2201"}, \
124+ {"Ascend950", "3510"}, \
125+ {"Ascend910", "1001"} \
126+ }; \
127+ map<string, string> socInfos; \
128+ map<string, string> aicoreSpec; \
129+ map<string, string> intrinsics; \
130+ map<string, string> socversions = { \
131+ {"NpuArch", socToArch[buildSocVersion]}, {"Short_SoC_version", buildSocVersion}}; \
132+ GetPlatFormInfos(compileInfoString.c_str(), socInfos, aicoreSpec, intrinsics); \
133+ auto tilingContext = contextHolder.GetContext(); \
134+ tilingContext->GetPlatformInfo()->SetPlatformRes("SoCInfo", socInfos); \
135+ tilingContext->GetPlatformInfo()->SetPlatformRes("AICoreSpec", aicoreSpec); \
136+ tilingContext->GetPlatformInfo()->SetCoreNumByCoreType("AICore"); \
137+ tilingContext->GetPlatformInfo()->SetPlatformRes("AICoreintrinsicDtypeMap", intrinsics); \
138+ tilingContext->GetPlatformInfo()->SetPlatformRes("version", socversions); \
139+ /* 3. get tiling func */ \
140+ auto spaceRegistry = gert::DefaultOpImplSpaceRegistryV2::GetInstance().GetSpaceRegistry(); \
141+ if (spaceRegistry == nullptr) { \
142+ throw std::invalid_argument("not found spaceRegistry"); \
143+ } \
144+ auto functionStruct = spaceRegistry->GetOpImpl(tilingContextPara.opName_.c_str()); \
145+ if (functionStruct == nullptr) { \
146+ throw std::invalid_argument("not found "+tilingContextPara.opName_); \
147+ } \
148+ auto tilingFunc =functionStruct->tiling; /* 4. check tiling func */ \
149+ /* 4. check tiling func */ \
150+ auto tilingRet = tilingFunc(tilingContext);
151+ 
152+template <typename T>
153+static string to_string(void* buf, size_t size) {
154+ string result;
155+ const T* data = reinterpret_cast<const T*>(buf);
156+ size_t len = size / sizeof(T);
157+ for (size_t i = 0; i < len; i++) {
158+ result += std::to_string(data[i]);
159+ result += " ";
160+ }
161+ return result;
162+}
163+ 
164+static void GetPlatFormInfos(const char* compileInfoStr, map<string, string>& socInfos, map<string, string>& aicoreSpec,
165+ map<string, string>& intrinsics) {
166+ string default_hardward_info = R"({
167+ "hardware_info": {"BT_SIZE": 0, "load3d_constraints": "1", "Intrinsic_fix_pipe_l0c2out": false,
168+ "Intrinsic_data_move_l12ub": true, "Intrinsic_data_move_l0c2ub": true,
169+ "Intrinsic_data_move_out2l1_nd2nz": false, "UB_SIZE": 262144, "L2_SIZE": 33554432,
170+ "L1_SIZE": 1048576, "L0A_SIZE": 65536, "L0B_SIZE": 65536, "L0C_SIZE": 262144,
171+ "CORE_NUM": 32}})";
172+ nlohmann::json compileInfoJson = nlohmann::json::parse(compileInfoStr);
173+ if (compileInfoJson.type() != nlohmann::json::value_t::object) {
174+ compileInfoJson = nlohmann::json::parse(default_hardward_info.c_str());
175+ }
176+ 
177+ map<string, string> socInfoKeys = {{"ai_core_cnt", "CORE_NUM"},
178+ {"l2_size", "L2_SIZE"},
179+ {"cube_core_cnt", "cube_core_cnt"},
180+ {"vector_core_cnt", "vector_core_cnt"},
181+ {"core_type_list", "core_type_list"}};
182+ socInfos["core_type_list"] = "AICore";
183+ 
184+ for (auto &t : socInfoKeys) {
185+ if (compileInfoJson.contains("hardware_info") && compileInfoJson["hardware_info"].contains(t.second)) {
186+ auto &objJson = compileInfoJson["hardware_info"][t.second];
187+ if (objJson.is_number_integer()) {
188+ socInfos[t.first] = to_string(compileInfoJson["hardware_info"][t.second].get<uint32_t>());
189+ } else if (objJson.is_string()) {
190+ socInfos[t.first] = objJson;
191+ }
192+ }
193+ }
194+ map<string, string> aicoreSpecKeys = {{"ub_size", "UB_SIZE"},
195+ {"l0_a_size", "L0A_SIZE"},
196+ {"l0_b_size", "L0B_SIZE"},
197+ {"l0_c_size", "L0C_SIZE"},
198+ {"l1_size", "L1_SIZE"},
199+ {"bt_size", "BT_SIZE"},
200+ {"load3d_constraints", "load3d_constraints"}};
201+ aicoreSpec["cube_freq"] = "cube_freq";
202+ for (auto &t : aicoreSpecKeys) {
203+ if (compileInfoJson.contains("hardware_info") && compileInfoJson["hardware_info"].contains(t.second)) {
204+ if (t.second == "load3d_constraints") {
205+ aicoreSpec[t.first] = compileInfoJson["hardware_info"][t.second].get<string>();
206+ } else {
207+ aicoreSpec[t.first] = to_string(compileInfoJson["hardware_info"][t.second].get<uint32_t>());
208+ }
209+ }
210+ }
211+ 
212+ std::string intrinsicsKeys[] = {"Intrinsic_data_move_l12ub", "Intrinsic_data_move_l0c2ub",
213+ "Intrinsic_fix_pipe_l0c2out", "Intrinsic_data_move_out2l1_nd2nz",
214+ "Intrinsic_matmul_ub_to_ub", "Intrinsic_conv_ub_to_ub",
215+ "Intrinsic_data_move_l12bt"};
216+ for (string key : intrinsicsKeys) {
217+ if (compileInfoJson.contains("hardware_info") && compileInfoJson["hardware_info"].contains(key) &&
218+ compileInfoJson["hardware_info"][key].get<bool>()) {
219+ intrinsics[key] = "float16";
220+ if (key.find("Intrinsic_data_move_l12bt") != string::npos) {
221+ intrinsics[key] = "bf16";
222+ }
223+ }
224+ }
225+}
226+ 
227+void ExecuteTestCase(const gert::TilingContextPara& tilingContextPara,
228+ ge::graphStatus expectResult,
229+ uint64_t expectTilingKey,
230+ const string& expectTilingData,
231+ const std::vector<size_t>& expectWorkspaces)
232+{
233+ DO_TILING(tilingContextPara);
234+ 
235+ // check tiling func
236+ EXPECT_EQ(tilingRet, expectResult);
237+ if (expectResult == ge::GRAPH_FAILED) {
238+ return;
239+ }
240+ 
241+ // check workspace
242+ size_t workspaceCount = tilingContext->GetWorkspaceNum();
243+ if (workspaceCount > 0) {
244+ ASSERT_EQ(workspaceCount, expectWorkspaces.size());
245+ auto workspaceSizes = tilingContext->GetWorkspaceSizes(workspaceCount);
246+ for (size_t i = 0; i < workspaceCount; i++) {
247+ ASSERT_EQ(workspaceSizes[i], expectWorkspaces[i]);
248+ }
249+ }
250+ 
251+ // check tiling key
252+ auto tilingKeyResult = tilingContext->GetTilingKey();
253+ ASSERT_EQ(tilingKeyResult, expectTilingKey);
254+ 
255+ // check tiling data
256+ if (expectTilingData == EMPTY_EXPECT_TILING_DATA) {
257+ return;
258+ }
259+ auto rawTilingData = tilingContext->GetRawTilingData();
260+ auto tilingDataResult = to_string<int64_t>(rawTilingData->GetData(), rawTilingData->GetDataSize());
261+ EXPECT_EQ(tilingDataResult, expectTilingData);
262+}
263+ 
264+void ExecuteTestCase(const gert::TilingContextPara& tilingContextPara,
265+ ge::graphStatus expectResult,
266+ uint64_t expectTilingKey,
267+ const std::vector<size_t>& expectWorkspaces)
268+{
269+ ExecuteTestCase(tilingContextPara, expectResult, expectTilingKey, EMPTY_EXPECT_TILING_DATA, expectWorkspaces);
270+}
271+ 
272+bool ExecuteTiling(const gert::TilingContextPara& tilingContextPara, TilingInfo& tilingInfo)
273+{
274+ DO_TILING(tilingContextPara);
275+ 
276+ if (tilingRet != ge::GRAPH_SUCCESS) {
277+ return false;
278+ }
279+
280+ tilingInfo.tilingKey = tilingContext->GetTilingKey();
281+ tilingInfo.blockNum = tilingContext->GetBlockDim();
282+ size_t workspaceCount = tilingContext->GetWorkspaceNum();
283+ if (workspaceCount > 0) {
284+ auto workSpaceSizes = tilingContext->GetWorkspaceSizes(workspaceCount);
285+ for (size_t i = 0; i < workspaceCount; i++) {
286+ tilingInfo.workspaceSizes.push_back(workSpaceSizes[i]);
287+ }
288+ }
289+ auto rawTilingData = tilingContext->GetRawTilingData();
290+ tilingInfo.tilingData = std::make_unique<uint8_t[]>(rawTilingData->GetDataSize());
291+ tilingInfo.tilingDataSize = rawTilingData->GetDataSize();
292+ std::memcpy(tilingInfo.tilingData.get(), rawTilingData->GetData(), rawTilingData->GetDataSize());
293+ 
294+ return true;
295+}
@@ -0,0 +1,31 @@
1+#ifndef OPS_MATH_DEV_TESTS_UT_COMMON_TILING_CASE_EXECUTOR_H
2+#define OPS_MATH_DEV_TESTS_UT_COMMON_TILING_CASE_EXECUTOR_H
3+ 
4+#include "tiling_context_faker.h"
5+ 
6+using namespace std;
7+ 
8+const string EMPTY_EXPECT_TILING_DATA = "EMPTY_EXPECT_TILING_DATA";
9+ 
10+struct TilingInfo {
11+ int64_t tilingKey = -1;
12+ std::vector<int64_t> workspaceSizes;
13+ std::unique_ptr<uint8_t[]> tilingData;
14+ size_t tilingDataSize = 0;
15+ size_t blockNum = 0;
16+};
17+ 
18+void ExecuteTestCase(const gert::TilingContextPara& tilingContextPara,
19+ ge::graphStatus expectResult = ge::GRAPH_FAILED,
20+ uint64_t expectTilingKey = 0,
21+ const string& expectTilingData = "",
22+ const std::vector<size_t>& expectWorkspaces = {});
23+ 
24+void ExecuteTestCase(const gert::TilingContextPara& tilingContextPara,
25+ ge::graphStatus expectResult,
26+ uint64_t expectTilingKey,
27+ const std::vector<size_t>& expectWorkspaces);
28+ 
29+bool ExecuteTiling(const gert::TilingContextPara& tilingContextPara, TilingInfo& tilingInfo);
30+ 
31+#endif // OPS_MATH_DEV_TESTS_UT_COMMON_TILING_CASE_EXECUTOR_H
@@ -0,0 +1,71 @@
1+#include "tiling_context_faker.h"
2+ 
3+namespace gert {
4+ 
5+TilingContextFaker& TilingContextFaker::SetOpType(const std::string opType)
6+{
7+ OpTilingContextBuilder::OpType(opType.c_str()).OpName(opType.c_str());
8+ return *this;
9+}
10+ 
11+TilingContextFaker& TilingContextFaker::NodeIoNum(size_t inputNum, size_t outputNum)
12+{
13+ OpTilingContextBuilder::IONum(inputNum, outputNum);
14+ return *this;
15+}
16+ 
17+TilingContextFaker& TilingContextFaker::IrInstanceNum(const std::vector<uint32_t>& inputInstanceNum,
18+ const std::vector<uint32_t>& outputInstanceNum)
19+{
20+ OpTilingContextBuilder::IOInstanceNum(inputInstanceNum, outputInstanceNum);
21+ return *this;
22+}
23+ 
24+TilingContextFaker& TilingContextFaker::InputTensors(const std::vector<Tensor *>& inputTensors)
25+{
26+ OpTilingContextBuilder::InputTensors(inputTensors);
27+ return *this;
28+}
29+ 
30+TilingContextFaker& TilingContextFaker::OutputTensors(const std::vector<Tensor *>& outputTensors)
31+{
32+ OpTilingContextBuilder::OutputTensors(outputTensors);
33+ return *this;
34+}
35+ 
36+TilingContextFaker& TilingContextFaker::CompileInfo(const void* compileInfo)
37+{
38+ OpTilingContextBuilder::CompileInfo(compileInfo);
39+ return *this;
40+}
41+ 
42+TilingContextFaker& TilingContextFaker::PlatformInfo(const void* platformInfo)
43+{
44+ OpTilingContextBuilder::PlatformInfo(platformInfo);
45+ return *this;
46+}
47+ 
48+TilingContextFaker& TilingContextFaker::DeterministicInfo(int32_t* deterministicInfo)
49+{
50+ OpTilingContextBuilder::Deterministic(*deterministicInfo);
51+ return *this;
52+}
53+ 
54+TilingContextFaker& TilingContextFaker::TilingData(const void* tilingData)
55+{
56+ OpTilingContextBuilder::TilingData(static_cast<const gert::TilingData *>(tilingData));
57+ return *this;
58+}
59+ 
60+TilingContextFaker& TilingContextFaker::Workspace(const ContinuousVector* workspace)
61+{
62+ OpTilingContextBuilder::Workspace(workspace);
63+ return *this;
64+}
65+ 
66+ContextHolder<TilingContext> TilingContextFaker::Build()
67+{
68+ return OpTilingContextBuilder::Build();
69+}
70+ 
71+} // namespace gert
@@ -0,0 +1,194 @@
1+#ifndef OPS_MATH_DEV_TESTS_UT_COMMON_TILING_CONTEXT_FAKER_H
2+#define OPS_MATH_DEV_TESTS_UT_COMMON_TILING_CONTEXT_FAKER_H
3+ 
4+#include "op_tiling_context_builder.h"
5+#include "any_value.h"
6+ 
7+namespace gert {
8+ 
9+class TilingContextPara {
10+public:
11+ class TensorDescription {
12+ public:
13+ TensorDescription(const gert::StorageShape& shape,
14+ ge::DataType dtype,
15+ ge::Format format,
16+ bool isConst = false,
17+ void* constValue = nullptr) :
18+ shape_(shape), dtype_(dtype), format_(format), isConst_(isConst), constValue_(constValue) {}
19+ public:
20+ gert::StorageShape shape_;
21+ ge::DataType dtype_ = ge::DT_FLOAT;
22+ ge::Format format_ = ge::FORMAT_ND;
23+ bool isConst_ = false;
24+ void* constValue_ = nullptr;
25+ };
26+ 
27+ class OpAttr {
28+ public:
29+ OpAttr(const std::string& attrName, const Ops::Math::AnyValue& attr) : attrName_(attrName), attr_(attr) {}
30+ public:
31+ std::string attrName_;
32+ Ops::Math::AnyValue attr_;
33+ };
34+public:
35+ TilingContextPara(const std::string& opName,
36+ const std::vector<TensorDescription>& inputTensorDesc,
37+ const std::vector<TensorDescription>& outputTensorDesc,
38+ const std::vector<OpAttr>& attrs,
39+ void* compileInfo = nullptr,
40+ uint64_t coreNum = 64,
41+ uint64_t ubSize = 262144,
42+ uint64_t tilingDataSize = 4096) :
43+ opName_(opName),
44+ inputInstanceNum_(),
45+ outputInstanceNum_(),
46+ inputTensorDesc_(inputTensorDesc),
47+ outputTensorDesc_(outputTensorDesc),
48+ attrs_(attrs),
49+ coreNum_(coreNum),
50+ ubSize_(ubSize),
51+ tilingDataSize_(tilingDataSize),
52+ compileInfo_(compileInfo) {}
53+ 
54+ TilingContextPara(const std::string& opName,
55+ const std::vector<TensorDescription>& inputTensorDesc,
56+ const std::vector<TensorDescription>& outputTensorDesc,
57+ void* compileInfo = nullptr,
58+ uint64_t coreNum = 64,
59+ uint64_t ubSize = 262144,
60+ uint64_t tilingDataSize = 4096) :
61+ opName_(opName),
62+ inputInstanceNum_(),
63+ outputInstanceNum_(),
64+ inputTensorDesc_(inputTensorDesc),
65+ outputTensorDesc_(outputTensorDesc),
66+ attrs_(),
67+ coreNum_(coreNum),
68+ ubSize_(ubSize),
69+ tilingDataSize_(tilingDataSize),
70+ compileInfo_(compileInfo) {}
71+ 
72+ TilingContextPara(const std::string& opName,
73+ const std::vector<TensorDescription>& inputTensorDesc,
74+ const std::vector<TensorDescription>& outputTensorDesc,
75+ const std::vector<OpAttr>& attrs,
76+ const std::vector<uint32_t>& inputInstanceNum,
77+ const std::vector<uint32_t>& outputInstanceNum,
78+ void* compileInfo = nullptr,
79+ uint64_t coreNum = 64,
80+ uint64_t ubSize = 262144,
81+ uint64_t tilingDataSize = 4096) :
82+ opName_(opName),
83+ inputInstanceNum_(inputInstanceNum),
84+ outputInstanceNum_(outputInstanceNum),
85+ inputTensorDesc_(inputTensorDesc),
86+ outputTensorDesc_(outputTensorDesc),
87+ attrs_(attrs),
88+ coreNum_(coreNum),
89+ ubSize_(ubSize),
90+ tilingDataSize_(tilingDataSize),
91+ compileInfo_(compileInfo) {}
92+ 
93+ TilingContextPara(const std::string& opName,
94+ const std::vector<TensorDescription>& inputTensorDesc,
95+ const std::vector<TensorDescription>& outputTensorDesc,
96+ const std::vector<uint32_t>& inputInstanceNum,
97+ const std::vector<uint32_t>& outputInstanceNum,
98+ void* compileInfo = nullptr,
99+ uint64_t coreNum = 64,
100+ uint64_t ubSize = 262144,
101+ uint64_t tilingDataSize = 4096) :
102+ opName_(opName),
103+ inputInstanceNum_(inputInstanceNum),
104+ outputInstanceNum_(outputInstanceNum),
105+ inputTensorDesc_(inputTensorDesc),
106+ outputTensorDesc_(outputTensorDesc),
107+ attrs_(),
108+ coreNum_(coreNum),
109+ ubSize_(ubSize),
110+ tilingDataSize_(tilingDataSize),
111+ compileInfo_(compileInfo) {}
112+public:
113+ std::string opName_;
114+ std::vector<uint32_t> inputInstanceNum_;
115+ std::vector<uint32_t> outputInstanceNum_;
116+ std::vector<TensorDescription> inputTensorDesc_;
117+ std::vector<TensorDescription> outputTensorDesc_;
118+ std::vector<OpAttr> attrs_;
119+ uint64_t coreNum_ = 64;
120+ uint64_t ubSize_ = 262144;
121+ uint64_t tilingDataSize_ = 4096;
122+ void* compileInfo_ = nullptr;
123+};
124+ 
125+class TilingContextFaker : public OpTilingContextBuilder {
126+public:
127+ TilingContextFaker& SetOpType(const std::string opType);
128+ 
129+ /* only one can be choosed from IrInstanceNum */
130+ TilingContextFaker& NodeIoNum(size_t inputNum, size_t outputNum);
131+ 
132+ /* can be used for dynamic inputs/outputs
133+ * only one can be choosed from NodeIoNum */
134+ TilingContextFaker& IrInstanceNum(const std::vector<uint32_t>& inputInstanceNum,
135+ const std::vector<uint32_t>& outputInstanceNum);
136+ 
137+ 
138+ 
139+ TilingContextFaker& Attr(const std::string& attrName, bool attr) {
140+ OpTilingContextBuilder::AppendAttr(attr);
141+ return *this;
142+ }
143+ TilingContextFaker& Attr(const std::string& attrName, int64_t attr) {
144+ OpTilingContextBuilder::AppendAttr(attr);
145+ return *this;
146+ }
147+ TilingContextFaker& Attr(const std::string& attrName, float attr) {
148+ OpTilingContextBuilder::AppendAttr(attr);
149+ return *this;
150+ }
151+ TilingContextFaker& Attr(const std::string& attrName, const ge::AscendString& attr) {
152+ OpTilingContextBuilder::AppendAttr(attr);
153+ return *this;
154+ }
155+ TilingContextFaker& Attr(const std::string& attrName, const std::vector<bool>& attr) {
156+ OpTilingContextBuilder::AppendAttr(attr);
157+ return *this;
158+ }
159+ TilingContextFaker& Attr(const std::string& attrName, const std::vector<int64_t>& attr) {
160+ OpTilingContextBuilder::AppendAttr(attr);
161+ return *this;
162+ }
163+ TilingContextFaker& Attr(const std::string& attrName, const std::vector<float>& attr) {
164+ OpTilingContextBuilder::AppendAttr(attr);
165+ return *this;
166+ }
167+ TilingContextFaker& Attr(const std::string& attrName, const std::vector<ge::AscendString>& attr) {
168+ OpTilingContextBuilder::AppendAttr(attr);
169+ return *this;
170+ }
171+ TilingContextFaker& Attr(const std::string& attrName, const std::vector<std::vector<int64_t>>& attr) {
172+ OpTilingContextBuilder::AppendAttr(attr);
173+ return *this;
174+ }
175+ 
176+ TilingContextFaker& InputTensors(const std::vector<Tensor *>& inputTensors);
177+ 
178+ TilingContextFaker& OutputTensors(const std::vector<Tensor *>& outputTensors);
179+ 
180+ TilingContextFaker& CompileInfo(const void* compileInfo);
181+ 
182+ TilingContextFaker& PlatformInfo(const void* platformInfo);
183+ 
184+ TilingContextFaker& DeterministicInfo(int32_t* deterministicInfo);
185+ 
186+ TilingContextFaker& TilingData(const void* tilingData);
187+ 
188+ TilingContextFaker& Workspace(const ContinuousVector* workspace);
189+ 
190+ ContextHolder<TilingContext> Build();
191+};
192+ 
193+} // namespace gert
194+#endif // OPS_MATH_DEV_TESTS_UT_COMMON_INFERSHAPE_CONTEXT_FAKER_H
@@ -0,0 +1,102 @@
1+set(OP_NAME tanh)
2+set(UT_EXE ${OP_NAME}_op_host_ut)
3+ 
4+if(NOT DEFINED SOC_VERSION)
5+ set(SOC_VERSION "Ascend910B")
6+endif()
7+ 
8+message(STATUS "Building UT for SOC_VERSION: ${SOC_VERSION}")
9+ 
10+set(OP_HOST_DIR ${CMAKE_CURRENT_SOURCE_DIR}/../../../op_host)
11+ 
12+set(TILING_SRC ${OP_HOST_DIR}/${OP_NAME}_tiling.cpp)
13+set(TILING_DATA_H ${CMAKE_CURRENT_SOURCE_DIR}/../../../op_kernel/${OP_NAME}_tiling_data.h)
14+ 
15+set(OP_HOST_SRCS
16+ ${OP_HOST_DIR}/${OP_NAME}_def.cpp
17+ ${OP_HOST_DIR}/${OP_NAME}_infershape.cpp
18+ ${TILING_SRC}
19+)
20+ 
21+set(UT_TEST_SRCS
22+ test_${OP_NAME}_tiling.cpp
23+)
24+ 
25+set(UT_MAIN_SRC ${CMAKE_CURRENT_SOURCE_DIR}/test_op_host_main.cpp)
26+ 
27+set(OP_HOST_SO ${OP_NAME}_op_host_ut_lib)
28+ 
29+add_library(${OP_HOST_SO} SHARED
30+ ${OP_HOST_SRCS}
31+)
32+ 
33+target_include_directories(${OP_HOST_SO} PRIVATE
34+ ${UT_COMMON_INCLUDE_DIRS}
35+ ${OP_HOST_DIR}
36+ ${CMAKE_CURRENT_SOURCE_DIR}/../../../op_kernel
37+)
38+ 
39+target_compile_definitions(${OP_HOST_SO} PRIVATE
40+ BUILD_SOC_VERSION=${SOC_VERSION}
41+ _GLIBCXX_USE_CXX11_ABI=0
42+)
43+ 
44+target_link_libraries(${OP_HOST_SO} PRIVATE
45+ -Wl,--no-as-needed
46+ -Wl,--whole-archive
47+ $ENV{ASCEND_HOME_PATH}/lib64/librt2_registry.a
48+ $ENV{ASCEND_HOME_PATH}/lib64/libtiling_api.a
49+ -Wl,--no-whole-archive
50+ $ENV{ASCEND_HOME_PATH}/lib64/libopp_registry.so
51+ $ENV{ASCEND_HOME_PATH}/lib64/libregister.so
52+ $ENV{ASCEND_HOME_PATH}/lib64/libmmpa.so
53+)
54+ 
55+add_executable(${UT_EXE}
56+ ${UT_MAIN_SRC}
57+ ${UT_TEST_SRCS}
58+ ${UT_COMMON_SRCS}
59+)
60+ 
61+target_include_directories(${UT_EXE} PRIVATE
62+ ${UT_COMMON_INCLUDE_DIRS}
63+ ${OP_HOST_DIR}
64+ ${CMAKE_CURRENT_SOURCE_DIR}/../../../op_kernel
65+)
66+ 
67+target_compile_definitions(${UT_EXE} PRIVATE
68+ BUILD_SOC_VERSION=${SOC_VERSION}
69+ _GLIBCXX_USE_CXX11_ABI=0
70+)
71+ 
72+target_compile_options(${UT_EXE} PUBLIC
73+ -fPIE
74+ -fno-access-control
75+)
76+ 
77+target_link_directories(${UT_EXE} PRIVATE
78+ ${UT_COMMON_LIB_DIRS}
79+ ${CMAKE_CURRENT_BINARY_DIR}
80+)
81+ 
82+target_link_libraries(${UT_EXE} PRIVATE
83+ gtest
84+ gtest_main
85+ -Wl,--no-as-needed
86+ $ENV{ASCEND_HOME_PATH}/lib64/libmetadef.so
87+ $ENV{ASCEND_HOME_PATH}/lib64/libunified_dlog.so
88+ $ENV{ASCEND_HOME_PATH}/lib64/libopp_registry.so
89+ $ENV{ASCEND_HOME_PATH}/lib64/libregister.so
90+ $ENV{ASCEND_HOME_PATH}/lib64/libgraph.so
91+ $ENV{ASCEND_HOME_PATH}/lib64/libgraph_base.so
92+ $ENV{ASCEND_HOME_PATH}/lib64/libplatform.so
93+ $ENV{ASCEND_HOME_PATH}/lib64/libc_sec.so
94+ $ENV{ASCEND_HOME_PATH}/lib64/libmmpa.so
95+ dl
96+ pthread
97+)
98+ 
99+set_target_properties(${UT_EXE} PROPERTIES
100+ INSTALL_RPATH "$ORIGIN;$ENV{ASCEND_HOME_PATH}/lib64"
101+ BUILD_WITH_INSTALL_RPATH TRUE
102+)
@@ -0,0 +1,50 @@
1+#include <filesystem>
2+#include <iostream>
3+#include <gtest/gtest.h>
4+#include "base/registry/op_impl_space_registry_v2.h"
5+ 
6+using namespace std;
7+ 
8+class OpHostUtEnvironment : public testing::Environment {
9+public:
10+ OpHostUtEnvironment(char** argv) : argv_(argv)
11+ {}
12+
13+ virtual void SetUp() {
14+ cout << "Global Environment SetUp." << endl;
15+ 
16+ std::filesystem::path currDir = argv_[0];
17+ if (currDir.is_relative()) {
18+ currDir = std::filesystem::weakly_canonical(std::filesystem::current_path() / currDir);
19+ } else {
20+ currDir = std::filesystem::canonical(currDir);
21+ }
22+ 
23+ string opHostSoPath = currDir.parent_path().string() + string("/libtanh_op_host_ut_lib.so");
24+ cout << "Loading op_host .so from: " << opHostSoPath << endl;
25+
26+ gert::OppSoDesc oppSoDesc({ge::AscendString(opHostSoPath.c_str())}, "op_host_so");
27+
28+ shared_ptr<gert::OpImplSpaceRegistryV2> opImplSpaceRegistryV2 = make_shared<gert::OpImplSpaceRegistryV2>();
29+ if (opImplSpaceRegistryV2->AddSoToRegistry(oppSoDesc) == ge::GRAPH_FAILED) {
30+ cerr << "Failed to add .so to registry." << endl;
31+ return;
32+ }
33+ 
34+ gert::DefaultOpImplSpaceRegistryV2::GetInstance().SetSpaceRegistry(opImplSpaceRegistryV2);
35+ cout << "OpImplSpaceRegistryV2 initialized successfully" << endl;
36+ }
37+ 
38+ virtual void TearDown() {
39+ cout << "Global Environment TearDown" << endl;
40+ }
41+ 
42+private:
43+ char** argv_;
44+};
45+ 
46+int main(int argc, char** argv) {
47+ testing::InitGoogleTest(&argc, argv);
48+ testing::AddGlobalTestEnvironment(new OpHostUtEnvironment(argv));
49+ return RUN_ALL_TESTS();
50+}
@@ -0,0 +1,85 @@
1+#include <iostream>
2+#include <gtest/gtest.h>
3+#include "tiling_context_faker.h"
4+#include "tiling_case_executor.h"
5+#include "tanh_tiling_data.h"
6+ 
7+namespace TanhUT {
8+using namespace std;
9+using namespace ge;
10+using namespace gert;
11+static const std::string OP_NAME = "Tanh";
12+ 
13+struct TanhTestParam {
14+ std::string caseName;
15+ std::initializer_list<int64_t> xShape;
16+ ge::DataType xDtype;
17+ ge::Format xFormat;
18+ std::initializer_list<int64_t> yShape;
19+ ge::DataType yDtype;
20+ ge::Format yFormat;
21+ std::string socVersion;
22+ ge::graphStatus status;
23+ uint64_t expectTilingKey;
24+ std::string expectTilingData;
25+ std::vector<size_t> expectWorkspaces;
26+ uint64_t maxAIVNum;
27+ uint64_t ubSize;
28+ uint64_t tilingDataMaxSize;
29+};
30+ 
31+// TODO: 以下期望值基于初始模板实现,修改 tiling 逻辑后请更新:
32+// expectTilingKey: 参考 op_kernel/tanh_tiling_key.h 和 op_host/tanh_tiling.cpp 中 tilingKey 的逻辑
33+// expectTilingData: 参考 op_host/tanh_tiling.cpp 中 TilingData 各字段的赋值
34+// expectWorkspaces: 参考 op_host/tanh_tiling.cpp 中 GetWorkspaceSize 的逻辑
35+static TanhTestParam testCases[] = {
36+ {"tanh_0", {8, 2048}, ge::DT_FLOAT, ge::FORMAT_ND, {8, 2048}, ge::DT_FLOAT, ge::FORMAT_ND, "Ascend910B", ge::GRAPH_SUCCESS, 1UL, "0 1 0 ", {0}, 64, 262144, 4096},
37+};
38+ 
39+class TanhTilingTest : public testing::TestWithParam<TanhTestParam> {
40+protected:
41+ static void SetUpTestCase() {
42+ std::cout << "TanhTilingTest SetUp." << std::endl;
43+ }
44+ static void TearDownTestCase() {
45+ std::cout << "TanhTilingTest TearDown." << std::endl;
46+ }
47+};
48+ 
49+struct TanhCompileInfo {} compileInfo;
50+ 
51+static void TestOneParamCase(const TanhTestParam &param)
52+{
53+ gert::StorageShape xShape = {param.xShape, param.xShape};
54+ gert::StorageShape yShape = {param.yShape, param.yShape};
55+ std::vector<gert::TilingContextPara::TensorDescription> inputTensorDesc_(
56+ {{xShape, param.xDtype, param.xFormat}});
57+ std::vector<gert::TilingContextPara::TensorDescription> outputTensorDesc_(
58+ {{yShape, param.yDtype, param.yFormat}});
59+ std::vector<gert::TilingContextPara::OpAttr> attrs_;
60+ 
61+ gert::TilingContextPara tilingContextPara(
62+ OP_NAME,
63+ inputTensorDesc_,
64+ outputTensorDesc_,
65+ attrs_,
66+ &compileInfo,
67+ param.maxAIVNum,
68+ param.ubSize,
69+ param.tilingDataMaxSize);
70+ ExecuteTestCase(tilingContextPara, param.status, param.expectTilingKey,
71+ param.expectTilingData, param.expectWorkspaces);
72+}
73+ 
74+TEST_P(TanhTilingTest, tiling_test)
75+{
76+ const TanhTestParam &param = GetParam();
77+ TestOneParamCase(param);
78+}
79+ 
80+INSTANTIATE_TEST_SUITE_P(
81+ TanhTilingTests,
82+ TanhTilingTest,
83+ testing::ValuesIn(testCases));
84+ 
85+}
@@ -0,0 +1,97 @@
1+set(OP_NAME tanh)
2+set(UT_EXE ${OP_NAME}_op_kernel_ut)
3+ 
4+message(STATUS "Building op_kernel UT for: ${OP_NAME}")
5+ 
6+if(DEFINED ENV{ASCEND_HOME_PATH} AND EXISTS "$ENV{ASCEND_HOME_PATH}")
7+ set(CANN_HOME "$ENV{ASCEND_HOME_PATH}")
8+elseif(EXISTS "/usr/local/Ascend/cann-8.5.0")
9+ set(CANN_HOME "/usr/local/Ascend/cann-8.5.0")
10+else()
11+ message(FATAL_ERROR "Cannot find CANN installation")
12+endif()
13+ 
14+message(STATUS "Using CANN_HOME: ${CANN_HOME}")
15+ 
16+set(SOC_VERSION "Ascend910B1")
17+ 
18+set(tikicpulib_DIR ${CANN_HOME}/tools/tikicpulib/lib/cmake)
19+find_package(tikicpulib REQUIRED)
20+ 
21+set(OP_KERNEL_DIR ${CMAKE_CURRENT_SOURCE_DIR}/../../../op_kernel)
22+ 
23+set(UT_TEST_SRCS
24+ test_${OP_NAME}.cpp
25+)
26+ 
27+set(UT_INCLUDE_DIRS
28+ ${CMAKE_CURRENT_SOURCE_DIR}
29+ ${OP_KERNEL_DIR}
30+ ${UT_COMMON_INCLUDE_DIRS}
31+ ${CANN_HOME}/tools/tikicpulib/lib/include
32+ ${CANN_HOME}/compiler/ascendc/include/highlevel_api
33+ ${CANN_HOME}/compiler/ascendc/include/basic_api/impl
34+ ${CANN_HOME}/compiler/ascendc/include/basic_api
35+ ${CANN_HOME}/compiler/ascendc/include
36+ ${CANN_HOME}/compiler/ascendc/include/ascendc/host_api/tiling
37+ ${CANN_HOME}/compiler/ascendc/include/ascendc/host_api
38+ ${CANN_HOME}/compiler/ascendc/include/ascendc
39+ ${CANN_HOME}/aarch64-linux/include
40+ ${CANN_HOME}/aarch64-linux/asc/include/basic_api
41+ ${CANN_HOME}/aarch64-linux/asc/include
42+ ${CANN_HOME}/aarch64-linux/asc/include/interface
43+ ${CANN_HOME}/aarch64-linux/asc/impl/basic_api
44+ ${CANN_HOME}/aarch64-linux/asc
45+ ${CANN_HOME}/compiler/tikcpp/tikcfw
46+ ${CANN_HOME}/compiler/tikcpp/tikcfw/impl
47+ ${CANN_HOME}/compiler/tikcpp/tikcfw/interface
48+ ${CANN_HOME}/aarch64-linux/include/graph
49+ ${CANN_HOME}/aarch64-linux/ascendc/include/highlevel_api
50+ ${CANN_HOME}/aarch64-linux/asc/impl/basic_api/utils
51+)
52+ 
53+add_executable(${UT_EXE}
54+ ${UT_TEST_SRCS}
55+)
56+ 
57+target_include_directories(${UT_EXE} PRIVATE
58+ ${UT_INCLUDE_DIRS}
59+)
60+ 
61+target_compile_definitions(${UT_EXE} PRIVATE
62+ _GLIBCXX_USE_CXX11_ABI=0
63+ __aicore__=
64+ __NPU_TILING__=1
65+)
66+ 
67+target_compile_options(${UT_EXE} PRIVATE
68+ -Wall
69+ -Wno-deprecated-declarations
70+ -Wno-unused-variable
71+ -Wno-unknown-pragmas
72+ -fno-access-control
73+ -std=c++17
74+)
75+ 
76+target_link_directories(${UT_EXE} PRIVATE
77+ ${UT_COMMON_LIB_DIRS}
78+ ${CANN_HOME}/lib64
79+)
80+ 
81+target_link_libraries(${UT_EXE} PRIVATE
82+ gtest_main
83+ gtest
84+ ${CANN_HOME}/lib64/libmmpa.so
85+ tikicpulib::${SOC_VERSION}
86+ dl
87+ pthread
88+)
89+ 
90+set_target_properties(${UT_EXE} PROPERTIES
91+ INSTALL_RPATH "$ORIGIN;${CANN_HOME}/lib64;${CANN_HOME}/tools/tikicpulib/lib;${CANN_HOME}/tools/tikicpulib/lib/Ascend910B1;${CANN_HOME}/aarch64-linux/simulator/Ascend910B1/lib"
92+ BUILD_WITH_INSTALL_RPATH TRUE
93+)
94+ 
95+message(STATUS " Kernel UT executable: ${UT_EXE}")
96+message(STATUS " CANN Home: ${CANN_HOME}")
97+message(STATUS " SOC Version: ${SOC_VERSION}")
@@ -0,0 +1,275 @@
1+#!/usr/bin/env python3
2+# -*- coding: utf-8 -*-
3+ 
4+import sys
5+import numpy as np
6+from ml_dtypes import bfloat16
7+import glob
8+import os
9+ 
10+curr_dir = os.path.dirname(os.path.realpath(__file__))
11+ 
12+ 
13+def get_threshold_by_dtype(dtype):
14+ """
15+ 根据数据类型获取通过阈值(社区标准)
16+
17+ 精度标准:
18+ - FLOAT16: threshold = 2^-10 ≈ 0.000977
19+ - BFLOAT16: threshold = 2^-7 ≈ 0.00781
20+ - FLOAT32: threshold = 2^-13 ≈ 0.000122
21+ - HiFLOAT32: threshold = 2^-11 ≈ 0.000488
22+ - FLOAT8 E4M3: threshold = 2^-3 ≈ 0.125
23+ - FLOAT8 E5M2: threshold = 2^-2 ≈ 0.25
24+ """
25+ dtype_str = str(dtype).lower().replace(' ', '').replace('_', '')
26+
27+ thresholds = {
28+ 'float16': 2 ** (-10),
29+ 'bfloat16': 2 ** (-7),
30+ 'float32': 2 ** (-13),
31+ 'float64': 2 ** (-13),
32+ 'hifloat32': 2 ** (-11),
33+ }
34+
35+ if 'float8e4m3' in dtype_str or 'fp8e4m3' in dtype_str:
36+ return 2 ** (-3)
37+ elif 'float8e5m2' in dtype_str or 'fp8e5m2' in dtype_str:
38+ return 2 ** (-2)
39+
40+ return thresholds.get(dtype_str, 2 ** (-13))
41+ 
42+ 
43+def calculate_mare(actual, golden):
44+ """计算最大相对误差(MARE)"""
45+ relative_errors = np.abs(actual - golden) / np.maximum(np.abs(golden), 1e-6)
46+ return np.max(relative_errors)
47+ 
48+ 
49+def calculate_mere(actual, golden):
50+ """计算平均相对误差(MERE)"""
51+ relative_errors = np.abs(actual - golden) / np.maximum(np.abs(golden), 1e-6)
52+ return np.mean(relative_errors)
53+ 
54+ 
55+def compare_data_float(golden_file_lists, output_file_lists, d_type):
56+ """
57+ 浮点类型精度比对
58+
59+ 通过标准:
60+ - MERE < threshold
61+ - MARE < 10 * threshold
62+ """
63+ if d_type == "float16":
64+ np_dtype = np.float16
65+ elif d_type == "float32":
66+ np_dtype = np.float32
67+ elif d_type == "float64":
68+ np_dtype = np.float64
69+ elif d_type == "bfloat16":
70+ np_dtype = bfloat16
71+ elif d_type == "hifloat32":
72+ np_dtype = np.float32
73+ elif d_type in ("fp8_e4m3fn", "fp8_e5m2"):
74+ np_dtype = np.uint8
75+ else:
76+ np_dtype = np.float32
77+
78+ threshold = get_threshold_by_dtype(d_type)
79+ mare_threshold = 10 * threshold
80+
81+ def _read_bin(path):
82+ raw = np.fromfile(path, np_dtype)
83+ if d_type == "fp8_e4m3fn":
84+ from ml_dtypes import float8_e4m3fn
85+ return raw.view(float8_e4m3fn).astype(np.float32)
86+ elif d_type == "fp8_e5m2":
87+ from ml_dtypes import float8_e5m2
88+ return raw.view(float8_e5m2).astype(np.float32)
89+ return raw
90+
91+ data_same = True
92+ # 当 golden 文件数 < output 文件数时,将多个 output 拼接后与单个 golden 比对
93+ if len(golden_file_lists) == 1 and len(output_file_lists) > 1:
94+ tmp_gold = _read_bin(golden_file_lists[0])
95+ tmp_out = np.concatenate([_read_bin(f) for f in output_file_lists])
96+ mere = calculate_mere(tmp_out, tmp_gold)
97+ mare = calculate_mare(tmp_out, tmp_gold)
98+
99+ mere_pass = mere < threshold
100+ mare_pass = mare < mare_threshold
101+ is_pass = mere_pass and mare_pass
102+
103+ if is_pass:
104+ print(f"PASSED! MERE={mere:.6f}, MARE={mare:.6f}")
105+ else:
106+ print(f"FAILED! MERE={mere:.6f} (threshold={threshold:.6f}), MARE={mare:.6f} (threshold={mare_threshold:.6f})")
107+ diff = np.abs(tmp_out - tmp_gold)
108+ diff_idx = np.argsort(diff)[-5:][::-1]
109+ for idx in diff_idx:
110+ print(f" index: {idx}, output: {tmp_out[idx]}, golden: {tmp_gold[idx]}")
111+ data_same = False
112+ else:
113+ for gold, out in zip(golden_file_lists, output_file_lists):
114+ tmp_out = _read_bin(out)
115+ tmp_gold = _read_bin(gold)
116+
117+ mere = calculate_mere(tmp_out, tmp_gold)
118+ mare = calculate_mare(tmp_out, tmp_gold)
119+
120+ mere_pass = mere < threshold
121+ mare_pass = mare < mare_threshold
122+ is_pass = mere_pass and mare_pass
123+
124+ if is_pass:
125+ print(f"PASSED! MERE={mere:.6f}, MARE={mare:.6f}")
126+ else:
127+ print(f"FAILED! MERE={mere:.6f} (threshold={threshold:.6f}), MARE={mare:.6f} (threshold={mare_threshold:.6f})")
128+ diff = np.abs(tmp_out - tmp_gold)
129+ diff_idx = np.argsort(diff)[-5:][::-1]
130+ for idx in diff_idx:
131+ print(f" index: {idx}, output: {tmp_out[idx]}, golden: {tmp_gold[idx]}")
132+ data_same = False
133+ return data_same
134+ 
135+ 
136+def compare_data_integer(golden_file_lists, output_file_lists, d_type):
137+ """
138+ 整数类型精度比对
139+
140+ 通过标准: 二进制一致 或 绝对误差为0
141+ """
142+ if d_type == "int32":
143+ np_dtype = np.int32
144+ elif d_type == "int8":
145+ np_dtype = np.int8
146+ elif d_type == "int16":
147+ np_dtype = np.int16
148+ elif d_type == "int64":
149+ np_dtype = np.int64
150+ elif d_type == "uint8":
151+ np_dtype = np.uint8
152+ elif d_type == "uint16":
153+ np_dtype = np.uint16
154+ elif d_type == "uint32":
155+ np_dtype = np.uint32
156+ elif d_type == "uint64":
157+ np_dtype = np.uint64
158+ elif d_type == "bool":
159+ np_dtype = np.bool_
160+ else:
161+ np_dtype = np.int32
162+
163+ data_same = True
164+ if len(golden_file_lists) == 1 and len(output_file_lists) > 1:
165+ tmp_gold = np.fromfile(golden_file_lists[0], np_dtype)
166+ tmp_out = np.concatenate([np.fromfile(f, np_dtype) for f in output_file_lists])
167+ bitwise_match = np.array_equal(tmp_out, tmp_gold)
168+ abs_error_zero = np.all(np.abs(tmp_out.astype(np.int64) - tmp_gold.astype(np.int64)) == 0)
169+ is_pass = bitwise_match or abs_error_zero
170+ if is_pass:
171+ print(f"PASSED! bitwise_match={bitwise_match}, abs_error_zero={abs_error_zero}")
172+ else:
173+ print(f"FAILED!")
174+ data_same = False
175+ else:
176+ for gold, out in zip(golden_file_lists, output_file_lists):
177+ tmp_out = np.fromfile(out, np_dtype)
178+ tmp_gold = np.fromfile(gold, np_dtype)
179+
180+ # 检查二进制一致
181+ bitwise_match = np.array_equal(tmp_out, tmp_gold)
182+ # 检查绝对误差为0
183+ abs_error_zero = np.all(np.abs(tmp_out.astype(np.int64) - tmp_gold.astype(np.int64)) == 0)
184+
185+ is_pass = bitwise_match or abs_error_zero
186+
187+ if is_pass:
188+ print(f"PASSED! bitwise_match={bitwise_match}, abs_error_zero={abs_error_zero}")
189+ else:
190+ print(f"FAILED!")
191+ diff_idx = np.where(tmp_out != tmp_gold)[0][:5]
192+ for idx in diff_idx:
193+ print(f" index: {idx}, output: {tmp_out[idx]}, golden: {tmp_gold[idx]}")
194+ data_same = False
195+ return data_same
196+ 
197+ 
198+def get_file_lists(dtype):
199+ golden_file_lists = sorted(glob.glob(curr_dir + "/*golden*.bin"))
200+ output_file_lists = sorted(glob.glob(curr_dir + "/*output*.bin"))
201+ return golden_file_lists, output_file_lists
202+ 
203+ 
204+def infer_dtype_from_filename():
205+ """从 golden 文件名推断 dtype"""
206+ golden_files = glob.glob(curr_dir + "/*golden*.bin")
207+ if not golden_files:
208+ return "float32"
209+
210+ filename = os.path.basename(golden_files[0])
211+ for dtype in ["float16", "float32", "float64", "bfloat16", "fp8_e4m3fn", "fp8_e5m2", "hifloat32", "int8", "int16", "int32", "int64", "uint8", "uint16", "uint32", "uint64", "bool"]:
212+ if filename.startswith(dtype + "_"):
213+ return dtype
214+ return "float32"
215+ 
216+ 
217+def infer_dtype_from_single_filename(filename):
218+ """从单个文件名推断 dtype"""
219+ basename = os.path.basename(filename)
220+ for dt in ["float16", "float32", "float64", "bfloat16", "fp8_e4m3fn", "fp8_e5m2", "hifloat32", "int8", "int16", "int32", "int64", "uint8", "uint16", "uint32", "uint64", "bool"]:
221+ if basename.startswith(dt + "_"):
222+ return dt
223+ return "float32"
224+ 
225+ 
226+def process(d_type):
227+ golden_file_lists, output_file_lists = get_file_lists(d_type)
228+
229+ if not golden_file_lists and not output_file_lists:
230+ print("No golden or output files found (no-output operator), skipping comparison")
231+ return True
232+
233+ if not golden_file_lists or not output_file_lists:
234+ print("ERROR: No golden or output files found")
235+ return False
236+
237+ # 单 golden 多 output:拼接比对
238+ if len(golden_file_lists) == 1 and len(output_file_lists) > 1:
239+ if d_type in ["int8", "int16", "int32", "int64", "uint8", "uint16", "uint32", "uint64", "bool"]:
240+ result = compare_data_integer(golden_file_lists, output_file_lists, d_type)
241+ else:
242+ result = compare_data_float(golden_file_lists, output_file_lists, d_type)
243+ print("compare result:", result)
244+ return result
245+ 
246+ if len(golden_file_lists) != len(output_file_lists):
247+ print(f"ERROR: file count mismatch: golden={len(golden_file_lists)}, output={len(output_file_lists)}")
248+ return False
249+
250+ # 逐对比对,每个 golden 文件独立推断 dtype
251+ all_pass = True
252+ for gold, out in zip(golden_file_lists, output_file_lists):
253+ file_dtype = infer_dtype_from_single_filename(gold)
254+ print(f"Comparing: {os.path.basename(gold)} vs {os.path.basename(out)} (dtype={file_dtype})")
255+ if file_dtype in ["int8", "int16", "int32", "int64", "uint8", "uint16", "uint32", "uint64", "bool"]:
256+ pair_result = compare_data_integer([gold], [out], file_dtype)
257+ else:
258+ pair_result = compare_data_float([gold], [out], file_dtype)
259+ if not pair_result:
260+ all_pass = False
261+
262+ print("compare result:", all_pass)
263+ return all_pass
264+ 
265+ 
266+if __name__ == '__main__':
267+ # 从文件名推断 dtype,或使用命令行参数
268+ if len(sys.argv) >= 2:
269+ d_type = sys.argv[1]
270+ else:
271+ d_type = infer_dtype_from_filename()
272+ print(f"从文件名推断 dtype: {d_type}")
273+
274+ ret = process(d_type)
275+ exit(0 if ret else 1)
@@ -0,0 +1,58 @@
1+#!/usr/bin/env python3
2+# -*- coding: utf-8 -*-
3+ 
4+import os
5+import glob
6+import numpy as np
7+from ml_dtypes import bfloat16
8+ 
9+def impl(x):
10+ # 验算函数:逐元素 tanh,float64 中间计算,最后一次性转回输入 dtype
11+ dtype = x.dtype
12+ return np.tanh(x.astype(np.float64)).astype(dtype)
13+ 
14+ 
15+if __name__ == "__main__":
16+ # 清理bin文件
17+ for f in glob.glob("*.bin"):
18+ os.remove(f)
19+
20+ # 从 JSON 第一个 case 获取参数
21+ d_type = "float32"
22+ d_type_dict = {
23+ "float32": np.float32,
24+ "float16": np.float16,
25+ "bfloat16": bfloat16,
26+ "float64": np.float64,
27+ "int8": np.int8,
28+ "int16": np.int16,
29+ "int32": np.int32,
30+ "int64": np.int64,
31+ "uint8": np.uint8,
32+ "uint16": np.uint16,
33+ "uint32": np.uint32,
34+ "uint64": np.uint64,
35+ "bool": np.bool_,
36+ "fp8_e4m3fn": np.uint8,
37+ "fp8_e5m2": np.uint8,
38+ }
39+ np_type = d_type_dict[d_type]
40+
41+ # 生成输入数据
42+ input_x = np.ones((8, 2048)).astype(d_type_dict["float32"])
43+ 
44+
45+ # 计算 golden 数据
46+ golden = impl(input_x)
47+
48+ # 保存数据到文件
49+ input_x.astype(d_type_dict["float32"]).tofile(f"{d_type}_input_tanh_x.bin")
50+ if golden is not None:
51+ if isinstance(golden, (list, tuple)):
52+ with open("float32_golden_tanh.bin", "wb") as _f:
53+ for _g in golden:
54+ _g.astype(d_type_dict["float32"]).tofile(_f)
55+ else:
56+ golden.astype(d_type_dict["float32"]).tofile("float32_golden_tanh.bin")
57+
58+ print(f"生成完成: dtype={d_type}")
@@ -0,0 +1,43 @@
1+/*!
2+ * \file tanh_tiling.h
3+ * \brief Tanh tiling 数据定义
4+ */
5+ 
6+#ifndef _I_TANH_TILING_H_
7+#define _I_TANH_TILING_H_
8+ 
9+#include <cstdint>
10+#include <cstring>
11+#include "../../../op_kernel/tanh_tiling_data.h"
12+#include "tikicpulib.h"
13+#include "kernel_operator.h"
14+#include "kernel_tiling/kernel_tiling.h"
15+#include "graph/c_types.h"
16+#include "ascendc/host_api/tiling/template_argument.h"
17+ 
18+#ifndef __aicore__
19+#define __aicore__
20+#endif
21+ 
22+#ifndef __gm__
23+#define __gm__
24+#endif
25+ 
26+#ifndef __ubuf__
27+#define __ubuf__
28+#endif
29+ 
30+inline void InitTilingData(uint8_t *tiling, TanhTilingData *constData)
31+{
32+ memcpy(constData, tiling, sizeof(TanhTilingData));
33+}
34+ 
35+#define GET_TILING_DATA_WITH_STRUCT(tilingStruct, tilingData, tilingArg) \
36+ tilingStruct tilingData; \
37+ InitTilingData(tilingArg, &tilingData)
38+ 
39+#define GET_TILING_DATA(tilingData, tilingArg) \
40+ TanhTilingData tilingData; \
41+ InitTilingData(tilingArg, &tilingData)
42+ 
43+#endif
@@ -0,0 +1,97 @@
1+/*!
2+ * \file test_tanh.cpp
3+ * \brief Tanh 算子 kernel UT 测试
4+ *
5+ * 独立运行,直接构造 tilingData,不依赖 op_host UT
6+ */
7+ 
8+#include "tanh_tiling.h"
9+#include "../../../op_kernel/tanh.cpp"
10+ 
11+#include <array>
12+#include <vector>
13+#include <iostream>
14+#include <cstdint>
15+#include <cstdlib>
16+#include <cstring>
17+#include <fstream>
18+#include "gtest/gtest.h"
19+#include "tikicpulib.h"
20+ 
21+using namespace std;
22+ 
23+static uint16_t FloatToHalf(float f) {
24+ uint32_t bits;
25+ memcpy(&bits, &f, sizeof(float));
26+ uint32_t sign = (bits >> 16) & 0x8000;
27+ int32_t exp = ((bits >> 23) & 0xff) - 127 + 15;
28+ uint32_t mant = (bits >> 13) & 0x3ff;
29+ if (exp <= 0) return sign;
30+ if (exp >= 31) return sign | 0x7c00;
31+ return sign | (exp << 10) | mant;
32+}
33+ 
34+static uint16_t FloatToBFloat16(float f) {
35+ uint32_t bits;
36+ memcpy(&bits, &f, sizeof(float));
37+ return (uint16_t)(bits >> 16);
38+}
39+ 
40+class TanhKernelTest : public testing::Test {
41+protected:
42+ static void SetUpTestCase()
43+ {
44+ cout << "TanhKernelTest SetUp" << endl;
45+ }
46+ static void TearDownTestCase()
47+ {
48+ cout << "TanhKernelTest TearDown" << endl;
49+ }
50+};
51+ 
52+TEST_F(TanhKernelTest, test_kernel_run)
53+{
54+ constexpr size_t size = 16384;
55+ constexpr size_t tilingDataSize = sizeof(TanhTilingData);
56+ constexpr uint32_t numBlocks = 1;
57+ 
58+ constexpr size_t xByteSize = 16384 * 4;
59+ constexpr size_t yByteSize = 16384 * 4;
60+ std::vector<float> xHost(16384, 1);
61+ std::vector<float> yHost(16384, 0);
62+
63+
64+ uint8_t* x = (uint8_t*)AscendC::GmAlloc(xByteSize);
65+ uint8_t* y = (uint8_t*)AscendC::GmAlloc(yByteSize);
66+ uint8_t* workspace = (uint8_t*)AscendC::GmAlloc(32);
67+ uint8_t* tiling = (uint8_t*)AscendC::GmAlloc(tilingDataSize);
68+
69+ memcpy(x, xHost.data(), xByteSize);
70+
71+ // TODO: 以下 tilingData 字段基于初始模板的 TilingData 结构。
72+ // 修改 TilingData 结构后请更新字段名和赋值:
73+ // 参考 op_kernel/tanh_tiling_data.h 中的字段定义
74+ // 参考 op_host/tanh_tiling.cpp 中的 tiling 计算逻辑
75+ TanhTilingData* tilingData = reinterpret_cast<TanhTilingData*>(tiling);
76+ tilingData->totalNum = size;
77+ tilingData->blockFactor = size;
78+ tilingData->ubFactor = size;
79+
80+ // TODO: tilingKey 应与 op_host/tanh_tiling.cpp 中 SetTilingKey 设置的值一致
81+ ICPU_SET_TILING_KEY(1);
82+ AscendC::SetKernelMode(KernelMode::AIV_MODE);
83+
84+ ICPU_RUN_KF((tanh<1>), numBlocks, x, y, workspace, tiling);
85+
86+ // 将动态输出的 packed buffer 拆回 individual buffers
87+
88+
89+ // 将 output 数据保存到 bin 文件供 compare_data.py 比对
90+ memcpy(yHost.data(), y, yByteSize);
91+ { std::ofstream _ofs("float32_output_tanh_0.bin", std::ios::binary); _ofs.write(reinterpret_cast<const char*>(yHost.data()), yByteSize); }
92+
93+ AscendC::GmFree(x);
94+ AscendC::GmFree(y);
95+ AscendC::GmFree(workspace);
96+ AscendC::GmFree(tiling);
97+}