已合并
add alg check #3840
梅国晗954517创建于 7月6日
add alg check #3840
已合并
共 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 | * \brief | 13 | * \brief |
| 14 | */ | 14 | */ |
| 15 | 15 | ||
| 16 | |||
| 16 | 17 | ||
| 17 | 18 | ||
| 18 | 19 | ||
| @@ -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 ExtractTensorValue | 74 | // GetData<>() on an empty tensor legally returns nullptr, which ExtractTensorValue |
| @@ -83,9 +100,8 @@ OpTilingConfig StatelessTruncatedNormalV2Tiling::BuildOpConfig() | |||
| 83 | 100 | ||
| 84 | ge::graphStatus StatelessTruncatedNormalV2Tiling::DoSimtBlockTiling() | 101 | ge::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 | |||
| 133 | IMPL_OP_OPTILING(StatelessTruncatedNormalV2) | 149 | IMPL_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 optiling | 153 | } // namespace optiling |