已合并
[社区任务] aclnnBernoulli低内存实现 #4248
hzw_rpap创建于 7月26日
[社区任务] aclnnBernoulli低内存实现 #4248
已合并
共 29 个文件变更+3371-0
| @@ -0,0 +1,36 @@ | |||
| 1 | +# ---------------------------------------------------------------------------- | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | +# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | +# ---------------------------------------------------------------------------- | ||
| 10 | + | ||
| 11 | +add_all_modules_sources( | ||
| 12 | + OPTYPE bernoulli_mask | ||
| 13 | + ACLNNTYPE aclnn_exclude | ||
| 14 | +) | ||
| 15 | + | ||
| 16 | +# A custom ACLNN package compiles most of its L0 dependency closure into | ||
| 17 | +# libcust_opapi.so. Keep those implementation symbols local to the custom | ||
| 18 | +# package so they cannot preempt identically named L0 symbols used by the | ||
| 19 | +# system libopapi.so. ACLNN_API declarations retain default visibility, so the | ||
| 20 | +# four scalar Bernoulli entry points remain exported. Do not change visibility | ||
| 21 | +# for the built-in opapi target: other custom packages may depend on its L0 ABI. | ||
| 22 | +if(ENABLE_CUSTOM AND TARGET ${OPHOST_NAME}_opapi_obj) | ||
| 23 | + # The scalar API keeps the upstream non-DAV2201 fallback for compatibility. | ||
| 24 | + # Compile that L0 wrapper into the custom package so libcust_opapi.so can be | ||
| 25 | + # loaded independently instead of relying on a feature library having | ||
| 26 | + # already populated the process-wide symbol scope. | ||
| 27 | + target_sources( | ||
| 28 | + ${OPHOST_NAME}_opapi_obj | ||
| 29 | + PRIVATE | ||
| 30 | + ${CMAKE_SOURCE_DIR}/random/stateless_bernoulli/op_api/stateless_bernoulli.cpp | ||
| 31 | + ) | ||
| 32 | + target_compile_options( | ||
| 33 | + ${OPHOST_NAME}_opapi_obj | ||
| 34 | + PRIVATE -fvisibility=hidden -fvisibility-inlines-hidden | ||
| 35 | + ) | ||
| 36 | +endif() | ||
| @@ -0,0 +1,182 @@ | |||
| 1 | +# BernoulliMask | ||
| 2 | + | ||
| 3 | +## 概述 | ||
| 4 | + | ||
| 5 | +本目录提供面向 Atlas A2/A3 训练系列产品的 `aclnnBernoulli` 和 | ||
| 6 | +`aclnnInplaceBernoulli` 低内存实现。算子保留 `DSAGenBitMask` 的 | ||
| 7 | +`seed/offset` 随机序列生成逻辑,由 Ascend C `BernoulliMask` Kernel 将压缩 | ||
| 8 | +bit mask 直接展开为目标数据类型的 `0` 或 `1`。 | ||
| 9 | + | ||
| 10 | +原实现中的全尺寸 `Fill`、`DropoutDoMask` 和部分 `Cast` 中间张量不再参与一般 | ||
| 11 | +概率路径。连续输出在满足存储条件时复用输出空间保存压缩 mask;小张量或非连续 | ||
| 12 | +输出使用独立 mask/连续结果缓冲区,以保持边界与视图语义。`prob=0` 和 `prob=1` | ||
| 13 | +分别沿用 `ZerosLike` 和 `OnesLike` 快速路径。 | ||
| 14 | + | ||
| 15 | +## 产品支持 | ||
| 16 | + | ||
| 17 | +| 产品 | 构建目标 | 支持状态 | | ||
| 18 | +| :--- | :---: | :---: | | ||
| 19 | +| Atlas A2 训练系列产品 | `ascend910b` | 支持 | | ||
| 20 | +| Atlas A3 训练系列产品 | `ascend910_93` | 支持 | | ||
| 21 | + | ||
| 22 | +## 功能与约束 | ||
| 23 | + | ||
| 24 | +- `self/out` 支持 FP16、FP32、FP64、BF16、UINT8、INT8、INT16、INT32、 | ||
| 25 | + INT64 和 BOOL; | ||
| 26 | +- `prob` 支持 FP16、FP32、FP64 和 BF16,必须为有限值且满足 | ||
| 27 | + `0 <= prob <= 1`; | ||
| 28 | +- 支持 0~8 维、标量、空 Tensor、连续 Tensor 和非连续 Tensor; | ||
| 29 | +- `self` 与 `out` 的 shape、dtype 必须一致; | ||
| 30 | +- `offset` 必须满足 `offset % 4 == 0`; | ||
| 31 | +- 相同的 `seed/offset` 生成可复现的随机序列。 | ||
| 32 | + | ||
| 33 | +## 接口 | ||
| 34 | + | ||
| 35 | +```cpp | ||
| 36 | +aclnnStatus aclnnBernoulliGetWorkspaceSize( | ||
| 37 | + const aclTensor* self, | ||
| 38 | + const aclScalar* prob, | ||
| 39 | + int64_t seed, | ||
| 40 | + int64_t offset, | ||
| 41 | + aclTensor* out, | ||
| 42 | + uint64_t* workspaceSize, | ||
| 43 | + aclOpExecutor** executor); | ||
| 44 | + | ||
| 45 | +aclnnStatus aclnnBernoulli( | ||
| 46 | + void* workspace, | ||
| 47 | + uint64_t workspaceSize, | ||
| 48 | + aclOpExecutor* executor, | ||
| 49 | + aclrtStream stream); | ||
| 50 | +``` | ||
| 51 | + | ||
| 52 | +`aclnnInplaceBernoulliGetWorkspaceSize` 和 `aclnnInplaceBernoulli` 使用相同的 | ||
| 53 | +参数语义,输出写回 `selfRef`。 | ||
| 54 | + | ||
| 55 | +## 实现说明 | ||
| 56 | + | ||
| 57 | +一般概率路径的数据流如下: | ||
| 58 | + | ||
| 59 | +```text | ||
| 60 | +self / prob / seed / offset | ||
| 61 | + -> DSAGenBitMask | ||
| 62 | + -> BernoulliMask(packed bit -> 目标 dtype 的 0/1) | ||
| 63 | + -> out | ||
| 64 | +``` | ||
| 65 | + | ||
| 66 | +对于容量和布局满足条件的连续输出,DSA packed mask 与输出复用同一块设备 | ||
| 67 | +内存。Kernel 从高地址到低地址分波展开,每一波完成后同步,避免输出覆盖尚未 | ||
| 68 | +读取的 mask。其他布局使用独立缓冲区;非连续输出通过 `ViewCopy` 写回目标 | ||
| 69 | +view。FP64 非连续写回使用 INT64 原始位模式 view,避免数值转换。 | ||
| 70 | + | ||
| 71 | +FP16/FP32 使用向量 `Select` 生成结果;整数类型和 BF16 通过向量类型转换输出; | ||
| 72 | +FP64 将 double 的两个 32 位字向量化构造后通过 `Gather` 交织。 | ||
| 73 | + | ||
| 74 | +## 目录结构 | ||
| 75 | + | ||
| 76 | +```text | ||
| 77 | +bernoulli_mask/ | ||
| 78 | +├── op_api/ ACLNN 接口、路径选择与 BernoulliMask L0 launcher | ||
| 79 | +├── op_host/ 算子定义、InferShape 与 Tiling | ||
| 80 | +├── op_kernel/ Ascend C Kernel、TilingData 与 TilingKey | ||
| 81 | +├── examples/ ACLNN 调用样例 | ||
| 82 | +└── tests/ | ||
| 83 | + ├── assets/ TTK golden | ||
| 84 | + ├── st/ ACLNN 系统测试 | ||
| 85 | + ├── ttk/ Kernel 通用及存储复用用例 | ||
| 86 | + └── ut/ Host UT | ||
| 87 | +``` | ||
| 88 | + | ||
| 89 | +## 构建 | ||
| 90 | + | ||
| 91 | +在 `ops-math` 仓库根目录执行: | ||
| 92 | + | ||
| 93 | +```bash | ||
| 94 | +# Atlas A2 | ||
| 95 | +bash build.sh --pkg --experimental --soc=ascend910b \ | ||
| 96 | + --ops=bernoulli_mask --build-type=Release | ||
| 97 | + | ||
| 98 | +# Atlas A3 | ||
| 99 | +bash build.sh --pkg --experimental --soc=ascend910_93 \ | ||
| 100 | + --ops=bernoulli_mask --build-type=Release | ||
| 101 | +``` | ||
| 102 | + | ||
| 103 | +运行包生成在 `build_out/`。安装后按照安装器提示加载 `custom_math` 的 | ||
| 104 | +`op_api/lib` 环境。 | ||
| 105 | + | ||
| 106 | +## 调用样例 | ||
| 107 | + | ||
| 108 | +构建并安装当前源码生成的自定义算子包后,在仓库根目录执行: | ||
| 109 | + | ||
| 110 | +```bash | ||
| 111 | +bash experimental/random/bernoulli_mask/examples/run.sh | ||
| 112 | +``` | ||
| 113 | + | ||
| 114 | +样例使用标准两段式 ACLNN 接口执行 FP32 用例,并检查输出值域、固定随机状态 | ||
| 115 | +的可复现性和样本均值。其他安装路径及仅编译方式见 | ||
| 116 | +[`examples/README.md`](examples/README.md)。 | ||
| 117 | + | ||
| 118 | +## 测试 | ||
| 119 | + | ||
| 120 | +### Host UT | ||
| 121 | + | ||
| 122 | +```bash | ||
| 123 | +bash build.sh -u --ophost --experimental --soc=ascend910b \ | ||
| 124 | + --ops=bernoulli_mask | ||
| 125 | +``` | ||
| 126 | + | ||
| 127 | +### ACLNN ST | ||
| 128 | + | ||
| 129 | +从当前 checkout 构建并安装自定义算子包后执行: | ||
| 130 | + | ||
| 131 | +```bash | ||
| 132 | +bash experimental/random/bernoulli_mask/tests/st/run.sh | ||
| 133 | +``` | ||
| 134 | + | ||
| 135 | +ST 覆盖全部输出 dtype、rank 0~8、空 Tensor、连续与非连续 view、 | ||
| 136 | +out-of-place/in-place、概率和 offset 边界、随机状态重现性,以及 | ||
| 137 | +mask/output 存储复用边界。环境变量和加载方式见 | ||
| 138 | +[`tests/st/README.md`](tests/st/README.md)。 | ||
| 139 | + | ||
| 140 | +### Kernel TTK | ||
| 141 | + | ||
| 142 | +`tests/ttk/bernoulli_mask.csv` 包含 26 个通用 Kernel 用例,覆盖全部输出 | ||
| 143 | +dtype、packed-bit 顺序、mask/tile 边界和多核大 shape; | ||
| 144 | +`tests/ttk/bernoulli_mask_alias.csv` 包含 8 个 fallback/alias 配对用例,覆盖 | ||
| 145 | +15、16、257 元素和百万元素多核 `SyncAll` 路径。 | ||
| 146 | + | ||
| 147 | +安装当前 checkout 构建的 Release 包后,从 TTK 3.0 根目录执行: | ||
| 148 | + | ||
| 149 | +```bash | ||
| 150 | +OPS_MATH_ROOT=/path/to/ops-math | ||
| 151 | + | ||
| 152 | +python -m ttk kernel \ | ||
| 153 | + -i "${OPS_MATH_ROOT}/experimental/random/bernoulli_mask/tests/ttk/bernoulli_mask.csv" \ | ||
| 154 | + -d false -b release \ | ||
| 155 | + --plugin "${OPS_MATH_ROOT}/experimental/random/bernoulli_mask/tests/assets/bernoulli_mask.py" \ | ||
| 156 | + --compare binary --seed 20260725 \ | ||
| 157 | + --device-whitelist=0 --pc=1 --proc-timeout=300 --proc-no-reuse \ | ||
| 158 | + -o bernoulli_mask_ttk_result.csv | ||
| 159 | + | ||
| 160 | +python -m ttk kernel \ | ||
| 161 | + -i "${OPS_MATH_ROOT}/experimental/random/bernoulli_mask/tests/ttk/bernoulli_mask_alias.csv" \ | ||
| 162 | + -d false -b release \ | ||
| 163 | + --plugin "${OPS_MATH_ROOT}/experimental/random/bernoulli_mask/tests/assets/bernoulli_mask.py" \ | ||
| 164 | + --compare binary --seed 20260725 \ | ||
| 165 | + --device-whitelist=0 --pc=1 --proc-timeout=300 --proc-no-reuse \ | ||
| 166 | + -o bernoulli_mask_alias_ttk_result.csv | ||
| 167 | +``` | ||
| 168 | + | ||
| 169 | +BF16 golden 依赖 `ml-dtypes==0.5.4`。如果设置了 | ||
| 170 | +`ASCEND_RT_VISIBLE_DEVICES`,`--device-whitelist` 使用映射后的逻辑 device; | ||
| 171 | +否则填写物理 device。 | ||
| 172 | + | ||
| 173 | +## 测试覆盖 | ||
| 174 | + | ||
| 175 | +随算子提交的测试覆盖以下内容: | ||
| 176 | + | ||
| 177 | +- 10 种输出 dtype 和 4 种 `prob` dtype; | ||
| 178 | +- rank 0~8、空 Tensor,以及 127/128/129、255/256/257 mask 边界; | ||
| 179 | +- 连续、转置/切片 view、非零 storage offset 和 in-place; | ||
| 180 | +- `prob=0/1`、一般概率、合法/非法 offset; | ||
| 181 | +- seed/offset 可复现性和大样本统计分布; | ||
| 182 | +- A2/A3 Kernel 通用路径与 mask/output 存储复用路径。 | ||
| @@ -0,0 +1,63 @@ | |||
| 1 | +# ---------------------------------------------------------------------------- | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | +# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | +# ---------------------------------------------------------------------------- | ||
| 10 | + | ||
| 11 | +cmake_minimum_required(VERSION 3.16) | ||
| 12 | +project(bernoulli_mask_example LANGUAGES CXX) | ||
| 13 | + | ||
| 14 | +set(CMAKE_CXX_STANDARD 17) | ||
| 15 | +set(CMAKE_CXX_STANDARD_REQUIRED ON) | ||
| 16 | + | ||
| 17 | +if(DEFINED ENV{ASCEND_HOME_PATH}) | ||
| 18 | + set(ASCEND_HOME "$ENV{ASCEND_HOME_PATH}") | ||
| 19 | +else() | ||
| 20 | + set(ASCEND_HOME "/usr/local/Ascend/cann") | ||
| 21 | +endif() | ||
| 22 | + | ||
| 23 | +if(DEFINED ENV{BERNOULLI_CUSTOM_VENDOR_ROOT}) | ||
| 24 | + set(CUSTOM_VENDOR_ROOT "$ENV{BERNOULLI_CUSTOM_VENDOR_ROOT}") | ||
| 25 | +else() | ||
| 26 | + set(CUSTOM_VENDOR_ROOT "${ASCEND_HOME}/opp/vendors/custom_math") | ||
| 27 | +endif() | ||
| 28 | + | ||
| 29 | +foreach(required | ||
| 30 | + "${ASCEND_HOME}/include/acl/acl.h" | ||
| 31 | + "${ASCEND_HOME}/lib64/libopapi.so" | ||
| 32 | + "${CUSTOM_VENDOR_ROOT}/op_api/include/aclnn_bernoulli.h" | ||
| 33 | + "${CUSTOM_VENDOR_ROOT}/op_api/lib/libcust_opapi.so") | ||
| 34 | + if(NOT EXISTS "${required}") | ||
| 35 | + message(FATAL_ERROR "Required CANN/custom package file is missing: ${required}") | ||
| 36 | + endif() | ||
| 37 | +endforeach() | ||
| 38 | + | ||
| 39 | +add_executable(test_aclnn_bernoulli test_aclnn_bernoulli.cpp) | ||
| 40 | +target_compile_options(test_aclnn_bernoulli PRIVATE -O2 -Wall -Wextra -Werror) | ||
| 41 | +target_include_directories(test_aclnn_bernoulli PRIVATE | ||
| 42 | + "${CUSTOM_VENDOR_ROOT}/op_api/include" | ||
| 43 | + "${ASCEND_HOME}/include" | ||
| 44 | + "${ASCEND_HOME}/include/aclnnop" | ||
| 45 | +) | ||
| 46 | +target_link_directories(test_aclnn_bernoulli PRIVATE | ||
| 47 | + "${CUSTOM_VENDOR_ROOT}/op_api/lib" | ||
| 48 | + "${ASCEND_HOME}/lib64" | ||
| 49 | +) | ||
| 50 | +# libcust_opapi intentionally reuses L0 operators (notably DSAGenBitMask) from | ||
| 51 | +# the installed ops-math libopapi, so --no-as-needed and this link order matter. | ||
| 52 | +target_link_options(test_aclnn_bernoulli PRIVATE -Wl,--no-as-needed) | ||
| 53 | +target_link_libraries(test_aclnn_bernoulli PRIVATE | ||
| 54 | + cust_opapi | ||
| 55 | + opapi | ||
| 56 | + ascendcl | ||
| 57 | + nnopbase | ||
| 58 | + dl | ||
| 59 | + pthread | ||
| 60 | +) | ||
| 61 | +set_target_properties(test_aclnn_bernoulli PROPERTIES | ||
| 62 | + BUILD_RPATH "${CUSTOM_VENDOR_ROOT}/op_api/lib;${ASCEND_HOME}/lib64" | ||
| 63 | +) | ||
| @@ -0,0 +1,31 @@ | |||
| 1 | +# aclnnBernoulli ACLNN 调用样例 | ||
| 2 | + | ||
| 3 | +`test_aclnn_bernoulli.cpp` 使用标准两段式 ACLNN 接口运行一个 FP32 | ||
| 4 | +`[256, 256]` 用例,并检查: | ||
| 5 | + | ||
| 6 | +- 输出只包含 `0` 和 `1`; | ||
| 7 | +- 相同 `seed=20260725`、`offset=4` 的两次调用逐元素一致; | ||
| 8 | +- `prob=0.35` 的样本均值位于二项分布均值的 6 倍标准差内; | ||
| 9 | +- 两次调用报告相同的 workspace 大小。 | ||
| 10 | + | ||
| 11 | +先构建并安装 `bernoulli_mask` 自定义算子包,再在 `ops-math` 仓库根目录运行: | ||
| 12 | + | ||
| 13 | +```bash | ||
| 14 | +bash experimental/random/bernoulli_mask/examples/run.sh | ||
| 15 | +``` | ||
| 16 | + | ||
| 17 | +也可以只编译样例: | ||
| 18 | + | ||
| 19 | +```bash | ||
| 20 | +bash experimental/random/bernoulli_mask/examples/run.sh --noexec | ||
| 21 | +``` | ||
| 22 | + | ||
| 23 | +默认使用 `${ASCEND_HOME_PATH}/opp/vendors/custom_math`;若安装在其他位置, | ||
| 24 | +通过 `BERNOULLI_CUSTOM_VENDOR_ROOT` 指定 vendor 根目录。构建输出默认放在 | ||
| 25 | +`build_out/bernoulli_mask_example/`,也可通过 `BERNOULLI_EXAMPLE_BUILD_DIR` | ||
| 26 | +覆盖。 | ||
| 27 | + | ||
| 28 | +自包含 `libcust_opapi.so` 会复用系统 `libopapi.so` 中的 | ||
| 29 | +`DSAGenBitMask`,样例 CMake 已显式链接两者。此文件是最小演示, | ||
| 30 | +完整 dtype、shape、非连续 view、in-place、边界与统计验收应使用仓库外层 | ||
| 31 | +的矩阵 runner;它不是运行本样例的前置依赖。 | ||
| @@ -0,0 +1,32 @@ | |||
| 1 | +#!/usr/bin/env bash | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | +# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + | ||
| 10 | +set -euo pipefail | ||
| 11 | + | ||
| 12 | +script_dir="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" | ||
| 13 | +ops_math_root="$(cd "${script_dir}/../../../.." && pwd)" | ||
| 14 | + | ||
| 15 | +if [[ -n "${ASCEND_HOME_PATH:-}" && -r "${ASCEND_HOME_PATH}/set_env.sh" ]]; then | ||
| 16 | + # shellcheck disable=SC1091 | ||
| 17 | + source "${ASCEND_HOME_PATH}/set_env.sh" | ||
| 18 | +elif [[ -r /usr/local/Ascend/cann/set_env.sh ]]; then | ||
| 19 | + # shellcheck disable=SC1091 | ||
| 20 | + source /usr/local/Ascend/cann/set_env.sh | ||
| 21 | +else | ||
| 22 | + echo "ERROR: CANN set_env.sh was not found; set ASCEND_HOME_PATH." >&2 | ||
| 23 | + exit 1 | ||
| 24 | +fi | ||
| 25 | + | ||
| 26 | +build_dir="${BERNOULLI_EXAMPLE_BUILD_DIR:-${ops_math_root}/build_out/bernoulli_mask_example}" | ||
| 27 | +cmake -S "${script_dir}" -B "${build_dir}" -DCMAKE_BUILD_TYPE=Release | ||
| 28 | +cmake --build "${build_dir}" -j"$(nproc)" | ||
| 29 | + | ||
| 30 | +if [[ "${1:-}" != "--noexec" ]]; then | ||
| 31 | + "${build_dir}/test_aclnn_bernoulli" | ||
| 32 | +fi | ||
| @@ -0,0 +1,163 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | +namespace { | ||
| 21 | +constexpr int32_t DEVICE_ID = 0; | ||
| 22 | +constexpr int64_t SEED = 20260725; | ||
| 23 | +constexpr int64_t OFFSET = 4; | ||
| 24 | +constexpr float PROBABILITY = 0.35F; | ||
| 25 | +constexpr int64_t ELEMENTS = 65536; | ||
| 26 | + | ||
| 27 | + | ||
| 28 | + do { \ | ||
| 29 | + const aclError checkAclStatus = (expr); \ | ||
| 30 | + if (checkAclStatus != ACL_SUCCESS) { \ | ||
| 31 | + std::fprintf(stderr, "%s failed: %d, %s\n", | ||
| 32 | + return 1; \ | ||
| 33 | + } \ | ||
| 34 | + } while (0) | ||
| 35 | + | ||
| 36 | +struct TensorResources { | ||
| 37 | + void* device = nullptr; | ||
| 38 | + aclTensor* tensor = nullptr; | ||
| 39 | + | ||
| 40 | + void Release() | ||
| 41 | + { | ||
| 42 | + if (tensor != nullptr) { | ||
| 43 | + (void)aclDestroyTensor(tensor); | ||
| 44 | + tensor = nullptr; | ||
| 45 | + } | ||
| 46 | + if (device != nullptr) { | ||
| 47 | + (void)aclrtFree(device); | ||
| 48 | + device = nullptr; | ||
| 49 | + } | ||
| 50 | + } | ||
| 51 | + | ||
| 52 | + ~TensorResources() { Release(); } | ||
| 53 | +}; | ||
| 54 | + | ||
| 55 | +int CreateContiguousFloatTensor(const std::vector<int64_t>& shape, const std::vector<float>& hostData, | ||
| 56 | + TensorResources& resources) | ||
| 57 | +{ | ||
| 58 | + const size_t bytes = std::max<size_t>(hostData.size() * sizeof(float), 1); | ||
| 59 | + CHECK_ACL(aclrtMalloc(&resources.device, bytes, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 60 | + if (!hostData.empty()) { | ||
| 61 | + CHECK_ACL(aclrtMemcpy(resources.device, bytes, hostData.data(), hostData.size() * sizeof(float), | ||
| 62 | + ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 63 | + } | ||
| 64 | + | ||
| 65 | + std::vector<int64_t> strides(shape.size(), 1); | ||
| 66 | + for (size_t reverse = shape.size(); reverse > 1; --reverse) { | ||
| 67 | + strides[reverse - 2] = strides[reverse - 1] * shape[reverse - 1]; | ||
| 68 | + } | ||
| 69 | + resources.tensor = aclCreateTensor(shape.data(), shape.size(), ACL_FLOAT, strides.data(), 0, ACL_FORMAT_ND, | ||
| 70 | + shape.data(), shape.size(), resources.device); | ||
| 71 | + if (resources.tensor == nullptr) { | ||
| 72 | + std::fprintf(stderr, "aclCreateTensor failed: %s\n", aclGetRecentErrMsg()); | ||
| 73 | + return 1; | ||
| 74 | + } | ||
| 75 | + return 0; | ||
| 76 | +} | ||
| 77 | + | ||
| 78 | +int RunOnce(aclrtStream stream, const aclTensor* self, aclTensor* out, const void* outDevice, | ||
| 79 | + const aclScalar* probability, std::vector<float>& hostOut, uint64_t& workspaceBytes) | ||
| 80 | +{ | ||
| 81 | + aclOpExecutor* executor = nullptr; | ||
| 82 | + const aclnnStatus status = aclnnBernoulliGetWorkspaceSize(self, probability, SEED, OFFSET, out, &workspaceBytes, | ||
| 83 | + &executor); | ||
| 84 | + if (status != ACL_SUCCESS) { | ||
| 85 | + std::fprintf(stderr, "aclnnBernoulliGetWorkspaceSize failed: %d, %s\n", status, aclGetRecentErrMsg()); | ||
| 86 | + return 1; | ||
| 87 | + } | ||
| 88 | + | ||
| 89 | + void* workspace = nullptr; | ||
| 90 | + if (workspaceBytes > 0) { | ||
| 91 | + CHECK_ACL(aclrtMalloc(&workspace, workspaceBytes, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 92 | + } | ||
| 93 | + const aclnnStatus launchStatus = aclnnBernoulli(workspace, workspaceBytes, executor, stream); | ||
| 94 | + if (launchStatus != ACL_SUCCESS) { | ||
| 95 | + std::fprintf(stderr, "aclnnBernoulli failed: %d, %s\n", launchStatus, aclGetRecentErrMsg()); | ||
| 96 | + if (workspace != nullptr) { | ||
| 97 | + (void)aclrtFree(workspace); | ||
| 98 | + } | ||
| 99 | + return 1; | ||
| 100 | + } | ||
| 101 | + CHECK_ACL(aclrtSynchronizeStream(stream)); | ||
| 102 | + CHECK_ACL(aclrtMemcpy(hostOut.data(), hostOut.size() * sizeof(float), outDevice, hostOut.size() * sizeof(float), | ||
| 103 | + ACL_MEMCPY_DEVICE_TO_HOST)); | ||
| 104 | + if (workspace != nullptr) { | ||
| 105 | + CHECK_ACL(aclrtFree(workspace)); | ||
| 106 | + } | ||
| 107 | + return 0; | ||
| 108 | +} | ||
| 109 | +} // namespace | ||
| 110 | + | ||
| 111 | +int main() | ||
| 112 | +{ | ||
| 113 | + CHECK_ACL(aclInit(nullptr)); | ||
| 114 | + CHECK_ACL(aclrtSetDevice(DEVICE_ID)); | ||
| 115 | + aclrtStream stream = nullptr; | ||
| 116 | + CHECK_ACL(aclrtCreateStream(&stream)); | ||
| 117 | + | ||
| 118 | + const std::vector<int64_t> shape = {256, 256}; | ||
| 119 | + const std::vector<float> input(ELEMENTS, 0.0F); | ||
| 120 | + std::vector<float> first(ELEMENTS, -1.0F); | ||
| 121 | + std::vector<float> second(ELEMENTS, -1.0F); | ||
| 122 | + TensorResources self; | ||
| 123 | + TensorResources out; | ||
| 124 | + if (CreateContiguousFloatTensor(shape, input, self) != 0 || CreateContiguousFloatTensor(shape, first, out) != 0) { | ||
| 125 | + return 1; | ||
| 126 | + } | ||
| 127 | + | ||
| 128 | + float probabilityValue = PROBABILITY; | ||
| 129 | + aclScalar* probability = aclCreateScalar(&probabilityValue, ACL_FLOAT); | ||
| 130 | + if (probability == nullptr) { | ||
| 131 | + std::fprintf(stderr, "aclCreateScalar failed: %s\n", aclGetRecentErrMsg()); | ||
| 132 | + return 1; | ||
| 133 | + } | ||
| 134 | + | ||
| 135 | + uint64_t firstWorkspace = 0; | ||
| 136 | + uint64_t secondWorkspace = 0; | ||
| 137 | + if (RunOnce(stream, self.tensor, out.tensor, out.device, probability, first, firstWorkspace) != 0 || | ||
| 138 | + RunOnce(stream, self.tensor, out.tensor, out.device, probability, second, secondWorkspace) != 0) { | ||
| 139 | + return 1; | ||
| 140 | + } | ||
| 141 | + | ||
| 142 | + const bool binary = std::all_of(first.begin(), first.end(), | ||
| 143 | + [](float value) { return value == 0.0F || value == 1.0F; }); | ||
| 144 | + const bool reproducible = first == second; | ||
| 145 | + const size_t ones = static_cast<size_t>(std::count(first.begin(), first.end(), 1.0F)); | ||
| 146 | + const double mean = static_cast<double>(ones) / static_cast<double>(first.size()); | ||
| 147 | + const double sigma = std::sqrt(static_cast<double>(PROBABILITY) * (1.0 - static_cast<double>(PROBABILITY)) / | ||
| 148 | + first.size()); | ||
| 149 | + const bool distributionOk = std::abs(mean - static_cast<double>(PROBABILITY)) <= 6.0 * sigma; | ||
| 150 | + const bool workspaceStable = firstWorkspace == secondWorkspace; | ||
| 151 | + | ||
| 152 | + std::printf("binary=%s reproducible=%s distribution=%s mean=%.8f expected=%.8f workspace=%llu bytes\n", | ||
| 153 | + binary ? "PASS" : "FAIL", reproducible ? "PASS" : "FAIL", distributionOk ? "PASS" : "FAIL", mean, | ||
| 154 | + static_cast<double>(PROBABILITY), static_cast<unsigned long long>(firstWorkspace)); | ||
| 155 | + | ||
| 156 | + (void)aclDestroyScalar(probability); | ||
| 157 | + self.Release(); | ||
| 158 | + out.Release(); | ||
| 159 | + (void)aclrtDestroyStream(stream); | ||
| 160 | + (void)aclrtResetDevice(DEVICE_ID); | ||
| 161 | + (void)aclFinalize(); | ||
| 162 | + return binary && reproducible && distributionOk && workspaceStable ? 0 : 1; | ||
| 163 | +} | ||
| @@ -0,0 +1,485 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2025 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | +/*! | ||
| 12 | + * \file aclnn_bernoulli.cpp | ||
| 13 | + * \brief | ||
| 14 | + */ | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 24 | + | ||
| 25 | + | ||
| 26 | + | ||
| 27 | + | ||
| 28 | + | ||
| 29 | + | ||
| 30 | + | ||
| 31 | + | ||
| 32 | + | ||
| 33 | + | ||
| 34 | + | ||
| 35 | + | ||
| 36 | + | ||
| 37 | + | ||
| 38 | + | ||
| 39 | + | ||
| 40 | + | ||
| 41 | +using namespace op; | ||
| 42 | + | ||
| 43 | +extern "C" { | ||
| 44 | + | ||
| 45 | + | ||
| 46 | +static const int64_t MAX_SHAPE_LENGTH = 8; | ||
| 47 | +static const int64_t VEC_BIT_NUMBER = 128; | ||
| 48 | +static const int64_t UINT8_BIT_NUMBER = 8; | ||
| 49 | + | ||
| 50 | +// 根据API定义,需要列出所能支持的所有dtype | ||
| 51 | +static const std::initializer_list<op::DataType> ASCEND910_DTYPE_SUPPORT_LIST = { | ||
| 52 | + op::DataType::DT_FLOAT, op::DataType::DT_INT32, op::DataType::DT_INT64, | ||
| 53 | + op::DataType::DT_FLOAT16, op::DataType::DT_INT16, op::DataType::DT_INT8, | ||
| 54 | + op::DataType::DT_UINT8, op::DataType::DT_DOUBLE, op::DataType::DT_BOOL}; | ||
| 55 | + | ||
| 56 | +static const std::initializer_list<op::DataType> ASCEND910_PROB_DTYPE_SUPPORT_LIST = { | ||
| 57 | + op::DataType::DT_FLOAT, op::DataType::DT_FLOAT16, op::DataType::DT_DOUBLE}; | ||
| 58 | + | ||
| 59 | +static const std::initializer_list<op::DataType> ASCEND910B_DTYPE_SUPPORT_LIST = { | ||
| 60 | + op::DataType::DT_FLOAT, op::DataType::DT_INT32, op::DataType::DT_INT64, op::DataType::DT_FLOAT16, | ||
| 61 | + op::DataType::DT_INT16, op::DataType::DT_INT8, op::DataType::DT_UINT8, op::DataType::DT_DOUBLE, | ||
| 62 | + op::DataType::DT_BOOL, op::DataType::DT_BF16}; | ||
| 63 | + | ||
| 64 | +static const std::initializer_list<op::DataType> ASCEND910B_PROB_DTYPE_SUPPORT_LIST = { | ||
| 65 | + op::DataType::DT_FLOAT, op::DataType::DT_FLOAT16, op::DataType::DT_DOUBLE, op::DataType::DT_BF16}; | ||
| 66 | + | ||
| 67 | +static const std::initializer_list<op::DataType> ARCH3510_DTYPE_SUPPORT_LIST = { | ||
| 68 | + op::DataType::DT_FLOAT, op::DataType::DT_INT32, op::DataType::DT_INT64, op::DataType::DT_FLOAT16, | ||
| 69 | + op::DataType::DT_INT16, op::DataType::DT_INT8, op::DataType::DT_UINT8, op::DataType::DT_UINT16, | ||
| 70 | + op::DataType::DT_UINT32, op::DataType::DT_DOUBLE, op::DataType::DT_BOOL, op::DataType::DT_UINT64, | ||
| 71 | + op::DataType::DT_BF16}; | ||
| 72 | + | ||
| 73 | +static const std::initializer_list<op::DataType> ARCH3510_PROB_DTYPE_SUPPORT_LIST = { | ||
| 74 | + op::DataType::DT_FLOAT, op::DataType::DT_FLOAT16, op::DataType::DT_DOUBLE, op::DataType::DT_BF16}; | ||
| 75 | + | ||
| 76 | +static const std::initializer_list<DataType> EMPTY_LIST = {}; | ||
| 77 | + | ||
| 78 | +static bool CheckNotNull(const aclTensor* self, const aclScalar* prob, const aclTensor* out) | ||
| 79 | +{ | ||
| 80 | + OP_CHECK_NULL(self, return false); | ||
| 81 | + OP_CHECK_NULL(prob, return false); | ||
| 82 | + OP_CHECK_NULL(out, return false); | ||
| 83 | + return true; | ||
| 84 | +} | ||
| 85 | + | ||
| 86 | +static const std::initializer_list<DataType>& GetOutDtypeSupportList() | ||
| 87 | +{ | ||
| 88 | + auto socVersion = GetCurrentPlatformInfo().GetSocVersion(); | ||
| 89 | + if (socVersion == SocVersion::ASCEND910) { | ||
| 90 | + return ASCEND910_DTYPE_SUPPORT_LIST; | ||
| 91 | + } else if (socVersion >= SocVersion::ASCEND910B && socVersion <= SocVersion::ASCEND910_93) { | ||
| 92 | + return ASCEND910B_DTYPE_SUPPORT_LIST; | ||
| 93 | + } else if (IsRegBase()) { | ||
| 94 | + return ARCH3510_DTYPE_SUPPORT_LIST; | ||
| 95 | + } else { | ||
| 96 | + OP_LOGW("Unknown SocVersion."); | ||
| 97 | + return EMPTY_LIST; | ||
| 98 | + } | ||
| 99 | +} | ||
| 100 | + | ||
| 101 | +static const std::initializer_list<DataType>& GetProbDtypeSupportList() | ||
| 102 | +{ | ||
| 103 | + auto socVersion = GetCurrentPlatformInfo().GetSocVersion(); | ||
| 104 | + if (socVersion == SocVersion::ASCEND910) { | ||
| 105 | + return ASCEND910_PROB_DTYPE_SUPPORT_LIST; | ||
| 106 | + } else if (socVersion >= SocVersion::ASCEND910B && socVersion <= SocVersion::ASCEND910_93) { | ||
| 107 | + return ASCEND910B_PROB_DTYPE_SUPPORT_LIST; | ||
| 108 | + } else if (IsRegBase()) { | ||
| 109 | + return ARCH3510_PROB_DTYPE_SUPPORT_LIST; | ||
| 110 | + } else { | ||
| 111 | + OP_LOGW("Unknown SocVersion."); | ||
| 112 | + return EMPTY_LIST; | ||
| 113 | + } | ||
| 114 | +} | ||
| 115 | + | ||
| 116 | +static bool IsDoubleEqual(double f1, double f2) { return std::abs(f1 - f2) <= std::numeric_limits<double>::epsilon(); } | ||
| 117 | + | ||
| 118 | +static bool CheckDtypeValid(const aclTensor* self, const aclScalar* prob, const aclTensor* out) | ||
| 119 | +{ | ||
| 120 | + // 检查self的数据类型是否在Bernoulli算子的支持列表内 | ||
| 121 | + const std::initializer_list<op::DataType> currentDtypeSupportList = GetOutDtypeSupportList(); | ||
| 122 | + const std::initializer_list<op::DataType> currentProbDtypeSupportList = GetProbDtypeSupportList(); | ||
| 123 | + | ||
| 124 | + // 检查self的数据类型是否在支持列表内 | ||
| 125 | + OP_CHECK_DTYPE_NOT_SUPPORT(self, currentDtypeSupportList, return false); | ||
| 126 | + | ||
| 127 | + // 检查prob的数据类型是否在支持列表内 | ||
| 128 | + OP_CHECK_DTYPE_NOT_SUPPORT(prob, currentProbDtypeSupportList, return false); | ||
| 129 | + | ||
| 130 | + // 检查self的数据类型是否和out的数据类型是否一致 | ||
| 131 | + OP_CHECK_DTYPE_NOT_MATCH(self, out->GetDataType(), return false); | ||
| 132 | + | ||
| 133 | + return true; | ||
| 134 | +} | ||
| 135 | + | ||
| 136 | +static bool CheckProb(const aclScalar* prob) | ||
| 137 | +{ | ||
| 138 | + const double probValue = prob->ToDouble(); | ||
| 139 | + if (!std::isfinite(probValue) || probValue > 1 || probValue < 0) { | ||
| 140 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "prob should be in range 0<=prob<=1 ."); | ||
| 141 | + return false; | ||
| 142 | + } | ||
| 143 | + | ||
| 144 | + return true; | ||
| 145 | +} | ||
| 146 | + | ||
| 147 | +static bool CheckOffset(int64_t offset) | ||
| 148 | +{ | ||
| 149 | + if (offset % 4 != 0) { | ||
| 150 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "offset must be a multiple of 4, but got %ld.", offset); | ||
| 151 | + return false; | ||
| 152 | + } | ||
| 153 | + return true; | ||
| 154 | +} | ||
| 155 | + | ||
| 156 | +static bool CheckFormat(const aclTensor* tensor, const char* parameterName) | ||
| 157 | +{ | ||
| 158 | + // Private layouts cannot be interpreted by the dense linear alias path or ViewCopy. | ||
| 159 | + if (op::IsPrivateFormat(tensor->GetStorageFormat())) { | ||
| 160 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Format only support ND、NCHW、NHWC、HWCN、NDHWC、NCDHW, %s [%s]", | ||
| 161 | + parameterName, ToString(tensor->GetStorageFormat()).GetString()); | ||
| 162 | + return false; | ||
| 163 | + } | ||
| 164 | + return true; | ||
| 165 | +} | ||
| 166 | + | ||
| 167 | +static bool CheckShape(const aclTensor* self, const aclTensor* out) | ||
| 168 | +{ | ||
| 169 | + OP_CHECK_MAX_DIM(self, MAX_SHAPE_LENGTH, return false); | ||
| 170 | + OP_CHECK_SHAPE_NOT_EQUAL(self, out, return false); | ||
| 171 | + return true; | ||
| 172 | +} | ||
| 173 | + | ||
| 174 | +static aclnnStatus CheckParams(const aclTensor* self, const aclScalar* prob, int64_t offset, const aclTensor* out) | ||
| 175 | +{ | ||
| 176 | + // 错误码等DFX方案细化后刷新,错误日志在check接口内打印 | ||
| 177 | + // 1. 检查参数是否为空指针 | ||
| 178 | + CHECK_RET(CheckNotNull(self, prob, out), ACLNN_ERR_PARAM_NULLPTR); | ||
| 179 | + | ||
| 180 | + // 2. 检查输入的数据类型是否在API支持的数据类型范围之内,需要根据api定义校验 | ||
| 181 | + CHECK_RET(CheckDtypeValid(self, prob, out), ACLNN_ERR_PARAM_INVALID); | ||
| 182 | + | ||
| 183 | + // 3. 检查输入的prob的值是否在范围之内,需要根据api定义校验 | ||
| 184 | + CHECK_RET(CheckProb(prob), ACLNN_ERR_PARAM_INVALID); | ||
| 185 | + | ||
| 186 | + // 4. 检查shape是否满足约束 | ||
| 187 | + CHECK_RET(CheckShape(self, out), ACLNN_ERR_PARAM_INVALID); | ||
| 188 | + | ||
| 189 | + // 5. 检查数据格式是否支持 | ||
| 190 | + CHECK_RET(CheckFormat(self, "self") && CheckFormat(out, "out"), ACLNN_ERR_PARAM_INVALID); | ||
| 191 | + | ||
| 192 | + // 6. 检查随机数偏移量是否满足接口约束 | ||
| 193 | + CHECK_RET(CheckOffset(offset), ACLNN_ERR_PARAM_INVALID); | ||
| 194 | + | ||
| 195 | + return ACLNN_SUCCESS; | ||
| 196 | +} | ||
| 197 | + | ||
| 198 | +static bool InferDSAOutShapeV2(const aclIntArray* shape, int64_t& packedBytes) | ||
| 199 | +{ | ||
| 200 | + uint64_t elements = 1; | ||
| 201 | + for (size_t index = 0; index < shape->Size(); index++) { | ||
| 202 | + const int64_t dim = (*shape)[index]; | ||
| 203 | + if (dim < 0) { | ||
| 204 | + return false; | ||
| 205 | + } | ||
| 206 | + if (dim == 0) { | ||
| 207 | + packedBytes = 0; | ||
| 208 | + return true; | ||
| 209 | + } | ||
| 210 | + const uint64_t unsignedDim = static_cast<uint64_t>(dim); | ||
| 211 | + if (elements > std::numeric_limits<uint64_t>::max() / unsignedDim) { | ||
| 212 | + return false; | ||
| 213 | + } | ||
| 214 | + elements *= unsignedDim; | ||
| 215 | + } | ||
| 216 | + | ||
| 217 | + const uint64_t maskBlockBytes = static_cast<uint64_t>(VEC_BIT_NUMBER) / static_cast<uint64_t>(UINT8_BIT_NUMBER); | ||
| 218 | + const uint64_t maskBlocks = (elements - 1) / static_cast<uint64_t>(VEC_BIT_NUMBER) + 1; | ||
| 219 | + if (maskBlocks > static_cast<uint64_t>(std::numeric_limits<int64_t>::max()) / maskBlockBytes) { | ||
| 220 | + return false; | ||
| 221 | + } | ||
| 222 | + packedBytes = static_cast<int64_t>(maskBlocks * maskBlockBytes); | ||
| 223 | + return true; | ||
| 224 | +} | ||
| 225 | + | ||
| 226 | +// The direct packed-mask path is used only on DAV_2201. Keep this mapping | ||
| 227 | +// scoped to the output dtypes supported by that architecture. | ||
| 228 | +static uint64_t GetDav2201OutputTypeBytes(DataType dtype) | ||
| 229 | +{ | ||
| 230 | + switch (dtype) { | ||
| 231 | + case DataType::DT_UINT8: | ||
| 232 | + case DataType::DT_INT8: | ||
| 233 | + case DataType::DT_BOOL: | ||
| 234 | + return 1; | ||
| 235 | + case DataType::DT_FLOAT16: | ||
| 236 | + case DataType::DT_BF16: | ||
| 237 | + case DataType::DT_INT16: | ||
| 238 | + return 2; | ||
| 239 | + case DataType::DT_FLOAT: | ||
| 240 | + case DataType::DT_INT32: | ||
| 241 | + return 4; | ||
| 242 | + case DataType::DT_DOUBLE: | ||
| 243 | + case DataType::DT_INT64: | ||
| 244 | + return 8; | ||
| 245 | + default: | ||
| 246 | + return 0; | ||
| 247 | + } | ||
| 248 | +} | ||
| 249 | + | ||
| 250 | +static bool HasDenseViewLayout(const aclTensor* tensor) | ||
| 251 | +{ | ||
| 252 | + if (tensor == nullptr || tensor->GetViewOffset() != 0 || tensor->GetStorageOffset() != 0) { | ||
| 253 | + return false; | ||
| 254 | + } | ||
| 255 | + const auto& viewShape = tensor->GetViewShape(); | ||
| 256 | + const auto& strides = tensor->GetViewStrides(); | ||
| 257 | + if (strides.size() != viewShape.GetDimNum()) { | ||
| 258 | + return false; | ||
| 259 | + } | ||
| 260 | + int64_t expectedStride = 1; | ||
| 261 | + for (int64_t i = static_cast<int64_t>(viewShape.GetDimNum()) - 1; i >= 0; --i) { | ||
| 262 | + if (strides[static_cast<size_t>(i)] != expectedStride) { | ||
| 263 | + return false; | ||
| 264 | + } | ||
| 265 | + const int64_t dim = static_cast<int64_t>(viewShape.GetDim(static_cast<size_t>(i))); | ||
| 266 | + if (dim < 0 || (dim != 0 && expectedStride > std::numeric_limits<int64_t>::max() / dim)) { | ||
| 267 | + return false; | ||
| 268 | + } | ||
| 269 | + expectedStride *= dim; | ||
| 270 | + } | ||
| 271 | + return true; | ||
| 272 | +} | ||
| 273 | + | ||
| 274 | +static bool CanWriteOutDirectly(const aclTensor* tensor) | ||
| 275 | +{ | ||
| 276 | + if (!HasDenseViewLayout(tensor)) { | ||
| 277 | + return false; | ||
| 278 | + } | ||
| 279 | + // Framework adapters may describe a dense N-D view over a flat 1-D | ||
| 280 | + // storage shape, so compare element counts rather than shape vectors. | ||
| 281 | + // The aliased kernel is tiled from the output storage shape and therefore | ||
| 282 | + // may write every backing element. Only use it when the dense logical view | ||
| 283 | + // covers the entire storage; a larger backing store must take ViewCopy so | ||
| 284 | + // bytes outside the logical view remain untouched. | ||
| 285 | + const int64_t storageElements = tensor->GetStorageShape().GetShapeSize(); | ||
| 286 | + const int64_t viewElements = tensor->GetViewShape().GetShapeSize(); | ||
| 287 | + return storageElements >= 0 && viewElements >= 0 && storageElements == viewElements; | ||
| 288 | +} | ||
| 289 | + | ||
| 290 | +static bool LaunchDSAGenBitMask(uint64_t count, uint64_t seed, uint64_t offset, const aclScalar* dropout, | ||
| 291 | + aclTensor* out, aclOpExecutor* executor) | ||
| 292 | +{ | ||
| 293 | + if (dropout == nullptr || out == nullptr || executor == nullptr) { | ||
| 294 | + return false; | ||
| 295 | + } | ||
| 296 | + L0_DFX(LaunchDSAGenBitMask, count, seed, offset, dropout); | ||
| 297 | + | ||
| 298 | + auto* args = op::GetOpArgContext(OP_INPUT(count, seed, offset, dropout), OP_OUTPUT(out), OP_ATTR(0)); | ||
| 299 | + if (args == nullptr) { | ||
| 300 | + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "Failed to create DSAGenBitMask argument context."); | ||
| 301 | + return false; | ||
| 302 | + } | ||
| 303 | + | ||
| 304 | + static const uint32_t dsaGenBitMaskOpType = op::GenOpTypeId("DSAGenBitMask"); | ||
| 305 | + CreatDSAKernelLauncher("DSAGenBitMask", dsaGenBitMaskOpType, DSAGenBitMaskTaskType, executor, args); | ||
| 306 | + return true; | ||
| 307 | +} | ||
| 308 | + | ||
| 309 | +static aclTensor* CreateInt64BitView(const aclTensor* tensor, aclOpExecutor* executor) | ||
| 310 | +{ | ||
| 311 | + auto bits = executor->CreateView(tensor, tensor->GetViewShape(), tensor->GetStorageShape(), | ||
| 312 | + tensor->GetViewStrides(), tensor->GetViewOffset()); | ||
| 313 | + CHECK_RET(bits != nullptr, nullptr); | ||
| 314 | + bits->SetDataType(DataType::DT_INT64); | ||
| 315 | + return bits; | ||
| 316 | +} | ||
| 317 | + | ||
| 318 | +static const aclTensor* ViewCopyWithDoubleSupport(const aclTensor* src, aclTensor* dst, aclOpExecutor* executor) | ||
| 319 | +{ | ||
| 320 | + if (src->GetDataType() != DataType::DT_DOUBLE) { | ||
| 321 | + return l0op::ViewCopy(src, dst, executor); | ||
| 322 | + } | ||
| 323 | + | ||
| 324 | + // The A2 AICore ViewCopy kernel moves 64-bit values through its INT64 | ||
| 325 | + // specialization but does not advertise DOUBLE. Reinterpret both tensors | ||
| 326 | + // as INT64 views so non-contiguous fp64 output preserves the exact bit | ||
| 327 | + // pattern without a numerical cast or an AICPU fallback. | ||
| 328 | + auto srcBits = CreateInt64BitView(src, executor); | ||
| 329 | + CHECK_RET(srcBits != nullptr, nullptr); | ||
| 330 | + auto dstBits = CreateInt64BitView(dst, executor); | ||
| 331 | + CHECK_RET(dstBits != nullptr, nullptr); | ||
| 332 | + return l0op::ViewCopy(srcBits, dstBits, executor); | ||
| 333 | +} | ||
| 334 | + | ||
| 335 | +aclnnStatus GetBernoulliByDSA(const aclTensor* input, const aclScalar* prob, int64_t seed, int64_t offset, | ||
| 336 | + aclTensor* directOut, const aclTensor*& doMaskOut, aclOpExecutor* executor) | ||
| 337 | +{ | ||
| 338 | + auto inputShape = op::ToShapeVector(input->GetViewShape()); | ||
| 339 | + auto inputSizeArray = executor->AllocIntArray(inputShape.data(), inputShape.size()); | ||
| 340 | + CHECK_RET(inputSizeArray != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 341 | + int64_t shapeSize = 0; | ||
| 342 | + CHECK_RET(InferDSAOutShapeV2(inputSizeArray, shapeSize) && shapeSize > 0, ACLNN_ERR_PARAM_INVALID); | ||
| 343 | + auto probScalar = executor->AllocScalar(static_cast<float>(1 - prob->ToDouble())); | ||
| 344 | + CHECK_RET(probScalar != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 345 | + | ||
| 346 | + const aclTensor* mask = nullptr; | ||
| 347 | + bool maskAliasesOut = false; | ||
| 348 | + const uint64_t outputTypeBytes = GetDav2201OutputTypeBytes(input->GetDataType()); | ||
| 349 | + const int64_t outputElementsSigned = input->GetViewShape().GetShapeSize(); | ||
| 350 | + CHECK_RET(outputElementsSigned > 0, ACLNN_ERR_PARAM_INVALID); | ||
| 351 | + const uint64_t outputElements = static_cast<uint64_t>(outputElementsSigned); | ||
| 352 | + const uint64_t dsaCount = static_cast<uint64_t>(shapeSize) * static_cast<uint64_t>(UINT8_BIT_NUMBER); | ||
| 353 | + if (directOut != nullptr && outputTypeBytes != 0 && | ||
| 354 | + outputElements <= std::numeric_limits<uint64_t>::max() / outputTypeBytes && | ||
| 355 | + outputElements * outputTypeBytes >= static_cast<uint64_t>(shapeSize)) { | ||
| 356 | + auto aliasedMask = executor->CreateView(directOut, op::Shape{shapeSize}, 0); | ||
| 357 | + CHECK_RET(aliasedMask != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 358 | + aliasedMask->SetDataType(op::DataType::DT_UINT8); | ||
| 359 | + | ||
| 360 | + const bool dsaLaunched = LaunchDSAGenBitMask(dsaCount, seed, offset, probScalar, aliasedMask, executor); | ||
| 361 | + CHECK_RET(dsaLaunched, ACLNN_ERR_INNER_NULLPTR); | ||
| 362 | + mask = aliasedMask; | ||
| 363 | + maskAliasesOut = true; | ||
| 364 | + } else { | ||
| 365 | + auto allocatedMask = executor->AllocTensor(op::Shape{shapeSize}, op::DataType::DT_UINT8); | ||
| 366 | + CHECK_RET(allocatedMask != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 367 | + CHECK_RET(LaunchDSAGenBitMask(dsaCount, seed, offset, probScalar, allocatedMask, executor), | ||
| 368 | + ACLNN_ERR_INNER_NULLPTR); | ||
| 369 | + mask = allocatedMask; | ||
| 370 | + CHECK_RET(mask != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 371 | + } | ||
| 372 | + if (directOut != nullptr) { | ||
| 373 | + doMaskOut = l0op::BernoulliMask(mask, directOut, maskAliasesOut, executor); | ||
| 374 | + } else { | ||
| 375 | + doMaskOut = l0op::BernoulliMask(mask, input, executor); | ||
| 376 | + } | ||
| 377 | + CHECK_RET(doMaskOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 378 | + | ||
| 379 | + return ACLNN_SUCCESS; | ||
| 380 | +} | ||
| 381 | + | ||
| 382 | +// This custom package overrides only the scalar-probability APIs named by the | ||
| 383 | +// community task. Tensor-probability APIs continue to resolve from the system | ||
| 384 | +// libopapi, preserving its full dtype/layout behavior instead of shadowing it | ||
| 385 | +// with an unrelated experimental dependency closure. | ||
| 386 | +static aclnnStatus BernoulliGetWorkspaceSizeCommon(const aclTensor* self, const aclScalar* prob, int64_t seed, | ||
| 387 | + int64_t offset, aclTensor* out, uint64_t* workspaceSize, | ||
| 388 | + aclOpExecutor** executor) | ||
| 389 | +{ | ||
| 390 | + // 固定写法,参数检查 | ||
| 391 | + auto ret = CheckParams(self, prob, offset, out); | ||
| 392 | + CHECK_RET(ret == ACLNN_SUCCESS, ret); | ||
| 393 | + | ||
| 394 | + // 固定写法,创建OpExecutor | ||
| 395 | + auto uniqueExecutor = CREATE_EXECUTOR(); | ||
| 396 | + CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); | ||
| 397 | + | ||
| 398 | + if (self->IsEmpty()) { | ||
| 399 | + // 根据实际支持情况补充 | ||
| 400 | + *workspaceSize = 0; | ||
| 401 | + uniqueExecutor.ReleaseTo(executor); | ||
| 402 | + return ACLNN_SUCCESS; | ||
| 403 | + } | ||
| 404 | + | ||
| 405 | + auto curArch = GetCurrentPlatformInfo().GetCurNpuArch(); | ||
| 406 | + const aclTensor* opOut = nullptr; | ||
| 407 | + if (curArch == NpuArch::DAV_2201) { | ||
| 408 | + // 调用DSAGenBitMask算子kernel | ||
| 409 | + const aclTensor* doMaskOut = nullptr; | ||
| 410 | + if (IsDoubleEqual(prob->ToDouble(), 0)) { | ||
| 411 | + doMaskOut = l0op::ZerosLike(self, uniqueExecutor.get()); | ||
| 412 | + } else if (IsDoubleEqual(prob->ToDouble(), 1)) { | ||
| 413 | + doMaskOut = l0op::OnesLike(self, uniqueExecutor.get()); | ||
| 414 | + } else { | ||
| 415 | + aclTensor* directOut = CanWriteOutDirectly(out) ? out : nullptr; | ||
| 416 | + auto executeResult = GetBernoulliByDSA(self, prob, seed, offset, directOut, doMaskOut, | ||
| 417 | + uniqueExecutor.get()); | ||
| 418 | + CHECK_RET(executeResult == ACLNN_SUCCESS, executeResult); | ||
| 419 | + } | ||
| 420 | + CHECK_RET(doMaskOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 421 | + opOut = doMaskOut; | ||
| 422 | + } else if (curArch == NpuArch::DAV_3510) { | ||
| 423 | + auto inputContiguous = l0op::Contiguous(self, uniqueExecutor.get()); | ||
| 424 | + CHECK_RET(inputContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 425 | + // 调用StatelessBernoulli算子kernel,ARCH3510统一转成float | ||
| 426 | + auto probScalar = uniqueExecutor->AllocScalar(static_cast<float>(prob->ToDouble())); | ||
| 427 | + CHECK_RET(probScalar != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 428 | + auto probTensor = uniqueExecutor.get()->ConvertToTensor(probScalar, probScalar->GetDataType()); | ||
| 429 | + CHECK_RET(probTensor != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 430 | + opOut = l0op::StatelessBernoulli(inputContiguous, probTensor, seed, offset, uniqueExecutor.get()); | ||
| 431 | + } else { | ||
| 432 | + auto inputContiguous = l0op::Contiguous(self, uniqueExecutor.get()); | ||
| 433 | + CHECK_RET(inputContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 434 | + auto probTensor = uniqueExecutor.get()->ConvertToTensor(prob, prob->GetDataType()); | ||
| 435 | + CHECK_RET(probTensor != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 436 | + opOut = l0op::StatelessBernoulli(inputContiguous, probTensor, seed, offset, uniqueExecutor.get()); | ||
| 437 | + } | ||
| 438 | + CHECK_RET(opOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 439 | + | ||
| 440 | + // 固定写法,将计算结果拷贝到输出out上,out可能是非连续的tensor | ||
| 441 | + if (opOut != out) { | ||
| 442 | + auto viewCopyResult = ViewCopyWithDoubleSupport(opOut, out, uniqueExecutor.get()); | ||
| 443 | + CHECK_RET(viewCopyResult != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 444 | + } | ||
| 445 | + | ||
| 446 | + // 固定写法,获取计算过程中需要使用的workspace大小 | ||
| 447 | + *workspaceSize = uniqueExecutor->GetWorkspaceSize(); | ||
| 448 | + uniqueExecutor.ReleaseTo(executor); // 需要把 uniqueExecutor持有executor转移给executor | ||
| 449 | + return ACLNN_SUCCESS; | ||
| 450 | +} | ||
| 451 | + | ||
| 452 | +aclnnStatus aclnnBernoulliGetWorkspaceSize(const aclTensor* self, const aclScalar* prob, int64_t seed, int64_t offset, | ||
| 453 | + aclTensor* out, uint64_t* workspaceSize, aclOpExecutor** executor) | ||
| 454 | +{ | ||
| 455 | + OP_CHECK_COMM_INPUT(workspaceSize, executor); | ||
| 456 | + L2_DFX_PHASE_1(aclnnBernoulli, DFX_IN(self, prob, seed, offset), DFX_OUT(out)); | ||
| 457 | + return BernoulliGetWorkspaceSizeCommon(self, prob, seed, offset, out, workspaceSize, executor); | ||
| 458 | +} | ||
| 459 | + | ||
| 460 | +aclnnStatus aclnnInplaceBernoulliGetWorkspaceSize(const aclTensor* selfRef, const aclScalar* prob, int64_t seed, | ||
| 461 | + int64_t offset, uint64_t* workspaceSize, aclOpExecutor** executor) | ||
| 462 | +{ | ||
| 463 | + OP_CHECK_COMM_INPUT(workspaceSize, executor); | ||
| 464 | + L2_DFX_PHASE_1(aclnnInplaceBernoulli, DFX_IN(selfRef, prob, seed, offset), DFX_OUT(selfRef)); | ||
| 465 | + auto out = const_cast<aclTensor*>(selfRef); | ||
| 466 | + return BernoulliGetWorkspaceSizeCommon(selfRef, prob, seed, offset, out, workspaceSize, executor); | ||
| 467 | +} | ||
| 468 | + | ||
| 469 | +aclnnStatus aclnnBernoulli(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, aclrtStream stream) | ||
| 470 | +{ | ||
| 471 | + L2_DFX_PHASE_2(aclnnBernoulli); | ||
| 472 | + // 固定写法,调用框架能力,完成计算 | ||
| 473 | + return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); | ||
| 474 | +} | ||
| 475 | + | ||
| 476 | +aclnnStatus aclnnInplaceBernoulli(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, aclrtStream stream) | ||
| 477 | +{ | ||
| 478 | + L2_DFX_PHASE_2(aclnnInplaceBernoulli); | ||
| 479 | + // 固定写法,调用框架能力,完成计算 | ||
| 480 | + return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); | ||
| 481 | +} | ||
| 482 | + | ||
| 483 | + | ||
| 484 | +} | ||
| 485 | + | ||
| @@ -0,0 +1,258 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2025 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | +extern "C" { | ||
| 18 | + | ||
| 19 | + | ||
| 20 | +/** | ||
| 21 | + * @brief aclnnBernoulli的第一段接口,根据具体的计算流程,计算workspace大小。 | ||
| 22 | + * @domain aclnn_rand | ||
| 23 | + * | ||
| 24 | + * 算子功能:从伯努利分布中提取二进制随机数 | ||
| 25 | + * 计算公式: | ||
| 26 | + * $$ out_i∼Bernoulli(self_i) $$ | ||
| 27 | + * | ||
| 28 | + * 实现说明: | ||
| 29 | + * api计算的基本路径: | ||
| 30 | + * ```mermaid | ||
| 31 | + * graph LR | ||
| 32 | + * A[(Self)] --> B([l0::Contiguous]) -->D([l0op::StatelessBernoulli]) --> I([l0op::ViewCopy]) --> J[(Out)] | ||
| 33 | + * K((p)) --> K0([ConvertToTensor]) --> D | ||
| 34 | + * E((seed)) --> D | ||
| 35 | + * F((offset)) --> D | ||
| 36 | + * ``` | ||
| 37 | + * | ||
| 38 | + * @param [in] self: npu | ||
| 39 | + * device侧的aclTensor,数据类型支持整型,浮点类型,支持非连续的Tensor,数据格式支持ND | ||
| 40 | + * @param [in] prob: host侧的aclScalar,浮点类型,需要满足$ 0≤p≤1 $ | ||
| 41 | + * @param [in] seed: host侧的aclScalar | ||
| 42 | + * @param [in] offset: host侧的aclScalar | ||
| 43 | + * @param [in] out: npu | ||
| 44 | + * device侧的aclTensor,数据类型支持整型,浮点类型 | ||
| 45 | + * @param [out] workspaceSize: 返回用户需要在npu device侧申请的workspace大小。 | ||
| 46 | + * @param [out] executor: 返回op执行器,包含算子计算流程。 | ||
| 47 | + * @return aclnnStatus: 返回状态码。 | ||
| 48 | + */ | ||
| 49 | +ACLNN_API aclnnStatus aclnnBernoulliGetWorkspaceSize(const aclTensor* self, const aclScalar* prob, int64_t seed, | ||
| 50 | + int64_t offset, aclTensor* out, uint64_t* workspaceSize, | ||
| 51 | + aclOpExecutor** executor); | ||
| 52 | + | ||
| 53 | +/** | ||
| 54 | + * @brief aclnnBernoulli的第二段接口,用于执行计算。 | ||
| 55 | + * | ||
| 56 | + * 算子功能:从伯努利分布中提取二进制随机数 | ||
| 57 | + * 计算公式: | ||
| 58 | + * $$ out_i∼Bernoulli(input_i) $$ | ||
| 59 | + * | ||
| 60 | + * 实现说明: | ||
| 61 | + * api计算的基本路径: | ||
| 62 | + * ```mermaid | ||
| 63 | + * graph LR | ||
| 64 | + * A[(Self)] --> B([l0::Contiguous]) -->D([l0op::StatelessBernoulli]) --> I([l0op::ViewCopy]) --> J[(Out)] | ||
| 65 | + * K((p)) --> K0([ConvertToTensor]) --> D | ||
| 66 | + * E((seed)) --> D | ||
| 67 | + * F((offset)) --> D | ||
| 68 | + * ``` | ||
| 69 | + * | ||
| 70 | + * @param [in] workspace: 在npu device侧申请的workspace内存起址。 | ||
| 71 | + * @param [in] workspace_size: 在npu device侧申请的workspace大小,由第一段接口aclnnAddGetWorkspaceSize获取。 | ||
| 72 | + * @param [in] executor: op执行器,包含了算子计算流程。 | ||
| 73 | + * @param [in] stream: acl stream流。 | ||
| 74 | + * @return aclnnStatus: 返回状态码。 | ||
| 75 | + */ | ||
| 76 | +ACLNN_API aclnnStatus aclnnBernoulli(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, | ||
| 77 | + aclrtStream stream); | ||
| 78 | + | ||
| 79 | +/** | ||
| 80 | + * @brief aclnnBernoulliTensor的第一段接口,根据具体的计算流程,计算workspace大小。 | ||
| 81 | + * @domain aclnn_rand | ||
| 82 | + * | ||
| 83 | + * 算子功能:从伯努利分布中提取二进制随机数 | ||
| 84 | + * 计算公式: | ||
| 85 | + * $$ out_i∼Bernoulli(self_i) $$ | ||
| 86 | + * | ||
| 87 | + * 实现说明: | ||
| 88 | + * api计算的基本路径: | ||
| 89 | + * ```mermaid | ||
| 90 | + * graph LR | ||
| 91 | + * A[(Self)] --> B([l0::Contiguous]) -->D([l0op::StatelessBernoulli]) --> I([l0op::ViewCopy]) --> J[(Out)] | ||
| 92 | + * K((p)) --> K0([ConvertToTensor]) --> D | ||
| 93 | + * E((seed)) --> D | ||
| 94 | + * F((offset)) --> D | ||
| 95 | + * ``` | ||
| 96 | + * | ||
| 97 | + * @param [in] self: npu | ||
| 98 | + * device侧的aclTensor,数据类型支持整型,浮点类型,支持非连续的Tensor,数据格式支持ND | ||
| 99 | + * @param [in] prob: npu | ||
| 100 | + * device侧的aclTensor,数据类型支持浮点类型,支持非连续的Tensor,数据格式支持ND | ||
| 101 | + * @param [in] seed: host侧的aclScalar | ||
| 102 | + * @param [in] offset: host侧的aclScalar | ||
| 103 | + * @param [in] out: npu | ||
| 104 | + * device侧的aclTensor,数据类型支持整型,浮点类型 | ||
| 105 | + * @param [out] workspaceSize: 返回用户需要在npu device侧申请的workspace大小。 | ||
| 106 | + * @param [out] executor: 返回op执行器,包含算子计算流程。 | ||
| 107 | + * @return aclnnStatus: 返回状态码。 | ||
| 108 | + */ | ||
| 109 | +ACLNN_API aclnnStatus aclnnBernoulliTensorGetWorkspaceSize(const aclTensor* self, const aclTensor* prob, int64_t seed, | ||
| 110 | + int64_t offset, aclTensor* out, uint64_t* workspaceSize, | ||
| 111 | + aclOpExecutor** executor); | ||
| 112 | + | ||
| 113 | +/** | ||
| 114 | + * @brief aclnnBernoulliTensor的第二段接口,用于执行计算。 | ||
| 115 | + * | ||
| 116 | + * 算子功能:从伯努利分布中提取二进制随机数 | ||
| 117 | + * 计算公式: | ||
| 118 | + * $$ out_i∼Bernoulli(input_i) $$ | ||
| 119 | + * | ||
| 120 | + * 实现说明: | ||
| 121 | + * api计算的基本路径: | ||
| 122 | + * ```mermaid | ||
| 123 | + * graph LR | ||
| 124 | + * A[(Self)] --> B([l0::Contiguous]) -->D([l0op::StatelessBernoulli]) --> I([l0op::ViewCopy]) --> J[(Out)] | ||
| 125 | + * K((p)) --> K0([ConvertToTensor]) --> D | ||
| 126 | + * E((seed)) --> D | ||
| 127 | + * F((offset)) --> D | ||
| 128 | + * ``` | ||
| 129 | + * | ||
| 130 | + * @param [in] workspace: 在npu device侧申请的workspace内存起址。 | ||
| 131 | + * @param [in] workspace_size: 在npu device侧申请的workspace大小,由第一段接口aclnnAddGetWorkspaceSize获取。 | ||
| 132 | + * @param [in] executor: op执行器,包含了算子计算流程。 | ||
| 133 | + * @param [in] stream: acl stream流。 | ||
| 134 | + * @return aclnnStatus: 返回状态码。 | ||
| 135 | + */ | ||
| 136 | +ACLNN_API aclnnStatus aclnnBernoulliTensor(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, | ||
| 137 | + aclrtStream stream); | ||
| 138 | + | ||
| 139 | +/** | ||
| 140 | + * @brief aclnnInplaceBernoulli的第一段接口,根据具体的计算流程,计算workspace大小。 | ||
| 141 | + * @domain aclnn_rand | ||
| 142 | + * | ||
| 143 | + * 算子功能:从伯努利分布中提取二进制随机数 | ||
| 144 | + * 计算公式: | ||
| 145 | + * $$ out_i∼Bernoulli(selfRef_i) $$ | ||
| 146 | + * | ||
| 147 | + * 实现说明: | ||
| 148 | + * api计算的基本路径: | ||
| 149 | + * ```mermaid | ||
| 150 | + * graph LR | ||
| 151 | + * A[(Self)] --> B([l0::Contiguous]) -->D([l0op::StatelessBernoulli]) --> I([l0op::ViewCopy]) --> J[(Self)] | ||
| 152 | + * K((p)) --> K0([ConvertToTensor]) --> D | ||
| 153 | + * E((seed)) --> D | ||
| 154 | + * F((offset)) --> D | ||
| 155 | + * ``` | ||
| 156 | + * | ||
| 157 | + * @param [in] selfRef: npu | ||
| 158 | + * device侧的aclTensor,数据类型支持整型,浮点类型,支持非连续的Tensor,数据格式支持ND | ||
| 159 | + * @param [in] prob: host侧的aclScalar,浮点类型,需要满足$ 0≤p≤1 $ | ||
| 160 | + * @param [in] seed: host侧的aclScalar | ||
| 161 | + * @param [in] offset: host侧的aclScalar | ||
| 162 | + * @param [out] workspaceSize: 返回用户需要在npu device侧申请的workspace大小。 | ||
| 163 | + * @param [out] executor: 返回op执行器,包含算子计算流程。 | ||
| 164 | + * @return aclnnStatus: 返回状态码。 | ||
| 165 | + */ | ||
| 166 | +ACLNN_API aclnnStatus aclnnInplaceBernoulliGetWorkspaceSize(const aclTensor* selfRef, const aclScalar* prob, | ||
| 167 | + int64_t seed, int64_t offset, uint64_t* workspaceSize, | ||
| 168 | + aclOpExecutor** executor); | ||
| 169 | + | ||
| 170 | +/** | ||
| 171 | + * @brief aclnnInplaceBernoulli的第二段接口,用于执行计算。 | ||
| 172 | + * | ||
| 173 | + * 算子功能:从伯努利分布中提取二进制随机数 | ||
| 174 | + * 计算公式: | ||
| 175 | + * $$ out_i∼Bernoulli(input_i) $$ | ||
| 176 | + * | ||
| 177 | + * 实现说明: | ||
| 178 | + * api计算的基本路径: | ||
| 179 | + * ```mermaid | ||
| 180 | + * graph LR | ||
| 181 | + * A[(Self)] --> B([l0::Contiguous]) -->D([l0op::StatelessBernoulli]) --> I([l0op::ViewCopy]) --> J[(Self)] | ||
| 182 | + * K((p)) --> K0([ConvertToTensor]) --> D | ||
| 183 | + * E((seed)) --> D | ||
| 184 | + * F((offset)) --> D | ||
| 185 | + * ``` | ||
| 186 | + * | ||
| 187 | + * @param [in] workspace: 在npu device侧申请的workspace内存起址。 | ||
| 188 | + * @param [in] workspace_size: 在npu device侧申请的workspace大小,由第一段接口aclnnAddGetWorkspaceSize获取。 | ||
| 189 | + * @param [in] executor: op执行器,包含了算子计算流程。 | ||
| 190 | + * @param [in] stream: acl stream流。 | ||
| 191 | + * @return aclnnStatus: 返回状态码。 | ||
| 192 | + */ | ||
| 193 | +ACLNN_API aclnnStatus aclnnInplaceBernoulli(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, | ||
| 194 | + aclrtStream stream); | ||
| 195 | + | ||
| 196 | +/** | ||
| 197 | + * @brief aclnnInplaceBernoulliTensor的第一段接口,根据具体的计算流程,计算workspace大小。 | ||
| 198 | + * @domain aclnn_rand | ||
| 199 | + * | ||
| 200 | + * 算子功能:从伯努利分布中提取二进制随机数 | ||
| 201 | + * 计算公式: | ||
| 202 | + * $$ out_i∼Bernoulli(selfRef_i) $$ | ||
| 203 | + * | ||
| 204 | + * 实现说明: | ||
| 205 | + * api计算的基本路径: | ||
| 206 | + * ```mermaid | ||
| 207 | + * graph LR | ||
| 208 | + * A[(Self)] --> B([l0::Contiguous]) -->D([l0op::StatelessBernoulli]) --> I([l0op::ViewCopy]) --> J[(Self)] | ||
| 209 | + * K((p)) --> K0([ConvertToTensor]) --> D | ||
| 210 | + * E((seed)) --> D | ||
| 211 | + * F((offset)) --> D | ||
| 212 | + * ``` | ||
| 213 | + * | ||
| 214 | + * @param [in] selfRef: npu | ||
| 215 | + * device侧的aclTensor,数据类型支持整型,浮点类型,支持非连续的Tensor,数据格式支持ND | ||
| 216 | + * @param [in] prob: npu | ||
| 217 | + * device侧的aclTensor,数据类型支持浮点类型,支持非连续的Tensor,数据格式支持ND | ||
| 218 | + * @param [in] seed: host侧的aclScalar | ||
| 219 | + * @param [in] offset: host侧的aclScalar | ||
| 220 | + * @param [out] workspaceSize: 返回用户需要在npu device侧申请的workspace大小。 | ||
| 221 | + * @param [out] executor: 返回op执行器,包含算子计算流程。 | ||
| 222 | + * @return aclnnStatus: 返回状态码。 | ||
| 223 | + */ | ||
| 224 | +ACLNN_API aclnnStatus aclnnInplaceBernoulliTensorGetWorkspaceSize(const aclTensor* selfRef, const aclTensor* prob, | ||
| 225 | + int64_t seed, int64_t offset, uint64_t* workspaceSize, | ||
| 226 | + aclOpExecutor** executor); | ||
| 227 | + | ||
| 228 | +/** | ||
| 229 | + * @brief aclnnInplaceBernoulliTensor的第二段接口,用于执行计算。 | ||
| 230 | + * | ||
| 231 | + * 算子功能:从伯努利分布中提取二进制随机数 | ||
| 232 | + * 计算公式: | ||
| 233 | + * $$ out_i∼Bernoulli(input_i) $$ | ||
| 234 | + * | ||
| 235 | + * 实现说明: | ||
| 236 | + * api计算的基本路径: | ||
| 237 | + * ```mermaid | ||
| 238 | + * graph LR | ||
| 239 | + * A[(Self)] --> B([l0::Contiguous]) -->D([l0op::StatelessBernoulli]) --> I([l0op::ViewCopy]) --> J[(Self)] | ||
| 240 | + * K((p)) --> K0([ConvertToTensor]) --> D | ||
| 241 | + * E((seed)) --> D | ||
| 242 | + * F((offset)) --> D | ||
| 243 | + * ``` | ||
| 244 | + * | ||
| 245 | + * @param [in] workspace: 在npu device侧申请的workspace内存起址。 | ||
| 246 | + * @param [in] workspace_size: 在npu device侧申请的workspace大小,由第一段接口aclnnAddGetWorkspaceSize获取。 | ||
| 247 | + * @param [in] executor: op执行器,包含了算子计算流程。 | ||
| 248 | + * @param [in] stream: acl stream流。 | ||
| 249 | + * @return aclnnStatus: 返回状态码。 | ||
| 250 | + */ | ||
| 251 | +ACLNN_API aclnnStatus aclnnInplaceBernoulliTensor(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, | ||
| 252 | + aclrtStream stream); | ||
| 253 | + | ||
| 254 | + | ||
| 255 | +} | ||
| 256 | + | ||
| 257 | + | ||
| 258 | + | ||
| @@ -0,0 +1,64 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | +using namespace op; | ||
| 19 | + | ||
| 20 | +namespace l0op { | ||
| 21 | +OP_TYPE_REGISTER(BernoulliMask); | ||
| 22 | + | ||
| 23 | +aclTensor* BernoulliMask(const aclTensor* mask, aclTensor* out, bool maskAliasesOut, aclOpExecutor* executor) | ||
| 24 | +{ | ||
| 25 | + if (mask == nullptr || out == nullptr || executor == nullptr) { | ||
| 26 | + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "BernoulliMask received a null argument."); | ||
| 27 | + return nullptr; | ||
| 28 | + } | ||
| 29 | + L0_DFX(BernoulliMask, mask, out); | ||
| 30 | + | ||
| 31 | + auto outputShape = op::ToShapeVector(out->GetViewShape()); | ||
| 32 | + auto outputShapeArray = executor->AllocIntArray(outputShape.data(), outputShape.size()); | ||
| 33 | + if (outputShapeArray == nullptr) { | ||
| 34 | + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "BernoulliMask failed to allocate output_shape."); | ||
| 35 | + return nullptr; | ||
| 36 | + } | ||
| 37 | + | ||
| 38 | + const int64_t maskAliasMode = maskAliasesOut ? 1 : 0; | ||
| 39 | + auto args = op::GetOpArgContext(OP_INPUT(mask), OP_OUTPUT(out), OP_ATTR(outputShapeArray, maskAliasMode)); | ||
| 40 | + if (args == nullptr) { | ||
| 41 | + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "BernoulliMask failed to create its argument context."); | ||
| 42 | + return nullptr; | ||
| 43 | + } | ||
| 44 | + auto ret = CreatAiCoreKernelLauncher("BernoulliMask", BernoulliMaskOpTypeId(), executor, args); | ||
| 45 | + if (ret != ACLNN_SUCCESS) { | ||
| 46 | + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "BernoulliMask ADD_TO_LAUNCHER_LIST_AICORE failed."); | ||
| 47 | + return nullptr; | ||
| 48 | + } | ||
| 49 | + return out; | ||
| 50 | +} | ||
| 51 | + | ||
| 52 | +const aclTensor* BernoulliMask(const aclTensor* mask, const aclTensor* like, aclOpExecutor* executor) | ||
| 53 | +{ | ||
| 54 | + if (mask == nullptr || like == nullptr || executor == nullptr) { | ||
| 55 | + return nullptr; | ||
| 56 | + } | ||
| 57 | + auto out = executor->AllocTensor(like->GetViewShape(), like->GetDataType(), op::Format::FORMAT_ND); | ||
| 58 | + if (out == nullptr) { | ||
| 59 | + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "BernoulliMask failed to allocate output."); | ||
| 60 | + return nullptr; | ||
| 61 | + } | ||
| 62 | + return BernoulliMask(mask, out, false, executor); | ||
| 63 | +} | ||
| 64 | +} // namespace l0op | ||
| @@ -0,0 +1,21 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | +namespace l0op { | ||
| 17 | +const aclTensor* BernoulliMask(const aclTensor* mask, const aclTensor* like, aclOpExecutor* executor); | ||
| 18 | +aclTensor* BernoulliMask(const aclTensor* mask, aclTensor* out, bool maskAliasesOut, aclOpExecutor* executor); | ||
| 19 | +} // namespace l0op | ||
| 20 | + | ||
| 21 | + | ||
| @@ -0,0 +1,41 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | +namespace ops { | ||
| 14 | +namespace { | ||
| 15 | + | ||
| 16 | +const std::vector<ge::DataType> kMaskDtypes(10, ge::DT_UINT8); | ||
| 17 | +const std::vector<ge::DataType> kOutputDtypes = {ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_DOUBLE, ge::DT_UINT8, | ||
| 18 | + ge::DT_INT8, ge::DT_INT16, ge::DT_INT32, ge::DT_INT64, | ||
| 19 | + ge::DT_BOOL, ge::DT_BF16}; | ||
| 20 | +const std::vector<ge::Format> kNdFormats(10, ge::FORMAT_ND); | ||
| 21 | + | ||
| 22 | +} // namespace | ||
| 23 | + | ||
| 24 | +class BernoulliMask : public OpDef { | ||
| 25 | +public: | ||
| 26 | + explicit BernoulliMask(const char* name) : OpDef(name) | ||
| 27 | + { | ||
| 28 | + this->Input("mask").ParamType(REQUIRED).DataType(kMaskDtypes).Format(kNdFormats).UnknownShapeFormat(kNdFormats); | ||
| 29 | + this->Output("out") | ||
| 30 | + .ParamType(REQUIRED) | ||
| 31 | + .DataType(kOutputDtypes) | ||
| 32 | + .Format(kNdFormats) | ||
| 33 | + .UnknownShapeFormat(kNdFormats); | ||
| 34 | + this->Attr("output_shape").AttrType(REQUIRED).ListInt(); | ||
| 35 | + this->Attr("mask_aliases_out").AttrType(OPTIONAL).Int(0); | ||
| 36 | + this->AICore().AddConfig("ascend910b").AddConfig("ascend910_93"); | ||
| 37 | + } | ||
| 38 | +}; | ||
| 39 | + | ||
| 40 | +OP_ADD(BernoulliMask); | ||
| 41 | +} // namespace ops | ||
| @@ -0,0 +1,58 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | +namespace ops { | ||
| 14 | +namespace { | ||
| 15 | +constexpr size_t OUTPUT_SHAPE_ATTR = 0; | ||
| 16 | +constexpr size_t MAX_SHAPE_LENGTH = 8; | ||
| 17 | +} // namespace | ||
| 18 | + | ||
| 19 | +static ge::graphStatus InferShapeBernoulliMask(gert::InferShapeContext* context) | ||
| 20 | +{ | ||
| 21 | + if (context == nullptr) { | ||
| 22 | + return ge::GRAPH_FAILED; | ||
| 23 | + } | ||
| 24 | + auto attrs = context->GetAttrs(); | ||
| 25 | + if (attrs == nullptr) { | ||
| 26 | + return ge::GRAPH_FAILED; | ||
| 27 | + } | ||
| 28 | + auto shapeAttr = attrs->GetListInt(OUTPUT_SHAPE_ATTR); | ||
| 29 | + if (shapeAttr == nullptr) { | ||
| 30 | + return ge::GRAPH_FAILED; | ||
| 31 | + } | ||
| 32 | + const size_t dimNum = shapeAttr->GetSize(); | ||
| 33 | + if (dimNum > MAX_SHAPE_LENGTH) { | ||
| 34 | + return ge::GRAPH_FAILED; | ||
| 35 | + } | ||
| 36 | + const auto* dims = shapeAttr->GetData(); | ||
| 37 | + if (dimNum > 0 && dims == nullptr) { | ||
| 38 | + return ge::GRAPH_FAILED; | ||
| 39 | + } | ||
| 40 | + for (size_t i = 0; i < dimNum; ++i) { | ||
| 41 | + if (dims[i] < 0) { | ||
| 42 | + return ge::GRAPH_FAILED; | ||
| 43 | + } | ||
| 44 | + } | ||
| 45 | + auto outShape = context->GetOutputShape(0); | ||
| 46 | + if (outShape == nullptr) { | ||
| 47 | + return ge::GRAPH_FAILED; | ||
| 48 | + } | ||
| 49 | + | ||
| 50 | + outShape->SetDimNum(dimNum); | ||
| 51 | + for (size_t i = 0; i < dimNum; ++i) { | ||
| 52 | + outShape->SetDim(i, dims[i]); | ||
| 53 | + } | ||
| 54 | + return ge::GRAPH_SUCCESS; | ||
| 55 | +} | ||
| 56 | + | ||
| 57 | +IMPL_OP_INFERSHAPE(BernoulliMask).InferShape(InferShapeBernoulliMask); | ||
| 58 | +} // namespace ops | ||
| @@ -0,0 +1,158 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | +namespace optiling { | ||
| 22 | +namespace { | ||
| 23 | +struct BernoulliMaskCompileInfo {}; | ||
| 24 | + | ||
| 25 | +constexpr uint64_t MASK_ALIGN_ELEMENTS = 256; | ||
| 26 | +constexpr uint64_t ASCENDC_RESERVED_UB_BYTES = 8 * 1024; | ||
| 27 | +constexpr uint64_t MAX_TILE_ELEMENTS = 16 * 1024; | ||
| 28 | +constexpr uint64_t WORK_BYTES_PER_ELEMENT = sizeof(float); | ||
| 29 | + | ||
| 30 | +uint64_t CeilDiv(uint64_t value, uint64_t divisor) { return divisor == 0 ? 0 : (value + divisor - 1) / divisor; } | ||
| 31 | + | ||
| 32 | +uint64_t AlignUp(uint64_t value, uint64_t alignment) { return CeilDiv(value, alignment) * alignment; } | ||
| 33 | + | ||
| 34 | +uint64_t AlignDown(uint64_t value, uint64_t alignment) | ||
| 35 | +{ | ||
| 36 | + return alignment == 0 ? value : value / alignment * alignment; | ||
| 37 | +} | ||
| 38 | + | ||
| 39 | +ge::graphStatus GetTilingKey(ge::DataType dtype, uint64_t& key) | ||
| 40 | +{ | ||
| 41 | + switch (dtype) { | ||
| 42 | + case ge::DT_FLOAT16: | ||
| 43 | + key = BernoulliMaskKey::FLOAT16; | ||
| 44 | + break; | ||
| 45 | + case ge::DT_FLOAT: | ||
| 46 | + key = BernoulliMaskKey::FLOAT; | ||
| 47 | + break; | ||
| 48 | + case ge::DT_DOUBLE: | ||
| 49 | + key = BernoulliMaskKey::DOUBLE; | ||
| 50 | + break; | ||
| 51 | + case ge::DT_UINT8: | ||
| 52 | + case ge::DT_BOOL: | ||
| 53 | + key = BernoulliMaskKey::UINT8_OR_BOOL; | ||
| 54 | + break; | ||
| 55 | + case ge::DT_INT8: | ||
| 56 | + key = BernoulliMaskKey::INT8; | ||
| 57 | + break; | ||
| 58 | + case ge::DT_INT16: | ||
| 59 | + key = BernoulliMaskKey::INT16; | ||
| 60 | + break; | ||
| 61 | + case ge::DT_INT32: | ||
| 62 | + key = BernoulliMaskKey::INT32; | ||
| 63 | + break; | ||
| 64 | + case ge::DT_INT64: | ||
| 65 | + key = BernoulliMaskKey::INT64; | ||
| 66 | + break; | ||
| 67 | + case ge::DT_BF16: | ||
| 68 | + key = BernoulliMaskKey::BFLOAT16; | ||
| 69 | + break; | ||
| 70 | + default: | ||
| 71 | + return ge::GRAPH_FAILED; | ||
| 72 | + } | ||
| 73 | + return ge::GRAPH_SUCCESS; | ||
| 74 | +} | ||
| 75 | + | ||
| 76 | +ge::graphStatus BernoulliTiling(gert::TilingContext* context) | ||
| 77 | +{ | ||
| 78 | + OP_CHECK_NULL_WITH_CONTEXT(context, context); | ||
| 79 | + auto outputShape = context->GetOutputShape(0); | ||
| 80 | + auto outputDesc = context->GetOutputDesc(0); | ||
| 81 | + OP_CHECK_NULL_WITH_CONTEXT(context, outputShape); | ||
| 82 | + OP_CHECK_NULL_WITH_CONTEXT(context, outputDesc); | ||
| 83 | + | ||
| 84 | + const int64_t signedElements = outputShape->GetStorageShape().GetShapeSize(); | ||
| 85 | + OP_CHECK_IF(signedElements < 0, OP_LOGE(context, "Output shape size must be non-negative."), | ||
| 86 | + return ge::GRAPH_FAILED); | ||
| 87 | + const uint64_t totalElements = static_cast<uint64_t>(signedElements); | ||
| 88 | + bool maskAliasesOut = false; | ||
| 89 | + auto attrs = context->GetAttrs(); | ||
| 90 | + OP_CHECK_NULL_WITH_CONTEXT(context, attrs); | ||
| 91 | + const int64_t* maskAliasesOutAttr = attrs->GetInt(1); | ||
| 92 | + if (maskAliasesOutAttr != nullptr) { | ||
| 93 | + maskAliasesOut = *maskAliasesOutAttr != 0; | ||
| 94 | + } | ||
| 95 | + | ||
| 96 | + uint64_t tilingKey = 0; | ||
| 97 | + OP_CHECK_IF(GetTilingKey(outputDesc->GetDataType(), tilingKey) != ge::GRAPH_SUCCESS, | ||
| 98 | + OP_LOGE(context, "Unsupported output dtype."), return ge::GRAPH_FAILED); | ||
| 99 | + | ||
| 100 | + uint32_t outputTypeBytes = 0; | ||
| 101 | + OP_CHECK_IF(!ge::TypeUtils::GetDataTypeLength(outputDesc->GetDataType(), outputTypeBytes) || outputTypeBytes == 0, | ||
| 102 | + OP_LOGE(context, "Failed to get output dtype length."), return ge::GRAPH_FAILED); | ||
| 103 | + | ||
| 104 | + auto platformInfo = context->GetPlatformInfo(); | ||
| 105 | + OP_CHECK_NULL_WITH_CONTEXT(context, platformInfo); | ||
| 106 | + auto platform = platform_ascendc::PlatformAscendC(platformInfo); | ||
| 107 | + uint64_t ubBytes = 0; | ||
| 108 | + platform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubBytes); | ||
| 109 | + uint32_t coreNum = platform.GetCoreNum(); | ||
| 110 | + OP_CHECK_IF(ubBytes <= ASCENDC_RESERVED_UB_BYTES || coreNum == 0, | ||
| 111 | + OP_LOGE(context, "Invalid platform UB/core count."), return ge::GRAPH_FAILED); | ||
| 112 | + | ||
| 113 | + // One packed-mask byte represents eight outputs. Charging one byte per element | ||
| 114 | + // is deliberately conservative and leaves room for queue alignment. | ||
| 115 | + const uint64_t bytesPerElement = WORK_BYTES_PER_ELEMENT + outputTypeBytes + 1; | ||
| 116 | + uint64_t tileElements = AlignDown((ubBytes - ASCENDC_RESERVED_UB_BYTES) / bytesPerElement, MASK_ALIGN_ELEMENTS); | ||
| 117 | + tileElements = std::min(tileElements, MAX_TILE_ELEMENTS); | ||
| 118 | + OP_CHECK_IF(tileElements == 0, OP_LOGE(context, "UB is too small for one aligned tile."), return ge::GRAPH_FAILED); | ||
| 119 | + | ||
| 120 | + uint64_t blockDim = std::min<uint64_t>(coreNum, std::max<uint64_t>(1, CeilDiv(totalElements, tileElements))); | ||
| 121 | + uint64_t elementsPerCore = AlignUp(CeilDiv(totalElements, blockDim), MASK_ALIGN_ELEMENTS); | ||
| 122 | + blockDim = std::max<uint64_t>(1, CeilDiv(totalElements, elementsPerCore)); | ||
| 123 | + | ||
| 124 | + auto tilingData = context->GetTilingData<BernoulliMaskTilingData>(); | ||
| 125 | + OP_CHECK_NULL_WITH_CONTEXT(context, tilingData); | ||
| 126 | + tilingData->totalElements = totalElements; | ||
| 127 | + tilingData->elementsPerCore = elementsPerCore; | ||
| 128 | + tilingData->tileElements = tileElements; | ||
| 129 | + tilingData->maskAliasesOut = maskAliasesOut ? 1 : 0; | ||
| 130 | + | ||
| 131 | + // The generated kernel ABI contains one workspace pointer even though the | ||
| 132 | + // kernel does not consume global workspace. Publish a zero-byte entry so | ||
| 133 | + // the standard L0 launcher preserves that ABI without charging this | ||
| 134 | + // kernel for additional device storage. | ||
| 135 | + auto workspace = context->GetWorkspaceSizes(1); | ||
| 136 | + OP_CHECK_NULL_WITH_CONTEXT(context, workspace); | ||
| 137 | + workspace[0] = 0; | ||
| 138 | + context->SetBlockDim(static_cast<uint32_t>(blockDim)); | ||
| 139 | + if (maskAliasesOut) { | ||
| 140 | + // The in-place expansion uses identical cross-core barrier counts | ||
| 141 | + // between safe output waves. | ||
| 142 | + OP_CHECK_IF(context->SetScheduleMode(1) != ge::GRAPH_SUCCESS, | ||
| 143 | + OP_LOGE(context, "Failed to enable the deterministic multi-core schedule."), | ||
| 144 | + return ge::GRAPH_FAILED); | ||
| 145 | + } | ||
| 146 | + context->SetTilingKey(tilingKey); | ||
| 147 | + return ge::GRAPH_SUCCESS; | ||
| 148 | +} | ||
| 149 | + | ||
| 150 | +ge::graphStatus BernoulliMaskTilingParse([[maybe_unused]] gert::TilingParseContext* context) | ||
| 151 | +{ | ||
| 152 | + return ge::GRAPH_SUCCESS; | ||
| 153 | +} | ||
| 154 | + | ||
| 155 | +} // namespace | ||
| 156 | + | ||
| 157 | +IMPL_OP_OPTILING(BernoulliMask).Tiling(BernoulliTiling).TilingParse<BernoulliMaskCompileInfo>(BernoulliMaskTilingParse); | ||
| 158 | +} // namespace optiling | ||
| @@ -0,0 +1,62 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | +extern "C" __global__ __aicore__ void bernoulli_mask(GM_ADDR mask, GM_ADDR out, GM_ADDR workspace, GM_ADDR tiling) | ||
| 15 | +{ | ||
| 16 | + // ProcessAliased uses the whole-vector-core SyncAll primitive. Ascend C | ||
| 17 | + // requires this mixed AIV task type even though no AIC task is launched. | ||
| 18 | + KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIV_1_0); | ||
| 19 | + REGISTER_TILING_DEFAULT(optiling::BernoulliMaskTilingData); | ||
| 20 | + GET_TILING_DATA_WITH_STRUCT(optiling::BernoulliMaskTilingData, tilingData, tiling); | ||
| 21 | + AscendC::TPipe pipe; | ||
| 22 | + | ||
| 23 | + // The AscendC precompiler requires numeric literals in TILING_KEY_IS. Keep | ||
| 24 | + // these values synchronized with bernoulli_mask_tiling_key.h. | ||
| 25 | + if (TILING_KEY_IS(1)) { | ||
| 26 | + BernoulliMask::KernelBernoulliMask<half> op; | ||
| 27 | + op.Init(mask, out, &tilingData, &pipe); | ||
| 28 | + op.Process(); | ||
| 29 | + } else if (TILING_KEY_IS(2)) { | ||
| 30 | + BernoulliMask::KernelBernoulliMask<float> op; | ||
| 31 | + op.Init(mask, out, &tilingData, &pipe); | ||
| 32 | + op.Process(); | ||
| 33 | + } else if (TILING_KEY_IS(3)) { | ||
| 34 | + BernoulliMask::KernelBernoulliMask<uint64_t, true> op; | ||
| 35 | + op.Init(mask, out, &tilingData, &pipe); | ||
| 36 | + op.Process(); | ||
| 37 | + } else if (TILING_KEY_IS(4)) { | ||
| 38 | + BernoulliMask::KernelBernoulliMask<uint8_t> op; | ||
| 39 | + op.Init(mask, out, &tilingData, &pipe); | ||
| 40 | + op.Process(); | ||
| 41 | + } else if (TILING_KEY_IS(5)) { | ||
| 42 | + BernoulliMask::KernelBernoulliMask<int8_t> op; | ||
| 43 | + op.Init(mask, out, &tilingData, &pipe); | ||
| 44 | + op.Process(); | ||
| 45 | + } else if (TILING_KEY_IS(6)) { | ||
| 46 | + BernoulliMask::KernelBernoulliMask<int16_t> op; | ||
| 47 | + op.Init(mask, out, &tilingData, &pipe); | ||
| 48 | + op.Process(); | ||
| 49 | + } else if (TILING_KEY_IS(7)) { | ||
| 50 | + BernoulliMask::KernelBernoulliMask<int32_t> op; | ||
| 51 | + op.Init(mask, out, &tilingData, &pipe); | ||
| 52 | + op.Process(); | ||
| 53 | + } else if (TILING_KEY_IS(8)) { | ||
| 54 | + BernoulliMask::KernelBernoulliMask<int64_t> op; | ||
| 55 | + op.Init(mask, out, &tilingData, &pipe); | ||
| 56 | + op.Process(); | ||
| 57 | + } else if (TILING_KEY_IS(9)) { | ||
| 58 | + BernoulliMask::KernelBernoulliMask<bfloat16_t> op; | ||
| 59 | + op.Init(mask, out, &tilingData, &pipe); | ||
| 60 | + op.Process(); | ||
| 61 | + } | ||
| 62 | +} | ||
| @@ -0,0 +1,250 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | +namespace BernoulliMask { | ||
| 19 | +using namespace AscendC; | ||
| 20 | + | ||
| 21 | +constexpr uint32_t BUFFER_NUM = 1; | ||
| 22 | +constexpr uint64_t BITS_PER_BYTE = 8; | ||
| 23 | +constexpr uint64_t GM_ALIGN_BYTES = 32; | ||
| 24 | +constexpr uint64_t MASK_ALIGN_ELEMENTS = 256; | ||
| 25 | +// A double 1.0 is stored as two little-endian fp32 words [0, 0x3ff00000]. | ||
| 26 | +// 0x3ff00000 is exactly the fp32 encoding of 1.875f, which lets the vector | ||
| 27 | +// unit build the high word without scalar bit manipulation. | ||
| 28 | +constexpr float DOUBLE_ONE_HIGH_WORD = 1.875f; | ||
| 29 | +constexpr uint64_t DOUBLE_WORK_PARTS = 4; | ||
| 30 | + | ||
| 31 | +__aicore__ inline uint64_t Min(uint64_t lhs, uint64_t rhs) { return lhs < rhs ? lhs : rhs; } | ||
| 32 | + | ||
| 33 | +__aicore__ inline uint64_t AlignUp(uint64_t value, uint64_t alignment) | ||
| 34 | +{ | ||
| 35 | + return (value + alignment - 1) / alignment * alignment; | ||
| 36 | +} | ||
| 37 | + | ||
| 38 | +template <typename T, bool IS_DOUBLE = false> | ||
| 39 | +class KernelBernoulliMask { | ||
| 40 | +public: | ||
| 41 | + __aicore__ inline KernelBernoulliMask() = default; | ||
| 42 | + | ||
| 43 | + __aicore__ inline void Init(GM_ADDR mask, GM_ADDR out, const optiling::BernoulliMaskTilingData* tilingData, | ||
| 44 | + TPipe* pipe) | ||
| 45 | + { | ||
| 46 | + pipe_ = pipe; | ||
| 47 | + totalElements_ = tilingData->totalElements; | ||
| 48 | + elementsPerCore_ = tilingData->elementsPerCore; | ||
| 49 | + tileElements_ = tilingData->tileElements; | ||
| 50 | + maskAliasesOut_ = tilingData->maskAliasesOut != 0; | ||
| 51 | + | ||
| 52 | + const uint64_t blockIdx = GetBlockIdx(); | ||
| 53 | + coreStart_ = blockIdx * elementsPerCore_; | ||
| 54 | + coreElements_ = coreStart_ < totalElements_ ? Min(elementsPerCore_, totalElements_ - coreStart_) : 0; | ||
| 55 | + | ||
| 56 | + maskGm_.SetGlobalBuffer(reinterpret_cast<__gm__ uint8_t*>(mask), | ||
| 57 | + (totalElements_ + BITS_PER_BYTE - 1) / BITS_PER_BYTE); | ||
| 58 | + outGm_.SetGlobalBuffer(reinterpret_cast<__gm__ T*>(out), totalElements_); | ||
| 59 | + | ||
| 60 | + const uint64_t maskBufferBytes = AlignUp((tileElements_ + BITS_PER_BYTE - 1) / BITS_PER_BYTE, GM_ALIGN_BYTES); | ||
| 61 | + const uint64_t outputBufferBytes = AlignUp(tileElements_ * sizeof(T), GM_ALIGN_BYTES); | ||
| 62 | + const uint64_t workBufferBytes = AlignUp(tileElements_ * sizeof(float), GM_ALIGN_BYTES); | ||
| 63 | + pipe_->InitBuffer(maskQueue_, BUFFER_NUM, maskBufferBytes); | ||
| 64 | + pipe_->InitBuffer(outQueue_, BUFFER_NUM, outputBufferBytes); | ||
| 65 | + pipe_->InitBuffer(workBuffer_, workBufferBytes); | ||
| 66 | + if constexpr (IS_DOUBLE) { | ||
| 67 | + doubleChunkElements_ = tileElements_ / DOUBLE_WORK_PARTS; | ||
| 68 | + BuildDoubleGatherOffsets(); | ||
| 69 | + } | ||
| 70 | + } | ||
| 71 | + | ||
| 72 | + __aicore__ inline void Process() | ||
| 73 | + { | ||
| 74 | + if (maskAliasesOut_) { | ||
| 75 | + ProcessAliased(); | ||
| 76 | + } else { | ||
| 77 | + ProcessRange(coreStart_, coreElements_); | ||
| 78 | + } | ||
| 79 | + } | ||
| 80 | + | ||
| 81 | +private: | ||
| 82 | + __aicore__ inline uint64_t CeilDiv(uint64_t value, uint64_t divisor) { return (value + divisor - 1) / divisor; } | ||
| 83 | + | ||
| 84 | + __aicore__ inline void ProcessRange(uint64_t rangeStart, uint64_t rangeElements) | ||
| 85 | + { | ||
| 86 | + if (rangeElements == 0) { | ||
| 87 | + return; | ||
| 88 | + } | ||
| 89 | + const uint64_t loopCount = (rangeElements + tileElements_ - 1) / tileElements_; | ||
| 90 | + for (uint64_t loop = 0; loop < loopCount; ++loop) { | ||
| 91 | + const uint64_t offset = rangeStart + loop * tileElements_; | ||
| 92 | + const uint32_t count = static_cast<uint32_t>(Min(tileElements_, rangeElements - loop * tileElements_)); | ||
| 93 | + CopyIn(offset, count); | ||
| 94 | + Compute(count); | ||
| 95 | + CopyOut(offset, count); | ||
| 96 | + } | ||
| 97 | + } | ||
| 98 | + | ||
| 99 | + __aicore__ inline void ProcessAliased() | ||
| 100 | + { | ||
| 101 | + const uint64_t blockIdx = GetBlockIdx(); | ||
| 102 | + const uint64_t blockNum = GetBlockNum(); | ||
| 103 | + uint64_t waveEnd = totalElements_; | ||
| 104 | + | ||
| 105 | + // The DSA packed mask initially occupies the beginning of out. For an | ||
| 106 | + // unconsumed prefix [0, waveEnd), output writes beginning at waveStart | ||
| 107 | + // are disjoint from every remaining mask byte when: | ||
| 108 | + // sizeof(T) * waveStart >= ceil(waveEnd / 8). | ||
| 109 | + // Process that safe suffix on all cores, synchronize, then recurse on | ||
| 110 | + // the smaller prefix. Alignment keeps every mask read byte-aligned. | ||
| 111 | + while (waveEnd > MASK_ALIGN_ELEMENTS) { | ||
| 112 | + const uint64_t firstSafe = CeilDiv(CeilDiv(waveEnd, BITS_PER_BYTE), sizeof(T)); | ||
| 113 | + const uint64_t waveStart = AlignUp(firstSafe, MASK_ALIGN_ELEMENTS); | ||
| 114 | + const uint64_t waveElements = waveEnd - waveStart; | ||
| 115 | + const uint64_t elementsPerCore = AlignUp(CeilDiv(waveElements, blockNum), MASK_ALIGN_ELEMENTS); | ||
| 116 | + const uint64_t coreStart = waveStart + blockIdx * elementsPerCore; | ||
| 117 | + const uint64_t coreElements = coreStart < waveEnd ? Min(elementsPerCore, waveEnd - coreStart) : 0; | ||
| 118 | + ProcessRange(coreStart, coreElements); | ||
| 119 | + SyncAll(); | ||
| 120 | + waveEnd = waveStart; | ||
| 121 | + } | ||
| 122 | + | ||
| 123 | + // All mask bytes for higher output indices have been consumed. Core 0 | ||
| 124 | + // copies the final small prefix into UB before its writes can overwrite | ||
| 125 | + // the aliased mask. Kernel completion provides the final global join. | ||
| 126 | + if (blockIdx == 0) { | ||
| 127 | + ProcessRange(0, waveEnd); | ||
| 128 | + } | ||
| 129 | + } | ||
| 130 | + | ||
| 131 | + __aicore__ inline void CopyIn(uint64_t offset, uint32_t count) | ||
| 132 | + { | ||
| 133 | + LocalTensor<uint8_t> maskLocal = maskQueue_.AllocTensor<uint8_t>(); | ||
| 134 | + const uint32_t maskBytes = (count + BITS_PER_BYTE - 1) / BITS_PER_BYTE; | ||
| 135 | + const uint8_t rightPadding = static_cast<uint8_t>(AlignUp(maskBytes, GM_ALIGN_BYTES) - maskBytes); | ||
| 136 | + DataCopyExtParams params{1, maskBytes, 0, 0, 0}; | ||
| 137 | + DataCopyPadExtParams<uint8_t> padParams{rightPadding != 0, 0, rightPadding, 0}; | ||
| 138 | + DataCopyPad(maskLocal, maskGm_[offset / BITS_PER_BYTE], params, padParams); | ||
| 139 | + maskQueue_.EnQue(maskLocal); | ||
| 140 | + } | ||
| 141 | + | ||
| 142 | + __aicore__ inline void SelectHalf(const LocalTensor<uint8_t>& maskLocal, const LocalTensor<half>& selected, | ||
| 143 | + uint32_t count) | ||
| 144 | + { | ||
| 145 | + Duplicate(selected, static_cast<half>(1.0f), count); | ||
| 146 | + Select(selected, maskLocal, selected, static_cast<half>(0.0f), SELMODE::VSEL_TENSOR_SCALAR_MODE, count); | ||
| 147 | + } | ||
| 148 | + | ||
| 149 | + __aicore__ inline void SelectFloat(const LocalTensor<uint8_t>& maskLocal, const LocalTensor<float>& selected, | ||
| 150 | + uint32_t count) | ||
| 151 | + { | ||
| 152 | + Duplicate(selected, 1.0f, count); | ||
| 153 | + Select(selected, maskLocal, selected, 0.0f, SELMODE::VSEL_TENSOR_SCALAR_MODE, count); | ||
| 154 | + } | ||
| 155 | + | ||
| 156 | + __aicore__ inline void BuildDoubleGatherOffsets() | ||
| 157 | + { | ||
| 158 | + // The work buffer is split into: | ||
| 159 | + // packed[0:2*S] | byteOffsets[0:2*S], S = tileElements / 4. | ||
| 160 | + // Gather offsets map the planar fp32 words | ||
| 161 | + // [low_0 ... low_S-1 | high_0 ... high_S-1] | ||
| 162 | + // to interleaved fp64 storage | ||
| 163 | + // [low_0, high_0, low_1, high_1, ...]. | ||
| 164 | + const int32_t outputWords = static_cast<int32_t>(2 * doubleChunkElements_); | ||
| 165 | + const int32_t chunkElements = static_cast<int32_t>(doubleChunkElements_); | ||
| 166 | + LocalTensor<int32_t> scratch = workBuffer_.Get<int32_t>(); | ||
| 167 | + LocalTensor<int32_t> halfIndex = scratch; | ||
| 168 | + LocalTensor<int32_t> byteOffsets = scratch[2 * doubleChunkElements_]; | ||
| 169 | + CreateVecIndex(byteOffsets, static_cast<int32_t>(0), outputWords); | ||
| 170 | + ShiftRight(halfIndex, byteOffsets, static_cast<int32_t>(1), outputWords); | ||
| 171 | + Muls(byteOffsets, byteOffsets, static_cast<int32_t>(4 * chunkElements), outputWords); | ||
| 172 | + Muls(halfIndex, halfIndex, static_cast<int32_t>(4 - 8 * chunkElements), outputWords); | ||
| 173 | + Add(byteOffsets, byteOffsets, halfIndex, outputWords); | ||
| 174 | + } | ||
| 175 | + | ||
| 176 | + __aicore__ inline void SelectDouble(const LocalTensor<uint8_t>& maskLocal, const LocalTensor<T>& outLocal, | ||
| 177 | + uint32_t count) | ||
| 178 | + { | ||
| 179 | + LocalTensor<float> work = workBuffer_.Get<float>(); | ||
| 180 | + LocalTensor<float> packed = work; | ||
| 181 | + LocalTensor<uint32_t> byteOffsets = work.template ReinterpretCast<uint32_t>()[2 * doubleChunkElements_]; | ||
| 182 | + LocalTensor<float> outputWords = outLocal.template ReinterpretCast<float>(); | ||
| 183 | + | ||
| 184 | + for (uint32_t done = 0; done < count; done += doubleChunkElements_) { | ||
| 185 | + const uint32_t chunk = static_cast<uint32_t>( | ||
| 186 | + Min(doubleChunkElements_, static_cast<uint64_t>(count - done))); | ||
| 187 | + Duplicate(packed, 0.0f, chunk); | ||
| 188 | + Duplicate(packed[doubleChunkElements_], DOUBLE_ONE_HIGH_WORD, chunk); | ||
| 189 | + Select(packed[doubleChunkElements_], maskLocal[done / BITS_PER_BYTE], packed[doubleChunkElements_], 0.0f, | ||
| 190 | + SELMODE::VSEL_TENSOR_SCALAR_MODE, chunk); | ||
| 191 | + PipeBarrier<PIPE_V>(); | ||
| 192 | + Gather(outputWords[2 * done], packed, byteOffsets, static_cast<uint32_t>(0), 2 * chunk); | ||
| 193 | + } | ||
| 194 | + } | ||
| 195 | + | ||
| 196 | + __aicore__ inline void Compute(uint32_t count) | ||
| 197 | + { | ||
| 198 | + LocalTensor<uint8_t> maskLocal = maskQueue_.DeQue<uint8_t>(); | ||
| 199 | + LocalTensor<T> outLocal = outQueue_.AllocTensor<T>(); | ||
| 200 | + | ||
| 201 | + if constexpr (std::is_same_v<T, half>) { | ||
| 202 | + SelectHalf(maskLocal, outLocal, count); | ||
| 203 | + } else if constexpr (std::is_same_v<T, float>) { | ||
| 204 | + SelectFloat(maskLocal, outLocal, count); | ||
| 205 | + } else if constexpr (std::is_same_v<T, int8_t> || std::is_same_v<T, uint8_t> || std::is_same_v<T, int16_t>) { | ||
| 206 | + LocalTensor<half> selected = workBuffer_.Get<half>(); | ||
| 207 | + SelectHalf(maskLocal, selected, count); | ||
| 208 | + Cast(outLocal, selected, RoundMode::CAST_RINT, count); | ||
| 209 | + } else if constexpr (std::is_same_v<T, bfloat16_t>) { | ||
| 210 | + LocalTensor<float> selected = workBuffer_.Get<float>(); | ||
| 211 | + SelectFloat(maskLocal, selected, count); | ||
| 212 | + Cast(outLocal, selected, RoundMode::CAST_RINT, count); | ||
| 213 | + } else if constexpr (std::is_same_v<T, int32_t> || std::is_same_v<T, int64_t>) { | ||
| 214 | + LocalTensor<float> selected = workBuffer_.Get<float>(); | ||
| 215 | + SelectFloat(maskLocal, selected, count); | ||
| 216 | + Cast(outLocal, selected, RoundMode::CAST_TRUNC, count); | ||
| 217 | + } else if constexpr (IS_DOUBLE) { | ||
| 218 | + SelectDouble(maskLocal, outLocal, count); | ||
| 219 | + } | ||
| 220 | + | ||
| 221 | + outQueue_.EnQue(outLocal); | ||
| 222 | + maskQueue_.FreeTensor(maskLocal); | ||
| 223 | + } | ||
| 224 | + | ||
| 225 | + __aicore__ inline void CopyOut(uint64_t offset, uint32_t count) | ||
| 226 | + { | ||
| 227 | + LocalTensor<T> outLocal = outQueue_.DeQue<T>(); | ||
| 228 | + DataCopyExtParams params{1, static_cast<uint32_t>(count * sizeof(T)), 0, 0, 0}; | ||
| 229 | + DataCopyPad(outGm_[offset], outLocal, params); | ||
| 230 | + outQueue_.FreeTensor(outLocal); | ||
| 231 | + } | ||
| 232 | + | ||
| 233 | +private: | ||
| 234 | + TPipe* pipe_ = nullptr; | ||
| 235 | + TQue<TPosition::VECIN, BUFFER_NUM> maskQueue_; | ||
| 236 | + TQue<TPosition::VECOUT, BUFFER_NUM> outQueue_; | ||
| 237 | + TBuf<TPosition::VECCALC> workBuffer_; | ||
| 238 | + GlobalTensor<uint8_t> maskGm_; | ||
| 239 | + GlobalTensor<T> outGm_; | ||
| 240 | + uint64_t totalElements_ = 0; | ||
| 241 | + uint64_t elementsPerCore_ = 0; | ||
| 242 | + uint64_t tileElements_ = 0; | ||
| 243 | + uint64_t coreStart_ = 0; | ||
| 244 | + uint64_t coreElements_ = 0; | ||
| 245 | + uint32_t doubleChunkElements_ = 0; | ||
| 246 | + bool maskAliasesOut_ = false; | ||
| 247 | +}; | ||
| 248 | +} // namespace BernoulliMask | ||
| 249 | + | ||
| 250 | + | ||
| @@ -0,0 +1,25 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | +namespace optiling { | ||
| 17 | +struct BernoulliMaskTilingData { | ||
| 18 | + uint64_t totalElements = 0; | ||
| 19 | + uint64_t elementsPerCore = 0; | ||
| 20 | + uint64_t tileElements = 0; | ||
| 21 | + uint64_t maskAliasesOut = 0; | ||
| 22 | +}; | ||
| 23 | +} // namespace optiling | ||
| 24 | + | ||
| 25 | + | ||
| @@ -0,0 +1,28 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | +namespace BernoulliMaskKey { | ||
| 17 | +constexpr uint64_t FLOAT16 = 1; | ||
| 18 | +constexpr uint64_t FLOAT = 2; | ||
| 19 | +constexpr uint64_t DOUBLE = 3; | ||
| 20 | +constexpr uint64_t UINT8_OR_BOOL = 4; | ||
| 21 | +constexpr uint64_t INT8 = 5; | ||
| 22 | +constexpr uint64_t INT16 = 6; | ||
| 23 | +constexpr uint64_t INT32 = 7; | ||
| 24 | +constexpr uint64_t INT64 = 8; | ||
| 25 | +constexpr uint64_t BFLOAT16 = 9; | ||
| 26 | +} // namespace BernoulliMaskKey | ||
| 27 | + | ||
| 28 | + | ||
| @@ -0,0 +1,11 @@ | |||
| 1 | +# ---------------------------------------------------------------------------- | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | +# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | +# ---------------------------------------------------------------------------- | ||
| 10 | + | ||
| 11 | +add_subdirectory(ut) | ||
| @@ -0,0 +1,91 @@ | |||
| 1 | +#!/usr/bin/env python3 | ||
| 2 | +# -*- coding: UTF-8 -*- | ||
| 3 | +# ---------------------------------------------------------------------------- | ||
| 4 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 5 | +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 6 | +# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 7 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 8 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 9 | +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 10 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 11 | +# ---------------------------------------------------------------------------- | ||
| 12 | + | ||
| 13 | +"""TTK golden and deterministic packed-mask inputs for BernoulliMask.""" | ||
| 14 | + | ||
| 15 | +import numpy as np | ||
| 16 | + | ||
| 17 | + | ||
| 18 | +__spec__ = {"bernoulli_mask": "BernoulliMaskTestSpec"} | ||
| 19 | + | ||
| 20 | + | ||
| 21 | +_DTYPE_MAP = { | ||
| 22 | + "float16": np.float16, | ||
| 23 | + "fp16": np.float16, | ||
| 24 | + "float32": np.float32, | ||
| 25 | + "fp32": np.float32, | ||
| 26 | + "double": np.float64, | ||
| 27 | + "float64": np.float64, | ||
| 28 | + "fp64": np.float64, | ||
| 29 | + "uint8": np.uint8, | ||
| 30 | + "int8": np.int8, | ||
| 31 | + "int16": np.int16, | ||
| 32 | + "int32": np.int32, | ||
| 33 | + "int64": np.int64, | ||
| 34 | + "bool": np.bool_, | ||
| 35 | +} | ||
| 36 | + | ||
| 37 | + | ||
| 38 | +def _output_dtype(kwargs): | ||
| 39 | + output_dtypes = kwargs.get("output_dtypes", ("float32",)) | ||
| 40 | + dtype_name = str(output_dtypes[0]).lower() | ||
| 41 | + if "bfloat16" in dtype_name or "bf16" in dtype_name: | ||
| 42 | + try: | ||
| 43 | + from ml_dtypes import bfloat16 | ||
| 44 | + except ImportError as exc: | ||
| 45 | + raise RuntimeError( | ||
| 46 | + "TTK bfloat16 cases require the optional ml-dtypes package" | ||
| 47 | + ) from exc | ||
| 48 | + return bfloat16 | ||
| 49 | + try: | ||
| 50 | + return _DTYPE_MAP[dtype_name] | ||
| 51 | + except KeyError as exc: | ||
| 52 | + raise ValueError( | ||
| 53 | + f"unsupported BernoulliMask output dtype: {dtype_name}" | ||
| 54 | + ) from exc | ||
| 55 | + | ||
| 56 | + | ||
| 57 | +class BernoulliMaskTestSpec: | ||
| 58 | + """Decode each packed byte from least-significant bit to most-significant bit.""" | ||
| 59 | + | ||
| 60 | + | ||
| 61 | + def customize_inputs(mask, **kwargs): | ||
| 62 | + del kwargs | ||
| 63 | + pattern = np.array( | ||
| 64 | + [0x00, 0x01, 0x02, 0x80, 0xA5, 0x5A, 0x7F, 0xFF], | ||
| 65 | + dtype=np.uint8, | ||
| 66 | + ) | ||
| 67 | + values = np.resize(pattern, mask.size).reshape(mask.shape) | ||
| 68 | + return (values,) | ||
| 69 | + | ||
| 70 | + | ||
| 71 | + def golden(mask, *, output_shape, **kwargs): | ||
| 72 | + shape = tuple(int(dim) for dim in output_shape) | ||
| 73 | + elements = int(np.prod(shape, dtype=np.int64)) if shape else 1 | ||
| 74 | + unpacked = np.unpackbits( | ||
| 75 | + np.asarray(mask, dtype=np.uint8).reshape(-1), | ||
| 76 | + bitorder="little", | ||
| 77 | + )[:elements] | ||
| 78 | + return [unpacked.reshape(shape).astype(_output_dtype(kwargs), copy=False)] | ||
| 79 | + | ||
| 80 | + tolerance = { | ||
| 81 | + "float16": {"standard": "binary_equal"}, | ||
| 82 | + "float32": {"standard": "binary_equal"}, | ||
| 83 | + "float64": {"standard": "binary_equal"}, | ||
| 84 | + "bfloat16": {"standard": "binary_equal"}, | ||
| 85 | + "uint8": {"standard": "binary_equal"}, | ||
| 86 | + "int8": {"standard": "binary_equal"}, | ||
| 87 | + "int16": {"standard": "binary_equal"}, | ||
| 88 | + "int32": {"standard": "binary_equal"}, | ||
| 89 | + "int64": {"standard": "binary_equal"}, | ||
| 90 | + "bool": {"standard": "binary_equal"}, | ||
| 91 | + } | ||
| @@ -0,0 +1,60 @@ | |||
| 1 | +# ---------------------------------------------------------------------------- | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | +# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | +# ---------------------------------------------------------------------------- | ||
| 10 | + | ||
| 11 | +cmake_minimum_required(VERSION 3.16) | ||
| 12 | +project(bernoulli_mask_st LANGUAGES CXX) | ||
| 13 | + | ||
| 14 | +set(CMAKE_CXX_STANDARD 17) | ||
| 15 | +set(CMAKE_CXX_STANDARD_REQUIRED ON) | ||
| 16 | + | ||
| 17 | +if(DEFINED ENV{ASCEND_HOME_PATH}) | ||
| 18 | + set(ASCEND_HOME "$ENV{ASCEND_HOME_PATH}") | ||
| 19 | +else() | ||
| 20 | + set(ASCEND_HOME "/usr/local/Ascend/cann") | ||
| 21 | +endif() | ||
| 22 | + | ||
| 23 | +if(DEFINED ENV{BERNOULLI_CUSTOM_VENDOR_ROOT}) | ||
| 24 | + set(CUSTOM_VENDOR_ROOT "$ENV{BERNOULLI_CUSTOM_VENDOR_ROOT}") | ||
| 25 | +else() | ||
| 26 | + set(CUSTOM_VENDOR_ROOT "${ASCEND_HOME}/opp/vendors/custom_math") | ||
| 27 | +endif() | ||
| 28 | + | ||
| 29 | +foreach(required | ||
| 30 | + "${ASCEND_HOME}/include/acl/acl.h" | ||
| 31 | + "${ASCEND_HOME}/lib64/libopapi.so" | ||
| 32 | + "${CUSTOM_VENDOR_ROOT}/op_api/lib/libcust_opapi.so") | ||
| 33 | + if(NOT EXISTS "${required}") | ||
| 34 | + message(FATAL_ERROR "Required CANN/custom package file is missing: ${required}") | ||
| 35 | + endif() | ||
| 36 | +endforeach() | ||
| 37 | + | ||
| 38 | +add_executable(test_aclnn_bernoulli_st test_aclnn_bernoulli_st.cpp) | ||
| 39 | +target_compile_options(test_aclnn_bernoulli_st PRIVATE -O2 -Wall -Wextra -Werror) | ||
| 40 | +target_include_directories(test_aclnn_bernoulli_st PRIVATE | ||
| 41 | + "${CMAKE_CURRENT_SOURCE_DIR}/../../op_api" | ||
| 42 | + "${ASCEND_HOME}/include" | ||
| 43 | + "${ASCEND_HOME}/include/aclnnop" | ||
| 44 | +) | ||
| 45 | +target_link_directories(test_aclnn_bernoulli_st PRIVATE | ||
| 46 | + "${CUSTOM_VENDOR_ROOT}/op_api/lib" | ||
| 47 | + "${ASCEND_HOME}/lib64" | ||
| 48 | +) | ||
| 49 | +target_link_options(test_aclnn_bernoulli_st PRIVATE -Wl,--no-as-needed) | ||
| 50 | +target_link_libraries(test_aclnn_bernoulli_st PRIVATE | ||
| 51 | + cust_opapi | ||
| 52 | + opapi | ||
| 53 | + ascendcl | ||
| 54 | + nnopbase | ||
| 55 | + dl | ||
| 56 | + pthread | ||
| 57 | +) | ||
| 58 | +set_target_properties(test_aclnn_bernoulli_st PROPERTIES | ||
| 59 | + BUILD_RPATH "${CUSTOM_VENDOR_ROOT}/op_api/lib;${ASCEND_HOME}/lib64" | ||
| 60 | +) | ||
| @@ -0,0 +1,52 @@ | |||
| 1 | +# aclnnBernoulli ST | ||
| 2 | + | ||
| 3 | +本目录是随算子提交的自包含 ACLNN 系统测试,不依赖竞赛工作区外层的 | ||
| 4 | +`tools/`。它覆盖: | ||
| 5 | + | ||
| 6 | +- 10 种输出 dtype 的一般概率路径; | ||
| 7 | +- 10 种 dtype 的非连续 view,分别验证 outplace 与 in-place,并检查 | ||
| 8 | + view 外 storage guard 未被改写; | ||
| 9 | +- 典型转置 view 的 outplace 与 in-place; | ||
| 10 | +- 10 种 dtype 的 257 元素 alias/`SyncAll` 边界和百万元素多核路径; | ||
| 11 | +- 1/2/4/8 字节输出的 alias/fallback workspace 阈值,以及 | ||
| 12 | + 127/128/129、255/256/257 mask 边界; | ||
| 13 | +- FP16/FP32/FP64/BF16 四种 `prob` 标量 dtype; | ||
| 14 | +- rank 0–8、空 Tensor、`prob=0/1` 和接近 0/1 的概率; | ||
| 15 | +- 相同 seed/offset 的逐字节重现性,以及 seed/offset 改变后的流变化; | ||
| 16 | +- 合法 offset `0/4/8`、非法 `offset % 4 != 0`,以及非法概率的参数拒绝。 | ||
| 17 | + | ||
| 18 | +大 shape 的 dense alias case 断言 workspace 不超过 4096 字节、不会随 | ||
| 19 | +元素数线性增长;1/2/4/8 字节阈值用例则成对断言 alias workspace 严格 | ||
| 20 | +小于 fallback。这个判据能区分旧 packed-mask workspace,同时允许 A2/A3 | ||
| 21 | +保留不同大小的 runtime 同步区。日志中的 `[METRIC]` 行会保留每个专项 | ||
| 22 | +case 的原始 workspace 字节数;本次 A2 实测 alias 为 1024、fallback | ||
| 23 | +为 1536。 | ||
| 24 | + | ||
| 25 | +先从当前 checkout 构建并安装 `bernoulli_mask` 包,然后执行: | ||
| 26 | + | ||
| 27 | +```bash | ||
| 28 | +bash experimental/random/bernoulli_mask/tests/st/run.sh | ||
| 29 | +``` | ||
| 30 | + | ||
| 31 | +只编译测试程序: | ||
| 32 | + | ||
| 33 | +```bash | ||
| 34 | +bash experimental/random/bernoulli_mask/tests/st/run.sh --noexec | ||
| 35 | +``` | ||
| 36 | + | ||
| 37 | +默认使用逻辑 device 0。若未设置可见设备映射,可通过 | ||
| 38 | +`BERNOULLI_ST_DEVICE_ID` 选择设备: | ||
| 39 | + | ||
| 40 | +```bash | ||
| 41 | +unset ASCEND_RT_VISIBLE_DEVICES | ||
| 42 | +BERNOULLI_ST_DEVICE_ID=6 \ | ||
| 43 | + bash experimental/random/bernoulli_mask/tests/st/run.sh | ||
| 44 | +``` | ||
| 45 | + | ||
| 46 | +默认从 `${ASCEND_HOME_PATH}/opp/vendors/custom_math` 加载已安装的自定义 | ||
| 47 | +opapi;非默认安装位置可设置 `BERNOULLI_CUSTOM_VENDOR_ROOT`。测试显式链接 | ||
| 48 | +`libcust_opapi.so` 与系统 `libopapi.so`,因为公共 ACLNN 实现复用了系统 | ||
| 49 | +`DSAGenBitMask` L0 符号。 | ||
| 50 | + | ||
| 51 | +Kernel 层的官方 TTK CSV 与 golden 位于相邻的 `../ttk/` 和 `../assets/` | ||
| 52 | +目录。TTK 3.0 的可复现命令见算子顶层 README。 | ||
| @@ -0,0 +1,32 @@ | |||
| 1 | +#!/usr/bin/env bash | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | +# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + | ||
| 10 | +set -euo pipefail | ||
| 11 | + | ||
| 12 | +script_dir="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" | ||
| 13 | +ops_math_root="$(cd "${script_dir}/../../../../.." && pwd)" | ||
| 14 | + | ||
| 15 | +if [[ -n "${ASCEND_HOME_PATH:-}" && -r "${ASCEND_HOME_PATH}/set_env.sh" ]]; then | ||
| 16 | + # shellcheck disable=SC1091 | ||
| 17 | + source "${ASCEND_HOME_PATH}/set_env.sh" | ||
| 18 | +elif [[ -r /usr/local/Ascend/cann/set_env.sh ]]; then | ||
| 19 | + # shellcheck disable=SC1091 | ||
| 20 | + source /usr/local/Ascend/cann/set_env.sh | ||
| 21 | +else | ||
| 22 | + echo "ERROR: CANN set_env.sh was not found; set ASCEND_HOME_PATH." >&2 | ||
| 23 | + exit 1 | ||
| 24 | +fi | ||
| 25 | + | ||
| 26 | +build_dir="${BERNOULLI_ST_BUILD_DIR:-${ops_math_root}/build_out/bernoulli_mask_st}" | ||
| 27 | +cmake -S "${script_dir}" -B "${build_dir}" -DCMAKE_BUILD_TYPE=Release | ||
| 28 | +cmake --build "${build_dir}" -j"$(nproc)" | ||
| 29 | + | ||
| 30 | +if [[ "${1:-}" != "--noexec" ]]; then | ||
| 31 | + "${build_dir}/test_aclnn_bernoulli_st" | ||
| 32 | +fi | ||
| @@ -0,0 +1,27 @@ | |||
| 1 | +testcase_name,op_name,input_shapes,input_dtypes,input_formats,output_shapes,output_dtypes,output_formats,attributes,input_data_ranges | ||
| 2 | +bernoulli_mask_fp16,bernoulli_mask,"((17,),)","('uint8',)","('ND',)","((129,),)","('float16',)","('ND',)","{'output_shape':[129]}","((0,255),)" | ||
| 3 | +bernoulli_mask_fp32,bernoulli_mask,"((17,),)","('uint8',)","('ND',)","((129,),)","('float32',)","('ND',)","{'output_shape':[129]}","((0,255),)" | ||
| 4 | +bernoulli_mask_fp64,bernoulli_mask,"((17,),)","('uint8',)","('ND',)","((129,),)","('double',)","('ND',)","{'output_shape':[129]}","((0,255),)" | ||
| 5 | +bernoulli_mask_uint8,bernoulli_mask,"((17,),)","('uint8',)","('ND',)","((129,),)","('uint8',)","('ND',)","{'output_shape':[129]}","((0,255),)" | ||
| 6 | +bernoulli_mask_int8,bernoulli_mask,"((17,),)","('uint8',)","('ND',)","((129,),)","('int8',)","('ND',)","{'output_shape':[129]}","((0,255),)" | ||
| 7 | +bernoulli_mask_int16,bernoulli_mask,"((17,),)","('uint8',)","('ND',)","((129,),)","('int16',)","('ND',)","{'output_shape':[129]}","((0,255),)" | ||
| 8 | +bernoulli_mask_int32,bernoulli_mask,"((17,),)","('uint8',)","('ND',)","((129,),)","('int32',)","('ND',)","{'output_shape':[129]}","((0,255),)" | ||
| 9 | +bernoulli_mask_int64,bernoulli_mask,"((17,),)","('uint8',)","('ND',)","((129,),)","('int64',)","('ND',)","{'output_shape':[129]}","((0,255),)" | ||
| 10 | +bernoulli_mask_bool,bernoulli_mask,"((17,),)","('uint8',)","('ND',)","((129,),)","('bool',)","('ND',)","{'output_shape':[129]}","((0,255),)" | ||
| 11 | +bernoulli_mask_bf16,bernoulli_mask,"((17,),)","('uint8',)","('ND',)","((129,),)","('bfloat16',)","('ND',)","{'output_shape':[129]}","((0,255),)" | ||
| 12 | +bernoulli_mask_n1,bernoulli_mask,"((1,),)","('uint8',)","('ND',)","((1,),)","('float32',)","('ND',)","{'output_shape':[1]}","((0,255),)" | ||
| 13 | +bernoulli_mask_n7,bernoulli_mask,"((1,),)","('uint8',)","('ND',)","((7,),)","('float32',)","('ND',)","{'output_shape':[7]}","((0,255),)" | ||
| 14 | +bernoulli_mask_n8,bernoulli_mask,"((1,),)","('uint8',)","('ND',)","((8,),)","('float32',)","('ND',)","{'output_shape':[8]}","((0,255),)" | ||
| 15 | +bernoulli_mask_n15,bernoulli_mask,"((2,),)","('uint8',)","('ND',)","((15,),)","('float32',)","('ND',)","{'output_shape':[15]}","((0,255),)" | ||
| 16 | +bernoulli_mask_n16,bernoulli_mask,"((2,),)","('uint8',)","('ND',)","((16,),)","('float32',)","('ND',)","{'output_shape':[16]}","((0,255),)" | ||
| 17 | +bernoulli_mask_n31,bernoulli_mask,"((4,),)","('uint8',)","('ND',)","((31,),)","('float32',)","('ND',)","{'output_shape':[31]}","((0,255),)" | ||
| 18 | +bernoulli_mask_n32,bernoulli_mask,"((4,),)","('uint8',)","('ND',)","((32,),)","('float32',)","('ND',)","{'output_shape':[32]}","((0,255),)" | ||
| 19 | +bernoulli_mask_n63,bernoulli_mask,"((8,),)","('uint8',)","('ND',)","((63,),)","('float32',)","('ND',)","{'output_shape':[63]}","((0,255),)" | ||
| 20 | +bernoulli_mask_n64,bernoulli_mask,"((8,),)","('uint8',)","('ND',)","((64,),)","('float32',)","('ND',)","{'output_shape':[64]}","((0,255),)" | ||
| 21 | +bernoulli_mask_n127,bernoulli_mask,"((16,),)","('uint8',)","('ND',)","((127,),)","('float32',)","('ND',)","{'output_shape':[127]}","((0,255),)" | ||
| 22 | +bernoulli_mask_n128,bernoulli_mask,"((16,),)","('uint8',)","('ND',)","((128,),)","('float32',)","('ND',)","{'output_shape':[128]}","((0,255),)" | ||
| 23 | +bernoulli_mask_n16383,bernoulli_mask,"((2048,),)","('uint8',)","('ND',)","((16383,),)","('float32',)","('ND',)","{'output_shape':[16383]}","((0,255),)" | ||
| 24 | +bernoulli_mask_n16384,bernoulli_mask,"((2048,),)","('uint8',)","('ND',)","((16384,),)","('float32',)","('ND',)","{'output_shape':[16384]}","((0,255),)" | ||
| 25 | +bernoulli_mask_n16385,bernoulli_mask,"((2049,),)","('uint8',)","('ND',)","((16385,),)","('float32',)","('ND',)","{'output_shape':[16385]}","((0,255),)" | ||
| 26 | +bernoulli_mask_rank8,bernoulli_mask,"((2,),)","('uint8',)","('ND',)","((1,2,1,2,1,2,1,2),)","('float32',)","('ND',)","{'output_shape':[1,2,1,2,1,2,1,2]}","((0,255),)" | ||
| 27 | +bernoulli_mask_large_multicore,bernoulli_mask,"((125000,),)","('uint8',)","('ND',)","((1000000,),)","('float32',)","('ND',)","{'output_shape':[1000000]}","((0,255),)" | ||
| @@ -0,0 +1,9 @@ | |||
| 1 | +testcase_name,op_name,input_shapes,input_dtypes,input_formats,output_shapes,output_dtypes,output_formats,attributes,input_data_ranges,output_inplace_indexes | ||
| 2 | +bernoulli_mask_fallback_fp32_n15,bernoulli_mask,"((2,),)","('uint8',)","('ND',)","((15,),)","('float32',)","('ND',)","{'output_shape':[15],'mask_aliases_out':0}","((0,255),)","()" | ||
| 3 | +bernoulli_mask_alias_fp32_n15,bernoulli_mask,"((60,),)","('uint8',)","('ND',)","((15,),)","('float32',)","('ND',)","{'output_shape':[15],'mask_aliases_out':1}","((0,255),)","(0,)" | ||
| 4 | +bernoulli_mask_fallback_fp32_n16,bernoulli_mask,"((2,),)","('uint8',)","('ND',)","((16,),)","('float32',)","('ND',)","{'output_shape':[16],'mask_aliases_out':0}","((0,255),)","()" | ||
| 5 | +bernoulli_mask_alias_fp32_n16,bernoulli_mask,"((64,),)","('uint8',)","('ND',)","((16,),)","('float32',)","('ND',)","{'output_shape':[16],'mask_aliases_out':1}","((0,255),)","(0,)" | ||
| 6 | +bernoulli_mask_fallback_fp32_n257,bernoulli_mask,"((33,),)","('uint8',)","('ND',)","((257,),)","('float32',)","('ND',)","{'output_shape':[257],'mask_aliases_out':0}","((0,255),)","()" | ||
| 7 | +bernoulli_mask_alias_fp32_n257,bernoulli_mask,"((1028,),)","('uint8',)","('ND',)","((257,),)","('float32',)","('ND',)","{'output_shape':[257],'mask_aliases_out':1}","((0,255),)","(0,)" | ||
| 8 | +bernoulli_mask_fallback_fp32_large_multicore,bernoulli_mask,"((125000,),)","('uint8',)","('ND',)","((1000000,),)","('float32',)","('ND',)","{'output_shape':[1000000],'mask_aliases_out':0}","((0,255),)","()" | ||
| 9 | +bernoulli_mask_alias_fp32_large_multicore,bernoulli_mask,"((4000000,),)","('uint8',)","('ND',)","((1000000,),)","('float32',)","('ND',)","{'output_shape':[1000000],'mask_aliases_out':1}","((0,255),)","(0,)" | ||
| @@ -0,0 +1,11 @@ | |||
| 1 | +# ---------------------------------------------------------------------------- | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | +# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | +# ---------------------------------------------------------------------------- | ||
| 10 | + | ||
| 11 | +add_subdirectory(op_host) | ||
| @@ -0,0 +1,14 @@ | |||
| 1 | +# ---------------------------------------------------------------------------- | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | +# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | +# ---------------------------------------------------------------------------- | ||
| 10 | + | ||
| 11 | +if(UT_TEST_ALL OR OP_HOST_UT) | ||
| 12 | + add_modules_ut_sources(UT_NAME ${OP_TILING_MODULE_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR}) | ||
| 13 | + add_modules_ut_sources(UT_NAME ${OP_INFERSHAPE_MODULE_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR}) | ||
| 14 | +endif() | ||
| @@ -0,0 +1,67 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | +namespace { | ||
| 20 | +void CheckInferShape(const std::vector<int64_t>& requestedShape) | ||
| 21 | +{ | ||
| 22 | + gert::StorageShape maskShape = {{128}, {128}}; | ||
| 23 | + gert::StorageShape outputShape = {}; | ||
| 24 | + gert::InfershapeContextPara context( | ||
| 25 | + "BernoulliMask", | ||
| 26 | + { | ||
| 27 | + {maskShape, ge::DT_UINT8, ge::FORMAT_ND}, | ||
| 28 | + }, | ||
| 29 | + { | ||
| 30 | + {outputShape, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 31 | + }, | ||
| 32 | + { | ||
| 33 | + gert::InfershapeContextPara::OpAttr("output_shape", | ||
| 34 | + Ops::Math::AnyValue::CreateFrom<std::vector<int64_t>>(requestedShape)), | ||
| 35 | + }); | ||
| 36 | + ExecuteTestCase(context, ge::GRAPH_SUCCESS, {requestedShape}); | ||
| 37 | +} | ||
| 38 | + | ||
| 39 | +void CheckInferShapeRejected(const std::vector<int64_t>& requestedShape) | ||
| 40 | +{ | ||
| 41 | + gert::StorageShape maskShape = {{128}, {128}}; | ||
| 42 | + gert::StorageShape outputShape = {}; | ||
| 43 | + gert::InfershapeContextPara context( | ||
| 44 | + "BernoulliMask", | ||
| 45 | + { | ||
| 46 | + {maskShape, ge::DT_UINT8, ge::FORMAT_ND}, | ||
| 47 | + }, | ||
| 48 | + { | ||
| 49 | + {outputShape, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 50 | + }, | ||
| 51 | + { | ||
| 52 | + gert::InfershapeContextPara::OpAttr("output_shape", | ||
| 53 | + Ops::Math::AnyValue::CreateFrom<std::vector<int64_t>>(requestedShape)), | ||
| 54 | + }); | ||
| 55 | + ExecuteTestCase(context, ge::GRAPH_FAILED); | ||
| 56 | +} | ||
| 57 | + | ||
| 58 | +TEST(BernoulliMaskInferShape, preserves_rank_eight_shape) { CheckInferShape({1, 2, 1, 2, 1, 2, 1, 2}); } | ||
| 59 | + | ||
| 60 | +TEST(BernoulliMaskInferShape, preserves_scalar_shape) { CheckInferShape({}); } | ||
| 61 | + | ||
| 62 | +TEST(BernoulliMaskInferShape, preserves_empty_dimension) { CheckInferShape({2, 0, 3}); } | ||
| 63 | + | ||
| 64 | +TEST(BernoulliMaskInferShape, rejects_rank_greater_than_eight) { CheckInferShapeRejected({1, 1, 1, 1, 1, 1, 1, 1, 1}); } | ||
| 65 | + | ||
| 66 | +TEST(BernoulliMaskInferShape, rejects_negative_dimension) { CheckInferShapeRejected({2, -1, 3}); } | ||
| 67 | +} // namespace | ||
| @@ -0,0 +1,174 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | +namespace { | ||
| 23 | +constexpr uint32_t A2_VECTOR_CORE_NUM = 40; | ||
| 24 | +constexpr uint64_t A2_UB_BYTES = 192 * 1024; | ||
| 25 | +constexpr uint64_t MAX_TILE_ELEMENTS = 16 * 1024; | ||
| 26 | +constexpr uint64_t DOUBLE_TILE_ELEMENTS = 14 * 1024; | ||
| 27 | +constexpr uint64_t MASK_ALIGN_ELEMENTS = 256; | ||
| 28 | + | ||
| 29 | +struct BernoulliMaskCompileInfoForTest {}; | ||
| 30 | + | ||
| 31 | +uint64_t CeilDiv(uint64_t value, uint64_t divisor) { return divisor == 0 ? 0 : (value + divisor - 1) / divisor; } | ||
| 32 | + | ||
| 33 | +uint64_t AlignUp(uint64_t value, uint64_t alignment) { return CeilDiv(value, alignment) * alignment; } | ||
| 34 | + | ||
| 35 | +gert::StorageShape MakeShape(const std::vector<int64_t>& dims) | ||
| 36 | +{ | ||
| 37 | + gert::StorageShape shape; | ||
| 38 | + for (int64_t dim : dims) { | ||
| 39 | + shape.MutableShape().AppendDim(dim); | ||
| 40 | + shape.MutableStorageShape().AppendDim(dim); | ||
| 41 | + } | ||
| 42 | + return shape; | ||
| 43 | +} | ||
| 44 | + | ||
| 45 | +gert::TilingContextPara MakeContext(int64_t elements, ge::DataType outputDtype, | ||
| 46 | + BernoulliMaskCompileInfoForTest* compileInfo, uint32_t coreNum = A2_VECTOR_CORE_NUM, | ||
| 47 | + uint64_t ubBytes = A2_UB_BYTES, bool maskAliasesOut = false) | ||
| 48 | +{ | ||
| 49 | + gert::StorageShape maskShape = MakeShape({elements}); | ||
| 50 | + gert::StorageShape outputShape = MakeShape({elements}); | ||
| 51 | + return gert::TilingContextPara( | ||
| 52 | + "BernoulliMask", | ||
| 53 | + { | ||
| 54 | + {maskShape, ge::DT_UINT8, ge::FORMAT_ND}, | ||
| 55 | + }, | ||
| 56 | + { | ||
| 57 | + {outputShape, outputDtype, ge::FORMAT_ND}, | ||
| 58 | + }, | ||
| 59 | + { | ||
| 60 | + gert::TilingContextPara::OpAttr("output_shape", | ||
| 61 | + Ops::Math::AnyValue::CreateFrom<std::vector<int64_t>>({elements})), | ||
| 62 | + gert::TilingContextPara::OpAttr("mask_aliases_out", | ||
| 63 | + Ops::Math::AnyValue::CreateFrom<int64_t>(maskAliasesOut ? 1 : 0)), | ||
| 64 | + }, | ||
| 65 | + compileInfo, coreNum, ubBytes); | ||
| 66 | +} | ||
| 67 | + | ||
| 68 | +void CheckSuccess(int64_t elements, ge::DataType dtype, uint64_t expectedKey, uint64_t expectedTileElements, | ||
| 69 | + uint32_t expectedBlockDim, uint64_t expectedElementsPerCore, bool maskAliasesOut = false) | ||
| 70 | +{ | ||
| 71 | + BernoulliMaskCompileInfoForTest compileInfo; | ||
| 72 | + auto context = MakeContext(elements, dtype, &compileInfo, A2_VECTOR_CORE_NUM, A2_UB_BYTES, maskAliasesOut); | ||
| 73 | + TilingInfo info; | ||
| 74 | + ASSERT_TRUE(ExecuteTiling(context, info)); | ||
| 75 | + ASSERT_EQ(info.tilingDataSize, sizeof(optiling::BernoulliMaskTilingData)); | ||
| 76 | + ASSERT_EQ(info.workspaceSizes.size(), 1U); | ||
| 77 | + EXPECT_EQ(info.workspaceSizes[0], 0U); | ||
| 78 | + EXPECT_EQ(info.tilingKey, expectedKey); | ||
| 79 | + EXPECT_EQ(info.blockNum, expectedBlockDim); | ||
| 80 | + | ||
| 81 | + const auto* data = reinterpret_cast<const optiling::BernoulliMaskTilingData*>(info.tilingData.get()); | ||
| 82 | + ASSERT_NE(data, nullptr); | ||
| 83 | + EXPECT_EQ(data->totalElements, static_cast<uint64_t>(elements)); | ||
| 84 | + EXPECT_EQ(data->elementsPerCore, expectedElementsPerCore); | ||
| 85 | + EXPECT_EQ(data->tileElements, expectedTileElements); | ||
| 86 | + EXPECT_EQ(data->maskAliasesOut, maskAliasesOut ? 1U : 0U); | ||
| 87 | +} | ||
| 88 | + | ||
| 89 | +struct DtypeCase { | ||
| 90 | + const char* name; | ||
| 91 | + ge::DataType dtype; | ||
| 92 | + uint64_t tilingKey; | ||
| 93 | + uint64_t tileElements; | ||
| 94 | +}; | ||
| 95 | + | ||
| 96 | +class BernoulliMaskDtypeTiling : public testing::TestWithParam<DtypeCase> {}; | ||
| 97 | + | ||
| 98 | +TEST_P(BernoulliMaskDtypeTiling, maps_supported_dtype_to_kernel_key) | ||
| 99 | +{ | ||
| 100 | + const auto& param = GetParam(); | ||
| 101 | + CheckSuccess(129, param.dtype, param.tilingKey, param.tileElements, 1, MASK_ALIGN_ELEMENTS); | ||
| 102 | +} | ||
| 103 | + | ||
| 104 | +INSTANTIATE_TEST_SUITE_P( | ||
| 105 | + Ascend910B, BernoulliMaskDtypeTiling, | ||
| 106 | + testing::Values(DtypeCase{"fp16", ge::DT_FLOAT16, BernoulliMaskKey::FLOAT16, MAX_TILE_ELEMENTS}, | ||
| 107 | + DtypeCase{"fp32", ge::DT_FLOAT, BernoulliMaskKey::FLOAT, MAX_TILE_ELEMENTS}, | ||
| 108 | + DtypeCase{"fp64", ge::DT_DOUBLE, BernoulliMaskKey::DOUBLE, DOUBLE_TILE_ELEMENTS}, | ||
| 109 | + DtypeCase{"uint8", ge::DT_UINT8, BernoulliMaskKey::UINT8_OR_BOOL, MAX_TILE_ELEMENTS}, | ||
| 110 | + DtypeCase{"int8", ge::DT_INT8, BernoulliMaskKey::INT8, MAX_TILE_ELEMENTS}, | ||
| 111 | + DtypeCase{"int16", ge::DT_INT16, BernoulliMaskKey::INT16, MAX_TILE_ELEMENTS}, | ||
| 112 | + DtypeCase{"int32", ge::DT_INT32, BernoulliMaskKey::INT32, MAX_TILE_ELEMENTS}, | ||
| 113 | + DtypeCase{"int64", ge::DT_INT64, BernoulliMaskKey::INT64, DOUBLE_TILE_ELEMENTS}, | ||
| 114 | + DtypeCase{"bool", ge::DT_BOOL, BernoulliMaskKey::UINT8_OR_BOOL, MAX_TILE_ELEMENTS}, | ||
| 115 | + DtypeCase{"bf16", ge::DT_BF16, BernoulliMaskKey::BFLOAT16, MAX_TILE_ELEMENTS}), | ||
| 116 | + [](const testing::TestParamInfo<DtypeCase>& info) { return std::string(info.param.name); }); | ||
| 117 | + | ||
| 118 | +class BernoulliMaskBoundaryTiling : public testing::TestWithParam<int64_t> {}; | ||
| 119 | + | ||
| 120 | +TEST_P(BernoulliMaskBoundaryTiling, handles_packed_mask_and_alignment_boundaries) | ||
| 121 | +{ | ||
| 122 | + const int64_t elements = GetParam(); | ||
| 123 | + const uint32_t blockDim = elements > static_cast<int64_t>(MAX_TILE_ELEMENTS) ? 2U : 1U; | ||
| 124 | + const uint64_t elementsPerCore = elements == 0 ? 0 : | ||
| 125 | + AlignUp(CeilDiv(static_cast<uint64_t>(elements), blockDim), | ||
| 126 | + MASK_ALIGN_ELEMENTS); | ||
| 127 | + CheckSuccess(elements, ge::DT_FLOAT, BernoulliMaskKey::FLOAT, MAX_TILE_ELEMENTS, blockDim, elementsPerCore); | ||
| 128 | +} | ||
| 129 | + | ||
| 130 | +INSTANTIATE_TEST_SUITE_P(BitAndTileEdges, BernoulliMaskBoundaryTiling, | ||
| 131 | + testing::Values(0, 1, 7, 8, 15, 16, 31, 32, 63, 64, 127, 128, 129, 16383, 16384, 16385)); | ||
| 132 | + | ||
| 133 | +TEST(BernoulliMaskTiling, fp64_tile_boundary_uses_two_cores) | ||
| 134 | +{ | ||
| 135 | + CheckSuccess(DOUBLE_TILE_ELEMENTS + 1, ge::DT_DOUBLE, BernoulliMaskKey::DOUBLE, DOUBLE_TILE_ELEMENTS, 2, | ||
| 136 | + AlignUp(CeilDiv(DOUBLE_TILE_ELEMENTS + 1, 2), MASK_ALIGN_ELEMENTS)); | ||
| 137 | +} | ||
| 138 | + | ||
| 139 | +TEST(BernoulliMaskTiling, large_shape_uses_all_a2_vector_cores) | ||
| 140 | +{ | ||
| 141 | + constexpr uint64_t elements = 1000000; | ||
| 142 | + constexpr uint64_t elementsPerCore = ((elements / A2_VECTOR_CORE_NUM + MASK_ALIGN_ELEMENTS - 1) / | ||
| 143 | + MASK_ALIGN_ELEMENTS) * | ||
| 144 | + MASK_ALIGN_ELEMENTS; | ||
| 145 | + CheckSuccess(elements, ge::DT_FLOAT, BernoulliMaskKey::FLOAT, MAX_TILE_ELEMENTS, A2_VECTOR_CORE_NUM, | ||
| 146 | + elementsPerCore); | ||
| 147 | +} | ||
| 148 | + | ||
| 149 | +TEST(BernoulliMaskTiling, records_mask_output_alias_mode) | ||
| 150 | +{ | ||
| 151 | + constexpr uint64_t elements = 1000000; | ||
| 152 | + constexpr uint64_t elementsPerCore = ((elements / A2_VECTOR_CORE_NUM + MASK_ALIGN_ELEMENTS - 1) / | ||
| 153 | + MASK_ALIGN_ELEMENTS) * | ||
| 154 | + MASK_ALIGN_ELEMENTS; | ||
| 155 | + CheckSuccess(elements, ge::DT_FLOAT, BernoulliMaskKey::FLOAT, MAX_TILE_ELEMENTS, A2_VECTOR_CORE_NUM, | ||
| 156 | + elementsPerCore, true); | ||
| 157 | +} | ||
| 158 | + | ||
| 159 | +TEST(BernoulliMaskTiling, rejects_unsupported_output_dtype) | ||
| 160 | +{ | ||
| 161 | + BernoulliMaskCompileInfoForTest compileInfo; | ||
| 162 | + auto context = MakeContext(128, ge::DT_COMPLEX64, &compileInfo); | ||
| 163 | + TilingInfo info; | ||
| 164 | + EXPECT_FALSE(ExecuteTiling(context, info)); | ||
| 165 | +} | ||
| 166 | + | ||
| 167 | +TEST(BernoulliMaskTiling, rejects_platform_with_insufficient_ub) | ||
| 168 | +{ | ||
| 169 | + BernoulliMaskCompileInfoForTest compileInfo; | ||
| 170 | + auto context = MakeContext(128, ge::DT_FLOAT, &compileInfo, A2_VECTOR_CORE_NUM, 8 * 1024); | ||
| 171 | + TilingInfo info; | ||
| 172 | + EXPECT_FALSE(ExecuteTiling(context, info)); | ||
| 173 | +} | ||
| 174 | +} // namespace | ||