已合并
为AclnnNormalTensorTensor接口增加广播功能 #2146
huairuchen创建于 4月9日
为AclnnNormalTensorTensor接口增加广播功能 #2146
已合并
huairuchen创建于 4月9日
已删除 :pr_broadcast合入到cann/ops-mathmaster
2 个文件变更+49-9
@@ -16,6 +16,7 @@
16#include "aclnn_kernels/cast.h"16#include "aclnn_kernels/cast.h"
17#include "conversion/view_copy/op_api/view_copy.h"17#include "conversion/view_copy/op_api/view_copy.h"
18#include "aclnn_kernels/contiguous.h"18#include "aclnn_kernels/contiguous.h"
19+#include "conversion/broadcast_to/op_api/broadcast_to.h"
19#include "aclnn/aclnn_base.h"20#include "aclnn/aclnn_base.h"
20#include "opdev/shape_utils.h"21#include "opdev/shape_utils.h"
21#include "aclnn_kernels/common/op_error_check.h"22#include "aclnn_kernels/common/op_error_check.h"
@@ -218,13 +219,38 @@ aclnnStatus CommonLogicGeneralNormal(
218 const aclTensor* addOut = nullptr;219 const aclTensor* addOut = nullptr;
219 220 
220 if(GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510 && self->GetDataType() != DataType::DT_DOUBLE){221 if(GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510 && self->GetDataType() != DataType::DT_DOUBLE){
221- // 调用normal_算子kernel function(AI Core算子)222+ // V3 kernel 要求 mean/std 参数为 DT_FLOAT
222- // V3 kernel 要求 mean/std 参数为 DT_FLOAT,与 InplaceNormal 保持一致223+ auto meanCasted = l0op::Cast(mean, DataType::DT_FLOAT, uniqueExecutor.get());
223- auto meanFP32 = l0op::Cast(mean, DataType::DT_FLOAT, uniqueExecutor.get());224+ CHECK_RET(meanCasted != nullptr, ACLNN_ERR_INNER_NULLPTR);
224- CHECK_RET(meanFP32 != nullptr, ACLNN_ERR_INNER_NULLPTR);225+ auto stdCasted = l0op::Cast(std, DataType::DT_FLOAT, uniqueExecutor.get());
225- auto stdFP32 = l0op::Cast(std, DataType::DT_FLOAT, uniqueExecutor.get());226+ CHECK_RET(stdCasted != nullptr, ACLNN_ERR_INNER_NULLPTR);
226- CHECK_RET(stdFP32 != nullptr, ACLNN_ERR_INNER_NULLPTR);227+ 
227- addOut = l0op::StatelessRandomNormalV3(self, keyArr, counterArr, meanFP32, stdFP32, uniqueExecutor.get());228+ // V3 kernel 不支持广播,需要将 shape out 不一致的多元素 tensor 广播到 out shape
229+ // Size()==1 的 tensor(scalar 转换而来)由 V3 kernel 的 scalar 路径原生处理,无需广播
230+ // self 需要单独广播以保持原始 dtype(决定 V3 kernel 的计算精度)
231+ bool selfNeedBcast = self->GetViewShape() != out->GetViewShape();
232+ bool meanNeedBcast = meanCasted->GetViewShape() != out->GetViewShape() && mean->Size() > 1;
233+ bool stdNeedBcast = stdCasted->GetViewShape() != out->GetViewShape() && std->Size() > 1;
234+ if (selfNeedBcast || meanNeedBcast || stdNeedBcast) {
235+ op::FVector<int64_t, op::MAX_DIM_NUM> outDims = op::ToShapeVector(out->GetViewShape());
236+ auto outShapeArray = uniqueExecutor.get()->AllocIntArray(outDims.data(), outDims.size());
237+ CHECK_RET(outShapeArray != nullptr, ACLNN_ERR_INNER_NULLPTR);
238+ 
239+ if (selfNeedBcast) {
240+ auto selfBroadcast = l0op::BroadcastTo(self, outShapeArray, uniqueExecutor.get());
241+ CHECK_RET(selfBroadcast != nullptr, ACLNN_ERR_INNER_NULLPTR);
242+ self = const_cast<aclTensor*>(selfBroadcast);
243+ }
244+ if (meanNeedBcast) {
245+ meanCasted = l0op::BroadcastTo(meanCasted, outShapeArray, uniqueExecutor.get());
246+ CHECK_RET(meanCasted != nullptr, ACLNN_ERR_INNER_NULLPTR);
247+ }
248+ if (stdNeedBcast) {
249+ stdCasted = l0op::BroadcastTo(stdCasted, outShapeArray, uniqueExecutor.get());
250+ CHECK_RET(stdCasted != nullptr, ACLNN_ERR_INNER_NULLPTR);
251+ }
252+ }
253+ addOut = l0op::StatelessRandomNormalV3(self, keyArr, counterArr, meanCasted, stdCasted, uniqueExecutor.get());
228 }254 }
229 else{255 else{
230 // 调用normal_算子kernel function(AI Cpu算子)256 // 调用normal_算子kernel function(AI Cpu算子)
@@ -93,16 +93,30 @@ ge::graphStatus StatelessRandomNormalV3Tiling::UniqueProcess()
93{93{
94 uint32_t v3KernelMode = 0;94 uint32_t v3KernelMode = 0;
95 95 
96+ auto outputShape = context_->GetOutputShape(OUTPUT_IDX_Y);
97+ OP_CHECK_NULL_WITH_CONTEXT(context_, outputShape);
98+ int64_t outputSize = outputShape->GetStorageShape().GetShapeSize();
99+ 
96 auto meanTensor = context_->GetInputTensor(INPUT_IDX_MEAN);100 auto meanTensor = context_->GetInputTensor(INPUT_IDX_MEAN);
97 OP_CHECK_NULL_WITH_CONTEXT(context_, meanTensor);101 OP_CHECK_NULL_WITH_CONTEXT(context_, meanTensor);
98- if (meanTensor->GetShapeSize() == 1) {102+ int64_t meanSize = meanTensor->GetShapeSize();
103+ if (meanSize == 1) {
99 v3KernelMode |= MEAN_SCALAR_FLAG;104 v3KernelMode |= MEAN_SCALAR_FLAG;
105+ } else if (meanSize != outputSize) {
106+ OP_LOGE(context_, "StatelessRandomNormalV3 does not support broadcast for mean, "
107+ "mean shapeSize: %ld != output shapeSize: %ld.", meanSize, outputSize);
108+ return ge::GRAPH_FAILED;
100 }109 }
101 110 
102 auto stdevTensor = context_->GetInputTensor(INPUT_IDX_STDEV);111 auto stdevTensor = context_->GetInputTensor(INPUT_IDX_STDEV);
103 OP_CHECK_NULL_WITH_CONTEXT(context_, stdevTensor);112 OP_CHECK_NULL_WITH_CONTEXT(context_, stdevTensor);
104- if (stdevTensor->GetShapeSize() == 1) {113+ int64_t stdevSize = stdevTensor->GetShapeSize();
114+ if (stdevSize == 1) {
105 v3KernelMode |= STDEV_SCALAR_FLAG;115 v3KernelMode |= STDEV_SCALAR_FLAG;
116+ } else if (stdevSize != outputSize) {
117+ OP_LOGE(context_, "StatelessRandomNormalV3 does not support broadcast for stdev, "
118+ "stdev shapeSize: %ld != output shapeSize: %ld.", stdevSize, outputSize);
119+ return ge::GRAPH_FAILED;
106 }120 }
107 121 
108 tilingData_.v3KernelMode = v3KernelMode;122 tilingData_.v3KernelMode = v3KernelMode;