已合并
为AclnnNormalTensorTensor接口增加广播功能 #2146
huairuchen创建于 4月9日
为AclnnNormalTensorTensor接口增加广播功能 #2146
已合并
从已删除 :pr_broadcast合入到cann/ops-mathmaster
共 2 个文件变更+49-9
| @@ -16,6 +16,7 @@ | |||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | + | ||
| 19 | 20 | ||
| 20 | 21 | ||
| 21 | 22 | ||
| @@ -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; |