已合并
[社区任务] aclnnBernoulli低内存实现 #4248
[社区任务] aclnnBernoulli低内存实现 #4248
已合并
hzw_rpap创建于 7月26日
共 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+#include <algorithm>
12+#include <cmath>
13+#include <cstdint>
14+#include <cstdio>
15+#include <vector>
16+ 
17+#include "acl/acl.h"
18+#include "aclnn_bernoulli.h"
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+#define CHECK_ACL(expr) \
28+ do { \
29+ const aclError checkAclStatus = (expr); \
30+ if (checkAclStatus != ACL_SUCCESS) { \
31+ std::fprintf(stderr, "%s failed: %d, %s\n", #expr, checkAclStatus, aclGetRecentErrMsg()); \
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+#include <cmath>
17+#include <cstdint>
18+#include <initializer_list>
19+#include <limits>
20+ 
21+#include "aclnn_bernoulli.h"
22+#include "aclnn_kernels/cast.h"
23+#include "aclnn_kernels/contiguous.h"
24+#include "bernoulli_mask.h"
25+#include "random/stateless_bernoulli/op_api/stateless_bernoulli.h"
26+#include "math/zero_op/op_api/zero_op.h"
27+#include "math/ones_like/op_api/ones_like.h"
28+#include "aclnn/aclnn_base.h"
29+#include "aclnn_kernels/common/op_error_check.h"
30+#include "op_api/aclnn_check.h"
31+#include "opdev/common_types.h"
32+#include "opdev/data_type_utils.h"
33+#include "opdev/format_utils.h"
34+#include "opdev/make_op_executor.h"
35+#include "opdev/op_dfx.h"
36+#include "opdev/op_executor.h"
37+#include "opdev/op_log.h"
38+#include "opdev/platform.h"
39+#include "opdev/tensor_view_utils.h"
40+ 
41+using namespace op;
42+#ifdef __cplusplus
43+extern "C" {
44+#endif
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+#ifdef __cplusplus
484+}
485+#endif
@@ -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+#ifndef OP_API_INC_BERNOULLI_H_
11+#define OP_API_INC_BERNOULLI_H_
12+ 
13+#include "aclnn/aclnn_base.h"
14+#include "aclnn_util.h"
15+ 
16+#ifdef __cplusplus
17+extern "C" {
18+#endif
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+#ifdef __cplusplus
255+}
256+#endif
257+ 
258+#endif // OP_API_INC_BERNOULLI_H_
@@ -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+#include "bernoulli_mask.h"
12+ 
13+#include "opdev/make_op_executor.h"
14+#include "opdev/op_dfx.h"
15+#include "opdev/op_log.h"
16+#include "opdev/shape_utils.h"
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+#ifndef OP_API_INC_LEVEL0_OP_BERNOULLI_MASK_H
12+#define OP_API_INC_LEVEL0_OP_BERNOULLI_MASK_H
13+ 
14+#include "opdev/op_executor.h"
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+#endif // OP_API_INC_LEVEL0_OP_BERNOULLI_MASK_H
@@ -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+#include "register/op_def_registry.h"
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+#include "register/op_impl_registry.h"
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+#include <algorithm>
12+#include <cstdint>
13+ 
14+#include <graph/utils/type_utils.h>
15+#include "log/log.h"
16+#include "register/op_impl_registry.h"
17+#include "tiling/platform/platform_ascendc.h"
18+#include "../op_kernel/bernoulli_mask_tiling_data.h"
19+#include "../op_kernel/bernoulli_mask_tiling_key.h"
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+#include "bernoulli_mask.h"
12+#include "bernoulli_mask_tiling_key.h"
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+#ifndef BERNOULLI_MASK_H
12+#define BERNOULLI_MASK_H
13+ 
14+#include <type_traits>
15+#include "kernel_operator.h"
16+#include "bernoulli_mask_tiling_data.h"
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+#endif // BERNOULLI_MASK_H
@@ -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+#ifndef BERNOULLI_MASK_TILING_DATA_H
12+#define BERNOULLI_MASK_TILING_DATA_H
13+ 
14+#include <cstdint>
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+#endif // BERNOULLI_MASK_TILING_DATA_H
@@ -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+#ifndef BERNOULLI_MASK_TILING_KEY_H
12+#define BERNOULLI_MASK_TILING_KEY_H
13+ 
14+#include <cstdint>
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+#endif // BERNOULLI_MASK_TILING_KEY_H
@@ -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+ @staticmethod
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+ @staticmethod
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+#include <cstdint>
12+#include <vector>
13+ 
14+#include <gtest/gtest.h>
15+ 
16+#include "infershape_case_executor.h"
17+#include "infershape_context_faker.h"
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+#include <cstdint>
12+#include <string>
13+#include <tuple>
14+#include <vector>
15+ 
16+#include <gtest/gtest.h>
17+ 
18+#include "tiling_case_executor.h"
19+#include "../../../op_kernel/bernoulli_mask_tiling_data.h"
20+#include "../../../op_kernel/bernoulli_mask_tiling_key.h"
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