Pull Request已成功合入, 合并人@CANN-robot
(感谢 hzw_rpap 的贡献)变更摘要
此 PR 为 Atlas A2/A3 训练系列产品新增 aclnnBernoulli 和 aclnnInplaceBernoulli 算子的低内存实现,代码位于 experimental/random/bernoulli_mask。核心思路是:保留 DSAGenBitMask 的 seed/offset 随机序列生成逻辑,新增 Ascend C BernoulliMask Kernel,将 packed bit mask 直接展开为目标数据类型的 0 或 1,从而避免原实现中全尺寸 Fill、DropoutDoMask 及部分 Cast 中间张量带来的额外内存占用。对于连续输出,压缩 mask 与输出复用同一块设备内存,Kernel 从高地址向低地址分波展开,在波次之间同步以避免覆盖尚未读取的 mask;非连续输出则通过 ViewCopy 写回。
主要改动
-
新增
BernoulliMaskAscend C Kernel:在op_kernel/bernoulli_mask.h中实现模板类KernelBernoulliMask<T>,支持 10 种输出数据类型(half/float/double/uint8_t/int8_t/int16_t/int32_t/int64_t/bfloat16_t/bool),通过模板特化和Select/Cast组合将 packed bit mask 按位展开为目标类型的 0 或 1;double类型使用Gather指令将两个float拼接为fp64。 -
新增内存复用路径
ProcessAliased:当输出满足连续稠密布局且可容纳压缩 mask 时(由CanWriteOutDirectly判定),BernoulliMaskKernel 将 mask 直接写入输出 buffer 的高地址段,通过ProcessAliased方法从高到低逐波展开并在波间执行SyncAll,最后在 Core 0 处理剩余前缀,实现 mask 与输出的零额外内存复用。 -
新增 API 层架构分支调度:
aclnn_bernoulli.cpp中BernoulliGetWorkspaceSizeCommon根据 NPU 架构分派:DAV_2201走DSAGenBitMask+BernoulliMask低内存路径,DAV_3510走StatelessBernoulli路径;prob=0和prob=1分别走ZerosLike/OnesLike快速路径;同时提供aclnnBernoulli和aclnnInplaceBernoulli的完整两段式 API。 -
新增 Tiling 与算子注册:
bernoulli_mask_tiling.cpp实现 UB 内存预算驱动的分 tile 策略——根据 UB 大小扣除预留量后按bytes_per_element计算tileElements,并分配给多核;bernoulli_mask_def.cpp注册算子,输入固定为DT_UINT8packed mask,输出支持 10 种数据类型;bernoulli_mask_infershape.cpp通过output_shape属性完成形状推导。 -
新增完整的测试与示例:包含 ACLNN 功能矩阵 ST 测试(
test_aclnn_bernoulli_st.cpp,覆盖 10 种输出 dtype、4 种概率 dtype、rank 0~8、连续/非连续/转置/空 Tensor、inplace/outplace、边界情况、可复现性等)、Kernel UT(shape 推导与 tiling 参数验证)、TTK golden 脚本(bernoulli_mask.py)以及示例程序(test_aclnn_bernoulli.cpp)。


