已合并
stateless normal update #3219
biabu创建于 6月9日
stateless normal update #3219
已合并
biabu创建于 6月9日
已删除 :master合入到cann/ops-mathmaster
9 个文件变更+90-179
Mmath/mul/op_api/mul.cpp+31-0
@@ -143,4 +143,35 @@ const aclTensor *Mul(const aclTensor *self, const aclTensor *other, aclOpExecuto
143 return MulAiCpu(self, other, mulOut, executor);143 return MulAiCpu(self, other, mulOut, executor);
144}144}
145 145 
146+const aclTensor *MulInplace(const aclTensor *self, const aclTensor *rfRes, aclOpExecutor *executor) {
147+ Shape broadcastShape;
148+ OP_CHECK_BROADCAST_AND_INFER_SHAPE(self, rfRes, broadcastShape, return nullptr);
149+ 
150+ // 校验输出tensor的shape和rfRes tensor一致
151+ if (broadcastShape != rfRes->GetViewShape()) {
152+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "self and rfRes broadcastShape [%s] not equal to rfRes shape [%s], do no support inplace from the 'rfRes' tensor!",
153+ op::ToString(broadcastShape).GetString(), op::ToString(rfRes->GetViewShape()).GetString());
154+ return nullptr;
155+ }
156+ 
157+ // 校验输出tensor的dtype和rfRes tensor一致
158+ bool isMixDataType = (self->GetDataType() == DataType::DT_FLOAT16 && rfRes->GetDataType() == DataType::DT_FLOAT) ||
159+ (self->GetDataType() == DataType::DT_FLOAT && rfRes->GetDataType() == DataType::DT_FLOAT16) ||
160+ (self->GetDataType() == DataType::DT_BF16 && rfRes->GetDataType() == DataType::DT_FLOAT) ||
161+ (self->GetDataType() == DataType::DT_FLOAT && rfRes->GetDataType() == DataType::DT_BF16);
162+ if (isMixDataType && (rfRes->GetDataType() == DataType::DT_FLOAT16 || rfRes->GetDataType() == DataType::DT_BF16)) {
163+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "out dtype DataType::DT_FLOAT not equal to rfRes dtype [%s], do no support inplace from the 'rfRes' tensor!",
164+ op::ToString(rfRes->GetDataType()).GetString());
165+ return nullptr;
166+ }
167+ 
168+ auto mulOut = const_cast<aclTensor*>(rfRes);
K

参数定义成const,后续又使用const_cast去掉了

