已合并
add alg check #3840
梅国晗954517创建于 7月6日
add alg check #3840
已合并
梅国晗954517创建于 7月6日
1 个文件变更+35-19
Mrandom/stateless_truncated_normal_v2/op_host/arch35/stateless_truncated_normal_v2_tiling_arch35.cpp+35-19
@@ -13,6 +13,7 @@
13 * \brief13 * \brief
14 */14 */
15#include "stateless_truncated_normal_v2_tiling_arch35.h"15#include "stateless_truncated_normal_v2_tiling_arch35.h"
16#include <string>
16#include "log/log.h"17#include "log/log.h"
17#include "platform/platform_ascendc.h"18#include "platform/platform_ascendc.h"
18#include "register/op_def_registry.h"19#include "register/op_def_registry.h"
@@ -37,21 +38,37 @@ OpTilingConfig StatelessTruncatedNormalV2Tiling::BuildOpConfig()
37{38{
38 OpTilingConfig config;39 OpTilingConfig config;
39 40 
40 config.inputCheckRules = {41 config.inputCheckRules = {{INPUT_IDX_SHAPE, {{ge::DT_INT32, ge::DT_INT64}, -1, {1}, nullptr}},
41 {INPUT_IDX_SHAPE, {{ge::DT_INT32, ge::DT_INT64}, -1, {1}, nullptr}},42 {INPUT_IDX_KEY, {{ge::DT_UINT64}, 1, {1}, nullptr}},
42 {INPUT_IDX_KEY, {{ge::DT_UINT64}, 1, {1}, nullptr}},43 {INPUT_IDX_COUNTER, {{ge::DT_UINT64}, 2, {1}, nullptr}},
43 {INPUT_IDX_COUNTER, {{ge::DT_UINT64}, 2, {1}, nullptr}},44 {INPUT_IDX_ALG, {{ge::DT_INT32}, 1, {0, 1}, [](gert::TilingContext* ctx) {
44 {INPUT_IDX_ALG, {{ge::DT_INT32}, 1, {0, 1}, nullptr}}};45 const auto* algTensor = ctx->GetInputTensor(INPUT_IDX_ALG);
45 config.outputCheckRules = {46 if (algTensor == nullptr) {
46 {OUTPUT_IDX_Y, {{ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16}, -1, {}, nullptr}}};47 return false;
47 config.attrCheckRules = {48 }
48 {INDEX_0, [](gert::TilingContext* ctx) {49 const int32_t* algData = algTensor->GetData<int32_t>();
49 const auto* attrs = ctx->GetAttrs();50 if (algData == nullptr) {
50 const int64_t* attrPtr = attrs ? attrs->GetAttrPointer<int64_t>(INDEX_0) : nullptr;51 return false;
51 const auto* outDesc = ctx->GetOutputDesc(OUTPUT_IDX_Y);52 }
52 return attrPtr != nullptr && outDesc != nullptr &&53 constexpr int32_t ALG_PHILOX = 1;
53 static_cast<ge::DataType>(*attrPtr) == outDesc->GetDataType();54 if (algData[0] != ALG_PHILOX) {
54 }}};55 std::string valueStr = std::to_string(algData[0]);
56 std::string reasonMsg = "Unsupported algorithm id: " + valueStr;
57 OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
58 ctx->GetNodeName(), "input alg", valueStr.c_str(),
59 reasonMsg.c_str());
60 return false;
61 }
62 return true;
63 }}}};
64 config.outputCheckRules = {{OUTPUT_IDX_Y, {{ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16}, -1, {}, nullptr}}};
65 config.attrCheckRules = {{INDEX_0, [](gert::TilingContext* ctx) {
66 const auto* attrs = ctx->GetAttrs();
67 const int64_t* attrPtr = attrs ? attrs->GetAttrPointer<int64_t>(INDEX_0) : nullptr;
68 const auto* outDesc = ctx->GetOutputDesc(OUTPUT_IDX_Y);
69 return attrPtr != nullptr && outDesc != nullptr &&
70 static_cast<ge::DataType>(*attrPtr) == outDesc->GetDataType();
71 }}};
55 config.getOutputSize = [](gert::TilingContext* ctx, int64_t& size) {72 config.getOutputSize = [](gert::TilingContext* ctx, int64_t& size) {
56 // 0-dim scalar output: the `shape` input is an empty tensor (ShapeSize==0).73 // 0-dim scalar output: the `shape` input is an empty tensor (ShapeSize==0).
57 // GetData<>() on an empty tensor legally returns nullptr, which ExtractTensorValue74 // GetData<>() on an empty tensor legally returns nullptr, which ExtractTensorValue
@@ -83,9 +100,8 @@ OpTilingConfig StatelessTruncatedNormalV2Tiling::BuildOpConfig()
83 100 
84ge::graphStatus StatelessTruncatedNormalV2Tiling::DoSimtBlockTiling()101ge::graphStatus StatelessTruncatedNormalV2Tiling::DoSimtBlockTiling()
85{102{
86 OP_CHECK_IF(103 OP_CHECK_IF((totalCoreNum_ <= 0), OP_LOGE(opName_, "totalCoreNum is less than or equal to 0. please check."),
87 (totalCoreNum_ <= 0), OP_LOGE(opName_, "totalCoreNum is less than or equal to 0. please check."),104 return ge::GRAPH_FAILED);
88 return ge::GRAPH_FAILED);
89 int64_t threadNum = Ops::Base::CeilAlign(simtTilingData_.outputSize, THREAD_DISPOSAL_NUM);105 int64_t threadNum = Ops::Base::CeilAlign(simtTilingData_.outputSize, THREAD_DISPOSAL_NUM);
90 int64_t coreNum = Ops::Base::CeilAlign(threadNum, MAX_THREAD_NUM);106 int64_t coreNum = Ops::Base::CeilAlign(threadNum, MAX_THREAD_NUM);
91 simtTilingData_.usedCoreNum = std::min(coreNum, totalCoreNum_);107 simtTilingData_.usedCoreNum = std::min(coreNum, totalCoreNum_);
@@ -133,5 +149,5 @@ static ge::graphStatus TilingPrepare4StatelessTruncatedNormalV2(gert::TilingPars
133IMPL_OP_OPTILING(StatelessTruncatedNormalV2)149IMPL_OP_OPTILING(StatelessTruncatedNormalV2)
134 .Tiling(Tiling4StatelessTruncatedNormalV2)150 .Tiling(Tiling4StatelessTruncatedNormalV2)
135 .TilingParse<RandomOperatorCompileInfo>(TilingPrepare4StatelessTruncatedNormalV2)151 .TilingParse<RandomOperatorCompileInfo>(TilingPrepare4StatelessTruncatedNormalV2)
136 .TilingInputsDataDependency({INPUT_IDX_SHAPE, INPUT_IDX_KEY, INPUT_IDX_COUNTER});152 .TilingInputsDataDependency({INPUT_IDX_SHAPE, INPUT_IDX_KEY, INPUT_IDX_COUNTER, INPUT_IDX_ALG});
137} // namespace optiling153} // namespace optiling