已合并
feat: mul/select 算子支持 PcieThrough 场景识别与回退 #5116
hahaha22创建于 8月27日
feat: mul/select 算子支持 PcieThrough 场景识别与回退 #5116
已合并
hahaha22创建于 8月27日
共 4 个文件变更+89-58
@@ -16,33 +16,40 @@
16#include "log/log.h"16#include "log/log.h"
17#include "infershape_broadcast_util.h"17#include "infershape_broadcast_util.h"
18#include "register/op_impl_registry.h"18#include "register/op_impl_registry.h"
19+#include "version/metadef_version.h"
19 20 
20using namespace ge;21using namespace ge;
21namespace ops {22namespace 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+#endif
54+ ;
55+} // namespace ops
@@ -21,6 +21,7 @@
21#include "register/op_def_registry.h"21#include "register/op_def_registry.h"
22#include "util/math_util.h"22#include "util/math_util.h"
23#include "op_host/math_tiling_templates_registry.h"23#include "op_host/math_tiling_templates_registry.h"
24+#include "version/metadef_version.h"
24 25 
25using namespace AscendC;26using namespace AscendC;
26using namespace ge;27using 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+#if defined(METADEF_VERSION_NUM) && METADEF_VERSION_NUM >= 90200000
111+ bool isPcieThrough = context_->GetPcieThroughFlag();
112+ OP_LOGD(context_->GetNodeName(), "IsPcieThrough: %s", (isPcieThrough ? "true" : "false"));
113+ return isPcieThrough;
114+#else
115+ OP_LOGD(context_->GetNodeName(), "IsPcieThrough: false");
116+ return false;
117+#endif
118+}
119+ 
107bool SelectSimtTiling::IsCapable()120bool 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 @@
26namespace optiling {26namespace optiling {
27 27 
28class SelectSimtTiling : public Ops::Base::TilingBaseClass {28class 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#ifdef DAVID_FPGA57#ifdef DAVID_FPGA
57- int64_t threadNum_ = 128;58+ int64_t threadNum_ = 128;
58#else59#else
59- int64_t threadNum_ = 2048;60+ int64_t threadNum_ = 2048;
60#endif61#endif
61};62};
62 63 
63-} // namespace optiling64+} // namespace optiling
64 65 
65-#endif // OPS_BUILD_IN_OP_TILING_RUNTIME_SELECT_SIMT_TILING_H66+#endif // OPS_BUILD_IN_OP_TILING_RUNTIME_SELECT_SIMT_TILING_H
@@ -14,18 +14,25 @@
14 */14 */
15#include "log/log.h"15#include "log/log.h"
16#include "register/op_impl_registry.h"16#include "register/op_impl_registry.h"
17+#include "version/metadef_version.h"
17 18 
18using namespace ge;19using namespace ge;
19namespace ops {20namespace 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+#if defined(METADEF_VERSION_NUM) && METADEF_VERSION_NUM >= 90200000
35+ .SetSupportPcieThrough()
36+#endif
37+ ;
38+} // namespace ops