代码审查
Now I have completed my thorough review of all 29 files. Let me write the closing summary.
审查总结
审查覆盖的 29 个文件
| 文件 | 结论 |
|---|---|
experimental/random/bernoulli_mask/CMakeLists.txt |
无问题 |
experimental/random/bernoulli_mask/README.md |
无问题 |
experimental/random/bernoulli_mask/examples/CMakeLists.txt |
无问题 |
experimental/random/bernoulli_mask/examples/README.md |
无问题 |
experimental/random/bernoulli_mask/examples/run.sh |
无问题 |
experimental/random/bernoulli_mask/examples/test_aclnn_bernoulli.cpp |
P3: executor 泄漏 |
experimental/random/bernoulli_mask/op_api/aclnn_bernoulli.cpp |
无问题 |
experimental/random/bernoulli_mask/op_api/aclnn_bernoulli.h |
无问题 |
experimental/random/bernoulli_mask/op_api/bernoulli_mask.cpp |
无问题 |
experimental/random/bernoulli_mask/op_api/bernoulli_mask.h |
无问题 |
experimental/random/bernoulli_mask/op_host/bernoulli_mask_def.cpp |
无问题 |
experimental/random/bernoulli_mask/op_host/bernoulli_mask_infershape.cpp |
无问题 |
experimental/random/bernoulli_mask/op_host/bernoulli_mask_tiling.cpp |
无问题 |
experimental/random/bernoulli_mask/op_kernel/bernoulli_mask.cpp |
无问题 |
experimental/random/bernoulli_mask/op_kernel/bernoulli_mask.h |
无问题 |
experimental/random/bernoulli_mask/op_kernel/bernoulli_mask_tiling_data.h |
无问题 |
experimental/random/bernoulli_mask/op_kernel/bernoulli_mask_tiling_key.h |
无问题 |
experimental/random/bernoulli_mask/tests/CMakeLists.txt |
无问题 |
experimental/random/bernoulli_mask/tests/assets/bernoulli_mask.py |
无问题 |
experimental/random/bernoulli_mask/tests/st/CMakeLists.txt |
无问题 |
experimental/random/bernoulli_mask/tests/st/README.md |
无问题 |
experimental/random/bernoulli_mask/tests/st/run.sh |
无问题 |
experimental/random/bernoulli_mask/tests/st/test_aclnn_bernoulli_st.cpp |
P3: executor 泄漏 |
experimental/random/bernoulli_mask/tests/ttk/bernoulli_mask.csv |
无问题 |
experimental/random/bernoulli_mask/tests/ttk/bernoulli_mask_alias.csv |
无问题 |
experimental/random/bernoulli_mask/tests/ut/CMakeLists.txt |
无问题 |
experimental/random/bernoulli_mask/tests/ut/op_host/CMakeLists.txt |
无问题 |
experimental/random/bernoulli_mask/tests/ut/op_host/test_bernoulli_mask_infershape.cpp |
无问题 |
experimental/random/bernoulli_mask/tests/ut/op_host/test_bernoulli_mask_tiling.cpp |
无问题 |
问题统计
- P0: 0
- P1: 0
- P2: 0
- P3: 2(均为测试/示例代码中
aclOpExecutor*资源泄漏)
整体风险评估
低风险。 核心算子逻辑(Kernel、Tiling、L0 Launcher、InferShape、算子定义)经过严格审查,未发现正确性、安全性或可靠性问题。整数溢出检查、空指针校验、返回值检查、同步原语使用、存储复用地址计算等均正确。两个 P3 问题仅涉及测试和示例代码中的 executor 资源泄漏,不影响生产行为。
⚠️ 已识别出整体风险,但无法提取行内评论,请参考整体评估。


描述
本 PR 为 Atlas A2/A3 训练系列产品提供
aclnnBernoulli和aclnnInplaceBernoulli的低内存实现,代码位于experimental/random/bernoulli_mask。一般概率路径保留
DSAGenBitMask的seed/offset随机序列生成逻辑,新增 Ascend CBernoulliMaskKernel,将 packed bit mask 直接展开为目标数据类型的0或1,避免原实现中全尺寸Fill、DropoutDoMask及部分Cast中间张量带来的额外内存占用。对于满足存储条件的连续输出,压缩 mask 与输出复用同一块设备内存,Kernel 从高地址向低地址分波展开,并在波次之间同步,避免覆盖尚未读取的 mask。小张量和非连续输出使用独立缓冲区,非连续结果通过
ViewCopy写回目标 view。prob=0和prob=1分别沿用ZerosLike和OnesLike快速路径。实现支持 Atlas A2/A3、10 种输出数据类型、4 种 probability 数据类型、标量、空 Tensor、rank 0~8、连续及非连续 Tensor,以及 out-of-place 和 inplace 调用。
关联的Issue
测试
测试环境为 CANN 8.5.2,测试平台为 Atlas A2 和 Atlas A3。功能、性能和内存测试基于源码提交 295d1e7fe793d3e5477d00cd1af86ebcea770f1d;后续变更仅包含 clang-format 格式整理、README 表述修订和提交历史合并,不改变算子实现语义。
测试结果如下:
Sanitizer 中 13 项 LIMITED 为 CANN 8.5 对 masked
Select指令的已知识别限制,未计入 PASS,全部检查均无失败项。文档更新
BernoulliMask算子 README,说明接口、约束、实现原理、构建和测试方法;类型标签