已开启
【代码侦探Challenge08】赛诺信致(北京)软件技术有限公司 - challenge08_tanhcustom #4054
小东创建于 20 天前
【代码侦探Challenge08】赛诺信致(北京)软件技术有限公司 - challenge08_tanhcustom #4054
已开启
共 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,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 | + | ||
| 2 | + | ||
| 3 | + | ||
| 4 | + | ||
| 5 | + | ||
| 6 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | + | ||
| 10 | + do { \ | ||
| 11 | + if (!(cond)) { \ | ||
| 12 | + return_expr; \ | ||
| 13 | + } \ | ||
| 14 | + } while (0) | ||
| 15 | + | ||
| 16 | + | ||
| 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 | + | ||
| 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 | + | ||
| 7 | + | ||
| 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 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | + | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 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 | + | ||
| 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 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | + | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 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 | + | ||
| @@ -0,0 +1,14 @@ | |||
| 1 | +/*! | ||
| 2 | + * \file tanh_tiling_data.h | ||
| 3 | + * \brief tiling data struct | ||
| 4 | + */ | ||
| 5 | + | ||
| 6 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | +struct TanhTilingData { | ||
| 10 | + int64_t totalNum = 0; // 总元素数量 | ||
| 11 | + int64_t blockFactor = 1; // 每个核处理的元素数量 | ||
| 12 | + int64_t ubFactor = 0; // 每次 UB 循环处理的元素数量 | ||
| 13 | +}; | ||
| 14 | + | ||
| @@ -0,0 +1,21 @@ | |||
| 1 | +/*! | ||
| 2 | + * \file tanh_tiling_key.h | ||
| 3 | + * \brief Tiling 模板参数定义 | ||
| 4 | + */ | ||
| 5 | + | ||
| 6 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | + | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 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 | + | ||
| @@ -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 | + | ||
| 2 | + | ||
| 3 | + | ||
| 4 | + | ||
| 5 | + | ||
| 6 | + | ||
| 7 | + | ||
| 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 | + | ||
| @@ -0,0 +1,101 @@ | |||
| 1 | + | ||
| 2 | + | ||
| 3 | + | ||
| 4 | + | ||
| 5 | + | ||
| 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 | + | ||
| 2 | + | ||
| 3 | + | ||
| 4 | + | ||
| 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 | + | ||
| @@ -0,0 +1,42 @@ | |||
| 1 | + | ||
| 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 | + | ||
| 2 | + | ||
| 3 | + | ||
| 4 | + | ||
| 5 | + | ||
| 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 | + | ||
| @@ -0,0 +1,295 @@ | |||
| 1 | + | ||
| 2 | + | ||
| 3 | + | ||
| 4 | + | ||
| 5 | + | ||
| 6 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | + | ||
| 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 | + | ||
| 2 | + | ||
| 3 | + | ||
| 4 | + | ||
| 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 | + | ||
| @@ -0,0 +1,71 @@ | |||
| 1 | + | ||
| 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 | + | ||
| 2 | + | ||
| 3 | + | ||
| 4 | + | ||
| 5 | + | ||
| 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 | + | ||
| @@ -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 | + | ||
| 2 | + | ||
| 3 | + | ||
| 4 | + | ||
| 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 | + | ||
| 2 | + | ||
| 3 | + | ||
| 4 | + | ||
| 5 | + | ||
| 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 ¶m) | ||
| 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 ¶m = 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 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | + | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 24 | + | ||
| 25 | + | ||
| 26 | + | ||
| 27 | + | ||
| 28 | + | ||
| 29 | + | ||
| 30 | +inline void InitTilingData(uint8_t *tiling, TanhTilingData *constData) | ||
| 31 | +{ | ||
| 32 | + memcpy(constData, tiling, sizeof(TanhTilingData)); | ||
| 33 | +} | ||
| 34 | + | ||
| 35 | + | ||
| 36 | + tilingStruct tilingData; \ | ||
| 37 | + InitTilingData(tilingArg, &tilingData) | ||
| 38 | + | ||
| 39 | + | ||
| 40 | + TanhTilingData tilingData; \ | ||
| 41 | + InitTilingData(tilingArg, &tilingData) | ||
| 42 | + | ||
| 43 | + | ||
| @@ -0,0 +1,97 @@ | |||
| 1 | +/*! | ||
| 2 | + * \file test_tanh.cpp | ||
| 3 | + * \brief Tanh 算子 kernel UT 测试 | ||
| 4 | + * | ||
| 5 | + * 独立运行,直接构造 tilingData,不依赖 op_host UT | ||
| 6 | + */ | ||
| 7 | + | ||
| 8 | + | ||
| 9 | + | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 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 | +} | ||