已合并
feat: mul/select 算子支持 PcieThrough 场景识别与回退 #5116
hahaha22创建于 8月27日
feat: mul/select 算子支持 PcieThrough 场景识别与回退 #5116
已合并
共 4 个文件变更+89-58
| @@ -16,33 +16,40 @@ | |||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | + | ||
| 19 | 20 | ||
| 20 | using namespace ge; | 21 | using namespace ge; |
| 21 | namespace ops { | 22 | namespace ops { |
| 22 | -static std::string ShapeCannotBroadcastMsg(const gert::Shape& shape1, const gert::Shape& shape2) { | 23 | +static std::string ShapeCannotBroadcastMsg(const gert::Shape& shape1, const gert::Shape& shape2) |
| 23 | - std::string res = "shape "; | 24 | +{ |
| 24 | - res += Ops::Base::ToString(shape1); | 25 | + std::string res = "shape "; |
| 25 | - res += " and "; | 26 | + res += Ops::Base::ToString(shape1); |
| 26 | - res += Ops::Base::ToString(shape2); | 27 | + res += " and "; |
| 27 | - res += " cannot broadcast!"; | 28 | + res += Ops::Base::ToString(shape2); |
| 28 | - return res; | 29 | + res += " cannot broadcast!"; |
| 30 | + return res; | ||
| 29 | } | 31 | } |
| 30 | 32 | ||
| 31 | -static ge::graphStatus InferShape4Broadcast(gert::InferShapeContext* context) { | 33 | +static ge::graphStatus InferShape4Broadcast(gert::InferShapeContext* context) |
| 32 | - auto in_shape1 = context->GetInputShape(0); | 34 | +{ |
| 33 | - OP_CHECK_NULL_WITH_CONTEXT(context, in_shape1); | 35 | + auto in_shape1 = context->GetInputShape(0); |
| 34 | - auto in_shape2 = context->GetInputShape(1); | 36 | + OP_CHECK_NULL_WITH_CONTEXT(context, in_shape1); |
| 35 | - OP_CHECK_NULL_WITH_CONTEXT(context, in_shape2); | 37 | + auto in_shape2 = context->GetInputShape(1); |
| 36 | - auto out_shape = context->GetOutputShape(0); | 38 | + OP_CHECK_NULL_WITH_CONTEXT(context, in_shape2); |
| 37 | - OP_CHECK_NULL_WITH_CONTEXT(context, out_shape); | 39 | + auto out_shape = context->GetOutputShape(0); |
| 40 | + OP_CHECK_NULL_WITH_CONTEXT(context, out_shape); | ||
| 38 | 41 | ||
| 39 | - OP_CHECK_IF((!Ops::Base::BroadcastShape(in_shape1, in_shape2, out_shape)), | 42 | + OP_CHECK_IF((!Ops::Base::BroadcastShape(in_shape1, in_shape2, out_shape)), |
| 40 | - OP_LOGE(context->GetNodeName(), "%s", ShapeCannotBroadcastMsg(*in_shape2, *in_shape1).c_str()), | 43 | + OP_LOGE(context->GetNodeName(), "%s", ShapeCannotBroadcastMsg(*in_shape2, *in_shape1).c_str()), |
| 41 | - return ge::GRAPH_FAILED); | 44 | + return ge::GRAPH_FAILED); |
| 42 | 45 | ||
| 43 | - return ge::GRAPH_SUCCESS; | 46 | + return ge::GRAPH_SUCCESS; |
| 44 | } | 47 | } |
| 45 | 48 | ||
| 46 | -IMPL_OP_INFERSHAPE(Mul).InferShape(InferShape4Broadcast); | 49 | +IMPL_OP_INFERSHAPE(Mul) |
| 47 | - | 50 | + .InferShape(InferShape4Broadcast) |
| 48 | -} | 51 | +#if defined(METADEF_VERSION_NUM) && METADEF_VERSION_NUM >= 90200000 |
| 52 | + .SetSupportPcieThrough() | ||
| 53 | + | ||
| 54 | + ; | ||
| 55 | +} // namespace ops | ||
| @@ -21,6 +21,7 @@ | |||
| 21 | 21 | ||
| 22 | 22 | ||
| 23 | 23 | ||
| 24 | + | ||
| 24 | 25 | ||
| 25 | using namespace AscendC; | 26 | using namespace AscendC; |
| 26 | using namespace ge; | 27 | using namespace ge; |
| @@ -104,8 +105,23 @@ bool SelectSimtTiling::XDtypeImprove() | |||
| 104 | return false; | 105 | return false; |
| 105 | } | 106 | } |
| 106 | 107 | ||
| 108 | +bool SelectSimtTiling::IsPcieThrough() | ||
| 109 | +{ | ||
| 110 | + | ||
| 111 | + bool isPcieThrough = context_->GetPcieThroughFlag(); | ||
| 112 | + OP_LOGD(context_->GetNodeName(), "IsPcieThrough: %s", (isPcieThrough ? "true" : "false")); | ||
| 113 | + return isPcieThrough; | ||
| 114 | + | ||
| 115 | + OP_LOGD(context_->GetNodeName(), "IsPcieThrough: false"); | ||
| 116 | + return false; | ||
| 117 | + | ||
| 118 | +} | ||
| 119 | + | ||
| 107 | bool SelectSimtTiling::IsCapable() | 120 | bool SelectSimtTiling::IsCapable() |
| 108 | { | 121 | { |
| 122 | + if (IsPcieThrough()) { | ||
| 123 | + return false; | ||
| 124 | + } | ||
| 109 | if (!IsMatchAB()) { | 125 | if (!IsMatchAB()) { |
| 110 | return false; | 126 | return false; |
| 111 | } | 127 | } |
| @@ -26,40 +26,41 @@ | |||
| 26 | namespace optiling { | 26 | namespace optiling { |
| 27 | 27 | ||
| 28 | class SelectSimtTiling : public Ops::Base::TilingBaseClass { | 28 | class SelectSimtTiling : public Ops::Base::TilingBaseClass { |
| 29 | - public: | 29 | +public: |
| 30 | - explicit SelectSimtTiling(gert::TilingContext* context) : Ops::Base::TilingBaseClass(context) {} | 30 | + explicit SelectSimtTiling(gert::TilingContext* context) : Ops::Base::TilingBaseClass(context) {} |
| 31 | 31 | ||
| 32 | - protected: | 32 | +protected: |
| 33 | - bool IsCapable() override; | 33 | + bool IsCapable() override; |
| 34 | - ge::graphStatus GetPlatformInfo() override; | 34 | + ge::graphStatus GetPlatformInfo() override; |
| 35 | - ge::graphStatus GetShapeAttrsInfo() override; | 35 | + ge::graphStatus GetShapeAttrsInfo() override; |
| 36 | - ge::graphStatus DoOpTiling() override; | 36 | + ge::graphStatus DoOpTiling() override; |
| 37 | - ge::graphStatus DoLibApiTiling() override; | 37 | + ge::graphStatus DoLibApiTiling() override; |
| 38 | - uint64_t GetTilingKey() const override; | 38 | + uint64_t GetTilingKey() const override; |
| 39 | - ge::graphStatus GetWorkspaceSize() override; | 39 | + ge::graphStatus GetWorkspaceSize() override; |
| 40 | - ge::graphStatus PostTiling() override; | 40 | + ge::graphStatus PostTiling() override; |
| 41 | 41 | ||
| 42 | - private: | 42 | +private: |
| 43 | - bool IsMatchAB(); | 43 | + bool IsPcieThrough(); |
| 44 | - bool XDtypeImprove(); | 44 | + bool IsMatchAB(); |
| 45 | + bool XDtypeImprove(); | ||
| 45 | 46 | ||
| 46 | - uint64_t tilingKey = 0; | 47 | + uint64_t tilingKey = 0; |
| 47 | - int64_t aivNum_; | 48 | + int64_t aivNum_; |
| 48 | - int64_t ubSize_ = 0; | 49 | + int64_t ubSize_ = 0; |
| 49 | - int64_t needCoreNum_ = 0; | 50 | + int64_t needCoreNum_ = 0; |
| 50 | - gert::Shape conditionShape_; | 51 | + gert::Shape conditionShape_; |
| 51 | - gert::Shape x1Shape_; | 52 | + gert::Shape x1Shape_; |
| 52 | - gert::Shape x2Shape_; | 53 | + gert::Shape x2Shape_; |
| 53 | - int64_t aSize_ = 1; | 54 | + int64_t aSize_ = 1; |
| 54 | - int64_t bSize_ = 1; | 55 | + int64_t bSize_ = 1; |
| 55 | - int64_t ySize_ = 1; | 56 | + int64_t ySize_ = 1; |
| 56 | 57 | ||
| 57 | - int64_t threadNum_ = 128; | 58 | + int64_t threadNum_ = 128; |
| 58 | 59 | ||
| 59 | - int64_t threadNum_ = 2048; | 60 | + int64_t threadNum_ = 2048; |
| 60 | 61 | ||
| 61 | }; | 62 | }; |
| 62 | 63 | ||
| 63 | -} // namespace optiling | 64 | +} // namespace optiling |
| 64 | 65 | ||
| 65 | -#endif // OPS_BUILD_IN_OP_TILING_RUNTIME_SELECT_SIMT_TILING_H | 66 | +#endif // OPS_BUILD_IN_OP_TILING_RUNTIME_SELECT_SIMT_TILING_H |
| @@ -14,18 +14,25 @@ | |||
| 14 | */ | 14 | */ |
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | + | ||
| 17 | 18 | ||
| 18 | using namespace ge; | 19 | using namespace ge; |
| 19 | namespace ops { | 20 | namespace ops { |
| 20 | -static graphStatus InferShape4Select(gert::InferShapeContext* context) { | 21 | +static graphStatus InferShape4Select(gert::InferShapeContext* context) |
| 21 | - const gert::Shape* x1_shape = context->GetInputShape(1); | 22 | +{ |
| 22 | - OP_CHECK_NULL_WITH_CONTEXT(context, x1_shape); | 23 | + const gert::Shape* x1_shape = context->GetInputShape(1); |
| 24 | + OP_CHECK_NULL_WITH_CONTEXT(context, x1_shape); | ||
| 23 | 25 | ||
| 24 | - gert::Shape* y_shape = context->GetOutputShape(0); | 26 | + gert::Shape* y_shape = context->GetOutputShape(0); |
| 25 | - OP_CHECK_NULL_WITH_CONTEXT(context, y_shape); | 27 | + OP_CHECK_NULL_WITH_CONTEXT(context, y_shape); |
| 26 | - *y_shape = *x1_shape; | 28 | + *y_shape = *x1_shape; |
| 27 | 29 | ||
| 28 | - return GRAPH_SUCCESS; | 30 | + return GRAPH_SUCCESS; |
| 29 | -} | ||
| 30 | -IMPL_OP_INFERSHAPE(Select).InferShape(InferShape4Select); | ||
| 31 | } | 31 | } |
| 32 | +IMPL_OP_INFERSHAPE(Select) | ||
| 33 | + .InferShape(InferShape4Select) | ||
| 34 | + | ||
| 35 | + .SetSupportPcieThrough() | ||
| 36 | + | ||
| 37 | + ; | ||
| 38 | +} // namespace ops | ||