From 32d3ce532b7346c3c954da34e8a3c06326a89b6d Mon Sep 17 00:00:00 2001 From: Yang Jiao Date: Wed, 8 Sep 2021 21:53:52 +0800 Subject: [PATCH] cluster standardnormal --- akg | 2 +- .../_extends/graph_kernel/model/model.py | 1 + .../graph_kernel/expanders/standardnormal.cc | 44 +++++++++++++++++++ .../graph_kernel/graph_kernel_expander.cc | 3 ++ .../optimizer/graph_kernel/model/op_node.cc | 5 +++ .../optimizer/graph_kernel/model/op_node.h | 11 +++++ .../graph_kernel/model/op_register.h | 1 + mindspore/core/base/core_ops.h | 3 ++ 8 files changed, 69 insertions(+), 1 deletion(-) create mode 100644 mindspore/ccsrc/backend/optimizer/graph_kernel/expanders/standardnormal.cc diff --git a/akg b/akg index 78c10ce87a7..545ebff8aff 160000 --- a/akg +++ b/akg @@ -1 +1 @@ -Subproject commit 78c10ce87a7edbf28b8ccd2b23028cd2126cba61 +Subproject commit 545ebff8aff5fb7877337f004652b36fe7ca515e diff --git a/mindspore/_extends/graph_kernel/model/model.py b/mindspore/_extends/graph_kernel/model/model.py index 5ebc4c59306..f64e4b39204 100644 --- a/mindspore/_extends/graph_kernel/model/model.py +++ b/mindspore/_extends/graph_kernel/model/model.py @@ -234,6 +234,7 @@ class PrimLib: 'Gather': Prim(OPAQUE), 'GatherNd': Prim(OPAQUE), 'UnsortedSegmentSum': Prim(OPAQUE), + 'StandardNormal': Prim(OPAQUE), 'UserDefined': Prim(OPAQUE), } diff --git a/mindspore/ccsrc/backend/optimizer/graph_kernel/expanders/standardnormal.cc b/mindspore/ccsrc/backend/optimizer/graph_kernel/expanders/standardnormal.cc new file mode 100644 index 00000000000..7cf0cadc983 --- /dev/null +++ b/mindspore/ccsrc/backend/optimizer/graph_kernel/expanders/standardnormal.cc @@ -0,0 +1,44 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include +#include + +#include "backend/optimizer/graph_kernel/expanders/expander_factory.h" + +namespace mindspore { +namespace opt { +namespace expanders { +class StandardNormal : public OpExpander { + public: + StandardNormal() { + std::initializer_list attrs{"seed", "seed2"}; + validators_.emplace_back(std::make_unique(attrs)); + } + ~StandardNormal() {} + NodePtrList Expand() override { + const auto &inputs = gb.Get()->inputs(); + const auto &input_x = inputs[0]; + auto shape = MakeValue(outputs_info_[0].shape); + auto result = + gb.Emit("StandardNormal", {input_x}, {{"shape", shape}, {"seed", attrs_["seed"]}, {"seed2", attrs_["seed2"]}}); + return {result}; + } +}; +OP_EXPANDER_REGISTER("StandardNormal", StandardNormal); +} // namespace expanders +} // namespace opt +} // namespace mindspore diff --git a/mindspore/ccsrc/backend/optimizer/graph_kernel/graph_kernel_expander.cc b/mindspore/ccsrc/backend/optimizer/graph_kernel/graph_kernel_expander.cc index 4b0c608837b..6339c021380 100644 --- a/mindspore/ccsrc/backend/optimizer/graph_kernel/graph_kernel_expander.cc +++ b/mindspore/ccsrc/backend/optimizer/graph_kernel/graph_kernel_expander.cc @@ -46,6 +46,7 @@ using context::OpLevel_1; constexpr size_t kAssignInputIdx = 1; constexpr size_t kLambOptimizerInputIdx = 12; constexpr size_t kLambWeightInputIdx = 4; +constexpr size_t kRandomInputIdx = 1; std::vector GetExpandOps() { std::vector> expand_ops_with_level = { @@ -93,6 +94,7 @@ std::vector GetExpandOps() { {kGPUDevice, OpLevel_0, prim::kPrimSquareSumAll}, {kGPUDevice, OpLevel_0, prim::kPrimIdentityMath}, {kGPUDevice, OpLevel_0, prim::kPrimOnesLike}, + {kGPUDevice, OpLevel_0, prim::kPrimStandardNormal}, }; const auto &flags = context::GraphKernelFlags::GetInstance(); std::vector expand_ops = GetValidOps(expand_ops_with_level, flags.fusion_ops_level); @@ -201,6 +203,7 @@ ExpanderPtr GraphKernelExpander::GetExpander(const AnfNodePtr &node) { {prim::kPrimAssignSub, std::make_shared(kAssignInputIdx)}, {prim::kLambApplyOptimizerAssign, std::make_shared(kLambOptimizerInputIdx)}, {prim::kLambApplyWeightAssign, std::make_shared(kLambWeightInputIdx)}, + {prim::kPrimStandardNormal, std::make_shared(kRandomInputIdx)}, }; for (auto &e : expanders) { diff --git a/mindspore/ccsrc/backend/optimizer/graph_kernel/model/op_node.cc b/mindspore/ccsrc/backend/optimizer/graph_kernel/model/op_node.cc index 26c06295959..400c191cb15 100644 --- a/mindspore/ccsrc/backend/optimizer/graph_kernel/model/op_node.cc +++ b/mindspore/ccsrc/backend/optimizer/graph_kernel/model/op_node.cc @@ -543,6 +543,11 @@ void ComplexOp::CheckType(const NodePtrList &inputs, const DAttrs &attrs) { MS_LOG(EXCEPTION) << "Complex's input[0] and inputs[1]'s type mismatch"; } } + +DShape StandardNormalOp::InferShape(const NodePtrList &inputs, const DAttrs &attrs) { + CHECK_ATTR(attrs, "shape"); + return GetListInt(attrs.find("shape")->second); +} } // namespace graphkernel } // namespace opt } // namespace mindspore diff --git a/mindspore/ccsrc/backend/optimizer/graph_kernel/model/op_node.h b/mindspore/ccsrc/backend/optimizer/graph_kernel/model/op_node.h index eca01f79c93..cc1704613c9 100644 --- a/mindspore/ccsrc/backend/optimizer/graph_kernel/model/op_node.h +++ b/mindspore/ccsrc/backend/optimizer/graph_kernel/model/op_node.h @@ -282,6 +282,17 @@ class ComplexOp : public ElemwiseOp { void CheckType(const NodePtrList &inputs, const DAttrs &attrs) override; TypeId InferType(const NodePtrList &inputs, const DAttrs &attrs) override { return TypeId::kNumberTypeComplex64; } }; + +class StandardNormalOp : public OpaqueOp { + public: + StandardNormalOp(const std::string &op, const std::string &node_name) : OpaqueOp("StandardNormal", node_name) {} + ~StandardNormalOp() = default; + + protected: + DShape InferShape(const NodePtrList &inputs, const DAttrs &attrs) override; + TypeId InferType(const NodePtrList &inputs, const DAttrs &attrs) override { return TypeId::kNumberTypeFloat32; } + DFormat InferFormat(const NodePtrList &inputs, const DAttrs &attrs) override { return kOpFormat_DEFAULT; } +}; } // namespace graphkernel } // namespace opt } // namespace mindspore diff --git a/mindspore/ccsrc/backend/optimizer/graph_kernel/model/op_register.h b/mindspore/ccsrc/backend/optimizer/graph_kernel/model/op_register.h index 9b1bb1b85a3..7ec7ccc7aa7 100644 --- a/mindspore/ccsrc/backend/optimizer/graph_kernel/model/op_register.h +++ b/mindspore/ccsrc/backend/optimizer/graph_kernel/model/op_register.h @@ -79,6 +79,7 @@ class OpRegistry { Register("CImag", OP_CREATOR(CImagOp)); Register("Complex", OP_CREATOR(ComplexOp)); Register("Opaque", OP_CREATOR(OpaqueOp)); + Register("StandardNormal", OP_CREATOR(StandardNormalOp)); } ~OpRegistry() = default; std::unordered_map> creators; diff --git a/mindspore/core/base/core_ops.h b/mindspore/core/base/core_ops.h index b24e3d0651e..f0b2f28a68d 100644 --- a/mindspore/core/base/core_ops.h +++ b/mindspore/core/base/core_ops.h @@ -702,6 +702,9 @@ inline const PrimitivePtr kPrimBroadcastGradientArgs = std::make_shared(kDynamicBroadcastGradientArgs); +// Random +inline const PrimitivePtr kPrimStandardNormal = std::make_shared("StandardNormal"); + class DoSignaturePrimitive : public Primitive { public: explicit DoSignaturePrimitive(const std::string &name, const ValuePtr &function)