likedislike
biabu
biabu
6月16日 评论:
169+ 
170+ if (isMixDataType || (IsAiCoreSupport(self) && IsAiCoreSupport(rfRes)) || IsDoubleSupport(self, rfRes)) {
171+ return MulAiCore(self, rfRes, mulOut, executor);
172+ }
173+ 
174+ return MulAiCpu(self, rfRes, mulOut, executor);
175+}
176+ 
146} // namespace l0op177} // namespace l0op
Mmath/mul/op_api/mul.h+2-1
@@ -15,7 +15,8 @@
15 15 
16namespace l0op {16namespace l0op {
17const aclTensor *Mul(const aclTensor *self, const aclTensor *other, aclOpExecutor *executor);17const aclTensor *Mul(const aclTensor *self, const aclTensor *other, aclOpExecutor *executor);
18- bool IsMulSupportNonContiguous(const aclTensor* self, const aclTensor *other);18+bool IsMulSupportNonContiguous(const aclTensor* self, const aclTensor *other);
19+const aclTensor *MulInplace(const aclTensor *self, const aclTensor *rfRes, aclOpExecutor *executor);
19}20}
20 21 
21#endif // OP_API_INC_LEVEL0_MUL_H_22#endif // OP_API_INC_LEVEL0_MUL_H_
Mrandom/dsa_random_normal/op_host/op_api/aclnn_normal.cpp+4-15
@@ -15,7 +15,6 @@
15#include "random/stateless_random_normal_v2/op_api/stateless_random_normal_v2.h"15#include "random/stateless_random_normal_v2/op_api/stateless_random_normal_v2.h"
16#include "random/stateless_random_normal_v3/op_api/stateless_random_normal_v3.h"16#include "random/stateless_random_normal_v3/op_api/stateless_random_normal_v3.h"
17#include "random/stateless_normal/op_api/stateless_normal.h"17#include "random/stateless_normal/op_api/stateless_normal.h"
18-#include "conversion/broadcast_to/op_api/broadcast_to.h"
19#include "dsa_random_normal.h"18#include "dsa_random_normal.h"
20#include "../../../../conversion/concat_d/op_api/concat_d.h"19#include "../../../../conversion/concat_d/op_api/concat_d.h"
21#include "opdev/platform.h"20#include "opdev/platform.h"
@@ -184,11 +183,6 @@ static const aclTensor* normalDavidPath(
184 auto oneStdTensor = executor->ConvertToTensor(183 auto oneStdTensor = executor->ConvertToTensor(
185 oneStdVector.data(), oneStdVector.size(), op::DataType::DT_FLOAT);184 oneStdVector.data(), oneStdVector.size(), op::DataType::DT_FLOAT);
186 185 
187- auto outDims = op::ToShapeVector(selfContiguous->GetViewShape());
188- auto outShapeArray = executor->AllocIntArray(outDims.data(), outDims.size());
189- zeroMeanTensor = l0op::BroadcastTo(zeroMeanTensor, outShapeArray, executor);
190- oneStdTensor = l0op::BroadcastTo(oneStdTensor, outShapeArray, executor);
191- 
192 auto normalOut = l0op::StatelessNormal(186 auto normalOut = l0op::StatelessNormal(
193 selfFloat, seed, offset, zeroMeanTensor, oneStdTensor, executor);187 selfFloat, seed, offset, zeroMeanTensor, oneStdTensor, executor);
194 CHECK_RET(normalOut != nullptr, nullptr);188 CHECK_RET(normalOut != nullptr, nullptr);
@@ -196,12 +190,12 @@ static const aclTensor* normalDavidPath(
196 // fp32 下做 z * std + mean(全程 fp32,无中间舍入)190 // fp32 下做 z * std + mean(全程 fp32,无中间舍入)
197 auto stdScalar = executor->AllocScalar(std);191 auto stdScalar = executor->AllocScalar(std);
198 auto stdTensor = executor->ConvertToTensor(stdScalar, op::DataType::DT_FLOAT);192 auto stdTensor = executor->ConvertToTensor(stdScalar, op::DataType::DT_FLOAT);
199- auto mulOut = l0op::Mul(normalOut, stdTensor, executor);193+ auto mulOut = l0op::MulInplace(stdTensor, normalOut, executor);
200 CHECK_RET(mulOut != nullptr, nullptr);194 CHECK_RET(mulOut != nullptr, nullptr);
201 195 
202 auto meanScalar = executor->AllocScalar(mean);196 auto meanScalar = executor->AllocScalar(mean);
203 auto meanTensor = executor->ConvertToTensor(meanScalar, op::DataType::DT_FLOAT);197 auto meanTensor = executor->ConvertToTensor(meanScalar, op::DataType::DT_FLOAT);
204- return l0op::Add(mulOut, meanTensor, executor);198+ return l0op::AddInplace(meanTensor, mulOut, executor);
205 } else {199 } else {
206 // V2 路径:seed 转化为 key,offset 转化为 counter200 // V2 路径:seed 转化为 key,offset 转化为 counter
207 FVector<int64_t, op::MAX_DIM_NUM> key_vec = {seed};201 FVector<int64_t, op::MAX_DIM_NUM> key_vec = {seed};
@@ -256,11 +250,6 @@ static const aclTensor* normalTensorDavidPath(
256 auto oneStdTensor = executor->ConvertToTensor(250 auto oneStdTensor = executor->ConvertToTensor(
257 oneStdVector.data(), oneStdVector.size(), op::DataType::DT_FLOAT);251 oneStdVector.data(), oneStdVector.size(), op::DataType::DT_FLOAT);
258 252 
259- auto outDims = op::ToShapeVector(selfContiguous->GetViewShape());
260- auto outShapeArray = executor->AllocIntArray(outDims.data(), outDims.size());
261- zeroMeanTensor = l0op::BroadcastTo(zeroMeanTensor, outShapeArray, executor);
262- oneStdTensor = l0op::BroadcastTo(oneStdTensor, outShapeArray, executor);
263- 
264 auto normalOut = l0op::StatelessNormal(253 auto normalOut = l0op::StatelessNormal(
265 selfFloat, seedTensor, resultAddOut, zeroMeanTensor, oneStdTensor, executor);254 selfFloat, seedTensor, resultAddOut, zeroMeanTensor, oneStdTensor, executor);
266 CHECK_RET(normalOut != nullptr, nullptr);255 CHECK_RET(normalOut != nullptr, nullptr);
@@ -268,12 +257,12 @@ static const aclTensor* normalTensorDavidPath(
268 // Step 2: fp32 下做 z * std + mean257 // Step 2: fp32 下做 z * std + mean
269 auto stdScalar = executor->AllocScalar(std);258 auto stdScalar = executor->AllocScalar(std);
270 auto stdTensor = executor->ConvertToTensor(stdScalar, op::DataType::DT_FLOAT);259 auto stdTensor = executor->ConvertToTensor(stdScalar, op::DataType::DT_FLOAT);
271- auto mulOut = l0op::Mul(normalOut, stdTensor, executor);260+ auto mulOut = l0op::MulInplace(stdTensor, normalOut, executor);
272 CHECK_RET(mulOut != nullptr, nullptr);261 CHECK_RET(mulOut != nullptr, nullptr);
273 262 
274 auto meanScalar = executor->AllocScalar(mean);263 auto meanScalar = executor->AllocScalar(mean);
275 auto meanTensor = executor->ConvertToTensor(meanScalar, op::DataType::DT_FLOAT);264 auto meanTensor = executor->ConvertToTensor(meanScalar, op::DataType::DT_FLOAT);
276- return l0op::Add(mulOut, meanTensor, executor);265+ return l0op::AddInplace(meanTensor, mulOut, executor);
277 } else {266 } else {
278 // V2 路径:保持原有 Cast/concat 逻辑267 // V2 路径:保持原有 Cast/concat 逻辑
279 auto normalSeedU64 = l0op::Cast(seedTensor, op::DataType::DT_UINT64, executor);268 auto normalSeedU64 = l0op::Cast(seedTensor, op::DataType::DT_UINT64, executor);
Mrandom/stateless_normal/op_graph/stateless_normal_proto.h+2-2
@@ -28,8 +28,8 @@ namespace ge {
28* @li shape: 1-D. The shape of the output tensor. Must be one of the following types: int64.28* @li shape: 1-D. The shape of the output tensor. Must be one of the following types: int64.
29* @li seed: 0-D. Seed for the Philox4x32-10 RNG algorithm. Must be one of the following types: int64.29* @li seed: 0-D. Seed for the Philox4x32-10 RNG algorithm. Must be one of the following types: int64.
30* @li offset: 0-D. Offset for the Philox4x32-10 RNG algorithm. Must be one of the following types: int64.30* @li offset: 0-D. Offset for the Philox4x32-10 RNG algorithm. Must be one of the following types: int64.
31-* @li mean: Scalar or tensor. Mean of the normal distribution. Must be one of the following types: float, float16, bfloat16.31+* @li mean: Scalar or tensor. Mean of the normal distribution. Must be one of the following types: float, float16, bfloat16. Only the 0th element is used for calculation if a tensor is input.
32-* @li std: Scalar or tensor. Standard deviation of the normal distribution. Must be one of the following types: float, float16, bfloat16. \n32+* @li std: Scalar or tensor. Standard deviation of the normal distribution. Must be one of the following types: float, float16, bfloat16. Only the 0th element is used for calculation if a tensor is input. \n
33 33 
34* @par Attributes:34* @par Attributes:
35* dtype: Output data type. Must be one of the following types: float16, bfloat16, float32.35* dtype: Output data type. Must be one of the following types: float16, bfloat16, float32.
Mrandom/stateless_normal/op_host/arch35/stateless_normal_tiling_arch35.cpp+0-30
@@ -81,36 +81,6 @@ StatelessNormalTiling::StatelessNormalTiling(gert::TilingContext* ctx)
81 81 
82ge::graphStatus StatelessNormalTiling::UniqueProcess()82ge::graphStatus StatelessNormalTiling::UniqueProcess()
83{83{
84- // L2 层已将 Size=1 的 mean/stdev 广播到 output shape,kernel 只需 BothTensor 路径
85- 
86- auto outputShape = context_->GetOutputShape(OUTPUT_IDX_Y);
87- OP_CHECK_NULL_WITH_CONTEXT(context_, outputShape);
88- int64_t outputSize = outputShape->GetStorageShape().GetShapeSize();
89- 
90- // 校验 mean/stdev shape 与 output 一致(L2 已完成广播)
91- auto meanTensor = context_->GetInputTensor(INPUT_IDX_MEAN);
92- OP_CHECK_NULL_WITH_CONTEXT(context_, meanTensor);
93- int64_t meanSize = meanTensor->GetShapeSize();
94- if (meanSize != outputSize) {
95- std::string valueStr = std::to_string(meanSize) + " and " + std::to_string(outputSize);
96- std::string reasonMsg = "StatelessNormal requires mean shapeSize == output shapeSize";
97- OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON(context_->GetNodeName(), "input mean", valueStr.c_str(), reasonMsg.c_str());
98- return ge::GRAPH_FAILED;
99- }
100- 
101- auto stdevTensor = context_->GetInputTensor(INPUT_IDX_STDEV);
102- OP_CHECK_NULL_WITH_CONTEXT(context_, stdevTensor);
103- int64_t stdevSize = stdevTensor->GetShapeSize();
104- if (stdevSize != outputSize) {
105- std::string valueStr = std::to_string(stdevSize) + " and " + std::to_string(outputSize);
106- std::string reasonMsg = "StatelessNormal requires stdev shapeSize == output shapeSize";
107- OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON(context_->GetNodeName(), "input stdev", valueStr.c_str(), reasonMsg.c_str());
108- return ge::GRAPH_FAILED;
109- }
110- 
111- // SplitUntil32Bit + totalThreads/counterOffset/kernelOffset 全部由基类 enableSplitBlocks 自动完成
112- // tilingKey 使用基类默认值 100(与 V3 一致),不再设置
113- 
114 return ge::GRAPH_SUCCESS;84 return ge::GRAPH_SUCCESS;
115}85}
116 86 
Mrandom/stateless_normal/op_kernel/arch35/stateless_normal.h+17-24
@@ -30,45 +30,38 @@ namespace StatelessNormalSimt {
30using namespace AscendC;30using namespace AscendC;
31using namespace RandomKernelBase;31using namespace RandomKernelBase;
32 32 
33+static constexpr uint16_t NUM_TWO = 2;
34+ 
33// Box-Muller normal transform functor35// Box-Muller normal transform functor
34// mean/stdev are always float* (DT_FLOAT from L2 layer), output is T* (may be bf16/fp16/fp32)36// mean/stdev are always float* (DT_FLOAT from L2 layer), output is T* (may be bf16/fp16/fp32)
35// For bf16/fp16: three-step rounding to match GPU normal_()→mul_(std)→add_(mean)37// For bf16/fp16: three-step rounding to match GPU normal_()→mul_(std)→add_(mean)
36-template <typename T>38+template <typename T, typename M_T, typename S_T>
37struct NormalTransform {39struct NormalTransform {
38- __gm__ volatile float* meanGM_;40+ __gm__ volatile M_T* meanGM_;
39- __gm__ volatile float* stdevGM_;41+ __gm__ volatile S_T* stdevGM_;
40 42 
41- __aicore__ NormalTransform(__gm__ volatile float* meanGM, __gm__ volatile float* stdevGM)43+ __aicore__ NormalTransform(__gm__ volatile M_T* meanGM, __gm__ volatile S_T* stdevGM)
42 : meanGM_(meanGM), stdevGM_(stdevGM) {}44 : meanGM_(meanGM), stdevGM_(stdevGM) {}
43 45 
44 __simt_callee__ __aicore__ inline void operator()(46 __simt_callee__ __aicore__ inline void operator()(
45 __gm__ volatile T* outputGm, uint64_t li, const uint32_t* results, uint32_t iStep,47 __gm__ volatile T* outputGm, uint64_t li, const uint32_t* results, uint32_t iStep,
46 [[maybe_unused]] uint32_t unroll = 1)48 [[maybe_unused]] uint32_t unroll = 1)
47 {49 {
48- uint32_t pairBase = (iStep / 2) * 2;50+ uint32_t pairBase = (iStep / NUM_TWO) * NUM_TWO;
49 float u1 = results[pairBase] * RAND_2POW32_INV + RAND_2POW32_INV_HALF;51 float u1 = results[pairBase] * RAND_2POW32_INV + RAND_2POW32_INV_HALF;
50 float u2 = results[pairBase + 1] * RAND_2POW32_INV + RAND_2POW32_INV_HALF;52 float u2 = results[pairBase + 1] * RAND_2POW32_INV + RAND_2POW32_INV_HALF;
51 float z0, z1;53 float z0, z1;
52 // 使用对齐torch版本的BoxMullerFloat,无需eps保护54 // 使用对齐torch版本的BoxMullerFloat,无需eps保护
53 BoxMullerFloat(u1, u2, &z0, &z1);55 BoxMullerFloat(u1, u2, &z0, &z1);
54- float z = (iStep % 2 == 0) ? z0 : z1;56+ float z = (iStep % NUM_TWO == 0) ? z0 : z1;
55- if constexpr (IsSameType<T, float>::value) {57+ 
56- // fp32: 无中间舍入,直接计算58+ outputGm[li] = static_cast<T>(z * static_cast<float>(stdevGM_[0]) + static_cast<float>(meanGM_[0]));
57- outputGm[li] = z * stdevGM_[li] + meanGM_[li];
58- } else {
59- // bf16/fp16: 严格对标 PyTorch GPU 三步舍入
60- // GPU 内部: output.normal_(0,1) → output.mul_(std) → output.add_(mean)
61- // 每步都在目标 dtype 上存储,产生中间舍入
62- T zT = static_cast<T>(z); // step1: normal_() 存为目标dtype
63- T mulT = static_cast<T>(static_cast<float>(zT) * stdevGM_[li]); // step2: mul_(std) 在目标dtype下
64- outputGm[li] = static_cast<T>(static_cast<float>(mulT) + meanGM_[li]); // step3: add_(mean) 在目标dtype下
65- }
66 }59 }
67};60};
68 61 
69// Launcher: wraps VF_CALL for each split block, adjusts GM pointers by gmOffset62// Launcher: wraps VF_CALL for each split block, adjusts GM pointers by gmOffset
70// mean/stdev pointers are always float* (DT_FLOAT), output is T*63// mean/stdev pointers are always float* (DT_FLOAT), output is T*
71-template <typename T>64+template <typename T, typename M_T, typename S_T>
72struct NormalLauncher {65struct NormalLauncher {
73 int64_t seed_;66 int64_t seed_;
74 int64_t realOffset_;67 int64_t realOffset_;
@@ -89,11 +82,11 @@ struct NormalLauncher {
89 {82 {
90 __gm__ volatile T* gmPtr = reinterpret_cast<__gm__ volatile T*>(yAddr_) + gmOffset;83 __gm__ volatile T* gmPtr = reinterpret_cast<__gm__ volatile T*>(yAddr_) + gmOffset;
91 // mean/stdev 始终按 float* 读取(L2 传入 DT_FLOAT tensor)84 // mean/stdev 始终按 float* 读取(L2 传入 DT_FLOAT tensor)
92- __gm__ volatile float* meanPtr = reinterpret_cast<__gm__ volatile float*>(meanAddr_) + gmOffset;85+ __gm__ volatile M_T* meanPtr = reinterpret_cast<__gm__ volatile M_T*>(meanAddr_);
93- __gm__ volatile float* stdevPtr = reinterpret_cast<__gm__ volatile float*>(stdevAddr_) + gmOffset;86+ __gm__ volatile S_T* stdevPtr = reinterpret_cast<__gm__ volatile S_T*>(stdevAddr_);
94 87 
95- NormalTransform<T> transform(meanPtr, stdevPtr);88+ NormalTransform<T, M_T, S_T> transform(meanPtr, stdevPtr);
96- Simt::VF_CALL<PhiloxSimtKernelDiscontinuous<T, NormalTransform<T>>>(89+ Simt::VF_CALL<PhiloxSimtKernelDiscontinuous<T, NormalTransform<T, M_T, S_T>>>(
97 Simt::Dim3(DEFAULT_SIMT_THREAD_NUM),90 Simt::Dim3(DEFAULT_SIMT_THREAD_NUM),
98 gmPtr, realOffset_ + kernelOffset, seed_, static_cast<uint64_t>(numel),91 gmPtr, realOffset_ + kernelOffset, seed_, static_cast<uint64_t>(numel),
99 policy.magic, policy.shift, static_cast<uint64_t>(totalThreads), transform);92 policy.magic, policy.shift, static_cast<uint64_t>(totalThreads), transform);
@@ -101,7 +94,7 @@ struct NormalLauncher {
101};94};
102 95 
103// Entry point: called from stateless_normal.cpp96// Entry point: called from stateless_normal.cpp
104-template <typename T>97+template <typename T, typename M_T, typename S_T>
105__aicore__ inline void Process(98__aicore__ inline void Process(
106 GM_ADDR seed, GM_ADDR offset,99 GM_ADDR seed, GM_ADDR offset,
107 GM_ADDR y, GM_ADDR mean, GM_ADDR stdev,100 GM_ADDR y, GM_ADDR mean, GM_ADDR stdev,
@@ -115,7 +108,7 @@ __aicore__ inline void Process(
115 int64_t realSeed = *(reinterpret_cast<__gm__ int64_t*>(seed));108 int64_t realSeed = *(reinterpret_cast<__gm__ int64_t*>(seed));
116 int64_t realOffset = *(reinterpret_cast<__gm__ int64_t*>(offset));109 int64_t realOffset = *(reinterpret_cast<__gm__ int64_t*>(offset));
117 110 
118- NormalLauncher<T> launcher(realSeed, realOffset, y, mean, stdev);111+ NormalLauncher<T, M_T, S_T> launcher(realSeed, realOffset, y, mean, stdev);
119 ProcessWithSplitBlocks(tilingData, launcher);112 ProcessWithSplitBlocks(tilingData, launcher);
120}113}
121 114 
Mrandom/stateless_normal/op_kernel/stateless_normal.cpp+1-1
@@ -22,6 +22,6 @@ extern "C" __global__ __aicore__ void stateless_normal(
22 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIV_1_0);22 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIV_1_0);
23 GET_TILING_DATA_WITH_STRUCT(RandomUnifiedSimtTilingDataStruct, tilingData, tiling);23 GET_TILING_DATA_WITH_STRUCT(RandomUnifiedSimtTilingDataStruct, tilingData, tiling);
24 if (TILING_KEY_IS(STATELESS_NORMAL_DEFAULT_TILING_KEY)) {24 if (TILING_KEY_IS(STATELESS_NORMAL_DEFAULT_TILING_KEY)) {
25- StatelessNormalSimt::Process<DTYPE_Y>(seed, offset, y, mean, stdev, &tilingData);25+ StatelessNormalSimt::Process<DTYPE_Y, DTYPE_MEAN, DTYPE_STDEV>(seed, offset, y, mean, stdev, &tilingData);
26 }26 }
27}27}
Mrandom/stateless_normal/tests/ut/op_host/arch35/test_stateless_normal_tiling.cpp+0-56
@@ -337,59 +337,3 @@ TEST_F(StatelessNormalTilingTest, test_both_tensor_fp16)
337 std::vector<size_t> expectWorkspaces = {0};337 std::vector<size_t> expectWorkspaces = {0};
338 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectWorkspaces);338 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectWorkspaces);
339}339}
340- 
341-// Test 11: Error case - stdev shape mismatch (not scalar, not equal to output)
342-TEST_F(StatelessNormalTilingTest, test_stdev_shape_mismatch_error)
343-{
344- optiling::RandomOperatorCompileInfo compileInfo = {40, 196608};
345- vector<int64_t> shapeValue = {1024};
346- int64_t seedValue = 2;
347- int64_t offsetValue = 0;
348- float meanValue = 0.0f;
349- std::vector<float> stdevData(256, 1.0f); // size 256 != output size 1024
350- gert::TilingContextPara tilingContextPara(
351- "StatelessNormal",
352- {
353- {{{1,}, {1,}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
354- {{{1,}, {1,}}, ge::DT_INT64, ge::FORMAT_ND, true, &seedValue},
355- {{{1,}, {1,}}, ge::DT_INT64, ge::FORMAT_ND, true, &offsetValue},
356- {{{1,}, {1,}}, ge::DT_FLOAT, ge::FORMAT_ND, true, &meanValue},
357- {{{256,}, {256,}}, ge::DT_FLOAT, ge::FORMAT_ND, true, stdevData.data()},
358- },
359- {
360- {{{1024,}, {1024,}}, ge::DT_FLOAT, ge::FORMAT_ND},
361- },
362- {
363- {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
364- },
365- &compileInfo);
366- ExecuteTestCase(tilingContextPara, ge::GRAPH_FAILED);
367-}
368- 
369-// Test 12: Error case - mean shape mismatch with BF16 (not scalar, not equal to output)
370-TEST_F(StatelessNormalTilingTest, test_mean_shape_mismatch_error)
371-{
372- optiling::RandomOperatorCompileInfo compileInfo = {40, 196608};
373- vector<int64_t> shapeValue = {1024};
374- int64_t seedValue = 1;
375- int64_t offsetValue = 0;
376- std::vector<float> meanData(512, 0.0f); // size 512 != output size 1024
377- float stdevValue = 1.0f;
378- gert::TilingContextPara tilingContextPara(
379- "StatelessNormal",
380- {
381- {{{1,}, {1,}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
382- {{{1,}, {1,}}, ge::DT_INT64, ge::FORMAT_ND, true, &seedValue},
383- {{{1,}, {1,}}, ge::DT_INT64, ge::FORMAT_ND, true, &offsetValue},
384- {{{512,}, {512,}}, ge::DT_FLOAT, ge::FORMAT_ND, true, meanData.data()},
385- {{{1,}, {1,}}, ge::DT_FLOAT, ge::FORMAT_ND, true, &stdevValue},
386- },
387- {
388- {{{1024,}, {1024,}}, ge::DT_FLOAT, ge::FORMAT_ND},
389- },
390- {
391- {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
392- },
393- &compileInfo);
394- ExecuteTestCase(tilingContextPara, ge::GRAPH_FAILED);
395-}
Mrandom/stateless_random_normal_v2/op_api/aclnn_normal_out.cpp+33-50
@@ -39,6 +39,14 @@ static constexpr size_t MAX_DIM_LEN = 8;
39static const std::initializer_list<op::DataType> DTYPE_SUPPORT_LIST = {39static const std::initializer_list<op::DataType> DTYPE_SUPPORT_LIST = {
40 op::DataType::DT_BF16, op::DataType::DT_FLOAT16, op::DataType::DT_FLOAT, op::DataType::DT_DOUBLE};40 op::DataType::DT_BF16, op::DataType::DT_FLOAT16, op::DataType::DT_FLOAT, op::DataType::DT_DOUBLE};
41 41 
42+// 表示mean,std的数据类型是Tensor还是Scalar
43+enum class ScalarMode {
44+ TensorTensor,
45+ TensorScalar,
46+ ScalarTensor,
47+ ScalarScalar
48+};
49+ 
42/* 查看TensorFloat的Dtype和Shape */50/* 查看TensorFloat的Dtype和Shape */
43static bool CheckTensorAndFloatDtype(const aclTensor* mean, const aclTensor* out)51static bool CheckTensorAndFloatDtype(const aclTensor* mean, const aclTensor* out)
44{52{
@@ -204,19 +212,12 @@ static aclnnStatus CheckTensorAndTensorParams(const aclTensor* mean, const aclTe
204aclnnStatus CommonLogicGeneralNormal(212aclnnStatus CommonLogicGeneralNormal(
205 const aclTensor* mean, const aclTensor* std, int64_t seed, int64_t offset, aclTensor* self, aclTensor* out,213 const aclTensor* mean, const aclTensor* std, int64_t seed, int64_t offset, aclTensor* self, aclTensor* out,
206 UniqueExecutor& uniqueExecutor, uint64_t* workspaceSize, aclOpExecutor** executor,214 UniqueExecutor& uniqueExecutor, uint64_t* workspaceSize, aclOpExecutor** executor,
207- bool isScalarMode = false)215+ ScalarMode scalarMode = ScalarMode::TensorTensor)
208{216{
209 const aclTensor* addOut = nullptr;217 const aclTensor* addOut = nullptr;
210 218 
211 if(GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510 && self->GetDataType() != DataType::DT_DOUBLE){219 if(GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510 && self->GetDataType() != DataType::DT_DOUBLE){
212- if (isScalarMode) {220+ if (scalarMode == ScalarMode::TensorTensor || scalarMode == ScalarMode::ScalarTensor) {
213- // V4 标量路径:对标 PyTorch normal_(scalar_mean, scalar_std) 单步 FMA
214- // Kernel 以 fp32 生成 N(0,1),L2 层 fp32 Mul+Add,最后 Cast 到目标 dtype
215- auto selfFloat = (self->GetDataType() != DataType::DT_FLOAT)
216- ? l0op::Cast(self, DataType::DT_FLOAT, uniqueExecutor.get())
217- : self;
218- CHECK_RET(selfFloat != nullptr, ACLNN_ERR_INNER_NULLPTR);
219- 
220 FVector<float> zeroMeanVector = {0.0f};221 FVector<float> zeroMeanVector = {0.0f};
221 auto zeroMeanTensor = uniqueExecutor.get()->ConvertToTensor(222 auto zeroMeanTensor = uniqueExecutor.get()->ConvertToTensor(
222 zeroMeanVector.data(), zeroMeanVector.size(), op::DataType::DT_FLOAT);223 zeroMeanVector.data(), zeroMeanVector.size(), op::DataType::DT_FLOAT);
@@ -224,48 +225,30 @@ aclnnStatus CommonLogicGeneralNormal(
224 auto oneStdTensor = uniqueExecutor.get()->ConvertToTensor(225 auto oneStdTensor = uniqueExecutor.get()->ConvertToTensor(
225 oneStdVector.data(), oneStdVector.size(), op::DataType::DT_FLOAT);226 oneStdVector.data(), oneStdVector.size(), op::DataType::DT_FLOAT);
226 227 
227- auto outDims = op::ToShapeVector(out->GetViewShape());
228- auto outShapeArray = uniqueExecutor.get()->AllocIntArray(outDims.data(), outDims.size());
229- zeroMeanTensor = l0op::BroadcastTo(zeroMeanTensor, outShapeArray, uniqueExecutor.get());
230- oneStdTensor = l0op::BroadcastTo(oneStdTensor, outShapeArray, uniqueExecutor.get());
231- 
232 auto normalOut = l0op::StatelessNormal(228 auto normalOut = l0op::StatelessNormal(
233- selfFloat, seed, offset, zeroMeanTensor, oneStdTensor, uniqueExecutor.get());229+ self, seed, offset, zeroMeanTensor, oneStdTensor, uniqueExecutor.get());
234 CHECK_RET(normalOut != nullptr, ACLNN_ERR_INNER_NULLPTR);230 CHECK_RET(normalOut != nullptr, ACLNN_ERR_INNER_NULLPTR);
235 231 
236- auto mulOut = l0op::Mul(normalOut, std, uniqueExecutor.get());232+ auto mulOut = l0op::MulInplace(std, normalOut, uniqueExecutor.get());
237 CHECK_RET(mulOut != nullptr, ACLNN_ERR_INNER_NULLPTR);233 CHECK_RET(mulOut != nullptr, ACLNN_ERR_INNER_NULLPTR);
238- addOut = l0op::Add(mulOut, mean, uniqueExecutor.get());234+ // ScalarTensor 模式下,mean的类型一定为f32,mulout数据类型可为f16\f32\bf16,但mean是scalar, 故不能用AddInplace
239- } else {235+ if (scalarMode == ScalarMode::ScalarTensor) {
240- // V4 Tensor 路径:kernel 以目标 dtype 输出并做三步舍入对齐 GPU tensor 路径236+ addOut = l0op::Add(mulOut, mean, uniqueExecutor.get());
241- const aclTensor* selfForKernel = self;237+ } else {
242- auto meanCasted = l0op::Cast(mean, DataType::DT_FLOAT, uniqueExecutor.get());238+ addOut = l0op::AddInplace(mean, mulOut, uniqueExecutor.get());
243- CHECK_RET(meanCasted != nullptr, ACLNN_ERR_INNER_NULLPTR);
244- auto stdCasted = l0op::Cast(std, DataType::DT_FLOAT, uniqueExecutor.get());
245- CHECK_RET(stdCasted != nullptr, ACLNN_ERR_INNER_NULLPTR);
246- 
247- bool selfNeedBcast = selfForKernel->GetViewShape() != out->GetViewShape();
248- bool meanNeedBcast = meanCasted->GetViewShape() != out->GetViewShape();
249- bool stdNeedBcast = stdCasted->GetViewShape() != out->GetViewShape();
250- if (selfNeedBcast || meanNeedBcast || stdNeedBcast) {
251- op::FVector<int64_t, op::MAX_DIM_NUM> outDims = op::ToShapeVector(out->GetViewShape());
252- auto outShapeArray = uniqueExecutor.get()->AllocIntArray(outDims.data(), outDims.size());
253- CHECK_RET(outShapeArray != nullptr, ACLNN_ERR_INNER_NULLPTR);
254- 
255- if (selfNeedBcast) {
256- selfForKernel = l0op::BroadcastTo(selfForKernel, outShapeArray, uniqueExecutor.get());
257- CHECK_RET(selfForKernel != nullptr, ACLNN_ERR_INNER_NULLPTR);
258- }
259- if (meanNeedBcast) {
260- meanCasted = l0op::BroadcastTo(meanCasted, outShapeArray, uniqueExecutor.get());
261- CHECK_RET(meanCasted != nullptr, ACLNN_ERR_INNER_NULLPTR);
262- }
263- if (stdNeedBcast) {
264- stdCasted = l0op::BroadcastTo(stdCasted, outShapeArray, uniqueExecutor.get());
265- CHECK_RET(stdCasted != nullptr, ACLNN_ERR_INNER_NULLPTR);
266- }
267 }239 }
268- addOut = l0op::StatelessNormal(selfForKernel, seed, offset, meanCasted, stdCasted, uniqueExecutor.get());240+ } else if (scalarMode == ScalarMode::TensorScalar) {
241+ FVector<float> zeroMeanVector = {0.0f};
242+ auto zeroMeanTensor = uniqueExecutor.get()->ConvertToTensor(
243+ zeroMeanVector.data(), zeroMeanVector.size(), op::DataType::DT_FLOAT);
244+ 
245+ auto normalOut = l0op::StatelessNormal(
246+ self, seed, offset, zeroMeanTensor, std, uniqueExecutor.get());
247+ CHECK_RET(normalOut != nullptr, ACLNN_ERR_INNER_NULLPTR);
248+ 
249+ addOut = l0op::AddInplace(mean, normalOut, uniqueExecutor.get());
250+ } else {
251+ addOut = l0op::StatelessNormal(self, seed, offset, mean, std, uniqueExecutor.get());
269 }252 }
270 }253 }
271 else{254 else{
@@ -363,7 +346,7 @@ aclnnStatus aclnnNormalTensorTensorGetWorkspaceSize(
363 // 拷贝mean副本346 // 拷贝mean副本
364 auto self = const_cast<aclTensor*>(meanCasted);347 auto self = const_cast<aclTensor*>(meanCasted);
365 return CommonLogicGeneralNormal(348 return CommonLogicGeneralNormal(
366- meanCasted, stdCasted, seed, offset, self, out, uniqueExecutor, workspaceSize, executor);349+ meanCasted, stdCasted, seed, offset, self, out, uniqueExecutor, workspaceSize, executor, ScalarMode::TensorTensor);
367}350}
368 351 
369// normal.Tensor_float_out352// normal.Tensor_float_out
@@ -403,7 +386,7 @@ aclnnStatus aclnnNormalTensorFloatGetWorkspaceSize(
403 // 拷贝mean副本386 // 拷贝mean副本
404 auto self = const_cast<aclTensor*>(out);387 auto self = const_cast<aclTensor*>(out);
405 return CommonLogicGeneralNormal(388 return CommonLogicGeneralNormal(
406- meanContiguous, stdContiguous, seed, offset, self, out, uniqueExecutor, workspaceSize, executor);389+ meanContiguous, stdContiguous, seed, offset, self, out, uniqueExecutor, workspaceSize, executor, ScalarMode::TensorScalar);
407}390}
408 391 
409// normal.float_Tensor_out392// normal.float_Tensor_out
@@ -443,7 +426,7 @@ aclnnStatus aclnnNormalFloatTensorGetWorkspaceSize(
443 // 拷贝std副本426 // 拷贝std副本
444 auto self = const_cast<aclTensor*>(stdContiguous);427 auto self = const_cast<aclTensor*>(stdContiguous);
445 return CommonLogicGeneralNormal(428 return CommonLogicGeneralNormal(
446- meanContiguous, stdContiguous, seed, offset, self, out, uniqueExecutor, workspaceSize, executor);429+ meanContiguous, stdContiguous, seed, offset, self, out, uniqueExecutor, workspaceSize, executor, ScalarMode::ScalarTensor);
447}430}
448 431 
449// normal.float_float_out432// normal.float_float_out
@@ -484,7 +467,7 @@ aclnnStatus aclnnNormalFloatFloatGetWorkspaceSize(
484 auto self = const_cast<aclTensor*>(out);467 auto self = const_cast<aclTensor*>(out);
485 return CommonLogicGeneralNormal(468 return CommonLogicGeneralNormal(
486 meanContiguous, stdContiguous, seed, offset, self, out, uniqueExecutor, workspaceSize, executor,469 meanContiguous, stdContiguous, seed, offset, self, out, uniqueExecutor, workspaceSize, executor,
487- true);470+ ScalarMode::ScalarScalar);
488}471}
489 472 
490aclnnStatus aclnnNormalTensorTensor(473aclnnStatus aclnnNormalTensorTensor(